{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Baseline with Lightning ⚡ MONAI + ResNet3D\n\nThis id followu-up of EDA: https://www.kaggle.com/code/jirkaborovec/spine-fracture-eda-loading-dicom-3d-browse\n\nand my previous 3D classification: https://www.kaggle.com/code/jirkaborovec/brain-tumor-classif-lightning-monai-resnet3d","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19"}},{"cell_type":"code","source":"!pip install -q kaggle_vol3d_classify -f ../input/cervical-spine-fracture-detection-npz-3d-volumes/frozen_packages --no-index\n# !pip install -qU \"pytorch-lightning>1.5.0\" --no-index\n!pip uninstall -y torchtext\n!pip list | grep -e lightning -e kaggle -e monai","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2023-10-27T13:27:27.89602Z","iopub.execute_input":"2023-10-27T13:27:27.896389Z","iopub.status.idle":"2023-10-27T13:28:10.818195Z","shell.execute_reply.started":"2023-10-27T13:27:27.896307Z","shell.execute_reply":"2023-10-27T13:28:10.817129Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%matplotlib inline\n\nimport os\nimport glob\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport torch\n\nPATH_DATASET = \"/kaggle/input/rsna-2022-cervical-spine-fracture-detection\"\nPATH_VOLUME_NPZ = \"/kaggle/input/cervical-spine-fracture-detection-npz-3d-volumes/train_volumes\"","metadata":{"execution":{"iopub.status.busy":"2023-10-27T13:28:10.820613Z","iopub.execute_input":"2023-10-27T13:28:10.821344Z","iopub.status.idle":"2023-10-27T13:28:12.457162Z","shell.execute_reply.started":"2023-10-27T13:28:10.821299Z","shell.execute_reply":"2023-10-27T13:28:12.456096Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_train = pd.read_csv(os.path.join(PATH_DATASET, \"train.csv\"))\ndisplay(df_train.head())","metadata":{"execution":{"iopub.status.busy":"2023-10-27T13:28:12.458588Z","iopub.execute_input":"2023-10-27T13:28:12.459295Z","iopub.status.idle":"2023-10-27T13:28:12.494307Z","shell.execute_reply.started":"2023-10-27T13:28:12.459254Z","shell.execute_reply":"2023-10-27T13:28:12.49343Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Converting volumes NPZ to PT","metadata":{}},{"cell_type":"code","source":"from tqdm.auto import tqdm\nfrom joblib import Parallel, delayed\n\n! mkdir /tmp/train_volumes\n\ndef convert_npz2torch(p_npz, out_dir: str = \"/tmp/train_volumes\"):\n    vol = np.load(p_npz)['arr_0']\n    name, _ = os.path.splitext(os.path.basename(p_npz))\n    torch.save(torch.tensor(vol), os.path.join(out_dir, f\"{name}.pt\"))\n\nls_npz = glob.glob(os.path.join(PATH_VOLUME_NPZ, \"*.npz\"))\nprint(f\"found: {len(ls_npz)}\")\n_= Parallel(n_jobs=4)(delayed(convert_npz2torch)(p_npz) for p_npz in tqdm(ls_npz))","metadata":{"execution":{"iopub.status.busy":"2023-10-27T13:28:12.496908Z","iopub.execute_input":"2023-10-27T13:28:12.497191Z","iopub.status.idle":"2023-10-27T13:30:12.20815Z","shell.execute_reply.started":"2023-10-27T13:28:12.497162Z","shell.execute_reply":"2023-10-27T13:30:12.206824Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Conver Test cases\n\nsee: https://www.kaggle.com/code/jirkaborovec/spine-fracture-convert-dicom-imgs-3d-volume","metadata":{}},{"cell_type":"code","source":"import cv2\nimport pydicom\nfrom PIL import Image\nfrom dipy.denoise.nlmeans import nlmeans\nfrom dipy.denoise.noise_estimate import estimate_sigma\nfrom pydicom.pixel_data_handlers import apply_voi_lut\nfrom kaggle_volclassif.utils import interpolate_volume\nfrom skimage import exposure\n\ndef convert_volume(dir_path: str, out_dir: str = \"test_volumes\", size = (224, 224, 224)):\n    ls_imgs = glob.glob(os.path.join(dir_path, \"*.dcm\"))\n    ls_imgs = sorted(ls_imgs, key=lambda p: int(os.path.splitext(os.path.basename(p))[0]))\n\n    imgs = []\n    for p_img in ls_imgs:\n        dicom = pydicom.dcmread(p_img)\n        img = apply_voi_lut(dicom.pixel_array, dicom)\n        img = cv2.resize(img, size[:2], interpolation=cv2.INTER_LINEAR)\n        imgs.append(img.tolist())\n    vol = torch.tensor(imgs, dtype=torch.float32)\n\n    vol = (vol - vol.min()) / float(vol.max() - vol.min())\n    vol = interpolate_volume(vol, size).numpy()\n    \n    # https://scikit-image.org/docs/stable/auto_examples/color_exposure/plot_adapt_hist_eq_3d.html\n    vol = exposure.equalize_adapthist(vol, kernel_size=np.array([64, 64, 64]), clip_limit=0.01)\n    # vol = exposure.equalize_hist(vol)\n    vol = np.clip(vol * 255, 0, 255).astype(np.uint8)\n    \n    path_pt = os.path.join(out_dir, f\"{os.path.basename(dir_path)}.pt\")\n    torch.save(torch.tensor(vol), path_pt)","metadata":{"execution":{"iopub.status.busy":"2023-10-27T13:30:12.209877Z","iopub.execute_input":"2023-10-27T13:30:12.210202Z","iopub.status.idle":"2023-10-27T13:30:13.077727Z","shell.execute_reply.started":"2023-10-27T13:30:12.210169Z","shell.execute_reply":"2023-10-27T13:30:13.076693Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"! rm -rf /tmp/test_volumes\n! mkdir /tmp/test_volumes\n\nls_dirs = [p for p in glob.glob(os.path.join(PATH_DATASET, \"test_images\", \"*\")) if os.path.isdir(p)]\nprint(f\"volumes: {len(ls_dirs)}\")\n\n_= Parallel(n_jobs=3)(delayed(convert_volume)(p_dir, out_dir=\"/tmp/test_volumes\") for p_dir in tqdm(ls_dirs))\n\n! ls -lh /tmp/test_volumes","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2023-10-27T13:30:13.07922Z","iopub.execute_input":"2023-10-27T13:30:13.079629Z","iopub.status.idle":"2023-10-27T13:30:55.780151Z","shell.execute_reply.started":"2023-10-27T13:30:13.079589Z","shell.execute_reply":"2023-10-27T13:30:55.778876Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Preparing Dataset with volumes ⛱️","metadata":{}},{"cell_type":"code","source":"import os\nfrom typing import Union, Optional, Tuple\nfrom torch.utils.data import Dataset\n\n\nclass SpineScansDataset(Dataset):\n\n    def __init__(\n        self,\n        volume_dir: str = 'train_volumes',\n        df_table: Union[str, pd.DataFrame] = 'train.csv',\n        mode: str = 'train',\n        split: float = 0.8,\n        in_memory: bool = False,\n        random_state=42,\n    ):\n        self.volume_dir = volume_dir\n        self.mode = mode\n        self.in_memory = in_memory\n\n        # set or load the config table\n        if isinstance(df_table, pd.DataFrame):\n            self.table = df_table\n        elif isinstance(df_table, str):\n            assert os.path.isfile(df_table), f\"missing file: {df_table}\"\n            self.table = pd.read_csv(df_table)\n        else:\n            raise ValueError(f'unrecognised input for DataFrame/CSV: {df_table}')\n\n        # shuffle data\n        self.table = self.table.sample(frac=1, random_state=random_state).reset_index(drop=True)\n\n        # split dataset\n        assert 0.0 <= split <= 1.0, f\"split {split} is out of range\"\n        frac = int(split * len(self.table))\n        self.table = self.table[:frac] if mode == 'train' else self.table[frac:]\n\n        # populate images/labels\n        self.label_names = sorted([c for c in self.table.columns if c.startswith(\"C\")])\n        self.labels = self.table[self.label_names].values if self.label_names else [None] * len(self.table)\n        self.volumes = [os.path.join(volume_dir, f\"{row['StudyInstanceUID']}.pt\") for _, row in self.table.iterrows()]\n        assert len(self.volumes) == len(self.labels)\n\n    def __getitem__(self, idx: int) -> dict:\n        label = self.labels[idx]\n        vol_ = self.volumes[idx]\n        if isinstance(vol_, str):\n            try:\n                vol = torch.load(vol_).to(torch.float32)\n            except (EOFError, RuntimeError):\n                print(f\"failed loading: {vol_}\")\n        else:\n            vol = vol_\n        if self.in_memory:\n            self.volumes[idx] = vol\n        # in case of predictions, return image name as label\n        label = label if label is not None else vol_\n        return {\"data\": vol.unsqueeze(0), \"label\": label}\n\n    def __len__(self) -> int:\n        return len(self.volumes)","metadata":{"execution":{"iopub.status.busy":"2023-10-27T13:30:55.78205Z","iopub.execute_input":"2023-10-27T13:30:55.782474Z","iopub.status.idle":"2023-10-27T13:30:55.803386Z","shell.execute_reply.started":"2023-10-27T13:30:55.78242Z","shell.execute_reply":"2023-10-27T13:30:55.802573Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def show_sample(vol, lbs):\n    vol = vol[0]\n    fig, axarr = plt.subplots(ncols=3, figsize=(9, 3))\n    print(f\"volume: {vol.shape} with labels: {lbs}\")\n    v_z, v_y, v_x = vol.shape\n    axarr[0].imshow(vol[:, :, v_x // 2], cmap=\"gray\")\n    axarr[1].imshow(vol[v_z // 2, :, :], cmap=\"gray\")\n    axarr[2].imshow(vol[:, v_y // 2, :], cmap=\"gray\")\n    fig.suptitle(lbs)\n    fig.tight_layout()\n    \nds = SpineScansDataset(volume_dir=\"/tmp/train_volumes\", df_table=df_train)\nfor i in range(2):\n    spl = ds[i * 10]\n    show_sample(spl[\"data\"], spl[\"label\"])","metadata":{"execution":{"iopub.status.busy":"2023-10-27T13:30:55.804772Z","iopub.execute_input":"2023-10-27T13:30:55.805072Z","iopub.status.idle":"2023-10-27T13:30:57.097107Z","shell.execute_reply.started":"2023-10-27T13:30:55.805044Z","shell.execute_reply":"2023-10-27T13:30:57.096136Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Lightning 🗃️ DataModule","metadata":{}},{"cell_type":"code","source":"import logging\nfrom functools import partial\nfrom typing import Union, Optional, Tuple, Dict, Any\nfrom pytorch_lightning import LightningDataModule\nfrom rising.loading import DataLoader\nfrom kaggle_volclassif.transforms import rising_resize\n\n\nclass SpineScansDM(LightningDataModule):\n\n    def __init__(\n        self,\n        data_dir: str = '.',\n        path_csv: str = 'train.csv',\n        vol_size: Optional[Tuple[int, int, int]] = (128, 128, 128),\n        in_memory: bool = False,\n        batch_size: int = 4,\n        num_workers: Optional[int] = None,\n        train_transforms=None,\n        valid_transforms=None,\n        split: float = 0.8,\n        **kwargs_dataloader,\n    ):\n        super().__init__()\n        # path configurations\n        assert os.path.isdir(data_dir), f\"missing folder: {data_dir}\"\n        self.train_dir = os.path.join(data_dir, 'train_volumes')\n        self.test_dir = os.path.join(data_dir, 'test_volumes')\n\n        if not os.path.isfile(path_csv):\n            path_csv = os.path.join(data_dir, path_csv)\n        assert os.path.isfile(path_csv), f\"missing table: {path_csv}\"\n        self.path_csv = path_csv\n\n        # other configs\n        self.vol_size = vol_size\n        self.batch_size = batch_size\n        self.split = split\n        self.in_memory = in_memory\n        self.num_workers = num_workers if num_workers is not None else os.cpu_count()\n        self.kwargs_dataloader = kwargs_dataloader\n\n        # need to be filled in setup()\n        self.test_table = []\n        self.train_dataset = None\n        self.valid_dataset = None\n        self.test_dataset = None\n        self.label_names = {}\n        self.train_transforms = train_transforms\n        self.valid_transforms = valid_transforms\n\n    @property\n    def dl_defaults(self) -> Dict[str, Any]:\n        return dict(\n            batch_size=self.batch_size,\n            num_workers=self.num_workers,\n            sample_transforms=partial(rising_resize, size=self.vol_size),\n        )\n\n    @property\n    def num_labels(self) -> int:\n        return len(self.label_names)\n\n    def setup(self, *_, **__) -> None:\n        \"\"\"Prepare datasets\"\"\"\n        ds_training = dict(\n            volume_dir=self.train_dir,\n            df_table=self.path_csv,\n            split=self.split,\n            in_memory=self.in_memory,\n        )\n        self.train_dataset = SpineScansDataset(**ds_training, mode='train')\n        logging.info(f\"training dataset: {len(self.train_dataset)}\")\n        self.valid_dataset = SpineScansDataset(**ds_training, mode='valid')\n        logging.info(f\"validation dataset: {len(self.valid_dataset)}\")\n        self.label_names = sorted(set(self.train_dataset.label_names + self.valid_dataset.label_names))\n\n        if not os.path.isdir(self.test_dir):\n            logging.warning(f\"Missing test folder: {self.test_dir}\")\n            return\n        ls_cases = [os.path.basename(p) for p in glob.glob(os.path.join(self.test_dir, '*'))]\n        self.test_table = [dict(StudyInstanceUID=os.path.splitext(n)[0]) for n in ls_cases]\n        self.test_dataset = SpineScansDataset(\n            self.test_dir,\n            df_table=pd.DataFrame(self.test_table),\n            split=0,\n            mode='test',\n        )\n        logging.info(f\"test dataset: {len(self.test_dataset)}\")\n\n    def train_dataloader(self) -> DataLoader:\n        return DataLoader(\n            self.train_dataset,\n            shuffle=True,\n            batch_transforms=self.train_transforms,\n            **self.dl_defaults,\n            **self.kwargs_dataloader,\n        )\n\n    def val_dataloader(self) -> DataLoader:\n        return DataLoader(\n            self.valid_dataset,\n            shuffle=False,\n            batch_transforms=self.valid_transforms,\n            **self.dl_defaults,\n            **self.kwargs_dataloader,\n        )\n\n    def test_dataloader(self) -> Optional[DataLoader]:\n        if not self.test_dataset:\n            logging.warning('no testing data found')\n            return\n        return DataLoader(\n            self.test_dataset,\n            shuffle=False,\n            batch_transforms=self.valid_transforms,\n            **self.dl_defaults,\n            **self.kwargs_dataloader,\n        )","metadata":{"execution":{"iopub.status.busy":"2023-10-27T13:30:57.098691Z","iopub.execute_input":"2023-10-27T13:30:57.099021Z","iopub.status.idle":"2023-10-27T13:30:59.436369Z","shell.execute_reply.started":"2023-10-27T13:30:57.098991Z","shell.execute_reply":"2023-10-27T13:30:59.435453Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import warnings\nwarnings.filterwarnings(\"ignore\", category=ResourceWarning) ","metadata":{"execution":{"iopub.status.busy":"2023-10-27T13:30:59.439906Z","iopub.execute_input":"2023-10-27T13:30:59.440488Z","iopub.status.idle":"2023-10-27T13:30:59.44517Z","shell.execute_reply.started":"2023-10-27T13:30:59.440456Z","shell.execute_reply":"2023-10-27T13:30:59.444062Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dm = SpineScansDM(\n    data_dir='/tmp',\n    path_csv=os.path.join(PATH_DATASET, \"train.csv\"),\n    vol_size=(128, 128, 128),\n    batch_size=9,\n    num_workers=2,\n)\ndm.setup()\n\nfor batch in dm.train_dataloader():\n    for i in range(2):\n        show_sample(batch[\"data\"][i], batch[\"label\"][i])\n    break","metadata":{"execution":{"iopub.status.busy":"2023-10-27T13:30:59.44702Z","iopub.execute_input":"2023-10-27T13:30:59.447374Z","iopub.status.idle":"2023-10-27T13:31:03.295614Z","shell.execute_reply.started":"2023-10-27T13:30:59.447337Z","shell.execute_reply":"2023-10-27T13:31:03.294395Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Load pretrained Medical model","metadata":{}},{"cell_type":"code","source":"from monai.networks.nets import ResNet, resnet10, resnet18\n\nresnet = resnet10(\n    pretrained=False, spatial_dims=3, n_input_channels=1, num_classes=dm.num_labels,\n)","metadata":{"execution":{"iopub.status.busy":"2023-10-27T13:31:03.297436Z","iopub.execute_input":"2023-10-27T13:31:03.297884Z","iopub.status.idle":"2023-10-27T13:31:06.499635Z","shell.execute_reply.started":"2023-10-27T13:31:03.297839Z","shell.execute_reply":"2023-10-27T13:31:06.498558Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def create_pretrained_medical_resnet(\n    pretrained_path: str,\n    model_constructor: callable = resnet18,\n    spatial_dims: int = 3,\n    n_input_channels: int = 1,\n    num_classes: int = 1,\n) -> ResNet:\n    \"\"\"This si specific constructor for MONAI ResNet module loading MedicalNEt weights.\"\"\"\n    net = model_constructor(\n        pretrained=False, spatial_dims=spatial_dims, n_input_channels=n_input_channels, num_classes=num_classes,\n    )\n    net_dict = net.state_dict()\n    pretrain = torch.load(pretrained_path)\n    pretrain['state_dict'] = {k.replace('module.', ''): v for k, v in pretrain['state_dict'].items()}\n    missing = tuple({k for k in net_dict.keys() if k not in pretrain['state_dict']})\n    print(f\"missing in pretrained: {len(missing)}\")\n    inside = tuple({k for k in pretrain['state_dict'] if k in net_dict.keys()})\n    print(f\"inside pretrained: {len(inside)}\")\n    unused = tuple({k for k in pretrain['state_dict'] if k not in net_dict.keys()})\n    print(f\"unused pretrained: {len(unused)}\")\n    pretrain['state_dict'] = {k: v for k, v in pretrain['state_dict'].items() if k in net_dict.keys()}\n    net.load_state_dict(pretrain['state_dict'], strict=False)\n    return net\n\n# resnet = create_pretrained_medical_resnet(\n#     \"../input/meidcalnet-pretrained-3d-resnet-weights/resnet_10.pth\",\n#     model_constructor=resnet10, num_classes=dm.num_labels,\n# )\n","metadata":{"execution":{"iopub.status.busy":"2023-10-27T13:31:06.500863Z","iopub.execute_input":"2023-10-27T13:31:06.501182Z","iopub.status.idle":"2023-10-27T13:31:06.513987Z","shell.execute_reply.started":"2023-10-27T13:31:06.501153Z","shell.execute_reply":"2023-10-27T13:31:06.512967Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Lightning 🍱 Module","metadata":{}},{"cell_type":"code","source":"from typing import Any, Optional, Sequence, Tuple, Type, Union\nimport torch.nn.functional as F\nfrom monai.networks.nets import EfficientNetBN, ResNet, resnet18\nfrom pytorch_lightning import LightningModule\nfrom torch import nn, Tensor\nfrom torch.optim import AdamW, Optimizer\nfrom torch.optim.lr_scheduler import CosineAnnealingLR\nfrom torchmetrics import F1 as F1Score\n\nclass LitNeckCT(LightningModule):\n\n    def __init__(\n        self, net: nn.Module, num_labels: int = 7, lr: float = 1e-3,\n        optimizer: Optional[Type[Optimizer]] = None,\n    ):\n        super().__init__()\n        self.net = net\n        self.name = net.__class__.__name__\n        for n, param in self.net.named_parameters():\n            param.requires_grad = True\n        self.learning_rate = lr\n        self.optimizer = optimizer or AdamW\n\n        self.train_f1_score = F1Score(num_classes=num_labels)\n        self.val_f1_score = F1Score(num_classes=num_labels)\n\n    def forward(self, x: Tensor) -> Tensor:\n        return torch.sigmoid(self.net(x))\n\n    @staticmethod\n    def compute_loss(y_hat: Tensor, y: Tensor):\n        # print(y_hat, y.to(y_hat.dtype))\n        return F.binary_cross_entropy_with_logits(y_hat, y.to(y_hat.dtype))\n\n    def training_step(self, batch, batch_idx):\n        img, y = batch[\"data\"], batch[\"label\"]\n        y_hat = self(img)\n        loss = self.compute_loss(y_hat, y)\n        self.log(\"train/loss\", loss, prog_bar=False)\n        self.log(\"train/f1\", self.train_f1_score(y_hat, y), prog_bar=True)\n        return loss\n\n    def validation_step(self, batch, batch_idx):\n        img, y = batch[\"data\"], batch[\"label\"]\n        y_hat = self(img)\n        loss = self.compute_loss(y_hat, y)\n        self.log(\"valid/loss\", loss, prog_bar=False)\n        self.log(\"valid/f1\", self.val_f1_score(y_hat, y), prog_bar=True)\n\n    def configure_optimizers(self):\n        optimizer = self.optimizer(self.net.parameters(), lr=self.learning_rate)\n        # print(self.trainer.max_epochs)\n        scheduler = CosineAnnealingLR(optimizer, self.trainer.max_epochs * 200, 0)\n        return [optimizer], [scheduler]","metadata":{"execution":{"iopub.status.busy":"2023-10-27T13:31:06.515862Z","iopub.execute_input":"2023-10-27T13:31:06.516326Z","iopub.status.idle":"2023-10-27T13:31:06.537362Z","shell.execute_reply.started":"2023-10-27T13:31:06.516289Z","shell.execute_reply":"2023-10-27T13:31:06.536293Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# net = resnet18(pretrained=False, spatial_dims=3, n_input_channels=1, num_classes=1)\nmodel = LitNeckCT(resnet, num_labels=dm.num_labels, lr=0.01)\n\nprint(model.net)","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2023-10-27T13:31:06.538731Z","iopub.execute_input":"2023-10-27T13:31:06.539489Z","iopub.status.idle":"2023-10-27T13:31:06.55803Z","shell.execute_reply.started":"2023-10-27T13:31:06.53945Z","shell.execute_reply":"2023-10-27T13:31:06.557061Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training ⚙️ model","metadata":{}},{"cell_type":"code","source":"import pytorch_lightning as pl\n\ntrainer = pl.Trainer(\n    max_epochs=5,\n    logger=pl.loggers.CSVLogger(save_dir='logs/'),\n    gpus=torch.cuda.device_count(),\n    precision=16 if torch.cuda.is_available() else 32,\n    accumulate_grad_batches=4,\n    gradient_clip_val=0.01,\n)","metadata":{"execution":{"iopub.status.busy":"2023-10-27T13:31:06.559347Z","iopub.execute_input":"2023-10-27T13:31:06.56136Z","iopub.status.idle":"2023-10-27T13:31:06.63176Z","shell.execute_reply.started":"2023-10-27T13:31:06.56133Z","shell.execute_reply":"2023-10-27T13:31:06.630647Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"trainer.fit(model, datamodule=dm)\n\n# Save the model!\ntrainer.save_checkpoint(\"image_classification_model.pt\")\n\n!ls -lh .","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2023-10-27T13:31:06.633125Z","iopub.execute_input":"2023-10-27T13:31:06.633423Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import seaborn as sn\n\nmetrics = pd.read_csv(f'{trainer.logger.log_dir}/metrics.csv')\ndel metrics[\"step\"]\nmetrics.set_index(\"epoch\", inplace=True)\n# display(metrics.dropna(axis=1, how=\"all\").head())\ng = sn.relplot(data=metrics, kind=\"line\")\nplt.gcf().set_size_inches(12, 4)\n# plt.gca().set_yscale('log')\nplt.grid()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Predictions 🔥 inference","metadata":{}},{"cell_type":"code","source":"for batch in dm.test_dataloader():\n    for i in range(2):\n        show_sample(batch[\"data\"][i], batch[\"label\"][i])\n    break","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model.eval()\npredictions = []\nfor batch in dm.test_dataloader():\n    with torch.no_grad():\n        preds = model(batch[\"data\"].to(model.device)).cpu()\n    for pred, fname in zip(preds.detach().numpy(), batch[\"label\"]):\n        study_id, _ = os.path.splitext(os.path.basename(fname))\n        pred_overall = max(pred)\n        predictions.append({\n            \"StudyInstanceUID\": study_id,\n            \"prediction_type\": \"patient_overall\",\n            \"fractured\": pred_overall,\n        })\n        predictions += [\n            {\"StudyInstanceUID\": study_id, \"prediction_type\": c, \"fractured\": v}\n            for c, v in zip(dm.label_names, pred)\n        ]\n\ndf_predictions = pd.DataFrame(predictions)\ndisplay(df_predictions.head(10))","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Finalize 📩 submission","metadata":{}},{"cell_type":"code","source":"df_test = pd.read_csv(os.path.join(PATH_DATASET, \"test.csv\"))\ndisplay(df_test.head())\n\n!cat ../input/rsna-2022-cervical-spine-fracture-detection/sample_submission.csv","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_submission = df_test.merge(df_predictions, how=\"left\", on=[\"StudyInstanceUID\", \"prediction_type\"]).set_index(\"row_id\")\ndf_submission.fillna(0.5, inplace=True)\ndisplay(df_submission.head())","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_submission[\"fractured\"].to_csv(\"submission.csv\")\n!head submission.csv","metadata":{"trusted":true},"execution_count":null,"outputs":[]}]}