{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":117682,"databundleVersionId":15062069,"sourceType":"competition"}],"dockerImageVersionId":31234,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"under development,\nsee discussion:\n- https://www.kaggle.com/competitions/vesuvius-challenge-surface-detection/discussion/651532#3382504\n- https://www.kaggle.com/competitions/vesuvius-challenge-surface-detection/discussion/651532#3383883","metadata":{}},{"cell_type":"markdown","source":"# 1. modeling","metadata":{}},{"cell_type":"code","source":"import math\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\n\nclass PoolInY(nn.Module):\n    def __init__(self, in_c: int, out_c: int):\n        super().__init__()\n        #self.proj = nn.Conv3d(in_c, out_c, kernel_size=1) # noisy\n        #self.proj = nn.Conv3d(in_c, out_c, kernel_size=3, padding=1)\n        self.proj = nn.Conv3d(in_c, out_c, kernel_size=5, padding=2)\n\n    def forward(self, feat3d: torch.Tensor) -> torch.Tensor:\n        # B,C,D,H,W --> B,C,D,W\n        feat3d = self.proj(feat3d)\n        B,C,D,H,W = feat3d.shape\n        feat2d = feat3d.permute(0, 1, 3, 2, 4).reshape(B,-1, D, W)  #stack of feature[y] as a \"token\"\n        feat2d = feat2d.contiguous()\n        return feat2d\n\nclass SurfaceHead(nn.Module):\n    def __init__(\n        self,\n        last_dim,\n        query_dim,\n        pool_dim,\n        d_model=128,\n        last_shape=(80,80,80),\n        #num_query=20,\n    ):\n        super().__init__()\n        #self.query = nn.Embedding(num_query, query_dim)\n        self.last_shape = last_shape\n        Df,Hf,Wf = last_shape\n\n        self.pool_in_y = PoolInY(last_dim, pool_dim)\n        h_pool_dim = Hf*pool_dim\n\n        # self.canonial = nn.Sequential(\n        #     nn.Conv2d(h_pool_dim, d_model, 3, padding=1),\n        #     nn.ReLU(inplace=True),\n        #     nn.Conv2d(d_model, 1, 1),\n        # )\n        self.y_scale = nn.Parameter(torch.tensor(1.0))\n\n        # --- Query -> channel weights (for delta y) ---\n        self.query_to_y_weight = nn.Sequential(\n            nn.Linear(query_dim, h_pool_dim//2),\n            nn.ReLU(inplace=True),\n            nn.Linear(h_pool_dim//2, h_pool_dim),\n        )\n        self.query_to_y_bias = nn.Linear(query_dim, 1)\n\n        self.query_to_visibility_weight = nn.Linear(query_dim, h_pool_dim)\n        self.query_to_visibility_bias = nn.Linear(query_dim, 1)\n\n        self.query_to_detect = nn.Linear(query_dim, 1) #detcted instance score\n\n    def forward(self,\n        query,\n        last\n    ):\n        \"\"\"\n        query_feats: (B,K,Cq)\n        feat2d:      (B,Cf,Hc,Wc)\n        \"\"\"\n\n        B = last.shape[0]\n        _, K, Cq = query.shape\n        _, Cf, Df, Hf, Wf = last.shape\n        assert Hf == self.last_shape[1]\n\n        pooled = self.pool_in_y(last)  # (B,Cf,Dc,Wc)\n        _,_,Dc,Wc = pooled.shape\n        #print(pooled.shape)\n\n\n        # surface = [z,y[z,x],x]\n        # y: per query linear combination of last feature channels\n        weight = self.query_to_y_weight(query)\n        bias   = self.query_to_y_bias(query)                # (B,K,1)\n        y      = torch.einsum(\"bkc,bcdw->bkdw\", weight, pooled)  + bias.unsqueeze(-1)    # (B,K,Dc,Wc)\n        y = F.softplus(self.y_scale * y)  #non zero y coordinate\n        #(Hf-1) * torch.sigmoid(y_raw)\n\n\n        # non ordered\n        surface_y  = y #(B,K,Dc,Wc)\n\n        # optional ordered y field\n        # canonial = self.canonial(pooled)  # (B,1,Dc,Wc)\n        #surface_y = canonial + torch.cumsum(dy, dim=1)\n        # we turn this off becuase ordering is implemented as post-processing in stage-2\n\n        # --- visibility ---\n        weight = self.query_to_visibility_weight(query)  # (B,K,c)\n        bias = self.query_to_visibility_bias(query)    # (B,K,1)\n        surface_vis = torch.einsum(\"bkc,bcdw->bkdw\", weight, pooled) + bias.unsqueeze(-1)    # (B,K,Dc,Wc)\n\n\n        # --- presence ---\n        detect = self.query_to_detect(query).squeeze(-1)\n\n\n        surface_z, surface_x = torch.meshgrid(\n            torch.linspace(0,Dc-1,Dc,dtype=last.dtype,device=last.device),\n            torch.linspace(0,Wc-1,Wc,dtype=last.dtype,device=last.device),\n            indexing=\"ij\"\n        )\n\n        return {\n            'surface_z':surface_z,\n            'surface_x':surface_x,\n            'surface_y':surface_y,\n            'surface_vis':surface_vis, #logit\n            'detect':detect, #logit\n        }\n \n####################################################################################33\n#decoder\n# def sine_positional_encoding_3d_by_sum(d, h, w, c, device):\n#     z = torch.linspace(-1, 1, steps=d, device=device)\n#     y = torch.linspace(-1, 1, steps=h, device=device)\n#     x = torch.linspace(-1, 1, steps=w, device=device)\n#     zz, yy, xx = torch.meshgrid(z, y, x, indexing='ij')\n#     coord = torch.stack([zz, yy, xx], dim=-1).reshape(-1, 3)\n\n#     half = c // 2\n#     freq = torch.exp(torch.linspace(0, math.log(10000), steps=half, device=device))\n#     freqs= 1.0 / freq\n\n#     phase = coord.sum(dim=-1, keepdim=True) * freqs.unsqueeze(0)\n#     pe = torch.cat([torch.sin(phase), torch.cos(phase)], dim=-1)\n#     if pe.shape[-1] < c:\n#         pe = F.pad(pe, (0, c - pe.shape[-1]))\n#     return pe[:, :c].unsqueeze(1)  # (T,1,C)\n\ndef sine_positional_encoding_1d(pos: torch.Tensor, dim: int) -> torch.Tensor:\n    \"\"\"\n    pos: (T,) float\n    dim: output channels for this axis (will be padded if odd)\n    returns: (T, dim)\n    \"\"\"\n    # make dim even for sin/cos pairs\n    dim_even = (dim // 2) * 2\n    if dim_even == 0:\n        return pos.new_zeros(pos.numel(), dim)\n\n    half = dim_even // 2\n    # classic transformer frequencies\n    div_term = torch.exp(\n        torch.arange(half, device=pos.device, dtype=pos.dtype) * (-math.log(10000.0) / half)\n    )  # (half,)\n    phase = pos[:, None] * div_term[None, :]  # (T, half)\n    pe = torch.cat([torch.sin(phase), torch.cos(phase)], dim=-1)  # (T, dim_even)\n    if dim_even < dim:\n        pe = F.pad(pe, (0, dim - dim_even))\n    return pe\n\ndef sine_positional_encoding_3d_by_cat (d: int, h: int, w: int, c: int, device):\n    \"\"\"\n    returns: (T, 1, C) where T = d*h*w\n    Encodes z,y,x separately and concatenates.\n    \"\"\"\n    z = torch.linspace(-1, 1, steps=d, device=device)\n    y = torch.linspace(-1, 1, steps=h, device=device)\n    x = torch.linspace(-1, 1, steps=w, device=device)\n    zz, yy, xx = torch.meshgrid(z, y, x, indexing=\"ij\")  # (d,h,w)\n\n    T = d * h * w\n    zz = zz.reshape(T)\n    yy = yy.reshape(T)\n    xx = xx.reshape(T)\n\n    # split channels across 3 axes\n    c_z = c // 3\n    c_y = c // 3\n    c_x = c - c_z - c_y  # remainder goes to x\n\n    pe_z = sine_positional_encoding_1d(zz, c_z)  # (T, c_z)\n    pe_y = sine_positional_encoding_1d(yy, c_y)  # (T, c_y)\n    pe_x = sine_positional_encoding_1d(xx, c_x)  # (T, c_x)\n\n    pe = torch.cat([pe_z, pe_y, pe_x], dim=-1)   # (T, C)\n    return pe.unsqueeze(1)  # (T,1,C)\n\n\nclass InstanceDecoder3D(nn.Module):\n\n    def __init__(\n        self,\n        feature_dim,\n        last_dim,\n        pool_dim=12,\n        num_query=20,\n        d_model=256,\n        dim_ff=1024,\n        num_head=8,\n        num_layer=6, #divided into n feature level\n        last_shape = (80,80,80)\n    ):\n        super().__init__()\n        self.num_feature = len(feature_dim)\n        self.num_head = num_head\n        self.num_layer = num_layer\n        self.d_model = d_model\n\n        # project intermediate features -> d_model\n        self.feature_project = nn.ModuleList([\n            nn.Conv3d(c, d_model, kernel_size=1) for c in feature_dim\n        ])\n        self.level_embed = nn.Embedding(self.num_feature, d_model)\n\n        # project final UNet feature -> d_model (used for Z decoding)\n        #self.last_project = nn.Conv3d(last_dim, d_model, kernel_size=1)\n\n        # learnable query\n        self.query     = nn.Embedding(num_query, d_model)\n        self.query_pos = nn.Embedding(num_query, d_model)\n\n        #------ transformer decoder block (dynamic query) --------------\n        self.cross_attn = nn.ModuleList([\n            nn.MultiheadAttention(d_model, num_head, dropout=0.0, batch_first=False)\n            for i in range(num_layer)\n        ])\n        self.self_attn = nn.ModuleList([\n            nn.MultiheadAttention(d_model, num_head, dropout=0.0, batch_first=False)\n            for i in range(num_layer)\n        ])\n        self.ffn = nn.ModuleList([\n            nn.Sequential(\n                nn.Linear(d_model, dim_ff),\n                nn.ReLU(inplace=True),\n                nn.Linear(dim_ff, d_model),\n            )\n            for i in range(num_layer)\n        ])\n        self.norm_q   = nn.ModuleList([nn.LayerNorm(d_model) for i in range(num_layer)])\n        self.norm_sa  = nn.ModuleList([nn.LayerNorm(d_model) for i in range(num_layer)])\n        self.norm_ffn = nn.ModuleList([nn.LayerNorm(d_model) for i in range(num_layer)])\n        self.final_norm = nn.LayerNorm(d_model)\n\n        #----\n        #todo: instance head\n        self.surface_head = SurfaceHead(\n            last_dim,\n            query_dim=d_model,\n            pool_dim=pool_dim,\n            d_model=128,\n            last_shape=last_shape,\n        )\n\n\n    ###########################################################################\n    def feature_to_token(self, feature_lvl, lvl):\n        \"\"\"\n        feature_lvl: (B,C,D,H,W) -> tokens (T,B,d), pos (T,B,d), size (D,H,W)\n        \"\"\"\n        B, _, D, H, W = feature_lvl.shape\n        x = self.feature_project[lvl](feature_lvl)  # (B,d,D,H,W)\n        x = x + self.level_embed.weight[lvl][None, :, None, None, None]\n        token = x.flatten(2).permute(2, 0, 1).contiguous()  # (T,B,d)\n\n        pos = sine_positional_encoding_3d_by_cat(D, H, W, self.d_model, device=feature_lvl.device)\n        pos = pos.repeat(1, B, 1)  # (T,B,d)\n        return token, pos, (D, H, W)\n\n\n\n    def forward(self, feature, last, need_weights=False):\n        B = feature[0].shape[0]\n\n        # static queries P0: (Q,B,d)\n        P     = self.query.weight.unsqueeze(1).repeat(1, B, 1)\n        P_pos = self.query_pos.weight.unsqueeze(1).repeat(1, B, 1)\n\n        # dynamic queries\n        attention=[]\n        for t in range(self.num_layer):\n            lvl = t % self.num_feature #Alternate schedule\n            # or lvl = (t * self.num_feature) // self.num_layer #Block schedule\n            mem, mem_pos, size = self.feature_to_token(feature[lvl], lvl)\n            attn_mask = None # maskformer2 attn mask not used here\n\n            q = P   + P_pos\n            k = mem + mem_pos\n            v = mem\n\n            P2, cross_attn_weight = self.cross_attn[t](q, k, v, attn_mask=attn_mask, need_weights=need_weights)\n            P = self.norm_q[t](P + P2)\n\n            # self-attn\n            q = P + P_pos\n            P2, self_attn_weight = self.self_attn[t](q, q, P, need_weights=need_weights)\n            P = self.norm_sa[t](P + P2)\n\n            # ffn\n            P2 = self.ffn[t](P)\n            P = self.norm_ffn[t](P + P2)\n\n            attention.append([cross_attn_weight, self_attn_weight])\n\n        query = self.final_norm(P).permute(1, 0, 2).contiguous()  # (B,Q,d)\n\n        #----\n        out = self.surface_head(query, last)\n\n        #optional: additional output for debug\n        out['query'] = query  \n        out['attention'] =  attention\n        return out","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-01-01T20:11:25.989607Z","iopub.execute_input":"2026-01-01T20:11:25.990061Z","iopub.status.idle":"2026-01-01T20:11:31.442490Z","shell.execute_reply.started":"2026-01-01T20:11:25.990020Z","shell.execute_reply":"2026-01-01T20:11:31.441440Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def run_check_model():\n\n    # ----- Fake decoder features (multi-scale) -----\n    B   = 2\n\n    #decoder ouput from unet\n    d0 = torch.randn(B, 256, 20, 20, 20)\n    d1 = torch.randn(B, 128, 40, 40, 40)\n    d2 = torch.randn(B,  64, 80, 80, 80)\n    d3 = torch.randn(B,  32,160,160,160) #not use\n    last = d2\n\n    decoder_dim = [256,128,64,32]\n\n    model = InstanceDecoder3D(\n        feature_dim=decoder_dim[:2],\n        last_dim=decoder_dim[2],\n        num_query=20,\n        d_model=256,\n        dim_ff=1024,\n        num_head=8,\n        num_layer=6,\n    )\n\n    out = model(\n        feature=[d0, d1],\n        last=last,\n        need_weights = True # output attention weight for debug\n    )\n\n    for k,v in out.items():\n        if k in ['attention']: continue\n        print(f'{k}:',v.shape)\n\n    print('attention:')\n    for i in range(len(out['attention'])):\n        print('\\t',i, out['attention'][i][0].shape, out['attention'][i][1].shape)\n    print('')\n\nrun_check_model()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-01T20:12:06.700842Z","iopub.execute_input":"2026-01-01T20:12:06.701208Z","iopub.status.idle":"2026-01-01T20:12:17.528569Z","shell.execute_reply.started":"2026-01-01T20:12:06.701176Z","shell.execute_reply":"2026-01-01T20:12:17.527494Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 2.Loss","metadata":{}},{"cell_type":"code","source":"from scipy.optimize import linear_sum_assignment\n\n@torch.no_grad()\ndef match_hungarian(\n    pred_y, pred_vis, pred_p,\n    truth_y, truth_vis, truth_exist, #for padding to same num. if instance in batch\n    w_reg=1.0, w_v=0.2, w_p=0.05\n):\n    #z_stride=4, i_stride=4, #for fast sampling\n    SUB=4\n\n    B, K, Z, X = pred_y.shape\n    _, N, _, _ = truth_y.shape\n\n    truth_vis = (truth_vis>0.5).float() #must be 0,1\n\n    # Subsample for cost\n    py = pred_y[:,:,::SUB]          # (B,K,Zs,Xs) # here we only sibsample in Z\n    pvis = pred_vis[:,:,::SUB]      # (B,K,Zs,Xs)\n    gy = truth_y[:,:,::SUB]         # (B,N,Zs,Xs)\n    gvis = truth_vis[:,:,::SUB]     # (B,N,Zs,Xs)\n    Zs, Xs = py.shape[-2], py.shape[-1]\n\n    match_idx = torch.zeros((B, N), dtype=torch.long, device=pred_y.device)\n\n    for b in range(B):\n        # index of real truth\n        i = torch.nonzero(truth_exist[b] > 0.5, as_tuple=False).squeeze(-1)\n        J = i.numel()\n\n        if J == 0:\n            continue\n\n        # Build cost matrix for this b: (J,K)\n        py_b = py[b,None]              # (1,K,Zs,Xs)\n        gy_b = gy[b, i,None]           # (J,1,Zs,Xs)\n        gvis_b = gvis[b, i, None]      # (J,1,Zs,Xs)\n\n        # regression cost\n        reg  = (py_b-gy_b).abs()          # (J,K,Zs,Xs)\n        mask = gvis_b.expand(J, K, Zs,Xs)  # (J,K,Zs,Xs)\n        den = mask.sum(dim=(-2, -1)).clamp_min(1.0)\n        reg_cost = (mask * reg).sum(dim=(-2, -1)) / den  # (J,K)\n\n        # visibility cost\n        pvis_b = pvis[b, None]    # (1,K,Zs,Xs)\n        gvis_b = gvis_b.float()   # (J,1,Zs,Xs)\n        bce = F.binary_cross_entropy_with_logits(\n            pvis_b.expand( J, K, Zs, Xs),\n            gvis_b.expand( J, K, Zs, Xs),\n            reduction='none')  # torch.Size([1, 5, 20, 20, 80])\n        bce_cost = bce.sum(dim=(-2, -1)) / (Zs * Xs)  # (J,K)\n\n        # presence cost\n        pb_b = pred_p[b,None]  # (1,K)\n        logp = -F.logsigmoid(pb_b)\n        logp_cost = logp.expand(J, K) # (J,K)\n\n        cost = w_reg * reg_cost + w_v * bce_cost + w_p * logp_cost  # (J,K)\n        cost = cost.float() .detach().cpu().numpy()\n        row_ind, col_ind = linear_sum_assignment(cost)  # row_ind in [0..J-1], col_ind in [0..K-1]\n        match_idx[b, i[row_ind]] = torch.as_tensor(col_ind, device=pred_y.device, dtype=torch.long)\n\n    return match_idx\n\n\n\ndef smoothness_l1(field, mask, eps=1e-6):\n\n    d1 = field[:, :, :, 1:] - field[:, :, :, :-1]\n    d2 = field[:, :, 1:, :] - field[:, :, :-1, :,]\n    m1 = mask[:, :, :, 1:] * mask[:, :, :, :-1]\n    m2 = mask[:, :, 1:, :] * mask[:, :, :-1, :]\n\n    loss1 = (d1.abs() * m1).sum() /  m1.sum().clamp_min(eps)\n    loss2 = (d2.abs() * m2).sum() /  m2.sum().clamp_min(eps)\n    loss = loss1 + loss2\n    return loss\n\n'''\n#todo \nis_unlabel is to be computed outside, like match_idx.\nis_unlabel to be either zero or one (according to scoring in the external function)\n\n'''\n\ndef surface_as_unlabelled(\n    \n):\n    is_unlabel=None\n    return is_unlabel\n\ndef compute_surface_loss(\n    pred_y, pred_vis, pred_p,\n    truth_y, truth_vis, truth_exist,\n    match_idx,\n    is_unlabel = None,\n    w_reg=1.0, w_vis=1.0, w_det=1.0, w_smooth=0.1,\n    w_unmatched=1.0,   # weight for unmatched/background presence penalty\n):\n    B, K, Z, X = pred_y.shape\n    _, N, _, _ = truth_y.shape\n    Y = X\n    device = pred_y.device\n\n    if is_unlabel is None:\n        is_unlabel = torch.ones((B,K), dtype=pred_y.dtype, device=device)\n\n\n    gy   = truth_y\n    gvis = (truth_vis > 0.5).float()\n    gb   = (truth_exist > 0.5).float()  # (B,N)\n\n    # match case only !!!!\n    mask_voxel    = gvis * gb[:, :, None, None]      # (B,N,Z,X)\n    mask_instance = gb[:, :, None, None]             # (B,N,1,1)\n    idx  = match_idx.clamp(0, K-1)  # (B,N)\n    py   = pred_y.gather(1, idx[:, :, None, None].expand(B, N, Z, X))\n    pvis = pred_vis.gather(1, idx[:, :, None, None].expand(B, N, Z, X))\n\n    # --- regression (masked L1) ---\n    loss_reg = (mask_voxel * (py - gy).abs()).sum() / mask_voxel.sum().clamp_min(1e-6)\n\n    # --- visibility BCE (masked by instance existence) ---\n    vis = F.binary_cross_entropy_with_logits(pvis, gvis, reduction=\"none\")  # (B,N,Z,X)\n    vis_sum = (mask_instance * vis).sum()\n    loss_vis = vis_sum / (mask_instance.sum().clamp_min(1.0) * (Z * X))     # <-- FIXED\n\n    # --- presence: matched positives + unmatched (maybe positive in unlabelled) ---\n    assigned = torch.zeros((B, K), dtype=torch.bool, device=device)\n    pb_pos_list = []\n    gb_pos_list = []\n\n    for b in range(B):\n        n_valid = torch.nonzero(gb[b] > 0.5, as_tuple=False).squeeze(-1)\n        if n_valid.numel() == 0:\n            continue\n        k_assigned = idx[b, n_valid]\n        assigned[b, k_assigned] = True\n\n        pb_pos_list.append(pred_p[b, k_assigned])\n        gb_pos_list.append(torch.ones_like(pred_p[b, k_assigned]))\n\n    if len(pb_pos_list) > 0:\n        pb_pos = torch.cat(pb_pos_list, dim=0)\n        gb_pos = torch.cat(gb_pos_list, dim=0)\n        loss_det_pos = F.binary_cross_entropy_with_logits(pb_pos, gb_pos)\n    else:\n        loss_det_pos = pred_p.sum() * 0.0\n\n    #-------------\n    # negatives: only for unassigned and NOT in unlabeled region\n    if (~assigned).any():\n        neg_logits = pred_p[~assigned]\n        neg_target = torch.zeros_like(neg_logits) #force as background\n\n        neg_w = (1.0 - is_unlabel)[~assigned]  # (num_neg,)\n        neg_bce = F.binary_cross_entropy_with_logits(neg_logits, neg_target, reduction=\"none\")\n\n        den = neg_w.sum()\n        if den < 1e-6:\n            loss_det_neg = pred_p.sum() * 0.0\n        else:\n            loss_det_neg = (neg_w * neg_bce).sum() / den\n    else:\n        loss_det_neg = pred_p.sum() * 0.0\n\n    loss_det = loss_det_pos + w_unmatched * loss_det_neg ###todo: check imbalance ratio between pos and neg\n\n    # --- smoothness on visible regions ---\n    loss_smooth = smoothness_l1(py, mask_voxel)\n\n    total = w_reg*loss_reg + w_vis*loss_vis + w_det*loss_det + w_smooth*loss_smooth\n    return dict(\n        loss_total=total,\n        loss_reg=loss_reg,\n        loss_vis=loss_vis,\n        loss_det=loss_det,\n        loss_smooth=loss_smooth,\n    )","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}