{"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":"This notebook implements semantic segmentation with a simple toy network. To goal is to explore what a simple network with very few layers can achieve.\n\nI tried to keep the code clean. If you see ways to improve, please let me know in the comments.\n\nCurrently, there is no proper validation split; I will add that later.","metadata":{}},{"cell_type":"code","source":"from pathlib import Path\nimport pickle\nimport cv2\nimport numpy as np\nimport torch\nfrom torch import nn\nimport matplotlib.pyplot as plt\nimport matplotlib as mpl\nfrom tqdm.auto import tqdm\nfrom statistics import mean, stdev\nimport pandas as pd\nimport itertools","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-05-14T06:00:42.954061Z","iopub.execute_input":"2022-05-14T06:00:42.955038Z","iopub.status.idle":"2022-05-14T06:00:42.960738Z","shell.execute_reply.started":"2022-05-14T06:00:42.954978Z","shell.execute_reply":"2022-05-14T06:00:42.959816Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data_dir = Path(\"/kaggle/input/sartorius-cell-instance-segmentation\")\nmask_dir = Path(\"/kaggle/input/cell-image-masks/train_masks\")   # my public dataset","metadata":{"execution":{"iopub.status.busy":"2022-05-14T06:00:42.980012Z","iopub.execute_input":"2022-05-14T06:00:42.98068Z","iopub.status.idle":"2022-05-14T06:00:42.985852Z","shell.execute_reply.started":"2022-05-14T06:00:42.980626Z","shell.execute_reply":"2022-05-14T06:00:42.9848Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Helper code","metadata":{}},{"cell_type":"code","source":"class MyImageDataset(torch.utils.data.Dataset):\n    def __init__(self, filenames, mask_dir):\n        self.filenames = filenames\n        self.mask_dir = mask_dir\n\n    def __getitem__(self, index):\n        filename = Path(self.filenames[index])\n        mask_filename = self.mask_dir.joinpath(f\"mask_{filename.stem}.pkl\")\n\n        if not filename.exists():\n            raise ValueError(f\"Image {filename} does not exists\")\n\n        if not mask_filename.exists():\n            raise ValueError(f\"Mask {mask_filename} does not exists\")\n\n        img_data = load_img(filename)\n\n        with mask_filename.open(\"rb\") as f:\n            mask_data = pickle.load(f)\n\n        img_tensor = torch.tensor(img_data).unsqueeze(0)\n        mask_tensor = torch.tensor(mask_data.astype(np.float32)).unsqueeze(0)\n\n        return img_tensor, mask_tensor\n\n    def __len__(self):\n        return len(self.filenames)\n    \n    \ndef load_img(filename):\n    \"\"\"returns x-125 as float32\"\"\"\n    if not filename.exists():\n        raise IOError(f\"File {filename} not found\")\n\n    img_data = cv2.imread(str(filename))\n    assert img_data is not None\n\n    return cv2.cvtColor(img_data, cv2.COLOR_BGR2GRAY).astype(np.float32) - 125","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-05-14T06:00:42.999312Z","iopub.execute_input":"2022-05-14T06:00:42.999965Z","iopub.status.idle":"2022-05-14T06:00:43.011098Z","shell.execute_reply.started":"2022-05-14T06:00:42.999934Z","shell.execute_reply":"2022-05-14T06:00:43.009903Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Load data","metadata":{}},{"cell_type":"code","source":"batch_size = 32\n\nimage_filenames = list(data_dir.joinpath(\"train\").glob(\"*.png\"))\n\ntrain_dataset = MyImageDataset(image_filenames, mask_dir)\n\ntrain_dataloader = torch.utils.data.DataLoader(\n    train_dataset,\n    shuffle=True,\n    batch_size=batch_size,\n)","metadata":{"execution":{"iopub.status.busy":"2022-05-14T06:00:43.017219Z","iopub.execute_input":"2022-05-14T06:00:43.017857Z","iopub.status.idle":"2022-05-14T06:00:43.030061Z","shell.execute_reply.started":"2022-05-14T06:00:43.017816Z","shell.execute_reply":"2022-05-14T06:00:43.029216Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Train simple network\n\nKeep W/H the same across all layers and do simple log-loss against full semantic segmentation mask.","metadata":{}},{"cell_type":"code","source":"n_epochs = 50\n\nNN = nn.Sequential(\n    nn.Conv2d(1, 20, kernel_size=5, padding=\"same\"),\n    nn.BatchNorm2d(20),\n    nn.ReLU(),\n    nn.Conv2d(20, 10, kernel_size=1),\n    \n    nn.Conv2d(10, 10, kernel_size=5, padding=\"same\"),\n    nn.BatchNorm2d(10),\n    nn.ReLU(),\n    nn.Conv2d(10, 1, kernel_size=1),\n    \n).to(\"cuda\")\n\n##########################################################\nlossfunc = nn.BCEWithLogitsLoss()\n\nlosses = []\n\noptimizer = torch.optim.Adam(NN.parameters())\n\nfor epoch in tqdm(range(1, n_epochs+1)):\n    for i, (X, Y) in enumerate(train_dataloader):\n        torch.cuda.empty_cache()\n        \n        X = X.to(\"cuda\")\n        Y = Y.to(\"cuda\")\n        \n        optimizer.zero_grad()\n        \n        pred=NN(X)\n        \n        loss = lossfunc(pred, Y)\n        \n        losses.append(float(loss))\n\n        loss.backward()\n        optimizer.step()\n\n    print(f\"{epoch:3}/{n_epochs}: {mean(losses[-10:]):.4g}\")\n    \ntorch.cuda.empty_cache()","metadata":{"execution":{"iopub.status.busy":"2022-05-14T06:00:43.046466Z","iopub.execute_input":"2022-05-14T06:00:43.046699Z","iopub.status.idle":"2022-05-14T06:16:32.416805Z","shell.execute_reply.started":"2022-05-14T06:00:43.046666Z","shell.execute_reply":"2022-05-14T06:16:32.415807Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"pd.Series(losses).rolling(21).mean().plot(title=\"loss\");","metadata":{"execution":{"iopub.status.busy":"2022-05-14T06:16:32.419064Z","iopub.execute_input":"2022-05-14T06:16:32.41962Z","iopub.status.idle":"2022-05-14T06:16:32.717594Z","shell.execute_reply.started":"2022-05-14T06:16:32.419577Z","shell.execute_reply":"2022-05-14T06:16:32.716656Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Find best prediction mask threshold optimizing IoU score","metadata":{}},{"cell_type":"code","source":"# unoptimized and slow; any way to speed up?\n\ndef get_threshold(Y, pred):\n    scores = list(pred.ravel())\n    mask = list(Y.ravel())\n    \n    idxs=np.argsort(scores)[::-1]\n    mask_sorted=np.array(mask)[idxs]\n    sum_mask_one=np.cumsum(mask_sorted)\n    IoU=sum_mask_one/(np.arange(1,len(mask_sorted)+1)+np.sum(mask_sorted)-sum_mask_one)\n    best_IoU_idx=IoU.argmax()\n    best_threshold=scores[idxs[best_IoU_idx]]\n    best_IoU=IoU[best_IoU_idx]\n\n    return best_threshold, best_IoU\n    \nimg_thresholds = []         # one for each image\nimg_IoUs = []\n\nN=3\nfor X, Y in tqdm(itertools.islice(train_dataloader, N), total=N):\n    X = X.to(\"cuda\")\n    Y = Y.detach().numpy()\n\n    with torch.no_grad():\n        pred=torch.sigmoid(NN(X)).cpu().detach().numpy()\n\n    for i in range(Y.shape[0]):\n        best_img_threshold, best_img_IoU = get_threshold(Y[i], pred[i])\n        img_thresholds.append(best_img_threshold)\n        img_IoUs.append(best_img_IoU)\n    \nbest_threshold = np.mean(img_thresholds)\nbest_threshold_spread = np.std(img_thresholds)\navg_IoU = mean(img_IoUs)\n\nprint(f\"Best threshold: {best_threshold:.3g} (+-{best_threshold_spread:.3g}), Avg. Train IoU: {avg_IoU:.3f}\")","metadata":{"execution":{"iopub.status.busy":"2022-05-14T06:16:32.719587Z","iopub.execute_input":"2022-05-14T06:16:32.720278Z","iopub.status.idle":"2022-05-14T06:17:05.962301Z","shell.execute_reply.started":"2022-05-14T06:16:32.72018Z","shell.execute_reply":"2022-05-14T06:17:05.961514Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Visualize predictions","metadata":{}},{"cell_type":"code","source":"threshold = best_threshold\n\n###########################\nX, Y = next(iter(train_dataloader))\nX = X.to(\"cuda\")\nY = Y.detach().numpy()\n\nwith torch.no_grad():\n    pred=torch.sigmoid(NN(X)).cpu().detach().numpy()\n    \npred_Y = (pred >= threshold)\n    \ncmap = mpl.colors.ListedColormap(['black', 'gray', 'orange', 'green'])\n\ndef plot(img_Y, img_pred):\n    output = np.zeros_like(img_Y)\n    output = np.where((img_Y == 0) & (img_pred == 1), 1, output)\n    output = np.where((img_Y == 1) & (img_pred == 0), 2, output)\n    output = np.where((img_Y == 1) & (img_pred == 1), 3, output)\n\n    plt.figure(figsize=(10,10))\n    plt.imshow(output, cmap=cmap)\n    plt.xticks([])\n    plt.yticks([]);\n    \n\nN = 5\nfor i in range(N):\n    img_Y = Y[i, 0]\n    img_pred = pred_Y[i, 0]\n    \n    plot(img_Y, img_pred)\n    plt.show()\n\n# green: correct prediction\n# gray: false positive (too much)\n# orange: false negative (missed)","metadata":{"execution":{"iopub.status.busy":"2022-05-14T06:17:05.964277Z","iopub.execute_input":"2022-05-14T06:17:05.965124Z","iopub.status.idle":"2022-05-14T06:17:07.183518Z","shell.execute_reply.started":"2022-05-14T06:17:05.965079Z","shell.execute_reply":"2022-05-14T06:17:07.18257Z"},"trusted":true},"execution_count":null,"outputs":[]}]}