{"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":"nvidiaTeslaT4","dataSources":[{"sourceId":71549,"databundleVersionId":8561470,"sourceType":"competition"},{"sourceId":125367,"sourceType":"modelInstanceVersion","modelInstanceId":91818,"modelId":116030},{"sourceId":126057,"sourceType":"modelInstanceVersion","modelInstanceId":91818,"modelId":116030},{"sourceId":130113,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":91818,"modelId":116030}],"dockerImageVersionId":30762,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# In one file for kaggle :3\nimport numpy as np\nimport pandas as pd\nimport tqdm\nimport glob\nimport cv2\nimport pydicom\nfrom scipy.ndimage import label, center_of_mass\n\nimport time\n\nimport matplotlib.pyplot as plt\n\nimport torch\nimport torch.nn as nn\nimport timm\ntorch.set_grad_enabled(False) # remove grad for the script\n\n##########################################################\n#\n#   USEFULL\n#\n##########################################################\n\ndef get_instance(path):\n    return int(path.split(\"/\")[-1].split('.')[0])\n\ndef load_torch_script_model(path, device):\n    return torch.jit.load(path, map_location=device).eval()\n\ndef z_score_normalize(scan):\n    mean, std = np.mean(scan), np.std(scan)\n    return (scan - mean) / std\n\n##########################################################\n#\n#   SLICES SELECTION INFERENCE\n#\n##########################################################\n\ndef get_top_k_consecutive(preds, k=3):\n    max_sum = -float('inf')\n    max_idx = -1\n    \n    for i in range(len(preds) - k + 1):\n        current_sum = preds[i:i + k].sum().item()\n        if current_sum > max_sum:\n            max_sum = current_sum\n            max_idx = i\n    \n    max_consecutive_indices = list(range(max_idx, max_idx + k))\n    return max_consecutive_indices\n\ndef load_dicom_series(dicom_files):\n    dicom_files = sorted(dicom_files, key=get_instance)\n    slices = [pydicom.dcmread(f).pixel_array for f in dicom_files]\n    slices_resized = [cv2.resize(s, (224, 224), interpolation=cv2.INTER_LINEAR) for s in slices]\n    return np.array(slices_resized).transpose(1, 2, 0)\n\ndef get_best_slice_selection(config, pathes):\n    volume = load_dicom_series(pathes)\n    normalized_volume = z_score_normalize(volume)\n    images = torch.tensor(normalized_volume.transpose(2, 0, 1)).unsqueeze(1).float()\n\n    preds = config['model_slice_selection'](images.to(config[\"device\"]).unsqueeze(0)).squeeze()\n    preds_sum = torch.sum(preds, dim=0)\n    \n    best_slices_by_level_top3 = []\n    best_slices_by_level_top5 = []\n    \n    for level in range(preds.shape[0]):\n        pred_level = preds[level, :]\n        \n        top_3_indices = get_top_k_consecutive(pred_level, 3)\n        top_5_indices = get_top_k_consecutive(pred_level, 5)\n        \n        best_slices_by_level_top3.append({\n            \"pathes\": [pathes[i] for i in top_3_indices],\n            \"values\": [pred_level[i].item() for i in top_3_indices]\n        })\n        \n        best_slices_by_level_top5.append({\n            \"pathes\": [pathes[i] for i in top_5_indices],\n            \"values\": [pred_level[i].item() for i in top_5_indices]\n        })\n    \n    best_overall_indices = get_top_k_consecutive(preds_sum, 3)\n    best_slices_overall = {\n        \"pathes\": [pathes[i] for i in best_overall_indices],\n        \"values\": [preds_sum[i].item() for i in best_overall_indices]\n    }\n    \n    return best_slices_by_level_top3, best_slices_by_level_top5, best_slices_overall\n\ndef find_best_slices_across_series(best_slices_per_series, num_levels=5):\n\n    best_slices_per_level_top3 = [None] * num_levels\n    best_slices_per_level_top5 = [None] * num_levels\n    for level in range(num_levels):\n        best_avg_top3 = -1\n        best_avg_top5 = -1\n\n        for slices in best_slices_per_series.values():\n            avg_top3 = np.mean(slices[\"best_per_level_top3\"][level]['values'])\n            avg_top5 = np.mean(slices[\"best_per_level_top5\"][level]['values'])\n            \n            if avg_top3 > best_avg_top3:\n                best_avg_top3 = avg_top3\n                best_slices_per_level_top3[level] = slices[\"best_per_level_top3\"][level]['pathes']\n            \n            if avg_top5 > best_avg_top5:\n                best_avg_top5 = avg_top5\n                best_slices_per_level_top5[level] = slices[\"best_per_level_top5\"][level]['pathes']\n\n    best_slices_overall = []\n    best_avg_overall = -1\n    for slices in best_slices_per_series.values():\n        avg_overall = np.mean(slices[\"best_overall\"]['values'])\n        if avg_overall > best_avg_overall:\n            best_avg_overall = avg_overall\n            best_slices_overall = slices[\"best_overall\"]['pathes']\n\n    return best_slices_per_level_top3, best_slices_per_level_top5, best_slices_overall\n\ndef get_slices_to_use(study_id, series_ids, config):\n    slices_by_series = {\n        s_id: sorted(glob.glob(f\"{config['input_images_folder']}/{study_id}/{s_id}/*.dcm\"), key=get_instance)\n        for s_id in series_ids\n    }\n\n    best_slices_per_series = {}\n    for s_id in slices_by_series:\n        best_per_level_top3, best_per_level_top5, best_overall = get_best_slice_selection(config, slices_by_series[s_id])\n        best_slices_per_series[s_id] = {\n            \"best_per_level_top3\": best_per_level_top3,\n            \"best_per_level_top5\": best_per_level_top5,\n            \"best_overall\": best_overall\n        }\n\n    return find_best_slices_across_series(best_slices_per_series)\n\n\n##########################################################\n#\n#   SEGMENTATION INFERENCE\n#\n##########################################################\n\ndef seg_load_dicom_series(dicom_files):\n    dicom_files = sorted(dicom_files, key=lambda x: get_instance(x))\n    slices = [cv2.resize(pydicom.dcmread(f).pixel_array, (384, 384), interpolation=cv2.INTER_LINEAR) for f in dicom_files]\n    return np.array(slices).transpose(1, 2, 0)\n\ndef find_center_of_largest_activation(mask: torch.tensor) -> tuple:\n    mask_np = (mask > 0.5).float().detach().cpu().numpy()\n    labeled_mask, num_features = label(mask_np)\n    if num_features == 0:\n        return None\n    largest_component = np.argmax(np.bincount(labeled_mask.ravel())[1:]) + 1  # Ignore the background\n    largest_component_center = center_of_mass(labeled_mask == largest_component)\n    center_coords = tuple(map(int, largest_component_center))\n    return (center_coords[1] / mask.shape[1], center_coords[0] / mask.shape[0])  # x, y normalized\n\ndef normalize_and_prepare_images(slices_path: list, config: dict) -> torch.tensor:\n    volume = seg_load_dicom_series(slices_path)\n    normalized_volume = z_score_normalize(volume)\n    return torch.tensor(normalized_volume.transpose(2, 0, 1)).float().to(config[\"device\"])\n\ndef get_segmentation_input(slices_to_use: list, config: dict):\n    if config['segmentation_slice_selection'] == \"best_overall\":\n        return normalize_and_prepare_images(slices_to_use['best_overall'], config)\n    elif config['segmentation_slice_selection'] == \"best_by_level\":\n        images = np.zeros((5, 3, *config['segmentation_input_shape']))\n        for level in range(5):\n            images[level] = normalize_and_prepare_images(slices_to_use['best3_by_level'][level], config).cpu().numpy()\n        return torch.tensor(images).float().to(config[\"device\"])\n    return None\n\ndef process_positions(masks: torch.tensor, positions: torch.tensor, config: dict) -> list:\n    position_by_level = []\n    for i in range(5):\n        activation_center = find_center_of_largest_activation(masks[i])\n        if activation_center is not None:\n            position_by_level.append(activation_center)\n        else:\n            x, y = positions[i][0].item(), positions[i][1].item()\n            print(f\"Segmentation fails on {config['condition']} level {i}, using coordinates ({x}, {y})\")\n            position_by_level.append((x, y))\n    return position_by_level\n\ndef get_position_by_level(slices_to_use: list, config: dict) -> dict:\n    inputs = get_segmentation_input(slices_to_use, config)\n    if config['segmentation_slice_selection'] == \"best_overall\":\n        masks, positions = config[\"model_segmentation\"](inputs.unsqueeze(0))  # model predicts 5 levels\n        positions = positions.view(5, 2)\n    else:\n        masks, positions = config[\"model_segmentation\"](inputs)  # model predicts 1 level, batch across levels\n    masks = masks.squeeze(1) if config['segmentation_slice_selection'] != \"best_overall\" else masks.squeeze()\n    return process_positions(masks, positions, config)\n\n##########################################################\n#\n#   CLASSIFICATION INFERENCE\n#\n##########################################################\n\ndef extract_centered_square_with_padding(array, center_x, center_y, sizeX, sizeY):\n    start_x = max(center_x - (sizeX // 2), 0)\n    end_x = min(center_x + (sizeX // 2), array.shape[0])\n    start_y = max(center_y - (sizeY // 2), 0)\n    end_y = min(center_y + (sizeY // 2), array.shape[1])\n    \n    out_start_x = (sizeX // 2) - (center_x - start_x)\n    out_end_x = out_start_x + (end_x - start_x)\n    out_start_y = (sizeY // 2) - (center_y - start_y)\n    out_end_y = out_start_y + (end_y - start_y)\n    \n    square = np.zeros((sizeX, sizeY), dtype=array.dtype)\n    square[out_start_x:out_end_x, out_start_y:out_end_y] = array[start_x:end_x, start_y:end_y]\n    return square\n\ndef cut_crops(slices_path, x, y, crop_size, image_resize, use_flip):\n    output_crops = np.zeros((len(slices_path), 128, 128))\n    for k, slice_path in enumerate(slices_path):\n        pixel_array = pydicom.dcmread(slice_path).pixel_array.astype(np.float32)\n        pixel_array = cv2.resize(pixel_array, image_resize, interpolation=cv2.INTER_LINEAR)\n        crop = extract_centered_square_with_padding(pixel_array, y, x, *crop_size) # x y reversed in array\n        if use_flip:\n            crop = cv2.flip(crop, 1)\n        crop = cv2.resize(crop, (128, 128), interpolation=cv2.INTER_LINEAR)\n        output_crops[k] = crop\n    return output_crops\n\ndef get_crops_by_level(slices_to_use, position_by_level, config):\n    crops_output = torch.empty((5, 5, 3, 128, 128))\n    for level, (slices_, position) in enumerate(zip(slices_to_use[\"best5_by_level\"], position_by_level)):\n        px = int(position[0] * config['classification_resize_image'][1])\n        py = int(position[1] * config['classification_resize_image'][0])\n        crops_level = cut_crops(slices_, px, py, config['classification_input_size'], config['classification_resize_image'], config['classification_use_flip'])\n        crops_level = z_score_normalize(crops_level)\n        crops_output[level] = torch.tensor(crops_level).unsqueeze(1).expand(5, 3, 128, 128)\n    crops_output = crops_output.float().to(config[\"device\"])\n    return crops_output\n\ndef get_classification(crops: torch.tensor, config):\n    preds = config['model_classification'](crops)\n    preds = torch.softmax(preds, dim=1)\n    return preds\n\n##########################################################\n#\n#   MAIN INFERENCE\n#\n##########################################################\n\ndef predict_lumbar(df_description: pd.DataFrame, config: dict, study_id: int) -> list:\n    \n    #print(\"-\" * 50)\n    #print(config['condition'])\n    \n    try:\n        series_ids = df_description[(df_description['study_id'] == study_id)\n                                & (df_description['series_description'] == config['description'])]['series_id'].to_list()\n        \n        #s = time.time()\n        \n        b3, b5, bo = get_slices_to_use(study_id, series_ids, config)\n        slices_to_use = dict(\n            best3_by_level=b3,\n            best5_by_level=b5,\n            best_overall=bo\n        )\n        \n        #print(\"Get best slices:\", time.time() - s)\n        #s = time.time()\n        \n        positions_by_level = get_position_by_level(slices_to_use, config)\n        \n        #print(\"Get positions by level:\", time.time() - s)\n        #s = time.time()\n        \n        crops_by_level = get_crops_by_level(slices_to_use, positions_by_level, config)\n        \n        \n        #print(\"Get crops:\", time.time() - s)\n        #s = time.time()\n        \n        classification_results = get_classification(crops_by_level, config)\n        \n        #print(\"Get classification:\", time.time() - s)\n\n    except Exception as e:\n        print(f\"Error {study_id} {config['condition']}:\", e)\n        classification_results = None\n\n    predictions = list()\n    row_id = f\"{study_id}_{config['condition'].lower().replace(' ', '_')}\"\n    for level_int, level_str in enumerate(['l1_l2', 'l2_l3', 'l3_l4', 'l4_l5', 'l5_s1']):\n        predictions.append(dict(\n            row_id = f\"{row_id}_{level_str}\",\n            normal_mild = classification_results[level_int][0].item() if classification_results is not None else 1/3,\n            moderate = classification_results[level_int][1].item() if classification_results is not None else 1/3,\n            severe = classification_results[level_int][2].item() if classification_results is not None else 1/3,\n        ))\n\n    return predictions\n\ndef configure_inference(slice_model_path, seg_model_path, class_model_path, input_folder, description, condition, class_input_size, class_resize_image, seg_mode, class_flip, device):\n    return dict(\n        device=device,\n\n        # slices infos\n        model_slice_selection=load_torch_script_model(slice_model_path, device),\n        slice_selection_input_shape=(224, 224),\n\n        # segmentation infos\n        segmentation_input_shape=(384, 384),\n        segmentation_slice_selection=seg_mode,\n        model_segmentation=load_torch_script_model(seg_model_path, device),\n\n        # classification infos\n        classification_input_size=class_input_size,\n        classification_resize_image=class_resize_image,\n        classification_sequence_lenght=5,\n        classification_use_flip=class_flip,\n        model_classification=load_torch_script_model(class_model_path, device),\n        \n        # general infos\n        description=description,\n        condition=condition,\n        input_images_folder=input_folder,\n    )\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-10-08T17:05:42.001908Z","iopub.execute_input":"2024-10-08T17:05:42.002238Z","iopub.status.idle":"2024-10-08T17:05:46.182285Z","shell.execute_reply.started":"2024-10-08T17:05:42.002205Z","shell.execute_reply":"2024-10-08T17:05:46.181521Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def compute_pipeline(input_images_folder, description_file, nb_studies_id=None):\n    df_description = pd.read_csv(description_file)\n    final_predictions = []\n\n    tasks = [\n        {\n            \"class_model_path\": \"/kaggle/input/lumbar_pipeline/pytorch/v0/23/classification_st1_left.ts\",\n            \"slice_model_path\": \"/kaggle/input/lumbar_pipeline/pytorch/v0/15/slice_selector_st1_left_metamodel.ts\",\n            \"seg_model_path\": \"/kaggle/input/lumbar_pipeline/pytorch/v0/14/segmentation_st1_left.ts\",\n            \"description\": \"Sagittal T1\",\n            \"condition\": \"Left Neural Foraminal Narrowing\",\n            \"class_input_size\": (80, 120),\n            \"class_resize_image\": (640, 640),\n            \"segmentation_slice_selection\": \"best_overall\",\n            \"class_flip\": False,\n            \"device\": torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\n        },\n        {\n            \"class_model_path\": \"/kaggle/input/lumbar_pipeline/pytorch/v0/23/classification_st1_right.ts\",\n            \"slice_model_path\": \"/kaggle/input/lumbar_pipeline/pytorch/v0/15/slice_selector_st1_right_metamodel.ts\",\n            \"seg_model_path\": \"/kaggle/input/lumbar_pipeline/pytorch/v0/14/segmentation_st1_right.ts\",\n            \"description\": \"Sagittal T1\",\n            \"condition\": \"Right Neural Foraminal Narrowing\",\n            \"class_input_size\": (80, 120),\n            \"class_resize_image\": (640, 640),\n            \"segmentation_slice_selection\": \"best_overall\",\n            \"class_flip\": False,\n            \"device\": torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\n        },\n        {\n            \"class_model_path\": \"/kaggle/input/lumbar_pipeline/pytorch/v0/23/classification_st2.ts\",\n            \"slice_model_path\": \"/kaggle/input/lumbar_pipeline/pytorch/v0/15/slice_selector_st2_metamodel.ts\",\n            \"seg_model_path\": \"/kaggle/input/lumbar_pipeline/pytorch/v0/14/segmentation_st2.ts\",\n            \"description\": \"Sagittal T2/STIR\",\n            \"condition\": \"Spinal Canal Stenosis\",\n            \"class_input_size\": (80, 120),\n            \"class_resize_image\": (640, 640),\n            \"segmentation_slice_selection\": \"best_overall\",\n            \"class_flip\": False,\n            \"device\": torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\n        },\n        {\n            \"class_model_path\": \"/kaggle/input/lumbar_pipeline/pytorch/v0/23/classification_axial_left.ts\",\n            \"slice_model_path\": \"/kaggle/input/lumbar_pipeline/pytorch/v0/15/slice_selector_ax_left_metamodel.ts\",\n            \"seg_model_path\": \"/kaggle/input/lumbar_pipeline/pytorch/v0/14/segmentation_ax_left.ts\",\n            \"description\": \"Axial T2\",\n            \"condition\": \"Left Subarticular Stenosis\",\n            \"class_input_size\": (64, 64),\n            \"class_resize_image\": (640, 640),\n            \"segmentation_slice_selection\": \"best_by_level\",\n            \"class_flip\": False,\n            \"device\": torch.device(\"cuda:1\" if torch.cuda.is_available() else \"cpu\")\n        },\n        {\n            \"class_model_path\": \"/kaggle/input/lumbar_pipeline/pytorch/v0/23/classification_axial_right.ts\",\n            \"slice_model_path\": \"/kaggle/input/lumbar_pipeline/pytorch/v0/15/slice_selector_ax_right_metamodel.ts\",\n            \"seg_model_path\": \"/kaggle/input/lumbar_pipeline/pytorch/v0/14/segmentation_ax_right.ts\",\n            \"description\": \"Axial T2\",\n            \"condition\": \"Right Subarticular Stenosis\",\n            \"class_input_size\": (64, 64),\n            \"class_resize_image\": (640, 640),\n            \"segmentation_slice_selection\": \"best_by_level\",\n            \"class_flip\": False,\n            \"device\": torch.device(\"cuda:1\" if torch.cuda.is_available() else \"cpu\")\n        }\n    ]\n    studies_id = df_description[\"study_id\"].unique()\n    task_configs = []\n    for task in tasks:\n        print(f'Loading models and configuring for task: {task[\"condition\"]}')\n        config = configure_inference(\n            task[\"slice_model_path\"],\n            task[\"seg_model_path\"],\n            task[\"class_model_path\"],\n            input_images_folder,\n            task[\"description\"],\n            task[\"condition\"],\n            task[\"class_input_size\"],\n            task[\"class_resize_image\"],\n            task[\"segmentation_slice_selection\"],\n            task['class_flip'],\n            task['device'],\n        )\n        task_configs.append(config)\n\n    if nb_studies_id is None:\n        nb_studies_id = len(studies_id)\n\n    print(f\"Launch pipeline on {nb_studies_id} studies !\")\n    for study_id in tqdm.tqdm(studies_id[:nb_studies_id], desc=\"Predicting for each study\"):\n        for config in task_configs:\n            _predictions = predict_lumbar(df_description, config, study_id)\n            final_predictions.extend(_predictions)\n    return pd.DataFrame(final_predictions)","metadata":{"execution":{"iopub.status.busy":"2024-10-08T17:05:46.187188Z","iopub.execute_input":"2024-10-08T17:05:46.187597Z","iopub.status.idle":"2024-10-08T17:05:46.205224Z","shell.execute_reply.started":"2024-10-08T17:05:46.187552Z","shell.execute_reply":"2024-10-08T17:05:46.204344Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"input_images_folder = \"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/test_images/\"\ndescription_file = \"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/test_series_descriptions.csv\"\ndf = compute_pipeline(input_images_folder, description_file, None)\ndf.to_csv(\"submission.csv\", index=False)\ndf","metadata":{"execution":{"iopub.status.busy":"2024-10-08T17:05:46.206488Z","iopub.execute_input":"2024-10-08T17:05:46.207282Z","iopub.status.idle":"2024-10-08T17:07:02.464184Z","shell.execute_reply.started":"2024-10-08T17:05:46.207228Z","shell.execute_reply":"2024-10-08T17:07:02.463237Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}