Spaces:
Build error
Build error
File size: 10,048 Bytes
9890a43 0c60647 9890a43 0c60647 9890a43 0c60647 9890a43 0c60647 9890a43 0c60647 9890a43 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 | # utils_regression.py
# This file contains all helper functions extracted from Notebook_2_Regression_Modeling.ipynb
import streamlit as st
import pandas as pd
import numpy as np
import plotly.graph_objects as go
import joblib
from sklearn.preprocessing import StandardScaler
from sklearn.metrics import mean_absolute_error, mean_squared_error, r2_score
import os
_this_file_dir = os.path.dirname(os.path.abspath(__file__))
_models_dir = os.path.join(_this_file_dir, 'models') # Path to the models folder
# --- DATA & MODEL LOADING ---
@st.cache_data
def get_data_splits(engineered_data, ts_data):
"""Creates consistent train/test splits for all models."""
ts_daily = ts_data.groupby('Date').agg({
'Number of insects': 'sum',
'Average Temperature': 'mean',
'Average Humidity': 'mean'
}).reset_index().sort_values('Date')
train_size = int(0.8 * len(ts_daily))
train_dates = ts_daily['Date'].iloc[:train_size]
ml_train = engineered_data[engineered_data['Date'] <= train_dates.max()].copy()
ml_test = engineered_data[engineered_data['Date'] > train_dates.max()].copy()
ts_train = ts_daily.iloc[:train_size].copy()
ts_test = ts_daily.iloc[train_size:].copy()
return ml_train, ml_test, ts_train, ts_test
@st.cache_resource
def load_regression_models():
"""Loads all pre-trained regression models and the scaler."""
try:
models = {
"ARIMAX": joblib.load(os.path.join(_models_dir, 'arimax_model.joblib')),
"SARIMAX": joblib.load(os.path.join(_models_dir, 'sarimax_model.joblib')),
"Prophet": joblib.load(os.path.join(_models_dir, 'prophet_model.joblib')),
"Random Forest": joblib.load(os.path.join(_models_dir, 'rf_model.joblib')),
"XGBoost": joblib.load(os.path.join(_models_dir, 'xgb_model.joblib')),
"LightGBM": joblib.load(os.path.join(_models_dir, 'lgb_model.joblib'))
}
scaler = joblib.load(os.path.join(_models_dir, 'scaler.joblib'))
return models, scaler
except FileNotFoundError as e:
st.error(f"Model file not found: {e}. Please ensure all .joblib files are in the 'models/' directory.")
return None, None
# --- UTILITY FUNCTIONS ---
def calculate_metrics(y_true, y_pred):
"""Calculate regression metrics."""
mae = mean_absolute_error(y_true, y_pred)
rmse = np.sqrt(mean_squared_error(y_true, y_pred))
r2 = r2_score(y_true, y_pred)
return {'MAE': mae, 'RMSE': rmse, 'R2': r2}
def ensure_non_negative_int(predictions):
"""Ensure predictions are non-negative integers."""
return np.maximum(np.round(predictions), 0).astype(int)
def generate_future_dates(last_date, days=7):
"""Generate future dates safely."""
return pd.date_range(start=last_date, periods=days + 1, freq='D')[1:]
def aggregate_ml_data_for_plotting(ml_data, y_pred):
"""Aggregate ML prediction data by date for clean plotting."""
df = ml_data[['Date', 'Number of insects']].copy()
df['Predicted'] = y_pred
agg_df = df.groupby('Date').agg({'Number of insects': 'sum', 'Predicted': 'sum'}).reset_index().sort_values('Date')
return agg_df
# --- PLOTTING FUNCTIONS ---
def create_continuous_forecast_plot(historical_actual, test_actual, test_pred, future_pred,
dates_hist, dates_test, dates_future, title,
confidence_lower=None, confidence_upper=None,
future_confidence_lower=None, future_confidence_upper=None):
"""
Create a comprehensive and continuous forecast visualization with confidence intervals.
"""
fig = go.Figure()
# Plot historical actual data
fig.add_trace(go.Scatter(x=dates_hist, y=historical_actual, mode='lines', name='Historical Data', line=dict(color='#1f77b4')))
# Plot test period actual data
fig.add_trace(go.Scatter(x=dates_test, y=test_actual, mode='lines+markers', name='Actual (Test Period)', line=dict(color='#2ca02c'), marker=dict(size=6)))
# Plot test period predictions
fig.add_trace(go.Scatter(x=dates_test, y=ensure_non_negative_int(test_pred), mode='lines+markers', name='Test Predictions', line=dict(color='#ff7f0e', dash='dash'), marker=dict(symbol='x', size=6)))
# Plot future forecast predictions
if future_pred is not None and dates_future is not None:
fig.add_trace(go.Scatter(x=dates_future, y=ensure_non_negative_int(future_pred), mode='lines+markers', name='7-Day Forecast', line=dict(color='#d62728', dash='dot'), marker=dict(symbol='star', size=6)))
# Add confidence intervals for test predictions
if confidence_lower is not None and confidence_upper is not None:
fig.add_trace(go.Scatter(x=dates_test, y=ensure_non_negative_int(confidence_upper), mode='lines', line=dict(width=0), showlegend=False))
fig.add_trace(go.Scatter(x=dates_test, y=ensure_non_negative_int(confidence_lower), mode='lines', line=dict(width=0), fill='tonexty', fillcolor='rgba(255, 127, 14, 0.2)', name='95% Confidence (Test)'))
# Add confidence intervals for future forecast
if future_confidence_lower is not None and future_confidence_upper is not None:
fig.add_trace(go.Scatter(x=dates_future, y=ensure_non_negative_int(future_confidence_upper), mode='lines', line=dict(width=0), showlegend=False))
fig.add_trace(go.Scatter(x=dates_future, y=ensure_non_negative_int(future_confidence_lower), mode='lines', line=dict(width=0), fill='tonexty', fillcolor='rgba(214, 39, 40, 0.2)', name='95% Confidence (Forecast)'))
fig.update_layout(title=title, xaxis_title='Date', yaxis_title='Number of Insects', template='plotly_white', height=500, hovermode='x unified', legend=dict(orientation="h", yanchor="bottom", y=1.02, xanchor="right", x=1))
return fig
def generate_ml_confidence_intervals(model, X_data, n_bootstrap=25):
"""Generate confidence intervals for ML models using optimized bootstrap sampling."""
predictions = []
n_samples = X_data.shape[0]
if n_samples == 0:
return np.array([]), np.array([])
for _ in range(n_bootstrap):
bootstrap_indices = np.random.choice(n_samples, size=n_samples, replace=True)
X_bootstrap = X_data[bootstrap_indices]
noise = np.random.normal(0, 0.01, X_bootstrap.shape)
X_noisy = X_bootstrap + noise
pred = model.predict(X_noisy)
predictions.append(pred)
predictions = np.array(predictions)
lower_bound = np.percentile(predictions, 2.5, axis=0)
upper_bound = np.percentile(predictions, 97.5, axis=0)
return ensure_non_negative_int(lower_bound), ensure_non_negative_int(upper_bound)
def create_champion_comparison_plot(full_actual_dates, full_actual_y, test_dates, test_actual_y,
pred1_y, pred1_name, pred2_y, pred2_name, title):
"""Creates a side-by-side plot for two champion models against actuals for the full timeline."""
fig = go.Figure()
# Full historical actuals
fig.add_trace(go.Scatter(x=full_actual_dates, y=full_actual_y, mode='lines', name='Actual Data', line=dict(color='#1f77b4', width=3)))
# Model 1 Predictions
fig.add_trace(go.Scatter(x=test_dates, y=pred1_y, mode='lines', name=pred1_name, line=dict(color='#ff7f0e', dash='dash')))
# Model 2 Predictions
fig.add_trace(go.Scatter(x=test_dates, y=pred2_y, mode='lines', name=pred2_name, line=dict(color='#2ca02c', dash='dot')))
fig.update_layout(title=title, xaxis_title='Date', yaxis_title='Number of Insects', template='plotly_white', height=500, hovermode='x unified', legend=dict(orientation="h", yanchor="bottom", y=1.02, xanchor="right", x=1))
return fig
# This is the main plotting function we will use everywhere
def create_full_forecast_plot(title, train_dates, train_y, test_dates, test_y, test_pred,
future_dates, future_pred, test_ci_lower=None, test_ci_upper=None,
future_ci_lower=None, future_ci_upper=None):
"""A master function to create a complete forecast plot with history, test predictions, and future forecast."""
fig = go.Figure()
# 1. Historical Data
fig.add_trace(go.Scatter(x=train_dates, y=train_y, mode='lines', name='Historical Data', line=dict(color='#1f77b4')))
# 2. Actual Test Data
fig.add_trace(go.Scatter(x=test_dates, y=test_y, mode='lines', name='Actual (Test)', line=dict(color='#2ca02c', width=2.5)))
# 3. Test Predictions
fig.add_trace(go.Scatter(x=test_dates, y=test_pred, mode='lines', name='Predicted (Test)', line=dict(color='#ff7f0e', dash='dash')))
if test_ci_lower is not None and test_ci_upper is not None:
fig.add_trace(go.Scatter(x=test_dates, y=test_ci_upper, mode='lines', line=dict(width=0), showlegend=False))
fig.add_trace(go.Scatter(x=test_dates, y=test_ci_lower, mode='lines', line=dict(width=0), fill='tonexty', fillcolor='rgba(255, 127, 14, 0.2)', name='95% Confidence'))
# 4. Future Forecast
if future_pred is not None:
fig.add_trace(go.Scatter(x=future_dates, y=future_pred, mode='lines', name='7-Day Forecast', line=dict(color='#d62728', dash='dot')))
if future_ci_lower is not None and future_ci_upper is not None:
fig.add_trace(go.Scatter(x=future_dates, y=future_ci_upper, mode='lines', line=dict(width=0), showlegend=False))
fig.add_trace(go.Scatter(x=future_dates, y=future_ci_lower, mode='lines', line=dict(width=0), fill='tonexty', fillcolor='rgba(214, 39, 40, 0.2)', name='Future Confidence'))
fig.update_layout(title=title, xaxis_title='Date', yaxis_title='Number of Insects', template='plotly_white', height=500, hovermode='x unified', legend=dict(orientation="h", yanchor="bottom", y=1.02, xanchor="right", x=1))
return fig
|