{"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 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-12T09:31:24.403772Z","iopub.execute_input":"2022-07-12T09:31:24.404328Z","iopub.status.idle":"2022-07-12T09:31:24.525740Z","shell.execute_reply.started":"2022-07-12T09:31:24.404223Z","shell.execute_reply":"2022-07-12T09:31:24.524983Z"},"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-12T09:31:24.526916Z","iopub.execute_input":"2022-07-12T09:31:24.527365Z","iopub.status.idle":"2022-07-12T09:31:24.624770Z","shell.execute_reply.started":"2022-07-12T09:31:24.527328Z","shell.execute_reply":"2022-07-12T09:31:24.624049Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_df","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-12T09:31:24.626948Z","iopub.execute_input":"2022-07-12T09:31:24.627433Z","iopub.status.idle":"2022-07-12T09:31:24.649269Z","shell.execute_reply.started":"2022-07-12T09:31:24.627390Z","shell.execute_reply":"2022-07-12T09:31:24.648516Z"},"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):\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        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):\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, 20)\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-12T09:31:24.650914Z","iopub.execute_input":"2022-07-12T09:31:24.651555Z","iopub.status.idle":"2022-07-12T09:31:24.667545Z","shell.execute_reply.started":"2022-07-12T09:31:24.651509Z","shell.execute_reply":"2022-07-12T09:31:24.666842Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_fts = get_features(test_df)","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":[],"execution":{"iopub.status.busy":"2022-07-12T09:31:24.668929Z","iopub.execute_input":"2022-07-12T09:31:24.669400Z","iopub.status.idle":"2022-07-12T09:31:24.692097Z","shell.execute_reply.started":"2022-07-12T09:31:24.669348Z","shell.execute_reply":"2022-07-12T09:31:24.691213Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from tqdm import tqdm\nimport sys, os\nfrom transformers import AutoModel, AutoTokenizer\nfrom torch.utils.data import DataLoader, Dataset\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\nclass MarkdownDataset(Dataset):\n\n    def __init__(self, df, model_name_or_path, total_max_len, md_max_len, fts):\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        self.tokenizer = AutoTokenizer.from_pretrained(model_name_or_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        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=23,\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_md\"]\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-12T09:31:24.694003Z","iopub.execute_input":"2022-07-12T09:31:24.694515Z","iopub.status.idle":"2022-07-12T09:31:32.954232Z","shell.execute_reply.started":"2022-07-12T09:31:24.694461Z","shell.execute_reply":"2022-07-12T09:31:32.953437Z"},"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(model_path, ckpt_path):\n    model = MarkdownModel(model_path)\n    model = model.cuda()\n    model.eval()\n    model.load_state_dict(torch.load(ckpt_path))\n    BS = 32\n    NW = 8\n    MAX_LEN = 64\n    \n    tmp_df = test_df\n    \n    test_df[\"pct_rank\"] = 0\n    test_ds = MarkdownDataset(test_df[test_df[\"cell_type\"] == \"markdown\"].reset_index(drop=True), md_max_len=64,total_max_len=512, model_name_or_path=model_path, fts=test_fts)\n    test_loader = DataLoader(test_ds, batch_size=BS, shuffle=False, num_workers=NW,\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-12T09:31:32.955530Z","iopub.execute_input":"2022-07-12T09:31:32.955781Z","iopub.status.idle":"2022-07-12T09:31:32.967200Z","shell.execute_reply.started":"2022-07-12T09:31:32.955748Z","shell.execute_reply":"2022-07-12T09:31:32.966533Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model_path = \"../input/codebert-base/codebert-base/\"\nckpt_path = \"../input/ai4codebaseline/model.bin\"\ny_test = predict(model_path, ckpt_path)","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-12T09:31:32.969927Z","iopub.execute_input":"2022-07-12T09:31:32.970557Z","iopub.status.idle":"2022-07-12T09:31:51.122180Z","shell.execute_reply.started":"2022-07-12T09:31:32.970518Z","shell.execute_reply":"2022-07-12T09:31:51.121373Z"},"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-12T09:31:51.124142Z","iopub.execute_input":"2022-07-12T09:31:51.124647Z","iopub.status.idle":"2022-07-12T09:31:51.132109Z","shell.execute_reply.started":"2022-07-12T09:31:51.124607Z","shell.execute_reply":"2022-07-12T09:31:51.130772Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_df.head()","metadata":{"execution":{"iopub.status.busy":"2022-07-12T09:31:51.133523Z","iopub.execute_input":"2022-07-12T09:31:51.133810Z","iopub.status.idle":"2022-07-12T09:31:51.179350Z","shell.execute_reply.started":"2022-07-12T09:31:51.133772Z","shell.execute_reply":"2022-07-12T09:31:51.178536Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"default_begin = \"# This Python 3 environment comes with many helpful analytics libraries installed\"\ntest_df[\"pred\"] = np.where(test_df[\"source\"].str.contains(default_begin), -10000, test_df[\"pred\"])","metadata":{"execution":{"iopub.status.busy":"2022-07-12T09:31:51.183365Z","iopub.execute_input":"2022-07-12T09:31:51.184003Z","iopub.status.idle":"2022-07-12T09:31:51.194598Z","shell.execute_reply.started":"2022-07-12T09:31:51.183959Z","shell.execute_reply":"2022-07-12T09:31:51.190545Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub_str1 = 'to_csv\\(\\\"submission.csv\\\"'\nsub_str2 = \"to_csv\\(\\'submission.csv\\'\"\ntest_df[\"pred\"] = np.where(test_df[\"source\"].str.contains(sub_str1), 10000, test_df[\"pred\"])\ntest_df[\"pred\"] = np.where(test_df[\"source\"].str.contains(sub_str2), 10000, test_df[\"pred\"])","metadata":{"execution":{"iopub.status.busy":"2022-07-12T09:38:28.202160Z","iopub.execute_input":"2022-07-12T09:38:28.202431Z","iopub.status.idle":"2022-07-12T09:38:28.209087Z","shell.execute_reply.started":"2022-07-12T09:38:28.202399Z","shell.execute_reply":"2022-07-12T09:38:28.208401Z"},"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-07-12T09:31:51.673952Z","iopub.status.idle":"2022-07-12T09:31:51.674614Z","shell.execute_reply.started":"2022-07-12T09:31:51.674365Z","shell.execute_reply":"2022-07-12T09:31:51.674390Z"},"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-12T09:31:51.675899Z","iopub.status.idle":"2022-07-12T09:31:51.676440Z","shell.execute_reply.started":"2022-07-12T09:31:51.676208Z","shell.execute_reply":"2022-07-12T09:31:51.676233Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_df = test_df.merge(test_df.groupby(\"id\")[\"cell_id\"].count().reset_index(name=\"cell_counts\"), on=\"id\", how=\"left\")\ntest_df = test_df[test_df[\"cell_counts\"] <= 200]\ntest_df","metadata":{"execution":{"iopub.status.busy":"2022-07-12T09:31:51.677522Z","iopub.status.idle":"2022-07-12T09:31:51.678062Z","shell.execute_reply.started":"2022-07-12T09:31:51.677834Z","shell.execute_reply":"2022-07-12T09:31:51.677858Z"},"trusted":true},"execution_count":null,"outputs":[]}]}