{"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":"markdown","source":"Training code: https://github.com/suicao/ai4code-baseline","metadata":{}},{"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\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-07-14T05:27:29.575145Z","iopub.execute_input":"2022-07-14T05:27:29.575418Z","iopub.status.idle":"2022-07-14T05:27:29.581239Z","shell.execute_reply.started":"2022-07-14T05:27:29.575386Z","shell.execute_reply":"2022-07-14T05:27:29.580241Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Config\nclass CFG_0:\n    dynamic_code_len = False\n    model_path = \"../input/codebert-base/codebert-base/\"\n    ckpt_path = \"../input/codebert-fts-full/model.bin\"\n    num_sample = 20\n    total_max_len = 512\n    md_max_len = 64\n    code_max_len = 23\n    limit = 200\n\nclass CFG_1:\n    dynamic_code_len = True\n    model_path = \"../input/distilbertbaseuncased/\"\n    ckpt_path = \"../input/distilbert-full/model.bin\"\n    num_sample = 45\n    total_max_len = 512\n    md_max_len = 64\n    code_max_len = 10\n    limit = 500\n\n\nclass CFG_2:\n    dynamic_code_len = True\n    model_path = \"../input/distilroberta-base/\"\n    ckpt_path = \"../input/distil-roberta-full/model.bin\"\n    num_sample = 45\n    total_max_len = 512\n    md_max_len = 64\n    code_max_len = 10\n    limit = 500\n    \nclass CFG_3:\n    dynamic_code_len = False\n    model_path = \"../input/deberta-v3-base/deberta-v3-base/\"\n    ckpt_path = \"../input/dv3b-fts-full/model.bin\"\n    num_sample = 20\n    total_max_len = 512\n    md_max_len = 64\n    code_max_len = 23\n    limit = 200\n    \nclass CFG_4:\n    dynamic_code_len = True\n    model_path = \"../input/deberta-v3-base/deberta-v3-base/\"\n    ckpt_path = \"../input/dv3b-fts-10k-dl-samp45-len10-lim500-ema/model.bin\"\n    num_sample = 45\n    total_max_len = 512\n    md_max_len = 64\n    code_max_len = 10\n    limit = 500\n\nclass CFG_5:\n    dynamic_code_len = True\n    model_path = \"../input/deberta-v3-base-10k-m40/\"\n    ckpt_path = \"../input/dv3b-10k-exclude-val-mlm3eps/model.bin\"\n    num_sample = 45\n    total_max_len = 512\n    md_max_len = 64\n    code_max_len = 10\n    limit = 500\n    \nclass CFG_6:\n    dynamic_code_len = True\n    model_path = \"../input/codebert-base/codebert-base/\"\n    ckpt_path = \"../input/codebert-dl-full/model.bin\"\n    num_sample = 45\n    total_max_len = 512\n    md_max_len = 64\n    code_max_len = 10\n    limit = 500\n\nclass CFG_7:\n    dynamic_code_len = True\n    model_path = \"../input/debertabase/deberta-base/\"\n    ckpt_path = \"../input/dv1b-fts-10k-dl-samp45-len10-lim500/model.bin\"\n    num_sample = 45\n    total_max_len = 512\n    md_max_len = 64\n    code_max_len = 10\n    limit = 500\n    \nclass CFG_8:\n    dynamic_code_len = True\n    model_path = \"../input/deberta-v3-base/deberta-v3-base/\"\n    ckpt_path = \"../input/dv3b-full-dl/model.bin\"\n    num_sample = 45\n    total_max_len = 512\n    md_max_len = 64\n    code_max_len = 10\n    limit = 500","metadata":{"execution":{"iopub.status.busy":"2022-07-14T05:27:29.587775Z","iopub.execute_input":"2022-07-14T05:27:29.588044Z","iopub.status.idle":"2022-07-14T05:27:29.600439Z","shell.execute_reply.started":"2022-07-14T05:27:29.588016Z","shell.execute_reply":"2022-07-14T05:27:29.599332Z"},"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\npaths_test = list((data_dir / 'test').glob('*.json'))\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-07-14T05:27:29.604952Z","iopub.execute_input":"2022-07-14T05:27:29.605196Z","iopub.status.idle":"2022-07-14T05:27:29.670060Z","shell.execute_reply.started":"2022-07-14T05:27:29.605167Z","shell.execute_reply":"2022-07-14T05:27:29.669210Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_df.tail(50)","metadata":{"papermill":{"duration":0.03602,"end_time":"2022-05-23T03:29:09.891968","exception":false,"start_time":"2022-05-23T03:29:09.855948","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-07-14T05:27:29.672131Z","iopub.execute_input":"2022-07-14T05:27:29.673038Z","iopub.status.idle":"2022-07-14T05:27:29.696513Z","shell.execute_reply.started":"2022-07-14T05:27:29.672994Z","shell.execute_reply":"2022-07-14T05:27:29.695757Z"},"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, limit):\n    print(n)\n    cells = [clean_code(cell) for cell in cells]\n    if n >= len(cells):\n        return [cell[:limit] for cell in cells]\n    else:\n        results = []\n        step = len(cells) / n\n        idx = 0\n        while int(np.round(idx)) < len(cells):\n            results.append(cells[int(np.round(idx))])\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, num_sample, limit):\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, num_sample, limit)\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-07-14T05:27:29.697988Z","iopub.execute_input":"2022-07-14T05:27:29.698439Z","iopub.status.idle":"2022-07-14T05:27:29.712206Z","shell.execute_reply.started":"2022-07-14T05:27:29.698397Z","shell.execute_reply":"2022-07-14T05:27:29.711284Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"papermill":{"duration":0.030138,"end_time":"2022-05-23T03:29:09.972433","exception":false,"start_time":"2022-05-23T03:29:09.942295","status":"completed"},"tags":[],"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        self.model = AutoModel.from_pretrained(model_path)\n        self.top = nn.Linear(769, 1)\n        \n    def forward(self, ids, mask, fts):\n        x = self.model(ids, mask)[0]\n        x = self.top(torch.cat((x[:, 0, :], fts),1))\n        return x\n\nfrom torch.utils.data import DataLoader, Dataset\nclass MarkdownDataset(Dataset):\n\n    def __init__(self, df, cfg, fts):\n        super().__init__()\n        self.df = df.reset_index(drop=True)\n        self.md_max_len = cfg.md_max_len\n        self.code_max_len = cfg.code_max_len\n        self.total_max_len = cfg.total_max_len  # maxlen allowed by model config\n        self.dynamic_code_len = cfg.dynamic_code_len\n        self.num_sample = cfg.num_sample\n        self.tokenizer = AutoTokenizer.from_pretrained(cfg.model_path)\n        self.fts = fts\n\n    def __getitem__(self, index):\n        row = self.df.iloc[index]\n\n        inputs = self.tokenizer.encode_plus(\n            row.source,\n            None,\n            add_special_tokens=True,\n            max_length=self.md_max_len,\n            padding=\"max_length\",\n            return_token_type_ids=True,\n            truncation=True\n        )\n        \n        num_code_cells = len(self.fts[row.id][\"codes\"])\n        if self.dynamic_code_len:\n            if num_code_cells < self.num_sample:\n                code_inputs = self.tokenizer.batch_encode_plus(\n                    [str(x) for x in self.fts[row.id][\"codes\"]],\n                    add_special_tokens=True,\n                    max_length=int(np.round(450/num_code_cells)),\n                    padding=\"max_length\",\n                    truncation=True\n                )\n            \n            else:\n                code_inputs = self.tokenizer.batch_encode_plus(\n                    [str(x) for x in self.fts[row.id][\"codes\"]],\n                    add_special_tokens=True,\n                    max_length=self.code_max_len,\n                    padding=\"max_length\",\n                    truncation=True\n                )\n        else:\n            code_inputs = self.tokenizer.batch_encode_plus(\n                [str(x) for x in self.fts[row.id][\"codes\"]],\n                add_special_tokens=True,\n                max_length=self.code_max_len,\n                padding=\"max_length\",\n                truncation=True\n            )\n        n_md = self.fts[row.id][\"total_md\"]\n        n_code = self.fts[row.id][\"total_code\"]\n        if n_md + n_code == 0:\n            fts = torch.FloatTensor([0])\n        else:\n            fts = torch.FloatTensor([n_md / (n_md + n_code)])\n\n        ids = inputs['input_ids']\n        for x in code_inputs['input_ids']:\n            ids.extend(x[:-1])\n        ids = ids[:self.total_max_len]\n        if len(ids) != self.total_max_len:\n            ids = ids + [self.tokenizer.pad_token_id, ] * (self.total_max_len - len(ids))\n        ids = torch.LongTensor(ids)\n\n        mask = inputs['attention_mask']\n        for x in code_inputs['attention_mask']:\n            mask.extend(x[:-1])\n        mask = mask[:self.total_max_len]\n        if len(mask) != self.total_max_len:\n            mask = mask + [self.tokenizer.pad_token_id, ] * (self.total_max_len - len(mask))\n        mask = torch.LongTensor(mask)\n\n        assert len(ids) == self.total_max_len\n\n        return ids, mask, fts, torch.FloatTensor([row.pct_rank])\n\n    def __len__(self):\n        return self.df.shape[0]","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-07-14T05:27:29.715079Z","iopub.execute_input":"2022-07-14T05:27:29.715653Z","iopub.status.idle":"2022-07-14T05:27:29.742732Z","shell.execute_reply.started":"2022-07-14T05:27:29.715607Z","shell.execute_reply":"2022-07-14T05:27:29.741769Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def read_data(data):\n    return tuple(d.cuda() for d in data[:-1]), data[-1].cuda()\n\n\ndef validate(model, val_loader):\n    model.eval()\n    \n    tbar = tqdm(val_loader, file=sys.stdout)\n    \n    preds = []\n    labels = []\n\n    with torch.no_grad():\n        for idx, data in enumerate(tbar):\n            inputs, target = read_data(data)\n\n            pred = model(*inputs)\n\n            preds.append(pred.detach().cpu().numpy().ravel())\n            labels.append(target.detach().cpu().numpy().ravel())\n    \n    return np.concatenate(labels), np.concatenate(preds)\n\ndef predict(cfg, test_fts):\n    model = MarkdownModel(cfg.model_path)\n    model = model.cuda()\n    model.eval()\n    model.load_state_dict(torch.load(cfg.ckpt_path))\n\n    test_df[\"pct_rank\"] = 0\n    test_ds = MarkdownDataset(test_df[test_df[\"cell_type\"] == \"markdown\"].reset_index(drop=True), cfg, fts=test_fts)\n    test_loader = DataLoader(test_ds, batch_size=32, shuffle=False, num_workers=2,\n                              pin_memory=False, drop_last=False)\n    _, y_test = validate(model, test_loader)\n    return y_test","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-07-14T05:27:29.744371Z","iopub.execute_input":"2022-07-14T05:27:29.744680Z","iopub.status.idle":"2022-07-14T05:27:29.760201Z","shell.execute_reply.started":"2022-07-14T05:27:29.744638Z","shell.execute_reply":"2022-07-14T05:27:29.759307Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_fts = get_features(test_df, CFG_1.num_sample, CFG_1.limit)    \ny_test_0 = predict(CFG_1, test_fts)\n\ntest_fts = get_features(test_df, CFG_6.num_sample, CFG_6.limit)    \ny_test_1 = predict(CFG_6, test_fts)\n\ntest_fts = get_features(test_df, CFG_3.num_sample, CFG_3.limit)    \ny_test_2 = predict(CFG_3, test_fts)\n\ntest_fts = get_features(test_df, CFG_8.num_sample, CFG_8.limit)\ny_test_3 = predict(CFG_8, test_fts)\n\n\n# test_fts = get_features(test_df, CFG_6.num_sample, CFG_6.limit)    \n# y_test = predict(CFG_6, test_fts)\n\ny_test = 0.07 * y_test_0 + 0.12 * y_test_1 + 0.31 * y_test_2 + 0.5 * y_test_3\n# 27 11 62 86479","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-07-14T05:27:29.761600Z","iopub.execute_input":"2022-07-14T05:27:29.762245Z","iopub.status.idle":"2022-07-14T05:27:49.845075Z","shell.execute_reply.started":"2022-07-14T05:27:29.762200Z","shell.execute_reply":"2022-07-14T05:27:49.843944Z"},"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-07-14T05:27:49.847237Z","iopub.execute_input":"2022-07-14T05:27:49.848271Z","iopub.status.idle":"2022-07-14T05:27:49.855716Z","shell.execute_reply.started":"2022-07-14T05:27:49.848220Z","shell.execute_reply":"2022-07-14T05:27:49.854881Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# test_df","metadata":{"execution":{"iopub.status.busy":"2022-07-14T05:27:49.857359Z","iopub.execute_input":"2022-07-14T05:27:49.857743Z","iopub.status.idle":"2022-07-14T05:27:49.868422Z","shell.execute_reply.started":"2022-07-14T05:27:49.857694Z","shell.execute_reply":"2022-07-14T05:27:49.867542Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# test_df.sort_values(\"pred\").groupby(\"id\")[\"cell_id\"]","metadata":{"execution":{"iopub.status.busy":"2022-07-14T05:27:49.871292Z","iopub.execute_input":"2022-07-14T05:27:49.871658Z","iopub.status.idle":"2022-07-14T05:27:49.881747Z","shell.execute_reply.started":"2022-07-14T05:27:49.871616Z","shell.execute_reply":"2022-07-14T05:27:49.880900Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\nsub_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-07-14T05:27:49.883112Z","iopub.execute_input":"2022-07-14T05:27:49.883534Z","iopub.status.idle":"2022-07-14T05:27:49.905428Z","shell.execute_reply.started":"2022-07-14T05:27:49.883492Z","shell.execute_reply":"2022-07-14T05:27:49.904519Z"},"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-07-14T05:27:49.908169Z","iopub.execute_input":"2022-07-14T05:27:49.908788Z","iopub.status.idle":"2022-07-14T05:27:49.915681Z","shell.execute_reply.started":"2022-07-14T05:27:49.908749Z","shell.execute_reply":"2022-07-14T05:27:49.914787Z"},"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":[]}]}