{"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":"# Google Al4Code\n**Setup**","metadata":{}},{"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.feature_extraction.text import TfidfVectorizer\n\npd.options.display.width = 180\npd.options.display.max_colwidth = 120","metadata":{"execution":{"iopub.status.busy":"2022-07-15T10:16:28.422572Z","iopub.execute_input":"2022-07-15T10:16:28.423001Z","iopub.status.idle":"2022-07-15T10:16:28.430958Z","shell.execute_reply.started":"2022-07-15T10:16:28.422959Z","shell.execute_reply":"2022-07-15T10:16:28.429911Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data_dir = Path('../input/AI4Code')","metadata":{"execution":{"iopub.status.busy":"2022-07-15T10:16:28.432159Z","iopub.execute_input":"2022-07-15T10:16:28.432937Z","iopub.status.idle":"2022-07-15T10:16:28.444574Z","shell.execute_reply.started":"2022-07-15T10:16:28.432906Z","shell.execute_reply":"2022-07-15T10:16:28.443552Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"NUM_TRAIN = 9999\n\n\ndef 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\n\npaths_train = list((data_dir / 'train').glob('*.json'))[:NUM_TRAIN]\nnotebooks_train = [\n    read_notebook(path) for path in tqdm(paths_train, desc='Train NBs')\n]\ndf = (\n    pd.concat(notebooks_train)\n    .set_index('id', append=True)\n    .swaplevel()\n    .sort_index(level='id', sort_remaining=False)\n)\ndf","metadata":{"execution":{"iopub.status.busy":"2022-07-15T10:16:28.445930Z","iopub.execute_input":"2022-07-15T10:16:28.446505Z","iopub.status.idle":"2022-07-15T10:17:34.662807Z","shell.execute_reply.started":"2022-07-15T10:16:28.446454Z","shell.execute_reply":"2022-07-15T10:17:34.661652Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import plotly.io as pio\npio.renderers.default='notebook'\nimport plotly.express as px\n\ndf_temp = df.reset_index()\npie_data = df_temp[\"cell_type\"].value_counts().reset_index()\npie_data.columns = [\"cell_type\", \"count\"]\n\nfig = px.pie(pie_data, values='count', names='cell_type', title='Code vs Markdown')\nfig.show()","metadata":{"execution":{"iopub.status.busy":"2022-07-15T10:17:34.665837Z","iopub.execute_input":"2022-07-15T10:17:34.666336Z","iopub.status.idle":"2022-07-15T10:17:34.948772Z","shell.execute_reply.started":"2022-07-15T10:17:34.666287Z","shell.execute_reply":"2022-07-15T10:17:34.947690Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"cell_analysis = df_temp.groupby([\"id\", \"cell_type\"])[\"cell_id\"].count().reset_index()\nscatter_data = pd.pivot(data=cell_analysis, index=\"id\", columns=\"cell_type\", values=\"cell_id\")\nscatter_data[\"size\"] = 30\n\nfig = px.scatter(scatter_data, x=\"code\", y=\"markdown\")\nfig.show()","metadata":{"execution":{"iopub.status.busy":"2022-07-15T10:17:34.950097Z","iopub.execute_input":"2022-07-15T10:17:34.950667Z","iopub.status.idle":"2022-07-15T10:17:35.275907Z","shell.execute_reply.started":"2022-07-15T10:17:34.950632Z","shell.execute_reply":"2022-07-15T10:17:35.274781Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Get an example notebook\nnb_id = df.index.unique('id')[6]\nprint('Notebook:', nb_id)\n\nprint(\"The disordered notebook:\")\nnb = df.loc[nb_id, :]\ndisplay(nb)\nprint()","metadata":{"execution":{"iopub.status.busy":"2022-07-15T10:17:35.277123Z","iopub.execute_input":"2022-07-15T10:17:35.277942Z","iopub.status.idle":"2022-07-15T10:17:35.301611Z","shell.execute_reply.started":"2022-07-15T10:17:35.277907Z","shell.execute_reply":"2022-07-15T10:17:35.300546Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_orders = pd.read_csv(\n    data_dir / 'train_orders.csv',\n    index_col='id',\n    squeeze=True,\n).str.split()  # Split the string representation of cell_ids into a list\n\ndf_orders.head(10)","metadata":{"execution":{"iopub.status.busy":"2022-07-15T10:17:35.303310Z","iopub.execute_input":"2022-07-15T10:17:35.303757Z","iopub.status.idle":"2022-07-15T10:17:37.661632Z","shell.execute_reply.started":"2022-07-15T10:17:35.303725Z","shell.execute_reply":"2022-07-15T10:17:37.660651Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Get the correct order\ncell_order = df_orders.loc[nb_id]\n\nprint(\"The ordered notebook:\")\nnb.loc[cell_order, :]","metadata":{"execution":{"iopub.status.busy":"2022-07-15T10:17:37.663067Z","iopub.execute_input":"2022-07-15T10:17:37.663361Z","iopub.status.idle":"2022-07-15T10:17:37.712902Z","shell.execute_reply.started":"2022-07-15T10:17:37.663334Z","shell.execute_reply":"2022-07-15T10:17:37.711582Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_ranks(base, derived):\n    return [base.index(d) for d in derived]\n\ncell_ranks = get_ranks(cell_order, list(nb.index))\nnb.insert(0, 'rank', cell_ranks)\n\nnb","metadata":{"execution":{"iopub.status.busy":"2022-07-15T10:17:37.714372Z","iopub.execute_input":"2022-07-15T10:17:37.714728Z","iopub.status.idle":"2022-07-15T10:17:37.729098Z","shell.execute_reply.started":"2022-07-15T10:17:37.714697Z","shell.execute_reply":"2022-07-15T10:17:37.728342Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from pandas.testing import assert_frame_equal\n\nassert_frame_equal(nb.loc[cell_order, :], nb.sort_values('rank'))","metadata":{"execution":{"iopub.status.busy":"2022-07-15T10:17:37.730408Z","iopub.execute_input":"2022-07-15T10:17:37.730971Z","iopub.status.idle":"2022-07-15T10:17:37.743548Z","shell.execute_reply.started":"2022-07-15T10:17:37.730939Z","shell.execute_reply":"2022-07-15T10:17:37.742480Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_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)}\n\ndf_ranks = (\n    pd.DataFrame\n    .from_dict(ranks, orient='index')\n    .rename_axis('id')\n    .apply(pd.Series.explode)\n    .set_index('cell_id', append=True)\n)\n\ndf_ranks","metadata":{"execution":{"iopub.status.busy":"2022-07-15T10:17:37.744927Z","iopub.execute_input":"2022-07-15T10:17:37.745236Z","iopub.status.idle":"2022-07-15T10:17:40.947284Z","shell.execute_reply.started":"2022-07-15T10:17:37.745209Z","shell.execute_reply":"2022-07-15T10:17:40.946122Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_ancestors = pd.read_csv(data_dir / 'train_ancestors.csv', index_col='id')\ndf_ancestors","metadata":{"execution":{"iopub.status.busy":"2022-07-15T10:17:40.949150Z","iopub.execute_input":"2022-07-15T10:17:40.949545Z","iopub.status.idle":"2022-07-15T10:17:41.148271Z","shell.execute_reply.started":"2022-07-15T10:17:40.949505Z","shell.execute_reply":"2022-07-15T10:17:41.147103Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from sklearn.model_selection import GroupShuffleSplit\n\nNVALID = 0.1  # size of validation set\n\nsplitter = GroupShuffleSplit(n_splits=1, test_size=NVALID, random_state=0)\n\n# Split, keeping notebooks with a common origin (ancestor_id) together\nids = df.index.unique('id')\nancestors = df_ancestors.loc[ids, 'ancestor_id']\nids_train, ids_valid = next(splitter.split(ids, groups=ancestors))\nids_train, ids_valid = ids[ids_train], ids[ids_valid]\n\ndf_train = df.loc[ids_train, :]\ndf_valid = df.loc[ids_valid, :]","metadata":{"execution":{"iopub.status.busy":"2022-07-15T10:17:41.149856Z","iopub.execute_input":"2022-07-15T10:17:41.150179Z","iopub.status.idle":"2022-07-15T10:17:41.288460Z","shell.execute_reply.started":"2022-07-15T10:17:41.150150Z","shell.execute_reply":"2022-07-15T10:17:41.287338Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from sklearn.feature_extraction.text import TfidfVectorizer\n\n# Training set\ntfidf = TfidfVectorizer(min_df=0.01)\nX_train = tfidf.fit_transform(df_train['source'].astype(str))\n# Rank of each cell within the notebook\ny_train = df_ranks.loc[ids_train].to_numpy()\n# Number of cells in each notebook\ngroups = df_ranks.loc[ids_train].groupby('id').size().to_numpy()","metadata":{"execution":{"iopub.status.busy":"2022-07-15T10:17:41.290189Z","iopub.execute_input":"2022-07-15T10:17:41.290627Z","iopub.status.idle":"2022-07-15T10:17:59.241010Z","shell.execute_reply.started":"2022-07-15T10:17:41.290580Z","shell.execute_reply":"2022-07-15T10:17:59.239971Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Add code cell ordering\nX_train = sparse.hstack((\n    X_train,\n    np.where(\n        df_train['cell_type'] == 'code',\n        df_train.groupby(['id', 'cell_type']).cumcount().to_numpy() + 1,\n        0,\n    ).reshape(-1, 1)\n))\nprint(X_train.shape)","metadata":{"execution":{"iopub.status.busy":"2022-07-15T10:17:59.242727Z","iopub.execute_input":"2022-07-15T10:17:59.243052Z","iopub.status.idle":"2022-07-15T10:17:59.523470Z","shell.execute_reply.started":"2022-07-15T10:17:59.243022Z","shell.execute_reply":"2022-07-15T10:17:59.522288Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from xgboost import XGBRanker\n\nmodel = XGBRanker(\n    min_child_weight=10,\n    subsample=0.5,\n    tree_method='hist',\n)\nmodel.fit(X_train, y_train, group=groups)","metadata":{"execution":{"iopub.status.busy":"2022-07-15T10:17:59.524965Z","iopub.execute_input":"2022-07-15T10:17:59.525317Z","iopub.status.idle":"2022-07-15T10:18:11.961668Z","shell.execute_reply.started":"2022-07-15T10:17:59.525286Z","shell.execute_reply":"2022-07-15T10:18:11.960669Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Validation set\nX_valid = tfidf.transform(df_valid['source'].astype(str))\n\n# The metric uses cell ids\ny_valid = df_orders.loc[ids_valid]\n\nX_valid = sparse.hstack((\n    X_valid,\n    np.where(\n        df_valid['cell_type'] == 'code',\n        df_valid.groupby(['id', 'cell_type']).cumcount().to_numpy() + 1,\n        0,\n    ).reshape(-1, 1)\n))","metadata":{"execution":{"iopub.status.busy":"2022-07-15T10:18:11.963018Z","iopub.execute_input":"2022-07-15T10:18:11.963632Z","iopub.status.idle":"2022-07-15T10:18:13.864004Z","shell.execute_reply.started":"2022-07-15T10:18:11.963600Z","shell.execute_reply":"2022-07-15T10:18:13.862743Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from catboost import FeaturesData","metadata":{"execution":{"iopub.status.busy":"2022-07-15T10:18:13.865684Z","iopub.execute_input":"2022-07-15T10:18:13.866022Z","iopub.status.idle":"2022-07-15T10:18:13.871243Z","shell.execute_reply.started":"2022-07-15T10:18:13.865991Z","shell.execute_reply":"2022-07-15T10:18:13.870153Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Each notebook has all the code cells given first with the markdown cells following. The code cells are in the correct relative order, while the markdown cells are shuffled. In the next section, we'll see how to recover the correct orderings for notebooks in the training set.","metadata":{}},{"cell_type":"code","source":"# Get an example notebook\nnb_id = df.index.unique('id')[6]\nprint('Notebook:', nb_id)\n\nprint(\"The disordered notebook:\")\nnb = df.loc[nb_id, :]\ndisplay(nb)\nprint()","metadata":{"execution":{"iopub.status.busy":"2022-07-15T10:18:13.873327Z","iopub.execute_input":"2022-07-15T10:18:13.874149Z","iopub.status.idle":"2022-07-15T10:18:13.897966Z","shell.execute_reply.started":"2022-07-15T10:18:13.874091Z","shell.execute_reply":"2022-07-15T10:18:13.896690Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_orders = pd.read_csv(\n    data_dir / 'train_orders.csv',\n    index_col='id',\n    squeeze=True,\n).str.split()  # Split the string representation of cell_ids into a list\n\ndf_orders","metadata":{"execution":{"iopub.status.busy":"2022-07-15T10:18:13.899825Z","iopub.execute_input":"2022-07-15T10:18:13.900538Z","iopub.status.idle":"2022-07-15T10:18:15.576671Z","shell.execute_reply.started":"2022-07-15T10:18:13.900498Z","shell.execute_reply":"2022-07-15T10:18:15.575374Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"cell_order = df_orders.loc[nb_id]\n\nprint(\"The ordered notebook:\")\nnb.loc[cell_order, :]","metadata":{"execution":{"iopub.status.busy":"2022-07-15T10:18:15.578623Z","iopub.execute_input":"2022-07-15T10:18:15.579359Z","iopub.status.idle":"2022-07-15T10:18:15.597411Z","shell.execute_reply.started":"2022-07-15T10:18:15.579310Z","shell.execute_reply":"2022-07-15T10:18:15.596237Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_ranks(base, derived):\n    return [base.index(d) for d in derived]\n\ncell_ranks = get_ranks(cell_order, list(nb.index))\nnb.insert(0, 'rank', cell_ranks)\n\nnb","metadata":{"execution":{"iopub.status.busy":"2022-07-15T10:18:15.598724Z","iopub.execute_input":"2022-07-15T10:18:15.599151Z","iopub.status.idle":"2022-07-15T10:18:15.615699Z","shell.execute_reply.started":"2022-07-15T10:18:15.599120Z","shell.execute_reply":"2022-07-15T10:18:15.614594Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from pandas.testing import assert_frame_equal\n\nassert_frame_equal(nb.loc[cell_order, :], nb.sort_values('rank'))","metadata":{"execution":{"iopub.status.busy":"2022-07-15T10:18:15.616922Z","iopub.execute_input":"2022-07-15T10:18:15.617574Z","iopub.status.idle":"2022-07-15T10:18:15.625326Z","shell.execute_reply.started":"2022-07-15T10:18:15.617541Z","shell.execute_reply":"2022-07-15T10:18:15.624455Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_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)}\n\ndf_ranks = (\n    pd.DataFrame\n    .from_dict(ranks, orient='index')\n    .rename_axis('id')\n    .apply(pd.Series.explode)\n    .set_index('cell_id', append=True)\n)\n\ndf_ranks","metadata":{"execution":{"iopub.status.busy":"2022-07-15T10:18:15.626612Z","iopub.execute_input":"2022-07-15T10:18:15.626953Z","iopub.status.idle":"2022-07-15T10:18:19.551623Z","shell.execute_reply.started":"2022-07-15T10:18:15.626923Z","shell.execute_reply":"2022-07-15T10:18:19.550669Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_ancestors = pd.read_csv(data_dir / 'train_ancestors.csv', index_col='id')\ndf_ancestors","metadata":{"execution":{"iopub.status.busy":"2022-07-15T10:18:19.552952Z","iopub.execute_input":"2022-07-15T10:18:19.553266Z","iopub.status.idle":"2022-07-15T10:18:19.733715Z","shell.execute_reply.started":"2022-07-15T10:18:19.553236Z","shell.execute_reply":"2022-07-15T10:18:19.732505Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from sklearn.model_selection import GroupShuffleSplit\n\nNVALID = 0.1  # size of validation set\n\nsplitter = GroupShuffleSplit(n_splits=1, test_size=NVALID, random_state=0)\n\n# Split, keeping notebooks with a common origin (ancestor_id) together\nids = df.index.unique('id')\nancestors = df_ancestors.loc[ids, 'ancestor_id']\nids_train, ids_valid = next(splitter.split(ids, groups=ancestors))\nids_train, ids_valid = ids[ids_train], ids[ids_valid]\n\ndf_train = df.loc[ids_train, :]\ndf_valid = df.loc[ids_valid, :]","metadata":{"execution":{"iopub.status.busy":"2022-07-15T10:18:19.735340Z","iopub.execute_input":"2022-07-15T10:18:19.735748Z","iopub.status.idle":"2022-07-15T10:18:19.875909Z","shell.execute_reply.started":"2022-07-15T10:18:19.735716Z","shell.execute_reply":"2022-07-15T10:18:19.874790Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from sklearn.feature_extraction.text import TfidfVectorizer\n\n# Training set\ntfidf = TfidfVectorizer(min_df=0.01)\nX_train = tfidf.fit_transform(df_train['source'].astype(str))\n# Rank of each cell within the notebook\ny_train = df_ranks.loc[ids_train].to_numpy()\n# Number of cells in each notebook\ngroups = df_ranks.loc[ids_train].groupby('id').size().to_numpy()","metadata":{"execution":{"iopub.status.busy":"2022-07-15T10:18:19.877877Z","iopub.execute_input":"2022-07-15T10:18:19.878310Z","iopub.status.idle":"2022-07-15T10:18:37.621944Z","shell.execute_reply.started":"2022-07-15T10:18:19.878265Z","shell.execute_reply":"2022-07-15T10:18:37.620742Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Add code cell ordering\nX_train = sparse.hstack((\n    X_train,\n    np.where(\n        df_train['cell_type'] == 'code',\n        df_train.groupby(['id', 'cell_type']).cumcount().to_numpy() + 1,\n        0,\n    ).reshape(-1, 1)\n))\nprint(X_train.shape)","metadata":{"execution":{"iopub.status.busy":"2022-07-15T10:18:37.623535Z","iopub.execute_input":"2022-07-15T10:18:37.623901Z","iopub.status.idle":"2022-07-15T10:18:37.902802Z","shell.execute_reply.started":"2022-07-15T10:18:37.623869Z","shell.execute_reply":"2022-07-15T10:18:37.901949Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from xgboost import XGBRanker\n\nmodel = XGBRanker(\n    min_child_weight=10,\n    subsample=0.5,\n    tree_method='hist',\n)\nmodel.fit(X_train, y_train, group=groups)","metadata":{"execution":{"iopub.status.busy":"2022-07-15T10:18:37.904018Z","iopub.execute_input":"2022-07-15T10:18:37.904537Z","iopub.status.idle":"2022-07-15T10:18:50.428283Z","shell.execute_reply.started":"2022-07-15T10:18:37.904504Z","shell.execute_reply":"2022-07-15T10:18:50.427540Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Validation set\nX_valid = tfidf.transform(df_valid['source'].astype(str))\n# The metric uses cell ids\ny_valid = df_orders.loc[ids_valid]\n\nX_valid = sparse.hstack((\n    X_valid,\n    np.where(\n        df_valid['cell_type'] == 'code',\n        df_valid.groupby(['id', 'cell_type']).cumcount().to_numpy() + 1,\n        0,\n    ).reshape(-1, 1)\n))","metadata":{"execution":{"iopub.status.busy":"2022-07-15T10:18:50.429268Z","iopub.execute_input":"2022-07-15T10:18:50.429820Z","iopub.status.idle":"2022-07-15T10:18:52.280940Z","shell.execute_reply.started":"2022-07-15T10:18:50.429783Z","shell.execute_reply":"2022-07-15T10:18:52.279880Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"y_pred = pd.DataFrame({'rank': model.predict(X_valid)}, index=df_valid.index)\ny_pred = (\n    y_pred\n    .sort_values(['id', 'rank'])  # Sort the cells in each notebook by their rank.\n                                  # The cell_ids are now in the order the model predicted.\n    .reset_index('cell_id')  # Convert the cell_id index into a column.\n    .groupby('id')['cell_id'].apply(list)  # Group the cell_ids for each notebook into a list.\n)\ny_pred.head(10)","metadata":{"execution":{"iopub.status.busy":"2022-07-15T10:18:52.282393Z","iopub.execute_input":"2022-07-15T10:18:52.282843Z","iopub.status.idle":"2022-07-15T10:18:52.469635Z","shell.execute_reply.started":"2022-07-15T10:18:52.282808Z","shell.execute_reply":"2022-07-15T10:18:52.468450Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"nb_id = df_valid.index.get_level_values('id').unique()[8]\n\ndisplay(df.loc[nb_id])\ndisplay(df.loc[nb_id].loc[y_pred.loc[nb_id]])","metadata":{"execution":{"iopub.status.busy":"2022-07-15T10:18:52.470801Z","iopub.execute_input":"2022-07-15T10:18:52.471107Z","iopub.status.idle":"2022-07-15T10:18:52.503957Z","shell.execute_reply.started":"2022-07-15T10:18:52.471079Z","shell.execute_reply":"2022-07-15T10:18:52.502706Z"},"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":{"execution":{"iopub.status.busy":"2022-07-15T10:18:52.505656Z","iopub.execute_input":"2022-07-15T10:18:52.506022Z","iopub.status.idle":"2022-07-15T10:18:52.515629Z","shell.execute_reply.started":"2022-07-15T10:18:52.505991Z","shell.execute_reply":"2022-07-15T10:18:52.514396Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"y_dummy = df_valid.reset_index('cell_id').groupby('id')['cell_id'].apply(list)\nkendall_tau(y_valid, y_dummy)","metadata":{"execution":{"iopub.status.busy":"2022-07-15T10:18:52.517727Z","iopub.execute_input":"2022-07-15T10:18:52.518246Z","iopub.status.idle":"2022-07-15T10:18:52.670678Z","shell.execute_reply.started":"2022-07-15T10:18:52.518201Z","shell.execute_reply":"2022-07-15T10:18:52.669417Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"kendall_tau(y_valid, y_pred)","metadata":{"execution":{"iopub.status.busy":"2022-07-15T10:18:52.672457Z","iopub.execute_input":"2022-07-15T10:18:52.673135Z","iopub.status.idle":"2022-07-15T10:18:52.774472Z","shell.execute_reply.started":"2022-07-15T10:18:52.673100Z","shell.execute_reply":"2022-07-15T10:18:52.773360Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"paths_test = list((data_dir / 'test').glob('*.json'))\nnotebooks_test = [\n    read_notebook(path) for path in tqdm(paths_test, desc='Test NBs')\n]\ndf_test = (\n    pd.concat(notebooks_test)\n    .set_index('id', append=True)\n    .swaplevel()\n    .sort_index(level='id', sort_remaining=False)\n)","metadata":{"execution":{"iopub.status.busy":"2022-07-15T10:18:52.775847Z","iopub.execute_input":"2022-07-15T10:18:52.776775Z","iopub.status.idle":"2022-07-15T10:18:52.818750Z","shell.execute_reply.started":"2022-07-15T10:18:52.776736Z","shell.execute_reply":"2022-07-15T10:18:52.817855Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"X_test = tfidf.transform(df_test['source'].astype(str))\nX_test = sparse.hstack((\n    X_test,\n    np.where(\n        df_test['cell_type'] == 'code',\n        df_test.groupby(['id', 'cell_type']).cumcount().to_numpy() + 1,\n        0,\n    ).reshape(-1, 1)\n))","metadata":{"execution":{"iopub.status.busy":"2022-07-15T10:18:52.820157Z","iopub.execute_input":"2022-07-15T10:18:52.820464Z","iopub.status.idle":"2022-07-15T10:18:52.841855Z","shell.execute_reply.started":"2022-07-15T10:18:52.820438Z","shell.execute_reply":"2022-07-15T10:18:52.840553Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"y_infer = pd.DataFrame({'rank': model.predict(X_test)}, index=df_test.index)\ny_infer = y_infer.sort_values(['id', 'rank']).reset_index('cell_id').groupby('id')['cell_id'].apply(list)\ny_infer","metadata":{"execution":{"iopub.status.busy":"2022-07-15T10:18:52.843742Z","iopub.execute_input":"2022-07-15T10:18:52.844082Z","iopub.status.idle":"2022-07-15T10:18:52.867703Z","shell.execute_reply.started":"2022-07-15T10:18:52.844052Z","shell.execute_reply":"2022-07-15T10:18:52.866733Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"y_sample = pd.read_csv(data_dir / 'sample_submission.csv', index_col='id', squeeze=True)\ny_sample","metadata":{"execution":{"iopub.status.busy":"2022-07-15T10:18:52.869187Z","iopub.execute_input":"2022-07-15T10:18:52.869844Z","iopub.status.idle":"2022-07-15T10:18:52.881653Z","shell.execute_reply.started":"2022-07-15T10:18:52.869796Z","shell.execute_reply":"2022-07-15T10:18:52.880702Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"y_submit = (\n    y_infer\n    .apply(' '.join)  # list of ids -> string of ids\n    .rename_axis('id')\n    .rename('cell_order')\n)\ny_submit","metadata":{"execution":{"iopub.status.busy":"2022-07-15T10:18:52.883458Z","iopub.execute_input":"2022-07-15T10:18:52.884237Z","iopub.status.idle":"2022-07-15T10:18:52.894015Z","shell.execute_reply.started":"2022-07-15T10:18:52.884195Z","shell.execute_reply":"2022-07-15T10:18:52.892837Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"y_submit.to_csv('submission.csv')","metadata":{"execution":{"iopub.status.busy":"2022-07-15T10:18:52.895957Z","iopub.execute_input":"2022-07-15T10:18:52.896709Z","iopub.status.idle":"2022-07-15T10:18:52.906041Z","shell.execute_reply.started":"2022-07-15T10:18:52.896661Z","shell.execute_reply":"2022-07-15T10:18:52.904790Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}