{"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":"MODEL_INP_PATH = '/kaggle/input/rsna23-train-stage2-final-4'\nEXTRA_EPOCHS = 10","metadata":{},"execution_count":null,"outputs":[]},{"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":{"editable":false}},{"cell_type":"code","source":"!pip -q install timm","metadata":{"editable":false,"execution":{"iopub.status.busy":"2023-10-15T06:20:14.141548Z","iopub.execute_input":"2023-10-15T06:20:14.142422Z","iopub.status.idle":"2023-10-15T06:20:26.284719Z","shell.execute_reply.started":"2023-10-15T06:20:14.142388Z","shell.execute_reply":"2023-10-15T06:20:26.283408Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"DEBUG = False","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","editable":false,"execution":{"iopub.status.busy":"2023-10-15T06:20:26.287097Z","iopub.execute_input":"2023-10-15T06:20:26.288168Z","iopub.status.idle":"2023-10-15T06:20:26.29355Z","shell.execute_reply.started":"2023-10-15T06:20:26.288128Z","shell.execute_reply":"2023-10-15T06:20:26.292117Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from fastcore.all import Path","metadata":{"editable":false,"execution":{"iopub.status.busy":"2023-10-15T06:20:26.295139Z","iopub.execute_input":"2023-10-15T06:20:26.2957Z","iopub.status.idle":"2023-10-15T06:20:26.336307Z","shell.execute_reply.started":"2023-10-15T06:20:26.295666Z","shell.execute_reply":"2023-10-15T06:20:26.335429Z"},"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\nimport sklearn\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') if torch.cuda.is_available() else torch.device('cpu')\ntorch.backends.cudnn.benchmark = True\nprint('device is', device)","metadata":{"execution":{"iopub.status.busy":"2023-10-15T06:20:26.338805Z","iopub.execute_input":"2023-10-15T06:20:26.339312Z","iopub.status.idle":"2023-10-15T06:20:29.013217Z","shell.execute_reply.started":"2023-10-15T06:20:26.339279Z","shell.execute_reply":"2023-10-15T06:20:29.012121Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Config","metadata":{"editable":false}},{"cell_type":"code","source":"CROPS = Path('')\n\nORGANS = ['liver', 'spleen', 'kidney', 'bowel']\nLABELS = [['liver_healthy', 'liver_low', 'liver_high'], \n          ['spleen_healthy', 'spleen_low', 'spleen_high'], \n          ['kidney_healthy', 'kidney_low', 'kidney_high'], \n          ['bowel_healthy', 'bowel_injury']\n         ]\nABS = [[0, 15], [15, 30], [30, 60], [60, 75]]\nN_ORGANS = len(ORGANS)","metadata":{"execution":{"iopub.status.busy":"2023-10-15T06:20:29.01466Z","iopub.execute_input":"2023-10-15T06:20:29.015664Z","iopub.status.idle":"2023-10-15T06:20:29.022253Z","shell.execute_reply.started":"2023-10-15T06:20:29.015624Z","shell.execute_reply":"2023-10-15T06:20:29.0212Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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 = 'tf_efficientnetv2_s_in21ft1k'\n\nimage_size = 224\nn_slice_per_c = 15\nin_chans = 6\n\ninit_lr = 23e-5\neta_min = 23e-6\nbatch_size = 4\ndrop_rate = 0.\ndrop_rate_last = 0.3\ndrop_path_rate = 0.\np_mixup = 0.5\np_rand_order_v1 = 0.2\n\n# data_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 = 75\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":"2023-10-15T06:31:36.064924Z","iopub.execute_input":"2023-10-15T06:31:36.065667Z","iopub.status.idle":"2023-10-15T06:31:36.072875Z","shell.execute_reply.started":"2023-10-15T06:31:36.065632Z","shell.execute_reply":"2023-10-15T06:31:36.07187Z"},"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_count":null,"outputs":[]},{"cell_type":"markdown","source":"# DataFrame","metadata":{"editable":false}},{"cell_type":"code","source":"INPUT = '/kaggle/input/rsna-2023-abdominal-trauma-detection'\ndef load_df(kind='train'):    \n    df = pd.read_parquet(os.path.join(INPUT, f'{kind}_dicom_tags.parquet'))\n    df['StudyInstanceUID'] = df.path.str.split('/').str[-2]\n\n    df = df[['StudyInstanceUID', 'path', 'PatientID']].drop_duplicates('StudyInstanceUID')\n    df['image_folder'] = INPUT + '/' + df.path.str.split('/').str[:-1].apply('/'.join)\n    df['study'] = df.StudyInstanceUID\n    df['patient'] = df.PatientID\n    df['patient_id'] = df.PatientID.astype(int)\n    return df\ndf = load_df('train')\ndf_train = pd.read_csv(os.path.join(INPUT, 'train.csv'))","metadata":{"execution":{"iopub.status.busy":"2023-10-15T06:20:29.057995Z","iopub.execute_input":"2023-10-15T06:20:29.058903Z","iopub.status.idle":"2023-10-15T06:20:39.133586Z","shell.execute_reply.started":"2023-10-15T06:20:29.05887Z","shell.execute_reply":"2023-10-15T06:20:39.132617Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"IN = Path('../input')","metadata":{"execution":{"iopub.status.busy":"2023-10-15T06:20:39.134905Z","iopub.execute_input":"2023-10-15T06:20:39.135787Z","iopub.status.idle":"2023-10-15T06:20:39.140531Z","shell.execute_reply.started":"2023-10-15T06:20:39.135749Z","shell.execute_reply":"2023-10-15T06:20:39.139613Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"cls_inp_paths = [x for x in IN.ls() if 's1-inf' in x.stem]\nlen(cls_inp_paths), cls_inp_paths[0]","metadata":{"execution":{"iopub.status.busy":"2023-10-15T06:20:39.14498Z","iopub.execute_input":"2023-10-15T06:20:39.145252Z","iopub.status.idle":"2023-10-15T06:20:39.15722Z","shell.execute_reply.started":"2023-10-15T06:20:39.14523Z","shell.execute_reply":"2023-10-15T06:20:39.156308Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for p in cls_inp_paths:\n    for f in p.ls(): \n        if not str(f)[-3:] == 'npy': continue\n#         print(f)\n        study = f.stem\n#         print(study)\n        df.loc[df['study'] == study, 'cls_inp_path'] = str(f)","metadata":{"execution":{"iopub.status.busy":"2023-10-15T06:20:39.158559Z","iopub.execute_input":"2023-10-15T06:20:39.158911Z","iopub.status.idle":"2023-10-15T06:20:45.596621Z","shell.execute_reply.started":"2023-10-15T06:20:39.15887Z","shell.execute_reply":"2023-10-15T06:20:45.595601Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(df.shape[0])\ndf = df[df.cls_inp_path.notnull()]\nprint(df.shape[0], 'size after removing studies with no inputs')","metadata":{"execution":{"iopub.status.busy":"2023-10-15T06:20:45.597954Z","iopub.execute_input":"2023-10-15T06:20:45.598877Z","iopub.status.idle":"2023-10-15T06:20:45.610225Z","shell.execute_reply.started":"2023-10-15T06:20:45.598839Z","shell.execute_reply":"2023-10-15T06:20:45.609194Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = df.merge(df_train, on='patient_id')\nassert df.isna().sum().sum() == 0","metadata":{"editable":false,"execution":{"iopub.status.busy":"2023-10-15T06:20:45.611659Z","iopub.execute_input":"2023-10-15T06:20:45.612368Z","iopub.status.idle":"2023-10-15T06:20:45.634396Z","shell.execute_reply.started":"2023-10-15T06:20:45.612326Z","shell.execute_reply":"2023-10-15T06:20:45.633527Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for organ, cols in zip(ORGANS, LABELS): \n    print(organ)\n    display(df[['patient', 'study'] + cols].head(2))","metadata":{"editable":false,"execution":{"iopub.status.busy":"2023-10-15T06:20:45.635655Z","iopub.execute_input":"2023-10-15T06:20:45.635979Z","iopub.status.idle":"2023-10-15T06:20:45.671107Z","shell.execute_reply.started":"2023-10-15T06:20:45.635947Z","shell.execute_reply":"2023-10-15T06:20:45.67011Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"kf = sklearn.model_selection.StratifiedGroupKFold()\ndf['fold'] = 0\nfor fold, (_, test_ind) in enumerate(kf.split(df, df['any_injury'], df['patient'])): \n#     print(len(_), len(test_ind))\n    df.iloc[test_ind, -1] = fold","metadata":{"editable":false,"execution":{"iopub.status.busy":"2023-10-15T06:20:45.67231Z","iopub.execute_input":"2023-10-15T06:20:45.673106Z","iopub.status.idle":"2023-10-15T06:20:46.56733Z","shell.execute_reply.started":"2023-10-15T06:20:45.673073Z","shell.execute_reply":"2023-10-15T06:20:46.566345Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df.groupby('fold')[list(df)[-15:]].sum()","metadata":{"editable":false,"execution":{"iopub.status.busy":"2023-10-15T06:20:46.568876Z","iopub.execute_input":"2023-10-15T06:20:46.569502Z","iopub.status.idle":"2023-10-15T06:20:46.588869Z","shell.execute_reply.started":"2023-10-15T06:20:46.569463Z","shell.execute_reply":"2023-10-15T06:20:46.587899Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Dataset","metadata":{"editable":false}},{"cell_type":"code","source":"from collections import defaultdict","metadata":{"execution":{"iopub.status.busy":"2023-10-15T06:20:46.590082Z","iopub.execute_input":"2023-10-15T06:20:46.590647Z","iopub.status.idle":"2023-10-15T06:20:46.595177Z","shell.execute_reply.started":"2023-10-15T06:20:46.590612Z","shell.execute_reply":"2023-10-15T06:20:46.594096Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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        \n        \n        image_full = np.load(row.cls_inp_path)\n        out = defaultdict(dict)\n        for organ, cols, (a, b) in zip(ORGANS, LABELS, ABS): \n            images = []\n            for image in image_full[a: b]: \n                image = image.transpose(1, 2, 0)\n                image = transforms_train(image=image)['image']\n                image = image.transpose(2, 0, 1)\n                images.append(image)\n            images = np.stack(images, 0)\n            if organ == 'kidney': \n                images = np.concatenate((images[:15, :, :, :], images[15:, :, :, :]), 2)\n            out[organ]['images'] = torch.tensor(images).float()\n            out[organ]['labels'] = torch.tensor([row[cols]] * n_slice_per_c).float()\n        return out\n\n","metadata":{"execution":{"iopub.status.busy":"2023-10-15T06:20:46.596716Z","iopub.execute_input":"2023-10-15T06:20:46.597524Z","iopub.status.idle":"2023-10-15T06:20:46.608027Z","shell.execute_reply.started":"2023-10-15T06:20:46.597491Z","shell.execute_reply":"2023-10-15T06:20:46.607182Z"},"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":{"editable":false,"execution":{"iopub.status.busy":"2023-10-15T06:20:46.609348Z","iopub.execute_input":"2023-10-15T06:20:46.610599Z","iopub.status.idle":"2023-10-15T06:20:46.623927Z","shell.execute_reply.started":"2023-10-15T06:20:46.610566Z","shell.execute_reply":"2023-10-15T06:20:46.622839Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"f, axarr = plt.subplots(2,4)\nfor p in range(2):\n    idx = p * 20\n    out = dataset_show[idx]\nf, axarr = plt.subplots(2,4)\nfor i, organ in enumerate(out.keys()): \n    print('*******', organ, '*******')\n    axarr[0, i].imshow(out[organ]['images'][7][:3].permute(1, 2, 0))\n    axarr[1, i].imshow(out[organ]['images'][7][-1])\n#     print(out[organ]['labels'])","metadata":{"execution":{"iopub.status.busy":"2023-10-15T06:20:46.625479Z","iopub.execute_input":"2023-10-15T06:20:46.626251Z","iopub.status.idle":"2023-10-15T06:20:51.186673Z","shell.execute_reply.started":"2023-10-15T06:20:46.626216Z","shell.execute_reply":"2023-10-15T06:20:51.185748Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model","metadata":{"editable":false}},{"cell_type":"code","source":"class TimmModel(nn.Module):\n    def __init__(self, backbone, pretrained=False, out_dim=3, h=image_size, w=image_size):\n        super(TimmModel, self).__init__()\n        self.h = h\n        self.w = w\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 'convnext' in backbone:\n            hdim = self.encoder.head.fc.in_features\n            self.encoder.head.fc = 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), # chacnged\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, self.h, self.w)\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, -1).contiguous()\n\n        return feat","metadata":{"execution":{"iopub.status.busy":"2023-10-15T06:20:51.187662Z","iopub.execute_input":"2023-10-15T06:20:51.187973Z","iopub.status.idle":"2023-10-15T06:20:51.202805Z","shell.execute_reply.started":"2023-10-15T06:20:51.187942Z","shell.execute_reply":"2023-10-15T06:20:51.201893Z"},"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":"2023-10-15T06:20:51.204454Z","iopub.execute_input":"2023-10-15T06:20:51.205286Z","iopub.status.idle":"2023-10-15T06:20:59.920531Z","shell.execute_reply.started":"2023-10-15T06:20:51.205249Z","shell.execute_reply":"2023-10-15T06:20:59.919584Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Loss & Metric","metadata":{"editable":false}},{"cell_type":"code","source":"bce = nn.BCEWithLogitsLoss(reduction='none')\ndef criterion(logits, targets, activated=False):\n    n_labels = targets.shape[-1]\n    if activated:\n        losses = nn.BCELoss(reduction='none')(logits.view(-1), targets.view(-1))\n    else:\n        losses = bce(logits.view(-1, n_labels), targets.view(-1, n_labels))\n    norm = torch.ones(logits.view(-1, n_labels).shape).to(device)\n    for i, weight in [[1, 2], [2, 4]]: \n        if i == n_labels: break\n        mask = (targets.view(-1, n_labels)[:, i] == 1)    \n        losses[mask] *= weight\n        norm[mask] *= weight\n    return losses.sum() / norm.sum()","metadata":{"execution":{"iopub.status.busy":"2023-10-15T06:20:59.921851Z","iopub.execute_input":"2023-10-15T06:20:59.922703Z","iopub.status.idle":"2023-10-15T06:20:59.931497Z","shell.execute_reply.started":"2023-10-15T06:20:59.922664Z","shell.execute_reply":"2023-10-15T06:20:59.9304Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Train & Valid func","metadata":{"editable":false}},{"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(models, loader_train, optimizers, scalers=None):\n    [model.train() for model in models]\n    train_loss = [[] for _ in range(N_ORGANS)]\n    bar = tqdm(loader_train)\n    for out in bar:\n        for i, (optimizer, organ, scaler, model) in enumerate(zip(optimizers, ORGANS, scalers, models)): \n            optimizer.zero_grad()\n            images = out[organ]['images'].to(device)\n            targets = out[organ]['labels'].to(device)\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[i].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(tl) for tl in train_loss]\n\n\ndef valid_func(models, loader_valid):\n    [model.train() for model in models]\n    valid_loss = [[] for _ in range(N_ORGANS)]\n    gts = [[] for _ in range(N_ORGANS)]\n    outputs = [[] for _ in range(N_ORGANS)]\n    bar = tqdm(loader_valid)\n    with torch.no_grad():\n        for out in bar:\n            for i, (organ, model) in enumerate(zip(ORGANS, models)):\n                images = out[organ]['images'].to(device) \n                targets = out[organ]['labels'].to(device)\n\n                logits = model(images)\n                loss = criterion(logits, targets)\n\n                gts[i].append(targets.cpu())\n                outputs[i].append(logits.cpu())\n                valid_loss[i].append(loss.item())\n\n#                 bar.set_description(f'smth:{np.mean(valid_loss[-30:]):.4f}')\n\n    outputs = [torch.cat(output) for output in outputs]\n    gts = [torch.cat(gt) for gt in gts]\n    valid_loss = [criterion(o, g).item() for o, g in zip(outputs, gts)]\n\n    return valid_loss","metadata":{"execution":{"iopub.status.busy":"2023-10-15T06:20:59.932926Z","iopub.execute_input":"2023-10-15T06:20:59.933699Z","iopub.status.idle":"2023-10-15T06:20:59.95325Z","shell.execute_reply.started":"2023-10-15T06:20:59.933665Z","shell.execute_reply":"2023-10-15T06:20:59.952255Z"},"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.CosineAnnealingLR(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":{"editable":false,"execution":{"iopub.status.busy":"2023-10-15T06:20:59.954688Z","iopub.execute_input":"2023-10-15T06:20:59.955058Z","iopub.status.idle":"2023-10-15T06:21:00.165702Z","shell.execute_reply.started":"2023-10-15T06:20:59.955006Z","shell.execute_reply":"2023-10-15T06:21:00.164719Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training","metadata":{"editable":false}},{"cell_type":"code","source":"n_epochs","metadata":{"execution":{"iopub.status.busy":"2023-10-15T06:29:55.32314Z","iopub.execute_input":"2023-10-15T06:29:55.324274Z","iopub.status.idle":"2023-10-15T06:29:55.330678Z","shell.execute_reply.started":"2023-10-15T06:29:55.324215Z","shell.execute_reply":"2023-10-15T06:29:55.329578Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def run(fold):\n\n    log_file = os.path.join(log_dir, f'{kernel_type}.txt')\n    model_files = [os.path.join(model_dir, f'{organ}_{kernel_type}_fold{fold}_best.pth') for organ in ORGANS]\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    models = [TimmModel(backbone, pretrained=True), \n             TimmModel(backbone, pretrained=True), \n             TimmModel(backbone, pretrained=True, h=image_size*2), \n             TimmModel(backbone, pretrained=True, out_dim=2), ]\n    models = [model.to(device) for model in models]\n\n    optimizers = [optim.AdamW(model.parameters(), lr=init_lr) for model in models]\n    scalers = [torch.cuda.amp.GradScaler() if use_amp else None for _ in range(N_ORGANS)]\n\n    metric_best = [np.inf for _ in range(N_ORGANS)]\n    epoch_best = [0 for _ in range(N_ORGANS)]\n#     loss_mins = [np.inf for _ in range(N_ORGANS)]\n\n    scheduler_cosines = [torch.optim.lr_scheduler.CosineAnnealingWarmRestarts(optimizer, n_epochs, eta_min=eta_min) for optimizer in optimizers]\n\n    print(len(dataset_train), len(dataset_valid))\n    for i, (organ, model, scaler, optimizer) in enumerate(zip(ORGANS, models, scalers, optimizers)):\n        sd_file = f'{MODEL_INP_PATH}/models/{organ}_0920_1bonev2_effv2s_224_15_6ch_augv2_mixupp5_drl3_rov1p2_bs8_lr23e5_eta23e6_50ep_fold0_last.pth'\n        sd = torch.load(sd_file,  map_location=torch.device('cpu'))\n        msd = sd['model_state_dict']\n        msd = {k[7:] if k.startswith('module.') else k: msd[k] for k in msd.keys()}\n        model.load_state_dict(msd, strict=True)\n        optimizer.load_state_dict(sd['optimizer_state_dict'])\n        scaler.load_state_dict(sd['scaler_state_dict'])\n        metric_best[i] = sd['score_best']\n        epoch_start = sd['epoch'] + 1\n    print(epoch_start, 'epoch start')\n        \n    print('next')\n    for epoch in range(epoch_start, epoch_start + EXTRA_EPOCHS):\n        print('************* EPOCH {epoch} *********************')\n        scheduler_cosine.step(epoch-1)\n\n        print(time.ctime(), 'Epoch:', epoch)\n\n        train_losss = train_func(models, loader_train, optimizers, scalers)\n        valid_losss = valid_func(models, loader_valid)\n        metrics = valid_losss\n\n        for organ, optimizer, train_loss, valid_loss in zip(ORGANS, optimizers, train_losss, valid_losss): \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: {(valid_loss):.6f}.'\n            print(content)\n            with open(log_file, 'a') as appender:\n                appender.write(content + '\\n')\n            \n        for i, (organ, metric, model_file, model, scaler, optimizer) in enumerate(zip(ORGANS, metrics, model_files, models, scalers, optimizers)): \n            if metric < metric_best[i]:\n                print(f'{organ} metric_best ({metric_best[i]:.6f} --> {metric:.6f}). Saving model ...')\n    #             if not DEBUG:\n                torch.save(model.state_dict(), model_file)\n                metric_best[i] = metric\n                epoch_best[i] = epoch\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[i],\n                        'epoch_best': epoch_best[i],\n                    },\n                    model_file.replace('_best', '_last')\n                )\n\n    del models\n    torch.cuda.empty_cache()\n    gc.collect()","metadata":{"execution":{"iopub.status.busy":"2023-10-15T06:26:15.158379Z","iopub.execute_input":"2023-10-15T06:26:15.158731Z","iopub.status.idle":"2023-10-15T06:26:15.176583Z","shell.execute_reply.started":"2023-10-15T06:26:15.158705Z","shell.execute_reply":"2023-10-15T06:26:15.175431Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"run(0)\n# run(1)\n# run(2)\n# run(3)\n# run(4)","metadata":{"execution":{"iopub.status.busy":"2023-10-15T06:34:09.373073Z","iopub.execute_input":"2023-10-15T06:34:09.373868Z","iopub.status.idle":"2023-10-15T06:34:38.972962Z","shell.execute_reply.started":"2023-10-15T06:34:09.373831Z","shell.execute_reply":"2023-10-15T06:34:38.970329Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}