{"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":"import numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nimport torchvision.models as models\nfrom torchvision import transforms\nimport cv2\nimport matplotlib.pyplot as plt\nimport os\nfrom tqdm.notebook import tqdm\nimport warnings\nwarnings.filterwarnings('ignore')\n# from sklearn.cross_validation import cross_val_score","metadata":{"execution":{"iopub.status.busy":"2021-10-31T13:31:47.480257Z","iopub.execute_input":"2021-10-31T13:31:47.480675Z","iopub.status.idle":"2021-10-31T13:31:49.588066Z","shell.execute_reply.started":"2021-10-31T13:31:47.480584Z","shell.execute_reply":"2021-10-31T13:31:49.586934Z"},"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/ranger/ranger')\nimport timm\nfrom ranger import Ranger\nimport albumentations as A\nfrom albumentations.pytorch.transforms import ToTensorV2","metadata":{"execution":{"iopub.status.busy":"2021-10-31T13:31:49.590564Z","iopub.execute_input":"2021-10-31T13:31:49.591054Z","iopub.status.idle":"2021-10-31T13:31:58.238339Z","shell.execute_reply.started":"2021-10-31T13:31:49.590943Z","shell.execute_reply":"2021-10-31T13:31:58.237057Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\n# !git clone https://github.com/shengliu66/ELR ../input/ELR_Loss\n    \n# sys.path.append('./ELR_Loss')\n","metadata":{"execution":{"iopub.status.busy":"2021-10-31T13:31:58.241027Z","iopub.execute_input":"2021-10-31T13:31:58.241441Z","iopub.status.idle":"2021-10-31T13:31:58.249400Z","shell.execute_reply.started":"2021-10-31T13:31:58.241382Z","shell.execute_reply":"2021-10-31T13:31:58.245444Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# !pip install -r requirements.txt\n# from ELR.model.loss import elr_loss","metadata":{"execution":{"iopub.status.busy":"2021-10-31T13:31:58.251349Z","iopub.execute_input":"2021-10-31T13:31:58.251779Z","iopub.status.idle":"2021-10-31T13:31:58.262461Z","shell.execute_reply.started":"2021-10-31T13:31:58.251731Z","shell.execute_reply":"2021-10-31T13:31:58.261204Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"cfg = {}\ncfg['epoch'] = 20\ncfg['batch_size'] = 16\ncfg['lr'] = 0.016\ncfg['image_size'] = 512 # (600, 800)\ncfg['device'] = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\ncfg['fold'] = 1\ncfg['split'] = 0.8\ncfg['smoothing_factor'] = 0.01\nTRAINDATAPATH = '../input/cassava-leaf-disease-classification/train_images'\nTESTDATAPATH = '../input/cassava-leaf-disease-classification/test_images'\nTRAINLABELPATH = '../input/cassava-leaf-disease-classification/train.csv'\nTESTSAMPLEPATH = '../input/cassava-leaf-disease-classification/sample_submission.csv'\nCLASSESJSONPATH = '../input/cassava-leaf-disease-classification/label_num_to_disease_map.json'\n\nclasses = {}\nnum_class = 0","metadata":{"execution":{"iopub.status.busy":"2021-10-31T13:31:58.267078Z","iopub.execute_input":"2021-10-31T13:31:58.267592Z","iopub.status.idle":"2021-10-31T13:31:58.330685Z","shell.execute_reply.started":"2021-10-31T13:31:58.267531Z","shell.execute_reply":"2021-10-31T13:31:58.329292Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"seed=31\n\ndef setSeed(seed=31):\n#     random.seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    np.random.seed(seed)\n    pd.core.common.random_state(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\n\nsetSeed(31)","metadata":{"execution":{"iopub.status.busy":"2021-10-31T13:31:58.332670Z","iopub.execute_input":"2021-10-31T13:31:58.333164Z","iopub.status.idle":"2021-10-31T13:31:58.347700Z","shell.execute_reply.started":"2021-10-31T13:31:58.333119Z","shell.execute_reply":"2021-10-31T13:31:58.346602Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## get the class","metadata":{}},{"cell_type":"code","source":"import json\n\nwith open(CLASSESJSONPATH) as f:\n    global classses, num_class\n    classes = json.load(f)\n    num_class = len(classes)","metadata":{"execution":{"iopub.status.busy":"2021-10-31T13:31:58.349207Z","iopub.execute_input":"2021-10-31T13:31:58.350588Z","iopub.status.idle":"2021-10-31T13:31:58.363992Z","shell.execute_reply.started":"2021-10-31T13:31:58.350541Z","shell.execute_reply":"2021-10-31T13:31:58.362763Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## utils","metadata":{}},{"cell_type":"code","source":"# read dataset\ntrain_df = pd.read_csv(TRAINLABELPATH)\n\n\ndef readImage(ID, path):\n    ID = ID.split('.')[0]\n    filepath = os.path.join(path, ID + '.jpg')\n    img = cv2.imread(filepath)\n    img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n    return img\n","metadata":{"execution":{"iopub.status.busy":"2021-10-31T13:31:58.365566Z","iopub.execute_input":"2021-10-31T13:31:58.368034Z","iopub.status.idle":"2021-10-31T13:31:58.405596Z","shell.execute_reply.started":"2021-10-31T13:31:58.367986Z","shell.execute_reply":"2021-10-31T13:31:58.404665Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# img = readImage('../input/cassava-leaf-disease-classification/train_images', '1000015157.jpg')\n","metadata":{"execution":{"iopub.status.busy":"2021-10-31T13:31:58.406741Z","iopub.execute_input":"2021-10-31T13:31:58.408032Z","iopub.status.idle":"2021-10-31T13:31:58.413229Z","shell.execute_reply.started":"2021-10-31T13:31:58.407985Z","shell.execute_reply":"2021-10-31T13:31:58.412171Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## show data/ pictures","metadata":{}},{"cell_type":"code","source":"# show data\nfrom collections import Counter\ndef showData(df):\n    display(df.info())\n    \n    c = Counter(df.label)\n    display(c)\n    \n    plt.clf()\n    plt.pie(c.values(), labels=c.keys(),autopct='%1.2f%%')\n    plt.show()\n    \n# showData(train_df)\ndef showExample(df):\n    num_show = 4\n    plt.clf()\n    f, axis = plt.subplots(num_class, num_show, figsize=(6 * num_show, 25))\n    for cls, name in classes.items():\n        idx = int(cls)\n#         print(cls)\n        sampleimgs = df[df['label'] == idx]['image_id']\n        rndidxs = np.random.randint(len(sampleimgs), size=num_show)\n        for i, randnum in enumerate(rndidxs):\n            axis[idx,i].set_title(name)\n            axis[idx,i].imshow(readImage(sampleimgs.iloc[i], TRAINDATAPATH))\n        \n    plt.show()\n","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2021-10-31T13:31:58.415366Z","iopub.execute_input":"2021-10-31T13:31:58.416161Z","iopub.status.idle":"2021-10-31T13:31:58.429225Z","shell.execute_reply.started":"2021-10-31T13:31:58.416107Z","shell.execute_reply":"2021-10-31T13:31:58.428230Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# showExample(train_df)","metadata":{"execution":{"iopub.status.busy":"2021-10-31T13:31:58.430362Z","iopub.execute_input":"2021-10-31T13:31:58.432436Z","iopub.status.idle":"2021-10-31T13:31:58.440298Z","shell.execute_reply.started":"2021-10-31T13:31:58.432389Z","shell.execute_reply":"2021-10-31T13:31:58.439271Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## dataset, dataloader","metadata":{}},{"cell_type":"code","source":"# dataset and data loader\nclass CLCdataset(Dataset):\n    def __init__(self, df,isTrain = True):\n        self.df = df\n        self.isTrain = isTrain\n        \n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, idx):\n        ID = self.df['image_id'].iloc[idx].split('.')[0]\n        img = readImage(ID, TRAINDATAPATH if self.isTrain else TESTDATAPATH)\n        if self.isTrain:\n            img = train_transform(image = img)['image']\n        else:\n            img = test_transform(image = img)['image']\n        if self.isTrain:\n            label = self.df['label'].iloc[idx]\n            return ID, img, label\n        return ID, img","metadata":{"execution":{"iopub.status.busy":"2021-10-31T13:31:58.442342Z","iopub.execute_input":"2021-10-31T13:31:58.442709Z","iopub.status.idle":"2021-10-31T13:31:58.453123Z","shell.execute_reply.started":"2021-10-31T13:31:58.442664Z","shell.execute_reply":"2021-10-31T13:31:58.452055Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# transform\n\ntrain_transform = A.Compose([\n    A.Flip(p=0.5),\n    A.Resize(cfg['image_size'], cfg['image_size'],p=1),\n    A.ShiftScaleRotate(p=0.5),\n    A.HueSaturationValue(hue_shift_limit=0.2, sat_shift_limit=0.2, val_shift_limit=0.2, p=0.5),\n    A.RandomBrightnessContrast(brightness_limit=(-0.1,0.1), contrast_limit=(-0.1, 0.1), p=0.5),\n    A.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225], max_pixel_value=255.0, p=1.0),\n    ToTensorV2() # (H,W,C) to (C,H,W)\n])\n\ntest_transform = A.Compose([\n    A.Resize(cfg['image_size'], cfg['image_size'],p=1),\n    A.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225], max_pixel_value=255.0, p=1.0),\n    ToTensorV2() # (H,W,C) to (C,H,W)\n])\n","metadata":{"execution":{"iopub.status.busy":"2021-10-31T13:31:58.454769Z","iopub.execute_input":"2021-10-31T13:31:58.455335Z","iopub.status.idle":"2021-10-31T13:31:58.468043Z","shell.execute_reply.started":"2021-10-31T13:31:58.455290Z","shell.execute_reply":"2021-10-31T13:31:58.466923Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## model","metadata":{}},{"cell_type":"code","source":"class CFCmodel(nn.Module):\n    def __init__(self, out_dim):\n        super(CFCmodel, self).__init__()\n        self.model = timm.create_model('efficientnet_b3', pretrained=True)\n        self.model.classifier = nn.Linear(in_features=1536,out_features=out_dim, bias=True)\n        self.model.eval()\n#         self.classifier = nn.Sequential(nn.Linear(1000, out_dim),\n#                                         nn.Softmax())\n#         self.classifier = nn.Sequential(nn.Linear(1000, out_dim))\n        # freeze attibute\n#         for p in self.features.parameters():\n#             p.requires_grad = False\n    def forward(self, input):\n        return self.model(input)","metadata":{"execution":{"iopub.status.busy":"2021-10-31T13:31:58.473625Z","iopub.execute_input":"2021-10-31T13:31:58.475035Z","iopub.status.idle":"2021-10-31T13:31:58.483147Z","shell.execute_reply.started":"2021-10-31T13:31:58.474984Z","shell.execute_reply":"2021-10-31T13:31:58.481706Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# import numpy as np\n# models = timm.create_model('efficientnet_b3', pretrained=False)\n# models.classifier = nn.Linear(in_features=1536,out_features=5, bias=True)\n# print(models)\n","metadata":{"execution":{"iopub.status.busy":"2021-10-31T13:31:58.485632Z","iopub.execute_input":"2021-10-31T13:31:58.486043Z","iopub.status.idle":"2021-10-31T13:31:58.493342Z","shell.execute_reply.started":"2021-10-31T13:31:58.485959Z","shell.execute_reply":"2021-10-31T13:31:58.491910Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# label smoothing loss\nclass LabelSmoothingLoss(nn.Module):\n    def __init__(self, classes, smoothing=0.01, dim=-1): \n        super(LabelSmoothingLoss, self).__init__() \n        self.confidence = 1.0 - smoothing \n        self.smoothing = smoothing \n        self.cls = classes\n        self.dim = dim \n    def forward(self, pred, target): \n        pred = pred.log_softmax(dim=self.dim)\n        with torch.no_grad(): \n            true_dist = torch.zeros_like(pred) \n            true_dist.fill_(self.smoothing / (self.cls - 1)) \n            true_dist.scatter_(1, target.data.unsqueeze(1), self.confidence) \n        return torch.mean(torch.sum(-true_dist * pred, dim=self.dim))\n    \nclass LabelSmoothingCrossEntropy(nn.Module):\n\n    def __init__(self, epsilon: float = 0.1):\n        super().__init__()\n        self.epsilon = epsilon\n\n    def linear_combination(self, x, y, epsilon):\n        return epsilon*x + (1-epsilon)*y\n\n    def reduce_loss(self, loss, reduction='mean'):\n        return loss.mean() if reduction == 'mean' else loss.sum() if reduction == 'sum' else loss\n\n    def forward(self, preds, target, reduction='mean'):\n        n = preds.size()[-1]\n        log_preds = F.log_softmax(preds, dim=-1)\n        loss = self.reduce_loss(-log_preds.sum(dim=-1), reduction=reduction)\n        nll = F.nll_loss(log_preds, target, reduction=reduction)\n        return self.linear_combination(loss/n, nll, self.epsilon)","metadata":{"execution":{"iopub.status.busy":"2021-10-31T13:31:58.495890Z","iopub.execute_input":"2021-10-31T13:31:58.496278Z","iopub.status.idle":"2021-10-31T13:31:58.513265Z","shell.execute_reply.started":"2021-10-31T13:31:58.496235Z","shell.execute_reply":"2021-10-31T13:31:58.511914Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_loss(name):\n    loss_dict = {\n        'CrossEntropy': F.cross_entropy,\n        'LabelSmoothingCrossEntropy':\n        LabelSmoothingCrossEntropy(epsilon=0.1)\n    }\n    return loss_dict[name]\n\nclass OUSMLoss(nn.Module):\n    '''\n    Implementation of \n    Loss with Online Uncertainty Sample Mining:\n    https://arxiv.org/pdf/1901.07759.pdf\n    # Params\n    k: num of samples to drop in a mini batch\n    loss: loss function name (see get_loss function above)\n    trigger: the epoch it starts to train on OUSM (please call `.update(epoch)` each epoch)\n    '''\n\n    def __init__(self, k=1, loss='LabelSmoothingCrossEntropy', trigger=2, ousm=False):\n        super(OUSMLoss, self).__init__()\n        self.k = k\n        self.loss_name = loss\n        self.loss = get_loss(loss)\n        self.trigger = trigger\n        self.ousm = ousm\n\n    def forward(self, logits, targets, indices=None):\n        bs = logits.shape[0]\n        if self.ousm and bs - self.k > 0:\n            losses = self.loss(logits, targets, reduction='none')\n            if len(losses.shape) == 2:\n                losses = losses.mean(1)\n            _, idxs = losses.topk(bs-self.k, largest=False)\n            losses = losses.index_select(0, idxs)\n            return losses.mean()\n        else:\n            return self.loss(logits, targets)\n\n    def update(self, current_epoch):\n        self.current_epoch = current_epoch\n        if current_epoch == self.trigger:\n            self.ousm = True\n            print('criterion: ousm is True.')\n\n    def __repr__(self):\n        return f'OUSM(loss={self.loss_name}, k={self.k}, trigger={self.trigger}, ousm={self.ousm})'","metadata":{"execution":{"iopub.status.busy":"2021-10-31T13:31:58.515366Z","iopub.execute_input":"2021-10-31T13:31:58.515983Z","iopub.status.idle":"2021-10-31T13:31:58.531977Z","shell.execute_reply.started":"2021-10-31T13:31:58.515899Z","shell.execute_reply":"2021-10-31T13:31:58.531003Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def accuracy_score(pred, label):\n    assert len(pred) == len(label), 'Shape not match! pred:{}, label:{}, batch size: {}'.format(np.shape(pred),np.shape(label),cfg['batch_size'])\n    lengh = len(pred)\n    tt_acc = 0\n    for (i, j) in zip(pred, label):\n        tt_acc += 1 if i == j else 0\n    return tt_acc / lengh","metadata":{"execution":{"iopub.status.busy":"2021-10-31T13:31:58.533566Z","iopub.execute_input":"2021-10-31T13:31:58.534067Z","iopub.status.idle":"2021-10-31T13:31:58.546471Z","shell.execute_reply.started":"2021-10-31T13:31:58.534010Z","shell.execute_reply":"2021-10-31T13:31:58.545350Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"## train / valid a epoch","metadata":{"execution":{"iopub.status.busy":"2021-10-31T13:31:58.548144Z","iopub.execute_input":"2021-10-31T13:31:58.549263Z","iopub.status.idle":"2021-10-31T13:31:58.556907Z","shell.execute_reply.started":"2021-10-31T13:31:58.549218Z","shell.execute_reply":"2021-10-31T13:31:58.555775Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train(model, dataloader, criterion, optimizer, scheduler):\n    model.train()\n    totalloss = 0\n    totalacc = 0\n    with tqdm(dataloader,unit='batch',desc='Train') as tqdm_loader:\n        for idx, (ID, img, label) in enumerate(tqdm_loader):\n            img = img.to(device=cfg['device'])\n            label = label.to(device=cfg['device'])\n            label = torch.tensor(label, dtype=torch.long) \n#             print(f'img:\\n{img}')\n            pred = model(img).to(device=cfg['device'])\n#             print(f'pred:\\n{pred} label:\\n{label}')\n            loss = criterion(pred, label)\n            pred = pred.cpu().detach().argmax(dim=1)\n\n            optimizer.zero_grad()\n            loss.backward()           \n            optimizer.step()\n#             scheduler.step()\n            \n            nowloss = loss.detach().item()\n            totalloss += nowloss\n            \n            acc = accuracy_score(pred, label.cpu())\n            totalacc += acc\n            \n            tqdm_loader.set_postfix(loss=nowloss,avgloss=totalloss/(idx+1),avgACC=totalacc/(idx+1) )\n\ndef valid(model, dataloader, certification, fold, epoch):\n    model.eval()\n    totalloss=0\n    totalacc=0\n    with torch.no_grad():\n        with tqdm(dataloader,unit='batch',desc='Valid') as tqdm_loader:\n            for idx, (ID, img, label) in enumerate(tqdm_loader):\n                img = img.to(device=cfg['device'])\n#                 label = label.to(device=cfg['device'])\n                label = torch.tensor(label, dtype=torch.long).to(device=cfg['device'])\n                \n                pred = model(img).to(device=cfg['device'])\n                \n                loss = certification(pred, label)\n                \n                pred = pred.cpu().detach().argmax(dim=1)\n                \n                nowloss = loss.detach().item()\n                totalloss += nowloss\n                \n                acc = accuracy_score(pred, label.cpu())\n                totalacc += acc\n\n                tqdm_loader.set_postfix(loss=nowloss,avgloss=totalloss/(idx+1),avgACC=totalacc/(idx+1) )\n            bestacc = totalacc/len(tqdm_loader)\n            bestloss = totalloss/len(tqdm_loader)\n    return bestacc, bestloss","metadata":{"execution":{"iopub.status.busy":"2021-10-31T13:31:58.558786Z","iopub.execute_input":"2021-10-31T13:31:58.559260Z","iopub.status.idle":"2021-10-31T13:31:58.578227Z","shell.execute_reply.started":"2021-10-31T13:31:58.559216Z","shell.execute_reply":"2021-10-31T13:31:58.576928Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\ndef plotLosses(losses):\n    plt.clf()\n    plt.title(\"Losses\")\n    plt.plot(losses)\n    plt.savefig(\"lossfig.jpg\")","metadata":{"execution":{"iopub.status.busy":"2021-10-31T13:31:58.579779Z","iopub.execute_input":"2021-10-31T13:31:58.580828Z","iopub.status.idle":"2021-10-31T13:31:58.591421Z","shell.execute_reply.started":"2021-10-31T13:31:58.580781Z","shell.execute_reply":"2021-10-31T13:31:58.590349Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Cross validation train","metadata":{}},{"cell_type":"code","source":"def cv_train(dataset):\n    frac = int(len(dataset) / cfg['fold'])\n    accs = []\n    for fold in range(cfg['fold']):\n        print(f'\\nFold {fold}')\n        if cfg['fold'] != 1:\n            train_rg = list(range(0, fold * frac)) + list(range((fold+1) * frac, len(dataset)) )  \n            valid_rg = list(range(fold * frac, (fold+1) * frac))\n        else:\n            train_rg = list(range(0, int(cfg['split'] * len(dataset))))\n            valid_rg = list(range(int(cfg['split'] * len(dataset)), len(dataset)))\n#         print(train_rg)\n#         print(valid_rg)\n        \n        train_dataset = torch.utils.data.Subset(dataset, train_rg)\n        train_dataset.isTrain = True\n        valid_dataset = torch.utils.data.Subset(dataset, valid_rg)\n        valid_dataset.isTrain = False\n#         print(train_dataset)\n#         print(valid_dataset)\n        train_dataloader = DataLoader(train_dataset, batch_size=cfg['batch_size'], shuffle=True, num_workers=2)\n        valid_dataloader = DataLoader(valid_dataset, batch_size=cfg['batch_size'], shuffle=False, num_workers=2)\n        \n        model = CFCmodel(num_class).to(cfg['device'])\n\n#         criterion = LabelSmoothingLoss(classes=num_class, smoothing=cfg['smoothing_factor'])\n        criterion = OUSMLoss()\n        optimizer = Ranger(model.parameters())\n        scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max= 10, eta_min=1e-6)\n#         optimizer = torch.optim.SGD(model.parameters(), lr=0.1, momentum=0.9)\n#         scheduler = torch.optim.lr_scheduler.OneCycleLR(optimizer,max_lr=0.1, steps_per_epoch=len(train_dataloader), epochs=cfg['epoch'])\n        best_fold_model = -1\n        best_fold_loss = 100\n        losses = []\n        for epoch in range(cfg['epoch']):\n            print(f'\\nEpoch {epoch}')\n            train(model, train_dataloader, criterion, optimizer, scheduler)\n            acc, loss = valid(model, valid_dataloader, criterion, fold, epoch)\n            if loss <= best_fold_loss:\n                best_fold_model = epoch\n                best_fold_loss = loss\n            torch.save(model.state_dict(), os.path.join('./', f'model{fold}_{epoch}.pth'))\n                \n            accs.append(acc)\n            losses.append(loss)\n            scheduler.step()\n            criterion.update(epoch)\n        print(f'fold {fold} best epoch: {best_fold_model}')\n        plotLosses(losses)\n#     print(accs)\n    print(f'ac_score: {np.mean(accs)}')\n    ","metadata":{"execution":{"iopub.status.busy":"2021-10-31T13:31:58.593348Z","iopub.execute_input":"2021-10-31T13:31:58.594500Z","iopub.status.idle":"2021-10-31T13:31:58.612505Z","shell.execute_reply.started":"2021-10-31T13:31:58.594453Z","shell.execute_reply":"2021-10-31T13:31:58.611315Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# torch.save(model.state_dict(), os.path.join('./', f'best_model0_9.pth'))","metadata":{"execution":{"iopub.status.busy":"2021-10-31T13:31:58.614879Z","iopub.execute_input":"2021-10-31T13:31:58.615222Z","iopub.status.idle":"2021-10-31T13:31:58.626459Z","shell.execute_reply.started":"2021-10-31T13:31:58.615152Z","shell.execute_reply":"2021-10-31T13:31:58.625317Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_dataset = CLCdataset(train_df)\ncv_train(train_dataset)","metadata":{"execution":{"iopub.status.busy":"2021-10-31T13:31:58.627816Z","iopub.execute_input":"2021-10-31T13:31:58.628142Z","iopub.status.idle":"2021-10-31T17:50:32.528099Z","shell.execute_reply.started":"2021-10-31T13:31:58.628110Z","shell.execute_reply":"2021-10-31T17:50:32.527069Z"},"trusted":true},"execution_count":null,"outputs":[]}]}