"""
SHAP Analysis - 30-Feature Model
=================================
基于30特征Ridge模型的SHAP全局和个案解释分析。
"""

import pandas as pd
import numpy as np
import matplotlib.pyplot as plt
import shap
import joblib
import os
import sys

sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
from fetal_weight_prediction_30features import create_30_features

OUTPUT_DIR = '/Users/wangkai/Desktop/kimi/2/30特征'
SCI_DIR = f'{OUTPUT_DIR}/SCI_Figures'
os.makedirs(SCI_DIR, exist_ok=True)

plt.rcParams['font.family'] = 'Arial'
plt.rcParams['axes.unicode_minus'] = False


def prepare_data():
    df = pd.read_csv('processed_data.csv')
    X = create_30_features(df)
    y = df['birth_weight_g']
    y_class = pd.cut(y, bins=[0, 2500, 4000, 10000], labels=[0, 1, 2]).astype(int)

    from sklearn.model_selection import train_test_split
    X_train, X_test, y_train, y_test, y_class_train, y_class_test = train_test_split(
        X, y, y_class, test_size=0.2, random_state=42, stratify=y_class)

    model = joblib.load(f'{OUTPUT_DIR}/best_model_30features.pkl')
    scaler = joblib.load(f'{OUTPUT_DIR}/scaler_30features.pkl')

    X_train_scaled = scaler.transform(X_train)
    X_test_scaled = scaler.transform(X_test)

    X_train_df = pd.DataFrame(X_train_scaled, columns=X_train.columns)
    X_test_df = pd.DataFrame(X_test_scaled, columns=X_test.columns)

    return X_train_df, X_test_df, y_test.values, y_class_test.values, model


def shap_global_analysis(X_train_df, X_test_df, model):
    print("[INFO] Computing global SHAP values...")
    explainer = shap.LinearExplainer(model, X_train_df)
    shap_values = explainer.shap_values(X_test_df)

    # Save SHAP values
    shap_df = pd.DataFrame(shap_values, columns=X_test_df.columns)
    shap_df.to_csv(f'{OUTPUT_DIR}/shap_values.csv', index=False)
    print(f"[INFO] Saved: {OUTPUT_DIR}/shap_values.csv")

    # Beeswarm (Figure 5)
    fig, ax = plt.subplots(figsize=(10, 12))
    shap.summary_plot(shap_values, X_test_df, plot_type="dot", show=False)
    plt.tight_layout()
    fig.savefig(f'{SCI_DIR}/Fig5_SHAP_Global_Beeswarm.png', dpi=300, bbox_inches='tight')
    fig.savefig(f'{SCI_DIR}/Fig5_SHAP_Global_Beeswarm.pdf', bbox_inches='tight')
    plt.close()
    print("[INFO] Saved Fig5_SHAP_Global_Beeswarm")

    # Bar (Supplementary S2)
    fig, ax = plt.subplots(figsize=(10, 12))
    shap.summary_plot(shap_values, X_test_df, plot_type="bar", show=False)
    plt.tight_layout()
    fig.savefig(f'{SCI_DIR}/FigS2_SHAP_Global_Bar.png', dpi=300, bbox_inches='tight')
    fig.savefig(f'{SCI_DIR}/FigS2_SHAP_Global_Bar.pdf', bbox_inches='tight')
    plt.close()
    print("[INFO] Saved FigS2_SHAP_Global_Bar")

    return shap_values


def shap_waterfall_cases(X_test_df, shap_values, y_test, y_class_test, model):
    print("[INFO] Generating waterfall plots for representative cases...")
    explainer = shap.LinearExplainer(model, X_test_df)

    # Find representative indices in test set (indices 0-248 within test set)
    test_indices = {
        'LBW': np.where(y_class_test == 0)[0],
        'Normal': np.where(y_class_test == 1)[0],
        'Macro': np.where(y_class_test == 2)[0],
    }

    # Pick cases close to median prediction error within each class
    y_pred = model.predict(X_test_df)
    for case_name, indices in test_indices.items():
        if len(indices) == 0:
            continue
        errors = np.abs(y_pred[indices] - y_test[indices])
        median_err_idx = indices[np.argmin(np.abs(errors - np.median(errors)))]

        fig, ax = plt.subplots(figsize=(10, 8))
        shap.waterfall_plot(
            shap.Explanation(
                values=shap_values[median_err_idx],
                base_values=explainer.expected_value,
                data=X_test_df.iloc[median_err_idx],
                feature_names=X_test_df.columns.tolist()
            ),
            show=False
        )
        plt.tight_layout()
        fname = f'{SCI_DIR}/FigS4_SHAP_Case_{case_name}'
        fig.savefig(f'{fname}.png', dpi=300, bbox_inches='tight')
        fig.savefig(f'{fname}.pdf', bbox_inches='tight')
        plt.close()
        print(f"  {case_name}: actual={y_test[median_err_idx]:.0f}g, predicted={y_pred[median_err_idx]:.0f}g -> {fname}.png")


def main():
    print("=" * 70)
    print("SHAP Analysis - 30-Feature Model")
    print("=" * 70)

    X_train_df, X_test_df, y_test, y_class_test, model = prepare_data()
    shap_values = shap_global_analysis(X_train_df, X_test_df, model)
    shap_waterfall_cases(X_test_df, shap_values, y_test, y_class_test, model)

    print("\n" + "=" * 70)
    print("SHAP analysis completed!")
    print("=" * 70)


if __name__ == '__main__':
    main()
