{"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":"!pip install timm","metadata":{"execution":{"iopub.status.busy":"2022-05-03T07:09:40.078095Z","iopub.execute_input":"2022-05-03T07:09:40.078467Z","iopub.status.idle":"2022-05-03T07:09:51.26008Z","shell.execute_reply.started":"2022-05-03T07:09:40.078384Z","shell.execute_reply":"2022-05-03T07:09:51.259092Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\nimport time\nimport os\nimport copy\nimport json\n\n# visualization modules\nimport seaborn as sns\nimport matplotlib.pyplot as plt\nfrom PIL import Image\n\n# pytorch modules\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.nn import functional as F\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import models\nimport torchvision.transforms as transforms\nfrom sklearn.metrics import f1_score\n\n\n# augmentation\nimport albumentations\nfrom albumentations.pytorch.transforms import ToTensorV2\n\nimport warnings\nwarnings.filterwarnings('ignore')\n%matplotlib inline\n\nimport sys\nfrom tqdm import tqdm\nimport time\nimport copy\n\nimport timm\nfrom timm.loss import LabelSmoothingCrossEntropy","metadata":{"execution":{"iopub.status.busy":"2022-05-03T07:09:51.263666Z","iopub.execute_input":"2022-05-03T07:09:51.2639Z","iopub.status.idle":"2022-05-03T07:09:59.780171Z","shell.execute_reply.started":"2022-05-03T07:09:51.263871Z","shell.execute_reply":"2022-05-03T07:09:59.77934Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"BASE_DIR = \"../input/cassava-leaf-disease-classification/\"\n\ntrain = pd.read_csv(BASE_DIR+'train.csv')\ntrain.head()","metadata":{"execution":{"iopub.status.busy":"2022-05-03T07:09:59.781779Z","iopub.execute_input":"2022-05-03T07:09:59.782043Z","iopub.status.idle":"2022-05-03T07:09:59.823898Z","shell.execute_reply.started":"2022-05-03T07:09:59.782009Z","shell.execute_reply":"2022-05-03T07:09:59.823184Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# loading mapping for target label\nwith open(BASE_DIR+'label_num_to_disease_map.json') as f:\n    mapping = json.loads(f.read())\n    mapping = {int(k): v for k, v in mapping.items()}\nmapping","metadata":{"execution":{"iopub.status.busy":"2022-05-03T07:09:59.826181Z","iopub.execute_input":"2022-05-03T07:09:59.826439Z","iopub.status.idle":"2022-05-03T07:09:59.837365Z","shell.execute_reply.started":"2022-05-03T07:09:59.826407Z","shell.execute_reply":"2022-05-03T07:09:59.836672Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train['label_name'] = train['label'].map(mapping)\ntrain.head()","metadata":{"execution":{"iopub.status.busy":"2022-05-03T07:09:59.83865Z","iopub.execute_input":"2022-05-03T07:09:59.838913Z","iopub.status.idle":"2022-05-03T07:09:59.856551Z","shell.execute_reply.started":"2022-05-03T07:09:59.838878Z","shell.execute_reply":"2022-05-03T07:09:59.855687Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def plot_images(class_id, label, total_images=6):\n    plot_list = train[train['label'] == class_id].sample(total_images)[\"image_id\"].tolist()\n    \n    # labels = [label] * total_images\n    size = int(np.sqrt(total_images))\n    if size*size < total_images:\n        size += 1\n    \n    plt.figure(figsize=(15,15))\n    \n    for i in range(total_images):\n        plt.subplot(size, size, i+1)\n        image = Image.open(str(BASE_DIR + \"train_images/\" + plot_list[i]))\n        plt.imshow(image)\n        plt.title(label, fontsize=14)\n        plt.axis(\"off\")\n        \n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2022-05-03T07:09:59.858373Z","iopub.execute_input":"2022-05-03T07:09:59.858642Z","iopub.status.idle":"2022-05-03T07:09:59.86735Z","shell.execute_reply.started":"2022-05-03T07:09:59.858586Z","shell.execute_reply":"2022-05-03T07:09:59.866427Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plot_images(0, mapping[0])","metadata":{"execution":{"iopub.status.busy":"2022-05-03T07:09:59.868997Z","iopub.execute_input":"2022-05-03T07:09:59.869282Z","iopub.status.idle":"2022-05-03T07:10:01.061268Z","shell.execute_reply.started":"2022-05-03T07:09:59.869214Z","shell.execute_reply":"2022-05-03T07:10:01.060328Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plot_images(1, mapping[1])","metadata":{"execution":{"iopub.status.busy":"2022-05-03T07:10:01.066736Z","iopub.execute_input":"2022-05-03T07:10:01.067133Z","iopub.status.idle":"2022-05-03T07:10:01.97395Z","shell.execute_reply.started":"2022-05-03T07:10:01.067088Z","shell.execute_reply":"2022-05-03T07:10:01.973201Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plot_images(2, mapping[2])","metadata":{"execution":{"iopub.status.busy":"2022-05-03T07:10:01.974938Z","iopub.execute_input":"2022-05-03T07:10:01.975265Z","iopub.status.idle":"2022-05-03T07:10:02.743958Z","shell.execute_reply.started":"2022-05-03T07:10:01.975235Z","shell.execute_reply":"2022-05-03T07:10:02.743288Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plot_images(3, mapping[3])","metadata":{"execution":{"iopub.status.busy":"2022-05-03T07:10:02.746699Z","iopub.execute_input":"2022-05-03T07:10:02.7474Z","iopub.status.idle":"2022-05-03T07:10:03.530447Z","shell.execute_reply.started":"2022-05-03T07:10:02.747363Z","shell.execute_reply":"2022-05-03T07:10:03.52979Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plot_images(4, mapping[4])","metadata":{"execution":{"iopub.status.busy":"2022-05-03T07:10:03.531606Z","iopub.execute_input":"2022-05-03T07:10:03.531956Z","iopub.status.idle":"2022-05-03T07:10:04.56971Z","shell.execute_reply.started":"2022-05-03T07:10:03.531924Z","shell.execute_reply":"2022-05-03T07:10:04.565587Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sns.countplot(x=train[\"label_name\"])\nplt.xticks(rotation=90)","metadata":{"execution":{"iopub.status.busy":"2022-05-03T07:10:04.571107Z","iopub.execute_input":"2022-05-03T07:10:04.57156Z","iopub.status.idle":"2022-05-03T07:10:04.804905Z","shell.execute_reply.started":"2022-05-03T07:10:04.571523Z","shell.execute_reply":"2022-05-03T07:10:04.804251Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"DIM = (224, 224)\nWIDTH, HEIGHT = DIM\nNUM_CLASSES = 5\nNUM_WORKERS = 24\nTRAIN_BATCH_SIZE = 32\nTEST_BATCH_SIZE = 32\nSEED = 1\n\nDEVICE = 'cuda'\n\nMEAN = [0.485, 0.456, 0.406]\nSTD = [0.229, 0.224, 0.225] ","metadata":{"execution":{"iopub.status.busy":"2022-05-03T07:10:04.806156Z","iopub.execute_input":"2022-05-03T07:10:04.806564Z","iopub.status.idle":"2022-05-03T07:10:04.812801Z","shell.execute_reply.started":"2022-05-03T07:10:04.806527Z","shell.execute_reply":"2022-05-03T07:10:04.812042Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_test_transforms(value='val'):\n    if value==\"train\":\n        return albumentations.Compose([\n            albumentations.Resize(WIDTH, HEIGHT),\n            albumentations.HorizontalFlip(p=0.5),\n            albumentations.Rotate(limit=(-90, 90)),\n            albumentations.VerticalFlip(p=0.5),\n            albumentations.Normalize(MEAN, STD, max_pixel_value=255.0, always_apply=True),\n            ToTensorV2(p=1.0)\n        ])\n    elif value==\"val\":\n        return albumentations.Compose([\n            albumentations.Resize(WIDTH, HEIGHT),\n            albumentations.Normalize(MEAN, STD, max_pixel_value=255.0, always_apply=True),\n            ToTensorV2(p=1.0)\n        ])","metadata":{"execution":{"iopub.status.busy":"2022-05-03T07:10:04.814189Z","iopub.execute_input":"2022-05-03T07:10:04.814471Z","iopub.status.idle":"2022-05-03T07:10:04.825399Z","shell.execute_reply.started":"2022-05-03T07:10:04.814436Z","shell.execute_reply":"2022-05-03T07:10:04.824639Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CassavaDataset(Dataset):\n    def __init__(self, image_ids, labels, dimension=None, augmentations=None, folder='train_images'):\n        super().__init__()\n        self.image_ids = image_ids\n        self.labels = labels\n        self.dim = dimension\n        self.augmentations = augmentations\n        self.folder = folder\n    \n    def __len__(self):\n        return len(self.image_ids)\n    \n    def __getitem__(self, idx):\n        img = Image.open(os.path.join(BASE_DIR, self.folder, self.image_ids[idx]))\n        \n        if self.dim:\n            img = img.resize(self.dim)\n        img = np.array(img)\n        if self.augmentations:\n            augmented = self.augmentations(image = img)\n            img = augmented['image']\n        \n        label = torch.tensor(self.labels[idx], dtype=torch.long)\n        return img, label","metadata":{"execution":{"iopub.status.busy":"2022-05-03T07:10:04.826649Z","iopub.execute_input":"2022-05-03T07:10:04.82691Z","iopub.status.idle":"2022-05-03T07:10:04.836464Z","shell.execute_reply.started":"2022-05-03T07:10:04.826878Z","shell.execute_reply":"2022-05-03T07:10:04.835786Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from sklearn.model_selection import train_test_split\nX_train, X_test, y_train, y_test = train_test_split(train['image_id'], train['label'], test_size=0.25, random_state=42)","metadata":{"execution":{"iopub.status.busy":"2022-05-03T07:10:04.837794Z","iopub.execute_input":"2022-05-03T07:10:04.838365Z","iopub.status.idle":"2022-05-03T07:10:04.850378Z","shell.execute_reply.started":"2022-05-03T07:10:04.838326Z","shell.execute_reply":"2022-05-03T07:10:04.849195Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# create Dataset for training and validation\ntrain_dataset = CassavaDataset(\n    image_ids = X_train.values,\n    labels = y_train.values,\n    augmentations=get_test_transforms('train'),\n    dimension= DIM\n)\n\ntrain_loader = DataLoader(\n    train_dataset,\n    batch_size = TRAIN_BATCH_SIZE,\n    num_workers = NUM_WORKERS,\n    shuffle = True\n)\n\n\nval_dataset = CassavaDataset(\n    image_ids = X_test.values,\n    labels = y_test.values,\n    augmentations=get_test_transforms('val'),\n    dimension= DIM\n)\n\nval_loader = DataLoader(\n    val_dataset,\n    batch_size = TRAIN_BATCH_SIZE,\n    num_workers = NUM_WORKERS,\n    shuffle = True\n)\n","metadata":{"execution":{"iopub.status.busy":"2022-05-03T07:10:04.852458Z","iopub.execute_input":"2022-05-03T07:10:04.852661Z","iopub.status.idle":"2022-05-03T07:10:04.859806Z","shell.execute_reply.started":"2022-05-03T07:10:04.852638Z","shell.execute_reply":"2022-05-03T07:10:04.859054Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"val_dataset[0]","metadata":{"execution":{"iopub.status.busy":"2022-05-03T07:10:04.8613Z","iopub.execute_input":"2022-05-03T07:10:04.862097Z","iopub.status.idle":"2022-05-03T07:10:04.972261Z","shell.execute_reply.started":"2022-05-03T07:10:04.862058Z","shell.execute_reply":"2022-05-03T07:10:04.971482Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dataloaders = {\n    \"train\": train_loader,\n    \"val\": val_loader\n}\n\ndataset_sizes = {\n    \"train\": len(X_train),\n    \"val\": len(X_test)\n}","metadata":{"execution":{"iopub.status.busy":"2022-05-03T07:10:04.973557Z","iopub.execute_input":"2022-05-03T07:10:04.973865Z","iopub.status.idle":"2022-05-03T07:10:04.978953Z","shell.execute_reply.started":"2022-05-03T07:10:04.973829Z","shell.execute_reply":"2022-05-03T07:10:04.977938Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(len(train_loader), len(val_loader))","metadata":{"execution":{"iopub.status.busy":"2022-05-03T07:10:04.980158Z","iopub.execute_input":"2022-05-03T07:10:04.980564Z","iopub.status.idle":"2022-05-03T07:10:04.989747Z","shell.execute_reply.started":"2022-05-03T07:10:04.980528Z","shell.execute_reply":"2022-05-03T07:10:04.988946Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"classes = list(mapping.values())\nprint(classes, len(classes))","metadata":{"execution":{"iopub.status.busy":"2022-05-03T07:10:04.991115Z","iopub.execute_input":"2022-05-03T07:10:04.991491Z","iopub.status.idle":"2022-05-03T07:10:04.999675Z","shell.execute_reply.started":"2022-05-03T07:10:04.991438Z","shell.execute_reply":"2022-05-03T07:10:04.998818Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# now, for the model\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\ndevice","metadata":{"execution":{"iopub.status.busy":"2022-05-03T07:10:05.001196Z","iopub.execute_input":"2022-05-03T07:10:05.001599Z","iopub.status.idle":"2022-05-03T07:10:05.07851Z","shell.execute_reply.started":"2022-05-03T07:10:05.001541Z","shell.execute_reply":"2022-05-03T07:10:05.077799Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# SWIN","metadata":{}},{"cell_type":"code","source":"HUB_URL = \"SharanSMenon/swin-transformer-hub:main\"\nMODEL_NAME = \"swin_tiny_patch4_window7_224\"\n# check hubconf for more models.\nmodel = torch.hub.load(HUB_URL, MODEL_NAME, pretrained=True) # load from torch hub","metadata":{"execution":{"iopub.status.busy":"2022-05-03T07:10:05.079845Z","iopub.execute_input":"2022-05-03T07:10:05.080126Z","iopub.status.idle":"2022-05-03T07:10:11.394532Z","shell.execute_reply.started":"2022-05-03T07:10:05.080069Z","shell.execute_reply":"2022-05-03T07:10:11.393771Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for param in model.parameters(): #freeze model\n    param.requires_grad = False\n\nn_inputs = model.head.in_features\nmodel.head = nn.Sequential(\n    nn.Linear(n_inputs, 512),\n    nn.ReLU(),\n    nn.Dropout(0.3),\n    nn.Linear(512, len(classes))\n)\nmodel = model.to(device)\nprint(model.head)","metadata":{"execution":{"iopub.status.busy":"2022-05-03T07:10:11.395802Z","iopub.execute_input":"2022-05-03T07:10:11.396124Z","iopub.status.idle":"2022-05-03T07:10:16.11605Z","shell.execute_reply.started":"2022-05-03T07:10:11.396087Z","shell.execute_reply":"2022-05-03T07:10:16.115325Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train_model(model, criterion, optimizer, scheduler, num_epochs=10):\n    since = time.time()\n    best_model_wts = copy.deepcopy(model.state_dict())\n    best_acc = 0.0\n    f1_value = 0.0\n    \n    for epoch in range(num_epochs):\n        print(f'Epoch {epoch+1}/{num_epochs}')\n        print(\"-\"*10)\n        \n        for phase in ['train', 'val']: # We do training and validation phase per epoch\n            if phase == 'train':\n                model.train() # model to training mode\n            else:\n                model.eval() # model to evaluate\n            \n            running_loss = 0.0\n            running_corrects = 0.0\n            \n            for inputs, labels in tqdm(dataloaders[phase]):\n                inputs = inputs.to(device)\n                labels = labels.to(device)\n                \n                optimizer.zero_grad()\n                \n                with torch.set_grad_enabled(phase == 'train'): # no autograd makes validation go faster\n                    outputs = model(inputs)\n                    _, preds = torch.max(outputs, 1) # used for accuracy\n                    loss = criterion(outputs, labels)\n                    \n                    if phase == 'train':\n                        loss.backward()\n                        optimizer.step()\n                running_loss += loss.item() * inputs.size(0)\n                running_corrects += torch.sum(preds == labels.data)\n                \n            if phase == 'train':\n                scheduler.step() # step at end of epoch\n            \n            epoch_loss = running_loss / dataset_sizes[phase]\n            epoch_acc =  running_corrects.double() / dataset_sizes[phase]\n            \n            print(\"{} Loss: {:.4f} Acc: {:.4f}\".format(phase, epoch_loss, epoch_acc))\n            \n            if phase == 'val' and epoch_acc > best_acc:\n                best_acc = epoch_acc\n                f1_value = f1_score(labels.cpu().data, preds.cpu(),average='micro')\n                best_model_wts = copy.deepcopy(model.state_dict()) # keep the best validation accuracy model\n        print()\n    time_elapsed = time.time() - since # slight error\n    print('Training complete in {:.0f}m {:.0f}s'.format(time_elapsed // 60, time_elapsed % 60))\n    print(\"Best Val Acc: {:.4f}\".format(best_acc))\n    print(\"F1 Score: {:.4f}\".format(f1_value))\n    \n    model.load_state_dict(best_model_wts)\n    return model","metadata":{"execution":{"iopub.status.busy":"2022-05-03T07:10:16.117541Z","iopub.execute_input":"2022-05-03T07:10:16.117791Z","iopub.status.idle":"2022-05-03T07:10:16.130439Z","shell.execute_reply.started":"2022-05-03T07:10:16.117756Z","shell.execute_reply":"2022-05-03T07:10:16.129725Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"criterion = LabelSmoothingCrossEntropy()\ncriterion = criterion.to(device)\noptimizer = optim.AdamW(model.head.parameters(), lr=0.001)","metadata":{"execution":{"iopub.status.busy":"2022-05-03T07:10:16.131687Z","iopub.execute_input":"2022-05-03T07:10:16.132092Z","iopub.status.idle":"2022-05-03T07:10:16.144707Z","shell.execute_reply.started":"2022-05-03T07:10:16.132055Z","shell.execute_reply":"2022-05-03T07:10:16.143981Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"exp_lr_scheduler = optim.lr_scheduler.StepLR(optimizer, step_size=3, gamma=0.97)","metadata":{"execution":{"iopub.status.busy":"2022-05-03T07:10:16.145955Z","iopub.execute_input":"2022-05-03T07:10:16.146502Z","iopub.status.idle":"2022-05-03T07:10:16.155173Z","shell.execute_reply.started":"2022-05-03T07:10:16.146465Z","shell.execute_reply":"2022-05-03T07:10:16.154495Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model_ft = train_model(model, criterion, optimizer, exp_lr_scheduler, num_epochs=7) # now it is a lot faster","metadata":{"execution":{"iopub.status.busy":"2022-05-03T07:10:16.15935Z","iopub.execute_input":"2022-05-03T07:10:16.159564Z","iopub.status.idle":"2022-05-03T07:45:20.328844Z","shell.execute_reply.started":"2022-05-03T07:10:16.159536Z","shell.execute_reply":"2022-05-03T07:45:20.327534Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"total_params = sum(p.numel() for p in model_ft.parameters())\nprint(\"Total Parameters :\", total_params)","metadata":{"execution":{"iopub.status.busy":"2022-05-03T09:09:06.006178Z","iopub.execute_input":"2022-05-03T09:09:06.006668Z","iopub.status.idle":"2022-05-03T09:09:06.012231Z","shell.execute_reply.started":"2022-05-03T09:09:06.006634Z","shell.execute_reply":"2022-05-03T09:09:06.011553Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"example = torch.rand(1, 3, 224, 224)\ntraced_script_module = torch.jit.trace(model.cpu(), example)\ntraced_script_module.save(\"./leafDiseaseClassifier.pt\")","metadata":{"execution":{"iopub.status.busy":"2022-05-03T07:45:20.330664Z","iopub.execute_input":"2022-05-03T07:45:20.330938Z","iopub.status.idle":"2022-05-03T07:45:23.321039Z","shell.execute_reply.started":"2022-05-03T07:45:20.330898Z","shell.execute_reply":"2022-05-03T07:45:23.320224Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = torch.jit.load('./leafDiseaseClassifier.pt')\n# model.eval()","metadata":{"execution":{"iopub.status.busy":"2022-05-03T07:45:23.322608Z","iopub.execute_input":"2022-05-03T07:45:23.322875Z","iopub.status.idle":"2022-05-03T07:45:23.531278Z","shell.execute_reply.started":"2022-05-03T07:45:23.322824Z","shell.execute_reply":"2022-05-03T07:45:23.530495Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"resizer = albumentations.Compose([\n            albumentations.Resize(WIDTH, HEIGHT),\n            albumentations.Normalize(MEAN, STD, max_pixel_value=255.0, always_apply=True),\n            ToTensorV2(p=1.0)\n        ])","metadata":{"execution":{"iopub.status.busy":"2022-05-03T07:45:23.532439Z","iopub.execute_input":"2022-05-03T07:45:23.532675Z","iopub.status.idle":"2022-05-03T07:45:23.539895Z","shell.execute_reply.started":"2022-05-03T07:45:23.532643Z","shell.execute_reply":"2022-05-03T07:45:23.539187Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from PIL import Image\ndef predict(model, path, resizer):\n    image = Image.open(path)\n    image = resizer(image=np.array(image))\n    image = image['image']\n    image = torch.unsqueeze(image, axis=0)\n    output = model(image)\n    output = torch.argmax(output)\n    output = mapping[output.item()]\n    return output","metadata":{"execution":{"iopub.status.busy":"2022-05-03T07:45:23.541169Z","iopub.execute_input":"2022-05-03T07:45:23.541519Z","iopub.status.idle":"2022-05-03T07:45:23.548805Z","shell.execute_reply.started":"2022-05-03T07:45:23.541484Z","shell.execute_reply":"2022-05-03T07:45:23.547652Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"predict(model, '../input/cassava-leaf-disease-classification/test_images/2216849948.jpg', resizer)","metadata":{"execution":{"iopub.status.busy":"2022-05-03T07:45:23.549987Z","iopub.execute_input":"2022-05-03T07:45:23.550752Z","iopub.status.idle":"2022-05-03T07:45:23.865584Z","shell.execute_reply.started":"2022-05-03T07:45:23.550696Z","shell.execute_reply":"2022-05-03T07:45:23.864741Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# VGG16","metadata":{}},{"cell_type":"code","source":"vgg16 = models.vgg16(pretrained=True)\n\nfor param in vgg16.parameters(): #freeze model\n    param.requires_grad = False\n    \nn_inputs = vgg16.classifier[0].in_features\nvgg16.classifier = nn.Sequential(\n    nn.Linear(n_inputs, 4096),\n    nn.ReLU(),\n    nn.Dropout(0.3),\n    nn.Linear(4096, 512),\n    nn.ReLU(),\n    nn.Dropout(0.3),\n    nn.Linear(512, len(classes))\n)\nvgg16 = vgg16.to(device)\n# print(vgg16.classifier)\n\ncriterion = LabelSmoothingCrossEntropy()\ncriterion = criterion.to(device)\noptimizer = optim.AdamW(vgg16.classifier.parameters(), lr=0.001)\n\nexp_lr_scheduler = optim.lr_scheduler.StepLR(optimizer, step_size=3, gamma=0.97)\n\nvgg16 = train_model(vgg16, criterion, optimizer, exp_lr_scheduler, num_epochs=10) # now it is a lot faster","metadata":{"execution":{"iopub.status.busy":"2022-05-03T09:10:31.416916Z","iopub.execute_input":"2022-05-03T09:10:31.417172Z","iopub.status.idle":"2022-05-03T10:00:11.979769Z","shell.execute_reply.started":"2022-05-03T09:10:31.417144Z","shell.execute_reply":"2022-05-03T10:00:11.978969Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"total_params = sum(p.numel() for p in vgg16.parameters())\nprint(\"Total Parameters :\", total_params)","metadata":{"execution":{"iopub.status.busy":"2022-05-03T09:08:52.687912Z","iopub.execute_input":"2022-05-03T09:08:52.688219Z","iopub.status.idle":"2022-05-03T09:08:52.693794Z","shell.execute_reply.started":"2022-05-03T09:08:52.688189Z","shell.execute_reply":"2022-05-03T09:08:52.693076Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# RESNET18","metadata":{}},{"cell_type":"code","source":"resnet18 = models.resnet18(pretrained=True)\nfor param in resnet18.parameters(): #freeze model\n    param.requires_grad = False\n    \nn_inputs = resnet18.fc.in_features\nresnet18.fc = nn.Linear(n_inputs, len(classes))\nresnet18 = resnet18.to(device)\n\ncriterion = LabelSmoothingCrossEntropy()\ncriterion = criterion.to(device)\noptimizer = optim.AdamW(resnet18.fc.parameters(), lr=0.001)\n\nexp_lr_scheduler = optim.lr_scheduler.StepLR(optimizer, step_size=3, gamma=0.97)\n\nresnet18 = train_model(resnet18, criterion, optimizer, exp_lr_scheduler, num_epochs=10)","metadata":{"execution":{"iopub.status.busy":"2022-05-03T10:00:11.981899Z","iopub.execute_input":"2022-05-03T10:00:11.982354Z","iopub.status.idle":"2022-05-03T10:48:55.631937Z","shell.execute_reply.started":"2022-05-03T10:00:11.982313Z","shell.execute_reply":"2022-05-03T10:48:55.631174Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"total_params = sum(p.numel() for p in resnet18.parameters())\nprint(\"Total Parameters :\", total_params)","metadata":{"execution":{"iopub.status.busy":"2022-05-03T11:15:45.489582Z","iopub.execute_input":"2022-05-03T11:15:45.490249Z","iopub.status.idle":"2022-05-03T11:15:45.496716Z","shell.execute_reply.started":"2022-05-03T11:15:45.490212Z","shell.execute_reply":"2022-05-03T11:15:45.495989Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# AlexNet","metadata":{}},{"cell_type":"code","source":"alexnet = models.alexnet(pretrained=True)\n\nfor param in alexnet.parameters(): #freeze model\n    param.requires_grad = False\n    \nn_inputs = alexnet.classifier[1].in_features\nalexnet.classifier = nn.Sequential(\n    nn.Linear(n_inputs, 4096),\n    nn.ReLU(),\n    nn.Dropout(0.3),\n    nn.Linear(4096, 512),\n    nn.ReLU(),\n    nn.Dropout(0.3),\n    nn.Linear(512, len(classes))\n)\nalexnet = alexnet.to(device)\nprint(alexnet.classifier)\n\ncriterion = LabelSmoothingCrossEntropy()\ncriterion = criterion.to(device)\noptimizer = optim.AdamW(alexnet.classifier.parameters(), lr=0.001)\n\nexp_lr_scheduler = optim.lr_scheduler.StepLR(optimizer, step_size=3, gamma=0.97)\n\nalexnet = train_model(alexnet, criterion, optimizer, exp_lr_scheduler, num_epochs=10) # now it is a lot faster","metadata":{"execution":{"iopub.status.busy":"2022-05-03T11:16:25.679689Z","iopub.execute_input":"2022-05-03T11:16:25.679949Z","iopub.status.idle":"2022-05-03T12:06:17.635493Z","shell.execute_reply.started":"2022-05-03T11:16:25.679921Z","shell.execute_reply":"2022-05-03T12:06:17.634237Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"total_params = sum(p.numel() for p in alexnet.parameters())\nprint(\"Total Parameters :\", total_params)","metadata":{"execution":{"iopub.status.busy":"2022-05-03T12:06:17.63773Z","iopub.execute_input":"2022-05-03T12:06:17.638172Z","iopub.status.idle":"2022-05-03T12:06:17.643771Z","shell.execute_reply.started":"2022-05-03T12:06:17.63813Z","shell.execute_reply":"2022-05-03T12:06:17.643098Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# DenseNet161","metadata":{}},{"cell_type":"code","source":"densenet = models.densenet161(pretrained=True)\n\nfor param in densenet.parameters(): #freeze model\n    param.requires_grad = False\n    \nn_inputs = densenet.classifier.in_features\ndensenet.classifier = nn.Linear(n_inputs, len(classes))\ndensenet = densenet.to(device)\nprint(densenet.classifier)\n\ncriterion = LabelSmoothingCrossEntropy()\ncriterion = criterion.to(device)\noptimizer = optim.AdamW(densenet.classifier.parameters(), lr=0.001)\n\nexp_lr_scheduler = optim.lr_scheduler.StepLR(optimizer, step_size=3, gamma=0.97)\n\ndensenet = train_model(densenet, criterion, optimizer, exp_lr_scheduler, num_epochs=10) ","metadata":{"execution":{"iopub.status.busy":"2022-05-03T12:06:17.64507Z","iopub.execute_input":"2022-05-03T12:06:17.645536Z","iopub.status.idle":"2022-05-03T12:58:45.662538Z","shell.execute_reply.started":"2022-05-03T12:06:17.645502Z","shell.execute_reply":"2022-05-03T12:58:45.661599Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"total_params = sum(p.numel() for p in densenet.parameters())\nprint(\"Total Parameters :\", total_params)","metadata":{"execution":{"iopub.status.busy":"2022-05-03T12:58:45.665189Z","iopub.execute_input":"2022-05-03T12:58:45.665828Z","iopub.status.idle":"2022-05-03T12:58:45.674145Z","shell.execute_reply.started":"2022-05-03T12:58:45.665795Z","shell.execute_reply":"2022-05-03T12:58:45.672935Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}