{"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":"# Training pipline with ViT using Pytorch \nThis is a pipeline on training with ViT using PyTorch. If anyone finds any improvement, please comment in the notebook!\n\nReferences:\n\n[paper](https://arxiv.org/abs/2010.11929)\n\n[Github](https://github.com/rwightman/pytorch-image-models)","metadata":{}},{"cell_type":"markdown","source":"# Install Timm","metadata":{}},{"cell_type":"markdown","source":"# Print GPU info","metadata":{}},{"cell_type":"code","source":"# gpu_info = !nvidia-smi\n# gpu_info = '\\n'.join(gpu_info)\n# if gpu_info.find('failed') >= 0:\n#     print('Select the Runtime > \"Change runtime type\" menu to enable a GPU accelerator, ')\n#     print('and then re-execute this cell.')\n# else:\n#     print(gpu_info)","metadata":{"execution":{"iopub.status.busy":"2022-12-24T10:43:53.244278Z","iopub.execute_input":"2022-12-24T10:43:53.244550Z","iopub.status.idle":"2022-12-24T10:43:53.250529Z","shell.execute_reply.started":"2022-12-24T10:43:53.244515Z","shell.execute_reply":"2022-12-24T10:43:53.249678Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install huggingface_hub\npth_timm = '../input/pytorchimagemodels/pytorch-image-models-main'\nimport sys\nsys.path.append(pth_timm)\nimport timm","metadata":{"execution":{"iopub.status.busy":"2022-12-24T10:43:53.254849Z","iopub.execute_input":"2022-12-24T10:43:53.255158Z","iopub.status.idle":"2022-12-24T10:44:14.330955Z","shell.execute_reply.started":"2022-12-24T10:43:53.255124Z","shell.execute_reply":"2022-12-24T10:44:14.329923Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Import 3rdparty**","metadata":{}},{"cell_type":"code","source":"import os\nimport pandas as pd\nimport numpy as np\nimport random\nimport cv2\n\nfrom glob import glob\n\nfrom skimage import io\nfrom datetime import datetime\nimport time\n\n\nimport sklearn\nimport warnings\nimport joblib\nfrom sklearn.metrics import roc_auc_score, log_loss\nfrom sklearn import metrics\nimport warnings\nimport pydicom\nfrom scipy.ndimage.interpolation import zoom\n\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import DataLoader, Dataset\nfrom torch.utils.data.sampler import SequentialSampler, RandomSampler\nfrom torch.nn.modules.loss import _WeightedLoss\nimport torch.nn.functional as F\nfrom torch.cuda.amp import autocast, GradScaler\nimport torchvision\nfrom torchvision import transforms\n\n\nfrom matplotlib import pyplot as plt\nfrom sklearn.model_selection import GroupKFold,StratifiedKFold\n\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nfrom tqdm import tqdm","metadata":{"execution":{"iopub.status.busy":"2022-12-24T10:44:14.333098Z","iopub.execute_input":"2022-12-24T10:44:14.333442Z","iopub.status.idle":"2022-12-24T10:44:15.743776Z","shell.execute_reply.started":"2022-12-24T10:44:14.333403Z","shell.execute_reply":"2022-12-24T10:44:15.742821Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def seed_everything(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    torch.backends.cudnn.benchmark = True","metadata":{"execution":{"iopub.status.busy":"2022-12-24T10:44:15.745251Z","iopub.execute_input":"2022-12-24T10:44:15.745593Z","iopub.status.idle":"2022-12-24T10:44:15.751988Z","shell.execute_reply.started":"2022-12-24T10:44:15.745555Z","shell.execute_reply":"2022-12-24T10:44:15.750967Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def seed_everything(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    torch.backends.cudnn.benchmark = True\n    \ndef get_img(path):\n    im_bgr = cv2.imread(path)\n    im_rgb = im_bgr[:, :, ::-1]\n    #print(im_rgb)\n    return im_rgb","metadata":{"execution":{"iopub.status.busy":"2022-12-24T10:44:15.753675Z","iopub.execute_input":"2022-12-24T10:44:15.754276Z","iopub.status.idle":"2022-12-24T10:44:15.762169Z","shell.execute_reply.started":"2022-12-24T10:44:15.754239Z","shell.execute_reply":"2022-12-24T10:44:15.761083Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Global config**","metadata":{}},{"cell_type":"code","source":"class Config:\n    seed = 42\n    data_dir = '../input/cassava-leaf-disease-classification/'\n    train_data_dir = data_dir + 'train_images/'\n    train_csv_path = data_dir + 'train.csv'\n    arch = 'maxxvit_rmlp_nano_rw_256' ## model name\n    device = 'cuda:0'\n    debug = True                 ##\n    \n    image_size = 256 \n    train_batch_size = 4\n    #16\n    val_batch_size = 32\n    epochs = 20                 ## total train epochs\n    freeze_bn_epochs = 5        ## freeze bn weights before epochs\n    \n    lr=1e-4                     ## init learning rate\n    min_lr = 1e-6               ## min learning rate\n    weight_decay = 1e-6\n    num_workers = 4\n    num_splits = 5             ## numbers splits\n    num_classes = 5            ## numbers classes\n    T_0 = 10\n    T_mult = 1\n    accum_iter = 2\n    verbose_step = 1\n    \n    criterion = 'SymmetricEntropyLoss' ## CrossEntropy, LabelSmoothingCrossEntropy\n    label_smoothing = 0.3\n    \n    train_id = [0,1,2,3,4]","metadata":{"execution":{"iopub.status.busy":"2022-12-24T10:44:15.766521Z","iopub.execute_input":"2022-12-24T10:44:15.766907Z","iopub.status.idle":"2022-12-24T10:44:15.774057Z","shell.execute_reply.started":"2022-12-24T10:44:15.766872Z","shell.execute_reply":"2022-12-24T10:44:15.773058Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"CFG_0 = {\n    'fold_num': 5,\n    'seed': 719,\n    'model_arch': 'maxxvit_rmlp_nano_rw_256',\n    'img_size': 256,\n    'epochs': 20,#50\n    'train_bs': 4,\n    #16\n    'valid_bs': 32,\n    'T_0': 10,\n    'lr': 1e-4,\n    'min_lr': 1e-6,\n    'weight_decay':1e-6,\n    'num_workers': 4,\n    'accum_iter': 2, # suppoprt to do batch accumulation for backprop with effectively larger batch size\n    'verbose_step': 1,\n    'device': 'cuda:0'\n}","metadata":{"execution":{"iopub.status.busy":"2022-12-24T10:44:15.775818Z","iopub.execute_input":"2022-12-24T10:44:15.776468Z","iopub.status.idle":"2022-12-24T10:44:15.785785Z","shell.execute_reply.started":"2022-12-24T10:44:15.776410Z","shell.execute_reply":"2022-12-24T10:44:15.785083Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Load Image**","metadata":{}},{"cell_type":"code","source":"def load_image(image_path):\n    img = cv2.imread(image_path)\n    img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n    return img","metadata":{"execution":{"iopub.status.busy":"2022-12-24T10:44:15.788795Z","iopub.execute_input":"2022-12-24T10:44:15.789118Z","iopub.status.idle":"2022-12-24T10:44:15.795483Z","shell.execute_reply.started":"2022-12-24T10:44:15.789090Z","shell.execute_reply":"2022-12-24T10:44:15.794644Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **CassavaDataset**","metadata":{}},{"cell_type":"code","source":"def rand_bbox(size, lam):\n    W = size[0]\n    H = size[1]\n    cut_rat = np.sqrt(1. - lam)\n    cut_w = np.int(W * cut_rat)\n    cut_h = np.int(H * cut_rat)\n\n    # uniform\n    cx = np.random.randint(W)\n    cy = np.random.randint(H)\n\n    bbx1 = np.clip(cx - cut_w // 2, 0, W)\n    bby1 = np.clip(cy - cut_h // 2, 0, H)\n    bbx2 = np.clip(cx + cut_w // 2, 0, W)\n    bby2 = np.clip(cy + cut_h // 2, 0, H)\n    return bbx1, bby1, bbx2, bby2\n\nclass CassavaDataset(Dataset):\n    def __init__(self, df, data_root, \n                 transforms=None, \n                 output_label=True, \n                 one_hot_label=False,\n                 do_fmix=False, \n                 fmix_params={\n                     'alpha': 1., \n                     'decay_power': 3., \n                     'shape': (CFG_0['img_size'], CFG_0['img_size']),\n                     'max_soft': True, \n                     'reformulate': False\n                 },\n                 do_cutmix=False,\n                 cutmix_params={\n                     'alpha': 1,\n                 }\n                ):\n        \n        super().__init__()\n        self.df = df.reset_index(drop=True).copy()\n        self.transforms = transforms\n        self.data_root = data_root\n        self.do_fmix = do_fmix\n        self.fmix_params = fmix_params\n        self.do_cutmix = do_cutmix\n        self.cutmix_params = cutmix_params\n        \n        self.output_label = output_label\n        self.one_hot_label = one_hot_label\n        \n        if output_label == True:\n            self.labels = self.df['label'].values\n            #print(self.labels)\n            \n            if one_hot_label is True:\n                self.labels = np.eye(self.df['label'].max()+1)[self.labels]\n                #print(self.labels)\n            \n    def __len__(self):\n        return self.df.shape[0]\n    \n    def __getitem__(self, index: int):\n        \n        # get labels\n        if self.output_label:\n            target = self.labels[index]\n          \n        img  = get_img(\"{}/{}\".format(self.data_root, self.df.loc[index]['image_id']))\n\n        if self.transforms:\n            img = self.transforms(image=img)['image']\n        \n        if self.do_fmix and np.random.uniform(0., 1., size=1)[0] > 0.5:\n            with torch.no_grad():\n                #lam, mask = sample_mask(**self.fmix_params)\n                \n                lam = np.clip(np.random.beta(self.fmix_params['alpha'], self.fmix_params['alpha']),0.6,0.7)\n                \n                # Make mask, get mean / std\n                mask = make_low_freq_image(self.fmix_params['decay_power'], self.fmix_params['shape'])\n                mask = binarise_mask(mask, lam, self.fmix_params['shape'], self.fmix_params['max_soft'])\n    \n                fmix_ix = np.random.choice(self.df.index, size=1)[0]\n                fmix_img  = get_img(\"{}/{}\".format(self.data_root, self.df.iloc[fmix_ix]['image_id']))\n\n                if self.transforms:\n                    fmix_img = self.transforms(image=fmix_img)['image']\n\n                mask_torch = torch.from_numpy(mask)\n                \n                # mix image\n                img = mask_torch*img+(1.-mask_torch)*fmix_img\n\n                #print(mask.shape)\n\n                #assert self.output_label==True and self.one_hot_label==True\n\n                # mix target\n                rate = mask.sum()/CFG_0['img_size']/CFG_0['img_size']\n                target = rate*target + (1.-rate)*self.labels[fmix_ix]\n                #print(target, mask, img)\n                #assert False\n        \n        if self.do_cutmix and np.random.uniform(0., 1., size=1)[0] > 0.5:\n            #print(img.sum(), img.shape)\n            with torch.no_grad():\n                cmix_ix = np.random.choice(self.df.index, size=1)[0]\n                cmix_img  = get_img(\"{}/{}\".format(self.data_root, self.df.iloc[cmix_ix]['image_id']))\n                if self.transforms:\n                    cmix_img = self.transforms(image=cmix_img)['image']\n                    \n                lam = np.clip(np.random.beta(self.cutmix_params['alpha'], self.cutmix_params['alpha']),0.3,0.4)\n                bbx1, bby1, bbx2, bby2 = rand_bbox((CFG_0['img_size'], CFG_0['img_size']), lam)\n\n                img[:, bbx1:bbx2, bby1:bby2] = cmix_img[:, bbx1:bbx2, bby1:bby2]\n\n                rate = 1 - ((bbx2 - bbx1) * (bby2 - bby1) / (CFG_0.img_size * CFG_0.img_size))\n                target = rate*target + (1.-rate)*self.labels[cmix_ix]\n                \n            #print('-', img.sum())\n            #print(target)\n            #assert False\n                            \n        # do label smoothing\n        #print(type(img), type(target))\n        if self.output_label == True:\n            return img, target\n        else:\n            return img","metadata":{"execution":{"iopub.status.busy":"2022-12-24T10:44:15.797823Z","iopub.execute_input":"2022-12-24T10:44:15.798238Z","iopub.status.idle":"2022-12-24T10:44:15.826095Z","shell.execute_reply.started":"2022-12-24T10:44:15.798206Z","shell.execute_reply":"2022-12-24T10:44:15.825419Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def prepare_dataloader(df, trn_idx, val_idx, data_root='../input/cassava-leaf-disease-classification/train_images/'):\n    \n    from catalyst.data.sampler import BalanceClassSampler\n    \n    train_ = df.loc[trn_idx,:].reset_index(drop=True)\n    valid_ = df.loc[val_idx,:].reset_index(drop=True)\n        \n    train_ds = CassavaDataset(train_, data_root, transforms=get_train_transforms(), output_label=True, one_hot_label=False, do_fmix=False, do_cutmix=False)\n    valid_ds = CassavaDataset(valid_, data_root, transforms=get_valid_transforms(), output_label=True)\n    \n    train_loader = torch.utils.data.DataLoader(\n        train_ds,\n        batch_size=CFG_0['train_bs'],\n        pin_memory=False,\n        drop_last=False,\n        shuffle=True,        \n        num_workers=CFG_0['num_workers'],\n        #sampler=BalanceClassSampler(labels=train_['label'].values, mode=\"downsampling\")\n    )\n    val_loader = torch.utils.data.DataLoader(\n        valid_ds, \n        batch_size=CFG_0['valid_bs'],\n        num_workers=CFG_0['num_workers'],\n        shuffle=False,\n        pin_memory=False,\n    )\n    return train_loader, val_loader","metadata":{"execution":{"iopub.status.busy":"2022-12-24T10:44:15.829034Z","iopub.execute_input":"2022-12-24T10:44:15.829307Z","iopub.status.idle":"2022-12-24T10:44:15.839625Z","shell.execute_reply.started":"2022-12-24T10:44:15.829279Z","shell.execute_reply":"2022-12-24T10:44:15.838769Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **CassavaClassifier**","metadata":{}},{"cell_type":"code","source":"class CassavaClassifier(nn.Module):\n    def __init__(self, model_arch, num_classes, pretrained=False):\n        super().__init__()\n        self.model = timm.create_model(model_arch, pretrained=pretrained)\n        ### vit\n        # num_features = self.model.head.in_features\n        # self.model.head = nn.Linear(num_features, num_classes)\n        \n        \n#         self.model.classifier = nn.Sequential(\n#             nn.Dropout(0.3),\n#             #nn.Linear(num_features, hidden_size,bias=True), nn.ELU(),\n#             nn.Linear(num_features, num_classes, bias=True)\n#         )\n        \n    def forward(self, x):\n        x = self.model(x)\n        return x","metadata":{"execution":{"iopub.status.busy":"2022-12-24T10:44:15.842755Z","iopub.execute_input":"2022-12-24T10:44:15.843088Z","iopub.status.idle":"2022-12-24T10:44:15.851172Z","shell.execute_reply.started":"2022-12-24T10:44:15.843060Z","shell.execute_reply":"2022-12-24T10:44:15.850426Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Train and Val transforms**","metadata":{}},{"cell_type":"code","source":"def get_train_transforms(CFG):\n    return A.Compose([\n            A.RandomResizedCrop(height=CFG.image_size, width=CFG.image_size, p=0.5),\n            A.Transpose(p=0.5),\n            A.HorizontalFlip(p=0.5),\n            A.VerticalFlip(p=0.5),\n            A.RandomRotate90(p=0.5),\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.CenterCrop(CFG.image_size, CFG.image_size),\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            A.CoarseDropout(p=0.5),\n            A.Cutout(p=0.5),\n            ToTensorV2(),\n        ],p=1.0)\n\ndef get_val_transforms(cfg):\n    return A.Compose([\n            A.CenterCrop(CFG.image_size, CFG.image_size, p=1.),\n            A.Resize(CFG.image_size, CFG.image_size),\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(),\n        ],p=1.0)","metadata":{"execution":{"iopub.status.busy":"2022-12-24T10:44:15.853845Z","iopub.execute_input":"2022-12-24T10:44:15.854193Z","iopub.status.idle":"2022-12-24T10:44:15.866605Z","shell.execute_reply.started":"2022-12-24T10:44:15.854140Z","shell.execute_reply":"2022-12-24T10:44:15.865699Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_train_transforms():\n    return A.Compose([\n            A.RandomResizedCrop(CFG_0['img_size'], CFG_0['img_size']),\n            A.Transpose(p=0.5),\n            A.HorizontalFlip(p=0.5),\n            A.VerticalFlip(p=0.5),\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            A.CoarseDropout(p=0.5),\n            A.Cutout(p=0.5),\n            ToTensorV2(p=1.0),\n        ], p=1.)\n  \n        \ndef get_valid_transforms():\n    return A.Compose([\n            A.CenterCrop(CFG_0['img_size'], CFG_0['img_size'], p=1.),\n            A.Resize(CFG_0['img_size'], CFG_0['img_size']),\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(p=1.0),\n        ], p=1.)","metadata":{"execution":{"iopub.status.busy":"2022-12-24T10:44:15.869464Z","iopub.execute_input":"2022-12-24T10:44:15.869757Z","iopub.status.idle":"2022-12-24T10:44:15.879625Z","shell.execute_reply.started":"2022-12-24T10:44:15.869732Z","shell.execute_reply":"2022-12-24T10:44:15.878765Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Train and Val data loader**","metadata":{}},{"cell_type":"code","source":"def load_dataloader(CFG, df, train_idx, val_idx):\n    df_train = df.loc[train_idx,:].reset_index(drop=True)\n    df_val = df.loc[val_idx,:].reset_index(drop=True)\n\n    train_dataset = CassavaDataset(\n        CFG.train_data_dir,\n        df_train,\n        transforms=get_train_transforms(CFG), \n        output_label=True)\n\n    val_dataset = CassavaDataset(\n        CFG.train_data_dir,\n        df_val,\n        transforms=get_val_transforms(CFG), \n        output_label=True)\n\n    train_loader = torch.utils.data.DataLoader(\n        train_dataset,\n        batch_size=CFG.train_batch_size,\n        pin_memory=False,\n        drop_last=False,\n        #shuffle=False,#True        \n        num_workers=CFG.num_workers,\n        sampler=BalanceClassSampler(labels=train_['label'].values, mode=\"downsampling\")\n    )\n\n    val_loader = torch.utils.data.DataLoader(\n        val_dataset, \n        batch_size=CFG.val_batch_size,\n        num_workers=CFG.num_workers,\n        shuffle=False,\n        pin_memory=False,\n    )\n    \n    return train_loader, val_loader","metadata":{"execution":{"iopub.status.busy":"2022-12-24T10:44:15.881734Z","iopub.execute_input":"2022-12-24T10:44:15.882225Z","iopub.status.idle":"2022-12-24T10:44:15.891643Z","shell.execute_reply.started":"2022-12-24T10:44:15.882189Z","shell.execute_reply":"2022-12-24T10:44:15.890801Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Train one epoch**","metadata":{}},{"cell_type":"code","source":"def train_one_epoch(epoch,model,loss_fn,optimizer,train_loader,device,scheduler=None,schd_batch_update=False):\n    model.train()\n    lr = optimizer.state_dict()['param_groups'][0]['lr']\n    \n    running_loss = None\n    pbar = tqdm(enumerate(train_loader),total=len(train_loader))\n    for step,(images,targets) in pbar:\n        images = images.to(device).float()\n        targets = targets.to(device).long()\n        \n        with autocast():\n            preds = model(images)\n            loss = loss_fn(preds,targets)\n        \n            scaler.scale(loss).backward()\n            if running_loss is None:\n                running_loss = loss.item()\n            else:\n                running_loss = running_loss* 0.99 + loss.item()*0.01\n                \n            if ((step + 1) % CFG.accum_iter == 0) or ((step + 1) == len(train_loader)):\n                scaler.step(optimizer)\n                scaler.update()\n                optimizer.zero_grad()\n                \n                if scheduler is not None and schd_batch_update:\n                    scheduler.step()\n            if ((step + 1) % CFG.accum_iter == 0) or ((step + 1) == len(train_loader)):\n                description = f'Train epoch {epoch} loss: {running_loss:.5f}'\n                pbar.set_description(description)\n                \n    if scheduler is not None and schd_batch_update:\n        scheduler.step()","metadata":{"execution":{"iopub.status.busy":"2022-12-24T10:44:15.893241Z","iopub.execute_input":"2022-12-24T10:44:15.893598Z","iopub.status.idle":"2022-12-24T10:44:15.907558Z","shell.execute_reply.started":"2022-12-24T10:44:15.893564Z","shell.execute_reply":"2022-12-24T10:44:15.906802Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Valid one epoch**","metadata":{}},{"cell_type":"code","source":"def valid_one_epoch(epoch,model,loss_fn,val_loader,device,scheduler=None,schd_loss_update=False):\n    model.eval()\n    \n    loss_sum = 0\n    sample_num = 0\n    preds_all = []\n    targets_all = []\n    scores = []\n    \n    pbar = tqdm(enumerate(val_loader),total=len(val_loader))\n    for step,(images,targets) in pbar:\n        images = images.to(device).float()\n        targets = targets.to(device).long()\n        preds = model(images)\n            \n        preds_all += [torch.argmax(preds,1).detach().cpu().numpy()]\n        targets_all += [targets.detach().cpu().numpy()]\n\n        loss = loss_fn(preds,targets)\n        loss_sum += loss.item()*targets.shape[0]\n        sample_num += targets.shape[0]\n           \n        if ((step + 1) % CFG.accum_iter == 0) or ((step + 1) == len(train_loader)):\n            description = f'Val epoch {epoch} loss: {loss_sum/sample_num:.5f}'\n            pbar.set_description(description)\n            \n    preds_all = np.concatenate(preds_all)\n    targets_all = np.concatenate(targets_all)\n    accuracy = (preds_all == targets_all).mean()\n    print(f'Validation multi-class accuracy = {accuracy:.5f}')\n    \n    if scheduler is not None:\n        if schd_loss_update:\n            scheduler.step(loss_sum/sample_num)\n        else:\n            scheduler.step()\n    \n    return accuracy","metadata":{"execution":{"iopub.status.busy":"2022-12-24T10:44:15.908860Z","iopub.execute_input":"2022-12-24T10:44:15.909224Z","iopub.status.idle":"2022-12-24T10:44:15.921639Z","shell.execute_reply.started":"2022-12-24T10:44:15.909190Z","shell.execute_reply":"2022-12-24T10:44:15.920819Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Freeze bn weights**","metadata":{}},{"cell_type":"code","source":"################ freeze bn \ndef freeze_batchnorm_stats(net):\n    try:\n        for m in net.modules():\n            if isinstance(m,nn.BatchNorm2d) or isinstance(m,nn.LayerNorm):\n                required_grad = False\n#                 m.eval()\n\n    except ValuError:\n        print('error with batchnorm2d or layernorm')\n        return","metadata":{"execution":{"iopub.status.busy":"2022-12-24T10:44:15.922994Z","iopub.execute_input":"2022-12-24T10:44:15.923612Z","iopub.status.idle":"2022-12-24T10:44:15.933575Z","shell.execute_reply.started":"2022-12-24T10:44:15.923566Z","shell.execute_reply":"2022-12-24T10:44:15.932789Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Label Smoothing Cross Entropy Loss**","metadata":{}},{"cell_type":"code","source":"class LabelSmoothingCrossEntropy(nn.Module):\n    \"\"\"\n    NLL loss with label smoothing.\n    \"\"\"\n    def __init__(self, smoothing=0.1):\n        \"\"\"\n        Constructor for the LabelSmoothing module.\n        :param smoothing: label smoothing factor\n        \"\"\"\n        super(LabelSmoothingCrossEntropy, self).__init__()\n        assert smoothing < 1.0\n        self.smoothing = smoothing\n        self.confidence = 1. - smoothing\n\n    def forward(self, x, target):\n        logprobs = F.log_softmax(x, dim=-1)\n        nll_loss = -logprobs.gather(dim=-1, index=target.unsqueeze(1))\n        nll_loss = nll_loss.squeeze(1)\n        smooth_loss = -logprobs.mean(dim=-1)\n        loss = self.confidence * nll_loss + self.smoothing * smooth_loss\n        return loss.mean()","metadata":{"execution":{"iopub.status.busy":"2022-12-24T10:44:15.936639Z","iopub.execute_input":"2022-12-24T10:44:15.936895Z","iopub.status.idle":"2022-12-24T10:44:15.946754Z","shell.execute_reply.started":"2022-12-24T10:44:15.936871Z","shell.execute_reply":"2022-12-24T10:44:15.946102Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **FocalCosineLoss**","metadata":{}},{"cell_type":"code","source":"class FocalCosineLoss(nn.Module):\n    def __init__(self, alpha=1, gamma=2, xent=.1):\n        super(FocalCosineLoss, self).__init__()\n        self.alpha = alpha\n        self.gamma = gamma\n\n        self.xent = xent\n\n        self.y = torch.Tensor([1]).cuda()\n\n    def forward(self, input, target, reduction=\"mean\"):\n        cosine_loss = F.cosine_embedding_loss(input, F.one_hot(target, num_classes=input.size(-1)), self.y, reduction=reduction)\n\n        cent_loss = F.cross_entropy(F.normalize(input), target, reduce=False)\n        pt = torch.exp(-cent_loss)\n        focal_loss = self.alpha * (1-pt)**self.gamma * cent_loss\n\n        if reduction == \"mean\":\n            focal_loss = torch.mean(focal_loss)\n\n        return cosine_loss + self.xent * focal_loss","metadata":{"execution":{"iopub.status.busy":"2022-12-24T10:44:15.949852Z","iopub.execute_input":"2022-12-24T10:44:15.950212Z","iopub.status.idle":"2022-12-24T10:44:15.960387Z","shell.execute_reply.started":"2022-12-24T10:44:15.950187Z","shell.execute_reply":"2022-12-24T10:44:15.959724Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **SymmetricEntropyLoss**","metadata":{}},{"cell_type":"code","source":"class SymmetricCrossEntropy(nn.Module):\n\n    def __init__(self, alpha=0.1, beta=1.0, num_classes=5):\n        super(SymmetricCrossEntropy, self).__init__()\n        self.alpha = alpha\n        self.beta = beta\n        self.num_classes = num_classes\n\n    def forward(self, logits, targets, reduction='mean'):\n        onehot_targets = torch.eye(self.num_classes)[targets].cuda()\n        ce_loss = F.cross_entropy(logits, targets, reduction=reduction)\n        rce_loss = (-onehot_targets*logits.softmax(1).clamp(1e-7, 1.0).log()).sum(1)\n        if reduction == 'mean':\n            rce_loss = rce_loss.mean()\n        elif reduction == 'sum':\n            rce_loss = rce_loss.sum()\n        return self.alpha * ce_loss + self.beta * rce_loss","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Main Loop**","metadata":{}},{"cell_type":"code","source":"if __name__ == '__main__':\n    \n    CFG = Config\n    train = pd.read_csv(CFG.train_csv_path)\n    \n#     if CFG.debug:\n#         CFG.epochs = 1\n#         train = train.sample(100,random_state=CFG.seed).reset_index(drop=True)\n    \n    print('CFG seed is ', CFG.seed)\n    if CFG.seed is not None:\n        seed_everything(CFG.seed)\n    \n    folds = StratifiedKFold(\n        n_splits=CFG.num_splits, \n        shuffle=True, \n        random_state=CFG.seed).split(np.arange(train.shape[0]), train.label.values)\n    \n    cross_accuracy = []\n    for fold,(train_idx,val_idx) in enumerate(folds):\n        print(fold)\n        if fold<=1:\n            continue\n        ########\n        # load data\n        #######\n        train_loader, val_loader = prepare_dataloader(train, train_idx, val_idx, data_root='../input/cassava-leaf-disease-classification/train_images/')\n#         train_loader,val_loader = load_dataloader(CFG, train, train_idx, val_idx)\n        \n        device = torch.device(CFG.device)\n#         assert(CFG.num_classes ==  train.label.nunique())\n        print(CFG.arch)\n        model = CassavaClassifier(CFG.arch, train.label.nunique(), pretrained=True).to(device)\n        \n        scaler = GradScaler()\n        optimizer = torch.optim.Adam(\n            model.parameters(), \n            lr=CFG.lr, \n            weight_decay=CFG.weight_decay)\n\n        scheduler = torch.optim.lr_scheduler.CosineAnnealingWarmRestarts(\n            optimizer, \n            T_0=CFG.T_0, \n            T_mult=CFG.T_mult, \n            eta_min=CFG.min_lr, \n            last_epoch=-1)\n    \n        ########\n        # criterion\n        #######\n        if CFG.criterion == 'LabelSmoothingCrossEntropy':  #### label smoothing cross entropy\n            loss_train = LabelSmoothingCrossEntropy(smoothing=CFG.label_smoothing)\n        elif CFG.criterion == 'FocalCosineLoss':\n            loss_train = FocalCosineLoss().to(device)\n        elif CFG.criterion == 'SymmetricEntropyLoss':\n            loss_train = SymmetricEntropyLoss().to(device)\n        else:\n            loss_train = nn.CrossEntropyLoss().to(device)\n        loss_val = nn.CrossEntropyLoss().to(device)\n        \n        best_accuracy = 0\n        best_epoch = 0\n        for epoch in range(CFG.epochs):\n            if epoch < CFG.freeze_bn_epochs:\n                freeze_batchnorm_stats(model)  \n            train_one_epoch(\n                epoch, \n                model, \n                loss_train, \n                optimizer, \n                train_loader, \n                device, \n                scheduler=scheduler, \n                schd_batch_update=False)\n\n            with torch.no_grad():\n                epoch_accuracy = valid_one_epoch(\n                    epoch, \n                    model, \n                    loss_val, \n                    val_loader, \n                    device, \n                    scheduler=None, \n                    schd_loss_update=False)\n\n            if epoch_accuracy > best_accuracy:\n                torch.save(model.state_dict(),'{}_fold{}_best.ckpt'.format(CFG.arch, fold))\n                best_accuracy = epoch_accuracy\n                best_epoch = epoch\n                print('Best model is saved')\n        cross_accuracy += [best_accuracy]\n        print('Fold{} best accuracy = {} in epoch {}'.format(fold,best_accuracy,best_epoch))\n        del model, optimizer, train_loader, val_loader, scaler, scheduler\n        torch.cuda.empty_cache()\n    print('{} folds cross validation CV = {:.5f}'.format(CFG.num_splits,np.average(cross_accuracy)))","metadata":{"execution":{"iopub.status.busy":"2022-12-24T10:44:15.963011Z","iopub.execute_input":"2022-12-24T10:44:15.963408Z","iopub.status.idle":"2022-12-24T17:45:02.702753Z","shell.execute_reply.started":"2022-12-24T10:44:15.963344Z","shell.execute_reply":"2022-12-24T17:45:02.699043Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Please Upvote if you liked the kernel ! Cheers.\n\nReferences :\nhttps://www.kaggle.com/khyeh0719/pytorch-efficientnet-baseline-train-amp-aug. \nPlease Upvote too.","metadata":{}}]}