{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.11","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":11305158,"sourceType":"competition"},{"sourceId":203900450,"sourceType":"kernelVersion"},{"sourceId":224157538,"sourceType":"kernelVersion"}],"dockerImageVersionId":31040,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"### Libs Needed for Data Loading","metadata":{}},{"cell_type":"code","source":"import os\nimport pickle\nimport joblib\nfrom tqdm import tqdm\n##################\nimport polars as pl\nimport numpy as np\nimport pandas as pd\n################","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-27T21:59:50.229017Z","iopub.execute_input":"2025-05-27T21:59:50.229797Z","iopub.status.idle":"2025-05-27T21:59:52.404845Z","shell.execute_reply.started":"2025-05-27T21:59:50.229748Z","shell.execute_reply":"2025-05-27T21:59:52.404252Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Libs for TabM Model","metadata":{}},{"cell_type":"code","source":"############################\n# Turn This on to Download TabM (if not download bf)\n!git clone https://github.com/yandex-research/tabm\n!pip install rtdl_num_embeddings","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-27T21:47:26.135728Z","iopub.execute_input":"2025-05-27T21:47:26.135892Z","iopub.status.idle":"2025-05-27T21:48:46.835383Z","shell.execute_reply.started":"2025-05-27T21:47:26.135877Z","shell.execute_reply":"2025-05-27T21:48:46.834566Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import sys\nsys.path.append('root/')\n#############################\nimport warnings\nwarnings.filterwarnings(\"ignore\")\n####################\n# PyTorch \nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n# import pytorch_lightning as pl\nfrom pytorch_lightning import (LightningDataModule, LightningModule, Trainer)\nfrom pytorch_lightning.callbacks import EarlyStopping, ModelCheckpoint, Timer\n########################\nfrom sklearn.metrics import r2_score\nfrom sklearn.model_selection import train_test_split\nfrom torch.utils.data import Dataset, DataLoader\n######################\n# Note: 'tabm_refernce' is a py script to funnel\n# the main TabM structure to a pre-defined PyTorch Model class\n# but leave some key parameters to be tuned\n# see 'github.com/yandex-research/tabm' for more details\nfrom tabm_reference import Model, make_parameter_groups\n#####################","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-27T22:00:00.268485Z","iopub.execute_input":"2025-05-27T22:00:00.268898Z","iopub.status.idle":"2025-05-27T22:00:15.916498Z","shell.execute_reply.started":"2025-05-27T22:00:00.268875Z","shell.execute_reply":"2025-05-27T22:00:15.915925Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 1. Data Load & Preprocess","metadata":{}},{"cell_type":"code","source":"# Convert the scatter integers to consecutive\ncategory_mappings = {\n    'feature_09': {\n        2: 0, 4: 1, 9: 2, 11: 3, 12: 4, 14: 5, \n        15: 6, 25: 7, 26: 8, 30: 9, 34: 10, 42: 11, \n        44: 12, 46: 13, 49: 14, 50: 15, 57: 16, 64: 17, \n        68: 18, 70: 19, 81: 20, 82: 21},\n    'feature_10': {\n        1: 0, 2: 1, 3: 2, 4: 3, 5: 4, \n        6: 5, 7: 6, 10: 7, 12: 8},\n    'feature_11': {\n        9: 0, 11: 1, 13: 2, 16: 3, 24: 4, 25: 5, 34: 6, \n        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, \n        214: 20, 230: 21, 261: 22, 297: 23, 336: 24, \n        376: 25, 388: 26, 410: 27, 522: 28, 534: 29, 539: 30},\n    'symbol_id': {\n        0: 0, 1: 1, 2: 2, 3: 3, 4: 4, 5: 5, 6: 6, 7: 7, 8: 8, 9: 9, \n        10: 10, 11: 11, 12: 12, 13: 13, 14: 14, 15: 15, 16: 16, \n        17: 17, 18: 18, 19: 19, 20: 20, 21: 21, 22: 22, 23: 23, \n        24: 24, 25: 25, 26: 26, 27: 27, 28: 28, 29: 29, 30: 30, 31: 31,\n        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\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, return_dtype=pl.Int16).alias(column)\n    )\n\n# Define feature, label, and weight names\n\nfeature_names = ['time_id', 'symbol_id', 'responder_6', 'responder_3', 'weight'\n                ] + [f\"feature_{i:02d}\" for i in range(79) if i != 61] + [\n                    f\"responder_{idx}_lag_1\" for idx in range(9)]\n\nlabel_name = 'responder_6'\nweight_name = 'weight'\nfeature_cat = ['feature_09', 'feature_10', 'feature_11', 'symbol_id', 'time_id']\nfeature_cont = [\n    col for col in feature_names if col not in feature_cat + [\n        label_name, weight_name]\n]\n\nfeature_cont_idx =  [feature_names.index(col) for col in feature_cont]\nfeature_cat_idx = [feature_names.index(col) for col in feature_cat]\nlabel_idx = feature_names.index(label_name)\nweight_idx = feature_names.index(weight_name)\n\n# Use the most voted Lag-1 preprocessed dataset from Kaggle Discussions\n\ninput_path = '/kaggle/input/js24-preprocessing-create-lags'\ntrain_original = pl.scan_parquet(\n    f\"{input_path}/training.parquet\").sort(['date_id', 'time_id', 'symbol_id'])\nvalid_original = pl.scan_parquet(\n    f\"{input_path}/validation.parquet\").sort(['date_id', 'time_id', 'symbol_id'])\n\nfor col in feature_cat:\n    train_original = encode_column(train_original, col, category_mappings[col])\n    valid_original = encode_column(valid_original, col, category_mappings[col])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-27T21:59:57.383420Z","iopub.execute_input":"2025-05-27T21:59:57.384169Z","iopub.status.idle":"2025-05-27T21:59:57.433211Z","shell.execute_reply.started":"2025-05-27T21:59:57.384144Z","shell.execute_reply":"2025-05-27T21:59:57.432587Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Load Data and Split\n\nUSE_FULLSIZE = False \n# for test runs | Turn this on for actual training\n# kaggle env may not have enough RAM for full training, \n# pls use private server\nDATE_AF = 1500\n\nif USE_FULLSIZE:\n    train_original = pl.concat([train_original, valid_original])\n    df = train_original\\\n        .filter(pl.col('date_id') >= 252)\\\n        .select(feature_names)\\\n        .collect().to_numpy()\n    df[np.isnan(df)] = 0\n    valid = None\nelse:\n    df = train_original\\\n        .filter(pl.col('date_id') >= DATE_AF)\\\n        .select(feature_names)\\\n        .collect().to_numpy()\n    df[np.isnan(df)] = 0\n    \n    valid = valid_original\\\n        .select(feature_names)\\\n        .collect().to_numpy()\n    valid[np.isnan(valid)] = 0","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-27T22:00:27.008655Z","iopub.execute_input":"2025-05-27T22:00:27.009661Z","iopub.status.idle":"2025-05-27T22:00:51.965601Z","shell.execute_reply.started":"2025-05-27T22:00:27.009634Z","shell.execute_reply":"2025-05-27T22:00:51.964759Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 2. Configuration & Model Define","metadata":{}},{"cell_type":"code","source":"import pytorch_lightning as pl \n# pl is taken by polars previously\n# now is changed to alias pytorch_lightning","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-27T22:02:21.428571Z","iopub.execute_input":"2025-05-27T22:02:21.429334Z","iopub.status.idle":"2025-05-27T22:02:21.433283Z","shell.execute_reply.started":"2025-05-27T22:02:21.429299Z","shell.execute_reply":"2025-05-27T22:02:21.432242Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class CFG():\n    use_gpu = True \n    # To finish in reasonale time, pls use GPU\n    # use GPU P100\n    gpu_id = 0\n    seed = 42\n    loader_workers = 12   \n    batch_size = 8192\n    lr = 1e-3\n    weight_decay = 8e-4\n    n_cont_features = len(feature_cont)\n    n_cat_features = 5\n    cat_cardinalities = None if feature_cat is None else  [23, 10, 32, 40, 969]\n    max_epochs = 7","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-27T22:02:24.508290Z","iopub.execute_input":"2025-05-27T22:02:24.508933Z","iopub.status.idle":"2025-05-27T22:02:24.513278Z","shell.execute_reply.started":"2025-05-27T22:02:24.508905Z","shell.execute_reply":"2025-05-27T22:02:24.512490Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Define Custom DataSet\n\nclass CustomDataset(Dataset):\n    def __init__(self, array):\n        self.features_cont = torch.FloatTensor(array[:, feature_cont_idx])\n        self.features_cat = torch.LongTensor(array[:, feature_cat_idx])\n        self.labels = torch.FloatTensor(array[:, label_idx] )\n        self.weights = torch.FloatTensor(array[:, weight_idx])\n    \n    def __len__(self):\n        return len(self.labels)\n    \n    def __getitem__(self, idx):\n        x_cont = self.features_cont[idx]\n        x_cat = self.features_cat[idx]\n        y = self.labels[idx]\n        w = self.weights[idx]\n        return x_cont, x_cat, y, w","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-27T22:02:26.973304Z","iopub.execute_input":"2025-05-27T22:02:26.973890Z","iopub.status.idle":"2025-05-27T22:02:26.978877Z","shell.execute_reply.started":"2025-05-27T22:02:26.973864Z","shell.execute_reply":"2025-05-27T22:02:26.978089Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Define TabM Model Specs\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, 512, 512],\n                    'dropout': 0.25 ,\n                },\n                bins=None,\n                num_embeddings= None,\n                arch_type='tabm',\n                k=self.k,\n            )\n\n        self.lr = lr\n        self.weight_decay = weight_decay\n        self.training_step_outputs = []\n        self.validation_step_outputs = []\n        self.n_classes = n_classes\n        self.loss_fn = R2Loss()\n        # self.loss_fn = nn.MSELoss()\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_cont,x_cat, y, w= batch\n        x_cont = x_cont + torch.randn_like(x_cont) * 0.02\n        y_hat = self(x_cont, x_cat)\n\n        if self.n_classes == 1:\n            loss = self.loss_fn(y_hat.flatten(0, 1), \n                                y.repeat_interleave(self.k))\n            self.training_step_outputs.append((y_hat.mean(1), y, w))\n        else:\n            loss = self.loss_fn(y_hat[:, :, 0].flatten(0, 1), \n                                 y.repeat_interleave(self.k))\n            self.training_step_outputs.append((y_hat[:, :, 0].mean(1), y, w))\n\n        self.log(\n            'train_loss', loss, on_step=True, on_epoch=True, \n            prog_bar=True, logger=True, batch_size=x_cont.size(0)\n        )\n        return loss\n\n    def validation_step(self, batch):\n        x_cont,x_cat, y, w = batch\n        if len(x_cat.size()) == 1:\n            x_cat = None\n        # x_cont = x_cont + torch.randn_like(x_cont) * 0.02\n        y_hat = self(x_cont, x_cat)\n\n        if self.n_classes == 1:\n            loss = self.loss_fn(y_hat.flatten(0, 1), \n                                y.repeat_interleave(self.k))\n            self.validation_step_outputs.append((y_hat.mean(1), y, w))\n        else:\n            loss = self.loss_fn(y_hat[:, :, 0].flatten(0, 1), \n                                 y.repeat_interleave(self.k))\n            self.validation_step_outputs.append((y_hat[:, :, 0].mean(1), y, w))\n\n        self.log(\n            'val_loss', loss, on_step=False, on_epoch=True, prog_bar=False, \n            logger=True, batch_size=x_cont.size(0))\n        return loss\n\n    def on_validation_epoch_end(self):\n        y = torch.cat([x[1] for x in self.validation_step_outputs]).cpu().numpy()\n        \n        if self.trainer.sanity_checking:\n            prob = torch.cat([x[0] for x in self.validation_step_outputs]).cpu().numpy()\n        else:\n            prob = torch.cat(\n                [x[0] for x in self.validation_step_outputs]).cpu().numpy()\n            weights = torch.cat(\n                [x[2] for x in self.validation_step_outputs]).cpu().numpy()\n            # r2_val\n            val_r_square = r2_val(y, prob, weights)\n\n            val_r_square_adj = 1 - (1 - r2_score(y, prob)) * (len(y) - 1) / (len(y) - 1 - 16)\n            self.log(\"val_r_square\", val_r_square, prog_bar=True, \n                     on_step=False, on_epoch=True)\n            self.log(\"val_r_square_adj\", val_r_square_adj, prog_bar=True, \n                     on_step=False, on_epoch=True)\n            \n        self.validation_step_outputs.clear()\n\n    def configure_optimizers(self):\n        optimizer = torch.optim.AdamW(\n            make_parameter_groups(self.model), \n            lr=self.lr, \n            weight_decay=self.weight_decay)\n\n        return {\n            'optimizer': optimizer,\n        }\n\n    def on_train_epoch_end(self):\n        if self.trainer.sanity_checking:\n            return\n\n        y = torch.cat([x[1] for x in self.training_step_outputs]).cpu().numpy()\n        prob = torch.cat([x[0] for x in self.training_step_outputs]).detach().cpu().numpy()\n        weights = torch.cat([x[2] for x in self.training_step_outputs]).cpu().numpy()\n        # r2_training\n        train_r_square = r2_val(y, prob, weights)\n        train_r_square_adj = 1 - (1 - r2_score(y, prob)) * (len(y) - 1) / (len(y) - 1 - 16)\n        # self.log(\"train_r_square\", train_r_square, prog_bar=True, on_step=False, on_epoch=True)\n        self.log(\"train_r_square_adj\", train_r_square_adj, prog_bar=True, \n                 on_step=False, on_epoch=True)\n        self.training_step_outputs.clear()\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:.7f}\" for k, v in metrics.items()}\n\n        formatted_metrics.pop('train_loss_step', None)\n        print(f\"Epoch {epoch}: {formatted_metrics}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-27T22:02:29.273577Z","iopub.execute_input":"2025-05-27T22:02:29.274320Z","iopub.status.idle":"2025-05-27T22:02:29.289690Z","shell.execute_reply.started":"2025-05-27T22:02:29.274294Z","shell.execute_reply":"2025-05-27T22:02:29.288920Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Loss Func\n\ndef r2_val(y_true, y_pred, sample_weight):\n    residuals = sample_weight * (y_true - y_pred) ** 2\n    nom = np.sum(residuals)\n    denom = np.sum(sample_weight * (y_true) ** 2)\n    r2 = 1 - nom/denom\n    return r2\n\n\nclass 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-38)\n        return loss","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-27T22:02:36.707561Z","iopub.execute_input":"2025-05-27T22:02:36.708118Z","iopub.status.idle":"2025-05-27T22:02:36.712922Z","shell.execute_reply.started":"2025-05-27T22:02:36.708093Z","shell.execute_reply":"2025-05-27T22:02:36.712087Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 3. Train","metadata":{}},{"cell_type":"code","source":"# Clean the RAM if NEEDED\n\n# import gc\n# ###########\n# if 'df' in globals():\n#     del df\n# ###########\n# gc.collect()","metadata":{"trusted":true,"execution":{"execution_failed":"2025-05-27T21:46:36.874Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Non-lazy/Direct Load Data\n\n####### Device Select\ndevice = torch.device(\n    f'cuda:{CFG.gpu_id}' if torch.cuda.is_available() and CFG.use_gpu else 'cpu')\naccelerator = 'gpu' if torch.cuda.is_available() and CFG.use_gpu else 'cpu'\nloader_device = 'cpu'\n\n#######\ntrain_ds = CustomDataset(df)\ntrain_dl = DataLoader(\n    train_ds, batch_size = CFG.batch_size, \n    shuffle=True, num_workers = CFG.loader_workers)\n\nif valid is not None:\n    valid_ds = CustomDataset(valid)\n    valid_dl = DataLoader(\n        valid_ds, batch_size = CFG.batch_size, \n        shuffle=True, num_workers = CFG.loader_workers)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-27T22:02:49.598299Z","iopub.execute_input":"2025-05-27T22:02:49.598578Z","iopub.status.idle":"2025-05-27T22:02:51.138787Z","shell.execute_reply.started":"2025-05-27T22:02:49.598558Z","shell.execute_reply":"2025-05-27T22:02:51.138135Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Init Model\n\nmodel = NN(\n    n_cont_features = CFG.n_cont_features,\n    cat_cardinalities = CFG.cat_cardinalities,\n    n_classes = 1,\n    lr = CFG.lr,\n    weight_decay = CFG.weight_decay\n)\n\n# Init Callbacks\n\ncheckpoint_callback = ModelCheckpoint(\n    monitor = 'val_r_square', mode = 'max', save_top_k = 1, \n    verbose = False, filename = f\"./models/tabm.model\") \n\nevery_heckpoint_callback = ModelCheckpoint(\n    every_n_epochs = 1, save_top_k = -1, verbose = False, \n    filename = \"./models/tabm_{epoch:02d}\") \n\ntimer = Timer()\n\n\n# Init PyTorchLightning Trainer\n\nprint(\"Training Epoch is \", CFG.max_epochs)\ntrainer = Trainer(\n    default_root_dir = 'root/',\n    max_epochs = CFG.max_epochs,\n    accelerator = accelerator,\n    devices = [CFG.gpu_id] if CFG.use_gpu else 'auto',\n    # devices = 'auto'\n    callbacks = [ checkpoint_callback, every_heckpoint_callback, timer],\n    enable_progress_bar = True,\n    val_check_interval = 0.5,\n)\n# Start Training\ntrainer.fit(model, train_dl, valid_dl)\n\nprint(f'\\Training completed in {timer.time_elapsed(\"train\"):.2f}s')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-27T22:02:53.808380Z","iopub.execute_input":"2025-05-27T22:02:53.808645Z","iopub.status.idle":"2025-05-27T22:23:32.058672Z","shell.execute_reply.started":"2025-05-27T22:02:53.808625Z","shell.execute_reply":"2025-05-27T22:23:32.057832Z"}},"outputs":[],"execution_count":null}]}