{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":12863,"databundleVersionId":788719,"sourceType":"competition"},{"sourceId":2233309,"sourceType":"datasetVersion","datasetId":1335671},{"sourceId":14366662,"sourceType":"datasetVersion","datasetId":9137949}],"dockerImageVersionId":31193,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"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\nprint('walking dir')\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","trusted":true,"execution":{"iopub.status.busy":"2026-01-01T23:23:28.875642Z","iopub.execute_input":"2026-01-01T23:23:28.875847Z","iopub.status.idle":"2026-01-01T23:23:30.781823Z","shell.execute_reply.started":"2026-01-01T23:23:28.875830Z","shell.execute_reply":"2026-01-01T23:23:30.781017Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## import os\nos.environ['PROTOCOL_BUFFERS_PYTHON_IMPLEMENTATION'] = 'python'\n","metadata":{"execution":{"iopub.status.busy":"2025-11-30T20:51:55.291431Z","iopub.execute_input":"2025-11-30T20:51:55.291823Z","iopub.status.idle":"2025-11-30T20:51:55.295081Z","shell.execute_reply.started":"2025-11-30T20:51:55.291805Z","shell.execute_reply":"2025-11-30T20:51:55.294372Z"}}},{"cell_type":"code","source":"print('unin proto')\n!pip uninstall -y tensorflow protobuf","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-01T06:39:16.095738Z","iopub.execute_input":"2026-01-01T06:39:16.096034Z","iopub.status.idle":"2026-01-01T06:39:38.315949Z","shell.execute_reply.started":"2026-01-01T06:39:16.096014Z","shell.execute_reply":"2026-01-01T06:39:38.315058Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print('import')\nfrom transformers import AutoTokenizer\nfrom transformers import AutoModelForQuestionAnswering\nfrom torch.optim import Adam\nimport torch\nimport pandas as pd\nimport json\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-01T06:39:38.317577Z","iopub.execute_input":"2026-01-01T06:39:38.317883Z","iopub.status.idle":"2026-01-01T06:39:48.673071Z","shell.execute_reply.started":"2026-01-01T06:39:38.317850Z","shell.execute_reply":"2026-01-01T06:39:48.672475Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print('load test')\ntest=[]\nwith open('/kaggle/input/tensorflow2-question-answering/simplified-nq-test.jsonl') as f:\n    for line in f.readlines():\n        test.append(json.loads(line))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-01T06:39:48.673777Z","iopub.execute_input":"2026-01-01T06:39:48.674225Z","iopub.status.idle":"2026-01-01T06:39:49.044446Z","shell.execute_reply.started":"2026-01-01T06:39:48.674199Z","shell.execute_reply":"2026-01-01T06:39:49.043941Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print('load model')\nfrom transformers import AutoTokenizer, DistilBertModel\n\nMODEL_DIR = \"/kaggle/input/huggingface-bert-variants/\"\ntokenizer = AutoTokenizer.from_pretrained('/kaggle/input/huggingface-bert-variants/distilbert-base-uncased/distilbert-base-uncased')\nmodel = DistilBertModel.from_pretrained('/kaggle/input/huggingface-bert-variants/distilbert-base-uncased/distilbert-base-uncased')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-01T06:39:49.046182Z","iopub.execute_input":"2026-01-01T06:39:49.046450Z","iopub.status.idle":"2026-01-01T06:39:54.194919Z","shell.execute_reply.started":"2026-01-01T06:39:49.046433Z","shell.execute_reply":"2026-01-01T06:39:54.194102Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_context_start(tokenizer, input_ids):\n    for i,v in enumerate(input_ids):\n        if tokenizer.decode(v) == '[SEP]':\n            return i+1","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-01T06:39:54.195757Z","iopub.execute_input":"2026-01-01T06:39:54.196147Z","iopub.status.idle":"2026-01-01T06:39:54.200510Z","shell.execute_reply.started":"2026-01-01T06:39:54.196128Z","shell.execute_reply":"2026-01-01T06:39:54.199642Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_char_pos_for_candidate(lc, orig_tok):    \n    start_idx= 0    \n    for i in range(0,lc['start_token']):\n        start_idx += len(orig_tok[i]) + 1\n    end_idx = start_idx\n    for i in range(lc['start_token'],lc['end_token']+1):\n        end_idx += len(orig_tok[i]) + 1        \n    return (start_idx, end_idx)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-01T06:39:54.201481Z","iopub.execute_input":"2026-01-01T06:39:54.201685Z","iopub.status.idle":"2026-01-01T06:39:54.222587Z","shell.execute_reply.started":"2026-01-01T06:39:54.201668Z","shell.execute_reply":"2026-01-01T06:39:54.221937Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_start_end_logit_indices(offset, start, end, context_start):\n    idx_start = -1\n    idx_end = -1    \n    for i in range(context_start, len(offset)):\n        if idx_start == -1:\n            res = offset[i]            \n            if res[0]<=start and start <= res[1]:                \n                idx_start = i\n        if idx_end == -1:\n            res = offset[i]\n            if res[0]<=end and end <=res[1]:                \n                idx_end = i\n    return idx_start, idx_end","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-01T06:39:54.223408Z","iopub.execute_input":"2026-01-01T06:39:54.223664Z","iopub.status.idle":"2026-01-01T06:39:54.242870Z","shell.execute_reply.started":"2026-01-01T06:39:54.223640Z","shell.execute_reply":"2026-01-01T06:39:54.242306Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def calc_best_end(end_logits, s, e):\n    best_end = {}\n    for i in range(s,e+1):\n        best_end[i] = i\n    best_end[e] = e\n    best_end_idx = e\n    for i in range (e,s, -1):\n        if(end_logits[best_end[i]] < end_logits[best_end_idx]):\n            best_end[i] = best_end_idx\n        best_end_idx = best_end[i]\n    return best_end","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-01T06:39:54.243524Z","iopub.execute_input":"2026-01-01T06:39:54.244046Z","iopub.status.idle":"2026-01-01T06:39:54.258760Z","shell.execute_reply.started":"2026-01-01T06:39:54.244012Z","shell.execute_reply":"2026-01-01T06:39:54.258154Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_best_short_answer_idx(s, e, st_logits, end_logits, best_end):    \n    best_score = st_logits[s] + end_logits[e]\n    bs = s\n    be = e    \n    for i in range(s,e+1):                \n        score = st_logits[i] + end_logits[best_end[i]]\n        if score > best_score:\n            bs = i\n            be = best_end[i]\n            best_score = score\n    return bs, be, best_score\n            ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-01T06:39:54.259545Z","iopub.execute_input":"2026-01-01T06:39:54.259758Z","iopub.status.idle":"2026-01-01T06:39:54.273480Z","shell.execute_reply.started":"2026-01-01T06:39:54.259738Z","shell.execute_reply":"2026-01-01T06:39:54.272800Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_best_short_answer(st, e, off, origtok, start):    \n    sc = off[st][0] + start\n    ec = off[e][1] + start    \n    idxs = -1\n    idxe = -1\n    cur = 0\n    for i in range (0, len(origtok)): \n        if idxs == -1 and cur>=sc:\n            idxs = i\n        if idxe == -1 and cur>=ec:    \n            idxe = i\n        cur+= len(origtok[i]) + 1\n    return (idxs, idxe)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-01T06:39:54.275101Z","iopub.execute_input":"2026-01-01T06:39:54.275473Z","iopub.status.idle":"2026-01-01T06:39:54.289101Z","shell.execute_reply.started":"2026-01-01T06:39:54.275457Z","shell.execute_reply":"2026-01-01T06:39:54.288529Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nclass MyModule(torch.nn.Module):\n    def __init__(self, model, tokenizer, classification_layer, start_span_layer, end_span_layer):\n        super(MyModule, self).__init__()\n        self.model = model\n        self.classification_layer = classification_layer\n        self.start_span_layer = start_span_layer\n        self.end_span_layer = end_span_layer\n        self.tokenizer = tokenizer\n\n    def forward(self, ex):\n        sl, el, off, cls = self._compute(ex)        \n        loss = self._calc_loss(ex, sl, el, off, cls) \n        return loss\n    \n    def _compute(self, ex):\n        chunk_size = 250\n        batch_size = 3\n        off = []\n        sl = []\n        el = [] \n        cls = []\n        context_start = None\n        for i in range(0, len(ex['document_text']), 250):\n            tokz = self.tokenizer(ex['question_text'], ex['document_text'][i:i+250], return_offsets_mapping=True)\n            inp = torch.tensor(tokz['input_ids']).unsqueeze(0).to(device)\n            if context_start is None:\n                context_start = get_context_start(tokenizer, tokz['input_ids'])    \n            atn = torch.tensor(tokz['attention_mask']).unsqueeze(0).to(device)\n            res = self.model(input_ids=inp, attention_mask=atn)\n            c = res.last_hidden_state[:, 0, :]\n                        \n            classification = self.classification_layer(c)\n            cls.append(torch.squeeze(classification).to(device))\n            \n            res.last_hidden_state.shape\n            \n            start_logits =  self.start_span_layer(res.last_hidden_state)\n            sl.append(torch.squeeze(start_logits[0,context_start:,0]).view(-1))\n            end_logits = self.end_span_layer(res.last_hidden_state)\n            el.append(torch.squeeze(end_logits[0,context_start:,0]).view(-1))\n            \n            temp = []    \n            for o in tokz['offset_mapping'][context_start:]:                \n                temp.append((o[0] + i, o[1] + i))\n\n            temp_tensor = torch.tensor(temp)\n            if temp_tensor.dim() == 1:\n                temp_tensor = temp_tensor.unsqueeze(0)\n            off.append(temp_tensor)            \n                \n        sl = torch.cat(sl, dim=-1)\n        el = torch.cat(el, dim=-1)\n        off = torch.cat(off, dim=0)\n        cls = torch.vstack(cls)\n        cls = torch.mean(cls, dim=0)        \n        return sl, el, off, cls\n\n    def _calc_loss(self, ex, sl, el, off, cls):    \n        #sl = sl.unsqueeze(0)\n        #el = el.unsqueeze(0)\n        anno = ex['annotations'][0]\n        cx = get_class(ex).to(device)\n        classification_loss = torch.nn.CrossEntropyLoss()\n        total = classification_loss(cls.to(device), cx)\n        \n        if cx > 0:\n            ls, le =  map_to_model_space(off, anno['long_answer']['start_token'], \n                                         anno['long_answer']['end_token'])      \n            long_loss = torch.nn.CrossEntropyLoss()\n            total += long_loss(sl, torch.tensor(ls).to(device))\n            total += long_loss(sl, torch.tensor(le).to(device))\n        \n        if cx == 1:\n            ss, se = map_to_model_space(off, ex['annotations'][0]['short_answers'][0]['start_token'], \n                               ex['annotations'][0]['short_answers'][0]['end_token'])    \n            short_loss = torch.nn.CrossEntropyLoss()\n            total += short_loss(sl, torch.tensor(ss).to(device))\n            total += short_loss(el, torch.tensor(se).to(device))\n        return total","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-01T06:39:54.289726Z","iopub.execute_input":"2026-01-01T06:39:54.289955Z","iopub.status.idle":"2026-01-01T06:39:54.305931Z","shell.execute_reply.started":"2026-01-01T06:39:54.289940Z","shell.execute_reply":"2026-01-01T06:39:54.305257Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"classification_layer  = torch.nn.Linear(768, 5)\nstart_span_layer = torch.nn.Linear(768,1)\nend_span_layer = torch.nn.Linear(768,1)\nmy_module = MyModule(model, tokenizer, classification_layer, start_span_layer, end_span_layer)\nmy_module.load_state_dict(torch.load('/kaggle/input/tf-qa-100-1/trained_model.pth', map_location=torch.device('cpu')), strict=False)\nmy_module = my_module.to(device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-01T06:39:54.306653Z","iopub.execute_input":"2026-01-01T06:39:54.306878Z","iopub.status.idle":"2026-01-01T06:39:55.693766Z","shell.execute_reply.started":"2026-01-01T06:39:54.306853Z","shell.execute_reply":"2026-01-01T06:39:55.693171Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"ids = []\npreds = []\ntorch.cuda.empty_cache()\ntorch.set_grad_enabled(False)\nimport numpy as np\n\nfor idx, t in enumerate(test):\n    ds=[]\n    de=[]\n    try:\n        ogtok = t['document_text'].split()\n        for index, lc in enumerate(t['long_answer_candidates']):\n            s,e =  get_char_pos_for_candidate(lc, ogtok)\n            ds.append(s)\n            de.append(e)\n        start = int(np.percentile(ds,15))\n        end = int(np.percentile(de, 75))        \n        if end-start > 15000:\n            median = np.median(ds + de)\n            #print(f\"media is {median}\")\n            start = max(median-7500, 0)\n            end = min(median +7500, len(t['document_text'])-1)\n        start = int(start)\n        end = int(end)    \n        #print(f\"processing example {idx} with index {start} and {end}\")\n        ex = {}\n        ex['question_text'] = t['question_text']\n        ex['document_text'] = t['document_text'][start:end+1]    \n        sl, el, off,cls = my_module._compute(ex)        \n        cls_index = torch.argmax(cls)\n        # no answer\n        if cls_index == 0:\n            print('exiting because class index is 0')\n            #print(cls)\n            ids.append(f\"{t['example_id']}_long\")\n            ids.append(f\"{t['example_id']}_short\")\n            preds.append('')\n            preds.append('')\n            continue\n            \n        max_score = 0\n        best_index = 0\n        best_long_st = -1\n        best_long_e = -1\n        \n        for index, lc in enumerate(t['long_answer_candidates']):\n            sc,ec = get_char_pos_for_candidate(lc, ogtok)\n            if sc < start or ec > end:\n                continue # outisde truncate doc range\n            delta = ec - sc\n            sc = sc-start\n            ec = sc + delta\n            st_l, end_l = get_start_end_logit_indices(off,sc,ec,0)\n            score = sl[st_l] + el[end_l]\n            if index == 0 or score > max_score:\n                max_score = score\n                best_index = index\n                best_long_st = st_l\n                best_long_e = end_l\n            \n        best_long_answer = t['long_answer_candidates'][best_index] \n        ids.append(f\"{t['example_id']}_long\")\n        preds.append(f\"{best_long_answer['start_token']}:{best_long_answer['end_token']}\")\n    \n        if cls_index == 4:\n            ids.append(f\"{t['example_id']}_short\")\n            preds.append('')\n            continue\n        elif cls_index == 2:\n            ids.append(f\"{t['example_id']}_short\")\n            preds.append('YES')\n        elif cls_index == 3:\n            ids.append(f\"{t['example_id']}_short\")\n            preds.append('NO')\n        else:\n            best_end = calc_best_end(el, best_long_st, best_long_e)\n            best_short_st, best_short_e, best_short_score = get_best_short_answer_idx(best_long_st, best_long_e, sl, el, best_end)        \n            best_short_answer = get_best_short_answer(best_short_st, best_short_e, off, ogtok, start)\n            ids.append(f\"{t['example_id']}_short\")\n            preds.append(f\"{best_short_answer[0]}:{best_short_answer[1]}\")\n        torch.cuda.empty_cache()\n    except Exception as e:\n        ids.append(f\"{t['example_id']}_long\")\n        ids.append(f\"{t['example_id']}_short\")\n        preds.append('')\n        preds.append('')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-01T06:42:15.148795Z","iopub.execute_input":"2026-01-01T06:42:15.149594Z","iopub.status.idle":"2026-01-01T06:42:29.523704Z","shell.execute_reply.started":"2026-01-01T06:42:15.149570Z","shell.execute_reply":"2026-01-01T06:42:29.523123Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print('convert')\nsub = {'example_id': ids, 'PredictionString': preds}\nframe = pd.DataFrame(sub)\nframe.to_csv('/kaggle/working/submission.csv', index=False)","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}