{"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":"# 1st Place Solution Training 2.5D Classification Type1\n\nHi all,\n\nI'm very exciting to writing this notebook and the summary of our solution here.\n\nThis is small version of training my final models (stage2 type1), using efficientnetv2_s as backbone, and 224x224 as input.\n\nAfter all stage1 models are trained, then we can use those model to predict 3D masks for all training samples (2k)\n\nThen use those predicted masks to crop out all vertebraes (2k * 7 = 14k)\n\nI'll skip the code of predicting 3D maks and cropping vertebraes, but just uploaded the dataset of cropped vertebraes (https://www.kaggle.com/datasets/haqishen/rsna-cropped-2d-224-0920-2m)\n\nNow let's use this dataset to train a 2.5D classification with LSTM (Type1)\n\n**NOTE: The training time is too long for Kaggle kernels so you should run it locally**\n\nTo see more details of my solution: https://www.kaggle.com/competitions/rsna-2022-cervical-spine-fracture-detection/discussion/362607\n\n* Train Stage1 Notebook: https://www.kaggle.com/code/haqishen/rsna-2022-1st-place-solution-train-stage1\n* Train Stage2 (Type1) Notebook: This notebook\n* Train Stage2 (Type2) Notebook: https://www.kaggle.com/code/haqishen/rsna-2022-1st-place-solution-train-stage2-type2\n* Inference Notebook: https://www.kaggle.com/code/haqishen/rsna-2022-1st-place-solution-inference\n\n\n**If you find these notebooks helpful please upvote. Thanks!**","metadata":{}},{"cell_type":"code","source":"!pip -q install timm","metadata":{"execution":{"iopub.status.busy":"2022-12-08T05:51:26.593672Z","iopub.execute_input":"2022-12-08T05:51:26.594112Z","iopub.status.idle":"2022-12-08T05:51:40.576786Z","shell.execute_reply.started":"2022-12-08T05:51:26.59401Z","shell.execute_reply":"2022-12-08T05:51:40.575578Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"DEBUG = False","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-12-08T05:51:40.581232Z","iopub.execute_input":"2022-12-08T05:51:40.581597Z","iopub.status.idle":"2022-12-08T05:51:40.58812Z","shell.execute_reply.started":"2022-12-08T05:51:40.581562Z","shell.execute_reply":"2022-12-08T05:51:40.587045Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport sys\nimport gc\nimport ast\nimport cv2\nimport time\nimport timm\nimport pickle\nimport random\nimport argparse\nimport warnings\nimport numpy as np\nimport pandas as pd\nfrom glob import glob\nfrom PIL import Image\nfrom tqdm import tqdm\nimport albumentations\nfrom pylab import rcParams\nimport matplotlib.pyplot as plt\nfrom sklearn.model_selection import KFold, StratifiedKFold\n\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nimport torch.cuda.amp as amp\nimport torch.nn.functional as F\nfrom torch.utils.data import DataLoader, Dataset\n\n%matplotlib inline\nrcParams['figure.figsize'] = 20, 8\ndevice = torch.device('cuda')\ntorch.backends.cudnn.benchmark = True","metadata":{"execution":{"iopub.status.busy":"2022-12-08T05:51:40.589946Z","iopub.execute_input":"2022-12-08T05:51:40.590315Z","iopub.status.idle":"2022-12-08T05:51:44.569938Z","shell.execute_reply.started":"2022-12-08T05:51:40.590281Z","shell.execute_reply":"2022-12-08T05:51:44.568914Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Config","metadata":{}},{"cell_type":"code","source":"kernel_type = '0920_1bonev2_effv2s_224_15_6ch_augv2_mixupp5_drl3_rov1p2_bs8_lr23e5_eta23e6_50ep'\nload_kernel = None\nload_last = True\n\nn_folds = 5\nbackbone = 'densenet121'\n\nimage_size = 224\nn_slice_per_c = 15\nin_chans = 6\n\ninit_lr = 23e-5\neta_min = 23e-6\nbatch_size = 8\ndrop_rate = 0.1  # 드롭아웃 비율\ndrop_rate_last = 0.3\ndrop_path_rate = 0.\np_mixup = 0.5\np_rand_order_v1 = 0.2\n\ndata_dir = '../input/rsna-cropped-2d-224-0920-2m/cropped_2d_224_15_ext0_5ch_0920_2m/cropped_2d_224_15_ext0_5ch_0920_2m'\nuse_amp = True\nnum_workers = 4\nout_dim = 1\n\nn_epochs = 20\n\nlog_dir = './logs'\nmodel_dir = './models'\nos.makedirs(log_dir, exist_ok=True)\nos.makedirs(model_dir, exist_ok=True)","metadata":{"execution":{"iopub.status.busy":"2022-12-08T05:51:44.572523Z","iopub.execute_input":"2022-12-08T05:51:44.574182Z","iopub.status.idle":"2022-12-08T05:51:44.58293Z","shell.execute_reply.started":"2022-12-08T05:51:44.574152Z","shell.execute_reply":"2022-12-08T05:51:44.581999Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"transforms_train = albumentations.Compose([\n    albumentations.Resize(image_size, image_size),\n    albumentations.HorizontalFlip(p=0.5),\n    albumentations.VerticalFlip(p=0.5),\n    albumentations.Transpose(p=0.5),\n    albumentations.RandomBrightness(limit=0.1, p=0.7),\n    albumentations.ShiftScaleRotate(shift_limit=0.3, scale_limit=0.3, rotate_limit=45, border_mode=4, p=0.7),\n\n    albumentations.OneOf([\n        albumentations.MotionBlur(blur_limit=3),\n        albumentations.MedianBlur(blur_limit=3),\n        albumentations.GaussianBlur(blur_limit=3),\n        albumentations.GaussNoise(var_limit=(3.0, 9.0)),\n    ], p=0.5),\n    albumentations.OneOf([\n        albumentations.OpticalDistortion(distort_limit=1.),\n        albumentations.GridDistortion(num_steps=5, distort_limit=1.),\n    ], p=0.5),\n\n    albumentations.Cutout(max_h_size=int(image_size * 0.5), max_w_size=int(image_size * 0.5), num_holes=1, p=0.5),\n])\n\ntransforms_valid = albumentations.Compose([\n    albumentations.Resize(image_size, image_size),\n])","metadata":{"execution":{"iopub.status.busy":"2022-12-08T05:51:44.584226Z","iopub.execute_input":"2022-12-08T05:51:44.584838Z","iopub.status.idle":"2022-12-08T05:51:44.600932Z","shell.execute_reply.started":"2022-12-08T05:51:44.5848Z","shell.execute_reply":"2022-12-08T05:51:44.599886Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# DataFrame","metadata":{}},{"cell_type":"code","source":"df = pd.read_csv(os.path.join(f'../input/rsna-cropped-2d-224-0920-2m/train_seg.csv'))\ndf = df.sample(16).reset_index(drop=True) if DEBUG else df\n\n\nsid = []\ncs = []\nlabel = []\nfold = []\nfor _, row in df.iterrows():\n    for i in [1,2,3,4,5,6,7]:\n        sid.append(row.StudyInstanceUID)\n        cs.append(i)\n        label.append(row[f'C{i}'])\n        fold.append(row.fold)\n\ndf = pd.DataFrame({\n    'StudyInstanceUID': sid,\n    'c': cs,\n    'label': label,\n    'fold': fold\n})\n\ndf.tail()","metadata":{"execution":{"iopub.status.busy":"2022-12-08T05:51:44.602501Z","iopub.execute_input":"2022-12-08T05:51:44.603152Z","iopub.status.idle":"2022-12-08T05:51:44.997785Z","shell.execute_reply.started":"2022-12-08T05:51:44.603117Z","shell.execute_reply":"2022-12-08T05:51:44.996841Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Dataset","metadata":{}},{"cell_type":"code","source":"class CLSDataset(Dataset):\n    def __init__(self, df, mode, transform):\n\n        self.df = df.reset_index()\n        self.mode = mode\n        self.transform = transform\n\n    def __len__(self):\n        return self.df.shape[0]\n\n    def __getitem__(self, index):\n        row = self.df.iloc[index]\n        cid = row.c\n        \n        images = []\n        \n        for ind in list(range(n_slice_per_c)):\n            filepath = os.path.join(data_dir, f'{row.StudyInstanceUID}_{cid}_{ind}.npy')\n            image = np.load(filepath)\n            image = self.transform(image=image)['image']\n            image = image.transpose(2, 0, 1).astype(np.float32) / 255.\n            images.append(image)\n        images = np.stack(images, 0)\n\n        if self.mode != 'test':\n            images = torch.tensor(images).float()\n            labels = torch.tensor([row.label] * n_slice_per_c).float()\n            \n            if self.mode == 'train' and random.random() < p_rand_order_v1:\n                indices = torch.randperm(images.size(0))\n                images = images[indices]\n\n            return images, labels\n        else:\n            return torch.tensor(images).float()","metadata":{"execution":{"iopub.status.busy":"2022-12-08T05:51:44.99929Z","iopub.execute_input":"2022-12-08T05:51:44.999921Z","iopub.status.idle":"2022-12-08T05:51:45.010496Z","shell.execute_reply.started":"2022-12-08T05:51:44.999883Z","shell.execute_reply":"2022-12-08T05:51:45.009506Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"rcParams['figure.figsize'] = 20,8\n\ndf_show = df\ndataset_show = CLSDataset(df_show, 'train', transform=transforms_train)\nloader_show = torch.utils.data.DataLoader(dataset_show, batch_size=batch_size, shuffle=True, num_workers=num_workers)","metadata":{"execution":{"iopub.status.busy":"2022-12-08T05:51:45.013875Z","iopub.execute_input":"2022-12-08T05:51:45.014143Z","iopub.status.idle":"2022-12-08T05:51:45.027831Z","shell.execute_reply.started":"2022-12-08T05:51:45.014119Z","shell.execute_reply":"2022-12-08T05:51:45.02659Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"f, axarr = plt.subplots(2,4)\nfor p in range(4):\n    idx = p * 20\n    imgs, lbl = dataset_show[idx]\n    axarr[0, p].imshow(imgs[7][:3].permute(1, 2, 0))\n    axarr[1, p].imshow(imgs[7][-1])","metadata":{"execution":{"iopub.status.busy":"2022-12-08T05:51:45.029081Z","iopub.execute_input":"2022-12-08T05:51:45.030151Z","iopub.status.idle":"2022-12-08T05:51:47.670472Z","shell.execute_reply.started":"2022-12-08T05:51:45.030113Z","shell.execute_reply":"2022-12-08T05:51:47.669274Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model","metadata":{}},{"cell_type":"code","source":"class TimmModel(nn.Module):\n    def __init__(self, backbone, pretrained=False):\n        super(TimmModel, self).__init__()\n\n        self.encoder = timm.create_model(\n            backbone,\n            in_chans=in_chans,\n            num_classes=out_dim,\n            features_only=False,\n            drop_rate=drop_rate,\n#             drop_path_rate=drop_path_rate,\n            pretrained=pretrained\n        )\n\n        if 'efficient' in backbone:\n            hdim = self.encoder.conv_head.out_channels\n            self.encoder.classifier = nn.Identity()\n        elif 'densenet' in backbone:\n            hdim = self.encoder.classifier.in_features\n            self.encoder.classifier = nn.Identity()\n\n\n        self.lstm = nn.LSTM(hdim, 256, num_layers=2, dropout=drop_rate, bidirectional=True, batch_first=True)\n        self.head = nn.Sequential(\n            nn.Linear(512, 256),\n            nn.BatchNorm1d(256),\n            nn.Dropout(drop_rate_last),\n            nn.LeakyReLU(0.1),\n            nn.Linear(256, out_dim),\n        )\n\n    def forward(self, x):  # (bs, nslice, ch, sz, sz)\n        bs = x.shape[0]\n        x = x.view(bs * n_slice_per_c, in_chans, image_size, image_size)\n        feat = self.encoder(x)\n        feat = feat.view(bs, n_slice_per_c, -1)\n        feat, _ = self.lstm(feat)\n        feat = feat.contiguous().view(bs * n_slice_per_c, -1)\n        feat = self.head(feat)\n        feat = feat.view(bs, n_slice_per_c).contiguous()\n\n        return feat\n","metadata":{"execution":{"iopub.status.busy":"2022-12-08T05:51:47.674423Z","iopub.execute_input":"2022-12-08T05:51:47.675071Z","iopub.status.idle":"2022-12-08T05:51:47.688326Z","shell.execute_reply.started":"2022-12-08T05:51:47.675031Z","shell.execute_reply":"2022-12-08T05:51:47.687123Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"m = TimmModel(backbone)\nm(torch.rand(2, n_slice_per_c, in_chans, image_size, image_size)).shape","metadata":{"execution":{"iopub.status.busy":"2022-12-08T05:51:47.6899Z","iopub.execute_input":"2022-12-08T05:51:47.6919Z","iopub.status.idle":"2022-12-08T05:51:56.067427Z","shell.execute_reply.started":"2022-12-08T05:51:47.691862Z","shell.execute_reply":"2022-12-08T05:51:56.066245Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Loss & Metric","metadata":{}},{"cell_type":"code","source":"bce = nn.BCEWithLogitsLoss(reduction='none')\n\n\ndef criterion(logits, targets, activated=False):\n    if activated:\n        losses = nn.BCELoss(reduction='none')(logits.view(-1), targets.view(-1))\n    else:\n        losses = bce(logits.view(-1), targets.view(-1))\n    losses[targets.view(-1) > 0] *= 2.\n    norm = torch.ones(logits.view(-1).shape[0]).to(device)\n    norm[targets.view(-1) > 0] *= 2\n    return losses.sum() / norm.sum()","metadata":{"execution":{"iopub.status.busy":"2022-12-08T05:51:56.069065Z","iopub.execute_input":"2022-12-08T05:51:56.06953Z","iopub.status.idle":"2022-12-08T05:51:56.079068Z","shell.execute_reply.started":"2022-12-08T05:51:56.069493Z","shell.execute_reply":"2022-12-08T05:51:56.077277Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Train & Valid func","metadata":{}},{"cell_type":"code","source":"def mixup(input, truth, clip=[0, 1]):\n    indices = torch.randperm(input.size(0))\n    shuffled_input = input[indices]\n    shuffled_labels = truth[indices]\n\n    lam = np.random.uniform(clip[0], clip[1])\n    input = input * lam + shuffled_input * (1 - lam)\n    return input, truth, shuffled_labels, lam\n\n\ndef train_func(model, loader_train, optimizer, scaler=None):\n    model.train()\n    train_loss = []\n    bar = tqdm(loader_train)\n    for images, targets in bar:\n        optimizer.zero_grad()\n        images = images.cuda()\n        targets = targets.cuda()\n        \n        do_mixup = False\n        if random.random() < p_mixup:\n            do_mixup = True\n            images, targets, targets_mix, lam = mixup(images, targets)\n\n        with amp.autocast():\n            logits = model(images)\n            loss = criterion(logits, targets)\n            if do_mixup:\n                loss11 = criterion(logits, targets_mix)\n                loss = loss * lam  + loss11 * (1 - lam)\n        train_loss.append(loss.item())\n        scaler.scale(loss).backward()\n        scaler.step(optimizer)\n        scaler.update()\n\n        bar.set_description(f'smth:{np.mean(train_loss[-30:]):.4f}')\n\n    return np.mean(train_loss)\n\n\ndef valid_func(model, loader_valid):\n    model.eval()\n    valid_loss = []\n    gts = []\n    outputs = []\n    bar = tqdm(loader_valid)\n    with torch.no_grad():\n        for images, targets in bar:\n            images = images.cuda()\n            targets = targets.cuda()\n\n            logits = model(images)\n            loss = criterion(logits, targets)\n            \n            gts.append(targets.cpu())\n            outputs.append(logits.cpu())\n            valid_loss.append(loss.item())\n            \n            bar.set_description(f'smth:{np.mean(valid_loss[-30:]):.4f}')\n\n    outputs = torch.cat(outputs)\n    gts = torch.cat(gts)\n    valid_loss = criterion(outputs, gts).item()\n\n    return valid_loss\n","metadata":{"execution":{"iopub.status.busy":"2022-12-08T05:51:56.080912Z","iopub.execute_input":"2022-12-08T05:51:56.081738Z","iopub.status.idle":"2022-12-08T05:51:56.0959Z","shell.execute_reply.started":"2022-12-08T05:51:56.08169Z","shell.execute_reply":"2022-12-08T05:51:56.095019Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"rcParams['figure.figsize'] = 20, 2\noptimizer = optim.AdamW(m.parameters(), lr=init_lr)\nscheduler_cosine = torch.optim.lr_scheduler.CosineAnnealingWarmRestarts(optimizer, n_epochs, eta_min=eta_min)\n\nlrs = []\nfor epoch in range(1, n_epochs+1):\n    scheduler_cosine.step(epoch-1)\n    lrs.append(optimizer.param_groups[0][\"lr\"])\nplt.plot(range(len(lrs)), lrs)","metadata":{"execution":{"iopub.status.busy":"2022-12-08T05:51:56.09723Z","iopub.execute_input":"2022-12-08T05:51:56.097601Z","iopub.status.idle":"2022-12-08T05:51:56.308386Z","shell.execute_reply.started":"2022-12-08T05:51:56.097567Z","shell.execute_reply":"2022-12-08T05:51:56.307463Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training","metadata":{}},{"cell_type":"code","source":"def run(fold):\n\n    log_file = os.path.join(log_dir, f'{kernel_type}.txt')\n    model_file = os.path.join(model_dir, f'{kernel_type}_fold{fold}_best.pth')\n\n    train_ = df[df['fold'] != fold].reset_index(drop=True)\n    valid_ = df[df['fold'] == fold].reset_index(drop=True)\n    dataset_train = CLSDataset(train_, 'train', transform=transforms_train)\n    dataset_valid = CLSDataset(valid_, 'valid', transform=transforms_valid)\n    loader_train = torch.utils.data.DataLoader(dataset_train, batch_size=batch_size, shuffle=True, num_workers=num_workers, drop_last=True)\n    loader_valid = torch.utils.data.DataLoader(dataset_valid, batch_size=batch_size, shuffle=False, num_workers=num_workers)\n\n    model = TimmModel(backbone, pretrained=True)\n    model = model.to(device)\n\n    optimizer = optim.AdamW(model.parameters(), lr=init_lr)\n    scaler = torch.cuda.amp.GradScaler() if use_amp else None\n\n    metric_best = np.inf\n    loss_min = np.inf\n\n    scheduler_cosine = torch.optim.lr_scheduler.CosineAnnealingWarmRestarts(optimizer, n_epochs, eta_min=eta_min)\n\n    print(len(dataset_train), len(dataset_valid))\n\n    for epoch in range(1, n_epochs+1):\n        scheduler_cosine.step(epoch-1)\n\n        print(time.ctime(), 'Epoch:', epoch)\n\n        train_loss = train_func(model, loader_train, optimizer, scaler)\n        valid_loss = valid_func(model, loader_valid)\n        metric = valid_loss\n\n        content = time.ctime() + ' ' + f'Fold {fold}, Epoch {epoch}, lr: {optimizer.param_groups[0][\"lr\"]:.7f}, train loss: {train_loss:.5f}, valid loss: {valid_loss:.5f}, metric: {(metric):.6f}.'\n        print(content)\n        with open(log_file, 'a') as appender:\n            appender.write(content + '\\n')\n\n        if metric < metric_best:\n            print(f'metric_best ({metric_best:.6f} --> {metric:.6f}). Saving model ...')\n#             if not DEBUG:\n            torch.save(model.state_dict(), model_file)\n            metric_best = metric\n\n        # Save Last\n        if not DEBUG:\n            torch.save(\n                {\n                    'epoch': epoch,\n                    'model_state_dict': model.state_dict(),\n                    'optimizer_state_dict': optimizer.state_dict(),\n                    'scaler_state_dict': scaler.state_dict() if scaler else None,\n                    'score_best': metric_best,\n                },\n                model_file.replace('_best', '_last')\n            )\n\n    del model\n    torch.cuda.empty_cache()\n    gc.collect()\n","metadata":{"execution":{"iopub.status.busy":"2022-12-08T05:51:56.309861Z","iopub.execute_input":"2022-12-08T05:51:56.31057Z","iopub.status.idle":"2022-12-08T05:51:56.323698Z","shell.execute_reply.started":"2022-12-08T05:51:56.310534Z","shell.execute_reply":"2022-12-08T05:51:56.322667Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"run(0)","metadata":{"execution":{"iopub.status.busy":"2022-12-08T05:51:56.326699Z","iopub.execute_input":"2022-12-08T05:51:56.327576Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# run(1)\n# run(2)\n# run(3)\n# run(4)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}