import pandas as pd
import numpy as np

from scipy.stats import friedmanchisquare, wilcoxon
from statsmodels.stats.multitest import multipletests
import scikit_posthocs as sp


# ============================================================
# 0. Display options
# ============================================================

pd.set_option("display.max_columns", None)
pd.set_option("display.width", 300)
pd.set_option("display.max_colwidth", None)


# ============================================================
# 1. Fold-level MAE results
# ============================================================

mae_results = pd.DataFrame({
    "Mean Regressor": [18.5520, 16.8971, 18.4479, 18.1730, 17.0881],
    "Random Regressor": [32.1466, 31.0180, 30.5458, 30.4683, 30.1428],
    "Ridge Regressor": [14.8569, 14.2710, 14.2194, 14.1562, 14.3466],
    "XGBoost Regressor": [14.6095, 13.9877, 15.0843, 14.9531, 14.9619],
    "LightGBM Regressor": [14.8531, 13.7020, 14.5231, 14.0990, 14.1353],
    "Support Vector Regressor": [14.9052, 14.4084, 14.1775, 14.5395, 14.2648],
    "MLP Regressor": [15.3575, 15.1828, 14.1775, 16.1870, 15.8075],
    "Random Forest Regressor": [15.0514, 14.0772, 15.0323, 14.9942, 14.3722],
    "KNN Regressor": [15.8270, 15.4652, 15.7302, 15.9742, 15.1512],
    "ElasticNet Regressor": [14.8680, 14.2976, 17.5640, 14.5305, 14.3521],
    "CatBoost Regressor": [14.1610, 13.7003, 14.2160, 13.9823, 14.1214],
    "Stacking Regressor": [14.2118, 13.6202, 14.1485, 13.8617, 13.9651],
})

alpha = 0.05
reference_model = "Stacking Regressor"


# ============================================================
# 2. Helper functions
# ============================================================

def significance_label(p):
    """
    Converts a p-value into a significance label.
    """
    if pd.isna(p):
        return "--"
    elif p < 0.01:
        return "***"
    elif p < 0.05:
        return "**"
    elif p < 0.1:
        return "*"
    else:
        return "n.s."


def export_latex_landscape_table(df, filename, caption, label):
    """
    Exports a dataframe as a LaTeX table in landscape format.
    This is useful for large appendix tables.
    """
    latex_table = df.to_latex(
        caption=caption,
        label=label,
        escape=False
    )

    with open(filename, "w") as f:
        f.write("\\begin{landscape}\n")
        f.write("\\begin{table}[H]\n")
        f.write("\\centering\n")
        f.write("\\scriptsize\n")
        f.write(latex_table)
        f.write("\\end{table}\n")
        f.write("\\end{landscape}\n")


# ============================================================
# 3. Descriptive results
# ============================================================

print("\n============================================================")
print("DESCRIPTIVE RESULTS")
print("============================================================")

mean_mae = mae_results.mean().sort_values()
std_mae = mae_results.std().loc[mean_mae.index]

descriptive_results = pd.DataFrame({
    "Mean MAE": mean_mae,
    "Std MAE": std_mae,
})

print(descriptive_results.round(4))


# ============================================================
# 4. Friedman test
# ============================================================

print("\n============================================================")
print("FRIEDMAN TEST")
print("============================================================")

friedman_stat, friedman_p = friedmanchisquare(
    *[mae_results[col] for col in mae_results.columns]
)

print(f"Friedman statistic: {friedman_stat:.4f}")
print(f"p-value: {friedman_p:.6f}")

if friedman_p < alpha:
    print(f"Result: significant at alpha = {alpha}")
else:
    print(f"Result: not significant at alpha = {alpha}")


# ============================================================
# 5. Fold-level ranks and average ranks
# ============================================================

print("\n============================================================")
print("FOLD-LEVEL RANKS")
print("============================================================")

fold_ranks = mae_results.rank(axis=1, method="average", ascending=True)
print(fold_ranks.round(2))

print("\n============================================================")
print("AVERAGE RANKS")
print("============================================================")

average_ranks = fold_ranks.mean().sort_values()
ordered_models = average_ranks.index.tolist()

print(average_ranks.round(4))


# ============================================================
# 6. Nemenyi post-hoc test
# ============================================================

print("\n============================================================")
print("NEMENYI POST-HOC TEST")
print("============================================================")

nemenyi_matrix = sp.posthoc_nemenyi_friedman(mae_results)

# Reorder rows and columns according to average rank
nemenyi_matrix = nemenyi_matrix.loc[ordered_models, ordered_models]

print("\nNemenyi p-value matrix:")
print(nemenyi_matrix.round(4))


# ============================================================
# 7. Nemenyi comparisons involving the Stacking Regressor
# ============================================================

