{"cells":[{"metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\n# for dirname, _, filenames in os.walk('/kaggle/input'):\n#     for filename in filenames:\n#         print(os.path.join(dirname, filename))\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"# FOR TPU \n# !curl https://raw.githubusercontent.com/pytorch/xla/master/contrib/scripts/env-setup.py -o pytorch-xla-env-setup.py\n# !python pytorch-xla-env-setup.py --apt-packages libomp5 libopenblas-dev\n\n# !curl https://raw.githubusercontent.com/pytorch/xla/master/contrib/scripts/env-setup.py -o pytorch-xla-env-setup.py\n# !python pytorch-xla-env-setup.py --version nightly --apt-packages libomp5 libopenblas-dev","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"# !pip install torch==1.5.0\n# !pip install torchvision\n# !pip install pytorch-lightning\n# !pip install pretrainedmodels\n!pip install timm","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"import torch\nfrom torch import nn\nimport torch.nn.functional as F\nfrom torchvision import transforms\n# import pretrainedmodels\nfrom torch.optim.swa_utils import AveragedModel, update_bn\n\n#TPU specific\n# import torch_xla.core.xla_model as xm\n# import torch_xla.distributed.parallel_loader as parallel_loader\n# import torch_xla.distributed.xla_multiprocessing as xmp\n\nimport pytorch_lightning as pl\nfrom pytorch_lightning.metrics.functional import accuracy\nfrom pytorch_lightning.callbacks import EarlyStopping,ModelCheckpoint,LearningRateMonitor\nimport timm\n\nfrom sklearn import metrics, model_selection, preprocessing\nfrom PIL import Image\nfrom collections import Counter,OrderedDict\nimport json\nimport torchvision\nimport time\nimport albumentations\nimport matplotlib.pyplot as plt\nimport cv2\n\nfrom joblib import Parallel, delayed\nimport warnings\nwarnings.filterwarnings(\"ignore\")\n\n%matplotlib inline","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"f = open('/kaggle/input/cassava-leaf-disease-classification/label_num_to_disease_map.json')\nmappings = json.load(f)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"mappings = { 0: 'Cassava Bacterial Blight (CBB)',\n             1: 'Cassava Brown Streak Disease (CBSD)',\n             2: 'Cassava Green Mottle (CGM)',\n             3: 'Cassava Mosaic Disease (CMD)',\n             4: 'Healthy'}","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"images_path = '/kaggle/input/cassava-leaf-disease-classification/train_images/'\npath = \"/kaggle/input/cassava-leaf-disease-classification/train.csv\"\ndf = pd.read_csv(path)\ndf['image_id'] = images_path + df['image_id']","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"#train-validation split \ndf_train,df_valid = model_selection.train_test_split(df,test_size=0.1,random_state=42,stratify=df.label.values)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"df_train.label.value_counts(ascending=True).plot(kind='bar')\ndf_train.label.value_counts()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"df_valid.label.value_counts(ascending=True).plot(kind='bar')\ndf_valid.label.value_counts(ascending=True)","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"## Things which did not work\n1. Weighted Sampling\n2. Removing the background -> Chirag pre-processing steps\n"},{"metadata":{"trusted":true},"cell_type":"code","source":"def weighted_sampling(y,effective_weights=False):\n    \n    samples_per_cls = np.array([len(np.where(y==t)[0]) for t in np.unique(y)])\n    \n    if effective_weights: #from Class balance loss based on Effective number of samples paper\n        effective_num = 1.0 - np.power(cfg.BETA,samples_per_cls)\n        weights = (1.0-cfg.BETA)/np.array(effective_num)\n        weights = weights / np.sum(weights) * len(samples_per_cls)\n        \n    else:\n         weights = 1./(samples_per_cls)\n    \n    samples_weights = torch.from_numpy(np.array([weights[t] for t in y]))\n    \n    #define a sampler\n    sampler = torch.utils.data.WeightedRandomSampler(samples_weights.type('torch.DoubleTensor'),len(samples_weights))\n    \n    return weights,samples_per_cls,samples_weights,sampler","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"weights,samples_per_class,samples_weights,weighted_sampler = weighted_sampling(df_train.label)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"import glob \npaths = [path for path in glob.glob('/kaggle/input/cassava-leaf-disease-classification/train_images/*.jpg')]\ndef remove_bg(image):\n    hsv = cv2.cvtColor(image, cv2.COLOR_RGB2HSV)\n\n    ## mask of green (25,25,25) ~ (86, 255,255)\n    mask = cv2.inRange(hsv, (15, 25, 25), (90, 255,255))\n    # mask = cv2.inRange(hsv, (36, 25, 25), (70, 255,255))\n\n    ## slice the green\n    imask = mask>0\n    green = np.zeros_like(image, np.uint8)\n    green[imask] = image[imask]\n    return green\n\nindex = np.random.randint(len(paths))\nprint(index)\nprint(paths[index ])\n\nimage = Image.open(paths[index]).convert('RGB')\nimage_arr = np.array(image,dtype=np.uint8)\nimage_arr = remove_bg(image_arr)\nimage = valid_aug(image=image_arr)['image']\nplt.imshow(image);","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"collapsed":true},"cell_type":"code","source":"print(f'[INFO] Samples per class is - {samples_per_class}')\nprint(f'[INFO] Weights per calss is - {weights}')\nCounter(samples_weights.numpy())","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"## Augmentations"},{"metadata":{"trusted":true},"cell_type":"code","source":"# from Abhishek Thakur's notebook \ntrain_aug = albumentations.Compose([\n            albumentations.augmentations.transforms.LongestMaxSize(max_size=256,always_apply=True),\n#             albumentations.RandomResizedCrop(256, 256),\n            albumentations.Resize(224,224,always_apply=True),\n            albumentations.Transpose(p=0.5),\n            albumentations.HorizontalFlip(p=0.5),\n            albumentations.VerticalFlip(p=0.5),\n            albumentations.ShiftScaleRotate(p=0.5),\n            albumentations.HueSaturationValue(\n                hue_shift_limit=0.2, \n                sat_shift_limit=0.2, \n                val_shift_limit=0.2, \n                p=0.5\n            ),\n            albumentations.RandomBrightnessContrast(\n                brightness_limit=(-0.1,0.1), \n                contrast_limit=(-0.1, 0.1), \n                p=0.5\n            ),\n            albumentations.CoarseDropout(p=0.5),\n            albumentations.Cutout(p=0.5)], p=1.)\n\nvalid_aug = albumentations.Compose([\n                                    albumentations.augmentations.transforms.LongestMaxSize(max_size=256,always_apply=True),\n                      #             albumentations.RandomResizedCrop(256, 256),\n                                    albumentations.Resize(224,224,always_apply=True)])","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"## Dataset"},{"metadata":{"trusted":true},"cell_type":"code","source":"class CasavaDataset(torch.utils.data.Dataset):\n    def __init__(self,df,augmentations=None,mode='train'):\n        self.df = df\n        self.colums = self.df.columns #['image_id','label','cls_names']\n        self.image_paths = self.df['image_id'].to_numpy()\n        self.targets = self.df['label'].to_numpy()\n        self.classes = self.df['label'].unique()\n        self.augmentations = augmentations\n        self.transforms = transforms.Compose([transforms.ToTensor(),\n                                             transforms.Normalize(mean =[0.485, 0.456, 0.406],\n                                                                  std = [0.229, 0.224, 0.225])])\n        self.mode = mode\n        \n    def __len__(self):\n            return len(self.image_paths)\n        \n    def __getitem__(self,index):\n        image_paths = self.image_paths[index]\n        image = Image.open(image_paths).convert('RGB')\n#         image = image.resize(size=(224,224))\n        image_arr = np.array(image,dtype=np.uint8)\n        \n        #chirag's code for removing background from the images only during training\n        if self.mode == 'train':\n            image_arr = self.remove_bg(image_arr)\n        \n        if self.augmentations:\n            image = self.augmentations(image=image_arr)['image']\n            \n        \n        image = self.transforms(image)\n        \n        labels = self.targets[index]\n\n        return image,labels,image_paths\n    \n    def remove_bg(self,image):\n        hsv = cv2.cvtColor(image, cv2.COLOR_RGB2HSV)\n\n        ## mask of green (25,25,25) ~ (86, 255,255)\n        mask = cv2.inRange(hsv, (15, 25, 25), (90, 255,255))\n        # mask = cv2.inRange(hsv, (36, 25, 25), (70, 255,255))\n\n        ## slice the green\n        imask = mask>0\n        green = np.zeros_like(image, np.uint8)\n        green[imask] = image[imask]\n        return green","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"# custom dataset\ntrain_dataset = CasavaDataset(df_train,augmentations=train_aug,mode='train')\nvalid_dataset = CasavaDataset(df_valid,augmentations=valid_aug,mode='valid')\n\n#dataloaders for testing\ntrain_loader = torch.utils.data.DataLoader(train_dataset,batch_size=1,shuffle=True,\n                                          num_workers = 1,sampler=None)\n\nvalid_loader = torch.utils.data.DataLoader(valid_dataset,batch_size=1,shuffle=False,\n                                          num_workers = 1)","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"## Lightning Model"},{"metadata":{"trusted":true},"cell_type":"code","source":"import math\n#using arcface as criterion + label smoothing\nclass ArcFace(nn.Module):\n    def __init__(self,in_features=2048,out_features=5,s=64.0,m=0.5,easy_margin=False):\n        super(ArcFace,self).__init__()\n        self.in_features = in_features\n        self.out_features = out_features\n        \n        self.s = s\n        self.m = m\n        \n        self.kernel = torch.nn.Parameter(torch.FloatTensor(self.in_features,self.out_features))\n        nn.init.xavier_normal_(self.kernel)\n        \n        self.easy_margin = easy_margin\n        self.cos_m = math.cos(m)\n        self.sin_m = math.sin(m)\n        self.th = math.cos(math.pi - m)\n        self.mm = math.sin(math.pi - m) * m\n        \n    def forward(self,embeddings,label):\n        #normalize the embeddings\n        embeddings = F.normalize(embeddings,p=2,dim=1)\n        \n        kernel_norm = F.normalize(self.kernel,p=2,dim=0)\n        \n        cos_theta = torch.mm(embeddings,kernel_norm)\n        cos_theta = cos_theta.clamp(-1,1)\n        \n        with torch.no_grad():\n            origin_cos = cos_theta.clone()\n        \n        target_logit = cos_theta[torch.arange(0,embeddings.size(0)),label].view(-1,1)\n        \n        sin_theta = torch.sqrt(1.0 - torch.pow(target_logit,2))\n        cos_theta_m = target_logit * self.cos_m - sin_theta * self.sin_m\n        \n        if self.easy_margin:\n            final_target_logit = torch.where(target_logit>0, cos_theta_m,target_logit)\n        else:\n            final_target_logit = torch.where(target_logit>self.th, cos_theta_m, target_logit - self.mm)\n            \n        cos_theta.scatter_(1,label.view(-1,1).long(), final_target_logit)\n        \n        output = cos_theta * self.s\n        \n        return output, origin_cos * self.s\n    \n    \nclass LabelSmoothingCrossEntropy(nn.Module):\n    def __init__(self,epsilon:float = 0.1,reduction='mean'):\n        super().__init__()\n        self.epsilon = epsilon\n        self.reduction = reduction\n        \n    def forward(self,preds,labels):\n        n = preds.size()[-1]\n        log_preds = F.log_softmax(preds,dim=-1)\n        loss = reduce_loss(-log_preds.sum(dim=-1),self.reduction)\n        nll = F.nll_loss(log_preds,labels,reduction=self.reduction)\n        return linear_combination(loss/n , nll, self.epsilon)\n\n    \ndef reduce_loss(loss,reduction='mean'):\n    return loss.mean() if reduction == 'mean' else loss.sum() if reduction=='sum' else loss\n\ndef linear_combination(x,y,epsilon):\n    return epsilon*x + (1-epsilon) * y","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"class LitModel(pl.LightningModule):\n    def __init__(self,classify,n_cls=5,criterion=nn.CrossEntropyLoss(),head=None,t_data=None,v_data=None):\n        super().__init__()\n        self.classify = classify\n        self.n_cls = n_cls\n        self.model = self.get_model()\n        self.criterion = criterion\n        self.learning_rate = 0.0009\n        self.t_data = t_data\n        self.v_data = v_data\n        self.batch_size = 112\n        \n        #head\n        self.head = head\n\n    def forward(self,x):\n        #returns the unnormalized embeddings\n        x = self.model(x)        \n        return x\n      \n    def get_model(self):\n        model = timm.create_model('legacy_seresnext50_32x4d',pretrained=False,num_classes=0)\n        state_dicts = torch.load('/kaggle/input/pre-trained-model/se_resnext50_32x4d-a260b3a4.pth')\n        \n        #remove last layer from the state dict\n        del state_dicts['last_linear.weight']\n        del state_dicts['last_linear.bias']\n        \n        #load the state_dicts to the model\n        model.load_state_dict(state_dicts)\n\n\n        return model\n\n    \n    def separate_bn(self):\n        all_params = self.model.parameters()\n        paras_only_bn = []\n        for pname,p in self.model.named_parameters():\n            if pname.find('bn') >= 0:\n                paras_only_bn.append(p)\n        paras_only_bn_id = list(map(id,paras_only_bn))\n        paras_wo_bn = list(filter(lambda p: id(p) not in paras_only_bn_id,all_params))\n        return paras_only_bn, paras_wo_bn\n        \n    \n    def configure_optimizers(self):\n        \n        with_bn,without_bn = self.separate_bn()\n        optimizer = torch.optim.AdamW([{'params': filter(lambda p: p.requires_grad, with_bn)},\n                               {'params': filter(lambda p: p.requires_grad, without_bn), 'weight_decay': 0.01},\n                                {'params': self.head.parameters(),'weight_decay': 0.01}],\n                                 lr = self.learning_rate)\n        \n        lr_scheduler = {'scheduler': torch.optim.lr_scheduler.CosineAnnealingWarmRestarts(optimizer,T_0=10,eta_min=1e-6),\n                        'name': 'step'\n                       }\n        \n        return [optimizer],[lr_scheduler]\n    \n        \n    def training_step(self,train_batch,batch_idx):\n        x,y,path = train_batch\n        embeddings = self(x)\n        y_pred,original_logits = self.head(embeddings,y)\n        \n        loss = self.criterion(y_pred,y)\n        self.log('train_loss',loss,prog_bar=True)\n        \n        preds = torch.argmax(original_logits.data,dim=1)\n        acc = accuracy(preds,y)\n        self.log('train_acc',acc,prog_bar=True)\n        \n        return loss\n    \n    def validation_step(self,val_batch,batch_idx):\n        x,y,path = val_batch\n        embeddings = self(x)\n        y_pred,original_logits = self.head(embeddings,y)\n        loss = self.criterion(y_pred,y)\n        \n        self.log('val_loss',loss,prog_bar=True)\n        \n        preds = torch.argmax(original_logits.data,dim=1)\n        acc = accuracy(preds,y)\n        self.log('val_acc',acc,prog_bar=True)\n        \n        return loss\n\n    def train_dataloader(self):\n        return torch.utils.data.DataLoader(self.t_data,self.batch_size,shuffle=True)\n\n    def val_dataloader(self):\n      return torch.utils.data.DataLoader(self.v_data,self.batch_size)\n","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"lit_model = LitModel(classify=True,n_cls=5,\n                     head=ArcFace(),\n                     criterion=nn.CrossEntropyLoss(),\n                     t_data=train_dataset,\n                     v_data=valid_dataset)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"checkpoint_callback_val_acc = ModelCheckpoint(monitor='val_acc')\ncheckpoint_callback_val_loss = ModelCheckpoint(monitor='val_loss')\ncheckpoint_callback_train_loss = ModelCheckpoint(monitor='train_loss')\ncheckpoint_callback_train_acc = ModelCheckpoint(monitor='train_acc')\nlr_monitor = LearningRateMonitor(logging_interval='step')\nearly_stopping = EarlyStopping(monitor='val_loss')\n\ncallbacks = [early_stopping,lr_monitor,\n             checkpoint_callback_val_acc,\n             checkpoint_callback_val_loss,\n             checkpoint_callback_train_loss,\n             checkpoint_callback_train_acc]","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"trainer = pl.Trainer(gpus = 1,min_epochs = 1,max_epochs = 10,\n                    callbacks = callbacks,\n                    progress_bar_refresh_rate = 10,\n                    check_val_every_n_epoch = 1,\n                    auto_scale_batch_size = None,\n                    auto_lr_find = False,\n                    resume_from_checkpoint= '/kaggle/working/epoch=1-v2.ckpt'\n                    )","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"# #lr_finder\ntrainer.tune(lit_model)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"trainer.fit(lit_model)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"!nvidia-smi","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"print(f'[INFO] Best Val acc model - {checkpoint_callback_val_acc.best_k_models}')\nprint(f'[INFO] Best Val loss model - {checkpoint_callback_val_loss.best_k_models}')\n\nprint(f'[INFO] Best train acc model - {checkpoint_callback_train_acc.best_k_models}')\nprint(f'[INFO] Best train loss model - {checkpoint_callback_train_loss.best_k_models}')","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"tes1 = torch.load('/kaggle/working/lightning_logs/version_0/checkpoints/epoch=1.ckpt')\ntes1.keys()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"!ls lightning_logs/version_0/checkpoints/","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"!cp lightning_logs/version_0/checkpoints/*.ckpt /kaggle/working/","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Loading the pre-trained model"},{"metadata":{"trusted":true},"cell_type":"code","source":"\n# trained_state_dicts = torch.load('/kaggle/input/exp1-models/epoch9-v0_best_val_loss.ckpt',map_location='cpu')['state_dict']\n# lit_model.load_state_dict(trained_state_dicts)","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Using SWA things -> Did not work\n   1. Using the concept from https://github.com/PyTorchLightning/pytorch-lightning/blob/master/notebooks/07-cifar10-baseline.ipynb"},{"metadata":{"trusted":true},"cell_type":"code","source":"class SWAResnet(pl.LightningModule):\n    def __init__(self,trained_model,lr=0.01):\n        super().__init__()\n        self.save_hyperparameters('lr')\n        self.model = trained_model\n        self.swa_model = AveragedModel(self.model)\n        self.v_data = valid_dataset\n        self.t_data = train_dataset\n        self.batch_size = 256\n        self.criterion = nn.CrossEntropyLoss()\n        \n        \n    def forward(self,x):\n        logit = self.swa_model(x)\n        return logit\n            \n    def training_epoch_end(self,training_step_outputs):\n        self.swa_model.update_parameters(self.model)\n        \n    def training_step(self,train_batch,batch_idx):\n        x,y,path = train_batch\n        logits = self.model(x)\n        preds = torch.argmax(logits,dim=1)\n\n        loss = self.criterion(logits,y)\n        self.log('train_loss',loss,prog_bar=True)\n\n        acc = accuracy(preds,y)\n        self.log('train_acc',acc,prog_bar=True)\n\n        return loss\n    \n    def validation_step(self,val_batch,batch_idx):\n        x,y,path = val_batch\n        logits = self.model(x)\n        loss = self.criterion(logits,y)\n        self.log('val_loss',loss,prog_bar=True)\n        \n        preds = torch.argmax(logits,dim=1)\n        acc = accuracy(preds,y)\n        self.log('val_acc',acc,prog_bar=True)\n        \n        return loss\n\n    def train_dataloader(self):\n        return torch.utils.data.DataLoader(self.t_data,self.batch_size)\n\n    def val_dataloader(self):\n      return torch.utils.data.DataLoader(self.v_data,self.batch_size)\n\n    def separate_bn(self):\n        all_params = self.model.parameters()\n        paras_only_bn = []\n        for pname,p in self.model.named_parameters():\n            if pname.find('bn') >= 0:\n                paras_only_bn.append(p)\n        paras_only_bn_id = list(map(id,paras_only_bn))\n        paras_wo_bn = list(filter(lambda p: id(p) not in paras_only_bn_id,all_params))\n        return paras_only_bn, paras_wo_bn\n\n\n    def configure_optimizers(self):\n        with_bn,without_bn = self.separate_bn()\n        optimizer = torch.optim.AdamW([{'params': filter(lambda p: p.requires_grad, with_bn)},\n                               {'params': filter(lambda p: p.requires_grad, without_bn),\n                                'weight_decay': 0.01}],\n                                 lr = self.hparams.lr)\n\n        lr_scheduler = {'scheduler': torch.optim.lr_scheduler.CosineAnnealingWarmRestarts(optimizer,T_0=10,eta_min=1e-6),\n                        'name': 'step'\n                       }\n        \n        return [optimizer],[lr_scheduler]\n\n        \n    def on_train_end(self):\n        update_bn(self.train_dataloader(),self.swa_model,device=self.device)\n        ","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"#load the trained model\nlit_model = LitModel.load_from_checkpoint('/kaggle/input/exp1-models/epoch9-v0_best_val_loss.ckpt',\n                                         classify=True,n_cls=5,pretrained=False)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"#swa_model things\nswa_model = SWAResnet(lit_model,lr=7.5e-08)\n\n#test_model for correctness\n# swa_model.eval()\n# x = torch.randn(3,3,224,224,device='cuda')\n# logits = swa_model(x)\n# logits.shape","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"swa_trainer = pl.Trainer(gpus = 1,min_epochs = 1,max_epochs = 5,precision=16,\n                    callbacks = callbacks,\n                    progress_bar_refresh_rate = 10,\n                    num_sanity_val_steps = -1,\n                    check_val_every_n_epoch = 1,\n                    auto_scale_batch_size = False,\n                     auto_lr_find = False)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"swa_trainer.fit(swa_model)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"print(f'[INFO] Best Val acc model - {checkpoint_callback_val_acc.best_k_models}')\nprint(f'[INFO] Best Val loss model - {checkpoint_callback_val_loss.best_k_models}')\n\nprint(f'[INFO] Best train acc model - {checkpoint_callback_train_acc.best_k_models}')\nprint(f'[INFO] Best train loss model - {checkpoint_callback_train_loss.best_k_models}')","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"!cp lightning_logs/version_0/checkpoints/*.ckpt /kaggle/working/","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Submissions"},{"metadata":{"trusted":true},"cell_type":"code","source":"#inference for test images\nsample_submission_df = pd.read_csv(\"/kaggle/input/cassava-leaf-disease-classification/sample_submission.csv\")\npath = \"/kaggle/input/cassava-leaf-disease-classification/test_images/\"\n\nlit_model.eval()\npredictions = []\nwith torch.no_grad():\n    for img_id in sample_submission_df.image_id:\n        img_path = path + img_id\n        image = Image.open(img_path).convert('RGB')\n        \n        image = transforms.Resize((224,224))(image)\n        image = transforms.ToTensor()(image)\n        image = transforms.Normalize(mean=[0.485,0.456,0.406],\n                                    std = [0.229, 0.224, 0.225])(image)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"torch.argmax(lit_model(image.unsqueeze(0).cuda()))","execution_count":null,"outputs":[]}],"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":4,"nbformat_minor":4}