{"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":"### import Modules","metadata":{}},{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\nimport time\nimport os\nimport copy\nimport json\n\n# importing plotting lib\nimport seaborn as sns\nimport matplotlib.pyplot as plt\nfrom PIL import Image\n\n# importing all torch lib\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\nfrom torchvision import transforms\nfrom tqdm import tqdm\n# for augmentaiton\nimport albumentations as A\nfrom albumentations.pytorch.transforms import ToTensorV2\n\n# for ignoring warnings\nimport warnings\nwarnings.filterwarnings('ignore')\n%matplotlib inline","metadata":{"execution":{"iopub.status.busy":"2022-07-09T11:31:31.358159Z","iopub.execute_input":"2022-07-09T11:31:31.358774Z","iopub.status.idle":"2022-07-09T11:31:31.370722Z","shell.execute_reply.started":"2022-07-09T11:31:31.358737Z","shell.execute_reply":"2022-07-09T11:31:31.369366Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Load the Data","metadata":{}},{"cell_type":"code","source":"BASE_DIR = '../input/cassava-leaf-disease-classification/'","metadata":{"execution":{"iopub.status.busy":"2022-07-09T10:49:39.854054Z","iopub.execute_input":"2022-07-09T10:49:39.854809Z","iopub.status.idle":"2022-07-09T10:49:39.86748Z","shell.execute_reply.started":"2022-07-09T10:49:39.854772Z","shell.execute_reply":"2022-07-09T10:49:39.866514Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train = pd.read_csv(BASE_DIR + 'train.csv')\ntrain.head()","metadata":{"execution":{"iopub.status.busy":"2022-07-09T10:49:40.77269Z","iopub.execute_input":"2022-07-09T10:49:40.773629Z","iopub.status.idle":"2022-07-09T10:49:40.803269Z","shell.execute_reply.started":"2022-07-09T10:49:40.773583Z","shell.execute_reply":"2022-07-09T10:49:40.802184Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"## loading mapping for target label\n\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-07-09T10:49:40.805608Z","iopub.execute_input":"2022-07-09T10:49:40.806049Z","iopub.status.idle":"2022-07-09T10:49:40.815505Z","shell.execute_reply.started":"2022-07-09T10:49:40.80601Z","shell.execute_reply":"2022-07-09T10:49:40.814493Z"},"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-07-09T10:49:40.817307Z","iopub.execute_input":"2022-07-09T10:49:40.817928Z","iopub.status.idle":"2022-07-09T10:49:40.830902Z","shell.execute_reply.started":"2022-07-09T10:49:40.817892Z","shell.execute_reply":"2022-07-09T10:49:40.829935Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Exploratory Data Analysis","metadata":{}},{"cell_type":"code","source":"def plot_images(class_id, label, total_images = 6):\n    # get image ids correcsponding to the target class id\n    plot_list = train[train['label'] == class_id].sample(total_images)['image_id'].tolist()\n    \n    labels = [label for i in range(total_images)]\n    size = int(np.sqrt(total_images))\n    if size*size < total_images:\n        size += 1\n    plt.figure(figsize =(15, 15))\n    \n    #plot the image in subplot\n    for index, (image_id, label) in enumerate(zip(plot_list, labels)):\n        plt.subplot(size, size, index + 1)\n        image = Image.open(str(BASE_DIR + \"train_images/\" +image_id))\n        plt.imshow(image)\n        plt.title(label, fontsize = 14)\n        plt.axis(\"off\")\n        \n    plt.show()\n        \n        \n    ","metadata":{"execution":{"iopub.status.busy":"2022-07-09T10:49:41.442956Z","iopub.execute_input":"2022-07-09T10:49:41.443639Z","iopub.status.idle":"2022-07-09T10:49:41.452646Z","shell.execute_reply.started":"2022-07-09T10:49:41.443598Z","shell.execute_reply":"2022-07-09T10:49:41.451653Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plot_images(0, mapping[0], 6)","metadata":{"execution":{"iopub.status.busy":"2022-07-09T10:49:41.637717Z","iopub.execute_input":"2022-07-09T10:49:41.638104Z","iopub.status.idle":"2022-07-09T10:49:42.415454Z","shell.execute_reply.started":"2022-07-09T10:49:41.638072Z","shell.execute_reply":"2022-07-09T10:49:42.414216Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# class distribution\nsns.countplot(train[\"label\"])","metadata":{"execution":{"iopub.status.busy":"2022-07-09T10:49:42.417695Z","iopub.execute_input":"2022-07-09T10:49:42.41803Z","iopub.status.idle":"2022-07-09T10:49:42.584746Z","shell.execute_reply.started":"2022-07-09T10:49:42.417999Z","shell.execute_reply":"2022-07-09T10:49:42.583731Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#### Obs: The label distribution is skewed. Hence We should use AUC for metric and Stratified KFold for data spliting","metadata":{}},{"cell_type":"markdown","source":"### configuration and Utility Functions","metadata":{}},{"cell_type":"code","source":"# All constant config\n\nDIM = (256, 256)\nWIDTH, HEIGHT = DIM\nNUM_CLASSES = 5\nNUM_WORKERS = 24\nTRAIN_BATCH_SIZE = 128\nTEST_BATCH_SIZE = 128\nSEED = 4\nDEVICE  = 'cuda'\nMEAN = (0.485,0.456, 0.406)\nSTD = (0.229, 0.224, 0.225)\nLR = 0.001","metadata":{"execution":{"iopub.status.busy":"2022-07-09T10:49:43.344151Z","iopub.execute_input":"2022-07-09T10:49:43.344488Z","iopub.status.idle":"2022-07-09T10:49:43.350572Z","shell.execute_reply.started":"2022-07-09T10:49:43.34446Z","shell.execute_reply":"2022-07-09T10:49:43.349538Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Augmentations","metadata":{}},{"cell_type":"code","source":"def get_test_transform(value = 'val'):\n    if value == 'train':\n        return A.Compose([\n            A.Resize(WIDTH, HEIGHT),\n            A.HorizontalFlip(p = 0.5),\n            A.Rotate(limit = (-90, 90)),\n            A.VerticalFlip(p = 0.5),\n            A.Normalize(MEAN, STD, max_pixel_value = 255.0, always_apply = True),\n            ToTensorV2(p=1.0) # returning tensor for all images\n        ])\n    elif value == 'val':\n        return A.Compose([\n            A.Resize(WIDTH, HEIGHT),\n            A.Normalize(MEAN, STD, max_pixel_value = 255.0, always_apply = True),\n            ToTensorV2(p=1.0)\n        ])\n        ","metadata":{"execution":{"iopub.status.busy":"2022-07-09T10:49:43.35251Z","iopub.execute_input":"2022-07-09T10:49:43.353221Z","iopub.status.idle":"2022-07-09T10:49:43.363079Z","shell.execute_reply.started":"2022-07-09T10:49:43.353179Z","shell.execute_reply":"2022-07-09T10:49:43.362195Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Dataset Loader Class","metadata":{}},{"cell_type":"code","source":"class CassavaDataset(Dataset):\n    def __init__(self, image_ids, labels, dim = None, aug = None, folder = 'train_images'):\n        super().__init__()\n        self.image_ids = image_ids\n        self.labels = labels\n        self.dim = dim\n        self.aug = aug\n        self.folder = folder\n        \n    def __len__(self):\n        return len(self.image_ids)\n    \n    def __getitem__(self, index):\n        image = self.image_ids[index]\n        img = Image.open(os.path.join(BASE_DIR, self.folder, image))\n        \n        if(self.dim):\n            img = img.resize(self.dim)\n            \n        img = np.array(img)\n        \n        if self.aug is not None:\n            augmented = self.aug(image = img)\n            img = augmented[\"image\"]\n            \n        label = torch.tensor(self.labels[index], dtype = torch.long)\n        return img, label\n            \n        \n        ","metadata":{"execution":{"iopub.status.busy":"2022-07-09T10:49:43.365188Z","iopub.execute_input":"2022-07-09T10:49:43.365748Z","iopub.status.idle":"2022-07-09T10:49:43.377573Z","shell.execute_reply.started":"2022-07-09T10:49:43.365715Z","shell.execute_reply":"2022-07-09T10:49:43.376621Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Create Folds","metadata":{}},{"cell_type":"code","source":"from sklearn.model_selection import StratifiedKFold\nkf = StratifiedKFold(n_splits = 5, shuffle = True, random_state = 20)\ntrain['kfold']  = -1\nfor fold, (train_idx, val_idx) in enumerate(kf.split(X = train['image_id'], y = train['label'])):\n    train.loc[val_idx, 'kfold'] = fold\ntrain.head()","metadata":{"execution":{"iopub.status.busy":"2022-07-09T10:49:44.499046Z","iopub.execute_input":"2022-07-09T10:49:44.499798Z","iopub.status.idle":"2022-07-09T10:49:44.526988Z","shell.execute_reply.started":"2022-07-09T10:49:44.499761Z","shell.execute_reply":"2022-07-09T10:49:44.525805Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Use Pretrained Model (Transfer learning)","metadata":{}},{"cell_type":"code","source":"def getmodel():\n    net = models.efficientnet_b0(pretrained = True)\n    \n    ## freeze all the layers in the network\n    for param in net.parameters():\n        param.requires_grad = False\n        \n    num_ftrs =1280 \n    # create last few layers\n    net.classifier[1] = nn.Sequential(\n        nn.Linear(num_ftrs, 256),\n        nn.ReLU(),\n        nn.Dropout(0.3),\n        nn.Linear(256, NUM_CLASSES),\n        nn.LogSoftmax(dim = 1)\n    )\n    # use gpu if any\n    net = net.to(device = DEVICE)\n    return net","metadata":{"execution":{"iopub.status.busy":"2022-07-09T10:49:44.707779Z","iopub.execute_input":"2022-07-09T10:49:44.70859Z","iopub.status.idle":"2022-07-09T10:49:44.715948Z","shell.execute_reply.started":"2022-07-09T10:49:44.708558Z","shell.execute_reply":"2022-07-09T10:49:44.7146Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = getmodel()","metadata":{"execution":{"iopub.status.busy":"2022-07-09T10:49:44.719407Z","iopub.execute_input":"2022-07-09T10:49:44.720404Z","iopub.status.idle":"2022-07-09T10:49:44.924779Z","shell.execute_reply.started":"2022-07-09T10:49:44.720364Z","shell.execute_reply":"2022-07-09T10:49:44.923813Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"criterion = nn.CrossEntropyLoss()\noptimizer = optim.Adam(model.parameters(), lr = LR)","metadata":{"execution":{"iopub.status.busy":"2022-07-09T10:49:46.079621Z","iopub.execute_input":"2022-07-09T10:49:46.079979Z","iopub.status.idle":"2022-07-09T10:49:46.087388Z","shell.execute_reply.started":"2022-07-09T10:49:46.079948Z","shell.execute_reply":"2022-07-09T10:49:46.086119Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# find total parameters in the model\ntotal_params = sum(p.numel() for p in model.parameters())\nprint(f\"total_params: {total_params:,}\")\ntrainable_params = sum(p.numel() for p in model.parameters() if p.requires_grad)\nprint(f\"train_params: {trainable_params:,}\")","metadata":{"execution":{"iopub.status.busy":"2022-07-09T10:49:46.089623Z","iopub.execute_input":"2022-07-09T10:49:46.090792Z","iopub.status.idle":"2022-07-09T10:49:46.100999Z","shell.execute_reply.started":"2022-07-09T10:49:46.090751Z","shell.execute_reply":"2022-07-09T10:49:46.099916Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Steps for Training and Validaiton","metadata":{}},{"cell_type":"code","source":"\n\ndef train_model(model, dataloaders, criterion, optimizer, num_epochs=5):\n    # set starting time\n    start_time = time.time()\n    \n    val_acc_history = []\n    \n    best_model_wts = copy.deepcopy(model.state_dict())\n    best_acc = 0.0\n    \n    for epoch in tqdm(range(num_epochs)):\n        print(f'Epoch {epoch}/{num_epochs-1}')\n        print('-'*15)\n        \n        # each epoch have training and validation phase\n        for phase in ['train', 'val']:\n            # set mode for model\n            if phase == 'train':\n                model.train() # set model to training mode\n            else:\n                model.eval() # set model to evaluate mode\n                \n            running_loss = 0.0\n            running_corrects = 0\n            fin_out = []\n            \n            # iterate over data\n            for inputs, labels in dataloaders[phase]:\n                # move data to corresponding hardware\n                inputs = inputs.to(DEVICE)\n                labels = labels.to(DEVICE)\n                \n                # reset (or) zero the parameter gradients\n                optimizer.zero_grad()\n                \n                # training (or) validation process\n                with torch.set_grad_enabled(phase=='train'):\n                    outputs = model(inputs)\n                    loss = criterion(outputs, labels)\n                    \n                    _, preds = torch.max(outputs, 1)\n                    \n                    # back propagation in the network\n                    if phase == 'train':\n                        loss.backward()\n                        optimizer.step()\n\n                        \n                running_loss += loss.item() * inputs.size(0)\n                running_corrects += torch.sum(preds == labels.data)\n                \n            # calculate loss and accuarcy for the epoch\n            epoch_loss = running_loss / len(dataloaders[phase].dataset)\n            epoch_acc = running_corrects.double() / len(dataloaders[phase].dataset)\n            \n            # print loss and acc for training & validation\n            print('{} Loss: {:.4f} Acc: {:.4f}'.format(phase, epoch_loss, epoch_acc))\n            \n            # update the best weights\n            if phase == 'val' and epoch_acc > best_acc:\n                best_acc = epoch_acc\n                best_model_wts = copy.deepcopy(model.state_dict())\n            if phase == 'val':\n                val_acc_history.append(epoch_acc)\n                \n        print()\n    end_time = time.time() - start_time\n    \n    print('Training completes in {:.0f}m {:.0f}s'.format(end_time // 60, end_time % 60))\n    print('Best Val Acc: {:.4f}'.format(best_acc))\n    \n    # load best model weights\n    model.load_state_dict(best_model_wts)\n    return model, val_acc_history\n\n","metadata":{"execution":{"iopub.status.busy":"2022-07-09T10:52:30.352907Z","iopub.execute_input":"2022-07-09T10:52:30.353675Z","iopub.status.idle":"2022-07-09T10:52:30.370118Z","shell.execute_reply.started":"2022-07-09T10:52:30.35364Z","shell.execute_reply":"2022-07-09T10:52:30.369183Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"## train the model\nfor fold in tqdm(range(5)):\n    print(f\"Fold: {fold}\")\n    # create train data and val data\n    train_data = train[train['kfold'] != fold]\n    val_data = train[train['kfold'] == fold]\n    \n    # create trainset, trainloader\n    train_dataset = CassavaDataset(\n    image_ids = train_data['image_id'].values, \n    labels = train_data['label'].values,\n    aug = get_test_transform('train'),\n    dim = DIM)\n\n    train_loader  = DataLoader(\n        train_dataset, \n        batch_size = TRAIN_BATCH_SIZE, \n        shuffle = False, \n        num_workers = NUM_WORKERS)\n    \n    #create the valset and valloader\n    val_dataset = CassavaDataset(\n    image_ids = val_data['image_id'].values, \n    labels = val_data['label'].values,\n    aug = get_test_transform('val'),\n    dim = DIM)\n\n    val_loader  = DataLoader(\n        val_dataset, \n        batch_size = TRAIN_BATCH_SIZE, \n        shuffle = False, \n        num_workers = NUM_WORKERS)\n    \n    loader = {\"train\":train_loader, \"val\":val_loader}\n    \n    model, accuracy = train_model(model = model, \n                                  dataloaders = loader, \n                                  criterion = criterion, \n                                  optimizer = optimizer, \n                                  num_epochs = 5)\n    torch.save(model, f'/kaggle/working/best_model_{fold}.h5')\n    torch.save(model.state_dict(), f'/kaggle/working/best_model_weights_{fold}')","metadata":{"execution":{"iopub.status.busy":"2022-07-09T10:52:32.756756Z","iopub.execute_input":"2022-07-09T10:52:32.757185Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.cuda.empty_cache()","metadata":{"execution":{"iopub.status.busy":"2022-07-09T10:51:42.177963Z","iopub.status.idle":"2022-07-09T10:51:42.180825Z","shell.execute_reply.started":"2022-07-09T10:51:42.180554Z","shell.execute_reply":"2022-07-09T10:51:42.180581Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}