{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[],"dockerImageVersionId":28755,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"\"\"\"\nvimdr_model_kaggle.py\n=====================\nVimDR+  ·  Kaggle T4-optimised model\n\"\"\"\n\nimport math, functools\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\n\n# ──────────────────────────────────────────────────────────────\n# 0.  Compile helper (PyTorch ≥ 2.0 only, silent fallback)\n# ──────────────────────────────────────────────────────────────\ndef _maybe_compile(module: nn.Module) -> nn.Module:\n    try:\n        return torch.compile(module, mode='reduce-overhead', fullgraph=False)\n    except Exception:\n        return module\n\n\n# ──────────────────────────────────────────────────────────────\n# 1.  Parallel Selective Scan  (O-M5: removed redundant clone)\n# ──────────────────────────────────────────────────────────────\nclass SelectiveScan1D(nn.Module):\n\n    def __init__(self, dim: int, d_state: int = 16, expand: int = 2):\n        super().__init__()\n        self.d_inner = dim * expand\n        self.d_state = d_state\n\n        self.in_proj  = nn.Linear(dim, self.d_inner * 2, bias=False)\n        self.conv1d   = nn.Conv1d(self.d_inner, self.d_inner, kernel_size=4,\n                                  padding=3, groups=self.d_inner, bias=True)\n        self.dt_proj  = nn.Linear(self.d_inner, self.d_inner, bias=True)\n        self.B_proj   = nn.Linear(self.d_inner, d_state, bias=False)\n        self.C_proj   = nn.Linear(self.d_inner, d_state, bias=False)\n        self.A_log    = nn.Parameter(\n            torch.log(torch.arange(1, d_state + 1).float())\n                .unsqueeze(0).repeat(self.d_inner, 1))\n        self.D        = nn.Parameter(torch.ones(self.d_inner))\n        self.out_proj = nn.Linear(self.d_inner, dim, bias=False)\n        self.norm     = nn.LayerNorm(self.d_inner)\n        nn.init.constant_(self.dt_proj.bias, math.log(math.expm1(1.0)))\n\n    @staticmethod\n    def _parallel_scan(dA: torch.Tensor, dBu: torch.Tensor) -> torch.Tensor:\n        B, L, D, S = dA.shape\n        L_pad = 1 << (L - 1).bit_length()\n        pad   = L_pad - L\n        if pad:\n            dA  = F.pad(dA,  (0, 0, 0, 0, 0, pad))\n            dBu = F.pad(dBu, (0, 0, 0, 0, 0, pad))\n\n        # O-M5: work directly on dA/dBu tensors; left child read before write\n        a, b = dA, dBu                      # no unnecessary clone\n\n        stride = 1\n        while stride < L_pad:\n            left_a  = a[:, stride - 1::2 * stride].clone()   # need clone here only\n            left_b  = b[:, stride - 1::2 * stride].clone()\n            right_a = a[:, 2 * stride - 1::2 * stride]\n            right_b = b[:, 2 * stride - 1::2 * stride]\n            b[:, 2 * stride - 1::2 * stride] = right_a * left_b + right_b\n            a[:, 2 * stride - 1::2 * stride] = right_a * left_a\n            stride *= 2\n\n        b[:, L_pad - 1] = 0\n        stride = L_pad // 2\n        while stride >= 1:\n            tmp_a = a[:, stride - 1::2 * stride].clone()\n            tmp_b = b[:, stride - 1::2 * stride].clone()\n            b[:, stride - 1::2 * stride] = b[:, 2 * stride - 1::2 * stride]\n            b[:, 2 * stride - 1::2 * stride] = tmp_a * b[:, 2 * stride - 1::2 * stride] + tmp_b\n            stride //= 2\n\n        h = dA * b + dBu\n        return h[:, :L]\n\n    def forward(self, x: torch.Tensor) -> torch.Tensor:\n        B, L, _ = x.shape\n        xr, z   = self.in_proj(x).chunk(2, dim=-1)\n        xr = F.silu(self.conv1d(xr.transpose(1, 2))[:, :, :L].transpose(1, 2))\n\n        dt  = F.softplus(self.dt_proj(xr))\n        Bm  = self.B_proj(xr)\n        Cm  = self.C_proj(xr)\n        A   = -torch.exp(self.A_log.float())\n\n        dA  = torch.exp(torch.einsum('bld,ds->blds', dt, A))\n        dBu = torch.einsum('bld,bls->blds', dt, Bm)\n\n        h   = self._parallel_scan(dA, dBu)\n        y   = (h * Cm.unsqueeze(2)).sum(-1) + xr * self.D\n        return self.out_proj(self.norm(y) * F.silu(z))\n\n\n# ──────────────────────────────────────────────────────────────\n# 2.  Cross-Directional SSM  (O-M3: cached diagonal indices)\n# ──────────────────────────────────────────────────────────────\nclass CrossDirectionalSSM(nn.Module):\n\n    def __init__(self, dim: int, d_state: int = 16):\n        super().__init__()\n        assert dim % 4 == 0\n        c = dim // 4\n        self.ssm_h  = SelectiveScan1D(c, d_state)\n        self.ssm_v  = SelectiveScan1D(c, d_state)\n        self.ssm_d  = SelectiveScan1D(c, d_state)\n        self.ssm_a  = SelectiveScan1D(c, d_state)\n        self.fusion = nn.Sequential(nn.Linear(dim, dim, bias=False), nn.LayerNorm(dim))\n        self.norm   = nn.LayerNorm(dim)\n        self._idx_cache: dict = {}          # (H, W, anti) → (idx, inv_idx)\n\n    def _get_diag_index(self, H: int, W: int, anti: bool,\n                        device: torch.device):\n        key = (H, W, anti, device)\n        if key not in self._idx_cache:\n            indices = []\n            for s in range(H + W - 1):\n                for i in range(max(0, s - W + 1), min(H, s + 1)):\n                    j = s - i\n                    if anti:\n                        j = W - 1 - j\n                    indices.append(i * W + j)\n            idx     = torch.tensor(indices, dtype=torch.long, device=device)\n            inv_idx = torch.argsort(idx)\n            self._idx_cache[key] = (idx, inv_idx)\n        return self._idx_cache[key]\n\n    def _diag_scan(self, x: torch.Tensor, H: int, W: int,\n                   anti: bool = False) -> torch.Tensor:\n        idx, inv_idx = self._get_diag_index(H, W, anti, x.device)\n        ssm = self.ssm_a if anti else self.ssm_d\n        return ssm(x[:, idx, :])[:, inv_idx, :]\n\n    def forward(self, x: torch.Tensor, H: int, W: int) -> torch.Tensor:\n        x   = self.norm(x)\n        c   = x.shape[-1] // 4\n        x1, x2, x3, x4 = (x[..., i * c:(i + 1) * c] for i in range(4))\n\n        out_h = self.ssm_h(x1)\n\n        B, L, Dc = x2.shape\n        out_v = (self.ssm_v(x2.view(B, H, W, Dc).permute(0, 2, 1, 3).reshape(B, L, Dc))\n                   .view(B, W, H, Dc).permute(0, 2, 1, 3).reshape(B, L, Dc))\n\n        out_d = self._diag_scan(x3, H, W, anti=False)\n        out_a = self._diag_scan(x4, H, W, anti=True)\n\n        return x + self.fusion(torch.cat([out_h, out_v, out_d, out_a], dim=-1))\n\n\n# ──────────────────────────────────────────────────────────────\n# 3.  Vessel Morphological Block\n# ──────────────────────────────────────────────────────────────\nclass VesselMorphologicalBlock(nn.Module):\n    def __init__(self, dim: int, d_state: int = 16, dropout: float = 0.1):\n        super().__init__()\n        self.vessel_prior = nn.ModuleList([\n            nn.Sequential(\n                nn.Conv2d(dim, dim // 4, 1, bias=False),\n                nn.Conv2d(dim // 4, dim // 4, 3, padding=r, dilation=r,\n                          groups=dim // 4, bias=False),\n                nn.BatchNorm2d(dim // 4), nn.GELU(),\n            ) for r in [1, 2, 4, 8]\n        ])\n        self.prior_proj = nn.Conv2d(dim, dim, 1, bias=False)\n        self.prior_gate = nn.Sigmoid()\n        self.ssm   = CrossDirectionalSSM(dim, d_state)\n        self.norm  = nn.LayerNorm(dim)\n        self.norm2 = nn.LayerNorm(dim)\n        self.ffn   = nn.Sequential(\n            nn.Linear(dim, dim * 4, bias=False), nn.GELU(),\n            nn.Dropout(dropout), nn.Linear(dim * 4, dim, bias=False),\n        )\n\n    def forward(self, x: torch.Tensor, H: int, W: int) -> torch.Tensor:\n        B, L, D = x.shape\n        feat  = x.transpose(1, 2).view(B, D, H, W)\n        prior = self.prior_gate(\n            self.prior_proj(torch.cat([b(feat) for b in self.vessel_prior], dim=1))\n        ).flatten(2).transpose(1, 2)\n        x = x + self.ssm(self.norm(x), H, W) * (1.0 + prior)\n        return x + self.ffn(self.norm2(x))\n\n\n# ──────────────────────────────────────────────────────────────\n# 4.  Adaptive Pathology Scale Detector\n# ──────────────────────────────────────────────────────────────\nclass AdaptivePathologyScaleDetector(nn.Module):\n    def __init__(self, dim: int):\n        super().__init__()\n        hd = dim // 4\n        self.scale_convs = nn.ModuleList([\n            nn.Sequential(\n                nn.Conv2d(dim, hd, k, padding=k // 2, groups=hd, bias=False),\n                nn.BatchNorm2d(hd), nn.GELU(),\n            ) for k in [1, 3, 5, 7]\n        ])\n        self.fusion     = nn.Conv2d(dim, dim, 1, bias=False)\n        self.norm_bn    = nn.BatchNorm2d(dim)\n        self.scale_attn = nn.Sequential(\n            nn.AdaptiveAvgPool2d(1), nn.Flatten(),\n            nn.Linear(dim, dim // 4), nn.ReLU(),\n            nn.Linear(dim // 4, 4), nn.Softmax(dim=-1),\n        )\n\n    def forward(self, x: torch.Tensor, H: int, W: int) -> torch.Tensor:\n        B, L, D = x.shape\n        feat = x.transpose(1, 2).view(B, D, H, W)\n        outs = [conv(feat) for conv in self.scale_convs]\n        out  = self.norm_bn(self.fusion(torch.cat(outs, dim=1)))\n        return x + out.flatten(2).transpose(1, 2)\n\n\n# ──────────────────────────────────────────────────────────────\n# 5.  Ordinal Uncertainty Head\n# ──────────────────────────────────────────────────────────────\nclass OrdinalUncertaintyHead(nn.Module):\n    def __init__(self, in_dim: int, num_classes: int = 5, dropout: float = 0.3):\n        super().__init__()\n        self.num_classes = num_classes\n        self.proj = nn.Sequential(\n            nn.Linear(in_dim, in_dim // 2), nn.GELU(), nn.Dropout(dropout),\n            nn.Linear(in_dim // 2, in_dim // 4), nn.GELU(), nn.Dropout(dropout),\n        )\n        self.ordinal_fc  = nn.Linear(in_dim // 4, num_classes - 1)\n        self.reg_fc      = nn.Linear(in_dim // 4, 1)\n        self.temperature = nn.Parameter(torch.ones(1))\n\n    def forward(self, x: torch.Tensor) -> dict:\n        feat = self.proj(x)\n        return {\n            'grade_logits': self.ordinal_fc(feat) / (self.temperature.abs() + 1e-6),\n            'reg_score':    self.reg_fc(feat),\n            'features':     feat,\n        }\n\n    def predict_grade(self, logits: torch.Tensor) -> torch.Tensor:\n        return torch.clamp(\n            torch.round(torch.sigmoid(logits).sum(dim=-1)), 0, self.num_classes - 1\n        ).long()\n\n    @staticmethod\n    def logits_to_probs(logits: torch.Tensor) -> torch.Tensor:\n        \"\"\"Correct P(grade=k) from ordinal thresholds.  FIX 4 from v1.\"\"\"\n        p    = torch.sigmoid(logits)\n        p0   = 1 - p[:, :1]\n        pmid = p[:, :-1] - p[:, 1:]\n        pK   = p[:, -1:]\n        return torch.cat([p0, pmid, pK], dim=-1).clamp(min=1e-7)\n\n\n# ──────────────────────────────────────────────────────────────\n# 6.  VimDR+  (O-M1/O-M2: depth 8, embed_dim 192)\n# ──────────────────────────────────────────────────────────────\nclass VimDRPlus(nn.Module):\n    def __init__(\n        self,\n        embed_dim:   int   = 192,       # O-M2\n        depth:       int   = 8,         # O-M1\n        num_classes: int   = 5,\n        d_state:     int   = 16,        # O-M2\n        dropout:     float = 0.3,\n        use_compile: bool  = True,\n    ):\n        super().__init__()\n        assert depth % 2 == 0\n\n        # Hierarchical patch embedding (stride 16 total)\n        self.patch_embed = nn.Sequential(\n            nn.Conv2d(3, embed_dim // 4, 4, stride=4, bias=False),\n            nn.BatchNorm2d(embed_dim // 4), nn.GELU(),\n            nn.Conv2d(embed_dim // 4, embed_dim // 2, 2, stride=2, bias=False),\n            nn.BatchNorm2d(embed_dim // 2), nn.GELU(),\n            nn.Conv2d(embed_dim // 2, embed_dim, 2, stride=2, bias=False),\n            nn.BatchNorm2d(embed_dim),\n        )\n        self.pos_drop = nn.Dropout(0.1)\n        self.blocks   = nn.ModuleList([\n            m for _ in range(depth // 2)\n            for m in (\n                AdaptivePathologyScaleDetector(embed_dim),\n                VesselMorphologicalBlock(embed_dim, d_state, dropout=0.1),\n            )\n        ])\n        self.norm = nn.LayerNorm(embed_dim)\n        self.pool = nn.AdaptiveAvgPool1d(1)\n        self.head = OrdinalUncertaintyHead(embed_dim, num_classes, dropout)\n\n        self.apply(self._init_weights)\n\n        # O-M7: torch.compile (no-op on PyTorch < 2.0 or if disabled)\n        if use_compile:\n            self.patch_embed = _maybe_compile(self.patch_embed)\n\n    @staticmethod\n    def _init_weights(m):\n        if isinstance(m, nn.Linear):\n            nn.init.trunc_normal_(m.weight, std=0.02)\n            if m.bias is not None: nn.init.zeros_(m.bias)\n        elif isinstance(m, (nn.LayerNorm, nn.BatchNorm2d)):\n            nn.init.ones_(m.weight); nn.init.zeros_(m.bias)\n        elif isinstance(m, nn.Conv2d):\n            nn.init.kaiming_normal_(m.weight, mode='fan_out')\n            if m.bias is not None: nn.init.zeros_(m.bias)\n\n    @staticmethod\n    def _pos_embed(H: int, W: int, D: int, device) -> torch.Tensor:\n        pe  = torch.zeros(1, H * W, D, device=device)\n        div = torch.exp(\n            torch.arange(0, D, 2, device=device, dtype=torch.float)\n            * (-math.log(10000.0) / D)\n        )\n        y  = torch.arange(H, device=device).float().unsqueeze(1)\n        xc = torch.arange(W, device=device).float().unsqueeze(1)\n        pe[0, :, 0::2] = (torch.sin(y * div).unsqueeze(1)\n                           + torch.sin(xc * div).unsqueeze(0)).reshape(H * W, -1)\n        pe[0, :, 1::2] = (torch.cos(y * div).unsqueeze(1)\n                           + torch.cos(xc * div).unsqueeze(0)).reshape(H * W, -1)\n        return pe\n\n    def forward(self, x: torch.Tensor) -> dict:\n        x        = self.patch_embed(x)\n        B, D, H, W = x.shape\n        tokens   = x.flatten(2).transpose(1, 2)\n        tokens   = self.pos_drop(tokens + self._pos_embed(H, W, D, x.device))\n        for block in self.blocks:\n            tokens = block(tokens, H, W)\n        pooled   = self.pool(self.norm(tokens).transpose(1, 2)).flatten(1)\n        return self.head(pooled)\n\n    # FIX 2: BN frozen / Dropout stochastic (O-M6: n_samples default 10)\n    @torch.inference_mode()\n    def predict_with_uncertainty(\n        self, x: torch.Tensor, n_samples: int = 10\n    ) -> dict:\n        def _set_modes(m):\n            if isinstance(m, (nn.BatchNorm1d, nn.BatchNorm2d, nn.LayerNorm)):\n                m.eval()\n            elif isinstance(m, (nn.Dropout, nn.Dropout2d)):\n                m.train()\n        self.apply(_set_modes)\n        # torch.inference_mode disables grad; we still get stochastic dropout\n        preds = torch.stack(\n            [self.head.predict_grade(self.forward(x)['grade_logits']).float()\n             for _ in range(n_samples)], dim=0\n        )\n        self.eval()\n        return {\n            'mean_grade':  torch.clamp(torch.round(preds.mean(0)), 0,\n                                       self.head.num_classes - 1).long(),\n            'uncertainty': preds.std(0),\n        }\n\n    def param_count(self) -> int:\n        return sum(p.numel() for p in self.parameters() if p.requires_grad)\n\n\n# ──────────────────────────────────────────────────────────────\n# 7.  Focal Ordinal Loss  (O-M4: vectorised KL term)\n# ──────────────────────────────────────────────────────────────\nclass FocalOrdinalLoss(nn.Module):\n\n    def __init__(self, num_classes=5, lambda_kl=0.2, gamma=2.0, label_smooth=0.1):\n        super().__init__()\n        self.K, self.lk, self.gamma, self.ls = num_classes, lambda_kl, gamma, label_smooth\n\n        # Pre-build ordinal offset tensor [-K//2 … K//2] for neighbour smoothing\n        self.register_buffer('_k_range',\n                             torch.arange(num_classes, dtype=torch.float))\n\n    def _smooth_dist(self, target: torch.Tensor) -> torch.Tensor:\n        \"\"\"Vectorised ordinal-smooth label distribution.  O-M4.\"\"\"\n        B  = target.shape[0]\n        K  = self.K\n        ls = self.ls\n        # All-zero base, then scatter mass\n        dist = torch.zeros(B, K, device=target.device)\n        g    = target.view(B, 1)               # [B, 1]\n        dist.scatter_(1, g, 1.0 - ls)\n        # Left neighbour\n        left = (g - 1).clamp(min=0)\n        left_mask = (g > 0).float()\n        dist.scatter_add_(1, left, left_mask * (ls / 2.0))\n        # Right neighbour\n        right = (g + 1).clamp(max=K - 1)\n        right_mask = (g < K - 1).float()\n        dist.scatter_add_(1, right, right_mask * (ls / 2.0))\n        return dist\n\n    def forward(self, out: dict, target: torch.Tensor):\n        logit = out['grade_logits']\n        reg   = out['reg_score'].squeeze(-1)\n        tf    = target.float()\n\n        # Focal MSE\n        mse       = (reg - tf) ** 2\n        focal_mse = ((1 - torch.exp(-mse)) ** self.gamma * mse).mean()\n\n        # Ordinal BCE\n        ordinal_tgt = torch.stack(\n            [(target > k).float() for k in range(self.K - 1)], dim=1\n        )\n        bce_loss = F.binary_cross_entropy_with_logits(logit, ordinal_tgt)\n\n        # KL smooth (O-M4)\n        log_pred = F.log_softmax(\n            -torch.abs(reg.unsqueeze(1) - self._k_range), dim=-1\n        )\n        dist    = self._smooth_dist(target)\n        kl_loss = F.kl_div(log_pred, dist, reduction='batchmean')\n\n        total = focal_mse + bce_loss + self.lk * kl_loss\n        return total, {'focal_mse': focal_mse.item(),\n                       'bce': bce_loss.item(), 'kl': kl_loss.item()}\n\n\n# ──────────────────────────────────────────────────────────────\n# 8.  GradCAM (unchanged from v1 — hooks identical)\n# ──────────────────────────────────────────────────────────────\nclass GradCAMVimDR:\n    def __init__(self, model: VimDRPlus):\n        self.model  = model\n        self._acts  = self._grads = None\n        self._hooks = []\n        target = next(\n            (b for b in reversed(list(model.blocks))\n             if isinstance(b, VesselMorphologicalBlock)), None\n        )\n        if target is None:\n            raise ValueError(\"No VesselMorphologicalBlock in model.\")\n        self._hooks.append(\n            target.register_forward_hook(\n                lambda m, i, o: setattr(self, '_acts', o.detach())))\n        self._hooks.append(\n            target.register_full_backward_hook(\n                lambda m, gi, go: setattr(self, '_grads', go[0].detach())))\n\n    def __call__(self, img: torch.Tensor, target_class: int,\n                 H: int, W: int) -> 'np.ndarray':\n        import cv2, numpy as np\n        self.model.eval()\n        self._acts = self._grads = None\n        img = img.clone().requires_grad_(True)\n        out = self.model(img)\n        score = (out['grade_logits'][0, target_class - 1]\n                 if target_class > 0 else -out['grade_logits'][0, 0])\n        self.model.zero_grad(); score.backward()\n        if self._acts is None or self._grads is None:\n            return np.zeros((img.shape[-2], img.shape[-1]), dtype=np.float32)\n        weights = self._grads.mean(dim=1, keepdim=True)\n        cam = F.relu((self._acts * weights).sum(-1))[0].view(H, W).detach().cpu().numpy()\n        cam -= cam.min()\n        if cam.max() > 0: cam /= cam.max()\n        return cv2.resize(cam.astype(np.float32), (img.shape[-1], img.shape[-2]))\n\n    def remove_hooks(self):\n        for h in self._hooks: h.remove()\n        self._hooks.clear()\n\n\n# ──────────────────────────────────────────────────────────────\n# 9.  Sanity check\n# ──────────────────────────────────────────────────────────────\nif __name__ == '__main__':\n    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n    model  = VimDRPlus(embed_dim=192, depth=8, num_classes=5, d_state=16).to(device)\n    x      = torch.randn(2, 3, 384, 384, device=device)   # 384 for speed test\n\n    with torch.no_grad():\n        out = model(x)\n\n    logits = out['grade_logits']\n    probs  = OrdinalUncertaintyHead.logits_to_probs(logits)\n    print(f\"Parameters  : {model.param_count():,}\")\n    print(f\"Logits shape: {logits.shape}\")\n    print(f\"Probs sum   : {probs.sum(-1)}\")\n    print(f\"Grades      : {model.head.predict_grade(logits)}\")\n    unc = model.predict_with_uncertainty(x[:1], n_samples=3)\n    print(f\"Uncertainty : {unc['uncertainty']}\")\n    print(\"Kaggle model sanity check passed.\")","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}