print("\n============================================================")
print("NEMENYI COMPARISONS INVOLVING STACKING REGRESSOR")
print("============================================================")

stacking_nemenyi = (
    nemenyi_matrix.loc[reference_model]
    .drop(reference_model)
    .sort_values()
)

nemenyi_stacking_table = pd.DataFrame({
    "Comparison": [f"{reference_model} vs {model}" for model in stacking_nemenyi.index],
    "p-value": stacking_nemenyi.values,
    "Significant at 0.05": stacking_nemenyi.values < alpha,
})

print(nemenyi_stacking_table.to_string(index=False, formatters={
    "p-value": "{:.4f}".format
}))


# ============================================================
# 8. All significant Nemenyi pairwise comparisons
# ============================================================

print("\n============================================================")
print("ALL SIGNIFICANT NEMENYI PAIRWISE COMPARISONS")
print("============================================================")

significant_pairs = []

for i, model_1 in enumerate(ordered_models):
    for model_2 in ordered_models[i + 1:]:
        p_value = nemenyi_matrix.loc[model_1, model_2]

        if p_value < alpha:
            significant_pairs.append({
                "Model 1": model_1,
                "Model 2": model_2,
                "p-value": p_value,
            })

significant_pairs_df = pd.DataFrame(significant_pairs)

if significant_pairs_df.empty:
    print("No significant pairwise comparisons were found.")
else:
    print(significant_pairs_df.to_string(index=False, formatters={
        "p-value": "{:.4f}".format
    }))


# ============================================================
# 9. Critical difference for Nemenyi test
# ============================================================

print("\n============================================================")
print("CRITICAL DIFFERENCE")
print("============================================================")

k = mae_results.shape[1]   # number of models
N = mae_results.shape[0]   # number of folds

# Approximate q_alpha for k = 12 models and alpha = 0.05.
q_alpha = 3.268

critical_difference = q_alpha * np.sqrt(k * (k + 1) / (6 * N))

print(f"Number of models (k): {k}")
print(f"Number of folds (N): {N}")
print(f"q_alpha: {q_alpha}")
print(f"Critical difference: {critical_difference:.4f}")

print("\nAverage rank differences from Stacking Regressor:")

stacking_rank = average_ranks[reference_model]

critical_difference_results = []

for model in ordered_models:
    if model != reference_model:
        rank_difference = average_ranks[model] - stacking_rank
        significant_by_cd = rank_difference > critical_difference

        critical_difference_results.append({
            "Comparison": f"{reference_model} vs {model}",
            "Rank difference": rank_difference,
            "Significant by CD": significant_by_cd,
        })

critical_difference_df = pd.DataFrame(critical_difference_results)

print(critical_difference_df.to_string(index=False, formatters={
    "Rank difference": "{:.4f}".format
}))


# ============================================================
# 10. Pairwise Wilcoxon signed-rank tests
#     Focused comparisons: Stacking vs selected models
# ============================================================

print("\n============================================================")
print("PAIRWISE WILCOXON SIGNED-RANK TESTS WITH HOLM CORRECTION")
print("============================================================")

comparison_models = [
    "CatBoost Regressor",
    "LightGBM Regressor",
    "Ridge Regressor",
    "Support Vector Regressor",
]

wilcoxon_results = []

print("Reference model:", reference_model)
print("Alternative hypothesis: Stacking Regressor has lower MAE than the compared model.\n")

for model in comparison_models:
    reference_values = mae_results[reference_model]
    compared_values = mae_results[model]

    differences = reference_values - compared_values

    statistic, raw_p_value = wilcoxon(
        reference_values,
        compared_values,
        alternative="less"
    )

    wilcoxon_results.append({
        "Comparison": f"{reference_model} vs {model}",
        "Wilcoxon statistic": statistic,
        "Raw p-value": raw_p_value,
        "Mean MAE Stacking": reference_values.mean(),
        "Mean MAE compared model": compared_values.mean(),
        "Mean difference (Stacking - compared)": differences.mean(),
        "Folds won by Stacking": (differences < 0).sum(),
        "Folds won by compared model": (differences > 0).sum(),
        "Ties": (differences == 0).sum(),
    })

wilcoxon_results_df = pd.DataFrame(wilcoxon_results)

raw_p_values = wilcoxon_results_df["Raw p-value"].values

reject, corrected_p_values, _, _ = multipletests(
    raw_p_values,
    alpha=alpha,
    method="holm"
)

wilcoxon_results_df["Holm-corrected p-value"] = corrected_p_values
wilcoxon_results_df["Significant after Holm correction"] = reject

print(wilcoxon_results_df.to_string(index=False, formatters={
    "Raw p-value": "{:.6f}".format,
    "Holm-corrected p-value": "{:.6f}".format,
    "Mean MAE Stacking": "{:.4f}".format,
    "Mean MAE compared model": "{:.4f}".format,
    "Mean difference (Stacking - compared)": "{:.4f}".format,
}))


