{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":113558,"databundleVersionId":14456136,"isSourceIdPinned":false,"sourceType":"competition"},{"sourceId":3414836,"sourceType":"datasetVersion","datasetId":2058261},{"sourceId":13636007,"sourceType":"datasetVersion","datasetId":8667355},{"sourceId":649836,"sourceType":"modelInstanceVersion","isSourceIdPinned":false,"modelInstanceId":490265,"modelId":505697}],"dockerImageVersionId":31193,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"\"\"\"\nMethod origin:\n'Copy-Move Detection in Optical Microscopy: A Segmentation Network and A Dataset'\n\nCode modified: \ntake more intermediate feature vectors for trying to find the copy-moved areas.\n\nSimilarity check:\n'Learned Perceptual image patch similarity' from lightning AI has been used for checking \nthe similarity between the cut-out patches\n\nDataset:\nseparate the data into categorical groups, in this notebook, it is named 'tissue'.\n\nConclusion:\n- Small copy-moved cells will likely be missed\n- Similarity check doesn't perform well even when forged images have been provided for \n  comparison\n\n\n\"\"\"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"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":"2025-11-20T10:15:57.093232Z","iopub.execute_input":"2025-11-20T10:15:57.094263Z","iopub.status.idle":"2025-11-20T10:16:07.432093Z","shell.execute_reply.started":"2025-11-20T10:15:57.094225Z","shell.execute_reply":"2025-11-20T10:16:07.43123Z"},"collapsed":true,"jupyter":{"outputs_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torchvision.models as models\nimport numpy as np\nimport torch\nimport torch.nn as nn\n\nfrom torch.utils.data import Dataset, DataLoader\n\nimport numpy as np\nimport math\nimport os\nimport cv2","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T10:16:07.433391Z","iopub.execute_input":"2025-11-20T10:16:07.433788Z","iopub.status.idle":"2025-11-20T10:16:16.242412Z","shell.execute_reply.started":"2025-11-20T10:16:07.433766Z","shell.execute_reply":"2025-11-20T10:16:16.24167Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"device = \"cuda\" if torch.cuda.is_available() else \"cpu\"\ndevice","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T10:16:16.243238Z","iopub.execute_input":"2025-11-20T10:16:16.243672Z","iopub.status.idle":"2025-11-20T10:16:16.25322Z","shell.execute_reply.started":"2025-11-20T10:16:16.243648Z","shell.execute_reply":"2025-11-20T10:16:16.252291Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Model","metadata":{}},{"cell_type":"code","source":"# -*- coding: utf-8 -*-\n\"\"\"\nCreated on Tue May 14 22:19:34 2024\n\n@author: user\n\"\"\"\n\nclass ZeroWindow:\n    def __init__(self):\n        self.store = {}\n\n    def __call__(self, x_in, h, w, rat_s=0.1):\n        sigma = h * rat_s, w * rat_s\n        #print(sigma)\n        #sigma = 32,32\n        # c = h * w\n        b, c, h2, w2 = x_in.shape\n        key = str(x_in.shape) + str(rat_s)\n        if key not in self.store:\n            ind_r = torch.arange(h2).float()\n            ind_c = torch.arange(w2).float()\n            ind_r = ind_r.view(1, 1, -1, 1).expand_as(x_in)\n            ind_c = ind_c.view(1, 1, 1, -1).expand_as(x_in)\n\n            # center\n            c_indices = torch.from_numpy(np.indices((h, w))).float()\n            c_ind_r = c_indices[0].reshape(-1)\n            c_ind_c = c_indices[1].reshape(-1)\n\n            cent_r = c_ind_r.reshape(1, c, 1, 1).expand_as(x_in)\n            cent_c = c_ind_c.reshape(1, c, 1, 1).expand_as(x_in)\n\n            def fn_gauss(x, u, s):\n                return torch.exp(-(x - u) ** 2 / (2 * s ** 2))\n            #print(sigma)\n            gaus_r = fn_gauss(ind_r, cent_r, sigma[0])\n            gaus_c = fn_gauss(ind_c, cent_c, sigma[1])\n            out_g = 1 - gaus_r * gaus_c\n            out_g = out_g.to(x_in.device)\n            self.store[key] = out_g\n        else:\n            out_g = self.store[key]\n        out = out_g * x_in\n        return out\n\ndef get_topk(x, k=10, dim=-3):\n    # b, c, h, w = x.shape\n    val, _ = torch.topk(x, k=k, dim=dim)\n    return val\n\nclass Corr(nn.Module):\n    def __init__(self, topk=3):\n        super().__init__()\n        self.topk = topk\n        self.zero_window = ZeroWindow()\n        self.alpha = nn.Parameter(torch.tensor(5., dtype=torch.float32))\n\n    def forward(self, x):\n        b, c, h1, w1 = x.shape\n        h2 = h1\n        w2 = w1\n\n        xn = F.normalize(x, p=2, dim=-3)\n        x_aff_o = torch.matmul(xn.permute(0, 2, 3, 1).view(b, -1, c), xn.view(b, c, -1))\n\n        x_aff = self.zero_window(x_aff_o.view(b, -1, h1, w1), h1, w1, rat_s=0.05).reshape(b, h1 * w1, h2 * w2)\n        x_c = F.softmax(x_aff * self.alpha, dim=-1) * F.softmax(x_aff * self.alpha, dim=-2)\n        x_c = x_c.reshape(b, h1, w1, h2, w2)\n\n        xc_o = x_c.view(b, h1 * w1, h2, w2)\n        val = get_topk(xc_o, k=self.topk, dim=-3)\n\n        return val\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T10:16:16.25543Z","iopub.execute_input":"2025-11-20T10:16:16.255685Z","iopub.status.idle":"2025-11-20T10:16:16.283208Z","shell.execute_reply.started":"2025-11-20T10:16:16.255664Z","shell.execute_reply":"2025-11-20T10:16:16.282293Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class MobileNetV2(nn.Module):\n\n    \"\"\"\n    from MobileNetV2 import MobileNetV2\n\n    net = MobileNetV2(n_class=1000)\n    state_dict = torch.load('mobilenetv2.pth.tar') # add map_location='cpu' if no gpu\n    net.load_state_dict(state_dict)\n    \"\"\"\n\n    def __init__(self, n_class=1000, input_size=512, width_mult=1.):\n        super(MobileNetV2, self).__init__()\n        block = InvertedResidual\n        input_channel = 32\n        last_channel = 1280\n        interverted_residual_setting = [\n            # t, c, n, s\n            [1, 16, 1, 1],\n            [6, 24, 2, 2],\n            [6, 32, 3, 2],\n            [6, 64, 4, 2],\n            [6, 96, 3, 1],\n            [6, 160, 3, 2],\n            [6, 320, 1, 1],\n        ]\n\n        # building first layer\n        assert input_size % 32 == 0\n        input_channel = int(input_channel * width_mult)\n        self.last_channel = int(last_channel * width_mult) if width_mult > 1.0 else last_channel\n        self.features = [conv_bn(3, input_channel, 2)]\n        # building inverted residual blocks\n        for t, c, n, s in interverted_residual_setting:\n            output_channel = int(c * width_mult)\n            for i in range(n):\n                if i == 0:\n                    self.features.append(block(input_channel, output_channel, s, expand_ratio=t))\n                else:\n                    self.features.append(block(input_channel, output_channel, 1, expand_ratio=t))\n                input_channel = output_channel\n        # building last several layers\n        self.features.append(conv_1x1_bn(input_channel, self.last_channel))\n        # make it nn.Sequential\n        self.features = nn.Sequential(*self.features)\n\n        # building classifier\n        self.classifier = nn.Sequential(\n            nn.Dropout(0.2),\n            nn.Linear(self.last_channel, n_class),\n        )\n\n        self._initialize_weights()\n\n    def forward(self, x):\n        x = self.features(x)\n        \n        #x = x.mean(3).mean(2)\n        #x = self.classifier(x)\n        \n        return x\n\n    def _initialize_weights(self):\n        for m in self.modules():\n            if isinstance(m, nn.Conv2d):\n                n = m.kernel_size[0] * m.kernel_size[1] * m.out_channels\n                m.weight.data.normal_(0, np.sqrt(2. / n))\n                if m.bias is not None:\n                    m.bias.data.zero_()\n            elif isinstance(m, nn.BatchNorm2d):\n                m.weight.data.fill_(1)\n                m.bias.data.zero_()\n            elif isinstance(m, nn.Linear):\n                n = m.weight.size(1)\n                m.weight.data.normal_(0, 0.01)\n                m.bias.data.zero_()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T10:16:16.284134Z","iopub.execute_input":"2025-11-20T10:16:16.28442Z","iopub.status.idle":"2025-11-20T10:16:16.304197Z","shell.execute_reply.started":"2025-11-20T10:16:16.284391Z","shell.execute_reply":"2025-11-20T10:16:16.303427Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class SpatialAttention(nn.Module):\n    def __init__(self, ):\n        super(SpatialAttention, self).__init__()\n\n        self.conv1 = nn.Conv2d(2, 1, 7, padding=3, bias=False)\n\n    def forward(self, x):\n        input = x\n        avg_out = torch.mean(x, dim=1, keepdim=True)\n        max_out, _ = torch.max(x, dim=1, keepdim=True)\n        x = torch.cat([avg_out, max_out], dim=1)\n        x = self.conv1(x)\n        y = torch.sigmoid(x)\n        return input + input * y        \nclass InvertedResidual(nn.Module):\n    def __init__(self, inp, oup, stride, expand_ratio):\n        super(InvertedResidual, self).__init__()\n        self.stride = stride\n        assert stride in [1, 2]\n\n        hidden_dim = round(inp * expand_ratio)\n        self.use_res_connect = self.stride == 1 and inp == oup\n\n        if expand_ratio == 1:\n            self.conv = nn.Sequential(\n                # dw\n                nn.Conv2d(hidden_dim, hidden_dim, 3, stride, 1, groups=hidden_dim, bias=False),\n                nn.BatchNorm2d(hidden_dim),\n                nn.ReLU6(inplace=True),\n                # pw-linear\n                nn.Conv2d(hidden_dim, oup, 1, 1, 0, bias=False),\n                nn.BatchNorm2d(oup),\n            )\n        else:\n            self.conv = nn.Sequential(\n                # pw\n                nn.Conv2d(inp, hidden_dim, 1, 1, 0, bias=False),\n                nn.BatchNorm2d(hidden_dim),\n                nn.ReLU6(inplace=True),\n                # dw\n                nn.Conv2d(hidden_dim, hidden_dim, 3, stride, 1, groups=hidden_dim, bias=False),\n                nn.BatchNorm2d(hidden_dim),\n                nn.ReLU6(inplace=True),\n                # pw-linear\n                nn.Conv2d(hidden_dim, oup, 1, 1, 0, bias=False),\n                nn.BatchNorm2d(oup),\n            )\n\n    def forward(self, x):\n        if self.use_res_connect:\n            return x + self.conv(x)\n        else:\n            return self.conv(x)\n\ndef conv_bn(inp, oup, stride):\n    return nn.Sequential(\n        nn.Conv2d(inp, oup, 3, stride, 1, bias=False),\n        nn.BatchNorm2d(oup),\n        nn.ReLU6(inplace=True)\n    )\n\n\ndef conv_1x1_bn(inp, oup):\n    return nn.Sequential(\n        nn.Conv2d(inp, oup, 1, 1, 0, bias=False),\n        nn.BatchNorm2d(oup),\n        nn.ReLU6(inplace=True)\n    )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T10:16:16.305139Z","iopub.execute_input":"2025-11-20T10:16:16.305449Z","iopub.status.idle":"2025-11-20T10:16:16.324445Z","shell.execute_reply.started":"2025-11-20T10:16:16.305422Z","shell.execute_reply":"2025-11-20T10:16:16.323524Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#Model\nclass UnetMobilenetV2(nn.Module):\n    def __init__(self, num_classes=1, num_filters=32, pretrained=False,\n                 Dropout=.2, ): #path=os.path.join(r'mobilenet_v2.pth.tar')\n        super(UnetMobilenetV2, self).__init__()\n        self.encoder = MobileNetV2(n_class=1000)\n\n        self.num_classes = num_classes\n\n        self.dconv1 = nn.ConvTranspose2d(1280, 96, 4, padding=1, stride=2)\n        self.invres1 = InvertedResidual(192, 96, 1, 6)\n\n        self.dconv2 = nn.ConvTranspose2d(96, 32, 4, padding=1, stride=2)\n        self.invres2 = InvertedResidual(64, 32, 1, 6)\n\n        self.dconv3 = nn.ConvTranspose2d(32, 24, 4, padding=1, stride=2)\n        self.invres3 = InvertedResidual(48, 24, 1, 6)\n\n        self.dconv4 = nn.ConvTranspose2d(24, 16, 4, stride=4, padding=0)\n        self.invres4 = InvertedResidual(32, 16, 1, 6)\n        \n        self.dconv5 = nn.ConvTranspose2d(16, 3, 4, padding=1, stride=2)\n        self.invres5 = InvertedResidual(6, 3, 1, 6)\n\n        self.dconv_extra_in = nn.ConvTranspose2d(96, 64, 3, padding=1, stride = 1)\n        self.invres_extra_in = InvertedResidual(128, 64, 1, 6)\n\n        self.dconv_extra_out = nn.ConvTranspose2d(64, 32, 4, padding = 1, stride = 2)\n        self.invres_extra_out = InvertedResidual(64, 32, 1, 6)\n        \n        self.trans = nn.ConvTranspose2d(in_channels=16, out_channels=16, kernel_size=4, stride=2, padding=1)\n        \n        self.conv_last = nn.Conv2d(16, 3, 1)\n\n        self.conv_score = nn.Conv2d(3, 1, 1)\n\n        # Refinement module to process the initial heatmap\n        self.refine = nn.Sequential(\n            nn.Conv2d(16, 32, 3, padding=1), # Initial convolution\n            nn.BatchNorm2d(32), # Batch normalization\n            nn.ReLU(inplace=True), # ReLU activation\n            nn.Conv2d(32, 16, 3, padding=1), # Second convolution\n            nn.BatchNorm2d(16), # Batch normalization\n            nn.ReLU(inplace=True), # ReLU activation\n            nn.Conv2d(16, 1, 1), # Final 1x1 convolution to output a single channel\n            # Removed Sigmoid here as BCEWithLogitsLoss is used\n        )\n\n        \n        #doesn't needed; obly for compatibility\n        self.dconv_final = nn.ConvTranspose2d(1, 1, 4, padding=1, stride=2)\n\n        \n        if pretrained:\n            state_dict = torch.load(path)\n            self.encoder.load_state_dict(state_dict)\n        else: self.encoder._initialize_weights()\n\n        #############################\n\n        self.corr16 = Corr(topk=16)\n\n        self.corr24 = Corr(topk=24)\n\n        self.corr32 = Corr(topk=32)\n\n        self.corr96 = Corr(topk=96)\n\n        self.corr64 = Corr(topk=64)  # extra\n\n        self.corr_up1 = Corr(topk=96) # last\n        \n        self.aspp1 = models.segmentation.deeplabv3.ASPP(in_channels=96, out_channels=96,atrous_rates=[4, 8, 12,16])\n        self.aspp2 = models.segmentation.deeplabv3.ASPP(in_channels=32,out_channels=32, atrous_rates=[4, 8, 12,16])\n        self.aspp3 = models.segmentation.deeplabv3.ASPP(in_channels=24,out_channels=24, atrous_rates=[4, 8, 12,16])\n        self.aspp4 = models.segmentation.deeplabv3.ASPP(in_channels=16,out_channels=16, atrous_rates=[4, 8, 12,16])\n\n        self.aspp_extra = models.segmentation.deeplabv3.ASPP(in_channels=64,out_channels=64, atrous_rates=[4, 8, 12,16])\n        self.aspp_up1 = models.segmentation.deeplabv3.ASPP(in_channels=96,out_channels=96, atrous_rates=[4, 8, 12,16])\n        \n        \n        self.sam1 = SpatialAttention()\n        self.sam2 = SpatialAttention()\n        self.sam3 = SpatialAttention()\n        self.sam4 = SpatialAttention()\n        #self.sam5 = SpatialAttention()\n        self.sam_extra = SpatialAttention()\n        self.sam_up1 = SpatialAttention()\n        #############################\n\n    def forward(self, x):\n        #print('x:',x.shape)\n        # correlation matrix between the pixels of the original image\n\n        \n        for n in range(0, 2):\n            x = self.encoder.features[n](x)\n        x1 = x\n        # x1 = self.corr16(x1)\n        #print(\"x1\",x1.shape)\n        x1 = self.aspp4(x1)\n        #print(\"x1_aspp4\",x1.shape)\n        x1 = self.sam1(x1)\n        \n        #print(\"x1_sam1\",x1.shape)\n        #print('x1:',x.shape)\n        \n        for n in range(2, 4):\n            x = self.encoder.features[n](x)\n        x2 = x\n        #print(\"x2\",x2.shape)\n        x2 = self.corr24(x2)\n        #print(\"x2_corrr\",x2.shape)\n        x2 = self.aspp3(x2)\n        #print(\"x2_aspp3\",x2.shape)\n        x2 = self.sam2(x2)\n        \n        #print(\"x2_sam2\",x2.shape)\n        #print('x2:',x.shape)\n\n        for n in range(4, 7):\n            x = self.encoder.features[n](x)\n            # print(f'n = {n}, x.shape = {x.shape}')\n        x3 = x\n        #print(\"x3\",x3.shape)\n        x3 = self.corr32(x3)\n        #print(\"x3_corrr\",x3.shape)\n        x3 = self.aspp2(x3)\n        #print(\"x3_aspp2\",x3.shape)\n        x3 = self.sam3(x3)\n        \n        # print(\"x3_sam3\",x3.shape)\n        #print('x3:',x.shape)\n\n        \"added part\"\n        for n in range(7, 11):\n            x = self.encoder.features[n](x)\n            # print(f'n = {n}, x.shape = {x.shape}')\n        \n        x_extra = x\n        x_extra = self.corr64(x_extra)\n        x_extra = self.aspp_extra(x_extra)\n        x_extra = self.sam_extra(x_extra)\n        # print(\"x_extra_sam\", x_extra.shape)\n        \"end of added part\"\n        for n in range(11, 14):\n            x = self.encoder.features[n](x)\n            # print(f'n = {n}, x.shape = {x.shape}')\n        x4 = x\n        #print(\"x4\",x4.shape)\n        x4 = self.corr96(x4)\n        #print(\"x4_corrr\",x4.shape)\n        x4 = self.aspp1(x4)\n        #print(\"x4_aspp1\",x4.shape)\n        x4 = self.sam4(x4)\n        \n        # print(\"x4_sam4\",x4.shape)\n\n        for n in range(14, 19):\n            x = self.encoder.features[n](x)\n        x5=x\n        # x5= self.corr96(x5)\n        # print('x5:',x.shape)\n\n        up1 = torch.cat([\n            x4,\n            self.dconv1(x5)\n        ], dim=1)\n        # print('up1:',up1.shape)\n        up1 = self.invres1(up1)\n        # print('up1_invres:',up1.shape)\n\n        up1 = self.corr_up1(up1)\n        up1 = self.aspp_up1(up1)\n        up1 = self.sam_up1(up1)\n        # print('up1:',up1.shape)\n\n        \"added part\"\n        # print(\"deconv_extra_in:\", self.dconv_extra_in(up1).shape)\n        up_extra_1 = torch.cat([\n            x_extra,\n            self.dconv_extra_in(up1)\n        ], dim = 1)\n        # print(\"up_extra_dconv:\", up_extra_1.shape)\n        up_extra_1 = self.invres_extra_in(up_extra_1)\n        # print('up_extra_invres:', up_extra_1.shape)\n\n\n        up_extra_2 = torch.cat([\n            x3,\n            self.dconv_extra_out(up_extra_1)\n        ],dim = 1)\n        # print(\"up_extra_2:\", up_extra_2.shape)\n        up_extra_2 = self.invres_extra_out(up_extra_2)\n        # print(\"up_extra_invres_2:\", up_extra_2.shape)\n        \n        # up2 = torch.cat([\n        #     x3,\n        #     self.dconv2(up1)\n        # ], dim=1)\n        # print('up2:',up2.shape)\n        # up2 = self.invres2(up2)\n        # print('up2_invres:',up2.shape)\n\n        up3 = torch.cat([\n            x2,\n            self.dconv3(up_extra_2)\n        ], dim=1)\n        # print('up3:',up3.shape)\n        up3 = self.invres3(up3) \n        # print('up3_invres:',up3.shape)\n        \n        # up3 = self.corr_up3(up3)\n        # up3 = self.aspp_up3(up3)\n        # up3 = self.sam_up3(up3)\n        # print('up3:',up3.shape)\n        \"end of added part\"\n        \n        up4 = torch.cat([\n            self.trans(x1),\n            self.dconv4(up3)\n        ], dim=1)\n        # print('up4:',up4.shape)\n        up4 = self.invres4(up4)\n        # print('up4_invres:',up4.shape)\n        \n        # x = self.conv_last(up4)\n        #print('x_last',x.shape)\n        \n        # x = self.aspp_last(x)\n        # x = self.sam_last(x)\n        \n        # x = self.conv_score(x)\n        \n        #print('x_score',x.shape)\n\n        x = self.refine(up4)\n        \n        return x","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T10:16:16.325468Z","iopub.execute_input":"2025-11-20T10:16:16.325805Z","iopub.status.idle":"2025-11-20T10:16:16.349061Z","shell.execute_reply.started":"2025-11-20T10:16:16.325774Z","shell.execute_reply":"2025-11-20T10:16:16.348336Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n\n# # Initialize the modified model\n# model = UnetMobilenetV2(pretrained=False)\n\n# # Generate a random input tensor with shape (1, 3, 256, 384)\n# input_tensor = torch.randn(2, 3, 256, 256)\n\n# # Perform a forward pass through the model\n# output = model(input_tensor)\n\n# # Print the shape of the output tensor\n# output_shape = output.shape\n# output_shape\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T10:16:16.349932Z","iopub.execute_input":"2025-11-20T10:16:16.350725Z","iopub.status.idle":"2025-11-20T10:16:16.365863Z","shell.execute_reply.started":"2025-11-20T10:16:16.350704Z","shell.execute_reply":"2025-11-20T10:16:16.365073Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Losses","metadata":{}},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\n\n\n\nclass SoftDiceLoss(nn.Module):\n    def __init__(self):\n        super(SoftDiceLoss, self).__init__()\n\n\n    def forward(self, logits, targets):\n        eps = 1e-5 \n        bs = targets.size(0)\n\n        probs = torch.sigmoid(logits) \n        m1 = probs.view(bs, -1)\n        m2 = targets.view(bs, -1)\n\n        intersection = (m1 * m2).sum(dim=1)\n        union = m1.sum(dim=1) * 10 + m2.sum(dim=1) + eps\n\n        score = (2. * intersection + eps)/ union\n        loss = 1 - score.mean() \n\n        return loss\n\n\nclass MixedLoss(nn.Module):\n    def __init__(self, dice_weight=0.5):\n        super(MixedLoss, self).__init__()\n        self.dice_weight = dice_weight\n        self.bce = nn.BCEWithLogitsLoss()\n        # self.soft_dice = SoftDiceLoss()\n        \n\n    def forward(self, y_pred, y_true):\n        bce_loss = self.bce(y_pred, y_true)\n        # dice_loss = self.soft_dice(y_pred, y_true)\n        dice_loss = self.dice_loss(y_pred, y_true)\n        \n        total_loss = bce_loss + self.dice_weight * dice_loss\n        return total_loss, dice_loss, bce_loss\n\n    def dice_loss(self, pred, target, smooth=1.0):\n        # Apply sigmoid to predictions for Dice loss since the model output is now logits\n        pred_sigmoid = torch.sigmoid(pred)\n        pred_flat = pred_sigmoid.contiguous().view(-1)\n        target_flat = target.contiguous().view(-1)\n        intersection = (pred_flat * target_flat).sum()\n        dice = (2. * intersection + smooth) / (pred_flat.sum() + target_flat.sum() + smooth)\n        # Return the Dice loss\n        return 1 - dice","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T10:16:16.366732Z","iopub.execute_input":"2025-11-20T10:16:16.366973Z","iopub.status.idle":"2025-11-20T10:16:16.387642Z","shell.execute_reply.started":"2025-11-20T10:16:16.366949Z","shell.execute_reply":"2025-11-20T10:16:16.386675Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nfrom PIL import Image\nimport numpy as np\nimport matplotlib.pyplot as plt","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T10:16:16.390689Z","iopub.execute_input":"2025-11-20T10:16:16.391043Z","iopub.status.idle":"2025-11-20T10:16:16.403433Z","shell.execute_reply.started":"2025-11-20T10:16:16.391022Z","shell.execute_reply":"2025-11-20T10:16:16.402672Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Transformers","metadata":{}},{"cell_type":"code","source":"import random\n\nclass RandomNoise:\n    def __init__(self, noise_level=0.05, p=0.1):\n        self.noise_level = noise_level\n        self.p = p  \n\n    def __call__(self, img_tensor):\n        if random.random() < self.p:  \n            with torch.no_grad(): \n                \n                noise = torch.randn_like(img_tensor) * self.noise_level\n          \n                noisy_img = img_tensor + noise\n                \n                \n                noisy_img = torch.clamp(noisy_img, 0, 1)\n            return noisy_img\n        else:\n            return img_tensor ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T10:16:16.404368Z","iopub.execute_input":"2025-11-20T10:16:16.405108Z","iopub.status.idle":"2025-11-20T10:16:16.419374Z","shell.execute_reply.started":"2025-11-20T10:16:16.405078Z","shell.execute_reply":"2025-11-20T10:16:16.418344Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from torchvision.transforms import v2\n\n\ninput_size = 256\n\ntrain_transform = v2.Compose([\n        v2.RandomVerticalFlip(),\n        v2.RandomHorizontalFlip(),\n        # v2.Rotate(limit=30, p=0.5),\n        # v2.ShiftScaleRotate(shift_limit=0.05, scale_limit=0.05, rotate_limit=0, p=0.2),\n    ])\nval_transform = v2.Compose([\n        v2.ToPILImage(),\n        v2.Resize((input_size, input_size)),\n        v2.ToDtype(torch.float32, scale = True),\n        \n    ])\n\n\naugment = v2.Compose([\n    RandomNoise(0.1,0.05),\n    v2.ColorJitter(brightness=.5, hue=.3),\n    # v2.GaussianBlur(kernel_size=(5, 9), sigma=(0.1, 5.))\n    ])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T10:16:16.420298Z","iopub.execute_input":"2025-11-20T10:16:16.420537Z","iopub.status.idle":"2025-11-20T10:16:16.549193Z","shell.execute_reply.started":"2025-11-20T10:16:16.420519Z","shell.execute_reply":"2025-11-20T10:16:16.548287Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# DataSet","metadata":{}},{"cell_type":"code","source":"import torchvision.transforms.functional as tF\n\nclass ForgeryDataset(Dataset):\n    def __init__(self, authentic_path, forged_path, masks_path, transform=None, \n                 augment = None, is_train=True, img_size = 256):\n        \n        self.transform = transform\n        self.augment = augment\n        \n        self.to_tensor = v2.Compose([\n            \n            v2.ToImage(),\n            v2.Resize((img_size, img_size)),\n            v2.ToDtype(torch.float32, scale = True)\n        ])\n\n        self.to_tensor_mask = v2.Compose([\n            \n            v2.ToImage(),\n            v2.Resize((img_size, img_size)),\n            v2.ToDtype(torch.float32, scale = False)\n        ])\n        \n        self.is_train = is_train\n       \n        # Collect all data samples\n        self.samples = []\n        \n        # Authentic images\n        # for file in os.listdir(authentic_path):\n        #     img_path = os.path.join(authentic_path, file)\n        #     base_name = file.split('.')[0]\n            \n            \n        #     self.samples.append({\n        #         'image_path': img_path,\n                \n        #         'is_forged': False,\n        #         'image_id': base_name\n        #     })\n        \n        # Forged images\n        for file in os.listdir(forged_path):\n            img_path = os.path.join(forged_path, file)\n            base_name = file.split('.')[0]\n            mask_path = os.path.join(masks_path, f\"{base_name}.npy\")\n            \n            self.samples.append({\n                'image_path': img_path,\n                'mask_path': mask_path,\n                'is_forged': True,\n                'image_id': base_name\n            })\n\n            \n    def __len__(self):\n        return len(self.samples)\n    \n    def __getitem__(self, idx):\n\n        sample = self.samples[idx]\n        # Load image\n        image = Image.open(sample['image_path']).convert('RGB')\n        image = np.array(image)\n        \n        # Load and process mask\n        if  sample['is_forged']:\n            mask = np.load(sample['mask_path'])\n            \n            # Handle multi-channel masks\n            if mask.ndim == 3:\n                if mask.shape[0] <= 10:  # channels first (C, H, W)\n                    mask = np.any(mask, axis=0)\n                elif mask.shape[-1] <= 10:  # channels last (H, W, C)\n                    mask = np.any(mask, axis=-1)\n                else:\n                    raise ValueError(f\"Ambiguous 3D mask shape: {mask.shape}\")\n            \n            mask = (mask > 0).astype(np.uint8)\n        else:\n            mask = np.zeros_like(image[:, :, 0], dtype=np.uint8)\n    \n        # Shape validation\n        assert image.shape[:2] == mask.shape, f\"Shape mismatch: img {image.shape},mask {mask.shape}, is_forged ({sample['is_forged']}),base_name {sample['image_id']}\"\n\n        \n        image, mask = self.to_tensor(image), self.to_tensor_mask(mask)\n        \n        if self.transform:\n        \n            # Apply transformations\n            combined = torch.cat((image, mask), dim = 0)\n            combined = self.transform(combined)\n    \n            image = combined[:3, :, :]\n            mask = combined[3:, :, :]\n\n        if self.augment:\n            \n            # image = tF.adjust_contrast(image, contrast_factor=1.5)\n            image = self.augment(image)\n           \n        return image, mask\n    ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T10:16:16.550258Z","iopub.execute_input":"2025-11-20T10:16:16.550693Z","iopub.status.idle":"2025-11-20T10:16:16.563801Z","shell.execute_reply.started":"2025-11-20T10:16:16.550664Z","shell.execute_reply":"2025-11-20T10:16:16.563045Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Create Dataset","metadata":{}},{"cell_type":"code","source":"img_size = 256\n# augment = None\n\nbase_path = '/kaggle/input/recodai-luc-scientific-image-forgery-detection'\nauthentic_path = os.path.join(base_path, 'train_images/authentic')\nforged_path = os.path.join(base_path, 'train_images/forged')\nmasks_path = os.path.join(base_path, 'train_masks')\n\ncorn_path = '/kaggle/input/'\nforged_tissue_path = os.path.join(corn_path, 'tissue/tissue')\n\n\ntissue_dataset = ForgeryDataset(authentic_path, forged_tissue_path, masks_path, \n                              transform = train_transform, augment = augment,\n                             img_size = img_size)\n\ntrain_size = int(0.9 * len(tissue_dataset))\nval_size = len(tissue_dataset) - train_size\ntrain_dataset, val_dataset = torch.utils.data.random_split(tissue_dataset, [train_size, val_size])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T10:16:16.564612Z","iopub.execute_input":"2025-11-20T10:16:16.56486Z","iopub.status.idle":"2025-11-20T10:16:16.582061Z","shell.execute_reply.started":"2025-11-20T10:16:16.564842Z","shell.execute_reply":"2025-11-20T10:16:16.581308Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# DataLoader","metadata":{}},{"cell_type":"code","source":"BATCH_SIZE = 16\n\ntrain_dataloader = DataLoader(train_dataset, batch_size = BATCH_SIZE, \n                                shuffle = True, num_workers = 4, pin_memory=True)\nval_dataloader = DataLoader(val_dataset, batch_size = BATCH_SIZE, \n                           shuffle = True, num_workers = 4, pin_memory=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T10:16:16.58296Z","iopub.execute_input":"2025-11-20T10:16:16.583252Z","iopub.status.idle":"2025-11-20T10:16:16.590777Z","shell.execute_reply.started":"2025-11-20T10:16:16.583229Z","shell.execute_reply":"2025-11-20T10:16:16.590042Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"mask_colors = ['red', 'blue', 'green', 'purple', 'orange', 'brown', 'pink', 'gray', 'olive', 'cyan']\n\ndef visualize_batch_samples(dataloader, model=None, device=device):\n    \"\"\"Visualize batch samples with predictions if model provided\"\"\"\n    images, targets = next(iter(dataloader))\n    \n    fig, axes = plt.subplots(2, 4, figsize=(20, 10))\n    \n    for i in range(min(4, len(images))):\n        # Original image\n        img = images[i].cpu().permute(1, 2, 0).numpy()\n        # img = img * np.array([0.229, 0.224, 0.225]) + np.array([0.485, 0.456, 0.406])\n        img = np.clip(img, 0, 1)\n        \n        # Ground truth mask\n        \n        mask = targets[i]\n        for c in range(mask.shape[0]):\n            mask[mask > 0] = c + 1\n            \n        mask = mask.cpu().squeeze(dim=0).numpy()\n        \n        levels = np.unique(mask)[:-1] + 0.5\n        \n        axes[0, i].imshow(img)\n        axes[0, i].contour(mask, levels=levels, colors=mask_colors, linewidths=1)\n        axes[0, i].set_title(f'Image {i+1}')\n        axes[0, i].axis('off')\n\n        \n        axes[1, i].imshow(mask, cmap='hot', alpha=0.7)\n        axes[1, i].set_title(f'Ground Truth Mask {i+1}')\n        axes[1, i].axis('off')\n    # print(images)\n    plt.tight_layout()\n    plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T10:16:16.591665Z","iopub.execute_input":"2025-11-20T10:16:16.591946Z","iopub.status.idle":"2025-11-20T10:16:16.602484Z","shell.execute_reply.started":"2025-11-20T10:16:16.591926Z","shell.execute_reply":"2025-11-20T10:16:16.601597Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# visualize_batch_samples(train_dataloader)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T10:16:16.60346Z","iopub.execute_input":"2025-11-20T10:16:16.603879Z","iopub.status.idle":"2025-11-20T10:16:16.620143Z","shell.execute_reply.started":"2025-11-20T10:16:16.603852Z","shell.execute_reply":"2025-11-20T10:16:16.619269Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"image, mask = next(iter(train_dataloader))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T10:16:16.621034Z","iopub.execute_input":"2025-11-20T10:16:16.621314Z","iopub.status.idle":"2025-11-20T10:16:18.965008Z","shell.execute_reply.started":"2025-11-20T10:16:16.621293Z","shell.execute_reply":"2025-11-20T10:16:18.964052Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"image.shape, mask.shape","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T10:16:18.966469Z","iopub.execute_input":"2025-11-20T10:16:18.967348Z","iopub.status.idle":"2025-11-20T10:16:18.974226Z","shell.execute_reply.started":"2025-11-20T10:16:18.967302Z","shell.execute_reply":"2025-11-20T10:16:18.973377Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Train Step","metadata":{}},{"cell_type":"code","source":"def train_step(model: torch.nn.Module,\n               dataloader: torch.utils.data.DataLoader,\n               loss_fn: torch.nn.Module,\n               optimizer: torch.optim.Optimizer):\n\n    # Put model in train mode\n    model.train()\n\n    scaler = GradScaler()\n\n    running_loss = 0.0\n    running_dice_loss = 0.0\n    running_bce_loss = 0.0\n    \n    # Setup train loss and train F1 values\n    train_loss, train_dice_loss, train_bce_loss = 0., 0., 0.\n\n    # Loop through data loader \n    for batch, (X, y) in enumerate(dataloader):\n        # send data to target device\n        X, y = X.to(device), y.to(device)\n\n        # automatically cast the inputs and model parameters to FP16 during the forward pass\n        with autocast(): \n            # 1. Forward pass\n            y_pred = model(X)\n\n            # 2. Calculate and accumulate loss\n            total_loss, dice_loss, bce_loss = loss_fn(y_pred, y)\n            \n\n        # 3. Optimizer zero grad\n        optimizer.zero_grad()\n\n        # 4. Loss backward\n        scaler.scale(dice_loss).backward()\n\n        # 5. Optimizer step\n        scaler.step(optimizer)\n        scaler.update()\n\n        running_loss += total_loss.item()\n        running_dice_loss += dice_loss.item()\n        running_bce_loss += bce_loss.item()\n        \n        \n\n    # Adjust metrics to get average loss and F1 per batch\n    train_loss = running_loss / len(dataloader)\n    train_dice_loss = running_dice_loss / len(dataloader)\n    train_bce_loss = running_bce_loss / len(dataloader)\n\n    return train_loss, train_dice_loss, train_bce_loss","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T10:16:18.975341Z","iopub.execute_input":"2025-11-20T10:16:18.975598Z","iopub.status.idle":"2025-11-20T10:16:18.996391Z","shell.execute_reply.started":"2025-11-20T10:16:18.975578Z","shell.execute_reply":"2025-11-20T10:16:18.995449Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Test Step","metadata":{}},{"cell_type":"code","source":"def test_step(model: torch.nn.Module,\n             dataloader: torch.utils.data.DataLoader,\n             loss_fn: torch.nn.Module):\n    # Put model in eval mode\n    model.eval()\n\n    # set up test loss and F1 values\n    running_loss = 0.0\n    running_dice_loss = 0.0\n    running_bce_loss = 0.0\n\n    # Turn on inference context manager\n    with torch.no_grad():\n        # Loop through DataLoader batches\n        for batch, (X, y) in enumerate(dataloader):\n            # Send data to target device\n            X, y = X.to(device), y.to(device)\n\n            # 1. Forward pass\n            test_pred = model(X)\n\n            # 2. Calculate and accumulate loss\n            total_loss, dice_loss, bce_loss = loss_fn(test_pred, y)\n            \n\n            # Calculate and accumulate accuracy metrics across all batches\n            running_loss += total_loss.item()\n            running_dice_loss += dice_loss.item()\n            running_bce_loss += bce_loss.item()\n\n    # Adjust metrics to get average loss and accuracy per batch\n    test_loss = running_loss / len(dataloader)\n    test_dice_loss = running_dice_loss / len(dataloader)\n    test_bce_loss = running_bce_loss / len(dataloader)\n\n    return test_loss, test_dice_loss, test_bce_loss","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T10:16:18.997495Z","iopub.execute_input":"2025-11-20T10:16:18.99783Z","iopub.status.idle":"2025-11-20T10:16:19.018908Z","shell.execute_reply.started":"2025-11-20T10:16:18.997803Z","shell.execute_reply":"2025-11-20T10:16:19.017904Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Training setup","metadata":{}},{"cell_type":"code","source":"from tqdm.auto import tqdm\n\n# 1. Take in various parameters required for training and test steps\ndef train(model: torch.nn.Module,\n         train_dataloader: torch.utils.data.DataLoader,\n         test_dataloader: torch.utils.data.DataLoader,\n         optimizer: torch.optim.Optimizer,\n         loss_fn: torch.nn.Module = nn.BCELoss(),\n         epochs: int = 1):\n\n    # 2. Create empty results dictionary\n    results = {\"train_loss\": [],\n              \"train_dice_loss\": [],\n               \"train_bce_loss\": [],\n              \"test_loss\": [],\n                \"test_dice_loss\": [],\n              \"test_bce_loss\": []}\n\n    # 3. Loop through training and testing steps for a number of epochs\n    for epoch in tqdm(range(epochs)):\n        train_loss, train_dice_loss, train_bce_loss = train_step(model = model,\n                                         dataloader = train_dataloader,\n                          \n                                          loss_fn = loss_fn,\n                                         optimizer = optimizer)\n\n        test_loss, test_dice_loss, test_bce_loss = test_step(model = model,\n                                      dataloader = test_dataloader,\n                                      loss_fn = loss_fn)\n\n        #4. print out what's happening\n        print(\n            f\"Epoch: {epoch + 1} |\"\n            f\"Epoch: {train_loss:.4f} |\"\n            f\"train_dice_loss: {train_dice_loss:.4f} |\"\n            f\"train_bce_loss: {train_bce_loss:.4f} |\"\n            f\"test_loss: {test_loss:.4f} |\"\n            f\"test_dice_loss: {test_dice_loss: .4f}|\"\n            f\"test_bce_loss: {test_bce_loss:.4f} |\"\n        )\n    \n        # 5. Update results dictionary\n        # Ensure all data is moved to CPU and converted to float for storage\n        results[\"train_loss\"].append(train_loss.item() if isinstance(train_loss, torch.Tensor) else train_loss)\n        results[\"train_dice_loss\"].append(train_dice_loss.item() if isinstance(train_dice_loss, torch.Tensor) else train_dice_loss)\n        results[\"train_bce_loss\"].append(train_bce_loss.item() if isinstance(train_bce_loss, torch.Tensor) else train_bce_loss)\n        results[\"test_loss\"].append(test_loss.item() if isinstance(test_loss, torch.Tensor) else test_loss)\n        results[\"test_dice_loss\"].append(test_dice_loss.item() if isinstance(test_dice_loss, torch.Tensor) else test_dice_loss)\n        results[\"test_bce_loss\"].append(test_bce_loss.item() if isinstance(test_bce_loss, torch.Tensor) else test_bce_loss)\n    return results","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T10:16:19.019975Z","iopub.execute_input":"2025-11-20T10:16:19.020307Z","iopub.status.idle":"2025-11-20T10:16:19.187429Z","shell.execute_reply.started":"2025-11-20T10:16:19.020279Z","shell.execute_reply":"2025-11-20T10:16:19.186701Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Initialize Model","metadata":{}},{"cell_type":"code","source":"from torch.cuda.amp import GradScaler, autocast\n\n# Set random seeds\ntorch.manual_seed(42)\ntorch.cuda.manual_seed(42)\n\n# Set learning rate\nLR = 2e-4\n\n# Set up loss function and optimizer\n# Loss \nloss_fn = MixedLoss()\n\n# Initialize the model\n\nmodel = UnetMobilenetV2().to(device)\n\n# Load the state_dict\nmodel.load_state_dict(torch.load(\"/kaggle/input/copy-move-forgery-tissue-trained-many-times/pytorch/default/1/model_tissue_preliminary_state_dict (1).pt\",\n                                map_location=torch.device('cpu')))\n\n# Freeze parameters\nfor param in model.parameters():\n    param.requires_grad = True\n\n# for param in model.conv_last.parameters():\n#     param.requires_grad = True\n\n# for param in model.conv_score.parameters():\n#     param.requires_grad = True\n\n# for param in model.invres4.parameters():\n#     param.requires_grad = True\n\noptimizer = torch.optim.Adam(model.parameters(), lr = LR)\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T10:16:19.188488Z","iopub.execute_input":"2025-11-20T10:16:19.188894Z","iopub.status.idle":"2025-11-20T10:16:19.76685Z","shell.execute_reply.started":"2025-11-20T10:16:19.188866Z","shell.execute_reply":"2025-11-20T10:16:19.765935Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from torchsummary import summary\n# summary(model, (3, 256, 256))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T10:16:19.767894Z","iopub.execute_input":"2025-11-20T10:16:19.768219Z","iopub.status.idle":"2025-11-20T10:16:19.776993Z","shell.execute_reply.started":"2025-11-20T10:16:19.768195Z","shell.execute_reply":"2025-11-20T10:16:19.776031Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Train Model","metadata":{}},{"cell_type":"code","source":"NUM_EPOCHS = 10\n\nmodel_results = train(model = model,\n                     train_dataloader = train_dataloader,\n                     test_dataloader = val_dataloader,\n                     optimizer = optimizer,\n                     loss_fn = loss_fn,\n                     epochs = NUM_EPOCHS)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T10:16:19.792559Z","iopub.status.idle":"2025-11-20T10:16:19.792859Z","shell.execute_reply.started":"2025-11-20T10:16:19.792689Z","shell.execute_reply":"2025-11-20T10:16:19.7927Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\"Save the model's state dictionary\"\n# torch.save(model.state_dict(), '/kaggle/working/model_tissue_preliminary_state_dict.pt')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T10:16:19.794408Z","iopub.status.idle":"2025-11-20T10:16:19.79478Z","shell.execute_reply.started":"2025-11-20T10:16:19.794557Z","shell.execute_reply":"2025-11-20T10:16:19.794574Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import gc\n \n# Invoke garbage collector\ngc.collect()\n\n# Clear GPU cache\ntorch.cuda.empty_cache()\nprint(f\"Memory allocated after clearing cache: {torch.cuda.memory_allocated()} bytes\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T10:16:19.796351Z","iopub.status.idle":"2025-11-20T10:16:19.797234Z","shell.execute_reply.started":"2025-11-20T10:16:19.797008Z","shell.execute_reply":"2025-11-20T10:16:19.79703Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.patches as patches\n\ndef visualize_predictions(images, masks, preds, model=None, device=device):\n    \"\"\"Visualize batch samples with predictions if model provided\"\"\"\n\n    fig, axes = plt.subplots(2, 4, figsize=(40, 20))\n    preds= torch.sigmoid(preds)\n    \n    for i in range(min(2, len(images))):\n        # Original image\n        img = images[i].cpu().permute(1, 2, 0).numpy()\n        img = np.clip(img, 0, 1)\n        mask =masks[i]\n        levels = np.unique(mask)[:-1] + 0.5\n\n        \n        for c in range(mask.shape[0]):\n            mask[mask > 0] = c + 1\n            \n        mask = mask.cpu().squeeze(dim=0).numpy()\n        \n        levels = np.unique(mask)[:-1] + 0.5\n        \n        axes[i, 0].imshow(img)\n        axes[i, 0].contour(mask, levels=levels, colors=mask_colors, linewidths=2)\n        axes[i, 0].set_title(f'Image {i+1}')\n        axes[i, 0].axis('off')\n        \n        # Ground truth mask\n        # mask = torch.zeros_like(images[i][0])\n        # for target_mask in targets[i]:\n        #     mask = torch.max(mask, target_mask.cpu())\n        axes[i, 1].imshow(mask, cmap='bwr', alpha=0.7)\n        # axes[i, 1].set_title(f'Ground Truth Mask {i+1}')\n        # axes[i, 1].imshow(pred, cmap='hot', alpha=0.7)\n        axes[i, 1].axis('off')\n\n        # Predicted mask\n        \n        pred = preds[i].cpu().permute(1,2,0).numpy()\n        pred = (pred>0.2).astype(np.float32)\n        axes[i, 2].imshow(pred, cmap='hot', alpha=0.7)\n        \n        axes[i, 2].set_title(f'Pedicted Mask {i+2}')\n        axes[i, 2].axis('off')\n\n        # Overlap between the Ground truth and the prediction\n        axes[i, 3].imshow(mask, cmap = 'bwr', alpha = 0.7)\n        axes[i, 3].imshow(pred, cmap = 'hot', alpha = 0.7)\n        axes[i, 3].set_title(f'Overlap between ground truth and predictions')\n        axes[i, 3].axis('off')\n        \n    plt.tight_layout()\n    plt.show()\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T10:16:31.138392Z","iopub.execute_input":"2025-11-20T10:16:31.13867Z","iopub.status.idle":"2025-11-20T10:16:31.149067Z","shell.execute_reply.started":"2025-11-20T10:16:31.138652Z","shell.execute_reply":"2025-11-20T10:16:31.147932Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Evaluate Model","metadata":{}},{"cell_type":"code","source":"# Load the state_dict\n# model.load_state_dict(torch.load(\"\",\n                                # map_location=torch.device('cpu')))\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T10:16:19.799808Z","iopub.status.idle":"2025-11-20T10:16:19.800094Z","shell.execute_reply.started":"2025-11-20T10:16:19.799962Z","shell.execute_reply":"2025-11-20T10:16:19.799977Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"img_train, mask_train = next(iter(train_dataloader))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T10:16:19.801893Z","iopub.status.idle":"2025-11-20T10:16:19.802487Z","shell.execute_reply.started":"2025-11-20T10:16:19.802295Z","shell.execute_reply":"2025-11-20T10:16:19.802314Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model.eval()\nwith torch.no_grad():\n    train_pred = model(img_train.to(device))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T10:16:19.803455Z","iopub.status.idle":"2025-11-20T10:16:19.803828Z","shell.execute_reply.started":"2025-11-20T10:16:19.80363Z","shell.execute_reply":"2025-11-20T10:16:19.803647Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"visualize_predictions(img_train[0:2], mask_train[0:2], train_pred[0:2])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T10:16:19.80503Z","iopub.status.idle":"2025-11-20T10:16:19.805626Z","shell.execute_reply.started":"2025-11-20T10:16:19.805423Z","shell.execute_reply":"2025-11-20T10:16:19.805441Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"img, ground_truth = next(iter(val_dataloader))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T11:38:22.347892Z","iopub.execute_input":"2025-11-20T11:38:22.348718Z","iopub.status.idle":"2025-11-20T11:38:23.405045Z","shell.execute_reply.started":"2025-11-20T11:38:22.34869Z","shell.execute_reply":"2025-11-20T11:38:23.403778Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model.eval()\nwith torch.no_grad():\n    test_pred = model(img.to(device))\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T11:38:24.504278Z","iopub.execute_input":"2025-11-20T11:38:24.505256Z","iopub.status.idle":"2025-11-20T11:38:45.938831Z","shell.execute_reply.started":"2025-11-20T11:38:24.505218Z","shell.execute_reply":"2025-11-20T11:38:45.937798Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"visualize_predictions(img[0:2], ground_truth[0:2], test_pred[0:2])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T10:16:19.810226Z","iopub.status.idle":"2025-11-20T10:16:19.810565Z","shell.execute_reply.started":"2025-11-20T10:16:19.810376Z","shell.execute_reply":"2025-11-20T10:16:19.81039Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"visualize_predictions(img[5:7], ground_truth[5:7], test_pred[5:7])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T10:16:19.811698Z","iopub.status.idle":"2025-11-20T10:16:19.812065Z","shell.execute_reply.started":"2025-11-20T10:16:19.811866Z","shell.execute_reply":"2025-11-20T10:16:19.811884Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"visualize_predictions(img[9:11], ground_truth[9:11], test_pred[9:11])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T10:16:19.813569Z","iopub.status.idle":"2025-11-20T10:16:19.813871Z","shell.execute_reply.started":"2025-11-20T10:16:19.813712Z","shell.execute_reply":"2025-11-20T10:16:19.813723Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"visualize_predictions(img[11:13], ground_truth[11:13], test_pred[11:13])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T10:16:19.815246Z","iopub.status.idle":"2025-11-20T10:16:19.81556Z","shell.execute_reply.started":"2025-11-20T10:16:19.81539Z","shell.execute_reply":"2025-11-20T10:16:19.815409Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"visualize_predictions(img[14:16], ground_truth[14:16], test_pred[14:16])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T10:16:19.816453Z","iopub.status.idle":"2025-11-20T10:16:19.816729Z","shell.execute_reply.started":"2025-11-20T10:16:19.816591Z","shell.execute_reply":"2025-11-20T10:16:19.816605Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class SoftDiceLoss(nn.Module):\n    def __init__(self):\n        super(SoftDiceLoss, self).__init__()\n\n    def forward(self, logits, targets):\n        eps = 1e-5 \n        bs = targets.size(0)\n\n        probs = torch.sigmoid(logits)\n        m1 = probs.view(bs, -1)\n        m2 = targets.view(bs, -1)\n\n        intersection = (m1 * m2).sum(dim=1)\n        union = m1.sum(dim=1) + m2.sum(dim=1) + eps\n\n        score = (2. * intersection + eps)/ union\n        loss = 1 - score.mean()  \n\n        return loss\n        ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T10:16:19.817994Z","iopub.status.idle":"2025-11-20T10:16:19.818259Z","shell.execute_reply.started":"2025-11-20T10:16:19.81813Z","shell.execute_reply":"2025-11-20T10:16:19.818141Z"},"jupyter":{"source_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"diceloss = SoftDiceLoss()\ndiceloss(test_pred, mask.to(device))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T10:16:19.8191Z","iopub.status.idle":"2025-11-20T10:16:19.819332Z","shell.execute_reply.started":"2025-11-20T10:16:19.819222Z","shell.execute_reply":"2025-11-20T10:16:19.819232Z"},"jupyter":{"source_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Feature-extracted dataset","metadata":{}},{"cell_type":"code","source":"import torchvision.transforms.functional as tF\n\nclass Potential_area_Dataset(Dataset):\n    def __init__(self, images, ground_truth, pred_logits,  \n                  is_train=True, img_size = 256, threshold = 0.2):\n        \n        self.images = images\n        self.ground_truth = ground_truth\n        \n        self.is_train = is_train\n        \n        self.threshold = threshold\n\n        self.pred_logits = pred_logits\n            \n    def __len__(self):\n        return len(self.images)\n    \n    def __getitem__(self, idx):\n\n        image = self.images[idx]\n        # print(type(image))\n        pred_mask = torch.sigmoid(self.pred_logits[idx]) > self.threshold  # persentage values to binaries\n        # Load image\n        boxes = self.mask_to_boxes(pred_mask)\n        \n        image_fragments = self.get_part_image(image, boxes)\n           \n        return image, idx, pred_mask, boxes, image_fragments\n\n    def mask_to_boxes(self, mask):\n        \"\"\"Convert segmentation mask to bounding boxes\"\"\"\n        if isinstance(mask, torch.Tensor):\n            mask_np = mask.permute(1,2,0).cpu().numpy().astype(np.uint8)\n            \n        else:\n            mask_np = mask.permute(1,2,0).astype(np.uint8)\n            \n        # Find contours in the mask\n        contours, _ = cv2.findContours(mask_np, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)\n        \n        boxes = []\n     \n        \n        for contour in contours:\n            if len(contour) > 0:\n                x, y, w, h = cv2.boundingRect(contour)\n                if w > 5 and h > 5:\n                    boxes.append([x, y, x + w, y + h])\n\n        return boxes\n\n    def get_part_image(self,image, boxes):\n        cut_outs = []\n       \n        for box in boxes:\n            # fragment = torch.zeros_like(image)\n            # x0, y0, x1, y1 = box[0], box[1], box[2], box[3]\n            # fragment[:,y0-3:y1+3, x0-3:x1+3] = 1                   \n            # cut_outs.append(fragment*image) # C * H * W\n            x0, y0, x1, y1 = box[0], box[1], box[2], box[3]\n            fragment = image[:, y0:y1, x0:x1]\n            cut_outs.append(fragment)      \n            \n        return cut_outs","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T10:17:12.056258Z","iopub.execute_input":"2025-11-20T10:17:12.056544Z","iopub.status.idle":"2025-11-20T10:17:12.066892Z","shell.execute_reply.started":"2025-11-20T10:17:12.056525Z","shell.execute_reply":"2025-11-20T10:17:12.066017Z"},"jupyter":{"source_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Possible forged areas","metadata":{}},{"cell_type":"code","source":"\npotential_forgeries_dataset = Potential_area_Dataset(img, ground_truth, test_pred)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T11:38:53.291739Z","iopub.execute_input":"2025-11-20T11:38:53.292135Z","iopub.status.idle":"2025-11-20T11:38:53.296487Z","shell.execute_reply.started":"2025-11-20T11:38:53.29211Z","shell.execute_reply":"2025-11-20T11:38:53.295682Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"BATCH_SIZE = 1\n\npotential_forged_dataloader = DataLoader(potential_forgeries_dataset, batch_size = BATCH_SIZE, \n                                shuffle = True, num_workers = 0, pin_memory=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T11:38:54.883871Z","iopub.execute_input":"2025-11-20T11:38:54.8847Z","iopub.status.idle":"2025-11-20T11:38:54.895353Z","shell.execute_reply.started":"2025-11-20T11:38:54.884671Z","shell.execute_reply":"2025-11-20T11:38:54.89432Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"image,idx, pred_mask, boxes, image_fragments = next(iter(potential_forged_dataloader))\nimage_fragments_np = [image.squeeze().permute(1,2,0).cpu().numpy() for image in image_fragments]\nvisualize_predictions(img[idx], ground_truth[idx], test_pred[idx])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T11:39:40.881084Z","iopub.execute_input":"2025-11-20T11:39:40.881434Z","iopub.status.idle":"2025-11-20T11:39:42.782429Z","shell.execute_reply.started":"2025-11-20T11:39:40.881409Z","shell.execute_reply":"2025-11-20T11:39:42.781483Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(idx)\nprint(len(boxes))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T11:10:18.943009Z","iopub.execute_input":"2025-11-20T11:10:18.943347Z","iopub.status.idle":"2025-11-20T11:10:18.949801Z","shell.execute_reply.started":"2025-11-20T11:10:18.943322Z","shell.execute_reply":"2025-11-20T11:10:18.948811Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"img_id_1 = np.random.randint(len(boxes))\nimg_id_2 = np.random.randint(len(boxes))\nimg_id_3 = np.random.randint(len(boxes))\nimg_id_4 = np.random.randint(len(boxes))\nimport matplotlib.pyplot as plt\nfig, axes = plt.subplots(2, 2, figsize=(10, 10))\naxes[0,0].imshow(image_fragments_np[img_id_1])\naxes[0,0].axis('off')\naxes[0,1].imshow(image_fragments_np[img_id_2])\naxes[0,1].axis('off')\naxes[1,0].imshow(image_fragments_np[img_id_3])\naxes[1,0].axis('off')\naxes[1,1].imshow(image_fragments_np[img_id_4])\naxes[1,1].axis('off')\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T11:39:53.507717Z","iopub.execute_input":"2025-11-20T11:39:53.508094Z","iopub.status.idle":"2025-11-20T11:39:53.665651Z","shell.execute_reply.started":"2025-11-20T11:39:53.508072Z","shell.execute_reply":"2025-11-20T11:39:53.664622Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Learned Perceptual Image Patch Similarity\n","metadata":{}},{"cell_type":"code","source":"\nfrom torchmetrics.image.lpip import LearnedPerceptualImagePatchSimilarity\nlpips = LearnedPerceptualImagePatchSimilarity(net_type='vgg')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T10:17:46.198097Z","iopub.execute_input":"2025-11-20T10:17:46.198445Z","iopub.status.idle":"2025-11-20T10:17:56.853605Z","shell.execute_reply.started":"2025-11-20T10:17:46.19842Z","shell.execute_reply":"2025-11-20T10:17:56.852632Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def imgshow(img,text=None,should_save=False):\n   \n    if type(img) == torch.Tensor:\n        npimg = img.numpy()\n        plt.axis(\"off\")\n        plt.imshow(np.transpose(npimg, (1, 2, 0)))\n        plt.show()   \n    else:\n        plt.axis(\"off\")\n        plt.imshow(img)\n        plt.show()  ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T10:18:49.388571Z","iopub.execute_input":"2025-11-20T10:18:49.388938Z","iopub.status.idle":"2025-11-20T10:18:49.394975Z","shell.execute_reply.started":"2025-11-20T10:18:49.388913Z","shell.execute_reply":"2025-11-20T10:18:49.393985Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from torchvision import transforms\ninput_size_squeeze = 256\nsimilarity_transform = v2.Compose([\n        \n    v2.ToTensor(),\n    transforms.Resize((input_size_squeeze, input_size_squeeze)),\n    # v2.Resize((input_size_squeeze, input_size_squeeze)),\n        \n        # v2.ToImage(),\n        # v2.ToDtype(torch.float32, scale = True),\n        \n    ])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T10:18:51.20535Z","iopub.execute_input":"2025-11-20T10:18:51.205675Z","iopub.status.idle":"2025-11-20T10:18:51.211696Z","shell.execute_reply.started":"2025-11-20T10:18:51.205652Z","shell.execute_reply":"2025-11-20T10:18:51.210743Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\"Example\"\n# >>> # LPIPS needs the images to be in the [-1, 1] range.\n# >>> img1 = (rand(10, 3, 100, 100) * 2) - 1\n# >>> img2 = (rand(10, 3, 100, 100) * 2) - 1\n# >>> lpips(img1, img2)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T10:16:19.83909Z","iopub.status.idle":"2025-11-20T10:16:19.839556Z","shell.execute_reply.started":"2025-11-20T10:16:19.839366Z","shell.execute_reply":"2025-11-20T10:16:19.839385Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"imgshow(image.squeeze(0))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T11:39:59.480495Z","iopub.execute_input":"2025-11-20T11:39:59.480867Z","iopub.status.idle":"2025-11-20T11:39:59.635168Z","shell.execute_reply.started":"2025-11-20T11:39:59.480843Z","shell.execute_reply":"2025-11-20T11:39:59.634093Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\nimg_id_1 = np.random.randint(len(boxes))\nimg_id_2 = np.random.randint(len(boxes))\nimg1 = image_fragments_np[3] \n\nimg1_simi = similarity_transform(img1).unsqueeze(0)\n\nimgshow(img1_simi.squeeze(0))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T11:41:49.680509Z","iopub.execute_input":"2025-11-20T11:41:49.680923Z","iopub.status.idle":"2025-11-20T11:41:49.801276Z","shell.execute_reply.started":"2025-11-20T11:41:49.680844Z","shell.execute_reply":"2025-11-20T11:41:49.8003Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"img2 = image_fragments_np[4] \nimg2_simi = similarity_transform(img2).unsqueeze(0)\nimgshow(img2_simi.squeeze(0))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T11:42:02.208354Z","iopub.execute_input":"2025-11-20T11:42:02.208696Z","iopub.status.idle":"2025-11-20T11:42:02.320254Z","shell.execute_reply.started":"2025-11-20T11:42:02.20867Z","shell.execute_reply":"2025-11-20T11:42:02.319213Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"lpips(img1_simi, img2_simi) # check the similarity between the two images","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T11:42:18.03868Z","iopub.execute_input":"2025-11-20T11:42:18.039043Z","iopub.status.idle":"2025-11-20T11:42:18.834642Z","shell.execute_reply.started":"2025-11-20T11:42:18.03902Z","shell.execute_reply":"2025-11-20T11:42:18.833675Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def find_matches(img_simi_list):\n\n    \n    scores_final = {}\n    \n    for idx in range(len(img_simi_list)):\n        \n        scores = []\n        img1_simi = img_simi_list[idx]\n\n        try:\n            for img2_simi in img_simi_list[idx+1:-1]:      \n            \n                scores.append(lpips(img1_simi, img2_simi).item())\n\n            scores_final[idx] = scores\n        except:\n            print('All items have been compared between each other.')\n            return scores_final\n        \n    return scores_final\n        ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T10:55:06.574559Z","iopub.execute_input":"2025-11-20T10:55:06.574927Z","iopub.status.idle":"2025-11-20T10:55:06.582018Z","shell.execute_reply.started":"2025-11-20T10:55:06.574892Z","shell.execute_reply":"2025-11-20T10:55:06.58109Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"img_simi_list = [similarity_transform(img_frag).unsqueeze(0) for img_frag in image_fragments_np]\n\nfind_matches(img_simi_list)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-20T11:40:20.27998Z","iopub.execute_input":"2025-11-20T11:40:20.280308Z","iopub.status.idle":"2025-11-20T11:40:20.29243Z","shell.execute_reply.started":"2025-11-20T11:40:20.280284Z","shell.execute_reply":"2025-11-20T11:40:20.291583Z"}},"outputs":[],"execution_count":null}]}