{"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":"code","source":"!pip install segmentation_models_pytorch","metadata":{"_kg_hide-output":true,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport copy\nimport cv2\nimport time\nimport random\n\nimport numpy as np\nimport pandas as pd\nfrom tqdm import tqdm\nimport matplotlib.pyplot as plt\nimport albumentations as A \nfrom albumentations.pytorch import ToTensorV2\n\nimport torch\nimport torchmetrics\nfrom torch import optim\nfrom torch.utils.data import (Dataset, \n                              DataLoader)\nfrom torchvision import transforms as tr\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom segmentation_models_pytorch.losses import DiceLoss\nimport segmentation_models_pytorch as smp\n\n\nfrom sklearn.model_selection import train_test_split\n\nplt.rcParams.update({'font.size': 15})","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Functions","metadata":{}},{"cell_type":"code","source":"# https://www.kaggle.com/paulorzp/rle-functions-run-length-encode-decode\ndef mask2rle(img):\n    '''\n    img: numpy array, 1 - mask, 0 - background\n    Returns run length as string formated\n    '''\n    pixels= img.T.flatten()\n    pixels = np.concatenate([[0], pixels, [0]])\n    runs = np.where(pixels[1:] != pixels[:-1])[0] + 1\n    runs[1::2] -= runs[::2]\n    return ' '.join(str(x) for x in runs)\n \ndef rle2mask(mask_rle, shape=(1600,256)):\n    '''\n    mask_rle: run-length as string formated (start length)\n    shape: (width,height) of array to return \n    Returns numpy array, 1 - mask, 0 - background\n\n    '''\n    s = mask_rle.split()\n    starts, lengths = [np.asarray(x, dtype=int) for x in (s[0:][::2], s[1:][::2])]\n    starts -= 1\n    ends = starts + lengths\n    img = np.zeros(shape[0]*shape[1], dtype=np.uint8)\n    for lo, hi in zip(starts, ends):\n        img[lo:hi] = 1\n    return img.reshape(shape).T\n\ndef enhancement_image(image):\n    # https://stackoverflow.com/questions/39308030/how-do-i-increase-the-contrast-of-an-image-in-python-opencv\n    # Enhance image contrast\n    sample = image.copy()\n    lab= cv2.cvtColor(sample, cv2.COLOR_BGR2LAB)\n    l_channel, a, b = cv2.split(lab)\n    # Applying CLAHE to L-channel\n    # feel free to try different values for the limit and grid size:\n    clahe = cv2.createCLAHE(clipLimit=5.0, tileGridSize=(12,12))\n    cl = clahe.apply(l_channel)\n\n    # merge the CLAHE enhanced L-channel with the a and b channel\n    limg = cv2.merge((cl,a,b))\n\n    # Converting image from LAB Color model to BGR color space\n    enhanced_img = cv2.cvtColor(limg, cv2.COLOR_LAB2BGR)\n\n    return enhanced_img\n\ndef read_image(image_path, enhancement=False):\n    image = cv2.imread(image_path)\n    image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n    \n    if enhancement:\n        image = enhancement_image(image)\n        \n    return image\n    \ndef read_mask(rle, image_shape):\n    mask = rle2mask(mask_rle=rle, \n                    shape=image_shape)\n    \n    return mask\n\ndef plot_result(result_dict, title, xlabel, ylabel):\n    plt.figure(figsize=(8, 6))\n\n    for key, values in result_dict.items():\n        if isinstance(values[0], torch.Tensor):\n            values = [value.cpu() for value in values]\n\n        plt.plot(values, label=key)\n    \n    plt.title(title)\n    plt.xlabel(xlabel)\n    plt.ylabel(ylabel)\n\n    plt.legend(loc=\"best\")\n\n    plt.show()  ","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def denormalize_image(image_tensor):\n    image = image_tensor.numpy().transpose((1, 2, 0))\n    image = std * image + mean\n    image = np.clip(image, 0, 1)\n    \n    return image\n\ndef show_images(images, masks, title):\n    fig, axs = plt.subplots(nrows=2, ncols=3, figsize=(14, 9))\n    st = plt.suptitle(title, fontsize=16)\n\n    for index, (image, mask) in enumerate(zip(images, masks)):\n        image = denormalize_image(image)\n        mask = mask.squeeze().numpy()\n        \n        # Show  Image\n        ax = axs[0][index]\n        ax.imshow(image)\n        ax.set_axis_off()\n\n        # Show Image + Mask\n        ax = axs[1][index]\n\n        ax.imshow(image)\n        ax.imshow(mask, cmap='seismic',alpha=0.4)\n        ax.set_axis_off()\n\n    plt.tight_layout()  \n    plt.show()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"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    \n    #the following line gives ~10% speedup\n    #but may lead to some stochasticity in the results \n    torch.backends.cudnn.benchmark = True\n    \nseed_everything(0)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Load DataFrame","metadata":{}},{"cell_type":"code","source":"data_base_path = os.path.join(\"..\", \"input\",\"hubmap-hpa-exploratory-analysis\")\n\n\ndf = pd.read_csv(os.path.join(data_base_path, \"clean_train.csv\"))\ndisplay(df.head())","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Define transforms\nREf : https://albumentations.ai/docs/examples/pytorch_classification/","metadata":{}},{"cell_type":"code","source":"image_size = 256\n\n# REF : https://www.kaggle.com/code/yassinealouini/mean-and-std-statistics-101/notebook\nmean = np.array([0.73074433, 0.70718829, 0.72277672])\nstd = np.array([0.27033332, 0.29168204, 0.27891807])\n\naug_transform = A.Compose([\n                    A.RandomResizedCrop(height=image_size, width=image_size, scale=(0.2, 1)),\n                    A.HorizontalFlip(),\n                    A.VerticalFlip(),\n                    A.RandomRotate90(),\n                    A.ShiftScaleRotate(scale_limit=0.1, rotate_limit=15,\n                                     border_mode=cv2.BORDER_REFLECT, p=.9),\n                    A.OneOf([\n                        A.OpticalDistortion(p=.3),\n                        A.GridDistortion(p=.1),\n                        A.PiecewiseAffine(p=.3),\n                    ], p=0.3),\n                    A.OneOf([\n                        A.HueSaturationValue(10,15,10),\n                        A.CLAHE(clip_limit=2),\n                        A.RandomBrightnessContrast(),            \n                    ], p=0.3),\n                    A.Normalize(mean=mean, std=std),\n                    ToTensorV2(),\n                    ])\n\nno_transform = A.Compose([\n                        A.SmallestMaxSize(max_size=image_size),\n                        A.CenterCrop(height=image_size, width=image_size),\n                        A.Normalize(mean=mean, std=std),\n                        ToTensorV2(),\n                        ])","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Define DataSet","metadata":{}},{"cell_type":"code","source":"class HuBMAPDataset(Dataset):\n    def __init__(self, df, transform=None, enhancement=False):\n        self.image_path_list = df[\"image_path\"].values\n        self.rle_list = df[\"rle\"].values\n        \n        self.transform = transform\n        self.enhancement = enhancement\n        \n    def __len__(self):\n        return len(self.image_path_list)\n    \n    def __getitem__(self, idx):\n        # Get Image \n        image_path = self.image_path_list[idx]\n        image = read_image(image_path, enhancement=self.enhancement)\n        \n        # Get Mask\n        rle = self.rle_list[idx]\n        mask = read_mask(rle, image_shape=image.shape[:2])\n        \n        if self.transform is not None:\n            result = self.transform(image=image, mask=mask)\n            image, mask = result[\"image\"], result[\"mask\"]\n        \n        return image, mask","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Raw Images visualization","metadata":{}},{"cell_type":"code","source":"train_dataset = HuBMAPDataset(df=df,\n                               transform=no_transform,\n                                enhancement=False)\ntrain_dataLoader = DataLoader(train_dataset, \n                              batch_size=3, \n                              shuffle=False)\n\n\nimages, masks = next(iter(train_dataLoader))\nshow_images(images, masks, title=\"Image without enhancement\")","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Enhanced Images visualization ","metadata":{}},{"cell_type":"code","source":"train_dataset = HuBMAPDataset(df=df,\n                               transform=no_transform,\n                                enhancement=True)\ntrain_dataLoader = DataLoader(train_dataset, \n                              batch_size=3, \n                              shuffle=False)\n\n\nimages, masks = next(iter(train_dataLoader))\nshow_images(images, masks, title=\"Image with enhancement\")","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Images with data augmentation","metadata":{}},{"cell_type":"code","source":"train_dataset = HuBMAPDataset(df=df,\n                               transform=aug_transform,\n                                enhancement=True)\ntrain_dataLoader = DataLoader(train_dataset, \n                              batch_size=3, \n                              shuffle=False)\n\n\nimages, masks = next(iter(train_dataLoader))\nshow_images(images, masks, title=\"Image with enhancement & data augmentation\")","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Import model architecture","metadata":{}},{"cell_type":"code","source":"def create_data_loader(batch_size):\n    train_df, val_df = train_test_split(df, test_size=0.2, stratify=df[\"organ\"], random_state=0)\n    \n    train_dataset = HuBMAPDataset(df=train_df, transform=aug_transform, enhancement=True)\n    train_dataloader = DataLoader(train_dataset, batch_size=batch_size, shuffle=True)\n    \n    val_dataset = HuBMAPDataset(df=val_df, transform=no_transform, enhancement=True)\n    val_dataloader = DataLoader(val_dataset, batch_size=batch_size, shuffle=True)\n\n    dataloaders = {\"train\" : train_dataloader,\n                   \"val\" : val_dataloader}\n    return dataloaders","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train_model(model, dataloaders, criterion, optimizer, num_epochs, scheduler=None):\n    history_dict = {\"train_loss\": [], \"train_metric\": [],\n                    \"val_loss\": [], \"val_metric\": [],\n                   \"lr\" : []}\n    \n    # Store model weight\n    best_model_wts = copy.deepcopy(model.state_dict())\n    best_metric = 0.0\n\n    for epoch in range(num_epochs):\n        # Each epoch has a training and validation phase\n        for phase in ['train', 'val']:\n            if phase == 'train':\n                model.train()  # Set model to training mode\n            else:\n                model.eval()   # Set model to evaluate mode\n            \n            mean_loss = torchmetrics.MeanMetric().to(device)\n            mean_metric = torchmetrics.MeanMetric().to(device)\n            \n            # Iterate over data.\n            for inputs, targets in tqdm(dataloaders[phase]):\n                inputs, targets = inputs.to(device), targets.to(device)\n\n                # zero the parameter gradients\n                optimizer.zero_grad()\n                \n                # forward\n                # track history if only in train\n                with torch.set_grad_enabled(phase == 'train'):\n                    outputs = model(inputs)\n                    \n                    loss = criterion(outputs, targets)\n                    mean_loss.update(loss)\n                    \n                    metric_result = metric(outputs, targets)\n                    mean_metric.update(metric_result)\n                    \n                    # backward + optimize only if in training phase\n                    if phase == 'train':\n                        loss.backward()\n                        optimizer.step()\n                \n                \n            epoch_loss = mean_loss.compute().item()\n            epoch_metric = mean_metric.compute().item()\n            \n            # deep copy the model\n            if phase == 'val': \n                if epoch_metric > best_metric:\n                    best_metric = epoch_metric\n                    \n                    best_model_wts = copy.deepcopy(model.state_dict())\n            \n            \n            if phase == 'val':\n                # Ajust learning rate\n                if scheduler is not None:\n                    history_dict['lr'].append(sheduler.get_last_lr())\n                    scheduler.step()\n                \n                \n                # Store result histories\n                history_dict[\"val_loss\"].append(epoch_loss)\n                history_dict[\"val_metric\"].append(epoch_metric)\n            else:\n                history_dict[\"train_loss\"].append(epoch_loss)\n                history_dict[\"train_metric\"].append(epoch_metric)\n                \n        \n        msg = \"[INFO] EPOCH: %d/%d\" % (epoch+1, num_epochs)\n        msg += \"\\t Metric train/Val: %.4f/%.4f \\tLoss train/Val: %.4f/%.4f\" % (history_dict[\"train_metric\"][-1], \n                                                                                history_dict[\"val_metric\"][-1], \n                                                                                history_dict[\"train_loss\"][-1], \n                                                                                history_dict[\"val_loss\"][-1])\n        print(msg)\n            \n    # load best model weights\n    print(\"Best metric %.4f\" % best_metric)\n    model.load_state_dict(best_model_wts)\n    \n    return model, history_dict","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\ndevice","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"REF : https://github.com/qubvel/segmentation_models.pytorch","metadata":{}},{"cell_type":"code","source":"model = smp.MAnet(encoder_name=\"resnext50_32x4d\",\n                        encoder_weights=\"imagenet\",\n                        classes=1)\nmodel = model.to(device)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"batch_size = 8\nn_epochs = 150\nlr = 1e-4\n\n\ncriterion = DiceLoss(mode='multilabel')\nmetric = torchmetrics.Dice(average='samples', ignore_index=0).to(device)\n\noptimizer = optim.Adam(model.parameters(), lr=lr)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Create Data Loader\ndata_loaders = create_data_loader(batch_size=batch_size)\n\n\n# Training model\nstart_time = time.time()\nmodel, history_dict = train_model(model, data_loaders, criterion, optimizer, n_epochs)\ntraining_duration = time.time() - start_time\n\nbest_metric = max(history_dict[\"val_metric\"])\nbest_loss = min(history_dict[\"val_loss\"])\n\ndel data_loaders","metadata":{"_kg_hide-output":true,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model_path = os.path.join(\"..\", \"working\", \"MAnet_v01.pt\")\n\ntorch.save(model, model_path)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Plot loss & Dice evolution","metadata":{}},{"cell_type":"code","source":"plot_result(result_dict={\"Train Dice\" : history_dict[\"train_metric\"],\n                        \"Validation Dice\" : history_dict[\"val_metric\"]},\n          title=\"Dice on Dataset\",\n          xlabel=\"Epoch\",\n          ylabel=\"Dice\")","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plot_result(result_dict={\"Train Loss\" : history_dict[\"train_loss\"],\n                        \"Validation Loss\" : history_dict[\"val_loss\"]},\n          title=\"Loss on Dataset\",\n          xlabel=\"Epoch\",\n          ylabel=\"Loss\")","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Evaluation of results","metadata":{}},{"cell_type":"code","source":"def predict(model, dataloader, criterion, metric):\n    mean_loss = torchmetrics.MeanMetric().to(device)\n    mean_metric = torchmetrics.MeanMetric().to(device)\n    \n    with torch.no_grad():\n        for images, targets in dataloader:\n            inputs, targets = images.to(device), targets.to(device)\n            \n            outputs = model(inputs)\n            \n            loss = criterion(outputs, targets)\n            mean_loss.update(loss)\n            \n            metric_result = metric(outputs, targets)\n            mean_metric.update(metric_result)\n            \n            \n    epoch_loss = mean_loss.compute().item()\n    epoch_metric = mean_metric.compute().item()\n    \n    return epoch_loss, epoch_metric","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df, val_df = train_test_split(df, test_size=0.2, stratify=df[\"organ\"], random_state=0)\n\nval_dataset = HuBMAPDataset(df=val_df, transform=no_transform, enhancement=True)\nval_dataloader = DataLoader(val_dataset, batch_size=3, shuffle=True)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"start_time = time.time()\n\nloss, metric_result = predict(model, val_dataloader, criterion, metric)\n\npredict_duration = time.time() - start_time\npredict_duration_per_frame = predict_duration / len(val_dataloader.dataset)\n\n\nprint(\"\\tValidation Dice \\t\\t: %.4f\" % metric_result)\nprint(\"\\tValidation Loss \\t\\t: %.6f\" %  loss)\nprint(\"\\tPredict Duration for 1 image \\t: %.4fs\" % predict_duration_per_frame)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Compute Metrics per organ","metadata":{}},{"cell_type":"code","source":"dataLoader_per_organ = {}\n\norgane_list = val_df[\"organ\"].unique()\nfor organe_name in organe_list:\n    filtered_df = val_df.query(\"organ == '%s'\" % organe_name)\n    \n    \n    print(\"\\n- %s\" % organe_name)\n    print(\"\\tNumber of images : %s\" % filtered_df.shape[0])\n    \n    \n    filtered_dataset = HuBMAPDataset(df=filtered_df, transform=no_transform, enhancement=True)\n    filtered_dataloader = DataLoader(filtered_dataset, batch_size=3, shuffle=True)\n    \n    dataLoader_per_organ[organe_name] = filtered_dataloader","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for organe_name, dataLoader in dataLoader_per_organ.items():\n    print(\"\\t - %s\" % organe_name)\n\n    loss, metric_result = predict(model, dataLoader, criterion, metric)\n    print(\"\\t\\tValidation Dice \\t\\t: %.4f\" % metric_result)\n    print(\"\\t\\tValidation Loss \\t\\t: %.6f\" %  loss)\n    \n    del dataLoader","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Show outputs Probability & Prediction mask","metadata":{}},{"cell_type":"code","source":"for epoch, (images, targets) in enumerate(val_dataloader):\n    fig, axs = plt.subplots(nrows=3, ncols=3, figsize=(14, 16))\n    st = plt.suptitle(\"One Batch\", fontsize=16)\n\n    \n    inputs, targets = images.to(device), targets.to(device)\n    \n    outputs = model(inputs)\n    \n    \n    for index, (image, target) in enumerate(zip(images, targets)):\n        image = denormalize_image(image)\n        target = target.detach().cpu().numpy()\n        \n        \n        output = outputs[index].squeeze()\n        output = torch.sigmoid(output)\n\n        output = output.detach().cpu().numpy()\n\n        \n        # Show Pred\n        ax = axs[0][index]\n\n        ax.set_title(\"Probability\")\n        ax.imshow(image)\n        ax.imshow(output, alpha=0.4)\n        ax.set_axis_off()\n        \n        \n        # Show Mask\n        ## Apply Threshold\n        mask = (output > 0.5) * 255\n        mask = mask.astype(np.uint8)\n\n        ax = axs[1][index]\n\n        ax.set_title(\"Prediction mask\")\n        ax.imshow(image)\n        ax.imshow(mask, alpha=0.5)\n        ax.set_axis_off()\n        \n        \n        \n        # Show Image + Target\n        ax = axs[2][index]\n        \n        ax.set_title(\"Reel\")\n        ax.imshow(image)\n        ax.imshow(target, cmap='seismic',alpha=0.4)\n        \n        ax.set_axis_off()\n        \n    plt.tight_layout()  \n    plt.show()\n    \n    \n    if epoch > 2:\n        break","metadata":{"trusted":true},"execution_count":null,"outputs":[]}]}