{"cells":[{"metadata":{"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 --apt-packages libomp5 libopenblas-dev","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"import torch_xla\nimport torch_xla.core.xla_model as xm","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"import torch_xla.distributed.xla_multiprocessing as xmp\nimport torch_xla.distributed.parallel_loader as pl","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"import transformers\n\n\nMAX_LEN = 128\nTRAIN_BATCH_SIZE = 256\nVALID_BATCH_SIZE = 64\nEPOCHS = 10\nBERT_PATH = \"../input/bertbasemultilingualuncased/\"\nMODEL_PATH = \"model.bin\"\nTOKENIZER = transformers.BertTokenizer.from_pretrained(\n    BERT_PATH,\n    do_lower_case=True)","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"cell_type":"code","source":"import torch\n\n\nclass BERTDataset:\n    def __init__(self, comment_text, target):\n        self.comment_text = comment_text\n        self.target = target\n        self.tokenizer = TOKENIZER\n        self.max_len = MAX_LEN\n    \n    def __len__(self):\n        return len(self.comment_text)\n    \n    def __getitem__(self, item):\n        comment_text = str(self.comment_text[item])\n        comment_text = \" \".join(comment_text.split())\n\n        inputs = self.tokenizer.encode_plus(\n            comment_text,\n            None,\n            add_special_tokens=True,\n            max_length=self.max_len,\n            pad_to_max_length=True\n        )\n\n        ids = inputs[\"input_ids\"]\n        mask = inputs[\"attention_mask\"]\n        token_type_ids = inputs[\"token_type_ids\"]\n\n        return {\n            'ids': torch.tensor(ids, dtype=torch.long),\n            'mask': torch.tensor(mask, dtype=torch.long),\n            'token_type_ids': torch.tensor(token_type_ids, dtype=torch.long),\n            'targets': torch.tensor(self.target[item], dtype=torch.float)\n        }","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"d629ff2d2480ee46fbb7e2d37f6b5fab8052498a","_cell_guid":"79c7e3d0-c299-4dcb-8224-4455121ee9b0","trusted":true},"cell_type":"code","source":"import transformers\nimport torch.nn as nn\n\n\nclass BERTBaseUncased(nn.Module):\n    def __init__(self):\n        super(BERTBaseUncased, self).__init__()\n        self.bert = transformers.BertModel.from_pretrained(BERT_PATH)\n        self.bert_drop = nn.Dropout(0.3)\n        self.out = nn.Linear(768*2, 1)\n    \n    def forward(self, ids, mask, token_type_ids):\n        o1, _ = self.bert(\n            ids, \n            attention_mask=mask,\n            token_type_ids=token_type_ids\n        )\n        mean_pooling = torch.mean(o1,1)\n        max_pooling,_ = torch.max(o1,1)\n        cat = torch.cat((mean_pooling,max_pooling),1)\n        \n        bo = self.bert_drop(cat)\n        output = self.out(bo)\n        return output","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"import torch\nimport torch.nn as nn\nfrom tqdm import tqdm\n\n\ndef loss_fn(outputs, targets):\n    return nn.BCEWithLogitsLoss()(outputs, targets.view(-1, 1))\n\n\ndef train_fn(data_loader, model, optimizer, device, scheduler):\n    model.train()\n\n    for bi, d in enumerate(data_loader):\n        ids = d[\"ids\"]\n        token_type_ids = d[\"token_type_ids\"]\n        mask = d[\"mask\"]\n        targets = d[\"targets\"]\n\n        ids = ids.to(device, dtype=torch.long)\n        token_type_ids = token_type_ids.to(device, dtype=torch.long)\n        mask = mask.to(device, dtype=torch.long)\n        targets = targets.to(device, dtype=torch.float)\n\n        optimizer.zero_grad()\n        outputs = model(\n            ids=ids,\n            mask=mask,\n            token_type_ids=token_type_ids\n        )\n\n        loss = loss_fn(outputs, targets)\n        loss.backward()\n        xm.optimizer_step(optimizer)\n        scheduler.step()\n        if bi % 10==0:\n            xm.master_print(f\"bi={bi},loss={loss.item()}\")\n\ndef eval_fn(data_loader, model, device):\n    model.eval()\n    fin_targets = []\n    fin_outputs = []\n    for bi, d in enumerate(data_loader):\n        ids = d[\"ids\"]\n        token_type_ids = d[\"token_type_ids\"]\n        mask = d[\"mask\"]\n        targets = d[\"targets\"]\n\n        ids = ids.to(device, dtype=torch.long)\n        token_type_ids = token_type_ids.to(device, dtype=torch.long)\n        mask = mask.to(device, dtype=torch.long)\n        targets = targets.to(device, dtype=torch.float)\n\n        outputs = model(\n            ids=ids,\n            mask=mask,\n            token_type_ids=token_type_ids\n        )\n        fin_targets.extend(targets.cpu().detach().numpy().tolist())\n        fin_outputs.extend(torch.sigmoid(outputs).cpu().detach().numpy().tolist())\n    return fin_outputs, fin_targets","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"import torch\nimport pandas as pd\nimport torch.nn as nn\nimport numpy as np\n\nfrom sklearn import model_selection\nfrom sklearn import metrics\nfrom transformers import AdamW\nfrom transformers import get_linear_schedule_with_warmup\n\ndef run():\n    df1 = pd.read_csv(\"../input/jigsaw-multilingual-toxic-comment-classification/jigsaw-toxic-comment-train.csv\",usecols=[\"comment_text\",\"toxic\"])\n    df2 = pd.read_csv(\"../input/jigsaw-multilingual-toxic-comment-classification/jigsaw-unintended-bias-train.csv\",usecols=[\"comment_text\",\"toxic\"])\n    df_train = pd.concat([df1,df2],axis=0).reset_index(drop=True)\n    \n    df_valid = pd.read_csv(\"../input/jigsaw-multilingual-toxic-comment-classification/validation.csv\")\n    \n\n    train_dataset = BERTDataset(\n        comment_text = df_train.comment_text.values,\n        target=df_train.toxic.values\n    )\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    \n\n    train_data_loader = torch.utils.data.DataLoader(\n        train_dataset,\n        batch_size=TRAIN_BATCH_SIZE,\n        num_workers=4,\n        sampler = train_sampler,\n        drop_last= True\n    )\n\n    valid_dataset = BERTDataset(\n        comment_text=df_valid.comment_text.values,\n        target=df_valid.toxic.values\n    )\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=True\n    )\n        \n        \n    valid_data_loader = torch.utils.data.DataLoader(\n        valid_dataset,\n        batch_size=VALID_BATCH_SIZE,\n        num_workers=1,\n        sampler = valid_sampler\n    )\n\n    device = xm.xla_device()\n    model = BERTBaseUncased()\n    model.to(device)\n    \n    param_optimizer = list(model.named_parameters())\n    no_decay = [\"bias\", \"LayerNorm.bias\", \"LayerNorm.weight\"]\n    optimizer_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\n    num_train_steps = int(len(df_train) / TRAIN_BATCH_SIZE/xm.xrt_world_size() * EPOCHS)\n    lr = 3e-5 * xm.xrt_world_size()\n    optimizer = AdamW(optimizer_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\n\n    best_accuracy = 0\n    for epoch in range(EPOCHS):\n        para_loader= pl.ParallelLoader(train_data_loader,[device])\n        train_fn(para_loader.per_device_loader(device), model, optimizer, device, scheduler)\n        para_loader= pl.ParallelLoader(valid_data_loader,[device])\n        outputs, targets = eval_fn(para_loader.per_device_loader(device), model, device)\n        targets = np.array(targets) >= 0.5\n        accuracy = metrics.roc_auc_score(targets, outputs)\n        print(f\"Auc Score = {accuracy}\")\n        if accuracy > best_accuracy:\n            xm.save(model.state_dict(),MODEL_PATH)\n            best_accuracy = accuracy\n\n\n","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def _multiprocessing_function(rank,flags):\n    torch.set_default_tensor_type(\"torch.FloatTensor\")\n    a = run()\n    \nxmp.spawn(_multiprocessing_funtion(), args=({},)),npross=1,start_method='fork')","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"run()","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}