{"cells":[{"metadata":{"_uuid":"d629ff2d2480ee46fbb7e2d37f6b5fab8052498a","_cell_guid":"79c7e3d0-c299-4dcb-8224-4455121ee9b0","trusted":true},"cell_type":"code","source":"import glob\n\n# separating out the filenames of training, validation and dataset separately in 3 different list variables\n\ntrain_files = glob.glob('/kaggle/input/tpu-getting-started/*/train/*.tfrec')\nval_files = glob.glob('/kaggle/input/tpu-getting-started/*/val/*.tfrec')\ntest_files = glob.glob('/kaggle/input/tpu-getting-started/*/test/*.tfrec')","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"import tensorflow as tf\n\nfor i in train_files:\n    train_image_dataset = tf.data.TFRecordDataset(i)\n\n    # Create a dictionary describing the features.\n    train_feature_description = {\n        'class': tf.io.FixedLenFeature([], tf.int64),\n        'id': tf.io.FixedLenFeature([], tf.string),\n        'image': tf.io.FixedLenFeature([], tf.string),\n    }\n\ndef _parse_image_function(example_proto):\n  # Parse the input tf.Example proto using the dictionary above.\n  return tf.io.parse_single_example(example_proto, train_feature_description)\n\ntrain_image_dataset = train_image_dataset.map(_parse_image_function)\n\n\ntrain_image_dataset","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"'''train_ids = [str(image_features['id'].numpy())[2:-1] for image_features in train_image_dataset]\n\ntrain_class = [int(image_features['class'].numpy()) for image_features in train_image_dataset]\n\ntrain_images = [image_features['image'].numpy() for image_features in train_image_dataset]'''","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"for i in val_files:\n    val_image_dataset = tf.data.TFRecordDataset(i)\n\n    # Create a dictionary describing the features.\n    val_feature_description = {\n        'class': tf.io.FixedLenFeature([], tf.int64),\n        'id': tf.io.FixedLenFeature([], tf.string),\n        'image': tf.io.FixedLenFeature([], tf.string),\n    }\n\ndef _parse_image_function(example_proto):\n  # Parse the input tf.Example proto using the dictionary above.\n  return tf.io.parse_single_example(example_proto, val_feature_description)\n\nval_image_dataset = val_image_dataset.map(_parse_image_function)\n\n\n'''val_ids = [str(image_features['id'].numpy())[2:-1] for image_features in val_image_dataset]\n\nval_class = [int(image_features['class'].numpy()) for image_features in val_image_dataset]\n\nval_images = [image_features['image'].numpy() for image_features in val_image_dataset]'''\n","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"for i in test_files:\n    test_image_dataset = tf.data.TFRecordDataset(i)\n\n    # Create a dictionary describing the features.\n    test_feature_description = {\n        'id': tf.io.FixedLenFeature([], tf.string),\n        'image': tf.io.FixedLenFeature([], tf.string),\n    }\n\ndef _parse_image_function(example_proto):\n  # Parse the input tf.Example proto using the dictionary above.\n  return tf.io.parse_single_example(example_proto, test_feature_description)\n\ntest_image_dataset = test_image_dataset.map(_parse_image_function)\n\n\ntest_ids = [str(image_features['id'].numpy())[2:-1] for image_features in test_image_dataset]\n\ntest_images = [image_features['image'].numpy() for image_features in test_image_dataset]\n\n","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"import IPython.display as display\n\ndisplay.display(display.Image(data=test_images[0]))","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"len(test_images)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"# defining dataset\n'''from PIL import Image\nimport cv2\nimport albumentations\nimport torch\nimport numpy as np\nimport io\n\nclass FlowerDataset:\n    def __init__(self, id , classes , image , img_height , img_width, mean , std , is_valid):\n        self.id = id\n        self.classes = classes\n        self.image = image\n        if is_valid == 1:\n            self.aug = albumentations.Compose([\n               albumentations.Resize(img_height , img_width, always_apply = True) ,\n               albumentations.Normalize(mean , std , always_apply = True) \n            ])\n        else:\n            self.aug = albumentations.Compose([\n                albumentations.Resize(img_height , img_width, always_apply = True) ,\n                albumentations.Normalize(mean , std , always_apply = True),\n                albumentations.ShiftScaleRotate(shift_limit = 0.0625,\n                                                scale_limit = 0.1 ,\n                                                rotate_limit = 5,\n                                                p = 0.9)\n            ]) \n        \n    def __len__(self):\n        return len(self.id)\n    \n    def __getitem__(self, index):\n        id = self.id[index]\n        img = np.array(Image.open(io.BytesIO(self.image[index]))) \n        img = cv2.resize(img, dsize=(128, 128), interpolation=cv2.INTER_CUBIC)\n        img = self.aug(image = img)['image']\n        img = np.transpose(img , (2,0,1)).astype(np.float32)\n       \n        \n        return {\n            'image' : torch.tensor(img, dtype = torch.float),\n            'class' : torch.tensor(self.classes[index], dtype = torch.long) \n        }\n    \n'''","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"'''dataset = FlowerDataset(id = train_ids, classes = train_class, image = val_images, \n                        img_height = 128 , img_width = 128, \n                        mean = (0.485, 0.456, 0.406),\n                        std = (0.229, 0.224, 0.225) , is_valid = 1)\n'''","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"# sanity check for FlowerDataset class created\n\n'''import matplotlib.pyplot as plt\n%matplotlib inline\n\nidx = 0\nimg = dataset[idx]['image']\n\nprint(dataset[idx]['class'])\n\nnpimg = img.numpy()\nplt.imshow(np.transpose(npimg, (1,2,0)))\n'''","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"!pip install efficientnet_pytorch","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"import torch.nn as nn\nimport efficientnet_pytorch\n\nclass EfficientNet(nn.Module):\n    def __init__(self):\n        super(EfficientNet, self).__init__()\n        self.base_model = efficientnet_pytorch.EfficientNet.from_pretrained('efficientnet-b7')\n        self.base_model._fc = nn.Linear(\n            in_features=2560, \n            out_features=104, \n            bias=True\n        )\n        \n    def forward(self, image, targets):\n        out = self.base_model(image)\n        loss = nn.CrossEntropyLoss()(out, targets.view(-1, 1).type_as(out))\n        return out, loss","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"# torch xla\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\n\n# import os\n# import collections\n# from datetime import datetime, timedelta\n\n# os.environ[\"XRT_TPU_CONFIG\"] = \"tpu_worker;0;10.0.0.2:8470\"\n\n# _VersionConfig = collections.namedtuple('_VersionConfig', 'wheels,server')\n# VERSION = \"torch_xla==nightly\"\n# CONFIG = {\n#     'torch_xla==nightly': _VersionConfig('nightly', 'XRT-dev{}'.format(\n#         (datetime.today() - timedelta(1)).strftime('%Y%m%d')))}[VERSION]\n\n# DIST_BUCKET = 'gs://tpu-pytorch/wheels'\n# TORCH_WHEEL = 'torch-{}-cp36-cp36m-linux_x86_64.whl'.format(CONFIG.wheels)\n# TORCH_XLA_WHEEL = 'torch_xla-{}-cp36-cp36m-linux_x86_64.whl'.format(CONFIG.wheels)\n# TORCHVISION_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":{"trusted":true},"cell_type":"code","source":"!pip install git+https://github.com/abhishekkrthakur/wtfml\n    \nfrom wtfml.utils import AverageMeter\nfrom wtfml.utils import EarlyStopping","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"from PIL import Image\nimport cv2\nimport albumentations\nimport torch\nimport numpy as np\nimport io\n\n\nfrom sklearn import metrics\nfrom sklearn import model_selection\nfrom torch.nn import functional as F\n\nimport torch_xla.core.xla_model as xm\nimport torch_xla.distributed.parallel_loader as pl\nimport torch_xla.distributed.xla_multiprocessing as xmp","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"# defining dataset\n\n'''try:\n    import torch_xla.core.xla_model as xm\n\n    _xla_available = True\nexcept ImportError:\n    _xla_available = False\n\nImageFile.LOAD_TRUNCATED_IMAGES = True\n'''\n\nclass FlowerDataset:\n    def __init__(self, ids , classes , image , augmentations = None):\n        self.ids = ids\n        self.classes = classes\n        self.image = image\n        self.augmentations = augmentations\n\n    def __len__(self):\n        return len(self.ids)\n\n    def __getitem__(self, item):\n        id = self.id[index]\n        img = np.array(Image.open(io.BytesIO(self.image[index]))) \n        img = cv2.resize(img, dsize=(128, 128), interpolation=cv2.INTER_CUBIC)\n        img = self.aug(image = img)['image']\n        img = np.transpose(img , (2,0,1)).astype(np.float32)\n       \n        \n        return {\n            'image' : torch.tensor(img, dtype = torch.float),\n            'class' : torch.tensor(self.classes[index], dtype = torch.long) \n        }\n\nclass ClassificationDataLoader:\n    def __init__(self, ids , classes , image , augmentations = None):\n        self.ids = ids\n        self.classes = classes\n        self.image = image\n        self.augmentations = augmentations\n        self.dataset = FlowerDataset(\n            ids=self.ids,\n            classes=self.classes,\n            resize=self.resize,\n            image=self.image,\n            augmentations=self.augmentations\n        )\n    \n    def fetch(self, batch_size, num_workers, drop_last=False, shuffle=True, tpu=False):\n        sampler = None\n        if tpu == True:\n            sampler = torch.utils.data.distributed.DistributedSampler(\n                self.dataset,\n                num_replicas=xm.xrt_world_size(),\n                rank=xm.get_ordinal(),\n                shuffle=shuffle\n            )\n\n        data_loader = torch.utils.data.DataLoader(\n            self.dataset,\n            batch_size=batch_size,\n            sampler=sampler,\n            drop_last=drop_last,\n            num_workers=num_workers\n        )\n        return data_loader","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"from tqdm import tqdm\nfrom wtfml.utils import AverageMeter\n\ntry:\n    import torch_xla.core.xla_model as xm\n    import torch_xla.distributed.parallel_loader as pl\n    _xla_available = True\nexcept ImportError:\n    _xla_available = False\n\ntry:\n    from apex import amp\n\n    _apex_available = True\nexcept ImportError:\n    _apex_available = False\n    \n\ndef reduce_fn(vals):\n    return sum(vals) / len(vals)\n\n\nclass Engine:\n    @staticmethod\n    def train(\n        data_loader,\n        model,\n        optimizer,\n        device,\n        scheduler=None,\n        accumulation_steps=1,\n        use_tpu=False,\n        fp16=False,\n    ):\n        if use_tpu and not _xla_available:\n            raise Exception(\n                \"You want to use TPUs but you dont have pytorch_xla installed\"\n            )\n        if fp16 and not _apex_available:\n            raise Exception(\"You want to use fp16 but you dont have apex installed\")\n        if fp16 and use_tpu:\n            raise Exception(\"Apex fp16 is not available when using TPUs\")\n        if fp16:\n            accumulation_steps = 1\n        losses = AverageMeter()\n        predictions = []\n        model.train()\n        if accumulation_steps > 1:\n            optimizer.zero_grad()\n        if use_tpu:\n            para_loader = pl.ParallelLoader(data_loader, [device])\n            tk0 = tqdm(\n                para_loader.per_device_loader(device), \n                total=len(data_loader)\n            )\n        else:\n            tk0 = tqdm(data_loader, total=len(data_loader))\n\n        for b_idx, data in enumerate(tk0):\n            for key, value in data.items():\n                data[key] = value.to(device)\n            if accumulation_steps == 1 and b_idx == 0:\n                optimizer.zero_grad()\n            _, loss = model(**data)\n\n            if not use_tpu:\n                with torch.set_grad_enabled(True):\n                    if fp16:\n                        with amp.scale_loss(loss, optimizer) as scaled_loss:\n                            scaled_loss.backward()\n                    else:\n                        loss.backward()\n                    if (b_idx + 1) % accumulation_steps == 0:\n                        optimizer.step()\n                        if scheduler is not None:\n                            scheduler.step()\n                        if b_idx > 0:\n                            optimizer.zero_grad()\n            else:\n                loss.backward()\n                xm.optimizer_step(optimizer)\n                if scheduler is not None:\n                    scheduler.step()\n                if b_idx > 0:\n                    optimizer.zero_grad()\n            if use_tpu:\n                reduced_loss = xm.mesh_reduce('loss_reduce', loss, reduce_fn)\n                losses.update(reduced_loss.item(), data_loader.batch_size)\n            else:\n                losses.update(loss.item(), data_loader.batch_size)\n            \n            tk0.set_postfix(loss=losses.avg)\n        return losses.avg","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"@staticmethod\ndef evaluate(data_loader, model, device, use_tpu=False):\n    losses = AverageMeter()\n    final_predictions = []\n    final_targets = []\n    model.eval()\n\n    with torch.no_grad():\n        if use_tpu:\n            para_loader = pl.ParallelLoader(data_loader, [device])\n            tk0 = tqdm(\n                para_loader.per_device_loader(device), \n                total=len(data_loader)\n            )\n        else:\n            tk0 = tqdm(data_loader, total=len(data_loader))\n        for b_idx, data in enumerate(tk0):\n            for key, value in data.items():\n                data[key] = value.to(device)\n            _, loss = model(**data)\n            if use_tpu:\n                reduced_loss = xm.mesh_reduce('loss_reduce', loss, reduce_fn)\n                losses.update(reduced_loss.item(), data_loader.batch_size)\n            else:\n                losses.update(loss.item(), data_loader.batch_size)\n            tk0.set_postfix(loss=losses.avg)\n    return losses.avg","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"# init model here\nMX = EfficientNet()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"import gc \n\ndef train():\n#     training_data_path = \"../input/siic-isic-224x224-images/train/\"\n#     df = pd.read_csv(\"/kaggle/working/train_folds.csv\")\n    \n    device = xm.xla_device()\n    \n    epochs = 5\n    train_bs = 32\n    valid_bs = 16\n#     fold = 0\n\n#     df_train = df[df.kfold != fold].reset_index(drop=True)\n#     df_valid = df[df.kfold == fold].reset_index(drop=True)\n\n    model = MX.to(device)\n\n    mean = (0.485, 0.456, 0.406)\n    std = (0.229, 0.224, 0.225)\n    train_aug = albumentations.Compose(\n        [\n            albumentations.Resize(128 , 128, always_apply = True) , #128 is img_height & img_width\n                albumentations.Normalize(mean , std , always_apply = True),\n                albumentations.ShiftScaleRotate(shift_limit = 0.0625,\n                                                scale_limit = 0.1 ,\n                                                rotate_limit = 5,\n                                                p = 0.9)\n        ]\n    )\n\n    valid_aug = albumentations.Compose(\n        [\n            albumentations.Resize(128 , 128, always_apply = True) ,\n            albumentations.Normalize(mean , std , always_apply = True)\n        ]\n    )\n\n#     train_images = df_train.image_name.values.tolist()\n#     train_images = [os.path.join(training_data_path, i + \".png\") for i in train_images]\n#     train_targets = df_train.target.values\n\n    train_ids = [str(image_features['id'].numpy())[2:-1] for image_features in train_image_dataset]\n\n    train_class = [int(image_features['class'].numpy()) for image_features in train_image_dataset]\n\n    train_images = [image_features['image'].numpy() for image_features in train_image_dataset]\n    \n\n#     valid_images = df_valid.image_name.values.tolist()\n#     valid_images = [os.path.join(training_data_path, i + \".png\") for i in valid_images]\n#     valid_targets = df_valid.target.values\n\n    val_ids = [str(image_features['id'].numpy())[2:-1] for image_features in val_image_dataset]\n\n    val_class = [int(image_features['class'].numpy()) for image_features in val_image_dataset]\n\n    val_images = [image_features['image'].numpy() for image_features in val_image_dataset]\n    \n#  id = train_ids, classes = train_class, image = val_images, \n#                         img_height = 128 , img_width = 128   \n\n    train_loader = ClassificationDataLoader(\n        ids = train_ids, \n        classes = train_class, \n        image = train_images, \n        augmentations=train_aug,\n    ).fetch(\n        batch_size=train_bs, \n        drop_last=True, \n        num_workers=0, \n        shuffle=True, \n        tpu=True\n    )\n\n    valid_loader = ClassificationDataLoader(\n        ids = val_ids, \n        classes = val_class, \n        image = val_images,\n        augmentations=valid_aug,\n    ).fetch(\n        batch_size=valid_bs, \n        drop_last=False, \n        num_workers=0, \n        shuffle=False, \n        tpu=True\n    )\n\n    optimizer = torch.optim.Adam(model.parameters(), lr=1e-4)\n    scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(\n        optimizer,\n        patience=3,\n        threshold=0.001,\n        mode=\"min\"\n    )\n\n    es = EarlyStopping(patience=5, mode=\"min\")\n\n    for epoch in range(epochs):\n        train_loss = Engine.train(\n            train_loader, \n            model, \n            optimizer, \n            device=device, \n            use_tpu=True)\n        \n        valid_loss = Engine.evaluate(\n            valid_loader, \n            model, \n            device=device, \n            use_tpu=True\n        )\n        xm.master_print(f\"Epoch = {epoch}, LOSS = {valid_loss}\")\n        scheduler.step(valid_loss)\n\n#         es(valid_loss, model, model_path=f\"model_fold_{fold}.bin\")\n        es(valid_loss, model)\n        if es.early_stop:\n            xm.master_print(\"Early stopping\")\n            break\n        gc.collect()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def _mp_fn(rank, flags):\n    torch.set_default_tensor_type('torch.FloatTensor')\n    a = train()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"FLAGS={}\nxmp.spawn(_mp_fn, args=(FLAGS,), nprocs=8, 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}