import numpy as np
import pandas as pd
import joblib
import os
import re

from sklearn.preprocessing import StandardScaler
from sklearn.metrics import mean_absolute_error, mean_squared_error, r2_score
from sklearn.model_selection import KFold

import physbo

# ---------- 基本インターフェース (全関数維持) ----------

def load_data(infile):
    """ExcelまたはCSVからデータを読み込む"""
    if infile.endswith('.xlsx'):
        return pd.read_excel(infile)
    return pd.read_csv(infile)

def save_data(X, y, filename, descriptor_names=None, objective_name="o:target"):
    """
    記述子Xと目的変数yを結合して保存する。
    o: 接頭辞を付けることで、このツールで再利用可能な形式にする。
    """
    df_X = pd.DataFrame(X, columns=descriptor_names)
    df_y = pd.DataFrame(y, columns=[objective_name])
    df_combined = pd.concat([df_X, df_y], axis=1)
    
    if filename.endswith('.xlsx'):
        df_combined.to_excel(filename, index=False)
    else:
        df_combined.to_csv(filename, index=False)
    print(f"Data saved to {filename}")

def transform(X, y):
    """記述子X, 目的変数yを標準化"""
    scaler_X = StandardScaler()
    scaler_y = StandardScaler()

    X_std = scaler_X.fit_transform(X)
    y_std = scaler_y.fit_transform(y.reshape(-1, 1)).ravel()

    scaler_dict = {
        "X_mean": scaler_X.mean_,
        "X_var": scaler_X.var_,
        "y_mean": scaler_y.mean_[0],
        "y_var": scaler_y.var_[0],
        "scaler_X": scaler_X,
        "scaler_y": scaler_y
    }
    return X_std, y_std, scaler_dict

def inverse_transform(y_std, scaler_dict):
    """標準化された目的変数を元のスケールに戻す"""
    scaler_y = scaler_dict["scaler_y"]
    return scaler_y.inverse_transform(y_std.reshape(-1, 1)).ravel()

def preprocess(df: pd.DataFrame):
    """DataFrameの前処理（o: 目的変数, -: 除外, 文字列はOne-hot）"""
    obj_cols = [c for c in df.columns if c.startswith("o:")]
    drop_cols = [c for c in df.columns if c.startswith("-:")]
    use_cols = [c for c in df.columns if c not in obj_cols + drop_cols]

    X = df[use_cols]
    y = df[obj_cols]

    # カテゴリ変数を one-hot 化
    X = pd.get_dummies(X, drop_first=True)
    # 欠損値除去
    X = X.dropna()
    y = y.loc[X.index]

    return X, y.values.ravel(), X.columns.tolist()

def model(args=None):
    """
    PHYSBO の設定を辞書で定義
    """
    default_args = {
        "num_rand_basis": 200, 
        "score_mode": 'EI', 
        "interval": 0, 
        "max_num_probes": 1
    }
    if args is not None:
        default_args.update(args)
    return default_args

def fit(X_train, y_train, X_all, params):
    """
    PHYSBOのPolicyを作成し学習を実行
    """
    # X_allの中における学習データのインデックスを特定（簡易的に先頭からマッチング）
    # 本来はnp.whereなどで厳密にマッチさせるが、ここでは初期データの順序を維持
    initial_idx = np.arange(len(X_train))
    
    # ポリシーの作成
    policy = physbo.search.discrete.policy(test_X=X_all, initial_data=(initial_idx, y_train))
    
    # 学習の実行 (ハイパーパラメータ更新含む)
    policy.bayes_search(
        max_num_probes=params.get("max_num_probes", 0),
        simulator=None, 
        score=params.get("score_mode", "EI"),
        interval=params.get("interval", 0),
        num_rand_basis=params.get("num_rand_basis", 200)
    )
    return policy

def predict(policy, X):
    """修正：常に(平均, 分散)を返す。回帰予測の本体。"""
    mean = policy.get_post_fmean(X)
    var = policy.get_post_fcov(X)
    return mean, var

def metrics(y_true, y_pred_mean):
    """評価指標（予測値は平均値を使用）"""
    return {
        "MAE": mean_absolute_error(y_true, y_pred_mean),
        "MSE": mean_squared_error(y_true, y_pred_mean),
        "R2": r2_score(y_true, y_pred_mean)
    }

