{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.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":71549,"databundleVersionId":8561470,"sourceType":"competition"},{"sourceId":9147692,"sourceType":"datasetVersion","datasetId":5483616},{"sourceId":9235934,"sourceType":"datasetVersion","datasetId":5586514},{"sourceId":9238150,"sourceType":"datasetVersion","datasetId":5588015},{"sourceId":9409759,"sourceType":"datasetVersion","datasetId":5577143},{"sourceId":9517627,"sourceType":"datasetVersion","datasetId":5481180},{"sourceId":9577221,"sourceType":"datasetVersion","datasetId":5481188},{"sourceId":9559645,"sourceType":"datasetVersion","datasetId":5825392},{"sourceId":192715298,"sourceType":"kernelVersion"},{"sourceId":193161758,"sourceType":"kernelVersion"},{"sourceId":193298353,"sourceType":"kernelVersion"},{"sourceId":193335312,"sourceType":"kernelVersion"},{"sourceId":193417638,"sourceType":"kernelVersion"},{"sourceId":199491752,"sourceType":"kernelVersion"},{"sourceId":199713407,"sourceType":"kernelVersion"},{"sourceId":114536,"sourceType":"modelInstanceVersion","modelInstanceId":64905,"modelId":89293}],"dockerImageVersionId":30747,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install -q --no-index --find-links /kaggle/input/ultralytics ultralytics","metadata":{"execution":{"iopub.status.busy":"2024-10-08T15:00:13.044328Z","iopub.execute_input":"2024-10-08T15:00:13.045162Z","iopub.status.idle":"2024-10-08T15:00:28.338623Z","shell.execute_reply.started":"2024-10-08T15:00:13.045120Z","shell.execute_reply":"2024-10-08T15:00:28.337437Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Install the ultralytics package from GitHub\n#!pip install git+https://github.com/ambisinistra/ultralyticsRSNA@non-nms","metadata":{"execution":{"iopub.status.busy":"2024-10-08T15:00:28.340977Z","iopub.execute_input":"2024-10-08T15:00:28.341438Z","iopub.status.idle":"2024-10-08T15:00:28.345743Z","shell.execute_reply.started":"2024-10-08T15:00:28.341405Z","shell.execute_reply":"2024-10-08T15:00:28.344829Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#!unzip -q -d / /kaggle/input/lsdc-get-all-images/images.zip","metadata":{"execution":{"iopub.status.busy":"2024-10-08T15:00:28.346973Z","iopub.execute_input":"2024-10-08T15:00:28.347327Z","iopub.status.idle":"2024-10-08T15:00:28.355495Z","shell.execute_reply.started":"2024-10-08T15:00:28.347297Z","shell.execute_reply":"2024-10-08T15:00:28.354761Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\"\"\"\nPatch Ultralytics\n\"\"\"\nimport ultralytics.engine.results\nimport ultralytics.utils.ops\n\ndef init(self, boxes, orig_shape) -> None:\n    \"\"\"\n    Initialize the Boxes class with detection box data and the original image shape.\n\n    Args:\n        boxes (torch.Tensor | np.ndarray): A tensor or numpy array with detection boxes.\n            Shape can be (num_boxes, 6), (num_boxes, 7), or (num_boxes, 6 + num_classes).\n            Columns should contain [x1, y1, x2, y2, confidence, class, (optional) track_id, (optional) class_conf_1, class_conf_2, ...].\n        orig_shape (tuple): The original image shape as (height, width). Used for normalization.\n\n    Returns:\n        (None)\n    \"\"\"\n\n    if boxes.ndim == 1:\n        boxes = boxes[None, :]\n    n = boxes.shape[-1]\n    super(ultralytics.engine.results.Boxes, self).__init__(boxes, orig_shape)\n    self.orig_shape = orig_shape\n    self.is_track = False\n    self.num_classes = 0\n\n    if n == 6:\n        self.format = 'xyxy_conf_cls'\n    elif n == 7:\n        self.format = 'xyxy_conf_cls_track'\n        self.is_track = True\n    else:\n        self.format = 'xyxy_conf_cls_classconf'\n        self.num_classes = n - 6\n\nultralytics.engine.results.Boxes.__init__ = init\n\nfrom ultralytics.utils.ops import xywh2xyxy, LOGGER, nms_rotated\nimport torch\nimport time\n\ndef non_max_suppression(\n    prediction,\n    conf_thres=0.2,\n    iou_thres=0.3,\n    classes=None,\n    agnostic=False,\n    multi_label=False,\n    labels=(),\n    max_det=300,\n    nc=0,  # number of classes (optional)\n    max_time_img=0.05,\n    max_nms=30000,\n    max_wh=7680,\n    in_place=True,\n    rotated=False,\n):\n    \"\"\"\n    Perform non-maximum suppression (NMS) on a set of boxes, with support for masks and multiple labels per box.\n    This version returns confidences for all classes.\n\n    Args:\n        (... same as before ...)\n\n    Returns:\n        (List[torch.Tensor]): A list of length batch_size, where each element is a tensor of\n            shape (num_boxes, 6 + num_classes + num_masks) containing the kept boxes, with columns\n            (x1, y1, x2, y2, confidence, class, class_conf_1, class_conf_2, ..., mask1, mask2, ...).\n    \"\"\"\n    import torchvision\n\n    # Checks and initialization (same as before)\n    assert 0 <= conf_thres <= 1, f\"Invalid Confidence threshold {conf_thres}, valid values are between 0.0 and 1.0\"\n    assert 0 <= iou_thres <= 1, f\"Invalid IoU {iou_thres}, valid values are between 0.0 and 1.0\"\n    if isinstance(prediction, (list, tuple)):\n        prediction = prediction[0]\n    if classes is not None:\n        classes = torch.tensor(classes, device=prediction.device)\n\n    bs = prediction.shape[0]  # batch size\n    nc = nc or (prediction.shape[1] - 4)  # number of classes\n    nm = prediction.shape[1] - nc - 4  # number of masks\n    mi = 4 + nc  # mask start index\n    xc = prediction[:, 4:mi].amax(1) > conf_thres  # candidates\n\n    # Settings\n    time_limit = 2.0 + max_time_img * bs  # seconds to quit after\n    multi_label &= nc > 1  # multiple labels per box (adds 0.5ms/img)\n\n    prediction = prediction.transpose(-1, -2)  # shape(1,84,6300) to shape(1,6300,84)\n    if not rotated:\n        if in_place:\n            prediction[..., :4] = xywh2xyxy(prediction[..., :4])  # xywh to xyxy\n        else:\n            prediction = torch.cat((xywh2xyxy(prediction[..., :4]), prediction[..., 4:]), dim=-1)  # xywh to xyxy\n\n    t = time.time()\n    output = [torch.zeros((0, 6 + nc + nm), device=prediction.device)] * bs\n    for xi, x in enumerate(prediction):  # image index, image inference\n        x = x[xc[xi]]  # confidence\n\n        # Cat apriori labels if autolabelling\n        if labels and len(labels[xi]) and not rotated:\n            lb = labels[xi]\n            v = torch.zeros((len(lb), nc + nm + 4), device=x.device)\n            v[:, :4] = xywh2xyxy(lb[:, 1:5])  # box\n            v[range(len(lb)), lb[:, 0].long() + 4] = 1.0  # cls\n            x = torch.cat((x, v), 0)\n\n        # If none remain process next image\n        if not x.shape[0]:\n            continue\n\n        # Detections matrix nx(4 + nc + nm) (xyxy, class_conf, cls, masks)\n        box, cls_conf, mask = x.split((4, nc, nm), 1)\n\n        # Confidence thresholding\n        conf, j = cls_conf.max(1, keepdim=True)\n        x = torch.cat((box, conf, j.float(), cls_conf, mask), 1)[conf.view(-1) > conf_thres]\n\n        # Filter by class\n        if classes is not None:\n            x = x[(x[:, 5:6] == classes).any(1)]\n\n        # Check shape\n        n = x.shape[0]  # number of boxes\n        if not n:  # no boxes\n            continue\n        if n > max_nms:  # excess boxes\n            x = x[x[:, 4].argsort(descending=True)[:max_nms]]  # sort by confidence and remove excess boxes\n\n        # Batched NMS\n        c = x[:, 5:6] * (0 if agnostic else max_wh)  # classes\n        scores = x[:, 4]  # scores\n        if rotated:\n            boxes = torch.cat((x[:, :2] + c, x[:, 2:4], x[:, -1:]), dim=-1)  # xywhr\n            i = nms_rotated(boxes, scores, iou_thres)\n        else:\n            boxes = x[:, :4] + c  # boxes (offset by class)\n            i = torchvision.ops.nms(boxes, scores, iou_thres)  # NMS\n        i = i[:max_det]  # limit detections\n\n        output[xi] = x[i]\n        if (time.time() - t) > time_limit:\n            LOGGER.warning(f\"WARNING ⚠️ NMS time limit {time_limit:.3f}s exceeded\")\n            break  # time limit exceeded\n\n    return output\n\nultralytics.utils.ops.non_max_suppression = non_max_suppression","metadata":{"execution":{"iopub.status.busy":"2024-10-08T15:00:28.357487Z","iopub.execute_input":"2024-10-08T15:00:28.357752Z","iopub.status.idle":"2024-10-08T15:00:36.034127Z","shell.execute_reply.started":"2024-10-08T15:00:28.357730Z","shell.execute_reply":"2024-10-08T15:00:36.033142Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%capture\n\"\"\"import ultralytics\n\nSCS_WEIGHTS = ['/kaggle/input/lsdc-train-yolo-scs/lsdc_yolov8/train/weights/best.pt']\n\nSS_WEIGHTS = ['/kaggle/input/lsdc-train-yolo-ss/lsdc_yolov8/train/weights/best.pt',]\n             #'/kaggle/input/lsdc-yolo-ssv3/best.pt']\n\nNFN_WEIGHTS = ['/kaggle/input/lsdc-train-yolo-nfn/lsdc_yolov8/train/weights/best.pt',]\n              #'/kaggle/input/lsdc-yolo-nfnv3/best.pt']\n\n# Load YOLO Model\nscs_models = []\nfor weight in SCS_WEIGHTS:\n    scs_models.append(ultralytics.YOLO(weight))\n    \nss_models = []\nfor weight in SS_WEIGHTS:\n    ss_models.append(ultralytics.YOLO(weight))\n    \nnfn_models = []\nfor weight in NFN_WEIGHTS:\n    nfn_models.append(ultralytics.YOLO(weight))\n    \nmodel = nfn_models[0]\ndir_name = \"/images/3852140407/2597842871\"\nresults = model(\"/images/3852140407/2597842871/4.jpg\", conf=0.01, verbose=False)\n#for img_name is os.listdir(dir_name):\nfor idx, res in enumerate(results):\n    xyxy_s = res.boxes.data[:, :4]\n    pred_conf = res.boxes.data[:, 4]\n    class_pred = res.boxes.data[:, 5]\n    probs = res.boxes.data[6:]\n\n    print (class_pred)\n\"\"\"","metadata":{"execution":{"iopub.status.busy":"2024-10-08T15:00:36.035371Z","iopub.execute_input":"2024-10-08T15:00:36.035831Z","iopub.status.idle":"2024-10-08T15:00:36.044483Z","shell.execute_reply.started":"2024-10-08T15:00:36.035799Z","shell.execute_reply":"2024-10-08T15:00:36.043458Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport pydicom\nfrom PIL import Image\nimport numpy as np\nfrom multiprocessing import Pool, cpu_count\n\nimport sklearn.metrics\nimport torch\nimport cv2\nimport numpy as np \nimport pandas as pd \nfrom tqdm.auto import tqdm","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-10-08T15:00:36.045584Z","iopub.execute_input":"2024-10-08T15:00:36.046147Z","iopub.status.idle":"2024-10-08T15:00:37.319785Z","shell.execute_reply.started":"2024-10-08T15:00:36.046018Z","shell.execute_reply":"2024-10-08T15:00:37.318987Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"EVAL = False # Change to True to compute the validation score\nIMG_DIR = '/images'\nFOLD = 0\nSAMPLE = False # True for quick debugging\nSEVERITIES = ['Normal/Mild', 'Moderate', 'Severe']\nLEVELS = ['l1_l2', 'l2_l3', 'l3_l4', 'l4_l5', 'l5_s1']\n\nSCS_WEIGHTS = ['/kaggle/input/lsdc-train-yolo-scs/lsdc_yolov8/train/weights/best.pt']\n\nSS_WEIGHTS = ['/kaggle/input/lsdc-train-yolo-ss/lsdc_yolov8/train/weights/best.pt',\n             '/kaggle/input/lsdc-yolo-ssv3/best.pt']\n\nNFN_WEIGHTS = ['/kaggle/input/yolo-041024-train/nfn_yolo.pt',\n              '/kaggle/input/lsdc-train-yolo-nfn-2-fold/lsdc_yolov8/train/weights/best.pt']\n              #'/kaggle/input/lsdc-yolo-nfnv3/best.pt']","metadata":{"execution":{"iopub.status.busy":"2024-10-08T15:00:37.320824Z","iopub.execute_input":"2024-10-08T15:00:37.321223Z","iopub.status.idle":"2024-10-08T15:00:37.326838Z","shell.execute_reply.started":"2024-10-08T15:00:37.321186Z","shell.execute_reply":"2024-10-08T15:00:37.325933Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!cp /kaggle/input/lsdc-train-yolo-scs/lsdc_yolov8/train/weights/best.pt scs_yolo.pt\n!cp /kaggle/input/lsdc-train-yolo-nfn/lsdc_yolov8/train/weights/best.pt nfn_yolo.pt\n!cp /kaggle/input/lsdc-train-yolo-ss/lsdc_yolov8/train/weights/best.pt ss_yolo.pt","metadata":{"execution":{"iopub.status.busy":"2024-10-08T15:00:37.327901Z","iopub.execute_input":"2024-10-08T15:00:37.328239Z","iopub.status.idle":"2024-10-08T15:00:41.423377Z","shell.execute_reply.started":"2024-10-08T15:00:37.328208Z","shell.execute_reply":"2024-10-08T15:00:41.422086Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if EVAL:\n    import sys\n    sys.path.append('/kaggle/input/lsdc-utils')\n    from metrics import score as lsdc_scoring","metadata":{"execution":{"iopub.status.busy":"2024-10-08T15:00:41.425526Z","iopub.execute_input":"2024-10-08T15:00:41.425899Z","iopub.status.idle":"2024-10-08T15:00:41.431175Z","shell.execute_reply.started":"2024-10-08T15:00:41.425867Z","shell.execute_reply":"2024-10-08T15:00:41.430355Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_val_df = pd.read_csv('/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train.csv')","metadata":{"execution":{"iopub.status.busy":"2024-10-08T15:00:41.435340Z","iopub.execute_input":"2024-10-08T15:00:41.435662Z","iopub.status.idle":"2024-10-08T15:00:41.487585Z","shell.execute_reply.started":"2024-10-08T15:00:41.435636Z","shell.execute_reply":"2024-10-08T15:00:41.486672Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if EVAL:\n    train_xy = pd.read_csv('/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_label_coordinates.csv')\n    des = pd.read_csv('/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_series_descriptions.csv')\nelse:    \n    des = pd.read_csv('/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/test_series_descriptions.csv')","metadata":{"execution":{"iopub.status.busy":"2024-10-08T15:00:41.489313Z","iopub.execute_input":"2024-10-08T15:00:41.489964Z","iopub.status.idle":"2024-10-08T15:00:41.500187Z","shell.execute_reply.started":"2024-10-08T15:00:41.489929Z","shell.execute_reply":"2024-10-08T15:00:41.499412Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def read_dcm(src_path):\n    dicom_data = pydicom.dcmread(src_path)\n    image = dicom_data.pixel_array\n    image = (image - image.min()) / (image.max() - image.min() +1e-6) * 255\n    return image\n\ndef convert_dcm_to_jpg(file_path):\n    try:\n        # Read the DICOM file\n        image_array = read_dcm(file_path)\n        \n        # Define the output path\n        relative_path = os.path.relpath(file_path, start=input_directory)\n        output_path = os.path.join(output_directory, relative_path)\n        output_path = output_path.replace('.dcm', '.jpg')\n                \n        # Create the output directory if it doesn't exist\n        os.makedirs(os.path.dirname(output_path), exist_ok=True)\n        \n        # Save the image as a JPEG file\n        cv2.imwrite(output_path, image_array)\n        \n        return output_path\n    except Exception as e:\n        print(f\"Error processing file {file_path}: {e}\")\n        return None\n\ndef process_files(dcm_files):\n    with Pool(cpu_count()) as pool:\n        # Wrap pool.map with tqdm to show the progress bar\n        list(tqdm(pool.imap(convert_dcm_to_jpg, dcm_files), total=len(dcm_files)))\n\ndef get_dcm_files(directory):\n    dcm_files = []\n    for root, dirs, files in os.walk(directory):\n        for file in files:\n            if file.endswith('.dcm'):\n                dcm_files.append(os.path.join(root, file))\n    return dcm_files    ","metadata":{"execution":{"iopub.status.busy":"2024-10-08T15:00:41.501538Z","iopub.execute_input":"2024-10-08T15:00:41.501846Z","iopub.status.idle":"2024-10-08T15:00:41.512898Z","shell.execute_reply.started":"2024-10-08T15:00:41.501822Z","shell.execute_reply":"2024-10-08T15:00:41.511884Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Replace these with your input and output directories\nif not EVAL:\n    input_directory = '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/test_images'\n\n    output_directory = IMG_DIR\n\n    # Get all .dcm files in the input directory\n    dcm_files = get_dcm_files(input_directory)\n\n    # Process the files using multiprocessing\n    process_files(dcm_files)\n\n    print(f\"Conversion completed. Images saved to {output_directory}\")\nelse:\n    if not os.path.exists(IMG_DIR):\n        print('Unziping data..')\n        !unzip -q -d / /kaggle/input/lsdc-get-all-images/images.zip\n        print('Done unziping data')","metadata":{"execution":{"iopub.status.busy":"2024-10-08T15:00:41.514154Z","iopub.execute_input":"2024-10-08T15:00:41.514493Z","iopub.status.idle":"2024-10-08T15:00:43.002987Z","shell.execute_reply.started":"2024-10-08T15:00:41.514469Z","shell.execute_reply":"2024-10-08T15:00:43.001911Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if EVAL:\n    fold_df = pd.read_csv('/kaggle/input/lsdc-fold-split/5folds.csv')\n    test_df = fold_df[fold_df.fold == FOLD]\n    \nelse:\n    test_df = os.listdir('/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/test_images')\n    test_df = pd.DataFrame(test_df, columns=['study_id'])\n    test_df['study_id'] = test_df['study_id'].astype(int)\n    \ntest_df = test_df.merge(des, on=['study_id'])","metadata":{"execution":{"iopub.status.busy":"2024-10-08T15:00:43.004414Z","iopub.execute_input":"2024-10-08T15:00:43.004717Z","iopub.status.idle":"2024-10-08T15:00:43.039210Z","shell.execute_reply.started":"2024-10-08T15:00:43.004691Z","shell.execute_reply":"2024-10-08T15:00:43.038260Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def gen_label_map(CONDITIONS):\n    label2id = {}\n    id2label = {}\n    i = 0\n    for cond in CONDITIONS:\n        for level in LEVELS:\n            for severity in SEVERITIES:\n                cls_ = f\"{cond.lower().replace(' ', '_')}_{level}_{severity.lower()}\"\n                label2id[cls_] = i\n                id2label[i] = cls_\n                i+=1\n    return label2id, id2label\n                \nscs_label2id, scs_id2label = gen_label_map(['Spinal Canal Stenosis'])\nss_label2id, ss_id2label = gen_label_map(['Left Subarticular Stenosis', 'Right Subarticular Stenosis'])\nnfn_label2id, nfn_id2label = gen_label_map(['Left Neural Foraminal Narrowing', 'Right Neural Foraminal Narrowing'])","metadata":{"execution":{"iopub.status.busy":"2024-10-08T15:00:43.040503Z","iopub.execute_input":"2024-10-08T15:00:43.040846Z","iopub.status.idle":"2024-10-08T15:00:43.049259Z","shell.execute_reply.started":"2024-10-08T15:00:43.040816Z","shell.execute_reply":"2024-10-08T15:00:43.048267Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from ultralytics import YOLO\n\n# Load YOLO Model\nscs_models = []\nfor weight in SCS_WEIGHTS:\n    scs_models.append(YOLO(weight))\n    \nss_models = []\nfor weight in SS_WEIGHTS:\n    ss_models.append(YOLO(weight))\n    \nnfn_models = []\nfor weight in NFN_WEIGHTS:\n    nfn_models.append(YOLO(weight))","metadata":{"execution":{"iopub.status.busy":"2024-10-08T15:00:43.050668Z","iopub.execute_input":"2024-10-08T15:00:43.051421Z","iopub.status.idle":"2024-10-08T15:00:45.012366Z","shell.execute_reply.started":"2024-10-08T15:00:43.051389Z","shell.execute_reply":"2024-10-08T15:00:45.011460Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"all_label_set = train_val_df.iloc[0, 1:].index.tolist()\nscs_label_set = all_label_set[:5]\nnfn_label_set = all_label_set[5:15]\nss_label_set = all_label_set[15:]","metadata":{"execution":{"iopub.status.busy":"2024-10-08T15:00:45.013809Z","iopub.execute_input":"2024-10-08T15:00:45.014367Z","iopub.status.idle":"2024-10-08T15:00:45.020090Z","shell.execute_reply.started":"2024-10-08T15:00:45.014332Z","shell.execute_reply":"2024-10-08T15:00:45.019215Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def gen_label_map(CONDITIONS):\n    label2id = {}\n    id2label = {}\n    i = 0\n    for cond in CONDITIONS:\n        for level in LEVELS:\n            for severity in [\"N\", \"M\", \"S\"]:\n                cls_ = f\"{cond.lower().replace(' ', '_')}_{level}_{severity.lower()}\"\n                label2id[cls_] = i\n                id2label[i] = cls_\n                i+=1\n    return label2id, id2label\n\np_scs_label2id, p_scs_id2label = gen_label_map(['SCS'])\np_ss_label2id, p_ss_id2label = gen_label_map(['L-SS', 'R-SS'])\np_nfn_label2id, p_nfn_id2label = gen_label_map(['L-NFN', 'R-NFN'])","metadata":{"execution":{"iopub.status.busy":"2024-10-08T15:00:45.021415Z","iopub.execute_input":"2024-10-08T15:00:45.022336Z","iopub.status.idle":"2024-10-08T15:00:45.030168Z","shell.execute_reply.started":"2024-10-08T15:00:45.022304Z","shell.execute_reply":"2024-10-08T15:00:45.029395Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"settings = [\n    ( 'Sagittal T2/STIR', scs_models, scs_id2label, p_scs_id2label, scs_label_set, 0.01),\n    ( 'Sagittal T1', nfn_models, nfn_id2label, p_nfn_id2label, nfn_label_set, 0.1),\n    ( 'Axial T2', ss_models, ss_id2label, p_ss_id2label, ss_label_set, 0.01),\n]","metadata":{"execution":{"iopub.status.busy":"2024-10-08T15:00:45.032340Z","iopub.execute_input":"2024-10-08T15:00:45.032645Z","iopub.status.idle":"2024-10-08T15:00:45.044691Z","shell.execute_reply.started":"2024-10-08T15:00:45.032593Z","shell.execute_reply":"2024-10-08T15:00:45.043773Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from collections import defaultdict","metadata":{"execution":{"iopub.status.busy":"2024-10-08T15:00:45.046166Z","iopub.execute_input":"2024-10-08T15:00:45.046537Z","iopub.status.idle":"2024-10-08T15:00:45.053169Z","shell.execute_reply.started":"2024-10-08T15:00:45.046507Z","shell.execute_reply":"2024-10-08T15:00:45.052385Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import re\n\n# Function to extract the numeric part of the filename\ndef extract_number(filename):\n    match = re.search(r'(\\d+)', filename)  # Extracts the first number found in the filename\n    return int(match.group(0)) if match else 0  # Returns 0 if no number is found\n\ndef get_filenames(dir_name):\n    # List all image files in the directory\n    image_files = [f for f in os.listdir(dir_name) if f.endswith(('jpg', 'jpeg', 'png', 'dcm'))]\n\n    # Sort the filenames numerically\n    sorted_images = sorted(image_files, key=extract_number) \n\n    return [dir_name + \"/\" + f for f in sorted_images]\n\npred_rows = list()\npred_dataframe = list()\nPLOT_DIR = \"\"\n\nfor modality, models, id2label, p_id2label, label_set, thresh in settings:\n    mod_df = test_df[test_df.series_description == modality]\n    \n    if SAMPLE:\n        mod_df = mod_df.sample(20, random_state=610)\n    \n    # for each study, at each level and condition, get the maximum probability score\n    for study_id, group in tqdm(mod_df.groupby('study_id')):\n        if PLOT_DIR:\n            study_pred_dir = os.path.join(PLOT_DIR, str(study_id), modality)\n            if not os.path.exists(study_pred_dir):\n                os.makedirs(study_pred_dir)\n        \n        predictions = defaultdict(list)\n        for i, row in group.iterrows():\n            # predict on all images from all the series\n            series_dir = os.path.join(IMG_DIR, str(row['study_id']), str(row['series_id']))\n            for model in models:\n                results = model(get_filenames(series_dir), conf=thresh, verbose=False)\n                for idx, res in enumerate(results):\n                    res.names = p_id2label\n                    if PLOT_DIR:\n                        res.plot(filename=os.path.join(study_pred_dir, '{0:02}'.format(idx + 1) + \".png\"), save=True) \n                    for datka in res.boxes.data:\n                        pred_row = dict()\n                        datka = datka.cpu().numpy()\n                        \n                        #pred_row[\"model\"] = model.__name__\n                        pred_row[\"study_id\"] = study_id\n                        pred_row[\"modality\"] = modality\n                        pred_row[\"instance_id\"] = idx\n       \n                        pred_row[\"confidence\"] = datka[4]\n                        pred_row[\"class\"] = id2label[int(datka[5])]\n            \n                        assert len(id2label) == len(datka[6:])\n                        for label, prob in zip(id2label.values(), datka[6:]):\n                            pred_row[label] = prob\n    \n                        pred_dataframe.append(pred_row)\n        \n        # aggregate the result on images to obtain study-level prediction\n        for condition in label_set:\n            res_dict = {'row_id': f'{study_id}_{condition}' }\n\n            score_vec = []\n            for severity in SEVERITIES:\n                severity = severity.lower()\n                key = f'{condition}_{severity}'\n                if len(predictions[key]) > 0:\n                    score = np.max(predictions[key])\n                else:\n                    score = thresh\n                score_vec.append(score)\n                \n            # normalize score to sum to 1\n            score_vec = torch.tensor(score_vec)\n            score_vec = score_vec / score_vec.sum()\n\n            for idx, severity in enumerate(SEVERITIES):\n                res_dict[severity.replace('/', '_').lower()] = score_vec[idx].item()\n\n            pred_rows.append(res_dict)\n","metadata":{"execution":{"iopub.status.busy":"2024-10-08T15:00:45.054586Z","iopub.execute_input":"2024-10-08T15:00:45.055401Z","iopub.status.idle":"2024-10-08T15:00:51.544141Z","shell.execute_reply.started":"2024-10-08T15:00:45.055370Z","shell.execute_reply":"2024-10-08T15:00:51.543237Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"pred_df = pd.DataFrame(pred_rows)\npred_df","metadata":{"execution":{"iopub.status.busy":"2024-10-08T15:00:51.545522Z","iopub.execute_input":"2024-10-08T15:00:51.545950Z","iopub.status.idle":"2024-10-08T15:00:51.568272Z","shell.execute_reply.started":"2024-10-08T15:00:51.545925Z","shell.execute_reply":"2024-10-08T15:00:51.567233Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# pred_df.to_csv('submission.csv', index=False)","metadata":{"execution":{"iopub.status.busy":"2024-10-08T15:00:51.569400Z","iopub.execute_input":"2024-10-08T15:00:51.569673Z","iopub.status.idle":"2024-10-08T15:00:51.574018Z","shell.execute_reply.started":"2024-10-08T15:00:51.569650Z","shell.execute_reply":"2024-10-08T15:00:51.573047Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#%%capture\n#!zip -r output.zip output\n#!rm -rf output\n# !rm submission.csv","metadata":{"execution":{"iopub.status.busy":"2024-10-08T15:00:51.575474Z","iopub.execute_input":"2024-10-08T15:00:51.575751Z","iopub.status.idle":"2024-10-08T15:00:51.582927Z","shell.execute_reply.started":"2024-10-08T15:00:51.575728Z","shell.execute_reply":"2024-10-08T15:00:51.582213Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"len(pred_dataframe)","metadata":{"execution":{"iopub.status.busy":"2024-10-08T15:00:51.584104Z","iopub.execute_input":"2024-10-08T15:00:51.584951Z","iopub.status.idle":"2024-10-08T15:00:51.595090Z","shell.execute_reply.started":"2024-10-08T15:00:51.584916Z","shell.execute_reply":"2024-10-08T15:00:51.594213Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"debug_df = pd.DataFrame(pred_dataframe)\n\ndebug_df.to_csv(\"row_yolos_predicts.csv\", index=False)","metadata":{"execution":{"iopub.status.busy":"2024-10-08T15:00:51.596158Z","iopub.execute_input":"2024-10-08T15:00:51.596457Z","iopub.status.idle":"2024-10-08T15:00:51.724984Z","shell.execute_reply.started":"2024-10-08T15:00:51.596433Z","shell.execute_reply":"2024-10-08T15:00:51.724186Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"yolo_preds = debug_df.fillna(0)\n\nyolo_preds.head()","metadata":{"execution":{"iopub.status.busy":"2024-10-08T15:00:51.726047Z","iopub.execute_input":"2024-10-08T15:00:51.726335Z","iopub.status.idle":"2024-10-08T15:00:51.759661Z","shell.execute_reply.started":"2024-10-08T15:00:51.726311Z","shell.execute_reply":"2024-10-08T15:00:51.758383Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def make_neighbors(num):\n    if num % 3 == 0:\n        return [num, num + 1, num + 2]\n    elif num % 3 == 1:\n        return [num - 1, num, num + 1]\n    else:\n        return [num - 2, num - 1, num]\n\nargmax_xd = yolo_preds.iloc[:, 5:].values.argmax(axis=1)\nargmax_yd = np.array([[num] * 3 for num, _ in enumerate(argmax_xd)]).flatten()\nargmax_xd = np.array([make_neighbors(num) for num in argmax_xd]).flatten() \n\nseverity = yolo_preds.iloc[:, 5:].values[argmax_yd, argmax_xd].reshape(-1, 3)\n\nyolo_preds = yolo_preds.iloc[:, :5]\n\nyolo_preds[\"class\"] = yolo_preds[\"class\"].apply(lambda x: \"_\".join(x.split(\"_\")[:-1]))\nyolo_preds[\"normal/mild\"] = severity[:, 0]\nyolo_preds[\"moderate\"] = severity[:, 1]\nyolo_preds[\"severe\"] = severity[:, 2]\n\nyolo_preds.head()","metadata":{"execution":{"iopub.status.busy":"2024-10-08T15:00:51.766607Z","iopub.execute_input":"2024-10-08T15:00:51.766932Z","iopub.status.idle":"2024-10-08T15:00:51.796167Z","shell.execute_reply.started":"2024-10-08T15:00:51.766906Z","shell.execute_reply":"2024-10-08T15:00:51.795132Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"norm_answers = yolo_preds.groupby([\"study_id\", \"class\"]).max().reset_index()\n\nnorm_answers[\"row_id\"] = norm_answers[\"study_id\"].apply(str) + \"_\" + norm_answers[\"class\"]\n\nnorm_answers[\"sum\"] = norm_answers[\"normal/mild\"] + norm_answers[\"moderate\"] + norm_answers[\"severe\"]\nnorm_answers[\"normal_mild\"] = norm_answers[\"normal/mild\"] / norm_answers[\"sum\"]\nnorm_answers[\"moderate\"] = norm_answers[\"moderate\"] / norm_answers[\"sum\"]\nnorm_answers[\"severe\"] = norm_answers[\"severe\"] / norm_answers[\"sum\"]\n\nnorm_answers = norm_answers[['row_id', 'normal_mild', 'moderate', 'severe']].set_index(\"row_id\")\n\npred_df = pred_df.set_index(\"row_id\")\n\nfor col in ['normal_mild', 'moderate', 'severe']:\n    pred_df[col] = norm_answers[col]\n\npred_df = pred_df.fillna(0)\n    \n# pred_df.to_csv('submission.csv') #, index=False)\n\npred_df.reset_index(inplace = True)\n\npred_df","metadata":{"execution":{"iopub.status.busy":"2024-10-08T15:00:51.797372Z","iopub.execute_input":"2024-10-08T15:00:51.797691Z","iopub.status.idle":"2024-10-08T15:00:51.841806Z","shell.execute_reply.started":"2024-10-08T15:00:51.797664Z","shell.execute_reply":"2024-10-08T15:00:51.840884Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def sample_weight(row):\n    if row['normal_mild'] == 1:\n        return 1\n    if row['moderate'] == 1:\n        return 2\n    if row['severe'] == 1:\n        return 4\n    raise ValueError('No such value')\n    \ndef get_class(row):\n    return np.argmax([row['normal_mild'], row['moderate'], row['severe']])","metadata":{"execution":{"iopub.status.busy":"2024-10-08T15:00:51.843160Z","iopub.execute_input":"2024-10-08T15:00:51.843542Z","iopub.status.idle":"2024-10-08T15:00:51.849707Z","shell.execute_reply.started":"2024-10-08T15:00:51.843510Z","shell.execute_reply":"2024-10-08T15:00:51.848718Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if EVAL:\n    gt_df = train_val_df.dropna().melt(id_vars=['study_id'], value_vars=all_label_set)\n    gt_df['row_id'] = gt_df['study_id'].astype(str) + '_' + gt_df['variable']\n    gt_df= gt_df[['row_id', 'value']]\n    gt_df = pd.get_dummies(gt_df, columns=['value'], dtype=int)\n    gt_df.columns = ['row_id', 'moderate', 'normal_mild', 'severe']\n    gt_df = gt_df[['row_id', 'normal_mild', 'moderate', 'severe']]\n    gt_df['sample_weight'] = gt_df.apply(sample_weight, axis=1)\n\n    gt_df1 = gt_df.merge(pred_df['row_id'], how='inner', on='row_id').sort_values('row_id').reset_index(drop=True)\n    pred_df1 = pred_df.merge(gt_df1['row_id'], how='inner', on='row_id').sort_values('row_id').reset_index(drop=True)\n    gt_df1['pred_cls'] = gt_df1.apply(get_class, axis=1)\n    pred_df1['pred_cls'] = pred_df1.apply(get_class, axis=1)\n\n    gt_df1[(gt_df1['pred_cls'] != pred_df1['pred_cls'])]\n    pred_df1[(gt_df1['pred_cls'] != pred_df1['pred_cls'])]\n    print('Label count:\\n', gt_df1['pred_cls'].value_counts(normalize=True))\n    print('Prediction accuracy:', (gt_df1['pred_cls'] == pred_df1['pred_cls']).mean())\n    print()\n\n    target_levels = ['normal_mild', 'moderate', 'severe']\n    loss = lsdc_scoring(gt_df1.drop(['pred_cls'], axis=1), pred_df1.drop(['pred_cls'], axis=1), row_id_column_name='row_id', any_severe_scalar=1)\n    print('Total weighted log loss:', loss)","metadata":{"execution":{"iopub.status.busy":"2024-10-08T15:00:51.850938Z","iopub.execute_input":"2024-10-08T15:00:51.851252Z","iopub.status.idle":"2024-10-08T15:00:51.862705Z","shell.execute_reply.started":"2024-10-08T15:00:51.851225Z","shell.execute_reply":"2024-10-08T15:00:51.861913Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install /kaggle/input/pip-install/humanfriendly-10.0-py2.py3-none-any.whl --no-index --find-links /kaggle/input/pip-install\n!pip install /kaggle/input/pip-install/coloredlogs-15.0.1-py2.py3-none-any.whl --no-index --find-links /kaggle/input/pip-install\n!pip install /kaggle/input/pip-install/onnxruntime-1.17.3-cp310-cp310-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl/onnxruntime-1.17.3-cp310-cp310-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl","metadata":{"execution":{"iopub.status.busy":"2024-10-08T15:00:51.863865Z","iopub.execute_input":"2024-10-08T15:00:51.864187Z","iopub.status.idle":"2024-10-08T15:01:50.428477Z","shell.execute_reply.started":"2024-10-08T15:00:51.864154Z","shell.execute_reply":"2024-10-08T15:01:50.427267Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nfrom torch.utils.data import Dataset, DataLoader\nimport pandas as pd\nimport numpy as np\nimport albumentations as A\nimport pydicom\nimport os\nfrom PIL import Image\nimport sys \nsys.path.append(\"/kaggle/input/preprocess\")\nfrom segmantation_inference import SegmentaionInference\nfrom detection_inference import DetectionInference, transforms\nimport torch.nn.functional as F\nfrom pathlib import Path\nimport re\nimport ast\nfrom Cross_Reference_Axial import CrossReferenceAxial\nimport matplotlib.pyplot as plt\nclass CFG():\n    AUG_PROB = 0.75\n    NOT_DEBUG = True\n    AUG = True\n    Axial_shape = (152, 152)\n    Sagittal_shape = (152, 152)\n    channel_size_sagittal_t1 = 12\n    channel_size_sagittal_t2 = 9\n    channel_size_sagittal = channel_size_sagittal_t1 + channel_size_sagittal_t2\n    channel_size_axial = 9\n    train_path = \"train_images\"\n    segmentation = SegmentaionInference(model_path=r\"/kaggle/input/models/simple_unet.onnx\")\n    DetectionInference = DetectionInference(model_path=r\"/kaggle/input/models/axial_detection_resnet18.pth\", transforms=transforms)\n    cross_reference = CrossReferenceAxial(\"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/test_images/\",\"/kaggle/input/models/simple_unet.onnx\")\n    label2id = {'Normal/Mild': 0, 'Moderate':1, 'Severe':2, np.nan: -100}\n    category2id = {\"L1\": 1, \"L2\": 2, \"L3\": 3, \"L4\": 4, \"L5\": 5, \"L5-S1\": 11, \"L4-L5\": 12, \"L3-L4\": 13, \"L2-L3\": 14, \"L1-L2\": 15}\n    skip_study_id = [2492114990, 2780132468, 3008676218]\n    two_classes_category = {11: 'L5-S1', 12: 'L4-L5', 13: 'L3-L4', 14: 'L2-L3', 15: 'L1-L2'}\n    model_name_axial = \"timm/vit_small_patch14_reg4_dinov2.lvd142m\"\n    model_name_sagittal = \"timm/convnext_nano.in12k\"\n\ncfg = CFG()\n\ntransforms_val = A.Compose([\n  A.Normalize(mean=[0.485], std=[0.229])\n])\n\nclass CustomDataset(Dataset):\n    def __init__(self, study_ids, labels_path, test_path, transform):\n        self.study_ids = study_ids\n        self.df_description = pd.read_csv(labels_path)\n        self.transform = transform\n        self.test_path = test_path\n    def __len__(self):\n        return len(self.study_ids)\n    \n    @staticmethod\n    def plot(stack,x = 5,y = 6):\n        fig, axes = plt.subplots(x, y, figsize=(15, 9))\n        for i, ax in enumerate(axes.flat):\n            ax.imshow(stack[..., i], cmap='gray')\n        plt.tight_layout()\n        plt.show()\n    \n    def load_dicom(self, path):\n        original_dicom = pydicom.dcmread(path).pixel_array\n        original_dicom = original_dicom.clip(np.percentile(original_dicom, 1), np.percentile(original_dicom, 99))\n        original_dicom = np.array(self.resize_image(original_dicom, (512, 512)))\n        original_dicom = (original_dicom - original_dicom.min()) / (original_dicom.max() - original_dicom.min() + 1e-6) * 255\n        return original_dicom.astype(np.uint8)\n\n    @staticmethod\n    def pad_images_list(images_list, max_len): # need to check this function\n        if len(images_list) < 0:\n            raise ValueError(\"images_list is empty\")\n        if len(images_list) == max_len:\n            return images_list\n        \n        n = len(images_list)\n        output_list = []\n        \n        # How many times should we duplicate each element minimally?\n        min_repeats = max_len // n\n        \n        # How many extra duplicates are needed beyond minimal repeats?\n        extra = max_len % n\n        \n        # Determine the central region to duplicate more\n        mid_point = n // 2\n        start_extra = mid_point - (extra // 2)\n        end_extra = start_extra + extra\n        \n        # Duplicate elements, adding extra repeats to central elements\n        for i in range(n):\n            repeats = min_repeats + 1 if start_extra <= i < end_extra else min_repeats\n            output_list.extend([images_list[i]] * repeats)\n\n        return output_list\n    \n    @staticmethod\n    def crop_axial_center(image, bbox):\n        min_x, max_x = 999999, -1\n        min_y, max_y = 999999, -1\n        for x, y, h, w in bbox:\n            min_x = min(min_x, x)\n            max_x = max(max_x, x+w)\n            min_y = min(min_y, y)\n            max_y = max(max_y, y+h)\n        if type(image) != Image.Image:\n            image = Image.fromarray(image)\n        \n        if min_x == 999999 or min_y == 999999:\n            \n            width, height = image.size\n            # Define the size of the crop\n            crop_size = 304\n\n            # Calculate coordinates for the middle crop\n            left = (width - crop_size) // 2\n            top = (height - crop_size) // 2\n            right = left + crop_size\n            bottom = top + crop_size\n            cropped_image = image.crop((left, top, right, bottom))\n            # print(\"left: \", left, \"top: \", top, \"right: \", right, \"bottom: \", bottom)\n            return (\n            cropped_image.crop((0, 0, 152, crop_size)),   # Adjusted coordinates\n            cropped_image.crop((76, 0, 228, crop_size)),  # Adjusted coordinates\n            cropped_image.crop((152, 0, 304, crop_size))  # Adjusted coordinates\n        )\n        \n        # Crop the center of the image\n        margin = 304 // 2 \n        \n        return (image.crop((min_x - margin, min_y - 50, min_x, min_y + 102)),\n                image.crop((min_x - 76 , min_y - 50, min_x + 76, min_y + 102)),\n                image.crop((min_x, min_y - 50, min_x + margin, min_y + 102)))\n\n    @staticmethod\n    def center_crop_by_categorys(original_dicom, bboxes, category, second_category): # need to check this function\n        bbox1 = bboxes[category]\n        bbox2 = bboxes[second_category]\n        x, y, h, w = bbox1[0]\n        x2, y2, h2, w2 = bbox2[0]\n        # Function to get the minimum value ignoring -1\n        def min_ignore_neg_one(a, b):\n            if a == -1:\n                return b\n            if b == -1:\n                return a\n            return min(a, b)\n\n        # Function to get the maximum value ignoring -1\n        min_x = min_ignore_neg_one(x, x2)\n        min_y = min_ignore_neg_one(y, y2)\n        max_x = max(x + w, x2 + w2)\n        max_y = max(y + h, y2 + h2)\n        image = Image.fromarray(original_dicom)\n        if (min_x == -1 and min_y == -1) or ((max_x - min_x < 35) or (max_y - min_y < 35)):\n            # shape = np.array(original_dicom).shape\n            # need to add plt.imshow() of the segmentation mask\n            return original_dicom\n        cropped_image = image.crop((min_x-10, min_y-10, max_x + 30, max_y))\n\n        return np.array(cropped_image)\n    \n    def center_crop_by_category(self, pixel_array, bboxes, category):\n        bbox = bboxes[category]\n        x, y, h, w = bbox[0]\n        image = Image.fromarray(pixel_array)\n        margin = cfg.Sagittal_shape[0] // 2 \n        cropped_image = image.crop((x - 20, y - margin, x + cfg.Sagittal_shape[0] - 20, y + margin))\n        return np.array(cropped_image)\n\n\n    @staticmethod\n    def unpad_images_list(images_list, max_len): # need to check this function\n        i = 0\n        while len(images_list) > max_len:\n            if i % 2 == 0:\n                images_list.pop(-1)\n            else:\n                images_list.pop(0)\n            i += 1\n        return images_list\n\n    @staticmethod\n    def resize_image(pixel_array, new_size):\n        if pixel_array.shape == new_size:\n            return pixel_array\n        else:\n            image = Image.fromarray(pixel_array)\n            return image.resize((new_size[1], new_size[0]))\n\n    @staticmethod\n    def extract_number(filename):\n            match = re.search(r'\\d+', filename)\n            return int(match.group()) if match else 0\n\n    def divide_Axiel(self, sub_set):\n        df_classes = pd.DataFrame(columns=['path', 'class_id'])\n        study_id = sub_set['study_id'].iloc[0]\n        series_id_axial = sub_set['series_id'].loc[sub_set['series_description'] == \"Axial T2\"].iloc[0]\n        list_ = os.listdir(os.path.join(self.train_path, str(study_id), str(series_id_axial)))\n        list_ = sorted(list_, key=self.extract_number)\n        divide_by_5 = len(list_) // 5\n        remainder = len(list_) % 5\n\n        class_ids = [\"L1\", \"L2\", \"L3\", \"L4\", \"L5\"]\n        start_idx = 0\n\n        for i, class_id in enumerate(class_ids):\n            end_idx = start_idx + divide_by_5 + (1 if i < remainder else 0)  # Add 1 to the first 'remainder' groups\n            for file in list_[start_idx:end_idx]:\n                df_classes.loc[len(df_classes)] = [os.path.join(self.train_path, str(study_id), str(series_id_axial), file), class_id]\n            start_idx = end_idx\n        \n        return df_classes\n    \n    def crop_sagittal_center(self, file, Sagittal_bboxes_scaled, Sagittal_path, category):\n\n            original_dicom = self.load_dicom(os.path.join(Sagittal_path, file))\n            if \"S1\" in category:\n                category1 = category.split(\"-\")[0]\n                category2 = category\n            else:\n                category1 = category.split(\"-\")[0]\n                category2 = category.split(\"-\")[1]\n            new_pixel_array = self.center_crop_by_category(original_dicom, Sagittal_bboxes_scaled,\n                                                            cfg.category2id[category])\n            # resized_pixel_array = self.resize_image(new_pixel_array, (cfg.Sagittal_shape[0], cfg.Sagittal_shape[1]))\n            new_shape = new_pixel_array.shape\n            # Ensure the padded array is large enough to hold new_pixel_array\n            padded_array = np.zeros((max(cfg.Sagittal_shape[0], new_shape[0]), max(cfg.Sagittal_shape[1], new_shape[1])))\n            \n            # Compute the starting indices for centering the new_pixel_array\n            start_x = (padded_array.shape[0] - new_shape[0]) // 2\n            start_y = (padded_array.shape[1] - new_shape[1]) // 2\n\n            # Place the new_pixel_array in the center of the padded_array\n            padded_array[start_x:start_x + new_shape[0], start_y:start_y + new_shape[1]] = new_pixel_array\n            return padded_array[:cfg.Sagittal_shape[0], :cfg.Sagittal_shape[1]].astype(np.uint8)\n\n\n    def create_stack(self, sagittal_stack, axial_stack,\n                     Sagittal_T2_files, Sagittal_T2_bboxes, Sagittal_T2_path,\n                        Sagittal_T1_files, Sagittal_T1_bboxes, Sagittal_T1_path,\n                     category, two_classes_category, df_classes):\n        k = 0\n        # the order of the images are mirrored\n        RIGHT_T1_files = Sagittal_T1_files[:len(Sagittal_T1_files)//2]\n        LEFT_T1_files = Sagittal_T1_files[len(Sagittal_T1_files)//2:]\n            \n\n        if len(RIGHT_T1_files) < cfg.channel_size_sagittal_t1 // 2:\n            RIGHT_T1_files = self.pad_images_list(RIGHT_T1_files, cfg.channel_size_sagittal_t1//2)\n        elif len(RIGHT_T1_files) > cfg.channel_size_sagittal_t1 // 2:\n            RIGHT_T1_files = RIGHT_T1_files[-cfg.channel_size_sagittal_t1 // 2:]\n        # elif len(RIGHT_T1_files) > cfg.channel_size_sagittal_t1 // 2:\n        #     RIGHT_T1_files = self.unpad_images_list(RIGHT_T1_files, cfg.channel_size_sagittal_t1//2)\n\n        if len(LEFT_T1_files) < cfg.channel_size_sagittal_t1 // 2:\n            LEFT_T1_files = self.pad_images_list(LEFT_T1_files, cfg.channel_size_sagittal_t1//2)\n        elif len(LEFT_T1_files) > cfg.channel_size_sagittal_t1 // 2:\n            LEFT_T1_files = LEFT_T1_files[:cfg.channel_size_sagittal_t1//2]\n        # elif len(LEFT_T1_files) > cfg.channel_size_sagittal_t1 // 2:\n        #     LEFT_T1_files = self.unpad_images_list(LEFT_T1_files, cfg.channel_size_sagittal_t1//2)\n\n        original_dicom = pydicom.dcmread(os.path.join(Sagittal_T1_path, Sagittal_T1_files[len(Sagittal_T1_files)//2])).pixel_array\n        Sagittal_T1_bboxes = cfg.segmentation.scale_bboxes(Sagittal_T1_bboxes, (512, 512), original_dicom.shape)\n        LEFT_T1_files = LEFT_T1_files[::-1] # reverse the order of the images\n        for file in LEFT_T1_files:\n            sagittal_stack[..., k] = self.crop_sagittal_center(file, Sagittal_T1_bboxes, Sagittal_T1_path, two_classes_category)\n            k += 1\n        \n        RIGHT_T1_files = RIGHT_T1_files[::-1] # reverse the order of the images\n        for file in RIGHT_T1_files:\n            sagittal_stack[..., k] = self.crop_sagittal_center(file, Sagittal_T1_bboxes, Sagittal_T1_path, two_classes_category)\n            k += 1\n\n        \n\n\n        if len(Sagittal_T2_files) < cfg.channel_size_sagittal_t2:\n            Sagittal_T2_files = self.pad_images_list(Sagittal_T2_files, cfg.channel_size_sagittal_t2)\n        elif len(Sagittal_T2_files) > cfg.channel_size_sagittal_t2:\n            Sagittal_T2_files = self.unpad_images_list(Sagittal_T2_files, cfg.channel_size_sagittal_t2)\n        original_dicom = pydicom.dcmread(os.path.join(Sagittal_T2_path, Sagittal_T2_files[len(Sagittal_T2_files)//2])).pixel_array\n        Sagittal_T2_bboxes = cfg.segmentation.scale_bboxes(Sagittal_T2_bboxes, (512, 512), original_dicom.shape)\n        for file in Sagittal_T2_files:\n            sagittal_stack[..., k] = self.crop_sagittal_center(file, Sagittal_T2_bboxes, Sagittal_T2_path, two_classes_category)\n            k += 1\n        \n\n        l = df_classes['path'].loc[df_classes['class_id'] == two_classes_category].unique() # df_classes['class_id'] == category)\n\n        l = l.tolist()\n        if len(l) == 0:\n            l = df_classes['path'].loc[(df_classes['class_id'] == category)].unique()\n            l = l.tolist()\n        l = sorted(l, key=self.extract_number)\n\n        if len(l) == 0:\n            return sagittal_stack, axial_stack\n        \n        if len(l) < 3:\n            l = self.pad_images_list(l, 3)\n\n        elif len(l) > 3:\n            l = self.unpad_images_list(l, 3)\n        \n        def crop_axial_image(pixel_array):\n            # Define the size of the crop\n            crop_size = 384\n            width, height = np.array(pixel_array).shape[0], np.array(pixel_array).shape[1]\n            # Calculate coordinates for the middle crop\n            left = (width - crop_size) // 2\n            top = (height - crop_size) // 2\n            right = left + crop_size\n            bottom = top + crop_size\n            cropped_image = pixel_array.crop((left, top, right, bottom))\n            return cropped_image\n        \n        p = 0\n        j = 3\n        o = 6\n        for file in l:\n            original_dicom = pydicom.dcmread(file).pixel_array\n            original_dicom = original_dicom.clip(np.percentile(original_dicom, 1), np.percentile(original_dicom, 99))\n            bbox = cfg.DetectionInference.inference(original_dicom, 512, 512)\n            original_dicom = (original_dicom - original_dicom.min()) / (original_dicom.max() - original_dicom.min() + 1e-6) * 255\n            original_dicom = original_dicom.astype(np.uint8)\n            resized_pixel_array = self.resize_image(original_dicom, (512,512))\n            if type(resized_pixel_array) != Image.Image:\n                resized_pixel_array = Image.fromarray(resized_pixel_array)\n            resized_pixel_array = resized_pixel_array.transpose(Image.FLIP_LEFT_RIGHT)\n            # cropped_image = crop_axial_image(resized_pixel_array)\n            # axial_stack[..., p] = cropped_image\n            # p += 1\n            # Crop the center of the DICOM image\n            cropped_left, cropped_middle, cropped_right = self.crop_axial_center(resized_pixel_array, bbox)\n            cropped_left = self.resize_image(np.array(cropped_left), (cfg.Axial_shape[0], cfg.Axial_shape[1]))\n            cropped_middle = self.resize_image(np.array(cropped_middle), (cfg.Axial_shape[0], cfg.Axial_shape[1]))\n            cropped_right = self.resize_image(np.array(cropped_right), (cfg.Axial_shape[0], cfg.Axial_shape[1]))\n            axial_stack[..., p] = np.array(cropped_left).astype(np.uint8)\n            sagittal_stack[..., k + p] = np.array(cropped_left).astype(np.uint8)\n            p += 1\n            axial_stack[..., j] = np.array(cropped_middle).astype(np.uint8)\n            sagittal_stack[..., k + j] = np.array(cropped_middle).astype(np.uint8)\n            j += 1\n            axial_stack[..., o] = np.array(cropped_right).astype(np.uint8)\n            sagittal_stack[..., k + o] = np.array(cropped_right).astype(np.uint8)\n            o += 1\n        \n        return sagittal_stack, axial_stack\n\n    \n    @staticmethod\n    def _is_dict_structure_correct(d):\n        required_keys = {1, 2, 3, 4, 5, 11, 12, 13, 14, 15}\n        if set(d.keys()) != required_keys:\n            return False\n        \n        for key in required_keys:\n            if not (isinstance(d[key], list) and len(d[key]) == 1 and d[key][0] == (-1, -1, -1, -1)):\n                return False\n    \n        return True\n    \n    @staticmethod\n    def _is_all_black(image_array):\n        return np.all(image_array == 0)\n    \n\n    @staticmethod\n    def _count_neg_ones(bboxes):\n        count = 0\n        for vals in bboxes.values():\n            count += vals.count((-1, -1, -1, -1))\n        return count\n  \n    def __getitem__(self, index):\n        sagittal_l1_l2 = np.zeros((cfg.Sagittal_shape[0], cfg.Sagittal_shape[1], cfg.channel_size_sagittal + cfg.channel_size_axial), dtype = np.uint8)\n        axial_l1_l2 = np.zeros((cfg.Axial_shape[0], cfg.Axial_shape[1], cfg.channel_size_axial), dtype = np.uint8)\n        sagittal_l2_l3 = np.zeros((cfg.Sagittal_shape[0], cfg.Sagittal_shape[1], cfg.channel_size_sagittal + cfg.channel_size_axial), dtype = np.uint8)\n        axial_l2_l3 = np.zeros((cfg.Axial_shape[0], cfg.Axial_shape[1], cfg.channel_size_axial), dtype = np.uint8)\n        sagittal_l3_l4 = np.zeros((cfg.Sagittal_shape[0], cfg.Sagittal_shape[1], cfg.channel_size_sagittal + cfg.channel_size_axial), dtype = np.uint8)\n        axial_l3_l4 = np.zeros((cfg.Axial_shape[0], cfg.Axial_shape[1], cfg.channel_size_axial), dtype = np.uint8)\n        sagittal_l4_l5 = np.zeros((cfg.Sagittal_shape[0], cfg.Sagittal_shape[1], cfg.channel_size_sagittal + cfg.channel_size_axial), dtype = np.uint8)\n        axial_l4_l5 = np.zeros((cfg.Axial_shape[0], cfg.Axial_shape[1], cfg.channel_size_axial), dtype = np.uint8)\n        sagittal_l5_s1 = np.zeros((cfg.Sagittal_shape[0], cfg.Sagittal_shape[1], cfg.channel_size_sagittal + cfg.channel_size_axial), dtype = np.uint8)\n        axial_l5_s1 = np.zeros((cfg.Axial_shape[0], cfg.Axial_shape[1], cfg.channel_size_axial), dtype = np.uint8)\n        study_id = self.study_ids[index]\n\n        \n        sub_set = self.df_description.loc[self.df_description.study_id == study_id]\n        Sagittal_T1_path = os.path.join(self.test_path, str(study_id), str(sub_set[\"series_id\"].loc[sub_set[\"series_description\"] == \"Sagittal T1\"].iloc[0]))\n        Sagittal_T2_STIR_path = os.path.join(self.test_path, str(study_id), str(sub_set[\"series_id\"].loc[sub_set[\"series_description\"] == \"Sagittal T2/STIR\"].iloc[0]))\n        Axial_path = os.path.join(self.test_path, str(study_id), str(sub_set[\"series_id\"].loc[sub_set[\"series_description\"] == \"Axial T2\"].iloc[0]))\n\n        Sagittal_T1_files = os.listdir(Sagittal_T1_path)\n        Sagittal_T1_files = sorted(Sagittal_T1_files, key=self.extract_number)\n        middle_index = len(Sagittal_T1_files) // 2\n        Sagittal_T1_bboxes = cfg.segmentation.inference(os.path.join(Sagittal_T1_path, Sagittal_T1_files[middle_index]))\n        if self._count_neg_ones(Sagittal_T1_bboxes) != 0:\n            pmax = 0\n            pmin = len(Sagittal_T1_files)\n            for p in range(pmin, pmax):\n                temp = cfg.segmentation.inference(os.path.join(Sagittal_T1_path, Sagittal_T1_files[p]))\n                for key, value in temp.items():\n                    if value != [(-1, -1, -1, -1)]:\n                        if key not in Sagittal_T1_bboxes or Sagittal_T1_bboxes[key] == [(-1, -1, -1, -1)]:\n                            Sagittal_T1_bboxes[key] = value\n                if self._count_neg_ones(Sagittal_T1_bboxes) == 0:\n                    break\n        \n            \n        Sagittal_T2_files = os.listdir(Sagittal_T2_STIR_path)\n        Sagittal_T2_files = sorted(Sagittal_T2_files, key=self.extract_number)\n        middle_index = len(Sagittal_T2_files) // 2\n        Sagittal_T2_bboxes = cfg.segmentation.inference(os.path.join(Sagittal_T2_STIR_path, Sagittal_T2_files[middle_index]))\n        if self._count_neg_ones(Sagittal_T2_bboxes) != 0:\n            pmax = 0\n            pmin = len(Sagittal_T2_files)\n            for p in range(pmin, pmax):\n                temp = cfg.segmentation.inference(os.path.join(Sagittal_T2_STIR_path, Sagittal_T2_files[p]))\n                for key, value in temp.items():\n                    if value != [(-1, -1, -1, -1)]:\n                        if key not in Sagittal_T2_bboxes or Sagittal_T2_bboxes[key] == [(-1, -1, -1, -1)]:\n                            Sagittal_T2_bboxes[key] = value\n                if self._count_neg_ones(Sagittal_T2_bboxes) == 0:\n                    break\n        \n        \n        # Axial_files = os.listdir(Axial_path)\n        # Axial_files = sorted(Axial_files, key=self.extract_number)\n        decription_df = self.df_description[(self.df_description['study_id'] == study_id)]\n        # df_classes = self.divide_Axiel(decription_df)\n        \n        df_classes = cfg.cross_reference.get_cross_reference_for_Axial(decription_df, \"test\")\n\n        \n        \n        sagittal_l1_l2, axial_l1_l2 = self.create_stack(sagittal_l1_l2, axial_l1_l2,\n                                                        Sagittal_T2_files, Sagittal_T2_bboxes, Sagittal_T2_STIR_path,\n                                                        Sagittal_T1_files, Sagittal_T1_bboxes, Sagittal_T1_path,\n                                                        \"L1\", \"L1-L2\", df_classes)\n        \n        sagittal_l2_l3, axial_l2_l3 = self.create_stack(sagittal_l2_l3, axial_l2_l3,\n                                                        Sagittal_T2_files, Sagittal_T2_bboxes, Sagittal_T2_STIR_path,\n                                                        Sagittal_T1_files, Sagittal_T1_bboxes, Sagittal_T1_path,\n                                        \"L2\", \"L2-L3\", df_classes)\n        \n        sagittal_l3_l4, axial_l3_l4 = self.create_stack(sagittal_l3_l4, axial_l3_l4,\n                                                        Sagittal_T2_files, Sagittal_T2_bboxes, Sagittal_T2_STIR_path,\n                                                        Sagittal_T1_files, Sagittal_T1_bboxes, Sagittal_T1_path,\n                                        \"L3\", \"L3-L4\", df_classes)\n                                        \n        sagittal_l4_l5, axial_l4_l5 = self.create_stack(sagittal_l4_l5, axial_l4_l5,\n                                                        Sagittal_T2_files, Sagittal_T2_bboxes, Sagittal_T2_STIR_path,\n                                                        Sagittal_T1_files, Sagittal_T1_bboxes, Sagittal_T1_path,\n                                        \"L4\", \"L4-L5\", df_classes)\n        \n        sagittal_l5_s1, axial_l5_s1 = self.create_stack(sagittal_l5_s1, axial_l5_s1,\n                                                        Sagittal_T2_files, Sagittal_T2_bboxes, Sagittal_T2_STIR_path,\n                                                        Sagittal_T1_files, Sagittal_T1_bboxes, Sagittal_T1_path,\n                                        \"L5\", \"L5-S1\", df_classes)\n#         self.plot(sagittal_l1_l2[...,], x = 4, y = 7)\n#         self.plot(sagittal_l2_l3[...,], x = 4, y = 7)\n#         self.plot(sagittal_l3_l4[...,], x = 4, y = 7)\n#         self.plot(sagittal_l4_l5[...,], x = 4, y = 7)\n#         self.plot(sagittal_l5_s1[...,], x = 4, y = 7)\n        \n        flag_l1_l2 = np.all(np.isclose(sagittal_l1_l2[:, :, -9:], 0.0))\n        flag_l2_l3 = np.all(np.isclose(sagittal_l2_l3[:, :, -9:], 0.0))\n        flag_l3_l4 = np.all(np.isclose(sagittal_l3_l4[:, :, -9:], 0.0))\n        flag_l4_l5 = np.all(np.isclose(sagittal_l4_l5[:, :, -9:], 0.0))\n        flag_l5_s1 = np.all(np.isclose(sagittal_l5_s1[:, :, -9:], 0.0))\n        if self.transform:\n            sagittal_l1_l2 = self.transform(image=sagittal_l1_l2)['image']\n            sagittal_l2_l3 = self.transform(image=sagittal_l2_l3)['image']\n            sagittal_l3_l4 = self.transform(image=sagittal_l3_l4)['image']\n            sagittal_l4_l5 = self.transform(image=sagittal_l4_l5)['image']\n            sagittal_l5_s1 = self.transform(image=sagittal_l5_s1)['image']\n            axial_l1_l2 = self.transform(image=axial_l1_l2)['image']\n            axial_l2_l3 = self.transform(image=axial_l2_l3)['image']\n            axial_l3_l4 = self.transform(image=axial_l3_l4)['image']\n            axial_l4_l5 = self.transform(image=axial_l4_l5)['image']\n            axial_l5_s1 = self.transform(image=axial_l5_s1)['image']\n        if flag_l1_l2:\n            sagittal_l1_l2[:, :, -9:] = 0\n        if flag_l2_l3:\n            sagittal_l2_l3[:, :, -9:] = 0\n        if flag_l3_l4:\n            sagittal_l3_l4[:, :, -9:] = 0\n        if flag_l4_l5:\n            sagittal_l4_l5[:, :, -9:] = 0\n        if flag_l5_s1:\n            sagittal_l5_s1[:, :, -9:] = 0\n\n\n        sagittal_l1_l2 = torch.tensor(sagittal_l1_l2).permute(2, 0, 1)\n        sagittal_l2_l3 = torch.tensor(sagittal_l2_l3).permute(2, 0, 1)\n        sagittal_l3_l4 = torch.tensor(sagittal_l3_l4).permute(2, 0, 1)\n        sagittal_l4_l5 = torch.tensor(sagittal_l4_l5).permute(2, 0, 1)\n        sagittal_l5_s1 = torch.tensor(sagittal_l5_s1).permute(2, 0, 1)\n        axial_l1_l2 = torch.tensor(axial_l1_l2).permute(2, 0, 1)\n        axial_l2_l3 = torch.tensor(axial_l2_l3).permute(2, 0, 1)\n        axial_l3_l4 = torch.tensor(axial_l3_l4).permute(2, 0, 1)\n        axial_l4_l5 = torch.tensor(axial_l4_l5).permute(2, 0, 1)\n        axial_l5_s1 = torch.tensor(axial_l5_s1).permute(2, 0, 1)\n        \n\n        return (study_id, sagittal_l1_l2, axial_l1_l2, sagittal_l2_l3,\n                 axial_l2_l3, sagittal_l3_l4, axial_l3_l4, sagittal_l4_l5, axial_l4_l5, sagittal_l5_s1,\n                   axial_l5_s1)\n            \n\n\ndef TestLoader(study_ids: list, labels_path: Path, test_path:Path) -> tuple[DataLoader, DataLoader]:\n    return CustomDataset(study_ids, labels_path, test_path, transforms_val)","metadata":{"execution":{"iopub.status.busy":"2024-10-08T15:01:50.430787Z","iopub.execute_input":"2024-10-08T15:01:50.431108Z","iopub.status.idle":"2024-10-08T15:01:54.766736Z","shell.execute_reply.started":"2024-10-08T15:01:50.431081Z","shell.execute_reply":"2024-10-08T15:01:54.765726Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\n\nimport torch.nn as nn\nimport timm\n\nclass CustomRain(nn.Module):\n    def __init__(self, num_classes=75, pretrained=True):\n        super().__init__()\n        self.stack_t1 = timm.create_model(\n                                    \"timm/edgenext_base.in21k_ft_in1k\",\n                                    pretrained=pretrained, \n                                    features_only=False,\n                                    in_chans=3,\n                                    num_classes=128,\n                                    )\n\n        self.stack_t2 = timm.create_model(\n                                    \"timm/edgenext_base.in21k_ft_in1k\",\n                                    pretrained=pretrained,\n                                    features_only=False,\n                                    in_chans=3,\n                                    num_classes=128,\n                                    )\n        self.stack_axial = timm.create_model(\n                                    \"timm/edgenext_base.in21k_ft_in1k\",\n                                    pretrained=pretrained,\n                                    features_only=False,\n                                    in_chans=3,\n                                    num_classes=128,\n                                    )\n        self.stack = timm.create_model(\n                                    \"timm/edgenext_base.in21k_ft_in1k\",\n                                    pretrained=pretrained,\n                                    features_only=False,\n                                    in_chans=30,\n                                    num_classes=128,)\n\n        \n        self.head_l = nn.Sequential(\n            nn.Linear(128*2, 128),\n            nn.ReLU(),\n        )\n        self.head_r = nn.Sequential(\n            nn.Linear(128*2, 128),\n            nn.ReLU(),\n        )\n        self.head_t2 = nn.Sequential(\n            nn.Linear(128*3, 128),\n            nn.ReLU(),\n        )\n\n        self.scs = nn.Sequential(\n            nn.Linear(128*4, 128),\n            nn.ReLU(),\n            nn.Linear(128, 3),\n        )\n        self.nfn_l = nn.Sequential(\n            nn.Linear(128*4, 128),\n            nn.ReLU(),\n            nn.Linear(128, 3),\n        )\n        self.nfn_r = nn.Sequential(\n            nn.Linear(128*4, 128),\n            nn.ReLU(),\n            nn.Linear(128, 3),\n        )\n        self.ss_left = nn.Sequential(\n            nn.Linear(128*4, 128),\n            nn.ReLU(),\n            nn.Linear(128, 3),\n        )\n        self.ss_right = nn.Sequential(\n            nn.Linear(128*4, 128),\n            nn.ReLU(),\n            nn.Linear(128, 3),\n        )\n        # self.ascension_callback = AscensionCallback(margin=0.15)\n        self.level_embeddings = nn.Embedding(5, 128)\n\n    def forward(self, stack, level_idx):\n        t1_l1 = stack[:, :3, :, :]\n        t1_l2 = stack[:, 3:6, :, :]\n        t1_r1 = stack[:, 6:9, :, :]\n        t1_r2 = stack[:, 9:12, :, :]\n\n        t2_1 = stack[:, 12:15, :, :]\n        t2_2 = stack[:, 15:18, :, :]\n        t2_3 = stack[:, 18:21, :, :]\n\n        axial_left = stack[:, 21:24, :, :]\n        axial_center = stack[:, 24:27, :, :]\n        axial_right = stack[:, 27:, :, :]\n        \n\n        x_t1_l1 = self.stack_t1(t1_l1)\n        x_t1_l2= self.stack_t1(t1_l2)\n\n        x_t1_r1 = self.stack_t1(t1_r1)\n        x_t1_r2 = self.stack_t1(t1_r2)\n\n        x_t1_l = torch.cat((x_t1_l1, x_t1_l2), dim=1)\n        x_t1_l = self.head_l(x_t1_l)\n\n        x_t1_r = torch.cat((x_t1_r1, x_t1_r2), dim=1)\n        x_t1_r = self.head_r(x_t1_r)\n\n        x_t2_1 = self.stack_t2(t2_1)\n        x_t2_2 = self.stack_t2(t2_2)\n        x_t2_3 = self.stack_t2(t2_3)\n\n        x_t2 = torch.cat((x_t2_1, x_t2_2, x_t2_3), dim=1)\n        x_t2 = self.head_t2(x_t2)\n\n        x_axial_left = self.stack_axial(axial_left)\n        x_axial_center = self.stack_axial(axial_center)\n        x_axial_right = self.stack_axial(axial_right)\n        # stack[:, :21, :, :] = torch.rot90(stack[:, :21, :, :], k=3, dims=[2, 3])\n        x_stack = self.stack(stack)\n        bs = x_t2.size(0)\n        level_embeddings = self.level_embeddings(level_idx)\n        level_embeddings = level_embeddings.expand(bs, -1)\n\n        # SCS NFN SS\n        SCS = torch.cat((x_t2, x_axial_center, x_stack, level_embeddings), dim=1)\n\n        LNFN = torch.cat((x_t1_l, x_axial_left, x_stack, level_embeddings), dim=1)\n        RNFN = torch.cat((x_t1_r, x_axial_right, x_stack, level_embeddings), dim=1)\n\n        SS_left = torch.cat((x_axial_left, x_t1_l, x_stack, level_embeddings), dim=1)\n        SS_right = torch.cat((x_axial_right, x_t1_r, x_stack, level_embeddings), dim=1)\n\n        x_scs = self.scs(SCS)\n\n        x_nfn_left = self.nfn_l(LNFN)\n        x_nfn_right = self.nfn_r(RNFN)\n\n        x_ss_left = self.ss_left(SS_left)\n        x_ss_right = self.ss_right(SS_right)\n\n\n        # self._ascension_callback()\n        x = torch.cat([x_scs, x_nfn_left, x_nfn_right, x_ss_left, x_ss_right], dim=1)\n        \n        # # Apply the ascension callback for cutpoint clipping\n        \n\n        # Final output (15 ordinal classes and combined output)\n        assert x.shape[1]==15, f\"Expected 15 classes, got {x.shape[1]}\"\n        return x","metadata":{"execution":{"iopub.status.busy":"2024-10-08T15:01:54.768620Z","iopub.execute_input":"2024-10-08T15:01:54.769234Z","iopub.status.idle":"2024-10-08T15:01:56.620364Z","shell.execute_reply.started":"2024-10-08T15:01:54.769179Z","shell.execute_reply":"2024-10-08T15:01:56.619557Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"CONDITIONS = [\n    'spinal_canal_stenosis', \n    'left_neural_foraminal_narrowing', \n    'right_neural_foraminal_narrowing',\n    'left_subarticular_stenosis',\n    'right_subarticular_stenosis'\n]\n\nLEVELS = [\n    'l1_l2',\n    'l2_l3',\n    'l3_l4',\n    'l4_l5',\n    'l5_s1',\n]","metadata":{"execution":{"iopub.status.busy":"2024-10-08T15:01:56.621416Z","iopub.execute_input":"2024-10-08T15:01:56.621691Z","iopub.status.idle":"2024-10-08T15:01:56.626389Z","shell.execute_reply.started":"2024-10-08T15:01:56.621666Z","shell.execute_reply":"2024-10-08T15:01:56.625346Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def reorder_labels(labels):\n    \n    # Create an empty tensor to hold the reversed labels\n    original_labels = torch.empty_like(labels)\n\n    # Number of total sets\n    num_sets = len(labels) // 3\n\n    # Reversing the order based on % 5 position\n    for set_index in range(num_sets):\n        source_start = set_index * 3\n        # Calculate destination start based on the modulus operation\n        dest_start = ((set_index % 5) * 15) + (set_index // 5 * 3)\n        original_labels[dest_start:dest_start + 3] = labels[source_start:source_start + 3]\n    return original_labels\n\nfrom tqdm import tqdm\nimport torch\nimport pandas as pd\nfrom torch.utils.data import DataLoader\n\nclass Settings:\n    number_of_classes = 15\n    batch_size = 1\n    pretrain = False\n    path = \"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/test_images\"\n    description_path = \"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/test_series_descriptions.csv\"\n    test_path = path\n    model_path = \"/kaggle/input/models/RainDrop_0.4910611528158188_fold_1.pt\"\n    N_LABELS = 25\n    LABELS = ['normal_mild','moderate','severe']\n\nsettings = Settings()\nmodel = CustomRain(settings.number_of_classes, settings.pretrain)\nmodel.load_state_dict(torch.load(settings.model_path))\n\ntest_df = pd.read_csv(settings.description_path)\nstudy_ids = list(test_df['study_id'].unique())\ndata_loader = TestLoader(study_ids, settings.description_path, settings.test_path)\ntest_loader = DataLoader(data_loader, batch_size=settings.batch_size, shuffle=False)\nsubmissions = pd.DataFrame()\nmodel.cuda()\nmodel.eval()\n\nrow_names = []\ny_preds = []\n\ndevice = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n\nfor (study_id, sagittal_l1_l2, axial_l1_l2, sagittal_l2_l3,\n                 axial_l2_l3, sagittal_l3_l4, axial_l3_l4, sagittal_l4_l5, axial_l4_l5, sagittal_l5_s1,\n                   axial_l5_s1,) in tqdm(test_loader):\n    with torch.no_grad():\n        sagittal_l1_l2 = sagittal_l1_l2.cuda()\n        sagittal_l2_l3 = sagittal_l2_l3.cuda()\n        sagittal_l3_l4 = sagittal_l3_l4.cuda()\n        sagittal_l4_l5 = sagittal_l4_l5.cuda()\n        sagittal_l5_s1 = sagittal_l5_s1.cuda()\n        axial_l1_l2 = axial_l1_l2.cuda()\n        axial_l2_l3 = axial_l2_l3.cuda()\n        axial_l3_l4 = axial_l3_l4.cuda()\n        axial_l4_l5 = axial_l4_l5.cuda()\n        axial_l5_s1 = axial_l5_s1.cuda()\n        \n        output1 = model(sagittal_l1_l2, torch.tensor(0, device=device)).squeeze()\n        output2 = model(sagittal_l2_l3, torch.tensor(1, device=device)).squeeze()\n        output3 = model(sagittal_l3_l4, torch.tensor(2, device=device)).squeeze()\n        output4 = model(sagittal_l4_l5, torch.tensor(3, device=device)).squeeze()\n        output5 = model(sagittal_l5_s1, torch.tensor(4, device=device)).squeeze()\n        output = torch.cat([output1, output2, output3, output4, output5], dim=0)\n        output = reorder_labels(output)\n        pred_per_study = np.zeros((25, 3))\n\n        for cond in CONDITIONS:\n            for level in LEVELS:\n                row_names.append(str(study_id.tolist()[0]) + '_' + cond + '_' + level)\n\n        for col in range(settings.N_LABELS):\n            pred = output[col*3:col*3+3]\n            y_pred = pred.float().softmax(0).cpu().numpy()\n            pred_per_study[col] += y_pred\n        y_preds.append(pred_per_study)\ny_preds = np.concatenate(y_preds, axis=0)\n\nsubmissions['row_id'] = row_names\nsubmissions[settings.LABELS] = y_preds\nsubmissions","metadata":{"execution":{"iopub.status.busy":"2024-10-08T15:01:56.627832Z","iopub.execute_input":"2024-10-08T15:01:56.628332Z","iopub.status.idle":"2024-10-08T15:02:11.404251Z","shell.execute_reply.started":"2024-10-08T15:01:56.628299Z","shell.execute_reply":"2024-10-08T15:02:11.403242Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install --quiet /kaggle/input/timm_3d_deps/other/initial/10/pydicom/pydicom/pydicom-2.4.4-py3-none-any.whl\n!pip install timm_3d --no-index --quiet --find-links=/kaggle/input/timm_3d_deps/other/initial/10/timm_3d/\n!pip install torchio --no-index --quiet --find-links=/kaggle/input/timm_3d_deps/other/initial/10/torchio/\n!pip install itk --no-index --quiet --find-links=/kaggle/input/timm_3d_deps/other/initial/10/itk/itk\n!pip install skorch --no-index --quiet --find-links=/kaggle/input/timm_3d_deps/other/initial/10/skorch/skorch\n!pip install spacecutter --no-index --quiet --find-links=/kaggle/input/timm_3d_deps/other/initial/10/spacecutter/\n!pip install open3d --no-index --quiet --find-links=/kaggle/input/timm_3d_deps/other/initial/10/open3d","metadata":{"execution":{"iopub.status.busy":"2024-10-08T15:02:11.405725Z","iopub.execute_input":"2024-10-08T15:02:11.406122Z","iopub.status.idle":"2024-10-08T15:04:41.563924Z","shell.execute_reply.started":"2024-10-08T15:02:11.406087Z","shell.execute_reply":"2024-10-08T15:04:41.562658Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data_path = \"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/\"","metadata":{"execution":{"iopub.status.busy":"2024-10-08T15:04:41.565656Z","iopub.execute_input":"2024-10-08T15:04:41.565983Z","iopub.status.idle":"2024-10-08T15:04:41.570871Z","shell.execute_reply.started":"2024-10-08T15:04:41.565955Z","shell.execute_reply":"2024-10-08T15:04:41.570012Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import pandas as pd\n\ndef retrieve_test_data(data_path):\n    test_df = pd.read_csv(data_path + 'test_series_descriptions.csv')\n\n    return test_df\n\nretrieve_test_data(data_path)","metadata":{"execution":{"iopub.status.busy":"2024-10-08T15:04:41.572047Z","iopub.execute_input":"2024-10-08T15:04:41.572347Z","iopub.status.idle":"2024-10-08T15:04:41.593737Z","shell.execute_reply.started":"2024-10-08T15:04:41.572325Z","shell.execute_reply":"2024-10-08T15:04:41.592781Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\n\ndef retrieve_image_paths(base_path, study_id, series_id):\n    series_dir = os.path.join(base_path, str(study_id), str(series_id))\n    images = os.listdir(series_dir)\n    image_paths = [os.path.join(series_dir, img) for img in images]\n    return image_paths","metadata":{"execution":{"iopub.status.busy":"2024-10-08T15:04:41.594906Z","iopub.execute_input":"2024-10-08T15:04:41.595244Z","iopub.status.idle":"2024-10-08T15:04:41.601076Z","shell.execute_reply.started":"2024-10-08T15:04:41.595189Z","shell.execute_reply":"2024-10-08T15:04:41.600140Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import open3d as o3d\nfrom pydicom import dcmread\nimport math\nimport numpy as np\nimport cv2\nimport copy\n\ndef read_study_as_pcd(dir_path, series_types_dict=None, downsampling_factor=1, img_size=(256, 256)):\n    pcd_overall = o3d.geometry.PointCloud()\n\n    for path in glob.glob(os.path.join(dir_path, \"**/*.dcm\"), recursive=True):\n        dicom_slice = dcmread(path)\n\n        series_id = os.path.basename(os.path.dirname(path))\n        study_id = os.path.basename(os.path.dirname(os.path.dirname(path)))\n        if series_types_dict is None or int(series_id) not in series_types_dict:\n            series_desc = dicom_slice.SeriesDescription\n        else:\n            series_desc = series_types_dict[int(series_id)]\n            series_desc = series_desc.split(\" \")[-1]\n\n        x_orig, y_orig = dicom_slice.pixel_array.shape\n        img = np.expand_dims(cv2.resize(dicom_slice.pixel_array, img_size, interpolation=cv2.INTER_AREA), -1)\n        x, y, z = np.where(img)\n\n        downsampling_factor_iter = max(downsampling_factor, int(math.ceil(len(x) / 6e6)))\n\n        index_voxel = np.vstack((x, y, z))[:, ::downsampling_factor_iter]\n        grid_index_array = index_voxel.T\n        pcd = o3d.geometry.PointCloud(o3d.utility.Vector3dVector(grid_index_array.astype(np.float64)))\n\n        vals = np.expand_dims(img[x, y, z][::downsampling_factor_iter], -1)\n        if series_desc == \"T1\":\n            vals = np.pad(vals, ((0, 0), (0, 2)))\n        elif series_desc == \"T2\":\n            vals = np.pad(vals, ((0, 0), (1, 1)))\n        elif series_desc == \"T2/STIR\":\n            vals = np.pad(vals, ((0, 0), (2, 0)))\n        else:\n            raise ValueError(f\"Unknown series desc: {series_desc}\")\n\n        pcd.colors = o3d.utility.Vector3dVector(vals.astype(np.float64))\n\n        dX, dY = dicom_slice.PixelSpacing\n        dZ = dicom_slice.SliceThickness\n\n        X = np.array(list(dicom_slice.ImageOrientationPatient[:3]) + [0]) * dX\n        Y = np.array(list(dicom_slice.ImageOrientationPatient[3:]) + [0]) * dY\n\n        for z in range(int(dZ)):\n            pos = list(dicom_slice.ImagePositionPatient)\n            if series_desc == \"T2\":\n                pos[-1] += z\n            else:\n                pos[0] += z\n            S = np.array(pos + [1])\n\n            transform_matrix = np.array([X, Y, np.zeros(len(X)), S]).T\n            transform_matrix = transform_matrix @ np.matrix(\n                [[0, y_orig / img_size[1], 0, 0],\n                 [x_orig / img_size[0], 0, 0, 0],\n                 [0, 0, 1, 0],\n                 [0, 0, 0, 1]]\n            )\n\n            pcd_overall += copy.deepcopy(pcd).transform(transform_matrix)\n\n    return pcd_overall\n\n\n\ndef read_study_as_voxel_grid(dir_path, series_type_dict=None, downsampling_factor=1, img_size=(256, 256)):\n    pcd_overall = read_study_as_pcd(dir_path,\n                                    series_types_dict=series_type_dict,\n                                    downsampling_factor=downsampling_factor,\n                                    img_size=img_size)\n    box = pcd_overall.get_axis_aligned_bounding_box()\n\n    max_b = np.array(box.get_max_bound())\n    min_b = np.array(box.get_min_bound())\n\n    pts = (np.array(pcd_overall.points) - (min_b)) * (\n                (img_size[0] - 1, img_size[0] - 1, img_size[0] - 1) / (max_b - min_b))\n    coords = np.round(pts).astype(np.int32)\n    vals = np.array(pcd_overall.colors, dtype=np.float16)\n\n    grid = np.zeros((3, img_size[0], img_size[0], img_size[0]), dtype=np.float16)\n    indices = coords[:, 0], coords[:, 1], coords[:, 2]\n\n    np.maximum.at(grid[0], indices, vals[:, 0])\n    np.maximum.at(grid[1], indices, vals[:, 1])\n    np.maximum.at(grid[2], indices, vals[:, 2])\n\n\n    return grid","metadata":{"execution":{"iopub.status.busy":"2024-10-08T15:04:41.602640Z","iopub.execute_input":"2024-10-08T15:04:41.602975Z","iopub.status.idle":"2024-10-08T15:04:43.529770Z","shell.execute_reply.started":"2024-10-08T15:04:41.602935Z","shell.execute_reply":"2024-10-08T15:04:43.528863Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np\nfrom torch.utils.data import Dataset, DataLoader\nimport torchio as tio\nimport torch.nn as nn\nimport pydicom\n\nCONDITIONS = {\n    \"Sagittal T2/STIR\": [\"Spinal Canal Stenosis\"],\n    \"Axial T2\": [\"Left Subarticular Stenosis\", \"Right Subarticular Stenosis\"],\n    \"Sagittal T1\": [\"Left Neural Foraminal Narrowing\", \"Right Neural Foraminal Narrowing\"],\n}\n\n\nclass PatientLevelTestset(Dataset):\n    def __init__(self,\n                 base_path: str,\n                 dataframe: pd.DataFrame,\n                 transform_3d=None):\n        self.base_path = base_path\n\n        self.dataframe = (dataframe[['study_id', \"series_id\", \"series_description\"]]\n                          .drop_duplicates())\n\n        self.subjects = self.dataframe[['study_id']].drop_duplicates().reset_index(drop=True)\n        self.series_descs = {e[0]: e[1] for e in self.dataframe[[\"series_id\", \"series_description\"]].drop_duplicates().values}\n\n        self.transform_3d = transform_3d\n\n    def __len__(self):\n        return len(self.subjects)\n\n    def __getitem__(self, index):\n        curr = self.subjects.iloc[index]\n        study_path = os.path.join(self.base_path, str(curr[\"study_id\"]))\n\n        study_images = read_study_as_voxel_grid(study_path, self.series_descs)\n\n        if self.transform_3d is not None:\n            study_images = self.transform_3d(torch.FloatTensor(study_images))  # .data\n            return study_images.to(torch.half), str(curr[\"study_id\"])\n\n        return torch.HalfTensor(study_images), str(curr[\"study_id\"])","metadata":{"execution":{"iopub.status.busy":"2024-10-08T15:04:43.530941Z","iopub.execute_input":"2024-10-08T15:04:43.531407Z","iopub.status.idle":"2024-10-08T15:04:44.323064Z","shell.execute_reply.started":"2024-10-08T15:04:43.531381Z","shell.execute_reply":"2024-10-08T15:04:44.322063Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"transform_3d = tio.Compose([\n    tio.RescaleIntensity([0, 1]),\n])","metadata":{"execution":{"iopub.status.busy":"2024-10-08T15:04:44.324354Z","iopub.execute_input":"2024-10-08T15:04:44.324681Z","iopub.status.idle":"2024-10-08T15:04:44.329157Z","shell.execute_reply.started":"2024-10-08T15:04:44.324655Z","shell.execute_reply":"2024-10-08T15:04:44.328253Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def create_subject_level_testset_and_loader(df: pd.DataFrame,\n                                             transform_3d,\n                                             base_path: str,\n                                             batch_size=1,\n                                             num_workers=0):\n    testset = PatientLevelTestset(base_path, df, transform_3d=transform_3d)\n    test_loader = DataLoader(testset, batch_size=batch_size, shuffle=False, num_workers=num_workers)\n\n    return testset, test_loader","metadata":{"execution":{"iopub.status.busy":"2024-10-08T15:04:44.330357Z","iopub.execute_input":"2024-10-08T15:04:44.330670Z","iopub.status.idle":"2024-10-08T15:04:44.341738Z","shell.execute_reply.started":"2024-10-08T15:04:44.330636Z","shell.execute_reply":"2024-10-08T15:04:44.340971Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\n\ndata = retrieve_test_data(data_path)\ndataset, dataloader = create_subject_level_testset_and_loader(data, transform_3d, os.path.join(data_path, \"test_images\"))","metadata":{"execution":{"iopub.status.busy":"2024-10-08T15:04:44.342817Z","iopub.execute_input":"2024-10-08T15:04:44.343088Z","iopub.status.idle":"2024-10-08T15:04:44.364800Z","shell.execute_reply.started":"2024-10-08T15:04:44.343065Z","shell.execute_reply":"2024-10-08T15:04:44.363936Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport glob \nimport torch\n\ngrid = dataset[0][0]\n\nfig, axs = plt.subplots(3, 3)\n\naxs[0, 0].imshow(grid[0, 128])\naxs[1, 0].imshow(grid[1, 128])\naxs[2, 0].imshow(grid[2, 128])\n\naxs[0, 1].imshow(grid[0, :, 128])\naxs[1, 1].imshow(grid[1, :, 128])\naxs[2, 1].imshow(grid[2, :, 128])\n\naxs[0, 2].imshow(grid[0, :, :, 128])\naxs[1, 2].imshow(grid[1, :, :, 128])\naxs[2, 2].imshow(grid[2, :, :, 128])\n\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-10-08T15:04:44.365954Z","iopub.execute_input":"2024-10-08T15:04:44.366261Z","iopub.status.idle":"2024-10-08T15:05:05.949480Z","shell.execute_reply.started":"2024-10-08T15:04:44.366237Z","shell.execute_reply":"2024-10-08T15:05:05.948458Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\ndevice = torch.device(\"cuda\") if torch.cuda.is_available() else \"cpu\"","metadata":{"execution":{"iopub.status.busy":"2024-10-08T15:05:05.950930Z","iopub.execute_input":"2024-10-08T15:05:05.951315Z","iopub.status.idle":"2024-10-08T15:05:05.956362Z","shell.execute_reply.started":"2024-10-08T15:05:05.951285Z","shell.execute_reply":"2024-10-08T15:05:05.955356Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch.nn as nn\nimport timm_3d\nfrom spacecutter import *\nfrom spacecutter.losses import *\nfrom spacecutter.models import *\nfrom spacecutter.callbacks import *\n\n\nclass CNN_Model_3D_Multihead(nn.Module):\n    def __init__(self,\n                 backbone=\"efficientnet_lite0\",\n                 in_chans=1,\n                 out_classes=5,\n                 cutpoint_margin=0.15,\n                 pretrained=False):\n        super(CNN_Model_3D_Multihead, self).__init__()\n        self.out_classes = out_classes\n\n        self.encoder = timm_3d.create_model(\n            backbone,\n            features_only=False,\n            drop_rate=0,\n            drop_path_rate=0,\n            pretrained=pretrained,\n            in_chans=in_chans,\n            global_pool=\"max\"\n        )\n        if \"efficientnet\" in backbone:\n            head_in_dim = self.encoder.classifier.in_features\n            self.encoder.classifier = nn.Sequential(\n                nn.LayerNorm(head_in_dim),\n                nn.Dropout(0),\n            )\n\n        elif \"vit\" in backbone:\n            self.encoder.head.drop = nn.Dropout(0)\n            head_in_dim = self.encoder.head.fc.in_features\n            self.encoder.head.fc = nn.Identity()\n\n        self.heads = nn.ModuleList(\n            [nn.Sequential(\n                nn.Linear(head_in_dim, 1),\n                LogisticCumulativeLink(3)\n            ) for i in range(out_classes)]\n        )\n\n        self.ascension_callback = AscensionCallback(margin=cutpoint_margin)\n\n    def forward(self, x):\n        feat = self.encoder(x)\n        return torch.swapaxes(torch.stack([head(feat) for head in self.heads]), 0, 1)\n\n    def _ascension_callback(self):\n        for head in self.heads:\n            self.ascension_callback.clip(head[-1])","metadata":{"execution":{"iopub.status.busy":"2024-10-08T15:05:05.957705Z","iopub.execute_input":"2024-10-08T15:05:05.958087Z","iopub.status.idle":"2024-10-08T15:05:06.137716Z","shell.execute_reply.started":"2024-10-08T15:05:05.958041Z","shell.execute_reply":"2024-10-08T15:05:06.136888Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = CNN_Model_3D_Multihead(backbone=\"maxvit_rmlp_small_rw_256\", in_chans=3, out_classes=25).to(device)\nmodel.load_state_dict(torch.load(\"/kaggle/input/models/maxvit_rmlp_small_rw_256_256_v2_fold_3_22.pt\"))","metadata":{"execution":{"iopub.status.busy":"2024-10-08T15:05:06.139154Z","iopub.execute_input":"2024-10-08T15:05:06.139554Z","iopub.status.idle":"2024-10-08T15:05:10.449298Z","shell.execute_reply.started":"2024-10-08T15:05:06.139519Z","shell.execute_reply":"2024-10-08T15:05:10.448261Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"CONDITIONS = {\n    \"Sagittal T2/STIR\": [\"spinal_canal_stenosis\"],\n    \"Axial T2\": [\"left_subarticular_stenosis\", \"right_subarticular_stenosis\"],\n    \"Sagittal T1\": [\"left_neural_foraminal_narrowing\", \"right_neural_foraminal_narrowing\"],\n}\n\nALL_CONDITIONS = sorted([\"spinal_canal_stenosis\", \"left_subarticular_stenosis\", \"right_subarticular_stenosis\", \"left_neural_foraminal_narrowing\", \"right_neural_foraminal_narrowing\"])\nLEVELS = [\"l1_l2\", \"l2_l3\", \"l3_l4\", \"l4_l5\", \"l5_s1\"]\n\nresults_df = pd.DataFrame({\"row_id\":[], \"normal_mild\": [], \"moderate\": [], \"severe\": []})\n\nALL_CONDITIONS","metadata":{"execution":{"iopub.status.busy":"2024-10-08T15:05:10.450636Z","iopub.execute_input":"2024-10-08T15:05:10.451019Z","iopub.status.idle":"2024-10-08T15:05:10.460498Z","shell.execute_reply.started":"2024-10-08T15:05:10.450986Z","shell.execute_reply":"2024-10-08T15:05:10.459508Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Pre-populate results df\nimport glob\nimport os\n\nstudy_ids = glob.glob(\"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/test_images/*\")\nstudy_ids = [os.path.basename(e) for e in study_ids]\n\nresults_df = pd.DataFrame({\"row_id\":[], \"normal_mild\": [], \"moderate\": [], \"severe\": []})\nfor study_id in study_ids:\n    for condition in ALL_CONDITIONS:\n        for level in LEVELS:\n            row_id = f\"{study_id}_{condition}_{level}\"\n            results_df = results_df._append({\"row_id\": row_id, \"normal_mild\": 1/3, \"moderate\": 1/3, \"severe\": 1/3}, ignore_index=True)","metadata":{"execution":{"iopub.status.busy":"2024-10-08T15:05:10.461868Z","iopub.execute_input":"2024-10-08T15:05:10.462379Z","iopub.status.idle":"2024-10-08T15:05:10.500965Z","shell.execute_reply.started":"2024-10-08T15:05:10.462345Z","shell.execute_reply":"2024-10-08T15:05:10.500140Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dataset.dataframe","metadata":{"execution":{"iopub.status.busy":"2024-10-08T15:05:10.502247Z","iopub.execute_input":"2024-10-08T15:05:10.502843Z","iopub.status.idle":"2024-10-08T15:05:10.512351Z","shell.execute_reply.started":"2024-10-08T15:05:10.502810Z","shell.execute_reply":"2024-10-08T15:05:10.511341Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nfrom torch.cuda.amp import autocast\nimport time\n\nstart_time = time.time()\n\nwith torch.no_grad():\n    with autocast(dtype=torch.float16):\n        model.eval()\n\n        for images, study_id in dataloader:\n            output = model(images.to(device))\n            for i, batch_out in enumerate(output):\n                batch_out = output.cpu().numpy()[i]\n                for index, level in enumerate(batch_out):\n                    row_id = f\"{study_id[i]}_{ALL_CONDITIONS[index // 5]}_{LEVELS[index % 5]}\"\n                    results_df.loc[results_df.row_id == row_id,'normal_mild'] = level[0]\n                    results_df.loc[results_df.row_id == row_id,'moderate'] = level[1]\n                    results_df.loc[results_df.row_id == row_id,'severe'] = level[2]\n                \nprint(\"--- %s seconds ---\" % (time.time() - start_time))","metadata":{"execution":{"iopub.status.busy":"2024-10-08T15:05:10.513961Z","iopub.execute_input":"2024-10-08T15:05:10.514321Z","iopub.status.idle":"2024-10-08T15:05:31.138764Z","shell.execute_reply.started":"2024-10-08T15:05:10.514291Z","shell.execute_reply":"2024-10-08T15:05:31.137752Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"results_df","metadata":{"execution":{"iopub.status.busy":"2024-10-08T15:05:31.140182Z","iopub.execute_input":"2024-10-08T15:05:31.141457Z","iopub.status.idle":"2024-10-08T15:05:31.157743Z","shell.execute_reply.started":"2024-10-08T15:05:31.141429Z","shell.execute_reply":"2024-10-08T15:05:31.156693Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submissions = submissions.set_index(\"row_id\")\npred_df = pred_df.set_index(\"row_id\")\nresults_df = results_df.set_index(\"row_id\")\n# new_order_pred_df = submissions.copy()\n# new_order_results_df = submissions.copy()\n\n# for col in [\"normal_mild\", \"moderate\", \"severe\"]:\n#     new_order_pred_df[col] = pred_df[col]\n    \n# for col in [\"normal_mild\", \"moderate\", \"severe\"]:\n#     new_order_results_df[col] = results_df[col]\n\n\nnew_csv = submissions[[\"normal_mild\", \"moderate\", \"severe\"]] * 0.4 + pred_df[[\"normal_mild\", \"moderate\", \"severe\"]] * 0.4 + results_df[[\"normal_mild\", \"moderate\", \"severe\"]] * 0.2\nnew_csv.to_csv('submission.csv')\npd.read_csv('submission.csv')","metadata":{"execution":{"iopub.status.busy":"2024-10-08T15:05:31.158983Z","iopub.execute_input":"2024-10-08T15:05:31.159322Z","iopub.status.idle":"2024-10-08T15:05:31.187707Z","shell.execute_reply.started":"2024-10-08T15:05:31.159290Z","shell.execute_reply":"2024-10-08T15:05:31.186618Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# new_order_pred_df.reset_index(inplace=True)\n# submissions.reset_index(inplace=True)\n# new_csv = pd.DataFrame(columns = submissions.columns)\n\n# for (index_1, row_1), (index_2, row_2) in zip(submissions.iterrows(), new_order_pred_df.iterrows()):\n#     if index_1 % 25 in [0,1,2,3,4]:\n#         new_csv.loc[len(new_csv)] = row_1\n#     else:\n#         new_csv.loc[len(new_csv)] = row_2\n\n# new_csv\n        ","metadata":{"execution":{"iopub.status.busy":"2024-10-08T15:05:31.189400Z","iopub.execute_input":"2024-10-08T15:05:31.189700Z","iopub.status.idle":"2024-10-08T15:05:31.193746Z","shell.execute_reply.started":"2024-10-08T15:05:31.189675Z","shell.execute_reply":"2024-10-08T15:05:31.192725Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# new_csv.to_csv('submission.csv', index=False)\n# pd.read_csv('submission.csv')","metadata":{"execution":{"iopub.status.busy":"2024-10-08T15:05:31.194869Z","iopub.execute_input":"2024-10-08T15:05:31.195176Z","iopub.status.idle":"2024-10-08T15:05:31.202763Z","shell.execute_reply.started":"2024-10-08T15:05:31.195151Z","shell.execute_reply":"2024-10-08T15:05:31.202025Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}