{"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":[{"sourceType":"competition","sourceId":84493,"databundleVersionId":9871156},{"sourceType":"datasetVersion","sourceId":9761005,"datasetId":5931288,"databundleVersionId":10001369},{"sourceType":"datasetVersion","sourceId":9870312,"datasetId":5982075,"databundleVersionId":10123315},{"sourceType":"datasetVersion","sourceId":9742909,"datasetId":4143466,"databundleVersionId":9980931},{"sourceType":"kernelVersion","sourceId":201579302},{"sourceType":"kernelVersion","sourceId":206342978}],"dockerImageVersionId":30787,"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":"# !pip uninstall scikit_learn -qq -y\n# !pip install --find-links /kaggle/input/utilities /kaggle/input/utilities/scikit_learn-1.3.2-cp310-cp310-manylinux_2_17_x86_64.manylinux2014_x86_64.whl -qq","metadata":{"execution":{"iopub.status.busy":"2024-11-11T09:40:59.647782Z","iopub.execute_input":"2024-11-11T09:40:59.648091Z","iopub.status.idle":"2024-11-11T09:40:59.652896Z","shell.execute_reply.started":"2024-11-11T09:40:59.648053Z","shell.execute_reply":"2024-11-11T09:40:59.651978Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport gc\nimport json\nimport random \nimport joblib\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-11-11T09:40:59.654152Z","iopub.execute_input":"2024-11-11T09:40:59.654430Z","iopub.status.idle":"2024-11-11T09:41:02.735770Z","shell.execute_reply.started":"2024-11-11T09:40:59.654399Z","shell.execute_reply":"2024-11-11T09:41:02.734779Z"},"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":{"execution":{"iopub.status.busy":"2024-11-11T09:41:02.737125Z","iopub.execute_input":"2024-11-11T09:41:02.737660Z","iopub.status.idle":"2024-11-11T09:41:02.744210Z","shell.execute_reply.started":"2024-11-11T09:41:02.737611Z","shell.execute_reply":"2024-11-11T09:41:02.743284Z"},"trusted":true},"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-11-11T09:41:02.747865Z","iopub.execute_input":"2024-11-11T09:41:02.748607Z","iopub.status.idle":"2024-11-11T09:41:02.784959Z","shell.execute_reply.started":"2024-11-11T09:41:02.748573Z","shell.execute_reply":"2024-11-11T09:41:02.784214Z"},"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-11-11T09:41:02.785959Z","iopub.execute_input":"2024-11-11T09:41:02.786266Z","iopub.status.idle":"2024-11-11T09:41:02.806255Z","shell.execute_reply.started":"2024-11-11T09:41:02.786234Z","shell.execute_reply":"2024-11-11T09:41:02.805299Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Load Model","metadata":{}},{"cell_type":"code","source":"from gatemlp import GateMLP","metadata":{"execution":{"iopub.status.busy":"2024-11-11T09:41:02.807686Z","iopub.execute_input":"2024-11-11T09:41:02.808050Z","iopub.status.idle":"2024-11-11T09:41:02.813881Z","shell.execute_reply.started":"2024-11-11T09:41:02.808008Z","shell.execute_reply":"2024-11-11T09:41:02.812866Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model_dir = Path('/kaggle/input/js-2024-neural-networks/v7_gate_mlp_tag3_time_embed')\n\nckpt_list = list(model_dir.glob('fold*.ckpt'))\n\nfeature_means = pd.read_csv(model_dir/\"feature_means.csv\")\n\nfeature_means_dict = feature_means.set_index('feature')['mean'].to_dict()","metadata":{"execution":{"iopub.status.busy":"2024-11-11T09:41:02.815166Z","iopub.execute_input":"2024-11-11T09:41:02.815924Z","iopub.status.idle":"2024-11-11T09:41:02.827893Z","shell.execute_reply.started":"2024-11-11T09:41:02.815883Z","shell.execute_reply":"2024-11-11T09:41:02.826892Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df_tag = pd.read_csv('/kaggle/input/jane-street-real-time-market-data-forecasting/features.csv')\n\nnum_features = df_tag[df_tag['tag_3']]['feature'].values.tolist() + ['weight'] + ['time_id_sin', 'time_id_cos', 'time_id_sin_half', 'time_id_cos_half']\n\nprint(f\"Number of numerical features: {len(num_features)}\")\nprint(num_features)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-11T09:41:02.829078Z","iopub.execute_input":"2024-11-11T09:41:02.829704Z","iopub.status.idle":"2024-11-11T09:41:02.838692Z","shell.execute_reply.started":"2024-11-11T09:41:02.829671Z","shell.execute_reply":"2024-11-11T09:41:02.837952Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# MLP_CONFIG = {\n#     'input_size': len(num_features),\n#     'hidden_sizes': [8, 8, 8, 8, 8],\n#     'output_size': 1\n# }\n\nMLP_CONFIG = {\n    'input_size': len(num_features),\n    'hidden_sizes': [32, 8, 32],\n    'output_size': 1\n}\n\nprint(\"MLP config: \")\nfor k, v in MLP_CONFIG.items():\n    print(\"  \", k, \":\", v)\n\nDEVICE = 'cuda' if torch.cuda.is_available() else 'cpu'\nprint(\"Use Device:\", DEVICE)","metadata":{"execution":{"iopub.status.busy":"2024-11-11T09:41:02.839704Z","iopub.execute_input":"2024-11-11T09:41:02.840468Z","iopub.status.idle":"2024-11-11T09:41:02.846437Z","shell.execute_reply.started":"2024-11-11T09:41:02.840412Z","shell.execute_reply":"2024-11-11T09:41:02.845806Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model_list = []\n\nfor ckpt in ckpt_list:\n\n    model = GateMLP(**MLP_CONFIG)\n    \n    state_dict = torch.load(ckpt, map_location=DEVICE, weights_only=False)\n\n    native_state_dict = OrderedDict()\n    for 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\n    model.load_state_dict(native_state_dict)\n    model.eval() \n    model.to(DEVICE)\n    \n    model_list.append(model)\n\n    # try a simple test\n    x_random = torch.rand(32, len(num_features)).float().to(DEVICE)\n    y_rand_pred = model(x_random)\n    print(ckpt.stem, y_rand_pred.shape)","metadata":{"execution":{"iopub.status.busy":"2024-11-11T09:41:02.847404Z","iopub.execute_input":"2024-11-11T09:41:02.848166Z","iopub.status.idle":"2024-11-11T09:41:03.105701Z","shell.execute_reply.started":"2024-11-11T09:41:02.848125Z","shell.execute_reply":"2024-11-11T09:41:03.104709Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Run Inference","metadata":{}},{"cell_type":"code","source":"import gc\n\ndef predict_func(x, models):\n    \n    result_all = []\n    \n    with torch.no_grad():\n    \n        x = torch.tensor(x, dtype=torch.float).to(DEVICE)\n        \n        for i, model in enumerate(models):    \n            predict = model(x)\n            result_all.append(predict)\n\n    y_ensemble = torch.stack(result_all, dim=0).mean(axis=0).cpu().numpy()\n\n    # if DEVICE == 'cuda':\n    #     gc.collect()\n    #     torch.cuda.empty_cache()\n            \n    return y_ensemble\n       ","metadata":{"execution":{"iopub.status.busy":"2024-11-11T09:41:03.107099Z","iopub.execute_input":"2024-11-11T09:41:03.107506Z","iopub.status.idle":"2024-11-11T09:41:03.114060Z","shell.execute_reply.started":"2024-11-11T09:41:03.107461Z","shell.execute_reply":"2024-11-11T09:41:03.113098Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class JaneStreetPredictor:\n    \n    def __init__(self, models, test_parquet, lag_parquet, feature_cols, pbar_length=0):\n        \n        # Initialize model and parameters\n        self.models = models\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        self.responder_cols = [f'responder_{i}' for i in range(9)]\n        self.lag_columns = [f'responder_{i}_lag_1' for i in range(9)]\n        self.lag_stats = None\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.to_pandas()\n        #     self.lag_stats = self.lags_.groupby(['date_id', 'symbol_id'])[self.lag_columns].last().reset_index()\n        #     self.lag_stats = self.lag_stats.rename(columns={k: '_'.join(k.split('_')[:-1]) for k in self.lag_columns})\n        \n        # Cache the test data for future use\n        # self.test_ = test\n        \n        global feature_means_dict\n\n        # test = test.to_pandas()\n        # test = pd.merge(test, self.lag_stats, on=['date_id', 'symbol_id'], how='left')\n\n        row_id = test['row_id'].to_numpy()\n\n        max_val = 968\n        col = 'time_id'\n        test = test.with_columns(\n            (2 * np.pi * pl.col(col) / max_val).sin().alias(col + '_sin'),\n            (2 * np.pi * pl.col(col) / max_val).cos().alias(col + '_cos'),\n            (2 * np.pi * pl.col(col) / (max_val // 2)).sin().alias(col + '_sin_half'),\n            (2 * np.pi * pl.col(col) / (max_val // 2)).cos().alias(col + '_cos_half')\n        )\n\n        col_missing = test.select(pl.all().is_null().any()).transpose().with_columns(\n            pl.Series(test.columns).alias('feature')\n            ).filter(pl.col('column_0'))['feature'].to_list()\n\n        test = test.with_columns(\n            *[pl.col(col).fill_null(feature_means_dict[col]) for col in col_missing]\n            ).select(self.feature_cols)\n        \n        x_stack = test.to_numpy()\n       \n        # Call the model's prediction function\n        predicts = predict_func(x_stack, self.models)  # Placeholder for the actual model prediction logic\n\n        # Create a DataFrame to store the predictions\n        predictions = pl.DataFrame({'row_id': row_id, 'responder_6': predicts})\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":{"execution":{"iopub.status.busy":"2024-11-11T09:41:03.115261Z","iopub.execute_input":"2024-11-11T09:41:03.115555Z","iopub.status.idle":"2024-11-11T09:41:03.132587Z","shell.execute_reply.started":"2024-11-11T09:41:03.115523Z","shell.execute_reply":"2024-11-11T09:41:03.131771Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"**Note: The actual score time is about 30 times longer than the inference time with the synthetic test data.**","metadata":{}},{"cell_type":"code","source":"js_predictor = JaneStreetPredictor(\n    model_list, \n    test_parquet, \n    lag_parquet, \n    feature_cols = num_features,\n    pbar_length = total_time_steps\n)\n\njs_predictor.run_inference_server()","metadata":{"execution":{"iopub.status.busy":"2024-11-11T09:41:03.133636Z","iopub.execute_input":"2024-11-11T09:41:03.133961Z","iopub.status.idle":"2024-11-11T09:41:47.161048Z","shell.execute_reply.started":"2024-11-11T09:41:03.133929Z","shell.execute_reply":"2024-11-11T09:41:47.160133Z"},"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    print(len(pl_sub))\n    display(pl_sub)","metadata":{"execution":{"iopub.status.busy":"2024-11-11T09:41:47.162397Z","iopub.execute_input":"2024-11-11T09:41:47.163120Z","iopub.status.idle":"2024-11-11T09:41:47.177761Z","shell.execute_reply.started":"2024-11-11T09:41:47.163065Z","shell.execute_reply":"2024-11-11T09:41:47.176412Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if USE_STYNTHETIC:\n    \n    # Custom R2 metric for LightGBM\n    def r2_lgb(y_true, y_pred, sample_weight):\n\n        a = np.average((y_pred - y_true) ** 2, weights=sample_weight)\n        b = np.average((y_true) ** 2, weights=sample_weight)\n        r2 = 1 - a / (b + 1e-38)\n\n        return 'r2', r2, True\n    \n    y_trues = []\n    y_preds = []\n    weights = []\n    \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', 'weight'])\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        y_true = pl_merge['responder_6_lag_1'].to_numpy()\n        y_pred = pl_merge['responder_6'].to_numpy()\n        weight = pl_merge['weight'].to_numpy()\n        \n        y_trues.append(y_true)\n        y_preds.append(y_pred)\n        weights.append(weight)\n    \n        _, r2_result, _ = r2_lgb(y_true, y_pred, weight)\n        \n        ax.plot(y_true[:200], 'b-', label='True')\n        ax.plot(y_pred[:200], 'r-', label='Pred')\n        ax.set_title(f\"Symbol {symbol_id} | R^2 = {r2_result:.6f}\")\n        ax.legend()\n        ax.grid(True, ls=\"--\")\n\n    fig.tight_layout()\n    plt.show()\n    \n    _, r2_result, _ = r2_lgb(np.concatenate(y_trues), np.concatenate(y_preds), np.concatenate(weights))\n    print(\"Overall R^2 score: \", r2_result)","metadata":{"execution":{"iopub.status.busy":"2024-11-11T09:41:47.179225Z","iopub.execute_input":"2024-11-11T09:41:47.180443Z","iopub.status.idle":"2024-11-11T09:41:50.242685Z","shell.execute_reply.started":"2024-11-11T09:41:47.180395Z","shell.execute_reply":"2024-11-11T09:41:50.241793Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}