{"cells":[{"metadata":{"trusted":true},"cell_type":"code","source":"!pip install efficientnet-pytorch\n!git clone https://github.com/AlanChou/Truncated-Loss\n!git clone https://github.com/Bjarten/early-stopping-pytorch\n!git clone https://github.com/ufoym/imbalanced-dataset-sampler\n\nimport pandas as pd\nimport torch.optim as optim\nimport torch.nn as nn\nimport os\nimport sys\nimport torch\nimport time\nimport pandas as pd\nfrom skimage import io, transform\nimport numpy as np\nimport matplotlib.pyplot as plt\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import transforms, utils\nfrom skimage import io\nfrom sklearn.model_selection import train_test_split\nfrom torch.optim.lr_scheduler import ReduceLROnPlateau\nfrom sklearn.model_selection import train_test_split\nfrom tqdm import tqdm\nimport albumentations as A\nfrom albumentations.augmentations.transforms import Rotate\nfrom datetime import datetime\nimport cv2\n\nsys.path.append('./imbalanced-dataset-sampler/')\nsys.path.append('./early-stopping-pytorch/')\nsys.path.append('./Truncated-Loss/')\n\nfrom pytorchtools import EarlyStopping\nfrom torchsampler import ImbalancedDatasetSampler\nfrom TruncatedLoss import TruncatedLoss\nfrom efficientnet_pytorch import EfficientNet\n\n# CUDA for PyTorch\nuse_cuda = torch.cuda.is_available()\ndevice = torch.device(\"cuda:0\" if use_cuda else \"cpu\")\ntorch.backends.cudnn.benchmark = True","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Train/Test Split"},{"metadata":{"trusted":true},"cell_type":"code","source":"# Train/Test Split\ndf = pd.read_csv('../input/cassava-leaf-disease-classification/train.csv')\ntrain_split, val_split = train_test_split(df, test_size = 0.3, random_state=42, shuffle=True)\ntrain_split.to_csv('./train_split.csv', index=False)\nval_split.to_csv('./val_split.csv', index=False)","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# init Dataset"},{"metadata":{"trusted":true},"cell_type":"code","source":"class LeavesDataset(Dataset):\n\n    def __init__(self, csv_file, root_dir, transform=None, TTA=False, num_TTA=0):\n        \"\"\"\n        Args:\n            annos(string): Path to the csv file with annotations.\n            root_dir (string): Directory with all the images.\n            transform (callable, optional): Optional transform to be applied\n                on a sample.\n        \"\"\"\n        self.annos = pd.read_csv(csv_file)\n        self.root_dir = root_dir\n        self.transform = transform\n        self.TTA = TTA\n        self.num_TTA = num_TTA\n\n    def __len__(self):\n        return len(self.annos)\n\n    def __getitem__(self, idx):\n        if torch.is_tensor(idx):\n            idx = idx.tolist()\n\n        img_name = os.path.join(self.root_dir, self.annos.iloc[idx, 0])\n        image = io.imread(img_name)\n        label = self.annos.iloc[idx, 1]\n\n        if self.TTA:\n            sample = {'image': [image], 'label': label, 'idx': idx}\n        else:\n            sample = {'image': image, 'label': label, 'idx': idx}\n\n        if self.transform:\n            if self.TTA:\n                for i in range(self.num_TTA):\n                    sample['image'].append(self.transform(image = image))\n            else:\n                sample['image'] = self.transform(image = image)\n        return sample","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# **Augmentation**"},{"metadata":{"trusted":true},"cell_type":"code","source":"class ToTensor(object):\n    def __call__(self, image, force_apply=True):\n        output = image.transpose((2, 0, 1))\n        return torch.from_numpy(output)","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Visualize"},{"metadata":{"trusted":true},"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport sys\n\ntransform = A.Compose([\n    A.CenterCrop(width = 512, height=512),\n    A.Flip(p=0.5),\n    A.Normalize(mean=(0.485, 0.456, 0.406), std=(0.229, 0.224, 0.225), max_pixel_value=255.0, p=1.0),\n    ToTensor()\n])\n\nimage = cv2.imread('../input/cassava-leaf-disease-classification/train_images/100042118.jpg')\nprint(image.nbytes)\n\nimage2 = transform(image = image).cpu().detach().numpy().transpose(2,1,0)\nplt.imshow(image2)\nprint(image2.nbytes)\n#plt.imshow(image2['image'])","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Configs"},{"metadata":{"trusted":true},"cell_type":"code","source":"class Configs:\n    PATH_model = None\n    batch_size = 4\n    lr = 1e-2\n    lr_patience = 1 # can be the number of batches (if its value > 100)\n    es_patience = 2 # earlystopping\n    lr_factor = 0.1\n    alpha = 0.99\n    eps = 1e-08\n    weight_decay = 1e-6\n    momentum = 0.0\n    num_classes = 5\n    max_epochs = 10\n    betas = (0.9, 0.999)\n    model_name = 'efficientnet-b4'\n    num_workers = 1\n    train_transform = A.Compose([\n        A.CenterCrop(width=512, height=512),\n        A.Normalize(mean=(0.485, 0.456, 0.406), std=(0.229, 0.224, 0.225), max_pixel_value=255.0, p=1.0),\n        A.HorizontalFlip(p=0.5),\n        A.VerticalFlip(p=0.5),\n        ToTensor()\n    ])\n    val_transform = A.Compose([\n        A.CenterCrop(width=512, height=512),\n        A.Normalize(mean=(0.485, 0.456, 0.406), std=(0.229, 0.224, 0.225), max_pixel_value=255.0, p=1.0),\n        ToTensor()\n    ])\n    root_dir = '../input/cassava-leaf-disease-classification/train_images'\n    train_annos = './train_split.csv'\n    val_annos = './val_split.csv'\n    print_when = 500\n    TTA = True\n    num_TTA = 2","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Training / Validation Functions"},{"metadata":{"trusted":true},"cell_type":"code","source":"def validation(valloader, model, loss_fn, num_classes):   ## calculate validation loss too\n    correct = 0\n    total = 0\n    confusion_matrix = torch.zeros(num_classes, num_classes)\n    model.eval()\n    loss_fn = loss_fn.to(device)\n    with torch.no_grad():\n        loss = 0\n        for data in valloader:\n            images_batch, labels_batch = data['image'].to(device), data['label'].to(device)\n            images_batch = images_batch.float()\n            outputs_batch = model(images_batch)\n            predicted = torch.max(outputs_batch.data, 1)[1]\n            loss += loss_fn(outputs_batch, labels_batch, data['idx'])\n            for i, p in enumerate(predicted):\n                confusion_matrix[p, labels_batch[i]] += 1 \n\n            total += labels_batch.size(0)\n            correct += (predicted == labels_batch).sum().item()\n\n        validation_acc = correct/total\n        avg_validation_loss = loss/total\n\n\n    return validation_acc, avg_validation_loss, confusion_matrix\n\ndef training(trainloader, valloader, Cfgs, model, loss_fn, optimizer, scheduler):\n    PATH_model = Cfgs.PATH_model\n    epochs = Cfgs.max_epochs\t\n    print_when = Cfgs.print_when\n    num_classes = Cfgs.num_classes\n    model_name = Cfgs.model_name\n    es_patience = Cfgs.es_patience\n    lr_patience = Cfgs.lr_patience\n    loss_fn = loss_fn.to(device)\n    best_acc = 0\n    # Loop over epochs\n    print('model: {}'.format(Cfgs.model_name))\n    if PATH_model is not None:\n        print('weight: existed')\n        print('downloaded from: {}'.format(PATH_model))\n    else:\n        print('weight: not existed')\n        print('downloadded from: {}'.format('https://github.com/lukemelas/EfficientNet-PyTorch'))\n    print('Training...')\n    log_training = {'training_acc':[], 'avg_training_loss':[]}\n    log_validation = {'validation_acc': [], 'avg_validation_loss': []}\n    scheduler = ReduceLROnPlateau(optimizer, mode='max', factor=0.1, patience=lr_patience, verbose=True)\n    early_stopping = EarlyStopping(patience=es_patience, verbose=True)\n    for epoch in range(epochs):\n        running_loss = 0.0\n        correct = 0\n        total = 0\n        # Training\n        loop = tqdm(enumerate(trainloader, 0), total=len(trainloader), leave=False, bar_format='{l_bar}{bar:10}{r_bar}{bar:-10b}')\n        for i, data in loop:\n            # Transfer to GPU\n            images_batch, labels_batch = data['image'].to(device), data['label'].to(device)\n            images_batch = images_batch.float()\n\n                # zero the parameter gradients\n            optimizer.zero_grad()\n\n                # forward + backward + optimize\n            outputs_batch = model(images_batch)\n            loss = loss_fn(outputs_batch, labels_batch, data['idx'])\n            loss.backward()\n            optimizer.step()\n\n            # caluculate training accuracy\n            outputs_batch = torch.max(outputs_batch.data, 1)[1]\n            total += labels_batch.size(0)\n            correct += (outputs_batch == labels_batch).sum().item()\n            training_acc = correct/total\n\n            # update progress bar\n            loop.set_description(f\"Epoch [{epoch + 1}/{epochs}]\")\n            loop.set_postfix(loss = loss.item(), acc = training_acc)\n\n\n        avg_training_loss = running_loss/total\n        log_training['training_acc'].append(training_acc)\n        log_training['avg_training_loss'].append(running_loss/total)\t\n        print('validating...')\n        validation_acc, avg_validation_loss, confusion_matrix = validation(valloader, model, loss_fn, num_classes)\n        log_validation['validation_acc'].append(validation_acc)\n        log_validation['avg_validation_loss'].append(avg_validation_loss)\n        print('validating complated...')\n        print('-'*50)\n        print(f'Epochs: {epoch}')\n        print('-'*50)\n        print('training_acc: {:.3f}   | avg_training_loss: {:.3f}'.format(training_acc, avg_training_loss))\n        print('validation_acc: {:.3f} | avg_validation_loss: {:.3f}'.format(validation_acc, avg_validation_loss.detach().cpu().numpy()))\n        print('-'*50)\n\n        scheduler.step(training_acc)\n        if PATH_model is None:\n            PATH_model = os.path.join(current_ver, model_name + '-e' + str(epochs) + '.pt')\n\n        if validation_acc > best_acc:\n            print('saving the model at: {}'.format(PATH_model))\n            best_acc = validation_acc\n        torch.save(model.state_dict(), PATH_model)\n\n        if early_stopping.early_stop:\n            print(f'Early stopping at epoch: {epoch}')\n            break\n    # save confusion matrix\n    confusion_matrix = confusion_matrix.detach().cpu().numpy()\n    np.savetxt(os.path.join(current_ver, 'confusion_matrix.csv'), confusion_matrix, delimiter=\",\")\n\n    # load the last checkpoint with the best model\n    model.load_state_dict(torch.load(PATH_model))\n\n    print('Training complete...')\n\n    return log_training, log_validation\n\ndef pick_model(Cfgs):\n    num_classes = Cfgs.num_classes\n    PATH_model = Cfgs.PATH_model\n    model_name = Cfgs.model_name\n    if PATH_model is None:\n        model = EfficientNet.from_pretrained(model_name, num_classes = num_classes).to(device)\n        model = model.float()\n    else:\n        model = EfficientNet.from_name(model_name, num_classes = num_classes).to(device)\n        model.load_state_dict(torch.load(PATH_model))\n\n    return model","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Dataloader"},{"metadata":{"trusted":true},"cell_type":"code","source":"train_dataset = LeavesDataset(csv_file= Configs.train_annos,\n                                           root_dir= Configs.root_dir,\n                                           transform = Configs.train_transform)\n\nval_dataset = LeavesDataset(csv_file= Configs.val_annos,\n                                           root_dir= Configs.root_dir,\n                                           transform= Configs.val_transform,\n                                           TTA=Configs.TTA,\n                                           num_TTA=Configs.num_TTA)\n\ntrainloader = DataLoader(train_dataset, \n                         sampler=ImbalancedDatasetSampler(train_dataset),\n                         batch_size= Configs.batch_size, \n                         num_workers= Configs.num_workers)\n\nvalloader = DataLoader(val_dataset, \n                       sampler=ImbalancedDatasetSampler(val_dataset), \n                       batch_size= Configs.batch_size, \n                       num_workers= Configs.num_workers)  ","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Define utility functions for training"},{"metadata":{"trusted":true},"cell_type":"code","source":"# pick model\nmodel = pick_model(Configs)\n\n# loss_fn \nloss_fn = TruncatedLoss(trainset_size=len(train_dataset))\n\n# optimizer\noptimizer = optim.SGD(model.parameters(), \n                      lr=Configs.lr, \n                      momentum=Configs.momentum, \n                      weight_decay=Configs.weight_decay)\n\n# lr_scheduler\nscheduler = ReduceLROnPlateau(optimizer,\n                              mode='max',\n                              factor=Configs.lr_factor,\n                              patience=Configs.lr_patience,\n                              verbose=True)","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Training"},{"metadata":{"trusted":true},"cell_type":"code","source":"log_training, log_validation = training(trainloader, \n                                        valloader,\n                                        Configs,\n                                        model,\n                                        loss_fn,\n                                        optimizer,\n                                        scheduler)\n\ntraining_acc = log_training['training_acc']\navg_training_loss = log_training['avg_training_loss']\nvalidation_acc = log_validation['validation_acc']\navg_validation_loss = log_validation['avg_validation_loss']","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Export "},{"metadata":{"trusted":true},"cell_type":"code","source":"result = pd.DataFrame(list(zip(training_acc, avg_training_loss,\n                               validation_acc, avg_validation_loss)),\n                      columns = ['training_acc', 'training_loss', 'validation_acc', 'validation_loss'])\nresult.to_csv(os.path.join(current_ver, Configs.model_name + '-e' + str(Configs.max_epochs) + '_loss_acc.csv'),\n              index=False)","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}