{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.14","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":71549,"databundleVersionId":8561470,"sourceType":"competition"},{"sourceId":9429410,"sourceType":"datasetVersion","datasetId":5728485},{"sourceId":9429420,"sourceType":"datasetVersion","datasetId":5728490},{"sourceId":9429425,"sourceType":"datasetVersion","datasetId":5728492},{"sourceId":9429547,"sourceType":"datasetVersion","datasetId":5728586},{"sourceId":9429572,"sourceType":"datasetVersion","datasetId":5728603},{"sourceId":9429628,"sourceType":"datasetVersion","datasetId":5728634},{"sourceId":10423778,"sourceType":"datasetVersion","datasetId":5728629},{"sourceId":197208471,"sourceType":"kernelVersion"},{"sourceId":225492,"sourceType":"modelInstanceVersion","modelInstanceId":101636,"modelId":125842},{"sourceId":225524,"sourceType":"modelInstanceVersion","modelInstanceId":190299,"modelId":118918},{"sourceId":299905,"sourceType":"modelInstanceVersion","modelInstanceId":256225,"modelId":277538},{"sourceId":329492,"sourceType":"modelInstanceVersion","modelInstanceId":196133,"modelId":218032},{"sourceId":329737,"sourceType":"modelInstanceVersion","modelInstanceId":256679,"modelId":278008},{"sourceId":330146,"sourceType":"modelInstanceVersion","modelInstanceId":256683,"modelId":278011},{"sourceId":376544,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":256273,"modelId":277590}],"dockerImageVersionId":30805,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"The goal of this notebook is to submit our models to the 2024 RSNA lumbar challenge. ","metadata":{}},{"cell_type":"code","source":"train = False ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T17:29:40.912310Z","iopub.execute_input":"2025-04-28T17:29:40.912554Z","iopub.status.idle":"2025-04-28T17:29:40.936778Z","shell.execute_reply.started":"2025-04-28T17:29:40.912526Z","shell.execute_reply":"2025-04-28T17:29:40.936208Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Cuda Setup ","metadata":{}},{"cell_type":"code","source":"import torch\n\n# Check if GPU is available and set the device accordingly\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\nif torch.cuda.is_available():\n    print(f\"GPU is available. Using: {torch.cuda.get_device_name(0)}\")\nelse:\n    print(\"GPU not available. Using CPU.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T17:29:40.938107Z","iopub.execute_input":"2025-04-28T17:29:40.938343Z","iopub.status.idle":"2025-04-28T17:29:44.428525Z","shell.execute_reply.started":"2025-04-28T17:29:40.938319Z","shell.execute_reply":"2025-04-28T17:29:44.427569Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Image SCT class ","metadata":{}},{"cell_type":"code","source":"import os\nimport numpy as np\nimport nibabel as nib\nimport logging\nfrom copy import deepcopy\n\nlogger = logging.getLogger(__name__)\n\nclass Image(object):\n    \"\"\"\n    Compact version of SCT's Image Class (https://github.com/spinalcordtoolbox/spinalcordtoolbox/blob/master/spinalcordtoolbox/image.py#L245)\n    Create an object that behaves similarly to nibabel's image object. Useful additions include: dims, change_orientation and getNonZeroCoordinates.\n    \"\"\"\n\n    def __init__(self, param=None, hdr=None, orientation=None, absolutepath=None, dim=None):\n        \"\"\"\n        :param param: string indicating a path to a image file or an `Image` object.\n        \"\"\"\n\n        # initialization of all parameters\n        self.affine = None\n        self.data = None\n        self._path = None\n        self.ext = \"\"\n\n        if absolutepath is not None:\n            self._path = os.path.abspath(absolutepath)\n        \n        # Case 1: load an image from file\n        if isinstance(param, str):\n            self.loadFromPath(param)\n        # Case 2: create a copy of an existing `Image` object\n        elif isinstance(param, type(self)):\n            self.copy(param)\n        # Case 3: create a blank image from a list of dimensions\n        elif isinstance(param, list):\n            self.data = np.zeros(param)\n            self.hdr = hdr.copy() if hdr is not None else nib.Nifti1Header()\n            self.hdr.set_data_shape(self.data.shape)\n        # Case 4: create an image from an existing data array\n        elif isinstance(param, (np.ndarray, np.generic)):\n            self.data = param\n            self.hdr = hdr.copy() if hdr is not None else nib.Nifti1Header()\n            self.hdr.set_data_shape(self.data.shape)\n        else:\n            raise TypeError('Image constructor takes at least one argument.')\n    \n        # Fix any mismatch between the array's datatype and the header datatype\n        self.fix_header_dtype()\n\n    @property\n    def dim(self):\n        return get_dimension(self)\n    \n    @property\n    def orientation(self):\n        return get_orientation(self)\n    \n    @property\n    def absolutepath(self):\n        \"\"\"\n        Storage path (either actual or potential)\n\n        Notes:\n\n        - As several tools perform chdir() it's very important to have absolute paths\n        - When set, if relative:\n\n          - If it already existed, it becomes a new basename in the old dirname\n          - Else, it becomes absolute (shortcut)\n\n        Usually not directly touched (use `Image.save`), but in some cases it's\n        the best way to set it.\n        \"\"\"\n        return self._path\n    \n    @absolutepath.setter\n    def absolutepath(self, value):\n        if value is None:\n            self._path = None\n            return\n        elif not os.path.isabs(value) and self._path is not None:\n            value = os.path.join(os.path.dirname(self._path), value)\n        elif not os.path.isabs(value):\n            value = os.path.abspath(value)\n        self._path = value\n    \n    @property\n    def header(self):\n        return self.hdr\n\n    @header.setter\n    def header(self, value):\n        self.hdr = value\n\n    def __deepcopy__(self, memo):\n        return type(self)(deepcopy(self.data, memo), deepcopy(self.hdr, memo), deepcopy(self.orientation, memo), deepcopy(self.absolutepath, memo), deepcopy(self.dim, memo))\n\n    def copy(self, image=None):\n        if image is not None:\n            self.affine = deepcopy(image.affine)\n            self.data = deepcopy(image.data)\n            self.hdr = deepcopy(image.hdr)\n            self._path = deepcopy(image._path)\n        else:\n            return deepcopy(self)\n\n    def loadFromPath(self, path):\n        \"\"\"\n        This function load an image from an absolute path using nibabel library\n\n        :param path: path of the file from which the image will be loaded\n        :return:\n        \"\"\"\n\n        self.absolutepath = os.path.abspath(path)\n        im_file = nib.load(self.absolutepath, mmap=True)\n        self.affine = im_file.affine.copy()\n        self.data = np.asanyarray(im_file.dataobj)\n        self.hdr = im_file.header.copy()\n        if path != self.absolutepath:\n            logger.debug(\"Loaded %s (%s) orientation %s shape %s\", path, self.absolutepath, self.orientation, self.data.shape)\n        else:\n            logger.debug(\"Loaded %s orientation %s shape %s\", path, self.orientation, self.data.shape)\n\n    def change_orientation(self, orientation, inverse=False):\n        \"\"\"\n        Change orientation on image (in-place).\n\n        :param orientation: orientation string (SCT \"from\" convention)\n\n        :param inverse: if you think backwards, use this to specify that you actually\\\n                        want to transform *from* the specified orientation, not *to*\\\n                        it.\n\n        \"\"\"\n        change_orientation(self, orientation, self, inverse=inverse)\n        return self\n    \n    def getNonZeroCoordinates(self, sorting=None, reverse_coord=False):\n        \"\"\"\n        This function return all the non-zero coordinates that the image contains.\n        Coordinate list can also be sorted by x, y, z, or the value with the parameter sorting='x', sorting='y', sorting='z' or sorting='value'\n        If reverse_coord is True, coordinate are sorted from larger to smaller.\n\n        Removed Coordinate object\n        \"\"\"\n        n_dim = 1\n        if self.dim[3] == 1:\n            n_dim = 3\n        else:\n            n_dim = 4\n        if self.dim[2] == 1:\n            n_dim = 2\n\n        if n_dim == 3:\n            X, Y, Z = (self.data > 0).nonzero()\n            list_coordinates = [[X[i], Y[i], Z[i], self.data[X[i], Y[i], Z[i]]] for i in range(0, len(X))]\n        elif n_dim == 2:\n            try:\n                X, Y = (self.data > 0).nonzero()\n                list_coordinates = [[X[i], Y[i], 0, self.data[X[i], Y[i]]] for i in range(0, len(X))]\n            except ValueError:\n                X, Y, Z = (self.data > 0).nonzero()\n                list_coordinates = [[X[i], Y[i], 0, self.data[X[i], Y[i], 0]] for i in range(0, len(X))]\n\n        if sorting is not None:\n            if reverse_coord not in [True, False]:\n                raise ValueError('reverse_coord parameter must be a boolean')\n\n            if sorting == 'x':\n                list_coordinates = sorted(list_coordinates, key=lambda el: el[0], reverse=reverse_coord)\n            elif sorting == 'y':\n                list_coordinates = sorted(list_coordinates, key=lambda el: el[1], reverse=reverse_coord)\n            elif sorting == 'z':\n                list_coordinates = sorted(list_coordinates, key=lambda el: el[2], reverse=reverse_coord)\n            elif sorting == 'value':\n                list_coordinates = sorted(list_coordinates, key=lambda el: el[3], reverse=reverse_coord)\n            else:\n                raise ValueError(\"sorting parameter must be either 'x', 'y', 'z' or 'value'\")\n\n        return list_coordinates\n    \n    def change_type(self, dtype):\n        \"\"\"\n        Change data type on image.\n\n        Note: the image path is voided.\n        \"\"\"\n        change_type(self, dtype, self)\n        return self\n    \n    def fix_header_dtype(self):\n        \"\"\"\n        Change the header dtype to the match the datatype of the array.\n        \"\"\"\n        # Using bool for nibabel headers is unsupported, so use uint8 instead:\n        # `nibabel.spatialimages.HeaderDataError: data dtype \"bool\" not supported`\n        dtype_data = self.data.dtype\n        if dtype_data == bool:\n            dtype_data = np.uint8\n\n        dtype_header = self.hdr.get_data_dtype()\n        if dtype_header != dtype_data:\n            logger.warning(f\"Image header specifies datatype '{dtype_header}', but array is of type \"\n                           f\"'{dtype_data}'. Header metadata will be overwritten to use '{dtype_data}'.\")\n            self.hdr.set_data_dtype(dtype_data)\n    \n    def save(self, path=None, dtype=None, verbose=1, mutable=False):\n        \"\"\"\n        Write an image in a nifti file\n\n        :param path: Where to save the data, if None it will be taken from the\\\n                     absolutepath member.\\\n                     If path is a directory, will save to a file under this directory\\\n                     with the basename from the absolutepath member.\n\n        :param dtype: if not set, the image is saved in the same type as input data\\\n                      if 'minimize', image storage space is minimized\\\n                        (2, 'uint8', np.uint8, \"NIFTI_TYPE_UINT8\"),\\\n                        (4, 'int16', np.int16, \"NIFTI_TYPE_INT16\"),\\\n                        (8, 'int32', np.int32, \"NIFTI_TYPE_INT32\"),\\\n                        (16, 'float32', np.float32, \"NIFTI_TYPE_FLOAT32\"),\\\n                        (32, 'complex64', np.complex64, \"NIFTI_TYPE_COMPLEX64\"),\\\n                        (64, 'float64', np.float64, \"NIFTI_TYPE_FLOAT64\"),\\\n                        (256, 'int8', np.int8, \"NIFTI_TYPE_INT8\"),\\\n                        (512, 'uint16', np.uint16, \"NIFTI_TYPE_UINT16\"),\\\n                        (768, 'uint32', np.uint32, \"NIFTI_TYPE_UINT32\"),\\\n                        (1024,'int64', np.int64, \"NIFTI_TYPE_INT64\"),\\\n                        (1280, 'uint64', np.uint64, \"NIFTI_TYPE_UINT64\"),\\\n                        (1536, 'float128', _float128t, \"NIFTI_TYPE_FLOAT128\"),\\\n                        (1792, 'complex128', np.complex128, \"NIFTI_TYPE_COMPLEX128\"),\\\n                        (2048, 'complex256', _complex256t, \"NIFTI_TYPE_COMPLEX256\"),\n\n        :param mutable: whether to update members with newly created path or dtype\n        \"\"\"\n        if mutable:  # do all modifications in-place\n            # Case 1: `path` not specified\n            if path is None:\n                if self.absolutepath:  # Fallback to the original filepath\n                    path = self.absolutepath\n                else:\n                    raise ValueError(\"Don't know where to save the image (no absolutepath or path parameter)\")\n            # Case 2: `path` points to an existing directory\n            elif os.path.isdir(path):\n                if self.absolutepath:  # Use the original filename, but save to the directory specified by `path`\n                    path = os.path.join(os.path.abspath(path), os.path.basename(self.absolutepath))\n                else:\n                    raise ValueError(\"Don't know where to save the image (path parameter is dir, but absolutepath is \"\n                                     \"missing)\")\n            # Case 3: `path` points to a file (or a *nonexistent* directory) so use its value as-is\n            #    (We're okay with letting nonexistent directories slip through, because it's difficult to distinguish\n            #     between nonexistent directories and nonexistent files. Plus, `nibabel` will catch any further errors.)\n            else:\n                pass\n\n            if os.path.isfile(path) and verbose:\n                logger.warning(\"File %s already exists. Will overwrite it.\", path)\n            if os.path.isabs(path):\n                logger.debug(\"Saving image to %s orientation %s shape %s\",\n                             path, self.orientation, self.data.shape)\n            else:\n                logger.debug(\"Saving image to %s (%s) orientation %s shape %s\",\n                             path, os.path.abspath(path), self.orientation, self.data.shape)\n\n            # Now that `path` has been set and log messages have been written, we can assign it to the image itself\n            self.absolutepath = os.path.abspath(path)\n\n            if dtype is not None:\n                self.change_type(dtype)\n\n            if self.hdr is not None:\n                self.hdr.set_data_shape(self.data.shape)\n                self.fix_header_dtype()\n\n            # nb. that copy() is important because if it were a memory map, save() would corrupt it\n            dataobj = self.data.copy()\n            affine = None\n            header = self.hdr.copy() if self.hdr is not None else None\n            nib.save(nib.nifti1.Nifti1Image(dataobj, affine, header), self.absolutepath)\n            if not os.path.isfile(self.absolutepath):\n                raise RuntimeError(f\"Couldn't save image to {self.absolutepath}\")\n        else:\n            # if we're not operating in-place, then make any required modifications on a throw-away copy\n            self.copy().save(path, dtype, verbose, mutable=True)\n        return self\n\n\nclass SlicerOneAxis(object):\n    \"\"\"\n    Image slicer to use when you don't care about the 2D slice orientation,\n    and don't want to specify them.\n    The slicer will just iterate through the right axis that corresponds to\n    its specification.\n\n    Can help getting ranges and slice indices.\n\n    Copied from https://github.com/spinalcordtoolbox/spinalcordtoolbox/image.py\n    \"\"\"\n\n    def __init__(self, im, axis=\"IS\"):\n        opposite_character = {'L': 'R', 'R': 'L', 'A': 'P', 'P': 'A', 'I': 'S', 'S': 'I'}\n        axis_labels = \"LRPAIS\"\n        if len(axis) != 2:\n            raise ValueError()\n        if axis[0] not in axis_labels:\n            raise ValueError()\n        if axis[1] not in axis_labels:\n            raise ValueError()\n        if axis[0] != opposite_character[axis[1]]:\n            raise ValueError()\n\n        for idx_axis in range(2):\n            dim_nr = im.orientation.find(axis[idx_axis])\n            if dim_nr != -1:\n                break\n        if dim_nr == -1:\n            raise ValueError()\n\n        # SCT convention\n        from_dir = im.orientation[dim_nr]\n        self.direction = +1 if axis[0] == from_dir else -1\n        self.nb_slices = im.dim[dim_nr]\n        self.im = im\n        self.axis = axis\n        self._slice = lambda idx: tuple([(idx if x in axis else slice(None)) for x in im.orientation])\n\n    def __len__(self):\n        return self.nb_slices\n\n    def __getitem__(self, idx):\n        \"\"\"\n\n        :return: an image slice, at slicing index idx\n        :param idx: slicing index (according to the slicing direction)\n        \"\"\"\n        if isinstance(idx, slice):\n            raise NotImplementedError()\n\n        if idx >= self.nb_slices:\n            raise IndexError(\"I just have {} slices!\".format(self.nb_slices))\n\n        if self.direction == -1:\n            idx = self.nb_slices - 1 - idx\n\n        return self.im.data[self._slice(idx)]\n\ndef get_dimension(im_file, verbose=1):\n    \"\"\"\n    Copied from https://github.com/spinalcordtoolbox/spinalcordtoolbox/\n\n    Get dimension from Image or nibabel object. Manages 2D, 3D or 4D images.\n\n    :param: im_file: Image or nibabel object\n    :return: nx, ny, nz, nt, px, py, pz, pt\n    \"\"\"\n    if not isinstance(im_file, (nib.nifti1.Nifti1Image, Image)):\n        raise TypeError(\"The provided image file is neither a nibabel.nifti1.Nifti1Image instance nor an Image instance\")\n    # initializating ndims [nx, ny, nz, nt] and pdims [px, py, pz, pt]\n    ndims = [1, 1, 1, 1]\n    pdims = [1, 1, 1, 1]\n    data_shape = im_file.header.get_data_shape()\n    zooms = im_file.header.get_zooms()\n    for i in range(min(len(data_shape), 4)):\n        ndims[i] = data_shape[i]\n        pdims[i] = zooms[i]\n    return *ndims, *pdims\n\n\ndef change_orientation(im_src, orientation, im_dst=None, inverse=False):\n    \"\"\"\n    Copied from https://github.com/spinalcordtoolbox/spinalcordtoolbox/\n\n    :param im_src: source image\n    :param orientation: orientation string (SCT \"from\" convention)\n    :param im_dst: destination image (can be the source image for in-place\n                   operation, can be unset to generate one)\n    :param inverse: if you think backwards, use this to specify that you actually\n                    want to transform *from* the specified orientation, not *to* it.\n    :return: an image with changed orientation\n\n    .. note::\n        - the resulting image has no path member set\n        - if the source image is < 3D, it is reshaped to 3D and the destination is 3D\n    \"\"\"\n\n    if len(im_src.data.shape) < 3:\n        pass  # Will reshape to 3D\n    elif len(im_src.data.shape) == 3:\n        pass  # OK, standard 3D volume\n    elif len(im_src.data.shape) == 4:\n        pass  # OK, standard 4D volume\n    elif len(im_src.data.shape) == 5 and im_src.header.get_intent()[0] == \"vector\":\n        pass  # OK, physical displacement field\n    else:\n        raise NotImplementedError(\"Don't know how to change orientation for this image\")\n\n    im_src_orientation = im_src.orientation\n    im_dst_orientation = orientation\n    if inverse:\n        im_src_orientation, im_dst_orientation = im_dst_orientation, im_src_orientation\n\n    perm, inversion = _get_permutations(im_src_orientation, im_dst_orientation)\n\n    if im_dst is None:\n        im_dst = im_src.copy()\n        im_dst._path = None\n\n    im_src_data = im_src.data\n    if len(im_src_data.shape) < 3:\n        im_src_data = im_src_data.reshape(tuple(list(im_src_data.shape) + ([1] * (3 - len(im_src_data.shape)))))\n\n    # Update data by performing inversions and swaps\n\n    # axes inversion (flip)\n    data = im_src_data[::inversion[0], ::inversion[1], ::inversion[2]]\n\n    # axes manipulations (transpose)\n    if perm == [1, 0, 2]:\n        data = np.swapaxes(data, 0, 1)\n    elif perm == [2, 1, 0]:\n        data = np.swapaxes(data, 0, 2)\n    elif perm == [0, 2, 1]:\n        data = np.swapaxes(data, 1, 2)\n    elif perm == [2, 0, 1]:\n        data = np.swapaxes(data, 0, 2)  # transform [2, 0, 1] to [1, 0, 2]\n        data = np.swapaxes(data, 0, 1)  # transform [1, 0, 2] to [0, 1, 2]\n    elif perm == [1, 2, 0]:\n        data = np.swapaxes(data, 0, 2)  # transform [1, 2, 0] to [0, 2, 1]\n        data = np.swapaxes(data, 1, 2)  # transform [0, 2, 1] to [0, 1, 2]\n    elif perm == [0, 1, 2]:\n        # do nothing\n        pass\n    else:\n        raise NotImplementedError()\n\n    # Update header\n\n    im_src_aff = im_src.hdr.get_best_affine()\n    aff = nib.orientations.inv_ornt_aff(\n        np.array((perm, inversion)).T,\n        im_src_data.shape)\n    im_dst_aff = np.matmul(im_src_aff, aff)\n\n    im_dst.header.set_qform(im_dst_aff)\n    im_dst.header.set_sform(im_dst_aff)\n    im_dst.header.set_data_shape(data.shape)\n    im_dst.data = data\n\n    return im_dst\n\n\ndef _get_permutations(im_src_orientation, im_dst_orientation):\n    \"\"\"\n    Copied from https://github.com/spinalcordtoolbox/spinalcordtoolbox/\n\n    :param im_src_orientation str: Orientation of source image. Example: 'RPI'\n    :param im_dest_orientation str: Orientation of destination image. Example: 'SAL'\n    :return: list of axes permutations and list of inversions to achieve an orientation change\n    \"\"\"\n\n    opposite_character = {'L': 'R', 'R': 'L', 'A': 'P', 'P': 'A', 'I': 'S', 'S': 'I'}\n\n    perm = [0, 1, 2]\n    inversion = [1, 1, 1]\n    for i, character in enumerate(im_src_orientation):\n        try:\n            perm[i] = im_dst_orientation.index(character)\n        except ValueError:\n            perm[i] = im_dst_orientation.index(opposite_character[character])\n            inversion[i] = -1\n\n    return perm, inversion\n\n\ndef get_orientation(im):\n    \"\"\"\n    Copied from https://github.com/spinalcordtoolbox/spinalcordtoolbox/\n\n    :param im: an Image\n    :return: reference space string (ie. what's in Image.orientation)\n    \"\"\"\n    res = \"\".join(nib.orientations.aff2axcodes(im.hdr.get_best_affine()))\n    return orientation_string_nib2sct(res)\n\n\ndef orientation_string_nib2sct(s):\n    \"\"\"\n    Copied from https://github.com/spinalcordtoolbox/spinalcordtoolbox/\n\n    :return: SCT reference space code from nibabel one\n    \"\"\"\n    opposite_character = {'L': 'R', 'R': 'L', 'A': 'P', 'P': 'A', 'I': 'S', 'S': 'I'}\n    return \"\".join([opposite_character[x] for x in s])\n\n\ndef change_type(im_src, dtype, im_dst=None):\n    \"\"\"\n    Change the voxel type of the image\n\n    :param dtype:    if not set, the image is saved in standard type\\\n                    if 'minimize', image space is minimize\\\n                    if 'minimize_int', image space is minimize and values are approximated to integers\\\n                    (2, 'uint8', np.uint8, \"NIFTI_TYPE_UINT8\"),\\\n                    (4, 'int16', np.int16, \"NIFTI_TYPE_INT16\"),\\\n                    (8, 'int32', np.int32, \"NIFTI_TYPE_INT32\"),\\\n                    (16, 'float32', np.float32, \"NIFTI_TYPE_FLOAT32\"),\\\n                    (32, 'complex64', np.complex64, \"NIFTI_TYPE_COMPLEX64\"),\\\n                    (64, 'float64', np.float64, \"NIFTI_TYPE_FLOAT64\"),\\\n                    (256, 'int8', np.int8, \"NIFTI_TYPE_INT8\"),\\\n                    (512, 'uint16', np.uint16, \"NIFTI_TYPE_UINT16\"),\\\n                    (768, 'uint32', np.uint32, \"NIFTI_TYPE_UINT32\"),\\\n                    (1024,'int64', np.int64, \"NIFTI_TYPE_INT64\"),\\\n                    (1280, 'uint64', np.uint64, \"NIFTI_TYPE_UINT64\"),\\\n                    (1536, 'float128', _float128t, \"NIFTI_TYPE_FLOAT128\"),\\\n                    (1792, 'complex128', np.complex128, \"NIFTI_TYPE_COMPLEX128\"),\\\n                    (2048, 'complex256', _complex256t, \"NIFTI_TYPE_COMPLEX256\"),\n    :return:\n\n    Copied from https://github.com/spinalcordtoolbox/spinalcordtoolbox/\n    \"\"\"\n\n    if im_dst is None:\n        im_dst = im_src.copy()\n        im_dst._path = None\n\n    if dtype is None:\n        return im_dst\n\n    # get min/max from input image\n    min_in = np.nanmin(im_src.data)\n    max_in = np.nanmax(im_src.data)\n\n    # find optimum type for the input image\n    if dtype in ('minimize', 'minimize_int'):\n        # warning: does not take intensity resolution into account, neither complex voxels\n\n        # check if voxel values are real or integer\n        isInteger = True\n        if dtype == 'minimize':\n            for vox in im_src.data.flatten():\n                if int(vox) != vox:\n                    isInteger = False\n                    break\n\n        if isInteger:\n            if min_in >= 0:  # unsigned\n                if max_in <= np.iinfo(np.uint8).max:\n                    dtype = np.uint8\n                elif max_in <= np.iinfo(np.uint16):\n                    dtype = np.uint16\n                elif max_in <= np.iinfo(np.uint32).max:\n                    dtype = np.uint32\n                elif max_in <= np.iinfo(np.uint64).max:\n                    dtype = np.uint64\n                else:\n                    raise ValueError(\"Maximum value of the image is to big to be represented.\")\n            else:\n                if max_in <= np.iinfo(np.int8).max and min_in >= np.iinfo(np.int8).min:\n                    dtype = np.int8\n                elif max_in <= np.iinfo(np.int16).max and min_in >= np.iinfo(np.int16).min:\n                    dtype = np.int16\n                elif max_in <= np.iinfo(np.int32).max and min_in >= np.iinfo(np.int32).min:\n                    dtype = np.int32\n                elif max_in <= np.iinfo(np.int64).max and min_in >= np.iinfo(np.int64).min:\n                    dtype = np.int64\n                else:\n                    raise ValueError(\"Maximum value of the image is to big to be represented.\")\n        else:\n            # if max_in <= np.finfo(np.float16).max and min_in >= np.finfo(np.float16).min:\n            #    type = 'np.float16' # not supported by nibabel\n            if max_in <= np.finfo(np.float32).max and min_in >= np.finfo(np.float32).min:\n                dtype = np.float32\n            elif max_in <= np.finfo(np.float64).max and min_in >= np.finfo(np.float64).min:\n                dtype = np.float64\n\n        dtype = to_dtype(dtype)\n    else:\n        dtype = to_dtype(dtype)\n\n        # if output type is int, check if it needs intensity rescaling\n        if \"int\" in dtype.name:\n            # get min/max from output type\n            min_out = np.iinfo(dtype).min\n            max_out = np.iinfo(dtype).max\n            # before rescaling, check if there would be an intensity overflow\n\n            if (min_in < min_out) or (max_in > max_out):\n                # This condition is important for binary images since we do not want to scale them\n                logger.warning(f\"To avoid intensity overflow due to convertion to +{dtype.name}+, intensity will be rescaled to the maximum quantization scale\")\n                # rescale intensity\n                data_rescaled = im_src.data * (max_out - min_out) / (max_in - min_in)\n                im_dst.data = data_rescaled - (data_rescaled.min() - min_out)\n\n    # change type of data in both numpy array and nifti header\n    im_dst.data = getattr(np, dtype.name)(im_dst.data)\n    im_dst.hdr.set_data_dtype(dtype)\n    return im_dst\n\n\ndef to_dtype(dtype):\n    \"\"\"\n    Take a dtypeification and return an np.dtype\n\n    :param dtype: dtypeification (string or np.dtype or None are supported for now)\n    :return: dtype or None\n\n    Copied from https://github.com/spinalcordtoolbox/spinalcordtoolbox/\n    \"\"\"\n    # TODO add more or filter on things supported by nibabel\n\n    if dtype is None:\n        return None\n    if isinstance(dtype, type):\n        if isinstance(dtype(0).dtype, np.dtype):\n            return dtype(0).dtype\n    if isinstance(dtype, np.dtype):\n        return dtype\n    if isinstance(dtype, str):\n        return np.dtype(dtype)\n\n    raise TypeError(\"data type {}: {} not understood\".format(dtype.__class__, dtype))\n\n\ndef zeros_like(img, dtype=None):\n    \"\"\"\n\n    :param img: reference image\n    :param dtype: desired data type (optional)\n    :return: an Image with the same shape and header, filled with zeros\n\n    Similar to numpy.zeros_like(), the goal of the function is to show the developer's\n    intent and avoid doing a copy, which is slower than initialization with a constant.\n\n    Copied from https://github.com/spinalcordtoolbox/spinalcordtoolbox/image.py\n    \"\"\"\n    zimg = Image(np.zeros_like(img.data), hdr=img.hdr.copy())\n    if dtype is not None:\n        zimg.change_type(dtype)\n    return zimg\n\n\ndef empty_like(img, dtype=None):\n    \"\"\"\n    :param img: reference image\n    :param dtype: desired data type (optional)\n    :return: an Image with the same shape and header, whose data is uninitialized\n\n    Similar to numpy.empty_like(), the goal of the function is to show the developer's\n    intent and avoid touching the allocated memory, because it will be written to\n    afterwards.\n\n    Copied from https://github.com/spinalcordtoolbox/spinalcordtoolbox/image.py\n    \"\"\"\n    dst = change_type(img, dtype)\n    return dst\n\n\ndef find_zmin_zmax(im, threshold=0.1):\n    \"\"\"\n    Find the min (and max) z-slice index below which (and above which) slices only have voxels below a given threshold.\n\n    :param im: Image object\n    :param threshold: threshold to apply before looking for zmin/zmax, typically corresponding to noise level.\n    :return: [zmin, zmax]\n\n    Copied from https://github.com/spinalcordtoolbox/spinalcordtoolbox/image.py\n    \"\"\"\n    slicer = SlicerOneAxis(im, axis=\"IS\")\n\n    # Make sure image is not empty\n    if not np.any(slicer):\n        logger.error('Input image is empty')\n\n    # Iterate from bottom to top until we find data\n    for zmin in range(0, len(slicer)):\n        if np.any(slicer[zmin] > threshold):\n            break\n\n    # Conversely from top to bottom\n    for zmax in range(len(slicer) - 1, zmin, -1):\n        if np.any(slicer[zmax] > threshold):\n            break\n\n    return zmin, zmax","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T17:29:44.430835Z","iopub.execute_input":"2025-04-28T17:29:44.431303Z","iopub.status.idle":"2025-04-28T17:29:44.785728Z","shell.execute_reply.started":"2025-04-28T17:29:44.431265Z","shell.execute_reply":"2025-04-28T17:29:44.784843Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# TotalSpineSeg","metadata":{}},{"cell_type":"code","source":"# Installation of  libraries\n!pip install --no-deps /kaggle/input/kaggle-wheels/gryds-0.0.9-py3-none-any.whl\n!pip install --no-deps /kaggle/input/kaggle-wheels/monai-1.3.2-py3-none-any.whl\n!pip install --no-deps /kaggle/input/kaggle-wheels/torchio-0.19.9-py2.py3-none-any.whl\n!pip install --no-deps /kaggle/input/cc3d-wheel/connected_components_3d-3.18.0-cp310-cp310-manylinux_2_17_x86_64.manylinux2014_x86_64.whl\n!pip install --no-deps /kaggle/input/kaggle-wheels/totalspineseg-20241005-py3-none-any.whl\n#!pip install --no-deps /kaggle/input/kaggle-wheels/gryds-0.0.9-py3-none-any.whl\n#!pip install --no-deps /kaggle/input/kaggle-wheels/monai-1.3.2-py3-none-any.whl\n#!pip install --no-deps /kaggle/input/kaggle-wheels/torchio-0.19.9-py2.py3-none-any.whl\n!pip install --no-deps /kaggle/input/kaggle-wheels/nnunetv2-2.5.1-py3-none-any.whl\n!pip install --no-deps /kaggle/input/acvl-utils/acvl_utils-0.2-py3-none-any.whl\n!pip install --no-deps /kaggle/input/batchgenerators/batchgenerators-0.25-py3-none-any.whl\n!pip install --no-deps /kaggle/input/nnunet-wheels/batchgeneratorsv2-0.2.1-py3-none-any.whl\n!pip install --no-deps /kaggle/input/nnunet-wheels/fft_conv_pytorch-1.2.0-py3-none-any.whl\n!pip install --no-deps /kaggle/input/dynamic-net/dynamic_network_architectures-0.3.1-py3-none-any.whl\n#!pip install --no-deps /kaggle/input/cc3d-wheel/connected_components_3d-3.18.0-cp310-cp310-manylinux_2_17_x86_64.manylinux2014_x86_64.whl\n!pip install --no-deps /kaggle/input/pillow-wheel/pillow-10.4.0-cp310-cp310-manylinux_2_28_x86_64.whl","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T17:29:44.786919Z","iopub.execute_input":"2025-04-28T17:29:44.787588Z","iopub.status.idle":"2025-04-28T17:36:34.428962Z","shell.execute_reply.started":"2025-04-28T17:29:44.787550Z","shell.execute_reply":"2025-04-28T17:36:34.427853Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"! totalspineseg -h","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T17:36:34.431809Z","iopub.execute_input":"2025-04-28T17:36:34.432781Z","iopub.status.idle":"2025-04-28T17:36:40.255538Z","shell.execute_reply.started":"2025-04-28T17:36:34.432737Z","shell.execute_reply":"2025-04-28T17:36:40.254688Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"os.makedirs(\"TotalSpineSeg\", exist_ok = True)\nos.system(\"cp -r /kaggle/input/m/simonqueric/totalspineseg/other/default/6/totalspineseg/totalspineseg TotalSpineSeg\")\nos.system(\"python3 -m venv /kaggle/TotalSpineSeg/venv\")\nos.system('bash -c \"source /kaggle/TotalSpineSeg/venv/bin/activate\"')\nos.makedirs('/kaggle/working/TotalSpineSeg/tss_input', exist_ok=True)\nos.system('export TOTALSPINESEG=\"$(realpath totalspineseg)\"')\nos.system('export TOTALSPINESEG_DATA=\"$(realpath data)\"')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T17:36:40.257010Z","iopub.execute_input":"2025-04-28T17:36:40.257626Z","iopub.status.idle":"2025-04-28T17:36:43.941880Z","shell.execute_reply.started":"2025-04-28T17:36:40.257582Z","shell.execute_reply":"2025-04-28T17:36:43.940979Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_subjects(source_dir, batch_size=50):\n    \"\"\"Get a list of batches of subjects filtered by the filter_func.\"\"\"\n    all_subjects = [sub for sub in os.listdir(source_dir)]\n    return all_subjects\n\n\ndef run_totalspineseg(source_dir):\n    \"\"\"Applies TotalSpineSeg to every scan in the source_dir and saves the segmentations.\"\"\"\n    # Define temporary directories\n    tss_temp_dir = \"TotalSpineSeg/tss_input\"\n    output_temp = \"temp_output_data\"\n    failed_subjects = []\n\n    # Get all batches of subjects\n    subjects = get_subjects(source_dir, batch_size=50)\n   \n    # Process each batch\n    \n    os.makedirs(tss_temp_dir, exist_ok=True)\n    os.makedirs(output_temp, exist_ok=True)\n\n    for subdir in subjects:\n        try:\n            anat_path = os.path.join(source_dir, subdir, 'anat')\n            if os.path.exists(anat_path):\n                for file in os.listdir(anat_path):\n                    file_path = os.path.join(anat_path, file)\n                    if os.path.isfile(file_path) and 'ax' not in file_path and 'total_seg' not in file_path and 'T2' not in file_path:\n                        shutil.copy(file_path, tss_temp_dir)\n                        print('File copied successfully.')\n        except Exception as e:\n            print(f\"Failed processing subject {subdir}: {e}\")\n            failed_subjects.append(subdir)\n\n    # Run TotalSpineSeg segmentation\n    os.system(f\"totalspineseg --data-dir /kaggle/working/{tss_temp_dir} /kaggle/working/{tss_temp_dir} /kaggle/working/{output_temp} --step1\")\n    \n\n    # Move segmentations back into original data structure\n    segmentations_into_anat(output_temp, source_dir)\n\n    # Clean up temporary directories\n    os.system(\"rm /kaggle/working/TotalSpineSeg/tss_input/*.nii.gz\")\n    shutil.rmtree(output_temp)\n\n\ndef segmentations_into_anat(output_folder, nii_folder):\n    \"\"\"Send the segmentations into the folder with the nii volumes.\"\"\"\n    seg_folder = os.path.join(output_folder, \"step1_output\")\n    segmentations = os.listdir(seg_folder)\n\n    for segmentation in segmentations:\n        id_patient = segmentation.split('_')[0]\n        patient_folder = os.path.join(nii_folder, id_patient, 'anat')\n\n        if os.path.exists(patient_folder):\n            source_path = os.path.join(seg_folder, segmentation)\n            modified_segmentation = segmentation.replace('.nii.gz', '_total_seg.nii.gz')\n            destination_path = os.path.join(patient_folder, modified_segmentation)\n            shutil.copy(source_path, destination_path)\n\n\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T17:36:43.943155Z","iopub.execute_input":"2025-04-28T17:36:43.943436Z","iopub.status.idle":"2025-04-28T17:36:43.952773Z","shell.execute_reply.started":"2025-04-28T17:36:43.943409Z","shell.execute_reply":"2025-04-28T17:36:43.951851Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Library imports         ","metadata":{}},{"cell_type":"code","source":"import os\nimport sys\nimport shutil\nimport csv\nimport subprocess\nimport glob\nfrom pathlib import Path\n\nimport numpy as np\nimport pandas as pd\nimport torch\nfrom torch.utils.data import DataLoader, Dataset, Subset, ConcatDataset\nimport torch.optim as optim\nfrom torch.nn import CrossEntropyLoss\n\nfrom tqdm import tqdm\nimport monai\nfrom monai.transforms import (\n    Compose, LoadImaged, EnsureChannelFirstd, ScaleIntensityd, ConcatItemsd,\n    ToTensord, RandRotate90d, RandFlipd, SpatialPadd, CenterSpatialCropd,\n    NormalizeIntensityd, RandScaleIntensityd, RandShiftIntensityd, RandRotated,\n    Spacingd, RandSpatialCropd, RandBiasFieldd, Flipd, SpatialCropd, Transform, \n    Compose, LoadImaged, EnsureChannelFirstd, ScaleIntensityd, ConcatItemsd,\n    ToTensord, SpatialPadd, CenterSpatialCropd, NormalizeIntensityd,\n    RandRotated, RandSpatialCropd, RandBiasFieldd, Lambdad, Transform,\n    RandGaussianNoised, RandAffined, RandZoomd, Rand3DElasticd, Spacingd\n)\nfrom monai.networks.nets import DenseNet201, ResNet\nfrom monai.data import Dataset\n\nimport seaborn as sns\nimport matplotlib.pyplot as plt\nimport nibabel as nib\nimport torchio as tio\nfrom sklearn.metrics import confusion_matrix\nimport argparse\nfrom scipy.ndimage import center_of_mass\nfrom skimage.measure import regionprops","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T17:36:43.954118Z","iopub.execute_input":"2025-04-28T17:36:43.954805Z","iopub.status.idle":"2025-04-28T17:37:18.154203Z","shell.execute_reply.started":"2025-04-28T17:36:43.954745Z","shell.execute_reply":"2025-04-28T17:37:18.153293Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# BIDS","metadata":{}},{"cell_type":"code","source":"import shutil \nimport os \n\nshutil.copytree('/kaggle/input/dcm2nii', '/kaggle/working/dcm2nii', dirs_exist_ok=True)\n\n\n! chmod +x /kaggle/working/dcm2nii/dcm2niix/build/bin/dcm2niix\nos.environ['PATH'] += ':/kaggle/working/dcm2nii/dcm2niix/build/bin/'","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T17:37:18.155365Z","iopub.execute_input":"2025-04-28T17:37:18.155968Z","iopub.status.idle":"2025-04-28T17:37:24.423286Z","shell.execute_reply.started":"2025-04-28T17:37:18.155938Z","shell.execute_reply":"2025-04-28T17:37:24.422141Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\"\"\"def resample_nifti(image, target_spacing=(2.6, 0.4, 0.4), mode='linear'):\n    \n    Resample une image NIfTI en 3D à un nouvel espacement spécifié.\n    \n    Args:\n        image (nibabel.Nifti1Image): Image NIfTI à resampler.\n        target_spacing (tuple): Nouvel espacement des voxels (en mm).\n    \n    Returns:\n        None\n    \n    # Appliquer le resampling\n    resampler = tio.Resample(target_spacing, image_interpolation=mode)\n    resampled_image = resampler(image)\n\n    return resampled_image\"\"\"\n\n# use a subprocess to convert the dicom images to nifti format, requires the output path\ndef convert_dicom_to_nifti(subject_id, series_uid, input_path, output_path):\n    \"\"\"\n    Convert DICOM images to NIfTI format using dcm2niix.\n    \n    Parameters:\n    subject_id (str): The subject identifier.\n    series_uid (str): The series instance UID.\n    input_path (str): Path to the DICOM images directory.\n    output_path (str): Path to the output directory for NIfTI files.\n    \"\"\"\n    input_file = os.path.join(input_path, subject_id, series_uid)\n    output_file = os.path.join(output_path, f\"{subject_id}-{series_uid}\")\n    if not os.path.exists(output_file):\n        os.makedirs(output_file)\n\n    dcm2niix_command = f\"dcm2niix -z y -m 2 -o {output_file} {input_file}\"\n    \n    try:\n        subprocess.run(dcm2niix_command, shell=True, check=True)\n    except subprocess.CalledProcessError as e:\n        print(f\"Error converting {input_file}: {e}\")\n        return None\n\n# do not apply this function to the axial acquisitions, as it will merge different acquisitions with different orientations\ndef merge_nifti_volumes(output_path, subject_id, series_uid):\n    \"\"\"\n    Merge NIfTI volumes in the Z direction and save the merged volume if more than one volume exists.\n    Otherwise, save the single volume directly.\n    Rename the merged volume to the specified format.\n    \n    Parameters:\n    output_filename (str): Output filename for the merged NIfTI volume.\n    output_path (str): Path to the output directory for merged NIfTI volume.\n    subject_id (str): The subject identifier.\n    series_uid (str): The series instance UID.\n    \"\"\"\n    output_path_for_merge = os.path.join(output_path, f\"{subject_id}-{series_uid}\")\n\n    filenames = glob.glob(os.path.join(output_path_for_merge, '*.nii.gz'))\n    filenames.sort()\n    new_paths = []\n    if len(filenames) > 1:\n        for filename in filenames: \n            merged_filename = f\"sub-{subject_id}_run-{series_uid}_{filename[-14:]}\"\n            merged_path = os.path.join(output_path, merged_filename)  # Changed to output_folder\n            os.rename(filename, merged_path)\n            new_paths.append(merged_path)\n    elif len(filenames) == 1:\n        merged_filename = f\"sub-{subject_id}_run-{series_uid}.nii.gz\"\n        merged_path = os.path.join(output_path, merged_filename)  # Changed to output_folder\n        os.rename(filenames[0], merged_path)\n        new_paths.append(merged_path)\n    return new_paths\n\n\n# reorient the image to a common orientation \"LPI\"\ndef reorient(image):\n\n    # Get image dtype from the image data (preferred over header dtype to avoid data loss)\n    image_data_dtype = getattr(np, np.asanyarray(image.dataobj).dtype.name)\n\n    # Rescale the image to the output dtype range if necessary\n    # Modified from https://github.com/spinalcordtoolbox/spinalcordtoolbox/blob/6.3/spinalcordtoolbox/image.py#L1217\n    if \"int\" in np.dtype(image_data_dtype).name:\n        image_data = np.asanyarray(image.dataobj).astype(np.float64)\n        image_min, image_max = image_data.min(), image_data.max()\n        dtype_min, dtype_max = np.iinfo(image_data_dtype).min, np.iinfo(image_data_dtype).max\n        if (image_min < dtype_min) or (dtype_max < image_max):\n            data_rescaled = image_data * (dtype_max - dtype_min) / (image_max - image_min)\n            image_data = data_rescaled - (data_rescaled.min() - dtype_min)\n            image = nib.Nifti1Image(image_data.astype(image_data_dtype), image.affine, image.header)\n\n    # Transform the image to the closest canonical orientation\n    output_image = nib.as_closest_canonical(image)\n\n    # Ensure correct image dtype, affine and header\n    output_image = nib.Nifti1Image(\n        np.asanyarray(output_image.dataobj).astype(image_data_dtype),\n        output_image.affine, output_image.header\n    )\n    output_image.set_data_dtype(image_data_dtype)\n    output_image.set_qform(output_image.affine)\n    output_image.set_sform(output_image.affine)\n\n    return output_image\n\n\n\n\n\ndef process_subject(subject_id, input_path, output_path, train, meta_obj):\n    \"\"\"\n    Process DICOM to NIfTI conversion, merge volumes, and save with corrected orientation if applicable.\n    \n    Parameters:\n    subject_id (str): The subject identifier.\n    input_path (str): Path to the DICOM images directory.\n    output_path (str): Path to the output directory for NIfTI files.\n    train (DataFrame): DataFrame containing series information.\n    meta_obj (dict): Metadata object with series information.\n    \"\"\"\n    if subject_id not in train['study_id'].astype(str).values:\n        return\n\n    filtered_series = train[train['study_id'] == int(subject_id)].iloc[0]\n    ptobj = meta_obj[str(filtered_series['study_id'])]\n\n    if ptobj is None:\n        return\n\n    # create output directories if not existing\n    os.makedirs(os.path.join(output_path, f'sub-{subject_id}'), exist_ok=True)\n    os.makedirs(os.path.join(output_path, f'sub-{subject_id}', 'anat'), exist_ok=True)\n\n    # process through each acquisition of the subject    \n    for idx, series_uid in enumerate(ptobj['SeriesInstanceUIDs']):\n        description = ptobj['SeriesDescriptions'][idx]\n\n        convert_dicom_to_nifti(subject_id, series_uid, input_path, output_path)\n        new_paths = merge_nifti_volumes(output_path, subject_id, series_uid)\n        if 'Axial' in description and 'T2' in description:\n                modality = 'T2w'\n                acq = 'ax'\n        elif 'Sagittal' in description and 'T1' in description:\n            modality = 'T1w'\n            acq = 'sag'\n        elif 'Sagittal' in description and 'T2' in description:\n            modality = 'T2w'\n            acq = 'sag'\n        else:\n            continue\n\n        corrected_nifti_path = os.path.join(output_path, f\"sub-{subject_id}/anat/sub-{subject_id}_acq-{acq}_rec{series_uid}_{modality}\")\n        if len(new_paths) > 1 : \n            for merged_nifti_path in new_paths : \n                anat_img = nib.load(merged_nifti_path)\n                anat_data = anat_img.get_fdata()\n                anat_affine = anat_img.affine\n                anat_header = anat_img.header\n\n                new_affine = np.copy(anat_affine)\n                anat_header.set_qform(new_affine, code=1)\n                anat_header.set_sform(new_affine, code=1)\n\n                base, ext = os.path.splitext(merged_nifti_path)\n\n                new_path = corrected_nifti_path + base[-11:] + ext\n\n                # reorient the image\n                image = nib.Nifti1Image(anat_data, new_affine, header=anat_header)\n                \n                oriented_image = reorient(image)\n                # then apply the resampling to the median values resolution for axial T2w images\n                if acq == 'ax': \n                    \"\"\"final_image = resample_nifti(oriented_image, target_spacing=(0.4, 0.4, 4.4), mode='linear')  \n                    \n                    nib.save(final_image, new_path)\"\"\"\n                    nib.save(oriented_image, new_path)\n                else:\n                    \"\"\"final_image = resample_nifti(oriented_image, target_spacing=(4.0, 0.4, 0.4), mode='linear')  \n\n                    nib.save(final_image, new_path)\"\"\"\n                    nib.save(oriented_image, new_path)\n\n        else : \n            for merged_nifti_path in new_paths : \n                anat_img = nib.load(merged_nifti_path)\n                anat_data = anat_img.get_fdata()\n                anat_affine = anat_img.affine\n                anat_header = anat_img.header\n\n                new_affine = np.copy(anat_affine)\n                anat_header.set_qform(new_affine, code=1)\n                anat_header.set_sform(new_affine, code=1)\n                new_path = corrected_nifti_path + '.nii.gz'\n                # reorient the image\n                image = nib.Nifti1Image(anat_data, new_affine, header=anat_header)\n                oriented_image = reorient(image)\n\n                # then apply the resampling to the median values resolution for axial T2w images\n                if acq == 'ax': \n                    \n                                \n                    \"\"\"final_image = resample_nifti(oriented_image, target_spacing=(0.4, 0.4, 4.4), mode='linear') \n                    nib.save(final_image, new_path)\"\"\"\n                    nib.save(oriented_image, new_path)\n                else:\n\n                    \"\"\"final_image = resample_nifti(oriented_image, target_spacing=(4.0, 0.4, 0.4), mode='linear')  \n                    nib.save(final_image, new_path)\"\"\"\n                    nib.save(oriented_image, new_path)\n    \n\n# Main function to run the processing\ndef BIDSification(input_folder, output_folder, csv_description):\n    \n    os.makedirs(output_folder, exist_ok=True)\n\n    ### Create the dictionary based on the CSV file ###\n    df_meta_f = pd.read_csv(csv_description, sep=',')\n    subject_ids = np.unique(df_meta_f[\"study_id\"].values)\n\n    # List out all of the Studies we have on patients.\n    part_1 = os.listdir(input_folder)\n    part_1 = list(filter(lambda x: x.find('.DS') == -1, part_1))\n\n    p1 = [(x, f\"{input_folder}/{x}\") for x in part_1]\n    meta_obj = { p[0]: { 'folder_path': p[1], \n                        'SeriesInstanceUIDs': [] \n                    } \n                for p in p1 }\n\n    for m in meta_obj:\n        meta_obj[m]['SeriesInstanceUIDs'] = list( \n            filter(lambda x: x.find('.DS') == -1, \n                os.listdir(meta_obj[m]['folder_path'])\n                )\n        )\n    # Grabs the corresponding series descriptions\n    for k in tqdm(meta_obj):\n        for s in meta_obj[k]['SeriesInstanceUIDs']:\n            if 'SeriesDescriptions' not in meta_obj[k]:\n                meta_obj[k]['SeriesDescriptions'] = []\n            try:\n                meta_obj[k]['SeriesDescriptions'].append(\n                    df_meta_f[(df_meta_f['study_id'] == int(k)) & \n                    (df_meta_f['series_id'] == int(s))]['series_description'].iloc[0])\n            except:\n                None\n\n    # Process subjects and set up directories\n    \n    for subject_id in tqdm(subject_ids):  # Adjust range as needed: 1975 subjects\n        try: \n            subject_id = str(subject_id)\n            \n            # Create specific directories\n            os.makedirs(os.path.join(output_folder, f'sub-{subject_id}'), exist_ok=True)\n            os.makedirs(os.path.join(output_folder, f'sub-{subject_id}', 'anat'), exist_ok=True)\n\n            # Process subject and set up directories\n            process_subject(subject_id, input_folder, output_folder, df_meta_f, meta_obj)\n        except: \n            print(f'failed preprocessing for {subject_id}')\n \n\n    for item in os.listdir(output_folder):\n        item_path = os.path.join(output_folder, item)\n        \n        # Check if the item is a directory and starts with \"sub\"\n        if os.path.isdir(item_path) and item.startswith(\"sub\"):\n            continue  # Skip deletion for folders starting with \"sub\"\n        \n        # Delete the item (file or directory)\n        if os.path.isfile(item_path):\n            os.remove(item_path)\n            print(f\"Deleted file: {item_path}\")\n        elif os.path.isdir(item_path):\n            shutil.rmtree(item_path)\n            print(f\"Deleted folder: {item_path}\")\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T17:37:24.425304Z","iopub.execute_input":"2025-04-28T17:37:24.425620Z","iopub.status.idle":"2025-04-28T17:37:24.451938Z","shell.execute_reply.started":"2025-04-28T17:37:24.425590Z","shell.execute_reply":"2025-04-28T17:37:24.451098Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Patch extraction ","metadata":{}},{"cell_type":"code","source":"def get_shifted_point_along_disk(disk_mask):\n    \"\"\"\n    Calcule un point décalé selon l'axe du disque en LPI.\n    \n    Args:\n        disk_mask: Masque binaire 3D du disque\n    \n    Returns:\n        point: numpy array des coordonnées (x,y,z) du point décalé\n        disk_radius: rayon calculé du disque selon son axe\n        direction_vector: vecteur normalisé indiquant la direction du disque\n    \"\"\"\n    # Trouver le centre du disque\n    centroid = center_of_mass(disk_mask)\n    # Trouver la slice sagittale contenant le centre du disque\n    sagittal_slice_idx = int(centroid[0])\n    sagittal_slice = disk_mask[sagittal_slice_idx, :, :]\n    \n    # Calculer l'orientation sur la slice 2D\n    props = regionprops(sagittal_slice.astype(int))[0]\n    orientation = props.orientation  # en radians\n\n    direction_vector = np.array([\n        0,  # x reste inchangé\n        -np.cos(orientation),   # y\n        -np.sin(orientation)    # z\n    ])\n    \n    # Normaliser le vecteur\n    direction_vector = direction_vector / np.linalg.norm(direction_vector)\n    \n    # Calculer le rayon du disque (projection sur le vecteur)\n    mask_points = np.array(np.where(disk_mask)).T\n    centered_points = mask_points - centroid\n    projections = np.abs(centered_points @ direction_vector)\n    disk_radius = np.max(projections)\n    \n    # Calculer le point décalé\n    shifted_point = centroid + direction_vector * disk_radius\n\n    return shifted_point","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T17:37:24.454958Z","iopub.execute_input":"2025-04-28T17:37:24.455230Z","iopub.status.idle":"2025-04-28T17:37:24.467904Z","shell.execute_reply.started":"2025-04-28T17:37:24.455197Z","shell.execute_reply":"2025-04-28T17:37:24.467039Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"## STEP 3 ##\n\n# this is last part of the preprocessing pipeline\n# its goal is to extract the patches from the nii volumes based on the segmentation \n\ndef process_directory_other(main_dir):\n    '''\n    Transform the segmentations in main_dir folder to the image space to have the same origin, spacing, direction and shape as the image.\n\n    Parameters\n    main_dir: where to fetch the segmentations\n    \n    '''\n\n    main_dir_path = Path(main_dir)\n    \n    # Iterate through each subdirectory (for each patient)\n    for dirpath, dirnames, filenames in os.walk(main_dir_path):\n        # Vérifier que nous sommes dans un dossier patient et qu'il y a un sous-dossier anat\n       \n        if \"anat\" in dirnames:\n            \n            anat_path = os.path.join(dirpath, \"anat\")\n            \n            # Obtenir la liste des fichiers dans le sous-dossier anat\n            anat_filenames = os.listdir(anat_path)\n            \n        \n            # Find sagittal T2w image\n            sag_files = [f for f in anat_filenames if \"acq-sag\" in f and \"T1w_total_seg\" in f]\n            if len(sag_files) == 0:\n                \n                continue\n            \n            sag_file = os.path.join(anat_path, sag_files[0])\n            \n            # Find and process all axial images\n            for ax_file in anat_filenames:\n                try: \n                    if \"acq-ax\" in ax_file and not \"seg\" in ax_file and not \"patch\" in ax_file:\n                        ax_file_path = os.path.join(anat_path, ax_file)\n                        \n                        output_file_path = ax_file_path.replace(\".nii.gz\", \"_total_seg.nii.gz\")\n                        \n                        \n                        # Call the transformation function\n                        _transform_seg2image(ax_file_path, sag_file, output_file_path)\n                except: \n                    print (ax_file)\n\ndef _transform_seg2image(\n        image_path,\n        seg_path,\n        output_seg_path,\n        override=False,\n    ):\n    '''\n    Wrapper function to handle IO.\n    '''\n    image_path = Path(image_path)\n    seg_path = Path(seg_path)\n    output_seg_path = Path(output_seg_path)\n\n    # If the output image already exists and we are not overriding it, return\n    if not override and output_seg_path.exists():\n        return\n\n    # Check if the segmentation file exists\n    if not seg_path.is_file():\n        output_seg_path.is_file() and output_seg_path.unlink()\n        return\n\n    image = nib.load(image_path)\n    seg = nib.load(seg_path)\n\n    output_seg = transform_seg2image(image, seg)\n\n    # Ensure correct segmentation dtype, affine and header\n    output_seg = nib.Nifti1Image(\n        np.asanyarray(output_seg.dataobj).round().astype(np.uint8),\n        output_seg.affine, output_seg.header\n    )\n    output_seg.set_data_dtype(np.uint8)\n    output_seg.set_qform(output_seg.affine)\n    output_seg.set_sform(output_seg.affine)\n\n    # Make sure output directory exists and save the segmentation\n    output_seg_path.parent.mkdir(parents=True, exist_ok=True)\n    nib.save(output_seg, output_seg_path)\n\n\n\n\ndef transform_seg2image(\n        image,\n        seg,\n    ):\n    '''\n    Transform the segmentation to the image space to have the same origin, spacing, direction and shape as the image.\n\n    Parameters\n    ----------\n    image : nibabel.Nifti1Image\n        Image.\n    seg : nibabel.Nifti1Image\n        Segmentation.\n\n    Returns\n    -------\n    nibabel.Nifti1Image\n        Output segmentation.\n    '''\n    image_data = np.asanyarray(image.dataobj).astype(np.float64)\n    seg_data = np.asanyarray(seg.dataobj).round().astype(np.uint8)\n\n    # Make TorchIO images\n    tio_img=tio.ScalarImage(tensor=image_data[None, ...], affine=image.affine)\n    tio_seg=tio.LabelMap(tensor=seg_data[None, ...], affine=seg.affine)\n\n    # Resample the segmentation to the image space\n    tio_output_seg = tio.Resample(tio_img)(tio_seg)\n    output_seg_data = tio_output_seg.data.numpy()[0, ...].astype(np.uint8)\n\n    output_seg = nib.Nifti1Image(output_seg_data, image.affine, seg.header)\n\n    return output_seg\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T17:37:24.469091Z","iopub.execute_input":"2025-04-28T17:37:24.469340Z","iopub.status.idle":"2025-04-28T17:37:24.484102Z","shell.execute_reply.started":"2025-04-28T17:37:24.469317Z","shell.execute_reply":"2025-04-28T17:37:24.483244Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# function for sagittal patches\ndef patch_extraction_foraminal(vol, mask, affine):\n    \"\"\"\n    Extract two 3D patches from an MRI volume centered around mask's centroid\n    \n    Parameters:\n    - vol: 3D numpy array representing the volume\n    - mask: 3D segmentation mask \n    - affine: Affine matrix from the NIfTI file\n    \n    Returns:\n    - patch1, patch2: Two 3D numpy array patches\n    \"\"\"\n    i = 0\n\n    D, H, W = vol.shape\n    # mask = torch.Tensor(mask)\n    # nonzero_indices = torch.nonzero(mask)\n    \n    # Calculate centroid of the mask and shift it along the disk axis\n    centroid = get_shifted_point_along_disk(mask).astype(int)\n\n    # Get voxel sizes from the affine matrix\n    voxel_sizes = np.abs(np.diag(affine)[:3])\n    \n    patch_size_mm = {\n        'd': 50,  # depth\n        'h': 50,  # height\n        'w': 50   # width\n    }\n\n    # Patch sizes (in voxels)\n    patch_sizes_voxels = {\n        'd': (patch_size_mm['d'] / voxel_sizes[0]).astype(int),\n        'h': (patch_size_mm['h'] / voxel_sizes[1]).astype(int),\n        'w': (patch_size_mm['w'] / voxel_sizes[2]).astype(int)\n    }\n\n    \n    # Extract patches centered on centroid with posterior displacement\n    patch1 = vol[\n        max(0, centroid[0] + 1):min(D, centroid[0] + patch_sizes_voxels['d']//2 +1),\n        max(0, centroid[1] - patch_sizes_voxels['h']//2):min(H, centroid[1] + patch_sizes_voxels['h']//2),\n        max(0, centroid[2] - patch_sizes_voxels['w']//2):min(W, centroid[2] + patch_sizes_voxels['w']//2)\n    ]\n    \n    patch2 = vol[\n        max(0, centroid[0] -1 -patch_sizes_voxels['d']//2):min(D, centroid[0]-1),\n        max(0, centroid[1] - patch_sizes_voxels['h']//2):min(H, centroid[1] + patch_sizes_voxels['h']//2),\n        max(0, centroid[2] - patch_sizes_voxels['w']//2):min(W, centroid[2] + patch_sizes_voxels['w']//2)\n    ]\n    \n    return patch1, patch2\n\n# uses lists of sagittal images and segmentations to extract patches for each disc\ndef extract_and_save_sagittal_patches(sagittal_images, sagittal_segmentations, nii_folder, output_folder):\n    # Match each axial image to its corresponding sagittal segmentation\n    for img_name, seg_sag_name in zip(sagittal_images, sagittal_segmentations):\n        if \"patch\" not in img_name:\n            img_path = os.path.join(nii_folder, img_name)\n            seg_sag_path = os.path.join(nii_folder, seg_sag_name)\n            affine_ex = nib.load(img_path).affine\n            \n            # Load the volumetric image and sagittal segmentation\n            vol = nib.load(img_path).get_fdata()\n            seg_sag = nib.load(seg_sag_path).get_fdata()\n\n            # Détection des disques dans la segmentation sagittale\n            #The values to check are based on the classes in totalspineseg \n            disc_l5 = np.isin(seg_sag, [100]).astype(int)\n            disc_l4 = np.isin(seg_sag, [95]).astype(int)\n            disc_l3 = np.isin(seg_sag, [94]).astype(int)\n            disc_l2 = np.isin(seg_sag, [93]).astype(int)\n            disc_l1 = np.isin(seg_sag, [92]).astype(int)\n\n            \n            discs_dict = {\n                \"L1_L2\": disc_l1,\n                \"L2_L3\": disc_l2,\n                \"L3_L4\": disc_l3,\n                \"L4_L5\": disc_l4,\n                \"L5_S1\": disc_l5\n            }    \n\n            # Extract and save patches for each disc\n            for disc_name, disc_mask in discs_dict.items():\n                if np.any(disc_mask):  # If the disc is found in the segmentation\n                    # Extract the patch using the segmentation mask\n                    \n                    patch_img_left, patch_img_right = patch_extraction_foraminal(vol, disc_mask, affine_ex)\n                    \n                    if patch_img_left is not None or patch_img_right is not None:  # Proceed only if patch extraction was successful\n\n                        # Construct the filename and file path\n                        patch_img_filename_left = f\"{img_name[:-7]}_{disc_name}_foramen_left_patch.nii.gz\"\n                        patch_img_filepath_left = os.path.join(output_folder, patch_img_filename_left)\n\n                        patch_img_filename_right = f\"{img_name[:-7]}_{disc_name}_foramen_right_patch.nii.gz\"\n                        patch_img_filepath_right = os.path.join(output_folder, patch_img_filename_right)\n\n                        # Use the affine from the original volume to create the patch NIfTI image\n                        original_affine = nib.load(img_path).affine\n                        original_header = nib.load(img_path).header.copy()\n                        patch_nifti_img_left = nib.Nifti1Image(patch_img_left, affine=original_affine)\n                        patch_nifti_img_right = nib.Nifti1Image(patch_img_right, affine=original_affine)\n\n\n                        q_code = int(original_header['qform_code'])\n                        s_code = int(original_header['sform_code'])\n\n                        patch_nifti_img_left.header.set_qform(original_affine, code=q_code)\n                        patch_nifti_img_left.header.set_sform(original_affine, code=s_code)\n                        patch_nifti_img_right.header.set_qform(original_affine, code=q_code)\n                        patch_nifti_img_right.header.set_sform(original_affine, code=s_code)\n\n                        # Save the patch to the specified location\n                        nib.save(patch_nifti_img_left, patch_img_filepath_left)\n                        nib.save(patch_nifti_img_right, patch_img_filepath_right)\n\n# uses lists of axial images and segmentations to extract patches for each disc\ndef extract_and_save_axial_patches(axial_images, axial_segmentations, nii_folder, output_folder):\n    # Match each axial image to its corresponding sagittal segmentation\n    for img_name, seg_sag_name in zip(axial_images, axial_segmentations):\n        if \"patch\" not in img_name:\n            img_path = os.path.join(nii_folder, img_name)\n            seg_sag_path = os.path.join(nii_folder, seg_sag_name)\n            \n            # Load the volumetric image and sagittal segmentation\n            vol = nib.load(img_path).get_fdata()\n            seg_sag = nib.load(seg_sag_path).get_fdata()\n            affine = nib.load(img_path).affine\n\n            # Détection des disques dans la segmentation sagittale\n            #The values to check are based on the classes in totalspineseg \n            disc_l5 = np.isin(seg_sag, [100]).astype(int)\n            disc_l4 = np.isin(seg_sag, [95]).astype(int)\n            disc_l3 = np.isin(seg_sag, [94]).astype(int)\n            disc_l2 = np.isin(seg_sag, [93]).astype(int)\n            disc_l1 = np.isin(seg_sag, [92]).astype(int)\n\n            \n            discs_dict = {\n                \"L1_L2\": disc_l1,\n                \"L2_L3\": disc_l2,\n                \"L3_L4\": disc_l3,\n                \"L4_L5\": disc_l4,\n                \"L5_S1\": disc_l5\n            }\n\n            # Extract and save patches for each disc\n            for disc_name, disc_mask in discs_dict.items():\n                if np.any(disc_mask):  # If the disc is found in the segmentation\n                    # Extract the patch using the segmentation mask\n                    \n                    patch_img = patch_extraction_volume(vol, disc_mask, affine)\n                    \n                    if patch_img is not None:  # Proceed only if patch extraction was successful\n\n                        # Construct the filename and file path\n                        patch_img_filename = f\"{img_name[:-7]}_{disc_name}_patch.nii.gz\"\n                        patch_img_filepath = os.path.join(output_folder, patch_img_filename)\n                        \n                        # Use the affine from the original volume to create the patch NIfTI image\n                        original_affine = nib.load(img_path).affine\n                        patch_nifti_img = nib.Nifti1Image(patch_img, affine=original_affine)\n\n                        original_header = nib.load(img_path).header.copy()\n\n                        q_code = int(original_header['qform_code'])\n                        s_code = int(original_header['sform_code'])\n\n                        patch_nifti_img.header.set_qform(original_affine, code=q_code)\n                        patch_nifti_img.header.set_sform(original_affine, code=s_code)\n\n                        # Save the patch to the specified location\n                        nib.save(patch_nifti_img, patch_img_filepath)\n\n# extract patches from the discs in the nii folder, for axial and sagittal patches\ndef extract_patches_from_discs(nii_folder, output_folder):\n    \"\"\"\n    Traverses a folder containing MRIs and associated sagittal segmentations.\n    For each axial image and associated sagittal segmentation, extracts patches for discs with labels 206 to 202.\n    Saves each patch in the corresponding folder structure within output_folder.\n\n    nii_folder : path to the folder containing MRIs and segmentations\n    output_folder : path to the folder where patches will be saved\n    \"\"\"\n    axial_images = []\n    axial_segmentations = []\n    sagittal_T2_segmentations = []\n    sagittal_T1_segmentations = []\n    sagittal_T1_images = []\n    sagittal_T2_images = []\n\n    \n    # Traverse files in the nii_folder\n    for filename in os.listdir(nii_folder):\n        if 'acq-ax' in filename and filename.endswith('.nii.gz') and not filename.endswith('_seg.nii.gz'):          \n            axial_images.append(filename)  # Axial images\n        elif 'acq-ax' in filename and 'T2w' in filename and 'total_seg.nii.gz' in filename:\n            axial_segmentations.append(filename)  # Sagittal segmentations\n        #elif 'acq-sag' in filename and 'T2w' in filename and 'total_seg.nii.gz' in filename:\n        #    sagittal_T2_segmentations.append(filename)\n        elif 'acq-sag' in filename and 'T1w' in filename and 'total_seg.nii.gz' in filename:\n            sagittal_T1_segmentations.append(filename)\n        #elif 'acq-sag' in filename and 'T2' in filename and filename.endswith('.nii.gz') and not filename.endswith('_seg.nii.gz'):          \n        #    sagittal_T2_images.append(filename)\n        elif 'acq-sag' in filename and 'T1' in filename and filename.endswith('.nii.gz') and not filename.endswith('_seg.nii.gz'):          \n            sagittal_T1_images.append(filename)\n\n    # Sort lists to ensure corresponding order\n    axial_segmentations.sort()\n    axial_images.sort()\n    sagittal_T2_segmentations.sort()\n    sagittal_T2_images.sort()\n    sagittal_T1_segmentations.sort()\n    sagittal_T1_images.sort()\n    sagittal_T2_segmentations.sort()\n    \n    print(len(sagittal_T2_segmentations),len(sagittal_T2_images))\n    print(len(sagittal_T1_segmentations),len(sagittal_T1_images))\n\n    extract_and_save_sagittal_patches(sagittal_T2_images, sagittal_T2_segmentations, nii_folder, output_folder)\n    extract_and_save_sagittal_patches(sagittal_T1_images, sagittal_T1_segmentations, nii_folder, output_folder)\n    extract_and_save_axial_patches(axial_images, axial_segmentations, nii_folder, output_folder)\n\n\n# function to extract patches from the discs in the nii folder for axial patches\ndef patch_extraction_volume(vol, mask, affine):\n    \"\"\"\n    Extract a 3D patch from an MRI volume with specific real-world dimensions.\n    \n    Parameters:\n    - vol: 3D numpy array representing the volume\n    - mask: 3D segmentation mask \n    - affine: Affine matrix from the NIfTI file\n    - header: Header from the NIfTI file\n    \n    Returns:\n    - patch: 3D numpy array with specified real-world dimensions\n    \"\"\"\n    # Convert mask to tensor for non-zero index extraction\n    mask = torch.Tensor(mask)\n    nonzero_indices = torch.nonzero(mask)\n    \n    # Calculate the centroid of the mask\n    centroid = nonzero_indices.float().mean(0).numpy().astype(int)\n    \n    # Get voxel sizes from the affine matrix\n    voxel_sizes = np.abs(np.diag(affine)[:3])\n    \n    # Calculate the number of voxels corresponding to 2.5 cm posterior displacement\n    posterior_displacement_cm = 20\n    posterior_displacement_voxels = (posterior_displacement_cm / voxel_sizes[1]).astype(int)\n    \n    # Compute the new centroid with posterior displacement\n    # Assuming the third dimension (index 2) is the posterior-anterior axis\n    displaced_centroid = centroid.copy()\n    displaced_centroid[1] -= posterior_displacement_voxels\n    \n    # Define desired patch sizes in cm\n    patch_sizes_cm = {\n        'RL': 60,  # Right-Left \n        'AP': 40,  # Anterior-Posterior\n        'SI': 30   # Superior-Inferior\n    }\n    \n    # Calculate patch size in voxels\n    patch_sizes_voxels = np.floor(np.array([\n        patch_sizes_cm['RL'] / voxel_sizes[0],\n        patch_sizes_cm['AP'] / voxel_sizes[1], \n        patch_sizes_cm['SI'] / voxel_sizes[2]\n    ])).astype(int)\n\n    # Extract patch\n    D, H, W = vol.shape\n    half_sizes = patch_sizes_voxels // 2\n    \n    patch = vol[\n        max(0, displaced_centroid[0] - half_sizes[0]):min(D, displaced_centroid[0] + half_sizes[0] + patch_sizes_voxels[0] % 2),\n        max(0, displaced_centroid[1] - half_sizes[1]):min(H, displaced_centroid[1] + half_sizes[1] + patch_sizes_voxels[1] % 2),\n        max(0, displaced_centroid[2] - half_sizes[2]):min(W, displaced_centroid[2] + half_sizes[2] + patch_sizes_voxels[2] % 2)\n    ]\n\n    return patch\n\ndef select_best_patches(folder_path):\n    discs = ['L1_L2', 'L2_L3', 'L3_L4', 'L4_L5', 'L5_S1']\n    disc_patches = {disc: [] for disc in discs}\n    \n    for filename in os.listdir(folder_path):\n        if filename.endswith('.nii.gz') and '_seg' not in filename:\n            for disc in discs:\n                if f\"{disc}_patch\" in filename:\n                    file_path = os.path.join(folder_path, filename)\n                    img = nib.load(file_path)\n                    resolution = img.header.get_zooms()\n                    voxel_volume = resolution[0] * resolution[1] * resolution[2]\n                    \n                    # Check for corresponding segmentation file\n                    seg_filename = filename.replace('.nii.gz', '_seg.nii.gz')\n                    seg_path = os.path.join(folder_path, seg_filename)\n                    \n                    if os.path.exists(seg_path):\n                        disc_patches[disc].append((file_path, seg_path, voxel_volume))\n                        \n    for disc, patches in disc_patches.items():\n        if len(patches) > 1:\n            # Sort patches by increasing voxel volume (resolution)\n            patches.sort(key=lambda x: x[2])\n            \n            # Keep the patch with the best resolution\n            best_patch = patches[0]\n            \n            # Remove other patches and their segmentations\n            for patch in patches[1:]:\n                os.remove(patch[0])\n                os.remove(patch[1])\n\n\n\ndef process_all_subjects_in_directory(root_dir, output_root_dir):\n    \"\"\"\n    Traverses all subdirectories in the root directory corresponding to subjects,\n    and applies the patch extraction function to each subdirectory.\n    \n    root_dir : root directory containing subject subdirectories\n    output_root_dir : root directory where output patches are stored\n    \"\"\"\n\n    sub_treated = 0\n    sub_failed = 0\n\n    for subject_folder in os.listdir(root_dir):\n        \n        subject_path = os.path.join(root_dir, subject_folder, \"anat\")\n        output_subject_path = os.path.join(output_root_dir, subject_folder, \"anat\")\n        \n        # Check if it is a subdirectory\n        if os.path.isdir(subject_path):\n            os.makedirs(output_subject_path, exist_ok=True)\n            try:\n                # Extract patches from discs\n                extract_patches_from_discs(subject_path, output_subject_path)\n                \n                # Select the best patches if there are multiple ones for the same disc\n                select_best_patches(output_subject_path)\n\n                print(f\"Processed subject {subject_folder}\")\n\n                sub_treated += 1\n\n            except Exception as e:\n                # Print a message indicating that an exception was raised\n                print(\"An exception was raised during patch processing.\")\n                \n                # Print the arguments passed to 'extract_patches_from_discs'\n                print(f\"Arguments for 'extract_patches_from_discs': subject_path={subject_path}, output_subject_path={output_subject_path}\")\n                \n                # Print the arguments passed to 'select_best_patches'\n                print(f\"Arguments for 'select_best_patches': output_subject_path={output_subject_path}\")\n                \n                # Print the type of exception and its details\n                print(f\"Error type: {type(e).__name__}, Details: {e}\")\n\n                sub_failed += 1\n    \n    print(f\"Processed {sub_treated} subjects, {sub_failed} subjects failed\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T17:37:24.485430Z","iopub.execute_input":"2025-04-28T17:37:24.485860Z","iopub.status.idle":"2025-04-28T17:37:24.519750Z","shell.execute_reply.started":"2025-04-28T17:37:24.485823Z","shell.execute_reply":"2025-04-28T17:37:24.519111Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# MONAI Transforms ","metadata":{}},{"cell_type":"code","source":"class ExtractSlicesD(Transform):\n    def __init__(self, keys=['image'], target_size=(384, 384), verbose=False):\n        self.keys = keys\n        self.target_size = target_size\n        self.resize = tio.Resize(target_shape=(*target_size, 1))\n        self.verbose = verbose\n\n    def __call__(self, data):\n        d = dict(data)\n        \n        for key in self.keys:\n            # Get image and remove channel dimension (1, X, Y, 6) -> (X, Y, 6)\n            image = d[key].squeeze(0)\n            for i in range(image.shape[2]):\n                # Extract slice, add channel dim for torchio,\n                # resize, then normalize\n                slice_2d = image[:, :, i]\n                slice_3d = slice_2d.unsqueeze(0).unsqueeze(-1)\n                if self.verbose:\n                    print(f\"Shape before resize: {slice_3d.shape}\")\n                slice_resized = self.resize(slice_3d)\n                if self.verbose:\n                    print(f\"Shape after resize: {slice_resized.shape}\")\n                # Remove the z dimension that we added\n                slice_final = slice_resized.squeeze(-1)\n                d[f'slice_{i}'] = slice_final\n                if self.verbose:\n                    print(f\"Final slice {i} shape: {slice_final.shape}\")\n        return d\n\nclass ExtractSlicesD_nfn(Transform):\n    def __init__(self, keys=['image'], target_size=(384, 384), verbose=False):\n        self.keys = keys\n        self.target_size = target_size\n        self.resize = tio.Resize(target_shape=(*target_size, 1))\n        self.verbose = verbose\n\n    def __call__(self, data):\n        d = dict(data)\n        \n        for key in self.keys:\n            # Get image and remove channel dimension (1, X, Y, 6) -> (X, Y, 6)\n            image = d[key].squeeze(0)\n            for i in range(image.shape[0]):\n                # Extract slice, add channel dim for torchio,\n                # resize, then normalize\n                slice_2d = image[i, :, :]\n                slice_3d = slice_2d.unsqueeze(0).unsqueeze(-1)\n                if self.verbose:\n                    print(f\"Shape before resize: {slice_3d.shape}\")\n                slice_resized = self.resize(slice_3d)\n                if self.verbose:\n                    print(f\"Shape after resize: {slice_resized.shape}\")\n                # Remove the z dimension that we added\n                slice_final = slice_resized.squeeze(-1)\n                d[f'slice_{i}'] = slice_final\n                if self.verbose:\n                    print(f\"Final slice {i} shape: {slice_final.shape}\")\n        return d","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T17:37:24.521311Z","iopub.execute_input":"2025-04-28T17:37:24.521626Z","iopub.status.idle":"2025-04-28T17:37:24.536088Z","shell.execute_reply.started":"2025-04-28T17:37:24.521591Z","shell.execute_reply":"2025-04-28T17:37:24.535334Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Subarticular Stenosis ","metadata":{}},{"cell_type":"code","source":"# regular transforms just creating the dataset\nregular_transforms = Compose([\n        LoadImaged(keys=['image']),\n        EnsureChannelFirstd(keys=[\"image\"]),\n    Spacingd(keys=['image'], pixdim=(0.4, 0.4, 4.4), mode=('bilinear'))  # Ré-échantillonnage de l'image\n    ])\n\nregular_transforms_flip = Compose([\n    LoadImaged(keys=['image']),\n    EnsureChannelFirstd(keys=[\"image\"]),\n    Spacingd(keys=['image'], pixdim=(0.4, 0.4, 4.4), mode=('bilinear')),  # Ré-échantillonnage de l'image\n    Flipd(keys=['image'], spatial_axis=0)\n])\n\ndef get_transforms_sas(mode='basic', side='left'):\n\n    if side == 'left':\n        reg_t = regular_transforms\n    elif side == 'right':\n        reg_t = regular_transforms_flip\n    \n    if mode == 'basic':\n        common_transforms = Compose([\n            SpatialCropd(keys=['image'], roi_start=(0, 0, 0), roi_end=(80, 100, 6)),  # crop pour récupérer la gauche\n            SpatialPadd(keys=['image'], spatial_size=(60, 80, 6)),  # Padding pour atteindre une taille fixe\n            CenterSpatialCropd(keys=['image'], roi_size=(60, 80, 6))  # Crop pour obtenir une taille fixe\n        ])\n\n    # Create list of transforms for processing 2D slices\n    slice_transforms = Compose([\n        # Custom transform to extract and resize slices\n        ExtractSlicesD(keys=['image'], target_size=(384, 384)),\n        # Scale and normalize\n        ScaleIntensityd(\n            keys=[f'slice_{i}' for i in range(6)]\n        ),\n        NormalizeIntensityd(\n            keys=[f'slice_{i}' for i in range(6)],\n            nonzero=True\n        ),\n        # Ensure all slices are tensors\n        ToTensord(\n            keys=[f'slice_{i}' for i in range(6)]\n        ),\n        # Concatenate all slices into a bag\n        ConcatItemsd(\n            keys=[f'slice_{i}' for i in range(6)],\n            name='bag',\n            dim=0\n        ),\n        # Add a transform to ensure bag has the correct shape\n        Lambdad(\n            keys=['bag'],\n            func=lambda x: x.reshape(6, 1, 384, 384)\n        )\n    ])\n\n    # Combine common_transforms with slice_transforms\n    transforms = Compose([reg_t, common_transforms, slice_transforms])\n\n    return transforms","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T17:37:24.537312Z","iopub.execute_input":"2025-04-28T17:37:24.537630Z","iopub.status.idle":"2025-04-28T17:37:24.556559Z","shell.execute_reply.started":"2025-04-28T17:37:24.537605Z","shell.execute_reply":"2025-04-28T17:37:24.555855Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def prepare_data_sas(list_subjects, data_dir, transform_left, transform_right):\n    data_right = []\n    data_left = []\n    \n    counter = 0\n    \n    # Dictionnaire de conversion des étiquettes\n    text2int = {\"Normal/Mild\": 0, \"Moderate\": 1, \"Severe\": 2}\n    \n    for subject in list_subjects:\n        \n        subject_dir = os.path.join(data_dir, f'sub-{subject}', 'anat')\n        \n        if os.path.isdir(subject_dir):\n            for file in os.listdir(subject_dir):\n                \n                 if '_patch.nii.gz' in file and 'foramen' not in file:\n                    image_path = os.path.join(subject_dir, file)\n                    \n                    parts = image_path.split('_')\n                    disk_level = f\"{parts[-3]}_{parts[-2]}\"\n\n                    if os.path.exists(image_path):\n                        \n                        subject_id = (subject.replace('sub-', ''))\n                        \n                        label_column = f'_subarticular_stenosis_{disk_level.lower()}'\n                        \n                        \n                        label_left = f\"{subject_id}_left{label_column}\"\n                        label_right = f\"{subject_id}_right{label_column}\"\n                        \n                        \n                        data_right.append({\"image\": image_path, \"label\": label_right})\n                        data_left.append({\"image\": image_path, \"label\": label_left})\n                        counter += 2\n\n    print(f\"Nombre de données chargées: {counter}\")\n    \n    return ConcatDataset([Dataset(data=data_left, transform=transform_left), Dataset(data=data_right, transform=transform_right)])\n\n\n\n                           ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T17:37:24.557523Z","iopub.execute_input":"2025-04-28T17:37:24.557763Z","iopub.status.idle":"2025-04-28T17:37:24.565412Z","shell.execute_reply.started":"2025-04-28T17:37:24.557740Z","shell.execute_reply":"2025-04-28T17:37:24.564735Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Spinal Canal Stenosis ","metadata":{}},{"cell_type":"code","source":"def get_transforms_scs(mode='basic'):\n\n\n    regular_transforms = Compose([\n        LoadImaged(keys=['image']),\n        EnsureChannelFirstd(keys=[\"image\"]),\n        Spacingd(keys=['image'], pixdim=(0.4, 0.4, 4.4), mode=('bilinear')),  # Ré-échantillonnage de l'image\n    ])\n    \n    if mode == 'basic':\n        common_transforms = Compose([\n            SpatialPadd(keys=['image'], spatial_size=(120, 80, 6)),\n            CenterSpatialCropd(\n                keys=['image'],\n                roi_size=(120, 80, 6)\n            ),\n        ])\n\n    elif mode == 'random':\n        # Same transforms but with random augmentations\n        common_transforms = Compose([\n            RandRotated(keys=['image'], prob=1.0, range_x=0.2),\n            RandAffined(keys=['image'], prob=1.0, shear_range=(0.3, 0.3, 0.3)),\n            Rand3DElasticd(keys=['image'], prob=0.5, sigma_range=(8, 12), magnitude_range=(100, 200)),\n            RandGaussianNoised(keys=['image'], prob=0.5, mean=0.0, std=0.1),\n            RandBiasFieldd(keys=['image'], prob=0.5, coeff_range=(0, 0.4)),\n            RandZoomd(keys=['image'], prob=0.5, min_zoom=0.95, max_zoom=1.15),\n            SpatialPadd(keys=['image'], spatial_size=(120, 80, 6)),\n            RandSpatialCropd(\n                keys=['image'],\n                roi_size=(120, 80, 6),\n                random_size=False\n            ),\n        ])\n\n    # Create list of transforms for processing 2D slices\n    slice_transforms = Compose([\n        # Custom transform to extract and resize slices\n        ExtractSlicesD(keys=['image'], target_size=(384, 384)),\n        # Scale and normalize\n        ScaleIntensityd(\n            keys=[f'slice_{i}' for i in range(6)]\n        ),\n        NormalizeIntensityd(\n            keys=[f'slice_{i}' for i in range(6)],\n            nonzero=True\n        ),\n        # Ensure all slices are tensors\n        ToTensord(\n            keys=[f'slice_{i}' for i in range(6)]\n        ),\n        # Concatenate all slices into a bag\n        ConcatItemsd(\n            keys=[f'slice_{i}' for i in range(6)],\n            name='bag',\n            dim=0\n        ),\n        # Add a transform to ensure bag has the correct shape\n        Lambdad(\n            keys=['bag'],\n            func=lambda x: x.reshape(6, 1, 384, 384)\n        )\n    ])\n\n    # Combine common_transforms with slice_transforms\n    transforms = Compose([regular_transforms, common_transforms, slice_transforms])\n\n    return transforms\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T17:37:24.566447Z","iopub.execute_input":"2025-04-28T17:37:24.566709Z","iopub.status.idle":"2025-04-28T17:37:24.578219Z","shell.execute_reply.started":"2025-04-28T17:37:24.566675Z","shell.execute_reply":"2025-04-28T17:37:24.577574Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def prepare_data_scs(list_subjects, data_dir, transform):\n    data = []\n    \n    counter = 0\n\n    # Dictionnaire de conversion des étiquettes\n    text2int = {\"Normal/Mild\": 0, \"Moderate\": 1, \"Severe\": 2}\n    \n    for subject in list_subjects:\n        \n        subject_dir = os.path.join(data_dir, f'sub-{subject}', 'anat')\n        if os.path.isdir(subject_dir):\n            for file in os.listdir(subject_dir):\n                \n                if '_patch.nii.gz' in file and 'foramen' not in file:\n                    image_path = os.path.join(subject_dir, file)\n                    \n                    parts = image_path.split('_')\n                    disk_level = f\"{parts[-3]}_{parts[-2]}\"\n\n                    if os.path.exists(image_path):\n                        \n                        \n                        subject_id = (subject.replace('sub-', ''))\n                        \n                        label_column = f'spinal_canal_stenosis_{disk_level.lower()}'\n                        \n                         \n                        \n                        counter += 1\n                        label = f\"{subject_id}_{label_column}\"\n                        data.append({\"image\": image_path, \"label\": label})\n\n\n    print(f\"Nombre de données chargées: {counter}\")\n    return Dataset(data=data, transform=transform)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T17:37:24.579157Z","iopub.execute_input":"2025-04-28T17:37:24.579683Z","iopub.status.idle":"2025-04-28T17:37:24.590295Z","shell.execute_reply.started":"2025-04-28T17:37:24.579642Z","shell.execute_reply":"2025-04-28T17:37:24.589573Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Neural Foraminal Narowing ","metadata":{}},{"cell_type":"code","source":"def get_transforms_nfn(mode='basic', side = \"right\"):\n\n    regular_transforms = Compose([\n        LoadImaged(keys=['image']),\n        EnsureChannelFirstd(keys=[\"image\"]),\n    ])\n\n    if side == \"left\": \n        regular_transforms = Compose([regular_transforms,Flipd(keys=['image'], spatial_axis=0)])\n\n\n    if mode == 'basic':\n        common_transforms = Compose([\n            Spacingd(keys=['image'], pixdim=(4.0, 0.4, 0.4), mode=('bilinear')),\n            SpatialPadd(keys=['image'], spatial_size=(6, 100, 100)),\n            CenterSpatialCropd(\n                keys=['image'],\n                roi_size=(6, 100, 100)\n            ),\n        ])\n\n    elif mode == 'random':\n        # Same transforms but with random augmentations\n        common_transforms = Compose([\n            Spacingd(keys=['image'], pixdim=(4.0, 0.4, 0.4), mode=('bilinear')),\n            RandRotated(keys=['image'], prob=0.8, range_x=0.2),\n            RandGaussianNoised(keys=['image'], prob=0.4, mean=0.0, std=0.1),\n            RandBiasFieldd(keys=['image'], prob=0.4, coeff_range=(0, 0.3)),\n            SpatialPadd(keys=['image'], spatial_size=(6, 100, 100)),\n            RandSpatialCropd(\n                keys=['image'],\n                roi_size=(6, 100, 100),\n                random_size=False\n            ),\n        ])\n\n    # Create list of transforms for processing 2D slices\n    slice_transforms = Compose([\n        # Custom transform to extract and resize slices\n        ExtractSlicesD_nfn(keys=['image'], target_size=(224, 224)),\n        # Scale and normalize\n        ScaleIntensityd(\n            keys=[f'slice_{i}' for i in range(6)]\n        ),\n        NormalizeIntensityd(\n            keys=[f'slice_{i}' for i in range(6)],\n            nonzero=True\n        ),\n        # Ensure all slices are tensors\n        ToTensord(\n            keys=[f'slice_{i}' for i in range(6)]\n        ),\n        # Concatenate all slices into a bag\n        ConcatItemsd(\n            keys=[f'slice_{i}' for i in range(6)],\n            name='bag',\n            dim=0\n        ),\n        # Add a transform to ensure bag has the correct shape\n        Lambdad(\n            keys=['bag'],\n            func=lambda x: x.reshape(6, 1, 224, 224)\n        )\n    ])\n\n    # Combine common_transforms with slice_transforms\n    transforms = Compose([regular_transforms, common_transforms, slice_transforms])\n\n    return transforms","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T17:37:24.591265Z","iopub.execute_input":"2025-04-28T17:37:24.591538Z","iopub.status.idle":"2025-04-28T17:37:24.600457Z","shell.execute_reply.started":"2025-04-28T17:37:24.591507Z","shell.execute_reply":"2025-04-28T17:37:24.599753Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def prepare_data_nfn(list_subject, data_dir, random=False):\n    data_right = []\n    data_left = []\n    \n\n    counter = 0\n    # Label conversion dictionary\n    text2int = {\"Normal/Mild\": 0, \"Moderate\": 1, \"Severe\": 2}\n\n    for subject in list_subjects:\n        \n        subject_dir = os.path.join(data_dir, f'sub-{subject}', 'anat')\n        if os.path.isdir(subject_dir):\n            for file in os.listdir(subject_dir):\n                if '_patch.nii.gz' in file and 'foramen' in file and 'T1' in file:\n                    image_path = os.path.join(subject_dir, file)\n                    parts = image_path.split('_')\n                    disk_level = f\"{parts[-5]}_{parts[-4]}\"\n\n                    if os.path.exists(image_path):\n                        \n                        subject_id = subject\n                        if 'left' in file:\n                            orientation = 'right'\n                        elif 'right' in file: \n                            orientation = 'left'\n                        label_column = (\n                            f'{orientation}_neural_foraminal_narrowing_{disk_level.lower()}'\n                        )\n                        # Get raw label\n                        label = f\"{subject_id}_{label_column}\"\n\n                        # Convert text label to numeric value\n                        \n                        counter += 1\n                        if \"left\" in image_path: \n                            data_right.append({\n                                \"image\": image_path,\n                                \"label\": label\n                            })\n                        if \"right\" in image_path: \n                            data_left.append({\n                                \"image\": image_path,\n                                \"label\": label\n                            })\n\n    print(f\"Number of loaded data: {counter}\")\n    return ConcatDataset([Dataset(data=data_left, transform=get_transforms_nfn(mode='random', side='left') if random else get_transforms_nfn(mode='basic',side='left')), Dataset(data=data_right, transform=get_transforms_nfn(mode='random', side='right') if random else get_transforms_nfn(mode='basic',side='right'))]) \n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T17:37:24.601251Z","iopub.execute_input":"2025-04-28T17:37:24.601499Z","iopub.status.idle":"2025-04-28T17:37:24.612720Z","shell.execute_reply.started":"2025-04-28T17:37:24.601458Z","shell.execute_reply":"2025-04-28T17:37:24.612118Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Model Loading ","metadata":{}},{"cell_type":"code","source":"'''\nFile to introduce a MIL model\nNote that loads of hyperparameters could be included as arguments\nIt could be avg pooling size, hidden dim, etc...\nAlso encoder could be changed to a different model\n'''\n\nimport torch\nimport torch.nn as nn\n\n# import timm for models\nimport timm\n\n\n# define a MIL model\nclass MILsection(nn.Module):\n    def __init__(self, input_dim, hidden_dim, num_classes, num_layers=1):\n        super(MILsection, self).__init__()\n        self.num_layers = num_layers\n        if num_layers > 0:\n            self.rnn = nn.GRU(input_dim, input_dim//2, num_layers=num_layers,\n                             batch_first=True, dropout=0.1, bidirectional=True)\n        self.aux_attention = nn.Sequential(\n            nn.Tanh(),\n            nn.Linear(input_dim, 1)\n        )\n        self.attention = nn.Sequential(\n            nn.Tanh(),\n            nn.Linear(input_dim, 1)\n        )\n\n    def forward(self, bags):\n        \"\"\"\n        Args:\n            bags: (batch_size, num_instances, input_dim)\n\n        Returns:\n            logits: (batch_size, num_classes)\n        \"\"\"\n        batch_size, num_instances, input_dim = bags.size()\n\n        if self.num_layers > 0:\n            bags_rnn, _ = self.rnn(bags)\n        else:\n            bags_rnn = bags\n        \n        # Main attention\n        attn_scores = self.attention(bags_rnn).squeeze(-1)  # [batch_size, num_instances]\n        attn_weights = torch.softmax(attn_scores, dim=-1)  # [batch_size, num_instances]\n        weighted_instances = torch.bmm(attn_weights.unsqueeze(1), bags_rnn).squeeze(1)  # [batch_size, input_dim]\n        \n        # Auxiliary attention - process each instance independently\n        aux_attn_scores = self.aux_attention(bags_rnn).squeeze(-1)  # [batch_size, num_instances]\n        aux_features = bags_rnn  # [batch_size, num_instances, input_dim]\n        \n        return weighted_instances, aux_features\n\n\n# here define the whole MIL model\n# uses the MILsection model and a ConvNext Small as a feature extractor\n# note that loads of hyperparameters could be included as arguments\nclass MILmodel(nn.Module):\n    def __init__(self, encoder, num_layers=1):\n        super(MILmodel, self).__init__()\n        # encoder\n        self.encoder = encoder\n        # flattening layer, applying pooling and flattening\n        # note here that we could try different pooling methods\n        self.flatten = nn.Sequential(\n            nn.AdaptiveAvgPool2d(1),\n            nn.Flatten(1)\n        )\n        self.feature_size = self.encoder.num_features\n        \n        # MIL section, loads of hyperparameters here also\n        self.mil_section = MILsection(input_dim=self.feature_size,\n                                    hidden_dim=1024, \n                                    num_classes=3,\n                                    num_layers=num_layers)\n        # classifier output\n        self.classifier = nn.Linear(self.feature_size, 3)\n        # aux classifier output - now takes each instance independently\n        self.aux_classifier = nn.Linear(self.feature_size, 3)\n\n    def forward(self, x):\n        # x shape: (batch_size, 6, 1, 384, 384)\n        batch_size, num_instances, channels, H, W = x.shape\n\n        # Reshape to process all instances through encoder\n        x = x.reshape(-1, channels, H, W)  # shape: (batch_size * 6, 1, 384, 384)\n        \n        # Pass through encoder\n        x = self.encoder.forward_features(x)  # shape: (batch_size * 6, feature_size, h', w')\n        \n        # Apply pooling and flatten\n        x = self.flatten(x)  # shape: (batch_size * 6, feature_size)\n        \n        # Reshape back to separate instances\n        x = x.reshape(batch_size, num_instances, self.feature_size)  # shape: (batch_size, 6, feature_size)\n        \n        # Pass through MIL section\n        weighted_instances, aux_features = self.mil_section(x)\n        # weighted_instances: (batch_size, feature_size)\n        # aux_features: (batch_size, num_instances, feature_size)\n        \n        # Main classification\n        main_output = self.classifier(weighted_instances)  # shape: (batch_size, 3)\n        \n        # Auxiliary classification - apply to each instance independently\n        aux_output = self.aux_classifier(aux_features)  # shape: (batch_size, num_instances, 3)\n        # Average the auxiliary predictions across instances\n        aux_output = aux_output.mean(dim=1)  # shape: (batch_size, 3)\n        \n        return main_output, aux_output\n\n\nconvnext_small_sas = timm.create_model('convnext_small.fb_in22k_ft_in1k_384',\n                                   in_chans=1, pretrained=False, num_classes=0)\nconvnext_small_scs = timm.create_model('convnext_small.fb_in22k_ft_in1k_384',\n                                   in_chans=1, pretrained=False, num_classes=0)\nconvnext_small_nfn = timm.create_model('convnext_small.fb_in22k_ft_in1k_384',\n                                   in_chans=1, pretrained=False, num_classes=0)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T17:37:24.613922Z","iopub.execute_input":"2025-04-28T17:37:24.614454Z","iopub.status.idle":"2025-04-28T17:37:28.146184Z","shell.execute_reply.started":"2025-04-28T17:37:24.614421Z","shell.execute_reply":"2025-04-28T17:37:28.145359Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def load_model(model_path, layers, patho, device='cuda'):\n    if patho == 'sas':\n        encoder = convnext_small_sas\n    elif patho == 'scs':\n        encoder = convnext_small_scs\n    elif patho == 'nfn':\n        encoder = convnext_small_nfn\n        \n    \"\"\"Load the trained MIL model from the checkpoint.\"\"\"\n    checkpoint = torch.load(model_path, map_location=device)\n    model = MILmodel(encoder=encoder, num_layers=layers).to(device)\n    model.load_state_dict(checkpoint['model_state_dict'])\n\n    model.eval()\n    return model\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T17:39:55.631192Z","iopub.execute_input":"2025-04-28T17:39:55.631959Z","iopub.status.idle":"2025-04-28T17:39:55.637053Z","shell.execute_reply.started":"2025-04-28T17:39:55.631922Z","shell.execute_reply":"2025-04-28T17:39:55.636149Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Subarticular stenosis ","metadata":{}},{"cell_type":"code","source":"model_sas = load_model('/kaggle/input/sas_mil/other/default/2/best_mil_model.pth',2, 'sas')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T17:39:58.912727Z","iopub.execute_input":"2025-04-28T17:39:58.913384Z","iopub.status.idle":"2025-04-28T17:39:59.904100Z","shell.execute_reply.started":"2025-04-28T17:39:58.913351Z","shell.execute_reply":"2025-04-28T17:39:59.903380Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Spinal Canal Stenosis ","metadata":{}},{"cell_type":"code","source":"model_scs = load_model('/kaggle/input/scs_mil/other/default/2/best_mil_model.pth' ,2, 'scs')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T17:40:05.644120Z","iopub.execute_input":"2025-04-28T17:40:05.644851Z","iopub.status.idle":"2025-04-28T17:40:10.335136Z","shell.execute_reply.started":"2025-04-28T17:40:05.644812Z","shell.execute_reply":"2025-04-28T17:40:10.334139Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Neural Foraminal Narowing ","metadata":{}},{"cell_type":"code","source":"model_nfn = load_model('/kaggle/input/nfn_mil/other/default/5/best_mil_model5861.pth' ,2, 'nfn')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T17:40:16.286153Z","iopub.execute_input":"2025-04-28T17:40:16.286832Z","iopub.status.idle":"2025-04-28T17:40:21.002935Z","shell.execute_reply.started":"2025-04-28T17:40:16.286796Z","shell.execute_reply":"2025-04-28T17:40:21.001846Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Preprocessing \nThe goal of this section is to create the function to process the subjects in the batch. ","metadata":{}},{"cell_type":"code","source":"def copy_subject_folders(list_subjects, train, target_dir):\n    \"\"\"\n    Copies folders matching the subject names from the source directory to the target directory.\n\n    :param list_subjects: List of subject folder names to copy (e.g., [\"01\", \"02\"]).\n    :param train: boolean to know if you want to fetch your subjects from train or test dataset\n    \"\"\"\n    if train: \n        source_dir = \"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_images\"\n    else: \n        source_dir = \"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/test_images\"\n\n   \n    \n    if not os.path.exists(target_dir):\n        os.makedirs(target_dir)\n\n    for subject in list_subjects:\n        subject_path = os.path.join(source_dir, subject)\n        target_path = os.path.join(target_dir, subject)\n        \n        if os.path.exists(subject_path):\n            shutil.copytree(subject_path, target_path)\n        else:\n            print(f\"Subject folder {subject} does not exist in {source_dir}\")\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T17:40:37.034417Z","iopub.execute_input":"2025-04-28T17:40:37.035211Z","iopub.status.idle":"2025-04-28T17:40:37.040525Z","shell.execute_reply.started":"2025-04-28T17:40:37.035176Z","shell.execute_reply":"2025-04-28T17:40:37.039626Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def preprocessing(list_subjects, train=False): \n    \n    ### Copy the dcom files \n    dcom_folder = \"/kaggle/working/dcom_data\"\n    copy_subject_folders(list_subjects, train, dcom_folder)\n    \n    ### niftification \n    nii_folder = \"/kaggle/working/nii_data\"\n\n          \n        \n    os.makedirs(nii_folder, exist_ok=True)\n\n    \n    # Process subjects and set up directories\n    \n    for subject_id in list_subjects:  # Adjust range as needed: 1975 subjects\n        try: \n        \n            # Create specific directories\n            os.makedirs(os.path.join(nii_folder, f'sub-{subject_id}'), exist_ok=True)\n            os.makedirs(os.path.join(nii_folder, f'sub-{subject_id}', 'anat'), exist_ok=True)\n\n            # Process subject and set up directories\n            process_subject(subject_id, input_folder, nii_folder, df_meta_f, meta_obj)\n        except: \n            print(f'failed preprocessing for {subject_id}')\n \n\n    for item in os.listdir(nii_folder):\n        item_path = os.path.join(nii_folder, item)\n        \n        # Check if the item is a directory and starts with \"sub\"\n        if os.path.isdir(item_path) and item.startswith(\"sub\"):\n            continue  # Skip deletion for folders starting with \"sub\"\n        \n        # Delete the item (file or directory)\n        if os.path.isfile(item_path):\n            os.remove(item_path)\n            print(f\"Deleted file: {item_path}\")\n        elif os.path.isdir(item_path):\n            shutil.rmtree(item_path)\n            print(f\"Deleted folder: {item_path}\")\n\n    ### Totalspineseg \n    os.system(\"cp -r /kaggle/input/nnunet_totalspineseg/other/20250108_totalspineseg_version/10/nnUNet /kaggle/working/TotalSpineSeg/tss_input\")\n\n    run_totalspineseg(nii_folder)\n    \n    ### Patch extraction \n    process_directory_other(nii_folder)\n    process_all_subjects_in_directory(nii_folder, nii_folder)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T17:40:39.790200Z","iopub.execute_input":"2025-04-28T17:40:39.791013Z","iopub.status.idle":"2025-04-28T17:40:39.798022Z","shell.execute_reply.started":"2025-04-28T17:40:39.790979Z","shell.execute_reply":"2025-04-28T17:40:39.797163Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Evaluation of models ","metadata":{}},{"cell_type":"code","source":"def eval_sas(list_subjects, data_dir):     \n    # Préparer les données\n    transform_left=get_transforms_sas(mode='basic',side='left')\n    transform_right=get_transforms_sas(mode='basic',side='right')\n    data = prepare_data_sas(list_subjects, data_dir, transform_left=transform_left, transform_right=transform_right)\n    data_loader = DataLoader(data, batch_size=16)\n    \n    pred = []\n    i=0\n    \n    with torch.no_grad():\n            for batch in tqdm(data_loader):\n                \n                \n                inputs = batch[\"bag\"].cuda()\n                labels = batch[\"label\"]\n                \n                main, aux = model_sas(inputs)\n\n                '''if i<3:\n                    print(\"VISUALIZING SAS BATCH\")\n                    visualize_batch(batch)\n                    i+=1'''\n                probs = torch.softmax(main, dim=1).cpu().numpy()\n                #adjusted_probs = probs.copy()\n                #max_class_indices = np.argmax(probs, axis=1)\n                #boost_mask = (max_class_indices == 2)\n            \n                # Appliquer le boost uniquement sur les lignes concernées\n                #adjusted_probs[boost_mask, 2] *= 1.3\n            \n                # Renormalisation\n                #row_sums = adjusted_probs.sum(axis=1, keepdims=True)\n                #normalized_probs = adjusted_probs / row_sums\n\n                #boost_mask_mid = (max_class_indices == 1)\n            \n                # Appliquer le boost uniquement sur les lignes concernées\n                #normalized_probs[boost_mask, 1] *= 1.15\n            \n                # Renormalisation\n                #row_sums_1 = normalized_probs.sum(axis=1, keepdims=True)\n                #normalized_probs2 = normalized_probs / row_sums_1\n\n                # change for normalized probs to apply boosting\n                outputs = list(probs)\n                \n                for i in range(len(labels)): \n                    label = labels [i]\n                    output = list(outputs[i])\n                    pred.append((label, output))\n        \n    return pred \n        ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T17:40:46.581309Z","iopub.execute_input":"2025-04-28T17:40:46.581638Z","iopub.status.idle":"2025-04-28T17:40:46.588882Z","shell.execute_reply.started":"2025-04-28T17:40:46.581610Z","shell.execute_reply":"2025-04-28T17:40:46.588006Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def visualize_batch(batch):\n    \"\"\"\n    Visualize a batch of images and save them to wandb\n    batch: dictionary containing 'bag' tensor of shape [B, 6, 1, 384, 384] and 'label'\n    epoch: current epoch number\n    \"\"\"\n    # Get the first batch and ensure it's on CPU\n    images = batch['bag'].cpu().detach()  # Shape: [B, 6, 1, 384, 384]\n    labels = batch['label']\n    \n    # Take only the first 4 samples to avoid too large figures\n    n_samples = min(4, images.shape[0])\n    \n    # Create a figure with subplots for each sample and its 6 slices\n    fig, axes = plt.subplots(n_samples, 6, figsize=(20, 4*n_samples))\n    if n_samples == 1:\n        axes = axes[None, :]  # Add dimension for consistent indexing\n    \n    for i in range(n_samples):\n        for j in range(6):\n            # Get the image slice and ensure it's a valid image\n            img = images[i, j, 0].numpy()\n            \n            # Normalize the image for better visualization\n            img = (img - img.min()) / (img.max() - img.min() + 1e-8)\n            \n            # Plot the image\n            axes[i, j].imshow(img, cmap='gray')\n            axes[i, j].axis('off')\n            \n            # Add title only to the first row\n            if i == 0:\n                axes[i, j].set_title(f'Slice {j+1}')\n    \n    plt.tight_layout()\n    plt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T17:40:50.844677Z","iopub.execute_input":"2025-04-28T17:40:50.845028Z","iopub.status.idle":"2025-04-28T17:40:50.851753Z","shell.execute_reply.started":"2025-04-28T17:40:50.844998Z","shell.execute_reply":"2025-04-28T17:40:50.850946Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def eval_scs(list_subjects, data_dir):     \n    # Préparer les données\n    transform=get_transforms_scs(mode='basic')\n    data = prepare_data_scs(list_subjects, data_dir, transform)\n    data_loader = DataLoader(data, batch_size=16)\n    pred = []\n    i = 0\n    \n    with torch.no_grad():\n            for batch in tqdm(data_loader):\n\n\n                \n                inputs = batch[\"bag\"].cuda()\n                labels = batch[\"label\"]\n                '''if i<3:\n                    print(\"VISUALIZING SCS BATCH\")\n                    visualize_batch(batch)\n                    i+=1'''\n                main, aux = model_scs(inputs)\n\n                out_arr = torch.softmax(main, dim=1).cpu().numpy()\n                #out_arr[:, 2] *= 1.25  # Multiplie la proba de la 3ème classe\n                # Renormalisation\n                #row_sums = out_arr.sum(axis=1, keepdims=True)\n                #out_norm = out_arr / row_sums\n\n                # change for out_norm to apply severe change\n                outputs = list(out_arr)\n                \n                for i in range(len(labels)): \n                    label = labels [i]\n                    output = list(outputs[i])\n                    pred.append((label, output))\n        \n    return pred \n \n        ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T17:40:54.151872Z","iopub.execute_input":"2025-04-28T17:40:54.152658Z","iopub.status.idle":"2025-04-28T17:40:54.161385Z","shell.execute_reply.started":"2025-04-28T17:40:54.152612Z","shell.execute_reply":"2025-04-28T17:40:54.160464Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def eval_nfn(list_subjects, data_dir): \n    \n        # Préparer les données\n        transform=get_transforms_nfn()\n        data = prepare_data_nfn(list_subjects, data_dir, random=False)\n        data_loader = DataLoader(data, batch_size=16)\n    \n        model_nfn.eval()\n        \n        pred = []\n        \n        with torch.no_grad():\n            for batch in tqdm(data_loader):\n                \n                \n                inputs = batch[\"bag\"].cuda()\n                labels = batch[\"label\"]\n                \n                main, aux = model_scs(inputs)\n\n                out_arr = torch.softmax(main, dim=1).cpu().numpy()\n                \n                outputs = list(out_arr)\n                \n                for i in range(len(labels)): \n                    label = labels [i]\n                    output = list(outputs[i])\n                    pred.append((label, output))\n        return pred ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T17:40:57.831519Z","iopub.execute_input":"2025-04-28T17:40:57.831878Z","iopub.status.idle":"2025-04-28T17:40:57.838068Z","shell.execute_reply.started":"2025-04-28T17:40:57.831847Z","shell.execute_reply":"2025-04-28T17:40:57.837226Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Global evaluation","metadata":{}},{"cell_type":"code","source":"def eval(list_subjects, data_dir): \n\n\n    pred_nfn = eval_nfn(list_subjects, data_dir)\n    pred_scs = eval_scs(list_subjects, data_dir)\n    pred_sas = eval_sas(list_subjects, data_dir)\n\n    for label, output in pred_nfn: \n        result.loc[result[\"row_id\"] == label, ['normal_mild', 'moderate', 'severe']] = output\n    for label, output in pred_scs: \n        result.loc[result[\"row_id\"] == label, ['normal_mild', 'moderate', 'severe']] = output\n    for label, output in pred_sas: \n        result.loc[result[\"row_id\"] == label, ['normal_mild', 'moderate', 'severe']] = output\n   \n\n    ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T17:41:00.898464Z","iopub.execute_input":"2025-04-28T17:41:00.899239Z","iopub.status.idle":"2025-04-28T17:41:00.903701Z","shell.execute_reply.started":"2025-04-28T17:41:00.899205Z","shell.execute_reply":"2025-04-28T17:41:00.902805Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Inference Code","metadata":{}},{"cell_type":"code","source":"def inference(list_subjects): \n    list_subjects = list(map(str, list_subjects))\n    preprocessing(list_subjects)\n    data_dir = \"/kaggle/working/nii_data\"\n    eval(list_subjects, data_dir)\n    shutil.rmtree(\"/kaggle/working/nii_data\")\n    shutil.rmtree(\"/kaggle/working/dcom_data\")\n    ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T17:41:04.886599Z","iopub.execute_input":"2025-04-28T17:41:04.887410Z","iopub.status.idle":"2025-04-28T17:41:04.891571Z","shell.execute_reply.started":"2025-04-28T17:41:04.887376Z","shell.execute_reply":"2025-04-28T17:41:04.890615Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Main ","metadata":{}},{"cell_type":"code","source":"if train: \n        csv_description = \"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_series_descriptions.csv\"\n        input_folder = '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_images'\nelse: \n        csv_description = \"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/test_series_descriptions.csv\"\n        input_folder = '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/test_images'\n\n\ndf_meta_f = pd.read_csv(csv_description, sep=',')\nsubject_ids = np.unique(df_meta_f[\"study_id\"].values)\n\n# List out all of the Studies we have on patients.\npart_1 = os.listdir(input_folder)\npart_1 = list(filter(lambda x: x.find('.DS') == -1, part_1))\n\np1 = [(x, f\"{input_folder}/{x}\") for x in part_1]\nmeta_obj = { p[0]: { 'folder_path': p[1], \n                    'SeriesInstanceUIDs': [] \n                } \n            for p in p1 }\n\nfor m in meta_obj:\n    meta_obj[m]['SeriesInstanceUIDs'] = list( \n        filter(lambda x: x.find('.DS') == -1, \n            os.listdir(meta_obj[m]['folder_path'])\n            )\n    )\n# Grabs the corresponding series descriptions\nfor k in tqdm(meta_obj):\n    for s in meta_obj[k]['SeriesInstanceUIDs']:\n        if 'SeriesDescriptions' not in meta_obj[k]:\n            meta_obj[k]['SeriesDescriptions'] = []\n        try:\n            meta_obj[k]['SeriesDescriptions'].append(\n                df_meta_f[(df_meta_f['study_id'] == int(k)) & \n                (df_meta_f['series_id'] == int(s))]['series_description'].iloc[0])\n        except:\n            None","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T17:41:08.963068Z","iopub.execute_input":"2025-04-28T17:41:08.963671Z","iopub.status.idle":"2025-04-28T17:41:09.018623Z","shell.execute_reply.started":"2025-04-28T17:41:08.963637Z","shell.execute_reply":"2025-04-28T17:41:09.017812Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df = pd.read_csv('/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/sample_submission.csv')\ndf['study_id'] = df.row_id.str.extract(r'([0-9]*)_').astype('int')\ndf['side'] = df.row_id.str.extract(r'[0-9]_([a-z]*)_')\ndf['level'] = df.row_id.str.extract(r'([a-z0-9]*)_[a-z0-9]*$')\ndf.side = df.side.map({'left':0,'right':1,'spinal':2})\ndf.level = df.level.map({'l1':0,'l2':1,'l3':2,'l4':3,'l5':4})\nresult = df[['row_id']]\nresult.loc[:, ['normal_mild', 'moderate', 'severe']] = 1/3","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T17:41:15.009183Z","iopub.execute_input":"2025-04-28T17:41:15.009527Z","iopub.status.idle":"2025-04-28T17:41:15.034582Z","shell.execute_reply.started":"2025-04-28T17:41:15.009497Z","shell.execute_reply":"2025-04-28T17:41:15.033677Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"result","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T17:41:18.517075Z","iopub.execute_input":"2025-04-28T17:41:18.517879Z","iopub.status.idle":"2025-04-28T17:41:18.532246Z","shell.execute_reply.started":"2025-04-28T17:41:18.517841Z","shell.execute_reply":"2025-04-28T17:41:18.531330Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#shutil.rmtree(\"/kaggle/working/nii_data\")\n#shutil.rmtree(\"/kaggle/working/dcom_data\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T17:41:25.971751Z","iopub.execute_input":"2025-04-28T17:41:25.972463Z","iopub.status.idle":"2025-04-28T17:41:26.577843Z","shell.execute_reply.started":"2025-04-28T17:41:25.972429Z","shell.execute_reply":"2025-04-28T17:41:26.576702Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"subject_ids = np.unique(df_meta_f[\"study_id\"].values).tolist()\n\nB = 48\nN = len(subject_ids)\nQ, R = N//B, N%B\n\nfor i in range(Q):\n    \n    list_subjects = subject_ids[i*B:(i+1)*B]\n \n    inference(list_subjects)\n    \n    \nif R!=0:\n    list_subjects = subject_ids[Q*B:]\n    inference(list_subjects)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T17:41:29.400206Z","iopub.execute_input":"2025-04-28T17:41:29.401040Z","iopub.status.idle":"2025-04-28T17:42:43.971813Z","shell.execute_reply.started":"2025-04-28T17:41:29.401006Z","shell.execute_reply":"2025-04-28T17:42:43.970832Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"sub = result.copy()\nsub = sub[['row_id','normal_mild','moderate','severe']]\nsub.to_csv('submission.csv',index=False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T17:42:53.165176Z","iopub.execute_input":"2025-04-28T17:42:53.165516Z","iopub.status.idle":"2025-04-28T17:42:53.174738Z","shell.execute_reply.started":"2025-04-28T17:42:53.165485Z","shell.execute_reply":"2025-04-28T17:42:53.173856Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"sub","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T17:42:55.807115Z","iopub.execute_input":"2025-04-28T17:42:55.807462Z","iopub.status.idle":"2025-04-28T17:42:55.820358Z","shell.execute_reply.started":"2025-04-28T17:42:55.807429Z","shell.execute_reply":"2025-04-28T17:42:55.819264Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Cleaning working folder\n\nlistdir = os.listdir(\"/kaggle/working\")\nfor pth in listdir:\n    if pth != \"submission.csv\":\n        if os.path.isdir(\"/kaggle/working/\"+pth):\n            os.system(\"rm -r /kaggle/working/\"+pth)\n        else:\n            os.system(\"rm /kaggle/working/\"+pth)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T17:37:34.061083Z","iopub.status.idle":"2025-04-28T17:37:34.061385Z","shell.execute_reply.started":"2025-04-28T17:37:34.061245Z","shell.execute_reply":"2025-04-28T17:37:34.061260Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}