{"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"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":52254,"databundleVersionId":6863140,"sourceType":"competition"},{"sourceId":3899138,"sourceType":"datasetVersion","datasetId":2316152},{"sourceId":4204135,"sourceType":"datasetVersion","datasetId":2478634}],"dockerImageVersionId":30301,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# 1st Place Solution Training 3D Semantic Segmentation (Stage1)\n\nGreetings,\n\nWe're thrilled to share this notebook along with a summary of our winning solution.\n\nThis represents the FULL version of our training for the final 3D semantic segmentation models in Stage 1. We use resnet18d for the backbone, unet for the decoder, and utilized an input dimension of 128x128x128.\n\nNOTE: **It's imperative to run this code locally, as Kaggle's notebook does not have the requisite RAM.**\n\nNOTE2: **To expedite data loading, it's highly advised to preprocess the 3D semantic segmentation training data and save it on your local machine.**\n\nOur brief summary of winning solution: https://www.kaggle.com/competitions/rsna-2023-abdominal-trauma-detection/discussion/447449\n\n\n**If you find these notebooks helpful please upvote. Thanks!**\n\n---\n\n\n**Acknowledge: As a Z by HP Global Data Science Ambassador, I used the Z8G4 workstation with 2xA6000 provided by Z by HP. This machine helped me immensely in the competition. **","metadata":{}},{"cell_type":"code","source":"!pip -q install monai\n!pip -q install segmentation-models-pytorch","metadata":{"execution":{"iopub.status.busy":"2023-10-16T03:42:16.949023Z","iopub.execute_input":"2023-10-16T03:42:16.949592Z","iopub.status.idle":"2023-10-16T03:42:35.625541Z","shell.execute_reply.started":"2023-10-16T03:42:16.949467Z","shell.execute_reply":"2023-10-16T03:42:35.624223Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"DEBUG = False\n\nimport os\nimport sys\nsys.path = [\n    '../input/covn3d-same',\n] + sys.path","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-10-16T03:42:35.627826Z","iopub.execute_input":"2023-10-16T03:42:35.628126Z","iopub.status.idle":"2023-10-16T03:42:35.634271Z","shell.execute_reply.started":"2023-10-16T03:42:35.628098Z","shell.execute_reply":"2023-10-16T03:42:35.633253Z"},"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 pydicom\nimport argparse\nimport warnings\nimport numpy as np\nimport pandas as pd\nfrom glob import glob\nimport nibabel as nib\nfrom PIL import Image\nfrom tqdm import tqdm\nimport albumentations\nfrom pylab import rcParams\nimport matplotlib.pyplot as plt\nimport segmentation_models_pytorch as smp\nfrom sklearn.model_selection import StratifiedGroupKFold\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\nfrom monai.transforms import Resize\nimport  monai.transforms as transforms\n\n%matplotlib inline\nrcParams['figure.figsize'] = 20, 8\ndevice = torch.device('cuda')\ntorch.backends.cudnn.benchmark = True","metadata":{"execution":{"iopub.status.busy":"2023-10-16T03:42:35.635684Z","iopub.execute_input":"2023-10-16T03:42:35.636868Z","iopub.status.idle":"2023-10-16T03:42:44.40387Z","shell.execute_reply.started":"2023-10-16T03:42:35.636828Z","shell.execute_reply":"2023-10-16T03:42:44.402767Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Config","metadata":{}},{"cell_type":"code","source":"kernel_type =  f'timm3d_{backbone}_unet4b_128_128_128_dsv2_flip12_shift333p7_gd1p5_bs4_lr3e4_20x50ep'\nload_kernel = None\nload_last = True\nn_blocks = 4\nn_folds = 5\nbackbone = 'leakyrelu'image_sizes = [128, 128, 128]\nR = Resize(image_sizes)\n\ninit_lr = 3e-3\nbatch_size = 4\ndrop_rate = 0.\ndrop_path_rate = 0.\nloss_weights = [1, 1]\np_mixup = 0.1\n\ndata_dir = '../input/rsna-2023-abdominal-trauma-detection'\nuse_amp = True\nnum_workers = 4\nout_dim = 6\n\nn_epochs = 1000\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-16T03:42:44.40704Z","iopub.execute_input":"2023-10-16T03:42:44.408943Z","iopub.status.idle":"2023-10-16T03:42:44.41659Z","shell.execute_reply.started":"2023-10-16T03:42:44.408898Z","shell.execute_reply":"2023-10-16T03:42:44.415488Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"transforms_train = transforms.Compose([\n    transforms.RandFlipd(keys=[\"image\", \"mask\"], prob=0.5, spatial_axis=1),\n    transforms.RandFlipd(keys=[\"image\", \"mask\"], prob=0.5, spatial_axis=2),\n    transforms.RandAffined(keys=[\"image\", \"mask\"], translate_range=[int(x*y) for x, y in zip(image_sizes, [0.3, 0.3, 0.3])], padding_mode='zeros', prob=0.7),\n    transforms.RandGridDistortiond(keys=(\"image\", \"mask\"), prob=0.5, distort_limit=(-0.01, 0.01), mode=\"nearest\"),    \n])\n\ntransforms_valid = transforms.Compose([\n])","metadata":{"execution":{"iopub.status.busy":"2023-10-16T03:42:44.418012Z","iopub.execute_input":"2023-10-16T03:42:44.418414Z","iopub.status.idle":"2023-10-16T03:42:44.435117Z","shell.execute_reply.started":"2023-10-16T03:42:44.418378Z","shell.execute_reply":"2023-10-16T03:42:44.434113Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# DataFrame","metadata":{}},{"cell_type":"code","source":"df_train = pd.read_csv(os.path.join(data_dir, 'train.csv'))\n\nall_samples = sorted(glob(os.path.join(data_dir, 'train_images', '*/*')))\nall_series_id = [sample.split('/')[-1] for sample in all_samples]\nall_patient_id = [sample.split('/')[-2] for sample in all_samples]\ndf_tmp = pd.DataFrame({\n    'patient_id': all_patient_id,\n    'series_id': all_series_id\n})\ndf_tmp['patient_id'] = df_tmp['patient_id'].astype(int)\ndf_train = df_tmp.merge(df_train, on='patient_id', how='left')\n\n\nmask_files = os.listdir(f'{data_dir}/segmentations')\n\ndf_mask = pd.DataFrame({\n    'mask_file': mask_files,\n    'has_mask': 1\n})\ndf_mask['series_id'] = df_mask['mask_file'].apply(lambda x: x[:-4])\n\ndf_mask['mask_file'] = df_mask['mask_file'].apply(lambda x: os.path.join(data_dir, 'segmentations', x))\ndf_train = df_train.merge(df_mask, on='series_id', how='left')\ndf_train['mask_file'].fillna('', inplace=True)\ndf_train['has_mask'] = df_train['has_mask'].fillna(0).astype(int)\n\ndf_train['image_folder'] = df_train.apply(lambda row: os.path.join(data_dir, 'train_images', str(row.patient_id), str(row.series_id)), axis=1)\n\n\ndf_seg = df_train.query('mask_file != \"\"').reset_index(drop=True)\ndf_seg['tmp'] = df_seg.apply(lambda row: str(row.bowel_healthy)+str(row.extravasation_healthy)+str(row.kidney_healthy)+str(row.kidney_low)+str(row.liver_healthy)+str(row.liver_low)+str(row.spleen_healthy)+str(row.spleen_low), axis=1)\n\nsgkf = StratifiedGroupKFold(5)\ndf_seg['fold'] = -1\nfor fold, (train_idx, valid_idx) in enumerate(sgkf.split(df_seg, df_seg['tmp'], groups=df_seg['patient_id'])):\n    df_seg.loc[valid_idx, 'fold'] = fold\n\ndel df_seg['tmp']\n\ndf_seg.tail()","metadata":{"execution":{"iopub.status.busy":"2023-10-16T03:42:44.436733Z","iopub.execute_input":"2023-10-16T03:42:44.437656Z","iopub.status.idle":"2023-10-16T03:42:51.296673Z","shell.execute_reply.started":"2023-10-16T03:42:44.437604Z","shell.execute_reply":"2023-10-16T03:42:51.295582Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Dataset","metadata":{}},{"cell_type":"code","source":"def load_dicom(path):\n    dicom = pydicom.read_file(path)\n    data = dicom.pixel_array\n    data = cv2.resize(data, (image_sizes[0], image_sizes[1]), interpolation = cv2.INTER_LINEAR)\n    return data\n\n\ndef load_dicom_line_par(path):\n\n    t_paths = sorted(glob(os.path.join(path, \"*\")),\n       key=lambda x: int(x.split('/')[-1].split(\".\")[0]))\n\n    n_scans = len(t_paths)\n    indices = np.quantile(list(range(n_scans)), np.linspace(0., 1., image_sizes[2])).round().astype(int)\n    t_paths = [t_paths[i] for i in indices]\n\n    images = []\n    for filename in t_paths:\n        images.append(load_dicom(filename))\n    images = np.stack(images, -1)\n    \n    images = images - np.min(images)\n    images = images / (np.max(images) + 1e-4)\n    images = (images * 255).astype(np.uint8)\n\n    return images\n\n\ndef load_sample(row):\n\n    image = load_dicom_line_par(row.image_folder)\n    if image.ndim < 4:\n        image = np.expand_dims(image, 0).repeat(3, 0)  # to 3ch\n\n    mask_org = nib.load(row.mask_file).get_fdata()\n    shape = mask_org.shape\n    mask_org = mask_org.transpose(1, 0, 2)[::-1, :, ::-1]  # (d, w, h)\n    mask = np.zeros((7, shape[0], shape[1], shape[2]))\n    for cid in range(7):\n        mask[cid] = (mask_org == (cid+1))\n    mask = mask.astype(np.uint8) * 255\n    mask = R(mask).numpy()\n\n    return image, mask\n\n\n\nclass SEGDataset(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        ### using local cache\n#         image_file = os.path.join(data_dir, f'{row.StudyInstanceUID}.npy')\n#         mask_file = os.path.join(data_dir, f'{row.StudyInstanceUID}_mask.npy')\n#         image = np.load(image_file).astype(np.float32)\n#         mask = np.load(mask_file).astype(np.float32)\n\n        image, mask = load_sample(row)\n\n        res = self.transform({'image':image, 'mask':mask})\n        image = res['image'] / 255.\n        mask = res['mask']\n        mask = (mask > 127).astype(np.float32)\n\n        image, mask = torch.tensor(image).float(), torch.tensor(mask).float()\n\n        return image, mask\n","metadata":{"execution":{"iopub.status.busy":"2023-10-16T03:42:51.298585Z","iopub.execute_input":"2023-10-16T03:42:51.298983Z","iopub.status.idle":"2023-10-16T03:42:51.315424Z","shell.execute_reply.started":"2023-10-16T03:42:51.29894Z","shell.execute_reply":"2023-10-16T03:42:51.314351Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"rcParams['figure.figsize'] = 20,8\n\ndf_show = df_seg\ndataset_show = SEGDataset(df_show, 'train', transform=transforms_train)","metadata":{"execution":{"iopub.status.busy":"2023-10-16T03:42:51.316719Z","iopub.execute_input":"2023-10-16T03:42:51.317703Z","iopub.status.idle":"2023-10-16T03:42:51.334369Z","shell.execute_reply.started":"2023-10-16T03:42:51.317663Z","shell.execute_reply":"2023-10-16T03:42:51.3334Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"f, axarr = plt.subplots(1,2)\nfor idx in range(2):\n    img, mask = dataset_show[idx]\n    img = img[:, :, :, 60]\n    mask = mask[:, :, :, 60]\n    mask[0] = mask[0] + mask[3] + mask[6]\n    mask[1] = mask[1] + mask[4]\n    mask[2] = mask[2] + mask[5]\n    mask = mask[:3]\n    img = img * 0.7 + mask * 0.3\n    axarr[p].imshow(img.transpose(0, 1).transpose(1,2).squeeze())","metadata":{"execution":{"iopub.status.busy":"2023-10-16T03:43:06.804158Z","iopub.execute_input":"2023-10-16T03:43:06.804797Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model","metadata":{}},{"cell_type":"code","source":"class TimmSegModel(nn.Module):\n    def __init__(self, backbone, segtype='unet', pretrained=False):\n        super(TimmSegModel, self).__init__()\n\n        self.encoder = timm.create_model(\n            backbone,\n            in_chans=3,\n            features_only=True,\n            drop_rate=drop_rate,\n            drop_path_rate=drop_path_rate,\n            pretrained=pretrained\n        )\n        g = self.encoder(torch.rand(1, 3, 64, 64))\n        encoder_channels = [1] + [_.shape[1] for _ in g]\n        decoder_channels = [256, 128, 64, 32, 16]\n        if segtype == 'unet':\n            self.decoder = smp.unet.decoder.UnetDecoder(\n                encoder_channels=encoder_channels[:n_blocks+1],\n                decoder_channels=decoder_channels[:n_blocks],\n                n_blocks=n_blocks,\n            )\n\n        self.segmentation_head = nn.Conv2d(decoder_channels[n_blocks-1], out_dim, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))\n\n    def forward(self,x):\n        global_features = [0] + self.encoder(x)[:n_blocks]\n        seg_features = self.decoder(*global_features)\n        seg_features = self.segmentation_head(seg_features)\n        return seg_features","metadata":{"execution":{"iopub.status.busy":"2022-10-29T06:02:59.058663Z","iopub.execute_input":"2022-10-29T06:02:59.059903Z","iopub.status.idle":"2022-10-29T06:02:59.070787Z","shell.execute_reply.started":"2022-10-29T06:02:59.05986Z","shell.execute_reply":"2022-10-29T06:02:59.069705Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from timm.models.layers.conv2d_same import Conv2dSame\nfrom conv3d_same import Conv3dSame\n\n\ndef convert_3d(module):\n\n    module_output = module\n    if isinstance(module, torch.nn.BatchNorm2d):\n        module_output = torch.nn.BatchNorm3d(\n            module.num_features,\n            module.eps,\n            module.momentum,\n            module.affine,\n            module.track_running_stats,\n        )\n        if module.affine:\n            with torch.no_grad():\n                module_output.weight = module.weight\n                module_output.bias = module.bias\n        module_output.running_mean = module.running_mean\n        module_output.running_var = module.running_var\n        module_output.num_batches_tracked = module.num_batches_tracked\n        if hasattr(module, \"qconfig\"):\n            module_output.qconfig = module.qconfig\n            \n    elif isinstance(module, Conv2dSame):\n        module_output = Conv3dSame(\n            in_channels=module.in_channels,\n            out_channels=module.out_channels,\n            kernel_size=module.kernel_size[0],\n            stride=module.stride[0],\n            padding=module.padding[0],\n            dilation=module.dilation[0],\n            groups=module.groups,\n            bias=module.bias is not None,\n        )\n        module_output.weight = torch.nn.Parameter(module.weight.unsqueeze(-1).repeat(1,1,1,1,module.kernel_size[0]))\n\n    elif isinstance(module, torch.nn.Conv2d):\n        module_output = torch.nn.Conv3d(\n            in_channels=module.in_channels,\n            out_channels=module.out_channels,\n            kernel_size=module.kernel_size[0],\n            stride=module.stride[0],\n            padding=module.padding[0],\n            dilation=module.dilation[0],\n            groups=module.groups,\n            bias=module.bias is not None,\n            padding_mode=module.padding_mode\n        )\n        module_output.weight = torch.nn.Parameter(module.weight.unsqueeze(-1).repeat(1,1,1,1,module.kernel_size[0]))\n\n    elif isinstance(module, torch.nn.MaxPool2d):\n        module_output = torch.nn.MaxPool3d(\n            kernel_size=module.kernel_size,\n            stride=module.stride,\n            padding=module.padding,\n            dilation=module.dilation,\n            ceil_mode=module.ceil_mode,\n        )\n    elif isinstance(module, torch.nn.AvgPool2d):\n        module_output = torch.nn.AvgPool3d(\n            kernel_size=module.kernel_size,\n            stride=module.stride,\n            padding=module.padding,\n            ceil_mode=module.ceil_mode,\n        )\n\n    for name, child in module.named_children():\n        module_output.add_module(\n            name, convert_3d(child)\n        )\n    del module\n\n    return module_output\n\n\nm = TimmSegModel(backbone)\nm = convert_3d(m)\nm(torch.rand(1, 3, 128,128,128)).shape","metadata":{"execution":{"iopub.status.busy":"2022-10-29T06:02:59.072642Z","iopub.execute_input":"2022-10-29T06:02:59.073051Z","iopub.status.idle":"2022-10-29T06:03:13.720715Z","shell.execute_reply.started":"2022-10-29T06:02:59.073017Z","shell.execute_reply":"2022-10-29T06:03:13.719595Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Loss & Metric","metadata":{}},{"cell_type":"code","source":"from typing import Any, Dict, Optional\n\n\ndef binary_dice_score(\n    y_pred: torch.Tensor,\n    y_true: torch.Tensor,\n    threshold: Optional[float] = None,\n    nan_score_on_empty=False,\n    eps: float = 1e-7,\n) -> float:\n\n    if threshold is not None:\n        y_pred = (y_pred > threshold).to(y_true.dtype)\n\n    intersection = torch.sum(y_pred * y_true).item()\n    cardinality = (torch.sum(y_pred) + torch.sum(y_true)).item()\n\n    score = (2.0 * intersection) / (cardinality + eps)\n\n    has_targets = torch.sum(y_true) > 0\n    has_predicted = torch.sum(y_pred) > 0\n\n    if not has_targets:\n        if nan_score_on_empty:\n            score = np.nan\n        else:\n            score = float(not has_predicted)\n    return score\n\n\ndef multilabel_dice_score(\n    y_true: torch.Tensor,\n    y_pred: torch.Tensor,\n    threshold=None,\n    eps=1e-7,\n    nan_score_on_empty=False,\n):\n    ious = []\n    num_classes = y_pred.size(0)\n    for class_index in range(num_classes):\n        iou = binary_dice_score(\n            y_pred=y_pred[class_index],\n            y_true=y_true[class_index],\n            threshold=threshold,\n            nan_score_on_empty=nan_score_on_empty,\n            eps=eps,\n        )\n        ious.append(iou)\n\n    return ious\n\n\ndef dice_loss(input, target):\n    input = torch.sigmoid(input)\n    smooth = 1.0\n    iflat = input.view(-1)\n    tflat = target.view(-1)\n    intersection = (iflat * tflat).sum()\n    return 1 - ((2.0 * intersection + smooth) / (iflat.sum() + tflat.sum() + smooth))\n\n\ndef bce_dice(input, target, loss_weights=loss_weights):\n    loss1 = loss_weights[0] * nn.BCEWithLogitsLoss()(input, target)\n    loss2 = loss_weights[1] * dice_loss(input, target)\n    return (loss1 + loss2) / sum(loss_weights)\n\ncriterion = bce_dice","metadata":{"execution":{"iopub.status.busy":"2022-10-29T06:03:13.72458Z","iopub.execute_input":"2022-10-29T06:03:13.724951Z","iopub.status.idle":"2022-10-29T06:03:13.73657Z","shell.execute_reply.started":"2022-10-29T06:03:13.72492Z","shell.execute_reply":"2022-10-29T06:03:13.7354Z"},"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, gt_masks in bar:\n        optimizer.zero_grad()\n        images = images.cuda()\n        gt_masks = gt_masks.cuda()\n\n        do_mixup = False\n        if random.random() < p_mixup:\n            do_mixup = True\n            images, gt_masks, gt_masks_sfl, lam = mixup(images, gt_masks)\n\n        with amp.autocast():\n            logits = model(images)\n            loss = criterion(logits, gt_masks)\n            if do_mixup:\n                loss2 = criterion(logits, gt_masks_sfl)\n                loss = loss * lam  + loss2 * (1 - lam)\n\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    outputs = []\n    ths = [0.2, 0.3, 0.4, 0.5, 0.6, 0.7, 0.8]\n    batch_metrics = [[]] * 7\n    bar = tqdm(loader_valid)\n    with torch.no_grad():\n        for images, gt_masks in bar:\n            images = images.cuda()\n            gt_masks = gt_masks.cuda()\n\n            logits = model(images)\n            loss = criterion(logits, gt_masks)\n            valid_loss.append(loss.item())\n            for thi, th in enumerate(ths):\n                pred = (logits.sigmoid() > th).float().detach()\n                for i in range(logits.shape[0]):\n                    tmp = multilabel_dice_score(\n                        y_pred=logits[i].sigmoid().cpu(),\n                        y_true=gt_masks[i].cpu(),\n                        threshold=0.5,\n                    )\n                    batch_metrics[thi].extend(tmp)\n            bar.set_description(f'smth:{np.mean(valid_loss[-30:]):.4f}')\n            \n    metrics = [np.mean(this_metric) for this_metric in batch_metrics]\n    print('best th:', ths[np.argmax(metrics)], 'best dc:', np.max(metrics))\n\n    return np.mean(valid_loss), np.max(metrics)\n","metadata":{"execution":{"iopub.status.busy":"2022-10-29T06:03:13.738511Z","iopub.execute_input":"2022-10-29T06:03:13.738907Z","iopub.status.idle":"2022-10-29T06:03:13.755925Z","shell.execute_reply.started":"2022-10-29T06:03:13.738871Z","shell.execute_reply":"2022-10-29T06:03:13.755044Z"},"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, 1000)\nlrs = []\nfor epoch in range(1, 1000+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-10-29T06:03:13.757188Z","iopub.execute_input":"2022-10-29T06:03:13.757646Z","iopub.status.idle":"2022-10-29T06:03:13.977069Z","shell.execute_reply.started":"2022-10-29T06:03:13.75761Z","shell.execute_reply":"2022-10-29T06:03:13.976109Z"},"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_seg[df_seg['fold'] != fold].reset_index(drop=True)\n    valid_ = df_seg[df_seg['fold'] == fold].reset_index(drop=True)\n    dataset_train = SEGDataset(train_, 'train', transform=transforms_train)\n    dataset_valid = SEGDataset(valid_, 'valid', transform=transforms_valid)\n    loader_train = torch.utils.data.DataLoader(dataset_train, batch_size=batch_size, shuffle=True, num_workers=num_workers)\n    loader_valid = torch.utils.data.DataLoader(dataset_valid, batch_size=batch_size, shuffle=False, num_workers=num_workers)\n\n    model = TimmSegModel(backbone, pretrained=True)\n    model = convert_3d(model)\n    model = model.to(device)\n\n    optimizer = optim.AdamW(model.parameters(), lr=init_lr)\n    scaler = torch.cuda.amp.GradScaler()\n    from_epoch = 0\n    metric_best = 0.\n    loss_min = np.inf\n\n    scheduler_cosine = torch.optim.lr_scheduler.CosineAnnealingWarmRestarts(optimizer, n_epochs)\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, metric = valid_func(model, loader_valid)\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            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-10-29T06:03:13.978679Z","iopub.execute_input":"2022-10-29T06:03:13.979257Z","iopub.status.idle":"2022-10-29T06:03:13.992437Z","shell.execute_reply.started":"2022-10-29T06:03:13.979202Z","shell.execute_reply":"2022-10-29T06:03:13.991378Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"run(0)\nrun(1)\nrun(2)\nrun(3)\nrun(4)","metadata":{"execution":{"iopub.status.busy":"2022-10-29T06:03:33.005515Z","iopub.execute_input":"2022-10-29T06:03:33.006143Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}