{"cells":[{"metadata":{},"cell_type":"markdown","source":"# Toxic Comment Classification\n\nReference: https://www.kaggle.com/tanlikesmath/xlm-roberta-pytorch-xla-tpu","execution_count":null},{"metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"cell_type":"code","source":"# Install PyTorch/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 --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":"import os\nos.environ['XLA_USE_BF16']=\"1\"\nos.environ['XLA_TENSOR_ALLOCATOR_MAXSIZE'] = '100000000'\n\nimport torch\nimport pandas as pd\nfrom scipy import stats\nimport numpy as np\n\nfrom tqdm import tqdm\nfrom collections import OrderedDict, namedtuple\nimport torch.nn as nn\nfrom torch.optim import lr_scheduler\nimport joblib\n\nimport logging\nimport transformers\nfrom transformers import AdamW, get_linear_schedule_with_warmup, get_constant_schedule, XLMRobertaTokenizer, XLMRobertaModel, XLMRobertaConfig\nimport sys\nfrom sklearn import metrics, model_selection\n\nimport warnings\nimport torch_xla\nimport torch_xla.debug.metrics as met\nimport torch_xla.distributed.parallel_loader as pl\nimport torch_xla.utils.utils as xu\nimport torch_xla.core.xla_model as xm\nimport torch_xla.distributed.xla_multiprocessing as xmp\nimport torch_xla.test.test_utils as test_utils\n\nwarnings.filterwarnings(\"ignore\")","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"class AverageMeter:\n    \"\"\"\n    Computes and stores the average and current value\n    \"\"\"\n    def __init__(self):\n        self.reset()\n\n    def reset(self):\n        self.val = 0\n        self.avg = 0\n        self.sum = 0\n        self.count = 0\n\n    def update(self, val, n=1):\n        self.val = val\n        self.sum += val * n\n        self.count += n\n        self.avg = self.sum / self.count","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Dataset Preprocessing","execution_count":null},{"metadata":{"trusted":true},"cell_type":"code","source":"# Class to create datasets from numpy arrays\nclass ArrayDataset(torch.utils.data.Dataset):\n    def __init__(self,*arrays):\n        assert all(arrays[0].shape[0] == array.shape[0] for array in arrays)\n        self.arrays = arrays\n    \n    def __getitem__(self, index):\n        return tuple(torch.from_numpy(np.array(array[index])) for array in self.arrays)\n    \n    def __len__(self):\n        return self.arrays[0].shape[0]","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"tokenized_path = '../input/comments-preprocessed/'","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"x_train = np.load(tokenized_path+'x_train.npy',mmap_mode='r')\ntrain_toxic = np.load(tokenized_path+'df_train_toxic.npy',mmap_mode='r')\n\nx_valid = np.load(tokenized_path+'x_valid.npy',mmap_mode='r')\nvalid_toxic = np.load(tokenized_path+'df_valid_toxic.npy',mmap_mode='r')\n\nx_train.shape, x_valid.shape","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"train_dataset = ArrayDataset(x_train, train_toxic)\nvalid_dataset = ArrayDataset(x_valid, valid_toxic)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"# Delete unused variables\ndel x_train, x_valid\nimport gc;gc.collect()\ngc.collect()","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Transformer Model","execution_count":null},{"metadata":{"trusted":true},"cell_type":"code","source":"class CustomTransformer(nn.Module):\n    def __init__(self):\n        super(CustomTransformer, self).__init__()\n        self.num_labels = 1\n        self.roberta = transformers.XLMRobertaModel.from_pretrained(\"xlm-roberta-large\", output_hidden_states=False, num_labels=1) # Choose a model from https://huggingface.co/transformers/pretrained_models.html\n        self.dropout = nn.Dropout(p=0.2)\n        self.classifier = nn.Linear(1024, self.num_labels)\n\n    def forward(self,\n                input_ids=None,\n                attention_mask=None,\n                position_ids=None,\n                head_mask=None,\n                inputs_embeds=None):\n\n        _, o2 = self.roberta(input_ids,\n                               attention_mask=attention_mask,\n                               position_ids=position_ids,\n                               head_mask=head_mask,\n                               inputs_embeds=inputs_embeds)\n\n        logits = self.classifier(o2)       \n        outputs = logits\n        return outputs\n    \nmx = CustomTransformer();\nmx","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Training","execution_count":null},{"metadata":{"trusted":true},"cell_type":"code","source":"import torch_xla.version as xv\nprint('PYTORCH:', xv.__torch_gitrev__)\nprint('XLA:', xv.__xla_gitrev__)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"!free -h","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def loss_fn(outputs, targets):\n    return nn.BCEWithLogitsLoss()(outputs, targets.view(-1, 1))","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def reduce_fn(vals):\n    return sum(vals) / len(vals)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def train_loop_fn(data_loader, model, optimizer, device, scheduler=None):\n    model.train()\n    \n    for bi, d in enumerate(data_loader):\n        ids = d[0]\n        targets = d[1]\n\n        ids = ids.to(device, dtype=torch.long)\n        targets = targets.to(device, dtype=torch.float)\n\n        optimizer.zero_grad()\n        outputs = model(\n            input_ids=ids,\n        )\n        loss = loss_fn(outputs, targets)\n        if bi % 50 == 0:\n            loss_reduced = xm.mesh_reduce('loss_reduce',loss,reduce_fn)\n            xm.master_print(f'bi={bi}, loss={loss_reduced}')\n        loss.backward()\n        xm.optimizer_step(optimizer)\n        if scheduler is not None:\n            scheduler.step()\n            \n    model.eval()\n    \ndef eval_loop_fn(data_loader, model, device):\n    fin_targets = []\n    fin_outputs = []\n    for bi, d in enumerate(data_loader):\n        ids = d[0]\n        targets = d[1]\n\n        ids = ids.to(device, dtype=torch.long)\n        targets = targets.to(device, dtype=torch.float)\n\n        outputs = model(\n            input_ids=ids,\n        )\n\n        targets_np = targets.cpu().detach().numpy().tolist()\n        outputs_np = outputs.cpu().detach().numpy().tolist()\n        fin_targets.extend(targets_np)\n        fin_outputs.extend(outputs_np)    \n        del targets_np, outputs_np\n        gc.collect()\n    return fin_outputs, fin_targets","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def _run():\n    MAX_LEN = 192\n    TRAIN_BATCH_SIZE = 16\n    EPOCHS = 1\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    train_data_loader = torch.utils.data.DataLoader(\n        train_dataset,\n        batch_size=TRAIN_BATCH_SIZE,\n        sampler=train_sampler,\n        drop_last=True,\n        num_workers=0,\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    valid_data_loader = torch.utils.data.DataLoader(\n        valid_dataset,\n        batch_size=4,\n        sampler=valid_sampler,\n        drop_last=False,\n        num_workers=0\n    )\n\n    device = xm.xla_device()\n    model = mx.to(device)\n    xm.master_print('The model is loaded onto the device')\n\n    param_optimizer = list(model.named_parameters())\n    no_decay = ['bias', 'LayerNorm.bias', 'LayerNorm.weight']\n    optimizer_grouped_parameters = [\n        {'params': [p for n, p in param_optimizer if not any(nd in n for nd in no_decay)], 'weight_decay': 0.001},\n        {'params': [p for n, p in param_optimizer if any(nd in n for nd in no_decay)], 'weight_decay': 0.0}]\n\n    lr = 0.5e-5 * xm.xrt_world_size()\n    num_train_steps = int(len(train_dataset) / TRAIN_BATCH_SIZE / xm.xrt_world_size() * EPOCHS)\n    \n    optimizer = AdamW(optimizer_grouped_parameters, lr=lr)\n    scheduler = get_linear_schedule_with_warmup(\n        optimizer,\n        num_warmup_steps=0,\n        num_training_steps=num_train_steps\n    )\n    xm.master_print(f'num_train_steps = {num_train_steps}, world_size={xm.xrt_world_size()}')\n\n\n    for epoch in tqdm(range(EPOCHS)):\n        gc.collect()\n        para_loader = pl.ParallelLoader(train_data_loader, [device])\n        xm.master_print('Parallel loader created... Training...')\n        gc.collect()\n        train_loop_fn(para_loader.per_device_loader(device), model, optimizer, device, scheduler=scheduler)\n        del para_loader\n        para_loader = pl.ParallelLoader(valid_data_loader, [device])\n        gc.collect()\n        o, t = eval_loop_fn(para_loader.per_device_loader(device), model, device)\n        del para_loader\n        gc.collect()\n        auc = metrics.roc_auc_score(np.array(t) >= 0.5, o)\n        auc_reduced = xm.mesh_reduce('auc_reduce',auc,reduce_fn)\n        xm.master_print(f'AUC = {auc_reduced}')\n        gc.collect()\n    xm.save(model.state_dict(), \"xlm_roberta_model.bin\")","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"import time\n\n# Start training processes\ndef _mp_fn(rank, flags):\n    a = _run()\n\nFLAGS={}\nstart_time = time.time()\nxmp.spawn(_mp_fn, args=(FLAGS,), nprocs=8, start_method='fork')\n\nprint('Time taken: ',time.time()-start_time)","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":1}