{"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":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        print(os.path.join(dirname, filename))\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-12-23T14:51:24.806958Z","iopub.execute_input":"2022-12-23T14:51:24.807923Z","iopub.status.idle":"2022-12-23T14:51:24.837647Z","shell.execute_reply.started":"2022-12-23T14:51:24.807877Z","shell.execute_reply":"2022-12-23T14:51:24.836502Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import json\nimport pandas as pd\nimport re\nimport numpy as np\nimport torch\nimport random\nfrom torch.utils.data import Dataset, DataLoader\nfrom transformers import BertTokenizerFast, BertForSequenceClassification, BertForTokenClassification\nfrom sklearn.metrics import accuracy_score, f1_score\nimport matplotlib.pyplot as plt\nfrom IPython.display import clear_output\nimport time","metadata":{"execution":{"iopub.status.busy":"2022-12-23T14:51:24.839724Z","iopub.execute_input":"2022-12-23T14:51:24.840156Z","iopub.status.idle":"2022-12-23T14:51:24.846834Z","shell.execute_reply.started":"2022-12-23T14:51:24.840114Z","shell.execute_reply":"2022-12-23T14:51:24.845661Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"bert_base_uncased_filename = '/kaggle/input/transformers/bert-base-uncased'\nbert_base_cased_filename = '/kaggle/input/bert-base-cased'","metadata":{"execution":{"iopub.status.busy":"2022-12-23T14:51:24.851056Z","iopub.execute_input":"2022-12-23T14:51:24.851713Z","iopub.status.idle":"2022-12-23T14:51:24.860530Z","shell.execute_reply.started":"2022-12-23T14:51:24.851677Z","shell.execute_reply":"2022-12-23T14:51:24.859411Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"tokenizer = BertTokenizerFast.from_pretrained(bert_base_cased_filename)\nlong_md = BertForSequenceClassification.from_pretrained(\n    bert_base_uncased_filename,\n    num_labels = 2\n)","metadata":{"execution":{"iopub.status.busy":"2022-12-23T14:51:24.865024Z","iopub.execute_input":"2022-12-23T14:51:24.865343Z","iopub.status.idle":"2022-12-23T14:51:26.153532Z","shell.execute_reply.started":"2022-12-23T14:51:24.865278Z","shell.execute_reply":"2022-12-23T14:51:26.152406Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torch.nn import DataParallel","metadata":{"execution":{"iopub.status.busy":"2022-12-23T14:51:26.155637Z","iopub.execute_input":"2022-12-23T14:51:26.156337Z","iopub.status.idle":"2022-12-23T14:51:26.161425Z","shell.execute_reply.started":"2022-12-23T14:51:26.156295Z","shell.execute_reply":"2022-12-23T14:51:26.160324Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class ShortAnswerModel:\n\n    def __init__(self, model, device):\n        self.model = model\n        self.device = device\n        self.model = DataParallel(self.model).to(device)\n\n    def __call__(self, tokenized_input, question_lens):\n\n        self.model.eval()\n        with torch.no_grad():\n\n            ids = tokenized_input['input_ids'].to(self.device).squeeze(dim=1)\n            mask = tokenized_input['attention_mask'].to(self.device).squeeze(dim=1)\n\n            output = self.model(input_ids=ids, attention_mask=mask)\n            logits = output['logits']\n            pred = torch.argmax(logits, axis=-1)\n\n            ind_preds = []\n            for i in range(len(pred)):\n                start, end = self.get_start_end_tokens(pred[i],\n                                                       tokenized_input['offset_mapping'][i],\n                                                       question_lens[i])\n                ind_preds.append([start, end])\n\n        return ind_preds\n\n    def get_start_end_tokens(self, bio_tags, offset_mapping, question_len):\n\n        token_mapping = self.get_tokenization_mapping(offset_mapping)\n\n        answer_inds = torch.where(bio_tags != 0)[0].cpu()\n\n        true_mask = []\n        for ind in range(len(answer_inds)):\n            if ind == 0:\n                true_mask.append(answer_inds[ind])\n                continue\n            if answer_inds[ind - 1] != answer_inds[ind] - 1:\n                break\n            true_mask.append(answer_inds[ind])\n        true_mask = np.array(true_mask)\n\n        answer_tokens = []\n        for ind in true_mask:\n            for group_ind, token_group in enumerate(token_mapping):\n                if ind in token_group:\n                    answer_tokens.append(group_ind)\n                    break\n        if answer_tokens:\n            start = answer_tokens[0] - question_len\n            end = answer_tokens[-1] - question_len + 1\n        else:\n            start = -1\n            end = -1\n        return start, end\n\n    # TODO: don't duplicate this function!\n    def get_tokenization_mapping(self, offset_mapping):\n        d = []\n        for i, pair in enumerate(offset_mapping):\n            if pair[0] == 0 and pair[1] == 0:\n                continue\n            if pair[0] == 0:\n                d.append([i])\n            else:\n                if len(d) == 0:\n                    d.append([i])\n                else:\n                    d[-1].append(i)\n        return d","metadata":{"execution":{"iopub.status.busy":"2022-12-23T14:51:26.163188Z","iopub.execute_input":"2022-12-23T14:51:26.163935Z","iopub.status.idle":"2022-12-23T14:51:26.181437Z","shell.execute_reply.started":"2022-12-23T14:51:26.163895Z","shell.execute_reply":"2022-12-23T14:51:26.180600Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\ndevice","metadata":{"execution":{"iopub.status.busy":"2022-12-23T14:51:26.184214Z","iopub.execute_input":"2022-12-23T14:51:26.184870Z","iopub.status.idle":"2022-12-23T14:51:26.196325Z","shell.execute_reply.started":"2022-12-23T14:51:26.184834Z","shell.execute_reply":"2022-12-23T14:51:26.195138Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"long_model_filename = '/kaggle/input/long-answer-model/long_answer_model_part_freeze_balance_small.pt'\nstate_dict = torch.load(long_model_filename, map_location=device)\nnew_state_dict = {}\nfor key, value in state_dict.items():\n    new_state_dict[key.replace('module.', '')] = value\nlong_md.load_state_dict(new_state_dict)","metadata":{"execution":{"iopub.status.busy":"2022-12-23T14:51:26.198869Z","iopub.execute_input":"2022-12-23T14:51:26.199288Z","iopub.status.idle":"2022-12-23T14:51:26.616750Z","shell.execute_reply.started":"2022-12-23T14:51:26.199238Z","shell.execute_reply":"2022-12-23T14:51:26.615727Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class LongAnswerModel():\n    \n    def __init__(self, model, device):\n        self.model = model\n        self.device = device\n#         self.model.to(device)\n        \n        # PARALLEL\n        self.model = DataParallel(self.model).to(device)\n        \n    def __call__(self, input_ids, attn_mask):\n        \n        self.model.eval()\n        all_logits = []\n        with torch.no_grad():\n            \n            n_p = input_ids.size()[0]\n#             print(n_p)\n            for i in range(0, n_p, 50):\n                self.model = self.model.to(device)\n                batch_index = slice(i, min(i+50, n_p))\n                input_ids_i = input_ids[batch_index].to(self.device)\n                attn_mask_i = attn_mask[batch_index].to(self.device)\n                output = self.model(input_ids=input_ids_i, attention_mask=attn_mask_i)\n                logits = output['logits']\n                all_logits.append(logits)\n                torch.cuda.empty_cache()\n            logits = torch.cat(all_logits, 0)\n#             print(logits)\n            active_logits = logits.view(-1, long_md.num_labels)\n            flattened_pred = torch.argmax(active_logits, axis=1)\n            if 1 in flattened_pred:\n                ind = torch.argmax(active_logits[:, 1]).item()\n                prediction = torch.zeros_like(flattened_pred)\n                prediction[ind] = 1\n                return prediction\n        return flattened_pred","metadata":{"execution":{"iopub.status.busy":"2022-12-23T14:51:26.618691Z","iopub.execute_input":"2022-12-23T14:51:26.619846Z","iopub.status.idle":"2022-12-23T14:51:26.631820Z","shell.execute_reply.started":"2022-12-23T14:51:26.619780Z","shell.execute_reply":"2022-12-23T14:51:26.630669Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"long_answer_model = LongAnswerModel(long_md, device=device)","metadata":{"execution":{"iopub.status.busy":"2022-12-23T14:51:26.634595Z","iopub.execute_input":"2022-12-23T14:51:26.635344Z","iopub.status.idle":"2022-12-23T14:51:26.756551Z","shell.execute_reply.started":"2022-12-23T14:51:26.635305Z","shell.execute_reply":"2022-12-23T14:51:26.755558Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_filename = '/kaggle/input/tensorflow2-question-answering/simplified-nq-test.jsonl'\nwith open(test_filename, 'r') as json_file:\n    json_list = list(json_file)\n    \ndata = []\nfor json_str in json_list:\n    data.append(json.loads(json_str))","metadata":{"execution":{"iopub.status.busy":"2022-12-23T14:51:26.758148Z","iopub.execute_input":"2022-12-23T14:51:26.758543Z","iopub.status.idle":"2022-12-23T14:51:27.260027Z","shell.execute_reply.started":"2022-12-23T14:51:26.758507Z","shell.execute_reply":"2022-12-23T14:51:27.258622Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class LongAnswerDatasetBase(Dataset):\n    HTML_PATTERN = re.compile(r'<.*?>')\n    \n    def _get_nq_tokens(self, simplified_nq_example):\n        if \"document_text\" not in simplified_nq_example:\n            raise ValueError(\"`get_nq_tokens` should be called on a simplified NQ\"\n                         \"example that contains the `document_text` field.\")\n\n        return simplified_nq_example[\"document_text\"].split(\" \")\n    \n    def _clean_token(self, token):\n        return re.sub(u\" \", \"_\", token[\"token\"])\n\n    def _remove_html_byte_offsets(self, span):\n        if \"start_byte\" in span:\n            del span[\"start_byte\"]\n\n        if \"end_byte\" in span:\n            del span[\"end_byte\"]\n\n        return span\n\n    def _clean_annotation(self, annotation):\n        annotation[\"long_answer\"] = self._remove_html_byte_offsets(\n            annotation[\"long_answer\"])\n        annotation[\"short_answers\"] = [\n            self._remove_html_byte_offsets(sa) for sa in annotation[\"short_answers\"]\n        ]\n        return annotation\n    \n    def _simplify_nq_example(self, nq_example):\n        text = \" \".join([self._clean_token(t) for t in nq_example[\"document_tokens\"]])\n\n        simplified_nq_example = {\n          \"question_text\": nq_example[\"question_text\"],\n          \"example_id\": nq_example[\"example_id\"],\n          \"document_url\": nq_example[\"document_url\"],\n          \"document_text\": text,\n          \"long_answer_candidates\": [\n              self._remove_html_byte_offsets(c)\n              for c in nq_example[\"long_answer_candidates\"]\n          ],\n          \"annotations\": [self._clean_annotation(a) for a in nq_example[\"annotations\"]]\n        }\n\n        if len(self._get_nq_tokens(simplified_nq_example)) != len(\n          nq_example[\"document_tokens\"]):\n            raise ValueError(\"Incorrect number of tokens.\")\n\n        return simplified_nq_example\n    \n    def _get_question_and_document(self, line):\n        question = line['question_text']\n        text = line['document_text'].split(' ')\n        example_id = line['example_id']\n\n        return question, text, example_id\n\n\n    def _get_long_candidate(self, i, candidate):\n        long_start = candidate['start_token']\n        long_end = candidate['end_token']\n\n        return long_start, long_end\n    \n    def _preprocess_data(self, data):\n        rows = []\n\n        for line in data:\n            if not self._kaggle_format:\n                line = self._simplify_nq_example(line)\n            question, text, example_id = self._get_question_and_document(line)\n            for i, candidate in enumerate(line['long_answer_candidates']):\n                long_start, long_end = self._get_long_candidate(i, candidate)\n                rows.append(\n                    self._form_data_row(example_id, question, text, long_start, long_end)\n                )\n\n        return pd.DataFrame(rows)\n    \n    def _remove_stopwords(self, sentence):\n        words = sentence.split()\n        words = [word for word in words if word not in stopwords.words('english')]\n\n        return ' '.join(words)\n\n    def _remove_html(self, sentence):\n        return  self.HTML_PATTERN.sub(r'', sentence)\n\n    def _clean_df_by_column(self, df, column):\n        # df[column] = df[column].apply(lambda x : self._remove_stopwords(x))\n        df[column] = df[column].apply(lambda x : self._remove_html(x))\n        return df\n\n    def _clean_df(self, df):\n        df = self._clean_df_by_column(df, 'long_answer')\n        df = self._clean_df_by_column(df, 'question')\n        return df\n    \n    def __getitem__(self, idx):\n        raise NotImplementedError('method __getitem__ is not implemented')\n    \n    def __len__(self):\n        raise NotImplementedError('method __len__ is not implemented')","metadata":{"execution":{"iopub.status.busy":"2022-12-23T14:51:27.265393Z","iopub.execute_input":"2022-12-23T14:51:27.265781Z","iopub.status.idle":"2022-12-23T14:51:27.291916Z","shell.execute_reply.started":"2022-12-23T14:51:27.265728Z","shell.execute_reply":"2022-12-23T14:51:27.290600Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class TestLongAnswerDataset(LongAnswerDatasetBase):\n    def __init__(self, data, tokenizer, max_len=150, kaggle_format=True):\n        self._tokenizer = tokenizer\n        self._max_len = max_len\n        self._kaggle_format = kaggle_format\n        \n        data = self._preprocess_data(data)\n        data = self._clean_df(data)\n        self._data = {index: question_df for index, (question, question_df) in enumerate(data.groupby('question'))}\n\n    def _form_data_row(self, example_id, question, text, long_start, long_end):\n        row = {\n            'example_id': example_id,\n            'question': question,\n            'long_answer': ' '.join(text[long_start:long_end]),\n            'long_start': long_start,\n            'long_end': long_end\n        }\n\n        return row\n\n    def __getitem__(self, idx):\n        texts, indices, answers = [], [], []\n        current_data = self._data[idx]\n        example_id, question = None, None\n        for i in range(current_data.shape[0]):\n            example_id = current_data.example_id.iloc[i]\n            question = current_data.question.iloc[i]\n            answer = current_data.long_answer.iloc[i]\n            start = current_data.long_start.iloc[i]\n            end = current_data.long_end.iloc[i]\n            \n            texts.append(question + self._tokenizer.sep_token + answer)\n            indices.append(f\"{start}:{end}\")\n            answers.append(answer)\n            \n        encoding = self._tokenizer(texts,\n                                   return_offsets_mapping=False,\n                                   return_token_type_ids=False,\n                                   padding='max_length',\n                                   truncation=True,\n                                   max_length=self._max_len,\n                                   return_tensors='pt')\n        return encoding, indices, example_id, question, answers \n   \n    def __len__(self):\n        return len(self._data)","metadata":{"execution":{"iopub.status.busy":"2022-12-23T14:51:27.297627Z","iopub.execute_input":"2022-12-23T14:51:27.298256Z","iopub.status.idle":"2022-12-23T14:51:27.312742Z","shell.execute_reply.started":"2022-12-23T14:51:27.298219Z","shell.execute_reply":"2022-12-23T14:51:27.311475Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# print(data[0])","metadata":{"execution":{"iopub.status.busy":"2022-12-23T14:51:27.314518Z","iopub.execute_input":"2022-12-23T14:51:27.315304Z","iopub.status.idle":"2022-12-23T14:51:27.325057Z","shell.execute_reply.started":"2022-12-23T14:51:27.315256Z","shell.execute_reply":"2022-12-23T14:51:27.323906Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"MAX_LEN = 500\ndataset = TestLongAnswerDataset(data, tokenizer, MAX_LEN)","metadata":{"execution":{"iopub.status.busy":"2022-12-23T14:51:27.326643Z","iopub.execute_input":"2022-12-23T14:51:27.327497Z","iopub.status.idle":"2022-12-23T14:51:27.978757Z","shell.execute_reply.started":"2022-12-23T14:51:27.327445Z","shell.execute_reply":"2022-12-23T14:51:27.977736Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"len(dataset)","metadata":{"execution":{"iopub.status.busy":"2022-12-23T14:51:27.980702Z","iopub.execute_input":"2022-12-23T14:51:27.981092Z","iopub.status.idle":"2022-12-23T14:51:27.988198Z","shell.execute_reply.started":"2022-12-23T14:51:27.981055Z","shell.execute_reply":"2022-12-23T14:51:27.986944Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from tqdm import tqdm","metadata":{"execution":{"iopub.status.busy":"2022-12-23T14:51:27.990082Z","iopub.execute_input":"2022-12-23T14:51:27.990472Z","iopub.status.idle":"2022-12-23T14:51:27.995837Z","shell.execute_reply.started":"2022-12-23T14:51:27.990437Z","shell.execute_reply":"2022-12-23T14:51:27.994834Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_inds(labels, indices):\n    if 1 not in labels:\n        return ''\n    ind = labels.index(1)\n    return indices[ind]","metadata":{"execution":{"iopub.status.busy":"2022-12-23T14:51:27.997046Z","iopub.execute_input":"2022-12-23T14:51:27.997814Z","iopub.status.idle":"2022-12-23T14:51:28.005095Z","shell.execute_reply.started":"2022-12-23T14:51:27.997746Z","shell.execute_reply":"2022-12-23T14:51:28.004216Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def process_example(long_answer, question):\n    long_answer = long_answer.split()\n    texts = []\n\n    i = 0\n    max_len = MAX_LEN - 150\n    while max_len * i < len(long_answer):\n        curr_text = long_answer[max_len * i: max_len * (i + 1)]\n        texts.append(curr_text)\n\n        i += 1\n\n    if long_answer[max_len * (i + 1):]:\n        texts.append(long_answer[max_len * (i + 1):])\n\n    questions = [question] * len(texts)\n    return texts, questions","metadata":{"execution":{"iopub.status.busy":"2022-12-23T14:51:28.006489Z","iopub.execute_input":"2022-12-23T14:51:28.006958Z","iopub.status.idle":"2022-12-23T14:51:28.014555Z","shell.execute_reply.started":"2022-12-23T14:51:28.006922Z","shell.execute_reply":"2022-12-23T14:51:28.013389Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def encode_short_ans(batch_long_answers, batch_questions):\n    encodings, question_len = [], []\n\n    for i in range(len(batch_questions)):\n        input_tokens = batch_questions[i].split()\n        input_tokens.append(short_tokenizer.sep_token)\n        question_len = len(input_tokens)\n        input_tokens.extend(batch_long_answers[i])\n        encoding = short_tokenizer(input_tokens,\n                                   is_split_into_words=True,\n                                   return_offsets_mapping=True,\n                                   padding='max_length',\n                                   truncation=True,\n                                   max_length=MAX_LEN,\n                                   return_tensors='pt')\n        encodings.append(encoding)\n\n    return encodings, question_len","metadata":{"execution":{"iopub.status.busy":"2022-12-23T14:51:28.016392Z","iopub.execute_input":"2022-12-23T14:51:28.016722Z","iopub.status.idle":"2022-12-23T14:51:28.025992Z","shell.execute_reply.started":"2022-12-23T14:51:28.016690Z","shell.execute_reply":"2022-12-23T14:51:28.025036Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def join_test_example_preds(pred_starts, pred_ends, long_start):\n    if len(pred_starts) == pred_starts.count(-1):\n        return ''\n        \n    for i in range(len(pred_starts)):\n        if pred_starts[i] != -1:\n            return f\"{long_start+pred_starts[i]}:{long_start+pred_ends[i]}\"\n    return ''","metadata":{"execution":{"iopub.status.busy":"2022-12-23T14:51:28.029196Z","iopub.execute_input":"2022-12-23T14:51:28.029441Z","iopub.status.idle":"2022-12-23T14:51:28.036220Z","shell.execute_reply.started":"2022-12-23T14:51:28.029419Z","shell.execute_reply":"2022-12-23T14:51:28.035323Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"batch_size = 1","metadata":{"execution":{"iopub.status.busy":"2022-12-23T14:51:28.037644Z","iopub.execute_input":"2022-12-23T14:51:28.038032Z","iopub.status.idle":"2022-12-23T14:51:28.045653Z","shell.execute_reply.started":"2022-12-23T14:51:28.037997Z","shell.execute_reply":"2022-12-23T14:51:28.044588Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from collections import defaultdict","metadata":{"execution":{"iopub.status.busy":"2022-12-23T14:51:28.046962Z","iopub.execute_input":"2022-12-23T14:51:28.047601Z","iopub.status.idle":"2022-12-23T14:51:28.066366Z","shell.execute_reply.started":"2022-12-23T14:51:28.047550Z","shell.execute_reply":"2022-12-23T14:51:28.061267Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"result = defaultdict(list)\ndataset2 = []","metadata":{"execution":{"iopub.status.busy":"2022-12-23T14:51:28.067987Z","iopub.execute_input":"2022-12-23T14:51:28.068624Z","iopub.status.idle":"2022-12-23T14:51:28.073944Z","shell.execute_reply.started":"2022-12-23T14:51:28.068589Z","shell.execute_reply":"2022-12-23T14:51:28.072882Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for i in tqdm(range(len(dataset))):\n    long_answer_model = LongAnswerModel(long_md, device=device)\n    encoding, indices, example_id, question, answers = dataset[i]\n    answer = long_answer_model(encoding['input_ids'], encoding['attention_mask'])\n    long_prediction = answer.cpu().tolist()\n    long_answer_indices = get_inds(long_prediction, indices)\n    if long_answer_indices == '':\n        result[f'{example_id}_long'] = ''\n        result[f'{example_id}_short'] = ''\n        continue\n    long_answer = answers[np.argmax(long_prediction)]\n    result[f'{example_id}_long'] = long_answer_indices\n    dataset2.append((long_answer, question, example_id))\n    torch.cuda.empty_cache()","metadata":{"execution":{"iopub.status.busy":"2022-12-23T14:51:28.075983Z","iopub.execute_input":"2022-12-23T14:51:28.076418Z","iopub.status.idle":"2022-12-23T15:04:58.126968Z","shell.execute_reply.started":"2022-12-23T14:51:28.076382Z","shell.execute_reply":"2022-12-23T15:04:58.125411Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# torch.cuda.empty_cache()","metadata":{"execution":{"iopub.status.busy":"2022-12-23T15:04:58.128353Z","iopub.execute_input":"2022-12-23T15:04:58.129047Z","iopub.status.idle":"2022-12-23T15:04:58.134040Z","shell.execute_reply.started":"2022-12-23T15:04:58.128999Z","shell.execute_reply":"2022-12-23T15:04:58.132831Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"short_tokenizer = BertTokenizerFast.from_pretrained(bert_base_cased_filename)\ncol_spans = []\nfor i in range(1, 10):\n    col_spans.append(f'<Td_colspan=\"{i}\">')\n    col_spans.append(f'<Th_colspan=\"{i}\">')\nshort_tokenizer.add_tokens(['</Td>', '<Td>', '</Tr>', '<Tr>', '<Th>', '</Th>', '<Li>', '</Li>', '<Ul>', '</Ul>', '<Table>', '</Table>'])\nshort_tokenizer.add_tokens(col_spans)","metadata":{"execution":{"iopub.status.busy":"2022-12-23T15:04:58.135639Z","iopub.execute_input":"2022-12-23T15:04:58.136024Z","iopub.status.idle":"2022-12-23T15:04:58.204645Z","shell.execute_reply.started":"2022-12-23T15:04:58.135979Z","shell.execute_reply":"2022-12-23T15:04:58.203640Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"short_md = BertForTokenClassification.from_pretrained(bert_base_cased_filename, num_labels=2)\nshort_md.resize_token_embeddings(len(short_tokenizer))","metadata":{"execution":{"iopub.status.busy":"2022-12-23T15:04:58.205937Z","iopub.execute_input":"2022-12-23T15:04:58.206276Z","iopub.status.idle":"2022-12-23T15:04:59.931718Z","shell.execute_reply.started":"2022-12-23T15:04:58.206244Z","shell.execute_reply":"2022-12-23T15:04:59.930674Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"short_model_filename = '/kaggle/input/short-model-qa/short_model_2.pt'\nshort_md.load_state_dict(torch.load(short_model_filename, map_location=device))","metadata":{"execution":{"iopub.status.busy":"2022-12-23T15:04:59.933476Z","iopub.execute_input":"2022-12-23T15:04:59.933863Z","iopub.status.idle":"2022-12-23T15:05:00.372108Z","shell.execute_reply.started":"2022-12-23T15:04:59.933827Z","shell.execute_reply":"2022-12-23T15:05:00.370817Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"short_answer_model = ShortAnswerModel(short_md, device=device)","metadata":{"execution":{"iopub.status.busy":"2022-12-23T15:05:00.373993Z","iopub.execute_input":"2022-12-23T15:05:00.374389Z","iopub.status.idle":"2022-12-23T15:05:00.490349Z","shell.execute_reply.started":"2022-12-23T15:05:00.374347Z","shell.execute_reply":"2022-12-23T15:05:00.489415Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for long_answer, question, example_id in tqdm(dataset2):    \n    batch_long_answers, batch_questions = process_example(long_answer, question)\n    encodings, question_len = encode_short_ans(batch_long_answers, batch_questions)\n    question_len = len(question)\n    pred_starts, pred_ends = [], []\n    for i, enc in enumerate(encodings):\n        s, e = short_answer_model(enc, [question_len] * batch_size)[0]\n        s, e = int(s), int(e)\n        if s != -1:\n            s += (MAX_LEN-150)*i\n            e += (MAX_LEN-150)*i\n        pred_starts.append(s)\n        pred_ends.append(e)\n    long_start, _ = long_answer_indices.split(':')\n    pred_ans = join_test_example_preds(pred_starts, pred_ends, int(long_start))\n    result[f'{example_id}_short'] = pred_ans","metadata":{"execution":{"iopub.status.busy":"2022-12-23T15:05:00.495451Z","iopub.execute_input":"2022-12-23T15:05:00.495733Z","iopub.status.idle":"2022-12-23T15:05:19.187749Z","shell.execute_reply.started":"2022-12-23T15:05:00.495706Z","shell.execute_reply":"2022-12-23T15:05:19.186744Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# for ind in tqdm([2, 50, 59, 72, 177, 208, 216, 221, 282, 289, 293, 337]):\n#     encoding, indices, example_id, question, answers = dataset[ind]\n#     result[f'{example_id}_long'] = ''\n#     result[f'{example_id}_short'] = ''","metadata":{"execution":{"iopub.status.busy":"2022-12-23T15:05:19.189804Z","iopub.execute_input":"2022-12-23T15:05:19.190467Z","iopub.status.idle":"2022-12-23T15:05:19.199464Z","shell.execute_reply.started":"2022-12-23T15:05:19.190429Z","shell.execute_reply":"2022-12-23T15:05:19.197979Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"res = pd.Series(result, name='PredictionString')\nres.index.name = 'example_id'","metadata":{"execution":{"iopub.status.busy":"2022-12-23T15:05:19.202791Z","iopub.execute_input":"2022-12-23T15:05:19.203546Z","iopub.status.idle":"2022-12-23T15:05:19.212448Z","shell.execute_reply.started":"2022-12-23T15:05:19.203511Z","shell.execute_reply":"2022-12-23T15:05:19.210825Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"res.reset_index().sample(10)","metadata":{"execution":{"iopub.status.busy":"2022-12-23T15:05:19.214193Z","iopub.execute_input":"2022-12-23T15:05:19.214848Z","iopub.status.idle":"2022-12-23T15:05:19.235583Z","shell.execute_reply.started":"2022-12-23T15:05:19.214812Z","shell.execute_reply":"2022-12-23T15:05:19.234689Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# res.unique()","metadata":{"execution":{"iopub.status.busy":"2022-12-23T15:05:19.236726Z","iopub.execute_input":"2022-12-23T15:05:19.241308Z","iopub.status.idle":"2022-12-23T15:05:19.245145Z","shell.execute_reply.started":"2022-12-23T15:05:19.241275Z","shell.execute_reply":"2022-12-23T15:05:19.244047Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"res.to_csv('submission.csv')","metadata":{"execution":{"iopub.status.busy":"2022-12-23T15:05:19.246364Z","iopub.execute_input":"2022-12-23T15:05:19.247313Z","iopub.status.idle":"2022-12-23T15:05:19.259641Z","shell.execute_reply.started":"2022-12-23T15:05:19.247275Z","shell.execute_reply":"2022-12-23T15:05:19.258543Z"},"trusted":true},"execution_count":null,"outputs":[]}]}