{"cells":[{"metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"cell_type":"code","source":"import cv2\nfrom tqdm import tqdm_notebook as tqdm\nimport fastai\nfrom fastai.vision import *\nimport os\nfrom mish_activation import *\nimport warnings\nwarnings.filterwarnings(\"ignore\")\nimport skimage.io\nimport numpy as np\nimport pandas as pd\nfrom bisect import bisect_right","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"sz = 128\nbs = 1\nN = 128\nnworkers = 2\n\n#DATA = '../input/prostate-cancer-grade-assessment/train_images/'\n#TEST = '../input/prostate-cancer-grade-assessment/train.csv'\nDATA = '../input/prostate-cancer-grade-assessment/test_images'\nTEST = '../input/prostate-cancer-grade-assessment/test.csv'\nSAMPLE = '../input/prostate-cancer-grade-assessment/sample_submission.csv'\nMODELS = [f'../input/panda-init-class-model1/RNXT50_128k_0_{i}.pth' for i in range(4)] + \\\n         [f'../input/panda-init-class-model1/RNXT50_128k_3_{i}.pth' for i in range(4)] + \\\n         [f'../input/panda-init-class-model1/RNXT50_128kr_1_{i}.pth' for i in range(4)] + \\\n         [f'../input/panda-init-class-model1/RNXT50_128kr_2_{i}.pth' for i in range(4)] + \\\n         [f'../input/panda-init-class-model1/RNXT50_128kr_2feature_{i}.pth' for i in range(4)] + \\\n         [f'../input/panda-init-class-model1/RNXT50_128kr_9_{i}.pth' for i in range(4)] + \\\n         [f'../input/panda-init-class-model1/RNXT50_128kr_7c_{i}.pth' for i in range(4)] + \\\n         [f'../input/panda-init-class-model1/RNXT50_128kr_3_{i}.pth' for i in range(4)]\nws = [1,1,1,6,1,1,1,1]\nws = [w for w in ws for k in range(4)]","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"from torchvision.models.resnet import ResNet, Bottleneck\n\nclass Model(nn.Module):\n    def __init__(self, arch='resnext50_32x4d', n=11, pre=True):\n        super().__init__()\n        m = ResNet(Bottleneck, [3, 4, 6, 3], groups=32, width_per_group=4)\n        self.enc = nn.Sequential(*list(m.children())[:-2])       \n        nc = list(m.children())[-1].in_features\n        self.head = nn.Sequential(AdaptiveConcatPool2d(),Flatten(),\n                                  nn.Linear(2*nc,512),Mish(),nn.GroupNorm(32,512),\n                                  nn.Dropout(0.5),nn.Linear(512,n))\n        \n    def forward(self, x):\n        shape = x.shape\n        n = shape[1]\n        x = x.view(-1,shape[2],shape[3],shape[4])\n        x = self.enc(x)\n        shape = x.shape\n        x = x.view(-1,n,shape[1],shape[2],shape[3]).permute(0,2,1,3,4).contiguous()\\\n          .view(-1,shape[1],shape[2]*n,shape[3])\n        x = self.head(x)\n        return x[:,:1]","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"models = []\nfor path in MODELS:\n    state_dict = torch.load(path,map_location=torch.device('cpu'))\n    model = Model()\n    model.load_state_dict(state_dict)\n    model.float()\n    model.eval()\n    model.cuda()\n    models.append(model)\n\ndel state_dict","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def tile(img):\n    shape = img.shape\n    pad0,pad1 = (sz - shape[0]%sz)%sz, (sz - shape[1]%sz)%sz\n    img = np.pad(img,[[pad0//2,pad0-pad0//2],[pad1//2,pad1-pad1//2],[0,0]],constant_values=255)\n    img = img.reshape(img.shape[0]//sz,sz,img.shape[1]//sz,sz,3)\n    img = img.transpose(0,2,1,3,4).reshape(-1,sz,sz,3)\n    if len(img) < N:\n        img = np.pad(img,[[0,N-len(img)],[0,0],[0,0],[0,0]],constant_values=255)\n    idxs = np.argsort(img.reshape(img.shape[0],-1).sum(-1))[:N]\n    img = img[idxs]\n    return img\n\nmean = torch.Tensor([1.0-0.85506157, 1.0-0.7035249, 1.0-0.80203127])\nstd = torch.Tensor([0.40011922, 0.52504386, 0.42675745])\n\nclass PandaDataset(Dataset):\n    def __init__(self, path, test):\n        self.path = path\n        self.names = list(pd.read_csv(test).image_id)\n\n    def __len__(self):\n        return len(self.names)\n\n    def __getitem__(self, idx):\n        name = self.names[idx]\n        img = skimage.io.MultiImage(os.path.join(DATA,name+'.tiff'))[1]\n        tiles = torch.Tensor((255 - tile(img))/255.0)\n        tiles = (tiles - mean)/std\n        return tiles.permute(0,3,1,2), name","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"ths = np.array([1.03125,1.03125,0.8125,0.90625,1.0859375]).cumsum()\n#ths = np.array([1.0,1.0,1.0,1.0,1.0]).cumsum()\nsub_df = pd.read_csv(SAMPLE)\nif os.path.exists(DATA):\n    bs=2\n    ds = PandaDataset(DATA,TEST)\n    dl = DataLoader(ds, batch_size=bs, num_workers=nworkers, shuffle=False)\n    names,preds = [],[]\n\n    with torch.no_grad():\n        for x,y in tqdm(dl):\n            x = x.cuda()\n            b = x.shape[0]\n            #dihedral TTA\n            #x = torch.stack([x,x.flip(-1),x.flip(-2),x.flip(-1,-2),x.transpose(-1,-2),\\\n            #  x.transpose(-1,-2).flip(-1), x.transpose(-1,-2).flip(-2),\\\n            #  x.transpose(-1,-2).flip(-1,-2)],1)\n            x = torch.stack([x,x.flip(-1),x.flip(-2),x.flip(-1,-2),x.transpose(-1,-2),\\\n              x.transpose(-1,-2).flip(-1)],1)\n            n_tta = 6\n            x = x.view(-1,N,3,sz,sz)\n            p = [model(x) for model in models]\n            p = torch.stack(p,1)\n            p = p.view(b,n_tta,len(models)).mean(1).cpu()\n            p = 6.0*torch.sigmoid(p)\n            \n            for i in range(b):\n                pred = []\n                for pi in p[i]: pred.append(bisect_right(ths, pi.numpy()))\n                preds.append(np.argmax(np.bincount(pred,ws)))\n           \n            names.append(y)\n    \n    names = np.concatenate(names)\n    sub_df = pd.DataFrame({'image_id': names, 'isup_grade': preds})","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"sub_df.to_csv(\"submission.csv\", index=False)\nsub_df.head()","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"d629ff2d2480ee46fbb7e2d37f6b5fab8052498a","collapsed":true,"_cell_guid":"79c7e3d0-c299-4dcb-8224-4455121ee9b0","trusted":false},"cell_type":"code","source":"","execution_count":null,"outputs":[]}],"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":4,"nbformat_minor":4}