{"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":10256203,"sourceType":"datasetVersion","datasetId":6344006},{"sourceId":214042412,"sourceType":"kernelVersion"},{"sourceType":"kernelVersion","sourceId":214054114}],"dockerImageVersionId":30786,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip -q install rtdl_num_embeddings delu","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2024-12-20T17:42:48.072100Z","iopub.execute_input":"2024-12-20T17:42:48.072447Z","iopub.status.idle":"2024-12-20T17:42:57.576144Z","shell.execute_reply.started":"2024-12-20T17:42:48.072412Z","shell.execute_reply":"2024-12-20T17:42:57.574918Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torch.optim\nimport rtdl_num_embeddings\nfrom rtdl_num_embeddings import compute_bins\nfrom torch.utils.data import TensorDataset, DataLoader, Dataset, ConcatDataset\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\nimport delu\nfrom tqdm import tqdm\nimport polars as pl\nfrom collections import OrderedDict\nimport sys\n\nfrom tabm_reference import Model, make_parameter_groups\n\n\nfrom torch import Tensor\nfrom typing import List, Callable, Union, Any, TypeVar, Tuple\n\nimport joblib\n\nimport gc","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-20T17:42:57.578127Z","iopub.execute_input":"2024-12-20T17:42:57.578428Z","iopub.status.idle":"2024-12-20T17:43:01.945580Z","shell.execute_reply.started":"2024-12-20T17:42:57.578397Z","shell.execute_reply":"2024-12-20T17:43:01.944896Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\ndevice = torch.device('cuda:0' if torch.cuda.is_available() else 'cpu')\n\nfeature_train_list = [f\"feature_{idx:02d}\" for idx in range(79) if idx != 61] \n\n\ntarget_col = \"responder_6\"\ntarget_col2 = \"responder_3\"\n\nfeature_train = feature_train_list \\\n                + [f\"responder_{idx}_lag_1\" for idx in range(9)]+ ['sin_time_id','cos_time_id','sin_time_id_halfday','cos_time_id_halfday', 'sin_feature_61', 'cos_feature_61']\n    \nstart_dt = 900\nend_dt = 1577\n\nfeature_cat = [\"feature_09\", \"feature_10\", \"feature_11\", 'symbol_id', 'time_id']\nfeature_cont = [item for item in feature_train if item not in feature_cat]\nstd_feature = [i for i in feature_train_list if i not in feature_cat] + [f\"responder_{idx}_lag_1\" for idx in range(9)]\n\n# batch_size = 2048\nbatch_size = 8192\nnum_epochs = 10\n\nmeans = joblib.load(\"/kaggle/input/js-gbdt-tabm/data_stats.pkl\")['mean']\nstds = joblib.load(\"/kaggle/input/js-gbdt-tabm/data_stats.pkl\")['std']","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-20T17:43:01.946461Z","iopub.execute_input":"2024-12-20T17:43:01.946819Z","iopub.status.idle":"2024-12-20T17:43:01.994423Z","shell.execute_reply.started":"2024-12-20T17:43:01.946793Z","shell.execute_reply":"2024-12-20T17:43:01.993615Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Load data","metadata":{}},{"cell_type":"code","source":"train_original = pl.scan_parquet(\"/kaggle/input/jane-street-data-preprocessing/training.parquet\").sort(['date_id', 'time_id', 'symbol_id'])                                                                                              \nvalid_original = pl.scan_parquet(\"/kaggle/input/jane-street-data-preprocessing/validation.parquet\").sort(['date_id', 'time_id', 'symbol_id'])                                    \n# all_original = pl.concat([train_original, valid_original])\n# def get_category_mapping(df, column):\n#     unique_values = df.select([column]).unique().collect().to_series()\n#     return {cat: idx for idx, cat in enumerate(unique_values)}\n\n# category_mappings = {col: get_category_mapping(all_original, col) for col in feature_cat + ['symbol_id']}\n\ncategory_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\n\ndef encode_column(df, column, mapping):\n    def encode_category(category):\n        return mapping.get(category, -1)  \n    \n    return df.with_columns(\n        pl.col(column).map_elements(encode_category, return_dtype=pl.Int16).alias(column)\n    )\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])\n\n\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-20T17:43:01.996812Z","iopub.execute_input":"2024-12-20T17:43:01.997383Z","iopub.status.idle":"2024-12-20T17:43:02.042618Z","shell.execute_reply.started":"2024-12-20T17:43:01.997355Z","shell.execute_reply":"2024-12-20T17:43:02.041813Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_data = train_original \\\n             .filter((pl.col(\"date_id\") >= start_dt) & (pl.col(\"date_id\") <= end_dt)) \\\n             .select(feature_train + [target_col, target_col2, 'weight', 'symbol_id', 'time_id'])\n\nvalid_data = valid_original \\\n             .filter(pl.col(\"date_id\") > end_dt)\\\n             .select(feature_train + [target_col, target_col2, 'weight', 'symbol_id', 'time_id'])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-20T17:43:02.043674Z","iopub.execute_input":"2024-12-20T17:43:02.043991Z","iopub.status.idle":"2024-12-20T17:43:02.050116Z","shell.execute_reply.started":"2024-12-20T17:43:02.043955Z","shell.execute_reply":"2024-12-20T17:43:02.049355Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%time\n\n\ntrain_data_tensor = torch.tensor(train_data.collect().to_numpy(), dtype=torch.float32)\ntrain_ds = TensorDataset(train_data_tensor)\ntrain_dl = DataLoader(train_ds, batch_size=batch_size, num_workers=4, pin_memory=True, shuffle=True)\ndel train_data_tensor\ngc.collect()\n\nprint(\"train data done\")\n\nvalid_data_tensor = torch.tensor(valid_data.collect().to_numpy(), dtype=torch.float32)\nvalid_ds = TensorDataset(valid_data_tensor)\nvalid_dl = DataLoader(valid_ds, batch_size=batch_size, num_workers=4, pin_memory=True, shuffle=False)\ndel valid_data_tensor\ngc.collect()\n\n\nall_data = False\nif all_data:\n    train_ds = ConcatDataset([train_ds, valid_ds])\n    train_dl = DataLoader(train_ds, batch_size=batch_size, num_workers=4, pin_memory=True, shuffle=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-20T17:43:02.051519Z","iopub.execute_input":"2024-12-20T17:43:02.051937Z","iopub.status.idle":"2024-12-20T17:44:37.256454Z","shell.execute_reply.started":"2024-12-20T17:43:02.051897Z","shell.execute_reply":"2024-12-20T17:44:37.255594Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Define Model","metadata":{}},{"cell_type":"code","source":"n_cont_features = len(feature_cont)\nn_cat_features = 5\nn_classes = 2\n# cat_cardinalities = [83, 13, 540, 40]\ncat_cardinalities = [23, 10, 32, 40, 969]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-20T17:44:37.257553Z","iopub.execute_input":"2024-12-20T17:44:37.257829Z","iopub.status.idle":"2024-12-20T17:44:37.262218Z","shell.execute_reply.started":"2024-12-20T17:44:37.257803Z","shell.execute_reply":"2024-12-20T17:44:37.261266Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## TabM Model","metadata":{}},{"cell_type":"code","source":"class LogCoshLoss(nn.Module):\n    def __init__(self):\n        super(LogCoshLoss, self).__init__()\n\n    def forward(self, y_pred, y_true):\n        loss = torch.log(torch.cosh(y_pred - y_true))\n        return torch.mean(loss)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-20T17:44:37.263210Z","iopub.execute_input":"2024-12-20T17:44:37.263448Z","iopub.status.idle":"2024-12-20T17:44:37.275682Z","shell.execute_reply.started":"2024-12-20T17:44:37.263425Z","shell.execute_reply":"2024-12-20T17:44:37.274770Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%time\n\n# TabM\narch_type = 'tabm'\nbins = None\n\n# TabM-mini with the piecewise-linear embeddings.\n# arch_type = 'tabm-mini'\n# bins_input = train_data_tensor[:, :-4][:, [col for col in range(train_data_tensor[:, :-4].shape[1]) if col not in [9, 10, 11]]]\n# bins = compute_bins(bins_input[torch.randperm(len(bins_input))[:1000000]], ...)\n\n# del bins_input\n# gc.collect()\n\nk = 16 # 集成输出的数量\nmodel = 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    #     if bins is None\n    #     else {\n    #         'type': 'PiecewiseLinearEmbeddings',\n    #         'd_embedding': 16,\n    #         'activation': True,\n    #         'version': 'B',\n    #     }\n    # ),\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=k,\n).to(device)\n\noptimizer = torch.optim.AdamW(\n    # Instead of model.parameters(),\n    make_parameter_groups(model),\n    lr=1e-4,\n    weight_decay=5e-3 ,\n)\n\n# loss_fn = nn.MSELoss()\n# loss_fn = nn.HuberLoss(delta=0.2)\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\n        \ndef weighted_mse_loss(inputs, targets, weights=None):\n    loss = (inputs - targets) ** 2\n    if weights is not None:\n        loss *= weights.expand_as(loss)\n    loss = torch.mean(loss)\n    return loss\n\nloss_fn = R2Loss()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-20T17:44:37.276485Z","iopub.execute_input":"2024-12-20T17:44:37.276751Z","iopub.status.idle":"2024-12-20T17:44:38.664542Z","shell.execute_reply.started":"2024-12-20T17:44:37.276727Z","shell.execute_reply":"2024-12-20T17:44:38.663627Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Training","metadata":{}},{"cell_type":"code","source":"timer = delu.tools.Timer()\npatience = 5\nearly_stopping = delu.tools.EarlyStopping(patience, mode=\"max\")\nbest = {\n    \"val\": -math.inf,\n    \"epoch\": -1,\n}\ntimer.run()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-20T17:44:38.666516Z","iopub.execute_input":"2024-12-20T17:44:38.666941Z","iopub.status.idle":"2024-12-20T17:44:38.671439Z","shell.execute_reply.started":"2024-12-20T17:44:38.666914Z","shell.execute_reply":"2024-12-20T17:44:38.670527Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def r2_val(y_true, y_pred, sample_weight):\n    residuals = sample_weight * (y_true - y_pred) ** 2\n    weighted_residual_sum = np.sum(residuals)\n\n    # Calculate weighted sum of squared true values (denominator)\n    weighted_true_sum = np.sum(sample_weight * (y_true) ** 2)\n\n    # Calculate weighted R2\n    r2 = 1 - weighted_residual_sum / weighted_true_sum\n\n    return r2\n\nfor epoch in range(num_epochs):\n    model.train()\n\n    # Training\n    train_pred_list = []\n    with tqdm(train_dl, total=len(train_dl), leave=True) as phar:\n        for train_tensor in phar:\n            optimizer.zero_grad()\n            X_input = train_tensor[0][:, :-5].to(device)\n            y1_input = train_tensor[0][:, -5].to(device)\n            y2_input = train_tensor[0][:, -4].to(device)\n            w_input = train_tensor[0][:, -3].to(device)\n\n            \n            symbol_input = train_tensor[0][:, -2].to(device)\n            time_input = train_tensor[0][:, -1].to(device)\n\n                \n            x_cont_input = X_input[:, [col for col in range(X_input.shape[1]) if col not in [9, 10, 11]]]\n            x_cont_input = x_cont_input + torch.randn_like(x_cont_input) * 0.035\n            \n            x_cat_input = X_input[:, [9, 10, 11]]\n            x_cat_input = (torch.concat([x_cat_input, symbol_input.unsqueeze(-1), time_input.unsqueeze(-1)], axis=1)).to(torch.int64)\n\n            \n\n            output = model(x_cont_input, x_cat_input, mask_ratio=0.4)\n\n            loss1 = loss_fn(output[:, :, 0].flatten(0, 1), y1_input.repeat_interleave(k))\n            loss2 = loss_fn(output[:, :, 1].flatten(0, 1), y2_input.repeat_interleave(k))\n\n            loss = (loss1+loss2) / 2\n\n            train_pred_list.append((output[:, :, 0].mean(1), y1_input, w_input))\n        \n            loss.backward()\n            optimizer.step()\n\n            phar.set_postfix(\n                OrderedDict(\n                    epoch=f'{epoch+1}/{num_epochs}',\n                    loss=f'{loss.item():.6f}',\n                    lr=f'{optimizer.param_groups[0][\"lr\"]:.3e}'\n                )\n            )\n            phar.update(1)\n\n    weights_train = torch.cat([x[2] for x in train_pred_list]).cpu().numpy()\n    y_train = torch.cat([x[1] for x in train_pred_list]).cpu().numpy()\n    prob_train = torch.cat([x[0] for x in train_pred_list]).detach().cpu().numpy()\n    train_r2 = r2_val(y_train, prob_train, weights_train)\n    \n    \n    model.eval()\n    valid_loss_list = []\n    valid_pred_list = []\n    for valid_tensor in tqdm(valid_dl):\n        X_valid = valid_tensor[0][:, :-5].to(device)\n        y1_valid = valid_tensor[0][:, -5].to(device)\n        y2_valid = valid_tensor[0][:, -4].to(device)\n        w_valid = valid_tensor[0][:, -3].to(device)\n        symbol_valid = valid_tensor[0][:, -2].to(device)\n        time_valid = valid_tensor[0][:, -1].to(device)\n\n        \n        x_cont_valid = X_valid[:, [col for col in range(X_valid.shape[1]) if col not in [9, 10, 11]]]\n        # x_cont_valid = x_cont_valid + torch.randn_like(x_cont_valid) * 0.035\n        \n        x_cat_valid = X_valid[:, [9, 10, 11]]\n        x_cat_valid = (torch.concat([x_cat_valid, symbol_valid.unsqueeze(-1),time_valid.unsqueeze(-1)], axis=1)).to(torch.int64)\n    \n        with torch.no_grad():\n            y_pred = model(x_cont_valid, x_cat_valid).squeeze(-1)\n    \n\n        val_loss1 = loss_fn(y_pred[:, :, 0].flatten(0, 1), y1_valid.repeat_interleave(k))\n        val_loss2 = loss_fn(y_pred[:, :, 1].flatten(0, 1), y2_valid.repeat_interleave(k))\n    \n        val_loss = (val_loss1+val_loss2) / 2    \n        \n        # val_loss = loss_fn(y_pred.flatten(0, 1), y_valid.repeat_interleave(k))\n        valid_loss_list.append(val_loss)\n        valid_pred_list.append((y_pred[:, :, 0].mean(1), y1_valid, w_valid))\n    \n    valid_loss_mean = sum(valid_loss_list) / len(valid_loss_list)\n    # val_r2 = r2_score(y_valid_data, torch.cat(valid_pred_list).numpy(), sample_weight=w_valid_data)\n\n    weights_eval = torch.cat([x[2] for x in valid_pred_list]).cpu().numpy()\n    y_eval = torch.cat([x[1] for x in valid_pred_list]).cpu().numpy()\n    prob_eval = torch.cat([x[0] for x in valid_pred_list]).cpu().numpy()\n    val_r2 = r2_val(y_eval, prob_eval, weights_eval)\n\n\n    \n    print(f\"Epoch {epoch + 1}: train_r2 = {train_r2:.6f}, val_loss_mean={valid_loss_mean:.6f}, val_r2={val_r2:.6f}, [time] {timer}\")\n\n\n    \n    \n    \n    if val_r2 > best[\"val\"]:\n        print(\"🌸 New best epoch! 🌸\")\n        best = {\"val\": val_r2, \"epoch\": epoch}\n        checkpoint = {\n            'epoch': epoch,\n            'model_state_dict': model.state_dict(),\n            'optimizer_state_dict': optimizer.state_dict(),\n            'r2': val_r2,\n        }\n        torch.save(checkpoint, f'epoch{epoch}_r2_{val_r2}.pt')\n    print()\n    \n    early_stopping.update(val_r2)\n    if early_stopping.should_stop():\n        print(\"Early stop\")\n        break\n\n\ncheckpoint = {\n    'model_state_dict': model.state_dict(),\n    'optimizer_state_dict': optimizer.state_dict(),\n    # 'r2': val_r2,\n}\n\ntorch.save(checkpoint, f'last_tabm.pt')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-20T17:44:38.672511Z","iopub.execute_input":"2024-12-20T17:44:38.672883Z","execution_failed":"2024-12-20T17:46:39.906Z"}},"outputs":[],"execution_count":null}]}