{"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":"Calculating the Cross Validation score quickly using Polars. For more info on local cross-validation see [Radek's notebook](https://www.kaggle.com/code/cabbage972/a-robust-local-validation-framework/edit).\n\nPlease note that the AIDs of each candidate are assumed to be deduplicated. Otherwise, the deduplication would slow down the calculation significantly.","metadata":{}},{"cell_type":"code","source":"! pip install -q polars","metadata":{"execution":{"iopub.status.busy":"2022-12-25T18:13:01.630826Z","iopub.execute_input":"2022-12-25T18:13:01.631363Z","iopub.status.idle":"2022-12-25T18:13:16.339358Z","shell.execute_reply.started":"2022-12-25T18:13:01.631317Z","shell.execute_reply":"2022-12-25T18:13:16.337821Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Imports","metadata":{}},{"cell_type":"code","source":"import pandas as pd\nimport polars as pl\nimport os","metadata":{"execution":{"iopub.status.busy":"2022-12-25T18:13:25.321997Z","iopub.execute_input":"2022-12-25T18:13:25.322466Z","iopub.status.idle":"2022-12-25T18:13:25.398994Z","shell.execute_reply.started":"2022-12-25T18:13:25.322425Z","shell.execute_reply":"2022-12-25T18:13:25.397717Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Implementation of the scoring function","metadata":{}},{"cell_type":"code","source":"def fast_scores(submission, ground_truth, num_cands=20):\n    submission = submission.with_columns([\n            pl.col('session_type').str.split('_').arr.first().alias('session'),\n            pl.col('session_type').str.split('_').arr.last().alias('type'),\n            pl.col('labels').arr.head(num_cands).alias('labels2')\n    ]).select([pl.col('session'), pl.col('type'), pl.col('labels2').alias('labels')])\n        \n    ground_truth = ground_truth.with_columns([\n            pl.col('session_type').str.split('_').arr.first().alias('session'),\n            pl.col('session_type').str.split('_').arr.last().alias('type'),\n            pl.col('labels').cast(pl.List(int)).alias('labels2')]\n    ).select([pl.col('session'), pl.col('type'), pl.col('labels2').alias('labels')])\n\n    submission_with_gt = submission.join(ground_truth, how='left', on=['session', 'type'])\n    submission_with_gt = submission_with_gt.drop_nulls('labels_right')\n    submission_with_gt = submission_with_gt.with_columns([\n        pl.col(\"labels\")\n        .arr.concat(pl.col('labels_right').arr.head(20))\n        .arr.eval(pl.element().filter(pl.count().over(pl.element()) == 2))\n        .arr.unique()\n        .alias('hits').apply(len),\n        pl.col('labels_right').apply(len).alias('gt_count')\n    ])\n     \n    agg_df = submission_with_gt.groupby('type').agg([pl.col('hits').sum(), pl.col('gt_count').sum()]).collect().to_pandas()\n    agg_df = agg_df.set_index('type')\n    agg_df['recall'] = agg_df['hits'] / agg_df['gt_count']\n    agg_df = agg_df.sort_values('type')\n    return agg_df","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-12-25T15:22:51.709217Z","iopub.execute_input":"2022-12-25T15:22:51.709709Z","iopub.status.idle":"2022-12-25T15:22:51.726002Z","shell.execute_reply.started":"2022-12-25T15:22:51.709672Z","shell.execute_reply":"2022-12-25T15:22:51.724515Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Load the labels file","metadata":{}},{"cell_type":"code","source":"label_df = pl.scan_parquet('/kaggle/input/otto-val-labels/label.parquet')\nlabel_df.schema","metadata":{"execution":{"iopub.status.busy":"2022-12-25T18:14:10.124046Z","iopub.execute_input":"2022-12-25T18:14:10.124476Z","iopub.status.idle":"2022-12-25T18:14:10.134021Z","shell.execute_reply.started":"2022-12-25T18:14:10.124442Z","shell.execute_reply":"2022-12-25T18:14:10.133194Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Load the submission file","metadata":{}},{"cell_type":"code","source":"sub_df = pl.scan_parquet(os.path.join('/kaggle/input/otto-val-preds-example', '*.parquet'))\nsub_df.schema","metadata":{"execution":{"iopub.status.busy":"2022-12-25T18:13:59.795973Z","iopub.execute_input":"2022-12-25T18:13:59.796426Z","iopub.status.idle":"2022-12-25T18:13:59.805713Z","shell.execute_reply.started":"2022-12-25T18:13:59.796390Z","shell.execute_reply":"2022-12-25T18:13:59.804915Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Calculate the score","metadata":{}},{"cell_type":"code","source":"%%time\nagg_df = fast_scores(sub_df, label_df, num_cands=20)\n\nprint(agg_df['recall'])\nfinal_score = (agg_df['recall'] * pd.Series({'clicks': 0.10, 'carts': 0.30, 'orders': 0.60})).sum()\nprint('final score: {:.4f}'.format(final_score))","metadata":{"execution":{"iopub.status.busy":"2022-12-25T15:22:53.208390Z","iopub.execute_input":"2022-12-25T15:22:53.208858Z","iopub.status.idle":"2022-12-25T15:23:04.774391Z","shell.execute_reply.started":"2022-12-25T15:22:53.208823Z","shell.execute_reply":"2022-12-25T15:23:04.772882Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}