{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.11","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":39763,"databundleVersionId":11756775,"sourceType":"competition"}],"dockerImageVersionId":31040,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import torch\nfrom torch import nn\nimport timm","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-30T22:39:54.423437Z","iopub.execute_input":"2025-06-30T22:39:54.423631Z","iopub.status.idle":"2025-06-30T22:40:05.258471Z","shell.execute_reply.started":"2025-06-30T22:39:54.423610Z","shell.execute_reply":"2025-06-30T22:40:05.257777Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def reconstruct_tensor(output_tensor, cls_token_first=True):\n    bs = output_tensor.shape[0]\n    # 1. Remove CLS token (first token)\n    if cls_token_first:\n        patches = output_tensor[:, 1:, :]  # Shape [1, 625, 196]\n    else:\n        patches = output_tensor\n        \n    # 2. Reshape to [1, 25, 25, 14, 14] (25x25 grid of 14x14 patches)\n    patches = patches.view(bs, CFG.image_size[0]//14, CFG.image_size[1]//14, 14, 14)\n\n    # 3. Re-arrange patches into spatial order\n    #    - Combine grid rows first, then patch rows\n    reconstructed = patches.permute(0, 1, 3, 2, 4).contiguous()\n    reconstructed = reconstructed.view(bs, 1, CFG.image_size[0], CFG.image_size[1])  # Final shape\n\n    return reconstructed\n\ndef attention_rope_forward(self, x: torch.Tensor, rope=None) -> torch.Tensor:\n        B, N, C = x.shape\n        qkv = self.qkv(x).reshape(B, N, 3, self.num_heads, self.head_dim).permute(2, 0, 3, 1, 4)\n        q, k, v = qkv.unbind(0)\n        q, k = self.q_norm(q), self.k_norm(k)\n        \n        if rope is not None:\n            npt = 1 #self.num_prefix_tokens\n            q = torch.cat([q[:, :, :npt, :], timm.layers.apply_rot_embed_cat(q[:, :, npt:, :], rope)], dim=2).type_as(v)\n            k = torch.cat([k[:, :, :npt, :], timm.layers.apply_rot_embed_cat(k[:, :, npt:, :], rope)], dim=2).type_as(v)\n        \n        if self.fused_attn:\n            x = nn.functional.scaled_dot_product_attention(\n                q, k, v,\n                dropout_p=self.attn_drop.p if self.training else 0.,\n            )\n        else:\n            q = q * self.scale\n            attn = q @ k.transpose(-2, -1)\n            attn = attn.softmax(dim=-1)\n            attn = self.attn_drop(attn)\n            x = attn @ v\n\n        x = x.transpose(1, 2).reshape(B, N, C)\n        x = self.proj(x)\n        x = self.proj_drop(x)\n        return x\n    \ndef block_rope_forward(self, x: torch.Tensor, rope) -> torch.Tensor:\n    #x = x + self.drop_path1(self.ls1(self.attn(self.norm1(x), rope)))\n    x = x + self.drop_path1(self.ls1(attention_rope_forward(self.attn, self.norm1(x), rope)))\n    x = x + self.drop_path2(self.ls2(self.mlp(self.norm2(x))))\n    return x\n\nclass Model(nn.Module):\n    def __init__(self, pretrained=True, drop=0.):\n        super(Model\n              , self).__init__()\n        \n        patch_size = 14\n        \n        self.encoder = timm.create_model(CFG.model_name, pretrained=pretrained, in_chans=1, img_size=CFG.image_size, num_classes=0, global_pool='')\n        self.decoder = timm.create_model(CFG.model_name, pretrained=pretrained, in_chans=1, img_size=CFG.image_size, num_classes=0, global_pool='')\n        \n        self.feats = self.encoder.num_features\n        \n        self.head = nn.Linear(self.feats, patch_size*patch_size)\n        #self.head = nn.Linear(self.feats, 16)\n        self.sigmoid = nn.Sigmoid()\n        \n        #st = torch.load('./data/pretrained/eva02_small_patch14_224.mim_in22k_pytorch_model.bin', map_location='cpu')\n        #del st['pos_embed'], st['patch_embed.proj.weight']\n        #self.encoder.load_state_dict(st, strict=False)\n        \n        self.rope = timm.layers.RotaryEmbeddingCat(\n                self.feats // 6,\n                in_pixels=False,\n                feat_shape=None,# if dynamic_img_size else self.patch_embed.grid_size,\n                ref_feat_shape=None,#ref_feat_shape,\n            )\n        \n        self.cls_token = nn.Parameter(torch.zeros(1, 1, self.feats))\n        self.reg_token = None#nn.Parameter(torch.zeros(1, 1, self.feats))\n        \n        #for block in self.encoder.blocks:\n            #block.attn = Attention(384, num_heads=6, qkv_bias=False, qk_norm=False, proj_bias=True, attn_drop=0., proj_drop=0., norm_layer=nn.LayerNorm)\n        \n    def _pos_embed(self, x):\n        dynamic_img_size = True\n        self.pos_embed = None\n        \n        if dynamic_img_size:\n            B, H, W, C = x.shape\n            if self.pos_embed is not None:\n                prev_grid_size = self.patch_embed.grid_size\n                pos_embed = resample_abs_pos_embed(\n                    self.pos_embed,\n                    new_size=(H, W),\n                    old_size=prev_grid_size,\n                    num_prefix_tokens=self.num_prefix_tokens,\n                )\n            else:\n                pos_embed = None\n            x = x.view(B, -1, C)\n            rot_pos_embed = self.rope.get_embed(shape=(H, W)) if self.rope is not None else None\n        else:\n            pos_embed = self.pos_embed\n            rot_pos_embed = self.rope.get_embed() if self.rope is not None else None\n\n        #print(x.shape, self.cls_token.shape)\n            \n        if self.cls_token is not None:\n            x = torch.cat((self.cls_token.expand(x.shape[0], -1, -1), x), dim=1)\n\n        if self.reg_token is not None:\n            to_cat = []\n            if self.cls_token is not None:\n                to_cat.append(self.cls_token.expand(x.shape[0], -1, -1))\n            to_cat.append(self.reg_token.expand(x.shape[0], -1, -1))\n            x = torch.cat(to_cat + [x], dim=1)\n        \n        #print(rot_pos_embed.shape)\n        \n        return x, rot_pos_embed\n        \n    def forward_features(self, model, x: torch.Tensor) -> torch.Tensor:\n        x = model.patch_embed(x)\n        #x = self.encoder._pos_embed(x)\n        x = x.reshape(x.shape[0], CFG.image_size[0]//14, CFG.image_size[1]//14, x.shape[2])\n        x, rope = self._pos_embed(x)\n        x = model.patch_drop(x)\n        x = model.norm_pre(x)\n        \n        for blk in model.blocks:\n            #x = blk(x)\n            x = block_rope_forward(blk, x, rope)\n        \n        x = model.norm(x)\n        return x\n        \n    def forward(self, inp):\n        inp = torch.nan_to_num(inp, 0, 0, 0)\n        \n        #features = self.encoder(inp)\n        #print(features.shape)\n        features = self.forward_features(self.encoder, inp)\n        features = self.head(features)\n        \n        interm_masks = reconstruct_tensor(features)\n        \n        features = self.forward_features(self.decoder, interm_masks)\n        features = self.head(features)\n        \n        masks = reconstruct_tensor(features)\n        \n        masks = nn.functional.interpolate(masks, (70, 70), mode='bilinear')\n        \n        masks = self.sigmoid(masks)\n        \n        masks = (masks*3000.)+1500.\n        \n        return None, masks","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-30T22:40:05.261671Z","iopub.execute_input":"2025-06-30T22:40:05.261826Z","iopub.status.idle":"2025-06-30T22:40:05.278117Z","shell.execute_reply.started":"2025-06-30T22:40:05.261810Z","shell.execute_reply":"2025-06-30T22:40:05.277329Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class CFG:\n    model_name = 'vit_small_patch14_dinov2.lvd142m'\n    image_size = (350, 350)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-30T22:40:48.938144Z","iopub.execute_input":"2025-06-30T22:40:48.938690Z","iopub.status.idle":"2025-06-30T22:40:48.942163Z","shell.execute_reply.started":"2025-06-30T22:40:48.938670Z","shell.execute_reply":"2025-06-30T22:40:48.941241Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model = Model()\nmodel","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-30T22:40:57.858127Z","iopub.execute_input":"2025-06-30T22:40:57.858395Z","iopub.status.idle":"2025-06-30T22:40:58.716441Z","shell.execute_reply.started":"2025-06-30T22:40:57.858378Z","shell.execute_reply":"2025-06-30T22:40:58.715497Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"with torch.no_grad():\n    inp = torch.zeros((2, 1, 350, 350))\n    _, output = model(inp)\noutput.shape","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-30T22:41:38.789659Z","iopub.execute_input":"2025-06-30T22:41:38.789931Z","iopub.status.idle":"2025-06-30T22:41:39.997377Z","shell.execute_reply.started":"2025-06-30T22:41:38.789913Z","shell.execute_reply":"2025-06-30T22:41:39.996172Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}