{"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":"gpu","dataSources":[{"sourceId":84493,"databundleVersionId":9871156,"sourceType":"competition"},{"sourceId":201579302,"sourceType":"kernelVersion"},{"sourceId":202331107,"sourceType":"kernelVersion"},{"sourceId":143573,"sourceType":"modelInstanceVersion","modelInstanceId":117345,"modelId":140574}],"dockerImageVersionId":30786,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Import package","metadata":{}},{"cell_type":"code","source":"import os\nimport gc\nimport random \nfrom pathlib import Path\nfrom tqdm import tqdm\n\nfrom collections import OrderedDict, defaultdict\n\nimport polars as pl\nimport pandas as pd \nimport numpy as np\nfrom matplotlib import pyplot as plt\nimport seaborn as sns\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\n\nimport kaggle_evaluation.jane_street_inference_server as js_server","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-10-22T20:43:14.484812Z","iopub.execute_input":"2024-10-22T20:43:14.485769Z","iopub.status.idle":"2024-10-22T20:43:17.601759Z","shell.execute_reply.started":"2024-10-22T20:43:14.485717Z","shell.execute_reply":"2024-10-22T20:43:17.600932Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Global constants\n\nDATA_DIR = Path('/kaggle/input/jane-street-real-time-market-data-forecasting')\n\nMETA_COLS = [\"date_id\", \"time_id\", 'symbol_id', 'weight']\nFEATURE_COLS = [f'feature_{x:02}' for x in range(79)]\nRESPONDER_COLS = [f'responder_{i}' for i in range(9)]\n\nSEQUENCE_LEN = 16\n\nRANDOM_SEED = 2024","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-10-22T20:43:17.602942Z","iopub.execute_input":"2024-10-22T20:43:17.603396Z","iopub.status.idle":"2024-10-22T20:43:17.609181Z","shell.execute_reply.started":"2024-10-22T20:43:17.603360Z","shell.execute_reply":"2024-10-22T20:43:17.608112Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def seed_everything(seed):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    if torch.cuda.is_available():\n        torch.cuda.manual_seed(seed)\n        torch.cuda.manual_seed_all(seed)\n        # Ensure deterministic behavior (may impact performance)\n        torch.backends.cudnn.deterministic = True\n        torch.backends.cudnn.benchmark = False\n\ndef lazy_load(par_path):\n    return pl.scan_parquet(par_path).select(\n        pl.int_range(pl.len(), dtype=pl.UInt64).alias(\"index\"),\n        pl.all()\n    )\n\nseed_everything(RANDOM_SEED)","metadata":{"execution":{"iopub.status.busy":"2024-10-22T20:43:17.610277Z","iopub.execute_input":"2024-10-22T20:43:17.610660Z","iopub.status.idle":"2024-10-22T20:43:17.652422Z","shell.execute_reply.started":"2024-10-22T20:43:17.610618Z","shell.execute_reply":"2024-10-22T20:43:17.651662Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Use synthetic test or real test?","metadata":{}},{"cell_type":"code","source":"USE_STYNTHETIC = False\n\nif USE_STYNTHETIC:\n    syn_dir = '/kaggle/input/js24-rmf-generate-synthetic-test-data'\n    test_parquet = f'{syn_dir}/synthetic_test.parquet'\n    lag_parquet = f'{syn_dir}/synthetic_lag.parquet'\n    total_time_steps = pl.scan_parquet(test_parquet).select(\n        (pl.col(\"date_id\")*10000+pl.col('time_id')).n_unique()   \n        ).collect().item()\nelse:\n    test_parquet = DATA_DIR / f'test.parquet'\n    lag_parquet =  DATA_DIR / f'lags.parquet'\n    total_time_steps = 1\n    \nprint(\"Test parquet:\", test_parquet)\nprint(\"Lag parquet:\", lag_parquet)\nprint(\"Total time steps:\", total_time_steps)","metadata":{"execution":{"iopub.status.busy":"2024-10-22T20:43:17.653539Z","iopub.execute_input":"2024-10-22T20:43:17.653825Z","iopub.status.idle":"2024-10-22T20:43:17.660753Z","shell.execute_reply.started":"2024-10-22T20:43:17.653794Z","shell.execute_reply":"2024-10-22T20:43:17.659752Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Load Model","metadata":{}},{"cell_type":"code","source":"from gru_encoder_model import GRURegressor\n\n# MODEL_CONFIG = {\n#     'input_dim':  79,\n#     'output_dim': 1,\n#     'head_hidden_dim': 256, \n#     'emb_dim':  256, \n#     'num_blocks': 8,\n#     'n_gru_layers': 2,\n#     'ff_hidden_dim': 512,\n#     'dropout': 0, \n#     'bidirectional': True,\n# }\n\nMODEL_CONFIG = {\n    'input_dim':  79,\n    'output_dim': 1,\n    'head_hidden_dim': 256, \n    'emb_dim':  256, \n    'num_blocks': 4,\n    'n_gru_layers': 2,\n    'ff_hidden_dim': 512,\n    'dropout': 0, \n    'bidirectional': True,\n}\n\nDEVICE = 'cuda' if torch.cuda.is_available() else 'cpu'\nprint(\"Use Device:\", DEVICE)\n\n# local_cv -0.05\nMODEL_PATH = \"/kaggle/input/jane-street-2024-market-forecasting-models/pytorch/gru/11/epoch46-val_metric_epoch-0.05-state_dict.ckpt\"\nprint(\"Use state_dict:\", MODEL_PATH)","metadata":{"execution":{"iopub.status.busy":"2024-10-22T20:43:17.663695Z","iopub.execute_input":"2024-10-22T20:43:17.663982Z","iopub.status.idle":"2024-10-22T20:43:17.673413Z","shell.execute_reply.started":"2024-10-22T20:43:17.663951Z","shell.execute_reply":"2024-10-22T20:43:17.672468Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model = GRURegressor(**MODEL_CONFIG)\n\nstate_dict = torch.load(MODEL_PATH, map_location=DEVICE, weights_only=False)\n\nnative_state_dict = OrderedDict()\nfor k, v in state_dict['state_dict'].items():\n    if k.startswith('model.'):\n        new_key = k.replace('model.', '')  # remove 'model.' prefix\n    else:\n        new_key = k\n    native_state_dict[new_key] = v\n\nmodel.load_state_dict(native_state_dict)\nmodel.eval() \nmodel.to(DEVICE)\n\n# try a simple test\nx_random = torch.rand(10, 100, 79).float().to(DEVICE)\ny_rand_pred = model(x_random)\nprint(y_rand_pred.shape)","metadata":{"execution":{"iopub.status.busy":"2024-10-22T20:43:17.674499Z","iopub.execute_input":"2024-10-22T20:43:17.675026Z","iopub.status.idle":"2024-10-22T20:43:18.514052Z","shell.execute_reply.started":"2024-10-22T20:43:17.674983Z","shell.execute_reply":"2024-10-22T20:43:18.513074Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Run Inference","metadata":{}},{"cell_type":"code","source":"def predict_func(x, model):\n    with torch.no_grad():\n        x = torch.tensor(x, dtype=torch.float).to(DEVICE)\n        predict = model(x)\n    return predict.cpu().detach().numpy()\n    # return np.zeros(len(x))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-10-22T20:43:18.515497Z","iopub.execute_input":"2024-10-22T20:43:18.516192Z","iopub.status.idle":"2024-10-22T20:43:18.521479Z","shell.execute_reply.started":"2024-10-22T20:43:18.516142Z","shell.execute_reply":"2024-10-22T20:43:18.520491Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class JaneStreetPredictor:\n    \n    def __init__(self, model, test_parquet, lag_parquet, sequence_len, feature_cols, pbar_length=0):\n        # Initialize model and parameters\n        self.model = model\n        self.sequence_len = sequence_len\n        self.feature_cols = feature_cols\n        \n        # Initialize parquet data for test and lag\n        self.test_parquet = test_parquet\n        self.lag_parquet = lag_parquet\n        \n        # Initialize global variables as class attributes\n        self.history_cache = {}\n        self.test_ = None\n        self.lags_ = None\n        self.time_step_count = 0\n\n        # setup pbar:\n        self.pbar = tqdm(total=pbar_length, disable=(pbar_length == 0))\n\n    def run_inference_server(self):\n\n        self.pbar.refresh()\n\n        inference_server = js_server.JSInferenceServer(self.predict)\n        \n        if os.getenv('KAGGLE_IS_COMPETITION_RERUN'):\n            inference_server.serve()\n        else:\n            inference_server.run_local_gateway((self.test_parquet, self.lag_parquet))\n\n        self.pbar.close()\n\n    def predict(self, test: pl.DataFrame, lags: pl.DataFrame | None) -> pl.DataFrame | pd.DataFrame:\n\n        # Handle lags if provided\n        if lags is not None:\n            self.lags_ = lags\n\n        # Cache the test data for future use\n        self.test_ = test\n\n        # Get the number of unique symbols (assets or stocks)\n        n_symbols = test['symbol_id'].n_unique()\n\n        # Initialize arrays for feature history and row ids\n        x_stack = np.zeros((n_symbols, self.sequence_len, len(self.feature_cols)))\n        row_ids = np.zeros(len(test))\n\n        # Partition the test data by 'symbol_id'\n        test_partition = test.partition_by('symbol_id', as_dict=True)\n\n        # Process each symbol's time-series data\n        for i, ((symbol_id, ), partition) in enumerate(test_partition.items()):\n            # Update or initialize the history cache for each symbol\n            if symbol_id in self.history_cache:\n                self.history_cache[symbol_id] = np.concatenate(\n                    (self.history_cache[symbol_id], partition[self.feature_cols].to_numpy()),\n                    axis=0\n                )\n            else:\n                self.history_cache[symbol_id] = partition[self.feature_cols].to_numpy()\n\n            # Trim the history cache to the most recent SEQUENCE_LEN steps\n            if len(self.history_cache[symbol_id]) > self.sequence_len + 1:\n                self.history_cache[symbol_id] = self.history_cache[symbol_id][-self.sequence_len-1:]\n            \n            # Extract the most recent feature history for the symbol\n            x = self.history_cache[symbol_id][-self.sequence_len:, :]\n            x_stack[i, -len(x):, :] = x\n            row_ids[i] = partition['row_id'].item()\n\n        # Fill NaN values in the feature array\n        x_stack = np.nan_to_num(x_stack, nan=0.0)\n\n        # Call the model's prediction function\n        predicts = predict_func(x_stack, self.model)  # Placeholder for the actual model prediction logic\n\n        # Create a DataFrame to store the predictions\n        predictions = pl.DataFrame({\n            'row_id': row_ids,\n            'responder_6': predicts\n        }).with_columns(pl.col('row_id').cast(test['row_id'].dtype))\n\n        # Join the predictions back to the test data for correct alignment\n        predictions = test.select('row_id').join(predictions, on='row_id', how='left').select('row_id', 'responder_6')\n        \n        # sanity check\n        assert isinstance(predictions, pl.DataFrame | pd.DataFrame)\n        assert list(predictions.columns) == ['row_id', 'responder_6']\n        assert len(predictions) == len(test)\n\n        # update time_step_count\n        self.time_step_count += 1\n        self.pbar.update(1)\n\n        return predictions","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-10-22T20:43:18.522833Z","iopub.execute_input":"2024-10-22T20:43:18.523203Z","iopub.status.idle":"2024-10-22T20:43:18.541829Z","shell.execute_reply.started":"2024-10-22T20:43:18.523161Z","shell.execute_reply":"2024-10-22T20:43:18.540944Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"js_predictor = JaneStreetPredictor(\n    model, \n    test_parquet, \n    lag_parquet, \n    sequence_len = SEQUENCE_LEN, \n    feature_cols = FEATURE_COLS, \n    pbar_length = total_time_steps\n)\n\njs_predictor.run_inference_server()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-10-22T20:43:18.543082Z","iopub.execute_input":"2024-10-22T20:43:18.544055Z","iopub.status.idle":"2024-10-22T20:43:18.651337Z","shell.execute_reply.started":"2024-10-22T20:43:18.544014Z","shell.execute_reply":"2024-10-22T20:43:18.650337Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if os.path.isfile('submission.parquet'):\n    pl_sub = pl.read_parquet('submission.parquet')\n    print(len(pl_sub))\n    display(pl_sub)","metadata":{"execution":{"iopub.status.busy":"2024-10-22T20:43:18.652783Z","iopub.execute_input":"2024-10-22T20:43:18.653206Z","iopub.status.idle":"2024-10-22T20:43:18.661802Z","shell.execute_reply.started":"2024-10-22T20:43:18.653148Z","shell.execute_reply":"2024-10-22T20:43:18.660998Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"","metadata":{}},{"cell_type":"code","source":"if USE_STYNTHETIC:\n    pl_test = pl.read_parquet(test_parquet).filter(pl.col('date_id')<3)\n    symbol_choice = np.random.choice(pl_test['symbol_id'].unique(), 10)\n\n    pl_lags = pl.read_parquet(lag_parquet).with_columns(pl.col('date_id')-1).filter(pl.col('date_id')>0)\n\n    fig, axes = plt.subplots(5, 2, figsize=(20, 20))\n\n    for ax, symbol_id in zip(axes.flatten(), symbol_choice):\n        pl_symbol = pl_test.filter(\n            pl.col('symbol_id')==symbol_id\n        ).sort(['row_id']).select(['row_id', 'date_id', 'time_id', 'symbol_id'])\n\n        pl_merge = pl_symbol.join(\n            pl_lags.select(['date_id', 'time_id', 'symbol_id', 'responder_6_lag_1']), \n            on=(['date_id', 'time_id', 'symbol_id']), \n            how='left'\n            ).join(pl_sub, on='row_id', how='left').drop_nulls()\n\n        ax.plot(pl_merge['responder_6_lag_1'][:200], 'b-', label='True')\n        ax.plot(pl_merge['responder_6'][:200], 'r-', label='Pred')\n        ax.set_title(f\"Symbol {symbol_id}\")\n        ax.legend()\n        ax.grid(True, ls=\"--\")\n\n    fig.tight_layout()\n    plt.show()        ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-10-22T20:43:18.662923Z","iopub.execute_input":"2024-10-22T20:43:18.663516Z","iopub.status.idle":"2024-10-22T20:43:18.673621Z","shell.execute_reply.started":"2024-10-22T20:43:18.663478Z","shell.execute_reply":"2024-10-22T20:43:18.672725Z"}},"outputs":[],"execution_count":null}]}