{"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":"gpu","dataSources":[{"sourceId":99552,"databundleVersionId":13190393,"sourceType":"competition"},{"sourceId":12706082,"sourceType":"datasetVersion","datasetId":8030314},{"sourceId":12749799,"sourceType":"datasetVersion","datasetId":8059640},{"sourceId":255206151,"sourceType":"kernelVersion"},{"sourceId":512001,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":405274,"modelId":423182}],"dockerImageVersionId":31090,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import sys\nimport matplotlib.pyplot as plt \nsys.path.append('/kaggle/input/rsna-iad-vesselfm-codebase')\nimport torch\nimport torch.nn.functional as F\nimport numpy as np\nfrom tqdm import tqdm\nfrom monai.inferers import SlidingWindowInfererAdapt\nfrom skimage.morphology import remove_small_objects\nfrom skimage.exposure import equalize_hist\nfrom utils.data import generate_transforms\nfrom utils.io import determine_reader_writer\nimport os\nfrom monai.transforms import LoadImaged, Spacingd, LoadImage\nfrom monai.networks.nets import DynUNet\nimport SimpleITK as sitk\nimport yaml\nimport torch.nn as nn\nfrom scipy.ndimage import label\nfrom tqdm import  tqdm\nimport pydicom\nfrom concurrent.futures import ThreadPoolExecutor\nfrom collections import Counter\nimport pandas as pd\nfrom matplotlib.widgets import Slider\nimport ipywidgets as widgets\nfrom IPython.display import HTML\nfrom matplotlib.animation import FuncAnimation\nimport ast\nfrom tqdm import tqdm","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-08-17T06:49:07.281034Z","iopub.execute_input":"2025-08-17T06:49:07.28131Z","iopub.status.idle":"2025-08-17T06:49:07.288045Z","shell.execute_reply.started":"2025-08-17T06:49:07.281274Z","shell.execute_reply":"2025-08-17T06:49:07.28722Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Brain Veseel Segmentation","metadata":{}},{"cell_type":"code","source":"yaml_path = '/kaggle/input/rsna-iad-vesselfm-codebase/configs/inference.yaml'\n\nwith open(yaml_path, 'r') as f:\n    config = yaml.safe_load(f)\n\n\nclass CFG:\n    ckpt_path = '/kaggle/input/rsna-iad-vesselfm-finetuned/pytorch/default/1/finetune_rsna_vesselfm-val_volumetric_recall0.7549.ckpt'\n    model_structure = {\n        'in_channels': 1,\n        'out_channels': 1,\n        'spatial_dims': 3,\n        'strides': [[1, 1, 1], [2, 2, 2], [2, 2, 2], [2, 2, 2], [2, 2, 2], [2, 2, 2]],\n        'kernel_size': [[3, 3, 3], [3, 3, 3], [3, 3, 3], [3, 3, 3], [3, 3, 3], [3, 3, 3]],\n        'upsample_kernel_size': [[2, 2, 2], [2, 2, 2], [2, 2, 2], [2, 2, 2], [2, 2, 2]],\n        'filters': [32, 64, 128, 256, 320, 320],\n        'res_block': True}\n    device = 'cuda:0'\n    thrd = 0.1\n\n    #sliding window\n    batch_size= 1\n    patch_size= [128, 128, 128]\n    overlap= 0.5\n    mode= \"constant\"\n    sigma_scale= 0.125\n    padding_mode= \"constant\"\n\n    #volume transform\n    transforms_config = config['transforms_config']\n    transforms_config.insert(2, {'Resize': {\n        'spatial_size': (None, 512, 512),\n        'mode': 'bilinear'\n    }},)\n    \n    tta = config['tta']\n    post = config['post']\n    merging = config['merging']","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-17T06:32:52.880549Z","iopub.execute_input":"2025-08-17T06:32:52.881279Z","iopub.status.idle":"2025-08-17T06:32:52.898503Z","shell.execute_reply.started":"2025-08-17T06:32:52.881258Z","shell.execute_reply":"2025-08-17T06:32:52.897678Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def load_model(cfg):\n    ckpt = torch.load(cfg.ckpt_path, map_location=cfg.device, weights_only=False)['state_dict']\n    ckpt = {k.replace(\"model.\", \"\"): v for k, v in ckpt.items()}\n    model = DynUNet(**CFG.model_structure)\n    model.load_state_dict(ckpt)\n    model.eval()\n    return model.to(cfg.device)\n\n\n# def load_series2vol(series_path):\n#     loader = LoadImage(image_only = True, reader = \"ITKReader\")\n#     volume = loader(series_path)\n#     return volume.permute(2, 0, 1)\n\n\n# def load_series2vol(series_path):\n#     reader = sitk.ImageSeriesReader()\n#     dicom_names = reader.GetGDCMSeriesFileNames(series_path)\n#     reader.SetFileNames(dicom_names)\n#     image = reader.Execute()\n#     volume = sitk.GetArrayFromImage(image)\n#     return volume\n\n\ndef load_series2vol(series_path, series_id=None, spacing_tolerance=1e-3, resample=False, default_thickness=1.0):\n    reader = sitk.ImageSeriesReader()\n    \n    # Get all series IDs\n    series_ids = reader.GetGDCMSeriesIDs(series_path)\n    if not series_ids:\n        raise RuntimeError(f\"No DICOM series found in {series_path}\")\n    \n    # Pick first if not specified\n    if series_id is None:\n        series_id = series_ids[0]\n    else:\n        series_id = str(series_id)\n    \n    # Get file names\n    all_files = reader.GetGDCMSeriesFileNames(series_path, series_id)\n    \n    # --- Filter files by consistent size ---\n    file_sizes = {}\n    for f in all_files:\n        img = sitk.ReadImage(f)\n        file_sizes.setdefault(img.GetSize(), []).append(f)\n    \n    # Pick the most common size\n    target_size = max(file_sizes, key=lambda k: len(file_sizes[k]))\n    files = file_sizes[target_size]\n    \n    reader.SetFileNames(files)\n    image = reader.Execute()\n    \n    # --- Fix zero thickness ---\n    spacing = list(image.GetSpacing())\n    if spacing[2] == 0:\n        spacing[2] = default_thickness\n        image.SetSpacing(spacing)\n    \n    # --- Optional resample ---\n    if resample and abs(spacing[2] - spacing[0]) > spacing_tolerance:\n        new_spacing = [spacing[0], spacing[1], spacing[0]]\n        new_size = [\n            int(round(image.GetSize()[0] * spacing[0] / new_spacing[0])),\n            int(round(image.GetSize()[1] * spacing[1] / new_spacing[1])),\n            int(round(image.GetSize()[2] * spacing[2] / new_spacing[2]))\n        ]\n        resampler = sitk.ResampleImageFilter()\n        resampler.SetOutputSpacing(new_spacing)\n        resampler.SetSize(new_size)\n        resampler.SetOutputDirection(image.GetDirection())\n        resampler.SetOutputOrigin(image.GetOrigin())\n        resampler.SetInterpolator(sitk.sitkLinear)\n        image = resampler.Execute(image)\n\n    volume = sitk.GetArrayFromImage(image)\n    \n    return volume\n\n\ndef resample(image, factor=None, target_shape=None):\n    if factor == 1:\n        return image\n\n    if target_shape:\n        _, _, new_d, new_h, new_w = target_shape\n    else:\n        _, _, d, h, w = image.shape\n        new_d, new_h, new_w = int(round(d / factor)), int(round(h / factor)), int(round(w / factor))\n    return F.interpolate(image, size=(new_d, new_h, new_w), mode=\"trilinear\", align_corners=False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-17T06:32:52.899436Z","iopub.execute_input":"2025-08-17T06:32:52.899781Z","iopub.status.idle":"2025-08-17T06:32:52.91474Z","shell.execute_reply.started":"2025-08-17T06:32:52.899743Z","shell.execute_reply":"2025-08-17T06:32:52.914097Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def read_dicom_info(path):\n    \"\"\"Read DICOM metadata needed for sorting and decoding.\"\"\"\n    ds = pydicom.dcmread(path, stop_before_pixels=True)\n    try:\n        z = float(ds.ImagePositionPatient[2])  # Preferred\n    except AttributeError:\n        z = float(ds.InstanceNumber)           # Fallback\n    return path, z, ds\n\ndef read_pixel_data(path, rescale_slope, rescale_intercept):\n    \"\"\"Read pixel data and apply rescale.\"\"\"\n    ds = pydicom.dcmread(path)  # full read with pixels\n    arr = ds.pixel_array.astype(np.float32)\n    if rescale_slope is not None and rescale_intercept is not None:\n        arr = arr * rescale_slope + rescale_intercept\n    return arr\n\ndef load_volume_fast(directory, max_workers=8):\n    # Step 1: List all .dcm files\n    dcm_files = [os.path.join(directory, f) for f in os.listdir(directory) if f.endswith('.dcm')]\n    if not dcm_files:\n        raise RuntimeError(f\"No DICOM files found in {directory}\")\n\n    # Step 2: Read metadata in parallel\n    with ThreadPoolExecutor(max_workers=max_workers) as executor:\n        meta_info = list(executor.map(read_dicom_info, dcm_files))\n\n    # Step 3: Sort by slice position\n    meta_info.sort(key=lambda x: x[1])\n    sorted_paths = [m[0] for m in meta_info]\n    first_ds = meta_info[0][2]\n\n    # Extract rescale params once\n    rescale_slope = getattr(first_ds, \"RescaleSlope\", None)\n    rescale_intercept = getattr(first_ds, \"RescaleIntercept\", None)\n\n    # Step 4: Read one slice to get shape and dimension\n    test_arr = pydicom.dcmread(sorted_paths[0]).pixel_array\n    if test_arr.ndim == 2:\n        depth = len(sorted_paths)\n        height, width = test_arr.shape\n        volume = np.zeros((depth, height, width), dtype=np.float32)\n    elif test_arr.ndim == 3:\n        # multi-frame DICOM case\n        depth = test_arr.shape[0]\n        height, width = test_arr.shape[1], test_arr.shape[2]\n        if len(sorted_paths) > 1:\n            raise ValueError(\"Multiple multi-frame DICOM files not supported yet\")\n        volume = np.zeros((depth, height, width), dtype=np.float32)\n    else:\n        raise ValueError(f\"Unexpected pixel array shape: {test_arr.shape}\")\n\n    # Step 5: Load pixel data in parallel directly into volume\n    def load_into_array(idx_path):\n        idx, path = idx_path\n        arr = read_pixel_data(path, rescale_slope, rescale_intercept)\n        if arr.ndim == 3 and volume.shape[0] == arr.shape[0]:\n            volume[:] = arr  # entire volume from single file\n        elif arr.ndim == 2:\n            volume[idx] = arr\n        else:\n            raise ValueError(f\"Shape mismatch loading {path}, arr shape: {arr.shape}, volume shape: {volume.shape}\")\n\n    with ThreadPoolExecutor(max_workers=max_workers) as executor:\n        executor.map(load_into_array, enumerate(sorted_paths))\n\n    return volume","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-17T06:32:52.915717Z","iopub.execute_input":"2025-08-17T06:32:52.915977Z","iopub.status.idle":"2025-08-17T06:32:52.937744Z","shell.execute_reply.started":"2025-08-17T06:32:52.91595Z","shell.execute_reply":"2025-08-17T06:32:52.936969Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"inferer = SlidingWindowInfererAdapt(\n        roi_size=CFG.patch_size, sw_batch_size=CFG.batch_size, overlap=CFG.overlap,\n        mode=CFG.mode, sigma_scale=CFG.sigma_scale, padding_mode=CFG.padding_mode\n    )\n\ntransforms = generate_transforms(CFG.transforms_config)\n\nimage_reader_writer = determine_reader_writer('nii')()\n\nmodel = load_model(CFG)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-17T06:32:52.939535Z","iopub.execute_input":"2025-08-17T06:32:52.939778Z","iopub.status.idle":"2025-08-17T06:32:56.528986Z","shell.execute_reply.started":"2025-08-17T06:32:52.93976Z","shell.execute_reply":"2025-08-17T06:32:56.528131Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def inference(path, load_series=False):\n    #vol = load_series2vol(path)\n\n    preds = []\n    for scale in CFG.tta['scales']:\n        if load_series:\n            image = load_volume_fast(path)\n        else:\n            image = image_reader_writer.read_images(path)[0]\n        print(image.shape)\n        image = transforms(image.astype(np.float32))[None].to(CFG.device)\n        #apply test time augmentation\n        if CFG.tta['invert']:\n            image = 1 - image if image.mean() > CFG.tta['invert_mean_thresh'] else image\n            \n        if CFG.tta['equalize_hist']:\n            image_np = image.cpu().squeeze().numpy()\n            image_equal_hist_np = equalize_hist(image_np, nbins=CFG.tta['hist_bins'])\n            image = torch.from_numpy(image_equal_hist_np).to(image.device)[None][None]\n\n        original_shape = image.shape\n        image = resample(image, factor=scale)\n        logits = inferer(image, model)\n        logits = resample(logits, target_shape=original_shape)\n        preds.append(logits.cpu().squeeze())\n    if CFG.merging['max']:\n        pred = torch.stack(preds).max(dim=0)[0].sigmoid()\n    else:\n        pred = torch.stack(preds).mean(dim=0).sigmoid()\n    # pred_thresh = (pred > CFG.thrd).numpy()\n\n    # # post-processing\n    # if CFG.post['apply']:\n    #     pred_thresh = remove_small_objects(\n    #         pred_thresh, min_size=CFG.post['small_objects_min_size'],\n    #         connectivity=CFG.post['small_objects_connectivity']\n    #     )\n    return  pred, image","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-17T06:32:56.52996Z","iopub.execute_input":"2025-08-17T06:32:56.530221Z","iopub.status.idle":"2025-08-17T06:32:56.537657Z","shell.execute_reply.started":"2025-08-17T06:32:56.530199Z","shell.execute_reply":"2025-08-17T06:32:56.53677Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df_loc = pd.read_csv('/kaggle/input/rsna-intracranial-aneurysm-detection/train_localizers.csv')\ndf_loc['x'] = df_loc['coordinates'].map(lambda x: ast.literal_eval(x)['x'])\ndf_loc['y'] = df_loc['coordinates'].map(lambda x: ast.literal_eval(x)['y'])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-17T06:49:41.872278Z","iopub.execute_input":"2025-08-17T06:49:41.872566Z","iopub.status.idle":"2025-08-17T06:49:41.944348Z","shell.execute_reply.started":"2025-08-17T06:49:41.872543Z","shell.execute_reply":"2025-08-17T06:49:41.943825Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"series_path = '/kaggle/input/rsna-intracranial-aneurysm-detection/series'","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-17T06:49:42.003279Z","iopub.execute_input":"2025-08-17T06:49:42.003535Z","iopub.status.idle":"2025-08-17T06:49:42.006973Z","shell.execute_reply.started":"2025-08-17T06:49:42.003517Z","shell.execute_reply":"2025-08-17T06:49:42.006315Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"series_maps = {}\nfor series_uid in tqdm(df_loc[\"SeriesInstanceUID\"].unique()):\n    series_path_ = f\"{series_path}/{series_uid}\"\n    files = sorted(os.listdir(series_path_))  # ensure consistent order\n    # strip .dcm to get SOPInstanceUID\n    series_maps[series_uid] = {\n        f.replace(\".dcm\", \"\"): idx for idx, f in enumerate(files)\n    }","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-17T06:49:42.664355Z","iopub.execute_input":"2025-08-17T06:49:42.6651Z","iopub.status.idle":"2025-08-17T06:49:44.284333Z","shell.execute_reply.started":"2025-08-17T06:49:42.665073Z","shell.execute_reply":"2025-08-17T06:49:44.283633Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Flatten series_maps into a dataframe\nmapping = [\n    {\"SeriesInstanceUID\": series_uid, \"SOPInstanceUID\": sop_uid, \"dcm_idx\": int(idx)}\n    for series_uid, uid_map in series_maps.items()\n    for sop_uid, idx in uid_map.items()\n]\ndf_map = pd.DataFrame(mapping)\n\n# Merge instead of apply\ndf_loc = df_loc.merge(df_map, on=[\"SeriesInstanceUID\", \"SOPInstanceUID\"], how=\"left\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-17T06:49:46.870302Z","iopub.execute_input":"2025-08-17T06:49:46.870855Z","iopub.status.idle":"2025-08-17T06:49:47.55891Z","shell.execute_reply.started":"2025-08-17T06:49:46.870832Z","shell.execute_reply":"2025-08-17T06:49:47.55815Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df_loc[df.SeriesInstanceUID == uid]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-17T06:51:59.233462Z","iopub.execute_input":"2025-08-17T06:51:59.233738Z","iopub.status.idle":"2025-08-17T06:51:59.245861Z","shell.execute_reply.started":"2025-08-17T06:51:59.233716Z","shell.execute_reply":"2025-08-17T06:51:59.245273Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"path = '/kaggle/input/rsna-intracranial-aneurysm-detection/series/1.2.826.0.1.3680043.8.498.10004044428023505108375152878107656647'","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-17T06:56:04.714194Z","iopub.execute_input":"2025-08-17T06:56:04.714506Z","iopub.status.idle":"2025-08-17T06:56:04.71807Z","shell.execute_reply.started":"2025-08-17T06:56:04.714485Z","shell.execute_reply":"2025-08-17T06:56:04.717459Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df = pd.read_csv('/kaggle/input/rsna-intracranial-aneurysm-detection/train.csv')\ndf_abnormal = pd.read_csv('/kaggle/input/addition-resource-rsna-iad/multiframe_dicoms.csv')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-17T06:32:56.538809Z","iopub.execute_input":"2025-08-17T06:32:56.539092Z","iopub.status.idle":"2025-08-17T06:32:56.589371Z","shell.execute_reply.started":"2025-08-17T06:32:56.539065Z","shell.execute_reply":"2025-08-17T06:32:56.58864Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"series_path = '/kaggle/input/rsna-intracranial-aneurysm-detection/series'\npaths = sorted(os.listdir(series_path))\npath = os.path.join(series_path, paths[0])\nlen(os.listdir(path))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-17T06:32:56.590439Z","iopub.execute_input":"2025-08-17T06:32:56.590712Z","iopub.status.idle":"2025-08-17T06:32:56.777611Z","shell.execute_reply.started":"2025-08-17T06:32:56.59069Z","shell.execute_reply":"2025-08-17T06:32:56.776762Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"uid = path.split('/')[-1]\n\ndf_sample = df[df.SeriesInstanceUID==uid].copy()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-17T06:32:56.778278Z","iopub.execute_input":"2025-08-17T06:32:56.77852Z","iopub.status.idle":"2025-08-17T06:32:56.790574Z","shell.execute_reply.started":"2025-08-17T06:32:56.778502Z","shell.execute_reply":"2025-08-17T06:32:56.789761Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"uid in df_abnormal.SeriesInstanceUID.tolist()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-17T06:32:56.791515Z","iopub.execute_input":"2025-08-17T06:32:56.79179Z","iopub.status.idle":"2025-08-17T06:32:56.80741Z","shell.execute_reply.started":"2025-08-17T06:32:56.791766Z","shell.execute_reply":"2025-08-17T06:32:56.806677Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df_sample","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-17T06:32:56.808427Z","iopub.execute_input":"2025-08-17T06:32:56.808796Z","iopub.status.idle":"2025-08-17T06:32:56.839124Z","shell.execute_reply.started":"2025-08-17T06:32:56.808771Z","shell.execute_reply":"2025-08-17T06:32:56.838366Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df.Modality.unique()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-17T06:32:56.840042Z","iopub.execute_input":"2025-08-17T06:32:56.840336Z","iopub.status.idle":"2025-08-17T06:32:56.84816Z","shell.execute_reply.started":"2025-08-17T06:32:56.840311Z","shell.execute_reply":"2025-08-17T06:32:56.847464Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"CFG.transforms_config","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-17T06:32:56.849009Z","iopub.execute_input":"2025-08-17T06:32:56.849279Z","iopub.status.idle":"2025-08-17T06:32:56.864545Z","shell.execute_reply.started":"2025-08-17T06:32:56.849257Z","shell.execute_reply":"2025-08-17T06:32:56.86375Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"with torch.no_grad():\n    p, im = inference(path, True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-17T06:32:56.867703Z","iopub.execute_input":"2025-08-17T06:32:56.867942Z","iopub.status.idle":"2025-08-17T06:33:22.057755Z","shell.execute_reply.started":"2025-08-17T06:32:56.867923Z","shell.execute_reply":"2025-08-17T06:33:22.057086Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"depth = p.shape[0]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-17T06:33:22.058461Z","iopub.execute_input":"2025-08-17T06:33:22.058731Z","iopub.status.idle":"2025-08-17T06:33:22.062588Z","shell.execute_reply.started":"2025-08-17T06:33:22.058705Z","shell.execute_reply":"2025-08-17T06:33:22.0619Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"depth","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-17T06:33:22.06327Z","iopub.execute_input":"2025-08-17T06:33:22.063472Z","iopub.status.idle":"2025-08-17T06:33:22.155517Z","shell.execute_reply.started":"2025-08-17T06:33:22.063457Z","shell.execute_reply":"2025-08-17T06:33:22.154857Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"im.shape","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-17T06:33:22.156256Z","iopub.execute_input":"2025-08-17T06:33:22.156569Z","iopub.status.idle":"2025-08-17T06:33:22.173929Z","shell.execute_reply.started":"2025-08-17T06:33:22.156538Z","shell.execute_reply":"2025-08-17T06:33:22.173252Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plt.imshow(p[depth//4], alpha = 0.8, cmap='jet')\nplt.imshow(im[0, 0, depth//4].cpu(), alpha=0.5, cmap='gray')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-17T06:33:22.17472Z","iopub.execute_input":"2025-08-17T06:33:22.174948Z","iopub.status.idle":"2025-08-17T06:33:22.523709Z","shell.execute_reply.started":"2025-08-17T06:33:22.17493Z","shell.execute_reply":"2025-08-17T06:33:22.523041Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Point Sampling & Extraction","metadata":{}},{"cell_type":"markdown","source":"### Sampling 3D Points with Ellipsoidal Mask and Gaussian Distribution\n\nLet the 3D volume dimensions be \\( Z, Y, X \\).\n\nDefine the volume center:\n$$\n\\mathbf{c} = \\left(\\frac{Z}{2}, \\frac{Y}{2}, \\frac{X}{2}\\right)\n$$\n\nDefine elliptical radius scales:\n$$\nr_z = \\frac{Z}{2} \\cdot s_z, \\quad s_z \\in (0,1]\n$$\n$$\nr_{xy} = \\frac{\\min(Y, X)}{2}\n$$\n\n## Ellipsoidal Mask Condition\n\nA point \\(\\mathbf{p} = (z, y, x)\\) lies inside the ellipsoid if:\n$$\n\\left(\\frac{z - c_z}{r_z}\\right)^2 + \\frac{(y - c_y)^2 + (x - c_x)^2}{r_{xy}^2} \\leq 1\n$$\n\n## Segmentation Mask Indicator\n\nDefine:\n$$\nM(\\mathbf{p}) = \n\\begin{cases}\n1 & \\text{if } \\mathbf{p} \\text{ is inside the segmentation (positive)} \\\\\n0 & \\text{otherwise}\n\\end{cases}\n$$\n\n## Sampling Procedure\n\n1. Generate candidate points \\(\\mathbf{q}_i\\), \\(i=1,\\dots,N_c\\), from a 3D Gaussian distribution centered at \\(\\mathbf{c}\\):\n$$\n\\mathbf{q}_i = \\mathbf{c} + \n\\begin{bmatrix}\nZ_i \\\\\nY_i \\\\\nX_i\n\\end{bmatrix}, \\quad \\text{where} \\quad\nZ_i \\sim \\mathcal{N}(0, \\sigma_z^2), \\quad\nY_i \\sim \\mathcal{N}(0, \\sigma_{xy}^2), \\quad\nX_i \\sim \\mathcal{N}(0, \\sigma_{xy}^2)\n$$\n\nwith\n$$\n\\sigma_z = \\text{std\\_scale} \\times r_z, \\quad \\sigma_{xy} = \\text{std\\_scale} \\times r_{xy}\n$$\n\n2. Keep only candidate points inside the ellipsoid:\n$$\nS = \\{ \\mathbf{q}_i : \\mathbf{q}_i \\text{ satisfies ellipsoid condition} \\}\n$$\n\n3. From \\(S\\), keep only points positive in the segmentation mask:\n$$\nS^+ = \\{ \\mathbf{q} \\in S : M(\\mathbf{q}) = 1 \\}\n$$\n\n4. Final sample points \\(P\\) are chosen by:\n$$\nP = \n\\begin{cases}\n\\text{UniformSample}(S^+, N), & |S^+| \\geq N \\\\\n\\text{UniformSample}(\\{ \\mathbf{p} : M(\\mathbf{p})=1 \\}, N), & |S^+| < N \\leq |\\{ \\mathbf{p} : M(\\mathbf{p})=1 \\}| \\\\\n\\text{RepeatedSample}(\\{ \\mathbf{p} : M(\\mathbf{p})=1 \\}, N), & 0 < |\\{ \\mathbf{p} : M(\\mathbf{p})=1 \\}| < N \\\\\n\\text{UniformSample}(\\{ \\mathbf{p} : \\text{ellipsoid condition} \\}, N), & \\text{otherwise}\n\\end{cases}\n$$\n","metadata":{}},{"cell_type":"code","source":"def sample_positive_points(segmentation, N, z_scale=1., std_scale=0.5, seed=42, xy_ratio = 2.5, z_ratio=2):\n    if seed is not None:\n        torch.manual_seed(seed)   # Fix the PyTorch RNG seed\n    \n    B, _, Z, Y, X = segmentation.shape\n    points_list = []\n\n    center_z, center_y, center_x = Z / 2.0, Y / 2.0, X / 2.0\n    radius_z = (Z / z_ratio) * z_scale\n    radius_xy = min(Y, X) / xy_ratio  # circle radius in XY plane\n\n    device = segmentation.device\n\n    # Create ellipsoid mask\n    zz, yy, xx = torch.meshgrid(\n        torch.arange(Z, dtype=torch.float32, device=device),\n        torch.arange(Y, dtype=torch.float32, device=device),\n        torch.arange(X, dtype=torch.float32, device=device),\n        indexing='ij'\n    )\n    ellipsoid_mask = (\n        ((zz - center_z) / radius_z) ** 2 +\n        ((yy - center_y) ** 2 + (xx - center_x) ** 2) / radius_xy ** 2\n    ) <= 1\n\n    for b in range(B):\n        seg_masked = segmentation[b, 0] * ellipsoid_mask\n\n        pos_idx = torch.nonzero(seg_masked, as_tuple=False)\n\n        num_candidates = max(N * 10, 1000)\n\n        std_z = radius_z * std_scale\n        std_xy = radius_xy * std_scale\n\n        gaussian_samples = torch.empty((num_candidates, 3), device=device).normal_(0, 1)\n        gaussian_samples[:, 0] *= std_z\n        gaussian_samples[:, 1] *= std_xy\n        gaussian_samples[:, 2] *= std_xy\n        gaussian_samples += torch.tensor([center_z, center_y, center_x], device=device)\n\n        gaussian_samples = gaussian_samples.round().long()\n        gaussian_samples[:, 0] = gaussian_samples[:, 0].clamp(0, Z - 1)\n        gaussian_samples[:, 1] = gaussian_samples[:, 1].clamp(0, Y - 1)\n        gaussian_samples[:, 2] = gaussian_samples[:, 2].clamp(0, X - 1)\n\n        gaussian_samples = torch.unique(gaussian_samples, dim=0)\n\n        mask_vals = ellipsoid_mask[\n            gaussian_samples[:, 0], gaussian_samples[:, 1], gaussian_samples[:, 2]\n        ]\n        valid_gaussian_points = gaussian_samples[mask_vals]\n\n        seg_vals = segmentation[b, 0][\n            valid_gaussian_points[:, 0], valid_gaussian_points[:, 1], valid_gaussian_points[:, 2]\n        ]\n        valid_pos_gaussian_points = valid_gaussian_points[seg_vals > 0]\n\n        if len(valid_pos_gaussian_points) >= N:\n            choice = torch.randperm(len(valid_pos_gaussian_points), device=device)[:N]\n            sampled = valid_pos_gaussian_points[choice]\n        elif len(pos_idx) >= N:\n            choice = torch.randperm(len(pos_idx), device=device)[:N]\n            sampled = pos_idx[choice]\n        elif len(pos_idx) > 0:\n            repeats = (N + len(pos_idx) - 1) // len(pos_idx)\n            repeated = pos_idx.repeat((repeats, 1))\n            choice = torch.randperm(len(repeated), device=device)[:N]\n            sampled = repeated[choice]\n        else:\n            valid_idx = torch.nonzero(ellipsoid_mask, as_tuple=False)\n            choice = torch.randperm(len(valid_idx), device=device)[:N]\n            sampled = valid_idx[choice]\n\n        points_list.append(sampled)\n\n    points = torch.stack(points_list, dim=0)  # (B, N, 3)\n    return points","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-17T06:33:22.524652Z","iopub.execute_input":"2025-08-17T06:33:22.524933Z","iopub.status.idle":"2025-08-17T06:33:22.53897Z","shell.execute_reply.started":"2025-08-17T06:33:22.524907Z","shell.execute_reply":"2025-08-17T06:33:22.538284Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"mask = p>0.1\nmask.sum()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-17T06:33:22.539787Z","iopub.execute_input":"2025-08-17T06:33:22.53997Z","iopub.status.idle":"2025-08-17T06:33:22.807681Z","shell.execute_reply.started":"2025-08-17T06:33:22.539947Z","shell.execute_reply":"2025-08-17T06:33:22.806988Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"N = 10000 #number of sampling point","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-17T06:33:22.808341Z","iopub.execute_input":"2025-08-17T06:33:22.808569Z","iopub.status.idle":"2025-08-17T06:33:22.812043Z","shell.execute_reply.started":"2025-08-17T06:33:22.808547Z","shell.execute_reply":"2025-08-17T06:33:22.811326Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"points = sample_positive_points(mask[None, None], N)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-17T06:33:22.812746Z","iopub.execute_input":"2025-08-17T06:33:22.813015Z","iopub.status.idle":"2025-08-17T06:33:24.481977Z","shell.execute_reply.started":"2025-08-17T06:33:22.812997Z","shell.execute_reply":"2025-08-17T06:33:24.481186Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(f\"sampling ratio: {points.shape[1] * 100/mask.sum():.4f}%\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-17T06:33:24.48283Z","iopub.execute_input":"2025-08-17T06:33:24.48309Z","iopub.status.idle":"2025-08-17T06:33:24.695429Z","shell.execute_reply.started":"2025-08-17T06:33:24.483067Z","shell.execute_reply":"2025-08-17T06:33:24.694766Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"points.shape[1]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-17T06:33:24.696139Z","iopub.execute_input":"2025-08-17T06:33:24.696437Z","iopub.status.idle":"2025-08-17T06:33:24.701252Z","shell.execute_reply.started":"2025-08-17T06:33:24.696409Z","shell.execute_reply":"2025-08-17T06:33:24.700726Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 2D projection","metadata":{}},{"cell_type":"code","source":"plt.scatter(points[0, :, 1], points[0, :, 2])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-17T06:33:24.701974Z","iopub.execute_input":"2025-08-17T06:33:24.702204Z","iopub.status.idle":"2025-08-17T06:33:25.061002Z","shell.execute_reply.started":"2025-08-17T06:33:24.70218Z","shell.execute_reply":"2025-08-17T06:33:25.060196Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 3D Visualization","metadata":{}},{"cell_type":"code","source":"from mpl_toolkits.mplot3d import Axes3D  # needed for 3D plotting","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-17T06:33:25.06186Z","iopub.execute_input":"2025-08-17T06:33:25.062123Z","iopub.status.idle":"2025-08-17T06:33:25.065859Z","shell.execute_reply.started":"2025-08-17T06:33:25.062099Z","shell.execute_reply":"2025-08-17T06:33:25.065156Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Extract coordinates\nx, y, z = points[0, :, 0], points[0, :, 1], points[0, :, 2]\n\n# Create 3D scatter plot\nfig = plt.figure(figsize=(8, 8))\nax = fig.add_subplot(111, projection='3d')\nax.scatter(x, y, z, s=1, alpha=0.5)  # s=1 for small dots\n\nax.set_xlabel('X')\nax.set_ylabel('Y')\nax.set_zlabel('Z')\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-17T06:33:25.066656Z","iopub.execute_input":"2025-08-17T06:33:25.067009Z","iopub.status.idle":"2025-08-17T06:33:25.360921Z","shell.execute_reply.started":"2025-08-17T06:33:25.066989Z","shell.execute_reply":"2025-08-17T06:33:25.360009Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Check point sampling cross all slices","metadata":{}},{"cell_type":"code","source":"# %matplotlib inline\n\n# def browse_slices(z):\n#     fig, ax = plt.subplots()\n#     idx = points[0, :, 0] == z\n#     ax.imshow(im[0, 0, z].cpu(), alpha=0.5, cmap='gray')\n#     if idx.any():\n#         ax.scatter(points[0, idx, 1], points[0, idx, 2], color='r', s=10)\n#     ax.set_title(f\"Z-slice: {z}\")\n#     plt.show()\n\n# widgets.interact(browse_slices, z=widgets.IntSlider(min=0, max=im[0,0].shape[0]-1, step=1, value=0))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-17T06:33:25.36171Z","iopub.execute_input":"2025-08-17T06:33:25.361937Z","iopub.status.idle":"2025-08-17T06:33:25.365616Z","shell.execute_reply.started":"2025-08-17T06:33:25.36192Z","shell.execute_reply":"2025-08-17T06:33:25.364759Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def create_point_overlay_animation(im, points, interval=200, figsize=(6, 6)):\n    \"\"\"\n    im: torch.Tensor (1, 1, Z, H, W) or numpy array (Z, H, W)\n    points: numpy array (N, 3) or (B, N, 3) with (z, y, x)\n    \"\"\"\n\n    # Convert tensors to numpy if needed\n    if hasattr(im, \"cpu\"):\n        im = im.cpu().numpy()\n    if im.ndim == 5:  # (B, C, Z, H, W)\n        im = im[0, 0]\n    elif im.ndim == 4:  # (C, Z, H, W)\n        im = im[0]\n    \n    if hasattr(points, \"cpu\"):\n        points = points.cpu().numpy()\n    if points.ndim == 3:\n        points = points[0]\n\n    Z = im.shape[0]\n\n    fig, ax = plt.subplots(figsize=figsize)\n    img_artist = ax.imshow(im[0], alpha=0.5, cmap='gray')\n    scatter_artist = ax.scatter([], [], color='r', s=10)\n    ax.set_title(f\"Z-slice: 0\")\n\n    def update(frame):\n        img_artist.set_array(im[frame])\n        idx = points[:, 0] == frame\n        if np.any(idx):\n            scatter_artist.set_offsets(points[idx][:, [2, 1]])  # (x, y)\n        else:\n            scatter_artist.set_offsets(np.empty((0, 2)))  # ensure 2D empty array\n        ax.set_title(f\"Z-slice: {frame}\")\n        return img_artist, scatter_artist\n    anim = FuncAnimation(fig, update, frames=Z, interval=interval, blit=True)\n    plt.close(fig)\n    return anim","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-17T06:33:25.366521Z","iopub.execute_input":"2025-08-17T06:33:25.366797Z","iopub.status.idle":"2025-08-17T06:33:25.381666Z","shell.execute_reply.started":"2025-08-17T06:33:25.366773Z","shell.execute_reply":"2025-08-17T06:33:25.380915Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"anim = create_point_overlay_animation(im, points, interval=200)\nHTML(anim.to_jshtml())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-17T06:33:25.382478Z","iopub.execute_input":"2025-08-17T06:33:25.382814Z","iopub.status.idle":"2025-08-17T06:33:33.224617Z","shell.execute_reply.started":"2025-08-17T06:33:25.382789Z","shell.execute_reply":"2025-08-17T06:33:33.223194Z"},"collapsed":true,"jupyter":{"outputs_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Extraction","metadata":{}},{"cell_type":"code","source":"class StopForward(Exception): pass\n\ndef extract_features_with_hook(model, image, stop_at_layer):\n    features = {}\n    def hook_fn(module, input, output):\n        features['feat'] = output\n        raise StopForward()\n\n    target_module = dict(model.named_modules())[stop_at_layer]\n    handle = target_module.register_forward_hook(hook_fn)\n    try:\n        _ = model(image)\n    except StopForward:\n        pass\n    finally:\n        handle.remove()\n    return features['feat']\n\n\ndef compute_slices(spatial_shape, roi_size, step):\n    slices_list = []\n    for start_d in range(0, spatial_shape[0], step[0]):\n        end_d = start_d + roi_size[0]\n        if end_d > spatial_shape[0]:\n            start_d = spatial_shape[0] - roi_size[0]\n            end_d = spatial_shape[0]\n\n        for start_h in range(0, spatial_shape[1], step[1]):\n            end_h = start_h + roi_size[1]\n            if end_h > spatial_shape[1]:\n                start_h = spatial_shape[1] - roi_size[1]\n                end_h = spatial_shape[1]\n\n            for start_w in range(0, spatial_shape[2], step[2]):\n                end_w = start_w + roi_size[2]\n                if end_w > spatial_shape[2]:\n                    start_w = spatial_shape[2] - roi_size[2]\n                    end_w = spatial_shape[2]\n\n                slices_list.append((\n                    (start_d, end_d),\n                    (start_h, end_h),\n                    (start_w, end_w)\n                ))\n\n    # Remove duplicates by converting to hashable tuples\n    slices_list = list(dict.fromkeys(slices_list))\n\n    # Convert tuples back to slices\n    slices_list = [\n        (slice(d[0], d[1]), slice(h[0], h[1]), slice(w[0], w[1]))\n        for d, h, w in slices_list\n    ]\n    return slices_list\n\n\ndef sliding_window_patches(image, roi_size, device=None, overlap=0.5):\n    batch_mode = (image.dim() == 5)\n    if batch_mode:\n        B = image.shape[0]\n        C = image.shape[1]\n        spatial_shape = image.shape[2:]\n    else:\n        B = 1\n        C = image.shape[0]\n        spatial_shape = image.shape[1:]\n        image = image.unsqueeze(0)  # add batch dim\n\n    step = [max(1, int(s * (1 - overlap))) for s in roi_size]\n    slices = compute_slices(spatial_shape, roi_size, step)\n\n    for sl in slices:\n        patch = image[(slice(None), slice(None)) + sl]  # (B, C, D_roi, H_roi, W_roi)\n        coord = (sl[0].start, sl[1].start, sl[2].start)\n        if batch_mode:\n            yield patch.to(device), coord\n        else:\n            yield patch[0].to(device), coord\n\n\ndef assign_point_features_with_sliding(model, image, points, stop_at=\"downsamples.3\", roi_size=(128,128,128), overlap=0.5):\n    B, N, _ = points.shape\n    device = image.device\n    point_feats_sum = None\n    point_counts = torch.zeros((B, N), device=device)\n\n    for patch, (start_z, start_y, start_x) in sliding_window_patches(image, roi_size, device=device, overlap=overlap):\n        with torch.no_grad():\n            feat_map = extract_features_with_hook(model, patch, stop_at)\n        C_feat = feat_map.shape[1]\n\n        if point_feats_sum is None:\n            point_feats_sum = torch.zeros((B, N, C_feat), device=device)\n\n        scale_z = feat_map.shape[2] / patch.shape[2]\n        scale_y = feat_map.shape[3] / patch.shape[3]\n        scale_x = feat_map.shape[4] / patch.shape[4]\n\n        for b in range(B):\n            mask_inside = (\n                (points[b, :, 0] >= start_z) & (points[b, :, 0] < start_z + patch.shape[2]) &\n                (points[b, :, 1] >= start_y) & (points[b, :, 1] < start_y + patch.shape[3]) &\n                (points[b, :, 2] >= start_x) & (points[b, :, 2] < start_x + patch.shape[4])\n            )\n            inside_idx = torch.nonzero(mask_inside, as_tuple=False).squeeze(1)\n            if len(inside_idx) == 0:\n                continue\n\n            local_z = ((points[b, inside_idx, 0] - start_z).float() * scale_z).long().clamp(0, feat_map.shape[2]-1)\n            local_y = ((points[b, inside_idx, 1] - start_y).float() * scale_y).long().clamp(0, feat_map.shape[3]-1)\n            local_x = ((points[b, inside_idx, 2] - start_x).float() * scale_x).long().clamp(0, feat_map.shape[4]-1)\n\n            feats = feat_map[b, :, local_z, local_y, local_x].permute(1,0)  # (num_pts, C_feat)\n            point_feats_sum[b, inside_idx] += feats\n            point_counts[b, inside_idx] += 1\n\n    point_feats = point_feats_sum / point_counts.clamp_min(1).unsqueeze(-1)\n    return point_feats","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-17T06:33:33.22539Z","iopub.status.idle":"2025-08-17T06:33:33.225779Z","shell.execute_reply.started":"2025-08-17T06:33:33.225582Z","shell.execute_reply":"2025-08-17T06:33:33.225599Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"target_layer = \"downsamples.2\"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-17T06:33:33.227249Z","iopub.status.idle":"2025-08-17T06:33:33.22755Z","shell.execute_reply.started":"2025-08-17T06:33:33.22739Z","shell.execute_reply":"2025-08-17T06:33:33.227423Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"point_feats = assign_point_features_with_sliding(model, im, points, stop_at=target_layer)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-17T06:33:33.228432Z","iopub.status.idle":"2025-08-17T06:33:33.22881Z","shell.execute_reply.started":"2025-08-17T06:33:33.228604Z","shell.execute_reply":"2025-08-17T06:33:33.22862Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"point_feats.shape","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-17T06:33:33.230032Z","iopub.status.idle":"2025-08-17T06:33:33.230342Z","shell.execute_reply.started":"2025-08-17T06:33:33.230181Z","shell.execute_reply":"2025-08-17T06:33:33.230195Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Graph Connection","metadata":{}},{"cell_type":"code","source":"!pip install /kaggle/input/pip-install-pyg-v2/torch_spline_conv-1.2.2+pt26cu124-cp311-cp311-linux_x86_64.whl\n!pip install /kaggle/input/pip-install-pyg-v2/torch_sparse-0.6.18+pt26cu124-cp311-cp311-linux_x86_64.whl\n!pip install /kaggle/input/pip-install-pyg-v2/pyg_lib-0.4.0+pt26cu124-cp311-cp311-linux_x86_64.whl\n!pip install /kaggle/input/pip-install-pyg-v2/torch_cluster-1.6.3+pt26cu124-cp311-cp311-linux_x86_64.whl\n!pip install /kaggle/input/pip-install-pyg-v2/torch_geometric-2.6.1-py3-none-any.whl","metadata":{"trusted":true,"_kg_hide-output":true,"_kg_hide-input":true,"execution":{"iopub.status.busy":"2025-08-17T06:33:33.231941Z","iopub.status.idle":"2025-08-17T06:33:33.232284Z","shell.execute_reply.started":"2025-08-17T06:33:33.23212Z","shell.execute_reply":"2025-08-17T06:33:33.232137Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from torch_geometric.nn import radius_graph\nfrom torch_cluster import knn_graph\nfrom torch_geometric.transforms import AddRandomWalkPE\nfrom torch_geometric.data import Data","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-17T06:33:33.233035Z","iopub.status.idle":"2025-08-17T06:33:33.233274Z","shell.execute_reply.started":"2025-08-17T06:33:33.233159Z","shell.execute_reply":"2025-08-17T06:33:33.233171Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"batch = torch.zeros(points.shape[0], dtype=torch.int64)\nedge_index = knn_graph(points[0], k=15, loop=False)\ndata= Data(points = points[0], x = point_feats[0], edge_index = edge_index, batch = batch)\nadd_pos = AddRandomWalkPE(walk_length=8, attr_name=None)\ndata = add_pos(data)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-17T06:33:33.234999Z","iopub.status.idle":"2025-08-17T06:33:33.235227Z","shell.execute_reply.started":"2025-08-17T06:33:33.235125Z","shell.execute_reply":"2025-08-17T06:33:33.235134Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"data.x.shape #feature_dim + walk_length","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-17T06:33:33.236503Z","iopub.status.idle":"2025-08-17T06:33:33.236746Z","shell.execute_reply.started":"2025-08-17T06:33:33.23663Z","shell.execute_reply":"2025-08-17T06:33:33.236643Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}