{"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"},{"sourceId":9999437,"sourceType":"datasetVersion","datasetId":6154742},{"sourceId":10231695,"sourceType":"datasetVersion","datasetId":6326342}],"dockerImageVersionId":30786,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import joblib\nimport polars as pl\nimport pandas as pd\nimport numpy as np\n#environment provided by competition hoster\nimport os \nimport kaggle_evaluation.jane_street_inference_server","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2024-12-18T03:39:28.581774Z","iopub.execute_input":"2024-12-18T03:39:28.582193Z","iopub.status.idle":"2024-12-18T03:39:28.588092Z","shell.execute_reply.started":"2024-12-18T03:39:28.582156Z","shell.execute_reply":"2024-12-18T03:39:28.586672Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class CONFIG:\n    seed = 42\n    target_col = \"responder_6\"\n    # feature_cols = [\"symbol_id\", \"time_id\"] + [f\"feature_{idx:02d}\" for idx in range(79)]+ [f\"responder_{idx}_lag_1\" for idx in range(9)]\n    feature_cols = [f\"feature_{idx:02d}\" for idx in range(79)]+ [f\"responder_{idx}_lag_1\" for idx in range(9)]\n    \n    model_paths = [\n        #\"/kaggle/input/js24-train-gbdt-model-with-lags-singlemodel/result.pkl\",\n        #\"/kaggle/input/js24-trained-gbdt-model/result.pkl\",\n        \"/kaggle/input/js-xs-nn-trained-model\",\n        \"/kaggle/input/js-with-lags-trained-xgb/result.pkl\",\n    ]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-18T03:39:30.746414Z","iopub.execute_input":"2024-12-18T03:39:30.746871Z","iopub.status.idle":"2024-12-18T03:39:30.753408Z","shell.execute_reply.started":"2024-12-18T03:39:30.746823Z","shell.execute_reply":"2024-12-18T03:39:30.752080Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 加载模型\nmodel = joblib.load('/kaggle/input/js-ridge-model-exp2/ridge_model2.pkl')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-18T03:39:34.577385Z","iopub.execute_input":"2024-12-18T03:39:34.577817Z","iopub.status.idle":"2024-12-18T03:39:34.588031Z","shell.execute_reply.started":"2024-12-18T03:39:34.577777Z","shell.execute_reply":"2024-12-18T03:39:34.586721Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"lags_ : pl.DataFrame | None = None\n    \ndef predict(test: pl.DataFrame, lags: pl.DataFrame | None) -> pl.DataFrame | pd.DataFrame:\n    global lags_\n    if lags is not None:\n        lags_ = lags\n\n    predictions = test.select(\n        'row_id',\n        pl.lit(0.0).alias('responder_6'),\n    )\n    symbol_ids = test.select('symbol_id').to_numpy()[:, 0]\n\n    if not lags is None:\n        lags = lags.group_by([\"date_id\", \"symbol_id\"], maintain_order=True).last() # pick up last record of previous date\n        test = test.join(lags, on=[\"date_id\", \"symbol_id\"],  how=\"left\")\n    else:\n        test = test.with_columns(\n            ( pl.lit(0.0).alias(f'responder_{idx}_lag_1') for idx in range(9) )\n        )\n    \n    preds = np.zeros((test.shape[0],))\n    test_input = test[CONFIG.feature_cols].to_pandas()\n    test_input = test_input.fillna(method = 'ffill').fillna(0)\n    preds += model.predict(test_input)\n    print(f\"predict> preds.shape =\", preds.shape)\n    \n    predictions = \\\n    test.select('row_id').\\\n    with_columns(\n        pl.Series(\n            name   = 'responder_6', \n            values = np.clip(preds, a_min = -5, a_max = 5),\n            dtype  = pl.Float64,\n        )\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    return predictions","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-18T03:41:26.049960Z","iopub.execute_input":"2024-12-18T03:41:26.050418Z","iopub.status.idle":"2024-12-18T03:41:26.061122Z","shell.execute_reply.started":"2024-12-18T03:41:26.050380Z","shell.execute_reply":"2024-12-18T03:41:26.059892Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"inference_server = kaggle_evaluation.jane_street_inference_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/input/jane-street-real-time-market-data-forecasting/test.parquet',\n            '/kaggle/input/jane-street-real-time-market-data-forecasting/lags.parquet',\n        )\n    )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-18T03:41:29.860821Z","iopub.execute_input":"2024-12-18T03:41:29.861249Z","iopub.status.idle":"2024-12-18T03:41:29.960427Z","shell.execute_reply.started":"2024-12-18T03:41:29.861209Z","shell.execute_reply":"2024-12-18T03:41:29.959160Z"}},"outputs":[],"execution_count":null}]}