{"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":"code","source":"\n!pip install git+https://github.com/Project-MONAI/MONAI#egg=monai\n!pip install python-gdcm\n!pip install pylibjpeg\n!pip install einops\n","metadata":{"execution":{"iopub.status.busy":"2022-10-19T21:38:15.217488Z","iopub.execute_input":"2022-10-19T21:38:15.217899Z","iopub.status.idle":"2022-10-19T21:38:15.224999Z","shell.execute_reply.started":"2022-10-19T21:38:15.217867Z","shell.execute_reply":"2022-10-19T21:38:15.224061Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import glob\nimport torch\nfrom  monai.data.image_reader import PydicomReader\nfrom monai.networks.nets import EfficientNetBN\nimport matplotlib.pyplot as plt\nfrom torch.utils.data import DataLoader, ConcatDataset, random_split, Dataset\nfrom monai.transforms import ScaleIntensityD, EnsureChannelFirstD, LoadImageD, ResizeWithPadOrCropD, RandRotateD, RandFlipD, CastToTypeD, Compose\nfrom pytorch_lightning import LightningDataModule, LightningModule\nfrom monai.networks.nets.vit import ViT\n\nfrom pytorch_lightning import Trainer\nfrom pytorch_lightning.callbacks import LearningRateMonitor, ModelCheckpoint\nimport pandas as pd\nimport numpy as np","metadata":{"execution":{"iopub.status.busy":"2022-10-19T21:38:15.231237Z","iopub.execute_input":"2022-10-19T21:38:15.232252Z","iopub.status.idle":"2022-10-19T21:38:15.238983Z","shell.execute_reply.started":"2022-10-19T21:38:15.232213Z","shell.execute_reply":"2022-10-19T21:38:15.238064Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def weighted_loss(y_pred_logit, y):\n    \"\"\"\n    Weighted loss\n    We reuse torch.nn.functional.binary_cross_entropy_with_logits here. pos_weight and weights combined give us necessary coefficients described in https://www.kaggle.com/competitions/rsna-2022-cervical-spine-fracture-detection/discussion/340392\n    See also this explanation: https://www.kaggle.com/code/samuelcortinhas/rsna-fracture-detection-in-depth-eda/notebook\n    \"\"\"\n\n    DEVICE = torch.device(y_pred_logit.get_device())\n    neg_weights = (torch.tensor([7., 1, 1, 1, 1, 1, 1, 1]) if y_pred_logit.shape[-1] == 8 else torch.ones(y_pred_logit.shape[-1])).to(DEVICE)\n    pos_weights = (torch.tensor([14., 2, 2, 2, 2, 2, 2, 2]) if y_pred_logit.shape[-1] == 8 else torch.ones(y_pred_logit.shape[-1]) * 2.).to(DEVICE)\n\n    loss = torch.nn.functional.binary_cross_entropy_with_logits(\n        y_pred_logit,\n        y,\n        reduction='none',\n    )\n\n    pos_weights = y * pos_weights.unsqueeze(0)\n    neg_weights = (1 - y) * neg_weights.unsqueeze(0)\n    all_weights = pos_weights + neg_weights\n\n    loss *= all_weights\n    \n    norm = torch.sum(all_weights, dim=1).unsqueeze(1)\n\n    loss /= norm\n   \n    loss = torch.sum(loss, dim=1)\n    return torch.mean(loss)","metadata":{"execution":{"iopub.status.busy":"2022-10-19T21:38:15.243978Z","iopub.execute_input":"2022-10-19T21:38:15.244934Z","iopub.status.idle":"2022-10-19T21:38:15.254841Z","shell.execute_reply.started":"2022-10-19T21:38:15.244897Z","shell.execute_reply":"2022-10-19T21:38:15.253882Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class RSNADataset(Dataset):\n    \n    def __init__(self):\n        self.all_files = glob.glob('../input/rsna-2022-cervical-spine-fracture-detection/train_images/*')\n        self.csv = pd.read_csv(\"../input/rsna-2022-cervical-spine-fracture-detection/train.csv\")\n        self.reader = PydicomReader(swap_ij=True)\n        \n    def __len__(self):\n        return len(self.all_files)\n    \n    def __getitem__(self, index):\n        series = self.all_files[index]\n        name = series.split('/')[-1]\n        image = self.reader.get_data(self.reader.read(series))[0]\n        data = {'img':image, 'lbl':torch.tensor(self.csv [self.csv ['StudyInstanceUID']==name].values.flatten().tolist()[1:])}\n        return self.transform(data)\n\n        ","metadata":{"execution":{"iopub.status.busy":"2022-10-19T21:38:15.259441Z","iopub.execute_input":"2022-10-19T21:38:15.260145Z","iopub.status.idle":"2022-10-19T21:38:15.271556Z","shell.execute_reply.started":"2022-10-19T21:38:15.260116Z","shell.execute_reply":"2022-10-19T21:38:15.270546Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class DataModule(LightningDataModule):\n    def __init__(self, batch_size: int = 2):\n        super().__init__()\n\n        # Transforms for the dataset\n        dict_keys=[\"img\"]\n        transform = Compose(\n\n            [\n                EnsureChannelFirstD(keys=dict_keys),\n                ScaleIntensityD(keys=dict_keys, minv=0.0, maxv=1.0),\n                ResizeWithPadOrCropD(keys=dict_keys, spatial_size=(256, 256, 256)),\n                RandRotateD(keys=dict_keys, range_x=np.pi/12, prob=0.5, keep_size=True),\n                RandFlipD(keys=dict_keys, spatial_axis=0, prob=0.5),\n                CastToTypeD(keys=dict_keys+['lbl'], dtype=torch.half)\n            ]\n        )\n\n        # load datasets and set the transforms\n        dataset = RSNADataset()\n        dataset.transform = transform\n\n        # split into train and validate; no test as dont care much\n        self.train, self.val = random_split(dataset, [int(len(dataset)*0.8), len(dataset) - int(len(dataset)*0.8) ], )\n        self.batch_size = batch_size\n\n    def train_dataloader(self):\n        return DataLoader(self.train, batch_size=self.batch_size, num_workers=2)\n\n    def val_dataloader(self):\n        return DataLoader(self.val, batch_size=self.batch_size, num_workers=2)\n","metadata":{"execution":{"iopub.status.busy":"2022-10-19T21:38:15.275617Z","iopub.execute_input":"2022-10-19T21:38:15.275907Z","iopub.status.idle":"2022-10-19T21:38:15.287415Z","shell.execute_reply.started":"2022-10-19T21:38:15.275881Z","shell.execute_reply":"2022-10-19T21:38:15.286203Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class VIT(LightningModule):\n    def __init__(self, lr):\n        super().__init__()\n\n        self.save_hyperparameters()\n        \n        self.model = ViT(\n            in_channels=1,\n            img_size = (256, 256, 256),\n            patch_size = (16, 16, 16),\n            classification= True,\n            num_classes = 8\n                        )\n        \n        self.criterion = weighted_loss\n\n    def forward(self, images):\n        out = self.model(images)\n        return out \n\n    def training_step(self, batch, batch_idx):\n        return self._common_step(batch, batch_idx, \"train\")\n\n    def validation_step(self, batch, batch_idx):\n        self._common_step(batch, batch_idx, \"val\")\n\n\n    def configure_optimizers(self):\n        optimizer = torch.optim.AdamW(self.parameters(), lr=self.hparams.lr)\n        lr_scheduler = {\n            'scheduler': torch.optim.lr_scheduler.CosineAnnealingWarmRestarts(optimizer, T_0=10 ,T_mult=2, eta_min=1e-6),\n            #'scheduler': torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer, factor=0.5),\n            'monitor': 'val_loss'\n        }\n        return [optimizer], [lr_scheduler]\n\n    def _prepare_batch(self, batch):\n        return batch['img'], batch['lbl']\n\n    def _common_step(self, batch, batch_idx, stage: str):\n        imgs, gt = self._prepare_batch(batch)\n        prediction = self.forward(imgs)[0]\n        clf_loss = self.criterion(prediction, gt)\n        self.log(f'{stage}_loss', clf_loss.item(), prog_bar=True)\n        return clf_loss\n\n\n\n","metadata":{"execution":{"iopub.status.busy":"2022-10-19T21:38:15.315149Z","iopub.execute_input":"2022-10-19T21:38:15.315481Z","iopub.status.idle":"2022-10-19T21:38:15.326999Z","shell.execute_reply.started":"2022-10-19T21:38:15.31545Z","shell.execute_reply":"2022-10-19T21:38:15.325827Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"lr_monitor = LearningRateMonitor(logging_interval='epoch')\ncheckpoint_callback = ModelCheckpoint(dirpath=\"./saved_models/style_loss/\", save_top_k=1, monitor=\"val_total\", save_last=True)\ntrainer = Trainer(\n    accelerator='gpu', \n    devices=[0],\n    max_epochs=500,\n    callbacks=[lr_monitor, checkpoint_callback],\n    #overfit_batches=1,\n    log_every_n_steps=1,\n    precision=16,\n    reload_dataloaders_every_n_epochs=5,\n\n)\n# effective batch size is double due to reversing\ntrainer.fit(\n    model=VIT(\n        lr=1e-4, # learning rate\n    ), \n    datamodule=DataModule(\n        batch_size=1, # but is actually double due to flipping\n    ),\n)\n","metadata":{"execution":{"iopub.status.busy":"2022-10-19T21:38:15.329098Z","iopub.execute_input":"2022-10-19T21:38:15.329795Z","iopub.status.idle":"2022-10-19T21:39:22.883279Z","shell.execute_reply.started":"2022-10-19T21:38:15.329754Z","shell.execute_reply":"2022-10-19T21:39:22.880967Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}