{"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":"# Vision Transformers are for Vision\nIn this notebook we will use an Vision Transformers (ViT) as described in the [paper](https://arxiv.org/abs/2010.11929) \"An Image is worth 16x16 words: Transformers for Image Recognition at scale\" by _Alexey Dosovitskiy et al_ to train a model on the APTOS 2019 blindness detection contest dataset using TPU's","metadata":{}},{"cell_type":"markdown","source":"# Diabetic Retinopathy\n_Diabetic retinopathy is a diabetes complication that affects eyes. It's caused by damage to the blood vessels of the light-sensitive tissue at the back of the eye (retina)._\n\n![DR](https://www.mdpi.com/applsci/applsci-10-07274/article_deploy/html/images/applsci-10-07274-g001.png)\n\nDiabetic Retinopathy is classified into 5 categories\n1. No Diabetic Retinopathy\n2. Mild Diabetic Retinopathy\n3. Moderate Diabetic Retinopathy\n4. Severe Diabetic Retinopathy\n5. Prolifertative Diabetic Retinopathy\n\nIn this notebook we will aim to build a model to categories eye scans one of 5 aformentioned classes.","metadata":{}},{"cell_type":"markdown","source":"___","metadata":{}},{"cell_type":"markdown","source":"# A Primer on TPU and XLA","metadata":{}},{"cell_type":"markdown","source":"![TPUv3](https://techcrunch.com/wp-content/uploads/2019/05/empowering-businesses-with-google-cloud-ai_2x-1.png?w=990&crop=1)\n\nTensor Processing Units (TPUs) are Google’s custom-developed application-specific integrated circuits (ASICs) used to accelerate machine learning workloads.\n\nTPU resources accelerate the performance of linear algebra computation, which is used heavily in machine learning applications. TPUs minimize the time-to-accuracy when you train large, complex neural network models. Models that previously took weeks to train on other hardware platforms can converge in hours on TPUs.","metadata":{}},{"cell_type":"markdown","source":"**CPU**\n\n![CPU Working](https://miro.medium.com/max/640/1*ljApYgMaCiPs80uGjRlBYQ.gif)\n\n**GPU**\n\n![GPU Working](https://miro.medium.com/max/640/1*-7vLF7dzLiDSZY1edgI76w.gif)\n\n**TPU**\n\n![TPU Working](https://miro.medium.com/max/640/1*YZU6oT8QQGsWRltkJQ91SA.gif)","metadata":{}},{"cell_type":"markdown","source":"As we can see from above, the lack of writing to memory after every operation, the superscalar architecture of TPU's and other such advances allow for very high throughput and thus lower training time. You can learn more about how TPU's work [here](https://storage.googleapis.com/nexttpu/index.html)  ","metadata":{}},{"cell_type":"markdown","source":"![Pytorch-XLA](https://miro.medium.com/max/1400/1*uR7-mfI6AE6cqGqG0rDE8w.png)\n\nXLA (accelerated linear algebra) is a compiler-based linear algebra execution engine. It is the backend that powers machine learning frameworks such as TensorFlow and JAX at Google, on a variety of devices including CPUs, GPUs, and TPUs.\n\nWe can leverate TPU's using PyTorch with XLA","metadata":{}},{"cell_type":"markdown","source":"The amazing Pytorch XLA Kernels by [Abhishek Thakur](https://www.kaggle.com/abhishek) helped me put this kernel together.","metadata":{}},{"cell_type":"markdown","source":"___","metadata":{}},{"cell_type":"code","source":"%%capture\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\n!pip install timm\n!pip install onnx","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2021-10-04T11:30:04.063014Z","iopub.execute_input":"2021-10-04T11:30:04.063288Z","iopub.status.idle":"2021-10-04T11:31:21.033495Z","shell.execute_reply.started":"2021-10-04T11:30:04.063260Z","shell.execute_reply":"2021-10-04T11:31:21.032374Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\nimport matplotlib.pyplot as plt\n\nimport torch\nimport torch.nn as nn\nimport torchvision.transforms as transforms\nimport torch.onnx as onnx\n\nimport torch_xla\nimport torch_xla.core.xla_model as xm\nimport torch_xla.distributed.xla_multiprocessing as xmp\nimport torch_xla.distributed.parallel_loader as pl\n\nimport timm\nimport onnx\n\nimport gc\nimport os\nimport time\nimport random\nfrom datetime import datetime\n\nfrom PIL import Image\nfrom tqdm.notebook import tqdm\nfrom sklearn import model_selection, metrics","metadata":{"execution":{"iopub.status.busy":"2021-10-04T11:31:21.035587Z","iopub.execute_input":"2021-10-04T11:31:21.035836Z","iopub.status.idle":"2021-10-04T11:31:23.803796Z","shell.execute_reply.started":"2021-10-04T11:31:21.035802Z","shell.execute_reply":"2021-10-04T11:31:23.802744Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# For parallelization in TPUs\nos.environ[\"XLA_USE_BF16\"] = \"1\"\nos.environ[\"XLA_TENSOR_ALLOCATOR_MAXSIZE\"] = \"100000000\"","metadata":{"execution":{"iopub.status.busy":"2021-10-04T11:31:23.805363Z","iopub.execute_input":"2021-10-04T11:31:23.805644Z","iopub.status.idle":"2021-10-04T11:31:23.811500Z","shell.execute_reply.started":"2021-10-04T11:31:23.805611Z","shell.execute_reply":"2021-10-04T11:31:23.810389Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def seed_everything(seed):\n    \"\"\"\n    Seeds basic parameters for reproductibility of results\n    \n    Arguments:\n        seed {int} -- Number of the seed\n    \"\"\"\n    random.seed(seed)\n    os.environ[\"PYTHONHASHSEED\"] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\n\n\nseed_everything(1001)","metadata":{"execution":{"iopub.status.busy":"2021-10-04T11:31:23.813301Z","iopub.execute_input":"2021-10-04T11:31:23.813787Z","iopub.status.idle":"2021-10-04T11:31:23.831770Z","shell.execute_reply.started":"2021-10-04T11:31:23.813759Z","shell.execute_reply":"2021-10-04T11:31:23.830632Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# general global variables\nDATA_PATH = \"../input/aptos2019-blindness-detection\"\nTRAIN_PATH = \"../input/aptos2019-blindness-detection/train_images\"\nTEST_PATH = \"../input/aptos2019-blindness-detection/test_images\"\nMODEL_PATH = (\n    \"../input/vit-base-models-pretrained-pytorch/jx_vit_base_p16_224-80ecf9dd.pth\"\n)\n\n# model specific global variables\nIMG_SIZE = 224\nBATCH_SIZE = 32\nLR = 2e-05\nGAMMA = 0.7\nN_EPOCHS = 20","metadata":{"execution":{"iopub.status.busy":"2021-10-04T11:31:23.833529Z","iopub.execute_input":"2021-10-04T11:31:23.833825Z","iopub.status.idle":"2021-10-04T11:31:23.841100Z","shell.execute_reply.started":"2021-10-04T11:31:23.833738Z","shell.execute_reply":"2021-10-04T11:31:23.840407Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = pd.read_csv(os.path.join(DATA_PATH, \"train.csv\"))\ndf.head()","metadata":{"execution":{"iopub.status.busy":"2021-10-04T11:31:23.842338Z","iopub.execute_input":"2021-10-04T11:31:23.843203Z","iopub.status.idle":"2021-10-04T11:31:23.886179Z","shell.execute_reply.started":"2021-10-04T11:31:23.843160Z","shell.execute_reply":"2021-10-04T11:31:23.885332Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df.info()","metadata":{"execution":{"iopub.status.busy":"2021-10-04T11:31:23.887815Z","iopub.execute_input":"2021-10-04T11:31:23.888175Z","iopub.status.idle":"2021-10-04T11:31:23.913474Z","shell.execute_reply.started":"2021-10-04T11:31:23.888134Z","shell.execute_reply":"2021-10-04T11:31:23.912680Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df.diagnosis.value_counts().plot(kind='bar')","metadata":{"execution":{"iopub.status.busy":"2021-10-04T11:31:23.914732Z","iopub.execute_input":"2021-10-04T11:31:23.915426Z","iopub.status.idle":"2021-10-04T11:31:24.166748Z","shell.execute_reply.started":"2021-10-04T11:31:23.915382Z","shell.execute_reply":"2021-10-04T11:31:24.165825Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"As we can see there's huge class imbalance. So we stratify sampling data to maintain roughly equal number of examples per class in each sample","metadata":{}},{"cell_type":"code","source":"train_df, valid_df = model_selection.train_test_split(\n    df, test_size=0.1, random_state=42, stratify=df.diagnosis.values\n)","metadata":{"execution":{"iopub.status.busy":"2021-10-04T11:31:24.167993Z","iopub.execute_input":"2021-10-04T11:31:24.168296Z","iopub.status.idle":"2021-10-04T11:31:24.182949Z","shell.execute_reply.started":"2021-10-04T11:31:24.168257Z","shell.execute_reply":"2021-10-04T11:31:24.181918Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class DiabeticRetinopathyDataset(torch.utils.data.Dataset):\n    \"\"\"\n    Helper Class to create the pytorch dataset\n    \"\"\"\n\n    def __init__(self, df, data_path=DATA_PATH, mode=\"train\", transforms=None):\n        super().__init__()\n        self.df_data = df.values\n        self.data_path = data_path\n        self.transforms = transforms\n        self.mode = mode\n        self.data_dir = \"train_images\" if mode == \"train\" else \"test_images\"\n\n    def __len__(self):\n        return len(self.df_data)\n\n    def __getitem__(self, index):\n        img_name, label = self.df_data[index]\n        img_path = os.path.join(self.data_path, self.data_dir, img_name+'.png')\n        img = Image.open(img_path).convert(\"RGB\")\n\n        if self.transforms is not None:\n            image = self.transforms(img)\n\n        return image, label","metadata":{"execution":{"iopub.status.busy":"2021-10-04T11:31:24.186132Z","iopub.execute_input":"2021-10-04T11:31:24.186482Z","iopub.status.idle":"2021-10-04T11:31:24.195819Z","shell.execute_reply.started":"2021-10-04T11:31:24.186442Z","shell.execute_reply":"2021-10-04T11:31:24.194951Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# create image augmentations\ntransforms_train = transforms.Compose(\n    [\n        transforms.Resize((IMG_SIZE, IMG_SIZE)),\n        transforms.RandomHorizontalFlip(p=0.3),\n        transforms.RandomVerticalFlip(p=0.3),\n        transforms.RandomResizedCrop(IMG_SIZE),\n        transforms.ToTensor(),\n        transforms.Normalize((0.485, 0.456, 0.406), (0.229, 0.224, 0.225)),\n    ]\n)\n\n\ntransforms_valid = transforms.Compose(\n    [\n        transforms.Resize((IMG_SIZE, IMG_SIZE)),\n        transforms.ToTensor(),\n        transforms.Normalize((0.485, 0.456, 0.406), (0.229, 0.224, 0.225)),\n    ]\n)","metadata":{"execution":{"iopub.status.busy":"2021-10-04T11:31:24.197439Z","iopub.execute_input":"2021-10-04T11:31:24.198223Z","iopub.status.idle":"2021-10-04T11:31:24.212275Z","shell.execute_reply.started":"2021-10-04T11:31:24.198187Z","shell.execute_reply":"2021-10-04T11:31:24.211420Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# timm (Torch Image Models)\nprint(\"Available Vision Transformer Models: \")\ntimm.list_models(\"vit*\")","metadata":{"execution":{"iopub.status.busy":"2021-10-04T11:31:24.213640Z","iopub.execute_input":"2021-10-04T11:31:24.213912Z","iopub.status.idle":"2021-10-04T11:31:24.228985Z","shell.execute_reply.started":"2021-10-04T11:31:24.213852Z","shell.execute_reply":"2021-10-04T11:31:24.228073Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class ViTBase16(nn.Module):\n    def __init__(self, n_classes, pretrained=False):\n\n        super(ViTBase16, self).__init__()\n\n        self.model = timm.create_model(\"vit_base_patch16_224\", pretrained=False)\n        if pretrained:\n            self.model.load_state_dict(torch.load(MODEL_PATH))\n\n        self.model.head = nn.Linear(self.model.head.in_features, n_classes)\n\n    def forward(self, x):\n        x = self.model(x)\n        return x\n\n    def train_one_epoch(self, train_loader, criterion, optimizer, device):\n        # keep track of training loss\n        epoch_loss = 0.0\n        epoch_accuracy = 0.0\n\n        ###################\n        # train the model #\n        ###################\n        self.model.train()\n        for i, (data, target) in enumerate(train_loader):\n            # move tensors to GPU if CUDA is available\n            if device.type == \"cuda\":\n                data, target = data.cuda(), target.cuda()\n            elif device.type == \"xla\":\n                data = data.to(device, dtype=torch.float32)\n                target = target.to(device, dtype=torch.int64)\n\n            # clear the gradients of all optimized variables\n            optimizer.zero_grad()\n            # forward pass: compute predicted outputs by passing inputs to the model\n            output = self.forward(data)\n            # calculate the batch loss\n            loss = criterion(output, target)\n            # backward pass: compute gradient of the loss with respect to model parameters\n            loss.backward()\n            # Calculate Accuracy\n            accuracy = (output.argmax(dim=1) == target).float().mean()\n            # update training loss and accuracy\n            epoch_loss += loss\n            epoch_accuracy += accuracy\n\n            # perform a single optimization step (parameter update)\n            if device.type == \"xla\":\n                xm.optimizer_step(optimizer)\n\n                if i % 20 == 0:\n                    xm.master_print(f\"\\tBATCH {i+1}/{len(train_loader)} - LOSS: {loss}\")\n\n            else:\n                optimizer.step()\n\n        return epoch_loss / len(train_loader), epoch_accuracy / len(train_loader)\n\n    def validate_one_epoch(self, valid_loader, criterion, device):\n        # keep track of validation loss\n        valid_loss = 0.0\n        valid_accuracy = 0.0\n\n        ######################\n        # validate the model #\n        ######################\n        self.model.eval()\n        for data, target in valid_loader:\n            # move tensors to GPU if CUDA is available\n            if device.type == \"cuda\":\n                data, target = data.cuda(), target.cuda()\n            elif device.type == \"xla\":\n                data = data.to(device, dtype=torch.float32)\n                target = target.to(device, dtype=torch.int64)\n\n            with torch.no_grad():\n                # forward pass: compute predicted outputs by passing inputs to the model\n                output = self.model(data)\n                # calculate the batch loss\n                loss = criterion(output, target)\n                # Calculate Accuracy\n                accuracy = (output.argmax(dim=1) == target).float().mean()\n                # update average validation loss and accuracy\n                valid_loss += loss\n                valid_accuracy += accuracy\n\n        return valid_loss / len(valid_loader), valid_accuracy / len(valid_loader)","metadata":{"execution":{"iopub.status.busy":"2021-10-04T11:31:24.230222Z","iopub.execute_input":"2021-10-04T11:31:24.230633Z","iopub.status.idle":"2021-10-04T11:31:24.248684Z","shell.execute_reply.started":"2021-10-04T11:31:24.230602Z","shell.execute_reply":"2021-10-04T11:31:24.247810Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def fit_tpu(\n    model, epochs, device, criterion, optimizer, train_loader, valid_loader=None\n):\n\n    valid_loss_min = np.Inf  # track change in validation loss\n\n    # keeping track of losses as it happen\n    train_losses = []\n    valid_losses = []\n    train_accs = []\n    valid_accs = []\n\n    for epoch in range(1, epochs + 1):\n        gc.collect()\n        para_train_loader = pl.ParallelLoader(train_loader, [device])\n\n        xm.master_print(f\"{'='*50}\")\n        xm.master_print(f\"EPOCH {epoch} - TRAINING...\")\n        train_loss, train_acc = model.train_one_epoch(\n            para_train_loader.per_device_loader(device), criterion, optimizer, device\n        )\n        xm.master_print(\n            f\"\\n\\t[TRAIN] EPOCH {epoch} - LOSS: {train_loss}, ACCURACY: {train_acc}\\n\"\n        )\n        train_losses.append(train_loss)\n        train_accs.append(train_acc)\n        gc.collect()\n\n        if valid_loader is not None:\n            gc.collect()\n            para_valid_loader = pl.ParallelLoader(valid_loader, [device])\n            xm.master_print(f\"EPOCH {epoch} - VALIDATING...\")\n            valid_loss, valid_acc = model.validate_one_epoch(\n                para_valid_loader.per_device_loader(device), criterion, device\n            )\n            xm.master_print(f\"\\t[VALID] LOSS: {valid_loss}, ACCURACY: {valid_acc}\\n\")\n            valid_losses.append(valid_loss)\n            valid_accs.append(valid_acc)\n            gc.collect()\n\n            # save model if validation loss has decreased\n            if valid_loss <= valid_loss_min and epoch != 1:\n                xm.master_print(\n                    \"Validation loss decreased ({:.4f} --> {:.4f}).  Saving model ...\".format(\n                        valid_loss_min, valid_loss\n                    )\n                )\n            #                 xm.save(model.state_dict(), 'best_model.pth')\n\n            valid_loss_min = valid_loss\n\n    return {\n        \"train_loss\": train_losses,\n        \"valid_losses\": valid_losses,\n        \"train_acc\": train_accs,\n        \"valid_acc\": valid_accs,\n    }","metadata":{"execution":{"iopub.status.busy":"2021-10-04T11:31:24.249952Z","iopub.execute_input":"2021-10-04T11:31:24.250345Z","iopub.status.idle":"2021-10-04T11:31:24.265462Z","shell.execute_reply.started":"2021-10-04T11:31:24.250318Z","shell.execute_reply":"2021-10-04T11:31:24.264807Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = ViTBase16(n_classes=5, pretrained=True)","metadata":{"execution":{"iopub.status.busy":"2021-10-04T11:31:24.266750Z","iopub.execute_input":"2021-10-04T11:31:24.267085Z","iopub.status.idle":"2021-10-04T11:31:30.097912Z","shell.execute_reply.started":"2021-10-04T11:31:24.267059Z","shell.execute_reply":"2021-10-04T11:31:30.097120Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def _run():\n    train_dataset = DiabeticRetinopathyDataset(train_df, transforms=transforms_train)\n    valid_dataset = DiabeticRetinopathyDataset(valid_df, transforms=transforms_valid)\n\n    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,\n    )\n\n    valid_sampler = torch.utils.data.distributed.DistributedSampler(\n        valid_dataset,\n        num_replicas=xm.xrt_world_size(),\n        rank=xm.get_ordinal(),\n        shuffle=False,\n    )\n\n    train_loader = torch.utils.data.DataLoader(\n        dataset=train_dataset,\n        batch_size=BATCH_SIZE,\n        sampler=train_sampler,\n        drop_last=True,\n        num_workers=8,\n    )\n\n    valid_loader = torch.utils.data.DataLoader(\n        dataset=valid_dataset,\n        batch_size=BATCH_SIZE,\n        sampler=valid_sampler,\n        drop_last=True,\n        num_workers=8,\n    )\n\n    criterion = nn.CrossEntropyLoss()\n    #     device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n    device = xm.xla_device()\n    model.to(device)\n\n    lr = LR * xm.xrt_world_size()\n    optimizer = torch.optim.Adam(model.parameters(), lr=lr)\n\n    xm.master_print(f\"INITIALIZING TRAINING ON {xm.xrt_world_size()} TPU CORES\")\n    start_time = datetime.now()\n    xm.master_print(f\"Start Time: {start_time}\")\n\n    logs = fit_tpu(\n        model=model,\n        epochs=N_EPOCHS,\n        device=device,\n        criterion=criterion,\n        optimizer=optimizer,\n        train_loader=train_loader,\n        valid_loader=valid_loader,\n    )\n\n    xm.master_print(f\"Execution time: {datetime.now() - start_time}\")\n\n    xm.master_print(\"Saving Model\")\n    xm.save(\n        model.state_dict(), f'model_5e_{datetime.now().strftime(\"%Y%m%d-%H%M\")}.pth'\n    )","metadata":{"execution":{"iopub.status.busy":"2021-10-04T11:31:30.099220Z","iopub.execute_input":"2021-10-04T11:31:30.099775Z","iopub.status.idle":"2021-10-04T11:31:30.112216Z","shell.execute_reply.started":"2021-10-04T11:31:30.099733Z","shell.execute_reply":"2021-10-04T11:31:30.111313Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Start training processes\ndef _mp_fn(rank, flags):\n    torch.set_default_tensor_type(\"torch.FloatTensor\")\n    a = _run()\n\n\n# _run()\nFLAGS = {}\nxmp.spawn(_mp_fn, args=(FLAGS,), nprocs=8, start_method=\"fork\")","metadata":{"execution":{"iopub.status.busy":"2021-10-04T11:31:30.113555Z","iopub.execute_input":"2021-10-04T11:31:30.113869Z","iopub.status.idle":"2021-10-04T12:38:13.753344Z","shell.execute_reply.started":"2021-10-04T11:31:30.113830Z","shell.execute_reply":"2021-10-04T12:38:13.750276Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install onnx\nimport torch.onnx as onnx\ntorch.save(model,'ViT4Vision.pth')\n\ninput_image = torch.zeros((1,3,224,224))\nonnx.export(model, input_image, 'ViT4Vision.onnx')","metadata":{"execution":{"iopub.status.busy":"2021-10-04T12:38:13.764082Z","iopub.execute_input":"2021-10-04T12:38:13.765852Z","iopub.status.idle":"2021-10-04T12:38:29.413545Z","shell.execute_reply.started":"2021-10-04T12:38:13.765789Z","shell.execute_reply":"2021-10-04T12:38:29.412707Z"},"trusted":true},"execution_count":null,"outputs":[]}]}