{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Introduction\n\nThis notebook is a merging from the quoted notebook listed below based on my personal studying and digesting. Thanks a lot to these kindly contributors since their contributions really help me a lot!\n\n* **<a href=\"https://www.kaggle.com/code/shnakazawa/semantic-segmentation-with-pytorch-and-u-net/notebook\" style=\"text-decoration:none\">Semantic Segmentation with PyTorch and U-Net</a>** (mainly)\n* <a href=\"https://www.kaggle.com/code/frozenwolf/sartorius-visualization-training-u-net\" style=\"text-decoration:none\">🦠 Sartorius : Visualization + Training U-Net</a>\n* \n* <a href=\"https://www.kaggle.com/code/nikmarker/sartorius-starter-torch-mask-r-cnn-lb-0-273\" style=\"text-decoration:none\">Sartorius - Starter Torch Mask R-CNN [LB=0.273]</a>\n* <a href=\"https://www.kaggle.com/code/frozenwolf/sartorius-visualization-training-maskr-cnn\" style=\"text-decoration:none\">🧬 Sartorius : Visualization + Training MaskR-CNN</a>\n\nAnd of course, as a noice in Kaggle, I will refernce most of the codes from the notebook above, rearrange and display them on my own way and style.\n\n\n<br>\nRead more about the UNet paper <a href=\"https://arxiv.org/abs/1505.04597\" style=\"text-decoration:none\">here</a>.","metadata":{}},{"cell_type":"markdown","source":"<br>","metadata":{}},{"cell_type":"markdown","source":"# **Import Libraries**","metadata":{}},{"cell_type":"code","source":"import os\nimport gc\nimport time\nimport random\nimport collections\nfrom pathlib import Path\nfrom tqdm.notebook import tqdm\n\nimport math\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\n\nfrom sklearn.model_selection import KFold\n# from sklearn.model_selection import train_test_split\n\nimport cv2\nfrom PIL import Image\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n# from albumentations import HorizontalFlip, VerticalFlip, ShiftScaleRotate, Normalize, Resize, Compose, GaussNoise\n\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\n\nimport torchvision\nfrom torchvision.transforms import ToPILImage\nfrom torchvision.transforms import functional as F\nfrom torchvision.models.detection.faster_rcnn import FastRCNNPredictor\nfrom torchvision.models.detection.mask_rcnn import MaskRCNNPredictor\n\n\nimport warnings\nwarnings.filterwarnings(\"ignore\")","metadata":{"execution":{"iopub.status.busy":"2023-04-23T05:22:12.333596Z","iopub.execute_input":"2023-04-23T05:22:12.334029Z","iopub.status.idle":"2023-04-23T05:22:16.609382Z","shell.execute_reply.started":"2023-04-23T05:22:12.333993Z","shell.execute_reply":"2023-04-23T05:22:16.608253Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(f'PyTorch version: {torch.__version__}')\nprint(f'CUDA avaliable: {torch.cuda.is_available()}')","metadata":{"execution":{"iopub.status.busy":"2023-04-23T05:22:16.611866Z","iopub.execute_input":"2023-04-23T05:22:16.612562Z","iopub.status.idle":"2023-04-23T05:22:16.699268Z","shell.execute_reply.started":"2023-04-23T05:22:16.612518Z","shell.execute_reply":"2023-04-23T05:22:16.698091Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<br>","metadata":{}},{"cell_type":"markdown","source":"# **Fix Randomness**","metadata":{}},{"cell_type":"code","source":"# Fix randomness\n\ndef fix_all_seeds(seed):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n    # torch.backends.cudnn.deterministic = True\n    # torch.backends.cudnn.benchmark = True\n    \n    os.environ['PYTHONHASHSEED'] = str(seed)\n\n\n# Set seed\nfix_all_seeds(2023)","metadata":{"execution":{"iopub.status.busy":"2023-04-23T05:22:16.701390Z","iopub.execute_input":"2023-04-23T05:22:16.702702Z","iopub.status.idle":"2023-04-23T05:22:16.731279Z","shell.execute_reply.started":"2023-04-23T05:22:16.702628Z","shell.execute_reply":"2023-04-23T05:22:16.730118Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Configurations**","metadata":{"execution":{"iopub.status.busy":"2023-04-22T09:41:52.564143Z","iopub.execute_input":"2023-04-22T09:41:52.564607Z","iopub.status.idle":"2023-04-22T09:41:52.570386Z","shell.execute_reply.started":"2023-04-22T09:41:52.564558Z","shell.execute_reply":"2023-04-22T09:41:52.569283Z"}}},{"cell_type":"code","source":"# Directory setting\nDATA_DIR = '/kaggle/input/sartorius-cell-instance-segmentation/'\n\nTRAIN_CSV = DATA_DIR + 'train.csv'\nTRAIN_PATH = DATA_DIR + 'train/'\nTEST_PATH = DATA_DIR + 'test/'\n\nMODEL_DIR = '/kaggle/working/'\n\nFOLDS = 3                # Kfold cross-validation, here using a low value of 3 just for demonstration.\nNUM_WORKERS = 0          # 2\nTHRESHOLD_MASK = 0.5     # Threshold for mask prediction\nBATCH_SIZE = 20\nEPOCHS = 8               # here using a low value of 8 just for demonstration.\n\nUSE_SCHEDULER = False  # Use a StepLR scheduler if True. Not tried yet.\nMOMENTUM = 0.9\nLEARNING_RATE = 0.0001\nWEIGHT_DECAY = 0.0001\n\n## Set device\n# device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') \nDEVICE = torch.device('cuda') if torch.cuda.is_available() else torch.device('cpu')\nprint(f'Using {DEVICE} device')","metadata":{"execution":{"iopub.status.busy":"2023-04-23T05:22:16.734770Z","iopub.execute_input":"2023-04-23T05:22:16.735549Z","iopub.status.idle":"2023-04-23T05:22:16.744463Z","shell.execute_reply.started":"2023-04-23T05:22:16.735508Z","shell.execute_reply":"2023-04-23T05:22:16.743420Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.cuda.get_device_properties(DEVICE)","metadata":{"execution":{"iopub.status.busy":"2023-04-23T05:22:16.746300Z","iopub.execute_input":"2023-04-23T05:22:16.747059Z","iopub.status.idle":"2023-04-23T05:22:16.773539Z","shell.execute_reply.started":"2023-04-23T05:22:16.747018Z","shell.execute_reply":"2023-04-23T05:22:16.772542Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<br>","metadata":{}},{"cell_type":"markdown","source":"# **Helper Functions**","metadata":{}},{"cell_type":"markdown","source":"## rle_decode()","metadata":{}},{"cell_type":"markdown","source":"Decode the RLE (annotation) of an particular cell instance in an image to its correspongding mask.","metadata":{}},{"cell_type":"code","source":"def rle_decode(rle, img_shape, color=1):\n    \"\"\"Decode the RLE (annotation) of an particular cell instance in an image to its correspongding mask.\n\n    Args:\n        rle (str): mask with run length encoding.\n        img_shape ((int, int)): (height, width) of the image, also the shape of mask np.ndarray to return.\n        color (int): brightness of the mask pixel. Default to 1.\n\n    Returns:\n        np.ndarray: 1 - mask, 0 - background.\n    \"\"\"\n    rle_list = rle.split()\n    starts, lengths = [np.asarray(x, dtype=int) for x in (rle_list[0:][::2], rle_list[1:][::2])]\n    starts -= 1\n    ends = starts + lengths\n    \n    mask = np.zeros(img_shape[0] * img_shape[1], dtype=np.uint8)\n    for lo, hi in zip(starts, ends):\n        mask[lo:hi] = color\n    \n    return mask.reshape(img_shape)","metadata":{"execution":{"iopub.status.busy":"2023-04-23T05:22:16.776781Z","iopub.execute_input":"2023-04-23T05:22:16.777051Z","iopub.status.idle":"2023-04-23T05:22:16.787656Z","shell.execute_reply.started":"2023-04-23T05:22:16.777025Z","shell.execute_reply":"2023-04-23T05:22:16.786702Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## rle_encode()","metadata":{}},{"cell_type":"code","source":"def rle_encode(predicted_img):\n    predicted_img = (predicted_img > THRESHOLD_MASK).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    \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    \n    return ' '.join(rle_str)","metadata":{"execution":{"iopub.status.busy":"2023-04-23T05:22:16.790021Z","iopub.execute_input":"2023-04-23T05:22:16.791279Z","iopub.status.idle":"2023-04-23T05:22:16.800412Z","shell.execute_reply.started":"2023-04-23T05:22:16.791237Z","shell.execute_reply":"2023-04-23T05:22:16.799350Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<br>","metadata":{}},{"cell_type":"markdown","source":"## build_image_mask\n\nDecode RLEs (annotations) of all cell instance in an image into one mask image.","metadata":{}},{"cell_type":"code","source":"def build_image_mask(img, rle_list):\n    \"\"\"Decode RLEs (annotations) of all cell instance in an image into one mask image.\n    \n    Args:\n        img (np.ndarray): image with single channel or multiple channels.\n        rle_list (list of str): rles of all cell instances in an image as a list.\n    \n    Returns:\n        np.ndarray: 1 - mask, 0 - background.\n    \"\"\"\n    img_shape = img.shape\n    h = img_shape[0]\n    w = img_shape[1]\n    \n    mask = np.zeros((h, w))\n    for rle in rle_list:\n        mask += rle_decode(rle, (h, w))\n    mask = mask.clip(0,1)\n    \n    return mask","metadata":{"execution":{"iopub.status.busy":"2023-04-23T05:22:16.802130Z","iopub.execute_input":"2023-04-23T05:22:16.802617Z","iopub.status.idle":"2023-04-23T05:22:16.815739Z","shell.execute_reply.started":"2023-04-23T05:22:16.802578Z","shell.execute_reply":"2023-04-23T05:22:16.814598Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<br>","metadata":{}},{"cell_type":"markdown","source":"## transforms()\n\nFor segmentation tasks, the segmentation masks also need to undergo transformations and augmentations.\n\nPlease also see <a href=\"https://albumentations.ai/docs/getting_started/mask_augmentation/\" style=\"text-decoration:none\">Albumentations Documentation/Mask augmentation for segmentation</a>. In this notebook, transforms are applied in the `CellDataset()` class.","metadata":{}},{"cell_type":"code","source":"def transforms(train_only=True):\n    \"\"\"converts the image, a PIL image, into a PyTorch Tensor\"\"\"\n    \n    transforms = [A.Resize(256, 256, p=1)]\n    \n    if train_only:\n        ## during training, randomly flip the training images and ground-truth for data augmentation\n        transforms.append(A.HorizontalFlip(p=0.5))\n        transforms.append(A.VerticalFlip(p=0.5))\n        transforms.append(A.Transpose(p=0.5))\n    else:\n        ## for validation phase and test phase\n        pass\n    \n    transforms.append(ToTensorV2())\n    \n    return A.Compose(transforms)","metadata":{"execution":{"iopub.status.busy":"2023-04-23T05:22:16.817494Z","iopub.execute_input":"2023-04-23T05:22:16.818032Z","iopub.status.idle":"2023-04-23T05:22:16.827962Z","shell.execute_reply.started":"2023-04-23T05:22:16.817929Z","shell.execute_reply":"2023-04-23T05:22:16.826999Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<br>","metadata":{}},{"cell_type":"markdown","source":"## plot_train_val_loss()","metadata":{}},{"cell_type":"code","source":"def plot_train_val_loss(train_loss_fold_list, valid_loss_fold_list, \n                        save=False, save_dir=\"./save_images/\", save_name='segmentation_train_val_loss.png'):\n    fig = plt.figure(figsize=(10,10))\n    for i in range(FOLDS):\n        train_loss = train_loss_fold_list[i]\n        valid_loss = valid_loss_fold_list[i]\n        \n        ax = fig.add_subplot(math.ceil(np.sqrt(FOLDS)), math.ceil(np.sqrt(FOLDS)), 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()","metadata":{"execution":{"iopub.status.busy":"2023-04-23T05:22:16.833840Z","iopub.execute_input":"2023-04-23T05:22:16.834771Z","iopub.status.idle":"2023-04-23T05:22:16.844408Z","shell.execute_reply.started":"2023-04-23T05:22:16.834740Z","shell.execute_reply":"2023-04-23T05:22:16.843181Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<br>","metadata":{}},{"cell_type":"markdown","source":"# **EDA (Exploratory Data Analysis)**\n\nFirst, We will do some EDA (Exploratory Data Analysis) of the `train.csv` so that we can better understand our training dataset.","metadata":{}},{"cell_type":"markdown","source":"## Explore `train.csv`","metadata":{}},{"cell_type":"code","source":"df_train = pd.read_csv(TRAIN_CSV)\ndf_train.head().append(df_train.tail())","metadata":{"execution":{"iopub.status.busy":"2023-04-23T05:22:16.846020Z","iopub.execute_input":"2023-04-23T05:22:16.846731Z","iopub.status.idle":"2023-04-23T05:22:17.471336Z","shell.execute_reply.started":"2023-04-23T05:22:16.846693Z","shell.execute_reply":"2023-04-23T05:22:17.470127Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_train.shape","metadata":{"execution":{"iopub.status.busy":"2023-04-23T05:22:17.473185Z","iopub.execute_input":"2023-04-23T05:22:17.473906Z","iopub.status.idle":"2023-04-23T05:22:17.481654Z","shell.execute_reply.started":"2023-04-23T05:22:17.473861Z","shell.execute_reply":"2023-04-23T05:22:17.480154Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_train.dtypes","metadata":{"execution":{"iopub.status.busy":"2023-04-23T05:22:17.483480Z","iopub.execute_input":"2023-04-23T05:22:17.484343Z","iopub.status.idle":"2023-04-23T05:22:17.495098Z","shell.execute_reply.started":"2023-04-23T05:22:17.484303Z","shell.execute_reply":"2023-04-23T05:22:17.494100Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<br>","metadata":{}},{"cell_type":"markdown","source":"### Cell Type","metadata":{}},{"cell_type":"code","source":"df_train.cell_type.value_counts()","metadata":{"execution":{"iopub.status.busy":"2023-04-23T05:22:17.496458Z","iopub.execute_input":"2023-04-23T05:22:17.497546Z","iopub.status.idle":"2023-04-23T05:22:17.515455Z","shell.execute_reply.started":"2023-04-23T05:22:17.497508Z","shell.execute_reply":"2023-04-23T05:22:17.514110Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_train.cell_type.unique()","metadata":{"execution":{"iopub.status.busy":"2023-04-23T05:22:17.517092Z","iopub.execute_input":"2023-04-23T05:22:17.517552Z","iopub.status.idle":"2023-04-23T05:22:17.536759Z","shell.execute_reply.started":"2023-04-23T05:22:17.517509Z","shell.execute_reply":"2023-04-23T05:22:17.535609Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<br>","metadata":{}},{"cell_type":"markdown","source":"### Annotation (RLE)\n\nThe segment of training data is provided with **Run length encoding (RLE)** in the `annotation` column.\n\n\nRLE is a lossless compression technique used to represent data that contains long sequences of repeated values or characters. <a href=\"https://www.kaggle.com/c/severstal-steel-defect-detection/discussion/102311\" style=\"text-decoration:none\">Check out this discussion for a better understanding</a>.\n\n\nThe current dataset has one segment per row.","metadata":{}},{"cell_type":"code","source":"# train.csv 中 id 列是图片编号，每一行的 annotation 都是该图片中的一个细胞实例的 mask 数据。\n# 若一张图片中有 395 个细胞实例的 mask，则这张图片会在表格中出现 395 行。\ndf_instances = df_train.groupby(['id']).agg({'annotation': 'count', 'cell_type': 'first'})\ndf_instances.head()","metadata":{"execution":{"iopub.status.busy":"2023-04-23T05:22:17.540461Z","iopub.execute_input":"2023-04-23T05:22:17.540793Z","iopub.status.idle":"2023-04-23T05:22:17.575389Z","shell.execute_reply.started":"2023-04-23T05:22:17.540764Z","shell.execute_reply":"2023-04-23T05:22:17.574226Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_instances.shape","metadata":{"execution":{"iopub.status.busy":"2023-04-23T05:22:17.577094Z","iopub.execute_input":"2023-04-23T05:22:17.577554Z","iopub.status.idle":"2023-04-23T05:22:17.587091Z","shell.execute_reply.started":"2023-04-23T05:22:17.577514Z","shell.execute_reply":"2023-04-23T05:22:17.585897Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 每一种细胞类型（一张图片只有一种细胞）它在一张图片中 instance segmentation mask\n# (亦即表中对应图片的 annotation 行数) 的总数的分布的分位数情况。\ndf_instances_pentiles = df_train.groupby(['id']).agg({'annotation': 'count', 'cell_type': 'first'})\ndf_instances_pentiles = df_instances_pentiles.groupby(\"cell_type\")[['annotation']]\\\n                                             .describe(percentiles=[0.1, 0.25, 0.75, 0.8, 0.85, 0.9, 0.95, 0.99]).astype(int)\\\n                                             .T.droplevel(level=0).T.drop(['count', '50%', 'std'], axis=1)\ndf_instances_pentiles","metadata":{"execution":{"iopub.status.busy":"2023-04-23T05:22:17.588894Z","iopub.execute_input":"2023-04-23T05:22:17.589363Z","iopub.status.idle":"2023-04-23T05:22:17.652889Z","shell.execute_reply.started":"2023-04-23T05:22:17.589278Z","shell.execute_reply":"2023-04-23T05:22:17.651468Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Trying with different strategies\ndf_instances_pentiles['90%'].to_dict()","metadata":{"execution":{"iopub.status.busy":"2023-04-23T05:22:17.654943Z","iopub.execute_input":"2023-04-23T05:22:17.655397Z","iopub.status.idle":"2023-04-23T05:22:17.663775Z","shell.execute_reply.started":"2023-04-23T05:22:17.655357Z","shell.execute_reply":"2023-04-23T05:22:17.662444Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<br>","metadata":{}},{"cell_type":"markdown","source":"### Pixels Per Mask Per `cell_type`","metadata":{}},{"cell_type":"code","source":"df_train['n_pixels'] = df_train.annotation.apply(lambda x: np.sum([int(e) for e in x.split()[1:][::2]]))\n# 得到每一个 annotation 所描述的 mask 的面积（像素点个数）。\n\n# 各类型细胞的一个 annotation 对应的一个 instance segmentation mask 的像素点个数的分布的分位数情况。\ndf_pixels = df_train.groupby(\"cell_type\")[['n_pixels']].describe(percentiles=[0.02, 0.05, 0.1, 0.9, 0.95, 0.98])\\\n                    .astype(int).T.droplevel(level=0).T.drop(['count', '50%', 'std'], axis=1)\ndf_pixels","metadata":{"execution":{"iopub.status.busy":"2023-04-23T05:22:17.665933Z","iopub.execute_input":"2023-04-23T05:22:17.666943Z","iopub.status.idle":"2023-04-23T05:22:19.291516Z","shell.execute_reply.started":"2023-04-23T05:22:17.666829Z","shell.execute_reply":"2023-04-23T05:22:19.290118Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<br>","metadata":{}},{"cell_type":"markdown","source":"### Image Shape and Format","metadata":{}},{"cell_type":"code","source":"df_train.width.unique()","metadata":{"execution":{"iopub.status.busy":"2023-04-23T05:22:19.294008Z","iopub.execute_input":"2023-04-23T05:22:19.294582Z","iopub.status.idle":"2023-04-23T05:22:19.305396Z","shell.execute_reply.started":"2023-04-23T05:22:19.294539Z","shell.execute_reply":"2023-04-23T05:22:19.303525Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_train.height.unique()","metadata":{"execution":{"iopub.status.busy":"2023-04-23T05:22:19.307770Z","iopub.execute_input":"2023-04-23T05:22:19.308320Z","iopub.status.idle":"2023-04-23T05:22:19.319967Z","shell.execute_reply.started":"2023-04-23T05:22:19.308275Z","shell.execute_reply":"2023-04-23T05:22:19.317720Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img_shapes = set()\nimg_exts = set()\nimg_paths = Path(TRAIN_PATH).glob(\"*\")\n\nbar = tqdm(img_paths, total=df_train.id.unique().shape[0])    # should not be: total=len(list(img_paths))\n\nfor img_path in bar:\n    img_bgr = cv2.imread(img_path.as_posix())                 # BGR mode\n    img_rgb = img_bgr[:, :, ::-1]                             # RGB mode\n    \n    img_shapes.add(img_rgb.shape)\n    img_exts.add(img_path.suffix)\nprint(f'Image shapes are {img_shapes}.')\nprint(f'Image extensions are {img_exts}.')","metadata":{"execution":{"iopub.status.busy":"2023-04-23T05:22:19.322352Z","iopub.execute_input":"2023-04-23T05:22:19.323321Z","iopub.status.idle":"2023-04-23T05:22:28.515347Z","shell.execute_reply.started":"2023-04-23T05:22:19.323276Z","shell.execute_reply":"2023-04-23T05:22:28.514081Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<br>","metadata":{}},{"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.\n\nAs the figure displayed below, the image signals appear to be quite uniform, which is characteristic of typical microscopy images. :)","metadata":{}},{"cell_type":"code","source":"img_paths = Path(TRAIN_PATH).glob(\"*\")\nbar = tqdm(img_paths, total=df_train.id.unique().shape[0])     # should not be: total=len(list(img_paths))\n\nplt.figure(figsize=(8,8))\nfor img_path in bar:\n    img_bgr = cv2.imread(img_path.as_posix())    # BGR mode\n    img_rgb = img_bgr[:, :, ::-1]                # RGB mode\n    \n    hist = cv2.calcHist([img_rgb], [0], None ,[256], [0,256])\n    plt.plot(hist)\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-04-23T05:22:28.517413Z","iopub.execute_input":"2023-04-23T05:22:28.518251Z","iopub.status.idle":"2023-04-23T05:22:37.406225Z","shell.execute_reply.started":"2023-04-23T05:22:28.518202Z","shell.execute_reply":"2023-04-23T05:22:37.405071Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<br>","metadata":{}},{"cell_type":"markdown","source":"## Display Some Masks","metadata":{}},{"cell_type":"markdown","source":"### shsy5y","metadata":{}},{"cell_type":"code","source":"plt.figure(figsize=(16, 8))\nimg_id = \"0030fd0e6378\"\nimg = cv2.imread(TRAIN_PATH + img_id + \".png\")         # BGR mode\n# df_train[df_train.id == img_id].annotation.tolist()\nax1 = plt.subplot(121)\nax1.imshow(img)\nax1.set_title(\"Shsy5y Original Image (BGR mode)\")\n\nmask = build_image_mask(img, df_train[df_train.id == img_id].annotation.tolist())\nax2 = plt.subplot(122)\nax2.imshow(mask)\n# plt.imshow(mask, cmap=\"gray\")\n# plt.imshow(mask, cmap = plt.cm.gray)\n# plt.imshow(mask, cmap = plt.cm.gray_r)\nax2.set_title(\"Shsy5y Mask\")\n\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-04-23T05:22:37.408198Z","iopub.execute_input":"2023-04-23T05:22:37.408997Z","iopub.status.idle":"2023-04-23T05:22:38.181466Z","shell.execute_reply.started":"2023-04-23T05:22:37.408949Z","shell.execute_reply":"2023-04-23T05:22:38.180436Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.figure(figsize=(16, 8))\nimg_id = \"0030fd0e6378\"\nimg = cv2.imread(TRAIN_PATH + img_id + \".png\")          # BGR mode\nplt.imshow(img)\nplt.title(\"Cort Original Image (BGR mode) and Mask\")\n\nmask = build_image_mask(img, df_train[df_train.id == img_id].annotation.tolist())\nplt.imshow(mask, alpha=0.2)\n\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-04-23T05:22:38.182642Z","iopub.execute_input":"2023-04-23T05:22:38.183025Z","iopub.status.idle":"2023-04-23T05:22:39.094531Z","shell.execute_reply.started":"2023-04-23T05:22:38.182984Z","shell.execute_reply":"2023-04-23T05:22:39.093492Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<br>","metadata":{}},{"cell_type":"markdown","source":"### astro","metadata":{}},{"cell_type":"code","source":"plt.figure(figsize=(16, 8))\nimg_id = \"0140b3c8f445\"\nimg = cv2.imread(TRAIN_PATH + img_id + \".png\")         # BGR mode\n# df_train[df_train.id == img_id].annotation.tolist()\nax1 = plt.subplot(121)\nax1.imshow(img)\nax1.set_title(\"Astro Original Image (BGR mode)\")\n\nmask = build_image_mask(img, df_train[df_train.id == img_id].annotation.tolist())\nax2 = plt.subplot(122)\nax2.imshow(mask)\n# plt.imshow(mask, cmap=\"gray\")\n# plt.imshow(mask, cmap = plt.cm.gray)\n# plt.imshow(mask, cmap = plt.cm.gray_r)\nax2.set_title(\"Astro Mask\")\n\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-04-23T05:22:39.095984Z","iopub.execute_input":"2023-04-23T05:22:39.096688Z","iopub.status.idle":"2023-04-23T05:22:39.744767Z","shell.execute_reply.started":"2023-04-23T05:22:39.096646Z","shell.execute_reply":"2023-04-23T05:22:39.743779Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.figure(figsize=(16, 8))\nimg_id = \"0140b3c8f445\"\nimg = cv2.imread(TRAIN_PATH + img_id + \".png\")          # BGR mode\nplt.imshow(img)\nplt.title(\"Cort Original Image (BGR mode) and Mask\")\n\nmask = build_image_mask(img, df_train[df_train.id == img_id].annotation.tolist())\nplt.imshow(mask, alpha=0.05)\n\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-04-23T05:22:39.751476Z","iopub.execute_input":"2023-04-23T05:22:39.752485Z","iopub.status.idle":"2023-04-23T05:22:40.555013Z","shell.execute_reply.started":"2023-04-23T05:22:39.752452Z","shell.execute_reply":"2023-04-23T05:22:40.554043Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<br>","metadata":{}},{"cell_type":"markdown","source":"### cort","metadata":{}},{"cell_type":"code","source":"plt.figure(figsize=(16, 8))\nimg_id = \"01ae5a43a2ab\"\nimg = cv2.imread(TRAIN_PATH + img_id + \".png\")         # BGR mode\n# df_train[df_train.id == img_id].annotation.tolist()\nax1 = plt.subplot(121)\nax1.imshow(img)\nax1.set_title(\"Cort Original Image (BGR mode)\")\n\nmask = build_image_mask(img, df_train[df_train.id == img_id].annotation.tolist())\nax2 = plt.subplot(122)\nax2.imshow(mask)\n# plt.imshow(mask, cmap=\"gray\")\n# plt.imshow(mask, cmap = plt.cm.gray)\n# plt.imshow(mask, cmap = plt.cm.gray_r)\nax2.set_title(\"Cort Mask\")\n\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-04-23T05:22:40.556587Z","iopub.execute_input":"2023-04-23T05:22:40.557551Z","iopub.status.idle":"2023-04-23T05:22:41.154598Z","shell.execute_reply.started":"2023-04-23T05:22:40.557512Z","shell.execute_reply":"2023-04-23T05:22:41.153688Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.figure(figsize=(16, 8))\nimg_id = \"01ae5a43a2ab\"\nimg = cv2.imread(TRAIN_PATH + img_id + \".png\")          # BGR mode\nplt.imshow(img)\nplt.title(\"Cort Original Image (BGR mode) and Mask\")\n\nmask = build_image_mask(img, df_train[df_train.id == img_id].annotation.tolist())\nplt.imshow(mask, alpha=0.2)\n\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-04-23T05:22:41.156051Z","iopub.execute_input":"2023-04-23T05:22:41.156685Z","iopub.status.idle":"2023-04-23T05:22:41.934702Z","shell.execute_reply.started":"2023-04-23T05:22:41.156648Z","shell.execute_reply":"2023-04-23T05:22:41.932908Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<br>","metadata":{}},{"cell_type":"markdown","source":"# Dataset and DataLoader","metadata":{}},{"cell_type":"code","source":"# 与前面定义的 df_instances 有些类似。\n\ndf_grouped = df_train.copy()\ndf_grouped['segment_count'] = 1    # train.csv 中 id 列是图片编号，每一行的 annotation 是该图片中一个细胞的 rle。\ndf_grouped = df_grouped.groupby(['id', 'width', 'height', 'cell_type']).count().reset_index()\ndf_grouped = df_grouped[['id', 'width', 'height', 'cell_type', 'segment_count']]\ndf_grouped","metadata":{"execution":{"iopub.status.busy":"2023-04-23T05:22:41.936339Z","iopub.execute_input":"2023-04-23T05:22:41.937367Z","iopub.status.idle":"2023-04-23T05:22:42.022858Z","shell.execute_reply.started":"2023-04-23T05:22:41.937326Z","shell.execute_reply":"2023-04-23T05:22:42.021489Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_train.id.unique().shape","metadata":{"execution":{"iopub.status.busy":"2023-04-23T05:22:42.024769Z","iopub.execute_input":"2023-04-23T05:22:42.025505Z","iopub.status.idle":"2023-04-23T05:22:42.040526Z","shell.execute_reply.started":"2023-04-23T05:22:42.025460Z","shell.execute_reply":"2023-04-23T05:22:42.038942Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<br>","metadata":{}},{"cell_type":"markdown","source":"## Define CellDataset()","metadata":{}},{"cell_type":"code","source":"class CellDataset(Dataset):\n    def __init__(self, img_ids, df, train_path, transforms=None, train_val=True):\n        super().__init__()\n        self.df = df\n        # self.img_ids = self.df.id.unique()    # 1D array\n        self.img_ids = img_ids\n        self.train_path = train_path\n        self.transforms = transforms\n        self.train_val = train_val\n    \n    \n    def __len__(self):\n        return self.img_ids.shape[0]\n    \n    \n    def __getitem__(self, index):\n        img_id = self.img_ids[index]\n        \n        # Read image\n        img_bgr = cv2.imread(self.train_path + img_id + \".png\")\n        img_rgb = img_bgr[:, :, ::-1].astype(np.float32)\n        img = cv2.cvtColor(img_rgb, cv2.COLOR_RGB2GRAY)       # 3 channels -> 1 channel\n        img /= 255.0                                          # normalization\n        \n        \n        if self.train_val:  # For training and validation\n            # Create mask\n            rle_list = self.df[self.df.id == img_id]['annotation'].tolist()\n            mask = build_image_mask(img, rle_list)\n            \n            # Transform image and mask\n            if self.transforms:\n                transformed = self.transforms(image=img, mask=mask)   # params: image, mask, bboxes, keypoints\n                img, mask = transformed['image'], transformed['mask']\n                \n            return img, mask\n        \n        else: # For testing\n            # Transform images\n            if self.transforms:\n                img =self.transforms(image=img)['image']\n            \n            return img, img_id","metadata":{"execution":{"iopub.status.busy":"2023-04-23T05:22:42.042438Z","iopub.execute_input":"2023-04-23T05:22:42.043411Z","iopub.status.idle":"2023-04-23T05:22:42.055689Z","shell.execute_reply.started":"2023-04-23T05:22:42.043366Z","shell.execute_reply":"2023-04-23T05:22:42.054504Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<br>","metadata":{}},{"cell_type":"markdown","source":"## Define train_val_dataloader()","metadata":{}},{"cell_type":"code","source":"def train_val_dataloader(df_grouped, df, train_idx, val_idx):\n    train_ = df_grouped.loc[trn_idx,:].reset_index(drop=True)\n    valid_ = df_grouped.loc[val_idx,:].reset_index(drop=True)\n    \n    # Dataset\n    train_dataset = CellDataset(train_[\"id\"].to_numpy(), \n                                df, \n                                TRAIN_PATH, \n                                transforms=transforms())\n    valid_dataset = CellDataset(valid_[\"id\"].to_numpy(), \n                                df, \n                                TRAIN_PATH, \n                                transforms=transforms(train_only=False))\n    \n    # DataLoader\n    train_loader = DataLoader(train_dataset, \n                              batch_size=BATCH_SIZE, \n                              num_workers=NUM_WORKERS, \n                              shuffle=True\n                             )\n    valid_loader = DataLoader(valid_dataset, \n                              batch_size=BATCH_SIZE, \n                              num_workers=NUM_WORKERS, \n                              shuffle=False\n                             )\n    \n    return train_loader, valid_loader","metadata":{"execution":{"iopub.status.busy":"2023-04-23T05:22:42.057217Z","iopub.execute_input":"2023-04-23T05:22:42.059655Z","iopub.status.idle":"2023-04-23T05:22:42.072049Z","shell.execute_reply.started":"2023-04-23T05:22:42.059578Z","shell.execute_reply":"2023-04-23T05:22:42.070892Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<br>","metadata":{}},{"cell_type":"markdown","source":"## Define the UNet Model\n\nIn this notebook, we will implement the **U-Net** architecture using PyTorch, without relying on pre-existing implementations or models.","metadata":{}},{"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    \n    def __init__(self, in_channels, out_channels):\n        super(DoubleConv, self).__init__()\n        \n        self.double_conv = nn.Sequential(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\n        return x","metadata":{"execution":{"iopub.status.busy":"2023-04-23T05:22:42.073961Z","iopub.execute_input":"2023-04-23T05:22:42.074874Z","iopub.status.idle":"2023-04-23T05:22:42.087307Z","shell.execute_reply.started":"2023-04-23T05:22:42.074830Z","shell.execute_reply":"2023-04-23T05:22:42.086385Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Down(nn.Module):\n    \"\"\"Downscaling\"\"\"\n    \n    def __init__(self, in_channels, out_channels):\n        super(Down, self).__init__()\n        \n        self.maxpool_conv = nn.Sequential(nn.MaxPool2d(2),\n                                          DoubleConv(in_channels, out_channels)\n                                         )\n        \n    def forward(self, x):\n        x = self.maxpool_conv(x)\n        \n        return x","metadata":{"execution":{"iopub.status.busy":"2023-04-23T05:22:42.088890Z","iopub.execute_input":"2023-04-23T05:22:42.090200Z","iopub.status.idle":"2023-04-23T05:22:42.102847Z","shell.execute_reply.started":"2023-04-23T05:22:42.090157Z","shell.execute_reply":"2023-04-23T05:22:42.101957Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class 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)   # 这里为什么会有 3 个参数？？？\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 (batch_size, channels, height, width): low-level feature from expanding path\n           x2 (batch_size, channels, height, width): high-level feature from contracting path\n        \"\"\"\n        x1 = self.up(x1)\n        \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                              )\n        x = torch.cat([x2, x1], dim=1)\n        x = self.conv(x)\n        \n        return x","metadata":{"execution":{"iopub.status.busy":"2023-04-23T05:22:42.104473Z","iopub.execute_input":"2023-04-23T05:22:42.105429Z","iopub.status.idle":"2023-04-23T05:22:42.116840Z","shell.execute_reply.started":"2023-04-23T05:22:42.105398Z","shell.execute_reply":"2023-04-23T05:22:42.115717Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<br>","metadata":{}},{"cell_type":"code","source":"class UNet(nn.Module):\n    def __init__(self, in_channels=1, n_classes=1, bilinear=False):\n        super(UNet, self).__init__()\n        # self.in_conv = DoubleConv(in_channels, 64)   # 报错：NotImplementedError: \n        self.in_conv = (DoubleConv(in_channels, 64))   # 报错：NotImplementedError: \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.out_conv = nn.Conv2d(64, n_classes, 1)\n    \n    \n    def forward(self, x):\n        x1 = self.in_conv(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.out_conv(x)\n        x = torch.sigmoid(x)\n        \n        return x","metadata":{"execution":{"iopub.status.busy":"2023-04-23T05:22:42.118324Z","iopub.execute_input":"2023-04-23T05:22:42.118801Z","iopub.status.idle":"2023-04-23T05:22:42.133853Z","shell.execute_reply.started":"2023-04-23T05:22:42.118759Z","shell.execute_reply":"2023-04-23T05:22:42.133034Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<br>","metadata":{}},{"cell_type":"markdown","source":"# 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 <a href=\"https://mlflow.org/\" style=\"text-decoration:none\">MLFlow</a>.\n\nBefore running the following cell, please set configs <del>in the **Configurations** section</del>: `RUN_TRAINING = True` and `TRAIN_ALL = (as you want)`.","metadata":{}},{"cell_type":"markdown","source":"## Training and Validation","metadata":{}},{"cell_type":"code","source":"%%time\n\nRUN_TRAINING = True\nTRAIN_ALL = False\n\nif RUN_TRAINING:\n    if TRAIN_ALL:  # Train with all data\n        folds = [['','']]\n    else:          # Cross-validation\n        folds = KFold(n_splits=FOLDS, shuffle=True, random_state=2023)\\\n                .split(np.arange(df_grouped.shape[0]), df_grouped['id'].to_numpy())\n        \n        train_loss_fold_list = []\n        valid_loss_fold_list = []\n    \n    for fold, (trn_idx, val_idx) in enumerate(folds):\n        # Load Data   \n        if TRAIN_ALL:\n            train_dataset = CellDataset(df_grouped['id'].to_numpy(), df_train, DATA_DIR+'/train/', transforms=transform_train())\n            train_loader = DataLoader(train_dataset, \n                                      batch_size=BATCH_SIZE, \n                                      num_workers=NUM_WORKERS, \n                                      shuffle=True\n                                      )\n        else:\n            print(\"=\"*42 + f' Cross-Validation Fold {fold+1} ' + \"=\"*42)   \n            train_loader, valid_loader = train_val_dataloader(df_grouped, df_train, trn_idx, val_idx)\n            \n            valid_loss_epoch_list = []\n        train_loss_epoch_list = []\n        \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=LEARNING_RATE, 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) in pbar:\n            # for (imgs, masks) in train_loader:\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)           # Average train loss of this epoch\n\n            # Validation\n            if TRAIN_ALL == False:\n                print(f'Epoch {epoch+1} Start Validation:')\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) in pbar:\n                    # for (imgs, masks) in valid_loader:\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)   # Average validation loss of this epoch\n                    \n            # print results from this epoch\n            exec_t = int((time.time() - time_start)/60)\n            if TRAIN_ALL:\n                print(\">\"*30 + f' train_loss: {train_loss:.4f}  (Exec time {exec_t} min)\\n')\n\n            else:\n                print(\">\"*30 + f' train_loss: {train_loss:.4f}      val_loss : {valid_loss:.4f}  (Exec time {exec_t} min)\\n')\n                train_loss_epoch_list.append(train_loss)\n                valid_loss_epoch_list.append(valid_loss)\n        \n        if TRAIN_ALL:\n            print(\">\"*30 + 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_fold_list.append(train_loss_epoch_list)\n            valid_loss_fold_list.append(valid_loss_epoch_list)\n            del model, optimizer, train_loader, valid_loader, train_loss_epoch_list, valid_loss_epoch_list\n        gc.collect()\n        torch.cuda.empty_cache()\n    \n    if TRAIN_ALL == False:\n        plot_train_val_loss(train_loss_fold_list, valid_loss_fold_list)\n\nelse:\n    print('RUN_TRAINING is False')","metadata":{"execution":{"iopub.status.busy":"2023-04-23T05:22:42.135404Z","iopub.execute_input":"2023-04-23T05:22:42.136420Z","iopub.status.idle":"2023-04-23T05:50:21.171516Z","shell.execute_reply.started":"2023-04-23T05:22:42.136377Z","shell.execute_reply":"2023-04-23T05:50:21.170404Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<br>","metadata":{}},{"cell_type":"markdown","source":"## Train All","metadata":{}},{"cell_type":"code","source":"%%time\n\nRUN_TRAINING = True\nTRAIN_ALL = True\nEPOCHS = 30\n\nif RUN_TRAINING:\n    if TRAIN_ALL:  # Train with all data\n        folds = [['','']]\n    else:          # Cross-validation\n        folds = KFold(n_splits=FOLDS, shuffle=True, random_state=2023)\\\n                .split(np.arange(df_grouped.shape[0]), df_grouped['id'].to_numpy())\n        \n        train_loss_fold_list = []\n        valid_loss_fold_list = []\n    \n    for fold, (trn_idx, val_idx) in enumerate(folds):\n        # Load Data   \n        if TRAIN_ALL:\n            train_dataset = CellDataset(df_grouped['id'].to_numpy(), df_train, TRAIN_PATH, transforms=transforms())\n            train_loader = DataLoader(train_dataset, \n                                      batch_size=BATCH_SIZE, \n                                      num_workers=NUM_WORKERS, \n                                      shuffle=True\n                                      )\n        else:\n            print(\"=\"*42 + f' Cross-Validation Fold {fold+1} ' + \"=\"*42)   \n            train_loader, valid_loader = train_val_dataloader(df_grouped, df_train, trn_idx, val_idx)\n            \n            valid_loss_epoch_list = []\n        \n        train_loss_epoch_list = []\n        \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=LEARNING_RATE, 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) in pbar:\n            # for (imgs, masks) in train_loader:\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)           # Average train loss of this epoch\n\n            # Validation\n            if TRAIN_ALL == False:\n                print(f'Epoch {epoch+1} Start Validation:')\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) in pbar:\n                    # for (imgs, masks) in valid_loader:\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)   # Average validation loss of this epoch\n                    \n            # print results from this epoch\n            exec_t = int((time.time() - time_start)/60)\n            if TRAIN_ALL:\n                print(\">\"*30 + f' train_loss: {train_loss:.4f}  (Exec time {exec_t} min)\\n')\n\n            else:\n                print(\">\"*30 + f' train_loss: {train_loss:.4f}      val_loss : {valid_loss:.4f}  (Exec time {exec_t} min)\\n')\n                train_loss_epoch_list.append(train_loss)\n                valid_loss_epoch_list.append(valid_loss)\n        \n        if TRAIN_ALL:\n            print(\">\"*30 + f'Save model trained with all data:')\n            os.makedirs(MODEL_DIR + \"save_models/\", exist_ok=True)\n            torch.save(model.state_dict(), MODEL_DIR + 'save_models/pytorch_unet.pth')\n            del model, optimizer, train_loader\n        else:\n            train_loss_fold_list.append(train_loss_epoch_list)\n            valid_loss_fold_list.append(valid_loss_epoch_list)\n            del model, optimizer, train_loader, valid_loader, train_loss_epoch_list, valid_loss_epoch_list\n        gc.collect()\n        torch.cuda.empty_cache()\n    \n    if TRAIN_ALL == False:\n        plot_train_val_loss(train_loss_fold_list, valid_loss_fold_list)\n\nelse:\n    print('RUN_TRAINING is False')","metadata":{"execution":{"iopub.status.busy":"2023-04-23T05:50:21.173569Z","iopub.execute_input":"2023-04-23T05:50:21.174354Z","iopub.status.idle":"2023-04-23T06:27:29.391129Z","shell.execute_reply.started":"2023-04-23T05:50:21.174311Z","shell.execute_reply":"2023-04-23T06:27:29.389875Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Inference of Test Data\n\n\n## Inference\n\nIn this task, the same Dataset, Image transformation, and model as we used for validation. So no need to define new functions.\n\nBefore running the following cells, please set `RUN_INFERENCE = True` <del>in the **Set Config** section</del>.","metadata":{}},{"cell_type":"code","source":"%%time\n\nRUN_INFERENCE = True\n\nif RUN_INFERENCE:\n    files = os.listdir(TEST_PATH)\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)，需要传递给 cv2.resize() 第二个参数 dsize  的形式。\n\n    # Load Data\n    test_dataset = CellDataset(image_ids, \n                               df_train, \n                               TEST_PATH, \n                               transforms=transforms(train_only=False), \n                               train_val=False)\n\n    # Data Loader\n    test_loader = DataLoader(test_dataset, \n                             batch_size=BATCH_SIZE, \n                             num_workers=NUM_WORKERS, \n                             shuffle=False\n                             )\n\n    # Load model, loss function, and optimizing algorithm\n    model = UNet().to(DEVICE)\n    model.load_state_dict(torch.load(MODEL_DIR + 'save_models/pytorch_unet.pth'))\n       \n    # Start Inference\n    print(\"=\"*42 + f' Start Inference ' + \"=\"*42)\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            print(\"step: \", step)\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                # 这里需要注意，cv2.resize() 第二个参数 dsize  的形式虽然是 (width, height) 的形式，此处为 (704, 520)。\n                # 但是 cv2.resize() 的输出图像仍然是 (height, width) 的形式，即此处为 (520, 704)\n                \n                rle_mask = rle_encode(predicted_mask)\n                ids.append(image_id)\n                rle_test_preds.append(rle_mask)\n    \n    df_submission = pd.DataFrame({'id': ids, \n                                  'predicted': rle_test_preds})\n    print(df_submission.head())\n\nelse:\n    print('RUN_INFERENCE is False')","metadata":{"execution":{"iopub.status.busy":"2023-04-23T06:27:29.393341Z","iopub.execute_input":"2023-04-23T06:27:29.393744Z","iopub.status.idle":"2023-04-23T06:27:30.388201Z","shell.execute_reply.started":"2023-04-23T06:27:29.393704Z","shell.execute_reply":"2023-04-23T06:27:30.386958Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Visualize Predicted Mask","metadata":{}},{"cell_type":"code","source":"predicted_mask_last = predicted_mask\nimage_id_last = image_id\n\nprint(type(predicted_mask_last))\nprint(predicted_mask_last.shape)\nprint(\"image id: \", image_id_last)","metadata":{"execution":{"iopub.status.busy":"2023-04-23T06:27:30.393703Z","iopub.execute_input":"2023-04-23T06:27:30.395379Z","iopub.status.idle":"2023-04-23T06:27:30.403452Z","shell.execute_reply.started":"2023-04-23T06:27:30.395332Z","shell.execute_reply":"2023-04-23T06:27:30.401706Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img = cv2.imread(TEST_PATH + image_id_last + \".png\")   # BGR mode\n\nplt.figure(figsize=(15, 20))\nax1 = plt.subplot(121)\nax1.imshow(img)\nax1.set_title(f\"image id: {image_id_last} (BGR mode)\",\n              fontdict={\"fontsize\":22})\n\nax2 = plt.subplot(122)\nax2.imshow(predicted_mask_last > THRESHOLD_MASK)\nax2.set_title(f\"Mask Threshold: {THRESHOLD_MASK}\",\n              fontdict={\"fontsize\":22})\n\nplt.tight_layout()\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-04-23T06:27:30.405924Z","iopub.execute_input":"2023-04-23T06:27:30.406346Z","iopub.status.idle":"2023-04-23T06:27:31.088054Z","shell.execute_reply.started":"2023-04-23T06:27:30.406305Z","shell.execute_reply":"2023-04-23T06:27:31.087087Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.figure(figsize=(10,10))\nplt.imshow(img)\nplt.imshow(predicted_mask_last > THRESHOLD_MASK, alpha=0.08)\nplt.title(f\"Image(BGR) + Predicted Mask({THRESHOLD_MASK})\", fontdict={\"fontsize\":15})\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-04-23T06:27:31.089796Z","iopub.execute_input":"2023-04-23T06:27:31.090459Z","iopub.status.idle":"2023-04-23T06:27:31.641103Z","shell.execute_reply.started":"2023-04-23T06:27:31.090420Z","shell.execute_reply":"2023-04-23T06:27:31.640093Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<br>","metadata":{}},{"cell_type":"code","source":"## Try another mask threshold value\nthreshold_mask = 0.25\nplt.figure(figsize=(15, 15))\nax1 = plt.subplot(211)\nax1.imshow(predicted_mask_last > threshold_mask)\nax1.set_title(f\"Mask Threshold: {threshold_mask}\", fontdict={\"fontsize\":15})\n\nax2 = plt.subplot(212)\nax2.imshow(img)\nax2.imshow(predicted_mask_last > threshold_mask, alpha=0.08)\nax2.set_title(f\"Image(BGR) + Predicted Mask({threshold_mask})\", fontdict={\"fontsize\":15})\nplt.tight_layout()\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-04-23T06:27:31.642894Z","iopub.execute_input":"2023-04-23T06:27:31.643633Z","iopub.status.idle":"2023-04-23T06:27:32.628343Z","shell.execute_reply.started":"2023-04-23T06:27:31.643591Z","shell.execute_reply":"2023-04-23T06:27:32.627034Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<br>","metadata":{}},{"cell_type":"markdown","source":"## Visualize Predicted RLE Mask\n\nThis confirms whether the output style is correct.","metadata":{}},{"cell_type":"code","source":"target = df_submission.iloc[0]\nimg = cv2.imread(f'{TEST_PATH}{target[\"id\"]}.png')\nrle_predicted = [target['predicted']]\nmask = build_image_mask(img, rle_predicted)\n\nplt.figure()\nplt.imshow(img)\n\nplt.figure()\nplt.imshow(mask)\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-04-23T06:27:32.630008Z","iopub.execute_input":"2023-04-23T06:27:32.630510Z","iopub.status.idle":"2023-04-23T06:27:33.253361Z","shell.execute_reply.started":"2023-04-23T06:27:32.630475Z","shell.execute_reply":"2023-04-23T06:27:33.252104Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<br>","metadata":{}},{"cell_type":"markdown","source":"# Submission","metadata":{}},{"cell_type":"code","source":"df_submission.to_csv(r'/kaggle/working/Unet_Pytorch_Mine.csv', index=False)","metadata":{"execution":{"iopub.status.busy":"2023-04-23T06:30:17.736752Z","iopub.execute_input":"2023-04-23T06:30:17.737747Z","iopub.status.idle":"2023-04-23T06:30:17.744806Z","shell.execute_reply.started":"2023-04-23T06:30:17.737705Z","shell.execute_reply":"2023-04-23T06:30:17.743443Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<br>","metadata":{}},{"cell_type":"markdown","source":"# To Improve the Model\n\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\n\n* **Optimize the model**\n    * A larger and more diverse training dataset\n    * Experimenting with different architectures and hyperparameters\n    * Choosing an appropriate Loss function, such as Dice Loss\n\n\n* **Modify the model output**\n    * Finding the best threshold\n    * Using an ensemble of models\n    * Postprocessing\n\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!","metadata":{}},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}