{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.15","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"}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# Import necessary libraries\n\nimport numpy as np\nimport polars as pl\nimport pandas as pd\nimport lightgbm as lgb\nimport xgboost as xgb\nimport os\nimport joblib\nimport kaggle_evaluation.jane_street_inference_server","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Set up constants\nTARGET = 'responder_6'\nFEAT_COLS = [f\"feature_{i:02d}\" for i in range(79)]\n\n# Function to load data with optional filtering\ndef load_data(date_id_range=None, time_id_range=None, columns=None, return_type='pl'):\n    data_dir = '../input/jane-street-real-time-market-data-forecasting'\n    data = pl.scan_parquet(f\"{data_dir}/train.parquet\")\n\n    if date_id_range is not None:\n        start_date, end_date = date_id_range\n        data = data.filter((pl.col(\"date_id\") >= start_date) & (pl.col(\"date_id\") <= end_date))\n\n    if time_id_range is not None:\n        start_time, end_time = time_id_range\n        data = data.filter((pl.col(\"time_id\") >= start_time) & (pl.col(\"time_id\") <= end_time))\n\n    if columns is not None:\n        data = data.select(columns)\n\n    if return_type == 'pd':\n        return data.collect().to_pandas()\n    else:\n        return data.collect()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-25T14:29:11.17497Z","iopub.status.idle":"2024-11-25T14:29:11.175316Z","shell.execute_reply.started":"2024-11-25T14:29:11.175134Z","shell.execute_reply":"2024-11-25T14:29:11.17516Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_range = (0,1698)\ntrain_data = load_data(date_id_range=train_range, columns=[\"date_id\", \"weight\"] + FEAT_COLS + [TARGET], return_type='pl')\ntrain_data = train_data.drop_nulls()\ntrain_data","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-25T14:29:11.176516Z","iopub.status.idle":"2024-11-25T14:29:11.176895Z","shell.execute_reply.started":"2024-11-25T14:29:11.176722Z","shell.execute_reply":"2024-11-25T14:29:11.176739Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"data_size = train_data.shape[0]\nfeature_size = len(FEAT_COLS)\nprint(data_size)\nprint(feature_size)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-25T14:29:11.177946Z","iopub.status.idle":"2024-11-25T14:29:11.178225Z","shell.execute_reply.started":"2024-11-25T14:29:11.178081Z","shell.execute_reply":"2024-11-25T14:29:11.178095Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"split_idx = int(0.8*data_size)\nX_train = train_data[FEAT_COLS][:split_idx]\nX_valid = train_data[FEAT_COLS][split_idx:]\nY_train = train_data[TARGET][:split_idx]\nY_valid = train_data[TARGET][split_idx:]\nweights_train = train_data[\"weight\"][:split_idx]\nweights_valid = train_data[\"weight\"][split_idx:]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-25T14:29:11.179225Z","iopub.status.idle":"2024-11-25T14:29:11.179502Z","shell.execute_reply.started":"2024-11-25T14:29:11.179362Z","shell.execute_reply":"2024-11-25T14:29:11.179376Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.optim as optim","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-25T14:29:11.180774Z","iopub.status.idle":"2024-11-25T14:29:11.181064Z","shell.execute_reply.started":"2024-11-25T14:29:11.180919Z","shell.execute_reply":"2024-11-25T14:29:11.180934Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def customized_loss(y_true, y_pred, weights):\n    numerator = torch.sum(weights * (y_true - y_pred) ** 2)\n    denominator = torch.sum(weights * (y_true ** 2))\n    r2_score = 1 - (numerator / denominator)\n    return r2_score","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-25T14:29:11.181751Z","iopub.status.idle":"2024-11-25T14:29:11.182028Z","shell.execute_reply.started":"2024-11-25T14:29:11.181893Z","shell.execute_reply":"2024-11-25T14:29:11.181907Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 定义线性回归模型\nclass LinearRegressionModel(nn.Module):\n    def __init__(self, input_dim):\n        super(LinearRegressionModel, self).__init__()\n        self.linear = nn.Linear(input_dim, 1)  # 线性层，输出为1\n\n    def forward(self, x):\n        return self.linear(x)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-25T14:29:11.182747Z","iopub.status.idle":"2024-11-25T14:29:11.183042Z","shell.execute_reply.started":"2024-11-25T14:29:11.182899Z","shell.execute_reply":"2024-11-25T14:29:11.182913Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-25T14:29:11.184216Z","iopub.status.idle":"2024-11-25T14:29:11.1845Z","shell.execute_reply.started":"2024-11-25T14:29:11.184354Z","shell.execute_reply":"2024-11-25T14:29:11.184367Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 初始化模型、损失函数和优化器\nmodel = LinearRegressionModel(input_dim=feature_size).to(device)\noptimizer = optim.Adam(model.parameters(), lr=0.01)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-25T14:29:11.185623Z","iopub.status.idle":"2024-11-25T14:29:11.185954Z","shell.execute_reply.started":"2024-11-25T14:29:11.185786Z","shell.execute_reply":"2024-11-25T14:29:11.185802Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 训练模型\nnum_epochs = 100\nbatch_size = 2**14\n\n\n# 早停参数\npatience = 5  # 当验证损失不降低时，容忍的 epoch 数量\nmin_delta = 1e-4  # 损失减少的最小值\nbest_loss = np.inf\ntrigger_times = 0\n\nX_valid = torch.tensor(X_valid.to_numpy(), dtype=torch.float32).to(device)\nY_valid = torch.tensor(Y_valid.to_numpy(), dtype=torch.float32).to(device)\nweights_valid = torch.tensor(weights_valid.to_numpy(), dtype=torch.float32).to(device)\n\n\nfor epoch in range(num_epochs):\n\n    model.train()\n    for i in range(0, len(X_train), batch_size):\n        print(f\"Epoch [{epoch+1}/{num_epochs}], Batch [{i//batch_size+1}/{len(X_train)//batch_size+1}]\")\n        X_batch = torch.tensor(X_train[i:i+batch_size].to_numpy(), dtype=torch.float32).to(device)\n        Y_batch = torch.tensor(Y_train[i:i+batch_size].to_numpy(), dtype=torch.float32).to(device)\n        weights_batch = torch.tensor(weights_train[i:i+batch_size].to_numpy(), dtype=torch.float32).to(device)\n        # 前向传播\n        Y_batch_pred = model(X_batch)\n\n        # 计算损失\n        loss = customized_loss(Y_batch, Y_batch_pred, weights_batch)\n\n        # 反向传播和优化\n        optimizer.zero_grad()\n        loss.backward()\n        optimizer.step()\n\n\n    # 验证过程\n    model.eval()\n    with torch.no_grad():\n        Y_valid_pred = model(X_valid)\n        val_loss = customized_loss(Y_valid, Y_valid_pred, weights_valid).item()\n    \n    print(f\"Epoch [{epoch+1}/{num_epochs}], Validation Loss: {val_loss:.4f}\")\n    \n    # Early Stopping 逻辑\n    if val_loss < best_loss - min_delta:\n        best_loss = val_loss\n        trigger_times = 0\n    else:\n        trigger_times += 1\n        if trigger_times >= patience:\n            print(\"Early stopping triggered.\")\n            break\n\n    # 每 10 个 epoch 计算一次加权 R² 分数\n    # if (epoch + 1) % 10 == 0:\n        # r2 = customized_loss(Y_batch_pred, Y_batch, weights_batch)\n        # print(f\"Epoch [{epoch+1}/{num_epochs}], Loss: {loss.item():.4f}, Weighted R²: {r2.item():.4f}\")\n\nprint(\"Training complete.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-25T14:29:11.186949Z","iopub.status.idle":"2024-11-25T14:29:11.187237Z","shell.execute_reply.started":"2024-11-25T14:29:11.187092Z","shell.execute_reply":"2024-11-25T14:29:11.187106Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Function to calculate R² score\ndef calculate_r2(y_true, y_pred, weights):\n    numerator = np.sum(weights * (y_true - y_pred) ** 2)\n    denominator = np.sum(weights * (y_true ** 2))\n    r2_score = 1 - (numerator / denominator)\n    return r2_score","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Function to evaluate the model\ndef evaluate_model(model, test_data):\n    y_pred = model.predict(test_data[FEAT_COLS])\n    y_true = test_data[TARGET].to_numpy() \n    weights = test_data['weight'].to_numpy()  \n    r2_score = calculate_r2(y_true, y_pred, weights)\n    print(f\"Sample weighted zero-mean R-squared score (R2) on test data: {r2_score}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}