{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"codemirror_mode":{"name":"ipython","version":3},"file_extension":".py","mimetype":"text/x-python","name":"python","nbconvert_exporter":"python","pygments_lexer":"ipython3","version":"3.10.14"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":84493,"databundleVersionId":9871156,"sourceType":"competition"},{"sourceId":9896473,"sourceType":"datasetVersion","datasetId":6043403},{"sourceId":9944321,"sourceType":"datasetVersion","datasetId":6033262},{"sourceId":207810664,"sourceType":"kernelVersion"}],"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true},"papermill":{"default_parameters":{},"duration":32.100605,"end_time":"2024-11-10T22:54:52.065810","environment_variables":{},"exception":null,"input_path":"__notebook__.ipynb","output_path":"__notebook__.ipynb","parameters":{},"start_time":"2024-11-10T22:54:19.965205","version":"2.6.0"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Jane Street Real-Time Market Data Forecasting with PyTorch\n\n## Overview\nThis notebook is designed for a real-time market data forecasting task using PyTorch. The primary goal is to predict a target variable (responder_6) based on 79 features. This notebook provides a comprehensive workflow for real-time market data forecasting using PyTorch. It includes data loading, model training, validation, and a submission pipeline for inference. The model is a simple feedforward neural network, and the notebook uses custom metrics and callbacks to enhance training and evaluation.\n\n## Key Components\n1. **Configuration**: Defines a `CFG` class to control training mode and batch size.\n2. **Data Loading**:\n    * **Feature Columns**: 79 features named `feature_00` to `feature_78`.\n    * **Target Column**: `responder_6`.\n    * **Training Data**: Loaded from Parquet files, concatenated into `X_train` and `y_train`.\n    * **Validation Data**: Loaded similarly into `X_val` and `y_val`.\n3. **Generators**: Functions to yield batches of data for training and validation.\n4. **Model Architecture**: Simple feedforward neural network with three hidden layers.\n5. **Predict Function**: Handles null values, makes predictions, and formats output.\n6. **Inference Server**: Sets up an inference server using `JSInferenceServer` for handling prediction requests.","metadata":{"papermill":{"duration":0.006975,"end_time":"2024-11-10T22:54:22.806656","exception":false,"start_time":"2024-11-10T22:54:22.799681","status":"completed"},"tags":[]}},{"cell_type":"code","source":"import os\nimport time\nimport pandas as pd\nimport polars as pl\nimport kaggle_evaluation.jane_street_inference_server\nimport numpy as np\nimport pickle\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nfrom tqdm.notebook import tqdm\nfrom sklearn.metrics import r2_score\nfrom torch.cuda.amp import GradScaler, autocast\nimport gc","metadata":{"execution":{"iopub.status.busy":"2024-11-19T04:37:47.756794Z","iopub.execute_input":"2024-11-19T04:37:47.757275Z","iopub.status.idle":"2024-11-19T04:37:47.762928Z","shell.execute_reply.started":"2024-11-19T04:37:47.757244Z","shell.execute_reply":"2024-11-19T04:37:47.761945Z"},"papermill":{"duration":5.904342,"end_time":"2024-11-10T22:54:28.717393","exception":false,"start_time":"2024-11-10T22:54:22.813051","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 1. Configuration\n\n**CFG Class**: Controls training mode and batch size.\n\n**Device**: Uses CUDA if available, otherwise CPU.","metadata":{"papermill":{"duration":0.006326,"end_time":"2024-11-10T22:54:28.730694","exception":false,"start_time":"2024-11-10T22:54:28.724368","status":"completed"},"tags":[]}},{"cell_type":"code","source":"class CFG:\n    is_training = False\n    batch_size = 4096\n    statistic_data_path = \"/kaggle/input/jane-street-rmf-calculate-statistics-values/\"\n    device_name = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n    feature_columns =  [f\"feature_{i:02d}\" for i in range(79)]\n    target_column = \"responder_6\"\n    use_symbol_id_embedding = False\ndevice = torch.device(CFG.device_name)","metadata":{"execution":{"iopub.status.busy":"2024-11-19T04:37:49.855106Z","iopub.execute_input":"2024-11-19T04:37:49.855745Z","iopub.status.idle":"2024-11-19T04:37:49.891038Z","shell.execute_reply.started":"2024-11-19T04:37:49.855711Z","shell.execute_reply":"2024-11-19T04:37:49.890050Z"},"papermill":{"duration":0.077714,"end_time":"2024-11-10T22:54:28.814494","exception":false,"start_time":"2024-11-10T22:54:28.736780","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 2. Load training and validation data\n\n**Feature Columns**: List of 79 feature columns.\n\n**Target Column**: responder_6.\n\n**Training Data**: Loaded from multiple Parquet files, concatenated, and normalized.\n\n**Validation Data**: Loaded similarly and concatenated.","metadata":{"papermill":{"duration":0.005968,"end_time":"2024-11-10T22:54:28.827025","exception":false,"start_time":"2024-11-10T22:54:28.821057","status":"completed"},"tags":[]}},{"cell_type":"code","source":"with open(f\"{CFG.statistic_data_path}mean.pkl\", \"rb\") as f:\n    mean_values = pickle.load(f)\nwith open(f\"{CFG.statistic_data_path}std.pkl\", \"rb\") as f:\n    std_values = pickle.load(f)\nwith open(f\"{CFG.statistic_data_path}min.pkl\", \"rb\") as f:\n    min_values = pickle.load(f)\nwith open(f\"{CFG.statistic_data_path}max.pkl\", \"rb\") as f:\n    max_values = pickle.load(f)\nwith open(f\"{CFG.statistic_data_path}mean_pandas.pkl\", \"rb\") as f:\n    mean_pandas = pickle.load(f)","metadata":{"execution":{"iopub.status.busy":"2024-11-19T04:37:52.319865Z","iopub.execute_input":"2024-11-19T04:37:52.320227Z","iopub.status.idle":"2024-11-19T04:37:52.357321Z","shell.execute_reply.started":"2024-11-19T04:37:52.320196Z","shell.execute_reply":"2024-11-19T04:37:52.356403Z"},"papermill":{"duration":0.062992,"end_time":"2024-11-10T22:54:28.896212","exception":false,"start_time":"2024-11-10T22:54:28.833220","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Load training data","metadata":{"papermill":{"duration":0.006135,"end_time":"2024-11-10T22:54:28.908907","exception":false,"start_time":"2024-11-10T22:54:28.902772","status":"completed"},"tags":[]}},{"cell_type":"code","source":"if CFG.is_training:\n    X_train = None\n    y_train = None\n    train_time_ids = None\n    train_symbol_ids  = None\n    train_weights = None\n    for i in range(9):\n        print(\"=\" * 30)\n        print(f\"Partition {i}\")\n        print(\"=\" * 30)\n        df = pd.read_parquet(f\"/kaggle/input/jane-street-real-time-market-data-forecasting/train.parquet/partition_id={i}/part-0.parquet\")\n        X_part = df[CFG.feature_columns].replace(np.NAN, mean_pandas)\n        X_part = (X_part - min_values) / (max_values - min_values)\n        X_part = X_part.astype(np.float16)\n        weights = df[[\"weight\"]].astype(np.float16)\n        time_ids = df[[\"time_id\"]].astype(np.int16)\n        symbol_ids = df[[\"symbol_id\"]].astype(np.int16)\n        if X_train is None:\n            X_train = X_part\n            y_train = df[[CFG.target_column]]\n            train_time_ids = time_ids\n            train_symbol_ids = symbol_ids\n            train_weights = weights\n        else:\n            X_train = pd.concat([X_train, X_part])\n            y_train = pd.concat([y_train, df[[CFG.target_column]]])\n            train_time_ids = pd.concat([train_time_ids, time_ids])\n            train_symbol_ids = pd.concat([train_symbol_ids, symbol_ids])\n            train_weights = pd.concat([train_weights, weights])\n        del X_part\n        del df\n        del time_ids\n        del symbol_ids\n        gc.collect()","metadata":{"execution":{"iopub.status.busy":"2024-11-19T04:37:54.420298Z","iopub.execute_input":"2024-11-19T04:37:54.420637Z","iopub.status.idle":"2024-11-19T04:37:54.428153Z","shell.execute_reply.started":"2024-11-19T04:37:54.420607Z","shell.execute_reply":"2024-11-19T04:37:54.427269Z"},"papermill":{"duration":0.019291,"end_time":"2024-11-10T22:54:28.934400","exception":false,"start_time":"2024-11-10T22:54:28.915109","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Load Validation data","metadata":{"papermill":{"duration":0.005999,"end_time":"2024-11-10T22:54:28.946885","exception":false,"start_time":"2024-11-10T22:54:28.940886","status":"completed"},"tags":[]}},{"cell_type":"code","source":"X_val = None\ny_val = None\nvalid_weights = None\nvalid_time_ids = None\nvalid_symbol_ids  = None\nfor i in range(9, 10):\n    df = pd.read_parquet(f\"/kaggle/input/jane-street-real-time-market-data-forecasting/train.parquet/partition_id={i}/part-0.parquet\")\n    X_part = df[CFG.feature_columns].fillna(mean_pandas)\n    X_part = (X_part - min_values) / (max_values - min_values)\n    weights = df[[\"weight\"]].astype(np.float16)\n    if X_val is None:\n        X_val = X_part\n        y_val = df[[CFG.target_column]]\n        valid_weights = weights\n        valid_time_ids = df[[\"time_id\"]]\n        valid_symbol_ids = df[[\"symbol_id\"]]\n    else:\n        X_val = pd.concat([X_val, X_part])\n        y_val = pd.concat([y_val, df[[CFG.target_column]]])\n        valid_weights = pd.concat([valid_weights, weights])\n        valid_time_ids = pd.concat([valid_time_ids, df[[\"time_id\"]]])\n        valid_symbol_ids = pd.concat([valid_symbol_ids, df[[\"symbol_id\"]]])\n    del df\n    del X_part\n    gc.collect()\nweights = weights.values.reshape(-1)\nprint(X_val.shape, y_val.shape, weights.shape)","metadata":{"execution":{"iopub.status.busy":"2024-11-19T04:37:57.077040Z","iopub.execute_input":"2024-11-19T04:37:57.077373Z","iopub.status.idle":"2024-11-19T04:38:13.266957Z","shell.execute_reply.started":"2024-11-19T04:37:57.077345Z","shell.execute_reply":"2024-11-19T04:38:13.266073Z"},"papermill":{"duration":13.173775,"end_time":"2024-11-10T22:54:42.126943","exception":false,"start_time":"2024-11-10T22:54:28.953168","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"unique_symbol_ids = sorted(valid_symbol_ids[\"symbol_id\"].unique())","metadata":{"execution":{"iopub.status.busy":"2024-11-19T04:38:16.236478Z","iopub.execute_input":"2024-11-19T04:38:16.237308Z","iopub.status.idle":"2024-11-19T04:38:16.278460Z","shell.execute_reply.started":"2024-11-19T04:38:16.237274Z","shell.execute_reply":"2024-11-19T04:38:16.277763Z"},"papermill":{"duration":0.050044,"end_time":"2024-11-10T22:54:42.183385","exception":false,"start_time":"2024-11-10T22:54:42.133341","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 3. Modeling\n**Model Definition**: Simple feedforward neural network with three hidden layers.","metadata":{"papermill":{"duration":0.006361,"end_time":"2024-11-10T22:54:42.196550","exception":false,"start_time":"2024-11-10T22:54:42.190189","status":"completed"},"tags":[]}},{"cell_type":"code","source":"class WeightedMSELoss(nn.Module):\n    def __init__(self):\n        super(WeightedMSELoss, self).__init__()\n\n    def forward(self, predictions, targets, weights):\n        # Ensure that predictions, targets, and weights are of the same shape\n        assert predictions.shape == targets.shape == weights.shape, \"Shape mismatch among inputs\"\n\n        # Calculate the squared differences\n        squared_diff = (predictions - targets) ** 2\n        \n        # Apply weights to the squared differences\n        weighted_squared_diff = weights * squared_diff\n        \n        # Take the mean of the weighted squared differences\n        loss = weighted_squared_diff.mean()\n        return loss\n\n\n\nclass JaneStreetModel(nn.Module):\n    def __init__(self, num_features, unique_symbol_ids, embedding_dim=16):\n        super(JaneStreetModel, self).__init__()\n        input_dim = num_features + (embedding_dim if CFG.use_symbol_id_embedding else 0)\n        self.fc1 = nn.Linear(input_dim, 512)\n        self.fc2 = nn.Linear(512, 128)\n        self.fc3 = nn.Linear(128, 1)\n        self.relu = nn.ReLU()\n        if CFG.use_symbol_id_embedding:\n            self.embedding = nn.Embedding(len(unique_symbol_ids), embedding_dim)\n\n    def forward(self, x, symbol_ids):\n        if CFG.use_symbol_id_embedding:\n            symbol_embeds = self.embedding(symbol_ids).squeeze(1)\n\n            x = torch.cat([x, symbol_embeds], dim=1)\n        x = self.fc1(x)\n        x = self.relu(x)\n        x = self.fc2(x)\n        x = self.relu(x)\n        return self.fc3(x)","metadata":{"execution":{"iopub.status.busy":"2024-11-19T04:38:18.491660Z","iopub.execute_input":"2024-11-19T04:38:18.492026Z","iopub.status.idle":"2024-11-19T04:38:18.500339Z","shell.execute_reply.started":"2024-11-19T04:38:18.491997Z","shell.execute_reply":"2024-11-19T04:38:18.499461Z"},"papermill":{"duration":0.023403,"end_time":"2024-11-10T22:54:42.226600","exception":false,"start_time":"2024-11-10T22:54:42.203197","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model = JaneStreetModel(len(CFG.feature_columns), unique_symbol_ids).to(device)\nprint(model)","metadata":{"execution":{"iopub.status.busy":"2024-11-19T04:38:20.986961Z","iopub.execute_input":"2024-11-19T04:38:20.988042Z","iopub.status.idle":"2024-11-19T04:38:21.181021Z","shell.execute_reply.started":"2024-11-19T04:38:20.987994Z","shell.execute_reply":"2024-11-19T04:38:21.179985Z"},"papermill":{"duration":0.202566,"end_time":"2024-11-10T22:54:42.435741","exception":false,"start_time":"2024-11-10T22:54:42.233175","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 4. Training or loading model\n\n**Generator Function**: Yields batches of data for training and validation.\n\n**R² Calculation**: Custom function to calculate R² score.\n\n**Model Evaluation**: Function to evaluate model on validation data.\n\n**Training Loop**: Trains the model with mixed precision and saves the best model based on R² score.","metadata":{"papermill":{"duration":0.00622,"end_time":"2024-11-10T22:54:42.448575","exception":false,"start_time":"2024-11-10T22:54:42.442355","status":"completed"},"tags":[]}},{"cell_type":"code","source":"def get_generator(X, y, weights, symbol_ids, shuffle=True, batch_size=4096):\n    def generator():\n        # Create a shuffled index for random access\n        if shuffle:\n            indices = np.random.permutation(len(X))\n        else:\n            indices = np.arange(len(X))\n        num_batch = len(indices) // batch_size + (1 if len(indices) % batch_size > 0 else 0)\n        for i in range(num_batch):\n            start_index = i * batch_size\n            end_index = min((i + 1) * batch_size, len(X))\n            current_indices = indices[start_index: end_index]\n            features = X.iloc[current_indices].values\n            label = y.iloc[current_indices].values\n            symbol_id_batch = symbol_ids.iloc[current_indices].values\n            weight_batch = weights.iloc[current_indices].values\n            yield torch.Tensor(features).to(device), torch.tensor(symbol_id_batch, dtype=torch.int).to(device), torch.tensor(weight_batch).to(device), torch.Tensor(label).to(device)\n    return generator\n\ndef calculate_r2(y_true, y_pred, weights):\n    \"\"\"\n    Calculate the sample weighted zero-mean R-squared score (R2).\n\n    Parameters:\n    - y_true (pd.Series or np.array): Ground truth values.\n    - y_pred (pd.Series or np.array): Predicted values.\n    - weights (pd.Series or np.array): Sample weights.\n\n    Returns:\n    - float: R2 score.\n    \"\"\"\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":{"execution":{"iopub.status.busy":"2024-11-19T04:38:23.236795Z","iopub.execute_input":"2024-11-19T04:38:23.237682Z","iopub.status.idle":"2024-11-19T04:38:23.248573Z","shell.execute_reply.started":"2024-11-19T04:38:23.237638Z","shell.execute_reply":"2024-11-19T04:38:23.247576Z"},"papermill":{"duration":0.02057,"end_time":"2024-11-10T22:54:42.475647","exception":false,"start_time":"2024-11-10T22:54:42.455077","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train_model(model, optimizer, criterion, epoch):\n    model.train()\n    train_loss = 0\n    counter = 1\n    train_batch_count = X_train.shape[0] // CFG.batch_size + (1 if X_train.shape[0] % CFG.batch_size > 0 else 0)\n    train_generator = get_generator(X_train, y_train, train_weights, train_symbol_ids, shuffle=True, batch_size=CFG.batch_size)\n    begin_time = time.time()\n    for X_batch, simbol_id_batch, weight_batch, y_batch in train_generator():\n        if counter % 1000 == 0:\n            elapsed = time.time() - begin_time\n            estimated_time = elapsed / counter * train_batch_count\n            print(f\"[Epoch {epoch+1}][Training] {elapsed:.4f}s / {estimated_time: .4f}s\")\n        counter += 1\n        optimizer.zero_grad()\n        # Mixed precision forward and backward passes\n        with torch.amp.autocast(CFG.device_name):\n            y_pred = model(X_batch, simbol_id_batch)\n            loss = criterion(y_pred, y_batch, weight_batch)\n        scaler.scale(loss).backward()\n        scaler.step(optimizer)\n        scaler.update()\n        train_loss += loss.item()\n        \n    train_loss /= train_batch_count\n    return {\n        \"train_loss\": train_loss\n    }\n\ndef evaluate_model(model):\n    criterion = WeightedMSELoss()\n    # Validation step\n    model.eval()\n    val_loss = 0\n    all_y_true = []\n    all_y_pred = []\n    val_batch_count = X_val.shape[0] // CFG.batch_size + (1 if X_val.shape[0] % CFG.batch_size > 0 else 0)\n    begin_time = time.time()\n    counter = 1\n    with torch.no_grad():\n        val_generator = get_generator(X_val, y_val, valid_weights, valid_symbol_ids, shuffle=False, batch_size=CFG.batch_size)\n        for X_batch, symbol_id_batch, weight_batch, y_batch in val_generator():\n            counter += 1\n            with torch.amp.autocast(CFG.device_name):\n                y_pred = model(X_batch, symbol_id_batch)\n                loss = criterion(y_pred, y_batch, weight_batch)\n            val_loss += loss.item()\n                \n            # Collect true and predicted values for R² calculation\n            all_y_true.append(y_batch)\n            all_y_pred.append(y_pred)\n        \n    val_loss /= val_batch_count\n        \n    # Calculate R² for the entire validation set\n    all_y_true = torch.cat(all_y_true, dim=0).cpu().numpy().reshape(-1)\n    all_y_pred = torch.cat(all_y_pred, dim=0).cpu().numpy().reshape(-1)\n    val_r2 = calculate_r2(all_y_true, all_y_pred, weights)\n    return {\n        \"val_loss\": val_loss,\n        \"val_r2\": val_r2\n    }","metadata":{"execution":{"iopub.status.busy":"2024-11-19T04:38:26.534987Z","iopub.execute_input":"2024-11-19T04:38:26.535632Z","iopub.status.idle":"2024-11-19T04:38:26.545093Z","shell.execute_reply.started":"2024-11-19T04:38:26.535590Z","shell.execute_reply":"2024-11-19T04:38:26.544231Z"},"papermill":{"duration":0.022571,"end_time":"2024-11-10T22:54:42.504935","exception":false,"start_time":"2024-11-10T22:54:42.482364","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"best_model_path = \"model.pth\"\nmodels = []\nif CFG.is_training:\n    model = JaneStreetModel(79, unique_symbol_ids).to(device)\n    criterion = WeightedMSELoss()\n    optimizer = optim.Adam(model.parameters(), lr=0.0001)\n    scaler = torch.amp.GradScaler(CFG.device_name)\n    # Variable to keep track of the best R² score\n    best_r2 = -float('inf')  # Initialize to negative infinity to ensure any positive R² will be better\n    # Training loop with mixed precision and R² tracking\n    num_epochs = 30\n    best_epoch = 0\n    early_stopping_round = 5\n    for epoch in range(num_epochs):\n        train_results = train_model(model, optimizer, criterion, epoch)\n        train_loss = train_results[\"train_loss\"]\n        # Validation step\n        results = evaluate_model(model)\n        val_r2 = results[\"val_r2\"]\n        val_loss = results[\"val_loss\"]\n        # Check if this is the best model so far and save it\n        torch.save(model.state_dict(), f\"model_{epoch}.pth\")\n        if val_r2 > best_r2:\n            best_epoch = epoch\n            best_r2 = val_r2\n            torch.save(model.state_dict(), best_model_path)\n            print(f\"New best model saved with R²: {best_r2:.4f}\")\n        else:\n            if epoch - best_epoch > early_stopping_round:\n                print(\"Model stops improving, stop training.\")\n                break\n        print(f\"Epoch {epoch+1}/{num_epochs}, Train Loss: {train_loss:.4f}, Validation Loss: {val_loss:.4f}, Validation R²: {val_r2:.4f}\")\n    models.append(model)\nelse:\n    # Load models from best epochs\n    for epoch in [2, 6]:\n        print(\"=\" * 30)\n        print(f\"Loading model from epoch {epoch + 1}\")\n        print(\"=\" * 30)\n        model = JaneStreetModel(79, unique_symbol_ids).to(device)\n        base_model_path = f\"/kaggle/input/jane-street-pytorch-rmf-model/model_{epoch}.pth\"\n        model.load_state_dict(torch.load(base_model_path, weights_only=True))\n        results = evaluate_model(model)\n        val_r2 = results[\"val_r2\"]\n        print(f\"Validation R2: {val_r2:.4f}\")\n        models.append(model)","metadata":{"execution":{"iopub.status.busy":"2024-11-19T04:41:23.690617Z","iopub.execute_input":"2024-11-19T04:41:23.691467Z","iopub.status.idle":"2024-11-19T04:41:35.784902Z","shell.execute_reply.started":"2024-11-19T04:41:23.691434Z","shell.execute_reply":"2024-11-19T04:41:35.783975Z"},"papermill":{"duration":7.733837,"end_time":"2024-11-10T22:54:50.245220","exception":false,"start_time":"2024-11-10T22:54:42.511383","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 5. Create submission pipeline\n\n\n**Predict Function**: Handles inference on test data, returns predictions in required format.\n\n**Inference Server**: Sets up the inference server for handling prediction requests during evaluation.","metadata":{"papermill":{"duration":0.006444,"end_time":"2024-11-10T22:54:50.258265","exception":false,"start_time":"2024-11-10T22:54:50.251821","status":"completed"},"tags":[]}},{"cell_type":"markdown","source":"The evaluation API requires that you set up a server which will respond to inference requests. We have already defined the server; you just need write the predict function. When we evaluate your submission on the hidden test set the client defined in `jane_street_gateway` will run in a different container with direct access to the hidden test set and hand off the data timestep by timestep.\n\n\n\nYour code will always have access to the published copies of the files.","metadata":{"papermill":{"duration":0.006277,"end_time":"2024-11-10T22:54:50.271005","exception":false,"start_time":"2024-11-10T22:54:50.264728","status":"completed"},"tags":[]}},{"cell_type":"code","source":"def make_prediction(models, X, symbol_ids):\n    y_preds = []\n    with torch.no_grad():\n        with torch.amp.autocast(CFG.device_name):\n            for model in models:\n                y_pred = model(torch.Tensor(X).to(device), torch.tensor(symbol_ids, dtype=torch.int).to(device))\n                y_preds.append(y_pred.cpu().numpy().flatten())\n    return np.mean(y_preds, axis=0)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-19T04:44:09.826504Z","iopub.execute_input":"2024-11-19T04:44:09.826882Z","iopub.status.idle":"2024-11-19T04:44:09.832512Z","shell.execute_reply.started":"2024-11-19T04:44:09.826847Z","shell.execute_reply":"2024-11-19T04:44:09.831526Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"lags_ : pl.DataFrame | None = None\n\n\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 10 minutes of the batch features being provided.\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 lags_\n    if lags is not None:\n        lags_ = lags\n    # 1. Select the required feature columns and convert to numpy array for Keras\n    X_test = test.select(CFG.feature_columns).to_numpy()\n    X_test = np.where(np.isnan(X_test), mean_values, X_test)\n    X_test = (X_test - min_values) / (max_values - min_values)\n    X_test = np.nan_to_num(X_test, nan=0)\n    symbol_ids = test.select([\"symbol_id\"]).to_numpy()\n    y_pred = make_prediction(models, X_test, symbol_ids)\n    # 3. Prepare the DataFrame for output\n    predictions = test.select('row_id').with_columns(\n        pl.Series(\"responder_6\", y_pred)\n    )\n    # The predict function must return a DataFrame\n    assert isinstance(predictions, pl.DataFrame | pd.DataFrame)\n    # with columns 'row_id', 'responer_6'\n    assert predictions.columns == ['row_id', 'responder_6']\n    # and as many rows as the test data.\n    assert len(predictions) == len(test)\n\n    return predictions","metadata":{"_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","execution":{"iopub.status.busy":"2024-11-19T04:44:29.844981Z","iopub.execute_input":"2024-11-19T04:44:29.845956Z","iopub.status.idle":"2024-11-19T04:44:29.852521Z","shell.execute_reply.started":"2024-11-19T04:44:29.845921Z","shell.execute_reply":"2024-11-19T04:44:29.851746Z"},"papermill":{"duration":0.019434,"end_time":"2024-11-10T22:54:50.296835","exception":false,"start_time":"2024-11-10T22:54:50.277401","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"When your notebook is run on the hidden test set, inference_server.serve must be called within 15 minutes of the notebook starting or the gateway will throw an error. If you need more than 15 minutes to load your model you can do so during the very first `predict` call, which does not have the usual 10 minute response deadline.","metadata":{"papermill":{"duration":0.006419,"end_time":"2024-11-10T22:54:50.309880","exception":false,"start_time":"2024-11-10T22:54:50.303461","status":"completed"},"tags":[]}},{"cell_type":"code","source":"inference_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":{"_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","execution":{"iopub.status.busy":"2024-11-19T04:44:33.103983Z","iopub.execute_input":"2024-11-19T04:44:33.104859Z"},"papermill":{"duration":0.31963,"end_time":"2024-11-10T22:54:50.635948","exception":false,"start_time":"2024-11-10T22:54:50.316318","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null}]}