{"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":{"trusted":true},"cell_type":"code","source":"!pip install git+https://github.com/abhishekkrthakur/wtfml\n!pip install efficientnet_pytorch","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"import gc\nimport os\nimport torch\nimport albumentations\n\nimport numpy as np\nimport pandas as pd\n\nimport torch.nn as nn\nfrom sklearn import metrics\nfrom sklearn import model_selection\nfrom torch.nn import functional as F\n\nfrom wtfml.utils import EarlyStopping\nfrom wtfml.data_loaders.image import ClassificationDataLoader\n\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\n\nimport efficientnet_pytorch","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"class EfficientNet(nn.Module):\n    def __init__(self):\n        super(EfficientNet, self).__init__()\n        self.base_model = efficientnet_pytorch.EfficientNet.from_pretrained('efficientnet-b0')\n        self.base_model._fc = nn.Linear(\n            in_features=1280, \n            out_features=1, \n            bias=True\n        )\n        \n    def forward(self, image, targets):\n        out = self.base_model(image)\n        loss = nn.BCEWithLogitsLoss()(out, targets.view(-1, 1).type_as(out))\n        return out, loss","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"# create folds\ndf = pd.read_csv(\"../input/siim-isic-melanoma-classification/train.csv\")\ndf[\"kfold\"] = -1    \ndf = df.sample(frac=1).reset_index(drop=True)\ny = df.target.values\nkf = model_selection.StratifiedKFold(n_splits=5)\n\nfor f, (t_, v_) in enumerate(kf.split(X=df, y=y)):\n    df.loc[v_, 'kfold'] = f\n\ndf.to_csv(\"train_folds.csv\", index=False)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"import torch\nfrom 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\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        bs=1\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        data_loader = pl.ParallelLoader(data_loader, [device]).per_device_loader(device)\n \n        for b_idx, data in enumerate(data_loader):\n            optimizer.zero_grad()\n            _, loss = model(**data)\n\n            loss.backward()\n            xm.optimizer_step(optimizer)\n            if scheduler is not None:\n                scheduler.step()\n            reduced_loss = xm.mesh_reduce('loss_reduce', loss, reduce_fn)\n            losses.update(reduced_loss.item(), bs)\n\n        return losses.avg\n\n    @staticmethod\n    def evaluate(data_loader, model, device, use_tpu=False, bs=1):\n        losses = AverageMeter()\n        final_predictions = []\n        final_targets = []\n        model.eval()\n        with torch.no_grad():\n            data_loader = pl.ParallelLoader(data_loader, [device]).per_device_loader(device)\n            for b_idx, data in enumerate(data_loader):\n                _, loss = model(**data)\n                reduced_loss = xm.mesh_reduce('loss_reduce', loss, reduce_fn)\n                losses.update(reduced_loss.item(), bs)\n\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":"def train():\n    training_data_path = \"../input/siic-isic-224x224-images/train/\"\n    df = pd.read_csv(\"/kaggle/working/train_folds.csv\")\n    device = xm.xla_device()\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.Normalize(mean, std, max_pixel_value=255.0, always_apply=True),\n            albumentations.ShiftScaleRotate(shift_limit=0.0625, scale_limit=0.1, rotate_limit=15),\n            albumentations.Flip(p=0.5)\n        ]\n    )\n\n    valid_aug = albumentations.Compose(\n        [\n            albumentations.Normalize(mean, std, max_pixel_value=255.0, 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    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    train_loader = ClassificationDataLoader(\n        image_paths=train_images,\n        targets=train_targets,\n        resize=None,\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        image_paths=valid_images,\n        targets=valid_targets,\n        resize=None,\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            bs=train_bs)\n        \n        valid_loss = Engine.evaluate(\n            valid_loader, \n            model, \n            device=device, \n            use_tpu=True,\n            bs=valid_bs\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        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}