Views
No views yet
timm/efficientnet_b3 backbone with a custom multi-task head.severity: 5 classes (0: No DR, 1: Mild, 2: Moderate, 3: Severe, 4: Proliferative).lesions: 5 classes (multi-label for various lesion types).regions: 5 classes (multi-label for affected anatomical regions).severity task by setting the loss weights for auxiliary tasks to zero. The auxiliary heads can still produce outputs for interpretability.1# Install required libraries
2pip install torch torchvision timm albumentations huggingface-hub numpy pillow opencv-python1import torch
2import torch.nn as nn
3import torch.nn.functional as F
4import timm
5from PIL import Image
6import numpy as np
7import albumentations as A
8from albumentations.pytorch import ToTensorV2
9from huggingface_hub import hf_hub_download
10
11# Define the model architecture
12class MultiTaskDRModel(nn.Module):
13 def __init__(self, model_name='efficientnet_b3', num_classes=5,
14 num_lesion_types=5, num_regions=5, pretrained=False):
15 super(MultiTaskDRModel, self).__init__()
16 self.backbone = timm.create_model(model_name, pretrained=pretrained, num_classes=0)
17 self.feature_dim = self.backbone.num_features
18
19 self.attention = nn.Sequential(
20 nn.AdaptiveAvgPool2d(1), nn.Flatten(),
21 nn.Linear(self.feature_dim, self.feature_dim // 8), nn.ReLU(inplace=True),
22 nn.Linear(self.feature_dim // 8, self.feature_dim), nn.Sigmoid()
23 )
24
25 self.feature_norm = nn.BatchNorm1d(self.feature_dim)
26 self.dropout = nn.Dropout(0.4)
27
28 self.severity_classifier = nn.Sequential(
29 nn.Linear(self.feature_dim, self.feature_dim // 2), nn.ReLU(inplace=True),
30 nn.Dropout(0.2), nn.Linear(self.feature_dim // 2, num_classes)
31 )
32
33 self.lesion_detector = nn.Sequential(
34 nn.Linear(self.feature_dim, self.feature_dim // 4), nn.ReLU(inplace=True),
35 nn.Dropout(0.2), nn.Linear(self.feature_dim // 4, num_lesion_types)
36 )
37
38 self.region_predictor = nn.Sequential(
39 nn.Linear(self.feature_dim, self.feature_dim // 4), nn.ReLU(inplace=True),
40 nn.Dropout(0.2), nn.Linear(self.feature_dim // 4, num_regions)
41 )
42
43 def forward(self, x):
44 features = self.backbone.forward_features(x)
45 pooled_features = F.adaptive_avg_pool2d(features, 1).flatten(1)
46 attention_weights = self.attention(pooled_features.unsqueeze(-1).unsqueeze(-1))
47 features = pooled_features * attention_weights
48 features = self.feature_norm(features)
49 features = self.dropout(features)
50
51 severity_logits = self.severity_classifier(features)
52 lesion_logits = self.lesion_detector(features)
53 region_logits = self.region_predictor(features)
54
55 return {
56 'severity': severity_logits,
57 'lesions': lesion_logits,
58 'regions': region_logits,
59 'features': features
60 }
61
62# Load the model
63device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
64model = MultiTaskDRModel()
65
66# Download and load the checkpoint
67model_path = hf_hub_download(
68 repo_id="dheeren-tejani/DiabeticRetinpathyClassifier",
69 filename="best_model_v2.pth"
70)
71checkpoint = torch.load(model_path, map_location=device, weights_only=False)
72model.load_state_dict(checkpoint['model_state_dict'])
73model.to(device)
74model.eval()
75
76print("Model loaded successfully!")
77
78# Preprocessing function
79def preprocess_image(image_path):
80 transforms = A.Compose([
81 A.Resize(512, 512),
82 A.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),
83 ToTensorV2(),
84 ])
85 image = np.array(Image.open(image_path).convert("RGB"))
86 image_tensor = transforms(image=image)['image'].unsqueeze(0)
87 return image_tensor
88
89# Example inference
90def predict_dr_severity(image_path):
91 image_tensor = preprocess_image(image_path).to(device)
92
93 with torch.no_grad():
94 outputs = model(image_tensor)
95
96 # Get severity prediction
97 severity_probs = torch.softmax(outputs['severity'], dim=1)
98 predicted_class = torch.argmax(severity_probs, dim=1).item()
99 confidence = severity_probs[0, predicted_class].item()
100
101 severity_labels = {
102 0: "No DR",
103 1: "Mild DR",
104 2: "Moderate DR",
105 3: "Severe DR",
106 4: "Proliferative DR"
107 }
108
109 return {
110 'predicted_severity': severity_labels[predicted_class],
111 'confidence': confidence,
112 'all_probabilities': severity_probs[0].cpu().numpy()
113 }
114
115# Example usage
116# result = predict_dr_severity("path/to/your/fundus_image.jpg")
117# print(f"Predicted: {result['predicted_severity']} (Confidence: {result['confidence']:.3f})")| Metric | Score |
|---|---|
| Quadratic Weighted Kappa (QWK) | 0.796 |
| Accuracy | 65.0% |
| F1-Score (Weighted) | 66.3% |
| F1-Score (Macro) | 53.5% |
1@misc{dheerentejani2025dr,
2 author = {Dheeren Tejani},
3 title = {Diabetic Retinopathy Grading Model V2},
4 year = {2025},
5 publisher = {Hugging Face},
6 journal = {Hugging Face Model Hub},
7 howpublished = {\url{https://huggingface.co/dheeren-tejani/DiabeticRetinpathyClassifier}},
8}