CoolFace
Apppublic

Mawadaa/Pilot_Cognitive_Workload

sourceHugging Facemitupdated 5mo agoView on Hugging Face
0likes
app.py299 linesDownload Raw Back to root
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)