{"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":"# Problem Analysis\nWe have to segment instances of microvascular structures like arterioles and stuff... given an image\nimages are ... PAS-stained histology images\n","metadata":{}},{"cell_type":"markdown","source":"# STARTER","metadata":{}},{"cell_type":"markdown","source":"## Explore","metadata":{}},{"cell_type":"code","source":"!pip install git+https://www.github.com/raufie/train_utils.git","metadata":{"execution":{"iopub.status.busy":"2023-07-21T17:27:16.090994Z","iopub.execute_input":"2023-07-21T17:27:16.091436Z","iopub.status.idle":"2023-07-21T17:27:35.112098Z","shell.execute_reply.started":"2023-07-21T17:27:16.091406Z","shell.execute_reply":"2023-07-21T17:27:35.110833Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from PIL import Image, ImageDraw\nimport json\nimport matplotlib.pyplot as plt\nimport cv2\nimport numpy as np\nimport random\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport tqdm\nimport torchvision.transforms as transforms\nimport pandas as pd","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-07-21T17:27:35.116010Z","iopub.execute_input":"2023-07-21T17:27:35.116367Z","iopub.status.idle":"2023-07-21T17:27:41.315379Z","shell.execute_reply.started":"2023-07-21T17:27:35.116335Z","shell.execute_reply":"2023-07-21T17:27:41.314349Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Data Manager","metadata":{"execution":{"iopub.status.busy":"2023-07-07T15:50:23.798989Z","iopub.execute_input":"2023-07-07T15:50:23.799482Z","iopub.status.idle":"2023-07-07T15:50:23.828086Z","shell.execute_reply.started":"2023-07-07T15:50:23.799440Z","shell.execute_reply":"2023-07-07T15:50:23.826621Z"}}},{"cell_type":"code","source":"class DataManager:\n    def __init__(self, polygon_folder, train_folder = \"/kaggle/input/hubmap-hacking-the-human-vasculature/train/\", train_val_split=0.2, transform=None):\n        self.train_folder = train_folder\n        self.folder = polygon_folder\n        self.instances = self.__get_instances(polygon_folder)\n        \n        \n        \n        self.tile_data = pd.read_csv(\"/kaggle/input/hubmap-hacking-the-human-vasculature/tile_meta.csv\")\n        self.wsi_data = pd.read_csv(\"/kaggle/input/hubmap-hacking-the-human-vasculature/wsi_meta.csv\")\n        self.unlabelled_instances = self.__get_unlabelled()\n        \n        \n        \n#         splits\n\n        self.train_val_split = train_val_split\n        self.train_indices, self.val_indices = self.__initialize_indices()\n        \n        if transform == None:\n            transforms.Compose([\n                transforms.PILToTensor()\n            ])\n        \n        self.transform = transform\n\n# HELPERS\n    def __get_unlabelled(self):\n        instances = list(self.tile_data[self.tile_data[\"dataset\"] == 3][\"id\"])\n        return instances\n    def __get_instances(self, folder):\n        \n        try:\n            f = open(folder, \"r\")\n            jsonl = list(f)\n            instances = []\n\n            for json_str in jsonl:\n                instances.append(json.loads(json_str))\n            \n            \n        \n            return instances\n        except Exception as E:\n            print(\"error loading folder\")\n            print(E)\n            return []\n    \n    def __initialize_id_to_instance(self):\n        hashmap = {}\n        for i, instance in enumerate(self.instances):\n            hashmap[instance[\"id\"]] = i\n        return hashmap\n    \n    def __initialize_indices(self):\n        \n        all_indices = list(range(len(self.instances)))\n        random.shuffle(all_indices)\n        split = int(self.train_val_split*len(self.instances))\n        \n#         \n        \n        return all_indices[:split], all_indices[split:]\n            \n#     REAL METHODS\n    def id_to_instance(self, _id):\n        index = self._id_to_instance.get(_id, -1)\n        if index == -1:\n            return {}\n        return self.instances[index]\n    \n    def view_instance(self, instance_index, object_type = \"blood_vessel\"):\n        instance = self.instances[instance_index]\n#         return instance\n        polygons_data = instance[\"annotations\"]\n        \n        image = cv2.imread(f\"{self.train_folder}{instance['id']}.tif\")\n        # Set axis limits\n        \n        for polygon in polygons_data:\n            \n#             if polygon['type'] != object_type and polygon['type'] != 'unsure':\n            if polygon['type'] != object_type:\n                continue\n                \n            for polygon_coords in polygon[\"coordinates\"]:\n                \n                pts = [[int(x), int(y)] for x, y in polygon_coords]\n                pts = np.array(pts, np.int32)\n                cv2.polylines(image, [pts], isClosed=True, color=(252, 255, 51), thickness=3)\n\n        return image\n    \n    def get_unlabelled_batch(self, batch_size):\n        instances = random.sample(self.unlabelled_instances, batch_size)\n        outs = []\n        for instance in instances:\n            instance = f\"{self.train_folder}{instance}.tif\"\n            outs.append(self.transform(Image.open(instance)))\n        outs = torch.cat(outs).unsqueeze(1)\n        return outs, outs\n        \n    \n    def get_batch(self, batch_size, dataset=\"train\", expert_only = False):\n#         @params: batch_size , dataset=\"train\" or \"val\"\n#         purpose: inputs: tensor(batch, 1,512,512), binary masks: tensor(batch, 1, 512, 512)\n        indices = self.train_indices\n#         if dataset==\"train\" and expert_only:\n#             self.tile_data[\"\"]\n        \n        if dataset==\"val\":\n            indices = self.val_indices\n            \n#         sample batch_size samples\n# instead of doing random.sample use a while loop that filters expert frmo no expert only\n        indices = random.sample(indices, batch_size) \n#     \n        inputs = []# list of tensors\n        binary_masks = [] # list of tensors\n        \n        \n        for i in indices:\n            instance = self[i]\n            inputs.append(self.transform(Image.open(instance[1])))\n            binary_masks.append(self.__polygon_to_mask(i))\n        \n        inputs = torch.cat(inputs).unsqueeze(1)\n        \n        binary_masks = torch.cat(binary_masks).unsqueeze(1)\n#         need N,C,H,W for conv layers\n        \n        \n        return inputs, binary_masks\n#\n    def __polygon_to_mask(self, ix, object_type=\"blood_vessel\"):\n        instance = self.instances[ix]\n        polygons_data = instance[\"annotations\"]\n        \n        mask = Image.new(\"L\", (512,512), 0)  \n\n        draw = ImageDraw.Draw(mask)\n        for polygon in polygons_data:\n#             actually the poly\n#             if polygon['type'] != object_type and polygon['type'] != 'unsure':\n            if polygon['type'] != object_type:\n                continue\n            \n            \n            \n            for polygon_coords in polygon[\"coordinates\"]:\n                points = [(poly[0], poly[1]) for poly in polygon_coords]\n\n                draw.polygon(points, fill=255)\n#                 for x, y in polygon_coords:\n#                     mask[y,x] = 1\n        mask = transforms.functional.to_tensor(mask)\n        return mask\n        \n    \n    def __getitem__(self, ix):\n        mask = self.__polygon_to_mask(ix)\n        return self.instances[ix], self.train_folder+f\"{self.instances[ix]['id']}.tif\", mask\n    \n    def __len__(self):\n        return len(self.instances)\n    \n    def overlay_mask(self, main_image, mask):\n#         convert to 3d\n        image_rgb = F.to_pil_image(image).convert('RGB')\n        mask_rgb = mask.repeat(3, 1, 1)\n#     add a colored mask\n        ","metadata":{"execution":{"iopub.status.busy":"2023-07-21T17:27:41.317118Z","iopub.execute_input":"2023-07-21T17:27:41.318146Z","iopub.status.idle":"2023-07-21T17:27:41.344647Z","shell.execute_reply.started":"2023-07-21T17:27:41.318107Z","shell.execute_reply":"2023-07-21T17:27:41.343787Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"transform = transforms.Compose([\n    transforms.PILToTensor(),\n    transforms.Grayscale(),\n    transforms.Lambda(lambda x: x/512.)\n\n])","metadata":{"execution":{"iopub.status.busy":"2023-07-21T17:27:41.348765Z","iopub.execute_input":"2023-07-21T17:27:41.349061Z","iopub.status.idle":"2023-07-21T17:27:41.364343Z","shell.execute_reply.started":"2023-07-21T17:27:41.349035Z","shell.execute_reply":"2023-07-21T17:27:41.363369Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dataManager = DataManager(\"/kaggle/input/hubmap-hacking-the-human-vasculature/polygons.jsonl\", transform=transform)","metadata":{"execution":{"iopub.status.busy":"2023-07-21T17:27:41.366284Z","iopub.execute_input":"2023-07-21T17:27:41.366869Z","iopub.status.idle":"2023-07-21T17:27:46.177548Z","shell.execute_reply.started":"2023-07-21T17:27:41.366766Z","shell.execute_reply":"2023-07-21T17:27:46.176519Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"mask = dataManager[79][2]\nplt.imshow(mask[0])","metadata":{"execution":{"iopub.status.busy":"2023-07-21T17:27:46.178942Z","iopub.execute_input":"2023-07-21T17:27:46.181393Z","iopub.status.idle":"2023-07-21T17:27:46.628315Z","shell.execute_reply.started":"2023-07-21T17:27:46.181355Z","shell.execute_reply":"2023-07-21T17:27:46.627398Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Data Analysis","metadata":{}},{"cell_type":"code","source":"df_tile = pd.read_csv(\"/kaggle/input/hubmap-hacking-the-human-vasculature/tile_meta.csv\")","metadata":{"execution":{"iopub.status.busy":"2023-07-21T17:27:46.630011Z","iopub.execute_input":"2023-07-21T17:27:46.630388Z","iopub.status.idle":"2023-07-21T17:27:46.649028Z","shell.execute_reply.started":"2023-07-21T17:27:46.630354Z","shell.execute_reply":"2023-07-21T17:27:46.648095Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"len(dataManager)","metadata":{"execution":{"iopub.status.busy":"2023-07-21T17:27:46.651476Z","iopub.execute_input":"2023-07-21T17:27:46.651867Z","iopub.status.idle":"2023-07-21T17:27:46.658328Z","shell.execute_reply.started":"2023-07-21T17:27:46.651834Z","shell.execute_reply":"2023-07-21T17:27:46.657280Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Segmentation Task\n\nSo we are going to output a pixel mask. So architectures like UNet... \n\nThen we'll convert the mask to polygons later and do whatever encoding the competition wants us to do.\n\nWe'll also have to write our loss functions and metrics like mAP... let's do it after writing a basic model with a loss function to get immediate feedback\n\nThis was my initial strategy to train a model for going through the whole model dev->submission pipeline to make it easier to experiment later... its done, now its time to cook jessie\n","metadata":{}},{"cell_type":"markdown","source":"## Using and analyzing Meta data","metadata":{}},{"cell_type":"markdown","source":"## Demons to beat\n- WSI bias (~422 pro annotations, ~1211 non-pro annotations)\n- overfitting","metadata":{}},{"cell_type":"markdown","source":"# Model","metadata":{}},{"cell_type":"markdown","source":"# Layers exploration ","metadata":{}},{"cell_type":"code","source":"class DoubleConv(nn.Module):\n    \"\"\"(convolution => [BN] => ReLU) * 2\"\"\"\n\n    def __init__(self, in_channels, out_channels, mid_channels=None):\n        super().__init__()\n        if not mid_channels:\n            mid_channels = out_channels\n        self.double_conv = nn.Sequential(\n            nn.Conv2d(in_channels, mid_channels, kernel_size=3, padding=1),\n            nn.BatchNorm2d(mid_channels),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(mid_channels, out_channels, kernel_size=3, padding=1),\n            nn.BatchNorm2d(out_channels),\n            nn.ReLU(inplace=True)\n        )\n\n    def forward(self, x):\n        return self.double_conv(x)\n\n\nclass Down(nn.Module):\n    \"\"\"Downscaling with maxpool then double conv\"\"\"\n\n    def __init__(self, in_channels, out_channels):\n        super().__init__()\n        self.maxpool_conv = nn.Sequential(\n            nn.MaxPool2d(2),\n            DoubleConv(in_channels, out_channels)\n        )\n\n    def forward(self, x):\n        return self.maxpool_conv(x)\n\n\nclass Up(nn.Module):\n    \"\"\"Upscaling then double conv\"\"\"\n\n    def __init__(self, in_channels, out_channels, bilinear=True):\n        super().__init__()\n\n        # if bilinear, use the normal convolutions to reduce the number of channels\n        if bilinear:\n            self.up = nn.Upsample(scale_factor=2, mode='bilinear', align_corners=True)\n            self.conv = DoubleConv(in_channels, out_channels, in_channels // 2)\n        else:\n            self.up = nn.ConvTranspose2d(in_channels , in_channels // 2, kernel_size=2, stride=2)\n            self.conv = DoubleConv(in_channels, out_channels)\n\n\n    def forward(self, x1, x2):\n        x1 = self.up(x1)\n        # input is CHW\n        diffY = x2.size()[2] - x1.size()[2]\n        diffX = x2.size()[3] - x1.size()[3]\n\n        x1 = F.pad(x1, [diffX // 2, diffX - diffX // 2,\n                        diffY // 2, diffY - diffY // 2])\n        # if you have padding issues, see\n        # https://github.com/HaiyongJiang/U-Net-Pytorch-Unstructured-Buggy/commit/0e854509c2cea854e247a9c615f175f76fbb2e3a\n        # https://github.com/xiaopeng-liao/Pytorch-UNet/commit/8ebac70e633bac59fc22bb5195e513d5832fb3bd\n        x = torch.cat([x2, x1], dim=1)\n        return self.conv(x)\n\n\nclass OutConv(nn.Module):\n    def __init__(self, in_channels, out_channels):\n        super(OutConv, self).__init__()\n        self.conv = nn.Conv2d(in_channels, out_channels, kernel_size=1)\n\n    def forward(self, x):\n        return self.conv(x)","metadata":{"execution":{"iopub.status.busy":"2023-07-21T17:27:46.660124Z","iopub.execute_input":"2023-07-21T17:27:46.660852Z","iopub.status.idle":"2023-07-21T17:27:46.679664Z","shell.execute_reply.started":"2023-07-21T17:27:46.660817Z","shell.execute_reply":"2023-07-21T17:27:46.678545Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class UNet(nn.Module):\n    def __init__(self, n_channels, n_classes, bilinear=True):\n        super(UNet, self).__init__()\n        self.n_channels = n_channels\n        self.n_classes = n_classes\n        self.bilinear = bilinear\n\n        self.inc = DoubleConv(n_channels, 64)\n        self.down1 = Down(64, 128)\n        self.down2 = Down(128, 256)\n        self.down3 = Down(256, 512)\n        factor = 2 if bilinear else 1\n        self.down4 = Down(512, 1024 // factor)\n        self.up1 = Up(1024, 512 // factor, bilinear)\n        self.up2 = Up(512, 256 // factor, bilinear)\n        self.up3 = Up(256, 128 // factor, bilinear)\n        self.up4 = Up(128, 64, bilinear)\n        self.outc = OutConv(64, n_classes)\n\n    def forward(self, x):\n        x1 = self.inc(x)\n        x2 = self.down1(x1)\n        x3 = self.down2(x2)\n        x4 = self.down3(x3)\n        x5 = self.down4(x4)\n        x = self.up1(x5, x4)\n        x = self.up2(x, x3)#skp x3\n        x = self.up3(x, x2)#skp x2\n        x = self.up4(x, x1)#skp x1\n        logits = self.outc(x)\n        return logits\n    \n","metadata":{"execution":{"iopub.status.busy":"2023-07-21T17:27:46.685943Z","iopub.execute_input":"2023-07-21T17:27:46.686232Z","iopub.status.idle":"2023-07-21T17:27:46.700445Z","shell.execute_reply.started":"2023-07-21T17:27:46.686208Z","shell.execute_reply":"2023-07-21T17:27:46.699329Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Loss and transforms","metadata":{}},{"cell_type":"code","source":"DEVICE = \"cuda\"","metadata":{"execution":{"iopub.status.busy":"2023-07-21T17:27:46.701862Z","iopub.execute_input":"2023-07-21T17:27:46.702816Z","iopub.status.idle":"2023-07-21T17:27:46.713318Z","shell.execute_reply.started":"2023-07-21T17:27:46.702789Z","shell.execute_reply":"2023-07-21T17:27:46.712451Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def to_mask(logits, threshold=0.5):\n    probs = torch.nn.functional.sigmoid(logits)\n    mask = probs > threshold\n    return mask","metadata":{"execution":{"iopub.status.busy":"2023-07-21T17:27:46.714745Z","iopub.execute_input":"2023-07-21T17:27:46.715765Z","iopub.status.idle":"2023-07-21T17:27:46.727131Z","shell.execute_reply.started":"2023-07-21T17:27:46.715730Z","shell.execute_reply":"2023-07-21T17:27:46.726301Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = UNet(1,1)","metadata":{"execution":{"iopub.status.busy":"2023-07-21T17:27:46.728292Z","iopub.execute_input":"2023-07-21T17:27:46.728605Z","iopub.status.idle":"2023-07-21T17:27:46.920928Z","shell.execute_reply.started":"2023-07-21T17:27:46.728579Z","shell.execute_reply":"2023-07-21T17:27:46.919889Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = nn.DataParallel(model)","metadata":{"execution":{"iopub.status.busy":"2023-07-21T17:27:46.922259Z","iopub.execute_input":"2023-07-21T17:27:46.923239Z","iopub.status.idle":"2023-07-21T17:27:47.019411Z","shell.execute_reply.started":"2023-07-21T17:27:46.923203Z","shell.execute_reply":"2023-07-21T17:27:47.018399Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model.to(DEVICE)","metadata":{"execution":{"iopub.status.busy":"2023-07-21T17:27:47.020898Z","iopub.execute_input":"2023-07-21T17:27:47.021242Z","iopub.status.idle":"2023-07-21T17:27:51.519387Z","shell.execute_reply.started":"2023-07-21T17:27:47.021215Z","shell.execute_reply":"2023-07-21T17:27:51.518421Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.cuda.empty_cache()","metadata":{"execution":{"iopub.status.busy":"2023-07-21T17:27:51.521053Z","iopub.execute_input":"2023-07-21T17:27:51.521705Z","iopub.status.idle":"2023-07-21T17:27:51.528343Z","shell.execute_reply.started":"2023-07-21T17:27:51.521670Z","shell.execute_reply":"2023-07-21T17:27:51.527321Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import train_utils.data.StatsManager as StatsManager","metadata":{"execution":{"iopub.status.busy":"2023-07-21T17:27:51.530684Z","iopub.execute_input":"2023-07-21T17:27:51.531368Z","iopub.status.idle":"2023-07-21T17:27:51.541175Z","shell.execute_reply.started":"2023-07-21T17:27:51.531331Z","shell.execute_reply":"2023-07-21T17:27:51.540154Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training","metadata":{}},{"cell_type":"code","source":"def dice_bce_loss(logits, targets, eps = 1):\n#     get probs\n    probs = F.sigmoid(logits).view(-1)\n    targets = targets.view(-1)\n#     1- 0.77 e.g for a pixel\n\n#     dice coeff = 2*intersection/(probs.sum()+targets.sum())\n    iou = (probs*targets).sum()\n    dice_coeff = (2*iou + eps)/(probs.sum()+targets.sum()+eps)\n    \n    BCE = F.binary_cross_entropy(probs, targets, reduction='mean')\n    dice_loss =  1- dice_coeff\n# \n    return dice_loss*0.5 + BCE*0.5","metadata":{"execution":{"iopub.status.busy":"2023-07-21T17:27:51.542565Z","iopub.execute_input":"2023-07-21T17:27:51.544019Z","iopub.status.idle":"2023-07-21T17:27:51.558782Z","shell.execute_reply.started":"2023-07-21T17:27:51.543982Z","shell.execute_reply":"2023-07-21T17:27:51.557799Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# loss_function = torch.nn.BCELoss()\nloss_function = dice_bce_loss\n# def loss_function(logits, targets, eps=1):\n#     return dice_loss(logits, targets) + nn.BCEWithLogitsLoss()(logits, targets)","metadata":{"execution":{"iopub.status.busy":"2023-07-21T17:39:43.448172Z","iopub.execute_input":"2023-07-21T17:39:43.448586Z","iopub.status.idle":"2023-07-21T17:39:43.453285Z","shell.execute_reply.started":"2023-07-21T17:39:43.448553Z","shell.execute_reply":"2023-07-21T17:39:43.452232Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"stats = StatsManager(\"UNET_FOUNDATION\")","metadata":{"execution":{"iopub.status.busy":"2023-07-21T17:39:45.226261Z","iopub.execute_input":"2023-07-21T17:39:45.227336Z","iopub.status.idle":"2023-07-21T17:39:45.231666Z","shell.execute_reply.started":"2023-07-21T17:39:45.227299Z","shell.execute_reply":"2023-07-21T17:39:45.230268Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"min_loss = float(\"inf\")\nval_losses = [1.0]*10","metadata":{"execution":{"iopub.status.busy":"2023-07-21T17:39:45.483999Z","iopub.execute_input":"2023-07-21T17:39:45.484331Z","iopub.status.idle":"2023-07-21T17:39:45.490214Z","shell.execute_reply.started":"2023-07-21T17:39:45.484302Z","shell.execute_reply":"2023-07-21T17:39:45.489201Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"best_model_state_dict = model.state_dict()","metadata":{"execution":{"iopub.status.busy":"2023-07-21T17:39:45.698928Z","iopub.execute_input":"2023-07-21T17:39:45.699248Z","iopub.status.idle":"2023-07-21T17:39:45.705656Z","shell.execute_reply.started":"2023-07-21T17:39:45.699221Z","shell.execute_reply":"2023-07-21T17:39:45.704558Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train(n_steps, batch_size, lr = 0.0001, device=\"cuda\", acc_ = 4, early_stopping_window=10, prev_min_loss=None):\n    \n    optimizer = torch.optim.Adam(model.parameters(), lr=lr)\n    min_loss = float(\"inf\") if prev_min_loss == None else prev_min_loss\n    val_losses = [1.0]*early_stopping_window\n    \n    \n    \n    for i in (pbar:=tqdm.tqdm(range(n_steps))):\n        inputs, binary_masks = dataManager.get_batch(batch_size, dataset=\"train\")\n        inputs = inputs.to(device)\n        binary_masks = binary_masks.to(device)\n        model.train()\n        out = model(inputs)\n        \n#         transform image\n#         loss = DiceCoeff.apply(out, binary_masks.to(device).float())\n        loss = loss_function(out, binary_masks.float())\n        \n        loss.backward()\n        \n        \n        if (i+1)%acc_==0:\n            \n            \n            optimizer.step()\n            optimizer.zero_grad()\n        \n        torch.cuda.empty_cache()\n        stats.add(\"loss\",loss.item())\n        \n#         CALCULATE VAL LOSS\n        model.eval()\n        inputs, binary_masks = dataManager.get_batch(batch_size, dataset=\"val\")\n        inputs = inputs.to(device)\n        binary_masks = binary_masks.to(device)\n        \n        \n        with torch.no_grad():\n            out = model(inputs)\n            val_loss = loss_function(out, binary_masks.float())\n            torch.cuda.empty_cache()\n        \n        stats.add(\"val_loss\",val_loss.item())\n        \n        val_losses[i%early_stopping_window] = val_loss.item()\n        avg_loss = sum(val_losses)/early_stopping_window\n        stats.add(\"avg_loss\", avg_loss)\n        \n        \n        if avg_loss < min_loss :\n            min_loss = avg_loss\n            best_model_state_dict = model.state_dict()\n            \n        pbar.set_description(f\"step: {i+1}/{n_steps} train_loss: {loss.item():.2f} val_loss: {val_loss.item():.2f} avg_loss: {avg_loss}\")\n    return best_model_state_dict","metadata":{"execution":{"iopub.status.busy":"2023-07-21T17:39:45.986604Z","iopub.execute_input":"2023-07-21T17:39:45.986943Z","iopub.status.idle":"2023-07-21T17:39:45.999725Z","shell.execute_reply.started":"2023-07-21T17:39:45.986916Z","shell.execute_reply":"2023-07-21T17:39:45.998755Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# TRAINING SEGMENTATION MDOEL","metadata":{}},{"cell_type":"code","source":"# loss_function = nn.MSELoss()\n# best_model = train(1000,16, acc_=4, lr=0.0001)\n","metadata":{"execution":{"iopub.status.busy":"2023-07-21T17:39:46.862204Z","iopub.execute_input":"2023-07-21T17:39:46.863345Z","iopub.status.idle":"2023-07-21T17:39:46.868032Z","shell.execute_reply.started":"2023-07-21T17:39:46.863301Z","shell.execute_reply":"2023-07-21T17:39:46.866871Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# model.load_state_dict(torch.load(\"/kaggle/working/filled_model_3.pt\"))","metadata":{"execution":{"iopub.status.busy":"2023-07-21T17:39:48.500347Z","iopub.execute_input":"2023-07-21T17:39:48.501451Z","iopub.status.idle":"2023-07-21T17:39:48.506548Z","shell.execute_reply.started":"2023-07-21T17:39:48.501406Z","shell.execute_reply":"2023-07-21T17:39:48.505312Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.cuda.empty_cache()","metadata":{"execution":{"iopub.status.busy":"2023-07-21T17:39:50.340381Z","iopub.execute_input":"2023-07-21T17:39:50.341098Z","iopub.status.idle":"2023-07-21T17:39:50.383882Z","shell.execute_reply.started":"2023-07-21T17:39:50.341062Z","shell.execute_reply":"2023-07-21T17:39:50.382732Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# freeze all layers except last\n","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"best_model = train(200,16, acc_=3, lr=0.0001)\n#prev_min_loss=min(stats.get(\"val_loss\"))","metadata":{"execution":{"iopub.status.busy":"2023-07-21T17:46:59.281569Z","iopub.execute_input":"2023-07-21T17:46:59.281951Z","iopub.status.idle":"2023-07-21T17:57:53.078105Z","shell.execute_reply.started":"2023-07-21T17:46:59.281916Z","shell.execute_reply":"2023-07-21T17:57:53.077002Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model.load_state_dict(best_model)","metadata":{"execution":{"iopub.status.busy":"2023-07-21T17:58:42.376345Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"with torch.no_grad():\n    inputs, binary_masks = dataManager.get_batch(16, dataset=\"val\")\n    inputs = inputs.to(DEVICE)\n    binary_masks = binary_masks.to(DEVICE)\n    out = model(inputs)\n    val_loss = loss_function(out, binary_masks.float())\n    torch.cuda.empty_cache()\n    print(val_loss)\n","metadata":{"execution":{"iopub.status.busy":"2023-07-21T17:59:10.761012Z","iopub.execute_input":"2023-07-21T17:59:10.761677Z","iopub.status.idle":"2023-07-21T17:59:11.611567Z","shell.execute_reply.started":"2023-07-21T17:59:10.761641Z","shell.execute_reply":"2023-07-21T17:59:11.610573Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.plot(stats.get(\"loss\"), label=\"train loss\")\nplt.plot(stats.get(\"val_loss\")[:], label=\"val loss\")\nplt.plot(stats.get(\"avg_loss\")[:], label=\"avg val loss\")\nplt.legend()","metadata":{"execution":{"iopub.status.busy":"2023-07-21T17:59:14.356598Z","iopub.execute_input":"2023-07-21T17:59:14.358800Z","iopub.status.idle":"2023-07-21T17:59:14.696278Z","shell.execute_reply.started":"2023-07-21T17:59:14.358755Z","shell.execute_reply":"2023-07-21T17:59:14.695389Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# model.load_state_dict(torch.load(\"UNET_BEST_1.pt\"))","metadata":{"execution":{"iopub.status.busy":"2023-07-21T17:59:25.183000Z","iopub.execute_input":"2023-07-21T17:59:25.183380Z","iopub.status.idle":"2023-07-21T17:59:25.188709Z","shell.execute_reply.started":"2023-07-21T17:59:25.183349Z","shell.execute_reply":"2023-07-21T17:59:25.187570Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model.load_state_dict(best_model)","metadata":{"execution":{"iopub.status.busy":"2023-07-20T09:35:22.350944Z","iopub.execute_input":"2023-07-20T09:35:22.351327Z","iopub.status.idle":"2023-07-20T09:35:22.364504Z","shell.execute_reply.started":"2023-07-20T09:35:22.351297Z","shell.execute_reply":"2023-07-20T09:35:22.363422Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.save(model.state_dict(), \"UNET_WITH_FOUNDATION.pt\")","metadata":{"execution":{"iopub.status.busy":"2023-07-21T18:00:51.619996Z","iopub.execute_input":"2023-07-21T18:00:51.620375Z","iopub.status.idle":"2023-07-21T18:00:51.807091Z","shell.execute_reply.started":"2023-07-21T18:00:51.620343Z","shell.execute_reply":"2023-07-21T18:00:51.806019Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"stats.save()","metadata":{"execution":{"iopub.status.busy":"2023-07-21T17:59:30.705578Z","iopub.execute_input":"2023-07-21T17:59:30.705963Z","iopub.status.idle":"2023-07-21T17:59:30.714596Z","shell.execute_reply.started":"2023-07-21T17:59:30.705926Z","shell.execute_reply":"2023-07-21T17:59:30.713366Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def to_mask(logits, threshold=0.5, debug=False):\n    probs = torch.nn.functional.sigmoid(logits)\n#     print(torch.where (probs > threshold))\n    mask = probs.float() >= threshold\n    if debug:\n        return logits\n    return mask\n    ","metadata":{"execution":{"iopub.status.busy":"2023-07-21T17:59:30.984411Z","iopub.execute_input":"2023-07-21T17:59:30.984773Z","iopub.status.idle":"2023-07-21T17:59:30.991525Z","shell.execute_reply.started":"2023-07-21T17:59:30.984745Z","shell.execute_reply":"2023-07-21T17:59:30.990396Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"inputs, binary_masks = dataManager.get_batch(2, dataset=\"val\")\nwith torch.no_grad():\n    model.eval()\n    outs = model(inputs)\n    plt.figure(figsize=(7,7))\n    out_mask = to_mask(outs[0].detach().to(\"cpu\"), threshold=0.8, debug=False).permute(1,2,0)\n    plt.imshow(out_mask)\n    plt.show()\n\n","metadata":{"execution":{"iopub.status.busy":"2023-07-21T18:00:13.804048Z","iopub.execute_input":"2023-07-21T18:00:13.804440Z","iopub.status.idle":"2023-07-21T18:00:14.230550Z","shell.execute_reply.started":"2023-07-21T18:00:13.804409Z","shell.execute_reply":"2023-07-21T18:00:14.229433Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.figure(figsize=(7,7))\n\nplt.imshow(binary_masks[0].permute(1,2,0))","metadata":{"execution":{"iopub.status.busy":"2023-07-21T18:00:14.232561Z","iopub.execute_input":"2023-07-21T18:00:14.232921Z","iopub.status.idle":"2023-07-21T18:00:14.584273Z","shell.execute_reply.started":"2023-07-21T18:00:14.232886Z","shell.execute_reply":"2023-07-21T18:00:14.583260Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.imshow(inputs[0].permute(1,2,0))","metadata":{"execution":{"iopub.status.busy":"2023-07-20T08:58:39.974951Z","iopub.execute_input":"2023-07-20T08:58:39.975688Z","iopub.status.idle":"2023-07-20T08:58:40.356231Z","shell.execute_reply.started":"2023-07-20T08:58:39.975649Z","shell.execute_reply":"2023-07-20T08:58:40.355361Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# FOUNDATION MODEL","metadata":{"execution":{"iopub.status.busy":"2023-07-13T16:19:31.884574Z","iopub.execute_input":"2023-07-13T16:19:31.884936Z","iopub.status.idle":"2023-07-13T16:19:31.891616Z","shell.execute_reply.started":"2023-07-13T16:19:31.884906Z","shell.execute_reply":"2023-07-13T16:19:31.890582Z"}}},{"cell_type":"code","source":"def train_foundation(n_steps, batch_size, lr = 0.0001, device=\"cuda\", acc_ = 4, early_stopping_window=10, prev_min_loss=None):\n    \n    optimizer = torch.optim.Adam(model.parameters(), lr=lr)\n    min_loss = float(\"inf\") if prev_min_loss == None else prev_min_loss\n    val_losses = [1.0]*early_stopping_window\n    \n    \n    \n    for i in (pbar:=tqdm.tqdm(range(n_steps))):\n        inputs, binary_masks = dataManager.get_unlabelled_batch(batch_size)\n        inputs = inputs.to(device)\n        binary_masks = binary_masks.to(device)\n        model.train()\n        out = model(inputs)\n        \n#         transform image\n#         loss = DiceCoeff.apply(out, binary_masks.to(device).float())\n        loss = loss_function(out, binary_masks.float())\n        \n        loss.backward()\n        \n        \n        if (i+1)%acc_==0:\n            \n            \n            optimizer.step()\n            optimizer.zero_grad()\n        \n        torch.cuda.empty_cache()\n        stats.add(\"loss\",loss.item())\n        \n#         CALCULATE VAL LOSS\n        model.eval()\n        inputs, binary_masks = dataManager.get_batch(batch_size, dataset=\"val\")\n        inputs = inputs.to(device)\n        binary_masks = binary_masks.to(device)\n        \n        \n        with torch.no_grad():\n            out = model(inputs)\n            val_loss = loss_function(out, binary_masks.float())\n            torch.cuda.empty_cache()\n        \n        stats.add(\"val_loss\",val_loss.item())\n        \n        val_losses[i%early_stopping_window] = val_loss.item()\n        avg_loss = sum(val_losses)/early_stopping_window\n        stats.add(\"avg_loss\", avg_loss)\n        \n        \n        if avg_loss < min_loss :\n            min_loss = avg_loss\n            best_model_state_dict = model.state_dict()\n            \n        pbar.set_description(f\"step: {i+1}/{n_steps} train_loss: {loss.item():.2f} val_loss: {val_loss.item():.2f} avg_loss: {avg_loss}\")\n    return best_model_state_dict","metadata":{"execution":{"iopub.status.busy":"2023-07-21T17:27:58.641552Z","iopub.execute_input":"2023-07-21T17:27:58.642273Z","iopub.status.idle":"2023-07-21T17:27:58.654570Z","shell.execute_reply.started":"2023-07-21T17:27:58.642240Z","shell.execute_reply":"2023-07-21T17:27:58.653358Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model.load_state_dict(torch.load(\"/kaggle/working/foundation.pt\"))","metadata":{"execution":{"iopub.status.busy":"2023-07-21T17:30:30.807738Z","iopub.execute_input":"2023-07-21T17:30:30.808579Z","iopub.status.idle":"2023-07-21T17:30:30.886196Z","shell.execute_reply.started":"2023-07-21T17:30:30.808538Z","shell.execute_reply":"2023-07-21T17:30:30.885146Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"loss_function = nn.MSELoss()\nbest_model = train_foundation(100,16, acc_=4, lr=0.0001)\n","metadata":{"execution":{"iopub.status.busy":"2023-07-21T17:30:58.220139Z","iopub.execute_input":"2023-07-21T17:30:58.220493Z","iopub.status.idle":"2023-07-21T17:36:57.949335Z","shell.execute_reply.started":"2023-07-21T17:30:58.220464Z","shell.execute_reply":"2023-07-21T17:36:57.948265Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.save(model.state_dict(), \"foundation.pt\")","metadata":{"execution":{"iopub.status.busy":"2023-07-21T17:39:22.440286Z","iopub.execute_input":"2023-07-21T17:39:22.440697Z","iopub.status.idle":"2023-07-21T17:39:22.625812Z","shell.execute_reply.started":"2023-07-21T17:39:22.440664Z","shell.execute_reply":"2023-07-21T17:39:22.623450Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.plot(stats.get(\"loss\"))\n# plt.plot(stats.get(\"val_loss\"))","metadata":{"execution":{"iopub.status.busy":"2023-07-21T17:37:18.417870Z","iopub.execute_input":"2023-07-21T17:37:18.418252Z","iopub.status.idle":"2023-07-21T17:37:18.702818Z","shell.execute_reply.started":"2023-07-21T17:37:18.418221Z","shell.execute_reply":"2023-07-21T17:37:18.701876Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"inputs, binary_masks = dataManager.get_batch(2, dataset=\"val\")\nwith torch.no_grad():\n    model.eval()\n    outs = model(inputs)\n    plt.figure(figsize=(7,7))\n    out_mask = outs[0].detach().to(\"cpu\").permute(1,2,0)\n    plt.imshow(out_mask)\n    plt.show()\n","metadata":{"execution":{"iopub.status.busy":"2023-07-21T17:38:36.227267Z","iopub.execute_input":"2023-07-21T17:38:36.227956Z","iopub.status.idle":"2023-07-21T17:38:36.874387Z","shell.execute_reply.started":"2023-07-21T17:38:36.227919Z","shell.execute_reply":"2023-07-21T17:38:36.873408Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.imshow(inputs[0].permute(1,2,0).view(512,512))","metadata":{"execution":{"iopub.status.busy":"2023-07-21T17:38:49.642081Z","iopub.execute_input":"2023-07-21T17:38:49.642462Z","iopub.status.idle":"2023-07-21T17:38:50.035458Z","shell.execute_reply.started":"2023-07-21T17:38:49.642433Z","shell.execute_reply":"2023-07-21T17:38:50.034439Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from IPython.display import FileLink, HTML","metadata":{"execution":{"iopub.status.busy":"2023-07-21T18:01:43.666737Z","iopub.execute_input":"2023-07-21T18:01:43.667124Z","iopub.status.idle":"2023-07-21T18:01:43.672477Z","shell.execute_reply.started":"2023-07-21T18:01:43.667093Z","shell.execute_reply":"2023-07-21T18:01:43.671459Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_download_link(model, filename):\n    os.chdir(r'/kaggle/working')\n    torch.save(model.state_dict(), filename)\n\n    return FileLink(filename)","metadata":{"execution":{"iopub.status.busy":"2023-07-21T18:01:43.874329Z","iopub.execute_input":"2023-07-21T18:01:43.874981Z","iopub.status.idle":"2023-07-21T18:01:43.880789Z","shell.execute_reply.started":"2023-07-21T18:01:43.874932Z","shell.execute_reply":"2023-07-21T18:01:43.879569Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os","metadata":{"execution":{"iopub.status.busy":"2023-07-21T18:01:44.028889Z","iopub.execute_input":"2023-07-21T18:01:44.029173Z","iopub.status.idle":"2023-07-21T18:01:44.035195Z","shell.execute_reply.started":"2023-07-21T18:01:44.029148Z","shell.execute_reply":"2023-07-21T18:01:44.034219Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"os.chdir(r'/kaggle/working')\nFileLink(\"/kaggle/working/UNET_WITH_FOUNDATION.pt\")","metadata":{"execution":{"iopub.status.busy":"2023-07-21T18:01:44.179885Z","iopub.execute_input":"2023-07-21T18:01:44.180229Z","iopub.status.idle":"2023-07-21T18:01:44.188551Z","shell.execute_reply.started":"2023-07-21T18:01:44.180202Z","shell.execute_reply":"2023-07-21T18:01:44.187544Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}