{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.11","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":84493,"databundleVersionId":11305158,"sourceType":"competition"},{"sourceId":10460392,"sourceType":"datasetVersion","datasetId":6468154},{"sourceId":10460393,"sourceType":"datasetVersion","datasetId":6468162},{"sourceId":215844687,"sourceType":"kernelVersion"}],"dockerImageVersionId":31041,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nos.system('pip install --force-reinstall /kaggle/input/janestreet2025-code/janestreet-0.1-py3-none-any.whl')\n\nimport time\nimport copy\n\nimport numpy as np\nimport pandas as pd\nimport polars as pl\nimport torch \n\nfrom kaggle_evaluation import jane_street_inference_server\n\nfrom janestreet.pipeline import FullPipeline, PipelineCV\nfrom janestreet.data_processor import DataProcessor\nfrom janestreet.config import PATH_DATA","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-28T11:19:02.011689Z","iopub.execute_input":"2025-05-28T11:19:02.012201Z","iopub.status.idle":"2025-05-28T11:19:05.009414Z","shell.execute_reply.started":"2025-05-28T11:19:02.012178Z","shell.execute_reply":"2025-05-28T11:19:05.008631Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"RUN_NAME = \"full\"\nMODEL_NAMES = [\"gru_2.0_700\", \"gru_2.1_700\", \"gru_2.2_700\", \"gru_3.0_700\", \"gru_3.1_700\", \"gru_3.2_700\"]\nWEIGHTS = np.array([1.0]*len(MODEL_NAMES))/ len(MODEL_NAMES)\nWEIGHTS = WEIGHTS/sum(WEIGHTS)\nN_ROLL = 1000\n\nprint(\"_\".join(MODEL_NAMES))\nprint(WEIGHTS)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-28T11:19:07.588870Z","iopub.execute_input":"2025-05-28T11:19:07.589528Z","iopub.status.idle":"2025-05-28T11:19:07.594288Z","shell.execute_reply.started":"2025-05-28T11:19:07.589504Z","shell.execute_reply":"2025-05-28T11:19:07.593408Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"data_processor = DataProcessor(MODEL_NAMES[0]).load()\n\npipelines = {}\nfor model_name in MODEL_NAMES:\n    pipeline = FullPipeline(\n        None,\n        run_name=RUN_NAME,\n        name=model_name,\n        load_model=True,\n        features=None,\n        save_to_disc=False\n    )\n    pipeline.fit(verbose=True)\n    pipelines[model_name] = pipeline\n\n    print(\"-\"*100)\n    print(model_name)\n    print(pipeline.model.get_params())\n    print(f\"Number of features: {len(pipeline.features)}\")\n    print(pipeline.model.model.num_resp)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-28T11:19:10.371139Z","iopub.execute_input":"2025-05-28T11:19:10.371865Z","iopub.status.idle":"2025-05-28T11:19:11.272832Z","shell.execute_reply.started":"2025-05-28T11:19:10.371842Z","shell.execute_reply":"2025-05-28T11:19:11.271955Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"MAX_DATE = 1698\nCOLS_ID = ['row_id', 'date_id', 'time_id', 'symbol_id', 'weight', 'is_scored']\n\ndf_raw = pl.scan_parquet(f\"{PATH_DATA}/train.parquet\")\ndf_raw = df_raw.filter(pl.col(\"date_id\")>=MAX_DATE-10)\ndf_raw = df_raw.collect()\ndf_raw = df_raw.with_columns(\n    pl.lit(-1).cast(pl.Int64).alias(\"row_id\"),\n    pl.lit(True).alias(\"is_scored\"),\n    (pl.col(\"date_id\")-MAX_DATE-1).alias(\"date_id\")\n)\ndf_raw = df_raw.select(COLS_ID + data_processor.COLS_FEATURES_INIT)\n\ndf_raw = (\n    df_raw.filter(pl.col(\"date_id\") >= -5)\n    .sort(['date_id', 'time_id', 'symbol_id'])\n)\n\nhidden_states = [None] * len(pipelines)\ndfs = []","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-28T11:19:17.627613Z","iopub.execute_input":"2025-05-28T11:19:17.628100Z","iopub.status.idle":"2025-05-28T11:19:17.924052Z","shell.execute_reply.started":"2025-05-28T11:19:17.628080Z","shell.execute_reply":"2025-05-28T11:19:17.923460Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import kaggle_evaluation.jane_street_inference_server as jsi\n\ndef fixed_run_local_gateway(self, data_paths=None, file_share_dir=None, *args, **kwargs):\n    self.server.start()\n    try:\n        # Call _get_gateway_for_test with ONLY data_paths to avoid argument error\n        self.gateway = self._get_gateway_for_test(data_paths)\n        self.gateway.run()\n    except Exception as err:\n        raise err from None\n    finally:\n        self.server.stop(0)\n\njsi.JSInferenceServer.run_local_gateway = fixed_run_local_gateway\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-28T11:19:22.575666Z","iopub.execute_input":"2025-05-28T11:19:22.576227Z","iopub.status.idle":"2025-05-28T11:19:22.580745Z","shell.execute_reply.started":"2025-05-28T11:19:22.576208Z","shell.execute_reply":"2025-05-28T11:19:22.579936Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"DEBUG = False\n\nCNT_DATES = 9\nCNT_DATES_NOT_SCORED = 4\n\ntime_start = time.time()\ntime_start_not_scored = time.time()\ntime_est = 0\ntime_est_not_scored = 0\ncnt_dates = 0\n\ndef predict(test: pl.DataFrame, lags: pl.DataFrame | None) -> pl.DataFrame | pd.DataFrame:\n    \"\"\"Make a prediction.\"\"\"\n    start_time = time.time()\n    \n    global df_raw\n    global hidden_states\n    global pipeline\n    global dfs\n    global time_est, time_start, time_start_not_scored, time_est_not_scored\n    global cnt_dates\n    \n    date_id = test[\"date_id\"][0]\n    time_id = test[\"time_id\"][0]\n    is_scored = test[\"is_scored\"][0]\n\n    # Count time for debug\n    if DEBUG:\n        if not is_scored:\n            time_est_not_scored = time.time()-time_start_not_scored\n        else:\n            time_est = time.time()-time_start\n\n        if time_id == 0:\n            print(\"-\" * 100)\n            \n            if date_id == 1:\n                time_start_not_scored = time.time()\n    \n            if date_id == CNT_DATES_NOT_SCORED: \n                time_start = time.time()\n\n    # Reset hidden states and collect data for weights update\n    if time_id == 0:\n        cnt_dates += 1\n        hidden_states = [None for _, p in pipelines.items()]\n        lags = lags.with_columns(\n            pl.col(\"responder_6_lag_1\").alias(\"responder_6\"),\n            pl.lit(date_id-1).cast(pl.Int16).alias(\"date_id\")\n        ).select([\"date_id\", \"time_id\", \"symbol_id\", \"responder_6\"])\n        if cnt_dates > 1:\n            df = pl.concat(dfs)\n            dfs = []\n            df = df.join(lags, on=[\"date_id\", \"time_id\", \"symbol_id\"], how=\"left\")\n            df = df.sort([\"date_id\", \"time_id\", \"symbol_id\"])\n\n    # Add data to raw dataframe\n    test = test.select(df_raw.columns)\n    df_raw = pl.concat([df_raw, test], how=\"vertical_relaxed\")\n    df_raw = df_raw.select(test.columns)\n    \n    # Cut raw data (keep last N_ROLL time_ids for each symbol)\n    df_raw = (\n        df_raw\n        .group_by([\"symbol_id\"])\n        .tail(N_ROLL)\n    )\n\n    # Calculate features and save\n    df_cur = data_processor.process_test_data(df_raw, fast=True, date_id=date_id, time_id=time_id, symbols=test[\"symbol_id\"])\n    df_cur = df_cur.sort([\"symbol_id\"])\n    dfs.append(df_cur)\n    df_cur = df_cur.with_columns(pl.lit(None).alias(\"responder_6\"))\n    \n    # Update model weights\n    if (time_id == 0) & (cnt_dates > 1):\n        if len(df) > 968:\n            for i, (name, pipeline) in enumerate(pipelines.items()):\n                pipeline.update(df)\n\n    # Make predictions\n    if is_scored:\n        preds = []\n        for i, (name, pipeline) in enumerate(pipelines.items()):\n            pred, hidden_states[i] = pipeline.predict(df_cur, hidden=hidden_states[i], n_times=1)\n            preds.append(pred)\n        pred = np.average(preds, axis=0, weights=WEIGHTS)\n    \n        df_cur = df_cur.with_columns(pl.Series(\"responder_6\", pred))\n        df_cur = test.select([\"date_id\", \"time_id\", \"symbol_id\"]).join(df_cur, on=[\"date_id\", \"time_id\", \"symbol_id\"], how=\"left\")\n        predictions = df_cur.select([\"row_id\", \"responder_6\"])\n    else:\n        predictions = test.select(\n            'row_id',\n            pl.lit(0.0).alias('responder_6'),\n        )\n\n    if DEBUG:\n        if time_id % 100 == 0:\n            n_nans = sum(sum(predictions.fill_nan(None).null_count().to_numpy()))\n            print(\n                f\"{date_id} {time_id:3.0f} (is_scored {is_scored}): \"\n                f\"time elps {time.time()-start_time:.4f}, # nans {n_nans}\"\n            )\n    else:\n        if (time_id==0)&(date_id==0):\n            print(predictions)\n            print((time_id, time.time()-start_time))\n    \n    return predictions\n\n    \ninference_server = jane_street_inference_server.JSInferenceServer(predict)\n\nif os.getenv('KAGGLE_IS_COMPETITION_RERUN'):\n    inference_server.serve()\nelse:\n    if not DEBUG:\n        inference_server.run_local_gateway(\n            [f'{PATH_DATA}/test.parquet', f'{PATH_DATA}/lags.parquet']\n        )\n\n    else:\n        inference_server.run_local_gateway(\n            [\n                '/kaggle/input/js24-rmf-submission-api-debug-with-synthetic-test/synthetic_test.parquet',\n                '/kaggle/input/js24-rmf-submission-api-debug-with-synthetic-test/synthetic_lag.parquet',\n            ]\n        )\n\n\nif DEBUG:\n    time_est_cur = time_est/(CNT_DATES-CNT_DATES_NOT_SCORED)*200/60/60\n    time_est_scored = time_est/(CNT_DATES-CNT_DATES_NOT_SCORED)*120/60/60\n    time_est_not_scored = time_est_not_scored/(CNT_DATES_NOT_SCORED-1)*240/60/60\n    time_est_final = time_est_scored + time_est_not_scored\n    print(\"-\"*100)\n    print(f\"Estimated current time: {time_est_cur:.4f}\")\n    print(f\"Estimated final time (is_score=True): {time_est_scored:.4f}\")\n    print(f\"Estimated final time (is_score=False): {time_est_not_scored:.4f}\")\n    print(f\"Estimated final time: {time_est_final:.4f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-28T11:19:25.890852Z","iopub.execute_input":"2025-05-28T11:19:25.891130Z","iopub.status.idle":"2025-05-28T11:19:25.966783Z","shell.execute_reply.started":"2025-05-28T11:19:25.891110Z","shell.execute_reply":"2025-05-28T11:19:25.965968Z"}},"outputs":[],"execution_count":null}]}