{"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":"# 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\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        pass\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","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-07-06T07:55:22.640986Z","iopub.execute_input":"2022-07-06T07:55:22.641560Z","iopub.status.idle":"2022-07-06T07:55:27.016961Z","shell.execute_reply.started":"2022-07-06T07:55:22.641476Z","shell.execute_reply":"2022-07-06T07:55:27.016107Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%pip install torchinfo","metadata":{"execution":{"iopub.status.busy":"2022-07-06T07:55:27.018739Z","iopub.execute_input":"2022-07-06T07:55:27.019129Z","iopub.status.idle":"2022-07-06T07:55:38.074550Z","shell.execute_reply.started":"2022-07-06T07:55:27.019090Z","shell.execute_reply":"2022-07-06T07:55:38.073399Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%pip install timm","metadata":{"execution":{"iopub.status.busy":"2022-07-06T07:55:38.076173Z","iopub.execute_input":"2022-07-06T07:55:38.078622Z","iopub.status.idle":"2022-07-06T07:55:48.762800Z","shell.execute_reply.started":"2022-07-06T07:55:38.078590Z","shell.execute_reply":"2022-07-06T07:55:48.761659Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport pandas as pd\nimport numpy as np\nimport seaborn as sns\nimport matplotlib.pyplot as plt\n\nimport torch\nimport torch.nn as nn\nimport torchvision.models as models\nimport torchvision.datasets as datasets\nfrom torchvision.utils import make_grid\nimport torchvision.transforms as transforms\nfrom torchmetrics import Accuracy, ConfusionMatrix\nfrom torch.utils.data import DataLoader, ConcatDataset, Dataset\nfrom torchvision.transforms import AutoAugment, AutoAugmentPolicy, InterpolationMode\n\nimport pytorch_lightning as pl\nfrom pytorch_lightning import Trainer\nfrom pytorch_lightning.loggers import WandbLogger\nfrom pytorch_lightning.callbacks import ModelCheckpoint, EarlyStopping\n\nfrom PIL import Image \nfrom sklearn.preprocessing import LabelEncoder\nfrom torch.utils.data import WeightedRandomSampler\nfrom sklearn.model_selection import train_test_split\n\nimport timm\nfrom torchinfo import summary","metadata":{"execution":{"iopub.status.busy":"2022-07-06T07:55:48.765983Z","iopub.execute_input":"2022-07-06T07:55:48.766281Z","iopub.status.idle":"2022-07-06T07:55:57.419623Z","shell.execute_reply.started":"2022-07-06T07:55:48.766253Z","shell.execute_reply":"2022-07-06T07:55:57.418787Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import wandb\n\ntry:\n    from kaggle_secrets import UserSecretsClient\n    user_secrets = UserSecretsClient()\n    secret_value_0 = user_secrets.get_secret(\"wandb_api\")\n    wandb.login(key = secret_value_0)\n    anony = None\nexcept:\n    anony = \"must\"\n    print('If you want to use your W&B account, \\\n          go to Add-ons -> Secrets and provide your W&B access token. \\\n          Use the Label name as wandb_api. \\nGet your W&B access token from here: https://wandb.ai/authorize')","metadata":{"execution":{"iopub.status.busy":"2022-07-06T07:55:57.421021Z","iopub.execute_input":"2022-07-06T07:55:57.421919Z","iopub.status.idle":"2022-07-06T07:55:59.323239Z","shell.execute_reply.started":"2022-07-06T07:55:57.421881Z","shell.execute_reply":"2022-07-06T07:55:59.322337Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Exploratory Data Analysis**","metadata":{}},{"cell_type":"code","source":"df = pd.read_csv('../input/paddy-disease-classification/train.csv')\ndf.head(5)","metadata":{"execution":{"iopub.status.busy":"2022-07-06T07:55:59.324437Z","iopub.execute_input":"2022-07-06T07:55:59.324813Z","iopub.status.idle":"2022-07-06T07:55:59.362243Z","shell.execute_reply.started":"2022-07-06T07:55:59.324776Z","shell.execute_reply":"2022-07-06T07:55:59.361556Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def pie_chart(df, col = 'label', size = (20,7)):\n\n    var = dict(df.groupby(col).size())\n    palette_color = sns.color_palette('bright')\n    fig = plt.figure(figsize = size)\n\n    plt.pie(x = var.values(),\n            labels = var.keys(), colors = palette_color,\n            autopct='%.0f%%')\n\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2022-07-06T07:55:59.363399Z","iopub.execute_input":"2022-07-06T07:55:59.364109Z","iopub.status.idle":"2022-07-06T07:55:59.371225Z","shell.execute_reply.started":"2022-07-06T07:55:59.364070Z","shell.execute_reply":"2022-07-06T07:55:59.370208Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"pie_chart(df)","metadata":{"execution":{"iopub.status.busy":"2022-07-06T07:55:59.372725Z","iopub.execute_input":"2022-07-06T07:55:59.373342Z","iopub.status.idle":"2022-07-06T07:55:59.581303Z","shell.execute_reply.started":"2022-07-06T07:55:59.373307Z","shell.execute_reply":"2022-07-06T07:55:59.580480Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fig = plt.figure(figsize = (26,7))\nsns.countplot(data = df, x = 'variety')","metadata":{"execution":{"iopub.status.busy":"2022-07-06T07:55:59.585359Z","iopub.execute_input":"2022-07-06T07:55:59.586048Z","iopub.status.idle":"2022-07-06T07:55:59.840430Z","shell.execute_reply.started":"2022-07-06T07:55:59.586005Z","shell.execute_reply":"2022-07-06T07:55:59.839665Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fig = plt.figure(figsize = (26,7))\nsns.countplot(data = df, x = 'age')","metadata":{"execution":{"iopub.status.busy":"2022-07-06T07:55:59.842717Z","iopub.execute_input":"2022-07-06T07:55:59.843434Z","iopub.status.idle":"2022-07-06T07:56:00.094234Z","shell.execute_reply.started":"2022-07-06T07:55:59.843394Z","shell.execute_reply":"2022-07-06T07:56:00.093479Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sns.set(rc={\"figure.figsize\":(8, 4)})\nsns.displot(df, x=\"age\", hue=\"label\", kind=\"kde\", multiple=\"stack\", height=7,  aspect=2)","metadata":{"execution":{"iopub.status.busy":"2022-07-06T07:56:00.095485Z","iopub.execute_input":"2022-07-06T07:56:00.096012Z","iopub.status.idle":"2022-07-06T07:56:00.816875Z","shell.execute_reply.started":"2022-07-06T07:56:00.095973Z","shell.execute_reply":"2022-07-06T07:56:00.816111Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"variety = list(df['variety'].unique())\nvariety_dict = {label: idx for idx, label in enumerate(df['variety'].unique())}\ndf['variety'] = df['variety'].map(variety_dict)\n\nsns.set(rc={\"figure.figsize\":(8, 4)}) \nfig = sns.displot(df, x=\"variety\", hue=\"label\", kind=\"kde\", multiple=\"stack\", height=7,  aspect=2)\nplt.xticks(range(0,len(variety)), variety)\nplt.show(fig)","metadata":{"execution":{"iopub.status.busy":"2022-07-06T07:56:00.817845Z","iopub.execute_input":"2022-07-06T07:56:00.818278Z","iopub.status.idle":"2022-07-06T07:56:01.745491Z","shell.execute_reply.started":"2022-07-06T07:56:00.818244Z","shell.execute_reply":"2022-07-06T07:56:01.744760Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Parameters and Hyperparameters","metadata":{}},{"cell_type":"code","source":"class CONFIG:\n    df_path = '../input/paddy-disease-classification/train.csv'\n    train_path = '../input/paddy-disease-classification/train_images'\n    test_path = '../input/paddy-disease-classification/test_images'\n    model_path = 'checkpoint/{epoch:02d}-{val_loss:.4f}-{val_acc:.4f}'\n    split = [0.8, 0.2]\n    batch_size = 16\n    weight_decay = 1e-4\n    learning_rate = 1e-4\n    lr_patience = 2\n    stop_patience = 3\n    layers = 2\n    classes = 10\n    gpus = (1 if torch.cuda.is_available() else 0)\n    num_epochs = 30\n    accum = 64\n    num_workers = 2\n    shuffle = False","metadata":{"execution":{"iopub.status.busy":"2022-07-06T07:56:15.272859Z","iopub.execute_input":"2022-07-06T07:56:15.273541Z","iopub.status.idle":"2022-07-06T07:56:15.278914Z","shell.execute_reply.started":"2022-07-06T07:56:15.273504Z","shell.execute_reply":"2022-07-06T07:56:15.278171Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = pd.read_csv(CONFIG.df_path)","metadata":{"execution":{"iopub.status.busy":"2022-07-06T07:56:16.144400Z","iopub.execute_input":"2022-07-06T07:56:16.145321Z","iopub.status.idle":"2022-07-06T07:56:16.160744Z","shell.execute_reply.started":"2022-07-06T07:56:16.145268Z","shell.execute_reply":"2022-07-06T07:56:16.159921Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Label Encoding","metadata":{}},{"cell_type":"code","source":"labelencoder = LabelEncoder()\ndf['label_name'] = df['label']\ndf['label'] = labelencoder.fit_transform(df['label'])\nprint('classes encoded:', labelencoder.classes_)\ndisplay(df.head(5))","metadata":{"execution":{"iopub.status.busy":"2022-07-06T07:56:18.798124Z","iopub.execute_input":"2022-07-06T07:56:18.799043Z","iopub.status.idle":"2022-07-06T07:56:18.815838Z","shell.execute_reply.started":"2022-07-06T07:56:18.798993Z","shell.execute_reply":"2022-07-06T07:56:18.814890Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Custom Dataset","metadata":{}},{"cell_type":"code","source":"class PaddyImageDataset(Dataset):\n    \n    def __init__(self, \n                 df = None, \n                 root_dir = None, \n                 transform = None):\n\n        self.df = df\n        self.root_dir = root_dir\n        self.transform = transform\n        \n        if df is not None:\n            self.imgs = self.df[\"image_id\"]\n            self.labels = self.df[\"label\"]\n\n        \n    def __len__(self):\n        \n        if self.df is not None:\n            return len(self.df)\n        else:\n            return len(os.listdir(self.root_dir))\n    \n    def __getitem__(self, index):\n\n        if self.df is not None:\n            img_id = self.imgs[index]\n            img_label = self.labels[index]\n            img_label_name = self.df['label_name'][index]\n            \n            img_path = os.path.join(\n                self.root_dir, img_label_name, img_id)\n            \n            with Image.open(img_path).convert('RGB') as image:               \n                if self.transform is not None:\n                    img = self.transform(image)\n                    \n            return img, torch.tensor(int(img_label))\n\n        else:\n            \n            img_id = os.listdir(self.root_dir)[index]\n            img_path = os.path.join(self.root_dir, img_id)\n            \n            with Image.open(\n                os.path.join(img_path)).convert('RGB') as image:               \n                if self.transform is not None:\n                    img = self.transform(image)\n                            \n            return img","metadata":{"execution":{"iopub.status.busy":"2022-07-06T07:56:20.240752Z","iopub.execute_input":"2022-07-06T07:56:20.241133Z","iopub.status.idle":"2022-07-06T07:56:20.252429Z","shell.execute_reply.started":"2022-07-06T07:56:20.241100Z","shell.execute_reply":"2022-07-06T07:56:20.251467Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# WeightedRandomSamplers","metadata":{}},{"cell_type":"code","source":"def getWeightedRandomSampler(df):\n    \n    counts = np.bincount(df['label'])\n    label_weights = 1. / counts\n    weights = label_weights[df['label']]\n    \n    return WeightedRandomSampler(\n        weights, num_samples = len(weights), replacement = True)","metadata":{"execution":{"iopub.status.busy":"2022-07-06T07:56:21.664040Z","iopub.execute_input":"2022-07-06T07:56:21.664399Z","iopub.status.idle":"2022-07-06T07:56:21.669796Z","shell.execute_reply.started":"2022-07-06T07:56:21.664369Z","shell.execute_reply":"2022-07-06T07:56:21.669013Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Data Loaders","metadata":{}},{"cell_type":"code","source":"def get_loader(dataset, batch_size, sampler = None):\n    \n    loader = DataLoader(\n            dataset, batch_size = batch_size)\n    \n    if sampler is not None:\n        loader.__dict__['sampler'] = sampler\n    \n    return loader","metadata":{"execution":{"iopub.status.busy":"2022-07-06T07:56:23.026222Z","iopub.execute_input":"2022-07-06T07:56:23.027071Z","iopub.status.idle":"2022-07-06T07:56:23.032995Z","shell.execute_reply.started":"2022-07-06T07:56:23.027032Z","shell.execute_reply":"2022-07-06T07:56:23.032209Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_train, df_val = train_test_split(\n    df, stratify = df[['label', 'variety']], \n    test_size = (1.0 - CONFIG.split[0]), \n    random_state = 42)\n\ndf_train.reset_index(inplace = True)\ndf_val.reset_index(inplace = True)","metadata":{"execution":{"iopub.status.busy":"2022-07-06T07:56:23.692973Z","iopub.execute_input":"2022-07-06T07:56:23.693587Z","iopub.status.idle":"2022-07-06T07:56:23.768608Z","shell.execute_reply.started":"2022-07-06T07:56:23.693550Z","shell.execute_reply":"2022-07-06T07:56:23.767761Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"display(df_train.head(5))\ndf_val.head(5)","metadata":{"execution":{"iopub.status.busy":"2022-07-06T07:56:24.306336Z","iopub.execute_input":"2022-07-06T07:56:24.306772Z","iopub.status.idle":"2022-07-06T07:56:24.327594Z","shell.execute_reply.started":"2022-07-06T07:56:24.306733Z","shell.execute_reply":"2022-07-06T07:56:24.326331Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_sampler = getWeightedRandomSampler(df_train)\nval_sampler = getWeightedRandomSampler(df_val)","metadata":{"execution":{"iopub.status.busy":"2022-07-06T07:56:25.164154Z","iopub.execute_input":"2022-07-06T07:56:25.164917Z","iopub.status.idle":"2022-07-06T07:56:25.171369Z","shell.execute_reply.started":"2022-07-06T07:56:25.164878Z","shell.execute_reply":"2022-07-06T07:56:25.170470Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"transform = transforms.Compose([\n    transforms.Resize((224, 224)),\n    transforms.CenterCrop(224),\n    transforms.ToTensor(),\n    transforms.Normalize(\n        mean = [0.485, 0.456, 0.406],\n        std = [0.229, 0.224, 0.225])])","metadata":{"execution":{"iopub.status.busy":"2022-07-06T07:56:25.753730Z","iopub.execute_input":"2022-07-06T07:56:25.754114Z","iopub.status.idle":"2022-07-06T07:56:25.759047Z","shell.execute_reply.started":"2022-07-06T07:56:25.754074Z","shell.execute_reply":"2022-07-06T07:56:25.758282Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# AutoAugmentation","metadata":{}},{"cell_type":"code","source":"autoaugment = transforms.Compose([\n    transforms.Resize((224, 224)),\n    transforms.CenterCrop(224),\n    transforms.AutoAugment(\n        policy = AutoAugmentPolicy.IMAGENET, \n        interpolation = InterpolationMode.BILINEAR,),\n    transforms.ToTensor(),\n    transforms.Normalize(\n        mean = [0.485, 0.456, 0.406],\n        std = [0.229, 0.224, 0.225])])","metadata":{"execution":{"iopub.status.busy":"2022-07-06T07:56:28.136520Z","iopub.execute_input":"2022-07-06T07:56:28.138273Z","iopub.status.idle":"2022-07-06T07:56:28.144919Z","shell.execute_reply.started":"2022-07-06T07:56:28.138225Z","shell.execute_reply":"2022-07-06T07:56:28.143848Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_dataset = PaddyImageDataset(\n    root_dir = CONFIG.train_path, \n    df = df_train, \n    transform = autoaugment)\n\nprint('Length of train dataset:', len(train_dataset))\n\nval_dataset = PaddyImageDataset(\n    root_dir = CONFIG.train_path, \n    df = df_val, \n    transform = autoaugment)\n\nprint('Length of val dataset:', len(val_dataset))\n\npredict_dataset = PaddyImageDataset(\n    root_dir = CONFIG.test_path, \n    transform = transform)\n\nprint('Length of predict dataset:', len(predict_dataset))","metadata":{"execution":{"iopub.status.busy":"2022-07-06T07:56:28.827400Z","iopub.execute_input":"2022-07-06T07:56:28.827757Z","iopub.status.idle":"2022-07-06T07:56:28.837531Z","shell.execute_reply.started":"2022-07-06T07:56:28.827728Z","shell.execute_reply":"2022-07-06T07:56:28.836554Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_dataloader = get_loader(\n    dataset = train_dataset, \n    batch_size = CONFIG.batch_size,\n    sampler = train_sampler)\n\nval_dataloader = get_loader(\n    dataset = val_dataset, \n    batch_size = CONFIG.batch_size,\n    sampler = val_sampler)\n\npredict_dataloader = get_loader(\n    dataset = predict_dataset, \n    batch_size = CONFIG.batch_size,\n    sampler = None)","metadata":{"execution":{"iopub.status.busy":"2022-07-06T07:56:29.443859Z","iopub.execute_input":"2022-07-06T07:56:29.444548Z","iopub.status.idle":"2022-07-06T07:56:29.450206Z","shell.execute_reply.started":"2022-07-06T07:56:29.444506Z","shell.execute_reply":"2022-07-06T07:56:29.449138Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Data Visualization","metadata":{}},{"cell_type":"code","source":"batch = next(iter(predict_dataloader))\n\ngrid_img = make_grid(\n    batch, \n    nrow = 4,\n    normalize = True)\n\nplt.figure(figsize = (15,30))\nplt.imshow(grid_img.permute(1, 2, 0))","metadata":{"execution":{"iopub.status.busy":"2022-07-06T07:56:31.068348Z","iopub.execute_input":"2022-07-06T07:56:31.068719Z","iopub.status.idle":"2022-07-06T07:56:32.143429Z","shell.execute_reply.started":"2022-07-06T07:56:31.068687Z","shell.execute_reply":"2022-07-06T07:56:32.142456Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"batch = next(iter(train_dataloader))\n\ngrid_img = make_grid(\n    batch[0], \n    nrow = 4,\n    normalize = True)\n\nplt.figure(figsize = (15,30))\nplt.imshow(grid_img.permute(1, 2, 0))","metadata":{"execution":{"iopub.status.busy":"2022-07-06T07:56:32.145010Z","iopub.execute_input":"2022-07-06T07:56:32.145502Z","iopub.status.idle":"2022-07-06T07:56:33.206789Z","shell.execute_reply.started":"2022-07-06T07:56:32.145465Z","shell.execute_reply":"2022-07-06T07:56:33.206066Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Lightning Model","metadata":{}},{"cell_type":"code","source":"class PaddyDiseaseClassification(pl.LightningModule):\n\n\n    def __init__(self,\n                learning_rate,\n                weight_decay,\n                layers,\n                patience,\n                pretrained = True,\n                classes = 10):\n\n        super(PaddyDiseaseClassification, self).__init__()\n\n\n        self.learning_rate = learning_rate\n        self.weight_decay = weight_decay\n        self.layers = layers\n        self.patience = patience\n\n        self.pretrained = pretrained\n        self.classes = classes\n        \n        self.model = timm.create_model(\n            \"swin_base_patch4_window7_224_in22k\", \n            pretrained = pretrained)\n        \n        self.model.head = nn.Linear(\n            self.model.head.in_features, classes)\n\n        self.save_hyperparameters()\n\n        self.loss = nn.CrossEntropyLoss()\n\n        self.train_acc = Accuracy(num_classes = classes)\n        self.val_acc = Accuracy(num_classes = classes)\n        \n        self.finetune()\n        \n    def finetune(self):\n\n        for param in list(self.model.children())[:-1]:\n            for p in param.parameters():\n                p.requires_grad = False\n    \n        for param in list(self.model.layers._modules['3'].blocks):\n            for p in param.parameters():\n                p.requires_grad = True\n\n        if self.layers is not 0:\n            for param in list(self.model.layers._modules['2'].blocks)[-self.layers:]:\n                for p in param.parameters():\n                    p.requires_grad = True\n\n\n    def configure_optimizers(self):\n\n        optimizer = torch.optim.Adam(\n            self.model.parameters(), \n            lr = self.learning_rate, \n            weight_decay = self.weight_decay)\n        \n        scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(\n            optimizer, \n            mode = 'max', \n            factor = 0.2, \n            patience = self.patience, \n            verbose = True)\n\n        return {\n           'optimizer': optimizer,\n           'lr_scheduler': scheduler,\n           'monitor': 'val_acc' \n       }\n\n    \n    def forward(self, images):\n        \n        output = self.model(images)\n        return output\n        \n\n    def training_step(self, batch, batch_idx):\n\n        images, labels = batch\n\n        output = self.model(images)\n        train_loss = self.loss(output, labels)\n\n        self.log(\n            name = 'train_loss',\n            value = train_loss,\n            on_step = False,\n            on_epoch = True,\n            prog_bar = True)\n\n        train_acc = self.train_acc(output, labels)\n        \n        self.log(\n            name = \"train_acc\",\n            value = train_acc,\n            on_step = False,\n            on_epoch = True,\n            prog_bar = True)\n\n        return train_loss\n\n\n    def validation_step(self, batch, batch_idx):\n\n        images, labels = batch\n        \n        output = self.model(images)\n        val_loss = self.loss(output, labels)\n\n        val_acc = self.val_acc(output, labels)\n\n        return {'val_loss': val_loss,\n                'val_acc':val_acc,\n                'log': {'val_loss': val_loss}}\n\n    def validation_epoch_end(self, outputs):\n\n        loss = torch.stack([o['val_loss'] for o in outputs], 0).mean()\n        acc = torch.stack([o['val_acc'] for o in outputs], 0).mean()\n\n        out = {'val_loss': loss,\n               'val_acc': acc}\n\n        self.log(\n            name = 'val_loss',\n            value = loss,\n            on_epoch = True,\n            prog_bar = True)\n\n        self.log(\n            name = 'val_acc',\n            value = acc,\n            on_epoch = True,\n            prog_bar = True)\n\n        return {**out, 'log': out}\n    \n    ","metadata":{"execution":{"iopub.status.busy":"2022-07-06T07:56:33.334658Z","iopub.execute_input":"2022-07-06T07:56:33.335294Z","iopub.status.idle":"2022-07-06T07:56:33.356338Z","shell.execute_reply.started":"2022-07-06T07:56:33.335248Z","shell.execute_reply":"2022-07-06T07:56:33.355624Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model Training","metadata":{}},{"cell_type":"code","source":"model = PaddyDiseaseClassification(\n    learning_rate = CONFIG.learning_rate,\n    weight_decay = CONFIG.weight_decay,\n    layers = CONFIG.layers,\n    patience = CONFIG.lr_patience,\n    classes = CONFIG.classes)","metadata":{"execution":{"iopub.status.busy":"2022-07-06T07:56:34.643982Z","iopub.execute_input":"2022-07-06T07:56:34.644552Z","iopub.status.idle":"2022-07-06T07:57:13.157291Z","shell.execute_reply.started":"2022-07-06T07:56:34.644516Z","shell.execute_reply":"2022-07-06T07:57:13.156308Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"summary(\n    model, \n    input_size = (16, 3, 224, 224))","metadata":{"execution":{"iopub.status.busy":"2022-07-06T07:57:13.160497Z","iopub.execute_input":"2022-07-06T07:57:13.160891Z","iopub.status.idle":"2022-07-06T07:57:24.240353Z","shell.execute_reply.started":"2022-07-06T07:57:13.160859Z","shell.execute_reply":"2022-07-06T07:57:24.239550Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"early_stop_callback = EarlyStopping(\n                        monitor = \"val_acc\",\n                        min_delta = 0.00,\n                        patience = CONFIG.stop_patience,\n                        verbose = False,\n                        mode = \"max\")\n\nmodel_checkpoint = ModelCheckpoint(\n                        monitor = 'val_acc',\n                        save_top_k = 1,\n                        save_weights_only = True,\n                        filename = CONFIG.model_path,\n                        verbose = False,\n                        mode = 'max')\n\nlogger = WandbLogger(project = \"swin-vit-model\")","metadata":{"execution":{"iopub.status.busy":"2022-07-06T07:57:24.241680Z","iopub.execute_input":"2022-07-06T07:57:24.242204Z","iopub.status.idle":"2022-07-06T07:57:29.092844Z","shell.execute_reply.started":"2022-07-06T07:57:24.242166Z","shell.execute_reply":"2022-07-06T07:57:29.091705Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"trainer = Trainer(\n    auto_lr_find = True,\n    gpus = CONFIG.gpus,\n    max_epochs = CONFIG.num_epochs,\n    accumulate_grad_batches = CONFIG.accum,\n    logger = logger,\n    callbacks = [model_checkpoint, \n                 early_stop_callback])","metadata":{"execution":{"iopub.status.busy":"2022-07-06T07:57:29.099064Z","iopub.execute_input":"2022-07-06T07:57:29.101296Z","iopub.status.idle":"2022-07-06T07:57:29.117410Z","shell.execute_reply.started":"2022-07-06T07:57:29.101250Z","shell.execute_reply":"2022-07-06T07:57:29.116171Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"trainer.fit(\n    model = model, \n    train_dataloaders = train_dataloader, \n    val_dataloaders = val_dataloader)","metadata":{"execution":{"iopub.status.busy":"2022-07-06T07:57:29.122462Z","iopub.execute_input":"2022-07-06T07:57:29.124938Z","iopub.status.idle":"2022-07-06T09:15:58.832760Z","shell.execute_reply.started":"2022-07-06T07:57:29.124901Z","shell.execute_reply":"2022-07-06T09:15:58.831937Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model Inference","metadata":{}},{"cell_type":"code","source":"predictions = trainer.predict(\n    model = model, dataloaders = predict_dataloader)","metadata":{"execution":{"iopub.status.busy":"2022-07-06T09:26:06.529697Z","iopub.execute_input":"2022-07-06T09:26:06.530084Z","iopub.status.idle":"2022-07-06T09:27:23.702166Z","shell.execute_reply.started":"2022-07-06T09:26:06.530052Z","shell.execute_reply":"2022-07-06T09:27:23.701325Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"flag = True\nfor tensor in predictions:\n    if flag == True:\n        stack = tensor\n        flag = False\n    else:\n        stack = torch.concat(\n            [stack,tensor], axis = 0)","metadata":{"execution":{"iopub.status.busy":"2022-07-06T09:27:26.092674Z","iopub.execute_input":"2022-07-06T09:27:26.093105Z","iopub.status.idle":"2022-07-06T09:27:26.103727Z","shell.execute_reply.started":"2022-07-06T09:27:26.093069Z","shell.execute_reply":"2022-07-06T09:27:26.102970Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission = pd.DataFrame()\nsubmission['image_id'] = pd.DataFrame(os.listdir(CONFIG.test_path))\nsubmission['label'] = pd.DataFrame(np.argmax(stack, axis = 1).numpy().astype('int'))","metadata":{"execution":{"iopub.status.busy":"2022-07-06T09:27:28.609647Z","iopub.execute_input":"2022-07-06T09:27:28.610287Z","iopub.status.idle":"2022-07-06T09:27:28.638784Z","shell.execute_reply.started":"2022-07-06T09:27:28.610249Z","shell.execute_reply":"2022-07-06T09:27:28.638054Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission['label'] = labelencoder.inverse_transform(submission['label'])\nsubmission.head(10)","metadata":{"execution":{"iopub.status.busy":"2022-07-06T09:27:30.105542Z","iopub.execute_input":"2022-07-06T09:27:30.106362Z","iopub.status.idle":"2022-07-06T09:27:30.144299Z","shell.execute_reply.started":"2022-07-06T09:27:30.106281Z","shell.execute_reply":"2022-07-06T09:27:30.143132Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission.to_csv('submission.csv', index = False)","metadata":{"execution":{"iopub.status.busy":"2022-07-06T09:27:40.184043Z","iopub.execute_input":"2022-07-06T09:27:40.184700Z","iopub.status.idle":"2022-07-06T09:27:40.203424Z","shell.execute_reply.started":"2022-07-06T09:27:40.184660Z","shell.execute_reply":"2022-07-06T09:27:40.202698Z"},"trusted":true},"execution_count":null,"outputs":[]}]}