{"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 sys\n\nsys.path.append('../input/monai-v081/')\n\nimport collections.abc\nimport pandas as pd\nimport numpy as np\nimport os\nfrom glob import glob\nimport gc\ngc.enable()\nimport matplotlib.pyplot as plt\nimport torch; print(\"\\n \\t...PyTorch Version: \", torch.__version__)\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.utils.data.dataloader import default_collate\nfrom torch.utils.tensorboard import SummaryWriter\nimport torchvision.models\nfrom torchvision import transforms as T\nimport torchvision; print(\"\\n\\t...TorchVision Version: \", torchvision.__version__)\nimport torch.nn.functional as F\nfrom torch.autograd import Variable\nfrom PIL import Image\nimport cv2\nimport albumentations as A\nimport time\nimport os\nimport copy\nimport tifffile\nfrom tqdm import tqdm\n\nfrom torch.cuda.amp import GradScaler, autocast\nfrom torch.autograd import Variable\nimport torch.nn.functional as F\n\nimport numba\nimport numpy as np; print(\"\\n\\t...numpy version: \", np.__version__)\nfrom math import sqrt\nfrom scipy.spatial.distance import directed_hausdorff\nfrom scipy.ndimage import convolve\nfrom scipy.ndimage.morphology import distance_transform_edt as edt\n\nfrom torch.optim.lr_scheduler import _LRScheduler\nfrom torch.optim.lr_scheduler import ReduceLROnPlateau\n\nimport matplotlib.pyplot as plt\nfrom matplotlib.patches import Rectangle\n\nimport warnings\nwarnings.filterwarnings(\"ignore\")\n\nimport pandas as pd\nimport numpy as np\nimport os\nimport unicodedata\nfrom glob import glob\nimport gc\nfrom sklearn.model_selection import KFold,GroupKFold,StratifiedKFold,StratifiedGroupKFold\n\nfrom monai.data import Dataset, load_decathlon_datalist, ITKReader\nfrom monai.data.image_reader import WSIReader, PILReader, ITKReader\nfrom monai.metrics import Cumulative, CumulativeAverage\nfrom monai.transforms import Transform, Compose, LoadImageD, RandFlipd, RandRotate90d, ScaleIntensityRangeD, Resized, ToTensord, LoadImage, ScaleIntensityRange\nfrom monai.apps.pathology.transforms import TileOnGridd\nfrom monai.networks.nets import milmodel\n\ngc.enable()\n\nprint(\"\\n\\t...Import Finished\")","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-09-06T04:04:00.628159Z","iopub.execute_input":"2022-09-06T04:04:00.628674Z","iopub.status.idle":"2022-09-06T04:04:00.646076Z","shell.execute_reply.started":"2022-09-06T04:04:00.62864Z","shell.execute_reply":"2022-09-06T04:04:00.644602Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class cfg:\n    num_classes  = 1\n    mil_mode     = \"att_trans\"\n    tile_count   = 5\n    tile_size    = 256\n    epochs       = 10\n    batch_size   = 4\n    optim_lr     = 2e-8\n    weight_decay = 0.1\n    amp          = True\n    val_every    = 1\n    workers      = 2\n    n_fold       = 5\n    device       = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\n    log_dir      = \"./\"\n    TRAIN        = True\n    existedModelPath = \"../input/multipleinstancelearningmonai/MIL_F1_E5.pt\"","metadata":{"execution":{"iopub.status.busy":"2022-09-06T04:04:00.648296Z","iopub.execute_input":"2022-09-06T04:04:00.648745Z","iopub.status.idle":"2022-09-06T04:04:00.671403Z","shell.execute_reply.started":"2022-09-06T04:04:00.648708Z","shell.execute_reply":"2022-09-06T04:04:00.670349Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!mkdir -p train_images_tiles\n!mkdir -p models\n\ntdf = [{\"image\": \"../input/mayo-clinic-strip-ai/train/00c058_0.tif\", \"label\": 0},\n       {\"image\": \"../input/mayo-clinic-strip-ai/train/029c68_0.tif\", \"label\": 1},\n       {\"image\": \"../input/mayo-clinic-strip-ai/train/006388_0.tif\", \"label\": 1}\n      ]\n\nscale_tool = ScaleIntensityRange(a_min=np.float32(255), a_max=np.float32(0))","metadata":{"execution":{"iopub.status.busy":"2022-09-06T04:04:00.674416Z","iopub.execute_input":"2022-09-06T04:04:00.675144Z","iopub.status.idle":"2022-09-06T04:04:02.764673Z","shell.execute_reply.started":"2022-09-06T04:04:00.675111Z","shell.execute_reply":"2022-09-06T04:04:02.763329Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from openslide.deepzoom import DeepZoomGenerator\nfrom openslide import OpenSlide, open_slide\nfrom random import sample\nfrom itertools import product\nimport math\n\ndef plotTiles(tiled, isDict = True, save = False, name = None):\n    plt.figure(figsize = (15,15))\n    row = 5\n    col = 6\n    plt.tight_layout()\n    for i in range(len(tiled)):\n        plt.subplot(row, col, i+1 )\n        plt.xticks([])\n        plt.yticks([])\n        if isDict:\n            plt.imshow(np.moveaxis(tiled[i]['image'], 0, -1))\n        else:\n            plt.imshow(np.moveaxis(tiled[i], 0, -1))\n    if save:\n        plt.savefig(f\"./{name}.jpg\")\n            \ndef takeRandomRegions(tiles, level_num, region_collection, lst, n):\n    ch = sample(lst, n)\n    region_collection.append(np.array(tiles.get_tile(level_num, ch[0])))\n    region_collection.append(np.array(tiles.get_tile(level_num, ch[1])))\n    region_collection.append(np.array(tiles.get_tile(level_num, ch[2])))\n    return\n            \ndef isMostBG(region, thre = 0.75):\n    h,w,c = region.shape\n    total = h*w*c*255\n    s = np.sum(region, axis = (0,1,2))\n    if ( s / total ) > thre:\n        return True\n    else:\n        return False\n\ndef generateTiles(slide_path, division = 9):\n    slide = open_slide(slide_path)\n    \n    # Divide the tif image to major regions\n    h, w = slide.dimensions\n    factor = min(h,w)\n    factor = factor // division\n    tiles = DeepZoomGenerator(slide, tile_size=factor, overlap=0, limit_bounds=False)\n    level_num = tiles.level_count - 1\n    shape_x, shape_y = tiles.level_tiles[level_num]\n    \n    #Take off the last col and row, bc of incomplete dimensionality\n    if shape_x > division:\n        shape_x -= 1\n    if shape_y > division:\n        shape_y -= 1\n    \n    # Set up for random choices of positions for tiling \n    choices_x = [i for i in range(shape_x)]\n    choices_y = [i for i in range(shape_y)]\n    choices = list(set(product(choices_x, choices_y)))\n    \n    # For large image we only pick 10% of regions for training; increase percent to 30% for testing\n    \"\"\"if h * w >= 600000000:\n        search_percent = 0.1\"\"\"\n\n    region_limit = 2  \n    duration = 5 # Setting find tile region time limit to 3 secs\n    \n   \n    #if isValid:\n    #    duration = 10\n        \n    start_time = time.time()\n    search_regions = 2\n    region_collection = []\n    while(len(region_collection) < region_limit):\n    \n        search_regions = len(choices) if len(choices) < search_regions else search_regions \n        \n        if search_regions == 0 or int(time.time() - start_time) >= duration:\n            #print(f\"Out of time-{int(time.time() - start_time)}s > {duration} / out of choices-{search_regions}\")\n            temp_sample = sorted(list(set(product(choices_x,choices_y))))\n            takeRandomRegions(tiles, level_num, region_collection, temp_sample, 3)\n            break\n        \n        #print(f\"Pick # choices : {len(choices)}, search_regions : {search_regions}\")\n        curr_choices = sample(choices, search_regions)\n        for p in curr_choices:\n            single_tile = tiles.get_tile(level_num, p)\n            np_tile = np.array(single_tile)\n            if isMostBG(np_tile):\n                choices.remove(p)\n                continue\n            region_collection.append(np_tile)\n            if len(region_collection) >= region_limit:\n                break\n    \n    #print(\"Tiling on # of regions : \", len(region_collection))\n    #print(\"Each region has dims : \", region_collection[0].shape)\n    t_count = 40\n    monai_tile = TileOnGridd(keys=[\"image\"],\n                        tile_count= t_count,\n                        tile_size = 256,\n                        random_offset=True,\n                        background_val=255,\n                        return_list_of_dicts=True)\n    \n    tile_collection = []\n    for region in region_collection:\n        channel_first = np.moveaxis(region, -1, 0) # prepare for monai method\n        tiled = monai_tile({'image' : channel_first})\n        tile_collection += tiled\n        \n    return tile_collection, search_regions\n\ndef selectTopTiles(tile_collection, total = 30):\n    \n    tiles_concat = np.expand_dims(tile_collection[0]['image'], axis = 0)\n    for i in range(1, len(tile_collection)):\n        curr_tile = np.expand_dims(tile_collection[i]['image'], axis = 0)\n        tiles_concat = np.concatenate([tiles_concat, curr_tile])\n    \n    idxs = np.argsort(tiles_concat.sum(axis=(1, 2, 3)))[: total]\n    tiles_ret = tiles_concat[idxs]\n    \n    return tiles_ret\n\ndef saveTiles(path, id_, tiles, isDict = True):\n    for idx, t in enumerate(tiles):\n        if isDict:\n            img = np.moveaxis(t['image'],0, -1)\n        else:\n            img = np.moveaxis(t, 0, -1)\n        cv2.imwrite(os.path.join(path, f\"{id_}_tile_{idx}.jpg\"), img)","metadata":{"execution":{"iopub.status.busy":"2022-09-06T04:04:02.766768Z","iopub.execute_input":"2022-09-06T04:04:02.767715Z","iopub.status.idle":"2022-09-06T04:04:02.79472Z","shell.execute_reply.started":"2022-09-06T04:04:02.767671Z","shell.execute_reply":"2022-09-06T04:04:02.793781Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def showSlide(path):\n    slide = OpenSlide(path) # opening a full slide\n    region = (0, 0) # location of the top left pixel\n    level = 0 # level of the picture (we have only 0)\n    size = (10000, 10000) # region size in pixels\n    #size = slide.dimensions\n    region = slide.read_region(region, level, size)\n    plt.figure(figsize=(8, 8))\n    plt.imshow(region)\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2022-09-06T04:04:02.799408Z","iopub.execute_input":"2022-09-06T04:04:02.800168Z","iopub.status.idle":"2022-09-06T04:04:02.808719Z","shell.execute_reply.started":"2022-09-06T04:04:02.800138Z","shell.execute_reply":"2022-09-06T04:04:02.807523Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def getStandardDf(df):\n    df_new = pd.DataFrame()\n    column_nms = [\"image_id\", \"patient_id\", \"image_num\", \"label\"]\n    for c_nm in column_nms:\n        df_new[c_nm] = df[c_nm]\n    df_new[\"image_path\"] = df_new[\"image_id\"].apply(lambda x: os.path.join(DATA_DIR, 'train', x+'.tif'))\n    df_new['label_num'] = df_new[\"label\"].apply(lambda x : 0 if x == 'CE' else 1)\n    return df_new\n\nDATA_DIR = '../input/mayo-clinic-strip-ai'\ntrain_df = pd.read_csv(os.path.join(DATA_DIR, 'train.csv'))\ntrain_df = getStandardDf(train_df)\ndisplay(train_df.head())","metadata":{"execution":{"iopub.status.busy":"2022-09-06T04:04:02.810538Z","iopub.execute_input":"2022-09-06T04:04:02.811331Z","iopub.status.idle":"2022-09-06T04:04:02.839761Z","shell.execute_reply.started":"2022-09-06T04:04:02.811263Z","shell.execute_reply":"2022-09-06T04:04:02.838689Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df.loc[train_df[\"image_path\"] == \"../input/mayo-clinic-strip-ai/train/1d5335_0.tif\"]","metadata":{"execution":{"iopub.status.busy":"2022-09-06T04:04:02.841748Z","iopub.execute_input":"2022-09-06T04:04:02.842155Z","iopub.status.idle":"2022-09-06T04:04:02.856798Z","shell.execute_reply.started":"2022-09-06T04:04:02.842114Z","shell.execute_reply":"2022-09-06T04:04:02.855511Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_df = train_df.copy()[1:8]\nstart_t = time.time()\nif not cfg.TRAIN:\n    pbar = tqdm(test_df.iterrows(), total = len(test_df))\n    for idx, row in pbar:\n        print(row.image_path)\n        tiles_collection, search_regions = generateTiles(row.image_path, 9)\n        #packed = packTiles(tiles_selected, label, isDict = True)\n        tiles_selected = selectTopTiles(tiles_collection)\n        #saveTiles(\"./train_images_tiles\", row.image_id, tiles_selected, isDict = False)\n        #print(len(tiles_selected))\n        plotTiles(tiles_selected, isDict = False, name = row.image_id)\n        #del tiles_collection, tiles_selected\n        #gc.collect()\nprint(f\"{int(time.time()-start_t)} seconds passed\")","metadata":{"execution":{"iopub.status.busy":"2022-09-06T04:04:02.858556Z","iopub.execute_input":"2022-09-06T04:04:02.859104Z","iopub.status.idle":"2022-09-06T04:04:02.868246Z","shell.execute_reply.started":"2022-09-06T04:04:02.85906Z","shell.execute_reply":"2022-09-06T04:04:02.867033Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def packTiles(tiles, label, isDict = False):\n    lst = []\n    for t in tiles:\n        d = dict()\n        d[\"image\"] = t.astype(np.float32) if not isDict else t[\"image\"].astype(np.float32)\n        d[\"label\"] = np.array([label], dtype=np.float32)\n        lst.append(d)   \n    return lst\n\ndef list_data_collate(batch: collections.abc.Sequence):\n    '''\n        Combine instances from a list of dicts into a single dict, by stacking them along first dim\n        [{'image' : 3xHxW}, {'image' : 3xHxW}, {'image' : 3xHxW}...] - > {'image' : Nx3xHxW}\n        followed by the default collate which will form a batch BxNx3xHxW\n    '''\n    for i, item in enumerate(batch):\n        data = item[0]\n        data[\"image\"] = torch.stack([ix[\"image\"] for ix in item], dim=0)\n        batch[i] = data\n    return default_collate(batch)\n\n\nclass BuildDataSet(torch.utils.data.Dataset):\n    def __init__(self, df, isValid, isTest):\n        self.df = df\n        self.label = df[\"label_num\"].values\n        self.img_path = df[\"image_path\"].values\n        self.scale = ScaleIntensityRangeD(keys=[\"image\"], a_min=np.float32(255), a_max=np.float32(0))\n        self.toTensor = ToTensord(keys=[\"image\", \"label\"])\n        self.isValid = isValid\n        self.isTest = isTest\n        \n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, index):\n        path = self.img_path[index]\n        label = self.label[index]\n        tiles_collection, _ = generateTiles(path)\n        if self.isValid or self.isTest:\n            packed = packTiles(tiles_collection, label, isDict = True)\n        else:\n            tiles_selected = selectTopTiles(tiles_collection)\n            packed = packTiles(tiles_selected, label)\n            \n        #print(\"tiles selected # : \", len(packed))\n        for i, p in enumerate(packed):\n            scaled = self.scale(p)\n            packed[i] = self.toTensor(scaled)\n        return packed","metadata":{"execution":{"iopub.status.busy":"2022-09-06T04:04:02.870329Z","iopub.execute_input":"2022-09-06T04:04:02.871164Z","iopub.status.idle":"2022-09-06T04:04:02.884099Z","shell.execute_reply.started":"2022-09-06T04:04:02.871129Z","shell.execute_reply":"2022-09-06T04:04:02.883213Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train_epoch(model, loader, optimizer, scaler):\n    model.train()\n    criterion = nn.BCEWithLogitsLoss()\n    \n    run_loss = CumulativeAverage()\n    run_acc = CumulativeAverage()\n\n    start_time = time.time()\n    loss, acc = 0.0, 0.0\n    \n    tqdm_bar = tqdm(enumerate(loader), total = len(loader))\n    \n    for idx, batch_data in tqdm_bar:\n        data, target = batch_data[\"image\"].to(cfg.device), batch_data[\"label\"].to(cfg.device)\n        \n        with autocast(enabled=cfg.amp):\n            logits = model(data) # [[],[],[]]\n            loss = criterion(logits, target)\n\n        if cfg.amp:\n            scaler.scale(loss).backward()\n            scaler.step(optimizer)\n            scaler.update()\n        else:\n            loss.backward()\n            optimizer.step()\n\n        acc = ((logits.sigmoid() > 0.5).float().detach() == target).float().mean()\n\n        run_loss.append(loss)\n        run_acc.append(acc)\n\n        loss = run_loss.aggregate()\n        acc = run_acc.aggregate()\n        \n        tqdm_bar.set_description(desc=f\"loss: {loss:.4f}, acc : {acc:.4f}, Learning rate : {optimizer.param_groups[0]['lr']}, time : {time.time() - start_time}\")\n    \n    del data, target\n    gc.collect()\n    torch.cuda.empty_cache()\n    return loss, acc","metadata":{"execution":{"iopub.status.busy":"2022-09-06T04:04:02.886867Z","iopub.execute_input":"2022-09-06T04:04:02.888531Z","iopub.status.idle":"2022-09-06T04:04:02.899457Z","shell.execute_reply.started":"2022-09-06T04:04:02.888496Z","shell.execute_reply":"2022-09-06T04:04:02.898522Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def saveModel(model, fold, epoch):\n    filename = f\"MIL_F{fold}_E{epoch+1}.pt\"\n    print(f\"Saving model {filename}\")\n    torch.save(model.state_dict(), \"./models/\"+filename)\n    return\n\n\ndef val_epoch(model, loader):\n    model.eval()\n    \n    criterion = nn.BCEWithLogitsLoss()\n    \n    run_loss = CumulativeAverage()\n    run_acc = CumulativeAverage()\n    PREDS = Cumulative()\n    TARGETS = Cumulative()\n    \n    start_time = time.time()\n    loss, acc = 0.0, 0.0\n    \n    with torch.no_grad():\n        \n        tqdm_bar = tqdm(enumerate(loader), total = len(loader))\n\n        for idx, batch_data in tqdm_bar:\n\n            data, target = batch_data[\"image\"].to(cfg.device), batch_data[\"label\"].to(cfg.device)\n            \n            #print(\"------\", data.shape)\n            \n            with autocast(enabled=cfg.amp):\n                logits = model(data)\n                loss = criterion(logits, target)\n                \n            pred = (logits.sigmoid() > 0.5).float().detach()\n            #target = target.argmax(1)\n            acc = (pred == target).float().mean()\n            \n            run_loss.append(loss)\n            run_acc.append(acc)\n            loss = run_loss.aggregate()\n            acc = run_acc.aggregate()\n\n            PREDS.extend(pred)\n            TARGETS.extend(target)\n            \n            tqdm_bar.set_description(desc=f\"loss: {loss:.4f}, acc : {acc:.4f}, time : {time.time() - start_time}\")\n            \n            \n        PREDS = PREDS.get_buffer().cpu().numpy()\n        TARGETS = TARGETS.get_buffer().cpu().numpy()\n    \n    del data, target\n    gc.collect()\n    torch.cuda.empty_cache()\n    return loss, acc","metadata":{"execution":{"iopub.status.busy":"2022-09-06T04:04:02.901327Z","iopub.execute_input":"2022-09-06T04:04:02.901767Z","iopub.status.idle":"2022-09-06T04:04:02.915517Z","shell.execute_reply.started":"2022-09-06T04:04:02.901701Z","shell.execute_reply":"2022-09-06T04:04:02.914493Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#train_df = train_df[:40]","metadata":{"execution":{"iopub.status.busy":"2022-09-06T04:04:02.916743Z","iopub.execute_input":"2022-09-06T04:04:02.917406Z","iopub.status.idle":"2022-09-06T04:04:02.927242Z","shell.execute_reply.started":"2022-09-06T04:04:02.917371Z","shell.execute_reply":"2022-09-06T04:04:02.92641Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if cfg.TRAIN:\n    SKF = StratifiedKFold(n_splits = cfg.n_fold)\n    fold = 1\n    \n    if cfg.existedModelPath:\n        model = milmodel.MILModel(num_classes=cfg.num_classes, pretrained=False, mil_mode=cfg.mil_mode)\n        model.load_state_dict(torch.load(cfg.existedModelPath))\n        print(\"Loaded model!\")\n    else:    \n        model = milmodel.MILModel(num_classes=cfg.num_classes, pretrained=True, mil_mode=cfg.mil_mode)\n    \n    model.to(cfg.device)\n    for train_index, test_index in SKF.split(train_df['image_id'], train_df['label']):\n        if fold == 2:\n            break\n        print(f\"FOLD {fold} / {cfg.n_fold}\")\n        sub_train_df = train_df.iloc[train_index].copy()\n        sub_valid_df = train_df.iloc[test_index].copy()\n\n        sub_train_dataset = BuildDataSet(sub_train_df, isValid = False, isTest = False)\n        sub_valid_dataset = BuildDataSet(sub_valid_df, isValid = True, isTest = False)\n\n        print(\"Dataset training:\", len(sub_train_dataset), \" validation:\", len(sub_valid_dataset))\n\n        train_loader = DataLoader(sub_train_dataset, batch_size = cfg.batch_size,\n                                 shuffle = True,\n                                 num_workers = cfg.workers,\n                                 pin_memory = True,\n                                 collate_fn = list_data_collate)\n\n        valid_loader = DataLoader(sub_valid_dataset,\n                                  batch_size=1,\n                                  shuffle=False,\n                                  num_workers=cfg.workers,\n                                  pin_memory=True,\n                                  collate_fn=list_data_collate)\n\n        params = model.parameters()\n\n        if  cfg.mil_mode in [\"att_trans\", \"att_trans_pyramid\"]:\n            m = model\n            params = [\n                {\"params\": list(m.attention.parameters()) + list(m.myfc.parameters()) + list(m.net.parameters())},\n                {\"params\": list(m.transformer.parameters()), \"lr\": 6e-6, \"weight_decay\": 0.1},\n            ]\n\n        optimizer = torch.optim.AdamW(params, lr=cfg.optim_lr, weight_decay=cfg.weight_decay)\n        scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=cfg.epochs, eta_min=0)\n\n        best_loss = 1.0\n        best_acc = 0.0\n\n        scaler = None\n        if cfg.amp:\n            scaler = GradScaler()\n\n        for epoch in range(cfg.epochs):\n            print(f\"Epoch {epoch+1} / {cfg.epochs}:\")\n            print(\"Training : \")\n            train_loss, train_acc = train_epoch(model, train_loader, optimizer, scaler = scaler)\n            print(\"loss: {:.4f}\".format(train_loss),\"acc: {:.4f}\".format(train_acc))\n            print(\"Validating : \")\n            val_loss, val_acc = val_epoch(model, valid_loader)\n            print(\"loss: {:.4f}\".format(val_loss), \"acc: {:.4f}\".format(val_acc))\n\n            if val_loss < best_loss:\n                saveModel(model, fold, epoch)\n                best_loss = val_loss\n\n            scheduler.step()\n\n        #del model\n        gc.collect()\n        fold += 1\n\n    print(\"All Done.\")","metadata":{"execution":{"iopub.status.busy":"2022-09-06T04:04:02.929673Z","iopub.execute_input":"2022-09-06T04:04:02.930586Z","iopub.status.idle":"2022-09-06T04:06:18.389861Z","shell.execute_reply.started":"2022-09-06T04:04:02.93056Z","shell.execute_reply":"2022-09-06T04:06:18.388161Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}