#!/usr/bin/env python3

import argparse
from pathlib import Path

import matplotlib.pyplot as plt
import numpy as np
import pandas as pd
from sklearn.metrics import mean_absolute_error, mean_squared_error


parser = argparse.ArgumentParser()
parser.add_argument("--root", required=True)
parser.add_argument("--outdir", required=True)
args = parser.parse_args()

root = Path(args.root)
outdir = Path(args.outdir)
outdir.mkdir(parents=True, exist_ok=True)

model_dirs = sorted(
    path
    for path in root.glob("model_*_*")
    if path.is_dir()
)

if not model_dirs:
    raise ValueError(f"没有找到{root}/model_*_*")

position_blocks = []
raw_length_blocks = []
raw_delta_blocks = []
reference = None
original_ensemble = None

outcome_columns = [f"outcome_pos{i}" for i in range(3, 11)]
raw_key_columns = [
    "target_id",
    "length_group_id",
    "target_30mer",
    "fold",
    "test_target",
    "outcome_index",
    "outcome_key",
] + outcome_columns

for model_dir in model_dirs:
    backbone = model_dir.name

    internal_dir = model_dir / "internal"
    delta_dir = model_dir / "length_delta"

    internal_position_file = (
        internal_dir
        / "abe8e_internal_transfer_predictions.tsv"
    )
    delta_position_file = (
        delta_dir
        / "abe8e_length_delta_predictions.tsv"
    )
    internal_raw_file = (
        internal_dir
        / "abe8e_internal_transfer_raw_outcomes.tsv"
    )
    delta_raw_file = (
        delta_dir
        / "abe8e_length_delta_raw_outcomes.tsv"
    )

    needed = [
        internal_position_file,
        delta_position_file,
        internal_raw_file,
        delta_raw_file,
    ]
    if not all(path.exists() for path in needed):
        print(f"Skip incomplete backbone: {backbone}")
        continue

    internal = pd.read_csv(internal_position_file, sep="\t")
    delta = pd.read_csv(delta_position_file, sep="\t")

    length_position = internal[
        internal["model"].eq("length_internal")
    ][
        [
            "record_id",
            "target_id",
            "length_group_id",
            "reported_position",
            "observed",
            "prediction",
            "fold",
        ]
    ].copy().rename(
        columns={"prediction": "length_prediction"}
    )

    delta_position = delta[
        [
            "record_id",
            "target_id",
            "length_group_id",
            "reported_position",
            "observed",
            "prediction",
            "fold",
        ]
    ].copy().rename(
        columns={"prediction": "delta_prediction"}
    )

    merged_position = length_position.merge(
        delta_position,
        on=[
            "record_id",
            "target_id",
            "length_group_id",
            "reported_position",
            "observed",
            "fold",
        ],
        how="inner",
        validate="one_to_one",
    )

    if len(merged_position) != 156:
        raise ValueError(
            f"{backbone}只有{len(merged_position)}条position预测，不是预期156条"
        )

    current_reference = merged_position[
        [
            "record_id",
            "target_id",
            "length_group_id",
            "reported_position",
            "observed",
            "fold",
        ]
    ].sort_values("record_id").reset_index(drop=True)

    if reference is None:
        reference = current_reference
    elif not reference.equals(current_reference):
        raise ValueError(
            f"{backbone}的record/fold定义与其它backbone不一致"
        )

    if original_ensemble is None:
        original_ensemble = internal[
            internal["model"].eq("original_ensemble")
        ][
            ["record_id", "prediction"]
        ].rename(
            columns={"prediction": "original_ensemble"}
        )

    merged_position["backbone"] = backbone
    position_blocks.append(merged_position)

    internal_raw = pd.read_csv(internal_raw_file, sep="\t")
    length_raw = internal_raw[
        internal_raw["model"].eq("length_internal")
    ].copy()

    delta_raw = pd.read_csv(delta_raw_file, sep="\t")

    for raw_name, raw_data in [
        ("length", length_raw),
        ("length_delta", delta_raw),
    ]:
        missing = [
            column
            for column in raw_key_columns
            + ["raw_pred_eff", "raw_pred_freq"]
            if column not in raw_data.columns
        ]
        if missing:
            raise ValueError(
                f"{backbone} {raw_name} raw文件缺少列：{missing}"
            )

        if raw_data.duplicated(raw_key_columns).any():
            raise ValueError(
                f"{backbone} {raw_name}存在重复outcome key"
            )

    length_raw["backbone"] = backbone
    delta_raw["backbone"] = backbone
    raw_length_blocks.append(length_raw)
    raw_delta_blocks.append(delta_raw)


