Spaces:
Runtime error
Runtime error
| """ | |
| 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) |