{"cells":[{"metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"cell_type":"code","source":"!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 --version nightly --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":"!mkdir -p /tmp/pip/cache/\n!cp ../input/segmentationmodelspytorch/segmentation_models/efficientnet_pytorch-0.6.3.xyz /tmp/pip/cache/efficientnet_pytorch-0.6.3.tar.gz\n!cp ../input/segmentationmodelspytorch/segmentation_models/pretrainedmodels-0.7.4.xyz /tmp/pip/cache/pretrainedmodels-0.7.4.tar.gz\n!cp ../input/segmentationmodelspytorch/segmentation_models/segmentation-models-pytorch-0.1.2.xyz /tmp/pip/cache/segmentation_models_pytorch-0.1.2.tar.gz\n!cp ../input/segmentationmodelspytorch/segmentation_models/timm-0.1.20-py3-none-any.whl /tmp/pip/cache/\n!cp ../input/segmentationmodelspytorch/segmentation_models/timm-0.2.1-py3-none-any.whl /tmp/pip/cache/\n!pip install --no-index --find-links /tmp/pip/cache/ efficientnet-pytorch\n!pip install --no-index --find-links /tmp/pip/cache/ segmentation-models-pytorch","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"import os, gc\nimport numpy as np\nimport pandas as pd\nfrom tqdm.notebook import tqdm\n\nimport tifffile as tiff\nimport matplotlib.pyplot as plt\nimport cv2\nimport random\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch import optim\nfrom torch.utils.data import Dataset, DataLoader\nimport torchvision\nfrom torchvision import transforms\nfrom torch.optim.lr_scheduler import CosineAnnealingWarmRestarts\n\nfrom albumentations import *\n\nfrom sklearn.model_selection import StratifiedKFold\n\nfrom segmentation_models_pytorch.unet import Unet\nfrom segmentation_models_pytorch.encoders import get_preprocessing_fn\n\nimport torch_xla\nimport torch_xla.core.xla_model as xm\nfrom torch.utils.data.distributed import DistributedSampler\nimport torch_xla.distributed.parallel_loader as pl\nimport torch_xla.distributed.xla_multiprocessing as xmp\n\nimport warnings\nwarnings.filterwarnings(\"ignore\")","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"import timm\nimport pretrainedmodels","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"torch.__version__","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def set_all_seeds(seed):\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    \nset_all_seeds(2021)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"DATA_PATH = '/kaggle/input/cassava-leaf-disease-classification'\nos.listdir(DATA_PATH)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"PATH_TRAIN = os.path.join(DATA_PATH, 'train_images')\nPATH_TEST = os.path.join(DATA_PATH, 'test_images')\n\nprint(f'No. of training images: {len(os.listdir(PATH_TRAIN))}')\nprint(f'No. of testing images: {len(os.listdir(PATH_TEST))}')","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"LIST_TRAIN = os.listdir(PATH_TRAIN)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"os.listdir(PATH_TRAIN)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"class CassavaLeafDataset(Dataset):\n    def __init__(self, \n                 df,\n                 data_path, \n                 preprocess_input = None,\n                 transforms = None,\n                 is_test = False\n                ):\n        \n        self.df = df.reset_index(drop=True).copy()\n        self.preprocess_input = preprocess_input\n        self.transforms = transforms\n        self.is_test = is_test\n        \n        if is_test == True:\n            self.image_ids = df['image_id'].values\n            self.data_path = data_path + '/test_images'\n            \n        else:\n            self.image_ids = df['image_id'].values\n            self.labels = self.df['label'].values\n            self.data_path = data_path + '/train_images'\n                \n        \n    def __len__(self):\n        return self.df.shape[0]\n        #return len(self.fnames)\n    \n    def __getitem__(self, idx):\n        img_path = os.path.join(self.data_path, self.image_ids[idx])\n        \n        img = cv2.cvtColor(cv2.imread(img_path), cv2.COLOR_BGR2RGB)\n        \n        if self.transforms: \n            img = self.transforms(image=img)[\"image\"]\n            \n            \n        # not sure if preprocess_input and img.transpose is necessary?\n        ####################################################\n        #if self.preprocess_input:\n        #    img = self.preprocess_input(image=img)['image']\n            \n        img = img.transpose((2, 0, 1))\n        ####################################################\n        \n        img = torch.from_numpy(img)\n        \n        label = self.labels[idx]\n        \n        if self.is_test:\n            return img,\n        else:\n            return img, label","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"fnames = np.array(os.listdir(PATH_TRAIN))\n\nskf = StratifiedKFold(n_splits = 5, random_state=0, shuffle=True)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"groups","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"train_df = pd.read_csv('../input/cassava-leaf-disease-classification/train.csv')","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"t_df.loc[[1,2]]","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"ENCODER_NAME = \"timm-efficientnet-b4\"\n\npreprocessing_fn = Lambda(image=get_preprocessing_fn(encoder_name=ENCODER_NAME,\n                                                     pretrained = 'imagenet'\n                                                    ))\n\ntransforms = Compose([HorizontalFlip(p=0.5),\n                      VerticalFlip(p=0.5),\n                      ShiftScaleRotate(p=0.5),\n                      OneOf([OpticalDistortion(p=0.3),\n                             GridDistortion(p=0.1),\n                             IAAPiecewiseAffine(p=0.3)\n                            ], p=0.3),\n                      OneOf([HueSaturationValue(10,15,10),\n                             CLAHE(clip_limit=2),\n                             RandomBrightnessContrast()\n                            ], p=0.3),\n                      Normalize(mean=[0.485, 0.456, 0.406], \n                                std=[0.229, 0.224, 0.225], \n                                max_pixel_value=255.0, p=1.0),\n                     ], p=1.0)\n","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"class CassavaLeafModel(nn.Module):\n    def __init__(self, model_arch, n_classes, pretrained=False):\n        super(CassavaLeafModel, self).__init__()\n        self.model = timm.create_model(model_arch, pretrained=pretrained)\n        n_features = self.model.classifier.in_features\n        self.model.classifier = nn.Linear(n_features, n_classes)\n        \n    def forward(self, images):\n        x = self.model(images)\n        return x","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def reduce(values):\n    return sum(values)/len(values)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"demo_model = CassavaLeafModel('tf_efficientnet_b4_ns', 5, True)\n\ndemo_optim = optim.SGD(demo_model.parameters(),\n                       lr = 0.3,\n                       momentum = 0.9\n                      )\n\nDEMO_EPOCHS = 60\n\ndemo_sched = CosineAnnealingWarmRestarts(demo_optim,\n                                         T_0=DEMO_EPOCHS//3,\n                                         T_mult=1,\n                                         eta_min=0,\n                                         last_epoch=-1,\n                                         verbose=False\n                                        )\n\nlrs = []\n\nfor i in range(DEMO_EPOCHS):\n    demo_optim.step()\n    lrs.append(demo_optim.param_groups[0][\"lr\"])\n    demo_sched.step()\n    \nplt.plot(lrs)\nplt.show()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def train_one_epoch(epoch_no,\n                    data_loader,\n                    model,\n                    optimizer,\n                    device,\n                    loss_fn,\n                    scheduler=None\n                   ):\n    model.train()\n    losses = []\n    \n    for i, (image_batch, label_batch) in enumerate(data_loader):\n        image_batch = image_batch.to(device)\n        label_batch = label_batch.to(device)\n        \n        optimizer.zero_grad()\n        \n        pred_label_batch = model(image_batch.float())\n        \n        loss = loss_fn(pred_label_batch, label_batch.float())\n        \n        loss.backward()\n        \n        \n        xm.optimizer_step(optimizer)\n        if scheduler is not None:\n            scheduler.step()\n            \n        loss_reduced = xm.mesh_reduce('train_loss_reduce',\n                                      loss,\n                                      reduce\n                                     )\n        \n        losses.append(loss_reduced.item())\n        \n        del image_batch, pred_label_batch, label_batch\n        gc.collect()\n        \n    xm.master_print(f'{epoch_no+1} - Loss : {reduce(losses): .4f}')\n    \n    \ndef eval_fn(data_loader, model, device, loss_fn):\n    \n    model.eval()\n    \n    losses = []\n    \n    for i, (image_batch, label_batch) in enumerate(data_loader):\n        image_batch, label_batch = image_batch.to(device), label_batch.to(device)\n        \n        pred_label_batch = model(image_batch.float())\n        \n        loss = loss_fn(pred_label_batch, label_batch.float())\n        \n        losses.append(xm.mesh_reduce('val_loss_reduce',\n                                     loss,\n                                     reduce\n                                    ).item())\n        \n    total_loss = reduce(losses)\n    \n    return total_loss","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def _mp_fn(rank, flags):\n    \n    device = xm.xla_device()\n    \n    train_sampler = DistributedSampler(dataset=flags['TRAIN_DS'],\n                                       num_replicas=xm.xrt_world_size(),\n                                       rank=xm.get_ordinal(),\n                                       shuffle=True\n                                      )\n    \n    train_dl = DataLoader(dataset=flags['TRAIN_DS'],\n                          batch_size=flags['BATCH_SIZE'],\n                          sampler=train_sampler,\n                          num_workers=0\n                         )\n    \n    val_sampler = DistributedSampler(dataset=flags['VAL_DS'],\n                                     num_replicas=xm.xrt_world_size(),\n                                     rank=xm.get_ordinal(),\n                                     shuffle=False\n                                    )\n    \n    val_dl = DataLoader(dataset=flags['VAL_DS'],\n                        batch_size=flags['BATCH_SIZE'],\n                        sampler=val_sampler,\n                        num_workers=0\n                       )\n    \n    del train_sampler, val_sampler\n    gc.collect()\n    \n    \n    print(\"Process is using\", xm.xla_real_devices([str(device)])[0])\n    \n    fold_model = flags['FOLD_MODEL']\n    fold_model.to(device)\n    \n    lr = flags['LR']\n    \n    optimizer = optim.SGD(model.parameters(),\n                          lr=lr,\n                          momentum=0.9\n                         )\n    \n    scheduler = CosineAnnealingWarmRestarts(optimizer,\n                                            T_0=flags['EPOCHS']//3,\n                                            T_mult=1,\n                                            eta_min=0,\n                                            last_epoch=-1,\n                                            verbose=False\n                                           )\n    \n    xm.master_print('Training now...')\n    \n    for e_no, epoch in enumerate(range(flags['EPOCHS'])):\n        \n        train_para_loader = pl.ParallelLoader(train_dl,\n                                              [device]\n                                             ).per_device_loader(device)\n        \n        train_one_epoch(e_no,\n                        train_para_loader,\n                        fold_model,\n                        optimizer,\n                        device,\n                        flags['LOSS_FN'],\n                        scheduler\n                       )\n        \n        del train_para_loader\n        gc.collect()\n        \n    xm.master_print('\\nValidating now...')\n    \n    val_para_loader = pl.ParallelLoader(val_dl,\n                                        [device]\n                                       ).per_device_loader(device)\n    \n    loss = eval_fn(val_para_loader,\n                   fold_model,\n                   device,\n                   flags['LOSS_FN']\n                  )\n    \n    del val_para_loader\n    gc.collect()\n    \n    xm.master_print(f'Val Loss : {loss: .4f}')\n    \n    xm.save(fold_model.state_dict(), f\"8core_fold_model{flags['FOLD_NO']}.pth\")","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"for fold, (t_idx, v_idx) in enumerate(skf.split(train_df['image_id'], train_df['label'])):\n    \n    print(f'Fold: {fold+1}')\n    print('-' * 40)\n    \n    t_df = train_df.loc[t_idx]\n    v_df = train_df.loc[v_idx]\n    \n    train_ds = CassavaLeafDataset(data_path=DATA_PATH,\n                                  df=t_df,\n                                  preprocess_input=preprocessing_fn,\n                                  transforms=transforms,\n                                  is_test=False\n                                 )\n    \n    val_ds = CassavaLeafDataset(data_path=DATA_PATH,\n                                df=v_df,\n                                preprocess_input=preprocessing_fn,\n                                transforms=None,\n                                is_test=False\n                               )\n    \n\n    lr = 0.3\n    loss_fn = nn.CrossEntropyLoss()\n    \n    FLAGS = {'FOLD_NO': fold,\n             'TRAIN_DS': train_ds,\n             'VAL_DS': val_ds,\n             'LR': lr,\n             'BATCH_SIZE': 32,\n             'EPOCHS': 30,\n             'MODEL_ARCH': 'tf_efficientnet_b4_ns',\n             'LOSS_FN': loss_fn\n            }\n    \n    model = CassavaLeafModel(model_arch=FLAGS['MODEL_ARCH'], \n                             n_classes=5, \n                             pretrained=True)\n    \n    FLAGS['FOLD_MODEL'] = model.float()\n    \n    model.float()\n    \n    xmp.spawn(fn=_mp_fn,\n              args=(FLAGS,),\n              nprocs=8)\n    \n    print('\\n')","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"!ls","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"!--distributed-world-size=8","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"","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}