Views
No views yet
1import torch
2import torch.nn.functional as F
3import torchvision
4import torchvision.transforms as transforms
5from transformers import AutoModelForImageClassification
6from matplotlib import pyplot as plt
7
8model_name = "Normal1919/THW"
9
10model = AutoModelForImageClassification.from_pretrained(model_name)
11model.eval()
12# model = torch.compile(model)
13
14image_transform = transforms.Compose([
15 transforms.ToPILImage(),
16 transforms.Resize((256, 256)),
17 transforms.ToTensor(),
18 transforms.Normalize(mean=[0.697, 0.633, 0.635], std=[0.3135, 0.320, 0.315])
19])
20
21with torch.no_grad():
22 image_raw = torchvision.io.read_image("test_img/c9f00dbb7e8fe20538fcc71b1dc0fbb913029959.png")
23 if image_raw.size()[0] == 1:
24 image_raw = torch.cat([image_raw]*3, 0)
25 if image_raw.size()[0] == 4:
26 image_raw = image_raw[:3]
27 edit_image_tensor: torch.Tensor = image_transform(image_raw)
28 edit_image_tensor = edit_image_tensor.unsqueeze(0)
29
30 outputs = model(pixel_values=edit_image_tensor)
31 logits = F.sigmoid(outputs.logits)[0]
32 ind = logits.argmax().item()
33 print(model.config.id2label[ind])
34
35 cha_names = [model.config.id2label[i] for i in range(146)]
36 cha_probs = logits.numpy()
37 names_probs = list(zip(cha_names, cha_probs))
38 names_probs = sorted(names_probs, key=lambda x: x[1], reverse=True)
39
40 print(names_probs)
41
42 top_k = 10
43 names_show = []
44 probs_show = []
45 for i in range(top_k):
46 names_show.append(names_probs[i][0])
47 probs_show.append(names_probs[i][1])
48
49 plt.rcParams['font.sans-serif'] = ['SimHei']
50 plt.figure(figsize=(12, 8))
51 plt.bar(names_show, probs_show)
52 plt.show()