{"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":10056408,"sourceType":"datasetVersion","datasetId":6196729},{"sourceId":10073595,"sourceType":"datasetVersion","datasetId":6209270},{"sourceId":10085548,"sourceType":"datasetVersion","datasetId":6218218},{"sourceId":210819973,"sourceType":"kernelVersion"}],"dockerImageVersionId":30786,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        print(os.path.join(dirname, filename))\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2024-12-03T10:22:08.088637Z","iopub.execute_input":"2024-12-03T10:22:08.089090Z","iopub.status.idle":"2024-12-03T10:22:08.119315Z","shell.execute_reply.started":"2024-12-03T10:22:08.089036Z","shell.execute_reply":"2024-12-03T10:22:08.118189Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import sys\nsys.path.append( \"/kaggle/input/tabm-reference\" )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-03T10:22:08.121735Z","iopub.execute_input":"2024-12-03T10:22:08.122197Z","iopub.status.idle":"2024-12-03T10:22:08.127379Z","shell.execute_reply.started":"2024-12-03T10:22:08.122147Z","shell.execute_reply":"2024-12-03T10:22:08.126311Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install rtdl_num_embeddings -q --no-index --find-links=/kaggle/input/jane-street-packages/rtdl_num_embeddings\n!pip install rtdl_revisiting_models -q --no-index --find-links=/kaggle/input/jane-street-packages/rtdl_revisiting_models","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-03T10:22:08.128938Z","iopub.execute_input":"2024-12-03T10:22:08.129427Z","iopub.status.idle":"2024-12-03T10:22:26.900397Z","shell.execute_reply.started":"2024-12-03T10:22:08.129379Z","shell.execute_reply":"2024-12-03T10:22:26.898958Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torch.optim\nfrom torch.utils.data import Dataset, DataLoader, TensorDataset\n\nfrom sklearn.model_selection import train_test_split\n\nfrom sklearn.metrics import r2_score\nimport pandas as pd\nimport math\nimport numpy as np\nfrom tqdm import tqdm\nimport polars as pl\nfrom collections import OrderedDict\nimport sys\nfrom rtdl_revisiting_models import FTTransformer\nfrom tabm_reference import Model, make_parameter_groups\n\nimport warnings\nwarnings.filterwarnings(\"ignore\")\n\nimport kaggle_evaluation.jane_street_inference_server\n\nimport os","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-03T10:22:26.902255Z","iopub.execute_input":"2024-12-03T10:22:26.903385Z","iopub.status.idle":"2024-12-03T10:22:26.909993Z","shell.execute_reply.started":"2024-12-03T10:22:26.903321Z","shell.execute_reply":"2024-12-03T10:22:26.908954Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"device = torch.device('cuda:0' if torch.cuda.is_available() else 'cpu')\nn_cont_features = 78\nn_cat_features = 3\nn_classes = None\ncat_cardinalities = [83, 13, 540]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-03T10:22:26.913114Z","iopub.execute_input":"2024-12-03T10:22:26.913493Z","iopub.status.idle":"2024-12-03T10:22:26.928421Z","shell.execute_reply.started":"2024-12-03T10:22:26.913460Z","shell.execute_reply":"2024-12-03T10:22:26.927092Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"checkpoint = torch.load('/kaggle/input/epoch1-tmmodel/epoch1_r2_0.0032494751620658624.pt',map_location=torch.device('cpu'))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-03T10:22:26.929718Z","iopub.execute_input":"2024-12-03T10:22:26.930019Z","iopub.status.idle":"2024-12-03T10:22:28.073471Z","shell.execute_reply.started":"2024-12-03T10:22:26.929990Z","shell.execute_reply":"2024-12-03T10:22:28.072350Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"d_out=1\nmodel = FTTransformer(\n    n_cont_features=n_cont_features,\n    cat_cardinalities=cat_cardinalities,\n    d_out=d_out,\n    n_blocks=3,\n    d_block=192,\n    attention_n_heads=8,\n    attention_dropout=0.2,\n    ffn_d_hidden=None,\n    ffn_d_hidden_multiplier=4 / 3,\n    ffn_dropout=0.1,\n    residual_dropout=0.0,\n).to(device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-03T10:22:28.075817Z","iopub.execute_input":"2024-12-03T10:22:28.076145Z","iopub.status.idle":"2024-12-03T10:22:28.097948Z","shell.execute_reply.started":"2024-12-03T10:22:28.076113Z","shell.execute_reply":"2024-12-03T10:22:28.097007Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model.load_state_dict(checkpoint['model_state_dict'])\nmodel.eval()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-03T10:22:28.099300Z","iopub.execute_input":"2024-12-03T10:22:28.099699Z","iopub.status.idle":"2024-12-03T10:22:28.111702Z","shell.execute_reply.started":"2024-12-03T10:22:28.099658Z","shell.execute_reply":"2024-12-03T10:22:28.110548Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"lags_ : pl.DataFrame | None = None\n\nfeature_test=['time_id', 'symbol_id'] + [f\"feature_{idx:02d}\" for idx in range(79)]\nfeature_cat = [\"feature_09\", \"feature_10\", \"feature_11\"]\nfeature_cont = [item for item in feature_test if item not in feature_cat]\n\nbatch_size = 2048\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\n\n    # print(lags)\n    # print(\"----------------------------\")\n    # print(test)\n    # raise(\"stop for debug\")\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    \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    test = test.with_columns([\n        pl.col(col).fill_null(0) for col in feature_test  \n    ])\n    \n    X_test = test[feature_test].to_numpy()\n    # 计算每列的均值（忽略 NaN 值）\n    col_means = np.nanmean(X_test, axis=0)    \n    # 找到 NaN 的位置并用均值填充\n    inds = np.where(np.isnan(X_test))\n    X_test[inds] = np.take(col_means, inds[1])\n    \n    X_test_tensor = torch.tensor(X_test, dtype=torch.float32).to(device)\n    X_test_tensor = torch.nan_to_num(X_test_tensor, nan=0.0, posinf=0.0, neginf=0.0)\n    \n    X_cat = X_test_tensor[:, [11, 12, 13]].to(torch.int64)\n    X_cont = X_test_tensor[:, [i for i in range(X_test_tensor.shape[1]) if i not in [11, 12, 13]]]\n\n    with torch.no_grad():\n        \n        outputs = model(X_cont, X_cat)\n        # Assuming the model outputs a tensor of shape (batch_size, 1)\n        preds = outputs.squeeze(-1).squeeze(-1).cpu().numpy()\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\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-03T10:22:28.113418Z","iopub.execute_input":"2024-12-03T10:22:28.113740Z","iopub.status.idle":"2024-12-03T10:22:28.127274Z","shell.execute_reply.started":"2024-12-03T10:22:28.113710Z","shell.execute_reply":"2024-12-03T10:22:28.126228Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"test_dir = '/kaggle/input/jane-street-real-time-market-data-forecasting/test.parquet'\nlags_dir = '/kaggle/input/jane-street-real-time-market-data-forecasting/lags.parquet'\n\ninference_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            test_dir,\n            lags_dir\n        )\n    )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-03T10:22:28.128636Z","iopub.execute_input":"2024-12-03T10:22:28.129046Z","iopub.status.idle":"2024-12-03T10:22:28.278350Z","shell.execute_reply.started":"2024-12-03T10:22:28.129002Z","shell.execute_reply":"2024-12-03T10:22:28.277132Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}