{"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":"# About this notebook\n- **Version 1 :**\n    + PyTorch CoAT starter code\n    + 5 folds\n    + OneCycleLR scheduler\n    \n\n- **Version 3 :**\n    + PyTorch CoAT Model starter code :\n    + fixing bugs\n    + 5 folds\n    + OneCycleLR scheduler\n\nIf this notebook is helpful, feel free to upvote :)","metadata":{}},{"cell_type":"code","source":"!pip install -qq torch==1.7.1+cu110 torchvision==0.8.2+cu110 torchaudio==0.7.2 -f https://download.pytorch.org/whl/torch_stable.html\n!pip install -qq git+https://github.com/qubvel/segmentation_models.pytorch\n!pip install -qq timm==0.4.12\n!pip install -qq einops","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_kg_hide-output":true,"execution":{"iopub.status.busy":"2022-09-08T02:22:02.806672Z","iopub.execute_input":"2022-09-08T02:22:02.807350Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install segmentation-models-pytorch","metadata":{"_kg_hide-output":true,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import sys\nsys.path.append(\"../input/timm-pytorch-image-models/pytorch-image-models-master\")\nsys.path.append(\"../input/pretrained-models-pytorch\")\nsys.path.append(\"../input/efficientnet-pytorch\")\nsys.path.append(\"../input/segmentation-models-pytorch\")\nimport segmentation_models_pytorch as smp\n\n\nprint(f\"Segmentation Models version: {smp.__version__}\")","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Data Loading","metadata":{}},{"cell_type":"code","source":"import os\n\nimport pandas as pd\n\nfrom matplotlib import pyplot as plt","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"os.listdir('../input/hubmap-organ-segmentation')","metadata":{"_kg_hide-output":true,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train = pd.read_csv('../input/hubmap-organ-segmentation/train.csv')\ntest = pd.read_csv('../input/hubmap-organ-segmentation/test.csv')\nTRAIN_PATH = '../input/mmsegmentation256x256/train/'\nTEST_PATH = '../input/hubmap-organ-segmentation/test_images/'\nMASK_PATH = '../input/mmsegmentation256x256/masks/'\ndisplay(train.head())\ndisplay(test.head())","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Library","metadata":{}},{"cell_type":"code","source":"import os\nimport gc\nimport sys\nimport glob\nimport random\nimport time\nimport math\n\nimport torch\nimport torch.nn as nn\nimport albumentations as A\nimport skimage\nfrom contextlib import contextmanager\n\nimport cv2\nimport pandas as pd\nimport numpy as np\nimport matplotlib.pyplot as plt\n\nimport tqdm\nimport segmentation_models_pytorch as smp\nfrom sklearn.model_selection import StratifiedKFold\nfrom segmentation_models_pytorch.encoders import encoders\n\nimport tifffile as tiff\n\nsys.path.append('../input/hubmap-coat/')\n\nfrom coat import *\nfrom daformer import *\nfrom helper import *\n\ntorch.backends.cudnn.benchmark = True","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Directory settings","metadata":{}},{"cell_type":"code","source":"# ====================================================\n# Directory settings\n# ====================================================\nimport os\n\nOUTPUT_DIR = './'\nif not os.path.exists(OUTPUT_DIR):\n    os.makedirs(OUTPUT_DIR)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Config","metadata":{}},{"cell_type":"code","source":"# ====================================================\n# CFG\n# ====================================================\nclass CFG:\n    train=True\n    debug=False\n    ENCODER = 'resnet50'\n    DECODER = 'UnetPlusPlus'\n    fold = 0\n    nfolds = 5\n    imsize = 384\n    BATCH_SIZE = 8\n    print_freq=100\n    DEVICE = ('cuda' if torch.cuda.is_available() else 'cpu')\n    EPOCHS = 11\n    NUM_WORKERS = 4\n    SEED = 24\n    trn_folds=[0, 1, 2, 3, 4]","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Utils","metadata":{}},{"cell_type":"code","source":"# ====================================================\n# Utils\n# ====================================================\nclass DiceCoef(nn.Module):\n    def __init__(self, weight=None, size_average=True):\n        super().__init__()\n\n    def forward(self, y_pred, y_true, smooth=1.):\n        y_true = y_true.view(-1)\n        y_pred = y_pred.view(-1)\n        \n        #Round off y_pred\n        y_pred = torch.round((y_pred - y_pred.min()) / (y_pred.max() - y_pred.min()))\n        \n        intersection = (y_true * y_pred).sum()\n        dice = (2.0*intersection + smooth)/(y_true.sum() + y_pred.sum() + smooth)\n        \n        return dice\n\n@contextmanager\ndef timer(name):\n    t0 = time.time()\n    LOGGER.info(f'[{name}] start')\n    yield\n    LOGGER.info(f'[{name}] done in {time.time() - t0:.0f} s.')\n\n\ndef init_logger(log_file=OUTPUT_DIR+'train.log'):\n    from logging import getLogger, INFO, FileHandler,  Formatter,  StreamHandler\n    logger = getLogger(__name__)\n    logger.setLevel(INFO)\n    handler1 = StreamHandler()\n    handler1.setFormatter(Formatter(\"%(message)s\"))\n    handler2 = FileHandler(filename=log_file)\n    handler2.setFormatter(Formatter(\"%(message)s\"))\n    logger.addHandler(handler1)\n    logger.addHandler(handler2)\n    return logger\n\nLOGGER = init_logger()\n\n\ndef seed_torch(seed=CFG.SEED):\n    random.seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\n\nseed_torch(seed=CFG.SEED)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# CV split","metadata":{}},{"cell_type":"code","source":"folds = train.copy()\nkf = StratifiedKFold(n_splits=CFG.nfolds,random_state=CFG.SEED,shuffle=True)\nfor fold, (_, val_idx) in enumerate(kf.split(folds, y=folds[\"organ\"])):\n    folds.loc[val_idx, \"fold\"] = fold","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Dataset","metadata":{}},{"cell_type":"code","source":"# ====================================================\n# Dataset\n# ====================================================\nclass HuBMAPDataset(torch.utils.data.Dataset):\n    def __init__(self, df, tfms=None):\n        self.df = df\n        ids = self.df.id.values\n        self.fnames = [fname for fname in os.listdir(TRAIN_PATH) if int(fname.split('_')[0]) in ids]\n        self.image_size = CFG.imsize\n        self.tfms = tfms\n        \n    def img2tensor(self, img,dtype:np.dtype=np.float32):\n        if img.ndim==2 : img = np.expand_dims(img,2)\n        img = np.transpose(img,(2,0,1)) # C , H , W\n        return torch.from_numpy(img.astype(dtype, copy=False))\n    \n    def __len__(self):\n        return len(self.fnames)\n    \n    def resize(self, img, interp):\n        return  cv2.resize(\n            img, (self.image_size, self.image_size), interpolation=interp)\n    \n    def __getitem__(self, idx):\n        fname = self.fnames[idx]\n        img = cv2.cvtColor(cv2.imread(TRAIN_PATH + fname), cv2.COLOR_BGR2RGB)\n        mask = cv2.imread((MASK_PATH + fname),cv2.IMREAD_GRAYSCALE)\n        if self.tfms is not None:\n            augmented = self.tfms(image=img,mask=mask)\n            img,mask = augmented['image'],augmented['mask']\n        \n        return self.img2tensor(self.resize(img , cv2.INTER_NEAREST)) , self.img2tensor(self.resize(mask , cv2.INTER_NEAREST))","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# ====================================================\n# Dataset\n# ====================================================\nclass HuBMAPTestDataset(torch.utils.data.Dataset):\n    def __init__(self, df, tfms=None):\n        self.df = df\n        ids = self.df.id.values\n        self.fnames = [fname for fname in os.listdir(TEST_PATH) ]\n        self.image_size = CFG.imsize\n        self.tfms = tfms\n        \n    def img2tensor(self, img,dtype:np.dtype=np.float32):\n        if img.ndim==2 : img = np.expand_dims(img,2)\n        img = np.transpose(img,(2,0,1)) # C , H , W\n        return torch.from_numpy(img.astype(dtype, copy=False))\n    \n    def __len__(self):\n        return len(self.fnames)\n    \n    def resize(self, img, interp):\n        return  cv2.resize(\n            img, (self.image_size, self.image_size), interpolation=interp)\n    \n    def __getitem__(self, idx):\n        fname = self.fnames[idx]\n        img = cv2.cvtColor(cv2.imread(TEST_PATH + fname), cv2.COLOR_BGR2RGB)\n        if self.tfms is not None:\n            augmented = self.tfms(image=img)\n            img = augmented['image']\n        \n        return self.img2tensor(self.resize(img , cv2.INTER_NEAREST)) ","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Transforms","metadata":{}},{"cell_type":"code","source":"# ====================================================\n# Transforms\n# ====================================================\ndef transformer(p=1.0):\n    return A.Compose([\n        A.HorizontalFlip(),\n        A.VerticalFlip(),\n        A.RandomRotate90(),\n        A.ShiftScaleRotate(shift_limit=0.0625, scale_limit=0.2, rotate_limit=15, p=0.9, \n                         border_mode=cv2.BORDER_REFLECT),\n        A.OneOf([\n            A.OpticalDistortion(p=0.3),\n            A.GridDistortion(p=.1),\n            A.PiecewiseAffine(p=0.3),\n        ], p=0.3),\n        A.OneOf([\n            A.HueSaturationValue(10,15,10),\n            A.CLAHE(clip_limit=2),\n            A.RandomBrightnessContrast(),            \n        ], p=0.3),\n    ], p=p)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Viewing Data","metadata":{}},{"cell_type":"code","source":"ds = HuBMAPDataset(df=folds,tfms=transformer())\ndl = torch.utils.data.DataLoader(ds,batch_size=64,shuffle=False,num_workers=CFG.NUM_WORKERS)\nit = iter(dl)\nimgs,masks = next(it)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.figure(figsize=(16,16))\nfor i,(img,mask) in enumerate(zip(imgs,masks)):\n    img = ((img.permute(1,2,0))).numpy().astype(np.uint8)  # H , W , C\n    plt.subplot(8,8,i+1)\n    plt.imshow(img,vmin=0,vmax=255)\n    plt.imshow(mask.squeeze().numpy(), alpha=0.2)\n    plt.axis('off')\n    plt.subplots_adjust(wspace=None, hspace=None)\n    \ndel ds,dl,imgs,masks","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# MODEL","metadata":{}},{"cell_type":"code","source":"class Net(nn.Module):\n    \n    def __init__(self,\n                 encoder=coat_lite_medium,\n                 decoder=daformer_conv3x3,\n                 encoder_cfg={},\n                 decoder_cfg={},\n                 ):\n        \n        super(Net, self).__init__()\n        decoder_dim = decoder_cfg.get('decoder_dim', 320)\n        self.decoder_dim = decoder_dim\n\n        self.encoder = encoder\n        \n        self.rgb = RGB()\n        \n        encoder_dim = self.encoder.embed_dims\n        # [64, 128, 320, 512]\n\n        self.decoder = decoder(\n            encoder_dim=encoder_dim,\n            decoder_dim=decoder_dim,\n        )\n        self.logit = nn.Sequential(\n            nn.Conv2d(decoder_dim, 1, kernel_size=1),\n            nn.Upsample(scale_factor = 4, mode='bilinear', align_corners=False),\n        )\n\n    def forward(self, batch):\n        x = self.rgb(batch)\n        B, C, H, W = x.shape\n        encoder = self.encoder(x) # [1, 512, 12, 12]\n        last, decoder = self.decoder(encoder) # [1, 320, 96, 96]\n        logit = self.logit(last)\n\n        output = {}\n        probability_from_logit = torch.sigmoid(logit)\n\n        output['probability'] = probability_from_logit\n        return output","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def init_model():\n    encoder = coat_lite_medium()\n    checkpoint = '../input/hubmap-coat-medium/coat_lite_medium_384x384_f9129688.pth'\n    checkpoint = torch.load(checkpoint, map_location=lambda storage, loc: storage)\n    state_dict = checkpoint['model']\n    encoder.load_state_dict(state_dict,strict=False)\n    \n    net = Net(encoder=encoder).cuda()\n\n    return net","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\"\"\"sample = ds[0]\nimg = sample[0]\nmask = sample[1]\nimg = img.unsqueeze(0)\"\"\"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"DECODERS = [\n    \"Unet\",\n    \"Linknet\",\n    \"FPN\",\n    \"PSPNet\",\n    \"DeepLabV3\",\n    \"DeepLabV3Plus\",\n    \"PAN\",\n    \"UnetPlusPlus\",\n]\nENCODERS = list(encoders.keys())\n\n\ndef define_model(\n    decoder_name,\n    encoder_name,\n    num_classes=1,\n    activation=None,\n    encoder_weights=\"imagenet\",\n):\n    \"\"\"\n    Loads a segmentation architecture.\n    Args:\n        decoder_name (str): Decoder name.\n        encoder_name (str): Encoder name.\n        num_classes (int, optional): Number of classes. Defaults to 1.\n        pretrained : pretrained original weights\n        activation (str or None, optional): Activation of the last layer. Defaults to None.\n        encoder_weights (str, optional): Pretrained weights. Defaults to \"imagenet\".\n    Returns:\n        torch model: Segmentation model.\n    \"\"\"\n    assert decoder_name in DECODERS, \"Decoder name not supported\"\n    assert encoder_name in ENCODERS, \"Encoder name not supported\"\n\n    decoder = getattr(smp, decoder_name)\n\n    model = decoder(\n        encoder_name,\n        encoder_weights=encoder_weights,\n        classes=num_classes,\n        activation=activation,\n    )\n    model.num_classes = num_classes\n\n    return model","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#model = define_model(encoder_name=CFG.ENCODER,decoder_name=CFG.DECODER)\n#res = model(img)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Loss","metadata":{}},{"cell_type":"code","source":"# ====================================================\n# Loss function\n# ====================================================\nclass CustomLoss(nn.Module):\n    def __init__(self):\n        super(CustomLoss,self).__init__()\n        self.diceloss = smp.losses.DiceLoss(mode='binary')\n        self.binloss = smp.losses.SoftBCEWithLogitsLoss(reduction = 'mean' , smooth_factor = 0.1)\n        \n    def forward(self, outputs, mask):\n        dice = self.diceloss(outputs,mask)\n        bce = self.binloss(outputs , mask)\n        loss = dice * 0.7 + bce * 0.3\n        return loss","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Helper functions","metadata":{}},{"cell_type":"code","source":"# ====================================================\n# Helper functions\n# ====================================================\nclass AverageMeter(object):\n    \"\"\"Computes and stores the average and current value\"\"\"\n    def __init__(self):\n        self.reset()\n\n    def reset(self):\n        self.val = 0\n        self.avg = 0\n        self.sum = 0\n        self.count = 0\n\n    def update(self, val, n=1):\n        self.val = val\n        self.sum += val * n\n        self.count += n\n        self.avg = self.sum / self.count\n\n\ndef asMinutes(s):\n    m = math.floor(s / 60)\n    s -= m * 60\n    return '%dm %ds' % (m, s)\n\n\ndef timeSince(since, percent):\n    now = time.time()\n    s = now - since\n    es = s / (percent)\n    rs = es - s\n    return '%s (remain %s)' % (asMinutes(s), asMinutes(rs))\n\ndef train_fn(train_loader, model, criterion, metric , optimizer, scheduler,epoch, DEVICE):\n    \n    batch_time = AverageMeter()\n    data_time = AverageMeter()\n    losses = AverageMeter()\n    scores = AverageMeter()\n    train_loss=0.0\n    score=0.0\n    # switch to train mode\n    model.train()\n    start = end = time.time()\n    for step, data in enumerate(train_loader):\n        # measure data loading time\n        data_time.update(time.time() - end)\n        \n        img, mask = data\n        img = img.to(DEVICE)\n        mask = mask.to(DEVICE)\n        batch_size = img.size(0)\n    \n        outputs = model(img)['probability']\n        loss = criterion(outputs, mask)\n        loss.backward()\n        optimizer.step()\n        scheduler.step()\n        optimizer.zero_grad()\n        #loss\n        #loss = loss.detach().item()\n        train_loss += loss.item()\n        score += metric(outputs,mask).item()\n        # record loss\n        #losses.update(loss.item(), batch_size)\n        #scores.update(metric(outputs,mask).item(), batch_size)\n        # measure elapsed time\n        batch_time.update(time.time() - end)\n        end = time.time()\n            \n        #if step % CFG.print_freq == 0 or step == (len(train_loader)-1):\n        #    print('Epoch: [{0}][{1}/{2}] '\n        #          'Elapsed {remain:s} '\n        #          'Loss: {loss.val:.4f}({loss.avg:.4f}) '\n        #          'Score: {score.val:.4f}'\n        #          #'LR: {lr:.6f}  '\n        #          .format(epoch+1, step, len(train_loader), \n        #                  remain=timeSince(start, float(step+1)/len(train_loader)),\n        #                  loss=losses,\n        #                  score=scores\n        #                  #lr=scheduler.get_lr()[0]\n        #                 ))\n    TRAIN_LOSS = train_loss / len(train_loader)\n    SCORE = score / len(train_loader)\n    return TRAIN_LOSS , SCORE\n    \ndef valid_fn(valid_loader, model, criterion,metric,epoch ,DEVICE):\n    \n    batch_time = AverageMeter()\n    data_time = AverageMeter()\n    losses = AverageMeter()\n    scores = AverageMeter()\n    preds=[]\n    valid_loss=0.0\n    val_score=0.0\n    # switch to evaluation mode\n    model.eval()\n    start = end = time.time()\n    \n    for step, data in enumerate(valid_loader):\n        # measure data loading time\n        data_time.update(time.time() - end)\n            \n        img, mask = data\n        img = img.to(DEVICE)\n        mask = mask.to(DEVICE)\n        batch_size = img.size(0)\n        # compute loss\n        with torch.no_grad():\n            outputs = model(img)['probability']\n        preds.append(outputs)\n        loss = criterion(outputs, mask)\n        \n        valid_loss += loss.item()\n        val_score += metric(outputs,mask).item()\n        #losses.update(loss.item(), batch_size)\n        #scores.update(metric(outputs,mask).item(), batch_size)\n        # measure elapsed time\n        batch_time.update(time.time() - end)\n        end = time.time()\n        #if step % CFG.print_freq == 0 or step == (len(valid_loader)-1):\n        #    print('EVAL: [{0}/{1}] '\n        #            'Elapsed {remain:s} '\n        #            'Loss: {loss.val:.4f}({loss.avg:.4f}) '\n        #            'Score: {score.val:.4f}'\n        #            .format(step, len(valid_loader),\n        #                    loss=losses,\n        #                    score=scores,\n        #                    remain=timeSince(start, float(step+1)/len(valid_loader))))\n    VALID_LOSS = valid_loss / len(valid_loader)\n    VALID_SCORE = val_score / len(valid_loader)\n    return VALID_LOSS, VALID_SCORE","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Train loop","metadata":{}},{"cell_type":"code","source":"# ====================================================\n# Train loop\n# ====================================================\ndef train_loop(folds, fold):\n    LOGGER.info(f\"========== fold: {fold} training ==========\")\n    \n    # ====================================================\n    # loader\n    # ====================================================\n    if CFG.debug:\n        train_folds = folds[folds['fold'] != fold].sample(50)\n        valid_folds = folds[folds['fold'] == fold].sample(50)\n        \n    else:\n        train_folds = folds[folds['fold'] != fold]\n        valid_folds = folds[folds['fold'] == fold]\n        \n    best_loss = 999\n    best_score = 0\n    \n    ds_train = HuBMAPDataset(df=train_folds, tfms=transformer())\n    ds_val = HuBMAPDataset(df=valid_folds)\n    \n    dataloader_train = torch.utils.data.DataLoader(ds_train,batch_size=CFG.BATCH_SIZE, shuffle=True,num_workers=CFG.NUM_WORKERS)\n    dataloader_val = torch.utils.data.DataLoader(ds_val,batch_size=CFG.BATCH_SIZE, shuffle=False,num_workers=CFG.NUM_WORKERS)\n    \n    #model = define_model(encoder_name=CFG.ENCODER,decoder_name=CFG.DECODER).to(CFG.DEVICE)\n    model = init_model().to(CFG.DEVICE)\n    optimizer = torch.optim.Adam([\n        {'params': model.decoder.parameters(), 'lr': 5e-5}, \n        {'params': model.encoder.parameters(), 'lr': 8e-5},  \n    ])\n    \n    scheduler = torch.optim.lr_scheduler.OneCycleLR(optimizer=optimizer, pct_start=0.1, div_factor=1e3, \n                                              max_lr=1e-3, epochs=CFG.EPOCHS, steps_per_epoch=len(dataloader_train))\n    \n    loss_func = CustomLoss()\n    dice_coe = DiceCoef()\n    \n    for epoch in range(CFG.EPOCHS):\n        start_time = time.time()\n        \n        # train\n        train_loss, train_score = train_fn(dataloader_train, model, loss_func, dice_coe, optimizer, scheduler, epoch, CFG.DEVICE)\n        \n        # eval\n        valid_loss, val_score = valid_fn(dataloader_val, model, loss_func, dice_coe, epoch, CFG.DEVICE)\n        \n        \n        \n        \n        # scoring\n        elapsed = time.time() - start_time\n\n        LOGGER.info(f'Epoch {epoch+1} - train_loss: {train_loss:.4f}  train_score: {train_score:.4f}  time: {elapsed/60:.2f}m')\n        LOGGER.info(f'Epoch {epoch+1} - valid_loss: {valid_loss:.4f}  val_score: {val_score:.4f}')\n        \n        if val_score > best_score:\n            best_score = val_score\n            torch.save(model.state_dict(), f\"{OUTPUT_DIR}FOLD{fold}_best_score.pth\")\n            print(f\"Saved model for best score : FOLD{fold}_best_score.pth\")\n            LOGGER.info(f\"Saved model for best score : FOLD{fold}_best_score.pth\")\n        \n        if valid_loss < best_loss:\n            best_loss = valid_loss\n            torch.save(model.state_dict(), f\"{OUTPUT_DIR}FOLD{fold}_best_loss.pth\")\n            print(f\"Saved model for best loss : FOLD{fold}_best_loss.pth\")   \n            LOGGER.info(f\"Saved model for best loss : FOLD{fold}_best_loss.pth\")","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Main","metadata":{}},{"cell_type":"code","source":"# ====================================================\n# main\n# ====================================================\ndef main():\n\n    \n    if CFG.train:\n        # train \n        oof_df = pd.DataFrame()\n        for fold in range(CFG.nfolds):\n            if fold in CFG.trn_folds:\n                _oof_df = train_loop(folds, fold)\n                oof_df = pd.concat([oof_df, _oof_df])\n                LOGGER.info(f\"========== fold: {fold} result ==========\")","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import gc\ngc.collect()\ntorch.cuda.empty_cache()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if __name__ == \"__main__\":\n    main()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}