""" 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) @staticmethod 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()))