{"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":"import time\nTOTAL_ALLOWABLE_RUNTIME=8*3600+50*60 # 8 HRS 40 MINS\nt0=time.time()\nprint(t0)","metadata":{"execution":{"iopub.status.busy":"2022-08-10T23:20:35.139440Z","iopub.execute_input":"2022-08-10T23:20:35.140268Z","iopub.status.idle":"2022-08-10T23:20:35.145727Z","shell.execute_reply.started":"2022-08-10T23:20:35.140228Z","shell.execute_reply":"2022-08-10T23:20:35.144438Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# The following is necessary if you want to use the fast tokenizer for deberta v2 or v3\n# This must be done before importing transformers\nimport shutil\nfrom pathlib import Path\n\ntransformers_path = Path(\"/opt/conda/lib/python3.7/site-packages/transformers\")\n\ninput_dir = Path(\"../input/deberta-v3-tokenizer-fast\")\n\nconvert_file = input_dir / \"convert_slow_tokenizer.py\"\nconversion_path = transformers_path/convert_file.name\n\nif conversion_path.exists():\n    conversion_path.unlink()\n\nshutil.copy(convert_file, transformers_path)\ndeberta_v2_path = transformers_path / \"models\" / \"deberta_v2\"\n\nfor filename in ['tokenization_deberta_v2.py', 'tokenization_deberta_v2_fast.py']:\n    filepath = deberta_v2_path/filename\n    \n    if filepath.exists():\n        filepath.unlink()\n\n    shutil.copy(input_dir/filename, filepath)","metadata":{"execution":{"iopub.status.busy":"2022-08-10T23:20:35.151289Z","iopub.execute_input":"2022-08-10T23:20:35.151990Z","iopub.status.idle":"2022-08-10T23:20:35.165686Z","shell.execute_reply.started":"2022-08-10T23:20:35.151944Z","shell.execute_reply":"2022-08-10T23:20:35.164952Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import json\nfrom pathlib import Path\n\nimport numpy as np\nimport pandas as pd\nfrom scipy import sparse\nfrom tqdm import tqdm\nfrom transformers.models.deberta_v2.tokenization_deberta_v2_fast import DebertaV2TokenizerFast\n\npd.options.display.width = 180\npd.options.display.max_colwidth = 120\ndata_dir = Path('../input/AI4Code')","metadata":{"papermill":{"duration":0.107446,"end_time":"2022-05-23T03:29:09.706238","exception":false,"start_time":"2022-05-23T03:29:09.598792","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-08-10T23:20:35.167517Z","iopub.execute_input":"2022-08-10T23:20:35.168208Z","iopub.status.idle":"2022-08-10T23:20:35.177318Z","shell.execute_reply.started":"2022-08-10T23:20:35.168172Z","shell.execute_reply":"2022-08-10T23:20:35.176290Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"SIMULATE_W_TRAIN=False\nCODE_MAX_LEN=26\nBS=2","metadata":{"execution":{"iopub.status.busy":"2022-08-10T23:20:35.179077Z","iopub.execute_input":"2022-08-10T23:20:35.179680Z","iopub.status.idle":"2022-08-10T23:20:35.187556Z","shell.execute_reply.started":"2022-08-10T23:20:35.179645Z","shell.execute_reply":"2022-08-10T23:20:35.186725Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def read_notebook(path):\n    return (\n        pd.read_json(\n            path,\n            dtype={'cell_type': 'category', 'source': 'str'})\n        .assign(id=path.stem)\n        .rename_axis('cell_id')\n    )\n\nif SIMULATE_W_TRAIN:\n    paths_test = list((data_dir / 'train').glob('*.json'))[:1000]\nelse:\n    paths_test = list((data_dir / 'test').glob('*.json'))\n    \nnotebooks_test = [\n    read_notebook(path) for path in tqdm(paths_test, desc='Test NBs')\n]\ntest_df = (\n    pd.concat(notebooks_test)\n    .set_index('id', append=True)\n    .swaplevel()\n    .sort_index(level='id', sort_remaining=False)\n).reset_index()\ntest_df[\"rank\"] = test_df.groupby([\"id\", \"cell_type\"]).cumcount()\ntest_df[\"pred\"] = test_df.groupby([\"id\", \"cell_type\"])[\"rank\"].rank(pct=True)","metadata":{"papermill":{"duration":0.114595,"end_time":"2022-05-23T03:29:09.832611","exception":false,"start_time":"2022-05-23T03:29:09.718016","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-08-10T23:20:35.190341Z","iopub.execute_input":"2022-08-10T23:20:35.190985Z","iopub.status.idle":"2022-08-10T23:20:42.771309Z","shell.execute_reply.started":"2022-08-10T23:20:35.190955Z","shell.execute_reply":"2022-08-10T23:20:42.770437Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Additional code cells\ndef clean_code(cell):\n    return str(cell).replace(\"\\\\n\", \"\\n\")\n\n\ndef sample_cells(cells, n, offset):\n    cells = [clean_code(cell) for cell in cells]\n    if n >= len(cells):\n        return [cell[:200] for cell in cells]\n    else:\n        results = []\n        step = len(cells) / n\n        idx = 0\n        if offset>0:\n            results=[cells[0]]\n        while int(np.round(idx)) < len(cells):\n            coor=min(int(np.round(idx))+offset,len(cells)-1)\n            results.append(cells[coor])\n            idx += step\n        assert cells[0] in results\n        if cells[-1] not in results:\n            results[-1] = cells[-1]\n        return results\n\n\ndef get_features(df, n_code, offset=0):\n    features = dict()\n    df = df.sort_values(\"rank\").reset_index(drop=True)\n    for idx, sub_df in tqdm(df.groupby(\"id\")):\n        features[idx] = dict()\n        total_md = sub_df[sub_df.cell_type == \"markdown\"].shape[0]\n        code_sub_df = sub_df[sub_df.cell_type == \"code\"]\n        total_code = code_sub_df.shape[0]\n        codes = sample_cells(code_sub_df.source.values, n_code, offset)\n        features[idx][\"total_code\"] = total_code\n        features[idx][\"total_md\"] = total_md\n        features[idx][\"codes\"] = codes\n    return features","metadata":{"papermill":{"duration":0.023767,"end_time":"2022-05-23T03:29:09.929422","exception":false,"start_time":"2022-05-23T03:29:09.905655","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-08-10T23:20:42.775943Z","iopub.execute_input":"2022-08-10T23:20:42.778266Z","iopub.status.idle":"2022-08-10T23:20:42.823312Z","shell.execute_reply.started":"2022-08-10T23:20:42.778224Z","shell.execute_reply":"2022-08-10T23:20:42.822397Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from tqdm import tqdm\nimport sys, os\nfrom transformers import AutoModel, AutoTokenizer\nimport torch.nn.functional as F\nimport torch.nn as nn\nimport torch\n\nclass MarkdownModel(nn.Module):\n    def __init__(self, model_path):\n        super(MarkdownModel, self).__init__()\n        model = AutoModel.from_pretrained(model_path)\n        self.model_path=model_path\n        self.embeddings=model.embeddings\n        self.encoder=model.encoder\n        self.top = nn.Linear(self.embeddings.word_embeddings.embedding_dim*2+5, self.embeddings.word_embeddings.embedding_dim)\n        self.gru = nn.GRU(self.embeddings.word_embeddings.embedding_dim,self.embeddings.word_embeddings.embedding_dim,bidirectional=True, dropout=0.2,batch_first=True)\n        self.top2 = nn.Linear(self.embeddings.word_embeddings.embedding_dim*2, 1)\n        #del self.model.pooler.dense.bias\n        #del self.model.pooler.dense.weight\n\n    def forward(self, ids, mask, fts, gather_ids, sample_ids, code_pct_ranks):\n        x=self.embeddings(ids)\n        # print(mask.shape)\n        # print(x.shape)\n        # print(mask)\n        #exit()\n        x=self.encoder(x,attention_mask=(mask==1),return_dict=False)[0]\n\n        preds=[]\n        code_preds=[]\n        start=0\n        for i in range(len(x)):\n\n            n_code=-(gather_ids[i]).min()\n            tmp=[]\n            for j in range(1,n_code+1):\n                vector=x[i][gather_ids[i]==-j]\n                mean_vector=vector.mean(0)\n                mean_vector=torch.cat((mean_vector,code_pct_ranks[start+j-1].reshape(1)))\n                tmp.append(mean_vector)\n            #code_preds.append(torch.stack(tmp))\n            code_preds=torch.stack(tmp)\n\n            # print(code_preds.shape)\n            # exit()\n\n            n_preds=gather_ids[i].max()\n            tmp=[]\n            start=0\n            for j in range(1,n_preds+1):\n                vector=x[i][gather_ids[i]==j]\n                # if return_vectors:\n                #     vectors.append(vector)\n                mean_vector=vector.mean(0)\n                mean_vector=torch.cat((mean_vector, fts[i]))\n\n                mean_vector=mean_vector.expand(len(code_preds),len(mean_vector)) #n_codexC\n                mean_vector=torch.cat([mean_vector,code_preds],1)\n                tmp.append(mean_vector)\n\n            tmp=torch.stack(tmp)\n            merged=F.relu(self.top(tmp))\n            merged=self.gru(merged)[0]\n            preds.append(self.top2(merged).squeeze(-1))\n\n\n\n            start+=len(code_preds)\n\n\n        return preds\n\nfrom torch.utils.data import DataLoader, Dataset\nfrom tqdm import tqdm\nimport numpy as np\nimport re\n\ndef preprocess_text(source):\n        # Remove all the special characters\n    source = re.sub(r'\\W', ' ', str(source))\n    #\n    # # remove all single characters\n    # document = re.sub(r'\\s+[a-zA-Z]\\s+', ' ', document)\n    #\n    # # Remove single characters from the start\n    # document = re.sub(r'\\^[a-zA-Z]\\s+', ' ', document)\n    #\n    # # Substituting multiple spaces with single space\n    # document = re.sub(r'\\s+', ' ', document, flags=re.I)\n    #\n    # Removing prefixed 'b'\n    source = re.sub(r'^b\\s+', '', source)\n    #\n    # # Converting to Lowercase\n    source = source.lower()\n    #pattern = r'\\<.*?\\>'\n    #document = re.sub(pattern, '', document)\n    # source=source.split('\\n')\n    # new_source=''\n    # for s in source:\n    #     new_source+=s[:128]\n    #     new_source+=' '\n    return source\n\n\nclass MarkdownDataset(Dataset):\n\n    def __init__(self, df, tokenizer, total_max_len, md_max_len, fts, preds_per_forward=5):\n        super().__init__()\n        self.df = df.reset_index(drop=True)\n        self.md_max_len = md_max_len\n        self.total_max_len = total_max_len  # maxlen allowed by model config\n\n        self.tokenizer=tokenizer\n\n        #self.tokenizer = AutoTokenizer.from_pretrained(model_name_or_path)\n        self.fts = fts\n        self.preds_per_forward=preds_per_forward\n        self.coordinates=[]\n        self.df_chuncks={}\n        self.lengths=[]\n        for group in tqdm(df.groupby('id')):\n\n\n        #for notebook_id in tqdm(df['id'].unique()):\n            #notebook_df=df[df['id']==notebook_id]\n            notebook_id=group[0]\n            notebook_df=group[1]\n            self.df_chuncks[notebook_id]=notebook_df\n            #self.coordinates.append(notebook_df)\n            n_preds=np.ceil(len(notebook_df)/self.preds_per_forward).astype('int')\n            for i in range(n_preds):\n                start=int(i*self.preds_per_forward)\n                end=min(int((i+1)*self.preds_per_forward),len(notebook_df))\n                self.coordinates.append({\"coor\":(start,end),\"id\":notebook_id})\n                df_chunck=self.df_chuncks[notebook_id].iloc[start:end]\n                \n                approx_length=0\n                \n#                 for source in df_chunck['source']:\n#                     approx_length+=min(64,len(source.split()))\n#                 for source in self.fts[notebook_id][\"codes\"]:\n#                     approx_length+=min(23,len(source.split()))    \n                \n                self.lengths.append(len(df_chunck['source'])*self.md_max_len+len(self.fts[notebook_id][\"codes\"])*CODE_MAX_LEN)\n                #self.lengths.append(approx_length)\n                #self.lengths.append(len(df_chunck['source'])+len(self.fts[notebook_id][\"codes\"]))\n                \n                \n        sorted_index=np.argsort(self.lengths)\n        self.coordinates=[self.coordinates[i] for i in sorted_index]\n                \n        #self.train=train\n        # print(df.iloc[0])\n        # print(self.coordinates[0])\n        # print(df.iloc[-1])\n        # print(self.coordinates[-1])\n        # print(len(self.coordinates))\n        # print(len(df))\n        # exit()\n\n    def __getitem__(self, index):\n        #row = self.df.iloc[index]\n\n        #df_chunck=self.coordinates[idx]\n        start,end=self.coordinates[index]['coor']\n        notebook_id=self.coordinates[index]['id']\n        df_chunck=self.df_chuncks[notebook_id].iloc[start:end]\n\n        #row=df_chunck.iloc[0]\n\n        # print(df_chunck)\n        # exit()\n        ids=[]\n        gather_ids=[]\n        mask=[]\n        for i,source in enumerate(df_chunck['source']):\n            inputs = self.tokenizer.encode_plus(\n                preprocess_text(source),\n                None,\n                add_special_tokens=True,\n                max_length=self.md_max_len,\n                padding=False,\n                return_token_type_ids=True,\n                truncation=True\n            )\n            ids+=inputs['input_ids']\n            ids+=[self.tokenizer.sep_token_id, ]\n            gather_ids+=[i+1]*len(inputs['input_ids'])\n            gather_ids+=[0]\n            mask+=inputs['attention_mask']\n            mask+=[1]\n\n\n        code_inputs = self.tokenizer.batch_encode_plus(\n            [preprocess_text(str(x)) for x in self.fts[notebook_id][\"codes\"]],\n            add_special_tokens=True,\n            max_length=CODE_MAX_LEN,\n            padding=False,\n            truncation=True\n        )\n\n\n        allowable_code_len=self.total_max_len-len(ids)\n\n        code_len=0\n        n_code=0\n        start=0\n\n        n_md = self.fts[notebook_id][\"total_md\"]\n        n_code = self.fts[notebook_id][\"total_md\"]\n\n        labels=list(df_chunck.pct_rank.values)\n        indices=list(df_chunck.index.values)\n\n        labels=torch.tensor(labels)\n        indices=torch.tensor(indices)\n\n        inputs=[]\n\n        for i in range(len(code_inputs['input_ids'])):\n            #print(i)\n            #print(len(code_inputs))\n            code_len+=len(code_inputs['input_ids'][i])\n            if code_len>allowable_code_len or i==(len(code_inputs['input_ids'])-1):\n                segment_ids=ids[:]\n                segment_gather_ids=gather_ids[:]\n                segment_mask=mask[:]\n                #print(start,i)\n                for j,x in enumerate(code_inputs['input_ids'][start:i+1]):\n                    segment_ids.extend(x[:-1])\n                    segment_gather_ids.extend(len(x[:-1])*[-j-1])\n                    segment_mask.extend(len(x[:-1])*[1])\n\n                if n_md + n_code == 0:\n                    fts = torch.FloatTensor([0,0,0,0])\n                else:\n                    fts = torch.FloatTensor([start/len(code_inputs['input_ids']), n_md/128, n_code/128, n_md / (n_md + n_code)])\n\n\n                segment_ids = torch.LongTensor(segment_ids[:self.total_max_len])\n                segment_gather_ids=torch.LongTensor(segment_gather_ids[:self.total_max_len])\n                segment_mask = torch.LongTensor(segment_mask[:self.total_max_len])\n                n_code=-segment_gather_ids.min()\n                code_pct_ranks = torch.tensor(np.arange(0,1,1/len(code_inputs['input_ids']))).float()[:n_code]\n\n                # print(code_pct_ranks)\n                # print(segment_gather_ids.min())\n                # exit()\n\n                inputs.append({\"segment_ids\":segment_ids,\"segment_gather_ids\":segment_gather_ids,\n                               \"segment_mask\":segment_mask,\"fts\":fts,\"labels\":labels,\n                               \"indices\":indices,\"code_pct_ranks\": code_pct_ranks})\n\n                start=i\n                code_len=len(code_inputs['input_ids'][i])\n                break\n            if len(inputs)>4:\n                break\n\n\n\n        return inputs\n\n    def __len__(self):\n        #return self.df.shape[0]\n        return len(self.coordinates)\n    \nclass CustomCollate:\n    def __init__(self,tokenizer):\n        self.tokenizer=tokenizer\n\n    def __call__(self,batch):\n        ids, mask, fts, gather_ids, indices, labels =[],[],[],[],[],[]\n        # print(len(batch))\n        # exit()\n        bs=len(batch)\n        lengths=[]\n        for i in range(bs):\n            for inputs in batch[i]:\n                lengths.append(len(inputs[\"segment_ids\"]))\n\n        max_len=max(lengths)\n\n        sample_ids=[]\n        code_pct_ranks=[]\n        for i,data in enumerate(batch):\n            for inputs in data:\n                sample_len=len(inputs['segment_ids'])\n                ids.append(torch.nn.functional.pad(inputs['segment_ids'],(0,max_len-sample_len),value=self.tokenizer.pad_token_id))\n                mask.append(torch.nn.functional.pad(inputs['segment_mask'],(0,max_len-sample_len),value=0))\n                gather_ids.append(torch.nn.functional.pad(inputs['segment_gather_ids'],(0,max_len-sample_len),value=0))\n                #mask.append(data[1])\n                fts.append(inputs['fts'])\n                #gather_ids.append(data[3])\n                sample_ids.append(i)\n\n            indices.append(inputs['indices'])\n            labels.append(inputs['labels'])\n            #code_pct_ranks.append(inputs['code_pct_ranks']) #lx1-1xd == lxd .min(1)\n            code_pct_ranks.append(inputs['code_pct_ranks']) #lx1-1xd == lxd .min(1)\n            #code_pct_ranks.append(torch.abs(inputs['code_pct_ranks'].reshape(-1,1)-inputs['labels'].reshape(1,-1)).min(1)[0]) #lx1-1xd == lxd .min(1)\n            # print(\"closest distance\")\n            # print(code_pct_ranks[-1])\n            # print(len(code_pct_ranks[-1]))\n            # print(\"code_pct_ranks\")\n            # print(inputs['code_pct_ranks'])\n            # print(len(inputs['code_pct_ranks']))\n            assert len(code_pct_ranks[-1])==len(inputs['code_pct_ranks'])\n        #exit()\n        ids=torch.stack(ids)#.cuda()\n        mask=torch.stack(mask)#.cuda()\n        fts=torch.stack(fts)#.cuda()\n        gather_ids=torch.stack(gather_ids)#.cuda()\n        indices=torch.cat(indices)\n        labels=torch.cat(labels)#.cuda()\n        sample_ids=torch.tensor(np.array(sample_ids))\n        code_pct_ranks=torch.cat(code_pct_ranks)\n\n        return ids, mask, fts, gather_ids, indices, labels, sample_ids, code_pct_ranks","metadata":{"papermill":{"duration":6.071788,"end_time":"2022-05-23T03:29:16.059249","exception":false,"start_time":"2022-05-23T03:29:09.987461","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-08-10T23:20:42.829468Z","iopub.execute_input":"2022-08-10T23:20:42.829927Z","iopub.status.idle":"2022-08-10T23:20:42.982295Z","shell.execute_reply.started":"2022-08-10T23:20:42.829894Z","shell.execute_reply":"2022-08-10T23:20:42.981339Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\n\ndef validate(model, val_loader):\n    model.eval()\n    \n    tbar = tqdm(val_loader, file=sys.stdout)\n    \n    preds = []\n    labels = []\n    rearrange_indices=[]\n    with torch.no_grad():\n        for idx, data in enumerate(tbar):\n            try:\n                ids, mask, fts, gather_ids, indices, target, sample_ids, code_pct_ranks = data\n                ids, mask, fts, gather_ids, sample_ids=ids.cuda(), mask.cuda(), fts.cuda(), gather_ids.cuda(), sample_ids.cuda()\n                code_pct_ranks = code_pct_ranks.cuda()\n                pred = model(ids, mask, fts, gather_ids, sample_ids, code_pct_ranks)\n\n                md_start=0\n                code_start=0\n\n                new_pred=[]\n                for p in pred:\n                    target_segment=target[md_start:md_start+p.shape[0]].expand(p.shape[1],p.shape[0]).permute(1,0)\n                    code_pct_ranks_segment=code_pct_ranks[code_start:code_start+p.shape[1]].expand(p.shape[0],p.shape[1])\n                    predicted=(code_pct_ranks_segment+p).mean(1)\n                    new_pred.append(predicted)\n                    md_start+=p.shape[0]\n                    code_start+=p.shape[1]\n\n                new_pred=torch.cat(new_pred)\n                pred=new_pred\n\n\n\n                preds.append(pred.detach().cpu().numpy().ravel())\n                labels.append(target.numpy().ravel())\n                rearrange_indices.append(indices.numpy().ravel())\n            except:\n                pass\n            \n            if (time.time()-t0)>TOTAL_ALLOWABLE_RUNTIME:\n                break\n            \n            \n            \n    labels,preds,rearrange_indices=np.concatenate(labels), np.concatenate(preds), np.concatenate(rearrange_indices)\n#     new_preds=[]\n#     for i in range(len(preds)):\n#         index=np.where(rearrange_indices==i)[0][0]\n#         #print(index)\n#         new_preds.append(preds[index])\n    n_preds=len(test_df[test_df[\"cell_type\"] == \"markdown\"])\n    new_preds=np.zeros(n_preds)\n    new_preds_count=np.zeros(n_preds)\n    \n    \n    for index,value in zip(rearrange_indices,preds):\n        new_preds[index]=value\n        new_preds_count[index]+=1\n    \n    return new_preds, new_preds_count\n\ndef predict(model_path, ckpt_path, total_max_len, TTA, n_code,preds_per_forward, md_max_len):\n    \n    test_fts = get_features(test_df,n_code)\n    model = MarkdownModel(model_path)\n    model = model.cuda()\n    model.eval()\n    model.load_state_dict(torch.load(ckpt_path))\n    #BS = 2\n    NW = 2\n    test_df[\"pct_rank\"] = 0\n    test_ds = MarkdownDataset(test_df[test_df[\"cell_type\"] == \"markdown\"].reset_index(drop=True).sample(frac=1,random_state=1000), \n                              md_max_len=md_max_len,total_max_len=total_max_len, tokenizer=tokenizer, \n                              fts=test_fts,preds_per_forward=preds_per_forward)\n    test_loader = DataLoader(test_ds, batch_size=BS, shuffle=False, num_workers=NW,\n                              pin_memory=False, drop_last=False,collate_fn=CustomCollate(tokenizer))\n    y_test, new_preds_count = validate(model, test_loader)\n    y_test_tta=[y_test]\n    y_test_cnt=[new_preds_count]\n    #print(y_test)\n    for i in range(TTA):\n        if (time.time()-t0)>TOTAL_ALLOWABLE_RUNTIME:\n            break\n        test_fts = get_features(test_df,n_code,i+1)\n        test_ds = MarkdownDataset(test_df[test_df[\"cell_type\"] == \"markdown\"].reset_index(drop=True).sample(frac=1,random_state=i+1),\n                                  md_max_len=md_max_len,total_max_len=total_max_len, tokenizer=tokenizer, \n                                  fts=test_fts,preds_per_forward=preds_per_forward)\n        test_loader = DataLoader(test_ds, batch_size=BS, shuffle=False, num_workers=NW,\n                              pin_memory=False, drop_last=False,collate_fn=CustomCollate(tokenizer))\n        y_test, new_preds_count = validate(model, test_loader)\n        y_test_tta.append(y_test)\n        y_test_cnt.append(new_preds_count)\n    #y_test_tta=np.stack(y_test_tta).mean(0)\n    y_test_tta=np.stack(y_test_tta).sum(0)/np.stack(y_test_cnt).sum(0)\n    #print(y_test_tta)\n    return y_test_tta","metadata":{"papermill":{"duration":0.027037,"end_time":"2022-05-23T03:29:16.100816","exception":false,"start_time":"2022-05-23T03:29:16.073779","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-08-10T23:20:42.983629Z","iopub.execute_input":"2022-08-10T23:20:42.984052Z","iopub.status.idle":"2022-08-10T23:20:43.010814Z","shell.execute_reply.started":"2022-08-10T23:20:42.984018Z","shell.execute_reply":"2022-08-10T23:20:43.010048Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model_path = \"../input/deberta-v3-large\"\nckpt_path = \"../input/ai4code-test32-20-preds-80-code/model0.bin\"\nif \"deberta-v3\" in model_path or \"deberta-v2\" in model_path:\n    tokenizer = DebertaV2TokenizerFast.from_pretrained(model_path)\nelse:\n    tokenizer = AutoTokenizer.from_pretrained(model_path, )\ny_test = predict(model_path, ckpt_path, total_max_len=4250, TTA=3, n_code=100, preds_per_forward=30, md_max_len=80)\n#y_test = predict(model_path, ckpt_path, total_max_len=3200, TTA=4, n_code=200, preds_per_forward=30, md_max_len=64)\ny_test","metadata":{"papermill":{"duration":20.374059,"end_time":"2022-05-23T03:29:36.523092","exception":false,"start_time":"2022-05-23T03:29:16.149033","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-08-10T23:20:43.012229Z","iopub.execute_input":"2022-08-10T23:20:43.012745Z","iopub.status.idle":"2022-08-10T23:20:57.016487Z","shell.execute_reply.started":"2022-08-10T23:20:43.012655Z","shell.execute_reply":"2022-08-10T23:20:57.014079Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"y_test","metadata":{"execution":{"iopub.status.busy":"2022-08-10T23:20:57.017702Z","iopub.status.idle":"2022-08-10T23:20:57.018283Z","shell.execute_reply.started":"2022-08-10T23:20:57.018031Z","shell.execute_reply":"2022-08-10T23:20:57.018061Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_df.loc[test_df[\"cell_type\"] == \"markdown\", \"pred\"] = y_test","metadata":{"papermill":{"duration":0.024872,"end_time":"2022-05-23T03:29:36.604626","exception":false,"start_time":"2022-05-23T03:29:36.579754","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-08-10T23:20:57.019520Z","iopub.status.idle":"2022-08-10T23:20:57.020083Z","shell.execute_reply.started":"2022-08-10T23:20:57.019845Z","shell.execute_reply":"2022-08-10T23:20:57.019874Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub_df = test_df.sort_values(\"pred\").groupby(\"id\")[\"cell_id\"].apply(lambda x: \" \".join(x)).reset_index()\nsub_df.rename(columns={\"cell_id\": \"cell_order\"}, inplace=True)\nsub_df.head()","metadata":{"papermill":{"duration":0.033756,"end_time":"2022-05-23T03:29:36.655157","exception":false,"start_time":"2022-05-23T03:29:36.621401","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-08-10T23:20:57.021197Z","iopub.status.idle":"2022-08-10T23:20:57.021780Z","shell.execute_reply.started":"2022-08-10T23:20:57.021542Z","shell.execute_reply":"2022-08-10T23:20:57.021568Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub_df.to_csv(\"submission.csv\", index=False)","metadata":{"papermill":{"duration":0.027227,"end_time":"2022-05-23T03:29:36.699558","exception":false,"start_time":"2022-05-23T03:29:36.672331","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-08-10T23:20:57.022888Z","iopub.status.idle":"2022-08-10T23:20:57.023459Z","shell.execute_reply.started":"2022-08-10T23:20:57.023209Z","shell.execute_reply":"2022-08-10T23:20:57.023237Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"papermill":{"duration":0.017111,"end_time":"2022-05-23T03:29:36.734096","exception":false,"start_time":"2022-05-23T03:29:36.716985","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]}]}