{"cells":[{"metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"cell_type":"code","source":"# installing the torch-xla nightly version\n!curl https://raw.githubusercontent.com/pytorch/xla/master/contrib/scripts/env-setup.py -o pytorch-xla-env-setup.py\n!python pytorch-xla-env-setup.py --apt-packages libomp5 libopenblas-dev","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"d629ff2d2480ee46fbb7e2d37f6b5fab8052498a","_cell_guid":"79c7e3d0-c299-4dcb-8224-4455121ee9b0","trusted":true},"cell_type":"code","source":"# torchcontrib module for Stocastic Weight averaging \n!pip install torchcontrib","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"# installing cutblur repository\n!git clone https://github.com/clovaai/cutblur.git","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"# to see model.summary() like in keras\n!pip install torch-summary","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"# downloading the pretrained model - efficientnet b7  \n!pip install efficientnet_pytorch","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"# xla imports\n\nimport warnings\nimport torch_xla\nimport torch_xla.debug.metrics as met\nimport torch_xla.distributed.data_parallel as dp\nimport torch_xla.distributed.parallel_loader as pl\nimport torch_xla.utils.utils as xu\nimport torch_xla.core.xla_model as xm\nimport torch_xla.distributed.xla_multiprocessing as xmp","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"from cutblur.augments import cutblur\n\nfrom torch.utils.data import DataLoader\nimport matplotlib.pyplot as plt\n\nfrom PIL import Image\nimport cv2\nimport albumentations\nimport torch\nimport numpy as np\nimport io\nfrom torch.utils.data import Dataset\n\nimport torch.nn as nn\nimport torch.nn.functional as F\n\n#for Stochastic Weight Averaging in PyTorch\nfrom torchcontrib.optim import SWA\n\nfrom torchsummary import summary\n\nimport efficientnet_pytorch\n\n# required imports\nimport pandas as pd\nimport numpy as np","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"# creating object for integrating torch-xla with tpu \ndevice = xm.xla_device()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"# reading train csv data using pandas\ntrain = pd.read_csv(\"../input/cassava-leaf-disease-classification/train.csv\")\nsample_submission = pd.read_csv(\"../input/cassava-leaf-disease-classification/sample_submission.csv\")","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"len(train['label'].unique())","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"# Making the dataset class for training and testing leaf images\n\nclass leafDataset(Dataset):\n    def __init__(self, ids, classes, image_id):\n        self.ids = ids\n        self.classes = classes\n        self.image_id = image_id\n        self.aug = albumentations.Compose([\n            albumentations.Transpose(p=0.5),\n            albumentations.HorizontalFlip(p=0.5),\n            albumentations.VerticalFlip(p=0.5),\n            albumentations.ShiftScaleRotate(p=0.5),\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.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(p=0.5)\n            ]) \n        \n    def __len__(self):\n        return len(self.ids)\n    \n    def __getitem__(self, index):\n        img = np.array(Image.open('../input/cassava-leaf-disease-classification/train_images/' + self.image_id[index]))\n        img = cv2.resize(img, dsize=(256, 256), interpolation=cv2.INTER_CUBIC)\n        img = self.aug(image = img)['image']\n        img = np.transpose(img , (2,0,1)).astype(np.float32) # 2,0,1 because pytorch excepts image channel first then dimension of image\n       \n        return torch.tensor(img, dtype = torch.float),torch.tensor(self.classes[index], dtype = torch.long)\n    \n# creating object for the dataset class \ntrain_dataset = leafDataset(ids = [i for i in range(len(train))],\n                              classes = train['label'],\n                              image_id = train['image_id'])","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"idx = 115\nimg = train_dataset[idx][0]\n\nprint(train_dataset[idx][1])\nnpimg = img.numpy()\nplt.imshow(np.transpose(npimg, (1,2,0)))","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"train_sampler = torch.utils.data.distributed.DistributedSampler(\n          train_dataset,\n          num_replicas=xm.xrt_world_size(),\n          rank=xm.get_ordinal(),\n          shuffle=True)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"# setting up the validation data loader\n\nTRAIN_BATCH_SIZE = 128\n\ntraining_dataloader = DataLoader(train_dataset,\n                        num_workers=4,\n                        batch_size=TRAIN_BATCH_SIZE,\n                        sampler=train_sampler,\n                        drop_last=False)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"# increasing few layers in our model\nclass EfficientNet_b4(nn.Module):\n    def __init__(self):\n        super(EfficientNet_b4, self).__init__()\n        self.model = efficientnet_pytorch.EfficientNet.from_pretrained('efficientnet-b4')\n        self.dropout = nn.Dropout(0.1)\n        self.final_layer = nn.Linear(1792 , 5)\n        \n    def forward(self, inputs):\n        batch_size, _, _, _ = inputs.shape\n        \n        x = self.model.extract_features(inputs)\n\n        # Pooling and final linear layer\n        x = self.model._avg_pooling(x)\n        \n        x = F.adaptive_avg_pool2d(x, 1).reshape(batch_size, -1)\n        outputs = self.final_layer(self.dropout(x))\n\n        return outputs\n    \nmodel = EfficientNet_b4()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"summary(model, (3, 256, 256))","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"EPOCHS = 10\nnum_train_steps = int(len(train_dataset) / TRAIN_BATCH_SIZE / xm.xrt_world_size() * EPOCHS)\n\n# printing the no of training steps for each epoch of our training dataloader  \nxm.master_print(f'num_train_steps = {num_train_steps}, world_size={xm.xrt_world_size()}')\n\nmodel = model.to(device)\nparams = list(model.final_layer.parameters())\n\nbase_optimizer = torch.optim.Adam(params, lr= 1e-3 * 0.95 * xm.xrt_world_size())\n\n\n\noptimizer = SWA(base_optimizer, swa_start=5, swa_freq=5, swa_lr=0.05)\n\nloss_fn = torch.nn.CrossEntropyLoss()\n\nscheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer, patience = 5, verbose = True)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"# defining the training loop\ndef train_loop_fn(data_loader, model, optimizer, device, scheduler=None):\n    running_loss = 0.0\n    running_corrects = 0\n    \n    model.train()\n    \n    for inputs,labels in data_loader:\n        inputs_HQ = inputs.to(device, dtype=torch.float)\n        labels = labels.to(device, dtype=torch.float)\n\n        # or you can apply random noise, jittering, etc..\n        inputs_LQ = F.interpolate(inputs_HQ, scale_factor=1/4, mode=\"bilinear\")\n        inputs = cutblur(inputs_HQ, inputs_LQ)\n        \n        optimizer.zero_grad()\n\n        outputs = model(inputs)\n        _, preds = torch.max(outputs, 1)\n        \n        loss = loss_fn(outputs, label)\n\n        loss.backward()\n        xm.optimizer_step(optimizer)\n\n        running_loss += loss.item()\n        running_corrects += torch.sum(preds == label.data)\n            \n    train_loss = running_loss / float(len(train_data))\n    train_acc = running_corrects.double() / float(len(train_data))\n    \n    scheduler.step(train_loss)\n    \n    xm.master_print('training Loss: {:.4f} & training accuracy : {:.4f}'.format(train_loss , train_acc))","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"# training the model in _run function\ndef _run():\n    for param in model.parameters():\n        param.requires_grad = False\n    \n    for param in params:\n        param.requires_grad = True\n    \n    for epoch in range(EPOCHS):\n        xm.master_print(f\"Epoch --> {epoch+1} / {EPOCHS}\")\n        xm.master_print(f\"-------------------------------\")\n        para_loader = pl.ParallelLoader(training_dataloader, [device])\n        train_loop_fn(para_loader.per_device_loader(device), model, optimizer, device, scheduler)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"# initializing the training of model\ndef _mp_fn(rank, flags):\n    torch.set_default_tensor_type('torch.FloatTensor')\n    a = _run()\n    optimizer.swap_swa_sgd()\n    \n# applying multiprocessing so that images get paralley trained in different cores of kaggle-tpu\nFLAGS={}\nxmp.spawn(_mp_fn, args=(FLAGS,), nprocs=1, start_method='fork')","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}