{"cells":[{"metadata":{},"cell_type":"markdown","source":"# Description\nWelcome to Prostate cANcer graDe Assessment (PANDA) Challenge. The task of this competition is classification of images with cancer tissue. The main challenge of this task is dealing with images of extremely high resolution and large areas of empty space. So, effective locating the areas of concern and zooming them in would be the key to reach high LB score.\n\nIn this competition I found a number of public kernels performing straightforward rescaling the input images to square. However, for this particular data such an approach is not very efficient because the aspect ratio and size of provided images are not consistent and vary in a wide range. As a result, the input images are deformed to large extend in a not consistent manner uppon rescaling that limits the ability of the model to learn. Moreover, the input consists of large empty areas leading to inefficient use of GPU memory and GPU time.\n\nIn this kernel I propose an alternative approach based on **Concatenate Tile pooling**. Instead of passing an entire image as an input, N tiles are selected from each image based on the number of tissue pixels (see [this kernel](https://www.kaggle.com/iafoss/panda-16x128x128-tiles) for description of data preparation and the [corresponding dataset](https://www.kaggle.com/iafoss/panda-16x128x128-tiles-data)) and passed independently through the convolutional part. The outputs of the convolutional part is concatenated in a large single map for each image preceding pooling and FC head (see image below). Since any spatial information is eliminated by the pooling layer, the Concat Tile pooling approach is nearly identical to passing an entire image through the convolutional part, excluding predictions for nearly empty regions, which do not contribute to the final prediction, and shuffle the remaining outputs into a square map of smaller size. Below I provide just a basic kernel only illustrating this approach. In my first trial I got 0.76 LB score, top 2 at the moment, and I believe it could be easily boosted to 0.80+. I hope you would enjoy my kernel, and please also check my submission kernel implementing the tile based approach.\n\n![](https://i.ibb.co/hF6LRVm/TILE.png)","execution_count":null},{"metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"_kg_hide-input":true,"_kg_hide-output":true},"cell_type":"code","source":"%reload_ext autoreload\n%autoreload 2\n%matplotlib inline\n\nimport fastai\nfrom fastai.vision import *\nfrom fastai.callbacks import SaveModelCallback\nimport os\n#from sklearn.model_selection import KFold\nfrom radam import *\nfrom csvlogger import *\nfrom mish_activation import *\nfrom sklearn.model_selection import StratifiedKFold\nfrom sklearn.metrics import cohen_kappa_score,confusion_matrix\nimport warnings\nwarnings.filterwarnings(\"ignore\")\n\nfastai.__version__","execution_count":null,"outputs":[]},{"metadata":{"_kg_hide-input":true,"_kg_hide-output":true,"trusted":true},"cell_type":"code","source":"# remove this cell if run locally\n!mkdir 'cache'\n!mkdir 'cache/torch'\n!mkdir 'cache/torch/checkpoints'\n!cp '../input/pytorch-pretrained-models/semi_supervised_resnext50_32x4-ddb3e555.pth' 'cache/torch/checkpoints/'\ntorch.hub.DEFAULT_CACHE_DIR = 'cache'","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"sz = 128\nbs = 32\nnfolds = 4\nSEED = 2020\nN = 16 #number of tiles per image\nTRAIN = '../input/panda-tiles-16x128x128/train/'\nLABELS = '../input/prostate-cancer-grade-assessment/train.csv'","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"","execution_count":null},{"metadata":{"_kg_hide-input":true,"trusted":true},"cell_type":"code","source":"def seed_everything(seed):\n    random.seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = True\n\nseed_everything(SEED)","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Data","execution_count":null},{"metadata":{},"cell_type":"markdown","source":"Use stratified KFold split.","execution_count":null},{"metadata":{"trusted":true},"cell_type":"code","source":"df = pd.read_csv(LABELS).set_index('image_id')\nfiles = sorted(set([p[:32] for p in os.listdir(TRAIN)]))\ndf = df.loc[files]\ndf = df.reset_index()\nsplits = StratifiedKFold(n_splits=nfolds, random_state=SEED, shuffle=True)\nsplits = list(splits.split(df,df.isup_grade))\nfolds_splits = np.zeros(len(df)).astype(np.int)\nfor i in range(nfolds): folds_splits[splits[i][1]] = i\ndf['split'] = folds_splits\ndf.head()","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"Check [this kernel](https://www.kaggle.com/iafoss/panda-16x128x128-tiles) for image stats. Since I use zero padding and background corresponds to 255, I invert images as 255-img when load them. Therefore, the mean value is computed as '1 - val'.","execution_count":null},{"metadata":{"trusted":true},"cell_type":"code","source":"mean = torch.tensor([1.0-0.90949707, 1.0-0.8188697, 1.0-0.87795304])\nstd = torch.tensor([0.36357649, 0.49984502, 0.40477625])","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"The code below (in the hidden cell) creates ImageItemList capable of loading multiple tiles of an image. It is specific for fast.ai, and pure Pytorch code would be much simpler.","execution_count":null},{"metadata":{"_kg_hide-input":true,"trusted":true},"cell_type":"code","source":"def open_image(fn:PathOrStr, div:bool=True, convert_mode:str='RGB', cls:type=Image,\n        after_open:Callable=None)->Image:\n    with warnings.catch_warnings():\n        warnings.simplefilter(\"ignore\", UserWarning) # EXIF warning from TiffPlugin\n        x = PIL.Image.open(fn).convert(convert_mode)\n    if after_open: x = after_open(x)\n    x = pil2tensor(x,np.float32)\n    if div: x.div_(255)\n    return cls(1.0-x) #invert image for zero padding\n\nclass MImage(ItemBase):\n    def __init__(self, imgs):\n        self.obj, self.data = \\\n          (imgs), [(imgs[i].data - mean[...,None,None])/std[...,None,None] for i in range(len(imgs))]\n    \n    def apply_tfms(self, tfms,*args, **kwargs):\n        for i in range(len(self.obj)):\n            self.obj[i] = self.obj[i].apply_tfms(tfms, *args, **kwargs)\n            self.data[i] = (self.obj[i].data - mean[...,None,None])/std[...,None,None]\n        return self\n    \n    def __repr__(self): return f'{self.__class__.__name__} {img.shape for img in self.obj}'\n    def to_one(self):\n        img = torch.stack(self.data,1)\n        img = img.view(3,-1,N,sz,sz).permute(0,1,3,2,4).contiguous().view(3,-1,sz*N)\n        return Image(1.0 - (mean[...,None,None]+img*std[...,None,None]))\n\nclass MImageItemList(ImageList):\n    def __init__(self, *args, **kwargs):\n        super().__init__(*args, **kwargs)\n    \n    def __len__(self)->int: return len(self.items) or 1 \n    \n    def get(self, i):\n        fn = Path(self.items[i])\n        fnames = [Path(str(fn)+'_'+str(i)+'.png')for i in range(N)]\n        imgs = [open_image(fname, convert_mode=self.convert_mode, after_open=self.after_open)\n               for fname in fnames]\n        return MImage(imgs)\n\n    def reconstruct(self, t):\n        return MImage([mean[...,None,None]+_t*std[...,None,None] for _t in t])\n    \n    def show_xys(self, xs, ys, figsize:Tuple[int,int]=(300,50), **kwargs):\n        rows = min(len(xs),8)\n        fig, axs = plt.subplots(rows,1,figsize=figsize)\n        for i, ax in enumerate(axs.flatten() if rows > 1 else [axs]):\n            xs[i].to_one().show(ax=ax, y=ys[i], **kwargs)\n        plt.tight_layout()\n        \n\n#collate function to combine multiple images into one tensor\ndef MImage_collate(batch:ItemsList)->Tensor:\n    result = torch.utils.data.dataloader.default_collate(to_data(batch))\n    if isinstance(result[0],list):\n        result = [torch.stack(result[0],1),result[1]]\n    return result","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def get_data(fold=0):\n    return (MImageItemList.from_df(df, path='.', folder=TRAIN, cols='image_id')\n      .split_by_idx(df.index[df.split == fold].tolist())\n      .label_from_df(cols=['isup_grade'])\n      .transform(get_transforms(flip_vert=True,max_rotate=15),size=sz,padding_mode='zeros')\n      .databunch(bs=bs,num_workers=4))\n\ndata = get_data(0)\ndata.show_batch()","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Model","execution_count":null},{"metadata":{},"cell_type":"markdown","source":"The code below implements Concat Tile pooling idea. As a backbone I use [Semi-Weakly Supervised ImageNet pretrained ResNeXt50 model](https://github.com/facebookresearch/semi-supervised-ImageNet1K-models), which worked for me quite well in a number of previous competitions.","execution_count":null},{"metadata":{"trusted":true},"cell_type":"code","source":"class Model(nn.Module):\n    def __init__(self, arch='resnext50_32x4d_ssl', n=6, pre=True):\n        super().__init__()\n        m = torch.hub.load('facebookresearch/semi-supervised-ImageNet1K-models', arch)\n        self.enc = nn.Sequential(*list(m.children())[:-2])       \n        nc = list(m.children())[-1].in_features\n        self.head = nn.Sequential(AdaptiveConcatPool2d(),Flatten(),nn.Linear(2*nc,512),\n                            Mish(),nn.BatchNorm1d(512), nn.Dropout(0.5),nn.Linear(512,n))\n        \n    def forward(self, *x):\n        shape = x[0].shape\n        n = len(x)\n        x = torch.stack(x,1).view(-1,shape[1],shape[2],shape[3])\n        #x: bs*N x 3 x 128 x 128\n        x = self.enc(x)\n        #x: bs*N x C x 4 x 4\n        shape = x.shape\n        #concatenate the output for tiles into a single map\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: bs x C x N*4 x 4\n        x = self.head(x)\n        #x: bs x n\n        \n        x = F.softmax(x,dim=1) # output probabilities\n        return x","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Training","execution_count":null},{"metadata":{"trusted":true},"cell_type":"code","source":"fname = 'RNXT50'\npred,prob,target = [],[],[]\nfor fold in range(2,nfolds):\n    data = get_data(fold)\n    model = Model()\n    learn = Learner(data, model, loss_func=nn.CrossEntropyLoss(), opt_func=Over9000, \n                metrics=[KappaScore(weights='quadratic')]).to_fp16()\n    logger = CSVLogger(learn, f'log_{fname}_{fold}')\n    learn.clip_grad = 1.0\n    learn.split([model.head])\n    learn.unfreeze()\n\n    cyc_len = 16\n    learn.fit_one_cycle(cyc_len, max_lr=1e-4, div_factor=100, pct_start=0.0, \n      callbacks = [SaveModelCallback(learn,name=f'model',monitor='kappa_score')])\n    torch.save(learn.model.state_dict(), f'{fname}_{fold}.pth')\n    \n    learn.model.eval()\n    with torch.no_grad():\n        for step, (x, y) in progress_bar(enumerate(data.dl(DatasetType.Valid)),\n                                     total=len(data.dl(DatasetType.Valid))):\n            p = learn.model(*x)\n            prob.append(p.float().cpu())\n            target.append(y.cpu())\n        \n     ","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"probs = torch.cat(prob,0).numpy()      \npreds = torch.argmax(torch.cat(prob,0),1).numpy()\nt = torch.cat(target)\nprint(probs,preds,t)\nprint(cohen_kappa_score(t,preds,weights='quadratic'))\nprint(confusion_matrix(t,preds))","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_kg_hide-input":true,"_kg_hide-output":true},"cell_type":"code","source":"!rm -r 'cache'","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}