{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","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"}],"dockerImageVersionId":30822,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"Since the files are very large, it is more reasonable to do a batch-by-batch analysis because we cannot analyze all the files together. It seems that setting batch_size = 500,000 is a way to process without loading the RAM, unless for ML models where I take batch_size = 100,000. Here I will present an example of this.\n\nIn this case, I observed NaN values and correlations. Then, I exemplified a CatBoost model that will work batch by batch.\n\n1. In correlations, other responders are of utmost importance; therefore, I am not sure if we see them when testing.\n2. Stock prices do not appear in the correlations, which is expected because only the raw close data is visible. This needs to be processed. So I saw that responder_6 has attention on certain stock returns (as expected).\n3. I grouped by stocks to know which stocks the responders react more to.\n4. Later, I tried to predict using CatBoost as an example. From here, predictions can be made with more advanced models based on stocks and other responders.","metadata":{}},{"cell_type":"markdown","source":"# Preliminaries","metadata":{}},{"cell_type":"markdown","source":"## First batch","metadata":{}},{"cell_type":"code","source":"import polars as pl\nimport pandas as pd\n\n# Define the path to the parquet file\nfile_path = '../input/jane-street-real-time-market-data-forecasting/train.parquet'\n\n# Load data lazily with Polars\ndata = pl.scan_parquet(file_path)\n\n# Sort data by date_id and time_id\nsorted_data = data.sort(['date_id', 'time_id'])\n\n# Define batch size for processing\nbatch_size = 500_000  # Size of each batch\n\n# Define start and end indices for the first batch\nstart_idx = 0\nend_idx = batch_size\n\nprint(\"Analyzing the first batch...\")\n\n# Collect the first batch using Polars\nfirst_batch_polars = sorted_data.slice(start_idx, end_idx - start_idx).collect()\n\n# Convert the Polars DataFrame to Pandas for detailed analysis\nfirst_batch = first_batch_polars.to_pandas()\n\n# Check for NaN, non-numeric, and other unexpected values\nnan_columns = {}\nunexpected_details = {}\n\nfor col in first_batch.columns:\n    if pd.api.types.is_numeric_dtype(first_batch[col]):  # Process only numeric columns\n        # Check for NaN values\n        nan_count = first_batch[col].isna().sum()\n        if nan_count > 0:\n            nan_columns[col] = nan_count\n        \n        # Check for values that are neither numeric nor NaN\n        unexpected_values = first_batch[~first_batch[col].apply(lambda x: pd.isna(x) or isinstance(x, (int, float)))]\n        if not unexpected_values.empty:\n            unexpected_details[col] = unexpected_values[col].unique().tolist()\n\n# Display NaN results\nif nan_columns:\n    print(\"\\nColumns with NaN values:\")\n    for col, count in nan_columns.items():\n        print(f\"Column '{col}' has {count} NaN values.\")\nelse:\n    print(\"No NaN values found in numeric columns of the first batch.\")\n\n# Display unexpected values\nif unexpected_details:\n    print(\"\\nUnexpected Values Found in Numeric Columns:\")\n    for col, values in unexpected_details.items():\n        print(f\"Column '{col}': {values}\")\nelse:\n    print(\"No unexpected values (other than numeric or NaN) found in numeric columns of the first batch.\")\n\n# Print sample rows from the batch\nprint(\"\\nSample rows from the first batch:\")\nprint(first_batch.sample(2))\nprint(\"\\n\" + \"=\" * 80 + \"\\n\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-02T00:13:32.166711Z","iopub.execute_input":"2025-01-02T00:13:32.167209Z","iopub.status.idle":"2025-01-02T00:15:17.147242Z","shell.execute_reply.started":"2025-01-02T00:13:32.167177Z","shell.execute_reply":"2025-01-02T00:15:17.145801Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Second Batch","metadata":{}},{"cell_type":"code","source":"import polars as pl\nimport pandas as pd\n\n# Define the path to the parquet file\nfile_path = '../input/jane-street-real-time-market-data-forecasting/train.parquet'\n\n# Load data lazily with Polars\ndata = pl.scan_parquet(file_path)\n\n# Sort data by date_id and time_id\nsorted_data = data.sort(['date_id', 'time_id'])\n\n# Define batch size for processing\nbatch_size = 500_000  # Size of each batch\n\n# Define start and end indices for the second batch\nstart_idx = batch_size\nend_idx = start_idx + batch_size\n\nprint(\"Analyzing the second batch...\")\n\n# Collect the second batch using Polars\nsecond_batch_polars = sorted_data.slice(start_idx, end_idx - start_idx).collect()\n\n# Convert the Polars DataFrame to Pandas for detailed analysis\nsecond_batch = second_batch_polars.to_pandas()\n\n# Check for NaN, non-numeric, and other unexpected values\nnan_columns = {}\nunexpected_details = {}\n\nfor col in second_batch.columns:\n    if pd.api.types.is_numeric_dtype(second_batch[col]):  # Process only numeric columns\n        # Check for NaN values\n        nan_count = second_batch[col].isna().sum()\n        if nan_count > 0:\n            nan_columns[col] = nan_count\n        \n        # Check for values that are neither numeric nor NaN\n        unexpected_values = second_batch[~second_batch[col].apply(lambda x: pd.isna(x) or isinstance(x, (int, float)))]\n        if not unexpected_values.empty:\n            unexpected_details[col] = unexpected_values[col].unique().tolist()\n\n# Display NaN results\nif nan_columns:\n    print(\"\\nColumns with NaN values:\")\n    for col, count in nan_columns.items():\n        print(f\"Column '{col}' has {count} NaN values.\")\nelse:\n    print(\"No NaN values found in numeric columns of the second batch.\")\n\n# Display unexpected values\nif unexpected_details:\n    print(\"\\nUnexpected Values Found in Numeric Columns:\")\n    for col, values in unexpected_details.items():\n        print(f\"Column '{col}': {values}\")\nelse:\n    print(\"No unexpected values (other than numeric or NaN) found in numeric columns of the second batch.\")\n\n# Print sample rows from the batch\nprint(\"\\nSample rows from the second batch:\")\nprint(second_batch.sample(2))\nprint(\"\\n\" + \"=\" * 80 + \"\\n\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-02T13:03:46.929663Z","iopub.execute_input":"2025-01-02T13:03:46.931061Z","iopub.status.idle":"2025-01-02T13:05:05.642007Z","shell.execute_reply.started":"2025-01-02T13:03:46.930979Z","shell.execute_reply":"2025-01-02T13:05:05.639826Z"},"collapsed":true,"jupyter":{"outputs_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## All patches","metadata":{}},{"cell_type":"raw","source":"","metadata":{}},{"cell_type":"markdown","source":"# Important correlations of the batches","metadata":{}},{"cell_type":"markdown","source":"## First Batch","metadata":{}},{"cell_type":"code","source":"import polars as pl\nimport pandas as pd\n\n# Define the path to the parquet file\nfile_path = '../input/jane-street-real-time-market-data-forecasting/train.parquet'\n\n# Load data lazily with Polars\ndata = pl.scan_parquet(file_path)\n\n# Sort data by date_id and time_id\nsorted_data = data.sort(['date_id', 'time_id'])\n\n# Define batch size for processing\nbatch_size = 500_000  # Size of each batch\n\n# Define start and end indices for the batch to analyze\nstart_idx = 0  # Adjust as needed for other batches\nend_idx = batch_size\n\nprint(f\"Analyzing statistics and correlations for batch {start_idx // batch_size + 1}...\")\n\n# Collect the batch using Polars\nbatch_polars = sorted_data.slice(start_idx, end_idx - start_idx).collect()\n\n# Convert the Polars DataFrame to Pandas for detailed analysis\nbatch = batch_polars.to_pandas()\n\n# if needed:\n# Compute basic statistics\n# print(\"\\nBasic Statistics:\")\n# basic_stats = batch.describe(include='all').T\n# basic_stats['NaN Count'] = batch.isna().sum()\n# basic_stats['Unique Count'] = batch.nunique()\n# print(basic_stats.head(10))  # Print only the first 10 rows for brevity\n\n# Calculate correlations and filter significant ones\ncorrelation_threshold = 0.3  # Threshold for significant correlations\nprint(\"\\nSignificant Correlations:\")\n\ncorrelations = batch.corr(method='pearson')  # Compute the correlation matrix\nif 'responder_6' in correlations.columns:\n    responder_6_corr = correlations['responder_6'].sort_values(ascending=False)\n    significant_corr = responder_6_corr[responder_6_corr.abs() > correlation_threshold]\n    print(\"\\nCorrelations with responder_6 (|correlation| > 0.3):\")\n    print(significant_corr)\nelse:\n    print(\"Column 'responder_6' not found in the data.\")\n\n# if needed:\n# Mask to show only significant correlations in the entire matrix\n# significant_matrix = correlations.where(correlations.abs() > correlation_threshold)\n# significant_matrix = significant_matrix.dropna(how='all', axis=1).dropna(how='all', axis=0)\n# print(\"\\nSignificant Correlation Matrix (|correlation| > 0.3):\")\n# print(significant_matrix)\n\n\n# Save filtered results to CSV for further exploration\n# basic_stats.to_csv('filtered_basic_statistics.csv', index=True)\n# significant_corr.to_csv('responder_6_significant_correlations.csv', index=True)\n# significant_matrix.to_csv('filtered_correlation_matrix.csv', index=True)\n# print(\"\\nFiltered statistics and correlations saved to:\")\n# print(\"  - 'filtered_basic_statistics.csv'\")\n# print(\"  - 'responder_6_significant_correlations.csv'\")\n# print(\"  - 'filtered_correlation_matrix.csv'\")\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-01-02T14:15:19.497632Z","iopub.execute_input":"2025-01-02T14:15:19.498304Z","iopub.status.idle":"2025-01-02T14:16:31.658427Z","shell.execute_reply.started":"2025-01-02T14:15:19.498230Z","shell.execute_reply":"2025-01-02T14:16:31.657136Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"","metadata":{}},{"cell_type":"markdown","source":"## For custom batches:","metadata":{}},{"cell_type":"code","source":"import polars as pl\nimport pandas as pd\n\nfile_path = '../input/jane-street-real-time-market-data-forecasting/train.parquet'\ndata = pl.scan_parquet(file_path)\nsorted_data = data.sort(['date_id', 'time_id'])\n\nbatch_size = 500_000\ncorrelation_threshold = 0.3\nstart_batch = 2\nend_batch = 3\n\nbatch_responder_6_corr = []\n\nfor batch_idx in range(start_batch, end_batch + 1):\n    start_idx = (batch_idx - 1) * batch_size\n    end_idx = start_idx + batch_size\n    batch_polars = sorted_data.slice(start_idx, end_idx - start_idx).collect()\n    batch = batch_polars.to_pandas()\n    correlations = batch.corr(method='pearson')\n    if 'responder_6' in correlations.columns:\n        responder_6_corr = correlations['responder_6'].drop('responder_6', errors='ignore')  # Exclude self-correlation\n        significant_corr = responder_6_corr[responder_6_corr.abs() > correlation_threshold]\n        batch_result = significant_corr.reset_index()\n        batch_result.columns = ['Feature', 'Correlation']\n        batch_result['Batch'] = f'batch_{batch_idx}'\n        batch_responder_6_corr.append(batch_result)\n\ncombined_results = pd.concat(batch_responder_6_corr, ignore_index=True)\n\naggregated_results = combined_results.groupby('Feature').agg(\n    Correlation_Mean=('Correlation', 'mean'),\n    Correlation_Std=('Correlation', 'std'),\n    Batch_Count=('Batch', 'count')\n).reset_index()\n\naggregated_results = aggregated_results[aggregated_results['Correlation_Mean'].abs() > correlation_threshold]\naggregated_results = aggregated_results.sort_values(by='Correlation_Mean', key=abs, ascending=False)\n\nprint(\"\\nImportant Correlations with responder_6 Across Batches:\")\nprint(aggregated_results)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-02T19:31:41.383931Z","iopub.execute_input":"2025-01-02T19:31:41.384259Z","iopub.status.idle":"2025-01-02T19:34:49.089329Z","shell.execute_reply.started":"2025-01-02T19:31:41.384230Z","shell.execute_reply":"2025-01-02T19:34:49.087725Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## All batches","metadata":{}},{"cell_type":"code","source":"import polars as pl\nimport pandas as pd\n\nfile_path = '../input/jane-street-real-time-market-data-forecasting/train.parquet'\ndata = pl.scan_parquet(file_path)\nsorted_data = data.sort(['date_id', 'time_id'])\n\nbatch_size = 500_000\ncorrelation_threshold = 0.3\n\nstart_batch = 2\nbatch_idx = start_batch  # Start processing from batch 2\n\naggregated_results_list = []\n\nwhile True:\n    start_idx = (batch_idx - 1) * batch_size\n    batch = sorted_data.slice(start_idx, batch_size).collect()\n    \n    if batch.height == 0:  # If no data is returned, end the loop\n        print(f\"No more data to process after batch {batch_idx - 1}.\")\n        break\n    \n    print(f\"Processing batch {batch_idx} (starting row {start_idx})...\")\n    \n    # Convert to pandas for correlation calculation\n    batch_df = batch.to_pandas()\n    \n    if 'responder_6' in batch_df.columns:\n        correlations = batch_df.corr()\n        responder_6_corr = correlations['responder_6'].drop('responder_6', errors='ignore')\n        significant_corr = responder_6_corr[abs(responder_6_corr) > correlation_threshold]\n\n        # Store and print results per batch\n        batch_results = []\n        for feature, corr in significant_corr.items():\n            result = {\n                'Feature': feature,\n                'Correlation': corr\n            }\n            batch_results.append(result)\n            aggregated_results_list.append(result)\n        \n        if batch_results:\n            print(\"\\nImportant Correlations with responder_6 for Batch\", batch_idx)\n            batch_df_results = pd.DataFrame(batch_results)\n            print(batch_df_results)\n    \n    batch_idx += 1\n    print(f\"Batch {batch_idx - 1} processed.\\n\")\n\n# Aggregate results\nresults_df = pd.DataFrame(aggregated_results_list)\naggregated_df = results_df.groupby('Feature').agg(\n    Correlation_Mean=('Correlation', 'mean'),\n    Correlation_Std=('Correlation', 'std')).reset_index()\n\naggregated_df = aggregated_df[aggregated_df['Correlation_Mean'].abs() > correlation_threshold]\naggregated_df = aggregated_df.sort_values(by='Correlation_Mean', key=abs, ascending=False)\n\nprint(\"\\nImportant Correlations with responder_6 Across All Batches:\")\nprint(aggregated_df)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-02T23:35:45.156828Z","iopub.execute_input":"2025-01-02T23:35:45.157096Z","execution_failed":"2025-01-02T23:47:20.193Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Stock sensitivity (group by stocks)","metadata":{}},{"cell_type":"markdown","source":"## Raw analysis","metadata":{}},{"cell_type":"code","source":"import polars as pl\nimport pandas as pd\n\nfile_path = '../input/jane-street-real-time-market-data-forecasting/train.parquet'\ndata = pl.scan_parquet(file_path)\nsorted_data = data.sort(['date_id', 'time_id'])\n\nbatch_size = 100_000\nstart_batch = 2\nfinish_batch = 5  # Define the last batch to process\ncorrelation_threshold = 0.3  # Set a threshold for important correlations\n\n# Define the features and the target\nfeatures = [f\"feature_{str(i).zfill(2)}\" for i in range(79)]\ntarget = 'responder_6'\n\n# Initialize correlation results storage\nsymbol_correlation_results = {}\n\nfor batch_idx in range(start_batch, finish_batch + 1):\n    start_idx = (batch_idx - 1) * batch_size\n    batch = sorted_data.slice(start_idx, batch_size).collect()\n\n    if batch.height == 0:\n        print(f\"No more data to process after batch {batch_idx - 1}.\")\n        break\n\n    print(f\"Processing batch {batch_idx} (starting row {start_idx}, {batch.height} rows)...\")\n\n    batch_df = batch.to_pandas()\n    if 'symbol_id' in batch_df.columns and set(features).issubset(batch_df.columns) and target in batch_df.columns:\n        grouped = batch_df.groupby('symbol_id')\n        for symbol, group in grouped:\n            correlation_matrix = group[features + [target]].corr()\n            responder_correlations = correlation_matrix[target].drop(target)\n            important_correlations = responder_correlations[(responder_correlations > correlation_threshold) | (responder_correlations < -correlation_threshold)]\n            \n            if not important_correlations.empty:\n                print(f\"Important Correlations with responder_6 for symbol {symbol} in Batch {batch_idx}:\")\n                print(important_correlations.sort_values(ascending=False))\n                \n            if symbol not in symbol_correlation_results:\n                symbol_correlation_results[symbol] = []\n            symbol_correlation_results[symbol].append(important_correlations)\n\n    print(f\"Batch {batch_idx} processed.\\n\")\n\n    if batch_idx == finish_batch:\n        print(f\"Stopped processing as per finish_batch parameter at batch {finish_batch}.\")\n\n# Aggregate correlation results across batches for each symbol\nprint(\"\\nAggregated Important Correlations with responder_6 across batches:\")\nfor symbol, correlations in symbol_correlation_results.items():\n    final_correlations = pd.concat(correlations, axis=1).mean(axis=1)\n    final_important_correlations = final_correlations[(final_correlations > correlation_threshold) | (final_correlations < -correlation_threshold)]\n    if not final_important_correlations.empty:\n        print(f\"\\nSymbol {symbol}:\")\n        print(final_important_correlations.sort_values(ascending=False))\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-03T14:31:18.437779Z","iopub.execute_input":"2025-01-03T14:31:18.438300Z","iopub.status.idle":"2025-01-03T14:36:46.639708Z","shell.execute_reply.started":"2025-01-03T14:31:18.438264Z","shell.execute_reply":"2025-01-03T14:36:46.638392Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Return correlation","metadata":{}},{"cell_type":"code","source":"import polars as pl\nimport pandas as pd\n\nfile_path = '../input/jane-street-real-time-market-data-forecasting/train.parquet'\ndata = pl.scan_parquet(file_path)\nsorted_data = data.sort(['date_id', 'time_id'])\n\nbatch_size = 100_000\nstart_batch = 2\nfinish_batch = 5  # Define the last batch to process\ncorrelation_threshold = 0.3  # Set a threshold for important correlations\n\n# Define the features and the target\nfeatures = [f\"feature_{str(i).zfill(2)}\" for i in range(79)]\ntarget = 'responder_6'\n\n# Initialize correlation results storage\nstock_correlation_results = {}\n\nfor batch_idx in range(start_batch, finish_batch + 1):\n    start_idx = (batch_idx - 1) * batch_size\n    batch = sorted_data.slice(start_idx, batch_size).collect()\n\n    if batch.height == 0:\n        print(f\"No more data to process after batch {batch_idx - 1}.\")\n        break\n\n    print(f\"Processing batch {batch_idx} (starting row {start_idx}, {batch.height} rows)...\")\n\n    batch_df = batch.to_pandas()\n    if 'symbol_id' in batch_df.columns and set(features).issubset(batch_df.columns) and target in batch_df.columns:\n        # Calculate returns for price-like features\n        for feature in features:\n            batch_df[f'{feature}_return'] = batch_df[feature].pct_change(fill_method=None)  # Handle NA values explicitly\n\n        # Group by symbol_id and calculate correlations\n        grouped = batch_df.groupby('symbol_id')\n        for symbol, group in grouped:\n            return_features = [f\"{feature}_return\" for feature in features if f\"{feature}_return\" in group.columns]\n            correlation_matrix = group[return_features + [target]].corr()\n            responder_correlations = correlation_matrix[target].drop(target)\n            important_correlations = responder_correlations[abs(responder_correlations) > correlation_threshold]\n\n            if not important_correlations.empty:\n                print(f\"Important Correlations with responder_6 for symbol {symbol} in Batch {batch_idx}:\")\n                print(important_correlations.sort_values(ascending=False))\n\n            # Store results by symbol\n            if symbol not in stock_correlation_results:\n                stock_correlation_results[symbol] = []\n            stock_correlation_results[symbol].append(important_correlations)\n\n    print(f\"Batch {batch_idx} processed.\\n\")\n\n    if batch_idx == finish_batch:\n        print(f\"Stopped processing as per finish_batch parameter at batch {finish_batch}.\")\n\n# Aggregate correlation results across batches for each symbol\nprint(\"\\nAggregated Important Correlations with responder_6 across batches:\")\nfor symbol, correlations in stock_correlation_results.items():\n    if correlations:\n        aggregated_correlations = pd.concat(correlations, axis=1).mean(axis=1)\n        final_important_correlations = aggregated_correlations[abs(aggregated_correlations) > correlation_threshold]\n        print(f\"\\nSymbol {symbol}:\")\n        print(final_important_correlations.sort_values(ascending=False))\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-03T20:02:50.963824Z","iopub.execute_input":"2025-01-03T20:02:50.965004Z","iopub.status.idle":"2025-01-03T20:08:07.103249Z","shell.execute_reply.started":"2025-01-03T20:02:50.964952Z","shell.execute_reply":"2025-01-03T20:08:07.100039Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# An example of a Catboost model","metadata":{}},{"cell_type":"code","source":"import polars as pl\nfrom catboost import CatBoostRegressor\nfrom sklearn.metrics import mean_squared_error\n\n# Load the data\nfile_path = '../input/jane-street-real-time-market-data-forecasting/train.parquet'\ndata = pl.scan_parquet(file_path)\nsorted_data = data.sort(['date_id', 'time_id'])\n\n# Set up parameters\nbatch_size = 100_000\nstart_batch = 2\nend_batch = 5  # Define the last batch to process explicitly\n\n# Define the features and the target\nexpected_features = [f\"feature_{str(i).zfill(2)}\" for i in range(79)] + ['symbol_id']\ntarget = 'responder_6'\n\n# Initialize the CatBoost model, specifying that symbol_id is a categorical feature\nmodel = CatBoostRegressor(\n    iterations=100, \n    learning_rate=0.1, \n    depth=6, \n    loss_function='RMSE', \n    cat_features=['symbol_id'],\n    verbose=0\n)\n\n# Prepare validation data\nvalidation_data = sorted_data.filter(pl.col('date_id') > 950)\nsorted_data = sorted_data.filter(pl.col('date_id') <= 950)\nvalidation_df = validation_data.collect().to_pandas()\n\nmodel_performance = []\n\nfor batch_idx in range(start_batch, end_batch + 1):\n    start_idx = (batch_idx - 1) * batch_size\n    batch = sorted_data.slice(start_idx, batch_size).collect()\n\n    if batch.height == 0:\n        print(f\"No more data to process after batch {batch_idx - 1}.\")\n        break\n\n    print(f\"Processing batch {batch_idx} (starting row {start_idx})...\")\n    \n    batch_df = batch.to_pandas()\n    if set(expected_features).issubset(batch_df.columns) and target in batch_df.columns:\n        X_train = batch_df[expected_features]\n        y_train = batch_df[target]\n\n        # Update the model incrementally\n        if batch_idx == start_batch:\n            model.fit(X_train, y_train)\n        else:\n            model.fit(X_train, y_train, init_model=model)\n\n        # Evaluate the model on the validation set\n        y_pred = model.predict(validation_df[expected_features])\n        mse = mean_squared_error(validation_df[target], y_pred)\n        model_performance.append(mse)\n        print(f\"Batch {batch_idx} processed. MSE: {mse:.4f}\")\n\n    else:\n        print(f\"Missing necessary features in batch {batch_idx}. Skipping...\")\n\n# Final evaluation\ny_pred_final = model.predict(validation_df[expected_features])\nfinal_mse = mean_squared_error(validation_df[target], y_pred_final)\nmodel_performance.append(final_mse)\n\nprint(\"Training complete. Model performance across batches:\")\nfor i, mse in enumerate(model_performance, start=start_batch):\n    print(f\"Batch {i}: MSE = {mse:.4f}\")\nprint(f\"Final MSE after all batches: {final_mse:.4f}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-03T15:32:29.780109Z","iopub.execute_input":"2025-01-03T15:32:29.781557Z","iopub.status.idle":"2025-01-03T15:37:29.681906Z","shell.execute_reply.started":"2025-01-03T15:32:29.781486Z","shell.execute_reply":"2025-01-03T15:37:29.679641Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# End","metadata":{}}]}