{"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_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"### This kernel is forked from https://www.kaggle.com/theoviel/bert-pytorch-huggingface-starter\n- Thanks for sharing https://www.kaggle.com/theoviel","metadata":{}},{"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","metadata":{"execution":{"iopub.status.busy":"2022-09-15T17:27:27.989153Z","iopub.execute_input":"2022-09-15T17:27:27.989567Z","iopub.status.idle":"2022-09-15T17:28:13.775423Z","shell.execute_reply.started":"2022-09-15T17:27:27.989452Z","shell.execute_reply":"2022-09-15T17:28:13.77426Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install transformers","metadata":{"execution":{"iopub.status.busy":"2022-09-15T17:28:13.779261Z","iopub.execute_input":"2022-09-15T17:28:13.779694Z","iopub.status.idle":"2022-09-15T17:28:22.058028Z","shell.execute_reply.started":"2022-09-15T17:28:13.779638Z","shell.execute_reply":"2022-09-15T17:28:22.057161Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import json\nimport torch\nimport os\nimport numpy as np\nfrom tqdm import tqdm","metadata":{"execution":{"iopub.status.busy":"2022-09-15T17:28:22.059404Z","iopub.execute_input":"2022-09-15T17:28:22.059716Z","iopub.status.idle":"2022-09-15T17:28:22.508269Z","shell.execute_reply.started":"2022-09-15T17:28:22.059678Z","shell.execute_reply":"2022-09-15T17:28:22.507517Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Data","metadata":{}},{"cell_type":"code","source":"def read_data(path):\n    with open(path, 'rb') as f:\n        coqa = json.load(f)\n    contexts = []\n    questions = []\n    conversation = []\n    prev_questions = []\n    prev_answers = []\n    answers = []\n    for data in tqdm(coqa['data']):\n        context = data['story']\n        for i in range(len(data['questions'])):\n            contexts.append(context)\n            questions.append(data['questions'][i]['input_text'])\n            answers.append(data['answers'][i])\n            if i != 0:\n                prev_question = []\n                prev_answer = []\n                for j in range (i):\n                    prev_question.append(data['questions'][j]['input_text'])\n                    prev_answer.append(data['answers'][j]['span_text'])\n                prev_questions.append(prev_question)\n                prev_answers.append(prev_answer)\n            else:\n                prev_questions.append(\" \")\n                prev_answers.append(\" \")\n    for i in range(len(prev_answers)):\n        if len(prev_answers[i][0]) >1:\n            hist = []\n            for j in range(len(prev_answers[i])):\n                hist.append(\"[Q] \" + prev_questions[i][j])\n                hist.append(\"[A] \" + prev_answers[i][j])\n            past = \" \".join(hist)\n        else:\n            past = \"\"\n        conversation.append(past + \" [Q] \" + questions[i])\n    return contexts, conversation, answers\n\ntrain_contexts, train_conversation, train_answers = read_data('../input/conversational-question-answering-dataset-coqa/coqa-train-v1.0.json')\nvalid_contexts, valid_conversation, valid_answers = read_data('../input/conversational-question-answering-dataset-coqa/coqa-dev-v1.0.json')","metadata":{"execution":{"iopub.status.busy":"2022-09-15T17:28:22.510738Z","iopub.execute_input":"2022-09-15T17:28:22.511488Z","iopub.status.idle":"2022-09-15T17:28:26.050045Z","shell.execute_reply.started":"2022-09-15T17:28:22.51144Z","shell.execute_reply":"2022-09-15T17:28:26.048816Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def take_gold_answers(path):\n    with open(path, 'rb') as f:\n        coqa = json.load(f)\n    gold_answers = []\n    for data in tqdm(coqa['data']):\n        for j in range(len(data['answers'])):\n            ans = []\n            ans.append(data['answers'][j]['span_text'])\n            for k in range(3):\n                ans.append(data['additional_answers'][str(k)][j]['span_text'])\n            gold_answers.append(ans)\n    return gold_answers\ndev_answer = take_gold_answers('../input/conversational-question-answering-dataset-coqa/coqa-dev-v1.0.json')","metadata":{"execution":{"iopub.status.busy":"2022-09-15T17:28:26.051881Z","iopub.execute_input":"2022-09-15T17:28:26.052261Z","iopub.status.idle":"2022-09-15T17:28:26.185743Z","shell.execute_reply.started":"2022-09-15T17:28:26.052214Z","shell.execute_reply":"2022-09-15T17:28:26.184937Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def process_ans(context, answer):\n    for i in tqdm(range(len(answer))):\n        data = answer[i]\n        if data[\"span_start\"] != -1:\n            if context[i][data[\"span_start\"]] == \" \":\n                data[\"span_start\"] += 1\n            if context[i][data[\"span_end\"]-1] == \" \":\n                data[\"span_end\"] -= 1\n        data['span_text'] = \" \".join(data['span_text'].split())\nprocess_ans(train_contexts, train_answers)\nprocess_ans(valid_contexts, valid_answers)","metadata":{"execution":{"iopub.status.busy":"2022-09-15T17:28:26.187392Z","iopub.execute_input":"2022-09-15T17:28:26.187709Z","iopub.status.idle":"2022-09-15T17:28:26.430279Z","shell.execute_reply.started":"2022-09-15T17:28:26.187672Z","shell.execute_reply":"2022-09-15T17:28:26.429316Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def pre_truncating(contexts, conversation):\n    for i in tqdm(range(len(contexts))):\n        pair_question = conversation[i].split(\"[Q]\")\n        token = contexts[i].split(\" \") \n        conver_stack = pair_question[-2:]\n        if len(pair_question) > 3:\n            j = len(pair_question) - 3\n            while len(token) + len(((\"\").join(conver_stack)).split(\" \")) < 512 and j > -1:\n                conver_stack.insert(0, pair_question[j])\n                j-=1\n        if conver_stack[0] == \"\": \n            conversation[i] = (\"[Q]\").join(conver_stack)\n        else:\n            conversation[i] = '[Q]' + (\"[Q]\").join(conver_stack)\n        checking = conversation[i].split(\" \")\n        if checking[0] == checking[1] == \"[Q]\":\n            conversation[i] = \" \".join(checking[1:])\n        \npre_truncating(train_contexts, train_conversation)\npre_truncating(valid_contexts, valid_conversation)","metadata":{"execution":{"iopub.status.busy":"2022-09-15T17:28:26.431673Z","iopub.execute_input":"2022-09-15T17:28:26.431951Z","iopub.status.idle":"2022-09-15T17:28:34.903812Z","shell.execute_reply.started":"2022-09-15T17:28:26.431918Z","shell.execute_reply":"2022-09-15T17:28:34.902906Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from transformers import AutoTokenizer\ntokenizer = AutoTokenizer.from_pretrained(\"deepset/roberta-base-squad2\")\n# Add new tokens for Question and Answer conversation\ntokenizer.add_tokens([\"[Q]\",\"[A]\"])\ntrain_encodings = tokenizer(train_conversation, train_contexts, truncation=True, padding=True)\nvalid_encodings = tokenizer(valid_conversation, valid_contexts, truncation=True, padding=True)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def add_token_positions(encodings, answers):\n  start_positions = []\n  end_positions = []\n  for i in range(len(answers)):\n    if answers[i]['span_start'] == -1:\n      start_positions.append(tokenizer.model_max_length)\n      end_positions.append(tokenizer.model_max_length)\n    else:\n      start_positions.append(encodings.char_to_token(i, answers[i]['span_start']))\n      end_positions.append(encodings.char_to_token(i, answers[i]['span_end'] -1))\n    if start_positions[-1] is None:\n      start_positions[-1] = tokenizer.model_max_length\n    if end_positions[-1] is None:\n      end_positions[-1] = tokenizer.model_max_length\n  encodings.update({'start_positions': start_positions, 'end_positions': end_positions})\n\n  print(count/len(answers))\nadd_token_positions(train_encodings, train_answers)\nadd_token_positions(valid_encodings, valid_answers)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CoQA_Dataset(torch.utils.data.Dataset):\n  def __init__(self, encodings):\n    self.encodings = encodings\n  def __getitem__(self, idx):\n    return {key: torch.tensor(val[idx]) for key, val in self.encodings.items()}\n  def __len__(self):\n    return len(self.encodings.input_ids)\ntrain_dataset = CoQA_Dataset(train_encodings)\nvalid_dataset = CoQA_Dataset(valid_encodings)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torch.utils.data import DataLoader\n\ntrain_loader = DataLoader(train_dataset, batch_size=16, shuffle=True)\nvalid_loader = DataLoader(valid_dataset, batch_size=1)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training","metadata":{}},{"cell_type":"code","source":"import torch_xla\nimport torch_xla.core.xla_model as xm","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from transformers import BertForQuestionAnswering\n\nmodel = BertForQuestionAnswering.from_pretrained(\"bert-base-uncased\")\nmodel.resize_token_embeddings(len(tokenizer))\ndevice = xm.xla_device()\nmodel = model.to(device)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from transformers import AdamW\n\nN_EPOCHS = 5\noptim = AdamW(model.parameters(), lr=5e-5)\n\nfor epoch in range(N_EPOCHS):\n    model.train()\n    process_loss = 0\n    iter = 0\n    print(f'Start training Epoch {epoch+1}:')\n    for batch in tqdm(train_loader):\n        iter += 1\n        optim.zero_grad()\n        input_ids = batch['input_ids'].to(device)\n        attention_mask = batch['attention_mask'].to(device)\n        start_positions = batch['start_positions'].to(device)\n        end_positions = batch['end_positions'].to(device)\n        outputs = model(input_ids, attention_mask=attention_mask, start_positions=start_positions, end_positions=end_positions)\n        loss = outputs[0]\n        loss.backward()\n        optim.step()\n        process_loss += loss.item()\n        if iter % 100 ==0 or iter == len(train_loader):\n            print(f'\\nEpoch {epoch+1}: Batch {iter}/{len(train_loader)}: Loss = {process_loss/iter}')\n\n#     model.save_pretrained('/content/drive/MyDrive/Colab Notebooks/ThanhDat/CoQA/model_span')\n    \n    model.eval()\n    index = 0\n    f1_score = 0\n    for batch in tqdm(valid_loader):\n        with torch.no_grad():\n            input_ids = batch['input_ids'].to(device)\n            attention_mask = batch['attention_mask'].to(device)\n            start_true = batch['start_positions'].to(device)\n            end_true = batch['end_positions'].to(device)\n            \n            outputs = model(input_ids, attention_mask=attention_mask)\n\n            start_pred = torch.argmax(outputs['start_logits'], dim=1)\n            end_pred = torch.argmax(outputs['end_logits'], dim=1)\n\n            if start_pred > end_pred or start_pred == end_pred == 512:\n                    pred = \"\"\n            else:\n                pred = tokenizer.decode(batch['input_ids'][0][start_pred:end_pred])\n            f1_score += max([compute_f1(pred, answer) for answer in dev_answer[index]])\n            index+=1\n\n    print(\"====================================\")\n    print(\"Evaluating on Dev set\")\n    print(\"F1 on Dev:\", f1_score/len(valid_loader))\n    print(\"\\n\")\nprint(\"=========Training End=========\")","metadata":{"trusted":true},"execution_count":null,"outputs":[]}]}