{"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\nimport yaml\nimport time\n# for 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-08-11T14:22:25.850647Z","iopub.execute_input":"2022-08-11T14:22:25.857280Z","iopub.status.idle":"2022-08-11T14:22:25.932927Z","shell.execute_reply.started":"2022-08-11T14:22:25.857230Z","shell.execute_reply":"2022-08-11T14:22:25.931373Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"start_time = time.time()\nmax_duration_s = 7.5*60*60\nstart_time","metadata":{"execution":{"iopub.status.busy":"2022-08-11T14:22:25.935235Z","iopub.execute_input":"2022-08-11T14:22:25.936137Z","iopub.status.idle":"2022-08-11T14:22:25.952236Z","shell.execute_reply.started":"2022-08-11T14:22:25.936093Z","shell.execute_reply":"2022-08-11T14:22:25.950379Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"os.environ[\"TOKENIZERS_PARALLELISM\"] = \"false\"\n\n!mkdir -p /kaggle/temp/\n!rm -rf /kaggle/temp/*\n!ls -lh /kaggle/temp/","metadata":{"execution":{"iopub.status.busy":"2022-08-11T14:22:25.954706Z","iopub.execute_input":"2022-08-11T14:22:25.955724Z","iopub.status.idle":"2022-08-11T14:22:32.149810Z","shell.execute_reply.started":"2022-08-11T14:22:25.955684Z","shell.execute_reply":"2022-08-11T14:22:32.147798Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from tqdm import tqdm\nimport transformers\n\ntransformers.__version__","metadata":{"execution":{"iopub.status.busy":"2022-08-11T14:22:32.155638Z","iopub.execute_input":"2022-08-11T14:22:32.157180Z","iopub.status.idle":"2022-08-11T14:22:32.493248Z","shell.execute_reply.started":"2022-08-11T14:22:32.157134Z","shell.execute_reply":"2022-08-11T14:22:32.489343Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import sys\nsys.path.append('../input/ai4code-src-v2/src')\n\nmodel1_f0_dir = '../input/ai4code-model344-f0-770/344_single_bert_l2_max_loss_0.01_0_770'\nmodel1_f1_dir = '../input/ai4code-model342-f1-770/342_single_bert_l2_scaled_att_1_770'\nmodel1_f2_dir = '../input/ai4code-model342-f2-770/342_single_bert_l2_scaled_att_2_770'\nmodel1_f3_dir = '../input/ai4code-model342-f3-770/342_single_bert_l2_scaled_att_3_770'\n\nmodel2_f0_dir = \"../input/ai4code-model356-f0-798/356_bert_mpnet_l2_madgrad_0_798\"\nmodel2_f1_dir = \"../input/ai4code-model356-f1-798/356_bert_mpnet_l2_madgrad_1_798\"\nmodel2_f2_dir = \"../input/ai4code-model356-f2-798/356_bert_mpnet_l2_madgrad_2_798\"\n\nmodel1_l2_dirs = [\n    \"../input/ai4code-l2-770/l2_500_l6_b64_w64_0_770\",\n    \"../input/ai4code-l2-770/l2_501_l6_b64_w64_1_770\",\n    \"../input/ai4code-l2-770/l2_502_l6_b64_w64_2_770\",\n    \"../input/ai4code-l2-770/l2_503_l6_b64_w64_3_770\",\n]\n\nmodel1_l2_extra_dirs = [\n    \"../input/ai4code-l2-600/l2_600_l6_b64_w64_0_770\",\n    \"../input/ai4code-l2-600/l2_601_l6_b64_w64_1_770\",\n    \"../input/ai4code-l2-600/l2_602_l6_b64_w64_2_770\",\n    \"../input/ai4code-l2-600/l2_603_l6_b64_w64_3_770\",\n]\n\nmodel1_l2_dir_fallback = \"../input/ai4code-l2-770/l2_700_l2_light_0_770\"\n\nmodel2_l2_dirs = [\n    \"../input/ai4code-l2-770/l2_510_l6_b64_w64_0_770\",\n    \"../input/ai4code-l2-770/l2_511_l6_b64_w64_1_770\",\n    \"../input/ai4code-l2-770/l2_512_l6_b64_w64_2_770\"\n]\n\nmodel2_l2_extra_dirs = [\n    \"../input/ai4code-l2-600/l2_610_l6_b64_w64_0_770\",\n    \"../input/ai4code-l2-600/l2_611_l6_b64_w64_1_770\",\n    \"../input/ai4code-l2-600/l2_612_l6_b64_w64_2_770\"\n]","metadata":{"execution":{"iopub.status.busy":"2022-08-11T14:22:32.496346Z","iopub.execute_input":"2022-08-11T14:22:32.496827Z","iopub.status.idle":"2022-08-11T14:22:32.510663Z","shell.execute_reply.started":"2022-08-11T14:22:32.496788Z","shell.execute_reply":"2022-08-11T14:22:32.508239Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!ls ../input/ai4code-l2-600/l2_610_l6_b64_w64_0_770","metadata":{"execution":{"iopub.status.busy":"2022-08-11T14:22:32.516264Z","iopub.execute_input":"2022-08-11T14:22:32.517001Z","iopub.status.idle":"2022-08-11T14:22:34.009909Z","shell.execute_reply.started":"2022-08-11T14:22:32.516936Z","shell.execute_reply":"2022-08-11T14:22:34.008346Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import copy\nimport gc\nimport os\nimport random\nfrom dataclasses import dataclass\nimport time\nimport json\n\nimport numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nimport yaml\nfrom torch.utils.data import DataLoader\nfrom tqdm import tqdm\n\nimport code_dataset\nimport config\nimport models_bert2\nimport models_l2\nimport restore_order","metadata":{"execution":{"iopub.status.busy":"2022-08-11T14:22:34.012898Z","iopub.execute_input":"2022-08-11T14:22:34.013570Z","iopub.status.idle":"2022-08-11T14:22:36.110096Z","shell.execute_reply.started":"2022-08-11T14:22:34.013497Z","shell.execute_reply":"2022-08-11T14:22:36.108776Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def seed_everything(seed):\n    random.seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = False\n    torch.backends.cudnn.benchmark = False\n\n\nseed_everything(seed=42)\n\n\ndef build_model_l2(cfg):\n    model_params = copy.copy(cfg['model_params'])\n    # if model_params['model_type'] == 'models_segmentation':\n    cls = models_l2.__dict__[model_params['model_cls']]\n    del model_params['model_cls']\n    del model_params['model_type']\n    model: nn.Module = cls(**model_params)\n    return model\n\n\ndef no_collate(batch):\n    return batch[0]\n\n\n@dataclass\nclass L1ModelInfo:\n    name: str\n    config: dict\n    is_separate: bool\n    model_path_code: str\n    model_path_code_tokenizer: str\n    model_path_md: str\n    model_path_md_tokenizer: str\n\n    def __init__(self, name, cfg, model_dir):\n        self.name = name\n        self.config = cfg\n        model_cls = cfg['model_params']['model_cls']\n\n        if model_cls == 'DualBertWithL2':\n            self.is_separate = True\n        elif model_cls == 'SingleBertWithL2':\n            self.is_separate = False\n        else:\n            raise RuntimeError(f\"Invalid model cls {model_cls}\")\n\n        self.model_path_code = f'{model_dir}/l1_code'\n        self.model_path_code_tokenizer = f'{model_dir}/l1_code_tokenizer'\n        self.model_path_md = f'{model_dir}/l1_md' if self.is_separate else f'{model_dir}/l1_code'\n        self.model_path_md_tokenizer = f'{model_dir}/l1_md_tokenizer'\n\n\n@dataclass\nclass L2ModelInfo:\n    name: str\n    config: dict\n    model_path_code: str\n    weight: float = 1.0\n\n\ndef load_config(fn) -> dict:\n    return yaml.load(open(fn), Loader=yaml.FullLoader)\n\n\ndef masked_avg(t, mask):\n    return (t * mask[:, :, None]).sum(dim=1) / mask.sum(dim=1)[:, None]\n\n\ndef combine_predictions(model: nn.Module, tokens: [code_dataset.TokenIdsBatch]):\n    predictions = []\n\n    for token in tokens:\n        # print(token.input_ids.shape, token.input_ids.dtype, token.attention_mask.shape, token.attention_mask.dtype)\n        x = model(input_ids=token.input_ids, attention_mask=token.attention_mask)\n        x = masked_avg(x.last_hidden_state, token.attention_mask)\n        predictions.append(x)\n\n    predictions = torch.cat(predictions, dim=0)\n    return predictions[code_dataset.restore_order_idx(tokens), :]\n\n#\n# def combine_predictions_onnx(model, tokens: [code_dataset.TokenIdsBatch]):\n#     predictions = []\n#\n#     for token in tokens:\n#         # print(token.input_ids.shape, token.input_ids.dtype, token.attention_mask.shape, token.attention_mask.dtype)\n#         x = model.run(None, dict(input_ids=token.input_ids.numpy(), attention_mask=token.attention_mask.numpy()))[0]\n#         # print(x)\n#         # x = model(input_ids=token.input_ids, attention_mask=token.attention_mask)\n#         predictions.append(x)\n#\n#     predictions = np.concatenate(predictions, axis=0)\n#     return predictions[code_dataset.restore_order_idx(tokens), :]\n\n\ndef predict_l1(data_loader, model_l1, model_l1_name, tmp_dir, start_time, max_duration_s):\n    for data in tqdm(data_loader):\n        elapsed = time.time() - start_time\n        if elapsed > max_duration_s:\n            print(f'Stop prediction due to execution time {elapsed / 60.0:0.1f}m exceeded time limit {max_duration_s / 60:0.1f}m')\n            break\n\n        tokens_code = [d.cuda() for d in data['tokens_code']]\n        tokens_md = [d.cuda() for d in data['tokens_md']]\n\n        # with torch.cuda.amp.autocast():\n        try:\n            pred_code = combine_predictions(model_l1.get_code_model(), tokens_code)\n            pred_md = combine_predictions(model_l1.get_md_model(), tokens_md)\n        except:\n            with torch.cuda.amp.autocast():\n                pred_code = combine_predictions(model_l1.get_code_model(), tokens_code)\n                pred_md = combine_predictions(model_l1.get_md_model(), tokens_md)\n\n        pred_code = pred_code.float().detach().cpu().numpy()\n        pred_md = pred_md.float().detach().cpu().numpy()\n\n        keys_all_sorted = data['keys_all_sorted']\n        keys_code = data['keys_code']\n        keys_md = data['keys_md']\n        item_id = data['item_id']\n\n        np.savez(\n            f'{tmp_dir}/l1/{model_l1_name}/{item_id}.npz',\n            activations_code=pred_code,\n            activations_md=pred_md,\n            keys_all_sorted=keys_all_sorted,\n            keys_code=keys_code,\n            keys_md=keys_md\n        )\n\n\ndef predict_l2(model2, test_ids, tmp_dir, l1_model_info, l2_model_info):\n    predicted_ids = set()\n\n    for item_id in tqdm(test_ids):\n        pred_fn = f'{tmp_dir}/l1/{l1_model_info.name}/{item_id}.npz'\n\n        if not os.path.exists(pred_fn):\n            continue\n\n        l1_pred = np.load(pred_fn)\n\n        activations_code = torch.from_numpy(l1_pred['activations_code']).float().cuda()\n        activations_md = torch.from_numpy(l1_pred['activations_md']).float().cuda()\n\n        weight = l2_model_info.weight\n\n        with torch.cuda.amp.autocast():\n            try:\n                pred = model2(activations_code, activations_md)\n            except:\n                # if run out of VRAM, lets home at least one simple model is successful\n                continue\n\n        pred_md_after_code = torch.sigmoid(pred['md_after_code'].detach()).cpu().numpy()\n        pred_md_after_md = torch.sigmoid(pred['md_after_md'].detach()).cpu().numpy()\n        pred_md_between_code = torch.softmax(pred['md_between_code'].detach().cpu().float(), dim=0).numpy()\n\n        keys_code = l1_pred['keys_code']\n        keys_md = l1_pred['keys_md']\n\n        np.savez(\n            f'{tmp_dir}/l2/{l2_model_info.name}/{item_id}.npz',\n            keys_code=keys_code,\n            keys_md=keys_md,\n            md_after_code=pred_md_after_code,\n            md_after_md=pred_md_after_md,\n            md_between_code=pred_md_between_code,\n            weight=weight\n        )\n\n        predicted_ids.add(item_id)\n    return predicted_ids\n\n\ndef predict(l1_model_info: L1ModelInfo, l2_model_infos: [L2ModelInfo], l2_model_fallback, test_dir, tmp_dir, test_ids, start_time, max_duration_s):\n    torch.set_grad_enabled(False)\n\n    os.makedirs(f'{tmp_dir}/l1/{l1_model_info.name}', exist_ok=True)\n\n    elapsed = time.time() - start_time\n    if elapsed > max_duration_s:\n        print(f'Skip model prediction due to execution time {elapsed / 60.0:0.1f}m exceeded time limit {max_duration_s / 60:0.1f}m')\n        return\n\n    l1_cfg = l1_model_info.config\n\n    dataset_valid = code_dataset.CodeDatasetTest(\n        test_data_dir=test_dir,\n        test_ids=test_ids,\n        code_tokenizer_name=l1_model_info.model_path_code_tokenizer,\n        md_tokenizer_name=l1_model_info.model_path_md_tokenizer\n    )\n\n    if l1_model_info.is_separate:\n        model_l1 = models_bert2.DualBertWithL2(\n            code_model_name=l1_model_info.model_path_code,\n            md_model_name=l1_model_info.model_path_md,\n            l2_name='None',\n            l2_params=None,\n            pool_mode=l1_cfg['model_params']['pool_mode']\n        )\n    else:\n        model_l1 = models_bert2.SingleBertWithL2(\n            code_model_name=l1_model_info.model_path_code,\n            l2_name='None',\n            l2_params=None,\n            pool_mode=l1_cfg['model_params']['pool_mode']\n        )\n    model_l1 = model_l1.cuda()\n    model_l1.eval()\n\n    data_loader = DataLoader(\n        dataset_valid,\n        num_workers=1,\n        shuffle=False,\n        batch_size=1,\n        collate_fn=no_collate\n    )\n\n    predict_l1(\n        data_loader=data_loader,\n        model_l1=model_l1,\n        model_l1_name=l1_model_info.name,\n        tmp_dir=tmp_dir,\n        start_time=start_time,\n        max_duration_s=max_duration_s)\n\n    del model_l1\n    del data_loader\n    del dataset_valid\n    gc.collect()\n\n    predicted_ids = set()\n\n    for l2_model_info in l2_model_infos:\n        os.makedirs(f'{tmp_dir}/l2/{l2_model_info.name}', exist_ok=True)\n\n        model2 = build_model_l2(l2_model_info.config)\n        model2 = model2.cuda()\n        model2.eval()\n\n        checkpoint = torch.load(f\"{l2_model_info.model_path_code}/l2.pt\")\n        model2.load_state_dict(checkpoint[\"model_state_dict\"])\n        del checkpoint\n\n        cur_predicted_ids = predict_l2(model2, test_ids, tmp_dir, l1_model_info, l2_model_info)\n        del model2\n        gc.collect()\n\n        predicted_ids.update(cur_predicted_ids)\n\n    if l2_model_fallback is not None and len(predicted_ids) != len(test_ids):\n        l2_model_info = l2_model_fallback\n        failed_ids = list(set(test_ids).difference(predicted_ids))\n        print(f'Using fallback model to predict {len(failed_ids)} items')\n\n        os.makedirs(f'{tmp_dir}/l2/{l2_model_info.name}', exist_ok=True)\n\n        model2 = build_model_l2(l2_model_info.config)\n        model2 = model2.cuda()\n        model2.eval()\n\n        checkpoint = torch.load(f\"{l2_model_info.model_path_code}/l2.pt\")\n        model2.load_state_dict(checkpoint[\"model_state_dict\"])\n        del checkpoint\n\n        predict_l2(model2, failed_ids, tmp_dir, l1_model_info, l2_model_info)\n        del model2\n        gc.collect()\n\ndef restore_order_sm(keys_code, keys_md, md_after_code, md_after_md, md_between_code):\n    nb_code = len(keys_code)\n    nb_md = len(keys_md)\n\n    md_after_code_soft = md_after_code.astype(np.float64)\n    md_after_code_cost_pos = -1 * md_after_code_soft\n    md_after_code_cost_neg = -1 * (1 - md_after_code_soft)\n\n    nb_bins = nb_code + 1\n    md_bins = [[] for _ in range(nb_bins)]  # places to put md from before the first code to after the last code\n    # md_after_md = np.mean(md_after_md, axis=1)\n\n    for md_idx in range(nb_md):\n        pos_costs = [\n            md_after_code_cost_pos[md_idx, :i+1].sum() + md_after_code_cost_neg[md_idx, i+1:].sum() # - md_between_code[md_idx, i]\n            for i in range(nb_bins)\n        ]\n        md_bins[np.argmin(pos_costs)].append(md_idx)\n\n    for bin_idx in range(nb_bins):\n        bin = md_bins[bin_idx]\n        if len(bin) > 1:\n            items_with_order = [(md_after_md[b, bin].mean(), b) for b in bin]\n            # items_with_order = [(md_after_md[b], b) for b in bin]\n            items_with_order = list(sorted(items_with_order))\n            md_bins[bin_idx] = [b[1] for b in items_with_order]\n\n    res_order = []\n    for md_bin_idx, bin in enumerate(md_bins):\n        for md_idx in bin:\n            res_order.append(keys_md[md_idx])\n\n        if md_bin_idx < nb_code:\n            res_order.append(keys_code[md_bin_idx])\n\n    return res_order\n\n\n\ndef combine_l2_predictions(l2_model_names, tmp_dir, test_ids):\n    l2_predictions = {}\n\n    # for item_id, l1_pred in tqdm(l1_predictions.items()):\n    for item_id in tqdm(test_ids):\n        predictions = []\n        total_weight = 0.0\n        for l2_model_name in l2_model_names:\n            fn = f'{tmp_dir}/l2/{l2_model_name}/{item_id}.npz'\n            if os.path.exists(fn):\n                data = np.load(fn)\n                predictions.append(data)\n                total_weight += float(data['weight'])\n\n        if total_weight == 0:\n            continue  # TODO: add fallback solution\n\n        keys_code = predictions[0]['keys_code']\n        keys_md = predictions[0]['keys_md']\n        md_after_code = predictions[0]['md_after_code'] * predictions[0]['weight'] / total_weight\n        md_after_md = predictions[0]['md_after_md'] * predictions[0]['weight'] / total_weight\n        md_between_code = predictions[0]['md_between_code'] * predictions[0]['weight'] / total_weight\n\n        for p in predictions[1:]:\n            md_after_code += p['md_after_code'] * p['weight'] / total_weight\n            md_after_md += p['md_after_md'] * p['weight'] / total_weight\n            md_between_code += p['md_between_code'] * p['weight'] / total_weight\n\n        keys_all_sorted_pred = restore_order_sm(\n            keys_code=keys_code,\n            keys_md=keys_md,\n            md_after_code=md_after_code,\n            md_after_md=md_after_md,\n            md_between_code=md_between_code\n        )\n\n        l2_predictions[item_id] = keys_all_sorted_pred\n\n    return l2_predictions\n","metadata":{"execution":{"iopub.status.busy":"2022-08-11T14:22:56.641847Z","iopub.execute_input":"2022-08-11T14:22:56.642428Z","iopub.status.idle":"2022-08-11T14:22:56.701195Z","shell.execute_reply.started":"2022-08-11T14:22:56.642395Z","shell.execute_reply":"2022-08-11T14:22:56.699880Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"m1_l1_cfg_str = \"\"\"\n###################\n## Model options\nmodel_params:\n  model_type: \"models_bert2\"\n  model_cls: \"SingleBertWithL2\"\n  code_model_name: 'microsoft/codebert-base'\n  enable_gradient_checkpointing: true\n  pool_mode: 'avg'\n  l2_name: 'L2Transformer'\n  l2_params:\n    nb_combine_cells_around: 1\n    nhead_code: 16\n    nhead_md: 16\n    num_decoder_layers: 2\n    dec_dim: 1024\n    dim_feedforward: 2048\n    encoder_code_dim: 768\n    encoder_md_dim: 768\n    combined_pos_enc: true\n    add_extra_outputs: false\n    rescale_att: true\n\n\ndataset_params:\n  code_tokenizer_name: 'microsoft/codebert-base'\n  md_tokenizer_name: 'microsoft/codebert-base'\n  max_code_tokens_number: 256\n  max_md_tokens_number: 256\n  nb_code_cells: 2048\n  nb_md_cells: 2048\n  batch_cost: 32768\n  max_size2: 524288\n  cell_prefix: ''\n  preprocess_md: ''\n  use_pos_between_code: True\n  low_case_md: False\n  low_case_code: False\n\"\"\"\n\nm1_l1_cfg=yaml.load(m1_l1_cfg_str, Loader=yaml.FullLoader)","metadata":{"execution":{"iopub.status.busy":"2022-08-11T14:22:59.253103Z","iopub.execute_input":"2022-08-11T14:22:59.253733Z","iopub.status.idle":"2022-08-11T14:22:59.267004Z","shell.execute_reply.started":"2022-08-11T14:22:59.253700Z","shell.execute_reply":"2022-08-11T14:22:59.265557Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"m2_l1_cfg_str = \"\"\"\n###################\n## Model options\nmodel_params:\n  model_type: \"models_bert2\"\n  model_cls: \"DualBertWithL2\"\n  code_model_name: 'microsoft/codebert-base'\n  md_model_name: 'sentence-transformers/paraphrase-multilingual-mpnet-base-v2'\n  enable_gradient_checkpointing: true\n  pool_mode: 'avg'\n  l2_name: 'L2Transformer'\n  l2_params:\n    nb_combine_cells_around: 1\n    nhead_code: 16\n    nhead_md: 16\n    num_decoder_layers: 2\n    dec_dim: 1024\n    dim_feedforward: 2048\n    encoder_code_dim: 768\n    encoder_md_dim: 768\n    combined_pos_enc: true\n    add_extra_outputs: false\n    rescale_att: true\n\n\ndataset_params:\n  code_tokenizer_name: 'microsoft/codebert-base'\n  md_tokenizer_name: 'sentence-transformers/paraphrase-multilingual-mpnet-base-v2'\n  max_code_tokens_number: 256\n  max_md_tokens_number: 256\n  nb_code_cells: 2048\n  nb_md_cells: 2048\n  batch_cost: 32768\n  max_size2: 524288\n  cell_prefix: ''\n  preprocess_md: ''\n  use_pos_between_code: True\n  low_case_md: False\n  low_case_code: False\n\"\"\"\n\nm2_l1_cfg=yaml.load(m2_l1_cfg_str, Loader=yaml.FullLoader)","metadata":{"execution":{"iopub.status.busy":"2022-08-11T14:22:59.906066Z","iopub.execute_input":"2022-08-11T14:22:59.907420Z","iopub.status.idle":"2022-08-11T14:22:59.921573Z","shell.execute_reply.started":"2022-08-11T14:22:59.907376Z","shell.execute_reply":"2022-08-11T14:22:59.919487Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"m1_l2_cfg = dict(\n        model_params=dict(\n            model_type=\"models_l2\",\n            model_cls=m1_l1_cfg['model_params']['l2_name'],\n            **m1_l1_cfg['model_params']['l2_params']\n        )\n    )\n\nm2_l2_cfg = dict(\n        model_params=dict(\n            model_type=\"models_l2\",\n            model_cls=m2_l1_cfg['model_params']['l2_name'],\n            **m2_l1_cfg['model_params']['l2_params']\n        )\n    )","metadata":{"execution":{"iopub.status.busy":"2022-08-11T14:23:00.571302Z","iopub.execute_input":"2022-08-11T14:23:00.572835Z","iopub.status.idle":"2022-08-11T14:23:00.582372Z","shell.execute_reply.started":"2022-08-11T14:23:00.572771Z","shell.execute_reply":"2022-08-11T14:23:00.580328Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"l2_cfg_str = \"\"\"\nmodel_params:\n  model_type: \"models_l2\"\n  model_cls: \"L2Transformer\"\n  nb_combine_cells_around: 1\n  nhead_code: 16\n  nhead_md: 16\n  num_decoder_layers: 6\n  dec_dim: 1024\n  dim_feedforward: 2048\n  encoder_code_dim: 768\n  encoder_md_dim: 768\n  combined_pos_enc: true\n\"\"\"\n\nextra_l2_cfg=yaml.load(l2_cfg_str, Loader=yaml.FullLoader)","metadata":{"execution":{"iopub.status.busy":"2022-08-11T14:23:01.190146Z","iopub.execute_input":"2022-08-11T14:23:01.190924Z","iopub.status.idle":"2022-08-11T14:23:01.199236Z","shell.execute_reply.started":"2022-08-11T14:23:01.190891Z","shell.execute_reply":"2022-08-11T14:23:01.197645Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"l2_fallback_cfg_str = \"\"\"\nmodel_params:\n  model_type: \"models_l2\"\n  model_cls: \"L2Transformer\"\n  nb_combine_cells_around: 1\n  nhead_code: 4\n  nhead_md: 4\n  num_decoder_layers: 2\n  dec_dim: 256\n  ca_mul: 1\n  dim_feedforward: 512\n  encoder_code_dim: 768\n  encoder_md_dim: 768\n  combined_pos_enc: true\n\"\"\"\n\nextra_l2_cfg_fallback=yaml.load(l2_cfg_str, Loader=yaml.FullLoader)","metadata":{"execution":{"iopub.status.busy":"2022-08-11T14:23:01.652364Z","iopub.execute_input":"2022-08-11T14:23:01.655324Z","iopub.status.idle":"2022-08-11T14:23:01.665055Z","shell.execute_reply.started":"2022-08-11T14:23:01.655277Z","shell.execute_reply":"2022-08-11T14:23:01.663451Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def split_small_large_notebooks(test_dir: str, test_ids, large_ratio=0.1):\n    test_ids_with_size = []\n    for item_id in test_ids:\n        data = json.load(open(f'{test_dir}/{item_id}.json'))\n        nb_code = 0\n        nb_md = 0\n        for cell_type in data['cell_type'].values():\n            if cell_type == 'code':\n                nb_code += 1\n            else:\n                nb_md += 1\n\n        test_ids_with_size.append(((nb_code+1)*(nb_md+1), item_id))\n\n    test_ids_with_size = list(sorted(test_ids_with_size, reverse=True))\n    test_ids_sorted = [item[1] for item in test_ids_with_size]\n\n    nb_large = int(len(test_ids) * large_ratio + 0.999)\n\n    return test_ids_sorted[nb_large:], test_ids_sorted[:nb_large]\n\n","metadata":{"execution":{"iopub.status.busy":"2022-08-11T14:23:04.996806Z","iopub.execute_input":"2022-08-11T14:23:04.997203Z","iopub.status.idle":"2022-08-11T14:23:05.007877Z","shell.execute_reply.started":"2022-08-11T14:23:04.997172Z","shell.execute_reply":"2022-08-11T14:23:05.006472Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_data_path = '../input/AI4Code/test'\nall_samples = [s[:-5] for s in sorted(os.listdir(test_data_path)) if s.endswith('.json')]\n\nall_samples[:4], len(all_samples)","metadata":{"execution":{"iopub.status.busy":"2022-08-11T14:23:05.670903Z","iopub.execute_input":"2022-08-11T14:23:05.671360Z","iopub.status.idle":"2022-08-11T14:23:05.685963Z","shell.execute_reply.started":"2022-08-11T14:23:05.671317Z","shell.execute_reply":"2022-08-11T14:23:05.684070Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"batch_size = 4096*5\n\ntmp_dir='/kaggle/temp'\n\n\ndef predict_model1(fold, all_samples):\n    model_dir = [\n        model1_f0_dir,\n        model1_f1_dir,\n        model1_f2_dir,\n        model1_f3_dir,\n    ][fold]\n\n    model_extra_dir = model1_l2_dirs[fold]\n    model_extra_v2_dir = model1_l2_extra_dirs[fold]\n\n    if fold == 0:\n        l2_model_fallback = L2ModelInfo(\n                    name=f'l2_model1_f0',  # re-use the model name\n                    config=extra_l2_cfg_fallback,\n                    model_path_code=model1_l2_dir_fallback,\n                    weight=0.01\n                )\n    else:\n        l2_model_fallback = None\n        \n    nb_steps = (len(all_samples) + batch_size - 1) // batch_size\n\n    for step in tqdm(range(nb_steps)):\n        predict(\n            l1_model_info=L1ModelInfo(\n                name=f'model1_f{fold}',\n                cfg=m1_l1_cfg,\n                model_dir=model_dir\n            ),\n            l2_model_infos=[\n                L2ModelInfo(\n                    name=f'l2_model1_f{fold}',\n                    config=m1_l2_cfg,\n                    model_path_code=f'{model_dir}/l2',\n                    weight=1.0\n                ),\n                L2ModelInfo(\n                    name=f'l2_model1_extra_f{fold}',\n                    config=extra_l2_cfg,\n                    model_path_code=f'{model_extra_dir}',\n                    weight=1.6\n                ),\n#                 L2ModelInfo(\n#                     name=f'l2_model1_extra2_f{fold}',\n#                     config=extra_l2_cfg,\n#                     model_path_code=model_extra_v2_dir,\n#                     weight=0.8\n#                 ),\n            ],\n            l2_model_fallback=l2_model_fallback,\n            test_dir=test_data_path,\n            test_ids=all_samples[step * batch_size:(step + 1) * batch_size],\n            tmp_dir=tmp_dir,\n            start_time=start_time,\n            max_duration_s=max_duration_s\n        )\n        gc.collect()\n\ndef predict_model2(fold, all_samples):\n    model_dir = [\n        model2_f0_dir,\n        model2_f1_dir,\n        model2_f2_dir,\n    ][fold]\n\n    model_extra_dir = model2_l2_dirs[fold]\n    model_extra_v2_dir = model1_l2_extra_dirs[fold]\n    \n    nb_steps = (len(all_samples) + batch_size - 1) // batch_size\n\n    for step in tqdm(range(nb_steps)):\n        predict(\n            l1_model_info=L1ModelInfo(\n                name=f'model2_f{fold}',\n                cfg=m2_l1_cfg,\n                model_dir=model_dir\n            ),\n            l2_model_infos=[\n                L2ModelInfo(\n                    name=f'l2_model2_f{fold}',\n                    config=m2_l2_cfg,\n                    model_path_code=f'{model_dir}/l2',\n                    weight=1.0\n                ),\n                L2ModelInfo(\n                    name=f'l2_model2_extra_f{fold}',\n                    config=extra_l2_cfg,\n                    model_path_code=f'{model_extra_dir}',\n                    weight=1.6\n                ),\n#                 L2ModelInfo(\n#                     name=f'l2_model2_extra2_f{fold}',\n#                     config=extra_l2_cfg,\n#                     model_path_code=model_extra_v2_dir,\n#                     weight=0.8\n#                 ),\n            ],\n            l2_model_fallback=None,\n            test_dir=test_data_path,\n            test_ids=all_samples[step * batch_size:(step + 1) * batch_size],\n            tmp_dir=tmp_dir,\n            start_time=start_time,\n            max_duration_s=max_duration_s\n        )\n        gc.collect()\n\nall_samples_small, all_samples_large = split_small_large_notebooks(test_dir=test_data_path, test_ids=all_samples, large_ratio=0.05)\n\npredict_model1(fold=0, all_samples=all_samples_large)\npredict_model2(fold=0, all_samples=all_samples_large)\npredict_model1(fold=1, all_samples=all_samples_large)\npredict_model2(fold=1, all_samples=all_samples_large)\npredict_model1(fold=2, all_samples=all_samples_large)\npredict_model2(fold=2, all_samples=all_samples_large)\npredict_model1(fold=3, all_samples=all_samples_large)\n\npredict_model1(fold=0, all_samples=all_samples_small)\npredict_model2(fold=0, all_samples=all_samples_small)\npredict_model1(fold=1, all_samples=all_samples_small)\npredict_model2(fold=1, all_samples=all_samples_small)\npredict_model1(fold=2, all_samples=all_samples_small)\npredict_model2(fold=2, all_samples=all_samples_small)\npredict_model1(fold=3, all_samples=all_samples_small)\n\ngc.collect()\n\nprediction = {}\nstep_prediction = combine_l2_predictions(\n    [\n        'l2_model1_f0', 'l2_model1_f1', 'l2_model1_f2', 'l2_model1_f3',\n        'l2_model2_f0', 'l2_model2_f1', 'l2_model2_f2',\n        \n        'l2_model1_extra_f0', 'l2_model1_extra_f1', 'l2_model1_extra_f2', 'l2_model1_extra_f3',\n        'l2_model2_extra_f0', 'l2_model2_extra_f1', 'l2_model2_extra_f2',\n        \n#         'l2_model1_extra2_f0', 'l2_model1_extra2_f1', 'l2_model1_extra2_f2', 'l2_model1_extra2_f3',\n#         'l2_model2_extra2_f0', 'l2_model2_extra2_f1', 'l2_model2_extra2_f2',\n    ],\n    test_ids=all_samples,\n    tmp_dir=tmp_dir)\n\nprediction.update(step_prediction)\n\n\nitem_ids = []\ncell_order = []\n\nfor item_id, pred in prediction.items():\n    item_ids.append(item_id)\n    cell_order.append(' '.join([p for p in pred if p != '']))\n\nres_df = pd.DataFrame(data={\n    'id': item_ids,\n    'cell_order': cell_order\n})\nres_df.to_csv('submission.csv', index=False)\n","metadata":{"execution":{"iopub.status.busy":"2022-08-11T14:25:11.621305Z","iopub.execute_input":"2022-08-11T14:25:11.621820Z","iopub.status.idle":"2022-08-11T14:32:08.241238Z","shell.execute_reply.started":"2022-08-11T14:25:11.621787Z","shell.execute_reply":"2022-08-11T14:32:08.239515Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"res_df.head()","metadata":{"execution":{"iopub.status.busy":"2022-08-11T14:32:08.244836Z","iopub.execute_input":"2022-08-11T14:32:08.245342Z","iopub.status.idle":"2022-08-11T14:32:08.273328Z","shell.execute_reply.started":"2022-08-11T14:32:08.245262Z","shell.execute_reply":"2022-08-11T14:32:08.271576Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}