position_all = pd.concat(position_blocks, ignore_index=True)
raw_length_all = pd.concat(raw_length_blocks, ignore_index=True)
raw_delta_all = pd.concat(raw_delta_blocks, ignore_index=True)

backbones = sorted(position_all["backbone"].unique())
n_backbones = len(backbones)

if n_backbones != 25:
    print(
        f"Warning: 当前完整backbone数量为{n_backbones}，不是25；"
        "仍会按现有backbone生成ensemble"
    )


# --------------------------------------------------
# A. 原先的position-first ensemble：
# 每个backbone先normalize并转position，再跨backbone平均
# --------------------------------------------------

position_first = (
    position_all
    .groupby(
        [
            "record_id",
            "target_id",
            "length_group_id",
            "reported_position",
            "observed",
            "fold",
        ],
        as_index=False,
    )
    .agg(
        length_position_first=("length_prediction", "mean"),
        delta_position_first=("delta_prediction", "mean"),
        length_between_model_sd=("length_prediction", "std"),
        delta_between_model_sd=("delta_prediction", "std"),
        n_backbones=("backbone", "nunique"),
    )
)

position_first = position_first.merge(
    original_ensemble,
    on="record_id",
    how="left",
    validate="one_to_one",
)


# --------------------------------------------------
# B. 官方CRISPRon-BE式raw-first ensemble：
# 25个raw outcome outputs先平均 -> clip负freq -> normalize ->
# 最后从outcome frequency求position marginal efficiency
# --------------------------------------------------

def raw_first_position_ensemble(raw_all, output_column):
    # 先对同一个condition/outcome跨backbone平均raw输出
    averaged = (
        raw_all
        .groupby(raw_key_columns, as_index=False)
        .agg(
            raw_pred_eff=("raw_pred_eff", "mean"),
            raw_pred_freq=("raw_pred_freq", "mean"),
            raw_pred_eff_sd=("raw_pred_eff", "std"),
            raw_pred_freq_sd=("raw_pred_freq", "std"),
            n_backbones=("backbone", "nunique"),
        )
    )

    if averaged["n_backbones"].min() != n_backbones:
        raise ValueError(
            f"{output_column}存在某些outcome没有覆盖全部{n_backbones}个backbone"
        )

    # 对应官方normalize_dataset：
    # 1) pred_freq<=0设为0
    # 2) 一个condition的pred_eff取所有outcome的均值
    # 3) 正freq按总和归一化，使其总和等于pred_eff
    averaged["positive_pred_freq"] = averaged["raw_pred_freq"].clip(lower=0)

    condition_columns = [
        "target_id",
        "length_group_id",
        "target_30mer",
        "fold",
        "test_target",
    ]

    condition_eff = (
        averaged
        .groupby(condition_columns)["raw_pred_eff"]
        .transform("mean")
    )
    freq_sum = (
        averaged
        .groupby(condition_columns)["positive_pred_freq"]
        .transform("sum")
    )

    if (freq_sum <= 0).any():
        bad = averaged.loc[
            freq_sum <= 0,
            ["target_id", "length_group_id"]
        ].drop_duplicates().to_dict("records")
        raise ValueError(
            f"{output_column}存在正pred_freq总和为0的condition：{bad}"
        )

    averaged["normalized_pred_freq"] = (
        averaged["positive_pred_freq"]
        / freq_sum
        * condition_eff
    )
    averaged["normalized_pred_eff"] = condition_eff

    # outcome -> position marginal
    position_rows = []
    for condition_key, group in averaged.groupby(condition_columns):
        target_id, length_group_id, target_30mer, fold, test_target = condition_key

        for position in range(3, 11):
            prediction = float(
                (
                    group["normalized_pred_freq"]
                    * group[f"outcome_pos{position}"]
                ).sum()
            )

            position_rows.append({
                "target_id": target_id,
                "length_group_id": length_group_id,
                "target_30mer": target_30mer,
                "fold": fold,
                "test_target": test_target,
                "reported_position": position,
                output_column: prediction,
            })

    position_table = pd.DataFrame(position_rows)

    # 只保留实验表中真正有label的A位置，并恢复record_id/observed
    reference_table = reference[
        [
            "record_id",
            "target_id",
            "length_group_id",
            "reported_position",
            "observed",
            "fold",
        ]
    ].copy()

    position_table = reference_table.merge(
        position_table,
        on=[
            "target_id",
            "length_group_id",
            "reported_position",
            "fold",
        ],
        how="left",
        validate="one_to_one",
    )

    if position_table[output_column].isna().any():
        raise ValueError(
            f"{output_column}有实验位置无法从raw outcome ensemble恢复"
        )

    return position_table, averaged


