{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# RSNA'23 | II. Predicting Segmentations with VNet\n\nThis notebook trains a VNet model with the training segmentations and predicts the segmentations for the remaining series in the training dataset.\n\nOther notebooks and datasets in this analysis:\n1. [RSNA'23 | I. Series Normalisation](https://www.kaggle.com/code/bvinning/rsna-23-i-series-normalisation) -> [RSNA'23 | Normalised Series](https://www.kaggle.com/datasets/bvinning/rsna-2023-normalised-series)\n2. [RSNA'23 | II. Predicting Segmentations with VNet](https://www.kaggle.com/code/bvinning/rsna-23-ii-predicting-segmentations-with-vnet) -> [RSNA'23 | Predicted Segmentations](https://www.kaggle.com/datasets/bvinning/rsna-2023-predicted-segmentations)\n3. [RSNA'23 | III. Applying Organ Segmentations](https://www.kaggle.com/code/bvinning/rsna-23-iii-applying-organ-segmentations) -> [RSNA'23 | Segmented Series](https://www.kaggle.com/datasets/bvinning/rsna-2023-segmented-series)\n4. [RSNA'23 | IV. Predicting Injuries with a 3D CNN](https://www.kaggle.com/code/bvinning/rsna-23-iv-predicting-injuries-with-a-3d-cnn) -> [RSNA'23 | Injury Classifier Model](https://www.kaggle.com/datasets/bvinning/rsna-2023-injury-classifier-model)\n5. [RSNA'23 | V. Final Submission](https://www.kaggle.com/code/bvinning/rsna-23-v-final-submission)","metadata":{}},{"cell_type":"code","source":"# Load libraries.\nimport os\nimport glob\nimport torch\n\nimport joblib\nimport pandas as pd\nimport numpy as np\nimport matplotlib.pyplot as plt\n\n\nRANDOM_SEED = 777\nRANDOM_STATE = np.random.seed(RANDOM_SEED)\ntorch.manual_seed(RANDOM_SEED)\nDEVICE = torch.device('cuda' if torch.has_cuda else 'cpu')","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import wandb\nfrom kaggle_secrets import UserSecretsClient\n\nuser_secrets = UserSecretsClient()\nwandb_api = user_secrets.get_secret(\"wandb-api-key\")\n\nwandb.login(key=wandb_api)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import json\n\napi = wandb.Api()\n\nWANDB_PROJECT = 'rsna-2023-vnet-sgd'\n\n\ndef get_best_hyperparameters():\n    best_run = None\n    best_score = np.inf\n    for run in api.runs(path=WANDB_PROJECT):\n        try:\n            score = run.history()['val_loss'].min()\n        except KeyError:\n            continue\n        if score < best_score:\n            best_run = run\n            score = best_score\n\n    return {k: v['value'] for k, v in json.loads(best_run.json_config).items()}\n\n\ndef get_run_hyperparameters(name):\n    for run in api.runs(path=WANDB_PROJECT):\n        if run.name == name:\n            best_run = run\n            break\n    return {k: v['value'] for k, v in json.loads(best_run.json_config).items()}\n\n\nHYPERPARAMETERS = get_run_hyperparameters('dark-cloud-44') \nprint(HYPERPARAMETERS)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import warnings\nimport shutil\n\n\nclass FileSystem:\n    def __init__(self):\n        self.working_dir = '/kaggle/working/'\n        self.cache_dir = os.path.join(self.working_dir, 'cache')\n        self.input_dir = '/kaggle/input/rsna-2023-abdominal-trauma-detection'\n        self.setup()\n    \n    @property\n    def local(self):\n        return os.path.join(self.working_dir, 'local')\n\n    @property\n    def local_series(self):\n        return os.path.join(self.local, 'series')\n\n    @property\n    def local_segmentations(self):\n        return os.path.join(self.local, 'segmentations')\n\n    @property\n    def local_segmented_series(self):\n        return os.path.join(self.local, 'segmented_series')\n\n    @property\n    def cached_series(self):\n        return os.path.join(self.cache_dir, 'series')\n\n    @property\n    def cached_segmentations(self):\n        return os.path.join(self.cache_dir, 'segmentations')\n\n    def setup(self):\n        self.link_transformed_train()\n        self.link_series_normalisation_config()\n\n    def teardown(self):\n        shutil.rmtree(self.cache_dir)\n\n    def link_series_normalisation_config(self):\n        os.makedirs(self.cache_dir, exist_ok=True)\n        src = '/kaggle/input/rsna-2023-normalised-series/config.yaml'\n        dst = os.path.join(self.cache_dir, 'series_normalisation_config.yaml')\n        os.symlink(src, dst)\n        self.series_normalisation_config = dst\n\n    def link_transformed_train(self):\n        # This includes training series and the training segmentations.\n        root_dir = '/kaggle/input/rsna-2023-normalised-series'\n        if not os.path.exists(root_dir):\n            warnings.warn(f'External data source `{root_dir}` does not exist, skipping symlinks.')\n            return\n\n        src_root_dir = self.cached_series\n        os.makedirs(src_root_dir, exist_ok=True)\n        for src in glob.glob(os.path.join(root_dir, 'series', '*.npz')):\n            dst = os.path.join(src_root_dir, os.path.basename(src))\n            os.symlink(src, dst)\n\n        src_root_dir = self.cached_segmentations\n        os.makedirs(src_root_dir, exist_ok=True)\n        for src in glob.glob(os.path.join(root_dir, 'segmentations', '*.npz')):\n            dst = os.path.join(src_root_dir, os.path.basename(src))\n            os.symlink(src, dst)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import functools\n\n\nclass SeriesInfo:\n    def __init__(self, fs):\n        self.fs = fs\n\n    @functools.cached_property\n    def train_segmentations(self):\n        return self.get_series_ids(root_dir=f'{FILESYSTEM.input_dir}/segmentations')\n\n    @functools.cached_property\n    def train_series(self):\n        paths = glob.glob('*/*', root_dir=os.path.join(self.fs.input_dir, 'train_images'))\n        return {int(os.path.split(path)[-1]) for path in paths}\n    \n    @functools.cached_property\n    def test_series(self):\n        paths = glob.glob('*/*', root_dir=os.path.join(self.fs.input_dir, 'test_images'))\n        return {int(os.path.split(path)[-1]) for path in paths}\n\n    @functools.cached_property\n    def series(self):\n        return self.train_series + self.test_series\n\n    @functools.cached_property\n    def cached_series(self):\n        return self.get_series_ids(root_dir=self.fs.cached_series)\n    \n    @functools.cached_property\n    def cached_segmentations(self):\n        return self.get_series_ids(root_dir=self.fs.cached_segmentations)\n\n    @property\n    def local_series(self):\n        return self.get_series_ids(root_dir=self.fs.local_series)\n\n    @property\n    def local_segmentations(self):\n        return self.get_series_ids(root_dir=self.fs.local_segmentations)\n\n    @property\n    def local_segmented_series(self):\n        return self.get_series_ids(root_dir=self.fs.local_segmented_series)\n    \n    @staticmethod\n    def get_series_ids(root_dir):\n        return {int(os.path.splitext(path)[0]) for path in glob.glob('*', root_dir=root_dir)}","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import yaml\nimport dataclasses\nimport functools\n\n\n@dataclasses.dataclass(frozen=True)\nclass SeriesNormalisationConfig:\n    _series_dtype_out: str = 'uint16'\n    _segmentation_dtype_out: str = 'uint8'\n    shape: tuple[int] = dataclasses.field(default=(128, 128, 128)) # the shape after resampling\n    roi: tuple[int] = dataclasses.field(default=(350, 300, 400)) # region of interest in mm\n\n    @classmethod\n    def load(cls, path):\n        with open(path, 'r') as f:\n            data = yaml.safe_load(f)\n        return cls(**data)\n\n    @functools.cached_property\n    def series_dtype_in(self):\n        return np.dtype(self._series_dtype_in).type\n    \n    @functools.cached_property\n    def series_dtype_out(self):\n        return np.dtype(self._series_dtype_out).type\n    \n    @functools.cached_property\n    def segmentation_dtype_out(self):\n        return np.dtype(self._segmentation_dtype_out).type\n\n    @functools.cached_property\n    def series_dtype_out_max(self):\n        return np.iinfo(self.series_dtype_out).max\n\n    def dump(self):\n        data = dataclasses.asdict(self)\n        with open('config.yaml', 'w') as f:\n            yaml.safe_dump(data, f)\n","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import yaml\nimport dataclasses\n\n\n@dataclasses.dataclass\nclass PredictedSegmentationsConfig:\n    series_mean: float\n    series_std: float\n        \n    @classmethod\n    def load(cls, path):\n        with open(path, 'r') as f:\n            return yaml.safe_load(f)\n    \n    def dump(self):\n        data = dataclasses.asdict(self)\n        with open('config.yaml', 'w') as f:\n            yaml.safe_dump(data, f)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"FILESYSTEM = FileSystem()\nSERIES_INFO = SeriesInfo(FILESYSTEM)\nSERIES_NORMALISATION_CONFIG = SeriesNormalisationConfig.load(FILESYSTEM.series_normalisation_config)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"CACHED_SERIES = SERIES_INFO.cached_series\nCACHED_SEGMENTATIONS = SERIES_INFO.cached_segmentations\n\n\ndef load_transformed(series_id, return_segs=True):\n    if series_id in SERIES_INFO.local_series:\n        data_dir = FILESYSTEM.local_series\n    elif series_id in CACHED_SERIES:\n        data_dir = FILESYSTEM.cached_series\n    else:\n        raise ValueError(f'Normalised series for `{series_id}` not found.')\n    \n    img_path = os.path.join(data_dir, f'{series_id}.npz')\n    with open(img_path, 'rb') as f:\n        img_t = np.load(f)\n\n    if not return_segs:\n        return img_t\n\n    if series_id in SERIES_INFO.local_segmentations:\n        data_dir = FILESYSTEM.local_segmentations\n    elif series_id in CACHED_SEGMENTATIONS:\n        data_dir = FILESYSTEM.cached_segmentations\n    else:\n        raise ValueError(f'Normalised segmentations for `{series_id}` not found.')\n\n    img_segs_path = os.path.join(data_dir, f'{series_id}.npz')\n    with open(img_segs_path, 'rb') as f:\n        img_segs_t = np.load(f)\n\n    return img_t, img_segs_t","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from enum import Enum\n\n\nclass SegmentationClasses(Enum):\n    OTHER = 0\n    LIVER = 1\n    SPLEEN = 2\n    KIDNEY_LEFT = 3\n    KIDNEY_RIGHT = 4\n    BOWEL = 5","metadata":{"_kg_hide-input":true,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport warnings\nfrom torch.nn.modules.loss import _Loss\nfrom collections.abc import Callable, Sequence, Collection, Hashable, Mapping\nfrom typing import Any\nfrom enum import Enum, EnumMeta\n\n\ndef look_up_option(\n    opt_str: Hashable,\n    supported: Collection | EnumMeta,\n    default: Any = \"no_default\",\n    print_all_options: bool = True,\n) -> Any:\n    \"\"\"\n    Look up the option in the supported collection and return the matched item.\n    Raise a value error possibly with a guess of the closest match.\n\n    Args:\n        opt_str: The option string or Enum to look up.\n        supported: The collection of supported options, it can be list, tuple, set, dict, or Enum.\n        default: If it is given, this method will return `default` when `opt_str` is not found,\n            instead of raising a `ValueError`. Otherwise, it defaults to `\"no_default\"`,\n            so that the method may raise a `ValueError`.\n        print_all_options: whether to print all available options when `opt_str` is not found. Defaults to True\n\n    Examples:\n\n    .. code-block:: python\n\n        from enum import Enum\n        from monai.utils import look_up_option\n        class Color(Enum):\n            RED = \"red\"\n            BLUE = \"blue\"\n        look_up_option(\"red\", Color)  # <Color.RED: 'red'>\n        look_up_option(Color.RED, Color)  # <Color.RED: 'red'>\n        look_up_option(\"read\", Color)\n        # ValueError: By 'read', did you mean 'red'?\n        # 'read' is not a valid option.\n        # Available options are {'blue', 'red'}.\n        look_up_option(\"red\", {\"red\", \"blue\"})  # \"red\"\n\n    Adapted from https://github.com/NifTK/NiftyNet/blob/v0.6.0/niftynet/utilities/util_common.py#L249\n    \"\"\"\n    if not isinstance(opt_str, Hashable):\n        raise ValueError(f\"Unrecognized option type: {type(opt_str)}:{opt_str}.\")\n    if isinstance(opt_str, str):\n        opt_str = opt_str.strip()\n    if isinstance(supported, EnumMeta):\n        if isinstance(opt_str, str) and opt_str in {item.value for item in supported}:  # type: ignore\n            # such as: \"example\" in MyEnum\n            return supported(opt_str)\n        if isinstance(opt_str, Enum) and opt_str in supported:\n            # such as: MyEnum.EXAMPLE in MyEnum\n            return opt_str\n    elif isinstance(supported, Mapping) and opt_str in supported:\n        # such as: MyDict[key]\n        return supported[opt_str]\n    elif isinstance(supported, Collection) and opt_str in supported:\n        return opt_str\n\n    if default != \"no_default\":\n        return default\n\n    # find a close match\n    set_to_check: set\n    if isinstance(supported, EnumMeta):\n        set_to_check = {item.value for item in supported}  # type: ignore\n    else:\n        set_to_check = set(supported) if supported is not None else set()\n    if not set_to_check:\n        raise ValueError(f\"No options available: {supported}.\")\n    edit_dists = {}\n    opt_str = f\"{opt_str}\"\n    for key in set_to_check:\n        edit_dist = damerau_levenshtein_distance(f\"{key}\", opt_str)\n        if edit_dist <= 3:\n            edit_dists[key] = edit_dist\n\n    supported_msg = f\"Available options are {set_to_check}.\\n\" if print_all_options else \"\"\n    if edit_dists:\n        guess_at_spelling = min(edit_dists, key=edit_dists.get)  # type: ignore\n        raise ValueError(\n            f\"By '{opt_str}', did you mean '{guess_at_spelling}'?\\n\"\n            + f\"'{opt_str}' is not a valid value.\\n\"\n            + supported_msg\n        )\n    raise ValueError(f\"Unsupported option '{opt_str}', \" + supported_msg)\n\n\ndef damerau_levenshtein_distance(s1: str, s2: str) -> int:\n    \"\"\"\n    Calculates the Damerau–Levenshtein distance between two strings for spelling correction.\n    https://en.wikipedia.org/wiki/Damerau–Levenshtein_distance\n    \"\"\"\n    if s1 == s2:\n        return 0\n    string_1_length = len(s1)\n    string_2_length = len(s2)\n    if not s1:\n        return string_2_length\n    if not s2:\n        return string_1_length\n    d = {(i, -1): i + 1 for i in range(-1, string_1_length + 1)}\n    for j in range(-1, string_2_length + 1):\n        d[(-1, j)] = j + 1\n\n    for i, s1i in enumerate(s1):\n        for j, s2j in enumerate(s2):\n            cost = 0 if s1i == s2j else 1\n            d[(i, j)] = min(\n                d[(i - 1, j)] + 1, d[(i, j - 1)] + 1, d[(i - 1, j - 1)] + cost  # deletion  # insertion  # substitution\n            )\n            if i and j and s1i == s2[j - 1] and s1[i - 1] == s2j:\n                d[(i, j)] = min(d[(i, j)], d[i - 2, j - 2] + cost)  # transposition\n\n    return d[string_1_length - 1, string_2_length - 1]\n\n\nclass StrEnum(str, Enum):\n    \"\"\"\n    Enum subclass that converts its value to a string.\n\n    .. code-block:: python\n\n        from monai.utils import StrEnum\n\n        class Example(StrEnum):\n            MODE_A = \"A\"\n            MODE_B = \"B\"\n\n        assert (list(Example) == [\"A\", \"B\"])\n        assert Example.MODE_A == \"A\"\n        assert str(Example.MODE_A) == \"A\"\n        assert monai.utils.look_up_option(\"A\", Example) == \"A\"\n    \"\"\"\n\n    def __str__(self):\n        return self.value\n\n    def __repr__(self):\n        return self.value\n    \n\nclass Weight(StrEnum):\n    \"\"\"\n    See also: :py:class:`monai.losses.dice.GeneralizedDiceLoss`\n    \"\"\"\n\n    SQUARE = \"square\"\n    SIMPLE = \"simple\"\n    UNIFORM = \"uniform\"\n\n\nclass LossReduction(StrEnum):\n    \"\"\"\n    See also:\n        - :py:class:`monai.losses.dice.DiceLoss`\n        - :py:class:`monai.losses.dice.GeneralizedDiceLoss`\n        - :py:class:`monai.losses.focal_loss.FocalLoss`\n        - :py:class:`monai.losses.tversky.TverskyLoss`\n    \"\"\"\n\n    NONE = \"none\"\n    MEAN = \"mean\"\n    SUM = \"sum\"\n    \n\ndef one_hot(labels: torch.Tensor, num_classes: int, dtype: torch.dtype = torch.float, dim: int = 1) -> torch.Tensor:\n    \"\"\"\n    For every value v in `labels`, the value in the output will be either 1 or 0. Each vector along the `dim`-th\n    dimension has the \"one-hot\" format, i.e., it has a total length of `num_classes`,\n    with a one and `num_class-1` zeros.\n    Note that this will include the background label, thus a binary mask should be treated as having two classes.\n\n    Args:\n        labels: input tensor of integers to be converted into the 'one-hot' format. Internally `labels` will be\n            converted into integers `labels.long()`.\n        num_classes: number of output channels, the corresponding length of `labels[dim]` will be converted to\n            `num_classes` from `1`.\n        dtype: the data type of the output one_hot label.\n        dim: the dimension to be converted to `num_classes` channels from `1` channel, should be non-negative number.\n\n    Example:\n\n    For a tensor `labels` of dimensions [B]1[spatial_dims], return a tensor of dimensions `[B]N[spatial_dims]`\n    when `num_classes=N` number of classes and `dim=1`.\n\n    .. code-block:: python\n\n        from monai.networks.utils import one_hot\n        import torch\n\n        a = torch.randint(0, 2, size=(1, 2, 2, 2))\n        out = one_hot(a, num_classes=2, dim=0)\n        print(out.shape)  # torch.Size([2, 2, 2, 2])\n\n        a = torch.randint(0, 2, size=(2, 1, 2, 2, 2))\n        out = one_hot(a, num_classes=2, dim=1)\n        print(out.shape)  # torch.Size([2, 2, 2, 2, 2])\n\n    \"\"\"\n\n    # if `dim` is bigger, add singleton dim at the end\n    if labels.ndim < dim + 1:\n        shape = list(labels.shape) + [1] * (dim + 1 - len(labels.shape))\n        labels = torch.reshape(labels, shape)\n\n    sh = list(labels.shape)\n\n    if sh[dim] != 1:\n        raise AssertionError(\"labels should have a channel with length equal to one.\")\n\n    sh[dim] = num_classes\n\n    o = torch.zeros(size=sh, dtype=dtype, device=labels.device)\n    labels = o.scatter_(dim=dim, index=labels.long(), value=1)\n\n    return labels\n\n\nclass GeneralizedDiceLoss(_Loss):\n    \"\"\"\n    Compute the generalised Dice loss defined in:\n\n        Sudre, C. et. al. (2017) Generalised Dice overlap as a deep learning\n        loss function for highly unbalanced segmentations. DLMIA 2017.\n\n    Adapted from:\n        https://github.com/NifTK/NiftyNet/blob/v0.6.0/niftynet/layer/loss_segmentation.py#L279\n    \"\"\"\n\n    def __init__(\n        self,\n        include_background: bool = True,\n        to_onehot_y: bool = False,\n        sigmoid: bool = False,\n        softmax: bool = False,\n        other_act: Callable | None = None,\n        w_type: Weight | str = Weight.SQUARE,\n        reduction: LossReduction | str = LossReduction.MEAN,\n        smooth_nr: float = 1e-5,\n        smooth_dr: float = 1e-5,\n        batch: bool = False,\n    ) -> None:\n        \"\"\"\n        Args:\n            include_background: If False channel index 0 (background category) is excluded from the calculation.\n            to_onehot_y: whether to convert the ``target`` into the one-hot format,\n                using the number of classes inferred from `input` (``input.shape[1]``). Defaults to False.\n            sigmoid: If True, apply a sigmoid function to the prediction.\n            softmax: If True, apply a softmax function to the prediction.\n            other_act: callable function to execute other activation layers, Defaults to ``None``. for example:\n                ``other_act = torch.tanh``.\n            w_type: {``\"square\"``, ``\"simple\"``, ``\"uniform\"``}\n                Type of function to transform ground truth volume to a weight factor. Defaults to ``\"square\"``.\n            reduction: {``\"none\"``, ``\"mean\"``, ``\"sum\"``}\n                Specifies the reduction to apply to the output. Defaults to ``\"mean\"``.\n\n                - ``\"none\"``: no reduction will be applied.\n                - ``\"mean\"``: the sum of the output will be divided by the number of elements in the output.\n                - ``\"sum\"``: the output will be summed.\n            smooth_nr: a small constant added to the numerator to avoid zero.\n            smooth_dr: a small constant added to the denominator to avoid nan.\n            batch: whether to sum the intersection and union areas over the batch dimension before the dividing.\n                Defaults to False, intersection over union is computed from each item in the batch.\n                If True, the class-weighted intersection and union areas are first summed across the batches.\n\n        Raises:\n            TypeError: When ``other_act`` is not an ``Optional[Callable]``.\n            ValueError: When more than 1 of [``sigmoid=True``, ``softmax=True``, ``other_act is not None``].\n                Incompatible values.\n\n        \"\"\"\n        super().__init__(reduction=LossReduction(reduction).value)\n        if other_act is not None and not callable(other_act):\n            raise TypeError(f\"other_act must be None or callable but is {type(other_act).__name__}.\")\n        if int(sigmoid) + int(softmax) + int(other_act is not None) > 1:\n            raise ValueError(\"Incompatible values: more than 1 of [sigmoid=True, softmax=True, other_act is not None].\")\n\n        self.include_background = include_background\n        self.to_onehot_y = to_onehot_y\n        self.sigmoid = sigmoid\n        self.softmax = softmax\n        self.other_act = other_act\n\n        self.w_type = look_up_option(w_type, Weight)\n\n        self.smooth_nr = float(smooth_nr)\n        self.smooth_dr = float(smooth_dr)\n        self.batch = batch\n\n    def w_func(self, grnd):\n        if self.w_type == str(Weight.SIMPLE):\n            return torch.reciprocal(grnd)\n        if self.w_type == str(Weight.SQUARE):\n            return torch.reciprocal(grnd * grnd)\n        return torch.ones_like(grnd)\n\n    def forward(self, input: torch.Tensor, target: torch.Tensor) -> torch.Tensor:\n        \"\"\"\n        Args:\n            input: the shape should be BNH[WD].\n            target: the shape should be BNH[WD].\n\n        Raises:\n            ValueError: When ``self.reduction`` is not one of [\"mean\", \"sum\", \"none\"].\n\n        \"\"\"\n        if self.sigmoid:\n            input = torch.sigmoid(input)\n        n_pred_ch = input.shape[1]\n        if self.softmax:\n            if n_pred_ch == 1:\n                warnings.warn(\"single channel prediction, `softmax=True` ignored.\")\n            else:\n                input = torch.softmax(input, 1)\n\n        if self.other_act is not None:\n            input = self.other_act(input)\n\n        if self.to_onehot_y:\n            if n_pred_ch == 1:\n                warnings.warn(\"single channel prediction, `to_onehot_y=True` ignored.\")\n            else:\n                target = one_hot(target, num_classes=n_pred_ch)\n\n        if not self.include_background:\n            if n_pred_ch == 1:\n                warnings.warn(\"single channel prediction, `include_background=False` ignored.\")\n            else:\n                # if skipping background, removing first channel\n                target = target[:, 1:]\n                input = input[:, 1:]\n\n        if target.shape != input.shape:\n            raise AssertionError(f\"ground truth has differing shape ({target.shape}) from input ({input.shape})\")\n\n        # reducing only spatial dimensions (not batch nor channels)\n        reduce_axis: list[int] = torch.arange(2, len(input.shape)).tolist()\n        if self.batch:\n            reduce_axis = [0] + reduce_axis\n        intersection = torch.sum(target * input, reduce_axis)\n\n        ground_o = torch.sum(target, reduce_axis)\n        pred_o = torch.sum(input, reduce_axis)\n\n        denominator = ground_o + pred_o\n\n        w = self.w_func(ground_o.float())\n        infs = torch.isinf(w)\n        if self.batch:\n            w[infs] = 0.0\n            w = w + infs * torch.max(w)\n        else:\n            w[infs] = 0.0\n            max_values = torch.max(w, dim=1)[0].unsqueeze(dim=1)\n            w = w + infs * max_values\n\n        final_reduce_dim = 0 if self.batch else 1\n        numer = 2.0 * (intersection * w).sum(final_reduce_dim, keepdim=True) + self.smooth_nr\n        denom = (denominator * w).sum(final_reduce_dim, keepdim=True) + self.smooth_dr\n        f: torch.Tensor = 1.0 - (numer / denom)\n\n        if self.reduction == LossReduction.MEAN.value:\n            f = torch.mean(f)  # the batch and channel average\n        elif self.reduction == LossReduction.SUM.value:\n            f = torch.sum(f)  # sum over the batch and channel dims\n        elif self.reduction == LossReduction.NONE.value:\n            # If we are not computing voxelwise loss components at least\n            # make sure a none reduction maintains a broadcastable shape\n            broadcast_shape = list(f.shape[0:2]) + [1] * (len(input.shape) - 2)\n            f = f.view(broadcast_shape)\n        else:\n            raise ValueError(f'Unsupported reduction: {self.reduction}, available options are [\"mean\", \"sum\", \"none\"].')\n\n        return f","metadata":{"jupyter":{"source_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train = pd.read_csv('/kaggle/input/rsna-2023-abdominal-trauma-detection/train.csv')\nmeta = pd.read_csv('/kaggle/input/rsna-2023-abdominal-trauma-detection/train_series_meta.csv')","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"assert train.patient_id.nunique() == len(train)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train = train.merge(meta.groupby('patient_id').incomplete_organ.any(), on='patient_id')\nassert train.patient_id.nunique() == len(train)\ntrain = train.merge((meta.groupby('patient_id').size() > 1).reset_index(name='is_multiphasic'), on='patient_id')\nassert train.patient_id.nunique() == len(train)\ntrain['stratify'] = np.ones(len(train), dtype=int)\ni = 0\nfor left in [True, False]:\n    for right in [True, False]:\n        mask = (train.is_multiphasic == left) & (train.incomplete_organ == right) \n        train.loc[mask, 'stratify'] = i\n        i += 1\nmeta['has_segmentations'] = meta.series_id.isin(SERIES_INFO.train_segmentations)\nassert train.patient_id.nunique() == len(train)\ntrain = train.merge(meta.groupby('patient_id').has_segmentations.any(), on='patient_id')\nassert train.patient_id.nunique() == len(train)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_with_segmentations = train[train.has_segmentations]\nmeta_with_segmentations = meta[meta.has_segmentations]","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from sklearn.model_selection import train_test_split\n\ntrain_patient_ids, valid_patient_ids, train_stratify, valid_stratify = train_test_split(\n    train_with_segmentations.patient_id.values,\n    train_with_segmentations.stratify.values,\n    stratify=train_with_segmentations.stratify.values,\n    random_state=RANDOM_STATE,\n    test_size=0.2,\n)\ntrain_series_ids = meta_with_segmentations.series_id[meta_with_segmentations.patient_id.isin(train_patient_ids)].values\nvalid_series_ids = meta_with_segmentations.series_id[meta_with_segmentations.patient_id.isin(valid_patient_ids)].values\n# train_stratify = meta.stratify[meta.patient_id.isin(train_patient_ids)].values\n# valid_stratify = meta.stratify[meta.patient_id.isin(valid_patient_ids)].values","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import gc\ndel train, meta, train_with_segmentations, meta_with_segmentations\ngc.collect()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"assert len(np.unique(train_patient_ids)) == len(train_patient_ids), (len(train_patient_ids), len(np.unique(train_patient_ids)))","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"TRANSLATION_SCALING_REL = 0.15\nTRANSLATION_SCALING_ABS = (np.array(SERIES_NORMALISATION_CONFIG.shape) * TRANSLATION_SCALING_REL).astype(int)\nINTENSITY_SCALING_REL = 0.25\n\n\ndef suggest_augmentation_parameters():\n    return {\n        'rotation_axis': np.random.randint(0, 3),\n        'rotation_n': np.random.randint(0, 4),\n        'should_flip': bool(np.random.randint(0, 2)),\n        'flip_axis': np.random.randint(0, 3),\n        'translation_axis': np.random.randint(0, 3),\n        'translation_direction': np.random.randint(-1, 2),\n        'scale_direction': np.random.randint(-1, 2),\n    }\n\n\ndef augment_img(\n    img,\n    rotation_axis,\n    rotation_n,\n    flip_axis,\n    should_flip,\n    translation_axis,\n    translation_direction,\n    scale_direction,\n    spatial_only=False\n):\n    spatial_axes = list(range(img.ndim))[img.ndim - 3:]\n\n    if rotation_n != 0:\n        dims = tuple(ax for i, ax in enumerate(spatial_axes) if i != rotation_axis)\n        img = img.rot90(rotation_n, dims)\n\n    if should_flip:\n        img = img.flip(spatial_axes[flip_axis])\n\n    if translation_direction != 0:\n        img = img.roll(translation_direction * TRANSLATION_SCALING_ABS[translation_axis], spatial_axes[translation_axis])\n        indexing = [slice(None)] * img.ndim\n\n        roll_axis = spatial_axes[translation_axis]\n        if translation_direction > 0:\n            indexing[roll_axis] = slice(0, TRANSLATION_SCALING_ABS[translation_axis])\n\n        else:\n            indexing[roll_axis] = slice(img.shape[translation_axis] - TRANSLATION_SCALING_ABS[translation_axis], None)\n        img[tuple(indexing)] = img.min()\n\n    if scale_direction != 0 and not spatial_only:\n        img *= (1 + (scale_direction * INTENSITY_SCALING_REL))\n\n    return img","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\n\nfrom sklearn.model_selection import train_test_split\nimport torch\nfrom torch.utils.data import Dataset\nfrom torchvision.transforms import Compose, Normalize, Lambda\nimport torch.nn.functional as F\n\n\nclass SegmentationsDataset(Dataset):\n    def __init__(self, series_ids, train, augment_train=False, normalisation=None, one_hot_target=True):\n        self.series_ids = series_ids\n        self.train = train\n        self.augment_train = augment_train\n        \n        if SERIES_NORMALISATION_CONFIG.series_dtype_out is np.uint8:\n            load_numpy_fn = lambda x: torch.as_tensor(x, dtype=torch.float32)\n        else:\n            load_numpy_fn = lambda x: torch.from_numpy(x.astype(np.float32))\n        transform = [\n            Lambda(load_numpy_fn)\n        ]\n        if normalisation:\n            transform += [Normalize(*normalisation, inplace=True), Lambda(lambda x: x.unsqueeze_(0))]\n\n        self.transform = Compose(transform)\n        \n        target_transform = [\n            Lambda(lambda x: torch.as_tensor(x, dtype=torch.long))\n        ]\n        if one_hot_target:\n            target_transform.append(\n                Lambda(lambda x: F.one_hot(x, len(SegmentationClasses)).permute((3, 0, 1, 2)))\n            )\n        self.target_transform = Compose(target_transform)\n        \n    def __len__(self):\n        return len(self.series_ids)\n\n    def __getitem__(self, idx):\n        series_id = self.series_ids[idx]\n        if self.train:\n            img, img_segs = load_transformed(series_id)\n        else:\n            img = load_transformed(series_id, return_segs=False)\n\n        if self.transform:\n            img = self.transform(img)\n\n        if not self.train:\n            return series_id, img\n\n        if self.target_transform:\n            img_segs = self.target_transform(img_segs)\n        if self.augment_train:\n            augmentation_kwargs = suggest_augmentation_parameters()\n            img = augment_img(img, spatial_only=False, **augmentation_kwargs)\n            img_segs = augment_img(img_segs, spatial_only=True, **augmentation_kwargs)\n        return series_id, img, img_segs\n\n\ndef get_mean_and_std(dataset):\n    imgs = torch.stack([img_t for _, img_t, _ in dataset], dim=3)\n    return imgs.view(1, -1).mean(dim=1), imgs.view(1, -1).std(dim=1)\n\n\n# Train: Series IDs in the training dataset that have segmentations but are used for training.\n# Valid: Series IDs in the training dataset that have segmentations but are put aside for testing.   \n# train_series_ids = list(SERIES_INFO.train_segmentations)\n# train_series_ids, valid_series_ids = train_test_split(train_series_ids, test_size=0.2, random_state=RANDOM_STATE)\n# Series IDs in the training dataset that don't have segmentations. \ntest_series_ids = list(SERIES_INFO.train_series)\ntrain_dataset = SegmentationsDataset(train_series_ids, train=True, augment_train=False)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"normalisation = get_mean_and_std(train_dataset)\nif HYPERPARAMETERS['augment_train']:\n    train_series_ids = train_series_ids * HYPERPARAMETERS['train_duplication_factor']\n    train_dataset = SegmentationsDataset(train_series_ids, train=True, augment_train=True, normalisation=normalisation)\nelse:\n    train_dataset = SegmentationsDataset(train_series_ids, train=True, augment_train=False, normalisation=normalisation)\nvalid_dataset = SegmentationsDataset(valid_series_ids, train=True, augment_train=False, normalisation=normalisation)\ntest_dataset = SegmentationsDataset(test_series_ids, train=False, normalisation=normalisation)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from matplotlib.patches import Patch\nfrom matplotlib.lines import Line2D\n\nSERIES_CMAP = plt.get_cmap('bone')\nSEGMENTATION_CMAP = plt.get_cmap('rainbow')\nSEGMENTATION_PALETTE = [SEGMENTATION_CMAP(x) for x in np.linspace(0, 1, len(SegmentationClasses) - 1)]\nBACKGROUND_COLOUR = [np.nan] * 4\n\n# Pick 10 random series with/without segmentations.\ndef show_series(img, img_segs=None, title='', width=10):\n    \n    fig, axes = plt.subplots(ncols=img.ndim)\n    centroids = np.true_divide(img.shape, 2).astype(int)\n    centroids = [slice((x - width) // 2, (x + width) // 2) for x in img.shape]\n    for i, ax in enumerate(axes):\n        indexing = [slice(None)] * img.ndim\n        indexing[i] = centroids[i]\n        indexing = tuple(indexing)\n        img_slice = img[indexing].sum(axis=i).T\n        ax.imshow(img_slice, cmap=SERIES_CMAP)\n\n        if img_segs is not None:\n            img_seg_slice = img_segs[indexing]\n            for y_class in SegmentationClasses:\n                if y_class is SegmentationClasses.OTHER:\n                    continue\n                palette = np.array([BACKGROUND_COLOUR, SEGMENTATION_PALETTE[y_class.value - 1]])\n                img_seg_slice_class = (img_seg_slice == y_class.value).any(axis=i).astype(int).T\n                img_seg_slice_coloured = palette[img_seg_slice_class]\n                ax.imshow(img_seg_slice_coloured, alpha=0.3)\n\n        ax.axis('off')\n        ax.set_title(f\"{['YZ', 'XZ', 'XY'][i]}\")\n\n    if img_segs is not None:\n        mask_levels = [x for x in np.unique(img_segs) if x != 0]\n        legend_elements = [\n            Patch(facecolor=SEGMENTATION_PALETTE[x - 1], edgecolor=SEGMENTATION_PALETTE[x - 1], label=SegmentationClasses(x).name.lower()) for x in mask_levels\n        ]\n        ax.legend(handles=legend_elements, loc='lower center', bbox_to_anchor=(0.5, 1.05))\n\n    fig.suptitle(title)\n    fig.tight_layout()\n    return fig\n\n\nTRUTH_COLOUR = [1.0, 0.0, 0.0, 1.0]\nPREDICTION_COLOUR = [0.0, 1.0, 0.16, 1.0]\nTRUTH_PALETTE = np.array([BACKGROUND_COLOUR, TRUTH_COLOUR])\nPREDICTION_PALETTE = np.array([BACKGROUND_COLOUR, PREDICTION_COLOUR])\n\n\ndef show_segmentations(series_id, x, y_pred, y, segmentation_class, epoch):\n    fig, axes = plt.subplots(ncols=y_pred.ndim)\n    titles = ['YZ', 'XZ', 'XY']\n    for i, (ax, title) in enumerate(zip(axes, titles)):\n        y_projection = y.any(axis=i).T\n        y_pred_projection = y_pred.any(axis=i).T\n        ax.imshow(x.sum(axis=i).T, cmap=SERIES_CMAP)\n        ax.imshow(TRUTH_PALETTE[y_projection.astype(int)], alpha=0.4)\n        ax.imshow(PREDICTION_PALETTE[y_pred_projection.astype(int)], alpha=0.4)\n\n        ax.axis('off')\n        ax.set_title(title)\n\n    legend_elements = [\n        Patch(\n            facecolor=TRUTH_COLOUR,\n            edgecolor=TRUTH_COLOUR,\n            label='Truth',\n        ),\n        Patch(\n            facecolor=PREDICTION_COLOUR,\n            edgecolor=PREDICTION_COLOUR,\n            label='Pred.',\n        ),\n    ]\n    ax.legend(handles=legend_elements, loc='lower center', bbox_to_anchor=(0.5, 1.05))\n    fig.suptitle(f'Epoch: {epoch}, Series ID: {series_id}, Seg. Class: {segmentation_class}')\n    fig.tight_layout()\n    return fig","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torch.utils.data import DataLoader\nfrom multiprocessing import cpu_count\n\n\ndataloader_kwargs = {\n    'batch_size': HYPERPARAMETERS['batch_size'],\n    'num_workers': cpu_count(),\n    'pin_memory': True,\n}\n\ntrain_dataloader = DataLoader(train_dataset, shuffle=True, **dataloader_kwargs)\nvalid_dataloader = DataLoader(valid_dataset, **dataloader_kwargs)\ntest_dataloader = DataLoader(test_dataset, **dataloader_kwargs)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from typing import Optional\n\nimport torch\nfrom torch import nn\nimport torch.nn.functional as F\n\n    \ndef passthrough(x, **kwargs):\n    return x\n\n\ndef get_logit_fn(name):\n    if name == 'dice':\n        out = F.softmax\n    elif name == 'nll':\n        out = F.log_softmax\n    elif name is None or name == 'none':\n        out = passthrough\n    else:\n        valid_names = ('dice', 'nll', 'none')\n        raise ValueError(f'Loss must be one of ({\", \".join(f\"`{valid_name}`\" for valid_name in valid_names)}, None), `{name}` was passed.')\n    return out\n\n\ndef get_activation(name, n_channels):\n    if name == 'elu':\n        out = nn.ELU(inplace=True)\n    elif name == 'relu':\n        out = nn.ReLU(inplace=True)\n    elif name == 'prelu':\n        out = nn.PReLU(n_channels)\n    else:\n        valid_names = ('elu', 'relu', 'prelu')\n        raise ValueError(f'Activation must be one of {\", \".join(f\"`{valid_name}`\" for valid_name in valid_names)}, `{name}` was passed.')\n    return out\n\n\nclass VNet(nn.Module):\n    '''\n    A PyTorch implementation of segmentation model described in F. Milletari, N. Navab and S.A. Ahmadi's\n    \"V-Net: Fully Convolutional Neural Networks for Volumetric Medical Image Segmentation\"\n    (https://arxiv.org/pdf/1606.04797.pdf)\n    \n    '''\n    def __init__(\n        self,\n        in_channels: int,\n        out_channels: int,\n        depth: int,\n        wf: int,\n        batch_norm: bool,\n        activation: str,\n        loss: str,\n        kaiming_normal: bool,\n    ):\n        super().__init__()\n        \n        down_blocks = []\n        down_in_channels = in_channels\n        for i in range(depth - 1):\n            if i > 0:\n                down_conv_out_channels = None\n                down_out_channels = 2 * down_in_channels\n            else:\n                down_conv_out_channels = 2 ** wf\n                down_out_channels = 2 * down_conv_out_channels\n                \n            down_block = VNetDownBlock(\n                in_channels=down_in_channels,\n                conv_out_channels=down_conv_out_channels,\n                out_channels=down_out_channels,\n                n_convs=min(i + 1, 3),\n                batch_norm=batch_norm,\n                activation=activation,\n                kaiming_normal=kaiming_normal,\n            )\n            down_blocks.append(down_block)\n            down_in_channels = down_out_channels\n        self.down_blocks = nn.ModuleList(down_blocks)\n\n        up_blocks = []\n        up_in_channels = down_out_channels\n        for i in range(depth - 1):\n            if i > 0:\n                up_out_channels = up_in_channels // 4\n                up_conv_out_channels = up_in_channels // 2\n            else:\n                up_out_channels = up_in_channels // 2\n                up_conv_out_channels = None\n            up_block = VNetUpBlock(\n                in_channels=up_in_channels,\n                conv_out_channels=up_conv_out_channels,\n                out_channels=up_out_channels,\n                n_convs=min(depth - i + 1, 3),\n                batch_norm=batch_norm,\n                activation=activation,\n                kaiming_normal=kaiming_normal,\n            )\n\n            up_in_channels = 2 * up_out_channels\n            up_blocks.append(up_block)\n            \n        self.up_blocks = nn.ModuleList(up_blocks)\n        \n        self.final_conv_block = VNetOutputBlock(\n            in_channels=2 * up_out_channels,\n            conv_out_channels=up_out_channels,\n            out_channels=out_channels,\n            batch_norm=batch_norm,\n            activation=activation,\n            loss=loss,\n            kaiming_normal=kaiming_normal,\n        )\n\n    def forward(self, x: torch.Tensor) -> torch.Tensor:\n        x_bridges = []\n        for i, down_block in enumerate(self.down_blocks):\n            #print(f'Down {i} in: ', x.shape)\n            x = down_block(x)\n            x_bridges.append(x)\n            x = down_block.pool(x)\n            #print(f'Down {i}, out (bridge):', x_bridges[-1].shape, ', out (pooled)', x.shape, '\\n')\n        for i, up_block in enumerate(self.up_blocks):\n            x_bridge = x_bridges.pop(-1) if i > 0 else None\n            #print(f'Up {i} in: ', x.shape, ', bridge in:', x_bridge.shape if x_bridge is not None else None)\n            x = up_block(x, x_bridge)\n            #print(f'Up {i} out: ', x.shape, '\\n')\n        \n        x = self.final_conv_block(x, x_bridges.pop(-1))\n        #print(f'Final out: ', x.shape, '\\n')\n        assert not x_bridges\n        return x\n\n\nclass VNetOutputBlock(nn.Module):\n    def __init__(self, in_channels: int, conv_out_channels: int, out_channels: int, batch_norm: bool, activation: str, loss: str, kaiming_normal: bool):\n        super().__init__()\n        self.conv_block = VNetNConvBlock(\n            in_channels=in_channels,\n            out_channels=conv_out_channels,\n            n_convs=1,\n            batch_norm=batch_norm,\n            activation=activation,\n            kaiming_normal=kaiming_normal,\n        )\n        conv = nn.Conv3d(\n            in_channels=conv_out_channels,\n            out_channels=out_channels,\n            kernel_size=1,\n        )\n        if kaiming_normal:\n            nn.init.kaiming_normal_(conv.weight)\n        self.final_conv = nn.Sequential(conv, get_activation(activation, out_channels))\n        self.logit_fn = get_logit_fn(loss)\n\n    def forward(self, x: torch.Tensor, x_bridge: torch.Tensor) -> torch.Tensor:\n        assert x.shape == x_bridge.shape, (x.shape, x_bridge.shape)\n        x_cat = torch.cat([x, x_bridge], dim=1)\n        x = self.conv_block(x_cat, x)\n        x = self.final_conv(x)\n        return self.logit_fn(x, dim=1)\n\n    \nclass VNetDownBlock(nn.Module):\n    def __init__(self, in_channels: int, out_channels: int, n_convs: int, batch_norm: bool, activation: str, kaiming_normal: bool, conv_out_channels: Optional[int] = None):\n        super().__init__()\n        if conv_out_channels is None:\n            conv_out_channels = in_channels\n        self.conv_block = VNetNConvBlock(\n            n_convs=n_convs,\n            in_channels=in_channels,\n            out_channels=conv_out_channels,\n            batch_norm=batch_norm,\n            activation=activation,\n            kaiming_normal=kaiming_normal,\n        )\n        conv = nn.Conv3d(in_channels=conv_out_channels, out_channels=out_channels, kernel_size=2, stride=2, padding=0) \n        if kaiming_normal:\n            nn.init.kaiming_normal_(conv.weight)\n        self.pool = nn.Sequential(conv, get_activation(activation, out_channels))\n    \n    def forward(self, x):\n        x = self.conv_block(x, x)\n        return x\n\n\nclass VNetUpBlock(nn.Module):\n    def __init__(self, in_channels: int, out_channels: int, n_convs: int, batch_norm: bool, activation: str, kaiming_normal: bool, conv_out_channels: Optional[int] = None):\n        super().__init__()\n        if conv_out_channels is None:\n            conv_out_channels = in_channels\n        self.conv_block = VNetNConvBlock(\n            in_channels=in_channels,\n            out_channels=conv_out_channels,\n            n_convs=n_convs,\n            batch_norm=batch_norm,\n            activation=activation,\n            kaiming_normal=kaiming_normal,\n        )\n        conv_t = nn.ConvTranspose3d(\n            in_channels=conv_out_channels,\n            out_channels=out_channels,\n            kernel_size=2,\n            stride=2\n        )\n        if kaiming_normal:\n            nn.init.kaiming_normal_(conv_t.weight)\n        self.up_conv = nn.Sequential(conv_t, get_activation(activation, out_channels))\n\n    def forward(self, x: torch.Tensor, x_bridge: Optional[torch.Tensor] = None):\n        if x_bridge is not None:\n            assert x.shape == x_bridge.shape, (x.shape, x_bridge.shape)\n        x_cat = torch.cat([x, x_bridge], dim=1) if x_bridge is not None else x\n        x = self.conv_block(x_cat, x)\n        return self.up_conv(x)\n\n\nclass VNetNConvBlock(nn.Module):\n    def __init__(self, in_channels: int, n_convs: int, batch_norm: bool, activation: str, kaiming_normal: bool, out_channels: Optional[int] = None):\n        super().__init__()\n        if out_channels is None:\n            out_channels = in_channels\n        \n        layers = []\n        for i in range(n_convs):\n            conv = nn.Conv3d(\n                in_channels=in_channels,\n                out_channels=out_channels,\n                kernel_size=5,\n                padding=2\n            )\n            if kaiming_normal:\n                nn.init.kaiming_normal_(conv.weight)\n            layer = [conv, get_activation(activation, out_channels)]\n            if batch_norm:\n                layer.append(nn.BatchNorm3d(out_channels))\n            in_channels = out_channels\n            layers.extend(layer)\n        self.layers = nn.Sequential(*layers)\n    \n    def forward(self, x: torch.Tensor, x_bridge: torch.Tensor):\n        return self.layers(x) + x_bridge","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport torch\nimport torchmetrics\nfrom torch import nn\nimport torch.nn.functional as F\nimport pytorch_lightning as pl\n\nplot_dir = 'plots'\nos.makedirs(plot_dir, exist_ok=True)\n\n\nclass VNetTrainer(pl.LightningModule):\n    def __init__(\n        self,\n        depth,\n        wf,\n        batch_norm,\n        activation,\n        lr,\n        momentum,\n        dampening,\n        weight_decay,\n        nesterov,\n        lr_step_size,\n        lr_gamma,\n        batch_size,\n        augment_train,\n        train_duplication_factor,\n    ):\n        super().__init__()\n        self.model = VNet(\n            in_channels=1,\n            out_channels=len(SegmentationClasses),\n            depth=depth,\n            wf=depth,\n            batch_norm=batch_norm,\n            activation=activation,\n            kaiming_normal=True,\n            loss='dice'\n        )\n        self.loss_fn = GeneralizedDiceLoss(include_background=True)\n\n        self.lr = lr\n        self.momentum = momentum\n        self.dampening = dampening\n        self.weight_decay = weight_decay\n        self.nesterov = nesterov\n\n        self.lr_step_size = lr_step_size\n        self.lr_gamma = lr_gamma\n\n        self.batch_size = batch_size\n        self.augment_train = augment_train\n        self.train_duplication_factor = train_duplication_factor\n\n        self.init_metrics()\n        self.save_hyperparameters()\n\n    def init_metrics(self):\n        prefix = 'val'\n        multiclass_kwargs = {'task': 'multiclass', 'average': 'weighted', 'num_classes': len(SegmentationClasses)}\n        binary_kwargs = {'task': 'binary', 'average': 'weighted'}\n        self.multiclass_metrics = {}\n        self.binary_metrics = {\n            y_class.name.lower(): {} for y_class in SegmentationClasses if y_class is not SegmentationClasses.OTHER\n        }\n        for metric_cls in [\n            torchmetrics.classification.Accuracy,\n            torchmetrics.classification.Precision,\n            torchmetrics.classification.Recall,\n        ]:\n            self.multiclass_metrics[metric_cls.__name__.lower()] = metric_cls(**multiclass_kwargs).to(DEVICE)\n        \n        for metric_cls in [\n            torchmetrics.classification.Accuracy,\n            torchmetrics.classification.Precision,\n            torchmetrics.classification.Recall,\n            torchmetrics.classification.Dice,\n        ]:\n            for y_class in SegmentationClasses:\n                if y_class is SegmentationClasses.OTHER:\n                    continue\n                \n                k = y_class.name.lower()\n                if k not in self.binary_metrics:\n                    self.binary_metrics[k] = {}\n                if metric_cls is torchmetrics.classification.Dice:\n                    metric = metric_cls()\n                else:\n                    metric = metric_cls(**binary_kwargs)\n                self.binary_metrics[k][metric_cls.__name__.lower()] = metric.to(DEVICE)\n\n    def update_metrics(self, y_proba, y):\n        y_pred = y_proba.argmax(dim=1)\n        y = y.argmax(dim=1)\n\n        for metric in self.multiclass_metrics.values():\n            metric(y_pred, y)\n            \n        for y_class in SegmentationClasses:\n            if y_class is SegmentationClasses.OTHER:\n                continue\n            curr_y_pred = y_pred == y_class.value \n            curr_y = y == y_class.value\n            class_metrics = self.binary_metrics[y_class.name.lower()]\n            for metric in class_metrics.values():\n                metric(curr_y_pred, curr_y)\n\n    def compute_metrics(self):\n        prefix = 'val'\n        out = {}\n        for metric_name, metric in self.multiclass_metrics.items():\n            out[f'{prefix}_{metric_name}'] = metric.compute()\n            metric.reset()\n        for organ, metrics in self.binary_metrics.items():\n            for metric_name, metric in metrics.items():\n                out[f'{prefix}_{organ}_{metric_name}'] = metric.compute()\n                metric.reset()\n        return out\n\n    def log_images(self, series_ids, x, y_proba, y):\n        x = x.squeeze(1).cpu().numpy()\n        y_pred = y_proba.argmax(dim=1).squeeze(1).cpu().numpy() # Pred. as dense class map.\n        y = y.cpu().numpy() # Truth as one-hot.\n    \n        images_by_class = {}\n        for series_id, x_, y_pred_, y_ in zip(series_ids, x, y_pred, y):\n            for y_class in SegmentationClasses:\n                path = os.path.join(plot_dir, f'epoch-{self.current_epoch}-sid-{series_id}-{y_class.name.lower()}.png')\n                with plt.ioff():\n                    fig = show_segmentations(\n                        series_id=series_id,\n                        x=x_,\n                        y_pred=y_pred_ == y_class.value,\n                        y=y_[y_class.value],\n                        segmentation_class=y_class.name.lower(),\n                        epoch=self.current_epoch,\n                    )\n                    fig.savefig(path)\n                plt.close(fig)\n                key = y_class.name.lower()\n                if key not in images_by_class:\n                    images_by_class[key] = []\n                images_by_class[key].append(path)\n                \n        for key, images in images_by_class.items():\n            self.logger.log_image(key=f\"segmentations_{key}\", images=images)\n\n    def training_step(self, batch, batch_idx):\n        _, x, y = batch\n        y_proba = self.model(x)\n        loss = self.loss_fn(y_proba, y)\n        self.log(\"train_loss\", loss, prog_bar=True)\n        return loss\n\n    def validation_step(self, batch, batch_idx):\n        series_id, x, y = batch\n        y_proba = self.model(x)\n        loss = self.loss_fn(y_proba, y)\n        self.log(\"val_loss\", loss, prog_bar=True)\n        self.update_metrics(y_proba, y)\n        if batch_idx == 0:\n            self.log_images(series_id, x, y_proba, y)\n        return loss\n\n    def on_validation_epoch_end(self):\n        metrics = self.compute_metrics()\n        self.log_dict(metrics)\n\n    def predict_step(self, batch, batch_idx, dataloader_idx=0):\n        # Save predictions as we go.\n        ix, x = batch\n        y_probas = self.model(x)\n        # One hot to class enumeration.\n        for series_id, y_proba in zip(ix, y_probas):\n            a = y_proba.cpu().detach().numpy().astype(np.float16)\n            path = os.path.join(FILESYSTEM.local_segmentations, f'{series_id}.npz')\n            with open(path, 'wb') as f: \n                np.save(f, a)\n\n    def configure_optimizers(self):\n        optimizer = torch.optim.SGD(\n            self.parameters(),\n            lr=self.lr,\n            momentum=self.momentum,\n            dampening=self.dampening,\n            weight_decay=self.weight_decay, \n            nesterov=self.nesterov,\n        )\n        scheduler = torch.optim.lr_scheduler.StepLR(\n            optimizer,\n            step_size=self.lr_step_size,\n            gamma=self.lr_gamma,\n        )\n        lr_scheduler_config = {\n            # REQUIRED: The scheduler instance\n            \"scheduler\": scheduler,\n            # The unit of the scheduler's step size, could also be 'step'.\n            # 'epoch' updates the scheduler on epoch end whereas 'step'\n            # updates it after a optimizer update.\n            \"interval\": \"epoch\",\n            # How many epochs/steps should pass between calls to\n            # `scheduler.step()`. 1 corresponds to updating the learning\n            # rate after every epoch/step.\n            \"frequency\": 1,\n            # Metric to to monitor for schedulers like `ReduceLROnPlateau`\n            \"monitor\": \"val_loss\",\n            # If set to `True`, will enforce that the value specified 'monitor'\n            # is available when the scheduler is updated, thus stopping\n            # training if not found. If set to `False`, it will only produce a warning\n            \"strict\": True,\n            # If using the `LearningRateMonitor` callback to monitor the\n            # learning rate progress, this keyword can be used to specify\n            # a custom logged name\n            \"name\": None,\n        }\n        return {'optimizer': optimizer, 'lr_scheduler': lr_scheduler_config}","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import shutil\n\n\nfrom pytorch_lightning.callbacks.early_stopping import EarlyStopping\nfrom pytorch_lightning.callbacks.model_checkpoint import ModelCheckpoint\nfrom pytorch_lightning.callbacks import LearningRateMonitor\nfrom pytorch_lightning.profilers import AdvancedProfiler\nfrom pytorch_lightning.loggers import WandbLogger\n\n\nmodel_dir = 'vnet'\nos.makedirs(model_dir, exist_ok=True)\n\nwandb_logger = WandbLogger(project=WANDB_PROJECT, save_dir=model_dir)\n\nvnet = VNetTrainer(\n    **HYPERPARAMETERS,\n)\nwandb_logger.watch(vnet, log=\"all\")\nprofiler = AdvancedProfiler(dirpath=model_dir, filename=\"perf_logs\")\ncheckpointer = ModelCheckpoint(\n    dirpath=model_dir,\n    filename='model-{epoch}-{val_loss:.3f}-{train_loss:.3f}',\n    verbose=True,\n    save_top_k=3,\n    monitor='val_loss',\n)\nearly_stopper = EarlyStopping(monitor=\"val_loss\", mode=\"min\", patience=15)\nlr_monitor = LearningRateMonitor(logging_interval='epoch')\ncallbacks = [\n    checkpointer,\n    early_stopper,\n    lr_monitor\n]\n\ntrainer = pl.Trainer(\n    max_time='00:03:00:00',\n    log_every_n_steps=1,\n    profiler=profiler,\n    callbacks=callbacks,\n    default_root_dir=model_dir,\n    logger=wandb_logger,\n    precision='16-mixed',\n)\ntrainer.fit(\n    model=vnet,\n    train_dataloaders=train_dataloader,\n    val_dataloaders=valid_dataloader,\n)\n\nif checkpointer.best_model_path:\n    shutil.copy(checkpointer.best_model_path, os.path.join(model_dir, 'best_model.ckpt'))\n    \nwandb_logger.experiment.unwatch(vnet)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"os.makedirs(FILESYSTEM.local_segmentations, exist_ok=True)\ntrainer.predict(\n    model=vnet,\n    dataloaders=test_dataloader,\n)\nwandb.finish()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"series_mean, series_std = normalisation\nCONFIG = PredictedSegmentationsConfig(\n    series_mean=float(series_mean.detach()),\n    series_std=float(series_std.detach()),\n)\nCONFIG.dump()\n!mv {FILESYSTEM.local_segmentations} ./\n!rm -rf {FILESYSTEM.local}\n!rm -r {plot_dir}\nFILESYSTEM.teardown()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}