sandy45/ChestViT-Explainable-XRay-AI
0
1"""2models/vit_model.py3-------------------4ViT-Base-16 fine-tuned for multi-label chest X-ray classification.5 6Architecture:7 - Backbone: google/vit-base-patch16-224-in21k (pre-trained on ImageNet-21k)8 - Head: Linear(768 → 14) — replaces the [CLS] token classifier9 - Activation: Sigmoid (multi-label, not softmax)10 - Attention: output_attentions=True exposes all 12 transformer layer11 attention weight tensors for Attention Rollout visualization12 13RTX 3050 Optimizations:14 - gradient_checkpointing: reduces VRAM by ~30% by recomputing activations15 during backprop instead of caching them16 - Mixed precision (fp16) used during training (handled in train.py)17 18Reference:19 Dosovitskiy et al., "An Image is Worth 16x16 Words", ICLR 202120 https://arxiv.org/abs/2010.1192921"""22 23from pathlib import Path24from typing import Optional, Tuple, List25 26import torch27import torch.nn as nn28from transformers import ViTModel, ViTConfig29 30 31class ChestViT(nn.Module):32 """33 ViT-Base-16 with a multi-label classification head.34 35 The model exposes two outputs:36 1. logits — shape (B, 14), raw pre-sigmoid scores37 2. attentions — list of 12 tensors, each (B, num_heads, seq_len, seq_len)38 seq_len = 197 = 1 (CLS) + 196 (14×14 patches)39 Only returned when output_attentions=True (set during initialization).40 41 Usage:42 model = ChestViT(num_classes=14)43 logits, attentions = model(pixel_values, output_attentions=True)44 probs = torch.sigmoid(logits)45 """46 47 def __init__(48 self,49 num_classes: int = 14,50 pretrained_name: str = "google/vit-base-patch16-224-in21k",51 dropout: float = 0.1,52 gradient_checkpointing: bool = True,53 ):54 super().__init__()55 self.num_classes = num_classes56 57 # ── Load pre-trained ViT backbone ─────────────────────────────────────58 print(f" Loading ViT backbone: {pretrained_name}")59 self.vit = ViTModel.from_pretrained(60 pretrained_name,61 add_pooling_layer=False, # We extract [CLS] ourselves62 output_attentions=True, # Always expose attention weights63 )64 65 # ── RTX 3050: gradient checkpointing ─────────────────────────────────66 if gradient_checkpointing:67 self.vit.gradient_checkpointing_enable()68 print(" Gradient checkpointing: ENABLED (saves ~30% VRAM)")69 70 # ── Multi-label classification head ───────────────────────────────────71 hidden_size = self.vit.config.hidden_size # 768 for ViT-Base72 self.dropout = nn.Dropout(dropout)73 self.classifier = nn.Linear(hidden_size, num_classes)74 75 # Initialize head with small weights (better multi-label convergence)76 nn.init.xavier_uniform_(self.classifier.weight)77 nn.init.zeros_(self.classifier.bias)78 79 print(f" Classification head: Linear({hidden_size} → {num_classes})")80 print(f" Total parameters: {self.count_parameters():,.0f}")81 print(f" Trainable parameters: {self.count_parameters(trainable_only=True):,.0f}")82 83 def forward(84 self,85 pixel_values: torch.Tensor,86 output_attentions: bool = True,87 ) -> Tuple[torch.Tensor, Optional[List[torch.Tensor]]]:88 """89 Forward pass.90 91 Args:92 pixel_values: (B, 3, 224, 224) normalized image tensor.93 output_attentions: Return attention weights for rollout visualization.94 95 Returns:96 logits: (B, 14) raw pre-sigmoid classification scores.97 attentions: List of 12 tensors (B, 12, 197, 197), or None.98 """99 outputs = self.vit(100 pixel_values=pixel_values,101 output_attentions=output_attentions,102 )103 104 # [CLS] token representation — shape: (B, 768)105 cls_output = outputs.last_hidden_state[:, 0, :]106 cls_output = self.dropout(cls_output)107 108 # Multi-label logits — shape: (B, 14)109 logits = self.classifier(cls_output)110 111 # Attention weights: tuple of 12 tensors, each (B, 12, 197, 197)112 attentions = outputs.attentions if output_attentions else None113 114 return logits, attentions115 116 def count_parameters(self, trainable_only: bool = False) -> int:117 if trainable_only:118 return sum(p.numel() for p in self.parameters() if p.requires_grad)119 return sum(p.numel() for p in self.parameters())120 121 def get_patch_size(self) -> int:122 """Return the patch size (16 for ViT-Base-16)."""123 return self.vit.config.patch_size # 16124 125 def get_num_patches(self) -> int:126 """Return number of patches per side (14 for 224/16)."""127 img_size = self.vit.config.image_size # 224128 patch_size = self.vit.config.patch_size # 16129 return img_size // patch_size # 14130 131 132def load_checkpoint(133 checkpoint_path: str | Path,134 device: torch.device,135 num_classes: int = 14,136) -> ChestViT:137 """138 Load a saved ChestViT checkpoint.139 140 Args:141 checkpoint_path: Path to .pt or .pth checkpoint file.142 device: Target device (cuda / cpu).143 num_classes: Must match the saved model.144 145 Returns:146 Loaded ChestViT model in eval mode.147 """148 checkpoint_path = Path(checkpoint_path)149 print(f" Loading checkpoint: {checkpoint_path}")150 151 checkpoint = torch.load(checkpoint_path, map_location=device)152 model = ChestViT(num_classes=num_classes)153 model.load_state_dict(checkpoint["model_state_dict"])154 model.to(device)155 model.eval()156 print(f" Checkpoint loaded from epoch {checkpoint.get('epoch', '?')} "157 f"(val_auc={checkpoint.get('val_auc', '?'):.4f})")158 return model159 160 161def save_checkpoint(162 model: ChestViT,163 optimizer: torch.optim.Optimizer,164 epoch: int,165 val_auc: float,166 save_path: str | Path,167) -> None:168 """Save a training checkpoint."""169 save_path = Path(save_path)170 save_path.parent.mkdir(parents=True, exist_ok=True)171 torch.save({172 "epoch": epoch,173 "val_auc": val_auc,174 "model_state_dict": model.state_dict(),175 "optimizer_state_dict": optimizer.state_dict(),176 }, save_path)177 print(f" Checkpoint saved → {save_path} (val_auc={val_auc:.4f})")178 