{"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":"markdown","source":"## Introduction\nOne of the most challenging aspects of this competition is how to exclude the background and incorporate the areas of interest into the model. This problem can be solved by extracting tissue tiles([-> My notebook: a fast tile generation](https://www.kaggle.com/code/analokamus/a-fast-tile-generation)).\n\nHowever, the next problem is how to associate multiple tiles that may not contain a signal with a single label.\n**Multi Instance Learning (MIL)** solves this problem.\n\n![](http://drive.google.com/uc?export=view&id=1gGoZEmxp23UluGWCRE829NRCe2gew-BA)","metadata":{}},{"cell_type":"code","source":"!pip install torch-summary\n!pip install timm","metadata":{"execution":{"iopub.status.busy":"2022-07-17T01:15:05.731394Z","iopub.execute_input":"2022-07-17T01:15:05.731788Z","iopub.status.idle":"2022-07-17T01:15:26.059767Z","shell.execute_reply.started":"2022-07-17T01:15:05.731757Z","shell.execute_reply":"2022-07-17T01:15:26.058379Z"},"_kg_hide-output":true,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch.nn as nn\nfrom timm import create_model\nfrom torchsummary import summary","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-07-17T01:25:28.453484Z","iopub.execute_input":"2022-07-17T01:25:28.453855Z","iopub.status.idle":"2022-07-17T01:25:28.463431Z","shell.execute_reply.started":"2022-07-17T01:25:28.453827Z","shell.execute_reply":"2022-07-17T01:25:28.462341Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Flatten(nn.Module):\n    def __init__(self, dim=1):\n        super().__init__()\n        self.dim = dim\n\n    def forward(self, x): \n        input_shape = x.shape\n        output_shape = [input_shape[i] for i in range(self.dim)] + [-1]\n        return x.view(*output_shape)","metadata":{"execution":{"iopub.status.busy":"2022-07-17T01:15:30.959106Z","iopub.execute_input":"2022-07-17T01:15:30.959512Z","iopub.status.idle":"2022-07-17T01:15:30.965782Z","shell.execute_reply.started":"2022-07-17T01:15:30.959471Z","shell.execute_reply":"2022-07-17T01:15:30.964829Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class SimpleMIL(nn.Module):\n\n    def __init__(\n        self, \n        model_name, \n        num_instances=4, \n        num_classes=1, \n        pretrained=True):\n        super().__init__()\n\n        self.num_instances = num_instances\n        self.encoder = create_model(\n            model_name, \n            pretrained=pretrained, \n            num_classes=num_classes)\n        enc_type = self.encoder.__class__.__name__\n        feature_dim = self.encoder.get_classifier().in_features\n        self.head = nn.Sequential(\n            nn.AdaptiveMaxPool2d(1), Flatten(),\n            nn.Linear(feature_dim, 256), nn.ReLU(inplace=True), \n            nn.Linear(256, num_classes)\n        )\n\n    def forward(self, x):\n        # x: bs x N x C x W x W\n        bs, _, ch, w, h = x.shape\n        x = x.view(bs*self.num_instances, ch, w, h) # x: N bs x C x W x W\n        x = self.encoder.forward_features(x) # x: N bs x C' x W' x W'\n\n        # Concat and pool\n        bs2, ch2, w2, h2 = x.shape\n        x = x.view(-1, self.num_instances, ch2, w2, h2).permute(0, 2, 1, 3, 4)\\\n            .contiguous().view(bs, ch2, self.num_instances*w2, h2) # x: bs x C' x N W'' x W''\n        x = self.head(x)\n\n        return x","metadata":{"execution":{"iopub.status.busy":"2022-07-17T01:24:59.101357Z","iopub.execute_input":"2022-07-17T01:24:59.101771Z","iopub.status.idle":"2022-07-17T01:24:59.111251Z","shell.execute_reply.started":"2022-07-17T01:24:59.101737Z","shell.execute_reply":"2022-07-17T01:24:59.110106Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"mil = SimpleMIL('resnet18', pretrained=False)","metadata":{"execution":{"iopub.status.busy":"2022-07-17T01:24:59.614445Z","iopub.execute_input":"2022-07-17T01:24:59.614822Z","iopub.status.idle":"2022-07-17T01:24:59.812919Z","shell.execute_reply.started":"2022-07-17T01:24:59.614791Z","shell.execute_reply":"2022-07-17T01:24:59.811803Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"_ = summary(mil, (4, 3, 256, 256))","metadata":{"execution":{"iopub.status.busy":"2022-07-17T01:26:10.495843Z","iopub.execute_input":"2022-07-17T01:26:10.496888Z","iopub.status.idle":"2022-07-17T01:26:10.914732Z","shell.execute_reply.started":"2022-07-17T01:26:10.49685Z","shell.execute_reply":"2022-07-17T01:26:10.912645Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}