Download classification_segmentation/model.py from Ishaank18/aifh: direct link, hf CLI and curl.
- Browser
- Download file 14.6 kB
-
https://huggingface.co/Ishaank18/aifh/resolve/main/classification_segmentation/model.py
- Command line
-
hf download hf://Ishaank18/aifh/classification_segmentation/model.py
-
curl -L -o model.py https://huggingface.co/Ishaank18/aifh/resolve/main/classification_segmentation/model.py
14.6 kB
| """ | |
| model.py -- Part A model definitions. | |
| Model 1 (classification): XRVDenseNet | |
| A DenseNet-121 pretrained on large public chest-X-ray corpora, loaded through | |
| TorchXRayVision. The convolutional backbone is FROZEN (requires_grad=False, | |
| eval mode so BatchNorm statistics stay fixed); only a fresh 4-way head is | |
| trained. This is a linear-probe / feature-extraction setup, exactly as the | |
| assignment specifies. | |
| Model 2 (classification): StudentDenseNet | |
| torchvision densenet121 trained by us, either from scratch or initialised | |
| from non-X-ray ImageNet weights (the choice is recorded in the checkpoint | |
| and printed at train time so the report can state it unambiguously). | |
| Segmentation: UNet | |
| Standard Ronneberger-style encoder/decoder with BatchNorm, configurable | |
| base width, bilinear or transposed-conv upsampling. | |
| """ | |
| from __future__ import annotations | |
| from typing import List, Optional | |
| import torch | |
| import torch.nn as nn | |
| import torch.nn.functional as F | |
| from common import NUM_CLASSES | |
| # --------------------------------------------------------------------------- | |
| # Model 1 -- TorchXRayVision pretrained DenseNet (frozen backbone) | |
| # --------------------------------------------------------------------------- | |
| XRV_DEFAULT_WEIGHTS = "densenet121-res224-all" | |
| def import_torchxrayvision(): | |
| """Import torchxrayvision safely from inside a file called `model.py`. | |
| torchxrayvision vendors a third-party baseline whose code contains the | |
| ABSOLUTE import `from model.utils import get_norm`. It works because the | |
| package inserts its own directory on sys.path, but Python resolves | |
| `sys.modules['model']` first -- and that entry is THIS file, because the | |
| assignment requires the script to be named model.py. The result is a | |
| confusing "No module named 'model.utils'; 'model' is not a package". | |
| We therefore hide our own `model` entry for the duration of the import and | |
| restore it afterwards. Nothing else in the process is affected. | |
| """ | |
| import sys | |
| shadowed = {k: sys.modules.pop(k) for k in list(sys.modules) | |
| if k == "model" or k.startswith("model.")} | |
| sys_path_saved = list(sys.path) | |
| try: | |
| # make sure our own directory cannot win the `model` lookup either | |
| here = str(__import__("pathlib").Path(__file__).resolve().parent) | |
| sys.path = [p for p in sys.path if p not in ("", ".", here)] | |
| import torchxrayvision as xrv | |
| return xrv | |
| finally: | |
| sys.path = sys_path_saved | |
| for k in [k for k in list(sys.modules) if k == "model" or k.startswith("model.")]: | |
| del sys.modules[k] | |
| sys.modules.update(shadowed) | |
| class XRVDenseNet(nn.Module): | |
| """Frozen chest-X-ray-pretrained DenseNet-121 + trainable 4-class head. | |
| Input : (N, 1, 224, 224) float tensor scaled to [-1024, 1024] | |
| (see common.to_model_tensor(..., model_kind="xrv")) | |
| Output: (N, 4) logits | |
| """ | |
| def __init__(self, weights: str = XRV_DEFAULT_WEIGHTS, | |
| num_classes: int = NUM_CLASSES, dropout: float = 0.2, | |
| freeze: bool = True, hidden: int = 0): | |
| super().__init__() | |
| try: | |
| xrv = import_torchxrayvision() | |
| except ImportError as e: # pragma: no cover | |
| raise ImportError( | |
| "torchxrayvision is required for Model 1. Install with:\n" | |
| " pip install torchxrayvision" | |
| ) from e | |
| self.backbone = xrv.models.DenseNet(weights=weights) | |
| self.weights_name = weights | |
| self.feat_dim = 1024 # densenet121 pooled feature width | |
| self.frozen = freeze | |
| if freeze: | |
| for p in self.backbone.parameters(): | |
| p.requires_grad = False | |
| layers: List[nn.Module] = [nn.Flatten(), nn.BatchNorm1d(self.feat_dim)] | |
| if hidden > 0: | |
| layers += [nn.Linear(self.feat_dim, hidden), nn.ReLU(inplace=True), | |
| nn.Dropout(dropout), nn.Linear(hidden, num_classes)] | |
| else: | |
| layers += [nn.Dropout(dropout), nn.Linear(self.feat_dim, num_classes)] | |
| self.head = nn.Sequential(*layers) | |
| # -- keep the frozen backbone in eval mode even when the module trains ---- | |
| def train(self, mode: bool = True): | |
| super().train(mode) | |
| if self.frozen: | |
| self.backbone.eval() | |
| return self | |
| def feature_maps(self, x: torch.Tensor) -> torch.Tensor: | |
| """Last conv feature map (N, 1024, 7, 7) -- the Grad-CAM target.""" | |
| return self.backbone.features(x) | |
| def pooled_features(self, x: torch.Tensor) -> torch.Tensor: | |
| fmap = self.feature_maps(x) | |
| out = F.relu(fmap, inplace=False) | |
| out = F.adaptive_avg_pool2d(out, (1, 1)) | |
| return torch.flatten(out, 1) | |
| def forward(self, x: torch.Tensor) -> torch.Tensor: | |
| return self.head(self.pooled_features(x)) | |
| def gradcam_target_layer(self) -> nn.Module: | |
| """Deepest conv block; hooking it gives 7x7 CAMs at 224 input.""" | |
| return self.backbone.features.denseblock4 | |
| def trainable_parameter_report(self) -> dict: | |
| tot = sum(p.numel() for p in self.parameters()) | |
| tr = sum(p.numel() for p in self.parameters() if p.requires_grad) | |
| return {"total_params": tot, "trainable_params": tr, | |
| "frozen_params": tot - tr, "backbone_weights": self.weights_name} | |
| # --------------------------------------------------------------------------- | |
| # Model 2 -- student-trained DenseNet-121 | |
| # --------------------------------------------------------------------------- | |
| class StudentDenseNet(nn.Module): | |
| """torchvision DenseNet-121 that WE train. | |
| init_from : "imagenet" -> non-X-ray ImageNet-1k weights (transfer learning) | |
| "scratch" -> random initialisation | |
| in_channels: 3 (grayscale replicated, default) or 1 (grayscale ablation). | |
| """ | |
| def __init__(self, num_classes: int = NUM_CLASSES, init_from: str = "imagenet", | |
| in_channels: int = 3, dropout: float = 0.2): | |
| super().__init__() | |
| from torchvision import models | |
| assert init_from in {"imagenet", "scratch"} | |
| self.init_from = init_from | |
| self.in_channels = in_channels | |
| if init_from == "imagenet": | |
| try: | |
| from torchvision.models import DenseNet121_Weights | |
| net = models.densenet121(weights=DenseNet121_Weights.IMAGENET1K_V1) | |
| except Exception: | |
| net = models.densenet121(pretrained=True) # older torchvision | |
| else: | |
| net = models.densenet121(weights=None) | |
| if in_channels != 3: | |
| old = net.features.conv0 | |
| new = nn.Conv2d(in_channels, old.out_channels, kernel_size=old.kernel_size, | |
| stride=old.stride, padding=old.padding, bias=False) | |
| with torch.no_grad(): | |
| # average the RGB filters -> a sensible 1-channel initialisation | |
| new.weight.copy_(old.weight.mean(dim=1, keepdim=True).repeat(1, in_channels, 1, 1)) | |
| net.features.conv0 = new | |
| self.features = net.features | |
| self.feat_dim = net.classifier.in_features # 1024 | |
| self.head = nn.Sequential(nn.Dropout(dropout), nn.Linear(self.feat_dim, num_classes)) | |
| def feature_maps(self, x: torch.Tensor) -> torch.Tensor: | |
| return self.features(x) | |
| def pooled_features(self, x: torch.Tensor) -> torch.Tensor: | |
| out = F.relu(self.feature_maps(x), inplace=False) | |
| out = F.adaptive_avg_pool2d(out, (1, 1)) | |
| return torch.flatten(out, 1) | |
| def forward(self, x: torch.Tensor) -> torch.Tensor: | |
| return self.head(self.pooled_features(x)) | |
| def gradcam_target_layer(self) -> nn.Module: | |
| return self.features.denseblock4 | |
| def trainable_parameter_report(self) -> dict: | |
| tot = sum(p.numel() for p in self.parameters()) | |
| tr = sum(p.numel() for p in self.parameters() if p.requires_grad) | |
| return {"total_params": tot, "trainable_params": tr, | |
| "frozen_params": tot - tr, "init_from": self.init_from} | |
| # --------------------------------------------------------------------------- | |
| # Segmentation -- U-Net | |
| # --------------------------------------------------------------------------- | |
| class DoubleConv(nn.Module): | |
| def __init__(self, cin: int, cout: int, mid: Optional[int] = None): | |
| super().__init__() | |
| mid = mid or cout | |
| self.block = nn.Sequential( | |
| nn.Conv2d(cin, mid, 3, padding=1, bias=False), nn.BatchNorm2d(mid), | |
| nn.ReLU(inplace=True), | |
| nn.Conv2d(mid, cout, 3, padding=1, bias=False), nn.BatchNorm2d(cout), | |
| nn.ReLU(inplace=True)) | |
| def forward(self, x): | |
| return self.block(x) | |
| class UNet(nn.Module): | |
| """U-Net (Ronneberger et al., MICCAI 2015) with BatchNorm and padded convs | |
| so the output resolution equals the input resolution.""" | |
| def __init__(self, in_channels: int = 1, num_classes: int = 1, | |
| base: int = 32, bilinear: bool = True): | |
| super().__init__() | |
| b = base | |
| self.inc = DoubleConv(in_channels, b) | |
| self.down1 = nn.Sequential(nn.MaxPool2d(2), DoubleConv(b, b * 2)) | |
| self.down2 = nn.Sequential(nn.MaxPool2d(2), DoubleConv(b * 2, b * 4)) | |
| self.down3 = nn.Sequential(nn.MaxPool2d(2), DoubleConv(b * 4, b * 8)) | |
| factor = 2 if bilinear else 1 | |
| self.down4 = nn.Sequential(nn.MaxPool2d(2), DoubleConv(b * 8, b * 16 // factor)) | |
| self.bilinear = bilinear | |
| if bilinear: | |
| self.up = nn.Upsample(scale_factor=2, mode="bilinear", align_corners=True) | |
| self.conv1 = DoubleConv(b * 16 // factor + b * 8, b * 8 // factor) | |
| self.conv2 = DoubleConv(b * 8 // factor + b * 4, b * 4 // factor) | |
| self.conv3 = DoubleConv(b * 4 // factor + b * 2, b * 2 // factor) | |
| self.conv4 = DoubleConv(b * 2 // factor + b, b) | |
| else: | |
| self.up1 = nn.ConvTranspose2d(b * 16, b * 8, 2, stride=2) | |
| self.up2 = nn.ConvTranspose2d(b * 8, b * 4, 2, stride=2) | |
| self.up3 = nn.ConvTranspose2d(b * 4, b * 2, 2, stride=2) | |
| self.up4 = nn.ConvTranspose2d(b * 2, b, 2, stride=2) | |
| self.conv1 = DoubleConv(b * 16, b * 8) | |
| self.conv2 = DoubleConv(b * 8, b * 4) | |
| self.conv3 = DoubleConv(b * 4, b * 2) | |
| self.conv4 = DoubleConv(b * 2, b) | |
| self.outc = nn.Conv2d(b, num_classes, 1) | |
| def _cat(x, skip): | |
| dy = skip.size(-2) - x.size(-2) | |
| dx = skip.size(-1) - x.size(-1) | |
| if dy or dx: | |
| x = F.pad(x, [dx // 2, dx - dx // 2, dy // 2, dy - dy // 2]) | |
| return torch.cat([skip, x], dim=1) | |
| def forward(self, x): | |
| x1 = self.inc(x) | |
| x2 = self.down1(x1) | |
| x3 = self.down2(x2) | |
| x4 = self.down3(x3) | |
| x5 = self.down4(x4) | |
| if self.bilinear: | |
| y = self.conv1(self._cat(self.up(x5), x4)) | |
| y = self.conv2(self._cat(self.up(y), x3)) | |
| y = self.conv3(self._cat(self.up(y), x2)) | |
| y = self.conv4(self._cat(self.up(y), x1)) | |
| else: | |
| y = self.conv1(self._cat(self.up1(x5), x4)) | |
| y = self.conv2(self._cat(self.up2(y), x3)) | |
| y = self.conv3(self._cat(self.up3(y), x2)) | |
| y = self.conv4(self._cat(self.up4(y), x1)) | |
| return self.outc(y) # raw logits, (N, 1, H, W) | |
| # --------------------------------------------------------------------------- | |
| # Losses | |
| # --------------------------------------------------------------------------- | |
| class DiceBCELoss(nn.Module): | |
| """BCE-with-logits + soft Dice. BCE gives stable pixel-wise gradients, | |
| Dice directly optimises the overlap metric we report.""" | |
| def __init__(self, bce_weight: float = 0.5, smooth: float = 1.0, | |
| pos_weight: Optional[torch.Tensor] = None): | |
| super().__init__() | |
| self.bce_weight = bce_weight | |
| self.smooth = smooth | |
| self.bce = nn.BCEWithLogitsLoss(pos_weight=pos_weight) | |
| def forward(self, logits: torch.Tensor, target: torch.Tensor) -> torch.Tensor: | |
| bce = self.bce(logits, target) | |
| p = torch.sigmoid(logits) | |
| num = 2 * (p * target).sum(dim=(1, 2, 3)) + self.smooth | |
| den = p.sum(dim=(1, 2, 3)) + target.sum(dim=(1, 2, 3)) + self.smooth | |
| dice = 1 - (num / den).mean() | |
| return self.bce_weight * bce + (1 - self.bce_weight) * dice | |
| class FocalLoss(nn.Module): | |
| """Multi-class focal loss, optional for the classification imbalance study.""" | |
| def __init__(self, gamma: float = 2.0, weight: Optional[torch.Tensor] = None): | |
| super().__init__() | |
| self.gamma = gamma | |
| self.weight = weight | |
| def forward(self, logits, target): | |
| ce = F.cross_entropy(logits, target, weight=self.weight, reduction="none") | |
| pt = torch.exp(-ce) | |
| return ((1 - pt) ** self.gamma * ce).mean() | |
| # --------------------------------------------------------------------------- | |
| # Factory | |
| # --------------------------------------------------------------------------- | |
| def build_model(name: str, **kw) -> nn.Module: | |
| name = name.lower() | |
| if name == "xrv": | |
| return XRVDenseNet(weights=kw.get("xrv_weights", XRV_DEFAULT_WEIGHTS), | |
| num_classes=kw.get("num_classes", NUM_CLASSES), | |
| freeze=kw.get("freeze", True), | |
| hidden=kw.get("hidden", 0)) | |
| if name == "student": | |
| return StudentDenseNet(num_classes=kw.get("num_classes", NUM_CLASSES), | |
| init_from=kw.get("init_from", "imagenet"), | |
| in_channels=kw.get("in_channels", 3)) | |
| if name == "unet": | |
| return UNet(in_channels=kw.get("in_channels", 1), num_classes=1, | |
| base=kw.get("base", 32), bilinear=kw.get("bilinear", True)) | |
| raise ValueError(f"unknown model '{name}' (expected xrv | student | unet)") | |
| def model_kind_for(name: str) -> str: | |
| """Which intensity convention the model input needs (see common.to_model_tensor).""" | |
| return "xrv" if name.lower() == "xrv" else "student" | |
| if __name__ == "__main__": | |
| # quick shape self-test (no pretrained download for the student/unet path) | |
| m = StudentDenseNet(init_from="scratch") | |
| x = torch.randn(2, 3, 224, 224) | |
| print("student logits", m(x).shape, m.trainable_parameter_report()) | |
| u = UNet(in_channels=1, base=16) | |
| print("unet out", u(torch.randn(2, 1, 224, 224)).shape, | |
| "params", sum(p.numel() for p in u.parameters())) | |