{"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":"markdown","source":"## thanks for Khoi Nguyen  https://www.kaggle.com/competitions/AI4Code/discussion/326970","metadata":{}},{"cell_type":"code","source":"import json\nfrom pathlib import Path\nimport regex as re\nimport numpy as np\nimport pandas as pd\nfrom scipy import sparse\nfrom tqdm import tqdm\nfrom bs4 import BeautifulSoup\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-16T01:24:03.971502Z","iopub.execute_input":"2022-07-16T01:24:03.971923Z","iopub.status.idle":"2022-07-16T01:24:04.394299Z","shell.execute_reply.started":"2022-07-16T01:24:03.971813Z","shell.execute_reply":"2022-07-16T01:24:04.393561Z"},"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-16T01:24:04.398928Z","iopub.execute_input":"2022-07-16T01:24:04.401007Z","iopub.status.idle":"2022-07-16T01:24:04.50641Z","shell.execute_reply.started":"2022-07-16T01:24:04.400968Z","shell.execute_reply":"2022-07-16T01:24:04.505561Z"},"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-16T01:24:04.510865Z","iopub.execute_input":"2022-07-16T01:24:04.512768Z","iopub.status.idle":"2022-07-16T01:24:04.540522Z","shell.execute_reply.started":"2022-07-16T01:24:04.51273Z","shell.execute_reply":"2022-07-16T01:24:04.53987Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Additional code cells\ndef preprocess_text(document):\n    # Remove all the special characters\n    #document = re.sub(r'\\W', ' ', str(document))\n\n    # remove all single characters\n    document = re.sub(r'\\s+[a-zA-Z]\\s+', ' ', document)\n\n    # Remove single characters from the start\n    document = re.sub(r'\\^[a-zA-Z]\\s+', ' ', document)\n\n    # Substituting multiple spaces with single space\n    document = re.sub(r'\\s+', ' ', document, flags=re.I)\n\n    # Removing prefixed 'b'\n    document = re.sub(r'^b\\s+', '', document)\n\n    # Converting to Lowercase\n    document = document.lower()\n    #return document\n\n    # Lemmatization\n    tokens = document.split()\n    #tokens = [stemmer.lemmatize(word) for word in tokens]\n    #tokens = [word for word in tokens if len(word) > 3]\n\n    preprocessed_text = ' '.join(tokens)\n    return preprocessed_text\n\ndef preprocess_text_0(document):\n    try:\n        #html = markdown(some_html_string)\n        soup = BeautifulSoup(document, features='html.parser')\n        document = soup.get_text()\n        #document = ''.join(BeautifulSoup(html).findAll(text=True))\n        # Remove all the special characters\n        #document = re.sub(r'\\W', ' ', str(document))\n    except:\n        document = document\n        # remove all single characters\n    document = re.sub(r'\\s+[a-zA-Z]\\s+', ' ', document)\n\n    # Remove single characters from the start\n    document = re.sub(r'\\^[a-zA-Z]\\s+', ' ', document)\n\n    # Substituting multiple spaces with single space\n    document = re.sub(r'\\s+', ' ', document, flags=re.I)\n\n    # Removing prefixed 'b'\n    document = re.sub(r'^b\\s+', '', document)\n\n    # Converting to Lowercase\n    document = document.lower()\n    #return document\n\n    # Lemmatization\n    tokens = document.split()\n    #tokens = [stemmer.lemmatize(word) for word in tokens]\n    #tokens = [word for word in tokens if len(word) > 3]\n\n    preprocessed_text = ' '.join(tokens)\n    return preprocessed_text.strip()\n\ndef sample_cells(cells, n):\n    cells = [preprocess_text(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\ndef sample_cells_0(cells, n):\n    cells = [preprocess_text_0(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\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\n\ndef get_features_0(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_0(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-16T01:24:04.545054Z","iopub.execute_input":"2022-07-16T01:24:04.546956Z","iopub.status.idle":"2022-07-16T01:24:04.580562Z","shell.execute_reply.started":"2022-07-16T01:24:04.546919Z","shell.execute_reply":"2022-07-16T01:24:04.579917Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_fts = get_features(test_df)\ntest_fts_0 = get_features_0(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-16T01:24:04.584861Z","iopub.execute_input":"2022-07-16T01:24:04.587777Z","iopub.status.idle":"2022-07-16T01:24:04.654009Z","shell.execute_reply.started":"2022-07-16T01:24:04.58774Z","shell.execute_reply":"2022-07-16T01:24:04.653364Z"},"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, 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_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-16T01:24:04.657589Z","iopub.execute_input":"2022-07-16T01:24:04.659494Z","iopub.status.idle":"2022-07-16T01:24:12.566407Z","shell.execute_reply.started":"2022-07-16T01:24:04.659458Z","shell.execute_reply":"2022-07-16T01:24:12.565535Z"},"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    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\n\ndef predict_0(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    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_0)\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-16T01:24:12.567722Z","iopub.execute_input":"2022-07-16T01:24:12.567965Z","iopub.status.idle":"2022-07-16T01:24:12.585329Z","shell.execute_reply.started":"2022-07-16T01:24:12.567929Z","shell.execute_reply":"2022-07-16T01:24:12.584438Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model_path = \"../input/codebert-base/codebert-base/\"\nckpt_path = \"../input/ai4codecodebert/model0.bin\"\ny_test_1 = predict_0(model_path, ckpt_path)","metadata":{"papermill":{"duration":0.019733,"end_time":"2022-05-23T03:29:16.134762","exception":false,"start_time":"2022-05-23T03:29:16.115029","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-07-16T01:24:12.586648Z","iopub.execute_input":"2022-07-16T01:24:12.587133Z","iopub.status.idle":"2022-07-16T01:24:32.514678Z","shell.execute_reply.started":"2022-07-16T01:24:12.587097Z","shell.execute_reply":"2022-07-16T01:24:32.513706Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model_path = \"../input/codebert-base/codebert-base/\"\nckpt_path = \"../input/ai4codecodebert/model.bin\"\ny_test_2 = 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-16T01:24:32.516696Z","iopub.execute_input":"2022-07-16T01:24:32.517465Z","iopub.status.idle":"2022-07-16T01:24:41.051518Z","shell.execute_reply.started":"2022-07-16T01:24:32.51742Z","shell.execute_reply":"2022-07-16T01:24:41.050683Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model_path = \"../input/codebert-base/codebert-base/\"\nckpt_path = \"../input/ai4codecodebert/model_1.bin\"\ny_test_3 = predict_0(model_path, ckpt_path)","metadata":{"execution":{"iopub.status.busy":"2022-07-16T01:24:41.055181Z","iopub.execute_input":"2022-07-16T01:24:41.055399Z","iopub.status.idle":"2022-07-16T01:24:47.889938Z","shell.execute_reply.started":"2022-07-16T01:24:41.055369Z","shell.execute_reply":"2022-07-16T01:24:47.889076Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model_path = \"../input/codebert-base/codebert-base/\"\nckpt_path = \"../input/ai4codecodebert/model_4.bin\"\ny_test_4 = predict_0(model_path, ckpt_path)","metadata":{"execution":{"iopub.status.busy":"2022-07-16T01:24:47.891947Z","iopub.execute_input":"2022-07-16T01:24:47.892448Z","iopub.status.idle":"2022-07-16T01:24:54.655731Z","shell.execute_reply.started":"2022-07-16T01:24:47.892406Z","shell.execute_reply":"2022-07-16T01:24:54.654618Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# y_test = (y_test_1 + y_test_2)/2\ny_test = (y_test_1 + y_test_2 + y_test_3 + y_test_4) /4","metadata":{"papermill":{"duration":0.023439,"end_time":"2022-05-23T03:29:36.563143","exception":false,"start_time":"2022-05-23T03:29:36.539704","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-07-16T01:24:54.657902Z","iopub.execute_input":"2022-07-16T01:24:54.65841Z","iopub.status.idle":"2022-07-16T01:24:54.663178Z","shell.execute_reply.started":"2022-07-16T01:24:54.658366Z","shell.execute_reply":"2022-07-16T01:24:54.662166Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_df[test_df['id']=='0009d135ece78d']","metadata":{"execution":{"iopub.status.busy":"2022-07-16T01:24:54.664572Z","iopub.execute_input":"2022-07-16T01:24:54.665078Z","iopub.status.idle":"2022-07-16T01:24:54.690245Z","shell.execute_reply.started":"2022-07-16T01:24:54.665041Z","shell.execute_reply":"2022-07-16T01:24:54.689293Z"},"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-16T01:24:54.691596Z","iopub.execute_input":"2022-07-16T01:24:54.691932Z","iopub.status.idle":"2022-07-16T01:24:54.701351Z","shell.execute_reply.started":"2022-07-16T01:24:54.691893Z","shell.execute_reply":"2022-07-16T01:24:54.700428Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_df[test_df['id']=='0009d135ece78d']","metadata":{"execution":{"iopub.status.busy":"2022-07-16T01:24:54.702822Z","iopub.execute_input":"2022-07-16T01:24:54.703181Z","iopub.status.idle":"2022-07-16T01:24:54.72423Z","shell.execute_reply.started":"2022-07-16T01:24:54.703082Z","shell.execute_reply":"2022-07-16T01:24:54.723529Z"},"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-16T01:24:54.725684Z","iopub.execute_input":"2022-07-16T01:24:54.72595Z","iopub.status.idle":"2022-07-16T01:24:54.74286Z","shell.execute_reply.started":"2022-07-16T01:24:54.725917Z","shell.execute_reply":"2022-07-16T01:24:54.742165Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub_df[sub_df['id']=='0009d135ece78d']","metadata":{"execution":{"iopub.status.busy":"2022-07-16T01:24:54.744414Z","iopub.execute_input":"2022-07-16T01:24:54.745061Z","iopub.status.idle":"2022-07-16T01:24:54.757167Z","shell.execute_reply.started":"2022-07-16T01:24:54.745027Z","shell.execute_reply":"2022-07-16T01:24:54.756127Z"},"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-16T01:24:54.758653Z","iopub.execute_input":"2022-07-16T01:24:54.758934Z","iopub.status.idle":"2022-07-16T01:24:54.770063Z","shell.execute_reply.started":"2022-07-16T01:24:54.758896Z","shell.execute_reply":"2022-07-16T01:24:54.769337Z"},"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":[]}]}