{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":30201,"databundleVersionId":2750748,"sourceType":"competition"}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"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":"2024-04-16T12:09:44.616486Z","iopub.execute_input":"2024-04-16T12:09:44.616832Z","iopub.status.idle":"2024-04-16T12:09:53.863311Z","shell.execute_reply.started":"2024-04-16T12:09:44.616803Z","shell.execute_reply":"2024-04-16T12:09:53.862476Z"},"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":"2024-04-16T12:09:53.865559Z","iopub.execute_input":"2024-04-16T12:09:53.866192Z","iopub.status.idle":"2024-04-16T12:09:53.898892Z","shell.execute_reply.started":"2024-04-16T12:09:53.866157Z","shell.execute_reply":"2024-04-16T12:09:53.897992Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":"2024-04-16T12:09:53.900127Z","iopub.execute_input":"2024-04-16T12:09:53.900426Z","iopub.status.idle":"2024-04-16T12:09:53.914808Z","shell.execute_reply.started":"2024-04-16T12:09:53.900403Z","shell.execute_reply":"2024-04-16T12:09:53.914Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":"2024-04-16T12:09:53.916961Z","iopub.execute_input":"2024-04-16T12:09:53.917311Z","iopub.status.idle":"2024-04-16T12:09:53.925507Z","shell.execute_reply.started":"2024-04-16T12:09:53.91728Z","shell.execute_reply":"2024-04-16T12:09:53.924586Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.cuda.get_device_properties(DEVICE)","metadata":{"execution":{"iopub.status.busy":"2024-04-16T12:09:53.928279Z","iopub.execute_input":"2024-04-16T12:09:53.928517Z","iopub.status.idle":"2024-04-16T12:09:53.963637Z","shell.execute_reply.started":"2024-04-16T12:09:53.928496Z","shell.execute_reply":"2024-04-16T12:09:53.962688Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":"2024-04-16T12:09:53.964895Z","iopub.execute_input":"2024-04-16T12:09:53.965464Z","iopub.status.idle":"2024-04-16T12:09:53.973122Z","shell.execute_reply.started":"2024-04-16T12:09:53.96543Z","shell.execute_reply":"2024-04-16T12:09:53.972167Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":"2024-04-16T12:09:53.974705Z","iopub.execute_input":"2024-04-16T12:09:53.975084Z","iopub.status.idle":"2024-04-16T12:09:53.985411Z","shell.execute_reply.started":"2024-04-16T12:09:53.97505Z","shell.execute_reply":"2024-04-16T12:09:53.984417Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":"2024-04-16T12:09:53.98664Z","iopub.execute_input":"2024-04-16T12:09:53.987039Z","iopub.status.idle":"2024-04-16T12:09:53.997996Z","shell.execute_reply.started":"2024-04-16T12:09:53.987004Z","shell.execute_reply":"2024-04-16T12:09:53.997072Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":"2024-04-16T12:09:53.999435Z","iopub.execute_input":"2024-04-16T12:09:53.999922Z","iopub.status.idle":"2024-04-16T12:09:54.00803Z","shell.execute_reply.started":"2024-04-16T12:09:53.999888Z","shell.execute_reply":"2024-04-16T12:09:54.007109Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":"2024-04-16T12:09:54.009327Z","iopub.execute_input":"2024-04-16T12:09:54.009636Z","iopub.status.idle":"2024-04-16T12:09:54.019601Z","shell.execute_reply.started":"2024-04-16T12:09:54.0096Z","shell.execute_reply":"2024-04-16T12:09:54.018493Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_train = pd.read_csv(TRAIN_CSV)\ndf_train.head()._append(df_train.tail())","metadata":{"execution":{"iopub.status.busy":"2024-04-16T12:09:54.020719Z","iopub.execute_input":"2024-04-16T12:09:54.02101Z","iopub.status.idle":"2024-04-16T12:09:54.692918Z","shell.execute_reply.started":"2024-04-16T12:09:54.020986Z","shell.execute_reply":"2024-04-16T12:09:54.691962Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_train.shape","metadata":{"execution":{"iopub.status.busy":"2024-04-16T12:09:54.694352Z","iopub.execute_input":"2024-04-16T12:09:54.694799Z","iopub.status.idle":"2024-04-16T12:09:54.701266Z","shell.execute_reply.started":"2024-04-16T12:09:54.694765Z","shell.execute_reply":"2024-04-16T12:09:54.700261Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_train.dtypes","metadata":{"execution":{"iopub.status.busy":"2024-04-16T12:09:54.70246Z","iopub.execute_input":"2024-04-16T12:09:54.702705Z","iopub.status.idle":"2024-04-16T12:09:54.711455Z","shell.execute_reply.started":"2024-04-16T12:09:54.702683Z","shell.execute_reply":"2024-04-16T12:09:54.710508Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_train.cell_type.value_counts()","metadata":{"execution":{"iopub.status.busy":"2024-04-16T12:09:54.715437Z","iopub.execute_input":"2024-04-16T12:09:54.715833Z","iopub.status.idle":"2024-04-16T12:09:54.743135Z","shell.execute_reply.started":"2024-04-16T12:09:54.715791Z","shell.execute_reply":"2024-04-16T12:09:54.74191Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_train.cell_type.unique()","metadata":{"execution":{"iopub.status.busy":"2024-04-16T12:09:54.746298Z","iopub.execute_input":"2024-04-16T12:09:54.74713Z","iopub.status.idle":"2024-04-16T12:09:54.761659Z","shell.execute_reply.started":"2024-04-16T12:09:54.747102Z","shell.execute_reply":"2024-04-16T12:09:54.760647Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":"2024-04-16T12:09:54.762979Z","iopub.execute_input":"2024-04-16T12:09:54.763276Z","iopub.status.idle":"2024-04-16T12:09:54.808521Z","shell.execute_reply.started":"2024-04-16T12:09:54.763249Z","shell.execute_reply":"2024-04-16T12:09:54.80759Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_instances.shape","metadata":{"execution":{"iopub.status.busy":"2024-04-16T12:09:54.809763Z","iopub.execute_input":"2024-04-16T12:09:54.8101Z","iopub.status.idle":"2024-04-16T12:09:54.816144Z","shell.execute_reply.started":"2024-04-16T12:09:54.810071Z","shell.execute_reply":"2024-04-16T12:09:54.815115Z"},"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":"2024-04-16T12:09:54.817617Z","iopub.execute_input":"2024-04-16T12:09:54.818057Z","iopub.status.idle":"2024-04-16T12:09:54.886975Z","shell.execute_reply.started":"2024-04-16T12:09:54.818024Z","shell.execute_reply":"2024-04-16T12:09:54.886017Z"},"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":"2024-04-16T12:09:54.888149Z","iopub.execute_input":"2024-04-16T12:09:54.888439Z","iopub.status.idle":"2024-04-16T12:09:54.89499Z","shell.execute_reply.started":"2024-04-16T12:09:54.888415Z","shell.execute_reply":"2024-04-16T12:09:54.893899Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":"2024-04-16T12:09:54.896143Z","iopub.execute_input":"2024-04-16T12:09:54.896487Z","iopub.status.idle":"2024-04-16T12:09:56.572325Z","shell.execute_reply.started":"2024-04-16T12:09:54.896429Z","shell.execute_reply":"2024-04-16T12:09:56.571259Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_train.width.unique()","metadata":{"execution":{"iopub.status.busy":"2024-04-16T12:09:56.573506Z","iopub.execute_input":"2024-04-16T12:09:56.573803Z","iopub.status.idle":"2024-04-16T12:09:56.58131Z","shell.execute_reply.started":"2024-04-16T12:09:56.573777Z","shell.execute_reply":"2024-04-16T12:09:56.580262Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_train.height.unique()","metadata":{"execution":{"iopub.status.busy":"2024-04-16T12:09:56.582863Z","iopub.execute_input":"2024-04-16T12:09:56.58332Z","iopub.status.idle":"2024-04-16T12:09:56.597154Z","shell.execute_reply.started":"2024-04-16T12:09:56.583282Z","shell.execute_reply":"2024-04-16T12:09:56.596267Z"},"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":"2024-04-16T12:09:56.598584Z","iopub.execute_input":"2024-04-16T12:09:56.599028Z","iopub.status.idle":"2024-04-16T12:10:06.872145Z","shell.execute_reply.started":"2024-04-16T12:09:56.598992Z","shell.execute_reply":"2024-04-16T12:10:06.871185Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":"2024-04-16T12:10:06.873535Z","iopub.execute_input":"2024-04-16T12:10:06.873844Z","iopub.status.idle":"2024-04-16T12:10:15.498266Z","shell.execute_reply.started":"2024-04-16T12:10:06.873818Z","shell.execute_reply":"2024-04-16T12:10:15.497234Z"},"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\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":"2024-04-16T12:10:15.499577Z","iopub.execute_input":"2024-04-16T12:10:15.499985Z","iopub.status.idle":"2024-04-16T12:10:16.345214Z","shell.execute_reply.started":"2024-04-16T12:10:15.499943Z","shell.execute_reply":"2024-04-16T12:10:16.344279Z"},"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":"2024-04-16T12:11:01.369704Z","iopub.execute_input":"2024-04-16T12:11:01.370457Z","iopub.status.idle":"2024-04-16T12:11:02.360718Z","shell.execute_reply.started":"2024-04-16T12:11:01.370427Z","shell.execute_reply":"2024-04-16T12:11:02.359729Z"},"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\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":"2024-04-16T12:11:13.913412Z","iopub.execute_input":"2024-04-16T12:11:13.914094Z","iopub.status.idle":"2024-04-16T12:11:14.680039Z","shell.execute_reply.started":"2024-04-16T12:11:13.914044Z","shell.execute_reply":"2024-04-16T12:11:14.679073Z"},"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":"2024-04-16T12:11:20.709422Z","iopub.execute_input":"2024-04-16T12:11:20.70985Z","iopub.status.idle":"2024-04-16T12:11:21.451496Z","shell.execute_reply.started":"2024-04-16T12:11:20.70982Z","shell.execute_reply":"2024-04-16T12:11:21.450497Z"},"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\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":"2024-04-16T12:11:32.448213Z","iopub.execute_input":"2024-04-16T12:11:32.448867Z","iopub.status.idle":"2024-04-16T12:11:33.322922Z","shell.execute_reply.started":"2024-04-16T12:11:32.448836Z","shell.execute_reply":"2024-04-16T12:11:33.321961Z"},"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":"2024-04-16T12:11:40.278462Z","iopub.execute_input":"2024-04-16T12:11:40.279137Z","iopub.status.idle":"2024-04-16T12:11:41.139885Z","shell.execute_reply.started":"2024-04-16T12:11:40.279105Z","shell.execute_reply":"2024-04-16T12:11:41.1389Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":"2024-04-16T12:11:55.270908Z","iopub.execute_input":"2024-04-16T12:11:55.271622Z","iopub.status.idle":"2024-04-16T12:11:55.358014Z","shell.execute_reply.started":"2024-04-16T12:11:55.271589Z","shell.execute_reply":"2024-04-16T12:11:55.357117Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_train.id.unique().shape","metadata":{"execution":{"iopub.status.busy":"2024-04-16T12:12:05.859738Z","iopub.execute_input":"2024-04-16T12:12:05.860588Z","iopub.status.idle":"2024-04-16T12:12:05.873647Z","shell.execute_reply.started":"2024-04-16T12:12:05.860551Z","shell.execute_reply":"2024-04-16T12:12:05.872559Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":"2024-04-16T12:12:16.462826Z","iopub.execute_input":"2024-04-16T12:12:16.463263Z","iopub.status.idle":"2024-04-16T12:12:16.475387Z","shell.execute_reply.started":"2024-04-16T12:12:16.463205Z","shell.execute_reply":"2024-04-16T12:12:16.474287Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":"2024-04-16T12:12:27.512301Z","iopub.execute_input":"2024-04-16T12:12:27.513061Z","iopub.status.idle":"2024-04-16T12:12:27.520469Z","shell.execute_reply.started":"2024-04-16T12:12:27.513029Z","shell.execute_reply":"2024-04-16T12:12:27.519462Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":"2024-04-16T12:12:34.771908Z","iopub.execute_input":"2024-04-16T12:12:34.772648Z","iopub.status.idle":"2024-04-16T12:12:34.779384Z","shell.execute_reply.started":"2024-04-16T12:12:34.772617Z","shell.execute_reply":"2024-04-16T12:12:34.77841Z"},"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\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)   # 这里为什么会有 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":"2024-04-16T12:12:46.03296Z","iopub.execute_input":"2024-04-16T12:12:46.033737Z","iopub.status.idle":"2024-04-16T12:12:46.046504Z","shell.execute_reply.started":"2024-04-16T12:12:46.033695Z","shell.execute_reply":"2024-04-16T12:12:46.045528Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":"2024-04-16T12:12:53.149001Z","iopub.execute_input":"2024-04-16T12:12:53.14943Z","iopub.status.idle":"2024-04-16T12:12:53.160083Z","shell.execute_reply.started":"2024-04-16T12:12:53.149399Z","shell.execute_reply":"2024-04-16T12:12:53.159062Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":"2024-04-16T12:13:07.758827Z","iopub.execute_input":"2024-04-16T12:13:07.75992Z","iopub.status.idle":"2024-04-16T12:38:36.590836Z","shell.execute_reply.started":"2024-04-16T12:13:07.759883Z","shell.execute_reply":"2024-04-16T12:38:36.589818Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":"2024-04-16T12:39:13.228583Z","iopub.execute_input":"2024-04-16T12:39:13.229426Z","iopub.status.idle":"2024-04-16T13:13:43.36164Z","shell.execute_reply.started":"2024-04-16T12:39:13.229391Z","shell.execute_reply":"2024-04-16T13:13:43.360481Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":"2024-04-16T13:14:16.088598Z","iopub.execute_input":"2024-04-16T13:14:16.08903Z","iopub.status.idle":"2024-04-16T13:48:34.127827Z","shell.execute_reply.started":"2024-04-16T13:14:16.088987Z","shell.execute_reply":"2024-04-16T13:48:34.126675Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":"2024-04-16T13:48:39.47128Z","iopub.execute_input":"2024-04-16T13:48:39.472287Z","iopub.status.idle":"2024-04-16T13:48:40.762929Z","shell.execute_reply.started":"2024-04-16T13:48:39.472248Z","shell.execute_reply":"2024-04-16T13:48:40.761724Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":"2024-04-16T13:49:05.23858Z","iopub.execute_input":"2024-04-16T13:49:05.239319Z","iopub.status.idle":"2024-04-16T13:49:05.244406Z","shell.execute_reply.started":"2024-04-16T13:49:05.239289Z","shell.execute_reply":"2024-04-16T13:49:05.24339Z"},"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":"2024-04-16T13:49:12.334568Z","iopub.execute_input":"2024-04-16T13:49:12.334991Z","iopub.status.idle":"2024-04-16T13:49:13.244246Z","shell.execute_reply.started":"2024-04-16T13:49:12.334961Z","shell.execute_reply":"2024-04-16T13:49:13.243278Z"},"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":"2024-04-16T13:49:18.871881Z","iopub.execute_input":"2024-04-16T13:49:18.872301Z","iopub.status.idle":"2024-04-16T13:49:19.46686Z","shell.execute_reply.started":"2024-04-16T13:49:18.87227Z","shell.execute_reply":"2024-04-16T13:49:19.465559Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":"2024-04-16T13:49:26.446408Z","iopub.execute_input":"2024-04-16T13:49:26.446812Z","iopub.status.idle":"2024-04-16T13:49:27.724261Z","shell.execute_reply.started":"2024-04-16T13:49:26.446783Z","shell.execute_reply":"2024-04-16T13:49:27.723214Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":"2024-04-16T13:49:38.522768Z","iopub.execute_input":"2024-04-16T13:49:38.523179Z","iopub.status.idle":"2024-04-16T13:49:39.233743Z","shell.execute_reply.started":"2024-04-16T13:49:38.523147Z","shell.execute_reply":"2024-04-16T13:49:39.232661Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_submission.to_csv(r'/kaggle/working/Unet_Pytorch_Mine.csv', index=False)","metadata":{"execution":{"iopub.status.busy":"2024-04-16T13:49:46.569609Z","iopub.execute_input":"2024-04-16T13:49:46.570296Z","iopub.status.idle":"2024-04-16T13:49:46.578557Z","shell.execute_reply.started":"2024-04-16T13:49:46.570263Z","shell.execute_reply":"2024-04-16T13:49:46.57753Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}