{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.14","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":84493,"databundleVersionId":9871156,"sourceType":"competition"}],"dockerImageVersionId":30786,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Introduction\n\nThis notebook aims to generate a synthetic test dataset to fully test the functionality of the API. The provided test.parquet contains only one batch of example data, which is not enough to cover all coner cases in the submission. By generating a synthetic test data using the train.parquet, we can better simulate the submission and identify potential bugs in the code. This is especially crucial for the time-series competition, as our code will go through a 3-month online testing without a chance to debug. \n\n**<span style=\"color:red\">- WARNING: DO NOT FORGET TO CHANGE THE SYNTHETIC TEST BACK TO THE REAL TEST BEFORE SUBMISSION !!! -</span>**","metadata":{}},{"cell_type":"markdown","source":"### Update to V3:\n\n**Added `is_scored=False` to the synthetic data**\n\nAdded several dates with `is_scored=False` in the synthetic data, allowing testing the inference pipeline using the `is_scored` flag. \n\nDuring the 6 months of forecasting phase, data provided for the public LB will be marked with `is_scored=False` and will not contribute to the score calculation. We can simply predict zeros for these part (~200 days) to save inference time.\n\n**Extended the date_offset**\n\nIn order to accommodate the non-scored dates, dates included in the dataset have been extended. The start date_id of the new version is from 1690, and only the last 5 days are marked with `is_scored=True`.\n\n**Improved history cache handling**\n\nThe previous cache handling is slow and can explode the memory.","metadata":{}},{"cell_type":"code","source":"import os\nimport polars as pl\nimport pandas as pd \nfrom pathlib import Path\nfrom tqdm import tqdm\n\nimport kaggle_evaluation.jane_street_inference_server as js_server\n\nDATA_DIR = Path('/kaggle/input/jane-street-real-time-market-data-forecasting')","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2025-01-02T19:37:40.188998Z","iopub.execute_input":"2025-01-02T19:37:40.189441Z","iopub.status.idle":"2025-01-02T19:37:42.022754Z","shell.execute_reply.started":"2025-01-02T19:37:40.189399Z","shell.execute_reply":"2025-01-02T19:37:42.021701Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"date_offset = 1690\n\nis_score_dates = 5\n\npl_all = pl.scan_parquet(DATA_DIR/\"train.parquet\").filter(pl.col(\"date_id\") >= date_offset-1).collect()","metadata":{"execution":{"iopub.status.busy":"2025-01-02T19:37:42.024774Z","iopub.execute_input":"2025-01-02T19:37:42.025259Z","iopub.status.idle":"2025-01-02T19:37:43.422977Z","shell.execute_reply.started":"2025-01-02T19:37:42.025220Z","shell.execute_reply":"2025-01-02T19:37:43.421667Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Make synthetic test dataset","metadata":{}},{"cell_type":"code","source":"# make syn_test \nsyn_test = pl_all.with_columns(\n    # pl.lit(True).alias(\"is_scored\"),\n    pl.col('date_id') - date_offset\n    ).with_row_index(name=\"row_id\", offset=0)\n\nsyn_test = syn_test.with_columns(\n    pl.when(pl.col('date_id')<is_score_dates-1).then(pl.lit(False)).otherwise(pl.lit(True)).alias(\"is_scored\")\n)\n\nsyn_test = syn_test.select(\n    ['row_id', 'date_id', 'time_id', 'symbol_id', 'weight', 'is_scored'] + [f'feature_{x:02}' for x in range(79)]\n)\n\nsyn_test_partition = syn_test.partition_by('date_id', maintain_order=True, as_dict=True)\n\noutput_dir = \"synthetic_test.parquet\"\nos.makedirs(output_dir, exist_ok=True)\n\nrow_id_offset = syn_test.filter(pl.col('date_id')<0).select('row_id').max().item()\nprint(\"row_id_offset:\", row_id_offset)\n\nfor key, _df in syn_test_partition.items():\n    if key[0] >= 0:\n        os.makedirs(f\"{output_dir}/date_id={key[0]}\", exist_ok=True)\n        _df = _df.with_columns(pl.col('row_id')-row_id_offset)\n        _df.write_parquet(f\"{output_dir}/date_id={key[0]}/part-0.parquet\")","metadata":{"execution":{"iopub.status.busy":"2025-01-02T19:37:43.424537Z","iopub.execute_input":"2025-01-02T19:37:43.424976Z","iopub.status.idle":"2025-01-02T19:37:44.287293Z","shell.execute_reply.started":"2025-01-02T19:37:43.424920Z","shell.execute_reply":"2025-01-02T19:37:44.285634Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"syn_test","metadata":{"execution":{"iopub.status.busy":"2025-01-02T19:37:44.289600Z","iopub.execute_input":"2025-01-02T19:37:44.289957Z","iopub.status.idle":"2025-01-02T19:37:44.315392Z","shell.execute_reply.started":"2025-01-02T19:37:44.289924Z","shell.execute_reply":"2025-01-02T19:37:44.314199Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# check is_score dates --> only last 6 days are scored\nsyn_test.group_by('date_id').agg(pl.col('is_scored').all()).sort('date_id').to_pandas()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-02T19:37:45.534275Z","iopub.execute_input":"2025-01-02T19:37:45.534713Z","iopub.status.idle":"2025-01-02T19:37:45.742722Z","shell.execute_reply.started":"2025-01-02T19:37:45.534676Z","shell.execute_reply":"2025-01-02T19:37:45.741409Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# List number of symbols per date\nsyn_test.group_by('date_id').agg(pl.col('symbol_id').n_unique()).sort('date_id').to_pandas()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-02T19:37:56.704666Z","iopub.execute_input":"2025-01-02T19:37:56.705851Z","iopub.status.idle":"2025-01-02T19:37:56.768519Z","shell.execute_reply.started":"2025-01-02T19:37:56.705796Z","shell.execute_reply":"2025-01-02T19:37:56.767304Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"pl_test = pl.read_parquet(DATA_DIR/\"test.parquet\", n_rows=10000)\npl_test","metadata":{"execution":{"iopub.status.busy":"2025-01-02T19:38:06.371284Z","iopub.execute_input":"2025-01-02T19:38:06.371710Z","iopub.status.idle":"2025-01-02T19:38:06.402497Z","shell.execute_reply.started":"2025-01-02T19:38:06.371673Z","shell.execute_reply":"2025-01-02T19:38:06.401170Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Make synthetic lags","metadata":{}},{"cell_type":"code","source":"# make syn_lag\n\nsyn_lag = pl_all.select(\n    ['date_id', 'time_id', 'symbol_id'] + [f'responder_{x}' for x in range(9)]\n).with_columns(pl.col('date_id')-date_offset)\n\nsyn_lag = syn_lag.rename({f'responder_{x}': f'responder_{x}_lag_1' for x in range(9)})\n\nsyn_lag_partition = syn_lag.partition_by('date_id', maintain_order=True, as_dict=True)\n\noutput_dir = \"synthetic_lag.parquet\"\nos.makedirs(output_dir, exist_ok=True)\n\nfor key, _df in syn_lag_partition.items():\n    os.makedirs(f\"{output_dir}/date_id={key[0]+1}\", exist_ok=True)\n    _df = _df.with_columns(pl.col('date_id')+1)\n    _df.write_parquet(f\"{output_dir}/date_id={key[0]+1}/part-0.parquet\")","metadata":{"execution":{"iopub.status.busy":"2025-01-02T19:38:12.621827Z","iopub.execute_input":"2025-01-02T19:38:12.623054Z","iopub.status.idle":"2025-01-02T19:38:12.767011Z","shell.execute_reply.started":"2025-01-02T19:38:12.622987Z","shell.execute_reply":"2025-01-02T19:38:12.765917Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"pl_lag = pl.read_parquet(DATA_DIR / 'lags.parquet')\npl_lag","metadata":{"execution":{"iopub.status.busy":"2025-01-02T19:38:12.936381Z","iopub.execute_input":"2025-01-02T19:38:12.937460Z","iopub.status.idle":"2025-01-02T19:38:12.958229Z","shell.execute_reply.started":"2025-01-02T19:38:12.937419Z","shell.execute_reply":"2025-01-02T19:38:12.957105Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Test submission using the synthetic test & lags","metadata":{}},{"cell_type":"markdown","source":"The model function is only a dummy function. The user should modify it based on their real model (e.g. some models requires .predict() method).","metadata":{}},{"cell_type":"code","source":"from collections import defaultdict\nimport numpy as np\n\ndef model(x):\n    return np.nanmean(x, axis=(-1,-2))","metadata":{"execution":{"iopub.status.busy":"2025-01-02T19:38:16.383977Z","iopub.execute_input":"2025-01-02T19:38:16.385535Z","iopub.status.idle":"2025-01-02T19:38:16.393477Z","shell.execute_reply.started":"2025-01-02T19:38:16.385468Z","shell.execute_reply":"2025-01-02T19:38:16.392188Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"lags_ : pl.DataFrame | None = None\ntest_ = None\n\nlags_cache = pl.DataFrame()\nhistory_cache = {}\nlook_back = 100\nfeature_cols = [f'feature_{x:02}' for x in range(79)]\n\ndatetime_steps = pl.scan_parquet('/kaggle/working/synthetic_test.parquet').select(\n    (pl.col(\"date_id\")*10000+pl.col('time_id')).n_unique()   \n    ).collect().item()\n\npbar = tqdm(total=datetime_steps)\npbar.refresh()\n\ndef predict(test: pl.DataFrame, lags: pl.DataFrame | None) -> pl.DataFrame | pd.DataFrame:\n    '''\n    All the responders from the previous day are passed in at time_id == 0. \n    We save them in a global variable for access at every time_id. \n    Use them as extra features, if you like.\n    Each batch of predictions (except the very first) must be returned within 10 minutes of the batch features being provided.\n    '''\n    \n    global lags_ , history_cache, pbar, test_\n    \n    if lags is not None:\n        lags_ = lags\n        \n    test_ = test\n\n    # only run inference if is_scored==True\n    if test['is_scored'].any():\n        try:\n            data_dict = defaultdict(list)\n            for (symbol_id,), batch in test.group_by('symbol_id', maintain_order=True):\n                \n                if symbol_id in history_cache.keys():\n                    history_cache[symbol_id] = np.concatenate(\n                        (history_cache[symbol_id], batch[feature_cols].to_numpy()), axis=0)\n                else:\n                    history_cache[symbol_id] = batch[feature_cols].to_numpy()\n\n                if len(history_cache[symbol_id]) > look_back:\n                    history_cache[symbol_id] = history_cache[symbol_id][-look_back:]\n        \n                x = history_cache[symbol_id][-look_back:]\n                data_dict['x'].append(x)\n                data_dict['symbol_id'].append(symbol_id)\n                data_dict['row_id'].append(batch['row_id'][-1])\n    \n            x_stack = np.stack(data_dict['x']) #(n_symbol, length, dim=79)\n    \n            predictions = pl.DataFrame({\n                'row_id': data_dict['row_id'], \n                'responder_6': model(x_stack)\n                })\n            \n        except Exception as e:\n            print(f\"An error occurred: {e}\")\n    else:\n        predictions = pl.DataFrame({\n            'row_id': test['row_id'], \n            'responder_6': 0\n            })\n\n    # The predict function must return a DataFrame\n    assert isinstance(predictions, pl.DataFrame | pd.DataFrame)\n    # with columns 'row_id', 'responer_6'\n    assert list(predictions.columns) == ['row_id', 'responder_6']\n    # and as many rows as the test data.\n    assert len(predictions) == len(test)\n    \n    pbar.update(1)\n\n    return predictions","metadata":{"execution":{"iopub.status.busy":"2025-01-02T19:38:16.610979Z","iopub.execute_input":"2025-01-02T19:38:16.611509Z","iopub.status.idle":"2025-01-02T19:38:16.650572Z","shell.execute_reply.started":"2025-01-02T19:38:16.611460Z","shell.execute_reply":"2025-01-02T19:38:16.649270Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"inference_server = js_server.JSInferenceServer(predict)\n\nif os.getenv('KAGGLE_IS_COMPETITION_RERUN'):\n    inference_server.serve()\nelse:\n    inference_server.run_local_gateway(\n        (\n            '/kaggle/working/synthetic_test.parquet',\n            '/kaggle/working/synthetic_lag.parquet',\n        )\n    )\n    \npbar.close()","metadata":{"execution":{"iopub.status.busy":"2025-01-02T19:38:20.928111Z","iopub.execute_input":"2025-01-02T19:38:20.928562Z","iopub.status.idle":"2025-01-02T19:41:45.854526Z","shell.execute_reply.started":"2025-01-02T19:38:20.928522Z","shell.execute_reply":"2025-01-02T19:41:45.853287Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if os.path.isfile('submission.parquet'):\n    pl_sub = pl.read_parquet('submission.parquet')\n    display(pl_sub)","metadata":{"execution":{"iopub.status.busy":"2025-01-02T19:41:45.856889Z","iopub.execute_input":"2025-01-02T19:41:45.857392Z","iopub.status.idle":"2025-01-02T19:41:45.877448Z","shell.execute_reply.started":"2025-01-02T19:41:45.857321Z","shell.execute_reply":"2025-01-02T19:41:45.876386Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"history_cache","metadata":{"execution":{"iopub.status.busy":"2025-01-02T19:41:45.878872Z","iopub.execute_input":"2025-01-02T19:41:45.879210Z","iopub.status.idle":"2025-01-02T19:41:45.910730Z","shell.execute_reply.started":"2025-01-02T19:41:45.879176Z","shell.execute_reply":"2025-01-02T19:41:45.909598Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"lags_","metadata":{"execution":{"iopub.status.busy":"2025-01-02T19:41:45.912822Z","iopub.execute_input":"2025-01-02T19:41:45.913219Z","iopub.status.idle":"2025-01-02T19:41:45.927893Z","shell.execute_reply.started":"2025-01-02T19:41:45.913182Z","shell.execute_reply":"2025-01-02T19:41:45.926693Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"test_","metadata":{"execution":{"iopub.status.busy":"2025-01-02T19:41:45.929257Z","iopub.execute_input":"2025-01-02T19:41:45.929721Z","iopub.status.idle":"2025-01-02T19:41:45.951534Z","shell.execute_reply.started":"2025-01-02T19:41:45.929683Z","shell.execute_reply":"2025-01-02T19:41:45.950378Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}