{"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":"USE_POSTPROCESS = True\n\nFP16 = True","metadata":{"execution":{"iopub.status.busy":"2022-08-11T10:38:15.261104Z","iopub.execute_input":"2022-08-11T10:38:15.261713Z","iopub.status.idle":"2022-08-11T10:38:15.289421Z","shell.execute_reply.started":"2022-08-11T10:38:15.261632Z","shell.execute_reply":"2022-08-11T10:38:15.288715Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nos.environ[\"TOKENIZERS_PARALLELISM\"] = \"false\"\n\nos.listdir(\"../input/\")","metadata":{"execution":{"iopub.status.busy":"2022-08-11T10:38:44.528091Z","iopub.execute_input":"2022-08-11T10:38:44.528374Z","iopub.status.idle":"2022-08-11T10:38:44.536856Z","shell.execute_reply.started":"2022-08-11T10:38:44.528342Z","shell.execute_reply":"2022-08-11T10:38:44.536175Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import pathlib, json\n\nMK_SORT_MODEL_PATH = pathlib.Path(\"../input/ai4code-sub1-mksort-deb3l-ema/model_best-strictave-dev.pt\").resolve()\n\nmodel_paths_and_ensemble_weights = list()\nmodel_paths_and_ensemble_weights.append({\n    \"model_path\": pathlib.Path('../input/ai4code-sub1-disr-64-128/model_abest-strict-pp.pt').resolve(),\n    \"ensemble_weight\": 0.20,\n})\nmodel_paths_and_ensemble_weights.append({\n    \"model_path\": pathlib.Path('../input/ai4code-sub1-disr-l250-40-40/model_abest-strict-pp.pt').resolve(),\n    \"ensemble_weight\": 0.25,\n})\nmodel_paths_and_ensemble_weights.append({\n    \"model_path\": pathlib.Path('../input/ai4code-sub1-deb3l-ema/model_abest-pp.pt').resolve(),\n    \"ensemble_weight\": 0.55,\n})\n","metadata":{"execution":{"iopub.status.busy":"2022-08-11T10:39:21.591223Z","iopub.execute_input":"2022-08-11T10:39:21.591535Z","iopub.status.idle":"2022-08-11T10:39:21.613143Z","shell.execute_reply.started":"2022-08-11T10:39:21.591500Z","shell.execute_reply":"2022-08-11T10:39:21.612476Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(\"ensemble_weights:\", [params[\"ensemble_weight\"] for params in model_paths_and_ensemble_weights], \"  sum:\", sum(params[\"ensemble_weight\"] for params in model_paths_and_ensemble_weights))\nfor params in model_paths_and_ensemble_weights:\n    model_path = params[\"model_path\"]\n    assert model_path.parent.exists()\n\n    print(\"************************************\")\n    print(model_path)\n    with open(model_path.parent/\"log.txt\") as f:\n        print(f.read())\n    with open(model_path.parent/f'{model_path.stem.replace(\"model_\",\"\")}_log.txt') as f:\n        print(f.read().strip().splitlines()[-1])\n        print()\n    with open(model_path.parent/\"config.json\") as f:\n        d = json.load(f)\n        print(d)\n        print()\n        if \"note\" in d:\n            print(\"note:\", d[\"note\"])\n            print()\n","metadata":{"execution":{"iopub.status.busy":"2022-08-11T10:39:22.268963Z","iopub.execute_input":"2022-08-11T10:39:22.269381Z","iopub.status.idle":"2022-08-11T10:39:22.329472Z","shell.execute_reply.started":"2022-08-11T10:39:22.269337Z","shell.execute_reply":"2022-08-11T10:39:22.328719Z"},"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\nimport shutil\nfrom pathlib import Path\n\ntransformers_path = Path(\"/opt/conda/lib/python3.7/site-packages/transformers\")\n\ninput_dir = Path(\"../input/ai4code-deberta-v2-3-fast-tokenizer\")\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', \"deberta__init__.py\"]:\n    if str(filename).startswith(\"deberta\"):\n        filepath = deberta_v2_path/str(filename).replace(\"deberta\", \"\")\n    else:\n        filepath = deberta_v2_path/filename\n    if filepath.exists():\n        filepath.unlink()\n\n    shutil.copy(input_dir/filename, filepath)\nprint(\"ok\")","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-08-11T10:39:29.087711Z","iopub.execute_input":"2022-08-11T10:39:29.088061Z","iopub.status.idle":"2022-08-11T10:39:29.129592Z","shell.execute_reply.started":"2022-08-11T10:39:29.088018Z","shell.execute_reply":"2022-08-11T10:39:29.128849Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import sys, os\nimport time\nimport json, gzip\nimport collections\nimport pathlib\nimport math\nimport tqdm\nimport dataclasses as D\nfrom typing import List, Tuple, Optional, Dict, Any\n\nimport pandas as pd\nimport numpy as np\nimport torch\nimport transformers\ntransformers.logging.set_verbosity_error()\nprint(\"torch\", torch.__version__)\nprint(\"transformers\", transformers.__version__)\n\nimport sys\nsys.path.append(\"../input/ybaseline/\")\nimport baseline as B\nsys.path.pop()\n\nDEVICE = \"cuda:0\"\nBATCH_SIZE = 8\nROOT_PATH = pathlib.Path(\"../input/AI4Code\")\n","metadata":{"execution":{"iopub.status.busy":"2022-08-11T10:39:29.571275Z","iopub.execute_input":"2022-08-11T10:39:29.572082Z","iopub.status.idle":"2022-08-11T10:39:35.137393Z","shell.execute_reply.started":"2022-08-11T10:39:29.572041Z","shell.execute_reply":"2022-08-11T10:39:35.136625Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def map_transformers_name(name):\n    if \"../input/ai4code-pretrained-transformers\" in name:\n        return name.replace(\"../input/ai4code-pretrained-transformers\", \"../input/ai4code-pretrained-tokenizers\")\n    elif \"../input/transformers-settings\" in name:\n        return name\n    elif \"bio-lm\" in name:\n        return \"../input/transformers-settings/bio-lm/RoBERTa-large-PM-M3-Voc\"\n    else:\n        return os.path.join(\"../input/transformers-settings\", name)\n\n@D.dataclass\nclass BertTokenizerConfig:\n    bert_name: str\n    added_special_tokens: List[str]\n    def to_key(self):\n        return json.dumps(D.astuple(self))\n    @classmethod\n    def from_key(cls, key):\n        return cls(*json.loads(key))\n    @classmethod\n    def from_model_config(cls, model_config):\n        return cls(bert_name=model_config.bert_name, added_special_tokens=model_config.added_special_tokens)\n\nclass TestLoader:\n    def __init__(self):\n        self.bert_tokenizers = dict()\n\n    def get_bert_tokenizer(self, bert_tokenizer_key):\n        if bert_tokenizer_key not in self.bert_tokenizers:\n            bert_tokenizer_config = BertTokenizerConfig.from_key(bert_tokenizer_key)\n            print(f'load new bert tokenizer: {bert_tokenizer_config.bert_name}', file=sys.stderr, flush=True)\n            self.bert_tokenizers[bert_tokenizer_key], _, _ = B.utils.transformers_utils.from_pretrained_tokenizer(map_transformers_name(bert_tokenizer_config.bert_name), additional_special_tokens=bert_tokenizer_config.added_special_tokens, force_deberta_v2_fast=True)\n        return self.bert_tokenizers[bert_tokenizer_key]\n\n    def encode_cell(self, cell, bert_tokenizer_key):\n        tokenizer = self.get_bert_tokenizer(bert_tokenizer_key)\n        return {\n            \"cell_id\": cell[\"cell_id\"],\n            \"cell_type\": cell[\"cell_type\"],\n            \"input_ids\": tokenizer.encode(cell[\"source\"], add_special_tokens=False),\n        }\n\n    def read_json_file(self, fname:pathlib.Path, bert_tokenizer_key:str, cache:dict):\n        notebook_id = fname.stem\n        key = notebook_id + \"_\" + bert_tokenizer_key\n        if key not in cache:\n            cells = pd.read_json(fname, dtype={\"cell_type\":str, \"source\":str}, convert_axes=False).reset_index(drop=False).rename({\"index\":\"cell_id\"}, axis=1)\n            # cells[\"cell_type\"] = cells[\"cell_type\"].astype(\"category\")\n            code_cells = cells.loc[cells[\"cell_type\"] == \"code\"]\n            markdown_cells = cells.loc[cells[\"cell_type\"] == \"markdown\"]\n            cache[key] = {\n                \"bert_tokenizer\": self.get_bert_tokenizer(bert_tokenizer_key),\n                \"notebook_id\": notebook_id,\n                \"codes\": [self.encode_cell(cell, bert_tokenizer_key=bert_tokenizer_key) for cell in code_cells.to_dict(orient=\"index\").values()],\n                \"markdowns\": [self.encode_cell(cell, bert_tokenizer_key=bert_tokenizer_key) for cell in markdown_cells.to_dict(orient=\"index\").values()],\n            }\n        return cache[key]\n","metadata":{"execution":{"iopub.status.busy":"2022-08-11T10:39:35.140617Z","iopub.execute_input":"2022-08-11T10:39:35.141213Z","iopub.status.idle":"2022-08-11T10:39:35.156624Z","shell.execute_reply.started":"2022-08-11T10:39:35.141183Z","shell.execute_reply":"2022-08-11T10:39:35.155401Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from bisect import bisect\nimport typing\n\ndef count_inversions(a):\n    inversions = 0\n    sorted_so_far = []\n    for i, u in enumerate(a):\n        j = bisect(sorted_so_far, u)\n        inversions += i - j\n        sorted_so_far.insert(j, u)\n    return inversions\n\nValueType = typing.TypeVar(\"ValueType\")\ndef kendall_tau(ground_truth:List[List[ValueType]], predictions:List[List[ValueType]]):\n    assert len(ground_truth) == len(predictions)\n    total_inversions = 0\n    total_2max = 0  # twice the maximum possible inversions across all instances\n    for gt, pred in zip(ground_truth, predictions):\n        # ranks = [gt.index(x) for x in pred]  # rank predicted order in terms of ground truth\n        value_to_ground_rank = {value:rank for rank,value in enumerate(gt)}\n        ranks = [value_to_ground_rank[x] for x in pred]  # rank predicted order in terms of ground truth\n        total_inversions += count_inversions(ranks)\n        n = len(gt)\n        total_2max += n * (n - 1)\n    return 1 - 4 * total_inversions / total_2max\n","metadata":{"execution":{"iopub.status.busy":"2022-08-11T10:39:35.157908Z","iopub.execute_input":"2022-08-11T10:39:35.158151Z","iopub.status.idle":"2022-08-11T10:39:35.184613Z","shell.execute_reply.started":"2022-08-11T10:39:35.158119Z","shell.execute_reply":"2022-08-11T10:39:35.183609Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"loader = TestLoader()","metadata":{"execution":{"iopub.status.busy":"2022-08-11T10:39:35.190731Z","iopub.execute_input":"2022-08-11T10:39:35.191085Z","iopub.status.idle":"2022-08-11T10:39:35.194653Z","shell.execute_reply.started":"2022-08-11T10:39:35.191048Z","shell.execute_reply":"2022-08-11T10:39:35.193875Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if USE_POSTPROCESS:\n    def postprocess_cosine(cosines, temperature):\n        # logits: [M, C]\n        with torch.cuda.amp.autocast(enabled=False):\n            probs = (cosines.float()*temperature).softmax(-1) # [M, C]\n            r = torch.arange(probs.shape[1]) # [0, 1, ..., C-1]\n            rxr = (r[None] - r[:,None]).abs() # [C,C]\n            regrets = (probs.unsqueeze(-2) * rxr.to(probs)).sum(-1) # [M, C]\n        return -regrets\n\n    def postprocess_prob_np(probs):\n        # logits: [M, C]\n        r = np.arange(probs.shape[1]) # [0, 1, ..., C-1]\n        rxr = abs(r[None] - r[:,None]) # [C,C]\n        regrets = (probs[...,None,:] * rxr.astype(probs.dtype)).sum(-1) # [M, C]\n        return -regrets\n\nelse:\n    def postprocess_cosine(cosines, *args, **kwargs):\n        return cosines\n\n    def postprocess_prob_np(probs):\n        return probs\n","metadata":{"execution":{"iopub.status.busy":"2022-08-11T10:39:35.196023Z","iopub.execute_input":"2022-08-11T10:39:35.196553Z","iopub.status.idle":"2022-08-11T10:39:35.211758Z","shell.execute_reply.started":"2022-08-11T10:39:35.196518Z","shell.execute_reply":"2022-08-11T10:39:35.209353Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class AttrDict(dict):\n    def __getattr__(self, name):\n        if name in self:\n            return self[name]\n        else:\n            raise AttributeError()\n\nclass Model7(torch.nn.Module):\n    def __init__(self, model_config, bert_config_path:Optional[str]=None):\n        super(Model7, self).__init__()\n        self.model_config = model_config\n        if bert_config_path is not None:\n            self.bert = transformers.AutoModel.from_config(transformers.AutoConfig.from_pretrained(str(bert_config_path)))\n        else:\n            self.bert = transformers.AutoModel.from_pretrained(self.model_config.bert_name)\n            self.bert.resize_token_embeddings(self.bert.config.vocab_size + self.model_config.num_added_tokens)\n        if self.model_config.pool == \"attention\":\n            self.att1 = B.torch.layers.MultiHeadReduction(self.bert.config.hidden_size, self.model_config.num_head)\n        self.fc1 = torch.nn.Linear(self.bert.config.hidden_size, self.bert.config.hidden_size)\n        self.bert._init_weights(self.fc1)\n        self.dropout_output = torch.nn.Dropout(self.model_config.dropout_output)\n        self.fc_output_mk = torch.nn.Linear(self.bert.config.hidden_size, self.model_config.output_dim)\n        self.fc_output_code = torch.nn.Linear(self.bert.config.hidden_size, self.model_config.output_dim)\n        self.bert._init_weights(self.fc_output_mk)\n        self.bert._init_weights(self.fc_output_code)\n\n        self._freeze_bert = False\n\n        if self.model_config.code_fusion_type == \"none\":\n            pass\n        elif self.model_config.code_fusion_type == \"lstm-fc\":\n            self.code_rnn1 = torch.nn.LSTM(self.bert.config.hidden_size, self.bert.config.hidden_size, num_layers=self.model_config.code_rnn_num_layers, batch_first=True, bidirectional=True)\n            self.fc_code_rnn = torch.nn.Linear(2*self.bert.config.hidden_size, self.bert.config.hidden_size)\n        elif self.model_config.code_fusion_type == \"lstm-cnn\":\n            self.code_rnn1 = torch.nn.LSTM(self.bert.config.hidden_size, self.bert.config.hidden_size, num_layers=self.model_config.code_rnn_num_layers, batch_first=True, bidirectional=True)\n            CNN_KERNEL_SIZE = 3\n            CNN_PADDING_SIZE = (CNN_KERNEL_SIZE-1) // 2\n            self.cnn_code_rnn = torch.nn.Conv1d(2*self.bert.config.hidden_size, self.bert.config.hidden_size, CNN_KERNEL_SIZE, padding=CNN_PADDING_SIZE)\n        elif self.model_config.code_fusion_type == \"lstm-fc-res\":\n            self.code_rnn1 = torch.nn.LSTM(self.bert.config.hidden_size, self.bert.config.hidden_size, num_layers=self.model_config.code_rnn_num_layers, batch_first=True, bidirectional=True)\n            self.fc_code_rnn = torch.nn.Linear(2*self.bert.config.hidden_size, self.bert.config.hidden_size)\n        else:\n            raise ValueError(self.model_config.code_fusion_type)\n\n        if self.model_config.mk_fusion_type == \"none\":\n            pass\n        elif self.model_config.mk_fusion_type == \"attention\":\n            self.mk_fusion = torch.nn.MultiheadAttention(self.bert.config.hidden_size, self.model_config.mk_fusion_num_head, batch_first=True)\n        elif self.model_config.mk_fusion_type == \"attention-res\":\n            self.mk_fusion = torch.nn.MultiheadAttention(self.bert.config.hidden_size, self.model_config.mk_fusion_num_head, batch_first=True)\n            self.fc_att_mk = torch.nn.Linear(self.bert.config.hidden_size, self.bert.config.hidden_size)\n        else:\n            raise ValueError(self.model_config.mk_fusion_type)\n\n        self.actual_logit_temperature = getattr(self.model_config, \"logit_temperature\", 25.0)\n        self.actual_pre_code_strategy = getattr(self.model_config, \"pre_code_strategy\", \"first\")\n        self.actual_post_code_strategy = getattr(self.model_config, \"post_code_strategy\", \"first\")\n        print(\"self.actual_logit_temperature:\", self.actual_logit_temperature)\n\n    def set_freeze_bert(self, value:bool):\n        self._freeze_bert = value\n        self.bert.requires_grad_(not value)\n\n    def get_bert_encoded(self, input_ids, attention_mask, token_type_ids=None):\n        assert token_type_ids is None, \"not implemented because of def predict\"\n        h = self.bert(input_ids=input_ids, attention_mask=attention_mask, token_type_ids=token_type_ids).last_hidden_state # [B, seq_len, dim]\n\n        if self.model_config.pool == \"attention\":\n            h = self.att1(h, attention_mask)\n        elif self.model_config.pool == \"average\":\n            h = B.torch.utils.average_pool(h, attention_mask, axis=-2)\n        elif self.model_config.pool == \"cls\":\n            h = h[:,0]\n        else:\n            raise ValueError(self.model_config.pool)\n\n        if self.model_config.bert_encode_hidden == \"fc+tanh\":\n            h = self.fc1(h).tanh() # [B, dim]\n        elif self.model_config.bert_encode_hidden == \"fc\":\n            h = self.fc1(h) # [B, dim]\n        elif self.model_config.bert_encode_hidden == \"none\":\n            pass\n        else:\n            raise ValueError((self.model_config.bert_encode_hidden, getattr(self.model_config.bert_encode_hidden)))\n        return h\n\n    def _normalize(self, vs, axis=0, eps=1e-4):\n        m = vs.mean(axis).unsqueeze(axis) # [dim]\n        v = vs.var(axis).unsqueeze(axis) # [dim]\n        normalized = (vs-m) / (v+eps).sqrt()\n        return normalized\n\n    def get_code_output_vector(self, h):\n        # h: [N, dim]\n        outs = dict()\n        outs[\"bert_attended\"] = h\n\n        if self.model_config.normalize_code_bert_reprs:\n            h = self._normalize(h, axis=0) # [N, dim]\n        h = self.dropout_output(h)\n\n        if self.model_config.code_fusion_type == \"none\":\n            pass\n        elif self.model_config.code_fusion_type == \"lstm-fc\":\n            h, _ = self.code_rnn1(h.unsqueeze(0))\n            h = h.squeeze(0) # [N, dim]\n            h = self.fc_code_rnn(h)\n        elif self.model_config.code_fusion_type == \"lstm-cnn\":\n            h, _ = self.code_rnn1(h.unsqueeze(0)) # [1, N, dim]\n            h = self.cnn_code_rnn(h.transpose(1,2)).transpose(1,2) # [1,N,dim]\n            h = h.squeeze(0) # [N, dim]\n        elif self.model_config.code_fusion_type == \"lstm-fc-res\":\n            h_rnn, _ = self.code_rnn1(h.unsqueeze(0))\n            h_rnn = h_rnn.squeeze(0) # [N, dim]\n            h_rnn = self.fc_code_rnn(h_rnn)\n            h = h_rnn + h\n        else:\n            raise ValueError(self.model_config.code_fusion_type)\n        outs[\"fusioned\"] = h\n\n        h = self.fc_output_code(h)\n        outs[\"pre_normalization\"] = h\n        h = torch.nn.functional.normalize(h, dim=-1)\n        outs[\"output\"] = h\n        return outs\n\n    def fusion_markdown_vectors(self, h):\n        # h: [N, dim]\n        h = self.dropout_output(h)\n\n        if self.model_config.mk_fusion_type == \"none\":\n            pass\n        elif self.model_config.mk_fusion_type == \"attention\":\n            h_expanded = h.unsqueeze(0)\n            h, _ = self.mk_fusion(h_expanded,h_expanded,h_expanded) # [1, N, dim]\n            h = h.squeeze(0) # [N, dim]\n        elif self.model_config.mk_fusion_type == \"attention-res\":\n            h_expanded = h.unsqueeze(0)\n            h_att, _ = self.mk_fusion(h_expanded,h_expanded,h_expanded) # [1, N, dim]\n            h_att = h_att.squeeze(0) # [N, dim]\n            h_att = self.fc_att_mk(h_att) # [N, dim]\n            h = h_att + h\n        else:\n            raise ValueError(self.model_config.mk_fusion_type)\n        return h\n\n    def get_markdown_output_vector(self, h):\n        # h: [N, dim]\n        outs = dict()\n        outs[\"bert_attended\"] = h\n\n        h = self.fusion_markdown_vectors(h)\n        outs[\"fusioned\"] = h\n\n        h = self.fc_output_mk(h)\n        outs[\"pre_normalization\"] = h\n        h = torch.nn.functional.normalize(h, dim=-1)\n        outs[\"output\"] = h\n        return outs\n\n    def get_cosine(self, bert_encoded, num_markdowns, chunk_markdown:bool=True):\n        CHUNK_MARKDOWN_SIZE = 500\n        markdown_bert_h, code_bert_h = torch.split(bert_encoded, [num_markdowns, len(bert_encoded)-num_markdowns]) # [M, dim], [C, dim]\n        if chunk_markdown and (len(markdown_bert_h) > CHUNK_MARKDOWN_SIZE):\n            order = list(range(len(markdown_bert_h)))\n            np.random.RandomState(424242).shuffle(order)\n            num_chunk = math.ceil(len(order) / CHUNK_MARKDOWN_SIZE)\n            chunk_size = math.ceil(len(order) / num_chunk)\n            reverse_order = sorted(range(len(order)), key=lambda x:order[x])\n            chunked_markdown_outs = collections.defaultdict(list)\n            for c in range(0, len(order), chunk_size):\n                order_chunk = order[c:c+chunk_size]\n                for key, value in self.get_markdown_output_vector(markdown_bert_h[order_chunk]).items():\n                    chunked_markdown_outs[key].append(value)\n            markdown_outs = {key:torch.cat(value, 0)[reverse_order] for key,value in chunked_markdown_outs.items()}\n        else:\n            markdown_outs = self.get_markdown_output_vector(markdown_bert_h)\n        code_outs = self.get_code_output_vector(code_bert_h)\n\n        # cosines = (markdown_h[:,None] * code_h[None,:]).sum(-1) # [M, C]\n        cosines = torch.matmul(markdown_outs[\"output\"], code_outs[\"output\"].T) # [M, C]\n        return {\n            \"cosines\": cosines,\n            \"markdowns\": markdown_outs,\n            \"codes\": code_outs,\n        }\n\n    def forward(self, num_markdowns, input_ids, attention_mask, token_type_ids=None):\n        bert_encoded = self.get_bert_encoded(input_ids, attention_mask, token_type_ids=token_type_ids) # [B, dim]\n        cosines = self.get_cosine(bert_encoded=bert_encoded, num_markdowns=num_markdowns)[\"cosines\"]\n        return cosines\n\n    def build_test_dataloader(self, data, batch_size, device):\n        # TODO: half_and_half, post_only\n        assert self.actual_pre_code_strategy == \"first\"\n        assert self.actual_post_code_strategy == \"first\"\n        assert self.model_config.omit_strategy == \"none\"\n        codes_with_sentinel = [{\"input_ids\":self.model_config.blank_input_ids}] + data[\"codes\"] + [{\"input_ids\":self.model_config.blank_input_ids}]\n        markdowns = data[\"markdowns\"]\n        tok:transformers.BertTokenizer = data[\"bert_tokenizer\"]\n        instances = list()\n        for mcell in markdowns:\n            if self.model_config.double_mk_len:\n                mksrc = mcell[\"input_ids\"][:2*self.model_config.each_max_length]\n            else:\n                mksrc = mcell[\"input_ids\"][:self.model_config.each_max_length]\n            instances.append({\"input_ids\": tok.build_inputs_with_special_tokens(mksrc)})\n        for code_idx in range(len(codes_with_sentinel)-1):\n            presrc = codes_with_sentinel[code_idx][\"input_ids\"][:self.model_config.each_max_length]\n            postsrc = codes_with_sentinel[code_idx+1][\"input_ids\"][:self.model_config.each_max_length]\n            instances.append({\"input_ids\": tok.build_inputs_with_special_tokens(presrc, postsrc)})\n\n        backet_order = sorted(range(len(instances)), key=lambda x:len(instances[x][\"input_ids\"]))\n        instances = [instances[i] for i in backet_order]\n        reverse_backet_order = sorted(range(len(instances)), key=lambda x:backet_order[x])\n\n        dataloader = B.torch.utils.SelectiveDataset(instances, [\n            {\"name\":\"input_ids\", \"padding\":True, \"padding_value\":0, \"padding_mask\":True, \"dtype\":torch.long, \"device\":device},\n            ]).dataloader(batch_size, False)\n\n        return dataloader, reverse_backet_order, markdowns, codes_with_sentinel\n\n    def predict(self, data, batch_size, device, fp16, postprocess_cosine, return_probs=False, **kwargs):\n        dataloader, reverse_backet_order, markdowns, codes_with_sentinel = self.build_test_dataloader(data=data, batch_size=batch_size, device=device)\n\n        bert_reprs = list()\n        for minibatch in dataloader:\n            if fp16:\n                minibatch[\"input_ids_mask\"] = minibatch[\"input_ids_mask\"].to(torch.float16)\n            bert_reprs.append(self.get_bert_encoded(input_ids=minibatch[\"input_ids\"], attention_mask=minibatch[\"input_ids_mask\"]))\n        bert_reprs = torch.cat(bert_reprs, 0)\n        bert_reprs = bert_reprs[reverse_backet_order]\n\n        cosines = self.get_cosine(bert_encoded=bert_reprs, num_markdowns=len(markdowns))[\"cosines\"]\n        if return_probs:\n            probs = (cosines.float()*self.actual_logit_temperature).softmax(-1).cpu().detach().clone().numpy() # [M, C]\n            return probs\n        # probs = (cosines.float()*self.actual_logit_temperature).softmax(-1).cpu().detach().clone().numpy() # [M, C]\n        preds = postprocess_cosine(cosines=cosines, temperature=self.actual_logit_temperature).argmax(1).tolist() # [M]\n\n        insertions = [list() for _ in range(len(codes_with_sentinel)-1)]\n        for mcell, code_idx in zip(markdowns, preds):\n            insertions[code_idx].append(mcell[\"cell_id\"])\n\n        prediction = list()\n        for code_idx, code_cell in enumerate(data[\"codes\"]):\n            prediction.extend(insertions[code_idx])\n            prediction.append(code_cell[\"cell_id\"])\n        prediction.extend(insertions[code_idx+1])\n\n        # return prediction, probs, [mcell[\"cell_id\"] for mcell in markdowns], [code_cell[\"cell_id\"] for code_cell in data[\"codes\"]]\n        return prediction\n\n    def predict_for_ensemble(self, data, batch_size, device, fp16, return_probs:bool=True, return_cell_ids:bool=True, return_markdown_vectors:bool=True):\n        dataloader, reverse_backet_order, markdowns, codes_with_sentinel = self.build_test_dataloader(data=data, batch_size=batch_size, device=device)\n\n        bert_reprs = list()\n        for minibatch in dataloader:\n            if fp16:\n                minibatch[\"input_ids_mask\"] = minibatch[\"input_ids_mask\"].to(torch.float16)\n            bert_reprs.append(self.get_bert_encoded(input_ids=minibatch[\"input_ids\"], attention_mask=minibatch[\"input_ids_mask\"]))\n        bert_reprs = torch.cat(bert_reprs, 0)\n        bert_reprs = bert_reprs[reverse_backet_order]\n\n        model_outs = self.get_cosine(bert_encoded=bert_reprs, num_markdowns=len(markdowns))\n\n        outs = dict()\n\n        if return_probs:\n            probs = (model_outs[\"cosines\"].float()*self.actual_logit_temperature).softmax(-1).cpu().detach().clone().numpy() # [M, C]\n            outs[\"probs\"] = probs\n\n        if return_cell_ids:\n            markdown_cell_ids = [mk_cell[\"cell_id\"] for mk_cell in markdowns]\n            code_cell_ids = [code_cell[\"cell_id\"] for code_cell in data[\"codes\"]]\n            outs[\"cell_ids\"] = {\n                \"markdowns\": markdown_cell_ids,\n                \"codes\": code_cell_ids,\n            }\n\n        if return_markdown_vectors:\n            outs[\"markdown_vectors\"] = model_outs[\"markdowns\"]\n\n        return outs\n","metadata":{"execution":{"iopub.status.busy":"2022-08-11T10:39:35.213525Z","iopub.execute_input":"2022-08-11T10:39:35.213848Z","iopub.status.idle":"2022-08-11T10:39:35.314078Z","shell.execute_reply.started":"2022-08-11T10:39:35.213814Z","shell.execute_reply":"2022-08-11T10:39:35.313356Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def build_prediction(markdown_cell_ids, code_cell_ids, mk_idx_to_inserting_idx:List[int], markdown_vectors:dict):\n    # mk_idx_to_inserting_idx: [M]\n    insertions = [list() for _ in range(len(code_cell_ids) + 1)]\n    for mk_idx, markdown_cell_id in enumerate(markdown_cell_ids):\n        insertions[mk_idx_to_inserting_idx[mk_idx]].append(markdown_cell_id)\n\n    predicted_cell_ids_sequence = list()\n    for code_idx, code_cell_id in enumerate(code_cell_ids):\n        predicted_cell_ids_sequence.extend(insertions[code_idx])\n        predicted_cell_ids_sequence.append(code_cell_id)\n    predicted_cell_ids_sequence.extend(insertions[code_idx+1])\n    return predicted_cell_ids_sequence\n","metadata":{"execution":{"iopub.status.busy":"2022-08-11T10:39:35.315180Z","iopub.execute_input":"2022-08-11T10:39:35.315585Z","iopub.status.idle":"2022-08-11T10:39:35.322137Z","shell.execute_reply.started":"2022-08-11T10:39:35.315551Z","shell.execute_reply":"2022-08-11T10:39:35.321318Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class MKSortModel(torch.nn.Module):\n    def __init__(self, model_config):\n        super().__init__()\n        self.model_config = model_config\n        self.att1 = torch.nn.MultiheadAttention(self.model_config.input_dim, 4, batch_first=True)\n        self.fc1 = torch.nn.Linear(self.model_config.input_dim, self.model_config.input_dim)\n        self.fc2 = torch.nn.Linear(self.model_config.input_dim, 1)\n\n    def forward(self, inputs, mask):\n        # inputs: [B, S, dim]\n        h1 = inputs\n        h2 = self.att1(h1, h1, h1, mask)[0]\n        h3 = h1 + self.fc1(h2)\n        pred = self.fc2(h3).squeeze(-1) # [B, S, 1] -> [B, S]\n        return pred\n\nwith open(MK_SORT_MODEL_PATH.parent / \"model_config.json\") as f:\n    sorting_model_config = AttrDict(json.load(f))\nsorting_model = MKSortModel(sorting_model_config)\nsorting_model.load_state_dict(torch.load(MK_SORT_MODEL_PATH, map_location=\"cpu\"))\nif FP16:\n    sorting_model.to(torch.float16)\nsorting_model.to(DEVICE)\nsorting_model.eval()\n\nprint(MK_SORT_MODEL_PATH)\nprint(sorting_model_config)\n\ndef build_prediction(markdown_cell_ids, code_cell_ids, mk_idx_to_inserting_idx:List[int], markdown_vectors:dict):\n    # mk_idx_to_inserting_idx: [M]\n    insertion_idxs = [list() for _ in range(len(code_cell_ids) + 1)]\n    for mk_idx, markdown_cell_id in enumerate(markdown_cell_ids):\n        insertion_idxs[mk_idx_to_inserting_idx[mk_idx]].append(mk_idx)\n\n    fusioned_vectors = markdown_vectors[\"fusioned\"] # [M, dim]\n    sorting_input_idxs = list()\n    sorting_slot_idxs = list()\n    for slot_idx, insertion_idxs_i in enumerate(insertion_idxs):\n        if len(insertion_idxs_i) >= 2:\n            sorting_input_idxs.append(insertion_idxs_i)\n            sorting_slot_idxs.append(slot_idx)\n    if len(sorting_input_idxs) > 0:\n        padded_sorting_input_idxs, sorting_input_insertion_idxs_mask = B.utils.pad(sorting_input_idxs, 0)\n        sorting_input_insertion_idxs_mask_tensor = torch.FloatTensor(sorting_input_insertion_idxs_mask).to(fusioned_vectors) # [B, max_slot_size]\n        fusioned_vectors = fusioned_vectors[[idx for slot_idxs in padded_sorting_input_idxs for idx in slot_idxs]].reshape(len(padded_sorting_input_idxs), len(padded_sorting_input_idxs[0]), sorting_model_config.input_dim) # [B, max_slot_size, dim]\n\n        CHUNK_SLOT_SIZE = 100\n        if fusioned_vectors.shape[1] <= CHUNK_SLOT_SIZE:\n            sorting_scores = sorting_model(fusioned_vectors, mask=sorting_input_insertion_idxs_mask_tensor).detach().clone().cpu() # [B, max_slot_size]\n        else:\n            num_chunk = math.ceil(fusioned_vectors.shape[1] / CHUNK_SLOT_SIZE)\n            chunk_size = math.ceil(fusioned_vectors.shape[1] / num_chunk)\n            partial_sorting_scores = list()\n            for c in range(num_chunk):\n                start = c*chunk_size\n                end = (c+1)*chunk_size\n                partial_sorting_scores.append(\n                    sorting_model(fusioned_vectors[:,start:end], mask=sorting_input_insertion_idxs_mask_tensor[:,start:end]).detach().clone().cpu()\n                )\n            sorting_scores = torch.cat(partial_sorting_scores, 1) # [B, max_slot_size]\n        sorting_scores = sorting_scores.tolist()\n\n        for slot_idx, sorting_score in zip(sorting_slot_idxs, sorting_scores):\n            sorted_idxs_with_enum = sorted(enumerate(insertion_idxs[slot_idx]), key=lambda x:sorting_score[x[0]])\n            insertion_idxs[slot_idx] = [mk_idx for _,mk_idx in sorted_idxs_with_enum]\n\n    insertions = [[markdown_cell_ids[idx] for idx in slot] for slot in insertion_idxs]\n\n    predicted_cell_ids_sequence = list()\n    for code_idx, code_cell_id in enumerate(code_cell_ids):\n        predicted_cell_ids_sequence.extend(insertions[code_idx])\n        predicted_cell_ids_sequence.append(code_cell_id)\n    predicted_cell_ids_sequence.extend(insertions[code_idx+1])\n    return predicted_cell_ids_sequence\n","metadata":{"execution":{"iopub.status.busy":"2022-08-11T10:39:35.323299Z","iopub.execute_input":"2022-08-11T10:39:35.323696Z","iopub.status.idle":"2022-08-11T10:39:41.141132Z","shell.execute_reply.started":"2022-08-11T10:39:35.323664Z","shell.execute_reply":"2022-08-11T10:39:41.139546Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def build_model(model_path:pathlib.Path, ensemble_weight:float):\n    print(\"***********************\")\n    model_path = model_path.resolve()\n    print(model_path)\n    with open(model_path.parent / \"model_config.json\") as f:\n        model_config_dict = json.load(f)\n        model_config = AttrDict(model_config_dict)\n    bert_config_path = model_path.parent / \"bert_config\"\n    bert_tokenizer_key = BertTokenizerConfig.from_model_config(model_config).to_key()\n    if model_config.id == \"main7\":\n        model = Model7(model_config, bert_config_path=bert_config_path)\n    model.load_state_dict(torch.load(model_path, map_location=\"cpu\"))\n    if FP16:\n        model.to(torch.float16)\n    model.to(DEVICE)\n    model.eval()\n    print(model_config_dict)\n    return model, bert_tokenizer_key, ensemble_weight\n\nmodels = list()\nfor params in model_paths_and_ensemble_weights:\n    models.append(build_model(**params))\n","metadata":{"execution":{"iopub.status.busy":"2022-08-11T10:39:41.142693Z","iopub.execute_input":"2022-08-11T10:39:41.142961Z","iopub.status.idle":"2022-08-11T10:40:15.596513Z","shell.execute_reply.started":"2022-08-11T10:39:41.142919Z","shell.execute_reply":"2022-08-11T10:40:15.595697Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fnames = (ROOT_PATH / \"test\").resolve().iterdir()\n# fnames = list((ROOT_PATH / \"test\").iterdir()) + [fn for _,fn in zip(range(2000), (ROOT_PATH / \"train\").iterdir())]\n\nids = list()\ncell_orders = list()\nwith torch.no_grad(), B.utils.closing_tqdm(fnames) as pbar:\n    for fname in pbar:\n        probs = list()\n        cache = dict()\n        for model, bert_tokenizer_key, ensemble_weight in models:\n            target = loader.read_json_file(fname, bert_tokenizer_key, cache=cache)\n            prediction_outs = model.predict_for_ensemble(target, batch_size=BATCH_SIZE, device=DEVICE, fp16=FP16)\n            probs.append(prediction_outs[\"probs\"] * ensemble_weight)\n        probs = np.sum(probs, 0) # [M, C]\n\n        mk_idx_to_insertion_idx = postprocess_prob_np(probs).argmax(1).tolist() # [M]\n        predicted_cell_ids_sequence = build_prediction(markdown_cell_ids=prediction_outs[\"cell_ids\"][\"markdowns\"], code_cell_ids=prediction_outs[\"cell_ids\"][\"codes\"], mk_idx_to_inserting_idx=mk_idx_to_insertion_idx, markdown_vectors=prediction_outs[\"markdown_vectors\"])\n\n        ids.append(target[\"notebook_id\"])\n        cell_orders.append(\" \".join(predicted_cell_ids_sequence))\n\nsubmission_df = pd.DataFrame({\"id\":ids, \"cell_order\":cell_orders})\nsubmission_df\n","metadata":{"execution":{"iopub.status.busy":"2022-08-11T10:40:15.598950Z","iopub.execute_input":"2022-08-11T10:40:15.599387Z","iopub.status.idle":"2022-08-11T10:40:20.892804Z","shell.execute_reply.started":"2022-08-11T10:40:15.599349Z","shell.execute_reply":"2022-08-11T10:40:20.891976Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission_df.to_csv(\"submission.csv\", index=False)","metadata":{"execution":{"iopub.status.busy":"2022-08-11T10:40:20.894087Z","iopub.execute_input":"2022-08-11T10:40:20.895113Z","iopub.status.idle":"2022-08-11T10:40:20.904652Z","shell.execute_reply.started":"2022-08-11T10:40:20.895073Z","shell.execute_reply":"2022-08-11T10:40:20.903875Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}