This is a Vision Transformer (ViT) model designed for detecting pneumonia in chest X-rays. The model takes a chest X-ray image as input, processes it through the ViT architecture, and outputs the predicted class probabilities.
Classes include:
1from PIL import Image
2import torch
3from torchvision import transforms
4from ViT_model import ViT # replace with the actual class name from ViT_model.py
5
6device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
7vit_model = ViT()
8vit_model.load_state_dict(torch.load("vit_with_generated.pth", map_location=device))
9vit_model = vit_model.to(device)
10vit_model.eval()
11
12def vit_inference(image: Image.Image) -> dict:
13
14
15 class_names = ["Normal", "Bacterial Pneumonia", "Viral Pneumonia", "COVID-19"]
16
17
18 # load and preprocess image
19 image = image.convert("L") # Convert to grayscale
20
21 # define preprocessing transform
22 transform = transforms.Compose([
23 transforms.Resize((64, 64)),
24 transforms.ToTensor(),
25 transforms.Normalize(mean=[0.5], std=[0.5])
26 ])
27
28
29 image_tensor = transform(image).unsqueeze(0) # Add batch dimension
30 image_tensor = image_tensor.to(device)
31
32 #inference
33 with torch.no_grad():
34 outputs, attn_weights_all = vit_model(image_tensor)
35 print('vit outputs:', outputs)
36 if torch.isnan(outputs).any(): #debugging
37 print('Warning: NaNs in outputs')
38
39 probabilities = torch.softmax(outputs, dim=1)
40 print('probabilities:', probabilities)
41
42 #probabilities for all classes and transfer to cpu and to numpy array
43 probs = probabilities[0].cpu().numpy()
44 print('probs numpy:', probs)
45
46 #list of class_name, probability and sort by probability descending
47 class_probs = [(class_names[i], float(probs[i])) for i in range(len(class_names))]
48 class_probs.sort(key=lambda x: x[1], reverse=True)
49
50 #top predicted class and confidence
51 predicted_class = class_probs[0][0]
52 confidence_score = class_probs[0][1]
53 print(predicted_class, confidence_score)
54
55
56image_path = "t_viral_test.jpg" # replace with your test image path
57image = Image.open(image_path)
58vit_inference(image)