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