{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"codemirror_mode":{"name":"ipython","version":3},"file_extension":".py","mimetype":"text/x-python","name":"python","nbconvert_exporter":"python","pygments_lexer":"ipython3","version":"3.10.12"},"vscode":{"interpreter":{"hash":"916dbcbb3f70747c44a77c7bcd40155683ae19c65e1c03b4aa3499c5328201f1"}},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":6799,"databundleVersionId":4225553,"isSourceIdPinned":false,"sourceType":"competition"},{"sourceId":1462296,"sourceType":"datasetVersion","datasetId":857191},{"sourceId":1618416,"sourceType":"datasetVersion","datasetId":955838},{"sourceId":2898182,"sourceType":"datasetVersion","datasetId":1775969},{"sourceId":9554295,"sourceType":"datasetVersion","datasetId":5821679},{"sourceId":227530414,"sourceType":"kernelVersion"},{"sourceId":284769,"sourceType":"modelInstanceVersion","isSourceIdPinned":false,"modelInstanceId":244042,"modelId":265660},{"sourceId":285894,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":245044,"modelId":266647}],"dockerImageVersionId":30919,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# Disable wandb logging by this environment variable\n%env WANDB_MODE=offline\n# %env ML_DATA_ROOT=/your/path/to/data","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-14T07:58:16.483547Z","iopub.execute_input":"2025-03-14T07:58:16.484144Z","iopub.status.idle":"2025-03-14T07:58:16.489352Z","shell.execute_reply.started":"2025-03-14T07:58:16.484112Z","shell.execute_reply":"2025-03-14T07:58:16.488484Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Universal_pretraining.py","metadata":{}},{"cell_type":"code","source":"%pip install --default-timeout=1000 medicalmultitaskmodeling[interactive,testing]","metadata":{"trusted":true,"scrolled":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%pip install asciitree==0.3.3 crc32c==2.7.1 Deprecated==1.2.18 donfig==0.8.1.post1 fasteners==0.19 fsspec==2025.2.0 imagecodecs==2024.12.30 numcodecs==0.15.1 numpy==2.2.3 packaging==24.2 pillow==11.1.0 PyYAML==6.0.2 tifffile==2025.2.18 tiffslide==2.2.0 typing_extensions==4.12.2 wrapt==1.17.2 zarr==2.14.2","metadata":{"trusted":true,"scrolled":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%pip install -U albumentations","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"pip install medmnist","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"pip install -r /kaggle/input/modell/pytorch/default/1/original_version/requirements.txt","metadata":{"trusted":true,"scrolled":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"pip install --upgrade --force-reinstall Pillow","metadata":{"trusted":true,"scrolled":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\n\n# Set a default cache path (change if needed)\nos.environ[\"ML_DATA_CACHE\"] = \"/kaggle/working/ml_data_cache\"\n","metadata":{"trusted":true,"scrolled":true,"execution":{"iopub.status.busy":"2025-03-14T07:58:19.485435Z","iopub.execute_input":"2025-03-14T07:58:19.486149Z","iopub.status.idle":"2025-03-14T07:58:19.490167Z","shell.execute_reply.started":"2025-03-14T07:58:19.486118Z","shell.execute_reply":"2025-03-14T07:58:19.489287Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import sys\nsys.path.append(\"/kaggle/input/modell/pytorch/default/1/original_version/mtl-torch\")\n\nfrom mtl_torch.data_loading.medical.chestxray import VinBigDataDetCohort\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-14T07:58:19.724552Z","iopub.execute_input":"2025-03-14T07:58:19.724798Z","iopub.status.idle":"2025-03-14T07:58:26.818439Z","shell.execute_reply.started":"2025-03-14T07:58:19.724777Z","shell.execute_reply":"2025-03-14T07:58:26.817505Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import json\nimport os\n\n# Define the directory and file path\ndirectory = \"job_configs\"\nfile_path = os.path.join(directory, \"universal_pretraining.jsonc\")\n\n# Create the directory if it doesn't exist\nos.makedirs(directory, exist_ok=True)\n\n# Define the content for the JSONC file\njson_content = {\n    \"wandb_project\": \"universal_pretraining_default\",\n    \"wandb_entity\": \"tissue-concepts\",\n    \"encoder_args\": {\n        \"hidden_dim\": 768,\n        \"norm_layer\": \"layernorm\",\n        \"model\": {\n            \"use_latent_layer\": False,\n            \"pretrained\": True,\n            \"variant\": \"small\"\n        }\n    },\n    \"decoder_args\": {\n        \"model\": {\n            \"pixel_embedding_dim\": 32\n        }\n    },\n    \"fcosdecoder_args\": {},\n    \"trainer_config\": {\n        \"max_epochs\": 150,\n        \"optim\": {\n            \"lr\": 0.001,\n            \"weight_decay\": 0.0001,\n            \"betas\": [0.9, 0.999]\n        },\n        \"mtl_train_loop\": {\n            \"max_steps\": 4000,\n            \"task_sampler\": {\n                \"mode\": \"infinite\"\n            }\n        },\n        \"mtl_val_loop\": {\n            \"max_steps\": 400,\n            \"task_sampler\": {\n                \"mode\": \"infinite\"\n            }\n        }\n    },\n    \"pretrain_tasks\": [\n        {\n            \"task_id\": \"cocodet\",\n            \"task_type\": \"det\",\n            \"cohort_config\": {\n                \"num_workers\": 2,\n                \"batch_size\": 2\n            }\n        },\n        {\n            \"task_id\": \"brats2020\",\n            \"task_type\": \"seg\",\n            \"cohort_config\": {\n                \"num_workers\": 2,\n                \"batch_size\": 64\n            },\n            \"subjects_splitting_seed\": 42,\n            \"train_ds_config\": {\n                \"drain_each_epoch\": False,\n                \"subcase_cache_size\": 1000\n            }\n        },\n        {\n            \"task_id\": \"kather100k\",\n            \"task_type\": \"clf\",\n            \"cohort_config\": {\n                \"num_workers\": 2,\n                \"batch_size\": 64\n            }\n        },\n        {\n            \"task_id\": \"conicmdet\",\n            \"task_type\": \"det\",\n            \"cohort_config\": {\n                \"num_workers\": 4\n            }\n        },\n        {\n            \"task_id\": \"chexpertmclf\",\n            \"task_type\": \"mclf\",\n            \"cohort_config\": {\n                \"batch_size\": 16,\n                \"num_workers\": 3\n            }\n        },\n        {\n            \"task_id\": \"ssimpneumothoraxseg\",\n            \"task_type\": \"seg\",\n            \"cohort_config\": {\n                \"batch_size\": 16\n            }\n        },\n        {\n            \"task_id\": \"vinbigdatadet\",\n            \"task_type\": \"det\",\n            \"cohort_config\": {\n                \"batch_size\": 8,\n                \"num_workers\": 4\n            }\n        },\n        {\n            \"task_id\": \"cyto\",\n            \"task_type\": \"clf\",\n            \"cohort_config\": {\n                \"batch_size\": 64,\n                \"num_workers\": 2\n            }\n        },\n        {\n            \"task_id\": \"prostategleasonavanitimclf\",\n            \"task_type\": \"mclf\",\n            \"cohort_config\": {\n                \"num_workers\": 1,\n                \"batch_size\": 64\n            },\n            \"train_config\": {\n                \"cache_cfg\": {\n                    \"subcase_cache_size\": 500,\n                    \"drain_each_epoch\": False\n                },\n                \"extractor\": {\n                    \"patch_sizes\": [256, 512, 1024]\n                }\n            },\n            \"val_config\": {\n                \"cache_cfg\": {\n                    \"subcase_cache_size\": 1,\n                    \"drain_each_epoch\": True\n                },\n                \"extractor\": {\n                    \"patch_sizes\": [512]\n                }\n            }\n        },\n        {\n            \"task_id\": \"cocoseg\",\n            \"task_type\": \"seg\",\n            \"cohort_config\": {\n                \"num_workers\": 2,\n                \"batch_size\": 2\n            }\n        },\n        {\n            \"task_id\": \"imgnetclf\",\n            \"task_type\": \"clf\",\n            \"name_overwrite\": \"imgnet_full\",\n            \"cohort_config\": {\n                \"num_workers\": 2,\n                \"batch_size\": 64,\n                \"src_config\": {\n                    \"only_for_classes\": None\n                }\n            }\n        },\n        {\n            \"task_id\": \"amos22\",\n            \"task_type\": \"seg\",\n            \"cohort_config\": {\n                \"only_modality\": \"ct\"\n            }\n        },\n        {\n            \"task_id\": \"radimagenet\",\n            \"task_type\": \"clf\",\n            \"cohort_config\": {\n                \"batch_size\": 64,\n                \"num_workers\": 2\n            },\n            \"split_seed\": 42\n        },\n        {\n            \"task_id\": \"pandaregionclf\",\n            \"task_type\": \"clf\",\n            \"cohort_config\": {\n                \"num_workers\": 1,\n                \"batch_size\": 64\n            }\n        },\n        {\n            \"task_id\": \"picalmclf\",\n            \"task_type\": \"mclf\",\n            \"cohort_config\": {\n                \"num_workers\": 1,\n                \"batch_size\": 64\n            },\n            \"subjects_splitting_seed\": 42,\n            \"train_cacheds\": {\n                \"drain_each_epoch\": False,\n                \"subcase_cache_size\": 500\n            },\n            \"pop_prob_lesion_slice\": 0.2,\n            \"pop_prob_benign_slice\": 0.8\n        },\n        {\n            \"task_id\": \"pandaregionseg\",\n            \"task_type\": \"seg\",\n            \"cohort_config\": {\n                \"num_workers\": 1,\n                \"batch_size\": 64\n            }\n        },\n        {\n            \"task_id\": \"cragseg\",\n            \"task_type\": \"seg\",\n            \"cohort_config\": {\n                \"num_workers\": 1,\n                \"batch_size\": 64\n            }\n        }\n    ],\n    \"pretraining_clf_config\": {\n        \"module_name\": \"clftemplate\"\n    },\n    \"pretraining_mclf_config\": {\n        \"module_name\": \"mclftemplate\"\n    },\n    \"pretraining_det_config\": {\n        \"module_name\": \"dettemplate\"\n    },\n    \"pretraining_seg_config\": {\n        \"module_name\": \"segtemplate\",\n        \"headdropout\": 0.2,\n        \"head_kernel_size\": 1\n    },\n    \"from_checkpoint_path\": \"\"\n}\n\n# Write the content to the JSONC file\nwith open(file_path, \"w\") as file:\n    json.dump(json_content, file, indent=4)\n\nprint(f\"File created at {file_path}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-14T07:58:26.819951Z","iopub.execute_input":"2025-03-14T07:58:26.820279Z","iopub.status.idle":"2025-03-14T07:58:26.836358Z","shell.execute_reply.started":"2025-03-14T07:58:26.820256Z","shell.execute_reply":"2025-03-14T07:58:26.835645Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import sys\nsys.path.append(\"/kaggle/input/modell/pytorch/default/1/original_version\")\n\nfrom histo_data.crag import CRAGSemSegCohort","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-14T07:58:31.440015Z","iopub.execute_input":"2025-03-14T07:58:31.44037Z","iopub.status.idle":"2025-03-14T07:58:31.452188Z","shell.execute_reply.started":"2025-03-14T07:58:31.440343Z","shell.execute_reply":"2025-03-14T07:58:31.451308Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from pathlib import Path\nimport pandas as pd\nfrom torch.utils.data import Dataset\nfrom typing import Dict, Any, List\n\nclass Brats2020Subjects(Dataset[Dict[str, Any]]):\n\n    def __init__(self, root: Path, data_type: str = \"Training\") -> None:\n        self.root = root\n        if data_type == \"Training\":\n            self.data_root = self.root / \"MICCAI_BraTS2020_TrainingData\"\n        elif data_type == \"Validation\":\n            self.data_root = self.root / \"MICCAI_BraTS2020_ValidationData\"\n        else:\n            raise ValueError(\"data_type should be either 'Training' or 'Validation'\")\n\n        print(f\"Looking for data at: {self.data_root}\")\n        assert self.data_root.exists(), f\"Path does not exist: {self.data_root}\"\n\n        # Load name_mapping and survival_info if they exist\n        name_mapping_path = self.data_root / \"name_mapping.csv\"\n        survival_info_path = self.data_root / \"survival_info.csv\"\n\n        if name_mapping_path.exists():\n            self.name_mapping = pd.read_csv(name_mapping_path)\n            self.name_mapping = self.name_mapping.set_index(\"BraTS_2020_subject_ID\")\n        else:\n            self.name_mapping = None\n\n        if survival_info_path.exists():\n            self.survival_info = pd.read_csv(survival_info_path)\n            self.survival_info = self.survival_info.set_index(\"Brats20ID\")\n        else:\n            self.survival_info = None\n\n        self.subjects = {}\n\n        for volpath in self.data_root.glob(\"**/*.nii\"):  # Change to .nii\n            subject_id = volpath.parts[-2]\n            if subject_id not in self.subjects:\n                self.subjects[subject_id] = {}\n\n            filename = volpath.stem  # No need for .replace(\".nii\", \"\")\n            self.subjects[subject_id][filename.split(\"_\")[-1]] = volpath  # Extract modality\n\n    def __len__(self) -> int:\n        return len(self.subjects)\n\n    def __getitem__(self, i: int) -> Dict[str, Any]:\n        subject_id = list(self.subjects.keys())[i]\n        subject = self.subjects[subject_id]\n        res = {\n            \"ds_index\": i,\n            \"paths\": subject,\n            \"subject_id\": subject_id,\n        }\n        return res\n\n    def get_subjects_with_modalities(self, required_modalities: List[str]) -> Dict[str, Dict[str, Path]]:\n        subjects_with_all_modalities = {}\n        for subject_id, modalities in self.subjects.items():\n            if all(modality in modalities for modality in required_modalities):\n                subjects_with_all_modalities[subject_id] = modalities\n        return subjects_with_all_modalities\n\n    def st_printinfo(self):\n        import streamlit as st\n\n        if self.survival_info is not None:\n            st.write(self.survival_info)\n        if self.name_mapping is not None:\n            st.write(self.name_mapping)\n        st.write(self.subjects)\n\n\n# Example usage\ndataset_path = Path(\"/kaggle/input/brats2020-correct-dataset-training-validation/BraTS2020 Dataset (Training + Validation)\")\n\nif dataset_path.exists():\n    print(f\"Dataset path exists: {dataset_path}\")\n    subjects = Brats2020Subjects(dataset_path, data_type=\"Training\")\n    print(f\"Number of subjects: {len(subjects)}\")\n\n    # Define required modalities\n    required_modalities = [\"flair\", \"t1\", \"t1ce\", \"t2\"]\n    subjects_with_modalities = subjects.get_subjects_with_modalities(required_modalities)\n    print(f\"Number of subjects with all required modalities ({', '.join(required_modalities)}): {len(subjects_with_modalities)}\")\n\n    # Print details of the first subject with all required modalities\n    first_subject_id = list(subjects_with_modalities.keys())[0]\n    print(f\"First subject with all required modalities: {first_subject_id}\")\n    print(subjects_with_modalities[first_subject_id])\nelse:\n    print(f\"Dataset path does not exist: {dataset_path}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-14T07:58:31.728523Z","iopub.execute_input":"2025-03-14T07:58:31.728773Z","iopub.status.idle":"2025-03-14T07:58:32.227908Z","shell.execute_reply.started":"2025-03-14T07:58:31.728751Z","shell.execute_reply":"2025-03-14T07:58:32.22703Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\n\n# Set a default empty JSON string for the missing config\nos.environ[\"MLOPS_JSON_universal_pretraining\"] = \"{}\"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-14T07:58:37.475018Z","iopub.execute_input":"2025-03-14T07:58:37.475875Z","iopub.status.idle":"2025-03-14T07:58:37.479848Z","shell.execute_reply.started":"2025-03-14T07:58:37.475841Z","shell.execute_reply":"2025-03-14T07:58:37.478915Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\"\"\"\nCollection of datasets for patch classification on diverse histological data.\n\"\"\"\nfrom pathlib import Path\nfrom typing import Dict\nfrom mtl_torch.data_loading.MTLDataset import mtl_collate\nimport tifffile\nfrom torch.utils.data import random_split\nimport torch.nn as nn\nfrom torch.utils.data.dataset import ConcatDataset\nimport torchvision.datasets\nimport torchvision.transforms as transforms\n\nfrom mtl_torch.augmentations import PILHistoPatchAug, get_histo_augs\nfrom mtl_torch.data_loading.ClassificationDataset import ClassificationDataset\nfrom mtl_torch.data_loading.TrainValCohort import TrainValCohort\nfrom mtl_torch.data_loading.utils import TransformedSubset, train_val_split\nfrom mtl_torch.interactive import pipes\nfrom mtl_torch.transforms import UnifySizes, flatten_list\n\n\nclass Kather100kClassificationCohort(TrainValCohort[ClassificationDataset]):\n    \"\"\"\n    Seems very curated, benchmark performance should be an accuracy of >94% according\n    to https://www.ncbi.nlm.nih.gov/pmc/articles/PMC6345440/\n\n    - 9 classes\n    - 100k patches extracted from 86 H&E-stained slides\n\n    They provide the data with and without normalization.\n    \"\"\"\n    class Config(TrainValCohort.Config):\n        use_train_augs: bool = True\n        batch_size: int = 8\n        val_split: float = 0.1 # Percentage of data to be used for validation\n\n    def __init__(self, p: Path, args: Config):\n        # Load the entire dataset\n        full_dataset = torchvision.datasets.ImageFolder(\n            str(p),\n            loader=tifffile.imread,\n            transform=transforms.Compose([\n                transforms.ToPILImage(),\n                PILHistoPatchAug() if args.use_train_augs else nn.Identity()\n            ]),\n        )\n\n        # Calculate the number of samples for train and validation sets\n        val_size = int(len(full_dataset) * args.val_split)\n        train_size = len(full_dataset) - val_size\n\n        # Split the dataset\n        train_dataset, val_dataset = random_split(full_dataset, [train_size, val_size])\n\n        # Wrap datasets with ClassificationDataset and set class names\n        class_names = full_dataset.classes\n        train_dataset = ClassificationDataset.from_torchvision(train_dataset, class_names=class_names)\n        val_dataset = ClassificationDataset.from_torchvision(val_dataset, class_names=class_names)\n        \n        super().__init__(args, train_dataset, val_dataset)\n\n        # Ensure class names are consistent\n        assert self.datasets[0].class_names == self.datasets[1].class_names\n\nKather100kClassificationCohortConfig = Kather100kClassificationCohort.Config\n\n\nclass SchoemigMarkiefkaClassificationTrainValCohort(TrainValCohort[ClassificationDataset]):\n    class Config(TrainValCohort.Config):\n        use_train_augs: bool = True\n        batch_size: int = 8\n        train_perc_per_subset: Dict[str, float] = {\n            '01_case_western_native': 0,\n            '02_training_native' : 1,\n            '03_wns_leica_native': 1,\n            '04_wns_hama_native': 1,\n            '05_wns_glis_native': 1,\n            '06_ukk_native': 1\n        }\n        split_seed: int = 42\n\n    def __init__(self, args: Config, root: Path) -> None:\n        source_datasets = {\n            name: (TransformedSubset(ds, indices=train_indices) if train_indices else None, TransformedSubset(ds, indices=val_indices) if val_indices else None)\n            for name, train_perc in args.train_perc_per_subset.items()\n            for ds in (torchvision.datasets.ImageFolder(str(root / 'prostate' / 'schoemig-markiefka' / name)),)\n            for train_indices, val_indices in ((None, list(range(len(ds)))) if train_perc == 0 else (list(range(len(ds))), None) if train_perc == 1 else train_val_split(list(range(len(ds))), perc=train_perc, seed=args.split_seed),)\n        }\n\n        class_names = ['nongland', 'gland', 'tumor']\n\n        unify_sizes = UnifySizes(max_pixels_in_batch=600*600*args.batch_size)\n\n        train_ds = ClassificationDataset(\n            ConcatDataset([ds[0] for ds in source_datasets.values() if ds[0] is not None]),\n            class_names=class_names,\n            src_transform=ClassificationDataset.TorchvisionToMTLClf,\n            batch_transform=pipes.Alb(get_histo_augs()) if args.use_train_augs else None,\n            collate_fn=transforms.Compose([unify_sizes, mtl_collate])\n        )\n        val_ds = ClassificationDataset(\n            ConcatDataset([ds[1] for ds in source_datasets.values() if ds[1] is not None]),\n            class_names=class_names,\n            src_transform=ClassificationDataset.TorchvisionToMTLClf,\n            collate_fn=transforms.Compose([unify_sizes, mtl_collate])\n        )\n\n        super().__init__(args, train_ds, val_ds)\n\n\nclass TolkachClassificationTrainValCohort(TrainValCohort[ClassificationDataset]):\n    class Config(TrainValCohort.Config):\n        use_train_augs: bool = True\n        batch_size: int = 8\n        split_seed: int = 42\n\n    def __init__(self, args: Config, root: Path) -> None:\n        class_names = ['norm', 'tu']\n\n        unify_sizes = UnifySizes(max_pixels_in_batch=600*600*args.batch_size)\n\n        train_ds = ClassificationDataset(\n            torchvision.datasets.ImageFolder(str(root / 'prostate' / 'tolkach' / 'val_dataset_1')),\n            class_names=class_names,\n            src_transform=ClassificationDataset.TorchvisionToMTLClf,\n            batch_transform=pipes.Alb(get_histo_augs()) if args.use_train_augs else None,\n            collate_fn=transforms.Compose([unify_sizes, mtl_collate])\n        )\n        val_ds = ClassificationDataset(\n            torchvision.datasets.ImageFolder(str(root / 'prostate' / 'tolkach' / 'val_dataset_2')),\n            class_names=class_names,\n            src_transform=ClassificationDataset.TorchvisionToMTLClf,\n            collate_fn=transforms.Compose([unify_sizes, mtl_collate])\n        )\n\n        super().__init__(args, train_ds, val_ds)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-14T07:58:37.763289Z","iopub.execute_input":"2025-03-14T07:58:37.763532Z","iopub.status.idle":"2025-03-14T07:58:37.784295Z","shell.execute_reply.started":"2025-03-14T07:58:37.763511Z","shell.execute_reply":"2025-03-14T07:58:37.783546Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import warnings\nimport random\nfrom typing import List, Tuple, Dict, Any, Optional, Literal\nfrom tqdm.auto import tqdm\nfrom functools import reduce\nimport imageio.v2 as imageio\nimport os\nimport json5 as json\nfrom pathlib import Path\nimport numpy as np\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import transforms, datasets\nimport torchvision.transforms.functional as F\nimport monai.transforms as monai_transforms\nimport pydicom\nfrom pydicom.pixel_data_handlers.util import apply_voi_lut\nimport numpy as np\nimport pandas as pd\nimport albumentations as A\n\nfrom mtl_torch.bucketizing import Bucket, BucketConfig\nfrom mtl_torch.BaseModel import BaseModel\nfrom mtl_torch.interactive import pipes\nfrom mtl_torch.data_loading.TrainValCohort import TrainValCohort\nfrom mtl_torch.data_loading.ClassificationDataset import ClassificationDataset\nfrom mtl_torch.data_loading.DetectionDataset import DetectionDataset\nfrom mtl_torch.data_loading.MultilabelClassificationDataset import MultilabelClassificationDataset\nfrom mtl_torch.data_loading.SemSegDataset import SemSegDataset\nfrom mtl_torch.data_loading.utils import TransformedSubset, download_and_extract_archive, train_val_split\nfrom mtl_torch.utils import disk_cacher\nfrom mtl_torch.augmentations import get_xray_augs\nfrom mtl_torch.logging.wandb_ext import remove_wandb_special_chars\nfrom mtl_torch.transforms import UnifySizes\nfrom mtl_torch.data_loading.MTLDataset import mtl_collate\n\n\n\nclass ChestXRay2017Cohort(TrainValCohort[ClassificationDataset]):\n    \"\"\"\n    This dataset is the Pneumonia dataset from medmnist with the original resolution, where the images are gray-scale,\n    and their sizes are (384-2,916)x(127-2,713).\n\n    The dataset can be downloaded via the link from the download function.\n\n    Associated publication: https://pubmed.ncbi.nlm.nih.gov/29474911/\n\n    Performance with Resnet50 (224x224) -> 88.4%\n    Performance of Google AutoML Vision -> 94.6%\n    \"\"\"\n\n    class Config(TrainValCohort.Config):\n        resize: Tuple[int, int] = (512, 512)\n        train_aug: bool = True\n\n    def __init__(self,\n                 data_root: Path,\n                 args: Config,\n                 download=False,\n                 folder_name=\"ChestXRay2017\"):\n        \"\"\"\n        data_root: Directory of the ChestXRay2017 dataset\n        \"\"\"\n        self.ds_root = data_root / folder_name\n        train_path, val_path = self.ds_root / 'chest_xray' / 'train', self.ds_root / 'chest_xray' / 'test'\n        if not train_path.exists():\n            if download:\n                self.download_cohort()\n            else:\n                raise Exception(\"Chest XRay 2017 cannot be located on disk\")\n\n        base_transform = pipes.Alb([\n            A.Resize(width=args.resize[0], height=args.resize[1])\n        ])\n        src_t = transforms.Compose([ClassificationDataset.TorchvisionToMTLClf, base_transform])\n        self.train_src = datasets.ImageFolder(str(train_path))\n        self.val_src = datasets.ImageFolder(str(val_path))\n        super().__init__(\n            args,\n            ClassificationDataset(\n                self.train_src,\n                src_transform=src_t,\n                class_names=self.train_src.classes,\n                batch_transform=pipes.Alb(get_xray_augs()) if args.train_aug else nn.Identity()),\n            ClassificationDataset(self.val_src, src_transform=src_t, class_names=self.val_src.classes),\n            cache_foldername=folder_name\n        )\n        assert self.datasets[0].class_names == self.datasets[1].class_names\n\n    def download_cohort(self) -> bool:\n        download_and_extract_archive(\n            self.ds_root,\n            \"https://md-datasets-public-files-prod.s3.eu-west-1.amazonaws.com/31ab5ede-ed34-46d4-b1bf-c63d70411497\",\n            archive_type='zip'\n        )\n        return True\n\n\nclass SIIMSegSource(Dataset[Dict[str, Any]]):\n    \"\"\"\n    Wraps data from https://www.kaggle.com/c/siim-acr-pneumothorax-segmentation.\n\n    Downloaded using `kaggle competitions download -c siim-acr-pneumothorax-segmentation`.\n\n    Data was not mappable, so use temporary fix provided by some random user:\n    ```\n    kaggle datasets download -d abhishek/siim-dicom-images\n    kaggle datasets download -d abhishek/siim-png-images\n    ```\n    \"\"\"\n\n    @staticmethod\n    @disk_cacher()\n    def load_im_paths(root: Path):\n        return {p.stem: p for p in root.glob(\"**/*.dcm\")}\n\n    @staticmethod\n    def mask2rle(img, width, height):\n        \"\"\"\n        Supplied with \"mask_functions.py\" in the dataset.\n        \"\"\"\n        rle = []\n        lastColor = 0\n        currentPixel = 0\n        runStart = -1\n        runLength = 0\n\n        for x in range(width):\n            for y in range(height):\n                currentColor = img[x][y]\n                if currentColor != lastColor:\n                    if currentColor == 255:\n                        runStart = currentPixel\n                        runLength = 1\n                    else:\n                        rle.append(str(runStart))\n                        rle.append(str(runLength))\n                        runStart = -1\n                        runLength = 0\n                        currentPixel = 0\n                elif runStart > -1:\n                    runLength += 1\n                lastColor = currentColor\n                currentPixel += 1\n\n        return \" \".join(rle)\n\n    @staticmethod\n    def rle2mask(rle, width, height, positive_value=255):\n        \"\"\"\n        Supplied with \"mask_functions.py\" in the dataset.\n        \"\"\"\n        mask = np.zeros(width * height)\n        if rle != \" -1\":\n            array = np.asarray([int(x) for x in rle.split()])\n            starts = array[0::2]\n            lengths = array[1::2]\n\n            current_position = 0\n            for index, start in enumerate(starts):\n                current_position += start\n                mask[current_position:current_position+lengths[index]] = positive_value\n                current_position += lengths[index]\n\n        return mask.reshape(width, height)\n\n    def __init__(self, root: Path) -> None:\n        self.ds_root = root / \"siim-dicom-images\"\n        self.df = pd.read_csv(self.ds_root / \"train-rle.csv\", index_col=0)\n        self.image_reader = monai_transforms.LoadImage()\n        self.intensity_scaler = monai_transforms.ScaleIntensity()\n        self.imgs_root = self.ds_root / \"siim-original\"\n        self.img_paths = self.load_im_paths(self.imgs_root)\n        self.image_ids = list(self.df.index.unique())\n        # self.train_paths = list((self.imgs_root / \"dicom-images-train\").glob(\"*\"))\n        # self.test_paths = list((self.imgs_root / \"dicom-images-test\").glob(\"*\"))\n\n    def __len__(self) -> int:\n        return len(self.image_ids)\n\n    def __getitem__(self, index) -> Dict[str, Any]:\n        dicom_id: str = self.image_ids[index]\n        anno_indices = list(np.argwhere(self.df.index == dicom_id).flatten())\n        rle_masks = [self.df.iloc[i][0] for i in anno_indices]\n\n        dcm_path = self.img_paths[dicom_id]\n        dcm_img = pydicom.dcmread(dcm_path)\n        masks = [SIIMSegSource.rle2mask(rle_mask, 1024, 1024, positive_value=1).T != 0. for rle_mask in rle_masks]\n        mask: np.ndarray = reduce(np.bitwise_or, masks)\n\n        pixels = dcm_img.pixel_array\n        return {\n            \"dicomid\": dicom_id,\n            \"image\": F.to_tensor(pixels),\n            \"label\": torch.from_numpy(mask).long(),\n            \"sex\": dcm_img[0x0010, 0x0040].value,\n            \"age\": dcm_img.PatientAge,\n            \"counts\": np.unique(mask, return_counts=True)\n        }\n\n\nclass SIIMPneumothoraxSegCohort(TrainValCohort[SemSegDataset]):\n    \"\"\"\n    https://www.kaggle.com/competitions/siim-acr-pneumothorax-segmentation/overview\n\n    Benchmark should reach about 0.86 mtlval FG IoU, compare with:\n    https://www.kaggle.com/competitions/siim-acr-pneumothorax-segmentation/leaderboard\n    \"\"\"\n\n    class Config(TrainValCohort.Config):\n        resize = (512, 512)\n        train_perc = 0.8\n        split_seed: int = 42\n\n    def __init__(self, data_root: Path, args: Config) -> None:\n        ds_src = SIIMSegSource(data_root)\n        base_transform = transforms.Compose([\n            pipes.Alb([\n                A.Resize(width=args.resize[0], height=args.resize[1])\n            ]),\n            pipes.ApplyToKey(lambda im: torch.concat([im, im, im], 0), key='image'),\n        ])\n        train_indices, val_indices = train_val_split(list(range(len(ds_src))),\n                                                     perc=args.train_perc,\n                                                     seed=args.split_seed)\n\n        super().__init__(\n            args,\n            SemSegDataset(TransformedSubset(ds_src, indices=train_indices),\n                          src_transform=base_transform,\n                          class_names=[\"BG\", \"FG\"],\n                          batch_transform=pipes.Alb(get_xray_augs())),\n            SemSegDataset(TransformedSubset(ds_src, indices=val_indices),\n                          src_transform=base_transform,\n                          class_names=[\"BG\", \"FG\"]),\n            cache_foldername=\"ssimpneumoseg\"\n        )\n\nclass CheXpertSource(Dataset[Dict[str, Any]]):\n    \"\"\"\n    CheXpert source wrapper.\n\n    Classes are labeled using uncertainty labels because classes are extracted from text reports using keyword search.\n\n    The training set consists of 223414 samples.\n    \"\"\"\n\n    def __init__(self,\n                 train: bool = True,\n                 folder_name: Literal[\"CheXpert-v1.0-small\"] = \"CheXpert-v1.0-small\") -> None:\n        root = Path(\"/kaggle/input/chexpert\")\n        self.ds_root = Path('/kaggle/input/chexpert')\n        \n        if train:\n            self.split_root = self.ds_root / \"train\"\n            self.df = pd.read_csv(Path('/kaggle/input/chexpert/CheXpert-v1.0-small/train.csv'))\n        else:\n            self.split_root = self.ds_root / \"valid\"\n            self.df = pd.read_csv(Path('/kaggle/input/chexpert/CheXpert-v1.0-small/valid.csv'))\n\n        # Fix all image paths to correct absolute paths on this environment\n        self.df[\"Path\"] = self.df[\"Path\"].apply(lambda v: self.ds_root / v)\n\n    def __len__(self):\n        return self.df.shape[0]\n\n    @staticmethod\n    def image_to_mtlformat(img: np.ndarray) -> torch.Tensor:\n        t = F.to_tensor(img)\n        return torch.cat([t, t, t], dim=0)\n\n    def __getitem__(self, index) -> Dict[str, Any]:\n        row = dict(self.df.iloc[index])  # type: ignore\n        row[\"Path\"] = Path(row[\"Path\"])\n        row[\"filestem\"] = row[\"Path\"].stem\n        row[\"study\"] = row[\"Path\"].parts[-2]\n        row[\"patient\"] = row[\"Path\"].parts[-3]\n        return {\n            \"image\": imageio.imread(row[\"Path\"]),\n            \"meta\": row\n        }\n\nclass CheXpertClassificationCohort(TrainValCohort[ClassificationDataset]):\n\n    class Config(TrainValCohort.Config):\n        resize = (512, 512)\n        for_column: Literal['Sex', 'Frontal/Lateral', 'AP/PA', 'No Finding', 'Enlarged Cardiomediastinum',\n                            'Cardiomegaly', 'Lung Opacity', 'Lung Lesion', 'Edema', 'Consolidation', 'Pneumonia', 'Atelectasis',\n                            'Pneumothorax', 'Pleural Effusion', 'Pleural Other', 'Fracture', 'Support Devices'] = \"Pneumonia\"\n        train_augs: bool = True\n\n    def __init__(self,\n                 data_root: Path,\n                 args: Config,\n                 folder_name: Literal[\"CheXpert-v1.0-small\"] = \"CheXpert-v1.0-small\",\n                 cache_foldername: str = \"\") -> None:\n        self.args: CheXpertClassificationCohort.Config\n        self.train_src = CheXpertSource(train=True, folder_name=folder_name)\n        print(f\"Training source initialized: {self.train_src}\")\n        self.val_src = CheXpertSource(train=False, folder_name=folder_name)\n        print(f\"Validation source initialized: {self.val_src}\")\n\n        print(f\"Initial training DataFrame shape: {self.train_src.df.shape}\")\n        print(f\"Initial validation DataFrame shape: {self.val_src.df.shape}\")\n\n        self.train_src.df = self.train_src.df[~self.train_src.df[args.for_column].isna()]\n        self.val_src.df = self.val_src.df[~self.val_src.df[args.for_column].isna()]\n\n        print(f\"Filtered training DataFrame shape: {self.train_src.df.shape}\")\n        print(f\"Filtered validation DataFrame shape: {self.val_src.df.shape}\")\n\n        class_names = list(map(lambda x: remove_wandb_special_chars(\n            str(x)), self.train_src.df[args.for_column].unique()))\n        print(f\"Class names: {class_names}\")\n\n        def extract_labels(d):\n            d['class'] = class_names.index(remove_wandb_special_chars(str(d['meta'][args.for_column])))\n            print(f\"Extracted labels: {d['class']}\")\n            return d\n\n        src_transforms = [\n            pipes.ApplyToKey(CheXpertSource.image_to_mtlformat, key=\"image\"),\n            pipes.Alb([A.Resize(width=args.resize[0], height=args.resize[1])]),\n            extract_labels,\n        ]\n\n        print(\"Source transforms initialized\")\n\n        super().__init__(\n            args,\n            ClassificationDataset(self.train_src,\n                                  src_transform=transforms.Compose(src_transforms),\n                                  class_names=class_names,\n                                  batch_transform=pipes.Alb(get_xray_augs()) if args.train_augs else None),\n            ClassificationDataset(self.val_src,\n                                  src_transform=transforms.Compose(src_transforms),\n                                  class_names=class_names),\n            cache_foldername=cache_foldername\n        )\n\n        print(\"CheXpertClassificationCohort initialized\")\n\n\nclass CheXpertMultilabelCohort(TrainValCohort[MultilabelClassificationDataset]):\n    class Config(TrainValCohort.Config):\n        # Ordered by their label quality\n        classes: List[str] = ['Sex',  # Good labels, train and val about 0.98\n                              'Edema',  # Ok; train 0.8, val over 0.7\n                              'Lung Opacity',  # Ok: train 0.9, val over 0.7\n                              'Consolidation',  # Good\n                              'Pneumonia',  # Good val acc because of high imbalance\n                              'Pneumothorax',  # Good val acc because of high imbalance\n                              'Pleural Effusion',  # Good\n                              'Enlarged Cardiomediastinum',  # Noisy but ok, val worse than train\n                              'Cardiomegaly',  # Tends to overfit but ok\n                              #   'Lung Lesion',  # Bad with low valacc of 0.15 but good val auc (0.7)\n                              #   'Support Devices' # Random val acc\n                              #   'Fracture', # not well represented in val set\n                              #   'Pleural Other', # bad\n                              #   'Atelectasis', # Seems random\n                              #   'Frontal/Lateral', # Very obvious, train and val 1.\n                              #   'No Finding',  # Bad, train acc 1., val acc static\n                              ]\n        source_folder: Literal[\"downsized\", \"fullsized\"] = \"fullsized\"\n        longest_edge: int = 512\n        batch_size = 8\n\n    def __init__(self,\n                 data_root: Path,\n                 args: Config,\n                 cache_foldername: str = \"\") -> None:\n        self.args: CheXpertMultilabelCohort.Config = args\n        folder_name = \"CheXpert-v1.0-small\"\n        self.train_src = CheXpertSource(train=True, folder_name=folder_name)\n        self.val_src = CheXpertSource(train=False, folder_name=folder_name)\n\n        self.class_names = args.classes\n        self.positives = {\n            'Sex': (\"Female\", [\"Unknown\"]),\n            'Frontal/Lateral': (\"1.0\", [\"nan\"]),\n            'No Finding': (\"1.0\", [\"nan\"]),\n            'Enlarged Cardiomediastinum': (\"1.0\", [\"nan\"]),\n            'Cardiomegaly': (\"1.0\", [\"nan\"]),\n            'Lung Opacity': (\"1.0\", [\"nan\"]),\n            'Lung Lesion': (\"1.0\", [\"nan\"]),\n            'Edema': (\"1.0\", [\"nan\"]),\n            'Consolidation': (\"1.0\", [\"nan\"]),\n            'Pneumonia': (\"1.0\", [\"nan\"]),\n            'Atelectasis': (\"1.0\", [\"nan\"]),\n            'Pneumothorax': (\"1.0\", [\"nan\"]),\n            'Pleural Effusion': (\"1.0\", [\"nan\"]),\n            'Pleural Other': (\"1.0\", [\"nan\"]),\n            'Fracture': (\"1.0\", [\"nan\"]),\n            'Support Devices': (\"1.0\", [\"nan\"])\n        }\n\n        def extract_labels(d):\n            pos_each_class = [str(d['meta'][col_name]) == self.positives[col_name][0] for col_name in self.class_names]\n            loss_weights = [str(d['meta'][col_name]) not in self.positives[col_name][1]\n                            for col_name in self.class_names]\n            d['class_labels'] = torch.Tensor(pos_each_class).float()\n            d['loss_weights'] = torch.Tensor(loss_weights).float()\n            return d\n\n        src_transforms = [\n            pipes.ApplyToKey(CheXpertSource.image_to_mtlformat, key=\"image\"),\n            extract_labels,\n        ]\n\n        unifysizes = UnifySizes(max_edge_len=self.args.longest_edge)\n\n        super().__init__(\n            args,\n            MultilabelClassificationDataset(\n                self.train_src,\n                src_transform=transforms.Compose(src_transforms),\n                class_names=self.class_names,\n                batch_transform=pipes.Alb(get_xray_augs()),\n                collate_fn=transforms.Compose([unifysizes, mtl_collate])),\n            MultilabelClassificationDataset(\n                self.val_src,\n                src_transform=transforms.Compose(src_transforms),\n                class_names=self.class_names,\n                collate_fn=transforms.Compose([unifysizes, mtl_collate])\n            ),\n            cache_foldername=cache_foldername\n        )\n\nclass VinBigDataSrc(Dataset[Dict[str, Any]]):\n    \"\"\"\n    Single class per image including a box for non-normal cases. 18000\n\n    Classes: \n\n    0 - Aortic enlargement\n    1 - Atelectasis\n    2 - Calcification\n    3 - Cardiomegaly\n    4 - Consolidation\n    5 - ILD\n    6 - Infiltration\n    7 - Lung Opacity\n    8 - Nodule/Mass\n    9 - Other lesion\n    10 - Pleural effusion\n    11 - Pleural thickening\n    12 - Pneumothorax\n    13 - Pulmonary fibrosis\n    \"\"\"\n\n    def __init__(self, root: Path, writable_root: Path, from_folder: Literal[\"dicom\", \"jpg\"]) -> None:\n        self.root = root / \"vinbigdata-chest-xray-abnormalities-detection\"\n        self.df = pd.read_csv(self.root / \"train.csv\", index_col=\"image_id\")\n        self.image_ids = list(set(list(self.df.index)))\n        self.class_names: List[str] = [remove_wandb_special_chars(x) for x in pd.unique(self.df['class_name'])]\n        self.dcm_folder = self.root / \"train\"\n        self.jpg_folder = writable_root / \"train_jpg\"\n        self.meta_folder = writable_root / \"meta\"\n        self.from_folder = from_folder\n\n        # self.src_transform = pipes.Alb([A.Resize(self.cfg.resize[0], self.cfg.resize[1])])\n        self.src_transform = nn.Identity()\n\n        # Do not use in-built caching mechanism because there are multiple cases with the same images\n        if self.from_folder == \"jpg\" and (\n            not self.jpg_folder.exists() or len(os.listdir(self.jpg_folder)) < len(os.listdir(self.dcm_folder))\n        ):\n            self.jpg_folder.mkdir(exist_ok=True)\n            self.meta_folder.mkdir(exist_ok=True)\n            self.copy_as_jpeg()\n\n    def __len__(self) -> int:\n        return len(self.image_ids)\n\n    def copy_as_jpeg(self):\n        \"\"\"\n        Needs write permissions, uses multiprocessing to preprocess all full-sized dicom images to a specific jpg size\n        \"\"\"\n        class FastDS:\n\n            def __init__(self, ls, src_transform, jpg_folder, meta_folder) -> None:\n                self.ls, self.src_transform, self.jpg_folder = ls, src_transform, jpg_folder\n                self.meta_folder = meta_folder\n\n            def __len__(self):\n                return len(self.ls)\n\n            def __getitem__(self, i):\n                jpg_path = self.jpg_folder / f'{self.ls[i].stem}.jpg'\n                meta_path = self.meta_folder / f'{self.ls[i].stem}.json'\n                if not jpg_path.exists():\n                    dcm = pydicom.dcmread(self.ls[i])\n                    src_img = VinBigDataSrc.read_xray(dcm)\n                    mtl_image = self.src_transform({\"image\": src_img})[\"image\"]\n                    if not meta_path.exists():\n                        with open(meta_path, \"w+\") as f:\n                            json.dump({\n                                \"imgsize\": src_img.shape\n                            }, f)\n                    imageio.imwrite(jpg_path, F.to_pil_image(mtl_image))\n                else:\n                    assert meta_path.exists()\n                return i\n\n        dcm_paths = FastDS(list(self.dcm_folder.glob(\"*.dicom\")), self.src_transform, self.jpg_folder,\n                           self.meta_folder)\n        for _ in tqdm(\n            DataLoader(dcm_paths, num_workers=24, batch_size=1, shuffle=True, collate_fn=lambda x: x)  # type: ignore\n        ):\n            ...\n\n    @staticmethod\n    def read_xray(dicom, voi_lut=True, fix_monochrome=True):\n        \"\"\"\n        Adapted from https://www.kaggle.com/code/raddar/convert-dicom-to-np-array-the-correct-way\n        \"\"\"\n        # VOI LUT (if available by DICOM device) is used to transform raw DICOM data to \"human-friendly\" view\n        if voi_lut:\n            data = apply_voi_lut(dicom.pixel_array, dicom)\n        else:\n            data = dicom.pixel_array\n\n        # depending on this value, X-ray may look inverted - fix that:\n        if fix_monochrome and dicom.PhotometricInterpretation == \"MONOCHROME1\":\n            data = np.amax(data) - data\n\n        data = data - np.min(data)\n        data = data / np.max(data)\n\n        torch_pixels = torch.Tensor(data)\n        return torch.stack([torch_pixels, torch_pixels, torch_pixels])\n\n    def __getitem__(self, index) -> Dict[str, Any]:\n        image_id = self.image_ids[index]\n        rows = [dict(row) for _, row in self.df[self.df.index == image_id].iterrows()]\n\n        if self.from_folder == \"dicom\":\n            with warnings.catch_warnings():\n                # Suppress warning that frequently occurs:\n                # The (0028,0101) 'Bits Stored' value (14-bit) doesn't match the JPEG 2000 data (16-bit).\n                warnings.simplefilter(\"ignore\")\n                dcm = pydicom.dcmread(self.dcm_folder / f'{image_id}.dicom')\n                mtl_image = self.read_xray(dcm)\n        else:\n            mtl_image = F.to_tensor(imageio.imread(self.jpg_folder / f'{image_id}.jpg'))\n\n        # If there are annotations from multiple radiologists, select one randomly\n        rad_ids = [r['rad_id'] for r in rows]\n        selected_rad = random.choice(rad_ids)\n        transforms.ToTensor\n        box_ls, box_labels_ls = [], []\n        for row in filter(lambda r: r['rad_id'] == selected_rad, rows):\n            if not np.isnan(row[\"x_min\"]):\n                if self.from_folder == \"jpg\":\n                    with open(self.meta_folder / f'{image_id}.json') as f:\n                        meta_dict = json.load(f)\n                        orig_size = meta_dict['imgsize']  # type: ignore\n                    x_factor, y_factor = orig_size[2] / mtl_image.shape[2], orig_size[1] / mtl_image.shape[1]\n                    row_box = torch.Tensor(\n                        [\n                            row[\"x_min\"] // x_factor,\n                            row[\"y_min\"] // y_factor,\n                            row[\"x_max\"] // x_factor,\n                            row[\"y_max\"] // y_factor\n                        ]\n                    ).float()\n                else:\n                    row_box = torch.Tensor(\n                        [\n                            row[\"x_min\"],\n                            row[\"y_min\"],\n                            row[\"x_max\"],\n                            row[\"y_max\"]\n                        ]\n                    ).float()\n\n                box_ls.append(row_box)\n                box_labels_ls.append(self.class_names.index(remove_wandb_special_chars(row['class_name'])))\n        if not box_ls:\n            boxes = torch.empty(0, 4).float()\n            box_labels = torch.empty((0,)).long()\n        else:\n            boxes = torch.stack(box_ls)\n            box_labels = torch.Tensor(box_labels_ls).long()\n\n        return {\n            \"image\": mtl_image,\n            # \"class\": self.class_names.index(rows[0]['class_name']),\n            \"boxes\": boxes,\n            \"labels\": box_labels,\n            \"meta\": {\n                \"row\": rows,\n                \"selected_radiologist\": selected_rad\n                # sex seems to be not available for all patients in this dataset\n                # \"sex\": dcm.PatientSex,\n            }\n        }\n\n\nclass VinBigDataDetCohort(TrainValCohort[DetectionDataset]):\n    class Config(TrainValCohort.Config):\n        split_seed: int = 42\n        max_edge: int = 544\n        batch_size: int = 4\n\n    def __init__(self, root: Path, writable_root: Path, cfg: Config, *args, **kwargs) -> None:\n        self.root, self.cfg = root, cfg\n        self.src_ds = VinBigDataSrc(self.root, writable_root, from_folder=\"dicom\")\n        train_indices, val_indices = train_val_split(list(range(len(self.src_ds))),\n                                                     perc=0.85,\n                                                     seed=self.cfg.split_seed)\n        collate = UnifySizes(max_edge_len=self.cfg.max_edge, support_boxes=True)\n        super().__init__(cfg,\n                         DetectionDataset(\n                             TransformedSubset(self.src_ds, indices=train_indices),\n                             class_names=self.src_ds.class_names[1:],\n                             batch_transform=pipes.AlbWithBoxes(get_xray_augs()),\n                             collate_fn=collate\n                         ),\n                         DetectionDataset(\n                             TransformedSubset(self.src_ds, indices=val_indices),\n                             class_names=self.src_ds.class_names[1:],\n                             collate_fn=collate\n                         ),\n                         *args,\n                         **kwargs\n                         )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-14T08:06:29.905847Z","iopub.execute_input":"2025-03-14T08:06:29.906391Z","iopub.status.idle":"2025-03-14T08:06:29.98103Z","shell.execute_reply.started":"2025-03-14T08:06:29.906349Z","shell.execute_reply":"2025-03-14T08:06:29.980022Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Imagenet","metadata":{}},{"cell_type":"code","source":"import json\nfrom collections import OrderedDict\nfrom mtl_torch.data_loading.SemSegDataset import SemSegDataset\nimport torch\nfrom torch.utils.data import Dataset\nimport pandas as pd\nfrom pathlib import Path\nfrom typing import Optional, Tuple, Dict, Any, Literal, List\nimport torchvision.transforms as transforms\nimport numpy as np\nimport torchvision.transforms.functional as F\nfrom torchvision.datasets.folder import default_loader\n\nfrom mtl_torch.BaseModel import BaseModel\nfrom mtl_torch.interactive import data, pipes, configs as cfs\nfrom mtl_torch.utils import disk_cacher\nfrom mtl_torch.data_loading.TrainValCohort import TrainValCohort\nfrom mtl_torch.data_loading.ClassificationDataset import ClassificationDataset\nfrom mtl_torch.data_loading.DetectionDataset import DetectionDataset\nfrom mtl_torch.augmentations import get_realworld_augs\nfrom pycocotools.coco import COCO\nfrom mtl_torch.transforms import UnifySizes\nfrom mtl_torch.data_loading.MTLDataset import mtl_collate\n\nimport albumentations as A\nfrom albumentations.pytorch.transforms import ToTensorV2\n\n\nclass ImageNetSrc(Dataset[Dict[str, Any]]):\n    \"\"\"\n    Some images are gray scale.\n    At least case 294960 seems to have 4 channel.\n    Cohorts need to fix that.\n    \"\"\"\n    @staticmethod\n    def folder_name() -> str:\n        return \"imagenet-object-localization-challenge\"\n\n    @staticmethod\n    @disk_cacher()\n    def get_imagenet_class_descriptions_and_ids(p: Path) -> OrderedDict:\n        class_descs = OrderedDict()\n        with open(p, \"r\") as f:\n            for line in f.readlines():\n                class_id, class_desc = line[:9], line[10:-1]\n                class_descs[class_id] = class_desc\n        return class_descs\n\n    @staticmethod\n    def load_imagenet_img_clf(img_searchpath: Path, df: pd.DataFrame, for_classes: List[str]):\n        classes_and_images = []\n        for class_id in for_classes:\n            for p in img_searchpath.glob(f\"{class_id}/*.JPEG\"):\n                classes_and_images.append((\n                    p.stem,\n                    p.parts[-2],\n                    p,\n                    None\n                ))\n        return classes_and_images\n\n    @staticmethod\n    def load_imagenet_img_boxes(img_root: Path, df: pd.DataFrame, for_classes: List[str], train: bool):\n        classes_and_images = []\n        for _, line in df.iterrows():\n            lineparts = line[\"PredictionString\"].split(\" \")\n            boxes = [(lineparts[i], lineparts[i+1:i+5]) for i in range(0, len(lineparts) - 1, 5)]\n            assert len(set([box[0] for box in boxes])) == 1\n            cls_id = boxes[0][0]\n            if cls_id in for_classes:\n                classes_and_images.append((\n                    line[\"ImageId\"],\n                    cls_id,\n                    img_root / cls_id / f\"{line['ImageId']}.JPEG\" if train else img_root / f\"{line['ImageId']}.JPEG\",\n                    boxes\n                ))\n        return classes_and_images\n\n    class Config(BaseModel):\n        only_for_classes: Optional[List[str]] = None\n\n    def __init__(self, root: Path,\n                 split: Literal[\"train\", \"val\", \"test\"],\n                 args: Config,\n                 for_detection: bool) -> None:\n        self.ds_root = Path('/kaggle/input/imagenet-object-localization-challenge')\n        self.all_classes = ImageNetSrc.get_imagenet_class_descriptions_and_ids(Path('/kaggle/input/imagenet-object-localization-challenge/LOC_synset_mapping.txt'))\n        self.for_class_ids = list(self.all_classes.keys()) if args.only_for_classes is None else args.only_for_classes\n        self.class_names = [self.all_classes[class_id] for class_id in self.for_class_ids]\n        self.df = pd.read_csv(self.ds_root / f\"LOC_{split}_solution.csv\")\n        # self.df_LOC_val_solution = pd.read_csv(self.ds_root / \"LOC_val_solution.csv\")\n        self.root_annotations = self.ds_root / \"ILSVRC\" / \"Annotations\" / \"CLS-LOC\" / split\n        self.root_data = self.ds_root / \"ILSVRC\" / \"Data\" / \"CLS-LOC\" / split\n        self.root_imagesets = self.ds_root / \"ILSVRC\" / \"ImageSets\" / \"CLS-LOC\"\n        assert self.root_annotations.exists() and self.root_data.exists() and self.root_imagesets.exists()\n\n        if split == 'train':\n            if for_detection:\n                images = ImageNetSrc.load_imagenet_img_boxes(self.root_data, self.df, self.for_class_ids, True)\n            else:\n                images = ImageNetSrc.load_imagenet_img_clf(self.root_data, self.df, self.for_class_ids)\n        elif split == 'val':\n            images = ImageNetSrc.load_imagenet_img_boxes(self.root_data, self.df, self.for_class_ids, False)\n        else:\n            raise NotImplementedError()\n\n        # To avoid copy-on-access problem, convert to array type!\n        self.imageids = [x[0] for x in images]\n        self.class_ids = [x[1] for x in images]\n        self.img_paths = [str(x[2]) for x in images]\n        # if split == 'val' or for_detection:\n        #     self.boxes = torch.Tensor([[[float(x) for x in box[1]] for box in x[3]] for x in images])\n        #     self.labels = torch.Tensor([[self.for_class_ids.index(box[0]) for box in x[3]] for x in images]).long()\n        # else:\n        #     self.boxes, self.labels = None, None\n\n    def get_description(self) -> str:\n        return \"\"\"\n        Downloaded from kaggle: https://www.kaggle.com/competitions/imagenet-object-localization-challenge/\n\n        Do not compare with kaggle scores, because the evaluation is skewed towards a low number of predictions.\n\n        Labels are synonym sets such as \"goldfish, carassius auratus\" (for all images) or boxes (for 544546 images).\n        \"\"\"\n\n    def __len__(self) -> int:\n        return len(self.imageids)\n\n    def __getitem__(self, index) -> Dict[str, Any]:\n        # imageid, class_id, img_path, boxes = self.images[index]\n        imageid, class_id = self.imageids[index], self.class_ids[index]\n        img_path = self.img_paths[index]\n        # boxes, box_labels =\n        # cls_index = np.where(self.for_class_ids_arr == self.class_ids[index])[0][0]\n        cls_index = self.for_class_ids.index(self.class_ids[index])\n\n        res: Dict[str, Any] = {\n            \"image\": F.to_tensor(default_loader(img_path)),\n            \"meta\": {\n                \"class_id\": class_id,\n                \"class_index\": cls_index,\n                \"class_desc\": self.class_names[cls_index],\n                \"image_id\": imageid\n            }\n        }\n        # if boxes is not None:\n        #     res[\"meta\"][\"boxes\"] =\n        #     res[\"meta\"][\"labels\"] = torch.Tensor([self.for_class_ids.index(box[0]) for box in boxes]).long()\n        return res\n\n\ndef extract_clf_label(d: Dict) -> Dict:\n    d['class'] = d['meta']['class_index']\n    if d['image'].shape[0] > 3:\n        d['image'] = d['image'][:3, ...]\n    elif d['image'].shape[0] < 3:\n        d['image'] = torch.concat([d['image']] * (3 // d['image'].shape[0]), 0)\n\n    return d\n\n\nclass ImgnetClassificationCohort(TrainValCohort[ClassificationDataset]):\n    \"\"\"\n    Only uses the original training split, because I do not know how to recover the classes of the val split.\n    \"\"\"\n    class Config(TrainValCohort.Config):\n        src_config: ImageNetSrc.Config = ImageNetSrc.Config()\n        resize: Tuple[int, int] = (256, 256)\n\n    def __init__(self, ds_root: Path, args: Config):\n        self.train_src = ImageNetSrc(ds_root, \"train\", args.src_config, for_detection=False)\n        self.val_src = ImageNetSrc(ds_root, \"val\", args.src_config, for_detection=False)\n\n        src_t = transforms.Compose([\n            extract_clf_label,\n            pipes.Alb([A.Resize(height=args.resize[0], width=args.resize[1])])\n        ])\n\n        super().__init__(\n            args,\n            ClassificationDataset(\n                self.train_src,\n                src_transform=src_t,\n                class_names=self.train_src.class_names,\n                batch_transform=pipes.Alb(get_realworld_augs())),\n            ClassificationDataset(self.val_src, src_transform=src_t, class_names=self.val_src.class_names),\n        )\n\n\ndef combine_mask_to_panoptic(img: torch.Tensor, masks: List[torch.Tensor]):\n    out = np.zeros_like(img[0])\n\n    # combine the disjoint masks into out\n    for mask in masks:\n        out[mask != 0] = mask[mask != 0]\n\n    return torch.from_numpy(out)\n\n\nclass CocoSrc(Dataset):\n\n    def __init__(self, \n                 img_root: Path, \n                 anno_path: Path,\n                 classes: Literal[\"supercategory\", \"category\"] = \"category\") -> None:\n        self.coco = COCO(anno_path)\n        self.img_root = img_root\n        self.img_ids = list(sorted(self.coco.imgs.keys()))\n        if classes == \"category\":\n            self.class_names = [v['name'] for v in self.coco.cats.values()]\n        else:\n            self.class_names = list(set([c['supercategory'] for c in self.coco.cats.values()]))\n        self.class_names.sort()\n        self.class_names = [\"bg\"] + self.class_names\n        self.id_to_class = {k: self.class_names.index(v['supercategory' if classes == \"supercategory\" else 'name']) \n                            for k, v in self.coco.cats.items()}\n\n    def __len__(self) -> int:\n        return len(self.coco.imgs)\n\n    def __getitem__(self, index: int) -> Dict[str, Any]:\n        # Load img\n        img_id = self.img_ids[index]\n\n        img_path = self.img_root / self.coco.imgs[img_id]['file_name']\n        img = F.to_tensor(default_loader(str(img_path)))\n\n        # Load annos\n        ann_ids = self.coco.getAnnIds(imgIds=img_id)\n        annos = self.coco.loadAnns(ann_ids)\n        anno_classes = [self.id_to_class[anno[\"category_id\"]] for anno in annos]\n\n        anno_masks = [self.coco.annToMask(anno) * anno_classes[i]\n                      for i, anno in enumerate(annos)]\n        mask = combine_mask_to_panoptic(img, anno_masks)\n        # Minus one because predicting background makes no sense for boxes\n        box_labels = torch.tensor(anno_classes, dtype=torch.long) - 1\n        boxes = torch.tensor([anno[\"bbox\"] for anno in annos], dtype=torch.float32)\n\n        # COCO format boxes in [x1, y1, width, height] to [x1, y1, x2, y2]\n        if boxes.numel() > 0:\n            boxes[:, 2] += boxes[:, 0]\n            boxes[:, 3] += boxes[:, 1]\n\n        return {\n            \"image\": img,\n            \"label\": mask,\n            \"boxes\": boxes,\n            \"labels\": box_labels,\n            \"meta\": {\n                \"imgid\": img_id,\n                \"anno_classes\": anno_classes,\n                \"anno_names\": [self.class_names[i] for i in anno_classes],\n            }\n        }\n\n\nclass CocoSegCohort(TrainValCohort[data.SemSegDataset]):\n    class Config(TrainValCohort.Config):\n        pass\n\n    def __init__(self, args: Config, coco_root: Path, max_pixels_in_batch=224*224*40*3) -> None:\n        coco_root = Path('/kaggle/input/coco-2017-dataset/coco2017')\n        coco_anno_root = coco_root / \"annotations\"\n        trainsrc = CocoSrc(coco_root / \"train2017\", coco_anno_root / \"instances_train2017.json\")\n        valsrc = CocoSrc(coco_root / \"val2017\", coco_anno_root / \"instances_val2017.json\")\n        unifysizes = UnifySizes(max_pixels_in_batch=max_pixels_in_batch)\n        # return CocoSrc(env.data_root / \"siim-covid19-detection\", self.cohort_config)\n        assert trainsrc.class_names == valsrc.class_names\n        super().__init__(\n            args,\n            data.SemSegDataset(trainsrc,\n                               class_names=trainsrc.class_names,\n                               #   src_transform=pipes.Alb([A.Resize(512, 512)]),\n                               batch_transform=pipes.Alb(get_realworld_augs()),\n                               collate_fn=transforms.Compose([unifysizes, mtl_collate]),\n                               ),\n            data.SemSegDataset(valsrc,\n                               class_names=valsrc.class_names,\n                               #   src_transform=pipes.Alb([A.Resize(512, 512)]),\n                               collate_fn=transforms.Compose([unifysizes, mtl_collate])\n                               )\n        )\n\n\nclass CocoDetCohort(TrainValCohort[data.DetectionDataset]):\n    class Config(TrainValCohort.Config):\n        pass\n\n    def __init__(self, args: Config, coco_root: Path, max_pixels_in_batch=224*224*40*3) -> None:\n        coco_root = Path('/kaggle/input/coco-2017-dataset/coco2017')\n        coco_anno_root = coco_root / \"annotations\"\n        trainsrc = CocoSrc(coco_root / \"train2017\", coco_anno_root / \"instances_train2017.json\")\n        valsrc = CocoSrc(coco_root / \"val2017\", coco_anno_root / \"instances_val2017.json\")\n        unifysizes = UnifySizes(max_pixels_in_batch=max_pixels_in_batch, support_boxes=True)\n        # return CocoSrc(env.data_root / \"siim-covid19-detection\", self.cohort_config)\n        assert trainsrc.class_names == valsrc.class_names\n        super().__init__(\n            args,\n            data.DetectionDataset(trainsrc,\n                                  class_names=trainsrc.class_names[1:],\n                                  #   src_transform=pipes.Alb([A.Resize(512, 512)]),\n                                  batch_transform=pipes.Alb(get_realworld_augs(), support_boxes=True),\n                                  collate_fn=transforms.Compose([unifysizes]),\n                                  ),\n            data.DetectionDataset(valsrc,\n                                  class_names=valsrc.class_names[1:],\n                                  #   src_transform=pipes.Alb([A.Resize(512, 512)]),\n                                  collate_fn=transforms.Compose([unifysizes])\n                                  )\n        )\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-14T08:06:41.813374Z","iopub.execute_input":"2025-03-14T08:06:41.813714Z","iopub.status.idle":"2025-03-14T08:06:41.855919Z","shell.execute_reply.started":"2025-03-14T08:06:41.813685Z","shell.execute_reply":"2025-03-14T08:06:41.854931Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Train","metadata":{}},{"cell_type":"code","source":"os.environ[\"WANDB_RESUME\"] = \"allow\" ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-14T08:06:48.740916Z","iopub.execute_input":"2025-03-14T08:06:48.741245Z","iopub.status.idle":"2025-03-14T08:06:48.747467Z","shell.execute_reply.started":"2025-03-14T08:06:48.741222Z","shell.execute_reply":"2025-03-14T08:06:48.746568Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nfrom mtl_torch.data_loading.medical.chestxray import VinBigDataDetCohort\nfrom mtl_torch.data_loading.medical.chestxray import TuberculosisShenzhenClfCohort\nfrom mtl_torch.data_loading.medical.msd import PicaiSubjects\nfrom mtl_torch.augmentations import MRI3DProcessor\nfrom mtl_torch.neural.modules.swinformer import TorchVisionSwinformer\nfrom mtl_torch.torch_ext import CachingSubCaseDS\nfrom mtl_torch.data_loading.medical.msd import Brats2020Volumes, Amos22Cohort\nfrom mtl_torch.data_loading.medical.radimagenet import RadImageNetSrc\nfrom mtl_torch.data_loading.medical.pandav2 import PandaRegionSegCohort\nfrom mtl_torch.data_loading.medical.pandav2 import PandaRegionClfCohort\nfrom mtl_torch.data_loading.medical.prostate import ArvanitiGleasonSemseg\nfrom histo_data.crag import CRAGSemSegCohort\nfrom mtl_torch.data_loading.medical.conic import ConicDetectionTrainValCohort\nfrom mtl_torch.data_loading.medical.cytology import CytoCohort\nimport numpy as np\nfrom pathlib import Path\nfrom typing_extensions import Annotated\nfrom pydantic import Field\nfrom typing import Dict, List, Optional, Literal, Union, Any\nimport logging\nimport torch\nimport albumentations as A\nimport torch.multiprocessing as mp\nimport torch.distributed as dist\nfrom mtl_torch.interactive import blocks, configs as cfs, data, pipes, tasks, training\nfrom mtl_torch.augmentations import get_weak_default_augs\n\nTASK_TYPE = Literal[\"clf\", \"mclf\", \"seg\", \"det\"]\n\n\nclass TaskBase(cfs.BaseModel):\n    \"\"\"\n    If task_config is left at None, the corresponding template will be chosen\n    \"\"\"\n    task_config: Optional[cfs.MTLTaskConfig] = Field(\n        default=None,\n        description=\"If none, the task config will be created froim the template of this task's type\"\n    )\n    task_type: TASK_TYPE\n    task_id: str\n    name_overwrite: Optional[str] = None\n\n    def build_task(self, env, shared_blocks, free_mem: int) -> tasks.MTLTask:\n        raise NotImplementedError()\n\nclass TaskImgnetClf(TaskBase):\n    task_id: Literal[\"imgnetclf\"] = \"imgnetclf\"\n    task_type: Literal[\"clf\"] = \"clf\"\n    cohort_config: ImgnetClassificationCohort.Config = ImgnetClassificationCohort.Config(\n        src_config=ImageNetSrc.Config(only_for_classes=[\"n12267677\", \"n12144580\"]),\n        batch_size=32)\n\n    def build_cohort(self, env, free_mem: int):\n        return ImgnetClassificationCohort(env.data_root, self.cohort_config)\n\n    def build_task(self, env, shared_blocks, free_mem: int) -> tasks.MTLTask:\n        cohort = self.build_cohort(env, free_mem)\n        self.task_config: cfs.ClassificationTaskConfig\n        return tasks.ClassificationTask(\n            shared_blocks[self.task_config.encoder_key].args.hidden_dim,\n            class_names=cohort.datasets[0].class_names,\n            args=self.task_config,\n            cohort=cohort\n        )\n\n\nclass TaskCocoSeg(TaskBase):\n    task_type: Literal['seg'] = 'seg'\n    task_id: Literal[\"cocoseg\"] = \"cocoseg\"\n    cohort_config: CocoSegCohort.Config = CocoSegCohort.Config(\n        num_workers=2,\n        batch_size=2\n    )\n\n    def build_cohort(self, env, free_mem: int) -> data.TrainValCohort:\n        self.cohort_config.batch_size = 8\n        logging.info(f\"Batchsize of {self.task_type}, {free_mem=}, {self.cohort_config.batch_size=}\")\n        return CocoSegCohort(self.cohort_config, env.data_root / \"coco\")\n\n    def build_task(self, env, shared_blocks, free_mem: int) -> tasks.MTLTask:\n        cohort = self.build_cohort(env, free_mem)\n        self.task_config: tasks.SemSegTask.Config\n        return tasks.SemSegTask(\n            class_names=cohort.datasets[0].class_names,\n            for_decoder=shared_blocks[self.task_config.decoder_key],\n            args=self.task_config,\n            cohort=cohort\n        )\n\n\nclass TaskCocoDet(TaskBase):\n    task_type: Literal['det'] = 'det'\n    task_id: Literal[\"cocodet\"] = \"cocodet\"\n    cohort_config: CocoDetCohort.Config = CocoDetCohort.Config(\n        num_workers=2,\n        batch_size=32\n    )\n\n    def build_cohort(self, env, free_mem: int) -> data.TrainValCohort:\n        self.cohort_config.batch_size = 8\n        logging.info(f\"Batchsize of {self.task_type}, {free_mem=}, {self.cohort_config.batch_size=}\")\n        return CocoDetCohort(self.cohort_config, env.data_root / \"coco\")\n\n    def build_task(self, env, shared_blocks, free_mem: int) -> tasks.MTLTask:\n        cohort = self.build_cohort(env, free_mem)\n        self.task_config: tasks.MMDetectionTask.Config\n        return tasks.MMDetectionTask(\n            args=self.task_config,\n            for_strides=shared_blocks[self.task_config.decoder_key].get_strides(),\n            in_channels=shared_blocks[self.task_config.decoder_key].args.intermediate_channels,\n            cohort=cohort\n        )\n# %%\n\n\nclass Kather100KClf(TaskBase):\n    task_type: Literal[\"clf\"] = \"clf\"\n    task_id: Literal['kather100k'] = 'kather100k'\n    cohort_config: Kather100kClassificationCohortConfig = Kather100kClassificationCohortConfig(\n        num_workers=1,\n        batch_size=32\n    )\n\n    def build_cohort(self, env):\n        return Kather100kClassificationCohort(Path('/kaggle/input/nct-crc-he-100k/NCT-CRC-HE-100K'), self.cohort_config)\n\n    def build_task(self, env, shared_blocks, free_mem: int) -> tasks.MTLTask:\n        cohort = self.build_cohort(env)\n        self.task_config: cfs.ClassificationTaskConfig\n        return tasks.ClassificationTask(\n            hidden_dim=shared_blocks[self.task_config.encoder_key].args.hidden_dim,\n            class_names=cohort.datasets[0].class_names,\n            args=self.task_config,\n            cohort=cohort\n        )\n\n\n# %%\n\n\nclass TaskBrats2020(TaskBase):\n    task_id: Literal[\"brats2020\"] = \"brats2020\"\n    task_type: Literal[\"seg\"] = \"seg\"\n    cohort_config: data.TrainValCohort.Config = data.TrainValCohort.Config(num_workers=1, batch_size=32)\n    subjects_splitting_seed: int = 42\n    train_ds_config: CachingSubCaseDS.Config = CachingSubCaseDS.Config(drain_each_epoch=False, subcase_cache_size=1000)\n\n    def build_cohort(self, env):\n        # Update the dataset path to match your dataset location\n        dataset_path = Path(\"/kaggle/input/brats2020-correct-dataset-training-validation/BraTS2020 Dataset (Training + Validation)\")\n        if dataset_path.exists():\n            print(f\"Dataset path exists: {dataset_path}\")\n            subjects = Brats2020Subjects(dataset_path)\n            print(f\"Number of subjects: {len(subjects)}\")\n        else:\n            raise ValueError(f\"Dataset path does not exist: {dataset_path}\")\n        \n        train_indices, val_indices = data.train_val_split(\n            list(range(len(subjects))),\n            perc=0.8,\n            seed=self.subjects_splitting_seed\n        )\n        \n        if not train_indices:\n            raise ValueError(\"Training split is empty.\")\n        if not val_indices:\n            raise ValueError(\"Validation split is empty.\")\n        \n        train_volumes = Brats2020Volumes(data.TransformedSubset(subjects, indices=train_indices))  # type: ignore\n        val_volumes = Brats2020Volumes(data.TransformedSubset(subjects, indices=val_indices))  # type: ignore\n        \n        class_names = [\"BG\", \"c1\", \"c2\", \"c3\", \"c4\"]\n        return data.TrainValCohort(\n            self.cohort_config,\n            data.SemSegDataset(\n                CachingSubCaseDS(\n                    data.TransformedSubset(\n                        train_volumes,\n                        transform=MRI3DProcessor(\n                            MRI3DProcessor.Config(),\n                            lambda k: MRI3DProcessor.get_3d_rotations(k) + MRI3DProcessor.base_volume_augs(k),\n                            with_segmask=True)),\n                    Brats2020Volumes.extract_slices,\n                    self.train_ds_config\n                ),\n                src_transform=Brats2020Volumes.get_2d_src_transform(train=True),\n                batch_transform=pipes.Alb(get_weak_default_augs()),\n                class_names=class_names\n            ),\n            data.SemSegDataset(\n                CachingSubCaseDS(\n                    data.TransformedSubset(\n                        val_volumes,\n                        transform=MRI3DProcessor(MRI3DProcessor.Config(), None, with_segmask=True)),\n                    Brats2020Volumes.extract_slices,\n                    CachingSubCaseDS.Config(subcase_cache_size=1)\n                ),\n                src_transform=Brats2020Volumes.get_2d_src_transform(train=False),\n                class_names=class_names)\n        )\n\n    def build_task(self, env, shared_blocks, free_mem: int) -> tasks.MTLTask:\n        cohort = self.build_cohort(env)\n        self.task_config: cfs.SemSegTaskConfig\n        return tasks.SemSegTask(\n            class_names=cohort.datasets[0].class_names,\n            for_decoder=shared_blocks[self.task_config.decoder_key],\n            args=self.task_config,\n            cohort=cohort\n        )\n\nTaskConfig = Union[tuple(TaskBase.__subclasses__())]\n\nclass TaskCheXpertMClf(TaskBase):\n    task_id: Literal[\"chexpertmclf\"] = \"chexpertmclf\"\n    task_type: Literal[\"mclf\"] = \"mclf\"\n    cohort_config = CheXpertMultilabelCohort.Config(\n        batch_size=32,\n        num_workers=1\n    )\n\n    def build_cohort(self, env):\n        return CheXpertMultilabelCohort(env.data_root, self.cohort_config)\n\n    def build_task(self, env, shared_blocks, free_mem: int) -> tasks.MTLTask:\n        cohort = self.build_cohort(env)\n        self.task_config: cfs.MultilabelClassificationTaskConfig\n        return tasks.MultilabelClassificationTask(\n            shared_blocks[self.task_config.encoder_key].args.hidden_dim,\n            args=self.task_config,\n            cohort=cohort\n        )\n\n\n# %%\nclass HyperParameters(cfs.ExperimentHyperParameters):\n    wandb_project: Literal[\"universal_pretraining_default\", \"universal_pretraining\",\n                           \"domain_pretraining\", \"universal_singletasks\"] = \"universal_pretraining_default\"\n    wandb_entity = \"tissue-concepts\"\n\n    encoder_args: blocks.PyramidEncoder.Config = blocks.PyramidEncoder.Config(\n        hidden_dim=768,\n        norm_layer=\"layernorm\",\n        model=TorchVisionSwinformer.Config(use_latent_layer=False, pretrained=True, variant=\"small\")\n    )\n    decoder_args: blocks.PyramidDecoder.Config = blocks.PyramidDecoder.Config(\n        model=blocks.SMPUnetDecoder.Config(pixel_embedding_dim=32)\n    )\n    fcosdecoder_args: blocks.FCOSDecoder.Config = blocks.FCOSDecoder.Config()\n    trainer_config = training.MTLTrainer.Config(\n        max_epochs=10,\n        optim=cfs.ExperimentHyperParameters().example_vit_optim,\n        mtl_train_loop=cfs.TrainLoopConfig(\n            max_steps=10,\n            task_sampler=training.CyclicTaskSampler.Config(mode=\"infinite\")\n        ),\n        mtl_val_loop=cfs.ValLoopConfig(\n            max_steps=10,\n            task_sampler=training.CyclicTaskSampler.Config(mode=\"infinite\")\n        )\n    )\n    pretrain_tasks: List[Annotated[TaskConfig, Field(discriminator=\"task_id\")]] = [\n        TaskCocoDet(),\n        TaskBrats2020(),\n        Kather100KClf(),\n        TaskCheXpertMClf(),\n        TaskCocoSeg(),\n        TaskImgnetClf(\n            name_overwrite=f\"imgnet_full\",\n            cohort_config=ImgnetClassificationCohort.Config(\n                num_workers=1,\n                batch_size=32,\n                src_config=ImageNetSrc.Config(only_for_classes=None)\n            )),\n    ]\n\n    pretraining_clf_config: cfs.ClassificationTaskConfig = cfs.ClassificationTaskConfig(\n        module_name=\"clftemplate\"\n    )\n    pretraining_mclf_config: cfs.MultilabelClassificationTaskConfig = cfs.MultilabelClassificationTaskConfig(\n        module_name=\"mclftemplate\"\n    )\n    pretraining_det_config: cfs.MMDetectionTaskConfig = cfs.MMDetectionTaskConfig(\n        module_name=\"dettemplate\",\n        # head_kernel_size=1\n    )\n    pretraining_seg_config: cfs.SemSegTaskConfig = cfs.SemSegTaskConfig(\n        module_name=\"segtemplate\",\n        headdropout=0.2,\n        head_kernel_size=1\n    )\n\n    from_checkpoint_path: str = \"\"\n\n# %%\n\n\ndef build_sh_blocks(\n    encoder_args: blocks.PyramidEncoder.Config,\n    decoder_args: blocks.PyramidDecoder.Config,\n    fcosdecoder_args: blocks.FCOSDecoder.Config\n) -> Dict[str, blocks.SharedBlock]:\n    encoder = blocks.PyramidEncoder(args=encoder_args)\n    decoder = blocks.PyramidDecoder(decoder_args,\n                                    encoder.get_feature_pyramid_channels(),\n                                    encoder.get_strides()[-1])\n    fcosdecoder = blocks.FCOSDecoder(fcosdecoder_args,\n                                     encoder.get_feature_pyramid_channels(),\n                                     encoder.get_strides())\n    return {b.get_name(): b for b in [encoder, decoder, fcosdecoder]}\n\n# %%\n\n\ndef get_task_config_template(config, task_type: Literal[\"clf\", \"mclf\", \"seg\", \"det\"]) -> cfs.MTLTaskConfig:\n    if task_type == \"clf\":\n        return config.pretraining_clf_config\n    elif task_type == \"mclf\":\n        return config.pretraining_mclf_config\n    elif task_type == \"det\":\n        return config.pretraining_det_config\n    elif task_type == \"seg\":\n        return config.pretraining_seg_config\n    else:\n        raise Exception(f\"No default task configs for task type: {task_type}\")\n\n\ndef build_pretraining_tasks(env, config, task: TaskBase, shared_blocks, free_mem: int) -> tasks.MTLTask:\n    if not task.task_config:\n        task.task_config = get_task_config_template(config, task.task_type).copy()\n        task.task_config.module_name = task.task_id if task.name_overwrite is None else task.name_overwrite\n    return task.build_task(env, shared_blocks, free_mem)\n\n\n# %%\ndef prepare_schema(env):\n    HyperParameters.update_schema(env)\n\n# %%\n\n\ndef build_trainer(\n        env,\n        config: HyperParameters,\n        rank: int,\n        world_size: int):\n    wandb_run = config.init_experiment(env, rank, world_size)\n    logfile = Path(wandb_run.dir) / f\"rank{rank}log.txt\"\n    print(f\"Rank {rank} logging to {logfile}\")\n    logging.basicConfig(\n        filename=logfile,\n        level=logging.DEBUG if rank == 0 else logging.INFO,\n        force=True\n    )\n    logging.info(env)\n    trainer = training.MTLTrainer(\n        config.trainer_config,\n        checkpoint_cache_path=env.data_output / config.get_checkpoint_foldername_suggestion(\n            wandb_run,\n            rank,\n            world_size\n        ),\n        rank=rank,\n        world_size=world_size\n    )\n    free_mem_before_sh, _ = torch.cuda.mem_get_info(rank)\n    sh_blocks = build_sh_blocks(config.encoder_args, config.decoder_args, config.fcosdecoder_args)\n    trainer.add_shared_blocks(sh_blocks.values())\n\n    # Get free GPU memory after allocating shared blocks in bytes\n    free_mem, total_mem = torch.cuda.mem_get_info(rank)\n    logging.info(f\"GPU memory on {rank}: {free_mem / (1024**3)=}/{total_mem=} GB free for tasks\")\n    logging.info(f\"GPU memory used by shared blocks: {(free_mem_before_sh - free_mem) / (1024**3)=} GB\")\n\n    for taskconfig in [t for i, t in enumerate(config.pretrain_tasks) if i % world_size == rank]:\n        taskconfig: TaskBase\n        task = build_pretraining_tasks(env, config, taskconfig, trainer.shared_blocks, free_mem)\n        trainer.add_mtl_task(task)\n        torch.cuda.empty_cache()\n        torch.cuda.ipc_collect() \n    logging.info(f\"Trainer has {len(trainer.mtl_tasks)} tasks\")\n\n    if config.from_checkpoint_path and not (trainer.checkpoint_cache_path / \"latest\").exists():\n        trainer.load_checkpoint(\n            Path(config.from_checkpoint_path),\n            load_meta=False,\n            load_tasks=True,\n            load_optim_state=False\n        )\n\n    return trainer\n\n\n# %%\n\ndef start_multigpu(rank: int, world_size: int) -> None:\n    # Dataloaders will spawn processes if this is not set\n    mp.set_start_method(\"fork\", force=True)\n    env = cfs.EnvByConvention(\"universal_pretraining\")\n\n    assert 'MASTER_ADDR' in os.environ, 'MASTER_ADDR not set'\n    assert 'MASTER_PORT' in os.environ, 'MASTER_PORT not set'\n    dist.init_process_group(\"nccl\", rank=rank, world_size=world_size)\n\n    prepare_schema(env)\n    config: HyperParameters = HyperParameters.load_config(env)\n    trainer = build_trainer(env, config, rank, world_size)\n    trainer.fit()\n\n    dist.destroy_process_group()\n\n\n# %%\nif __name__ == \"__main__\":\n    import shutil\n    from mtl_torch.utils import remove_folder_blocking_if_exists\n    env = cfs.EnvByConvention(\"universal_pretraining\")\n    \n    # Remove default checkpoint folder if it exists\n    remove_folder_blocking_if_exists(env.data_output / \"default\")\n\n    # Force using only 1 GPU\n    n_gpus =  torch.cuda.device_count()\n\n    if n_gpus > 1:\n        os.environ['MASTER_ADDR'] = '0.0.0.0'\n        os.environ['MASTER_PORT'] = '12356'\n        logging.info(f\"Running on {n_gpus} GPUs.\")\n        mp.spawn(start_multigpu,\n                 args=(n_gpus,),\n                 nprocs=n_gpus,\n                 start_method=\"spawn\",\n                 join=True)\n    else:\n        os.environ['WANDB_MODE']=\"offline\"\n        torch.cuda.set_device(0)  # Ensure GPU 0 is used\n        prepare_schema(env)\n        config: HyperParameters = HyperParameters.load_config(env)\n        trainer = build_trainer(env, config, 0, 1)\n        trainer.fit()\n\n# %%","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-14T09:40:36.038569Z","iopub.execute_input":"2025-03-14T09:40:36.039622Z","iopub.status.idle":"2025-03-14T09:40:37.060236Z","shell.execute_reply.started":"2025-03-14T09:40:36.039582Z","shell.execute_reply":"2025-03-14T09:40:37.058402Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install pathos","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-14T09:37:05.992337Z","iopub.execute_input":"2025-03-14T09:37:05.992949Z","iopub.status.idle":"2025-03-14T09:37:09.88863Z","shell.execute_reply.started":"2025-03-14T09:37:05.992906Z","shell.execute_reply":"2025-03-14T09:37:09.887359Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"pip install pytorch-lightning==1.8.6","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-14T08:47:01.047846Z","iopub.execute_input":"2025-03-14T08:47:01.048226Z","iopub.status.idle":"2025-03-14T08:47:06.813838Z","shell.execute_reply.started":"2025-03-14T08:47:01.048196Z","shell.execute_reply":"2025-03-14T08:47:06.812759Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!python /kaggle/input/train/pytorch/default/1/train.py","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-14T09:25:37.632962Z","iopub.execute_input":"2025-03-14T09:25:37.633387Z","iopub.status.idle":"2025-03-14T09:25:38.372274Z","shell.execute_reply.started":"2025-03-14T09:25:37.633356Z","shell.execute_reply":"2025-03-14T09:25:38.370975Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"pip install huggingface_hub==0.15.1","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-14T08:48:52.023258Z","iopub.execute_input":"2025-03-14T08:48:52.024011Z","iopub.status.idle":"2025-03-14T08:48:55.863929Z","shell.execute_reply.started":"2025-03-14T08:48:52.02398Z","shell.execute_reply":"2025-03-14T08:48:55.86283Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from torch.optim.lr_scheduler import LRScheduler\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-14T08:42:24.16298Z","iopub.execute_input":"2025-03-14T08:42:24.163766Z","iopub.status.idle":"2025-03-14T08:42:24.204988Z","shell.execute_reply.started":"2025-03-14T08:42:24.163734Z","shell.execute_reply":"2025-03-14T08:42:24.203797Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(np.unique(metrics[\"targets\"][:, i], return_counts=True))\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-14T08:34:23.596229Z","iopub.status.idle":"2025-03-14T08:34:23.59656Z","shell.execute_reply.started":"2025-03-14T08:34:23.596413Z","shell.execute_reply":"2025-03-14T08:34:23.596435Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nos.environ[\"PYTORCH_CUDA_ALLOC_CONF\"] = \"max_split_size_mb:64\"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\n\ntorch.cuda.empty_cache()  # Releases all unoccupied cached memory\ntorch.cuda.ipc_collect()  # Helps reclaim unused memory from CUDA IPC\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"help(trainer)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"trainer","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"config","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"help(trainer)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install --upgrade ipywidgets","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import ipywidgets as widgets\nwidgets.HBox([widgets.IntSlider(), widgets.IntSlider()])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-14T07:58:01.799125Z","iopub.execute_input":"2025-03-14T07:58:01.799739Z","iopub.status.idle":"2025-03-14T07:58:01.881199Z","shell.execute_reply.started":"2025-03-14T07:58:01.799708Z","shell.execute_reply":"2025-03-14T07:58:01.880327Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install ipywidgets\n!jupyter nbextension enable --py widgetsnbextension","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!jupyter lab build\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!ls /kaggle/input/coco-2017-dataset/coco/annotations","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import wandb\nwandb.finish()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from mtl_torch.interactive.configs import TaskConfig","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"dir(mtl_torch.interactive.configs)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!ls /data_root/coco","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from pathlib import Path\n\nenv.data_root = Path('/kaggle/input/model/pytorch/default/1/original_version')\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!ls /data_root/coco/train2017","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Create the required directories\n!mkdir -p /data_root/coco/annotations\n!mkdir -p /data_root/coco/train2017\n!mkdir -p /data_root/coco/val2017","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"env.interactive_environment=True","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"env","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nfrom mtl_torch.interactive import blocks, configs as cfs, data, pipes, tasks, training\n\n# Set a default empty JSON string for the missing config\nos.environ[\"MLOPS_JSON_universal_pretraining\"] = \"{}\"\n# Now run your script\n\n\nimport json\n\ndefault_config = {}  # Define your default hyperparameters here\nwith open(\"job_configs/universal_pretraining.jsonc\", \"w\") as f:\n    json.dump(default_config, f)\n\n\nos.environ[\"MLOPS_JSON_training\"] = \"{}\"\n\nschema_dir = \"/kaggle/working/.vscode\"\nos.makedirs(schema_dir, exist_ok=True)  # Ensure the directory exists\n\n# Modify the path in the script\nenv.env_name = \"universal_pretraining\"  # Ensure env_name is set\nschema_path = f\"/kaggle/working/.vscode/{env.env_name}.schema.json\"\nprint(f\"Schema will be saved at: {schema_path}\")\n\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Eval.ipynb","metadata":{}},{"cell_type":"code","source":"import logging\nimport os\nfrom abc import abstractmethod\nfrom pydantic import Field\nfrom typing import Dict, List, Optional, Callable, Literal, Union, Any\nfrom typing_extensions import Annotated\nfrom pathlib import Path\nimport albumentations as A\nimport torchvision.transforms as T\nimport torch.nn as nn\nimport wandb\nimport numpy as np\nimport torch.nn.functional as F\nimport copy\n\nimport torch\n\nfrom mtl_torch.interactive import blocks, configs as cfs, data, pipes, tasks, training\nfrom universal_pretraining import build_sh_blocks, HyperParameters as UniversalPretrainingHyperParameters\n\nenv = cfs.EnvByConvention(env_name=\"singletask\")\nprint(env)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class Evaltask(cfs.BaseModel):\n    accumulationsteps: int = 1\n    fixed_spatial_size: int | None = None\n    cohort_config: data.TrainValCohort.Config\n\n    @abstractmethod\n    def build_cohort(self):\n        raise NotImplementedError\n\n    def build_eval_cohort(self):\n        cohort = self.build_cohort()\n        if self.fixed_spatial_size is not None:\n            # To the existing datasets we add a transform that resizes the images\n            cohort.datasets[0].batch_transform = T.Compose([\n                cohort.datasets[0].batch_transform if cohort.datasets[0].batch_transform is not None else nn.Identity(),\n                pipes.Alb(\n                    [A.Resize(self.fixed_spatial_size, self.fixed_spatial_size)],\n                    support_boxes=isinstance(cohort.datasets[0], data.DetectionDataset)\n                )\n            ])\n            cohort.datasets[1].batch_transform = T.Compose([\n                cohort.datasets[1].batch_transform if cohort.datasets[1].batch_transform is not None else nn.Identity(),\n                pipes.Alb(\n                    [A.Resize(self.fixed_spatial_size, self.fixed_spatial_size)],\n                    support_boxes=isinstance(cohort.datasets[1], data.DetectionDataset)\n                )\n            ])\n        return cohort\n\n    def build_task(self, shared_blocks) -> tasks.MTLTask:\n        cohort = self.build_eval_cohort()\n        # The child class should have a task_config attribute\n        self.task_config: Any\n        if isinstance(cohort.datasets[0], data.ClassificationDataset):\n            return tasks.ClassificationTask(\n                shared_blocks[self.task_config.encoder_key].args.hidden_dim,\n                class_names=cohort.datasets[0].class_names,\n                args=self.task_config,\n                cohort=cohort\n            )\n        elif isinstance(cohort.datasets[0], data.SemSegDataset):\n            return tasks.SemSegTask(\n                class_names=cohort.datasets[0].class_names,\n                for_decoder=shared_blocks[self.task_config.decoder_key],\n                args=self.task_config,\n                cohort=cohort\n            )\n        elif isinstance(cohort.datasets[0], data.DetectionDataset):\n            return tasks.MMDetectionTask(\n                args=self.task_config,\n                for_strides=shared_blocks[self.task_config.decoder_key].get_strides(),\n                in_channels=shared_blocks[self.task_config.decoder_key].args.intermediate_channels,\n                cohort=cohort\n            )\n        else:\n            raise NotImplementedError\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Chest X-ray","metadata":{},"attachments":{}},{"cell_type":"code","source":"from mtl_torch.data_loading.medical.chestxray import SIIMCovidClf, SIIMPneumothoraxSegCohort\n\nclass TaskSIIMPneumothoraxSegCohort(Evaltask):\n    dataset_id: Literal[\"ssimpneumothoraxseg\"] = \"ssimpneumothoraxseg\"\n    task_type: Literal[\"seg\"] = \"seg\"\n    cohort_config = SIIMPneumothoraxSegCohort.Config(batch_size=32)\n    task_config: tasks.SemSegTask.Config = tasks.SemSegTask.Config(\n        module_name=\"ssimpneumothoraxseg\",\n        head_kernel_size=1,\n        headdropout=0.2\n    )\n\n    def build_cohort(self):\n        return SIIMPneumothoraxSegCohort(env.data_root / \"siim-acr-pneumothorax-segmentation\", self.cohort_config)\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from mtl_torch.data_loading.medical.chestxray import TuberculosisShenzhenClfCohort\n\n\nclass TaskTuberculosisShenzhen(Evaltask):\n    dataset_id: Literal['tuberculosisshenzhenfinding'] = 'tuberculosisshenzhenfinding'\n    task_type: Literal[\"clf\"] = \"clf\"\n    cohort_config: TuberculosisShenzhenClfCohort.Config = TuberculosisShenzhenClfCohort.Config(\n        target_column=\"findings\",\n        batch_size=32\n    )\n    task_config: tasks.ClassificationTask.Config = tasks.ClassificationTask.Config(\n        module_name=\"tuberculosisshenzhenfinding\"\n    )\n\n    def build_cohort(self):\n        return TuberculosisShenzhenClfCohort(env.data_root, self.cohort_config)\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from mtl_torch.data_loading.medical.chestxray import ChestXRay2017Cohort\n\nclass TaskChestXRay2017(Evaltask):\n    dataset_id: Literal['chestxray2017'] = 'chestxray2017'\n    task_type: Literal[\"clf\"] = \"clf\"\n    cohort_config: ChestXRay2017Cohort.Config = ChestXRay2017Cohort.Config(batch_size=32)\n    task_config: tasks.ClassificationTask.Config = tasks.ClassificationTask.Config(\n        module_name=\"chestxray2017\"\n    )\n\n    def build_cohort(self):\n        return ChestXRay2017Cohort(env.data_root, self.cohort_config)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Tomographic","metadata":{},"attachments":{}},{"cell_type":"code","source":"from mtl_torch.data_loading.medical.msd import BrainTumorClf, Amos22Cohort\nfrom mtl_torch.torch_ext import CachingSubCaseDS\n\nclass TaskBrainTumorClf(Evaltask):\n    dataset_id: Literal['braintumorclf'] = 'braintumorclf'\n    task_type: Literal[\"clf\"] = \"clf\"\n    cohort_config: data.TrainValCohort.Config = data.TrainValCohort.Config(num_workers=5, batch_size=32)\n    task_config: tasks.ClassificationTask.Config = tasks.ClassificationTask.Config(\n        module_name=\"braintumorclf\"\n    )\n\n    def build_cohort(self):\n        root = env.data_root / \"brain_tumor_kagglesets\" / \"brain-tumor-mri-dataset\"\n        return BrainTumorClf.build_recommended_downstream_cohort(root, self.cohort_config, train_aug=True)\n\n\nclass TaskAmos22(Evaltask):\n    accumulationsteps: int = 2\n    dataset_id: Literal[\"amos22ct\"] = \"amos22ct\"\n    task_type: Literal[\"seg\"] = \"seg\"\n    cohort_config: Amos22Cohort.Config = Amos22Cohort.Config(\n        only_modality=\"ct\",\n        batch_size=2, # enable training on the small GPUs\n        num_workers=1,\n        train_cache=CachingSubCaseDS.Config(\n            drain_each_epoch=True,\n            subcase_cache_size=500\n        )\n    )\n    task_config: tasks.SemSegTask.Config = tasks.SemSegTask.Config(\n        module_name=\"amosct\",\n        head_kernel_size=1,\n        headdropout=0.2\n    )\n\n    def build_cohort(self):\n        return Amos22Cohort(self.cohort_config, env.data_root / \"amos22\")\n\nclass TaskAmos22MRI(Evaltask):\n    accumulationsteps: int = 2\n    dataset_id: Literal[\"amos22mri\"] = \"amos22mri\"\n    task_type: Literal[\"seg\"] = \"seg\"\n    cohort_config: Amos22Cohort.Config = Amos22Cohort.Config(\n        only_modality=\"mri\",\n        batch_size=2, # enable training on the small GPUs\n        num_workers=1,\n        train_cache=CachingSubCaseDS.Config(\n            drain_each_epoch=True,\n            subcase_cache_size=500\n        )\n    )\n    task_config: tasks.SemSegTask.Config = tasks.SemSegTask.Config(\n        module_name=\"amos22mri\",\n        head_kernel_size=5,\n        headdropout=0.2\n    )\n\n    def build_cohort(self):\n        return Amos22Cohort(self.cohort_config, env.data_root / \"amos22\")\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Histo","metadata":{},"attachments":{}},{"cell_type":"code","source":"from mtl_torch.data_loading.medical.histo_patchclassification import Kather100kClassificationCohort, Kather100kClassificationCohortConfig\n\n\nclass Kather100KClfTaskConfig(Evaltask):\n    dataset_id: Literal['kather100k'] = 'kather100k'\n    task_type: Literal[\"clf\"] = \"clf\"\n    cohort_config: Kather100kClassificationCohortConfig = Kather100kClassificationCohortConfig(\n        num_workers=5,\n        batch_size=32)\n    task_config: tasks.ClassificationTask.Config = tasks.ClassificationTask.Config(module_name=\"kather100k\")\n\n    def build_cohort(self):\n        return Kather100kClassificationCohort(env.histo_root, self.cohort_config)\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from histo_data.bach import Bach2018ClassificationCohort, Bach2018ClassificationCohortConfig\n\n\nclass BachTaskConfig(Evaltask):\n    dataset_id: Literal['bach'] = 'bach'\n    task_type: Literal[\"clf\"] = \"clf\"\n    cohort_config: Bach2018ClassificationCohortConfig = Bach2018ClassificationCohortConfig(\n        num_workers=4)\n    task_config: tasks.ClassificationTask.Config = tasks.ClassificationTask.Config(module_name=\"bach\")\n\n    def build_cohort(self):\n        return Bach2018ClassificationCohort(env.histo_root, self.cohort_config)\n\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from histo_data.breakhis import BreakHisCohort, BreakHisCohortConfig\n\n\nclass BreakHisTaskConfig(Evaltask):\n    dataset_id: Literal['breakhis'] = 'breakhis'\n    task_type: Literal[\"clf\"] = \"clf\"\n    cohort_config: BreakHisCohortConfig = BreakHisCohortConfig(\n        batch_size=32,\n        num_workers=5)\n    task_config: tasks.ClassificationTask.Config = tasks.ClassificationTask.Config(module_name=\"breakhis\")\n\n    def build_cohort(self):\n        return BreakHisCohort(env.histo_root, self.cohort_config)\n\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Parameters","metadata":{},"attachments":{}},{"cell_type":"code","source":"class FinetuneConfig(cfs.BaseModel):\n    stages: List[Literal[\"frozen\", \"unfrozen\"]] = [\"frozen\", \"unfrozen\"]\n    train_fraction: float = 1.\n    reduce_seed: int = 42\n\n\nEvalMode = FinetuneConfig\nEvaltasks = Union[tuple(Evaltask.__subclasses__())]  # type:ignore","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"DEFAULT_EPOCHS = 200\n# With small training datasets only the validation epochs are expensive, so reduce them\nVAL_EPOCH_DEFAULTS = [\n    0, \n    1, \n    DEFAULT_EPOCHS // 2 - 2, \n    DEFAULT_EPOCHS // 2 - 1, \n    DEFAULT_EPOCHS // 2, \n    DEFAULT_EPOCHS // 2 + 1, \n    DEFAULT_EPOCHS // 2 + 2, \n    DEFAULT_EPOCHS - 2,\n    DEFAULT_EPOCHS - 1,\n    DEFAULT_EPOCHS]\nVAL_EPOCH_DEFAULTS","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class HyperParameters(cfs.ExperimentHyperParameters):\n    wandb_project: Literal[\"evaldump\", \"univeval\"] = \"evaldump\"\n    # wandb_entity = \"tissue-concepts\"\n\n    encoder_args: blocks.PyramidEncoder.Config = UniversalPretrainingHyperParameters().encoder_args\n    decoder_args: blocks.PyramidDecoder.Config = UniversalPretrainingHyperParameters().decoder_args\n    fcosdecoder_args: blocks.FCOSDecoder.Config = UniversalPretrainingHyperParameters().fcosdecoder_args\n\n    trainer_config = training.MTLTrainer.Config(\n        max_epochs=DEFAULT_EPOCHS,\n        optim=cfs.MTLOptimizerConfig(\n            optim_config=cfs.OptimizerAdamWConfig(lr=1e-5),\n            lr_factor_tasks=10,\n        ),\n        mtl_validation_selector=cfs.FixedEventSelector(at_iterations=VAL_EPOCH_DEFAULTS),\n    )\n\n    task: Annotated[Evaltasks, Field(discriminator=\"dataset_id\")] = TaskTuberculosisShenzhen()\n\n    from_checkpoint_path: str = Field(\n        default=\"\",\n        description=\"Checkpoint to load from. If None, the model will be initialized according to config.\"\n    )\n\n    # eval_mode: Annotated[EvalMode, Field(discriminator=\"eval_type\")] = FinetuneConfig()\n    eval_mode: FinetuneConfig = FinetuneConfig()\n\n    # for grouping in wandb\n    experimentid: str = \"\"\n    modelid: str = \"\"\n\n\nHyperParameters.update_schema(env)\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"config = HyperParameters.load_config(env)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### W&B link","metadata":{},"attachments":{}},{"cell_type":"code","source":"wandb_run = config.init_experiment(env)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"shared_blocks = build_sh_blocks(config.encoder_args, config.decoder_args, config.fcosdecoder_args)\n\ntrainer = training.MTLTrainer(\n    config.trainer_config,\n    checkpoint_cache_path=env.data_output / wandb_run.id\n)\ntrainer.args.mtl_train_loop.multistep_mode.factor = config.task.accumulationsteps\ntrainer.stages = config.eval_mode.stages\ntrainer.add_shared_blocks(shared_blocks.values())","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Just load this everytime, if there is a latest checkpoint it will be loaded later\nif config.from_checkpoint_path:\n    trainer.load_checkpoint(\n        Path(config.from_checkpoint_path),\n        load_meta=False,\n        load_tasks=False,\n        load_optim_state=False\n    )","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"evaltask = config.task.build_task(shared_blocks)\nif config.eval_mode.train_fraction < 1.:\n    evaltask.cohort.datasets[0].set_indices_by_fraction(\n        config.eval_mode.train_fraction, \n        seed=config.eval_mode.reduce_seed\n    )","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"_ = trainer.add_mtl_task(evaltask)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def fine_tune_task(i: int, stage: Literal[\"frozen\", \"unfrozen\"]):\n    # Every stage gets min epochs\n    if trainer.args.early_stopping is not None:\n        logging.info(f\"Setting {trainer.args.early_stopping.min_train_loops=} to {trainer._epoch + trainer.args.early_stopping.min_train_loops}\")\n        trainer.args.early_stopping.min_train_loops = trainer._epoch + trainer.args.early_stopping.min_train_loops\n    \n    if stage == \"frozen\":\n        for _, block in [(\"\", shared_blocks[\"encoder\"])]: # shared_blocks.items():\n            logging.warning(f\"Froze {block.get_name()}\")\n            block.freeze_all_parameters(True)\n    else:\n        for _, block in shared_blocks.items():\n            logging.warning(f\"Unfroze {block.get_name()}\")\n            block.freeze_all_parameters(False)\n    \n    max_loops_for_stage = config.trainer_config.max_epochs // len(config.eval_mode.stages)\n    logging.info(f\"Starting stage {stage} for {trainer.mtl_tasks[0].get_name()} with {max_loops_for_stage=}\")\n\n    if evaltask.cohort.datasets[0].get_dataset_style() is data.DatasetStyle.MapStyle:\n        trainlen = len(trainer.mtl_tasks[0].cohort.datasets[0])\n        vallen = len(trainer.mtl_tasks[0].cohort.datasets[1])\n    else:\n        trainlen = len(trainer.mtl_tasks[0].cohort.datasets[0].src_ds.supercase_ds) # type:ignore\n        vallen = len(trainer.mtl_tasks[0].cohort.datasets[1].src_ds.supercase_ds) # type:ignore\n        \n    logging.info(f\"Train: {trainlen}, {trainer.mtl_tasks[0].cohort.datasets[0].get_dataset_style()=}\")\n    logging.info(f\"Test: {vallen}, {trainer.mtl_tasks[0].cohort.datasets[1].get_dataset_style()=}\")\n    wandb.log({\"stage\": stage, \"stage_num\": i})\n    trainer._stage_index = i\n    stage_until_epoch = max_loops_for_stage * (i + 1)\n    logging.info(f\"Training this stage until epoch {stage_until_epoch}\")\n    if trainer.mtl_optimizer is None:\n        trainer.init_optimizer()\n    # This avoids loading the last checkpoint which overwrites the stage index\n    trainer._fit(stage_until_epoch)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if wandb_run.resumed:\n    trainer.load_latest_checkpoint()\nfor i, stage in enumerate(config.eval_mode.stages):\n    # If the job is resumed, we need to skip stages until the latest one\n    if trainer._stage_index > i:\n        logging.info(f\"Skipping stage {stage} {i}\")\n        continue\n    fine_tune_task(i, stage)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"wandb.finish()","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}