Ishaank18's picture
Upload via upload_to_hf.py
d0518d9 verified
Raw History Blame Contribute Delete
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)
@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()))