{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.14","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":84969,"databundleVersionId":10033515,"sourceType":"competition"},{"sourceId":10083161,"sourceType":"datasetVersion","datasetId":6157138},{"sourceId":146264,"sourceType":"modelInstanceVersion","modelInstanceId":124048,"modelId":147101}],"dockerImageVersionId":30787,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"\"Om Namah Shivaya!! 🙏🙏\"","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2024-12-07T17:45:59.196852Z","iopub.execute_input":"2024-12-07T17:45:59.197688Z","iopub.status.idle":"2024-12-07T17:45:59.207518Z","shell.execute_reply.started":"2024-12-07T17:45:59.197644Z","shell.execute_reply":"2024-12-07T17:45:59.206882Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install --quiet git+https://github.com/facebookresearch/segment-anything-2/","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-07T17:45:59.213122Z","iopub.execute_input":"2024-12-07T17:45:59.213376Z","iopub.status.idle":"2024-12-07T17:48:20.059989Z","shell.execute_reply.started":"2024-12-07T17:45:59.213351Z","shell.execute_reply":"2024-12-07T17:48:20.058730Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Environment Set-up","metadata":{}},{"cell_type":"code","source":"from pathlib import Path\n\nimport numpy as np\nfrom PIL import Image\nimport pandas as pd\nimport pydicom\nimport matplotlib.pyplot as plt\n\nimport math\nimport torch","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-07T18:25:35.321151Z","iopub.execute_input":"2024-12-07T18:25:35.321477Z","iopub.status.idle":"2024-12-07T18:25:35.325901Z","shell.execute_reply.started":"2024-12-07T18:25:35.321447Z","shell.execute_reply":"2024-12-07T18:25:35.324965Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# select the device for computation\nif torch.cuda.is_available():\n    device = torch.device(\"cuda\")\n\nif device.type == \"cuda\":\n    # use bfloat16 for the entire notebook\n    if False:\n        torch.autocast(\"cuda\", dtype=torch.bfloat16).__enter__()\n    # turn on tfloat32 for Ampere GPUs (https://pytorch.org/docs/stable/notes/cuda.html#tensorfloat-32-tf32-on-ampere-devices)\n    if torch.cuda.get_device_properties(0).major >= 8:\n        torch.backends.cuda.matmul.allow_tf32 = True\n        torch.backends.cudnn.allow_tf32 = True","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-07T17:48:23.642121Z","iopub.execute_input":"2024-12-07T17:48:23.642502Z","iopub.status.idle":"2024-12-07T17:48:23.697005Z","shell.execute_reply.started":"2024-12-07T17:48:23.642471Z","shell.execute_reply":"2024-12-07T17:48:23.696148Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Load Dataset","metadata":{}},{"cell_type":"code","source":"DATA_DIR = Path(\"/kaggle/input/czii-cryoet-630x630-png-dataset-816bit\")\nIMG_DIR = DATA_DIR / \"train_png_normalized_8bit\"\nBBOX_PATH = DATA_DIR / \"train_bounding_boxes.csv\"\n\nMASKS_DIR = Path(\"/kaggle/working/generated-masks-png\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-07T17:48:23.699385Z","iopub.execute_input":"2024-12-07T17:48:23.700007Z","iopub.status.idle":"2024-12-07T17:48:23.703868Z","shell.execute_reply.started":"2024-12-07T17:48:23.699978Z","shell.execute_reply":"2024-12-07T17:48:23.702933Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df = pd.read_csv(BBOX_PATH)\ndf = df.set_index([\"exp_name\", \"frame\"])[[\"x_center\", \"y_center\", \"height\", \"width\", \"label\"]]\ndf = df.sort_index()\n\ndf.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-07T17:48:23.704959Z","iopub.execute_input":"2024-12-07T17:48:23.705283Z","iopub.status.idle":"2024-12-07T17:48:23.824294Z","shell.execute_reply.started":"2024-12-07T17:48:23.705251Z","shell.execute_reply":"2024-12-07T17:48:23.823575Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## End-to-end batched inference","metadata":{}},{"cell_type":"code","source":"from sam2.sam2_image_predictor import SAM2ImagePredictor\n\npredictor = SAM2ImagePredictor.from_pretrained(\"facebook/sam2.1-hiera-large\", device=device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-07T17:48:23.825149Z","iopub.execute_input":"2024-12-07T17:48:23.825365Z","iopub.status.idle":"2024-12-07T17:48:52.089421Z","shell.execute_reply.started":"2024-12-07T17:48:23.825342Z","shell.execute_reply":"2024-12-07T17:48:52.088587Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_img_path(exp_name, frame, IMG_DIR, type):\n    img_path = IMG_DIR / exp_name / type / f\"{frame:03d}.png\"\n    return img_path\n\ndef load_img(img_path):\n    image = Image.open(img_path)\n    image = np.array(image.convert(\"RGB\"), dtype=np.float32)\n    image /= 255.\n    return image\n    \n\n# https://github.com/facebookresearch/syegment-anything-2/blob/main/notebooks/image_predictor_example.ipynb\ndef xyhw_to_xyxy(box, H=630, W=630):\n    xmin = ((box[:, 0] - box[:, 2]/2) * H).astype(int)\n    ymin = ((box[:, 1] - box[:, 3]/2) * W).astype(int)\n    xmax = ((box[:, 0] + box[:, 2]/2) * H).astype(int)\n    ymax = ((box[:, 1] + box[:, 3]/2) * W).astype(int)\n    return np.stack([xmin, ymin, xmax, ymax], axis=-1)\n    \ndef show_mask(mask, ax):\n    color = np.array([30/255, 144/255, 255/255, 0.6])\n    h, w = mask.shape[-2:]\n    mask = mask.astype(np.uint8)\n    mask_image =  mask.reshape(h, w, 1) * color.reshape(1, 1, -1)\n    ax.imshow(mask_image)\n\ndef show_box(box, ax):\n    x0, y0 = box[0], box[1]\n    w, h = box[2] - box[0], box[3] - box[1]\n    ax.add_patch(plt.Rectangle((x0, y0), w, h, edgecolor='green', facecolor=(0, 0, 0, 0), lw=2))   \n\n# Utils to get and save mask\ndef get_mask_to_save(masks, labels):\n    h, w = masks.shape[-2:]\n    output_mask = np.zeros((h,w), dtype=np.uint8)\n    # loop over all labels and merge them\n    for label in np.unique(labels):\n        mask = np.any(masks[labels == label, ...], axis=0)[0]\n        output_mask[mask] = label\n\n    return output_mask\n\ndef save_mask(mask, mask_path):\n    mask_path.parent.mkdir(parents=True, exist_ok=True)\n    mask = Image.fromarray(mask)\n    mask.save(mask_path)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-07T17:48:52.090506Z","iopub.execute_input":"2024-12-07T17:48:52.090920Z","iopub.status.idle":"2024-12-07T17:48:52.100665Z","shell.execute_reply.started":"2024-12-07T17:48:52.090891Z","shell.execute_reply":"2024-12-07T17:48:52.099857Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"generated_masks = {}\n\nfor (exp_name, frame), chunk in df.groupby(level=[0, 1]):\n    \n    img_path = get_img_path(exp_name, frame, IMG_DIR, type=\"images\")\n    img = load_img(img_path)\n    H, W, _ = img.shape\n    \n    input_boxes = chunk.values[:, :4]\n    input_boxes = xyhw_to_xyxy(input_boxes, H, W)\n\n    # SAM Set image and predict\n    predictor.set_image(img)\n    \n    masks, scores, _ = predictor.predict(\n        box=input_boxes,\n        multimask_output=False,\n    )\n\n    # Store all masks in generated_masks\n    generated_masks[exp_name] = generated_masks.get(exp_name, {})\n    generated_masks[exp_name][frame] = masks.astype(np.uint8)\n\n    # Save masks\n    labels = chunk.values[:, 4]\n    output_mask = get_mask_to_save(masks, labels) \n\n    mask_path = get_img_path(exp_name, frame, MASKS_DIR, type=\"masks\")\n    save_mask(output_mask, mask_path)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-07T18:30:01.073012Z","iopub.execute_input":"2024-12-07T18:30:01.073654Z","iopub.status.idle":"2024-12-07T18:38:27.120332Z","shell.execute_reply.started":"2024-12-07T18:30:01.073619Z","shell.execute_reply":"2024-12-07T18:38:27.119598Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Score Generated Masks","metadata":{}},{"cell_type":"code","source":"from sklearn.metrics import mean_absolute_error, mean_squared_error\n\nmse_all = {}\nmae_all = {}\nfor (exp_name, frame), chunk in df.groupby(level=[0, 1]):\n\n    # get masks and labels\n    masks = generated_masks[exp_name][frame]\n    bboxes = chunk.values[:, :4]\n    labels = chunk.values[:, 4]\n\n    xc_preds, yc_preds = [], []\n    for i, mask in enumerate(masks):\n        if not mask.any():\n            xc_preds.append(0)\n            yc_preds.append(0)\n            continue\n\n        # Calculate xc, yc of the preds\n        rows, cols = np.where(mask.squeeze() != 0)\n        xmin, xmax = np.min(cols), np.max(cols)\n        ymin, ymax = np.min(rows), np.max(rows)\n        \n        xc_pred, yc_pred = (xmin + xmax) / 2, (ymin + ymax) / 2\n        xc_true, yc_true = bboxes[i][0] * H, bboxes[i][1] * W\n        h, w = bboxes[i][2] * H, bboxes[i][3] * W\n\n        xc_preds.append(xc_pred)\n        yc_preds.append(yc_pred)\n\n    # Calculate regression metrics\n    for label in np.unique((labels)):\n        mse = mean_squared_error(bboxes[:, 0][labels == label] * H, np.array(xc_preds)[labels == label]) \\\n                + mean_squared_error(bboxes[:, 1][labels == label] * W, np.array(yc_preds)[labels == label])\n        mae = mean_absolute_error(bboxes[:, 0][labels == label] * H, np.array(xc_preds)[labels == label]) \\\n                + mean_absolute_error(bboxes[:, 1][labels == label] * W, np.array(yc_preds)[labels == label])\n\n        # Append label wise scores\n        mse_all[label] = mse_all.get(label, []) + [mse]\n        mae_all[label] = mae_all.get(label, []) + [mae]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-07T18:53:56.541573Z","iopub.execute_input":"2024-12-07T18:53:56.542298Z","iopub.status.idle":"2024-12-07T18:54:27.400827Z","shell.execute_reply.started":"2024-12-07T18:53:56.542260Z","shell.execute_reply":"2024-12-07T18:54:27.400056Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"#\"*4, \"Observed errors between the target and generated Coordinates.\")\n\npd.DataFrame(\n    {\n        \"label\": mse_all.keys(),\n        \"mae\": [sum(e) / len(e) / 640 for _, e in mae_all.items()],\n        \"mse\": [sum(e) / len(e) / 640**2 for _, e in mse_all.items()],\n    },\n).sort_values('label').set_index('label')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-07T19:26:00.406974Z","iopub.execute_input":"2024-12-07T19:26:00.407390Z","iopub.status.idle":"2024-12-07T19:26:00.424026Z","shell.execute_reply.started":"2024-12-07T19:26:00.407354Z","shell.execute_reply":"2024-12-07T19:26:00.422750Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Visualizations","metadata":{}},{"cell_type":"code","source":"EXP_NAME = \"TS_5_4\"\nEXP_MASKS = generated_masks[EXP_NAME]\n\n# Rnadomly show 6 frames\nframes = np.random.choice(list(EXP_MASKS.keys()), size=6)\nfor j in range(3): \n    # Create a plt figure\n    plt.figure(figsize=(15, 10))\n    for i in range(2):\n        plt.subplot(1, 2, i+1)\n        frame = frames[j*2 + i]\n        plt.title(f'Exp: {EXP_NAME} / {frame}')\n        # Show image\n        img = load_img(get_img_path(EXP_NAME, frame, IMG_DIR, type=\"images\"))\n        plt.imshow(img)\n        # Show mask\n        masks = EXP_MASKS[frame]\n        for mask in masks:\n            show_mask(mask.squeeze(), plt.gca())\n        # Show bbox\n        boxes = xyhw_to_xyxy(df.loc[(EXP_NAME, frame)].values[:, :4])\n        for box in boxes:\n            show_box(box, plt.gca())\n\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-07T17:53:48.316650Z","iopub.status.idle":"2024-12-07T17:53:48.316968Z","shell.execute_reply.started":"2024-12-07T17:53:48.316823Z","shell.execute_reply":"2024-12-07T17:53:48.316839Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"EXP_NAME = \"TS_69_2\"\nEXP_MASKS = generated_masks[EXP_NAME]\n\n# Rnadomly show 6 frames\nframes = np.random.choice(list(EXP_MASKS.keys()), size=6)\nfor j in range(3): \n    # Create a plt figure\n    plt.figure(figsize=(15, 10))\n    for i in range(2):\n        plt.subplot(1, 2, i+1)\n        frame = frames[j*2 + i]\n        plt.title(f'Exp: {EXP_NAME} / {frame}')\n        # Show image\n        img = load_img(get_img_path(EXP_NAME, frame, IMG_DIR, type=\"images\"))\n        plt.imshow(img)\n        # Show mask\n        masks = EXP_MASKS[frame]\n        for mask in masks:\n            show_mask(mask.squeeze(), plt.gca())\n        # Show bbox\n        boxes = xyhw_to_xyxy(df.loc[(EXP_NAME, frame)].values[:, :4])\n        for box in boxes:\n            show_box(box, plt.gca())\n\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-07T17:53:48.318650Z","iopub.status.idle":"2024-12-07T17:53:48.319119Z","shell.execute_reply.started":"2024-12-07T17:53:48.318884Z","shell.execute_reply":"2024-12-07T17:53:48.318910Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"EXP_NAME = \"TS_6_4\"\nEXP_MASKS = generated_masks[EXP_NAME]\n\n# Rnadomly show 6 frames\nframes = np.random.choice(list(EXP_MASKS.keys()), size=6)\nfor j in range(3): \n    # Create a plt figure\n    plt.figure(figsize=(15, 10))\n    for i in range(2):\n        plt.subplot(1, 2, i+1)\n        frame = frames[j*2 + i]\n        plt.title(f'Exp: {EXP_NAME} / {frame}')\n        # Show image\n        img = load_img(get_img_path(EXP_NAME, frame, IMG_DIR, type=\"images\"))\n        plt.imshow(img)\n        # Show mask\n        masks = EXP_MASKS[frame]\n        for mask in masks:\n            show_mask(mask.squeeze(), plt.gca())\n        # Show bbox\n        boxes = xyhw_to_xyxy(df.loc[(EXP_NAME, frame)].values[:, :4])\n        for box in boxes:\n            show_box(box, plt.gca())\n\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-07T17:53:48.320712Z","iopub.status.idle":"2024-12-07T17:53:48.321169Z","shell.execute_reply.started":"2024-12-07T17:53:48.320941Z","shell.execute_reply":"2024-12-07T17:53:48.320966Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"EXP_NAME = \"TS_6_6\"\nEXP_MASKS = generated_masks[EXP_NAME]\n\n# Rnadomly show 6 frames\nframes = np.random.choice(list(EXP_MASKS.keys()), size=6)\nfor j in range(3): \n    # Create a plt figure\n    plt.figure(figsize=(15, 10))\n    for i in range(2):\n        plt.subplot(1, 2, i+1)\n        frame = frames[j*2 + i]\n        plt.title(f'Exp: {EXP_NAME} / {frame}')\n        # Show image\n        img = load_img(get_img_path(EXP_NAME, frame, IMG_DIR, type=\"images\"))\n        plt.imshow(img)\n        # Show mask\n        masks = EXP_MASKS[frame]\n        for mask in masks:\n            show_mask(mask.squeeze(), plt.gca())\n        # Show bbox\n        boxes = xyhw_to_xyxy(df.loc[(EXP_NAME, frame)].values[:, :4])\n        for box in boxes:\n            show_box(box, plt.gca())\n\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-07T17:53:48.322413Z","iopub.status.idle":"2024-12-07T17:53:48.323440Z","shell.execute_reply.started":"2024-12-07T17:53:48.323181Z","shell.execute_reply":"2024-12-07T17:53:48.323212Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"EXP_NAME = \"TS_86_3\"\nEXP_MASKS = generated_masks[EXP_NAME]\n\n# Rnadomly show 6 frames\nframes = np.random.choice(list(EXP_MASKS.keys()), size=6)\nfor j in range(3): \n    # Create a plt figure\n    plt.figure(figsize=(15, 10))\n    for i in range(2):\n        plt.subplot(1, 2, i+1)\n        frame = frames[j*2 + i]\n        plt.title(f'Exp: {EXP_NAME} / {frame}')\n        # Show image\n        img = load_img(get_img_path(EXP_NAME, frame, IMG_DIR, type=\"images\"))\n        plt.imshow(img)\n        # Show mask\n        masks = EXP_MASKS[frame]\n        for mask in masks:\n            show_mask(mask.squeeze(), plt.gca())\n        # Show bbox\n        boxes = xyhw_to_xyxy(df.loc[(EXP_NAME, frame)].values[:, :4])\n        for box in boxes:\n            show_box(box, plt.gca())\n\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-07T17:53:48.325018Z","iopub.status.idle":"2024-12-07T17:53:48.325493Z","shell.execute_reply.started":"2024-12-07T17:53:48.325257Z","shell.execute_reply":"2024-12-07T17:53:48.325281Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"EXP_NAME = \"TS_99_9\"\nEXP_MASKS = generated_masks[EXP_NAME]\n\n# Rnadomly show 6 frames\nframes = np.random.choice(list(EXP_MASKS.keys()), size=6)\nfor j in range(3): \n    # Create a plt figure\n    plt.figure(figsize=(15, 10))\n    for i in range(2):\n        plt.subplot(1, 2, i+1)\n        frame = frames[j*2 + i]\n        plt.title(f'Exp: {EXP_NAME} / {frame}')\n        # Show image\n        img = load_img(get_img_path(EXP_NAME, frame, IMG_DIR, type=\"images\"))\n        plt.imshow(img)\n        # Show mask\n        masks = EXP_MASKS[frame]\n        for mask in masks:\n            show_mask(mask.squeeze(), plt.gca())\n        # Show bbox\n        boxes = xyhw_to_xyxy(df.loc[(EXP_NAME, frame)].values[:, :4])\n        for box in boxes:\n            show_box(box, plt.gca())\n\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-07T17:53:48.328257Z","iopub.status.idle":"2024-12-07T17:53:48.328576Z","shell.execute_reply.started":"2024-12-07T17:53:48.328424Z","shell.execute_reply":"2024-12-07T17:53:48.328441Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}