{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.11","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,"sourceType":"competition"},{"sourceId":366018,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":303561,"modelId":324048}],"dockerImageVersionId":31011,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\n# We will be using these utils provided by PyTorch to simplify training loop code\n!pip install -qU pycocotools\nos.system(\"wget https://raw.githubusercontent.com/pytorch/vision/main/references/detection/engine.py\")\nos.system(\"wget https://raw.githubusercontent.com/pytorch/vision/main/references/detection/utils.py\")\nos.system(\"wget https://raw.githubusercontent.com/pytorch/vision/main/references/detection/coco_utils.py\")\nos.system(\"wget https://raw.githubusercontent.com/pytorch/vision/main/references/detection/coco_eval.py\")\nos.system(\"wget https://raw.githubusercontent.com/pytorch/vision/main/references/detection/transforms.py\")\nimport utils\nfrom engine import train_one_epoch, evaluate\n\nimport re\nimport itertools\nimport numpy as np\nimport pandas as pd\nfrom datetime import datetime\nfrom sklearn.model_selection import train_test_split\n\nimport seaborn as sns\nimport matplotlib.pyplot as plt\nimport matplotlib.patches as patches\nfrom matplotlib.colors import ListedColormap\n\nimport pydicom\nfrom tqdm import tqdm\nfrom PIL import Image\nfrom glob import glob\nfrom pydicom.pixel_data_handlers.util import apply_voi_lut\n\nimport torch\nimport torchvision\nfrom torchvision import tv_tensors\nfrom torchvision.transforms import v2 as T\nfrom torchvision.tv_tensors import BoundingBoxes\nfrom torchvision.models.detection.faster_rcnn import FastRCNNPredictor","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-30T10:32:43.949574Z","iopub.execute_input":"2025-04-30T10:32:43.949762Z","iopub.status.idle":"2025-04-30T10:33:03.086308Z","shell.execute_reply.started":"2025-04-30T10:32:43.949746Z","shell.execute_reply":"2025-04-30T10:33:03.085702Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"device = torch.device('cuda') if torch.cuda.is_available() else torch.device('cpu')\ndevice","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-30T10:33:19.087085Z","iopub.execute_input":"2025-04-30T10:33:19.087357Z","iopub.status.idle":"2025-04-30T10:33:19.184345Z","shell.execute_reply.started":"2025-04-30T10:33:19.087338Z","shell.execute_reply":"2025-04-30T10:33:19.183558Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"MODEL_DIR='/kaggle/working/models'\nPRETRAINED_MODEL_FILE = \"/kaggle/input/fasterrcnn_resnet50_fpn_v2_colab/pytorch/default/1/model_dict.pt\"\nCROP_DIR = '/kaggle/working/crops'\nDATA_DIR = \"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/\"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-30T10:33:23.581535Z","iopub.execute_input":"2025-04-30T10:33:23.581834Z","iopub.status.idle":"2025-04-30T10:33:23.585570Z","shell.execute_reply.started":"2025-04-30T10:33:23.581795Z","shell.execute_reply":"2025-04-30T10:33:23.584927Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_df = pd.read_csv(DATA_DIR + \"train.csv\")\ntrain_lcoord_df = pd.read_csv(DATA_DIR + \"train_label_coordinates.csv\")\ntrain_sdesc_df = pd.read_csv(DATA_DIR + \"train_series_descriptions.csv\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-30T10:33:37.092078Z","iopub.execute_input":"2025-04-30T10:33:37.092351Z","iopub.status.idle":"2025-04-30T10:33:37.328601Z","shell.execute_reply.started":"2025-04-30T10:33:37.092331Z","shell.execute_reply":"2025-04-30T10:33:37.327870Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Map file paths\nfrom glob import glob\nstudy_paths = glob(f\"{DATA_DIR}train_images/*\")\nimages_dict = {\n    \"study_id\" : [],\n    \"series_id\" : [],\n    \"instance_number\" : [],\n    \"image_path\": []\n}\n\nseries_instance_count = {\n    \"study_id\" : [],\n    \"series_id\" : [],\n    \"instance_count\" : []\n}\nfor i, study_path in enumerate(tqdm(study_paths)):\n    instance_study_count = 0\n    study_id = study_path.split(\"/\")[-1]\n    series_paths = glob(f\"{study_path}/*\")\n    for series_path in series_paths:\n        instance_count = 0\n        series_id = series_path.split(\"/\")[-1]\n        instance_paths = glob(f\"{series_path}/*\")\n        for instance_path in instance_paths:\n            instance_count+=1\n            instance_id = instance_path.split(\"/\")[-1].split(\".\")[0]\n\n            images_dict[\"study_id\"].append(int(study_id))\n            images_dict[\"series_id\"].append(int(series_id))\n            images_dict[\"instance_number\"].append(int(instance_id.split(\" \")[0]))\n            images_dict[\"image_path\"].append(instance_path)\n\n        series_instance_count[\"study_id\"].append(int(study_id))\n        series_instance_count[\"series_id\"].append(int(series_id))\n        series_instance_count[\"instance_count\"].append(\n          instance_count\n        )\n\n\nimages_df = pd.DataFrame(images_dict)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-30T10:33:41.841691Z","iopub.execute_input":"2025-04-30T10:33:41.842421Z","iopub.status.idle":"2025-04-30T10:36:05.257117Z","shell.execute_reply.started":"2025-04-30T10:33:41.842395Z","shell.execute_reply":"2025-04-30T10:36:05.256297Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"images_desc_df = pd.merge(images_df,\n  train_sdesc_df,\n  on=['study_id','series_id']\n  )\nimages_desc_df.head()\n\n# Merge to add condition, level, x, y \nfull_train_df = pd.merge(\n    train_lcoord_df,\n    images_desc_df,\n    on=['study_id','series_id', 'instance_number']\n)\n\n# Filter Sagittal T2/STIR images\ndf = full_train_df[(full_train_df['series_description']=='Sagittal T2/STIR')][['study_id', 'series_id', 'instance_number', 'image_path', 'x', 'y', 'level']]\n\n# Drop duplicate levels in every series.\n## One (X, Y) per level is enough!\n## It will be transported to other instances if necessary with little impact to performance.\ndf = df.drop_duplicates(subset=['study_id', 'series_id', 'level'])\n\n\n# Replace L1/L2 format to L1_L2 to prevent path conflict\ndf['level'] = df['level'].apply(lambda x: x.replace(\"/\", \"_\"))\ndf","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-30T10:36:11.022247Z","iopub.execute_input":"2025-04-30T10:36:11.022515Z","iopub.status.idle":"2025-04-30T10:36:11.175550Z","shell.execute_reply.started":"2025-04-30T10:36:11.022495Z","shell.execute_reply":"2025-04-30T10:36:11.174879Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Load image as tensor\ndef load_image_from_path(path):\n    img = pydicom.dcmread(path)\n    img = img.pixel_array\n    img = (img - img.min()) / (img.max() - img.min() +1e-6) * 255 # Pixel value between 0-255\n    img = tv_tensors.Image(img) # [CHANNEL, HEIGHT, WIDTH]\n    return img.double()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-30T10:36:19.661960Z","iopub.execute_input":"2025-04-30T10:36:19.662227Z","iopub.status.idle":"2025-04-30T10:36:19.666928Z","shell.execute_reply.started":"2025-04-30T10:36:19.662207Z","shell.execute_reply":"2025-04-30T10:36:19.666213Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Disc mapping \nDISC_LABELS = {\n    \"L1_L2\": [1],\n    \"L2_L3\": [2],\n    \"L3_L4\": [3],\n    \"L4_L5\": [4],\n    \"L5_S1\": [5]\n}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-30T10:36:23.541425Z","iopub.execute_input":"2025-04-30T10:36:23.541749Z","iopub.status.idle":"2025-04-30T10:36:23.545614Z","shell.execute_reply.started":"2025-04-30T10:36:23.541720Z","shell.execute_reply":"2025-04-30T10:36:23.545094Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class RSNAMultipleBBoxesDataset(torch.utils.data.Dataset):\n    def __init__(self, df, transforms=None, limit=None):\n        # Limit for debugging\n        if limit:\n            df = df.iloc[0:limit]\n        \n        self.df = df\n        \n        # Unique image_paths\n        self.images_df = df[\n          ['study_id', 'series_id', 'instance_number', 'image_path']\n        ].drop_duplicates().reset_index(drop=True)\n\n\n        self.transforms = transforms\n\n    def __getitem__(self, idx):\n        image = self.images_df.iloc[idx]\n        target = {}\n\n        # Image\n        img = load_image_from_path(image['image_path'])\n        w_orig, h_orig = img.shape[-1], img.shape[-2]\n        target['img'] = img\n        target['image_id'] = idx\n        target['series_id'] = image['series_id']\n        target['study_id'] = image['study_id']\n        target['instance_number'] = image['instance_number']\n\n        # Transform\n        if self.transforms:\n            img = self.transforms(img)\n\n        w_resize, h_resize = img.shape[-1], img.shape[-2]\n        w_ratio = w_resize / w_orig\n        h_ratio = h_resize/ h_orig\n\n\n        target['boxes'] = []\n        target['area'] = []\n        target['labels'] = []\n        target['iscrowd'] = []\n        series_df = self.df[self.df['series_id'] == image['series_id']]\n\n        for i,row in series_df.iterrows():\n            # Label\n            target['labels'].append(DISC_LABELS[row['level']])\n\n            # BBox\n            ##############################################\n            # You can play with this block of code to \n            # modify the bounding box generation.\n            if row['level'] == 'L5_S1':\n                w = 70\n                h = 30\n            else:\n                w = 70\n                h = 20\n            # Here, I'm dislocating the box in such a way\n            # that upper level discs are closer to the \n            # bottom of it's box and lower level discs \n            # are closer to the top of it's box.\n            level = int(row['level'][1])\n            x0 = (row['x'])*w_ratio - w*5/6\n            x1 = (row['x'])*w_ratio + w*1/6\n            y0 = (row['y'])*h_ratio - h*(6-level)/6\n            y1 = (row['y'])*h_ratio + h*level/6\n            ##############################################\n\n            target['boxes'].append(\n              [x0, y0, x1, y1]\n            )\n            # Box area\n            target['area'].append((x1-x0) * (y1-y0))\n\n            # Instances with iscrowd=True will be ignored during evaluation.\n            target['iscrowd'].append(False)\n\n        target['area'] = torch.tensor(target['area'])\n        target['labels'] = torch.tensor(target['labels']).squeeze(dim=-1)\n        target['iscrowd'] = torch.tensor(target['iscrowd'])\n        target['boxes'] = BoundingBoxes(\n            target['boxes'],\n            format='XYXY',\n            dtype=torch.float32,\n            canvas_size=img.shape[-2:]\n        )\n        return img, target\n\n    def __len__(self):\n        return len(self.images_df)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-30T10:36:42.581892Z","iopub.execute_input":"2025-04-30T10:36:42.582389Z","iopub.status.idle":"2025-04-30T10:36:42.592045Z","shell.execute_reply.started":"2025-04-30T10:36:42.582367Z","shell.execute_reply":"2025-04-30T10:36:42.591438Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_transform():\n    transforms = []\n    transforms.append(T.Resize((250, 250), antialias=True))\n    transforms.append(T.ToDtype(torch.float, scale=True))\n    transforms.append(T.ToPureTensor())\n    return T.Compose(transforms)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-30T10:36:51.831867Z","iopub.execute_input":"2025-04-30T10:36:51.832124Z","iopub.status.idle":"2025-04-30T10:36:51.835959Z","shell.execute_reply.started":"2025-04-30T10:36:51.832105Z","shell.execute_reply":"2025-04-30T10:36:51.835267Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Trying data loader\ntmp_ds = RSNAMultipleBBoxesDataset(df, transforms=get_transform())\ntmp_dl = torch.utils.data.DataLoader(\n  tmp_ds,\n  batch_size=1,\n  shuffle=True,\n  collate_fn=utils.collate_fn\n)\n\nfig, ax = plt.subplots(nrows=5, ncols=1, figsize=(30,30))\nfor i, (img, t) in enumerate(tmp_dl):\n    if i==5:break\n    img = img[0]\n    t = t[0]\n    y = img.squeeze().numpy()\n    ax[i].imshow(y)\n    for j, box in enumerate(t['boxes']):\n        x0, y0, x1, y1 = box.numpy()\n        w = x1 - x0\n        h = y1 - y0\n        ax[i].add_patch(patches.Rectangle((x0, y0), w, h, linewidth=1, edgecolor='r', facecolor='none'))\ndel tmp_ds, tmp_dl","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-30T10:36:59.171714Z","iopub.execute_input":"2025-04-30T10:36:59.172447Z","iopub.status.idle":"2025-04-30T10:37:00.679240Z","shell.execute_reply.started":"2025-04-30T10:36:59.172421Z","shell.execute_reply":"2025-04-30T10:37:00.678389Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_model_box_predictor(load=\"latest\"):\n    # Load a model pre-trained on COCO\n    model = torchvision.models.detection.fasterrcnn_resnet50_fpn_v2(weights=\"DEFAULT\")\n\n    # Replace the classifier with a new one, that has Num_classes\n    num_classes = 6  # 5 classes (discs) + background\n\n    # Get number of input features for the classifier\n    in_features = model.roi_heads.box_predictor.cls_score.in_features\n\n    # Replace the pre-trained head with a new one\n    model.roi_heads.box_predictor = FastRCNNPredictor(in_features, num_classes)\n\n    if load:\n        # Get all matching files\n        \n        if load ==\"latest\":\n            file_pattern = f'{MODEL_DIR}/*/model_dict.pt'\n            files = glob(file_pattern)\n        elif os.path.isfile(load):\n            files = glob(load)\n        else:\n            files=[]\n\n    # Check if any files were found\n    if not files:\n        print(\"Model not found. Creating a new one...\")\n    else:\n        # Get the latest file based on modification time\n        latest_file = max(files, key=os.path.getmtime)\n        print(f\"Loading latest model: {latest_file}\")\n        model.load_state_dict(torch.load(latest_file))\n\n    return model","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-30T10:37:15.901569Z","iopub.execute_input":"2025-04-30T10:37:15.901869Z","iopub.status.idle":"2025-04-30T10:37:15.907843Z","shell.execute_reply.started":"2025-04-30T10:37:15.901846Z","shell.execute_reply":"2025-04-30T10:37:15.907280Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model = get_model_box_predictor(load=\"latest\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-30T10:37:27.731989Z","iopub.execute_input":"2025-04-30T10:37:27.732261Z","iopub.status.idle":"2025-04-30T10:37:31.692022Z","shell.execute_reply.started":"2025-04-30T10:37:27.732241Z","shell.execute_reply":"2025-04-30T10:37:31.691230Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Create dataset and dataloader\ndataset = RSNAMultipleBBoxesDataset(df ,get_transform())\ndata_loader = torch.utils.data.DataLoader(\n  dataset,\n  batch_size=2,\n  shuffle=True,\n  collate_fn=utils.collate_fn\n)\n\n# Get first input from dataloader\nimages, targets = next(iter(data_loader))\nimages = list(image for image in images)\ntargets = [{k: v for k, v in t.items()} for t in targets]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-30T10:37:37.101995Z","iopub.execute_input":"2025-04-30T10:37:37.102261Z","iopub.status.idle":"2025-04-30T10:37:37.153088Z","shell.execute_reply.started":"2025-04-30T10:37:37.102243Z","shell.execute_reply":"2025-04-30T10:37:37.152250Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Inference\nmodel.eval()\nwith torch.inference_mode():\n    predictions = model(images)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-30T10:37:44.411486Z","iopub.execute_input":"2025-04-30T10:37:44.412002Z","iopub.status.idle":"2025-04-30T10:37:53.163296Z","shell.execute_reply.started":"2025-04-30T10:37:44.411981Z","shell.execute_reply":"2025-04-30T10:37:53.162678Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Inspect output\nprint(predictions[0].keys())\nprint(predictions[0]['boxes'].shape)  # [100 boxes, 4 dim (x0, y0, x1, y1)]\nprint(predictions[0]['boxes'][0].tolist()) # (x0, y0, x1, y1) of the first box\nprint(predictions[0]['labels'][0]) # label of the first box\nprint(predictions[0]['scores'][0]) # score of the first box","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-30T10:38:06.931905Z","iopub.execute_input":"2025-04-30T10:38:06.932627Z","iopub.status.idle":"2025-04-30T10:38:06.938534Z","shell.execute_reply.started":"2025-04-30T10:38:06.932605Z","shell.execute_reply":"2025-04-30T10:38:06.937874Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# use our dataset and defined transformations\ntrain_data, test_data = train_test_split(df, shuffle=False)\n\ndataset = RSNAMultipleBBoxesDataset(train_data ,get_transform())\ndataset_test = RSNAMultipleBBoxesDataset(test_data ,get_transform())\n\n# define training and validation data loaders\ndata_loader = torch.utils.data.DataLoader(\n    dataset,\n    batch_size=10,\n    shuffle=True,\n    collate_fn=utils.collate_fn,\n    num_workers=os.cpu_count()\n)\n\ndata_loader_test = torch.utils.data.DataLoader(\n    dataset_test,\n    batch_size=1,\n    shuffle=False,\n    collate_fn=utils.collate_fn,\n    num_workers=os.cpu_count()\n)\n\nmodel = get_model_box_predictor(load=\"latest\")\n\n# move model to the right device\nmodel.to(device)\n\n# construct an optimizer\nparams = [p for p in model.parameters() if p.requires_grad]\noptimizer = torch.optim.Adam(\n    params,\n    lr=0.0001,\n)\n\n# and a learning rate scheduler\nlr_scheduler = torch.optim.lr_scheduler.StepLR(\n    optimizer,\n    step_size=3,\n    gamma=0.1\n)\n\n# train it for 3 epochs\nnum_epochs = 5\n\nfor epoch in range(num_epochs):\n    # train for one epoch, printing every 10 iterations\n    train_one_epoch(model, optimizer, data_loader, device, epoch, print_freq=10)\n    # update the learning rate\n    lr_scheduler.step()\n    # evaluate on the test dataset\n    evaluate(model, data_loader_test, device=device)\n\n    now = datetime.now()\n    now = now.strftime(\"%Y_%m_%d__%H_%M_%S\")\n    dirname = f'{MODEL_DIR}/{now}'\n    os.makedirs(dirname, exist_ok=True,)\n    fname = f'{dirname}/model_dict.pt'\n    torch.save(model.state_dict(), fname)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-30T10:38:16.401646Z","iopub.execute_input":"2025-04-30T10:38:16.401962Z","iopub.status.idle":"2025-04-30T11:34:10.011077Z","shell.execute_reply.started":"2025-04-30T10:38:16.401941Z","shell.execute_reply":"2025-04-30T11:34:10.010146Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"LABELS_DICT = {\n    1: \"L1_L2\",\n    2: \"L2_L3\",\n    3: \"L3_L4\",\n    4: \"L4_L5\",\n    5: \"L5_S1\"\n}\n\ndef get_best_boxes(pred):\n    best_boxes = {}\n\n    for box, label, score in zip(pred['boxes'], pred['labels'], pred['scores']):\n        if label.item() not in best_boxes or score > best_boxes[label.item()]['score']:\n            best_boxes[label.item()] = {'box': box.tolist(), 'score': score.item()}\n\n    result = {\n        'boxes': [entry['box'] for entry in best_boxes.values()],\n        'labels': list(best_boxes.keys()),\n        'scores': [entry['score'] for entry in best_boxes.values()]\n    }\n\n    return result\n\ndef plot_prediction(x, pred):\n    x = x[0, :]\n    fig, ax = plt.subplots(nrows=1, ncols=1, figsize=(12,8))\n    ax.imshow(x, cmap=\"bone\")\n    pred = get_best_boxes(pred)\n    for i in range(len(pred['boxes'])):\n        x0, y0, x1, y1 = pred['boxes'][i]\n        label = pred['labels'][i]\n        score = pred['scores'][i]\n        h = y1 - y0\n        w = x1 - x0\n        ax.add_patch(patches.Rectangle((x0, y0), w, h, linewidth=1, edgecolor='r', facecolor='none'))\n        ax.text(x0+w+10, y0+h/2, f\"{LABELS_DICT[label]} ({'{:.2f}'.format(score)})\", color='r',fontsize=14)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-30T11:34:31.284127Z","iopub.execute_input":"2025-04-30T11:34:31.284414Z","iopub.status.idle":"2025-04-30T11:34:31.292073Z","shell.execute_reply.started":"2025-04-30T11:34:31.284392Z","shell.execute_reply":"2025-04-30T11:34:31.291383Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model = get_model_box_predictor(load=PRETRAINED_MODEL_FILE)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-30T11:34:41.144050Z","iopub.execute_input":"2025-04-30T11:34:41.144315Z","iopub.status.idle":"2025-04-30T11:34:43.255242Z","shell.execute_reply.started":"2025-04-30T11:34:41.144298Z","shell.execute_reply":"2025-04-30T11:34:43.254679Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_data, test_data = train_test_split(df, shuffle=False)\ndataset_test = RSNAMultipleBBoxesDataset(test_data, get_transform())\n\ndata_loader_test = torch.utils.data.DataLoader(\n    dataset_test,\n    batch_size=5,\n    shuffle=False,\n    collate_fn=utils.collate_fn\n)\n\nimages, targets = next(iter(data_loader_test))\nimages = list(image.to(device) for image in images)\ntargets = [{k: v.to(device) if isinstance(v, torch.Tensor) else v for k, v in t.items()} for t in targets]\n\nmodel.to(device)\nmodel.eval()\nwith torch.inference_mode():\n    predictions = model(images)\n\nfor i in range(len(images)):\n    plot_prediction(images[i].cpu(), predictions[i])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-30T11:34:53.254044Z","iopub.execute_input":"2025-04-30T11:34:53.254309Z","iopub.status.idle":"2025-04-30T11:34:55.837509Z","shell.execute_reply.started":"2025-04-30T11:34:53.254291Z","shell.execute_reply":"2025-04-30T11:34:55.836680Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def crop_bbox(image, bbox):\n    x0, y0, x1, y1 = bbox\n\n    cropped_img = torchvision.transforms.functional.crop(\n        image,\n        top=round(int(y0)),\n        left=round(int(x0)),\n        height=round(int(y1 - y0)),\n        width=round(int(x1 - x0))\n    )\n    return cropped_img\n\n\ndef plot_crop(image, bboxes):\n    fig, ax = plt.subplots(nrows=5, ncols=1, figsize=(4,3))\n    plt.subplots_adjust(top=2)\n\n    for i in range(len(bboxes['boxes'])):\n        label_i = bboxes['labels'][i] - 1\n        label = LABELS_DICT[label_i + 1]\n        score = bboxes['scores'][i]\n        bbox = bboxes['boxes'][i]\n\n        cropped_img = crop_bbox(image, bbox)\n        cropped_img = cropped_img[0, :]\n\n        ax[label_i].set_axis_off()\n        ax[label_i].imshow(cropped_img, cmap=\"bone\")\n        ax[label_i].set_title(f\"{label} ({'{:.2f}'.format(score)})\")\n        \n\ndef save_crop(image, bboxes, target):\n    series_id = target['series_id']\n    study_id = target['study_id']\n    instance_number = target['instance_number']\n\n\n    for i in range(len(bboxes['boxes'])):\n        label = LABELS_DICT[bboxes['labels'][i]]\n\n        dirname = f'{CROP_DIR}/train_images/{series_id}/{study_id}/{label}'\n        os.makedirs(dirname, exist_ok=True)\n        filepath = os.path.join(dirname, f'{instance_number}.pt')\n\n        bbox = bboxes['boxes'][i]\n\n        cropped_img = crop_bbox(image, bbox)\n        torch.save(cropped_img, filepath)\n\n    return","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-30T11:35:09.164099Z","iopub.execute_input":"2025-04-30T11:35:09.164842Z","iopub.status.idle":"2025-04-30T11:35:09.171920Z","shell.execute_reply.started":"2025-04-30T11:35:09.164800Z","shell.execute_reply":"2025-04-30T11:35:09.171365Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"LIMIT=5 # Alter limit if needed\n\ndataset = RSNAMultipleBBoxesDataset(df ,get_transform(), limit=LIMIT)\ndata_loader = torch.utils.data.DataLoader(\n    dataset,\n    batch_size=10,\n    shuffle=True,\n    collate_fn=utils.collate_fn\n)\n\nmodel.eval()\nwith torch.inference_mode():\n    for i, (images, targets) in enumerate(tqdm(data_loader)):\n        images = list(image.to(device) for image in images)\n        targets = [{k: v.to(device) if isinstance(v, torch.Tensor) else v for k, v in t.items()} for t in targets]\n        predictions = model(images)\n\n    for i in range(len(images)):\n        bboxes = get_best_boxes(predictions[i])\n        # plot_prediction(images[i], predictions[i])\n        # plot_crop(images[i], bboxes)\n        save_crop(images[i].cpu(), bboxes, targets[i])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-30T11:35:16.793765Z","iopub.execute_input":"2025-04-30T11:35:16.794318Z","iopub.status.idle":"2025-04-30T11:35:17.037911Z","shell.execute_reply.started":"2025-04-30T11:35:16.794294Z","shell.execute_reply":"2025-04-30T11:35:17.037188Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def load_crop(study_id, series_id, plot=True):\n    filepattern = f'{CROP_DIR}/train_images/{study_id}/{series_id}/**/*.pt'\n    files = glob(filepattern, recursive=True)\n    crops = []\n    for file in files:\n        crop = torch.load(file)\n        crops.append(crop)\n        if plot:\n            plt.imshow(crop[0], cmap=\"bone\")\n            plt.show()\n    return crops","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-30T11:35:39.784201Z","iopub.execute_input":"2025-04-30T11:35:39.784716Z","iopub.status.idle":"2025-04-30T11:35:39.789446Z","shell.execute_reply.started":"2025-04-30T11:35:39.784694Z","shell.execute_reply":"2025-04-30T11:35:39.788858Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"loaded_crops = load_crop(702807833, 4003253)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-30T11:35:46.163691Z","iopub.execute_input":"2025-04-30T11:35:46.163979Z","iopub.status.idle":"2025-04-30T11:35:46.679566Z","shell.execute_reply.started":"2025-04-30T11:35:46.163959Z","shell.execute_reply":"2025-04-30T11:35:46.678955Z"}},"outputs":[],"execution_count":null}]}