{"cells":[{"metadata":{"trusted":true},"cell_type":"code","source":"!pip install ../input/efficientnet-pytorch-070/efficientnet_pytorch-0.7.0-py3-none-any.whl","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"import os\nimport time\nimport glob\nimport copy\nimport time\nimport json\nimport pprint\n\nfrom tqdm.notebook import tqdm\nimport albumentations\nimport numpy as np\nimport pandas as pd\nfrom PIL import Image\nimport matplotlib.pyplot as plt\nfrom sklearn.model_selection import KFold, StratifiedKFold\nfrom sklearn.utils import class_weight\nfrom sklearn.metrics import classification_report, confusion_matrix\nimport seaborn as sn\n\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import models, transforms\n\nfrom efficientnet_pytorch import EfficientNet","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"# config\neffnet_config = {\n    'DATA': {\n        'IMAGES': \"../input/cassava-leaf-disease-classification/train_images\",\n        'LABELS': \"../input/cassava-leaf-disease-classification/train.csv\",\n        'SUB_IMAGES': \"../input/cassava-leaf-disease-classification/test_images\",\n        'SUB_LABELS': \"../input/cassava-leaf-disease-classification/sample_submission.csv\",\n        'SUB_OUTPUT': \"../input/cassava-leaf-disease-classification/submission.csv\"\n    },\n    'DEVICE': \"cuda\",\n    'NUM_GPU': torch.cuda.device_count(),\n    'TRAIN_BATCH_SIZE': 16,\n    'VAL_BATCH_SIZE': 8,\n    'CLASSES': 5,\n    'CV_FOLDS': 5,\n    'NUM_EPOCHS': 15,\n    'MODEL_PATH': \"model.pth\",\n    'SGD': {\n        'LR': 0.001,\n        'MOMENTUM': 0.9,\n        'WEIGHT_DECAY': 0.001\n    },\n    'COS_ANN_LR': {\n        'ETA_MIN': 0.00001\n    },\n    'MODEL_TYPE': 'EFFICIENT_NET_B4',\n}\n\nresnet_config = {\n    'DATA': {\n        'IMAGES': \"../input/cassava-leaf-disease-classification/train_images\",\n        'LABELS': \"../input/cassava-leaf-disease-classification/train.csv\",\n        'SUB_IMAGES': \"../input/cassava-leaf-disease-classification/test_images\",\n        'SUB_LABELS': \"../input/cassava-leaf-disease-classification/sample_submission.csv\",\n        'SUB_OUTPUT': \"../input/cassava-leaf-disease-classification/submission.csv\"\n    },\n    'DEVICE': \"cuda\",\n    'NUM_GPU': torch.cuda.device_count(),\n    'TRAIN_BATCH_SIZE': 32,\n    'VAL_BATCH_SIZE': 16,\n    'CLASSES': 5,\n    'CV_FOLDS': 5,\n    'NUM_EPOCHS': 15,\n    'MODEL_PATH': \"model.pth\",\n    'SGD': {\n        'LR': 0.0005,\n        'MOMENTUM': 0.9,\n        'WEIGHT_DECAY': 0.001\n    },\n    'COS_ANN_LR': {\n        'ETA_MIN': 0.00001\n    },\n    'MODEL_TYPE': 'RESNET_50'\n}\n\n# change this to change the models\nconfig = effnet_config","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"# # create log dir\n# current_time = time.strftime(\"%m.%d.%y_%H:%M:%S\", time.localtime())\n# log_folder = 'output/kfold_%s' % current_time\n# os.mkdir(log_folder)\n\n# # write config to log dir\n# with open(os.path.join(log_folder, 'config.json'), 'w') as fp:\n#     json.dump(config, fp, indent=4)","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Loading Data"},{"metadata":{"trusted":true},"cell_type":"code","source":"labels = pd.read_csv(config['DATA']['LABELS'])\nprint(\"Total labels:\", len(labels))","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Exploring data label distribution"},{"metadata":{"trusted":true},"cell_type":"code","source":"print(labels['label'].value_counts() / len(labels) * 100)\nlabels['label'].value_counts().plot.bar(x='label', y='count')","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Defining Augmentations"},{"metadata":{"trusted":true},"cell_type":"code","source":"# modded\ntrain_aug = albumentations.Compose([\n                albumentations.RandomResizedCrop(512, 512, scale=(0.5, 1.0)),\n                albumentations.Transpose(p=0.8),\n                albumentations.HorizontalFlip(p=0.5),\n                albumentations.VerticalFlip(p=0.5),\n                albumentations.ShiftScaleRotate(p=0.8),\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.CLAHE(p=0.5),\n                albumentations.Normalize(\n                    mean=[0.485, 0.456, 0.406], \n                    std=[0.229, 0.224, 0.225], \n                    max_pixel_value=255.0, \n                    p=1.0\n                ),\n                albumentations.CoarseDropout(max_holes=20, max_height=30, max_width=30, p=0.5),\n                albumentations.Cutout(num_holes=20, max_h_size=30, max_w_size=30, p=0.5)\n            ], p=1.)\n\nvalid_aug = albumentations.Compose([\n                albumentations.CenterCrop(512, 512, p=1.),\n                albumentations.Normalize(\n                    mean=[0.485, 0.456, 0.406], \n                    std=[0.229, 0.224, 0.225], \n                    max_pixel_value=255.0, \n                    p=1.0\n                )\n            ], p=1.)","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Defining Dataset Class and Visualize function"},{"metadata":{"trusted":true},"cell_type":"code","source":"def visualize_samples(dataset, count):\n    fig = plt.figure()\n\n    for i in range(len(dataset)):\n        sample = dataset[i]\n\n        ax = plt.subplot(1, count+1, i + 1)\n        plt.tight_layout()\n        ax.set_title('Label: {}'.format(sample[1]))\n        ax.axis('off')\n        plt.imshow(sample[0])\n\n        if i == count:\n            plt.show()\n            break\n\n            \nclass CassavaLeafDiseaseDataset(Dataset):\n    def __init__(self, data_images, data_labels, img_dir, transform=None):\n        self.data_images = data_images\n        self.data_labels = data_labels\n        self.img_dir = img_dir\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.data_images)\n\n    def __getitem__(self, idx):\n        if torch.is_tensor(idx):\n            idx = idx.tolist()\n\n        img_name = os.path.join(self.img_dir, self.data_images[idx])\n        image = np.array(Image.open(img_name))\n        \n        if self.transform:\n            image = self.transform(image=image)[\"image\"]\n            \n        image = transforms.ToTensor()(image)\n        label = self.data_labels[idx]\n        \n        return (image, label)\n    \n    \n# https://www.kaggle.com/ar2017/pytorch-efficientnet-train-aug-cutmix-fmix/notebook\n# https://www.kaggle.com/virajbagal/mixup-cutmix-fmix-visualisations/output\n\ndef rand_bbox(size, lam):\n    W = size[2]\n    H = size[3]\n    cut_rat = np.sqrt(1. - lam)\n    cut_w = np.int(W * cut_rat)\n    cut_h = np.int(H * cut_rat)\n\n    # uniform\n    cx = np.random.randint(W)\n    cy = np.random.randint(H)\n\n    bbx1 = np.clip(cx - cut_w // 2, 0, W)\n    bby1 = np.clip(cy - cut_h // 2, 0, H)\n    bbx2 = np.clip(cx + cut_w // 2, 0, W)\n    bby2 = np.clip(cy + cut_h // 2, 0, H)\n    return bbx1, bby1, bbx2, bby2\n\ndef cutmix(data, target, alpha):\n    indices = torch.randperm(data.size(0))\n    shuffled_data = data[indices]\n    shuffled_target = target[indices]\n\n    lam = np.clip(np.random.beta(alpha, alpha),0.3,0.4)\n    bbx1, bby1, bbx2, bby2 = rand_bbox(data.size(), lam)\n    new_data = data.clone()\n    new_data[:, :, bby1:bby2, bbx1:bbx2] = data[indices, :, bby1:bby2, bbx1:bbx2]\n    # adjust lambda to exactly match pixel ratio\n    lam = 1 - ((bbx2 - bbx1) * (bby2 - bby1) / (data.size()[-1] * data.size()[-2]))\n    targets = (target, shuffled_target, lam)\n\n    return new_data, targets","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Creating criterion and optimizer"},{"metadata":{"trusted":true},"cell_type":"code","source":"# Loading ResNet from PyTorch and update it's fc layer output to 5 classes\nmodel = None\nif config['MODEL_TYPE'] == \"EFFICIENT_NET_B4\":\n    model = model = EfficientNet.from_pretrained('efficientnet-b4', num_classes=config['CLASSES'])\nelif config['MODEL_TYPE'] == \"RESNET_50\":\n    model = models.resnext50_32x4d(pretrained=True)\n    model.fc = nn.Linear(2048, config['CLASSES'])\n    \nif config['DEVICE'] == 'cuda' and config['NUM_GPU'] > 1:\n    model = nn.DataParallel(model)\n\n# calculating class weights - class_weights = n_samples / (n_classes * np.bincount(y))\nclass_weights = {index: (len(labels)/(config['CLASSES'] * value)) for index, value in labels['label'].value_counts().items()}\nclass_weight_tensor = torch.FloatTensor([class_weights[label] for label in range(config['CLASSES'])])\nprint(\"Class Weights:\", class_weight_tensor)\n\n# Creating the criterion\n# criterion = nn.CrossEntropyLoss(weight=class_weight_tensor.to(config['DEVICE']))\ncriterion = nn.CrossEntropyLoss()\n\n# Creating the optimizer and lr scheduler\noptimizer = optim.SGD(\n    params=model.parameters(),\n    lr=config['SGD']['LR'],\n    momentum=config['SGD']['MOMENTUM'],\n    weight_decay=config['SGD']['WEIGHT_DECAY']\n)","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Training"},{"metadata":{"scrolled":false,"trusted":true},"cell_type":"code","source":"def train_model(fold_no, model, dataloaders, criterion, optimizer, num_epochs=10, device='cpu'):\n    model = model.to(device)\n    since = time.time()\n\n    best_model_wts = copy.deepcopy(model.state_dict())\n    best_acc = 0.0\n    \n    train_loss_history = []\n    train_acc_history = []\n    val_loss_history = []\n    val_acc_history = []\n    \n    # creating tqdm progress bar\n    progress_bar = tqdm(dynamic_ncols=True)\n    \n    # reset lr scheduler to base lr\n    lr_scheduler = optim.lr_scheduler.CosineAnnealingLR(optimizer,T_max=num_epochs, eta_min=config['COS_ANN_LR']['ETA_MIN'])\n            \n    for epoch in range(num_epochs):\n                \n        # Each epoch has a training and validation phase\n        for phase in ['train', 'val']:\n            \n            # storage for preds and target labels\n            val_preds = torch.tensor([]).to(device)\n            val_labels = torch.tensor([]).to(device)\n        \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            \n            # update the progress bar total\n            progress_bar.reset(total=len(dataloaders[phase]))\n\n            # Iterate over data.\n            for inputs, labels in dataloaders[phase]:\n                inputs = inputs.to(device)\n                labels = labels.to(device)\n                \n                if phase == \"train\":\n                    mix_decision = np.random.rand()\n                    if mix_decision < 0.5:\n                        inputs, labels = cutmix(inputs, labels, 1.)\n                else:\n                    mix_decision = 1.0\n                    \n                # zero the parameter gradients\n                optimizer.zero_grad()\n\n                # forward\n                # track history if only in train\n                with torch.set_grad_enabled(phase == 'train'):\n                    # Get model outputs and calculate loss\n                    # Special case for inception because in training it has an auxiliary output. In train\n                    #   mode we calculate the loss by summing the final output and the auxiliary output\n                    #   but in testing we only consider the final output.\n                    \n                    outputs = model(inputs)\n                    \n                    if mix_decision < 0.50:\n                        loss = criterion(outputs, labels[0]) * labels[2] + criterion(outputs, labels[1]) * (1. - labels[2])\n                    else:\n                        loss = criterion(outputs, labels)\n\n                    _, preds = torch.max(outputs, 1)\n                    \n                    # backward + optimize only if in training phase\n                    if phase == 'train':\n                        loss.backward()\n                        optimizer.step()\n                    else:\n                        val_preds = torch.cat((val_preds, preds),dim=0)\n                        val_labels = torch.cat((val_labels, labels),dim=0)\n\n                # statistics\n                running_loss += loss.item() * inputs.size(0)\n                if mix_decision < 0.5:\n                    if labels[2] >= 0.5:\n                        running_corrects += torch.sum(preds==labels[0].data)\n                    else:\n                        running_corrects += torch.sum(preds==labels[1].data)\n                else:\n                    running_corrects += torch.sum(preds == labels.data)\n                \n                # update tqdm progress bar\n                progress_bar.set_description_str(f'Epoch: {epoch}/{num_epochs - 1} Phase: {phase} Itr_Loss: {loss.item():.4f}')\n                progress_bar.update(1)\n            \n            # decaying the lr\n            if phase == 'train':\n                lr_scheduler.step()\n            \n            # calculating the epoch loss and accuracy\n            epoch_loss = running_loss / len(dataloaders[phase].dataset)\n            epoch_acc = running_corrects.double() / len(dataloaders[phase].dataset)\n            \n            print('Epoch: {} | Phase: {} | Loss: {:.4f} | Acc: {:.4f}'.format(epoch, phase, epoch_loss, epoch_acc))\n            \n            if phase =='val':\n                # calculating the precision, recall, f1-score and support\n                metric_report = classification_report(val_labels.cpu(), val_preds.cpu(), output_dict=True)\n                pprint.pprint(metric_report)\n\n                # calculating and plotting the confusion matrix\n                cmt = confusion_matrix(val_labels.cpu(), val_preds.cpu())\n                df_cm = pd.DataFrame(cmt, range(config['CLASSES']), range(config['CLASSES']))\n                plt.figure(figsize=(10,7))\n                sn.set(font_scale=1.4)\n                sn.heatmap(df_cm, annot=True, annot_kws={\"size\": 16}, fmt='d')\n                plt.savefig(os.path.join(log_folder, 'confusion_matrix.png'))\n                plt.show()\n            \n            # deep copy the model\n            if phase == 'val' and epoch_acc > best_acc:\n                best_acc = epoch_acc\n                best_model_wts = copy.deepcopy(model.state_dict())\n                \n                # Save the model dict\n                save_model(copy.deepcopy(model), suffix=config['MODEL_TYPE']+\"_\"+str(fold_no)+\"_\"+str(epoch))\n                \n                print(\"Saved the model\")\n            if phase == 'val':\n                val_acc_history.append(epoch_acc)\n                val_loss_history.append(epoch_loss)\n            else:\n                train_acc_history.append(epoch_acc)\n                train_loss_history.append(epoch_loss)\n        \n        print()\n    \n    # closing the progress bar \n    progress_bar.close()\n    \n    time_elapsed = time.time() - since\n    print('Training complete in {:.0f}m {:.0f}s'.format(time_elapsed // 60, time_elapsed % 60))\n    print('Best val accuracy: {:4f}'.format(best_acc))\n\n    # load best model weights\n    model.load_state_dict(best_model_wts)\n    return model, train_loss_history, val_loss_history, train_acc_history, val_acc_history\n\n\ndef save_model(model, suffix=\"\"):\n    model = model.to(\"cpu\")\n    try:\n        state_dict = model.module.state_dict()\n    except AttributeError:\n        state_dict = model.state_dict()\n        \n    torch.save(state_dict, suffix+\".pth\")\n\n\n# ~~~~~~~~~~~~~~~ Training code ~~~~~~~~~~~~~~~ #\n\n# separating data into images and labels\nall_images = labels['image_id'].to_numpy()\nall_labels = labels['label'].to_numpy()\n\n# creating stratified k folds\nskf = StratifiedKFold(n_splits=config['CV_FOLDS'], shuffle=True, random_state=10)\nskf.get_n_splits(all_labels)\n\n# make a copy of pretrained imagenet model weights\nmodel_weights = copy.deepcopy(model.state_dict())\n\nfor fold, (train_index, val_index) in enumerate(skf.split(all_images, all_labels)):\n    if fold < 2:\n        continue\n    print(\"Fold: {} \".format(fold))\n    \n    # printing distribution\n    print(\"Train dist:\", pd.DataFrame(all_labels[train_index],columns=['labels']).value_counts())\n    print(\"Val dist:\", pd.DataFrame(all_labels[val_index],columns=['labels']).value_counts())\n    \n    # creating datasets\n    train_dataset = CassavaLeafDiseaseDataset(data_images=all_images[train_index],\n                                              data_labels=all_labels[train_index],\n                                              img_dir=config['DATA']['IMAGES'], \n                                              transform=train_aug_mod)\n    val_dataset = CassavaLeafDiseaseDataset(data_images=all_images[val_index],\n                                            data_labels=all_labels[val_index], \n                                            img_dir=config['DATA']['IMAGES'], \n                                            transform=valid_aug_mod)\n    \n    # Creating the dataloaders\n    train_dataloader = DataLoader(train_dataset, batch_size=config['TRAIN_BATCH_SIZE'], shuffle=True, num_workers=4)\n    val_dataloader = DataLoader(val_dataset, batch_size=config['VAL_BATCH_SIZE'], shuffle=False, num_workers=4)\n    \n    trained_model, train_loss_history, val_loss_history, train_acc_history, val_acc_history = train_model(\n        fold_no=fold,\n        model=model,\n        dataloaders={'train': train_dataloader, 'val': val_dataloader},\n        criterion=criterion,\n        optimizer=optimizer,\n        num_epochs=config['NUM_EPOCHS'],\n        device = config['DEVICE']\n    )\n    \n    # Plot train and val loss\n    epochs = np.arange(0, len(train_loss_history))\n    plt.plot(epochs, train_loss_history, 'r-')\n    plt.plot(epochs, val_loss_history, 'b-')\n    plt.legend(['Train Loss', 'Val Loss'])\n    plt.xlabel('Epoch')\n    plt.ylabel('Loss')\n#     plt.savefig(os.path.join(log_folder, str(fold)+'_train_val_loss.png'))\n    plt.show();\n    # Plot train and val accuracy\n    plt.plot(epochs, train_acc_history, 'r-')\n    plt.plot(epochs, val_acc_history, 'b-')\n    plt.legend(['Train Acc', 'Val Acc'])\n    plt.xlabel('Epoch')\n    plt.ylabel('Accuracy')\n#     plt.savefig(os.path.join(log_folder, str(fold)+'_train_val_acc.png'))\n    plt.show();\n    \n    # reload the saved pretrained imagenet model weights\n    model.load_state_dict(model_weights)\n    \n# create avg metric report to log dir\n# logs = glob.glob(f'{log_folder}/*_metric_report.json')\n# avg_label_logs = {\n#     'precision': {\n#         '0.0': 0.0,\n#         '1.0': 0.0,\n#         '2.0': 0.0,\n#         '3.0': 0.0,\n#         '4.0': 0.0\n#     },\n#     'recall': {\n#         '0.0': 0.0,\n#         '1.0': 0.0,\n#         '2.0': 0.0,\n#         '3.0': 0.0,\n#         '4.0': 0.0\n#     },\n#     'f1-score': {\n#         '0.0': 0.0,\n#         '1.0': 0.0,\n#         '2.0': 0.0,\n#         '3.0': 0.0,\n#         '4.0': 0.0\n#     },\n#     'accuracy': 0.0\n# }\n\n# for log_file_path in logs:\n#     with open(log_file_path) as log_file:\n#         log = json.load(log_file)\n\n#         for label, label_log in log.items():\n#             if label == 'macro avg' or label == 'weighted avg':\n#                 continue\n#             if label == 'accuracy':\n#                 avg_label_logs[label] += label_log\n#             else:       \n#                 avg_label_logs['precision'][label] += label_log['precision']\n#                 avg_label_logs['recall'][label] += label_log['recall']\n#                 avg_label_logs['f1-score'][label] += label_log['f1-score']\n\n# avg_label_logs['precision'] = {key: label /len(logs) for key, label in avg_label_logs['precision'].items()}\n# avg_label_logs['recall'] = {key: label /len(logs) for key, label in avg_label_logs['recall'].items()}\n# avg_label_logs['f1-score'] = {key: label /len(logs) for key, label in avg_label_logs['f1-score'].items()}\n# avg_label_logs['accuracy'] /= len(logs)\n# avg_label_logs[\"avg_precision\"] = sum(value for _, value in avg_label_logs['precision'].items()) / len(logs)\n# avg_label_logs[\"avg_recall\"] = sum(value for _, value in avg_label_logs['recall'].items()) / len(logs)\n# avg_label_logs[\"avg_f1-score\"] = sum(value for _, value in avg_label_logs['f1-score'].items()) / len(logs)\n\n# pprint.pprint(avg_label_logs)\n\n# write avg metric report to log dir\n# with open(os.path.join(log_folder, 'metric_report.json'), 'w') as fp:\n#     json.dump(avg_label_logs, fp, indent=4)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"","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}