{"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":10253875,"sourceType":"datasetVersion","datasetId":6297065},{"sourceId":204479873,"sourceType":"kernelVersion"},{"sourceId":207787842,"sourceType":"kernelVersion"},{"sourceId":213144305,"sourceType":"kernelVersion"},{"sourceId":214280531,"sourceType":"kernelVersion"}],"dockerImageVersionId":30787,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"Train code : https://www.kaggle.com/code/i2nfinit3y/jane-street-tabm-ft-transformer-training/comments","metadata":{}},{"cell_type":"code","source":"!pip install rtdl_num_embeddings -q --no-index --find-links=/kaggle/input/jane-street-import/rtdl_num_embeddings","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-23T04:57:51.625601Z","iopub.execute_input":"2024-12-23T04:57:51.625869Z","iopub.status.idle":"2024-12-23T04:58:00.534954Z","shell.execute_reply.started":"2024-12-23T04:57:51.625842Z","shell.execute_reply":"2024-12-23T04:58:00.533974Z"}},"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 tanm_reference import Model, make_parameter_groups\n\nimport warnings\nwarnings.filterwarnings(\"ignore\")\n\nimport kaggle_evaluation.jane_street_inference_server\n\nimport os\n\nimport joblib\n\nfrom pytorch_lightning import (LightningDataModule, LightningModule, Trainer)\nfrom pytorch_lightning.callbacks import Callback\nimport gc","metadata":{"_uuid":"f573766f-0a4b-4a41-a873-d3e78e56afaf","_cell_guid":"8936e699-b090-4d2b-9bf2-b77f90fbefdb","trusted":true,"execution":{"iopub.status.busy":"2024-12-23T04:58:00.536726Z","iopub.execute_input":"2024-12-23T04:58:00.537045Z","iopub.status.idle":"2024-12-23T04:58:08.153191Z","shell.execute_reply.started":"2024-12-23T04:58:00.537012Z","shell.execute_reply":"2024-12-23T04:58:08.152449Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"feature_list = [f\"feature_{idx:02d}\" for idx in range(79) if idx != 61]\n\ntarget_col = \"responder_6\" \n\nfeature_test = feature_list \\\n                + [f\"responder_{idx}_lag_1\" for idx in range(9)]\n\nfeature_test_ol = feature_list \\\n                + [f\"responder_{idx}_lag_1\" for idx in range(9)] + ['symbol_id', 'time_id']\n\nfeature_cat = [\"feature_09\", \"feature_10\", \"feature_11\", \"symbol_id\", \"time_id\"]\nfeature_cont = [item for item in feature_test if item not in feature_cat]\n\nbatch_size = 8192\n\nstd_feature = [i for i in feature_list if i not in feature_cat] + [f\"responder_{idx}_lag_1\" for idx in range(9)]\n\ndata_stats = joblib.load(\"/kaggle/input/my-own-js/data_stats.pkl\")\nmeans = data_stats['mean']\nstds = data_stats['std']\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    ])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-23T04:58:08.154213Z","iopub.execute_input":"2024-12-23T04:58:08.154622Z","iopub.status.idle":"2024-12-23T04:58:08.168563Z","shell.execute_reply.started":"2024-12-23T04:58:08.154595Z","shell.execute_reply":"2024-12-23T04:58:08.167780Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"category_mappings = {'feature_09': {2: 0, 4: 1, 9: 2, 11: 3, 12: 4, 14: 5, 15: 6, 25: 7, 26: 8, 30: 9, 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': {0: 0, 1: 1, 2: 2, 3: 3, 4: 4, 5: 5, 6: 6, 7: 7, 8: 8, 9: 9, 10: 10, 11: 11, 12: 12, 13: 13, 14: 14, 15: 15, 16: 16, 17: 17, 18: 18, 19: 19,\n  20: 20, 21: 21, 22: 22, 23: 23, 24: 24, 25: 25, 26: 26, 27: 27, 28: 28, 29: 29, 30: 30, 31: 31, 32: 32, 33: 33, 34: 34, 35: 35, 36: 36, 37: 37, 38: 38},\n 'time_id' : {i : i for i in range(968)}}\n\ndef encode_column(df, column, mapping):\n    max_value = max(mapping.values())  \n\n    def encode_category(category):\n        return mapping.get(category, max_value + 1)  \n    \n    return df.with_columns(\n        pl.col(column).map_elements(encode_category).alias(column)\n    )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-23T04:58:08.169517Z","iopub.execute_input":"2024-12-23T04:58:08.169842Z","iopub.status.idle":"2024-12-23T04:58:08.180194Z","shell.execute_reply.started":"2024-12-23T04:58:08.169815Z","shell.execute_reply":"2024-12-23T04:58:08.179330Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"all_cols = [\"date_id\", \"symbol_id\", \"time_id\", \"weight\"] + [f\"feature_{idx:02d}\" for idx in range(79)]+ [f\"responder_{idx}_lag_1\" for idx in range(9)]\\\n            + [target_col]\n\npl_train = pl.scan_parquet(\"/kaggle/input/jane-street-data-preprocessing/validation.parquet\") \\\n             .sort([\"date_id\", \"time_id\", 'symbol_id'])\\\n             .select(all_cols).collect().sample(fraction=0.015)\n# pl_train = pl.scan_parquet(\"/kaggle/input/jane-street-data-preprocessing/training.parquet\") \\\n#              .filter(pl.col('date_id') >= 1457)\\\n#              .sort([\"date_id\", \"time_id\", 'symbol_id'])\\\n#              .select(all_cols).collect().sample(fraction=0.005)\n\npl_train = pl_train.with_row_count(name=\"row_id\")\npl_train = pl_train.with_columns(pl.col(\"row_id\").cast(pl.Int64))  \n\n\nfor col in feature_cat + ['symbol_id', 'time_id']:\n     pl_train = encode_column(pl_train, col, category_mappings[col])\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-23T04:58:08.181979Z","iopub.execute_input":"2024-12-23T04:58:08.182225Z","iopub.status.idle":"2024-12-23T04:58:16.426884Z","shell.execute_reply.started":"2024-12-23T04:58:08.182201Z","shell.execute_reply":"2024-12-23T04:58:16.425954Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# TabM","metadata":{}},{"cell_type":"code","source":"class R2Loss(nn.Module):\n    def __init__(self):\n        super(R2Loss, self).__init__()\n\n    def forward(self, y_pred, y_true):\n        mse_loss = torch.sum((y_pred - y_true) ** 2)\n        var_y = torch.sum(y_true ** 2)\n        loss = mse_loss / (var_y + 1e-8)\n        return loss\n\nclass ReturnModelCallback(Callback):\n    def __init__(self):\n        self.model = None\n\n    def on_fit_end(self, trainer, pl_module):\n        self.model = pl_module\n\nreturn_model_callback = ReturnModelCallback()\n\nclass NN(LightningModule):\n    def __init__(self, n_cont_features, cat_cardinalities, n_classes, lr, weight_decay):\n        super().__init__()\n        self.save_hyperparameters()\n        self.k = 16\n        self.model = 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,\n                },\n                bins=None,\n                num_embeddings= None,\n                arch_type='tabm',\n                k=self.k,\n            )\n        self.lr = lr\n        self.weight_decay = weight_decay\n        self.loss_fn = R2Loss()\n        # self.loss_fn = nn.MSELoss()\n        # self.loss_fn = nn.HuberLoss()\n        self.automatic_optimization = True\n\n\n    def forward(self, x_cont, x_cat):\n        return self.model(x_cont, x_cat).squeeze(-1)\n\n    def training_step(self, batch):\n        X_data, y_ol = batch\n        X_ol = X_data[:, :-2]\n        symbol_ol = X_data[:, -2]\n        time_ol = X_data[:, -1]\n\n        x_cont_ol = X_ol[:, [col for col in range(X_ol.shape[1]) if col not in [9, 10, 11]]]\n        x_cont_ol = x_cont_ol + torch.randn_like(x_cont_ol) * 0.02\n\n        x_cat_ol = X_ol[:, [9, 10, 11]]\n        x_cat_ol = (torch.concat([x_cat_ol, symbol_ol.unsqueeze(-1), time_ol.unsqueeze(-1)], axis=1)).to(torch.int64)\n\n        y_hat = self(x_cont_ol, x_cat_ol)\n\n        loss = self.loss_fn(y_hat.flatten(0, 1), y_ol.repeat_interleave(self.k))\n\n        self.log('train_loss', loss, on_step=True, on_epoch=True, prog_bar=True, logger=True, batch_size=x_cont_ol.size(0))\n\n        return loss\n\n\n    def configure_optimizers(self):\n        optimizer = torch.optim.AdamW(make_parameter_groups(self.model), lr=self.lr, weight_decay=self.weight_decay, eps=1e-4)\n        return {\n            'optimizer': optimizer,\n\n        }\n\n    def on_train_epoch_end(self):\n        if self.trainer.sanity_checking:\n            return\n\n        epoch = self.trainer.current_epoch\n        metrics = {k: v.item() if isinstance(v, torch.Tensor) else v for k, v in self.trainer.logged_metrics.items()}\n        formatted_metrics = {k: f\"{v:.5f}\" for k, v in metrics.items()}\n        print(f\"Epoch {epoch}: {formatted_metrics}\")\n        \n\ndevice = torch.device('cuda:0' if torch.cuda.is_available() else 'cpu')\n\npl_model = NN.load_from_checkpoint('/kaggle/input/my-own-js/tabm_epochepoch03.ckpt', lr=5e-5, weight_decay=1e-3)\n# model = pl_model.model.to(device)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-23T04:58:16.428263Z","iopub.execute_input":"2024-12-23T04:58:16.428548Z","iopub.status.idle":"2024-12-23T04:58:17.131625Z","shell.execute_reply.started":"2024-12-23T04:58:16.428523Z","shell.execute_reply":"2024-12-23T04:58:17.130625Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"cache = None\ncache_list = []\nbatch_count = 0\nday_count = 0\nretrain = True\nnew_model = None\n\nlags_ : pl.DataFrame | None = None\n# Replace this function with your inference code.\n# You can return either a Pandas or Polars dataframe, though Polars is recommended.\n# Each batch of predictions (except the very first) must be returned within 1 minute of the batch features being provided.\n\n\ndef predict(test: pl.DataFrame, lags: pl.DataFrame | None) -> pl.DataFrame | pd.DataFrame:\n    \"\"\"Make a prediction.\"\"\"\n    # All the responders from the previous day are passed in at time_id == 0. We save them in a global variable for access at every time_id.\n    # Use them as extra features, if you like.\n    global cache          # Declare the global cache\n    global batch_count\n    global day_count\n    global lags_\n    global cache_list\n    global pl_model\n    global return_model_callback\n    global new_model\n\n\n    for col in feature_cat + ['symbol_id', 'time_id']:\n         test = encode_column(test, col, category_mappings[col])\n         if (lags is not None) and (col == 'symbol_id' or col == 'time_id'):\n             lags = encode_column(lags, col, category_mappings[col])\n             \n\n    if lags is not None:\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_ = lags\n        day_count += 1\n        # print(day_count)\n\n        \n\n    symbol_ids = test.select('symbol_id').to_numpy()[:, 0]\n    \n    time_id = test.select(\"time_id\").to_numpy()[0]\n    timie_id_array = test.select(\"time_id\").to_numpy()[:, 0]\n\n\n    if not lags_ is None:\n        # lags_feature = lags_.group_by([\"date_id\", \"symbol_id\"], maintain_order=True).last() # pick up last record of previous date\n        # lags_feature = lags_feature.drop([\"time_id\"])\n        # test = test.join(lags_feature, on=[\"date_id\", \"symbol_id\"],  how=\"left\")\n        # lags_feature = lags_feature.drop([\"date_id\", \"time_id\"])\n        # test = test.join(lags_feature, on=[\"symbol_id\"],  how=\"left\")\n\n        lags_ol = lags_.filter(pl.col(\"time_id\") == time_id)\n        lags_ol = lags_ol.drop('time_id')\n        test = test.join(lags_ol, on=[\"date_id\", \"symbol_id\"],  how=\"left\")\n\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\n\n  \n    # re-train a model on the fly every N days\n    if retrain and day_count % (5 + 1 )== 0 and day_count>=(5 + 1) and time_id == 0:\n        \n        print(\"---------------------------------------------------------------------------------------------\")\n        print(\"Using cache data to retrain the model\")\n        if cache is not None:\n            cache_update = pl.concat(cache_list, rechunk=True)\n            cache = cache.sample(fraction=0.4)\n            cache = pl.concat([cache, cache_update], rechunk=True)\n        else:\n            cache = pl.concat(cache_list, rechunk=True)\n            \n        labels = cache[['date_id', 'time_id', 'symbol_id', 'responder_6_lag_1']]\n        \n        lag_cols_rename = {\"responder_6_lag_1\": \"responder_6\"}\n        labels = labels.rename(lag_cols_rename)\n\n        labels = labels.with_columns(\n            date_id = pl.col('date_id') - 1,  # lagged by 1 day\n        )\n\n        train = cache \n        train = train.join(labels, on=[\"date_id\", 'time_id', \"symbol_id\"],  how=\"left\")\n        train = train.drop_nulls(subset=[\"responder_6\"])\n        \n        train = train.drop([\"is_scored\",\"weight\"])\n        train = train.sample(fraction=0.7)\n\n        hist_data = pl_train.select(train.columns)\n\n\n        # Recasting columns of df1 to match the column types of df2\n        train = train.select([\n            pl.col(col).cast(hist_data.schema[col]) for col in hist_data.columns\n        ])\n\n        train = pl.concat([hist_data, train], rechunk=True)\n        # print(train)\n\n\n        X_train = train[feature_test_ol].to_numpy()\n        y_train = train.select(target_col).to_numpy().flatten()\n   \n\n        ol_ds = TensorDataset(torch.tensor(X_train, dtype=torch.float32), torch.tensor(y_train, dtype=torch.float32))\n        ol_dl = DataLoader(ol_ds, batch_size=8192, num_workers=4, pin_memory=True, shuffle=True)\n\n\n        print(\"Online Learning Start\")\n        \n        trainer = Trainer(\n            max_epochs=2,\n            accelerator = 'gpu',\n            devices=[0],\n            enable_progress_bar=False,\n            callbacks=[return_model_callback],\n            # precision = 16,\n        )\n\n        if new_model is None:\n        \n            pl_model.train()\n            trainer.fit(pl_model, ol_dl)\n    \n            new_model = return_model_callback.model.to(device)\n        else:\n            new_model.train()\n            trainer.fit(new_model, ol_dl)\n    \n            new_model = return_model_callback.model.to(device)\n        \n        print(\"Online Learning Done\")\n        # reset counter otherwise we will retrain for each time_id of the same day\n        day_count = 1\n        # empty cache list\n        cache_list = []\n\n\n\n    test = test.with_columns([\n                     pl.col(col).fill_null(0) for col in feature_test + ['symbol_id', 'time_id']])\n    \n    test = standardize(test, std_feature, means, stds)\n\n    \n    \n    X_test = test[feature_test].to_numpy()\n    X_test_tensor = torch.tensor(X_test, dtype=torch.float32).to(device)\n    \n    symbol_tensor = torch.tensor(symbol_ids, dtype=torch.float32).to(device)\n    time_tensor = torch.tensor(timie_id_array, dtype=torch.float32).to(device)\n    X_cat = X_test_tensor[:, [9, 10, 11]]\n    X_cont = X_test_tensor[:, [i for i in range(X_test_tensor.shape[1]) if i not in [9, 10, 11]]]\n    # X_cont = X_cont + torch.randn_like(X_cont) * 0.02\n\n\n    X_cat = (torch.concat([X_cat, symbol_tensor.unsqueeze(-1), time_tensor.unsqueeze(-1)], axis=1)).to(torch.int64)\n\n\n\n    if new_model is None:\n        pl_model.eval()\n        with torch.no_grad():\n            outputs = pl_model(X_cont, X_cat)\n            preds = outputs.cpu().numpy()\n            preds = preds.mean(1)\n\n    else:\n        pl_model.eval()\n        new_model.eval()\n        with torch.no_grad():\n            outputs_old = pl_model(X_cont, X_cat)\n            outputs_new = new_model(X_cont, X_cat)\n            \n            preds_old = outputs_old.cpu().numpy()\n            preds_old = preds_old.mean(1)\n            preds_new = outputs_new.cpu().numpy()\n            preds_new = preds_new.mean(1)\n        preds = 0.4 * preds_old + 0.6 * preds_new\n\n    \n    # print(f\"predict> preds.shape =\", preds.shape)\n    \n    if retrain:\n        # print(f\"Filling cache for batch count {batch_count}\")\n        cache_list.append(test)\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    if isinstance(predictions, pl.DataFrame):\n        assert predictions.columns == ['row_id', 'responder_6']\n    elif isinstance(predictions, pd.DataFrame):\n        assert (predictions.columns == ['row_id', 'responder_6']).all()\n    else:\n        raise TypeError('The predict function must return a DataFrame')\n    # Confirm has as many rows as the test data.\n    assert len( predictions) == len(test)\n    \n    batch_count+=1\n    \n    return predictions","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-23T04:58:17.132872Z","iopub.execute_input":"2024-12-23T04:58:17.133195Z","iopub.status.idle":"2024-12-23T04:58:17.154232Z","shell.execute_reply.started":"2024-12-23T04:58:17.133167Z","shell.execute_reply":"2024-12-23T04:58:17.153344Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%time\n\nEVAL = False\nif EVAL:\n    test_dir = '/kaggle/input/janestreet-updated-simulator-for-time-series-api/debug/test.parquet'\n    lags_dir = '/kaggle/input/janestreet-updated-simulator-for-time-series-api/debug/lags.parquet'\nelse:\n    test_dir = '/kaggle/input/jane-street-real-time-market-data-forecasting/test.parquet'\n    lags_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-23T04:58:17.155242Z","iopub.execute_input":"2024-12-23T04:58:17.155530Z","execution_failed":"2024-12-23T05:17:46.509Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def weighted_zero_mean_r2(y_true, y_pred, weights):\n    \"\"\"\n    Calculate the sample weighted zero-mean R-squared score.\n\n    Parameters:\n    y_true (numpy.ndarray): Ground-truth values for responder_6.\n    y_pred (numpy.ndarray): Predicted values for responder_6.\n    weights (numpy.ndarray): Sample weight vector.\n\n    Returns:\n    float: The weighted zero-mean R-squared score.\n    \"\"\"\n    numerator = np.sum(weights * (y_true - y_pred)**2)\n    denominator = np.sum(weights * y_true**2)\n    \n    r2_score = 1 - numerator / denominator\n    return r2_score","metadata":{"trusted":true,"execution":{"execution_failed":"2024-12-23T05:17:46.509Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if EVAL:\n\n    submission_file = pd.read_parquet('/kaggle/working/submission.parquet')\n    y_pred = submission_file['responder_6']\n    \n    valid_df = pl.read_parquet(\"/kaggle/input/janestreet-updated-simulator-for-time-series-api/valid_df.parquet\")\n    \n    y_true = valid_df.select(\"responder_6\").to_numpy().reshape(-1)\n    \n    weights = valid_df.select(\"weight\").to_numpy().reshape(-1)\n    \n    print(weighted_zero_mean_r2(y_true, y_pred, weights))","metadata":{"trusted":true,"execution":{"execution_failed":"2024-12-23T05:17:46.510Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!rm -rf lightning_logs","metadata":{"trusted":true,"execution":{"execution_failed":"2024-12-23T05:17:46.510Z"}},"outputs":[],"execution_count":null}]}