# ============================================================
# 11. Full pairwise Nemenyi significance matrix
# ============================================================

print("\n============================================================")
print("NEMENYI PAIRWISE SIGNIFICANCE MATRIX")
print("============================================================")

nemenyi_significance_matrix = nemenyi_matrix.map(significance_label)

print(nemenyi_significance_matrix)


# ============================================================
# 12. Full pairwise Wilcoxon signed-rank tests
#     All models compared 2 by 2
# ============================================================

print("\n============================================================")
print("FULL PAIRWISE WILCOXON SIGNED-RANK TESTS")
print("============================================================")

models = ordered_models
wilcoxon_pairs = []
raw_p_values = []

for i, model_1 in enumerate(models):
    for j, model_2 in enumerate(models):
        if j <= i:
            continue

        statistic, raw_p_value = wilcoxon(
            mae_results[model_1],
            mae_results[model_2],
            alternative="two-sided"
        )

        wilcoxon_pairs.append((model_1, model_2))
        raw_p_values.append(raw_p_value)

# Holm correction across all pairwise Wilcoxon tests
reject, corrected_p_values, _, _ = multipletests(
    raw_p_values,
    alpha=alpha,
    method="holm"
)

# Create matrices
wilcoxon_raw_p_matrix = pd.DataFrame(
    np.nan,
    index=models,
    columns=models
)

wilcoxon_corrected_p_matrix = pd.DataFrame(
    np.nan,
    index=models,
    columns=models
)

wilcoxon_significance_matrix = pd.DataFrame(
    "--",
    index=models,
    columns=models
)

for (model_1, model_2), raw_p, corrected_p in zip(
    wilcoxon_pairs,
    raw_p_values,
    corrected_p_values
):
    wilcoxon_raw_p_matrix.loc[model_1, model_2] = raw_p
    wilcoxon_raw_p_matrix.loc[model_2, model_1] = raw_p

    wilcoxon_corrected_p_matrix.loc[model_1, model_2] = corrected_p
    wilcoxon_corrected_p_matrix.loc[model_2, model_1] = corrected_p

    label = significance_label(corrected_p)

    wilcoxon_significance_matrix.loc[model_1, model_2] = label
    wilcoxon_significance_matrix.loc[model_2, model_1] = label

print("\nWilcoxon raw p-value matrix:")
print(wilcoxon_raw_p_matrix.round(4))

print("\nWilcoxon Holm-corrected p-value matrix:")
print(wilcoxon_corrected_p_matrix.round(4))

print("\nWilcoxon significance matrix:")
print(wilcoxon_significance_matrix)


# ============================================================
# 13. Short names for appendix tables
# ============================================================

short_names = {
    "Stacking Regressor": "Stacking",
    "CatBoost Regressor": "CatBoost",
    "LightGBM Regressor": "LightGBM",
    "Ridge Regressor": "Ridge",
    "Support Vector Regressor": "SVR",
    "XGBoost Regressor": "XGBoost",
    "ElasticNet Regressor": "ElasticNet",
    "Random Forest Regressor": "Random Forest",
    "MLP Regressor": "MLP",
    "KNN Regressor": "KNN",
    "Mean Regressor": "Mean",
    "Random Regressor": "Random",
}

nemenyi_significance_matrix_short = nemenyi_significance_matrix.rename(
    index=short_names,
    columns=short_names
)

wilcoxon_significance_matrix_short = wilcoxon_significance_matrix.rename(
    index=short_names,
    columns=short_names
)

nemenyi_p_values_matrix_short = nemenyi_matrix.rename(
    index=short_names,
    columns=short_names
)

wilcoxon_raw_p_matrix_short = wilcoxon_raw_p_matrix.rename(
    index=short_names,
    columns=short_names
)

wilcoxon_corrected_p_matrix_short = wilcoxon_corrected_p_matrix.rename(
    index=short_names,
    columns=short_names
)


# ============================================================
# 14. Summary
# ============================================================

print("\n============================================================")
print("SUMMARY")
print("============================================================")

print(f"Friedman test: statistic = {friedman_stat:.4f}, p-value = {friedman_p:.6f}")

print("\nBest average ranks:")
print(average_ranks.head(5).round(4))

print("\nNemenyi comparisons involving Stacking Regressor:")
for _, row in nemenyi_stacking_table.iterrows():
    print(
        f"{row['Comparison']}: "
        f"p = {row['p-value']:.4f}, "
        f"significant = {row['Significant at 0.05']}"
    )

print("\nWilcoxon + Holm comparisons:")
for _, row in wilcoxon_results_df.iterrows():
    print(
        f"{row['Comparison']}: "
        f"raw p = {row['Raw p-value']:.6f}, "
        f"Holm-corrected p = {row['Holm-corrected p-value']:.6f}, "
        f"significant = {row['Significant after Holm correction']}"
    )


