{"metadata":{"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":52279,"databundleVersionId":5822112,"sourceType":"competition"},{"sourceId":5901549,"sourceType":"datasetVersion","datasetId":3389674},{"sourceId":8790596,"sourceType":"datasetVersion","datasetId":5285091}],"dockerImageVersionId":30733,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true},"kernelspec":{"display_name":"Python 3 (ipykernel)","language":"python","name":"python3"},"language_info":{"codemirror_mode":{"name":"ipython","version":3},"file_extension":".py","mimetype":"text/x-python","name":"python","nbconvert_exporter":"python","pygments_lexer":"ipython3","version":"3.11.5"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"Key features:\n- Model: UNet\n    - simple Unet structure with skip connections, with 2 conv block per layer\n- Dataset: Filtered out 11 sample that did not include the target 'blood_vessel' class.\n- Cyclic training: Based on Leslie N. Smith's paper (https://arxiv.org/pdf/1506.01186) and torch implementation (https://github.com/davidtvs/pytorch-lr-finder)\n- Gradient Accumulation: Although I didn't experienced memory saturation with batch size of 16 or smaller, still it can enhance training speed and reduce memory usage.\n- Data augmentation\n- Combined loss (BCE + Dice loss)","metadata":{}},{"cell_type":"markdown","source":"## Imports","metadata":{}},{"cell_type":"code","source":"import os\nimport gc\nimport cv2\nimport json\nimport pandas as pd\nimport numpy as np\nimport matplotlib.pyplot as plt\nfrom matplotlib.colors import ListedColormap\nfrom datetime import datetime\nfrom tqdm import tqdm, notebook\nimport albumentations as A\n\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import DataLoader, Dataset\nfrom torchvision import models\nimport torch.nn.functional as F\nfrom torchinfo import summary\n\npd.set_option('display.max_rows', None)\npd.set_option('display.max_columns', None)\npd.set_option('display.width', None)\npd.set_option('display.max_colwidth', 999)\n\ndevice=torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(device)","metadata":{"execution":{"iopub.execute_input":"2024-06-28T07:57:42.027866Z","iopub.status.busy":"2024-06-28T07:57:42.027419Z","iopub.status.idle":"2024-06-28T07:57:49.455727Z","shell.execute_reply":"2024-06-28T07:57:49.454606Z","shell.execute_reply.started":"2024-06-28T07:57:42.027824Z"}},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# !pip install /kaggle/input/torch-lr-finder/torch_lr_finder-0.2.1-py3-none-any.whl\nfrom torch_lr_finder import LRFinder","metadata":{"execution":{"iopub.execute_input":"2024-06-28T07:57:49.458111Z","iopub.status.busy":"2024-06-28T07:57:49.457560Z","iopub.status.idle":"2024-06-28T07:58:23.013689Z","shell.execute_reply":"2024-06-28T07:58:23.012814Z","shell.execute_reply.started":"2024-06-28T07:57:49.458077Z"}},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Data loading","metadata":{}},{"cell_type":"code","source":"class HubMap_Dataset(Dataset):\n    def __init__(self, img_path, labels_file, metadata_file, augmentation=None):\n        self.image_dir = img_path\n        self.labels_file = labels_file\n        self.metadata_file = metadata_file\n        self.augmentation = augmentation\n\n        # Processing JSON data to create a pandas df        \n        with open(labels_file) as json_file:\n            json_list = list(json_file)\n\n        dataset = []\n        for json_str in notebook.tqdm(json_list, desc=\"Reading Json Data\"):\n            result = json.loads(json_str)\n            annotations = result['annotations']\n            \n            for ann in annotations:\n                if ann[\"type\"] != \"blood_vessel\":\n                    continue\n                row = {\n                    \"id\": result[\"id\"],\n                    \"coordinates\": ann[\"coordinates\"],\n                }\n                dataset.append(row)\n        \n        self.dataset = pd.DataFrame(dataset, columns=[\"id\",\"coordinates\"])\n    \n    def coordinates_to_mask(self, image_id):\n        \"\"\" Create a combined mask containing all polygons. \"\"\"\n        all_filled_masks = np.zeros((512, 512, 1), dtype=np.uint8)\n        annotations = self.dataset[self.dataset['id'] == image_id]\n        for _, row in annotations.iterrows():\n            coordinates = np.array(row[\"coordinates\"])\n            all_filled_masks = cv2.fillPoly(all_filled_masks, [coordinates], 1)\n        return all_filled_masks\n\n    def __len__(self):\n        return len(self.dataset['id'].unique())\n\n    def __getitem__(self, idx):\n        unique_ids = self.dataset['id'].unique()\n        img_id = unique_ids[idx]\n        image_path = f\"{self.image_dir}/{img_id}.tif\"\n        image = cv2.imread(image_path, cv2.COLOR_BGR2RGB) # cv2 reads in h,w,c                        \n        mask = self.coordinates_to_mask(image_id=img_id)\n                        \n        # Augment data\n        if self.augmentation is not None:\n            augmented = self.augmentation(image=image, mask=mask)\n            image, mask = augmented['image'], augmented['mask']\n\n        mask = torch.tensor(mask, dtype=torch.float32).permute(2, 0, 1)\n        image = np.array(image) / 255.0\n        image = torch.tensor(image, dtype=torch.float32).permute(2, 0, 1) # c,h,w\n        \n        return image, mask, img_id\n","metadata":{"execution":{"iopub.execute_input":"2024-06-28T07:58:23.015353Z","iopub.status.busy":"2024-06-28T07:58:23.015050Z","iopub.status.idle":"2024-06-28T07:58:23.029602Z","shell.execute_reply":"2024-06-28T07:58:23.028588Z","shell.execute_reply.started":"2024-06-28T07:58:23.015321Z"}},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"base_path = \"./data/hubmap_data\"\nimages_folder = base_path + \"/train/\"\nlabels_path = base_path + \"/polygons.jsonl\"\nmetadata_path = base_path + \"/tile_meta.csv\"\n\naugs = A.Compose([\n    A.HorizontalFlip(p=0.5),\n    A.VerticalFlip(p=0.5),\n    A.RandomRotate90(p=0.5),\n    A.OneOf([\n        A.PiecewiseAffine(scale=(0.01, 0.05), p=1),\n        A.GaussianBlur(blur_limit=(3, 3), p=1),\n    ], p=0.5),\n    A.OneOf([\n        A.HueSaturationValue(hue_shift_limit=10, sat_shift_limit=10, val_shift_limit=10, p=1),\n        A.RandomBrightnessContrast(brightness_limit=0.1, contrast_limit=0.1, p=1),\n        A.CLAHE(clip_limit=2.0, p=1), # equalize colors based on histogram\n    ], p=0.5),\n])\n\nHMD = HubMap_Dataset(img_path=images_folder, labels_file=labels_path, metadata_file=metadata_path, augmentation=augs)\nprint(f\"Number of ID-s removed, where 'blood_vessel' were missing: {1633-len(HMD)}\")\nval_len = int(0.2*len(HMD))\nlengths = [len(HMD)-val_len, val_len]\ntrain_set, val_set = torch.utils.data.random_split(HMD, lengths)\nprint(f\"train set: {len(train_set)}, validation set: {len(val_set)}\")\n","metadata":{"execution":{"iopub.execute_input":"2024-06-28T07:58:23.033251Z","iopub.status.busy":"2024-06-28T07:58:23.032518Z","iopub.status.idle":"2024-06-28T07:58:27.459352Z","shell.execute_reply":"2024-06-28T07:58:27.458491Z","shell.execute_reply.started":"2024-06-28T07:58:23.033226Z"}},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Visualize samples\nsamples = [HMD[i] for i in range(40,45)]\n\nfor sample in samples:\n    masks = sample[1].squeeze(0)\n    image = sample[0].permute(1,2,0).numpy()\n\n    fig, axs = plt.subplots(1,3,figsize=(10,8))\n    axs[0].imshow(image)    \n    axs[0].set_title(f\"id: {sample[2]}\")\n    axs[1].imshow(masks)\n    axs[2].imshow(image)\n    axs[2].imshow(masks, alpha=0.6)\n    plt.show()","metadata":{"execution":{"iopub.execute_input":"2024-06-28T07:58:27.460917Z","iopub.status.busy":"2024-06-28T07:58:27.460545Z","iopub.status.idle":"2024-06-28T07:58:31.331432Z","shell.execute_reply":"2024-06-28T07:58:31.330550Z","shell.execute_reply.started":"2024-06-28T07:58:27.460885Z"}},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Model","metadata":{}},{"cell_type":"code","source":"# # could be better than default(uniform)\ndef init_weights(m):\n    if isinstance(m, nn.Conv2d):\n        torch.nn.init.kaiming_uniform_(m.weight, nonlinearity='relu')\n        if m.bias is not None:\n            m.bias.data.fill_(0.01)\n\nclass ConvBlock(nn.Module):\n    def __init__(self, in_channels, out_channels):\n        super().__init__()\n        self.conv = nn.Sequential(\n            nn.Conv2d(in_channels, out_channels, kernel_size=3, padding=1),\n            nn.ReLU(inplace=True),\n            nn.BatchNorm2d(out_channels),\n            \n            nn.Conv2d(out_channels, out_channels, kernel_size=3, padding=1),\n            nn.ReLU(inplace=True),\n            nn.BatchNorm2d(out_channels),\n        )\n\n    def forward(self, x):\n        return self.conv(x)\n\nclass UNet(nn.Module):\n    def __init__(self, in_channels=3, in_filters=16, num_classes=4):\n        super().__init__()\n\n        self.num_classes = num_classes\n\n        # Donwsample\n        self.encoder1 = ConvBlock(in_channels, in_filters)\n        self.pool1 = nn.MaxPool2d(2)\n        self.encoder2 = ConvBlock(in_filters, in_filters * 2)\n        self.pool2 = nn.MaxPool2d(2)\n        self.encoder3 = ConvBlock(in_filters * 2, in_filters * 4)\n        self.pool3 = nn.MaxPool2d(2)\n        self.encoder4 = ConvBlock(in_filters * 4, in_filters * 8)\n        self.pool4 = nn.MaxPool2d(2)\n        self.encoder5 = ConvBlock(in_filters * 8, in_filters * 16)\n        self.pool5 = nn.MaxPool2d(2)\n\n        # Bottom layer\n        self.bottleneck = ConvBlock(in_filters * 16, in_filters * 32)\n\n        # Upsample\n        self.upconv5 = nn.ConvTranspose2d(in_filters * 32, in_filters * 16, 2, stride=2, output_padding=0)\n        self.decoder5 = ConvBlock(in_filters * 32, in_filters * 16)\n        self.upconv4 = nn.ConvTranspose2d(in_filters * 16, in_filters * 8, 2, stride=2)\n        self.decoder4 = ConvBlock(in_filters * 16, in_filters * 8)\n        self.upconv3 = nn.ConvTranspose2d(in_filters * 8, in_filters * 4, 2, stride=2)\n        self.decoder3 = ConvBlock(in_filters * 8, in_filters * 4)\n        self.upconv2 = nn.ConvTranspose2d(in_filters * 4, in_filters * 2, 2, stride=2)\n        self.decoder2 = ConvBlock(in_filters * 4, in_filters * 2)\n        self.upconv1 = nn.ConvTranspose2d(in_filters * 2, in_filters, 2, stride=2)\n        self.decoder1 = ConvBlock(in_filters * 2, in_filters)\n\n        # Classifier\n        self.conv = nn.Conv2d(in_filters, num_classes, 1)\n        self.apply(init_weights)\n\n    def forward(self, x):\n        enc1 = self.encoder1(x)\n        pool1 = self.pool1(enc1)\n        enc2 = self.encoder2(pool1)\n        pool2 = self.pool2(enc2)\n        enc3 = self.encoder3(pool2)\n        pool3 = self.pool3(enc3)\n        enc4 = self.encoder4(pool3)\n        pool4 = self.pool4(enc4)\n        enc5 = self.encoder5(pool4)\n        pool5 = self.pool5(enc5)\n\n        bottleneck = self.bottleneck(pool5)\n\n        upconv5 = self.upconv5(bottleneck)\n        cat5 = torch.cat([upconv5, enc5], dim=1)\n        dec5 = self.decoder5(cat5)\n        upconv4 = self.upconv4(dec5)\n        cat4 = torch.cat([upconv4, enc4], dim=1)\n        dec4 = self.decoder4(cat4)        \n        upconv3 = self.upconv3(dec4)\n        cat3 = torch.cat([upconv3, enc3], dim=1)\n        dec3 = self.decoder3(cat3)\n        upconv2 = self.upconv2(dec3)\n        cat2 = torch.cat([upconv2, enc2], dim=1)\n        dec2 = self.decoder2(cat2)\n        upconv1 = self.upconv1(dec2)\n        cat1 = torch.cat([upconv1, enc1], dim=1)\n        dec1 = self.decoder1(cat1)\n\n        return self.conv(dec1)","metadata":{"execution":{"iopub.execute_input":"2024-06-28T07:58:31.333100Z","iopub.status.busy":"2024-06-28T07:58:31.332662Z","iopub.status.idle":"2024-06-28T07:58:31.352124Z","shell.execute_reply":"2024-06-28T07:58:31.351138Z","shell.execute_reply.started":"2024-06-28T07:58:31.333073Z"}},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"## setup config\nclass cfg:\n    img_size = 512\n    in_ch = 3\n    num_classes = 1\n    filters = 16\n    mixed_precision = True\n    batch_size = 8\n    init_lr = 1e-3\n    weight_decay = 1e-4\n    early_stopping = 3\n    epochs = 12\n    repeat = 6\n\n# Create Unet\nunet_model = UNet(in_filters=cfg.filters, num_classes=cfg.num_classes)\nunet_model.to(device)\n\nsummary(unet_model, input_size=(1,3,512,512), col_names=['input_size','output_size','num_params','trainable'])","metadata":{"execution":{"iopub.execute_input":"2024-06-28T07:58:31.353529Z","iopub.status.busy":"2024-06-28T07:58:31.353253Z","iopub.status.idle":"2024-06-28T07:58:32.545707Z","shell.execute_reply":"2024-06-28T07:58:32.544702Z","shell.execute_reply.started":"2024-06-28T07:58:31.353507Z"}},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Training","metadata":{}},{"cell_type":"code","source":"train_loader = DataLoader(train_set, batch_size=cfg.batch_size, shuffle=True)\nval_loader = DataLoader(val_set, batch_size=cfg.batch_size, shuffle=False)\noneb = next(iter(train_loader))\nims, lbs, ids = oneb\nprint(f\"images: {ims.shape, ims.dtype, ims.device}\\nmasks: {lbs.shape, lbs.dtype, lbs.device}\")","metadata":{"execution":{"iopub.execute_input":"2024-06-28T07:58:32.547922Z","iopub.status.busy":"2024-06-28T07:58:32.547263Z","iopub.status.idle":"2024-06-28T07:58:33.655549Z","shell.execute_reply":"2024-06-28T07:58:33.654627Z","shell.execute_reply.started":"2024-06-28T07:58:32.547889Z"}},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def find_best_lr(model, dataloader, loss_fn, optimizer, num_iter=100):\n    \"\"\" Performs a learning rate range test to find optimal lr for next training. \"\"\"\n    lr_finder = LRFinder(model, optimizer, loss_fn, device='cuda')\n    lr_finder.range_test(dataloader, start_lr=1e-7, end_lr=1, num_iter=num_iter, step_mode='exp')\n    sug_lr = lr_finder.plot(suggest_lr=True)\n    lr_finder.reset()\n    torch.cuda.empty_cache()\n    gc.collect()\n    return sug_lr\n\ndef calc_stats(labels, preds):\n    \"\"\" Calculate IoU and dice score. \"\"\"\n    pred_class = torch.sigmoid(preds).round()\n    num_correct = (pred_class == labels).sum().item()\n    num_pixels = labels.numel()\n    intersection = (pred_class * labels).sum().item()\n    union = (pred_class + labels).sum().item() - intersection\n\n    if union > 0:\n        pixel_acc = num_correct / num_pixels * 100\n        iou = intersection / union * 100\n        dice_score = (2 * intersection) / ((pred_class + labels).sum().item() + 1e-8)\n    else:\n        pixel_acc = 0\n        iou = 0        \n        dice_score = 0\n\n    return dice_score, iou, pixel_acc\n\ndef training(train_loader, valid_loader, model, loss_fn, optimizer, epochs, early_stopping, history=None, prev_val_loss=None, set_scaler=False, device=\"cuda\"):\n    \n    scaler = torch.cuda.amp.GradScaler() if set_scaler else None\n    early_stopping_counter = 0\n    ga_target = cfg.batch_size * 3 # update in every 3 batches\n    if prev_val_loss is None:\n        prev_val_loss = np.Inf\n    if history is None:        \n        history = {\n            \"epoch\": [],\n            \"train_loss\": [],\n            \"valid_loss\": [],\n            \"p_acc\": [],\n            \"d_score\": [],\n            \"iou\": [],\n        }    \n    \n    for e in range(epochs):\n        # TRAINING\n        model.train()\n        print(\"-\" * 40)\n        print(f\"EPOCH {e + 1}/{epochs}. training step...\")\n        train_time_1 = datetime.now()\n        loss_list = []\n        ga_counter = 0\n\n        for i, data in enumerate(train_loader):\n            im = data[0].to(device)\n            seg_gt = data[1].to(device)\n\n            with torch.cuda.amp.autocast(enabled=set_scaler):\n                seg_pred = model.forward(im)\n                loss = loss_fn(seg_pred, seg_gt) / ga_target # need normalized loss\n\n            # Backward and accumulate gradients\n            if set_scaler:\n                scaler.scale(loss.mean()).backward()\n            else:\n                loss.backward()\n            \n            # Update GA counter\n            ga_counter += cfg.batch_size\n\n            # Perform the optimization step if the target accumulation size is reached\n            if ga_counter >= ga_target:\n                if set_scaler:\n                    scaler.step(optimizer)\n                    scaler.update()\n                else:\n                    optimizer.step()\n                \n                optimizer.zero_grad()\n                ga_counter = 0\n                \n            # Collect loss\n            loss_list.append(loss.item() * ga_target) # append accumulated loss\n            \n        train_time_2 = datetime.now()- train_time_1\n        print(f\" - epoch training time: (hh:mm:ss:ms): {train_time_2}\\n\")\n\n        # VALIDATION\n        model.eval()\n        print(f\"EPOCH {e + 1}/{epochs}. validation step...\")\n        val_loss_list = []\n        total_dice_scores = 0\n        total_iou = 0\n        total_pixel_acc = 0\n        \n        for i, data in enumerate(valid_loader):\n            im = data[0].to(device)\n            seg_gt = data[1].to(device)\n\n            with torch.no_grad():\n                seg_pred = model.forward(im)\n                val_loss = loss_fn(seg_pred, seg_gt)\n                val_loss_list.append(val_loss.item())\n                                \n                # Stats calculation\n                dice_score, iou, pixel_acc = calc_stats(labels=seg_gt, preds=seg_pred)\n                total_dice_scores += dice_score\n                total_iou += iou\n                total_pixel_acc += pixel_acc\n\n        #  Take the avg of the collected losses and metrrics\n        avg_train_loss = np.mean(loss_list)\n        avg_valid_loss = np.mean(val_loss_list)\n        avg_dice_scores = total_dice_scores / len(valid_loader)\n        avg_iou = total_iou / len(valid_loader)\n        avg_pixel_acc = total_pixel_acc / len(valid_loader)\n        \n        # Print out stats\n        print(f\"Results:\")\n        print(f\" - Train epoch mean loss: {avg_train_loss:.4f}\")\n        print(f\" - Valid epoch mean loss: {avg_valid_loss:.4f}\\n\")        \n        print(f\" - Dice score: {avg_dice_scores:.4f}\")\n        print(f\" - Iou: {avg_iou:.2f}%\")\n        print(f\" - Pixel accuracy: {avg_pixel_acc:.2f}%\")\n        history[\"epoch\"].append(e + 1)\n        history[\"train_loss\"].append(avg_train_loss)\n        history[\"valid_loss\"].append(avg_valid_loss)\n        history[\"p_acc\"].append(avg_pixel_acc)\n        history[\"d_score\"].append(avg_dice_scores)\n        history[\"iou\"].append(avg_iou)\n                    \n        # Compare losses and save model with lower loss\n        if avg_valid_loss <= prev_val_loss:\n            print(f\" - Lower validation loss achieved ({prev_val_loss:.4f}-->{avg_valid_loss:.4f}). Saving model...\")\n            best_model_state = model.state_dict()\n            torch.save(best_model_state, f'./best_model_state.pth')\n            prev_val_loss = avg_valid_loss\n            early_stopping_counter = 0\n        else:\n            early_stopping_counter +=1\n            print(f\" - Validation loss did not improved. Early stopping counter {early_stopping_counter}/{early_stopping}\")\n            if early_stopping == early_stopping_counter:\n                break\n\n        # Clear cache after epoch\n        torch.cuda.empty_cache()\n        gc.collect()\n        \n    return history, prev_val_loss\n\n\ndef model_trainer(train_loader, valid_loader, model, loss_fn, optimizer, epochs, early_stopping, set_scaler=False, device=\"cuda\", repeat=5):\n    \n    assert (repeat>=1) & (repeat<=10), \"The 'repeat' parameter must be an integer between 1 and 10.\"\n    print(f\"There will be {repeat} training iterations with {epochs} epochs each.\")\n    start_time = datetime.now()\n    \n    # Train model repeatably\n    for i in range(repeat):\n        print(f\"\\nTRAINING ITERATION {i+1}/{repeat}\")\n        \n        if i == 0:\n            history, best_val_loss = training(train_loader, valid_loader, model, loss_fn, optimizer, epochs, early_stopping, history=None, prev_val_loss=None, set_scaler=set_scaler, device=device)\n        else:\n            history, best_val_loss = training(train_loader, valid_loader, model, loss_fn, optimizer, epochs, early_stopping, history=history, prev_val_loss=best_val_loss, set_scaler=set_scaler, device=device)\n        \n        # Load current best model\n        saved_model = \"./best_model_state.pth\"\n        model.load_state_dict(torch.load(saved_model))\n        \n        if i != repeat-1:\n            # Check the optimal lr for next training\n            print()\n            suggested_lr = find_best_lr(model, train_loader, loss_fn, optimizer, num_iter=100)\n\n            if isinstance(suggested_lr, tuple) and len(suggested_lr) > 1:\n                lr_value = suggested_lr[1]\n            elif isinstance(suggested_lr, (float, int)):\n                lr_value = suggested_lr\n            else:\n                print(\"Unexpected format for suggested learning rate.\")\n                lr_value = 1e-4\n                \n            # Update optimizer\n            optimizer = torch.optim.AdamW(model.parameters(), lr=lr_value, weight_decay=cfg.weight_decay)\n            print(f\"\\nLearning rate for the next training: {lr_value}\")\n\n    end_time = datetime.now()-start_time\n    print(f\"\\nAll training has finished. Time spent (hh:mm:ss:ms) - {end_time}\")\n            \n    return model, history\n","metadata":{"execution":{"iopub.execute_input":"2024-06-28T07:58:33.657054Z","iopub.status.busy":"2024-06-28T07:58:33.656780Z","iopub.status.idle":"2024-06-28T07:58:33.690219Z","shell.execute_reply":"2024-06-28T07:58:33.689209Z","shell.execute_reply.started":"2024-06-28T07:58:33.657029Z"}},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class DiceLoss(nn.Module):\n    def __init__(self, smooth=1e-8):\n        super().__init__()\n        self.smooth = smooth\n\n    def forward(self, outputs, labels):\n        pred_class = torch.sigmoid(outputs)\n        intersection = (pred_class * labels).sum()\n        union = pred_class.sum() + labels.sum() - intersection\n        dice_score = (2 * intersection + self.smooth) / (union + self.smooth)\n        dice_loss = 1 - dice_score\n        return dice_loss\n    \nclass CombinedLoss(nn.Module):\n    \"\"\" Combined loss function, utilizing BCE and Dice losses. \"\"\"\n\n    def __init__(self, weight=None, smooth=1e-8):\n        super().__init__()\n        self.weight = weight\n        self.smooth = smooth\n\n    def forward(self, outputs, labels):\n        # Calculate BCE\n        bce_loss = F.binary_cross_entropy_with_logits(outputs, labels, weight=self.weight)\n\n        # Calculate Dice loss\n        pred_class = torch.sigmoid(outputs)\n        intersection = (pred_class * labels).sum()\n        union = pred_class.sum() + labels.sum() - intersection\n        dice_score = (2 * intersection + self.smooth) / (union + self.smooth)\n        dice_loss = 1 - dice_score\n\n        # Combine Dice loss and BCE loss\n        combined_loss = bce_loss + dice_loss\n        return combined_loss","metadata":{"execution":{"iopub.execute_input":"2024-06-28T07:58:33.693489Z","iopub.status.busy":"2024-06-28T07:58:33.693169Z","iopub.status.idle":"2024-06-28T07:58:33.705425Z","shell.execute_reply":"2024-06-28T07:58:33.704532Z","shell.execute_reply.started":"2024-06-28T07:58:33.693465Z"}},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"loss_fn = CombinedLoss()\noptimizer = torch.optim.AdamW(unet_model.parameters(), lr=cfg.init_lr, weight_decay=cfg.weight_decay)\n\nunet_model, history = model_trainer(\n    train_loader=train_loader,\n    valid_loader=val_loader,\n    model=unet_model,\n    loss_fn=loss_fn,\n    optimizer=optimizer,\n    epochs=cfg.epochs,\n    early_stopping=cfg.early_stopping,\n    set_scaler=cfg.mixed_precision,\n    device=device,\n    repeat=cfg.repeat,\n)","metadata":{"execution":{"iopub.execute_input":"2024-06-28T07:58:33.707564Z","iopub.status.busy":"2024-06-28T07:58:33.707262Z","iopub.status.idle":"2024-06-28T08:18:50.921557Z","shell.execute_reply":"2024-06-28T08:18:50.919980Z","shell.execute_reply.started":"2024-06-28T07:58:33.707541Z"}},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"mhist = pd.DataFrame(history)\nmhist","metadata":{"execution":{"iopub.status.busy":"2024-06-28T08:18:50.922783Z","iopub.status.idle":"2024-06-28T08:18:50.923248Z","shell.execute_reply":"2024-06-28T08:18:50.923047Z","shell.execute_reply.started":"2024-06-28T08:18:50.923029Z"}},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# best dice score\nbest_dice = mhist.loc[(mhist['valid_loss'] == min(mhist['valid_loss']))]['d_score']\nbd_index = best_dice.index.item()\nbd_value = best_dice.item()\nprint(f'Best dice score {bd_value:.4f} ({bd_index}. row)')","metadata":{"execution":{"iopub.status.busy":"2024-06-28T08:18:50.924759Z","iopub.status.idle":"2024-06-28T08:18:50.925095Z","shell.execute_reply":"2024-06-28T08:18:50.924948Z","shell.execute_reply.started":"2024-06-28T08:18:50.924934Z"}},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def plot_loss(history):\n    plt.plot(history['valid_loss'], label='valid', marker='o')\n    plt.plot( history['train_loss'], label='train', marker='o')\n    plt.title('Loss per epoch')\n    plt.ylabel('loss')\n    plt.xlabel('epoch')\n    plt.legend()\n    plt.grid()\n    plt.show()\n\ndef plot_metrics(history):\n    plt.plot(history['p_acc'], label='accuracy', marker='s')\n    plt.plot(history['d_score']*100, label='dice', marker='o')\n    plt.plot(history['iou'], label='iou', marker='*')\n    plt.title('Metrics')\n    plt.ylabel('percentage')\n    plt.xlabel('epoch')\n    plt.legend()\n    plt.grid()\n    plt.show()\n\n    \nplot_loss(mhist)\nplot_metrics(mhist)","metadata":{"execution":{"iopub.status.busy":"2024-06-28T08:18:50.926748Z","iopub.status.idle":"2024-06-28T08:18:50.927081Z","shell.execute_reply":"2024-06-28T08:18:50.926936Z","shell.execute_reply.started":"2024-06-28T08:18:50.926922Z"}},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Inference","metadata":{}},{"cell_type":"code","source":"def predict_on_batch(batch_of_data, eval_model, model_name=\"\", threshold=0.5):\n\n    for img, lbl, imid in zip(batch_of_data[0], batch_of_data[1], batch_of_data[2]):\n        img = img.to(device)\n        lbl = lbl.to(device)\n\n        img_to_pred = torch.unsqueeze(img, 0).to(torch.float32)\n        pred_probs = torch.sigmoid(eval_model(img_to_pred)).squeeze(0)\n        pred_class = (pred_probs > threshold).float()\n        \n        ## Metrics\n        dice, iou, pa = calc_stats(pred_class, lbl)\n       \n        ## Visualize\n        pred_probs = pred_probs.detach().cpu().numpy().swapaxes(0, -1)        \n        pred_class = pred_class.cpu().numpy().swapaxes(0, -1)\n        image = img.cpu().numpy().swapaxes(0, -1)\n        label = lbl.cpu().numpy().swapaxes(0, -1)\n        print(f\"Metrics of {imid}:\")\n#         print(f\"- pixel accuracy: {pa:.2f}%\")\n        print(f\"- iou: {iou:.2f}%\")\n        print(f\"- dice score: {dice:.2f}\")\n        \n        fig, axs = plt.subplots(1, 2, figsize=(10, 7))\n        axs[0].imshow(image)\n        axs[0].imshow(label, alpha=0.5)\n        axs[0].set_title(\"Mask & image\")\n        axs[1].imshow(image)\n        axs[1].imshow(pred_class, alpha=0.5)\n        axs[1].set_title(f\"Binary preds ({threshold})\")\n        axs[1].set_axis_off()\n        plt.show()","metadata":{"execution":{"iopub.status.busy":"2024-06-28T08:18:50.929134Z","iopub.status.idle":"2024-06-28T08:18:50.929478Z","shell.execute_reply":"2024-06-28T08:18:50.929326Z","shell.execute_reply.started":"2024-06-28T08:18:50.929313Z"}},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"## Run infernce and visualize\nbest_model_path = \"./best_model_state.pth\"\ninf_model = UNet(in_filters=cfg.filters, num_classes=cfg.num_classes)\ninf_model.load_state_dict(torch.load(best_model_path))\ninf_model.eval().to(device)\n\nitera = iter(val_loader)\nvbatch_1 = next(itera)\nvbatch_2 = next(itera)\npredict_on_batch(vbatch_1, inf_model, threshold=0.5)\npredict_on_batch(vbatch_2, inf_model, threshold=0.5)\n","metadata":{"execution":{"iopub.status.busy":"2024-06-28T08:18:50.930840Z","iopub.status.idle":"2024-06-28T08:18:50.931145Z","shell.execute_reply":"2024-06-28T08:18:50.931007Z","shell.execute_reply.started":"2024-06-28T08:18:50.930995Z"}},"execution_count":null,"outputs":[]}]}