""" step4_model.py -------------- Multi-task CNN model: - Backbone : ResNet18 (pretrained ImageNet, lightweight for CPU) - Shared : FC(512 -> 256) + ReLU + Dropout(0.4) - Head 1 : Biomarker regression -> FC(256 -> 4) - Head 2 : Severity classification -> FC(256 -> 5) Usage: from step4_model import MultiTaskCT, MultiTaskLoss """ import torch import torch.nn as nn from torchvision import models class MultiTaskCT(nn.Module): def __init__(self, num_biomarkers=4, num_severity_classes=5, dropout=0.4): super(MultiTaskCT, self).__init__() # ── Backbone: ResNet18 ──────────────────────────────────────────────── backbone = models.resnet18(weights=models.ResNet18_Weights.IMAGENET1K_V1) # Remove final classification layer — keep feature extractor only self.backbone = nn.Sequential(*list(backbone.children())[:-1]) backbone_out_dim = 512 # ── Shared layer ────────────────────────────────────────────────────── self.shared = nn.Sequential( nn.Linear(backbone_out_dim, 256), nn.ReLU(), nn.Dropout(dropout), ) # ── Head 1: Biomarker regression (4 continuous values) ──────────────── self.biomarker_head = nn.Sequential( nn.Linear(256, 128), nn.ReLU(), nn.Linear(128, num_biomarkers), ) # ── Head 2: Severity classification (5 classes: 0-4) ───────────────── self.severity_head = nn.Sequential( nn.Linear(256, 128), nn.ReLU(), nn.Linear(128, num_severity_classes), ) def forward(self, x): # x: [B, 3, 224, 224] features = self.backbone(x) # [B, 512, 1, 1] features = features.view(features.size(0), -1) # [B, 512] shared = self.shared(features) # [B, 256] biomarkers = self.biomarker_head(shared) # [B, 4] severity = self.severity_head(shared) # [B, 5] logits return biomarkers, severity class MultiTaskLoss(nn.Module): """ Combined loss: Total = w1 * MSE(biomarkers) + w2 * CrossEntropy(severity) Uses learnable uncertainty weights (Kendall et al. 2018) to automatically balance the two tasks during training. """ def __init__(self): super(MultiTaskLoss, self).__init__() self.log_var_bm = nn.Parameter(torch.zeros(1)) # biomarker weight self.log_var_sev = nn.Parameter(torch.zeros(1)) # severity weight self.mse = nn.MSELoss() self.ce = nn.CrossEntropyLoss() def forward(self, pred_bm, true_bm, pred_sev, true_sev): loss_bm = self.mse(pred_bm, true_bm) loss_sev = self.ce(pred_sev, true_sev) precision_bm = torch.exp(-self.log_var_bm) precision_sev = torch.exp(-self.log_var_sev) total = (precision_bm * loss_bm + self.log_var_bm + precision_sev * loss_sev + self.log_var_sev) return total, loss_bm.item(), loss_sev.item() # ── Quick test ──────────────────────────────────────────────────────────────── if __name__ == "__main__": print("=" * 55) print(" Model Architecture Test — Step 4") print("=" * 55) model = MultiTaskCT() loss_fn = MultiTaskLoss() # Count parameters total_params = sum(p.numel() for p in model.parameters()) train_params = sum(p.numel() for p in model.parameters() if p.requires_grad) print(f" Total parameters : {total_params:,}") print(f" Trainable parameters: {train_params:,}") # Dummy forward pass dummy_img = torch.randn(4, 3, 224, 224) # batch of 4 dummy_bm = torch.randn(4, 4) dummy_sev = torch.randint(0, 5, (4,)) bm_pred, sev_pred = model(dummy_img) total, l_bm, l_sev = loss_fn(bm_pred, dummy_bm, sev_pred, dummy_sev) print(f"\n Forward pass:") print(f" Input shape : {dummy_img.shape}") print(f" Biomarker output : {bm_pred.shape}") print(f" Severity output : {sev_pred.shape}") print(f"\n Loss:") print(f" Biomarker loss (MSE) : {l_bm:.4f}") print(f" Severity loss (CE) : {l_sev:.4f}") print(f" Total loss : {total.item():.4f}") print("\n Model OK") print("=" * 55)