{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","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":10290183,"sourceType":"datasetVersion","datasetId":6368446},{"sourceId":10386223,"sourceType":"datasetVersion","datasetId":6434311},{"sourceId":203900450,"sourceType":"kernelVersion"},{"sourceId":214286693,"sourceType":"kernelVersion"}],"dockerImageVersionId":30823,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import polars as pl\n!pip install rtdl_num_embeddings -q --no-index --find-links=/kaggle/input/testnew/rtdl_num_embeddings","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-07T01:54:46.920878Z","iopub.execute_input":"2025-01-07T01:54:46.921292Z","iopub.status.idle":"2025-01-07T01:54:50.567057Z","shell.execute_reply.started":"2025-01-07T01:54:46.921261Z","shell.execute_reply":"2025-01-07T01:54:50.565798Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os, sys, gc\nimport enum\nimport datetime\nimport pickle\nimport dill\nimport numpy as np\nimport pandas as pd\n\nfrom sklearn.metrics import r2_score\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import DataLoader, Dataset\nfrom pytorch_lightning import LightningModule\nimport kaggle_evaluation.jane_street_inference_server\nfrom tanm_reference import Model, make_parameter_groups\n\nimport warnings\nwarnings.filterwarnings('ignore')\npd.options.display.max_columns = None\n\n\n@enum.unique\nclass DataEnum(enum.IntEnum):\n    Train = 0\n    Valid = 1\n    Test = 2\n    Infer = 3\n\nis_debug = False\nis_rerun = os.environ.get('KAGGLE_IS_COMPETITION_RERUN', \"\") != \"\" \nis_local = os.environ.get(\"DOCKER_USING\", \"\") == \"LOCAL\"\nnum_workers = 4\n\nif is_rerun:\n    is_debug = False\n\ndef load_from_dill(model_name, model_path=None, file_ext='.dill'):\n    model_object = None\n    with open(f\"{model_path}/{model_name}{file_ext}\", \"rb\") as file_handle:\n        model_object = dill.load(file_handle)\n    return model_object","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-07T01:54:50.568754Z","iopub.execute_input":"2025-01-07T01:54:50.569105Z","iopub.status.idle":"2025-01-07T01:54:50.578225Z","shell.execute_reply.started":"2025-01-07T01:54:50.569067Z","shell.execute_reply":"2025-01-07T01:54:50.577033Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model_path = '/kaggle/input/js-2024' + ('/' if is_local else '-') + '19-02-1/last_tabm.pt'\nstats_path = '/kaggle/input/js-2024' + ('/' if is_local else '-') + '19-02-1/data_stats.dill'\ndevice = torch.device('cuda:0') # if torch.cuda.is_available() else \"cpu\"\n\ntarget_col = \"responder_6\"\nnecessary_cols = [target_col, 'weight']\nfeat_clear_categ = [\"feature_09\", \"feature_10\", \"feature_11\"]\nfeature_categ = feat_clear_categ + ['symbol_id', 'time_id']\nfeature_cols = [f\"feature_{idx:02d}\" for idx in range(79) if idx not in [9, 10, 11, 61]]\nresponder_cols = [f\"responder_{idx}_lag_1\" for idx in range(9)] \nfeature_cont = feature_cols + responder_cols\ndataset_cols = feature_cont + necessary_cols + feature_categ\nstd_feature = [i for i in feature_cont]\n\nbatch_size = 8192\nn_cont_features = len(feature_cont)\nn_cat_features = len(feature_categ)\nn_classes = None\ncat_cardinalities = [23, 10, 32, 40, 969]\n# TabM\narch_type = 'tabm'\nbins = None\nmodel_koef = 32\n\nprint(n_cont_features, n_cat_features, len(dataset_cols))\n\ncategory_mappings = {\n    'feature_09': {2: 0, 4: 1, 9: 2, 11: 3, 12: 4, 14: 5, 15: 6, 25: 7, 26: 8, 30: 9, \n        34: 10, 42: 11, 44: 12, 46: 13, 49: 14, 50: 15, 57: 16, 64: 17, 68: 18, 70: 19, 81: 20, 82: 21},\n    'feature_10': {1: 0, 2: 1, 3: 2, 4: 3, 5: 4, 6: 5, 7: 6, 10: 7, 12: 8},\n    'feature_11': {9: 0, 11: 1, 13: 2, 16: 3, 24: 4, 25: 5, 34: 6, 40: 7, 48: 8, 50: 9, 59: 10, 62: 11, 63: 12, 66: 13,\n        76: 14, 150: 15, 158: 16, 159: 17, 171: 18, 195: 19, 214: 20, 230: 21, 261: 22, 297: 23, 336: 24, 376: 25, 388: 26, 410: 27, 522: 28, 534: 29, 539: 30},\n    'symbol_id': {i : i for i in range(39)},\n    'time_id' : {i : i for i in range(968)}\n}\n\ndef standardize(df, feature_cols, means, stds):\n    return df.with_columns([\n        ((pl.col(col) - means[col]) / stds[col]).alias(col) for col in feature_cols\n    ])\n\ndef encode_column(df, column, mapping):\n    max_value = max(mapping.values())\n    def encode_category(category):\n        return mapping.get(category, max_value + 1)\n    return df.with_columns(pl.col(column).map_elements(encode_category, return_dtype=pl.Int64).alias(column))\n\nmodel_tabm = Model(\n    n_num_features=n_cont_features,\n    cat_cardinalities=cat_cardinalities,\n    n_classes=n_classes,\n    backbone={\n        'type': 'MLP',\n        'n_blocks': 3 ,\n        'd_block': 512,\n        'dropout': 0.25,\n    },\n    bins=bins,\n    num_embeddings=(\n        None\n        # {\n        #     'type': 'PeriodicEmbeddings',\n        #     'd_embedding': 16,\n        #     'lite':True,\n        # }\n    ),\n    arch_type=arch_type,\n    k=model_koef,\n).to(device)\n\nwith open(stats_path, \"rb\") as file_handle:\n    data_stats = dill.load(file_handle)\nmeans, stds = data_stats['means'], data_stats['stds']\n\ncheckpoint = torch.load(model_path, weights_only=True)\nmodel_tabm.load_state_dict(checkpoint['model_state_dict'])\nmodel_tabm.to(device)\n\nlags_history = None\n\ndef predict_tabm(test: pl.DataFrame, lags: pl.DataFrame | None) -> pl.DataFrame:\n    global lags_history\n    for col in feature_categ:\n        test = encode_column(test, col, category_mappings[col])\n    predictions = test.select(\n        'row_id',\n        pl.lit(0.0).alias('responder_6'),\n    )\n    time_id = test.select(\"time_id\").to_numpy()[0]\n    if time_id == 0:\n        lags = lags.with_columns(pl.col('time_id').cast(pl.Int64))\n        lags = lags.with_columns(pl.col('symbol_id').cast(pl.Int64))    \n        lags_history = lags\n    lags = lags_history.clone().group_by([\"date_id\", \"symbol_id\"], maintain_order=True).last() # pick up last record of previous date\n    test = test.join(lags, on=[\"time_id\", \"symbol_id\"], how=\"left\")\n    \n    test = test.select(pl.all().forward_fill())\n    test = test.with_columns([\n        pl.col(col).fill_null(0) for col in feature_cont + feat_clear_categ\n    ])\n    test = standardize(test, std_feature, means, stds)\n\n    model_tabm.eval()\n\n    with torch.no_grad():\n        X_cont = torch.tensor(test[feature_cont].to_numpy(), dtype=torch.float32).to(device)\n        X_categ = torch.tensor(test[feature_categ].to_numpy(), dtype=torch.int64).to(device)\n        outputs = model_tabm(X_cont, X_categ)\n        # Assuming the model outputs a tensor of shape (batch_size, 1)\n        preds = outputs.squeeze(-1).cpu().numpy()\n        preds = preds.mean(1)\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)\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    return predictions","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-01-07T01:54:50.580341Z","iopub.execute_input":"2025-01-07T01:54:50.580628Z","iopub.status.idle":"2025-01-07T01:54:50.673551Z","shell.execute_reply.started":"2025-01-07T01:54:50.580605Z","shell.execute_reply":"2025-01-07T01:54:50.672664Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"is_first = False\n\ndef predict(test:pl.DataFrame, lags:pl.DataFrame | None) -> pl.DataFrame | pd.DataFrame:\n    global is_first\n\n    pd_tabm = predict_tabm(test,lags).to_pandas()\n    pd_tabm = pd_tabm.rename(columns={'responder_6': 'col_tabm'})\n\n    if not is_first:\n        display(pd_tabm)\n        is_first = True\n\n    predictions = test.select('row_id', pl.lit(0.0).alias('responder_6'))\n    pred = pd_tabm['col_tabm'].to_numpy()\n    predictions = predictions.with_columns(pl.Series('responder_6', pred.ravel()))\n    return predictions\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            '/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":"2025-01-07T01:54:50.674523Z","iopub.execute_input":"2025-01-07T01:54:50.674777Z","iopub.status.idle":"2025-01-07T01:54:50.768237Z","shell.execute_reply.started":"2025-01-07T01:54:50.674753Z","shell.execute_reply":"2025-01-07T01:54:50.766917Z"}},"outputs":[],"execution_count":null}]}