# ============================================================
# 15. Export results to CSV files
# ============================================================

descriptive_results.to_csv("descriptive_mae_results.csv")
fold_ranks.to_csv("fold_level_ranks_mae.csv", index=False)
average_ranks.to_csv("average_ranks_mae.csv")

nemenyi_matrix.to_csv("nemenyi_p_values_mae.csv")
nemenyi_significance_matrix.to_csv("nemenyi_significance_matrix_mae.csv")
nemenyi_significance_matrix_short.to_csv("nemenyi_significance_matrix_mae_short.csv")
nemenyi_p_values_matrix_short.to_csv("nemenyi_p_values_matrix_mae_short.csv")

nemenyi_stacking_table.to_csv("nemenyi_stacking_comparisons_mae.csv", index=False)
critical_difference_df.to_csv("critical_difference_results_mae.csv", index=False)

wilcoxon_results_df.to_csv("wilcoxon_holm_results_mae.csv", index=False)
wilcoxon_raw_p_matrix.to_csv("wilcoxon_raw_p_values_matrix_mae.csv")
wilcoxon_corrected_p_matrix.to_csv("wilcoxon_holm_p_values_matrix_mae.csv")
wilcoxon_significance_matrix.to_csv("wilcoxon_significance_matrix_mae.csv")

wilcoxon_raw_p_matrix_short.to_csv("wilcoxon_raw_p_values_matrix_mae_short.csv")
wilcoxon_corrected_p_matrix_short.to_csv("wilcoxon_holm_p_values_matrix_mae_short.csv")
wilcoxon_significance_matrix_short.to_csv("wilcoxon_significance_matrix_mae_short.csv")


# ============================================================
# 16. Export LaTeX tables for Overleaf appendix
# ============================================================

export_latex_landscape_table(
    nemenyi_significance_matrix_short,
    filename="nemenyi_significance_matrix_mae.tex",
    caption="Pairwise Nemenyi post-hoc comparisons between regression models.",
    label="tab:nemenyi_pairwise_significance"
)

export_latex_landscape_table(
    wilcoxon_significance_matrix_short,
    filename="wilcoxon_significance_matrix_mae.tex",
    caption="Pairwise Wilcoxon signed-rank comparisons between regression models with Holm correction.",
    label="tab:wilcoxon_pairwise_significance"
)

export_latex_landscape_table(
    nemenyi_p_values_matrix_short.round(4),
    filename="nemenyi_p_values_matrix_mae.tex",
    caption="Pairwise Nemenyi post-hoc p-values between regression models.",
    label="tab:nemenyi_pairwise_p_values"
)

export_latex_landscape_table(
    wilcoxon_corrected_p_matrix_short.round(4),
    filename="wilcoxon_holm_p_values_matrix_mae.tex",
    caption="Pairwise Wilcoxon signed-rank Holm-corrected p-values between regression models.",
    label="tab:wilcoxon_pairwise_holm_p_values"
)


# ============================================================
# 17. Final output
# ============================================================

print("\n============================================================")
print("CSV FILES EXPORTED")
print("============================================================")
print("- descriptive_mae_results.csv")
print("- fold_level_ranks_mae.csv")
print("- average_ranks_mae.csv")
print("- nemenyi_p_values_mae.csv")
print("- nemenyi_significance_matrix_mae.csv")
print("- nemenyi_significance_matrix_mae_short.csv")
print("- nemenyi_p_values_matrix_mae_short.csv")
print("- nemenyi_stacking_comparisons_mae.csv")
print("- critical_difference_results_mae.csv")
print("- wilcoxon_holm_results_mae.csv")
print("- wilcoxon_raw_p_values_matrix_mae.csv")
print("- wilcoxon_holm_p_values_matrix_mae.csv")
print("- wilcoxon_significance_matrix_mae.csv")
print("- wilcoxon_raw_p_values_matrix_mae_short.csv")
print("- wilcoxon_holm_p_values_matrix_mae_short.csv")
print("- wilcoxon_significance_matrix_mae_short.csv")

print("\n============================================================")
print("LATEX FILES EXPORTED")
print("============================================================")
print("- nemenyi_significance_matrix_mae.tex")
print("- wilcoxon_significance_matrix_mae.tex")
print("- nemenyi_p_values_matrix_mae.tex")
print("- wilcoxon_holm_p_values_matrix_mae.tex")

print("\n============================================================")
print("SIGNIFICANCE LEGEND")
print("============================================================")
print("*** : p < 0.01")
print("**  : p < 0.05")
print("*   : p < 0.1")
print("n.s.: p >= 0.1")
print("--  : same model")