{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"class SmpUnetDecodernewlrf(nn.Module):\n    def __init__(self, in_channel, skip_channel, out_channel):\n        super().__init__()\n        self.center = nn.Identity()\n\n        i_channel = [in_channel,] + out_channel[:-1]\n        s_channel = skip_channel\n        o_channel = out_channel\n        block = [\n            DecoderBlocklrf(i, s, o, use_batchnorm=True, attention_type=None)\n            for i, s, o in zip(i_channel, s_channel, o_channel)\n        ]\n        self.block = nn.ModuleList(block)\n\n    def forward(self, feature, skip):\n        d = self.center(feature)\n        decode = []\n        for i, block in enumerate(self.block):\n            s = skip[i]\n            d = block(d, s)\n            decode.append(d)\n\n        last = d\n        return last, decode\n\n\nclass DecoderBlocklrf(nn.Module):\n    def __init__(\n        self,\n        in_channels,\n        skip_channels,\n        out_channels,\n        use_batchnorm=True,\n        attention_type=None,\n    ):\n        super(DecoderBlocklrf, self).__init__()\n        self.conv1 = nn.Conv2d(\n            in_channels + skip_channels,\n            out_channels,\n            kernel_size=3,\n            padding=1,\n        )\n        self.conv2 = nn.Conv2d(out_channels, out_channels, kernel_size=3, padding=1)\n        self.relu = nn.ReLU(inplace=True)\n        self.batchnorm = nn.BatchNorm2d(out_channels) if use_batchnorm else None\n        self.attention = (\n            AttentionModule(in_channels, skip_channels[-1], attention_type)\n            if attention_type\n            else None\n        )\n\n    def forward(self, x, skip):\n        if self.attention:\n            x = self.attention(x, skip)\n        x = F.interpolate(x, scale_factor=2, mode=\"bilinear\", align_corners=False)\n        x = torch.cat([x, skip], dim=1)\n        x = self.conv1(x)\n        x = self.relu(x)\n        x = self.conv2(x)\n        x = self.relu(x)\n        if self.batchnorm:\n            x = self.batchnorm(x)\n        return x\n\n\nclass AttentionModule(nn.Module):\n    def __init__(self, in_channels, skip_channels, attention_type):\n        super(AttentionModule, self).__init__()\n        self.attention_type = attention_type\n        if attention_type == \"add\":\n            self.conv = nn.Conv2d(skip_channels, in_channels, kernel_size=1)\n        elif attention_type == \"concat\":\n            self.conv = nn.Conv2d(in_channels + skip_channels, in_channels, kernel_size=1)\n\n    def forward(self, x, skip):\n        if self.attention_type == \"add\":\n            skip = self.conv(skip)\n            x = x + skip\n        elif self.attention_type == \"concat\":\n            x = torch.cat([x, skip], dim=1)\n            x = self.conv(x)\n        return x\n\n\nclass Net_stacked2_silu_lrf(nn.Module):\n    def __init__(self):\n        super(Net_stacked2_silu_lrf, self).__init__()\n\n        conv_dim = 64\n\n        encoder1_dim = [conv_dim] + [256, 512, 1024, 2048]\n        self.encoder1 = timm.create_model(\n            CFG.backbone, pretrained=False, in_chans=CFG.in_chans\n        )\n        decoder1_dim = [256, 128, 64, 64]\n        self.decoder1 = SmpUnetDecodernewlrf(\n            in_channel=encoder1_dim[-1],\n            skip_channel=encoder1_dim[:-1][::-1],\n            out_channel=decoder1_dim,\n        )\n        self.logit1 = nn.Conv2d(decoder1_dim[-1], 1, kernel_size=1)\n\n        encoder2_dim = [conv_dim] + [256, 512, 1024, 2048]\n        self.encoder2 = timm.create_model(\n            CFG.backbone, pretrained=False, in_chans=decoder1_dim[-1]\n        )\n        decoder2_dim = [256, 128, 64]\n        self.decoder2 = SmpUnetDecodernewlrf(\n            in_channel=encoder2_dim[-1],\n            skip_channel=encoder2_dim[:-1][::-1],\n            out_channel=decoder2_dim,\n        )\n        self.logit2 = nn.Conv2d(decoder2_dim[-1], 1, kernel_size=1)\n        #self.pooling = nn.AdaptiveAvgPool2d(output_size=(1, 1))\n        self.silu = nn.SiLU()\n        # Modify convolutional layers for larger receptive fields on z-axis\n        # Modify convolutional layers for larger receptive fields on z-axis\n        for name, module in self.encoder1.named_modules():\n            if isinstance(module, nn.Conv2d):\n                kernel_size = module.kernel_size\n                if len(kernel_size) < 3:\n                    kernel_size = _triple(kernel_size[0])\n                kernel_size = (kernel_size[0], kernel_size[1], kernel_size[2] + 2)\n                module.kernel_size = kernel_size\n\n        for name, module in self.encoder2.named_modules():\n            if isinstance(module, nn.Conv2d):\n                kernel_size = module.kernel_size\n                if len(kernel_size) < 3:\n                    kernel_size = _triple(kernel_size[0])\n                kernel_size = (kernel_size[0], kernel_size[1], kernel_size[2] + 2)\n                module.kernel_size = kernel_size\n    def forward(self, image):\n        v = image\n        B, C, H, W = v.shape\n        vv = [v[:, i : i + CFG.in_chans] for i in [0, 2, 4]]\n        K = len(vv)\n        x = torch.cat(vv, 0)\n\n        encoder = []\n        e = self.encoder1\n        x = e.conv1(x)\n        x = e.bn1(x)\n        x = self.silu(x)           #e.act1(x)\n        encoder.append(x)\n        x = F.avg_pool2d(x, kernel_size=2, stride=2)\n        #x = self.pooling(x)\n        x = e.layer1(x)\n        encoder.append(x)\n        x = e.layer2(x)\n        encoder.append(x)\n        x = e.layer3(x)\n        encoder.append(x)\n        x = e.layer4(x)\n        encoder.append(x)\n\n        for i in range(len(encoder)):\n            e = encoder[i]\n            f = e\n            _, c, h, w = e.shape\n            f = rearrange(f, \" (K B) c  h w -> B K c h w\", K=K, B=B, h=h, w=w)\n            e = rearrange(e, \"(K B) c  h w -> B K c h w\", K=K, B=B, h=h, w=w)\n            w = F.softmax(f, 1)\n            e = (w * e).sum(1)\n            encoder[i] = e\n\n        feature = encoder[-1]\n        skip = encoder[:-1][::-1]\n        last, decoder = self.decoder1(feature, skip)\n        logit1 = self.logit1(last)\n\n        x = last\n        encoder = []\n        e = self.encoder2\n        x = e.layer1(x)\n        encoder.append(x)\n        x = e.layer2(x)\n        encoder.append(x)\n        x = e.layer3(x)\n        encoder.append(x)\n        x = e.layer4(x)\n        encoder.append(x)\n\n        feature = encoder[-1]\n        skip = encoder[:-1][::-1]\n        last, decoder = self.decoder2(feature, skip)\n        logit2 = self.logit2(last)\n        logit2 = F.interpolate(\n            logit2, size=(H, W), mode=\"bilinear\", align_corners=False\n        )\n        output = logit2\n        return output","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}