{"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":"%%capture\n!pip install itk --no-index --find-links=file:///kaggle/input/monaiitk/itk -q\n!pip install monai --no-index --find-links=file:///kaggle/input/monaiitk/monai -q","metadata":{"execution":{"iopub.status.busy":"2022-08-24T14:18:26.431206Z","iopub.execute_input":"2022-08-24T14:18:26.43255Z","iopub.status.idle":"2022-08-24T14:19:04.865859Z","shell.execute_reply.started":"2022-08-24T14:18:26.432425Z","shell.execute_reply":"2022-08-24T14:19:04.8646Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import monai\nprint(monai.__version__)","metadata":{"execution":{"iopub.status.busy":"2022-08-24T14:19:04.869124Z","iopub.execute_input":"2022-08-24T14:19:04.86989Z","iopub.status.idle":"2022-08-24T14:19:11.586287Z","shell.execute_reply.started":"2022-08-24T14:19:04.869818Z","shell.execute_reply":"2022-08-24T14:19:11.584445Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"DEBUG = False","metadata":{"execution":{"iopub.status.busy":"2022-08-24T14:19:11.58891Z","iopub.execute_input":"2022-08-24T14:19:11.589639Z","iopub.status.idle":"2022-08-24T14:19:11.595933Z","shell.execute_reply.started":"2022-08-24T14:19:11.589596Z","shell.execute_reply":"2022-08-24T14:19:11.593404Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from kaggle_secrets import UserSecretsClient\nuser_secrets = UserSecretsClient()\nsecret_value_0 = user_secrets.get_secret(\"wandb\")","metadata":{"execution":{"iopub.status.busy":"2022-08-24T14:19:11.597657Z","iopub.execute_input":"2022-08-24T14:19:11.598424Z","iopub.status.idle":"2022-08-24T14:19:11.930118Z","shell.execute_reply.started":"2022-08-24T14:19:11.598386Z","shell.execute_reply":"2022-08-24T14:19:11.929198Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np \nimport pandas as pd \nimport os, sys, glob, random, copy, gc\nfrom tqdm.auto import tqdm\nimport matplotlib.pyplot as plt\nfrom typing import List, Tuple\nfrom collections import defaultdict\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torch.backends.cudnn as cudnn\nfrom torch.utils.data import DataLoader, Dataset\nfrom torch.optim.lr_scheduler import CosineAnnealingLR\nfrom torch.optim import Adam\nfrom torch.cuda.amp import autocast, GradScaler\n\nimport monai as mn","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-08-24T14:19:11.934369Z","iopub.execute_input":"2022-08-24T14:19:11.934893Z","iopub.status.idle":"2022-08-24T14:19:11.941547Z","shell.execute_reply.started":"2022-08-24T14:19:11.934825Z","shell.execute_reply":"2022-08-24T14:19:11.940428Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ROOT_DIR = '/kaggle/input/rsna-2022-cervical-spine-fracture-detection'\n\nTRAIN_PATH = \"../input/rsna-2022-cervical-spine-fracture-detection/train.csv\"\nTEST_PATH = \"../input/rsna-2022-cervical-spine-fracture-detection/test.csv\"\n\nTRAIN_DIR = \"../input/rsna-2022-cervival-spine-3d-numy-train/train_arrays\"\nTEST_DIR = \"../input/rsna-2022-cervical-spine-fracture-detection/test_images\"\n\nSAMPLE_SUB = \"../input/rsna-2022-cervical-spine-fracture-detection/sample_submission.csv\"\n\nLABELS_COLS = [\"patient_overall\", \"C1\", \"C2\", \"C3\", \"C4\", \"C5\", \"C6\", \"C7\"]\n\nTO_EXCLUDE = \"1.2.826.0.1.3680043.20574\"","metadata":{"execution":{"iopub.status.busy":"2022-08-24T14:19:11.943073Z","iopub.execute_input":"2022-08-24T14:19:11.943893Z","iopub.status.idle":"2022-08-24T14:19:11.952406Z","shell.execute_reply.started":"2022-08-24T14:19:11.943838Z","shell.execute_reply":"2022-08-24T14:19:11.951499Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(device)\n\nSEED = 42\nIMG_WIDTH = 128\nIMG_HEIGHT = 128\nIMG_DEPTH = 128\nWINDOW_WIDTH = 1800\nWINDOW_LEVEL = 400\nBATCH_SIZE = 2\nTRAIN_SIZE, VALID_SIZE, TEST_SIZE = [8, 1, 1]\nNUM_WORKERS = 2\nMODEL_NAME = \"resnet10\"\nLR = 1e-5\nEPOCHS = 15\nVER=2","metadata":{"execution":{"iopub.status.busy":"2022-08-24T14:19:11.955613Z","iopub.execute_input":"2022-08-24T14:19:11.956816Z","iopub.status.idle":"2022-08-24T14:19:12.02107Z","shell.execute_reply.started":"2022-08-24T14:19:11.956774Z","shell.execute_reply":"2022-08-24T14:19:12.019875Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import wandb\n\n!wandb login $secret_value_0\n\nwandb.config = {\n    'batch_size' : BATCH_SIZE,\n    'model_name' : MODEL_NAME,\n    'learning_rate' : LR,\n    'epochs' : EPOCHS,\n    'img_width' : 128,\n    'img_height' : 128,\n    'img_depth' : 128\n}\n\nwandb.init(project=\"RSNA 2022 Cervical Spine Monai\", entity=\"barteksadlej\", name=\"resnet10\")","metadata":{"execution":{"iopub.status.busy":"2022-08-24T14:19:12.022631Z","iopub.execute_input":"2022-08-24T14:19:12.023125Z","iopub.status.idle":"2022-08-24T14:19:18.848Z","shell.execute_reply.started":"2022-08-24T14:19:12.022985Z","shell.execute_reply":"2022-08-24T14:19:18.846762Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def set_seed(seed=SEED):\n    cudnn.benchmark = True\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    random.seed(seed)\n    \ndef seed_worker(worker_id):\n    worker_seed = torch.initial_seed() % 2**32\n    np.random.seed(worker_seed)\n    random.seed(worker_seed)\n\nset_seed()\n    \ng = torch.Generator()\ng.manual_seed(SEED)","metadata":{"execution":{"iopub.status.busy":"2022-08-24T14:19:18.853491Z","iopub.execute_input":"2022-08-24T14:19:18.853776Z","iopub.status.idle":"2022-08-24T14:19:18.878717Z","shell.execute_reply.started":"2022-08-24T14:19:18.853747Z","shell.execute_reply":"2022-08-24T14:19:18.877889Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df = pd.read_csv(TRAIN_PATH)\ntrain_df = train_df[train_df.StudyInstanceUID != TO_EXCLUDE]\n\nif DEBUG:\n    train_df = train_df.sample(10).reset_index(drop=True)\n    \nfrom sklearn.preprocessing import LabelEncoder\nfrom sklearn.model_selection import StratifiedKFold\n\ntrain_df['multilabel'] = LabelEncoder().fit_transform([str(x) for x in train_df[LABELS_COLS].values])\ntrain_df.head()\nprint(train_df['multilabel'].unique().shape)\n\nskf = StratifiedKFold(n_splits = np.sum([TRAIN_SIZE, VALID_SIZE, TEST_SIZE]), shuffle=True, random_state=SEED)\n\ntrain_idx = []\nvalid_idx = []\ntest_idx = []\n\nfor split, (_, indexes) in enumerate(skf.split(train_df, train_df['multilabel'])):\n    \n    if split < TRAIN_SIZE:\n        train_idx.extend(indexes)\n    elif split < TRAIN_SIZE + VALID_SIZE:\n        valid_idx.extend(indexes)\n    else:\n        test_idx.extend(indexes)\n        \nprint(f\"train: {len(train_idx)}, valid: {len(valid_idx)}, test: {len(test_idx)}\")\ntrain_df.drop(['multilabel'], axis=1, inplace=True)\n\ntrain = train_df.iloc[train_idx]\nvalid = train_df.iloc[valid_idx]\ntest = train_df.iloc[test_idx]\n\nprint(f\"train: {len(train)}, valid: {len(valid)}, test: {len(test)}\")","metadata":{"execution":{"iopub.status.busy":"2022-08-24T14:19:18.880225Z","iopub.execute_input":"2022-08-24T14:19:18.880551Z","iopub.status.idle":"2022-08-24T14:19:19.036072Z","shell.execute_reply.started":"2022-08-24T14:19:18.880518Z","shell.execute_reply":"2022-08-24T14:19:19.034233Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"basic_transforms = mn.transforms.Compose([\n    mn.transforms.LoadImageD(keys=[\"image\"], reader=[\"NumpyReader\"]),\n#     mn.transforms.AddChannelD(keys=[\"image\"]),\n#     mn.transforms.ScaleIntensityRangeD(keys=[\"image\"],a_min=WINDOW_LEVEL-(WINDOW_WIDTH/2), a_max=WINDOW_LEVEL+(WINDOW_WIDTH/2),b_min=0.0, b_max=1.0, clip=True),\n#     mn.transforms.Resized(keys=[\"image\"], spatial_size=(IMG_WIDTH, IMG_HEIGHT, IMG_DEPTH)),\n    mn.transforms.RandGaussianNoised(keys=[\"image\"], prob=0.5),\n    mn.transforms.RandBiasFieldd(keys=[\"image\"], prob=0.3),\n    mn.transforms.RandCoarseDropoutd(keys=[\"image\"], holes=20, spatial_size=10, fill_value=0, prob=0.3),\n    mn.transforms.RandRotated(keys=[\"image\"], prob=0.3 ),\n    mn.transforms.Zoomd(keys=[\"image\"], prob=0.3 , zoom=1.3),\n    mn.transforms.ToTensorD(keys=[\"image\"]),\n])\n\ndef gather_labels(df):\n    \n    for i,item in enumerate(df):\n        labels = np.empty(shape=(len(LABELS_COLS)), dtype=np.float32)\n        for index, label in enumerate(LABELS_COLS):\n            labels[index] = item[label]\n            del item[label]\n        item['label'] = labels\n    \n    return df\n            \ndef get_dataset(df, training=False):\n    \n    df = df.to_dict(orient=\"records\")\n    for i,item in enumerate(df):\n        df[i][\"image\"] = os.path.join(TRAIN_DIR if training else TEST_DIR, f'{item[\"StudyInstanceUID\"]}.npy')\n        \n    if training:\n        df = gather_labels(df)\n        transforms = mn.transforms.Compose([\n            basic_transforms,\n            mn.transforms.ToTensorD(keys=[\"label\"], dtype=torch.float),\n        ])\n    else:\n        transforms = basic_transforms\n        \n    return mn.data.Dataset(data=df, transform=transforms)\n\ndef make_dataloader(ds, training=False):\n    \n    return DataLoader(\n        ds, \n        batch_size = BATCH_SIZE,\n        num_workers=NUM_WORKERS,\n        shuffle=training,\n        worker_init_fn=seed_worker,\n        generator=g)","metadata":{"execution":{"iopub.status.busy":"2022-08-24T14:19:19.037535Z","iopub.execute_input":"2022-08-24T14:19:19.037894Z","iopub.status.idle":"2022-08-24T14:19:19.057035Z","shell.execute_reply.started":"2022-08-24T14:19:19.037841Z","shell.execute_reply":"2022-08-24T14:19:19.055988Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_ds = get_dataset(train, training=True)\ntrain_dl = make_dataloader(train_ds, training=True)\n\nvalid_ds = get_dataset(valid, training=True)\ntest_ds = get_dataset(test, training=True)","metadata":{"execution":{"iopub.status.busy":"2022-08-24T14:19:19.058607Z","iopub.execute_input":"2022-08-24T14:19:19.059045Z","iopub.status.idle":"2022-08-24T14:19:19.118866Z","shell.execute_reply.started":"2022-08-24T14:19:19.059009Z","shell.execute_reply":"2022-08-24T14:19:19.118068Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class RSNA2022Model(nn.Module):\n    \n    def __init__(self):\n        super().__init__()\n        self.backbone = monai.networks.nets.resnet10(spatial_dims=3, n_input_channels=1, n_classes=len(LABELS_COLS))\n        \n    def forward(self, x):\n    \n        return self.backbone(x)","metadata":{"execution":{"iopub.status.busy":"2022-08-24T14:19:19.120304Z","iopub.execute_input":"2022-08-24T14:19:19.120642Z","iopub.status.idle":"2022-08-24T14:19:19.126819Z","shell.execute_reply.started":"2022-08-24T14:19:19.120609Z","shell.execute_reply":"2022-08-24T14:19:19.125894Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch.nn as nn\n\nloss_fn = nn.BCEWithLogitsLoss(reduction='none')\n\ncompetition_weights = {\n    '-' : torch.tensor([7, 1, 1, 1, 1, 1, 1, 1], dtype=torch.float, device=device),\n    '+' : torch.tensor([14, 2, 2, 2, 2, 2, 2, 2], dtype=torch.float, device=device),\n}\n\n# y_hat.shape = (batch_size, num_classes)\n# y.shape = (batch_size, num_classes)\ndef competiton_loss(y_hat, y):\n    loss = loss_fn(y_hat, y)\n    weights = y * competition_weights['+'] + (1 - y) * competition_weights['-']\n    loss = (loss * weights).sum(axis=1).mean()\n    \n    return loss / weights.sum()","metadata":{"execution":{"iopub.status.busy":"2022-08-24T14:19:19.131001Z","iopub.execute_input":"2022-08-24T14:19:19.131265Z","iopub.status.idle":"2022-08-24T14:19:22.082921Z","shell.execute_reply.started":"2022-08-24T14:19:19.131238Z","shell.execute_reply":"2022-08-24T14:19:22.0819Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class AvgMeter():\n    \n    def __init__(self, name=None):\n        self.reset()\n        self.name = name\n        \n    def reset(self):\n        self.count = 0\n        self.sum = 0\n        self.avg = 0\n        \n    def update(self, x, n=1):\n        self.count +=  n\n        self.sum += x*n\n        self.avg = self.sum / self.count\n#         wandb.log({f\"{self.name}\" : x})\n        \n    def __str__(self):\n        return f\"{self.name} : {self.avg}\"\n    \n    def log(self):\n        wandb.log({f\"{self.name}\" : self.avg})","metadata":{"execution":{"iopub.status.busy":"2022-08-24T14:19:22.084483Z","iopub.execute_input":"2022-08-24T14:19:22.084825Z","iopub.status.idle":"2022-08-24T14:19:22.09356Z","shell.execute_reply.started":"2022-08-24T14:19:22.084792Z","shell.execute_reply":"2022-08-24T14:19:22.092506Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = RSNA2022Model().to(device)\nwandb.watch(model)\noptimizer = Adam(model.parameters(), lr=LR)\nscheduler = None #  CosineAnnealingLR(optimizer, EPOCHS)\n\n# scaler = GradScaler()\n\ndef train_one_epoch(model, optimizer, train_dl, device, scheduler=None):\n    \n    train_loss = AvgMeter(\"train loss\")\n    model = model.train(True)\n    for batch in tqdm(train_dl):\n        \n        ct = batch['image'].to(device)\n        target = batch[\"label\"].to(device)\n        optimizer.zero_grad()\n        \n        y_hat = model(ct)\n\n        loss = competiton_loss(y_hat, target)\n        train_loss.update(loss.item(), target.shape[0])\n        loss.backward()\n        optimizer.step()\n        \n        if scheduler is not None:\n            scheduler.step()\n            \n    return train_loss\n\n@torch.no_grad()\ndef eval_model(model, dl, device, loss_name=None):\n    \n    loss_meter = AvgMeter(loss_name)\n    model = model.train(False)\n    \n    for batch in tqdm(dl): \n        ct = batch['image'].to(device)\n        target = batch[\"label\"].to(device)\n        y_hat = model(ct)\n        loss = competiton_loss(y_hat, target)\n        loss_meter.update(loss.item(), target.shape[0])\n    \n    return loss_meter\n\nclass ModelCheckpoint():\n    \n    def __init__(self, path=f\"{MODEL_NAME}_VER_{VER}.pt\"):\n        \n        self.best_values = defaultdict(lambda : np.inf)\n        self.path = path\n        \n    def make_checkpoint(self, model, values: List[Tuple[str, float]]):\n        \n        for name, value in values:\n            \n            if value < self.best_values[name]:\n                \n                print(f\"new best {name}, saving model...\")\n                self.best_values[name] = value\n                torch.save(model.state_dict(), f\"{name}_{value:.6f}_{self.path}\")\n                print(\"model saved!\")\n\nmckpt = ModelCheckpoint()\n\nfor epoch in range(EPOCHS):\n    \n    _ = gc.collect()\n    \n    print(\"=\" * 25)\n    print(f\"epoch [{epoch + 1}]\")\n    print(\"=\" * 25)\n    \n    \n    train_loss = train_one_epoch(model, optimizer, train_dl, device, scheduler)\n    print(train_loss)\n    train_loss.log()\n    _ = gc.collect()\n    \n    valid_dl = make_dataloader(valid_ds, training=True)\n    valid_loss = eval_model(model, valid_dl, device, loss_name=f\"valid loss\")\n    del valid_dl\n    print(valid_loss)\n    valid_loss.log()\n    _ = gc.collect()\n    \n    test_dl = make_dataloader(test_ds, training=True)\n    test_loss = eval_model(model, test_dl, device, loss_name=f\"test loss\")\n    del test_dl \n    print(test_loss)\n    test_loss.log()\n    _ = gc.collect()\n    \n    mckpt.make_checkpoint(model, list(map(lambda loss : (loss.name, loss.avg), [train_loss, valid_loss, test_loss])))","metadata":{"execution":{"iopub.status.busy":"2022-08-24T14:19:22.095169Z","iopub.execute_input":"2022-08-24T14:19:22.095969Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}