length_raw_position, length_raw_averaged = raw_first_position_ensemble(
    raw_length_all,
    "length_raw_first",
)

delta_raw_position, delta_raw_averaged = raw_first_position_ensemble(
    raw_delta_all,
    "delta_raw_first",
)

raw_first = length_raw_position.merge(
    delta_raw_position[
        ["record_id", "delta_raw_first"]
    ],
    on="record_id",
    how="inner",
    validate="one_to_one",
)

ensemble = position_first.merge(
    raw_first[
        ["record_id", "length_raw_first", "delta_raw_first"]
    ],
    on="record_id",
    how="inner",
    validate="one_to_one",
)


# --------------------------------------------------
# QC：L20的length和length+Δdistance理论上完全一致
# --------------------------------------------------

l20 = ensemble[
    ensemble["length_group_id"].eq("L20")
]

l20_position_first_diff = float(
    np.max(
        np.abs(
            l20["length_position_first"]
            - l20["delta_position_first"]
        )
    )
)

l20_raw_first_diff = float(
    np.max(
        np.abs(
            l20["length_raw_first"]
            - l20["delta_raw_first"]
        )
    )
)


# --------------------------------------------------
# Metrics
# --------------------------------------------------

def metric_row(model_name, scope, data, column):
    observed = data["observed"].to_numpy(dtype=float)
    predicted = data[column].to_numpy(dtype=float)

    if (
        len(data) >= 2
        and np.std(observed) > 0
        and np.std(predicted) > 0
    ):
        pearson = np.corrcoef(observed, predicted)[0, 1]
        spearman = pd.Series(observed).corr(
            pd.Series(predicted),
            method="spearman",
        )
    else:
        pearson = np.nan
        spearman = np.nan

    return {
        "model": model_name,
        "scope": scope,
        "n": len(data),
        "pearson": pearson,
        "spearman": spearman,
        "mae": mean_absolute_error(observed, predicted),
        "rmse": mean_squared_error(observed, predicted) ** 0.5,
        "observed_mean": observed.mean(),
        "predicted_mean": predicted.mean(),
        "mean_error": (predicted - observed).mean(),
    }


model_columns = {
    "original_ensemble": "original_ensemble",
    "length_position_first_ensemble": "length_position_first",
    "length_delta_position_first_ensemble": "delta_position_first",
    "length_raw_first_ensemble": "length_raw_first",
    "length_delta_raw_first_ensemble": "delta_raw_first",
}

metric_rows = []

for model_name, column in model_columns.items():
    metric_rows.append(
        metric_row(
            model_name,
            "all_lengths",
            ensemble,
            column,
        )
    )

    for length_group, subset in ensemble.groupby(
        "length_group_id",
        sort=False,
    ):
        metric_rows.append(
            metric_row(
                model_name,
                str(length_group),
                subset,
                column,
            )
        )

metrics = pd.DataFrame(metric_rows)


# --------------------------------------------------
# per-target
# --------------------------------------------------

per_target_rows = []

