covid-severity / step4_model.py
Ranjith445's picture
Initial deployment - COVID severity AI
0366462
Raw
History Blame Contribute Delete
4.72 kB
"""
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)