{"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":"It's based on \n* https://www.kaggle.com/code/suicaokhoailang/stronger-baseline-with-code-cells\n* https://www.kaggle.com/code/aerdem4/ai4code-pytorch-distilbert-baseline\n* https://www.kaggle.com/code/ryanholbrook/getting-started-with-ai4code\n","metadata":{"execution":{"iopub.status.busy":"2022-07-19T00:03:17.014023Z","iopub.execute_input":"2022-07-19T00:03:17.014844Z","iopub.status.idle":"2022-07-19T00:03:17.048802Z","shell.execute_reply.started":"2022-07-19T00:03:17.014747Z","shell.execute_reply":"2022-07-19T00:03:17.047348Z"}}},{"cell_type":"code","source":"import json\nfrom pathlib import Path\nimport random\nimport os\nimport sys\n\nimport numpy as np\nimport pandas as pd\nfrom scipy import sparse\nfrom tqdm import tqdm\n\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\ntotal_max_len = 512\ndata_dir = Path('../input/AI4Code')\nmodel_path = \"../input/graphcodebert-base-model\"\nckpt_path = \"../input/20220718/model-graphcodeberttest2.bin\"\ntokenizer_path = \"../input/graphcodebert-base-tokenizer\"","metadata":{"execution":{"iopub.execute_input":"2022-07-18T14:28:21.909045Z","iopub.status.busy":"2022-07-18T14:28:21.908381Z","iopub.status.idle":"2022-07-18T14:28:22.000674Z","shell.execute_reply":"2022-07-18T14:28:21.999943Z"},"papermill":{"duration":0.11544,"end_time":"2022-07-18T14:28:22.002783","exception":false,"start_time":"2022-07-18T14:28:21.887343","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def seed_everything(seed=42):\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 = True\n    \nseed_everything(42)","metadata":{"execution":{"iopub.execute_input":"2022-07-18T14:28:22.029048Z","iopub.status.busy":"2022-07-18T14:28:22.028564Z","iopub.status.idle":"2022-07-18T14:28:22.032361Z","shell.execute_reply":"2022-07-18T14:28:22.031727Z"},"papermill":{"duration":0.018661,"end_time":"2022-07-18T14:28:22.034253","exception":false,"start_time":"2022-07-18T14:28:22.015592","status":"completed"},"tags":[]},"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":{"execution":{"iopub.execute_input":"2022-07-18T14:28:22.060580Z","iopub.status.busy":"2022-07-18T14:28:22.060158Z","iopub.status.idle":"2022-07-18T14:28:22.146973Z","shell.execute_reply":"2022-07-18T14:28:22.145184Z"},"papermill":{"duration":0.102206,"end_time":"2022-07-18T14:28:22.148760","exception":false,"start_time":"2022-07-18T14:28:22.046554","status":"completed"},"tags":[]},"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":{"execution":{"iopub.execute_input":"2022-07-18T14:28:22.203448Z","iopub.status.busy":"2022-07-18T14:28:22.203149Z","iopub.status.idle":"2022-07-18T14:28:22.220015Z","shell.execute_reply":"2022-07-18T14:28:22.219330Z"},"papermill":{"duration":0.044369,"end_time":"2022-07-18T14:28:22.222327","exception":false,"start_time":"2022-07-18T14:28:22.177958","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_fts = get_features(test_df)","metadata":{"execution":{"iopub.execute_input":"2022-07-18T14:28:22.259791Z","iopub.status.busy":"2022-07-18T14:28:22.259528Z","iopub.status.idle":"2022-07-18T14:28:22.275641Z","shell.execute_reply":"2022-07-18T14:28:22.274612Z"},"papermill":{"duration":0.034506,"end_time":"2022-07-18T14:28:22.279762","exception":false,"start_time":"2022-07-18T14:28:22.245256","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class 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, tokenizer_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(tokenizer_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":{"execution":{"iopub.execute_input":"2022-07-18T14:28:22.309899Z","iopub.status.busy":"2022-07-18T14:28:22.309371Z","iopub.status.idle":"2022-07-18T14:28:29.936556Z","shell.execute_reply":"2022-07-18T14:28:29.935705Z"},"papermill":{"duration":7.645045,"end_time":"2022-07-18T14:28:29.938891","exception":false,"start_time":"2022-07-18T14:28:22.293846","status":"completed"},"tags":[]},"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, tokenizer_path):\n    model = MarkdownModel(model_path)\n    model = model.to('cuda:0')\n    model.eval()\n    model.load_state_dict(torch.load(ckpt_path, map_location='cuda:0'))\n    BS = 32\n    NW = 8\n    MAX_LEN = 64\n    test_df[\"pct_rank\"] = 0\n    test_ds = MarkdownDataset(test_df[test_df[\"cell_type\"] == \"markdown\"].reset_index(drop=True), \n                              md_max_len=64,\n                              tokenizer_path = tokenizer_path,\n                              total_max_len = total_max_len,\n                              fts=test_fts)\n    test_loader = DataLoader(test_ds, batch_size=BS, \n                             shuffle=False,\n                             num_workers=NW,\n                             pin_memory=False,\n                             drop_last=False)\n    _, y_test = validate(model, test_loader)\n    return y_test","metadata":{"execution":{"iopub.execute_input":"2022-07-18T14:28:29.972300Z","iopub.status.busy":"2022-07-18T14:28:29.972092Z","iopub.status.idle":"2022-07-18T14:28:29.983387Z","shell.execute_reply":"2022-07-18T14:28:29.982513Z"},"papermill":{"duration":0.029813,"end_time":"2022-07-18T14:28:29.985330","exception":false,"start_time":"2022-07-18T14:28:29.955517","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"y_test = predict(model_path, ckpt_path, tokenizer_path)","metadata":{"execution":{"iopub.execute_input":"2022-07-18T14:28:30.057159Z","iopub.status.busy":"2022-07-18T14:28:30.056958Z","iopub.status.idle":"2022-07-18T14:28:47.096437Z","shell.execute_reply":"2022-07-18T14:28:47.094862Z"},"papermill":{"duration":17.057781,"end_time":"2022-07-18T14:28:47.098909","exception":false,"start_time":"2022-07-18T14:28:30.041128","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_df.loc[test_df[\"cell_type\"] == \"markdown\", \"pred\"] = y_test","metadata":{"execution":{"iopub.execute_input":"2022-07-18T14:28:47.172781Z","iopub.status.busy":"2022-07-18T14:28:47.172553Z","iopub.status.idle":"2022-07-18T14:28:47.179054Z","shell.execute_reply":"2022-07-18T14:28:47.178351Z"},"papermill":{"duration":0.025652,"end_time":"2022-07-18T14:28:47.180736","exception":false,"start_time":"2022-07-18T14:28:47.155084","status":"completed"},"tags":[]},"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":{"execution":{"iopub.execute_input":"2022-07-18T14:28:47.213341Z","iopub.status.busy":"2022-07-18T14:28:47.213125Z","iopub.status.idle":"2022-07-18T14:28:47.232423Z","shell.execute_reply":"2022-07-18T14:28:47.231620Z"},"papermill":{"duration":0.038273,"end_time":"2022-07-18T14:28:47.234777","exception":false,"start_time":"2022-07-18T14:28:47.196504","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub_df.to_csv(\"submission.csv\", index=False)","metadata":{"execution":{"iopub.execute_input":"2022-07-18T14:28:47.269505Z","iopub.status.busy":"2022-07-18T14:28:47.269265Z","iopub.status.idle":"2022-07-18T14:28:47.276090Z","shell.execute_reply":"2022-07-18T14:28:47.275442Z"},"papermill":{"duration":0.025563,"end_time":"2022-07-18T14:28:47.277782","exception":false,"start_time":"2022-07-18T14:28:47.252219","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]}]}