{"cells":[{"metadata":{},"cell_type":"markdown","source":"## Using all 8 cores of a v3 TPU using pytorch/XLA"},{"metadata":{},"cell_type":"markdown","source":"## setup for pytorch/xla on TPU\nthanks to this [kernel](https://www.kaggle.com/byrachonok/pytorch-xla-for-tpuv3)"},{"metadata":{"trusted":true,"_kg_hide-output":true},"cell_type":"code","source":"import os\nimport collections\nfrom datetime import datetime, timedelta\n\nos.environ[\"XRT_TPU_CONFIG\"] = \"tpu_worker;0;10.0.0.2:8470\"\n\n_VersionConfig = collections.namedtuple('_VersionConfig', 'wheels,server')\nVERSION = \"torch_xla==nightly\"\nCONFIG = {\n    'torch_xla==nightly': _VersionConfig('nightly', 'XRT-dev{}'.format(\n        (datetime.today() - timedelta(1)).strftime('%Y%m%d')))}[VERSION]\n\nDIST_BUCKET = 'gs://tpu-pytorch/wheels'\nTORCH_WHEEL = 'torch-{}-cp36-cp36m-linux_x86_64.whl'.format(CONFIG.wheels)\nTORCH_XLA_WHEEL = 'torch_xla-{}-cp36-cp36m-linux_x86_64.whl'.format(CONFIG.wheels)\nTORCHVISION_WHEEL = 'torchvision-{}-cp36-cp36m-linux_x86_64.whl'.format(CONFIG.wheels)\n\n!export LD_LIBRARY_PATH=/usr/local/lib:$LD_LIBRARY_PATH\n!apt-get install libomp5 -y\n!apt-get install libopenblas-dev -y\n\n!pip uninstall -y torch torchvision\n!gsutil cp \"$DIST_BUCKET/$TORCH_WHEEL\" .\n!gsutil cp \"$DIST_BUCKET/$TORCH_XLA_WHEEL\" .\n!gsutil cp \"$DIST_BUCKET/$TORCHVISION_WHEEL\" .\n!pip install \"$TORCH_WHEEL\"\n!pip install \"$TORCH_XLA_WHEEL\"\n!pip install \"$TORCHVISION_WHEEL\"","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"## Imports"},{"metadata":{"trusted":true},"cell_type":"code","source":"import os\nimport re\nimport cv2\nimport time\nimport tensorflow\nimport collections\nimport numpy as np\nimport pandas as pd\nfrom tqdm import tqdm\nfrom glob import glob\nfrom PIL import Image\nimport requests, threading\nimport matplotlib.pyplot as plt\nfrom datetime import datetime, timedelta\n\nimport torch\nimport torchvision\nimport torch.nn as nn\nimport torch.optim as optim\nimport torch.nn.functional as F\nfrom torchvision import datasets\nfrom torchvision import transforms\nfrom torch.autograd import Variable\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.optim.lr_scheduler import OneCycleLR\n\nimport torch_xla\nimport torch_xla.utils.utils as xu\nimport torch_xla.core.xla_model as xm\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.distributed.xla_multiprocessing as xmp\n\nimport warnings\nwarnings.filterwarnings(\"ignore\")\n\ntorch.manual_seed(42)\ntorch.set_default_tensor_type('torch.FloatTensor')","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"# do not uncomment see https://github.com/pytorch/xla/issues/1587\n\n# xm.get_xla_supported_devices()\n# xm.xrt_world_size() # 1","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"## Dataset"},{"metadata":{"trusted":true},"cell_type":"code","source":"DATASET_DIR = '/kaggle/input/104-flowers-garden-of-eden/jpeg-512x512'\nTRAIN_DIR  = DATASET_DIR + '/train'\nVAL_DIR  = DATASET_DIR + '/val'\nTEST_DIR  = DATASET_DIR + '/test'\nBATCH_SIZE = 16 # per core \nNUM_EPOCH = 25","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"normalize = transforms.Normalize(mean=[0.485, 0.456, 0.406],std=[0.229, 0.224, 0.225])\n\ntrain_transform = transforms.Compose([transforms.RandomResizedCrop(224),\n                                      transforms.RandomHorizontalFlip(0.5),\n                                      transforms.ToTensor(),\n                                      normalize])\n\nvalid_transform = transforms.Compose([transforms.Resize((224,224)),\n                                      transforms.ToTensor(),\n                                      normalize])","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"train = datasets.ImageFolder(TRAIN_DIR, transform=train_transform)\nvalid = datasets.ImageFolder(VAL_DIR, transform=valid_transform)\n\n# print out some data stats\nprint('Num training images: ', len(train))\nprint('Num test images: ', len(valid))","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"## Model"},{"metadata":{"trusted":true},"cell_type":"code","source":"class MyModel(nn.Module):\n\n    def __init__(self):\n        super(MyModel, self).__init__()\n        \n        self.base_model = torchvision.models.densenet201(pretrained=True)\n        self.base_model.classifier = nn.Identity()\n        self.fc = torch.nn.Sequential(\n                    torch.nn.Linear(1920, 1024, bias = True),\n                    torch.nn.BatchNorm1d(1024),\n                    torch.nn.ReLU(inplace=True),\n                    torch.nn.Dropout(0.3),\n                    torch.nn.Linear(1024, 512, bias = True),\n                    torch.nn.BatchNorm1d(512),\n                    torch.nn.ReLU(inplace=True),\n                    torch.nn.Dropout(0.3),\n                    torch.nn.Linear(512, 104))\n        \n    def forward(self, inputs):\n        x = self.base_model(inputs)\n        return self.fc(x)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_kg_hide-output":true},"cell_type":"code","source":"model = MyModel()\nprint(model)\ndel model","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"## Training"},{"metadata":{"trusted":true},"cell_type":"code","source":"def train_model():\n    global train, valid\n    \n    torch.manual_seed(42)\n    \n    train_sampler = torch.utils.data.distributed.DistributedSampler(\n        train,\n        num_replicas=xm.xrt_world_size(),\n        rank=xm.get_ordinal(),\n        shuffle=True)\n    \n    train_loader = torch.utils.data.DataLoader(\n        train,\n        batch_size=BATCH_SIZE,\n        sampler=train_sampler,\n        num_workers=0,\n        drop_last=True) # print(len(train_loader))\n    \n    valid_loader = torch.utils.data.DataLoader(\n        valid,\n        batch_size=BATCH_SIZE,\n        shuffle=False,\n        num_workers=0,\n        drop_last=True)\n    \n    xm.master_print(f\"Train for {len(train_loader)} steps per epoch\")\n    # Scale learning rate to num cores\n    learning_rate = 0.0001 * xm.xrt_world_size()\n\n    # Get loss function, optimizer, and model\n    device = xm.xla_device()\n\n    model = MyModel()\n    \n    for param in model.base_model.parameters(): # freeze some layers\n        param.requires_grad = False\n    \n    model = model.to(device)\n    loss_fn =  nn.CrossEntropyLoss()\n    optimizer = optim.Adam(model.parameters(), lr=learning_rate, weight_decay=5e-4)\n    scheduler = OneCycleLR(optimizer, \n                           learning_rate, \n                           div_factor=10.0, \n                           final_div_factor=50.0, \n                           epochs=NUM_EPOCH,\n                           steps_per_epoch=len(train_loader))\n    \n    \n    \n    def train_loop_fn(loader):\n        tracker = xm.RateTracker()\n        model.train()\n        for x, (data, target) in enumerate(loader):\n            optimizer.zero_grad()\n            output = model(data)\n            loss = loss_fn(output, target)\n            loss.backward()\n            xm.optimizer_step(optimizer)\n            tracker.add(data.shape[0])\n            scheduler.step()\n            if x % 30 == 0:\n                print('[xla:{}]({})\\tLoss={:.3f}\\tRate={:.2f}\\tGlobalRate={:.2f}'.format(\n                    xm.get_ordinal(), x, loss.item(), tracker.rate(),\n                    tracker.global_rate()), flush=True)\n\n    def test_loop_fn(loader):\n        with torch.no_grad():\n            total_samples, correct = 0, 0\n            model.eval()\n            for data, target in loader:\n                output = model(data)\n                pred = output.max(1, keepdim=True)[1]\n                correct += pred.eq(target.view_as(pred)).sum().item()\n                total_samples += data.size()[0]\n            accuracy = 100.0 * correct / total_samples\n            print('[xla:{}] Accuracy={:.2f}%'.format(xm.get_ordinal(), accuracy), flush=True)\n            model.train()\n        return accuracy\n\n    # Train and eval loops\n    accuracy = []\n    for epoch in range(1, NUM_EPOCH + 1):\n        start = time.time()\n        para_loader = pl.ParallelLoader(train_loader, [device])\n        train_loop_fn(para_loader.per_device_loader(device))\n        para_loader = pl.ParallelLoader(valid_loader, [device])\n        accuracy.append(test_loop_fn(para_loader.per_device_loader(device)))\n        xm.master_print(\"Finished training epoch {} acc {:.2f} in {:.2f} sec\"\\\n                        .format(epoch, accuracy[-1], time.time() - start))        \n        xm.save(model.state_dict(), \"./model.pt\")\n        \n#         if epoch == 15: #unfreeze\n#                 for param in model.base_model.parameters():\n#                     param.requires_grad = True\n\n    return accuracy","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"# Start training processes\ndef _mp_fn(rank, flags):\n    global acc_list\n    torch.set_default_tensor_type('torch.FloatTensor')\n    res = train_model()\n\nFLAGS={}\nxmp.spawn(_mp_fn, args=(FLAGS,), nprocs=8, start_method='fork')","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"### Doing validation every epoch increases time per epoch significantly. faster version of this notebook is in comments."},{"metadata":{"trusted":true},"cell_type":"code","source":"","execution_count":null,"outputs":[]}],"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.6.4","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"}},"nbformat":4,"nbformat_minor":4}