{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":71549,"databundleVersionId":8561470,"isSourceIdPinned":false,"sourceType":"competition"},{"sourceId":13545566,"sourceType":"datasetVersion","datasetId":8602567}],"dockerImageVersionId":31154,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# Cell 0: install libs (if needed) and imports\n# On Kaggle you may have internet; if not, skip pip and ensure required packages are available.\ntry:\n    import segment_anything as sam\nexcept Exception:\n    # try installing segment-anything (only works if internet is enabled)\n    !pip install -q git+https://github.com/facebookresearch/segment-anything.git\n    import segment_anything as sam\n\n# common libs\nimport os, sys, math, json\nimport numpy as np\nimport pandas as pd\nimport cv2\nimport pydicom\nfrom tqdm import tqdm\nimport matplotlib.pyplot as plt\nfrom pathlib import Path\nimport torch\n\nprint(\"torch:\", torch.__version__)\nprint(\"segment_anything imported\")\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-10-29T15:21:00.493526Z","iopub.execute_input":"2025-10-29T15:21:00.494083Z","iopub.status.idle":"2025-10-29T15:21:08.483052Z","shell.execute_reply.started":"2025-10-29T15:21:00.494059Z","shell.execute_reply":"2025-10-29T15:21:08.482346Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Cell 1: utils for DICOM read + window + mask selection + visualize\nimport numpy as np\nimport pydicom, cv2\nfrom skimage.measure import regionprops, label\n\ndef read_windowed_dcm(path, window_center=None, window_width=None):\n    \"\"\"\n    Read DICOM and apply basic windowing. Returns uint8 0..255 image.\n    \"\"\"\n    ds = pydicom.dcmread(path)\n    arr = ds.pixel_array.astype(np.float32)\n\n    # use DICOM window if present, else fallback to percentiles\n    try:\n        if window_center is None:\n            wc = float(ds.WindowCenter)\n        else:\n            wc = float(window_center)\n        if window_width is None:\n            ww = float(ds.WindowWidth)\n        else:\n            ww = float(window_width)\n    except Exception:\n        wc = np.median(arr)\n        ww = np.percentile(arr, 99) - np.percentile(arr, 1)\n\n    mn = wc - ww/2.0\n    mx = wc + ww/2.0\n    img = np.clip(arr, mn, mx)\n    img = ((img - mn) / max(1e-6, (mx - mn)) * 255.0).astype(np.uint8)\n    return img, ds\n\ndef save_mask_png(mask, out_path):\n    # mask: boolean or 0/1 array\n    mask_u8 = (mask.astype(np.uint8) * 255)\n    cv2.imwrite(out_path, mask_u8)\n\ndef mask_to_bbox(mask):\n    \"\"\"Return x1,y1,x2,y2 for nonzero area of mask. If empty -> None.\"\"\"\n    ys, xs = np.where(mask)\n    if len(xs)==0:\n        return None\n    x1, x2 = xs.min(), xs.max()\n    y1, y2 = ys.min(), ys.max()\n    return int(x1), int(y1), int(x2)+1, int(y2)+1\n\ndef enlarge_bbox(x1,y1,x2,y2, img_w, img_h, pad=0.2):\n    \"\"\"Pad bbox by fraction pad (0.2 => 20%), clamp to image size.\"\"\"\n    w = x2 - x1; h = y2 - y1\n    padx = int(round(w * pad)); pady = int(round(h * pad))\n    nx1 = max(0, x1 - padx)\n    ny1 = max(0, y1 - pady)\n    nx2 = min(img_w, x2 + padx)\n    ny2 = min(img_h, y2 + pady)\n    return nx1, ny1, nx2, ny2\n\ndef mask_centroid(mask):\n    props = regionprops(mask.astype(np.uint8))\n    if not props:\n        return None\n    p = props[0]\n    return (p.centroid[1], p.centroid[0])  # x,y\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-29T15:22:35.787991Z","iopub.execute_input":"2025-10-29T15:22:35.788281Z","iopub.status.idle":"2025-10-29T15:22:35.996004Z","shell.execute_reply.started":"2025-10-29T15:22:35.788257Z","shell.execute_reply":"2025-10-29T15:22:35.995190Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Cell 2: load SAM / MedSAM model\nCHECKPOINT_PATH = \"/kaggle/input/medsam-checkpoint/samcheckpoint.pth\"  # <-- place your checkpoint here\n\nif not os.path.exists(CHECKPOINT_PATH):\n    raise FileNotFoundError(f\"MedSAM/SAM checkpoint not found at {CHECKPOINT_PATH}. \"\n                            \"Upload the .pth file to /kaggle/working/ or change CHECKPOINT_PATH.\")\n\nfrom segment_anything import sam_model_registry, SamAutomaticMaskGenerator, SamPredictor\n\n# choose model_type matching checkpoint, e.g., \"vit_b\", \"vit_l\", \"vit_h\"\nMODEL_TYPE = \"vit_b\"   # change if your checkpoint is different (vit_b, vit_l, vit_h)\n\nprint(\"Loading SAM model - may take a moment...\")\nsam_model = sam_model_registry[MODEL_TYPE](checkpoint=CHECKPOINT_PATH)\ndevice = \"cuda\" if torch.cuda.is_available() else \"cpu\"\nsam_model.to(device)\nprint(\"Model loaded on\", device)\n\n# create automatic mask generator (for full-slice masks)\nmask_generator = SamAutomaticMaskGenerator(sam_model,\n                                           points_per_batch=64,   # tune for speed\n                                           pred_iou_thresh=0.3,\n                                           stability_score_thresh=0.5,\n                                           min_mask_region_area=100)  # smallest mask in px\n\n# create predictor (for point prompt)\npredictor = SamPredictor(sam_model)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-29T15:34:03.958864Z","iopub.execute_input":"2025-10-29T15:34:03.959145Z","iopub.status.idle":"2025-10-29T15:34:08.949748Z","shell.execute_reply.started":"2025-10-29T15:34:03.959124Z","shell.execute_reply":"2025-10-29T15:34:08.948898Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os, io, zipfile, cv2, pydicom\nimport numpy as np, pandas as pd\nfrom tqdm import tqdm\nfrom skimage.measure import regionprops\n\n# ================\n# Paths and setup\n# ================\nLABELS_CSV = \"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_label_coordinates.csv\"\nBASE_DIR = \"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_images\"\nOUT_ZIP_MASKS = \"/kaggle/working/medsam_masks.zip\"\nOUT_ZIP_CROPS = \"/kaggle/working/medsam_crops.zip\"\n\nCROP_SIZE = 224\nBOX_PAD = 0.18\nUSE_POINT_PROMPT = True\nMAX_SAMPLES = None\n\ndf = pd.read_csv(LABELS_CSV).astype({\"study_id\": str, \"series_id\": str})\nrecords = []\n\n# Create writable ZIP files\nzip_masks = zipfile.ZipFile(OUT_ZIP_MASKS, \"w\", compression=zipfile.ZIP_DEFLATED)\nzip_crops = zipfile.ZipFile(OUT_ZIP_CROPS, \"w\", compression=zipfile.ZIP_DEFLATED)\n\n# ================\n# Helper functions\n# ================\ndef read_windowed_dcm(path):\n    ds = pydicom.dcmread(path)\n    img = ds.pixel_array.astype(np.float32)\n    if hasattr(ds, \"RescaleSlope\"): img *= ds.RescaleSlope\n    if hasattr(ds, \"RescaleIntercept\"): img += ds.RescaleIntercept\n    img = np.clip(img, np.percentile(img, 0.5), np.percentile(img, 99.5))\n    img = (img - img.min()) / (img.max() - img.min() + 1e-5)\n    img = (img * 255).astype(np.uint8)\n    return img, ds\n\ndef mask_to_bbox(mask):\n    ys, xs = np.where(mask)\n    if len(xs) == 0: return None\n    return np.min(xs), np.min(ys), np.max(xs), np.max(ys)\n\ndef enlarge_bbox(x1, y1, x2, y2, w, h, pad=0.15):\n    bw, bh = x2 - x1, y2 - y1\n    pad_w, pad_h = int(bw * pad), int(bh * pad)\n    return (\n        max(0, x1 - pad_w),\n        max(0, y1 - pad_h),\n        min(w, x2 + pad_w),\n        min(h, y2 + pad_h),\n    )\n\ndef save_mask_to_zip(mask, fname, zip_handle):\n    ok, buf = cv2.imencode(\".png\", (mask * 255).astype(np.uint8))\n    if ok:\n        zip_handle.writestr(fname, buf.tobytes())\n\ndef save_crop_to_zip(crop, fname, zip_handle):\n    ok, buf = cv2.imencode(\".png\", crop)\n    if ok:\n        zip_handle.writestr(fname, buf.tobytes())\n\n# ================\n# Main loop\n# ================\ncount = 0\nfor idx, row in tqdm(df.iterrows(), total=len(df)):\n    if MAX_SAMPLES and count >= MAX_SAMPLES:\n        break\n\n    study, series, inst = str(row.study_id), str(row.series_id), int(row.instance_number)\n    x, y = float(row.x), float(row.y)\n    dcm_path = os.path.join(BASE_DIR, study, series, f\"{inst}.dcm\")\n    if not os.path.exists(dcm_path):\n        continue\n\n    img, ds = read_windowed_dcm(dcm_path)\n    h, w = img.shape\n    rgb = cv2.cvtColor(img, cv2.COLOR_GRAY2BGR)\n\n    # === segmentation ===\n    predictor.set_image(rgb)\n    input_point = np.array([[x, y]])\n    input_label = np.array([1])\n    masks, _, _ = predictor.predict(point_coords=input_point, point_labels=input_label, multimask_output=False)\n    mask = masks if masks.ndim == 2 else masks[0]\n    mask = mask.astype(bool)\n\n    # === crop ===\n    bbox = mask_to_bbox(mask)\n    if bbox is None: continue\n    x1, y1, x2, y2 = enlarge_bbox(*bbox, w, h, pad=BOX_PAD)\n    crop = img[y1:y2, x1:x2]\n    if crop.size == 0: continue\n    crop_resized = cv2.resize(crop, (CROP_SIZE, CROP_SIZE))\n\n    # === save into zips ===\n    mask_fname = f\"{study}_{series}_{inst}_{idx}_mask.png\"\n    crop_fname = f\"{study}_{series}_{inst}_{idx}_crop.png\"\n    save_mask_to_zip(mask, mask_fname, zip_masks)\n    save_crop_to_zip(crop_resized, crop_fname, zip_crops)\n\n    records.append({\n        \"filename_crop\": crop_fname,\n        \"filename_mask\": mask_fname,\n        \"study_id\": study,\n        \"series_id\": series,\n        \"instance_number\": inst,\n        \"x\": x, \"y\": y,\n        \"condition\": row.condition,\n        \"level\": row.level\n    })\n    count += 1\n\n# Close and save zips\nzip_masks.close()\nzip_crops.close()\n\n# Save metadata CSV\nout_df = pd.DataFrame(records)\nout_df.to_csv(\"/kaggle/working/medsam_crops_metadata.csv\", index=False)\n\nprint(f\"✅ Done! Saved {len(records)} crops and masks directly into ZIP files.\")\nprint(f\"📦 Crops ZIP size: {os.path.getsize(OUT_ZIP_CROPS)/1e9:.2f} GB\")\nprint(f\"📦 Masks ZIP size: {os.path.getsize(OUT_ZIP_MASKS)/1e9:.2f} GB\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-29T15:58:18.745779Z","iopub.execute_input":"2025-10-29T15:58:18.746044Z","iopub.status.idle":"2025-10-29T15:58:23.801316Z","shell.execute_reply.started":"2025-10-29T15:58:18.746023Z","shell.execute_reply":"2025-10-29T15:58:23.800266Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}