{"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":"# Semantic Segmentation with PyTorch and U-Net\n\n**Author:** Shingo Nakazawa ([@shnakazawa](https://twitter.com/shnakazawa))\n\n**Objective:** In this notebook, we will create a semantic segmentation model using the [**U-Net architecture**](https://arxiv.org/abs/1505.04597). ([A reference article (Japanese)](https://zenn.dev/aidemy/articles/a43ebe82dfbb8b))\n\nU-Net is specifically designed for image segmentation tasks, which involve identifying different objects or regions within an image and labeling them with different classes. In this notebook, we will implement the U-Net architecture using [**PyTorch**](https://pytorch.org/), **without relying on pre-existing implementations or models**.\n\nPlease see also\n\n- [Image Classification with PyTorch and EfficientNetV2](https://www.kaggle.com/code/shnakazawa/image-classification-with-pytorch-and-efficientnet)\n- [Object Detection with PyTorch and DETR](https://www.kaggle.com/code/shnakazawa/object-detection-with-pytorch-and-detr)\n\n# Import Modules","metadata":{"papermill":{"duration":0.014352,"end_time":"2023-03-15T03:18:58.499204","exception":false,"start_time":"2023-03-15T03:18:58.484852","status":"completed"},"tags":[]}},{"cell_type":"code","source":"import os\nimport numpy as np\nimport pandas as pd\nimport math\nimport time\nimport random\nimport gc\nfrom pathlib import Path\nimport cv2\nfrom tqdm.notebook import tqdm\nfrom sklearn.model_selection import KFold\n\n# Image augmentation\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n\n# Modeling\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\n# Visualization\nimport matplotlib.pyplot as plt\n%matplotlib inline\n\nprint(f'PyTorch version {torch.__version__}')\nprint(f'Albumentations version {A.__version__}')","metadata":{"execution":{"iopub.execute_input":"2023-03-15T03:18:58.521027Z","iopub.status.busy":"2023-03-15T03:18:58.520561Z","iopub.status.idle":"2023-03-15T03:19:02.976292Z","shell.execute_reply":"2023-03-15T03:19:02.975155Z"},"papermill":{"duration":4.467064,"end_time":"2023-03-15T03:19:02.978792","exception":false,"start_time":"2023-03-15T03:18:58.511728","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Set Configs\n\nSeparately configuring settings such as objectives for running, file paths, hyperparameters, and other parameters can greatly aid in maintaining a clear and organized workflow.","metadata":{"papermill":{"duration":0.007289,"end_time":"2023-03-15T03:19:02.993908","exception":false,"start_time":"2023-03-15T03:19:02.986619","status":"completed"},"tags":[]}},{"cell_type":"code","source":"RUN_EDA = True\nRUN_TRAINING = True\nTRAIN_ALL = False # If true, train with all data and output a single model. If False, run cross-validation and output multiple models.\nFOLD_NUM = 5 # For cross-validation\nEPOCHS = 20 # Training cycle\nRUN_INFERENCE = False\n\n# Directory setting\nDATA_DIR = '/kaggle/input/sartorius-cell-instance-segmentation/'\nMODEL_DIR = '/kaggle/working/'\nIMG_SAVE_DIR = '/kaggle/working/'\n\n# PyTorch variables\nSEED = 42\nNUM_WORKERS = 2\nBATCH_SIZE = 8\nWEIGHT_DECAY = 0.0001\nLR = 0.0001\nMOMENTUM = 0.9\n\n# Threshold for mask prediction\nTHRESHOLD = 0.3\n\n# Set device\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu') \nprint(f'Using {device} device')","metadata":{"execution":{"iopub.execute_input":"2023-03-15T03:19:03.011116Z","iopub.status.busy":"2023-03-15T03:19:03.009953Z","iopub.status.idle":"2023-03-15T03:19:03.080838Z","shell.execute_reply":"2023-03-15T03:19:03.079494Z"},"papermill":{"duration":0.081772,"end_time":"2023-03-15T03:19:03.083165","exception":false,"start_time":"2023-03-15T03:19:03.001393","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Define Helper Functions\n\nDefining reusable functions at the beginning of a Jupyter Notebook can result in code that is cleaner, more organized, and more efficient.","metadata":{"papermill":{"duration":0.007338,"end_time":"2023-03-15T03:19:03.098143","exception":false,"start_time":"2023-03-15T03:19:03.090805","status":"completed"},"tags":[]}},{"cell_type":"code","source":"def seed_everything(seed):\n    random.seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = True\n\n# Set seed\nseed_everything(SEED)\n\n\ndef show_gpu_memory(device):\n    print(f\"Allocated GPU memory: {torch.cuda.memory_allocated(device) / 1024 / 1024:.2f} MB\")\n    print(f\"Cached GPU memory: {torch.cuda.memory_cached(device) / 1024 / 1024:.2f} MB\")    \n\n    \ndef load_img(path):\n    img_bgr = cv2.imread(path)\n    img_rgb = img_bgr[:, :, ::-1]\n    return img_rgb\n\n\ndef group_bboxes(df):\n    df_ = df.copy()\n    df_['segment_count'] = 1\n    df_ = df_.groupby(['id', 'width', 'height', 'cell_type']).count().reset_index()\n    return_df = df_[['id', 'width', 'height', 'cell_type', 'segment_count']]\n    return return_df\n\n\ndef create_gallery(array, ncols=3):\n    \"\"\"Display multiple images in a gallery style.\n    Source: https://www.amazon.co.jp/Data-Analysis-Machine-Learning-Kaggle-ebook/dp/B09F3STL34/\n    \n    Args:\n        array (numpy.ndarray): array of images.\n        ncols (int, optional): Num of columns. Defaults to 3.\n\n    Returns:\n        numpy.ndarray: One concatenated image.\n    \"\"\"    \n    nindex, height, width, intensity = array.shape\n    nrows = nindex//ncols\n    assert nindex == nrows * ncols\n    result = (array.reshape(nrows, ncols, height, width, intensity)\n        .swapaxes(1,2)\n        .reshape(height*nrows, width*ncols, intensity))\n    return result\n\n\ndef decode_rle(rle, height, width):\n    \"\"\"RLE to image\n    modified from: https://www.kaggle.com/paulorzp/run-length-encode-and-decode\n\n    Args:\n        rle (str): mask with run length encoding.\n        height (int): return image height.\n        width (int): return image width.\n        brightness (int): brightness of the pixel. Default to 1.\n\n    Returns:\n        np.ndarray: 1(b) - mask, 0 - background.\n    \"\"\"    \n    s = 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(height * width, dtype=np.uint8)\n    for lo, hi in zip(starts, ends):\n        img[lo:hi] = 1 # brightness\n    return img.reshape((height, width)) \n\n\ndef create_mask_image(image, masks):\n    \"\"\"Create a mask image from RLE.\n\n    Args:\n        image (numpy.ndarray): array of images.\n        masks (list): List with RLE-encoded mask information.\n        b (int): brightness of the pixel. Default to 1.\n    \n    Returns:\n        numpy.ndarray: 1(b) - mask, 0 - background.\n    \"\"\"    \n    \n    s = image.shape\n    h = s[0]\n    w = s[1]\n    mask_image = np.zeros((h,w))\n    for mask in masks:\n        mask_image += decode_rle(mask, h, w)\n    mask_image = mask_image.clip(0, 1)\n    return mask_image\n\n\ndef show_validation_score(train_loss_list, valid_loss_list, save=False, save_dir=IMG_SAVE_DIR, save_name='segmentation_validation_score.png'):\n    fig = plt.figure(figsize=(10,10))\n    for i in range(FOLD_NUM):\n        train_loss = train_loss_list[i]\n        valid_loss = valid_loss_list[i]\n        \n        ax = fig.add_subplot(math.ceil(np.sqrt(FOLD_NUM)), math.ceil(np.sqrt(FOLD_NUM)), i+1, title=f'Fold {i+1}')\n        ax.plot(range(EPOCHS), train_loss, c='orange', label='train')\n        ax.plot(range(EPOCHS), valid_loss, c='blue', label='valid')\n        ax.set_xlabel('epoch')\n        ax.set_ylabel('loss')\n        ax.legend()\n    \n    plt.tight_layout()\n    if save:\n        os.makedirs(save_dir, exist_ok=True)\n        plt.savefig(save_dir+save_name)\n    else:\n        plt.show()\n\n\ndef encode_rle(predicted_img):\n    predicted_img = (predicted_img > THRESHOLD).astype(int)\n    height, width = predicted_img.shape\n    \n    # Get the index of the masked pixel\n    pixels = predicted_img.copy()\n    pixels_list = []\n    for y in range(height):\n        for x in range(width):\n            if pixels[y][x] != 0:\n                pixels_list.append(y * width + x)\n\n    # RLE encoding\n    rle_list = []\n    start = pixels_list[0]\n    count = 1\n    for i in range(1, len(pixels_list)):\n        if pixels_list[i] == pixels_list[i-1] + 1:\n            count += 1\n        else:\n            rle_list.extend([start, count])\n            start = pixels_list[i]\n            count = 1\n    rle_list.extend([start, count])\n    \n    rle_str = [str(x) for x in rle_list]\n    return ' '.join(rle_str)","metadata":{"execution":{"iopub.execute_input":"2023-03-15T03:19:03.115229Z","iopub.status.busy":"2023-03-15T03:19:03.114934Z","iopub.status.idle":"2023-03-15T03:19:03.138264Z","shell.execute_reply":"2023-03-15T03:19:03.137361Z"},"papermill":{"duration":0.0344,"end_time":"2023-03-15T03:19:03.140323","exception":false,"start_time":"2023-03-15T03:19:03.105923","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Load and Reshape a Table","metadata":{"papermill":{"duration":0.007346,"end_time":"2023-03-15T03:19:03.155300","exception":false,"start_time":"2023-03-15T03:19:03.147954","status":"completed"},"tags":[]}},{"cell_type":"code","source":"df = pd.read_csv(DATA_DIR + 'train.csv')\ndf.head()","metadata":{"execution":{"iopub.execute_input":"2023-03-15T03:19:03.171970Z","iopub.status.busy":"2023-03-15T03:19:03.171133Z","iopub.status.idle":"2023-03-15T03:19:03.778504Z","shell.execute_reply":"2023-03-15T03:19:03.777491Z"},"papermill":{"duration":0.618279,"end_time":"2023-03-15T03:19:03.781074","exception":false,"start_time":"2023-03-15T03:19:03.162795","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"The segment of training data is provided with **Run length encoding (RLE)**.\n\nRLE is a lossless compression technique used to represent data that contains long sequences of repeated values or characters. [Check out this discussion for a better understanding](https://www.kaggle.com/c/severstal-steel-defect-detection/discussion/102311).\n\nIn the context of image segmentation, RLE is used to represent the mask or label of an object or region in an image.\n\nThe labels are provided in various formats depending on the dataset. For instance, some datasets may provide the mask or label as an image with colored regions while others may provide it as coco-format. If the label is not provided in the desired format, it may be necessary to **preprocess the data to create the appropriate masks**. Various tools and libraries are available for this purpose, such as [OpenCV](https://opencv.org/) and [scikit-image](https://scikit-image.org/).\n\nThe current dataset has **one segment per row**. To better work with the data, let's reshape the table.","metadata":{"papermill":{"duration":0.008194,"end_time":"2023-03-15T03:19:03.797785","exception":false,"start_time":"2023-03-15T03:19:03.789591","status":"completed"},"tags":[]}},{"cell_type":"code","source":"grouped_df = group_bboxes(df)\ngrouped_df.head()","metadata":{"execution":{"iopub.execute_input":"2023-03-15T03:19:03.814967Z","iopub.status.busy":"2023-03-15T03:19:03.814670Z","iopub.status.idle":"2023-03-15T03:19:03.879060Z","shell.execute_reply":"2023-03-15T03:19:03.877926Z"},"papermill":{"duration":0.07577,"end_time":"2023-03-15T03:19:03.881761","exception":false,"start_time":"2023-03-15T03:19:03.805991","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Exploratory Data Analysis (EDA)\n\nEDA for images is typically simpler than for tabular data.\n\n**Please set `RUN_EDA = True` in the `Set Config` section.**\n\n## Inspect Representative Images\n\nPerforming a visual inspection of images is critical. If the images appear noisy or unusual, preprocessing may be necessary before training a model.","metadata":{"papermill":{"duration":0.007741,"end_time":"2023-03-15T03:19:03.897560","exception":false,"start_time":"2023-03-15T03:19:03.889819","status":"completed"},"tags":[]}},{"cell_type":"code","source":"if RUN_EDA:\n    img_names = Path(DATA_DIR+'train/').glob('*.png')\n    img_list = []\n    for i, img_name in enumerate(img_names):\n        img_list.append(load_img(img_name.as_posix()))\n        print(img_name.name)\n        if i == 5: \n            break\n    plt.figure(figsize=(10,10))\n    plt.imshow(create_gallery(np.array(img_list), ncols=3))\nelse:\n    print('RUN_EDA is False')","metadata":{"execution":{"iopub.execute_input":"2023-03-15T03:19:03.914388Z","iopub.status.busy":"2023-03-15T03:19:03.914088Z","iopub.status.idle":"2023-03-15T03:19:04.877936Z","shell.execute_reply":"2023-03-15T03:19:04.876614Z"},"papermill":{"duration":0.978167,"end_time":"2023-03-15T03:19:04.883525","exception":false,"start_time":"2023-03-15T03:19:03.905358","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Plot Mask\n\nLet's see an example of mask images.\n\nAs mentioned earlier, mask images are often encoded using the RLE format, and thus it is necessary to decode them.","metadata":{"papermill":{"duration":0.011816,"end_time":"2023-03-15T03:19:04.907726","exception":false,"start_time":"2023-03-15T03:19:04.895910","status":"completed"},"tags":[]}},{"cell_type":"code","source":"if RUN_EDA:\n    image_id = grouped_df['id'][0]\n    img = load_img(f'{DATA_DIR}train/{image_id}.png')\n    masks = df[df['id'] == image_id]['annotation'].tolist()\n    masked_img = create_mask_image(img, masks)\n    plt.figure()\n    plt.imshow(img)\n    plt.figure()\n    plt.imshow(masked_img)\nelse:\n    print('RUN_EDA is False')","metadata":{"execution":{"iopub.execute_input":"2023-03-15T03:19:04.933042Z","iopub.status.busy":"2023-03-15T03:19:04.932738Z","iopub.status.idle":"2023-03-15T03:19:05.735382Z","shell.execute_reply":"2023-03-15T03:19:05.734350Z"},"papermill":{"duration":0.818995,"end_time":"2023-03-15T03:19:05.738742","exception":false,"start_time":"2023-03-15T03:19:04.919747","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Check Image Shape\n\nExamining the size and color of the images in the dataset is another crucial step.\n\nGenerating histograms may also provide useful insights.","metadata":{"papermill":{"duration":0.01754,"end_time":"2023-03-15T03:19:05.774734","exception":false,"start_time":"2023-03-15T03:19:05.757194","status":"completed"},"tags":[]}},{"cell_type":"code","source":"if RUN_EDA:\n    img_shape = set()\n    img_ext = set()\n    img_names = Path(DATA_DIR+'train/').glob('*')\n    pbar = tqdm(img_names, total=len(grouped_df))\n    for img_name in pbar:\n        img = load_img(img_name.as_posix())\n        img_shape.add(img.shape)\n        img_ext.add(img_name.suffix)\n    print(f'Image shapes are {img_shape}.')\n    print(f'Image extensions are {img_ext}.')","metadata":{"execution":{"iopub.execute_input":"2023-03-15T03:19:05.811318Z","iopub.status.busy":"2023-03-15T03:19:05.810379Z","iopub.status.idle":"2023-03-15T03:19:14.007930Z","shell.execute_reply":"2023-03-15T03:19:14.006665Z"},"papermill":{"duration":8.218327,"end_time":"2023-03-15T03:19:14.010152","exception":false,"start_time":"2023-03-15T03:19:05.791825","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"We are confident that all images are 520 x 704 RGB png.","metadata":{"papermill":{"duration":0.017237,"end_time":"2023-03-15T03:19:14.045334","exception":false,"start_time":"2023-03-15T03:19:14.028097","status":"completed"},"tags":[]}},{"cell_type":"markdown","source":"## Plot histogram of pixel values\n\nGenerating a pixel value histogram can aid in identifying outlier images, such as those containing entirely zero-valued pixels.","metadata":{"papermill":{"duration":0.017194,"end_time":"2023-03-15T03:19:14.080061","exception":false,"start_time":"2023-03-15T03:19:14.062867","status":"completed"},"tags":[]}},{"cell_type":"code","source":"if RUN_EDA:\n    img_names = Path(DATA_DIR+'train/').glob('*')\n    plt.figure(figsize=(10,10))\n    pbar = tqdm(img_names, total=len(grouped_df))\n    for img_name in pbar:\n        img = load_img(img_name.as_posix())\n        hist = cv2.calcHist([img],[0],None,[256],[0,256])\n        plt.plot(hist)\n    plt.show()\nelse:\n    print('RUN_EDA is False')","metadata":{"execution":{"iopub.execute_input":"2023-03-15T03:19:14.116513Z","iopub.status.busy":"2023-03-15T03:19:14.116144Z","iopub.status.idle":"2023-03-15T03:19:21.809023Z","shell.execute_reply":"2023-03-15T03:19:21.807985Z"},"papermill":{"duration":7.714378,"end_time":"2023-03-15T03:19:21.811631","exception":false,"start_time":"2023-03-15T03:19:14.097253","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"The image signals appear to be quite uniform, which is characteristic of typical microscopy images. :)","metadata":{"papermill":{"duration":0.018575,"end_time":"2023-03-15T03:19:21.849297","exception":false,"start_time":"2023-03-15T03:19:21.830722","status":"completed"},"tags":[]}},{"cell_type":"markdown","source":"# Define Model Components\n\nBefore training a model using PyTorch, the following steps need to be completed:\n\n1. Define Image Transformation and Augmentation\n2. Define the Dataset\n3. Define the DataLoader\n4. Define the Model\n\n## Define Image Transformation and Augmentation\n\n**For segmentation tasks, the segmentation masks also need to undergo transformations and augmentations.**\n\nPlease also see [Albumentations Documentation/Mask augmentation for segmentation](https://albumentations.ai/docs/getting_started/mask_augmentation/). In this notebook, transforms are applied in the `CellDataset()` class.","metadata":{"papermill":{"duration":0.018068,"end_time":"2023-03-15T03:19:21.886118","exception":false,"start_time":"2023-03-15T03:19:21.868050","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# Image Transformation & Augmentation\ndef transform_train():\n    transforms = [\n        A.Resize(256,256,p=1),\n        A.HorizontalFlip(p=0.5),\n        A.Transpose(p=0.5),\n        ToTensorV2(p=1)\n    ]\n    return A.Compose(transforms)\n\n\n# Validation images undergo only resizing.\ndef transform_valid():\n    transforms = [\n        A.Resize(256,256,p=1),\n        ToTensorV2(p=1)\n    ]\n    return A.Compose(transforms)","metadata":{"execution":{"iopub.execute_input":"2023-03-15T03:19:21.924384Z","iopub.status.busy":"2023-03-15T03:19:21.924028Z","iopub.status.idle":"2023-03-15T03:19:21.930474Z","shell.execute_reply":"2023-03-15T03:19:21.929480Z"},"papermill":{"duration":0.028262,"end_time":"2023-03-15T03:19:21.932669","exception":false,"start_time":"2023-03-15T03:19:21.904407","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Define Datasets","metadata":{"papermill":{"duration":0.018092,"end_time":"2023-03-15T03:19:21.968822","exception":false,"start_time":"2023-03-15T03:19:21.950730","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# Dataset\nclass CellDataset(Dataset):\n    def __init__(self, image_ids, dataframe, data_root, transforms=None, stage='train'):\n        super().__init__()\n        self.image_ids = image_ids\n        self.dataframe = dataframe\n        self.data_root = data_root\n        self.transforms = transforms\n        self.stage = stage\n\n    def __len__(self):\n        return self.image_ids.shape[0]\n    \n    def __getitem__(self, index):\n        image_id = self.image_ids[index]\n        # Load images\n        image  = load_img(f'{self.data_root}{image_id}.png').astype(np.float32)\n\n        # 3 channels to 1 channel\n        image = cv2.cvtColor(image, cv2.COLOR_RGB2GRAY)\n        image /= 255.0 # normalization\n\n        # For training and validation\n        if self.stage == 'train':\n            # masks\n            masks = self.dataframe[self.dataframe['id'] == image_id]['annotation'].tolist()\n            mask_image = create_mask_image(image, masks)\n\n            # Transform images and masks\n            if self.transforms:\n                transformed = self.transforms(image=image, mask=mask_image)\n                image, mask_image = transformed['image'], transformed['mask']\n            return image, mask_image, image_id\n        \n        # For test\n        else:\n            # Transform images\n            if self.transforms:\n                image =self.transforms(image=image)['image']\n            \n            return image, image_id","metadata":{"execution":{"iopub.execute_input":"2023-03-15T03:19:22.007856Z","iopub.status.busy":"2023-03-15T03:19:22.006882Z","iopub.status.idle":"2023-03-15T03:19:22.016892Z","shell.execute_reply":"2023-03-15T03:19:22.015869Z"},"papermill":{"duration":0.03197,"end_time":"2023-03-15T03:19:22.019071","exception":false,"start_time":"2023-03-15T03:19:21.987101","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Define the DataLoader","metadata":{"papermill":{"duration":0.018433,"end_time":"2023-03-15T03:19:22.055449","exception":false,"start_time":"2023-03-15T03:19:22.037016","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# DataLoader\ndef create_dataloader(grouped_df, df, trn_idx, val_idx):\n    train_ = grouped_df.loc[trn_idx,:].reset_index(drop=True)\n    valid_ = grouped_df.loc[val_idx,:].reset_index(drop=True)\n\n    # Dataset\n    train_datasets = CellDataset(train_['id'].to_numpy(), df, DATA_DIR+'train/', transforms=transform_train())\n    valid_datasets = CellDataset(valid_['id'].to_numpy(), df, DATA_DIR+'train/', transforms=transform_valid())\n\n    # Data Loader\n    train_loader = DataLoader(train_datasets, batch_size=BATCH_SIZE, num_workers=NUM_WORKERS, shuffle=True, multiprocessing_context='fork')\n    valid_loader = DataLoader(valid_datasets, batch_size=BATCH_SIZE, num_workers=NUM_WORKERS, shuffle=False, multiprocessing_context='fork')\n\n    return train_loader, valid_loader","metadata":{"execution":{"iopub.execute_input":"2023-03-15T03:19:22.093176Z","iopub.status.busy":"2023-03-15T03:19:22.092864Z","iopub.status.idle":"2023-03-15T03:19:22.099687Z","shell.execute_reply":"2023-03-15T03:19:22.098626Z"},"papermill":{"duration":0.028628,"end_time":"2023-03-15T03:19:22.102120","exception":false,"start_time":"2023-03-15T03:19:22.073492","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Define the model\n\nIn this notebook, we will implement the U-Net architecture using PyTorch, without relying on pre-existing implementations or models.","metadata":{"papermill":{"duration":0.018484,"end_time":"2023-03-15T03:19:22.138527","exception":false,"start_time":"2023-03-15T03:19:22.120043","status":"completed"},"tags":[]}},{"cell_type":"code","source":"class DoubleConv(nn.Module):\n    \"\"\"DoubleConv is a basic building block of the encoder and decoder components. \n    Consists of two convolutional layers followed by a ReLU activation function.\n    \"\"\"    \n    def __init__(self, in_channels, out_channels):\n        super(DoubleConv, self).__init__()\n        self.double_conv = nn.Sequential(\n            nn.Conv2d(in_channels, out_channels, kernel_size=3, padding=1),\n            nn.BatchNorm2d(out_channels),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(out_channels, out_channels, 3, padding=1),\n            nn.BatchNorm2d(out_channels),\n            nn.ReLU(inplace=True)\n        )\n\n    def forward(self, x):\n        x = self.double_conv(x)\n        return x\n\n\nclass Down(nn.Module):\n    \"\"\"Downscaling.\n    Consists of two consecutive DoubleConv blocks followed by a max pooling operation.\n    \"\"\"    \n    def __init__(self, in_channels, out_channels):\n        super(Down, self).__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        x = self.maxpool_conv(x)\n        return x\n\n\nclass Up(nn.Module):\n    \"\"\"Upscaling.\n    Performed using transposed convolution and concatenation of feature maps from the corresponding \"Down\" operation.\n    \"\"\"\n    def __init__(self, in_channels, out_channels, bilinear=True):\n        super(Up, self).__init__()\n        \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    def forward(self, x1, x2):\n        x1 = self.up(x1)\n\n        # input tensor shape: (batch_size, channels, height, width)\n        diffY = x2.size()[2] - x1.size()[2]\n        diffX = x2.size()[3] - x1.size()[3]\n\n        x1 = nn.functional.pad(x1, [diffX // 2, diffX - diffX // 2,\n                        diffY // 2, diffY - diffY // 2])\n        x = torch.cat([x2, x1], dim=1)\n        x = self.conv(x)\n        return x\n    \n\nclass UNet(nn.Module):\n    def __init__(self, n_channels=1, n_classes=1, bilinear=False):\n        super(UNet, self).__init__()\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        \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 = nn.Conv2d(64, n_classes, 1)\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)\n        x = self.up3(x, x2)\n        x = self.up4(x, x1)\n        x = self.outc(x)\n        x = torch.sigmoid(x)\n        return x","metadata":{"execution":{"iopub.execute_input":"2023-03-15T03:19:22.177182Z","iopub.status.busy":"2023-03-15T03:19:22.176863Z","iopub.status.idle":"2023-03-15T03:19:22.194131Z","shell.execute_reply":"2023-03-15T03:19:22.193130Z"},"papermill":{"duration":0.039494,"end_time":"2023-03-15T03:19:22.196231","exception":false,"start_time":"2023-03-15T03:19:22.156737","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Run Training\n\nAs we did in the previous notebooks, we will perform cross-validation to evaluate the settings and determine **the best combinations of architectures and hyperparameters**. \n\nTo effectively track the model's performance, I recommend utilizing MLOps tools like [**MLFlow**](https://mlflow.org/).\n\n**Before running the following cell, please set configs in the `Set Config` section: `RUN_TRAINING = True` and `TRAIN_ALL = ` as you want.**","metadata":{"papermill":{"duration":0.01867,"end_time":"2023-03-15T03:19:22.233333","exception":false,"start_time":"2023-03-15T03:19:22.214663","status":"completed"},"tags":[]}},{"cell_type":"code","source":"if RUN_TRAINING:\n    if TRAIN_ALL:\n        # Train with all data\n        folds = [['','']]\n    else: \n        # Cross-validation\n        folds = KFold(n_splits=FOLD_NUM, shuffle=True, random_state=SEED)\\\n                .split(np.arange(grouped_df.shape[0]), grouped_df['id'].to_numpy())\n        \n        # For Visualization\n        train_loss_list = []\n        valid_loss_list = []\n    \n\n    for fold, (trn_idx, val_idx) in enumerate(folds):\n        # Load Data   \n        if TRAIN_ALL:\n            train_datasets = CellDataset(grouped_df['id'].to_numpy(), df, DATA_DIR+'train/', transforms=transform_train())\n            train_loader = DataLoader(train_datasets, batch_size=BATCH_SIZE, num_workers=NUM_WORKERS, shuffle=True, multiprocessing_context='fork')\n        else:\n            print(f'==========Cross-Validation Fold {fold+1}==========')   \n            train_loader, valid_loader = create_dataloader(grouped_df, df, trn_idx, val_idx)\n            # For Visualization\n            valid_losses = []\n\n        train_losses = []\n        # Load model, loss function, and optimizing algorithm\n        model = UNet().to(device)\n        criterion = nn.BCELoss().to(device)\n        optimizer = optim.SGD(model.parameters(), weight_decay=WEIGHT_DECAY, lr = LR, momentum=MOMENTUM)\n        \n        # Start training\n        best_loss = 10**5\n        for epoch in range(EPOCHS):\n            time_start = time.time()\n            print(f'==========Epoch {epoch+1} Start Training==========')\n            model.train()\n            train_loss = 0\n            pbar = tqdm(enumerate(train_loader), total=len(train_loader))\n            for step, (imgs, masks, image_ids) in pbar:\n                imgs = imgs.to(device).float()\n                # imgs = torch.squeeze(imgs)\n                masks = masks.to(device).float()\n                masks = masks.view(imgs.shape[0], -1, 256, 256)\n\n                optimizer.zero_grad()\n                \n                output = model(imgs)\n                loss = criterion(output, masks)\n                loss.backward()\n                optimizer.step()\n                \n                train_loss += loss.item()\n            train_loss /= len(train_loader)\n\n            # Validation\n            if TRAIN_ALL == False:\n                print(f'==========Epoch {epoch+1} Start Validation==========')\n                \n                with torch.no_grad():\n                    valid_loss = 0\n                    preds = []\n                    pbar = tqdm(enumerate(valid_loader), total=len(valid_loader))\n                    for step, (imgs, masks, image_ids) in pbar:\n                        imgs = imgs.to(device).float()\n                        # imgs = torch.squeeze(imgs)\n                        masks = masks.to(device).float()\n                        masks = masks.view(imgs.shape[0], -1, 256, 256)\n                \n                        val_output = model(imgs)\n                        val_loss = criterion(val_output, masks)\n                        \n                        valid_loss += val_loss.item()\n                    valid_loss /= len(valid_loader)\n                    \n            # print results from this epoch\n            exec_t = int((time.time() - time_start)/60)\n            if TRAIN_ALL:\n                print(f'Epoch : {epoch+1} - loss : {train_loss:.4f} / Exec time {exec_t} min\\n')\n\n            else:\n                print(\n                    f'Epoch : {epoch+1} - loss : {train_loss:.4f} - val_loss : {valid_loss:.4f} / Exec time {exec_t} min\\n'\n                )\n                # For visualization\n                train_losses.append(train_loss)\n                valid_losses.append(valid_loss)\n        \n        if TRAIN_ALL:\n            print(f'Save model trained with all data')\n            os.makedirs(MODEL_DIR, exist_ok=True)\n            torch.save(model.state_dict(), MODEL_DIR+'segmentation.pth')\n            del model, optimizer, train_loader\n        else:\n            train_loss_list.append(train_losses)\n            valid_loss_list.append(valid_losses)\n            del model, optimizer, train_loader, valid_loader, train_losses, valid_losses\n        gc.collect()\n        torch.cuda.empty_cache()\n    \n    if TRAIN_ALL == False:\n        show_validation_score(train_loss_list, valid_loss_list)\n\nelse:\n    print('RUN_TRAINING is False')","metadata":{"execution":{"iopub.execute_input":"2023-03-15T03:19:22.273012Z","iopub.status.busy":"2023-03-15T03:19:22.272636Z","iopub.status.idle":"2023-03-15T04:23:09.163404Z","shell.execute_reply":"2023-03-15T04:23:09.162395Z"},"papermill":{"duration":3826.914058,"end_time":"2023-03-15T04:23:09.165719","exception":false,"start_time":"2023-03-15T03:19:22.251661","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Although the learning has not yet converged, we will move on to the next step.\n\nRun this cell again with the `TRAIN_ALL` parameter set to `True` to obtain the `.pth` file.","metadata":{}},{"cell_type":"markdown","source":"# Run Inference of Test Data\n\nIn this task, the same Dataset, Image transformation, and model as we used for validation. So no need to define new functions.\n\n**Before running the following cells, please set `RUN_INFERENCE = True` in the `Set Config` section.**","metadata":{"papermill":{"duration":0.053536,"end_time":"2023-03-15T04:23:09.275336","exception":false,"start_time":"2023-03-15T04:23:09.221800","status":"completed"},"tags":[]}},{"cell_type":"code","source":"if RUN_INFERENCE:\n    files = os.listdir(DATA_DIR+'test/')\n    image_ids = np.array([os.path.splitext(file)[0] for file in files])\n    ids = []\n    rle_test_preds = []\n    original_size = (704, 520) # (width, height)\n\n    # Load Data\n    test_datasets = CellDataset(image_ids, df, DATA_DIR+'test/', transforms=transform_valid(), stage='test')\n\n    # Data Loader\n    test_loader = DataLoader(test_datasets, batch_size=BATCH_SIZE, num_workers=NUM_WORKERS, shuffle=False, multiprocessing_context='fork')\n\n    # Load model, loss function, and optimizing algorithm\n    model = UNet().to(device)\n    model.load_state_dict(torch.load(MODEL_DIR+'segmentation.pth'))\n       \n    # Start Inference\n    print(f'==========Start Inference==========')\n    with torch.no_grad():\n        test_preds = []\n        pbar = tqdm(enumerate(test_loader), total=len(test_loader))\n        for step, (imgs, image_ids) in pbar:\n            imgs = imgs.to(device).float()\n            output = model(imgs)\n\n            # Convert the output from PyTorch to np.array\n            output = output.detach().cpu().numpy()\n            \n            # run length encoding\n            for image_id, predicted_mask in zip(image_ids, output):\n                predicted_mask = np.squeeze(predicted_mask)\n                \n                # resize\n                predicted_mask = cv2.resize(predicted_mask, original_size)\n                \n                rle_mask = encode_rle(predicted_mask)\n                ids.append(image_id)\n                rle_test_preds.append(rle_mask)\n    \n    submission_df = pd.DataFrame({\n        'id': ids, 'predicted': rle_test_preds\n    })\n    print(submission_df.head())\n\nelse:\n    print('RUN_INFERENCE is False')","metadata":{"execution":{"iopub.execute_input":"2023-03-15T04:23:09.384334Z","iopub.status.busy":"2023-03-15T04:23:09.383752Z","iopub.status.idle":"2023-03-15T04:23:09.396700Z","shell.execute_reply":"2023-03-15T04:23:09.395764Z"},"papermill":{"duration":0.070695,"end_time":"2023-03-15T04:23:09.399219","exception":false,"start_time":"2023-03-15T04:23:09.328524","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Visualize One prediction example\n\nLet's see one output example at the current threshold. In practice, the output results are compared at various thresholds to select the threshold with the best evaluation metrics.","metadata":{"papermill":{"duration":0.0529,"end_time":"2023-03-15T04:23:09.505606","exception":false,"start_time":"2023-03-15T04:23:09.452706","status":"completed"},"tags":[]}},{"cell_type":"code","source":"if RUN_INFERENCE:\n    plt.imshow(predicted_mask > THRESHOLD)\nelse:\n    print('RUN_INFERENCE is False')","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Save inference result","metadata":{"papermill":{"duration":null,"end_time":null,"exception":null,"start_time":null,"status":"pending"},"tags":[]}},{"cell_type":"code","source":"if RUN_INFERENCE:\n    submission_df.to_csv('submission.csv', index=False)\nelse:\n    print('RUN_INFERENCE is False')","metadata":{"execution":{"iopub.execute_input":"2023-03-15T03:09:54.288780Z","iopub.status.busy":"2023-03-15T03:09:54.288075Z","iopub.status.idle":"2023-03-15T03:09:54.297263Z","shell.execute_reply":"2023-03-15T03:09:54.296250Z","shell.execute_reply.started":"2023-03-15T03:09:54.288740Z"},"papermill":{"duration":null,"end_time":null,"exception":null,"start_time":null,"status":"pending"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Visualize one submission example\n\nThis confirms whether the output style is correct.","metadata":{"papermill":{"duration":null,"end_time":null,"exception":null,"start_time":null,"status":"pending"},"tags":[]}},{"cell_type":"code","source":"if RUN_INFERENCE:\n    target = submission_df.iloc[0]\n    img = load_img(f'{DATA_DIR}test/{target[\"id\"]}.png')\n    masks = [target['predicted']]\n    masked_img = create_mask_image(img, masks)\n    plt.figure()\n    plt.imshow(img)\n    plt.figure()\n    plt.imshow(masked_img)\n\nelse:\n    print('RUN_INFERENCE is False')","metadata":{"execution":{"iopub.execute_input":"2023-03-15T03:09:54.301053Z","iopub.status.busy":"2023-03-15T03:09:54.300779Z","iopub.status.idle":"2023-03-15T03:09:54.940768Z","shell.execute_reply":"2023-03-15T03:09:54.939766Z","shell.execute_reply.started":"2023-03-15T03:09:54.301027Z"},"papermill":{"duration":null,"end_time":null,"exception":null,"start_time":null,"status":"pending"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# To Improve the Model\n\nOur segmentation model is now on the starting point. There are many techniques we can use to improve it.\n\nSome potential approaches to consider include:\n\n- **Preprocess input images**\n    - Variety of augmentation such as flipping, rotating, and scaling the images.\n- **Optimize the model**\n    - A larger and more diverse training dataset\n    - Experimenting with different architectures and hyperparameters\n    - Experiment with different architectures and hyperparameters\n    - Choosing an appropriate Loss function, such as Dice Loss\n- **Modify the model output**\n    - Finding the best threshold\n    - Using an ensemble of models\n    - Postprocessing\n\nBy implementing these techniques, it's possible to build a highly accurate semantic segmentation model.\n\nI trust that this notebook provided you with useful information for building an semantic segmentation model. Wishing you success in all your future modeling projects!\n\n# Acknowledgement\n\nI thank to [Matchan from YouTube Channel - Engawa AI Research Institute](https://www.youtube.com/channel/UCRwO-ewBHhNiC4qBEppi_JQ), Mr. Masaaki Aiba, and [Mr. Yuto Shinahara](https://twitter.com/snhrytdesu) for valuable discussion.\n\n# References\n\n- [U-Net: Convolutional Networks for Biomedical Image Segmentation](https://arxiv.org/abs/1505.04597) ... original papaer of U-Net\n- [PyTorch](https://pytorch.org/)\n\t- [PyTorch - SAVING AND LOADING MODELS](https://pytorch.org/tutorials/beginner/saving_loading_models.html)\n- [OpenCV](https://opencv.org/)\n- [scikit-image](https://scikit-image.org/)\n- [MLFlow](https://mlflow.org/) \n- [Albumentations Documentation/Mask augmentation for segmentation](https://albumentations.ai/docs/getting_started/mask_augmentation/)\n- [Kaggle - I dont get the discription of the expression of defect pixels](https://www.kaggle.com/c/severstal-steel-defect-detection/discussion/102311)\n","metadata":{"papermill":{"duration":null,"end_time":null,"exception":null,"start_time":null,"status":"pending"},"tags":[]}},{"cell_type":"markdown","source":"","metadata":{}}]}