{"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":"gpu","dataSources":[{"sourceType":"competition","sourceId":22990,"datasetId":1136396,"databundleVersionId":2048213},{"sourceType":"modelInstanceVersion","sourceId":769634,"databundleVersionId":15877836,"modelInstanceId":587935,"modelId":600257}],"dockerImageVersionId":31286,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        print(os.path.join(dirname, filename))\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-03-01T15:43:36.813540Z","iopub.execute_input":"2026-03-01T15:43:36.814105Z","iopub.status.idle":"2026-03-01T15:43:37.150132Z","shell.execute_reply.started":"2026-03-01T15:43:36.814073Z","shell.execute_reply":"2026-03-01T15:43:37.149405Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\n\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames[:5]:\n        print(os.path.join(dirname, filename))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-01T15:43:40.476814Z","iopub.execute_input":"2026-03-01T15:43:40.477651Z","iopub.status.idle":"2026-03-01T15:43:40.494989Z","shell.execute_reply.started":"2026-03-01T15:43:40.477619Z","shell.execute_reply":"2026-03-01T15:43:40.494237Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install imagecodecs -q","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-01T15:43:43.618511Z","iopub.execute_input":"2026-03-01T15:43:43.619078Z","iopub.status.idle":"2026-03-01T15:43:50.566077Z","shell.execute_reply.started":"2026-03-01T15:43:43.619049Z","shell.execute_reply":"2026-03-01T15:43:50.565357Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nfrom PIL import Image\nimport tifffile as tiff\nimport cv2\nimport os\nfrom tqdm.notebook import tqdm\nimport zipfile\nimport rasterio\nfrom rasterio.windows import Window\nfrom torch.utils.data import Dataset\nimport gc","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-01T15:43:56.865524Z","iopub.execute_input":"2026-03-01T15:43:56.866213Z","iopub.status.idle":"2026-03-01T15:43:59.408877Z","shell.execute_reply.started":"2026-03-01T15:43:56.866177Z","shell.execute_reply":"2026-03-01T15:43:59.408340Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"sz = 256   #the size of tiles\nreduce = 4 #reduce the original images by 4 times \nMASKS = '/kaggle/input/competitions/hubmap-kidney-segmentation/train.csv'\nDATA = '/kaggle/input/competitions/hubmap-kidney-segmentation/train/'\nOUT_TRAIN = 'train.zip'\nOUT_MASKS = 'masks.zip'","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-01T15:44:03.312197Z","iopub.execute_input":"2026-03-01T15:44:03.312665Z","iopub.status.idle":"2026-03-01T15:44:03.316805Z","shell.execute_reply.started":"2026-03-01T15:44:03.312638Z","shell.execute_reply":"2026-03-01T15:44:03.316113Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def enc2mask(encs, shape):\n    img = np.zeros(shape[0]*shape[1], dtype=np.uint8)\n    for m,enc in enumerate(encs):\n        if isinstance(enc, float) and np.isnan(enc): continue\n        s = enc.split()\n        for i in range(len(s)//2):\n            start = int(s[2*i]) - 1\n            length = int(s[2*i+1])\n            img[start:start+length] = 1 + m\n    return img.reshape(shape).T\n\ndef mask2enc(mask, n=1):\n    pixels = mask.T.flatten()\n    encs = []\n    for i in range(1,n+1):\n        p = (pixels == i).astype(np.int8)\n        if p.sum() == 0: encs.append(np.nan)\n        else:\n            p = np.concatenate([[0], p, [0]])\n            runs = np.where(p[1:] != p[:-1])[0] + 1\n            runs[1::2] -= runs[::2]\n            encs.append(' '.join(str(x) for x in runs))\n    return encs\n\ndf_masks = pd.read_csv(MASKS).set_index('id')\ndf_masks.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-01T15:44:32.752265Z","iopub.execute_input":"2026-03-01T15:44:32.753067Z","iopub.status.idle":"2026-03-01T15:44:33.089489Z","shell.execute_reply.started":"2026-03-01T15:44:32.753035Z","shell.execute_reply":"2026-03-01T15:44:33.088878Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"s_th = 40  #saturation blancking threshold\np_th = 1000*(sz//256)**2 #threshold for the minimum number of pixels\n\n\nclass HuBMAPDataset(Dataset):\n    def __init__(self, idx, sz=sz, reduce=reduce, encs=None):\n        self.data = rasterio.open(os.path.join(DATA,idx+'.tiff'),num_threads='all_cpus')\n        # some images have issues with their format \n        # and must be saved correctly before reading with rasterio\n        if self.data.count != 3:\n            subdatasets = self.data.subdatasets\n            self.layers = []\n            if len(subdatasets) > 0:\n                for i, subdataset in enumerate(subdatasets, 0):\n                    self.layers.append(rasterio.open(subdataset))\n        self.shape = self.data.shape\n        self.reduce = reduce\n        self.sz = reduce*sz\n        self.pad0 = (self.sz - self.shape[0]%self.sz)%self.sz\n        self.pad1 = (self.sz - self.shape[1]%self.sz)%self.sz\n        self.n0max = (self.shape[0] + self.pad0)//self.sz\n        self.n1max = (self.shape[1] + self.pad1)//self.sz\n        self.mask = enc2mask(encs,(self.shape[1],self.shape[0])) if encs is not None else None\n        \n    def __len__(self):\n        return self.n0max*self.n1max\n    \n    def __getitem__(self, idx):\n        # the code below may be a little bit difficult to understand,\n        # but the thing it does is mapping the original image to\n        # tiles created with adding padding (like in the previous version of the kernel)\n        # then the tiles are loaded with rasterio\n        # n0,n1 - are the x and y index of the tile (idx = n0*self.n1max + n1)\n        n0,n1 = idx//self.n1max, idx%self.n1max\n        # x0,y0 - are the coordinates of the lower left corner of the tile in the image\n        # negative numbers correspond to padding (which must not be loaded)\n        x0,y0 = -self.pad0//2 + n0*self.sz, -self.pad1//2 + n1*self.sz\n\n        # make sure that the region to read is within the image\n        p00,p01 = max(0,x0), min(x0+self.sz,self.shape[0])\n        p10,p11 = max(0,y0), min(y0+self.sz,self.shape[1])\n        img = np.zeros((self.sz,self.sz,3),np.uint8)\n        mask = np.zeros((self.sz,self.sz),np.uint8)\n        # mapping the loade region to the tile\n        if self.data.count == 3:\n            img[(p00-x0):(p01-x0),(p10-y0):(p11-y0)] = np.moveaxis(self.data.read([1,2,3],\n                window=Window.from_slices((p00,p01),(p10,p11))), 0, -1)\n        else:\n            for i,layer in enumerate(self.layers):\n                img[(p00-x0):(p01-x0),(p10-y0):(p11-y0),i] =\\\n                  layer.read(1,window=Window.from_slices((p00,p01),(p10,p11)))\n        if self.mask is not None: mask[(p00-x0):(p01-x0),(p10-y0):(p11-y0)] = self.mask[p00:p01,p10:p11]\n        \n        if self.reduce != 1:\n            img = cv2.resize(img,(self.sz//reduce,self.sz//reduce),\n                             interpolation = cv2.INTER_AREA)\n            mask = cv2.resize(mask,(self.sz//reduce,self.sz//reduce),\n                             interpolation = cv2.INTER_NEAREST)\n        #check for empty imges\n        hsv = cv2.cvtColor(img, cv2.COLOR_BGR2HSV)\n        h,s,v = cv2.split(hsv)\n        #return -1 for empty images\n        return img, mask, (-1 if (s>s_th).sum() <= p_th or img.sum() <= p_th else idx)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-01T15:44:35.932591Z","iopub.execute_input":"2026-03-01T15:44:35.933135Z","iopub.status.idle":"2026-03-01T15:44:35.946404Z","shell.execute_reply.started":"2026-03-01T15:44:35.933106Z","shell.execute_reply":"2026-03-01T15:44:35.945639Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"x_tot,x2_tot = [],[]\nwith zipfile.ZipFile(OUT_TRAIN, 'w') as img_out,\\\n zipfile.ZipFile(OUT_MASKS, 'w') as mask_out:\n    for index, encs in tqdm(df_masks.iterrows(),total=len(df_masks)):\n        #image+mask dataset\n        ds = HuBMAPDataset(index,encs=encs)\n        for i in range(len(ds)):\n            im,m,idx = ds[i]\n            if idx < 0: continue\n                \n            x_tot.append((im/255.0).reshape(-1,3).mean(0))\n            x2_tot.append(((im/255.0)**2).reshape(-1,3).mean(0))\n            \n            #write data   \n            im = cv2.imencode('.png',cv2.cvtColor(im, cv2.COLOR_RGB2BGR))[1]\n            img_out.writestr(f'{index}_{idx:04d}.png', im)\n            m = cv2.imencode('.png',m)[1]\n            mask_out.writestr(f'{index}_{idx:04d}.png', m)\n        \n#image stats\nimg_avr =  np.array(x_tot).mean(0)\nimg_std =  np.sqrt(np.array(x2_tot).mean(0) - img_avr**2)\nprint('mean:',img_avr, ', std:', img_std)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-01T15:44:38.873244Z","iopub.execute_input":"2026-03-01T15:44:38.874077Z","iopub.status.idle":"2026-03-01T15:51:46.878509Z","shell.execute_reply.started":"2026-03-01T15:44:38.874044Z","shell.execute_reply":"2026-03-01T15:51:46.877874Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"columns, rows = 4,4\nidx0 = 20\nfig=plt.figure(figsize=(columns*4, rows*4))\nwith zipfile.ZipFile(OUT_TRAIN, 'r') as img_arch, \\\n     zipfile.ZipFile(OUT_MASKS, 'r') as msk_arch:\n    fnames = sorted(img_arch.namelist())[8:]\n    for i in range(rows):\n        for j in range(columns):\n            idx = i+j*columns\n            img = cv2.imdecode(np.frombuffer(img_arch.read(fnames[idx0+idx]), \n                                             np.uint8), cv2.IMREAD_COLOR)\n            img = cv2.cvtColor(img, cv2.COLOR_RGB2BGR)\n            mask = cv2.imdecode(np.frombuffer(msk_arch.read(fnames[idx0+idx]), \n                                              np.uint8), cv2.IMREAD_GRAYSCALE)\n    \n            fig.add_subplot(rows, columns, idx+1)\n            plt.axis('off')\n            plt.imshow(Image.fromarray(img))\n            plt.imshow(Image.fromarray(mask), alpha=0.2)\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-01T15:52:05.144166Z","iopub.execute_input":"2026-03-01T15:52:05.144728Z","iopub.status.idle":"2026-03-01T15:52:06.802441Z","shell.execute_reply.started":"2026-03-01T15:52:05.144693Z","shell.execute_reply":"2026-03-01T15:52:06.801390Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\n\nfor dirname, _, filenames in os.walk('/kaggle/working'):\n    for filename in filenames[:5]:\n        print(os.path.join(dirname, filename))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-01T15:52:48.407650Z","iopub.execute_input":"2026-03-01T15:52:48.408509Z","iopub.status.idle":"2026-03-01T15:52:48.413220Z","shell.execute_reply.started":"2026-03-01T15:52:48.408467Z","shell.execute_reply":"2026-03-01T15:52:48.412516Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import zipfile\nimport torch\nimport numpy as np\nimport cv2\nfrom torch.utils.data import Dataset\n\nclass ZipDataset(Dataset):\n    def __init__(self, img_zip, mask_zip):\n        self.img_zip = zipfile.ZipFile(img_zip)\n        self.mask_zip = zipfile.ZipFile(mask_zip)\n        self.fnames = sorted(self.img_zip.namelist())\n\n    def __len__(self):\n        return len(self.fnames)\n\n    def __getitem__(self, idx):\n\n        fname = self.fnames[idx]\n\n        # Load image\n        img = cv2.imdecode(\n            np.frombuffer(self.img_zip.read(fname), np.uint8),\n            cv2.IMREAD_COLOR\n        )\n        img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n        img = img / 255.0\n\n        # Load mask\n        mask = cv2.imdecode(\n            np.frombuffer(self.mask_zip.read(fname), np.uint8),\n            cv2.IMREAD_GRAYSCALE\n        )\n        mask = (mask > 0).astype(np.float32)\n\n        # Convert to tensor\n        img = torch.tensor(img).permute(2, 0, 1).float()\n        mask = torch.tensor(mask).unsqueeze(0).float()\n\n        return img, mask","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-01T15:54:55.718051Z","iopub.execute_input":"2026-03-01T15:54:55.718970Z","iopub.status.idle":"2026-03-01T15:54:55.725425Z","shell.execute_reply.started":"2026-03-01T15:54:55.718938Z","shell.execute_reply":"2026-03-01T15:54:55.724793Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\ndataset = ZipDataset('/kaggle/working/train.zip',\n                     '/kaggle/working/masks.zip',)\n\nloader = torch.utils.data.DataLoader(\n    dataset,\n    batch_size=8,\n    shuffle=True,\n    num_workers=0\n)\n\nprint(\"Total patches:\", len(dataset))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-01T15:55:13.110277Z","iopub.execute_input":"2026-03-01T15:55:13.110876Z","iopub.status.idle":"2026-03-01T15:55:13.192217Z","shell.execute_reply.started":"2026-03-01T15:55:13.110846Z","shell.execute_reply":"2026-03-01T15:55:13.191601Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\nclass DilatedResBlock(nn.Module):\n    def __init__(self, in_channels, out_channels, dilation=1):\n        super(DilatedResBlock, self).__init__()\n\n        self.conv1 = nn.Conv2d(in_channels, out_channels, 3,\n                               padding=dilation, dilation=dilation)\n        self.bn1 = nn.BatchNorm2d(out_channels)\n\n        self.conv2 = nn.Conv2d(out_channels, out_channels, 3,\n                               padding=dilation, dilation=dilation)\n        self.bn2 = nn.BatchNorm2d(out_channels)\n\n        self.skip = nn.Conv2d(in_channels, out_channels, 1)\n\n    def forward(self, x):\n        identity = self.skip(x)\n\n        out = F.relu(self.bn1(self.conv1(x)))\n        out = self.bn2(self.conv2(out))\n\n        out += identity\n        return F.relu(out)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-01T15:55:16.840327Z","iopub.execute_input":"2026-03-01T15:55:16.841169Z","iopub.status.idle":"2026-03-01T15:55:16.847527Z","shell.execute_reply.started":"2026-03-01T15:55:16.841139Z","shell.execute_reply":"2026-03-01T15:55:16.846705Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class TransformerBlock(nn.Module):\n    def __init__(self, dim, num_heads=4):\n        super(TransformerBlock, self).__init__()\n\n        self.norm1 = nn.LayerNorm(dim)\n        self.attn = nn.MultiheadAttention(dim, num_heads)\n        self.norm2 = nn.LayerNorm(dim)\n\n        self.mlp = nn.Sequential(\n            nn.Linear(dim, dim*4),\n            nn.GELU(),\n            nn.Linear(dim*4, dim)\n        )\n\n    def forward(self, x):\n        B, C, H, W = x.shape\n\n        x = x.view(B, C, -1).permute(2, 0, 1)  # (HW, B, C)\n\n        x2 = self.norm1(x)\n        attn_out, _ = self.attn(x2, x2, x2)\n        x = x + attn_out\n\n        x2 = self.norm2(x)\n        x = x + self.mlp(x2)\n\n        x = x.permute(1, 2, 0).view(B, C, H, W)\n        return x","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-01T15:55:19.443117Z","iopub.execute_input":"2026-03-01T15:55:19.443442Z","iopub.status.idle":"2026-03-01T15:55:19.449575Z","shell.execute_reply.started":"2026-03-01T15:55:19.443415Z","shell.execute_reply":"2026-03-01T15:55:19.448882Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class ADCResTransXNet(nn.Module):\n    def __init__(self):\n        super(ADCResTransXNet, self).__init__()\n\n        # Encoder\n        self.enc1 = DilatedResBlock(3, 64, dilation=1)\n        self.pool1 = nn.MaxPool2d(2)\n\n        self.enc2 = DilatedResBlock(64, 128, dilation=2)\n        self.pool2 = nn.MaxPool2d(2)\n\n        self.enc3 = DilatedResBlock(128, 256, dilation=4)\n        self.pool3 = nn.MaxPool2d(2)\n\n        # Transformer\n        self.transformer = TransformerBlock(256)\n\n        # Decoder (FIXED CHANNELS)\n\n        self.up1 = nn.ConvTranspose2d(256, 128, 2, stride=2)\n        self.dec1 = DilatedResBlock(384, 128)   # 128 + 256 = 384\n\n        self.up2 = nn.ConvTranspose2d(128, 64, 2, stride=2)\n        self.dec2 = DilatedResBlock(192, 64)    # 64 + 128 = 192\n\n        self.up3 = nn.ConvTranspose2d(64, 32, 2, stride=2)\n        self.dec3 = DilatedResBlock(96, 32)     # 32 + 64 = 96\n\n        self.final = nn.Conv2d(32, 1, kernel_size=1)\n\n    def forward(self, x):\n\n        e1 = self.enc1(x)\n        p1 = self.pool1(e1)\n\n        e2 = self.enc2(p1)\n        p2 = self.pool2(e2)\n\n        e3 = self.enc3(p2)\n        p3 = self.pool3(e3)\n\n        bottleneck = self.transformer(p3)\n\n        d1 = self.up1(bottleneck)\n        d1 = torch.cat([d1, e3], dim=1)\n        d1 = self.dec1(d1)\n\n        d2 = self.up2(d1)\n        d2 = torch.cat([d2, e2], dim=1)\n        d2 = self.dec2(d2)\n\n        d3 = self.up3(d2)\n        d3 = torch.cat([d3, e1], dim=1)\n        d3 = self.dec3(d3)\n\n        return self.final(d3)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-01T15:55:21.716136Z","iopub.execute_input":"2026-03-01T15:55:21.716815Z","iopub.status.idle":"2026-03-01T15:55:21.725140Z","shell.execute_reply.started":"2026-03-01T15:55:21.716786Z","shell.execute_reply":"2026-03-01T15:55:21.724519Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model = ADCResTransXNet().cuda()\noptimizer = torch.optim.Adam(model.parameters(), lr=1e-5)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-01T15:55:25.034154Z","iopub.execute_input":"2026-03-01T15:55:25.034489Z","iopub.status.idle":"2026-03-01T15:55:29.271945Z","shell.execute_reply.started":"2026-03-01T15:55:25.034459Z","shell.execute_reply":"2026-03-01T15:55:29.271378Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"bce = nn.BCEWithLogitsLoss()\n\ndef dice_loss(pred, target):\n    smooth = 1e-8\n    intersection = (pred * target).sum()\n    return 1 - (2 * intersection + smooth) / (pred.sum() + target.sum() + smooth)\n\ndef combined_loss(output, target):\n    pred = torch.sigmoid(output)\n    return bce(output, target) + dice_loss(pred, target)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-01T15:58:16.707061Z","iopub.execute_input":"2026-03-01T15:58:16.707381Z","iopub.status.idle":"2026-03-01T15:58:16.712200Z","shell.execute_reply.started":"2026-03-01T15:58:16.707353Z","shell.execute_reply":"2026-03-01T15:58:16.711592Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for epoch in range(20):\n    model.train()\n    total_loss = 0\n\n    for imgs_batch, masks_batch in loader:\n        imgs_batch = imgs_batch.cuda()\n        masks_batch = masks_batch.cuda()\n\n        outputs = model(imgs_batch)\n        loss = combined_loss(outputs, masks_batch)\n\n        optimizer.zero_grad()\n        loss.backward()\n        optimizer.step()\n\n        total_loss += loss.item()\n\n    print(f\"Epoch {epoch+1}, Loss: {total_loss/len(loader)}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-01T15:58:33.319018Z","iopub.execute_input":"2026-03-01T15:58:33.319811Z","iopub.status.idle":"2026-03-01T17:20:32.729759Z","shell.execute_reply.started":"2026-03-01T15:58:33.319781Z","shell.execute_reply":"2026-03-01T17:20:32.729032Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"torch.save(model.state_dict(), \"/kaggle/working/adc_res_transxnet_seg.pth\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-28T15:49:01.079175Z","iopub.execute_input":"2026-02-28T15:49:01.079404Z","iopub.status.idle":"2026-02-28T15:49:01.117244Z","shell.execute_reply.started":"2026-02-28T15:49:01.079382Z","shell.execute_reply":"2026-02-28T15:49:01.116589Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\n\ndef compute_metrics(model, loader, threshold):\n    model.eval()\n    \n    total_dice = 0\n    total_iou = 0\n    total_acc = 0\n    n_batches = 0\n    \n    with torch.no_grad():\n        for imgs, masks in loader:\n            imgs = imgs.cuda()\n            masks = masks.cuda()\n            \n            outputs = model(imgs)\n            probs = torch.sigmoid(outputs)\n            preds = (probs > threshold).float()\n            \n            # Flatten\n            preds_flat = preds.view(-1)\n            masks_flat = masks.view(-1)\n            \n            intersection = (preds_flat * masks_flat).sum()\n            union = preds_flat.sum() + masks_flat.sum()\n            \n            # Dice\n            dice = (2 * intersection) / (union + 1e-8)\n            \n            # IoU\n            union_iou = ((preds_flat + masks_flat) > 0).float().sum()\n            iou = intersection / (union_iou + 1e-8)\n            \n            # Accuracy\n            correct = (preds_flat == masks_flat).sum()\n            acc = correct / masks_flat.numel()\n            \n            total_dice += dice.item()\n            total_iou += iou.item()\n            total_acc += acc.item()\n            n_batches += 1\n    \n    print(\"Dice:\", total_dice / n_batches)\n    print(\"IoU:\", total_iou / n_batches)\n    print(\"Pixel Accuracy:\", total_acc / n_batches)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-01T17:21:37.180728Z","iopub.execute_input":"2026-03-01T17:21:37.181436Z","iopub.status.idle":"2026-03-01T17:21:37.187761Z","shell.execute_reply.started":"2026-03-01T17:21:37.181408Z","shell.execute_reply":"2026-03-01T17:21:37.187200Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"compute_metrics(model, loader, threshold=0.5)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-01T17:21:47.412717Z","iopub.execute_input":"2026-03-01T17:21:47.413275Z","iopub.status.idle":"2026-03-01T17:23:26.114804Z","shell.execute_reply.started":"2026-03-01T17:21:47.413248Z","shell.execute_reply":"2026-03-01T17:23:26.114070Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"torch.save(model.state_dict(), \"/kaggle/working/adc_res_transxnet_seg.pth\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-01T17:24:30.989709Z","iopub.execute_input":"2026-03-01T17:24:30.990072Z","iopub.status.idle":"2026-03-01T17:24:31.047589Z","shell.execute_reply.started":"2026-03-01T17:24:30.990046Z","shell.execute_reply":"2026-03-01T17:24:31.046841Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}