def set_params(params_dict, new_params: dict):
    """パラメータ設定を更新"""
    params_dict.update(new_params)

def get_params(params_dict):
    """パラメータを取得"""
    return params_dict

def cross_validate(X, y, params, cv=5):
    """交差検証 (PHYSBOの構造上、内部でPolicyを再生成)"""
    scores = []
    kf = KFold(n_splits=cv, shuffle=True, random_state=0)
    for train_idx, test_idx in kf.split(X):
        # CV内では、全データ=今回の分割における全データとする
        p = physbo.search.discrete.policy(test_X=X, initial_data=(train_idx, y[train_idx]))
        p.bayes_search(max_num_probes=0, num_rand_basis=params["num_rand_basis"])
        y_pred, _ = predict(p, X[test_idx])
        scores.append(r2_score(y[test_idx], y_pred))
    return {"mean": np.mean(scores), "std": np.std(scores), "scores": scores}

def feature_importance(policy):
    """GPRのスケール長を返す"""
    # physboのモデルからthetaを取得
    model_internal = policy.predictor_config.model
    if hasattr(model_internal, "theta"):
        return {"lengthscale": model_internal.theta}
    return {"info": "not available"}

def sensitivity_analysis(policy, scaler_dict, X_columns, var_index, var_range):
    """1つの変数だけを変化させ、他は平均値に固定して応答を調べる"""
    scaler_X = scaler_dict["scaler_X"]
    scaler_y = scaler_dict["scaler_y"]
    mean_vec = scaler_dict["X_mean"]

    results = []
    for val in var_range:
        x_vec = mean_vec.copy()
        x_vec[var_index] = val
        x_std = scaler_X.transform([x_vec])
        y_std, _ = predict(policy, x_std)
        y_orig = scaler_y.inverse_transform(y_std.reshape(-1, 1)).ravel()
        results.append(y_orig[0])
    return var_range, results

def save_model(policy, file_path):
    """モデルを保存"""
    joblib.dump(policy, file_path)

def load_model(file_path):
    """モデルを読み込み"""
    return joblib.load(file_path)

# ---------- メイン処理例 ----------

if __name__ == "__main__":
    infile = "result.xlsx"
    if os.path.exists(infile):
        # 1. データの読み込みと前処理
        df = load_data(infile)
        X_raw, y_raw, feat_names = preprocess(df)
        
        # 2. 標準化
        X_std, y_std, s_dict = transform(X_raw.values, y_raw)
        
        # 3. モデル設定と学習
        config = model({"num_rand_basis": 500, "max_num_probes": 1})
        trained_policy = fit(X_std, y_std, X_std, config) # 回帰のためX_allもX_stdを使用
        
        # 4. 予測と評価
        y_pred_std, _ = predict(trained_policy, X_std)
        y_pred_orig = inverse_transform(y_pred_std, s_dict)
        print("Metrics:", metrics(y_raw, y_pred_orig))
        # パラメータ確認
        print("Params:", get_params(m))
        
        # 5. 保存
        save_data(X_raw, y_pred_orig, "predicted_output.xlsx", 
                  descriptor_names=feat_names, objective_name="o:predicted")

        """
        # 交差検証
        print("CV:", cross_validate(m, X, y))

        # 特徴量重要度 (参考値)
        print("Feature importance:", feature_importance(m, X))

        # 保存・読み込み
        save_model(m, "physbo_gpr.pkl")
        loaded = load_model("physbo_gpr.pkl")
        print("Loaded model params:", get_params(loaded))
        """



"""
# 前処理
X, y = preprocess(df)

# 標準化
X_std, y_std, scaler_dict = transform(X, y)

# モデル学習
m = model()
m = fit(m, X_std, y_std)

# 感度解析: 0番目の変数を -2～+2 の範囲で変化
var_range = np.linspace(-2, 2, 20)
vals, responses = sensitivity_analysis(m, scaler_dict, X, var_index=0, var_range=var_range)

# 可視化
import matplotlib.pyplot as plt
plt.plot(vals, responses)
plt.xlabel("Variable 0 (original scale)")
plt.ylabel("Predicted objective")
plt.title("Sensitivity analysis")
plt.show()


"""