{"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":"# !mkdir 'raw' 'tfrec'\n# !pip install fast_map\n# %env TOKENIZERS_PARALLELISM=true","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-07-12T11:38:14.271909Z","iopub.execute_input":"2022-07-12T11:38:14.272386Z","iopub.status.idle":"2022-07-12T11:38:14.292015Z","shell.execute_reply.started":"2022-07-12T11:38:14.272296Z","shell.execute_reply":"2022-07-12T11:38:14.291135Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!mkdir 'raw' 'tfrec'\n!pip install fast_map","metadata":{"execution":{"iopub.status.busy":"2022-07-12T11:40:56.240251Z","iopub.execute_input":"2022-07-12T11:40:56.240758Z","iopub.status.idle":"2022-07-12T11:41:11.150235Z","shell.execute_reply.started":"2022-07-12T11:40:56.240657Z","shell.execute_reply":"2022-07-12T11:41:11.148530Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import glob\nimport json\nimport os\nimport random\nfrom typing import List\n\nimport numpy as np\nimport pandas as pd\nimport tensorflow as tf\nimport transformers\nfrom sklearn.model_selection import GroupKFold\nfrom sklearn.utils import shuffle\nfrom tqdm.notebook import tqdm\nfrom fast_map import fast_map\nfrom operator import itemgetter","metadata":{"execution":{"iopub.status.busy":"2022-07-12T11:41:11.154229Z","iopub.execute_input":"2022-07-12T11:41:11.154670Z","iopub.status.idle":"2022-07-12T11:41:19.217557Z","shell.execute_reply.started":"2022-07-12T11:41:11.154633Z","shell.execute_reply":"2022-07-12T11:41:19.216141Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"RANDOM_STATE = 42\nMD_MAX_LEN = 64\nTOTAL_MAX_LEN = 512\nK_FOLDS = 5\nFILES_PER_FOLD = 16\nLIMIT = None\nMODEL_NAME = \"microsoft/codebert-base\"\nTOKENIZER = transformers.AutoTokenizer.from_pretrained(MODEL_NAME)\nINPUT_PATH = \"../input/AI4Code\"","metadata":{"execution":{"iopub.status.busy":"2022-07-12T11:41:19.219493Z","iopub.execute_input":"2022-07-12T11:41:19.220462Z","iopub.status.idle":"2022-07-12T11:41:27.751509Z","shell.execute_reply.started":"2022-07-12T11:41:19.220418Z","shell.execute_reply":"2022-07-12T11:41:27.750200Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def read_notebook(path: str) -> pd.DataFrame:\n    return (\n        pd.read_json(path, dtype={\"cell_type\": \"category\", \"source\": \"str\"})\n        .assign(id=os.path.basename(path).split(\".\")[0])\n        .rename_axis(\"cell_id\")\n    )\n\n\ndef clean_code(cell: str) -> str:\n    return str(cell).replace(\"\\\\n\", \"\\n\")\n\ndef sample_cells(cells, n):\n    if n >= len(cells):\n        return [clean_code(cell) for cell in cells]\n    \n    index_list = random.sample(range(len(cells)), n)\n    index_list.sort()\n    result = list(itemgetter(*index_list)(cells))\n    return [clean_code(cell) for cell in result]\n        \n\ndef get_features(df: pd.DataFrame) -> dict:\n    features = {}\n    for i, sub_df in tqdm(df.groupby(\"id\"), desc=\"Features\"):\n        features[i] = {}\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[i][\"total_code\"] = total_code\n        features[i][\"total_md\"] = total_md\n        features[i][\"codes\"] = codes\n    return features\n\n\ndef tokenize(df: pd.DataFrame, fts: dict) -> dict:\n    input_ids = np.zeros((len(df), TOTAL_MAX_LEN), dtype=np.int32)\n    attention_mask = np.zeros((len(df), TOTAL_MAX_LEN), dtype=np.int32)\n    features = np.zeros((len(df),), dtype=np.float32)\n    labels = np.zeros((len(df),), dtype=np.float32)\n\n    for i, row in tqdm(\n        df.reset_index(drop=True).iterrows(), desc=\"Tokens\", total=len(df)\n    ):\n        row_fts = fts[row.id]\n\n        inputs = TOKENIZER.encode_plus(\n            row.source,\n            None,\n            add_special_tokens=True,\n            max_length=MD_MAX_LEN,\n            padding=\"max_length\",\n            return_token_type_ids=True,\n            truncation=True,\n        )\n        code_inputs = TOKENIZER.batch_encode_plus(\n            [str(x) for x in row_fts[\"codes\"]] or [\"\"],\n            add_special_tokens=True,\n            padding=\"do_not_pad\",\n            truncation=False,\n        )\n\n#         CODE_MAX_LEN = (TOTAL_MAX_LEN - len(inputs[\"input_ids\"])) // len(row_fts[\"codes\"])\n#         if CODE_MAX_LEN < 10:\n#             print(CODE_MAX_LEN)\n            \n        \n#         def trunc1(x):\n#             return x[:CODE_MAX_LEN]\n        \n#         def trunc2(x):\n#             t = (len(x)-CODE_MAX_LEN) // 2\n#             return x[t:len(x)-t]\n        \n#         code_inputs['input_ids'] = fast_map(trunc1, code_inputs['input_ids'])\n#         code_inputs[\"attention_mask\"] = fast_map(trunc1, code_inputs[\"attention_mask\"])\n            \n        ids = inputs[\"input_ids\"]\n        for x in code_inputs[\"input_ids\"]:\n            ids.extend(x[:-1])\n        ids = ids[:TOTAL_MAX_LEN]\n        if len(ids) != TOTAL_MAX_LEN:\n            ids = ids + [\n                TOKENIZER.pad_token_id,\n            ] * (TOTAL_MAX_LEN - len(ids))\n\n        mask = inputs[\"attention_mask\"]\n        for x in code_inputs[\"attention_mask\"]:\n            mask.extend(x[:-1])\n        mask = mask[:TOTAL_MAX_LEN]\n        if len(mask) != TOTAL_MAX_LEN:\n            mask = mask + [\n                TOKENIZER.pad_token_id,\n            ] * (TOTAL_MAX_LEN - len(mask))\n\n        input_ids[i] = ids\n        attention_mask[i] = mask\n        features[i] = (\n            row_fts[\"total_md\"] / (row_fts[\"total_md\"] + row_fts[\"total_code\"]) or 1\n        )\n        labels[i] = row.pct_rank\n\n    return {\n        \"input_ids\": input_ids,\n        \"attention_mask\": attention_mask,\n        \"features\": features,\n        \"labels\": labels,\n    }\n\n\ndef get_ranks(base: pd.Series, derived: List[str]) -> List[str]:\n    return [base.index(d) for d in derived]\n\n\ndef _serialize_sample(\n    input_ids: np.array,\n    attention_mask: np.array,\n    feature: np.float64,\n    label: np.float64,\n) -> bytes:\n    feature = {\n        \"input_ids\": tf.train.Feature(int64_list=tf.train.Int64List(value=input_ids)),\n        \"attention_mask\": tf.train.Feature(\n            int64_list=tf.train.Int64List(value=attention_mask)\n        ),\n        \"feature\": tf.train.Feature(float_list=tf.train.FloatList(value=[feature])),\n        \"label\": tf.train.Feature(float_list=tf.train.FloatList(value=[label])),\n    }\n    sample = tf.train.Example(features=tf.train.Features(feature=feature))\n    return sample.SerializeToString()\n\n\ndef serialize(\n    input_ids: np.array,\n    attention_mask: np.array,\n    features: np.array,\n    labels: np.array,\n    path: str,\n) -> None:\n    with tf.io.TFRecordWriter(path) as writer:\n        for args in zip(input_ids, attention_mask, features, labels):\n            writer.write(_serialize_sample(*args))","metadata":{"execution":{"iopub.status.busy":"2022-07-12T11:41:58.532547Z","iopub.execute_input":"2022-07-12T11:41:58.533078Z","iopub.status.idle":"2022-07-12T11:41:58.564148Z","shell.execute_reply.started":"2022-07-12T11:41:58.533045Z","shell.execute_reply":"2022-07-12T11:41:58.562776Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"paths = glob.glob(os.path.join(INPUT_PATH, \"train\", \"*.json\"))\nif LIMIT is not None:\n    paths = paths[:LIMIT]\ndf = (\n    pd.concat([read_notebook(x) for x in tqdm(paths, desc=\"Concat\")])\n    .set_index(\"id\", append=True)\n    .swaplevel()\n    .sort_index(level=\"id\", sort_remaining=False)\n)\n\ndf_orders = pd.read_csv(\n    os.path.join(INPUT_PATH, \"train_orders.csv\"),\n    index_col=\"id\",\n    squeeze=True,\n).str.split()\ndf_orders_ = df_orders.to_frame().join(\n    df.reset_index(\"cell_id\").groupby(\"id\")[\"cell_id\"].apply(list),\n    how=\"right\",\n)\n\nranks = {}\nfor id_, cell_order, cell_id in df_orders_.itertuples():\n    ranks[id_] = {\"cell_id\": cell_id, \"rank\": get_ranks(cell_order, cell_id)}\ndf_ranks = (\n    pd.DataFrame.from_dict(ranks, orient=\"index\")\n    .rename_axis(\"id\")\n    .apply(pd.Series.explode)\n    .set_index(\"cell_id\", append=True)\n)\n\ndf_ancestors = pd.read_csv(\n    os.path.join(INPUT_PATH, \"train_ancestors.csv\"), index_col=\"id\"\n)\ndf = (\n    df.reset_index()\n    .merge(df_ranks, on=[\"id\", \"cell_id\"])\n    .merge(df_ancestors, on=[\"id\"])\n)\n\ndf[\"pct_rank\"] = df[\"rank\"] / df.groupby(\"id\")[\"cell_id\"].transform(\"count\")\ndf = df.sort_values(\"pct_rank\").reset_index(drop=True)\n\nfeatures = get_features(df)\n\ndf = df[df[\"cell_type\"] == \"markdown\"]\ndf = df.drop([\"rank\", \"parent_id\", \"cell_type\"], axis=1).dropna()","metadata":{"execution":{"iopub.status.busy":"2022-07-12T11:41:59.655355Z","iopub.execute_input":"2022-07-12T11:41:59.655835Z","iopub.status.idle":"2022-07-12T11:42:17.312962Z","shell.execute_reply.started":"2022-07-12T11:41:59.655790Z","shell.execute_reply":"2022-07-12T11:42:17.311125Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df.to_csv(\"data.csv\")\nwith open(\"features.json\", \"w\") as file:\n    json.dump(features, file)","metadata":{"execution":{"iopub.status.busy":"2022-07-12T11:42:25.598792Z","iopub.execute_input":"2022-07-12T11:42:25.600068Z","iopub.status.idle":"2022-07-12T11:42:25.629063Z","shell.execute_reply.started":"2022-07-12T11:42:25.600021Z","shell.execute_reply":"2022-07-12T11:42:25.626933Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = shuffle(df, random_state=RANDOM_STATE)\n\nfor fold, (_, split) in enumerate(\n    GroupKFold(K_FOLDS).split(df, groups=df[\"ancestor_id\"])\n):\n    print(\"=\" * 36, f\"Fold {fold}\", \"=\" * 36)\n    fold_dir = f\"tfrec/{fold}\"\n    if not os.path.exists(fold_dir):\n        os.mkdir(fold_dir)\n\n    data = tokenize(df.iloc[split], features)\n\n    np.savez_compressed(\n        f\"raw/{fold}.npz\",\n        input_ids=data[\"input_ids\"],\n        attention_mask=data[\"attention_mask\"],\n        features=data[\"features\"],\n        labels=data[\"labels\"],\n    )\n\n    for split, index in tqdm(\n        enumerate(np.array_split(np.arange(data[\"labels\"].shape[0]), FILES_PER_FOLD)),\n        desc=f\"Saving\",\n        total=FILES_PER_FOLD,\n    ):\n        serialize(\n            input_ids=data[\"input_ids\"][index],\n            attention_mask=data[\"attention_mask\"][index],\n            features=data[\"features\"][index],\n            labels=data[\"labels\"][index],\n            path=os.path.join(fold_dir, f\"{split:02d}-{len(index):06d}.tfrec\"),\n        )","metadata":{"execution":{"iopub.status.busy":"2022-07-12T11:42:33.518993Z","iopub.execute_input":"2022-07-12T11:42:33.519414Z","iopub.status.idle":"2022-07-12T11:42:33.564327Z","shell.execute_reply.started":"2022-07-12T11:42:33.519384Z","shell.execute_reply":"2022-07-12T11:42:33.561946Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!zip -r tfrec.zip tfrec ","metadata":{"execution":{"iopub.status.busy":"2022-07-12T11:38:27.112836Z","iopub.status.idle":"2022-07-12T11:38:27.113374Z","shell.execute_reply.started":"2022-07-12T11:38:27.113091Z","shell.execute_reply":"2022-07-12T11:38:27.113116Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!ls","metadata":{"execution":{"iopub.status.busy":"2022-07-12T11:38:27.114969Z","iopub.status.idle":"2022-07-12T11:38:27.115500Z","shell.execute_reply.started":"2022-07-12T11:38:27.115220Z","shell.execute_reply":"2022-07-12T11:38:27.115246Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}