Team Ai
Apppublic

sandy45/ChestViT-Explainable-XRay-AI

sourceHugging Faceupdated 4mo agoView on Hugging Face
0likes
vit_model.py178 linesDownload Raw Back to models
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