Mawadaa/Pilot_Cognitive_Workload
0
1import gradio as gr2import matplotlib3import matplotlib.pyplot as plt4import numpy as np5import joblib6import shap7import json8import tensorflow as tf9import warnings10warnings.filterwarnings("ignore")11 12matplotlib.use("Agg")13 14# ── Constants ─────────────────────────────────────────────────15CNN_CHANNEL_NAMES = [16 "EEG_FP1", "EEG_F7", "EEG_F8", "EEG_Fz", "EEG_T3",17 "EEG_T4", "EEG_P3", "EEG_Pz", "EEG_O1", "ECG", "GSR"18]19 20# ── Load XGBoost artifacts ────────────────────────────────────21xgb_model = joblib.load("xgb_model.pkl")22scaler = joblib.load("scaler.pkl")23X_train = np.load("X_train.npy")24feature_means = np.load("feature_means.npy")25feature_stds = np.load("feature_stds.npy")26best_workload_full = np.load("best_workload.npy")27best_baseline_full = np.load("best_baseline.npy")28WORKLOAD_PRESET = np.load("workload_preset.npy").tolist()29BASELINE_PRESET = np.load("baseline_preset.npy").tolist()30 31with open("features.json") as f:32 feat_data = json.load(f)33 34ALL_FEATURES = feat_data["ALL_FEATURES"]35TOP_UI_FEATURES = feat_data["TOP_UI_FEATURES"]36TOP_UI_INDICES = feat_data["TOP_UI_INDICES"]37BEST_T = feat_data["BEST_T"]38 39# ── CNN-LSTM custom loss + load ───────────────────────────────40def loss_fn(y_true, y_pred):41 bce = tf.keras.losses.BinaryCrossentropy()42 return bce(y_true, y_pred)43 44CNN_THRESHOLD = 0.0445n_cnn_features = 1146seq_len = 128047cnn_feat_mins = [-5000.0] * 1148cnn_feat_maxs = [5000.0] * 1149cnn_feat_means = [0.0] * 1150CNN_LOADED = False51cnn_model = None52 53try:54 import os55 print(f"Files: {os.listdir('/app')}")56 try:57 cnn_model = tf.keras.models.load_model(58 "cnn_lstm_model.keras",59 custom_objects={"loss_fn": loss_fn}60 )61 print("Loaded .keras")62 except Exception as e1:63 print(f"keras failed: {e1}")64 cnn_model = tf.keras.models.load_model(65 "cnn_lstm_model.h5",66 custom_objects={"loss_fn": loss_fn}67 )68 print("Loaded .h5")69 70 with open("feature_stats.json") as f:71 stats = json.load(f)72 73 n_cnn_features = stats.get("n_cnn_features", 11)74 seq_len = stats.get("seq_len", 1280)75 cnn_feat_mins = stats.get("mins", [-5000.0] * n_cnn_features)76 cnn_feat_maxs = stats.get("maxs", [5000.0] * n_cnn_features)77 cnn_feat_means = stats.get("means", [0.0] * n_cnn_features)78 CNN_LOADED = True79 print(f"CNN ready | features={n_cnn_features} | seq={seq_len}")80 81except Exception as e:82 import traceback83 print(f"CNN FAILED: {e}")84 print(traceback.format_exc())85 86# ── SHAP explainer ────────────────────────────────────────────87explainer_shap = shap.TreeExplainer(xgb_model)88current_background = {"vector": best_workload_full.copy()}89 90# ── XGBoost functions ─────────────────────────────────────────91def predict_workload(*args):92 x = current_background["vector"].copy()93 for feat_idx, val in zip(TOP_UI_INDICES, args):94 x[feat_idx] = float(val)95 prob = xgb_model.predict_proba(x.reshape(1, -1))[0, 1]96 if prob >= 0.70:97 risk = "🔴 HIGH WORKLOAD — Intervention recommended"98 elif prob >= BEST_T:99 risk = "🟡 ELEVATED WORKLOAD — Monitor pilot"100 else:101 risk = "🟢 BASELINE — Normal cognitive load"102 pred = "Workload" if prob >= BEST_T else "Baseline"103 return f"{prob:.3f}", risk, f"{pred} (threshold={BEST_T:.2f})"104 105def explain_prediction(*args):106 x = current_background["vector"].copy()107 for feat_idx, val in zip(TOP_UI_INDICES, args):108 x[feat_idx] = float(val)109 shap_vals = explainer_shap.shap_values(x.reshape(1, -1))[0]110 top_idx = np.argsort(np.abs(shap_vals))[::-1][:10]111 feat_names = [ALL_FEATURES[i] for i in top_idx]112 feat_shap = shap_vals[top_idx]113 colors = ["#FF0000" if v > 0 else "#378ADD" for v in feat_shap]114 fig, ax = plt.subplots(figsize=(9, 6))115 ax.barh(feat_names[::-1], feat_shap[::-1], color=colors[::-1], edgecolor="none")116 ax.axvline(0, color="black", linewidth=0.8)117 prob = xgb_model.predict_proba(x.reshape(1, -1))[0, 1]118 ax.set_title(119 f"SHAP Explanation — Workload Probability: {prob:.3f}\n"120 f"🔴 Red = pushes toward Workload | 🔵 Blue = pushes toward Baseline",121 fontsize=11)122 ax.set_xlabel("SHAP Value")123 ax.spines[["top", "right"]].set_visible(False)124 plt.tight_layout()125 return fig126 127def load_workload():128 current_background["vector"] = best_workload_full.copy()129 return WORKLOAD_PRESET130 131def load_baseline():132 current_background["vector"] = best_baseline_full.copy()133 return BASELINE_PRESET134 135def make_slider(feat_idx, feat_name, default_val):136 mu = float(feature_means[feat_idx])137 sd = float(feature_stds[feat_idx])138 return gr.Slider(139 minimum=round(mu - 3 * sd, 4),140 maximum=round(mu + 3 * sd, 4),141 value=round(default_val, 4),142 step=round(max(sd / 20, 1e-5), 5),143 label=feat_name,144 info=f"Mean={mu:.3f} | Std={sd:.3f}"145 )146 147# ── CNN-LSTM functions ────────────────────────────────────────148def predict_cnn(*args):149 if not CNN_LOADED or cnn_model is None:150 return "CNN-LSTM model not loaded", "N/A", None151 152 x_row = np.array([float(a) for a in args], dtype=float)153 x_seq = np.tile(x_row, (seq_len, 1))[np.newaxis, :, :]154 x_tf = tf.constant(x_seq, dtype=tf.float32)155 156 with tf.GradientTape() as tape:157 tape.watch(x_tf)158 pred = cnn_model(x_tf, training=False)159 160 prob = float(pred.numpy().flatten()[0])161 162 if prob >= 0.70:163 risk = "🔴 HIGH WORKLOAD — Intervention recommended"164 elif prob >= CNN_THRESHOLD:165 risk = "🟡 ELEVATED WORKLOAD — Monitor pilot"166 else:167 risk = "🟢 BASELINE — Normal cognitive load"168 169 grads = tape.gradient(pred, x_tf)170 if grads is None:171 feat_sal = np.ones(n_cnn_features)172 else:173 feat_sal = tf.abs(grads).numpy()[0].mean(axis=0)174 175 n = min(len(feat_sal), len(CNN_CHANNEL_NAMES))176 pairs = sorted(zip(CNN_CHANNEL_NAMES[:n], feat_sal[:n]), key=lambda x: x[1])177 names = [p[0] for p in pairs]178 values = [p[1] for p in pairs]179 180 fig, ax = plt.subplots(figsize=(9, 6))181 ax.barh(names, values, color="#E74C3C", alpha=0.85, edgecolor="none")182 ax.set_xlabel("Mean |Gradient| — higher = CNN-LSTM focused here", fontsize=10)183 ax.set_title(184 f"CNN-LSTM Gradient Saliency\n"185 f"{risk} | Prob={prob:.3f} | Threshold={CNN_THRESHOLD}",186 fontsize=11, fontweight="bold"187 )188 ax.spines[["top", "right"]].set_visible(False)189 plt.tight_layout()190 return f"{prob:.3f}", risk, fig191 192def make_cnn_slider(i):193 mu = float(cnn_feat_means[i])194 mn = float(cnn_feat_mins[i])195 mx = float(cnn_feat_maxs[i])196 rng = mx - mn197 return gr.Slider(198 minimum=round(mn, 4),199 maximum=round(mx, 4),200 value=round(mu, 4),201 step=round(max(rng / 100, 1e-5), 5),202 label=CNN_CHANNEL_NAMES[i],203 info=f"Mean={mu:.3f}"204 )205 206# ── Build app ─────────────────────────────────────────────────207with gr.Blocks(title="✈️ Pilot Cognitive Workload Monitor", theme=gr.themes.Soft()) as demo:208 209 gr.Markdown("""210 # ✈️ Pilot Cognitive Workload Monitor211 **EEG-Based Real-Time Detection | NASA LOFT Dataset | Emirates Aviation University**212 > Load a preset or adjust sliders, then click Predict or Explain.213 214 | Model | Architecture | AUC | Accuracy | XAI Method |215 |-------|-------------|-----|----------|------------|216 | Tab 1 & 2 | XGBoost (31 engineered features) | 0.764 | 82.9% | SHAP TreeExplainer |217 | Tab 3 | CNN-LSTM (raw EEG sequences) | 0.729 | 81.2% | Gradient Saliency |218 """)219 220 # ── TAB 1 ─────────────────────────────────────────────────221 with gr.Tab("Tab 1: Real-Time Workload Prediction"):222 gr.Markdown("### Input EEG Features (Top 8 by SHAP importance)\nClick a preset to load real EEG values, or adjust sliders manually.")223 with gr.Row():224 btn_workload = gr.Button("Load Workload Example (prob=0.985)", variant="secondary")225 btn_baseline = gr.Button("Load Baseline Example", variant="secondary")226 with gr.Row():227 with gr.Column(scale=2):228 sliders_1 = [229 make_slider(idx, name, WORKLOAD_PRESET[i])230 for i, (idx, name) in enumerate(zip(TOP_UI_INDICES, TOP_UI_FEATURES))231 ]232 with gr.Column(scale=1):233 gr.Markdown("### Prediction Results")234 out_prob = gr.Textbox(label="Workload Probability (0–1)", interactive=False)235 out_risk = gr.Textbox(label="Risk Level", interactive=False)236 out_class = gr.Textbox(label="Classification", interactive=False)237 gr.Markdown("---\n**Risk Levels:**\n- 🟢 < threshold → Baseline\n- 🟡 threshold–0.70 → Elevated\n- 🔴 > 0.70 → High Workload")238 predict_btn = gr.Button("🔍 Predict Workload Level", variant="primary", size="lg")239 predict_btn.click(fn=predict_workload, inputs=sliders_1, outputs=[out_prob, out_risk, out_class])240 btn_workload.click(fn=load_workload, outputs=sliders_1)241 btn_baseline.click(fn=load_baseline, outputs=sliders_1)242 gr.Markdown(f"---\n**Decision threshold:** {BEST_T:.2f} | **Workload F1:** 0.370 | **Accuracy:** 82.9%")243 244 # ── TAB 2 ─────────────────────────────────────────────────245 with gr.Tab("Tab 2: XAI Explanation (SHAP)"):246 gr.Markdown("""247 ### Why did the model make this prediction?248 - 🔴 **Red bars** → push toward **Workload**249 - 🔵 **Blue bars** → push toward **Baseline**250 """)251 with gr.Row():252 btn_workload_2 = gr.Button("Load Workload Example", variant="secondary")253 btn_baseline_2 = gr.Button("Load Baseline Example", variant="secondary")254 sliders_2 = [255 make_slider(idx, name, WORKLOAD_PRESET[i])256 for i, (idx, name) in enumerate(zip(TOP_UI_INDICES, TOP_UI_FEATURES))257 ]258 explain_btn = gr.Button("Generate SHAP Explanation", variant="secondary", size="lg")259 shap_plot = gr.Plot(label="SHAP Feature Attribution")260 explain_btn.click(fn=explain_prediction, inputs=sliders_2, outputs=shap_plot)261 btn_workload_2.click(fn=load_workload, outputs=sliders_2)262 btn_baseline_2.click(fn=load_baseline, outputs=sliders_2)263 gr.Markdown("---\n**Top SHAP features:** EEG_T3_beta · EEG_T4_diff_mean · EEG_P3_diff_mean · ECG_std · GSR_std\n**Glass Box validated** — SHAP + LIME cross-confirmed ✓")264 265 # ── TAB 3 ─────────────────────────────────────────────────266 with gr.Tab("Tab 3: CNN-LSTM + Gradient Saliency"):267 gr.Markdown(f"""268 ### CNN-LSTM Raw Sequence Classifier269 **Model B** | AUC = 0.729 | Accuracy = 81.2% | Threshold = {CNN_THRESHOLD}270 271 Uses **11 raw EEG/physiological channel amplitudes** as a 5-second temporal272 sequence (1,280 timesteps at 256 Hz). Gradient saliency shows which channels273 the CNN-LSTM focused on — complementary to XGBoost SHAP 274 275 > Adjust the EEG channel sliders below and click **Predict** to get a real-time CNN-LSTM assessment.276 """)277 with gr.Row():278 with gr.Column(scale=2):279 gr.Markdown("#### Raw EEG Channel Inputs")280 cnn_sliders = [make_cnn_slider(i) for i in range(n_cnn_features)]281 with gr.Column(scale=1):282 gr.Markdown("### Prediction Results")283 cnn_out_prob = gr.Textbox(label="Workload Probability (0–1)", interactive=False)284 cnn_out_risk = gr.Textbox(label="Risk Level", interactive=False)285 cnn_plot = gr.Plot(label="🔍 Gradient Saliency — Channel Importance")286 gr.Markdown("---\n**How to read:** Longer bar = CNN-LSTM focused more on that channel.\nGSR & EEG_FP1 dominate raw sequences — complementary to XGBoost's EEG_T3_beta")287 cnn_btn = gr.Button("Predict Workload (CNN-LSTM)", variant="primary", size="lg")288 cnn_btn.click(fn=predict_cnn, inputs=cnn_sliders, outputs=[cnn_out_prob, cnn_out_risk, cnn_plot])289 gr.Markdown("---\n**Architecture:** Conv1D(64) → Conv1D(128) → LSTM(64) → Dense(1)\n**Training:** Weighted BCE (pos_weight=7.71) · Early stopping · ReduceLR\n**LOSO:** Mean AUC = 0.547")290 291 # ── Footer ────────────────────────────────────────────────292 gr.Markdown("""293 ---294 **Dataset:** NASA LOFT EEG | 10 pilot sessions | 10.3M rows | 256 Hz295 **Pipeline:** PySpark ETL (6.7× speedup) → XGBoost + CNN-LSTM → SHAP + Gradient Saliency296 **Module EM06DS** — Big Data Analytics | Emirates Aviation University | MSc AI & Data Science297 """)298 299demo.launch(server_name="0.0.0.0", server_port=7860)