{"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":"Private LB 0.61","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Net_dynamicpool(nn.Module):\n        def __init__(self,):\n                super().__init__()\n\n                conv_dim=64\n                #encoder_dim  = [conv_dim] + [64, 128, 256, 512 ]\n                encoder_dim  = [conv_dim] + [256, 512,1024,2048 ]\n                self.encoder = timm.create_model(CFG.backbone,pretrained=False,in_chans=CFG.in_chans)\n\n                self.decoder = SmpUnetDecoder(\n                        encoder_channels=[0] + encoder_dim,\n                        decoder_channels=[256, 128, 64, 32, 16],\n                        n_blocks=5,\n                        use_batchnorm=True,\n                        center=False,\n                        attention_type=None,\n                )\n                self.logit = nn.Conv2d(16,1,kernel_size=1)\n                #for dim in encoder_dim:\n                #    print (dim)\n                #-- pool attention weight\n                self.weight = nn.ModuleList([nn.Sequential(nn.Conv2d(dim, dim, kernel_size=3, padding=1),nn.ReLU(inplace=True),) for dim in encoder_dim])\n                #self.weight = nn.ModuleList([nn.Sequential(nn.Conv2d(dim, , kernel_size=3, padding=1),nn.ReLU(inplace=True),) for dim in encoder_dim])\n                #self.weight = nn.ModuleList([nn.Sequential(nn.Conv2d(dim, dim, kernel_size=3, padding=1, groups=dim),nn.ReLU(inplace=True),) for dim in encoder_dim])\n        def forward(self, image):\n                v = image\n                #v = batch['volume']\n                #print (v.shape)\n                B,C,H,W = v.shape  #size()\n                vv = [v[:, i:i+CFG.in_chans] for i in [0, 2, 4]]\n                #for i, tensor in enumerate(vv):\n                #     print(f\"Shape of vv[{i}]:\", tensor.shape)\n                #vv = [v[:, :, :, i:i+CFG.in_chans] for i in [0, 2, 4]]\n                #vv = [v[:, i:i+CFG.in_chans] for i in [0, 2, 4]]\n\n                #vv = [tensor for tensor in vv if tensor.size(0) == vv[0].size(0)]\n                #vv = torch.split(v, CFG.in_chans, dim=1)\n                #vv = torch.split(v, CFG.in_chans, dim=1)\n                K = len(vv)\n                x = torch.cat(vv, 0)\n                #x = v\n\n                # ----\n                encoder = []\n                x = self.encoder.conv1(x)\n                x = self.encoder.bn1(x)\n                x = self.encoder.act1(x)   ; encoder.append(x)\n                x = F.avg_pool2d(x,kernel_size=2,stride=2)\n                #x = self.encoder.maxpool(x)\n                x = self.encoder.layer1(x) ; encoder.append(x)\n                x = self.encoder.layer2(x) ; encoder.append(x)\n                x = self.encoder.layer3(x) ; encoder.append(x)\n                x = self.encoder.layer4(x) ; encoder.append(x)\n                #print ('encoder', [f.shape for f in encoder])\n\n                #encode pooling -------\n                #<todo> add positional encode (z slice no.)\n                for i in range(len(encoder)):\n                        e = encoder[i]\n                        #print (self.weight[i].shape, encoder[i])\n                        f = self.weight[i](encoder[i])\n                        _, c, h, w = f.shape\n                        f = rearrange(f, ' (K B) c  h w -> B K c h w', K=K, B=B, h=h, w=w) #f.reshape(B, K, c, h, w)\n                        e = rearrange(e, '(K B) c  h w -> B K c h w', K=K, B=B, h=h, w=w) #e.reshape(B, K, c, h, w)\n                        #f = f.reshape(B, K, c, h, w)\n                        #e = e.reshape(B, K, c, h, w)\n                        w = F.softmax(f, 1)\n                        e = (w * e).mean(dim=1)\n                        encoder[i] = e\n\n                # ---\n                last, decoder = self.decoder(encoder)\n                #print('decoder',[f.shape for f in decoder])\n                #print('last',last.shape)\n                logit = self.logit(last)\n                #output = torch.sigmoid(logit)\n                output = logit\n                #output = {\n        #               'ink' : torch.sigmoid(logit),\n        #       } \n                return output","metadata":{},"execution_count":null,"outputs":[]}]}