for target_id, subset in ensemble.groupby("target_id"):
    for ensemble_type, length_col, delta_col in [
        (
            "position_first",
            "length_position_first",
            "delta_position_first",
        ),
        (
            "raw_first",
            "length_raw_first",
            "delta_raw_first",
        ),
    ]:
        length_row = metric_row(
            "length",
            target_id,
            subset,
            length_col,
        )
        delta_row = metric_row(
            "length_delta",
            target_id,
            subset,
            delta_col,
        )

        per_target_rows.append({
            "target_id": target_id,
            "ensemble_type": ensemble_type,
            "n": len(subset),
            "length_mae": length_row["mae"],
            "delta_mae": delta_row["mae"],
            "mae_improvement": (
                length_row["mae"] - delta_row["mae"]
            ),
            "length_rmse": length_row["rmse"],
            "delta_rmse": delta_row["rmse"],
            "length_pearson": length_row["pearson"],
            "delta_pearson": delta_row["pearson"],
        })

per_target = pd.DataFrame(per_target_rows)


# --------------------------------------------------
# window
# --------------------------------------------------

window = (
    ensemble
    .groupby(
        ["length_group_id", "reported_position"],
        as_index=False,
    )
    .agg(
        n=("observed", "size"),
        observed_mean=("observed", "mean"),
        original_mean=("original_ensemble", "mean"),
        length_position_first_mean=("length_position_first", "mean"),
        delta_position_first_mean=("delta_position_first", "mean"),
        length_raw_first_mean=("length_raw_first", "mean"),
        delta_raw_first_mean=("delta_raw_first", "mean"),
    )
)


# --------------------------------------------------
# 保存
# --------------------------------------------------

ensemble.to_csv(
    outdir / "abe8e_ensemble_predictions.tsv",
    sep="\t",
    index=False,
)

metrics.to_csv(
    outdir / "abe8e_ensemble_metrics.tsv",
    sep="\t",
    index=False,
)

per_target.to_csv(
    outdir / "abe8e_ensemble_per_target.tsv",
    sep="\t",
    index=False,
)

window.to_csv(
    outdir / "abe8e_ensemble_window.tsv",
    sep="\t",
    index=False,
)

length_raw_averaged.to_csv(
    outdir / "abe8e_length_raw_outcome_ensemble.tsv",
    sep="\t",
    index=False,
)

delta_raw_averaged.to_csv(
    outdir / "abe8e_length_delta_raw_outcome_ensemble.tsv",
    sep="\t",
    index=False,
)


# --------------------------------------------------
# 图：正式图默认展示官方raw-first逻辑
# --------------------------------------------------

for length_group in window["length_group_id"].drop_duplicates():
    subset = window[
        window["length_group_id"].eq(length_group)
    ].sort_values("reported_position")

    plt.figure(figsize=(7, 5))
    plt.plot(
        subset["reported_position"],
        subset["observed_mean"],
        marker="o",
        label="Observed",
    )
    plt.plot(
        subset["reported_position"],
        subset["original_mean"],
        marker="o",
        label="Original CRISPRon-ABE",
    )
    plt.plot(
        subset["reported_position"],
        subset["length_raw_first_mean"],
        marker="o",
        label="Length ensemble",
    )
    plt.plot(
        subset["reported_position"],
        subset["delta_raw_first_mean"],
        marker="o",
        label="Length + Δdistance ensemble",
    )

    plt.xlabel("Protospacer position")
    plt.ylabel("A-to-G efficiency (%)")
    plt.title(f"ABE8e raw-outcome ensemble: {length_group}")
    plt.xticks(range(3, 11))
    plt.ylim(0, 100)
    plt.legend()
    plt.tight_layout()

    safe_name = str(length_group).replace("/", "_").replace(" ", "_")
    plt.savefig(
        outdir / f"abe8e_raw_first_ensemble_{safe_name}.png",
        dpi=300,
    )
    plt.close()


print(f"\nComplete backbones: {n_backbones}")
print(
    "L20 position-first length-vs-delta max difference:",
    f"{l20_position_first_diff:.6g}",
)
print(
    "L20 raw-first length-vs-delta max difference:",
    f"{l20_raw_first_diff:.6g}",
)

print("\nEnsemble metrics:")
print(metrics.to_string(index=False))

raw_target = per_target[
    per_target["ensemble_type"].eq("raw_first")
]

print(
    "\nRaw-first targets improved after adding Δdistance:",
    int((raw_target["mae_improvement"] > 0).sum()),
    "/",
    len(raw_target),
)

print("\nOutput directory:", outdir)
