{"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":"# AI4Code Pytorch DistilBert \n\nI used a lot of code from Kaggle's starter notebook here: https://www.kaggle.com/code/ryanholbrook/getting-started-with-ai4code\nand here: https://www.kaggle.com/code/aerdem4/ai4code-pytorch-distilbert-baseline\n\nI replaced their model with a DistilBert model.","metadata":{"papermill":{"duration":0.031568,"end_time":"2022-05-12T10:15:13.890382","exception":false,"start_time":"2022-05-12T10:15:13.858814","status":"completed"},"tags":[]}},{"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\nfrom sklearn.metrics import mean_squared_error\n\npd.options.display.width = 180\npd.options.display.max_colwidth = 120\n\nBERT_PATH = \"../input/huggingface-bert-variants/distilbert-base-uncased/distilbert-base-uncased\"","metadata":{"papermill":{"duration":0.122804,"end_time":"2022-05-12T10:15:14.04297","exception":false,"start_time":"2022-05-12T10:15:13.920166","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-07-16T03:07:39.873188Z","iopub.execute_input":"2022-07-16T03:07:39.873555Z","iopub.status.idle":"2022-07-16T03:07:40.431153Z","shell.execute_reply.started":"2022-07-16T03:07:39.873449Z","shell.execute_reply":"2022-07-16T03:07:40.430206Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df = pd.read_csv('../input/preprocessed-data/train.csv')\nval_df = pd.read_csv('../input/preprocessed-data/val.csv')\ntrain_df","metadata":{"papermill":{"duration":0.27219,"end_time":"2022-05-12T10:16:50.495502","exception":false,"start_time":"2022-05-12T10:16:50.223312","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-07-16T03:07:40.433258Z","iopub.execute_input":"2022-07-16T03:07:40.433606Z","iopub.status.idle":"2022-07-16T03:08:35.766985Z","shell.execute_reply.started":"2022-07-16T03:07:40.433566Z","shell.execute_reply":"2022-07-16T03:08:35.766156Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"id_na = list(train_df.loc[train_df['source'].isna() | train_df['pct_rank'].isna(), 'id'].unique())\ntrain_df = train_df[~train_df['id'].isin(id_na)].reset_index(drop=True)\nid_na = list(val_df.loc[val_df['source'].isna() | val_df['pct_rank'].isna(), 'id'].unique())\nval_df = val_df[~val_df['id'].isin(id_na)].reset_index(drop=True)","metadata":{"execution":{"iopub.status.busy":"2022-07-16T03:08:35.768864Z","iopub.execute_input":"2022-07-16T03:08:35.769955Z","iopub.status.idle":"2022-07-16T03:08:38.532086Z","shell.execute_reply.started":"2022-07-16T03:08:35.769913Z","shell.execute_reply":"2022-07-16T03:08:38.531002Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from bisect import bisect\n\n\ndef count_inversions(a):\n    inversions = 0\n    sorted_so_far = []\n    for i, u in enumerate(a):\n        j = bisect(sorted_so_far, u)\n        inversions += i - j\n        sorted_so_far.insert(j, u)\n    return inversions\n\n\ndef kendall_tau(ground_truth, predictions):\n    total_inversions = 0\n    total_2max = 0  # twice the maximum possible inversions across all instances\n    for gt, pred in zip(ground_truth, predictions):\n        ranks = [gt.index(x) for x in pred]  # rank predicted order in terms of ground truth\n        total_inversions += count_inversions(ranks)\n        n = len(gt)\n        total_2max += n * (n - 1)\n    return 1 - 4 * total_inversions / total_2max","metadata":{"papermill":{"duration":0.262837,"end_time":"2022-05-12T10:16:51.011588","exception":false,"start_time":"2022-05-12T10:16:50.748751","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-07-16T03:08:38.535362Z","iopub.execute_input":"2022-07-16T03:08:38.535669Z","iopub.status.idle":"2022-07-16T03:08:38.543399Z","shell.execute_reply.started":"2022-07-16T03:08:38.535618Z","shell.execute_reply":"2022-07-16T03:08:38.542382Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df_mark = train_df[train_df[\"cell_type\"] == \"markdown\"].reset_index(drop=True)\n\nval_df_mark = val_df[val_df[\"cell_type\"] == \"markdown\"].reset_index(drop=True)","metadata":{"papermill":{"duration":0.371916,"end_time":"2022-05-12T10:16:52.797271","exception":false,"start_time":"2022-05-12T10:16:52.425355","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-07-16T03:08:38.545311Z","iopub.execute_input":"2022-07-16T03:08:38.545971Z","iopub.status.idle":"2022-07-16T03:08:39.935565Z","shell.execute_reply.started":"2022-07-16T03:08:38.545867Z","shell.execute_reply":"2022-07-16T03:08:39.934608Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from tqdm import tqdm\nimport sys, os\nfrom transformers import DistilBertModel, DistilBertTokenizer\nimport torch.nn.functional as F\nimport torch.nn as nn\nimport torch\n\nMAX_LEN = 128\n    \nclass MarkdownModel(nn.Module):\n    def __init__(self):\n        super(MarkdownModel, self).__init__()\n        self.distill_bert = DistilBertModel.from_pretrained(BERT_PATH)\n        self.top = nn.Linear(768, 1)\n        \n    def forward(self, ids, mask):\n        x = self.distill_bert(ids, mask)[0]\n        x = self.top(x[:, 0, :])\n        x = torch.sigmoid(x)\n        return x","metadata":{"papermill":{"duration":7.145711,"end_time":"2022-05-12T10:17:00.757077","exception":false,"start_time":"2022-05-12T10:16:53.611366","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-07-16T03:08:39.938050Z","iopub.execute_input":"2022-07-16T03:08:39.938826Z","iopub.status.idle":"2022-07-16T03:08:47.759834Z","shell.execute_reply.started":"2022-07-16T03:08:39.938779Z","shell.execute_reply":"2022-07-16T03:08:47.758894Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torch.utils.data import DataLoader, Dataset\n\n\n\nclass MarkdownDataset(Dataset):\n    \n    def __init__(self, df, max_len):\n        super().__init__()\n        self.df = df.reset_index(drop=True)\n        self.max_len = max_len\n        self.tokenizer = DistilBertTokenizer.from_pretrained(BERT_PATH, do_lower_case=True)\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.max_len,\n            padding=\"max_length\",\n            return_token_type_ids=True,\n            truncation=True\n        )\n        ids = torch.LongTensor(inputs['input_ids'])\n        mask = torch.LongTensor(inputs['attention_mask'])\n\n        return ids, mask, torch.FloatTensor([row.pct_rank])\n\n    def __len__(self):\n        return self.df.shape[0]\n    \ntrain_ds = MarkdownDataset(train_df_mark, max_len=MAX_LEN)\nval_ds = MarkdownDataset(val_df_mark, max_len=MAX_LEN)\n\nval_ds[0]","metadata":{"papermill":{"duration":0.474499,"end_time":"2022-05-12T10:17:01.487031","exception":false,"start_time":"2022-05-12T10:17:01.012532","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-07-16T03:08:47.761778Z","iopub.execute_input":"2022-07-16T03:08:47.762169Z","iopub.status.idle":"2022-07-16T03:08:48.162863Z","shell.execute_reply.started":"2022-07-16T03:08:47.762118Z","shell.execute_reply":"2022-07-16T03:08:48.161780Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def adjust_lr(optimizer, epoch):\n    if epoch < 1:\n        lr = 5e-5\n    elif epoch < 2:\n        lr = 4e-5\n    elif epoch < 5:\n        lr = 3e-5\n    else:\n        lr = 2e-5\n\n    for p in optimizer.param_groups:\n        p['lr'] = lr\n    return lr\n    \ndef get_optimizer(net):\n    optimizer = torch.optim.Adam(filter(lambda p: p.requires_grad, net.parameters()), lr=3e-4, betas=(0.9, 0.999),\n                                 eps=1e-08)\n    return optimizer","metadata":{"papermill":{"duration":0.265988,"end_time":"2022-05-12T10:17:02.580374","exception":false,"start_time":"2022-05-12T10:17:02.314386","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-07-16T03:08:48.165112Z","iopub.execute_input":"2022-07-16T03:08:48.165680Z","iopub.status.idle":"2022-07-16T03:08:48.173689Z","shell.execute_reply.started":"2022-07-16T03:08:48.165635Z","shell.execute_reply":"2022-07-16T03:08:48.172663Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"BS = 128\nNW = 2\n\ntrain_loader = DataLoader(train_ds, batch_size=BS, shuffle=True, num_workers=NW,\n                          pin_memory=False, drop_last=True)\nval_loader = DataLoader(val_ds, batch_size=BS, shuffle=False, num_workers=NW,\n                          pin_memory=False, drop_last=False)","metadata":{"papermill":{"duration":0.298424,"end_time":"2022-05-12T10:17:03.132","exception":false,"start_time":"2022-05-12T10:17:02.833576","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-07-16T03:08:48.174990Z","iopub.execute_input":"2022-07-16T03:08:48.175978Z","iopub.status.idle":"2022-07-16T03:08:48.195298Z","shell.execute_reply.started":"2022-07-16T03:08:48.175929Z","shell.execute_reply":"2022-07-16T03:08:48.194002Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"val_df.groupby([\"id\", \"cell_type\"])[\"rank\"].rank(pct=True) - val_df['pct_rank']","metadata":{"execution":{"iopub.status.busy":"2022-07-16T03:08:48.199518Z","iopub.execute_input":"2022-07-16T03:08:48.200220Z","iopub.status.idle":"2022-07-16T03:08:48.424798Z","shell.execute_reply.started":"2022-07-16T03:08:48.200188Z","shell.execute_reply":"2022-07-16T03:08:48.423775Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data_dir = Path('../input/AI4Code')\ndf_orders = pd.read_csv(\n    data_dir / 'train_orders.csv',\n    index_col='id',\n    squeeze=True,\n).str.split()  ","metadata":{"execution":{"iopub.status.busy":"2022-07-16T06:47:51.101851Z","iopub.execute_input":"2022-07-16T06:47:51.102835Z","iopub.status.idle":"2022-07-16T06:47:54.201224Z","shell.execute_reply.started":"2022-07-16T06:47:51.102796Z","shell.execute_reply":"2022-07-16T06:47:54.199990Z"},"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[0], inputs[1])\n\n            preds.append(pred.detach().cpu().numpy().ravel())\n            labels.append(target.detach().cpu().numpy().ravel())\n    return np.concatenate(labels), np.concatenate(preds)\n\ndef train(model, train_loader, val_loader, epochs):\n    np.random.seed(0)\n    \n    optimizer = get_optimizer(model)\n\n    criterion = torch.nn.L1Loss()\n    \n    for e in range(epochs):   \n        model.train()\n        tbar = tqdm(train_loader, file=sys.stdout)\n        \n        lr = adjust_lr(optimizer, e)\n        \n        loss_list = []\n        preds = []\n        labels = []\n\n        for idx, data in enumerate(tbar):\n            inputs, target = read_data(data)\n\n            optimizer.zero_grad()\n            pred = model(inputs[0], inputs[1])\n\n            loss = criterion(pred, target)\n            loss.backward()\n            optimizer.step()\n            \n            loss_list.append(loss.detach().cpu().item())\n            preds.append(pred.detach().cpu().numpy().ravel())\n            labels.append(target.detach().cpu().numpy().ravel())\n            \n            avg_loss = np.round(np.mean(loss_list), 4)\n\n            tbar.set_description(f\"Epoch {e+1} Loss: {avg_loss} lr: {lr}\")\n            \n        y_val, y_pred = validate(model, val_loader)\n            \n        val_df[\"pred\"] = val_df.groupby([\"id\", \"cell_type\"])[\"rank\"].rank(pct=True)\n        val_df.loc[val_df[\"cell_type\"] == \"markdown\", \"pred\"] = y_pred\n\n        y_dummy = val_df.sort_values(\"pred\").groupby('id')['cell_id'].apply(list)\n        score = kendall_tau(df_orders.loc[y_dummy.index], y_dummy)\n\n        print('score : ', score)\n\n        output_model_file = f\"./my_own_model_file_{e}_{np.round(score, 5)}.bin\"\n        model_to_save = model.module if hasattr(model, 'module') else model\n        torch.save(model_to_save.state_dict(), output_model_file)\n                        \n        print(\"Validation MSE:\", np.round(mean_squared_error(y_val, y_pred), 4))\n        print()\n    return model, y_pred\n\nmodel = MarkdownModel()\nmodel = model.cuda()\n# model.load_state_dict(torch.load('../input/model-markdown-3/model.bin'))\nmodel, y_pred = train(model, train_loader, val_loader, epochs=1)\n\n\ntorch.save(model.state_dict(), './model.bin')","metadata":{"papermill":{"duration":987.160977,"end_time":"2022-05-12T10:33:30.548236","exception":false,"start_time":"2022-05-12T10:17:03.387259","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-07-16T03:08:48.426590Z","iopub.execute_input":"2022-07-16T03:08:48.426905Z","iopub.status.idle":"2022-07-16T06:38:24.280477Z","shell.execute_reply.started":"2022-07-16T03:08:48.426848Z","shell.execute_reply":"2022-07-16T06:38:24.278785Z"},"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)\ntest_df[\"pct_rank\"] = 0\ntest_ds = MarkdownDataset(test_df[test_df[\"cell_type\"] == \"markdown\"].reset_index(drop=True), max_len=MAX_LEN)\ntest_loader = DataLoader(test_ds, batch_size=BS, shuffle=False, num_workers=NW,\n                          pin_memory=False, drop_last=False)\nmodel = MarkdownModel()\nmodel = model.cuda()\ny_test = validate(model, test_loader)[1]\ntest_df.loc[test_df[\"cell_type\"] == \"markdown\", \"pred\"] = y_test\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()\nsub_df.to_csv(\"submission.csv\", index=False)","metadata":{"execution":{"iopub.status.busy":"2022-07-16T06:46:44.309925Z","iopub.execute_input":"2022-07-16T06:46:44.310218Z","iopub.status.idle":"2022-07-16T06:46:45.846403Z","shell.execute_reply.started":"2022-07-16T06:46:44.310189Z","shell.execute_reply":"2022-07-16T06:46:45.845431Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}