#!/usr/bin/env python3

import argparse
from pathlib import Path

import matplotlib.pyplot as plt
import numpy as np
import pandas as pd
import tensorflow as tf
from sklearn.metrics import mean_absolute_error, mean_squared_error
from sklearn.model_selection import GroupShuffleSplit
from tensorflow import keras
from tensorflow.keras import layers, Model


parser = argparse.ArgumentParser()
parser.add_argument("--position-table", required=True)
parser.add_argument("--crisproff", required=True)
parser.add_argument("--crispron", required=True)
parser.add_argument("--pretrained-model", required=True)
parser.add_argument("--outdir", required=True)
parser.add_argument("--seed", type=int, default=42)
args = parser.parse_args()

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

np.random.seed(args.seed)
tf.random.set_seed(args.seed)


# --------------------------------------------------
# 1.读取20nt位置级真实数据
# --------------------------------------------------

df = pd.read_csv(args.position_table, sep="\t")

required_columns = [
    "seq_id",
    "seq30",
    "position",
    "true_position_efficiency",
    "pred_position_efficiency",
]

missing = [
    column for column in required_columns
    if column not in df.columns
]

if missing:
    raise ValueError(f"position table缺少列：{missing}")

df["seq30"] = (
    df["seq30"]
    .astype(str)
    .str.upper()
    .str.strip()
)

df["position"] = pd.to_numeric(
    df["position"],
    errors="raise",
).astype(int)

df["true_position_efficiency"] = pd.to_numeric(
    df["true_position_efficiency"],
    errors="raise",
)

df["pred_position_efficiency"] = pd.to_numeric(
    df["pred_position_efficiency"],
    errors="raise",
)

if not df["seq30"].str.fullmatch(r"[ACGT]{30}").all():
    raise ValueError("存在非法30mer")

if not df["position"].between(3, 10).all():
    raise ValueError("position存在3–10以外的位置")

if df["true_position_efficiency"].isna().any():
    raise ValueError("true_position_efficiency存在缺失值")

if df["pred_position_efficiency"].isna().any():
    raise ValueError("pred_position_efficiency存在缺失值")


# --------------------------------------------------
# 2.合并CRISPRoff
# --------------------------------------------------

off = pd.read_csv(args.crisproff, sep="\t")

required_off = {
    "guideSeq",
    "CRISPRoff_score",
}

missing_off = required_off - set(off.columns)

if missing_off:
    raise ValueError(
        f"CRISPRoff文件缺少列：{sorted(missing_off)}"
    )

off["guideSeq"] = (
    off["guideSeq"]
    .astype(str)
    .str.upper()
    .str.strip()
    .str.replace("U", "T", regex=False)
)

off["CRISPRoff_score"] = pd.to_numeric(
    off["CRISPRoff_score"],
    errors="raise",
)

score_range = (
    off.groupby("guideSeq")["CRISPRoff_score"]
    .agg(lambda x: x.max() - x.min())
)

conflicts = score_range[score_range > 1e-8]

if not conflicts.empty:
    raise ValueError(
        "同一guideSeq存在不同CRISPRoff_score。"
        f"示例：{conflicts.index[:10].tolist()}"
    )

off_score = (
    off.groupby("guideSeq", sort=False)["CRISPRoff_score"]
    .first()
)

df["target_23mer"] = df["seq30"].str.slice(4, 27)

df["CRISPRoff_score"] = (
    df["target_23mer"]
    .map(off_score)
)

if df["CRISPRoff_score"].isna().any():
    missing_guides = (
        df.loc[
            df["CRISPRoff_score"].isna(),
            "target_23mer",
        ]
        .drop_duplicates()
        .head(10)
        .tolist()
    )

    raise ValueError(
        "部分30mer没有CRISPRoff score："
        f"{missing_guides}"
    )


# --------------------------------------------------
# 3.合并CRISPRon
# --------------------------------------------------

on = pd.read_csv(args.crispron)

required_on = {
    "30mer",
    "CRISPRon",
}

missing_on = required_on - set(on.columns)

if missing_on:
    raise ValueError(
        f"crispron.csv缺少列：{sorted(missing_on)}"
    )

on["30mer"] = (
    on["30mer"]
    .astype(str)
    .str.upper()
    .str.strip()
)

on["CRISPRon"] = pd.to_numeric(
    on["CRISPRon"],
    errors="raise",
)

on_range = (
    on.groupby("30mer")["CRISPRon"]
    .agg(lambda x: x.max() - x.min())
)

on_conflicts = on_range[on_range > 1e-8]

if not on_conflicts.empty:
    raise ValueError(
        "同一30mer存在不同CRISPRon值。"
        f"示例：{on_conflicts.index[:10].tolist()}"
    )

on_score = (
    on.groupby("30mer", sort=False)["CRISPRon"]
    .first()
)

df["CRISPRon_score"] = df["seq30"].map(on_score)

if df["CRISPRon_score"].isna().any():
    missing_targets = (
        df.loc[
            df["CRISPRon_score"].isna(),
            "seq30",
        ]
        .drop_duplicates()
        .head(10)
        .tolist()
    )

    raise ValueError(
        "部分30mer没有CRISPRon score："
        f"{missing_targets}"
    )


# --------------------------------------------------
# 4.构建5类原始模型输入
# --------------------------------------------------

nt_index = {
    "A": 0,
    "T": 1,
    "G": 2,
    "C": 3,
}

x_seq = np.zeros(
    (len(df), 30, 4),
    dtype=np.float32,
)

for row_index, sequence in enumerate(df["seq30"]):
    for sequence_index, nucleotide in enumerate(sequence):
        x_seq[
            row_index,
            sequence_index,
            nt_index[nucleotide],
        ] = 1.0


# 8维one-hot：当前预测的是protospacer第3–10位中的哪个位置
x_position = np.zeros(
    (len(df), 8),
    dtype=np.float32,
)

position_index = (
    df["position"].to_numpy(dtype=int) - 3
)

x_position[
    np.arange(len(df)),
    position_index,
] = 1.0


x_energy = (
    df["CRISPRoff_score"]
    .to_numpy(dtype=np.float32)
    .reshape(-1, 1)
)

x_cas9 = (
    df["CRISPRon_score"]
    .to_numpy(dtype=np.float32)
    .reshape(-1, 1)
)


# ABE的第5个dataset对应Kissling ABE8e
x_dataset = np.tile(
    np.array(
        [[0, 0, 0, 0, 1]],
        dtype=np.float32,
    ),
    (len(df), 1),
)

y = (
    df["true_position_efficiency"]
    .to_numpy(dtype=np.float32)
    .reshape(-1, 1)
)


# --------------------------------------------------
# 5.按30mer分train/validation
# --------------------------------------------------

splitter = GroupShuffleSplit(
    n_splits=1,
    test_size=0.20,
    random_state=args.seed,
)

train_index, valid_index = next(
    splitter.split(
        df,
        y,
        groups=df["seq30"],
    )
)

train_targets = set(
    df.iloc[train_index]["seq30"]
)

valid_targets = set(
    df.iloc[valid_index]["seq30"]
)

overlap = train_targets & valid_targets

if overlap:
    raise ValueError(
        f"train/validation存在{len(overlap)}个重复30mer"
    )

X_train = [
    x_seq[train_index],
    x_position[train_index],
    x_energy[train_index],
    x_cas9[train_index],
    x_dataset[train_index],
]

X_valid = [
    x_seq[valid_index],
    x_position[valid_index],
    x_energy[valid_index],
    x_cas9[valid_index],
    x_dataset[valid_index],
]

y_train = y[train_index]
y_valid = y[valid_index]


# --------------------------------------------------
# 6.加载原CRISPRon-ABE模型
# --------------------------------------------------

pretrained = keras.models.load_model(
    args.pretrained_model,
    compile=False,
)

expected_input_shapes = [
    (None, 30, 4),
    (None, 8),
    (None, 1),
    (None, 1),
    (None, 5),
]

observed_input_shapes = [
    tuple(model_input.shape)
    for model_input in pretrained.inputs
]

if observed_input_shapes != expected_input_shapes:
    raise ValueError(
        "预训练模型输入shape与预期不一致。\n"
        f"Observed: {observed_input_shapes}\n"
        f"Expected: {expected_input_shapes}"
    )

for layer_name in [
    "Dense2",
    "Dense3",
    "Dropout3",
    "Output",
]:
    try:
        pretrained.get_layer(layer_name)
    except ValueError as exc:
        raise ValueError(
            f"预训练模型缺少必要层：{layer_name}"
        ) from exc

print("\nPretrained model loaded:")
print(args.pretrained_model)

print("\nInput shapes:")
for model_input in pretrained.inputs:
    print(
        model_input.name,
        tuple(model_input.shape),
    )


# --------------------------------------------------
# 7.建立position-level transfer model
# --------------------------------------------------

# 使用原模型Dense3后的200维表示
hidden = pretrained.get_layer(
    "Dropout3"
).output

position_output_layer = layers.Dense(
    1,
    name="PositionOutput",
)

position_output = position_output_layer(hidden)

model = Model(
    inputs=pretrained.inputs,
    outputs=position_output,
)


# 用原Output的第2列(pred_freq分支)初始化新的position输出层
old_kernel, old_bias = (
    pretrained
    .get_layer("Output")
    .get_weights()
)

if old_kernel.shape[1] != 2 or old_bias.shape[0] != 2:
    raise ValueError(
        "原Output不是2维输出，无法按pred_freq分支初始化PositionOutput"
    )

position_output_layer.set_weights([
    old_kernel[:, 1:2],
    old_bias[1:2],
])


# --------------------------------------------------
# 8.第一阶段：只训练PositionOutput
# --------------------------------------------------

for layer in model.layers:
    layer.trainable = False

model.get_layer(
    "PositionOutput"
).trainable = True

model.compile(
    optimizer=keras.optimizers.Adam(
        learning_rate=1e-3
    ),
    loss="mae",
)

head_only_trainable = int(
    sum(
        np.prod(weight.shape)
        for weight in model.trainable_weights
    )
)

print("\nHead-only training:")
print(
    "trainable parameters:",
    head_only_trainable,
)

head_early_stop = keras.callbacks.EarlyStopping(
    monitor="val_loss",
    patience=20,
    restore_best_weights=True,
)

head_history = model.fit(
    X_train,
    y_train,
    validation_data=(
        X_valid,
        y_valid,
    ),
    epochs=150,
    batch_size=32,
    verbose=2,
    callbacks=[head_early_stop],
)

head_only_prediction = (
    model.predict(
        X_valid,
        verbose=0,
    )
    .reshape(-1)
)

head_only_prediction = np.clip(
    head_only_prediction,
    0,
    100,
)

model.save(
    outdir / "model_1_1_position_head_only"
)


# --------------------------------------------------
# 9.第二阶段：解冻Dense2/Dense3，小学习率fine-tune
# --------------------------------------------------

model.get_layer("Dense2").trainable = True
model.get_layer("Dense3").trainable = True

model.compile(
    optimizer=keras.optimizers.Adam(
        learning_rate=1e-4
    ),
    loss="mae",
)

finetune_trainable = int(
    sum(
        np.prod(weight.shape)
        for weight in model.trainable_weights
    )
)

print("\nDense2/Dense3 fine-tuning:")
print(
    "trainable parameters:",
    finetune_trainable,
)

fine_tune_early_stop = keras.callbacks.EarlyStopping(
    monitor="val_loss",
    patience=30,
    restore_best_weights=True,
)

finetune_history = model.fit(
    X_train,
    y_train,
    validation_data=(
        X_valid,
        y_valid,
    ),
    epochs=300,
    batch_size=32,
    verbose=2,
    callbacks=[fine_tune_early_stop],
)

finetune_prediction = (
    model.predict(
        X_valid,
        verbose=0,
    )
    .reshape(-1)
)

finetune_prediction = np.clip(
    finetune_prediction,
    0,
    100,
)

model.save(
    outdir / "model_1_1_position_finetune"
)


# --------------------------------------------------
# 10.整理验证集结果
# --------------------------------------------------

validation = (
    df.iloc[valid_index]
    .copy()
    .reset_index(drop=True)
)

validation[
    "head_only_pred_position_efficiency"
] = head_only_prediction

validation[
    "finetune_pred_position_efficiency"
] = finetune_prediction

validation[
    "head_only_error"
] = (
    validation[
        "head_only_pred_position_efficiency"
    ]
    - validation[
        "true_position_efficiency"
    ]
)

validation[
    "finetune_error"
] = (
    validation[
        "finetune_pred_position_efficiency"
    ]
    - validation[
        "true_position_efficiency"
    ]
)

validation[
    "original_absolute_error"
] = (
    validation[
        "pred_position_efficiency"
    ]
    - validation[
        "true_position_efficiency"
    ]
).abs()

validation[
    "head_only_absolute_error"
] = (
    validation[
        "head_only_pred_position_efficiency"
    ]
    - validation[
        "true_position_efficiency"
    ]
).abs()

validation[
    "finetune_absolute_error"
] = (
    validation[
        "finetune_pred_position_efficiency"
    ]
    - validation[
        "true_position_efficiency"
    ]
).abs()


# --------------------------------------------------
# 11.计算整体指标
# --------------------------------------------------

observed = validation[
    "true_position_efficiency"
].to_numpy(dtype=float)

original_prediction = validation[
    "pred_position_efficiency"
].to_numpy(dtype=float)

head_prediction = validation[
    "head_only_pred_position_efficiency"
].to_numpy(dtype=float)

finetune_prediction = validation[
    "finetune_pred_position_efficiency"
].to_numpy(dtype=float)


def calculate_metrics(
    model_name,
    observed_values,
    predicted_values,
):
    observed_values = np.asarray(
        observed_values,
        dtype=float,
    )

    predicted_values = np.asarray(
        predicted_values,
        dtype=float,
    )

    if (
        len(observed_values) >= 2
        and np.std(observed_values) > 0
        and np.std(predicted_values) > 0
    ):
        pearson = np.corrcoef(
            observed_values,
            predicted_values,
        )[0, 1]

        spearman = (
            pd.Series(observed_values)
            .corr(
                pd.Series(predicted_values),
                method="spearman",
            )
        )
    else:
        pearson = np.nan
        spearman = np.nan

    return {
        "model": model_name,
        "n": len(observed_values),
        "pearson": pearson,
        "spearman": spearman,
        "mae": mean_absolute_error(
            observed_values,
            predicted_values,
        ),
        "rmse": mean_squared_error(
            observed_values,
            predicted_values,
        ) ** 0.5,
        "observed_mean": observed_values.mean(),
        "predicted_mean": predicted_values.mean(),
    }


metrics = pd.DataFrame([
    calculate_metrics(
        "original_crispron",
        observed,
        original_prediction,
    ),
    calculate_metrics(
        "head_only",
        observed,
        head_prediction,
    ),
    calculate_metrics(
        "dense23_finetune",
        observed,
        finetune_prediction,
    ),
])


# --------------------------------------------------
# 12.逐位置指标
# --------------------------------------------------

position_metric_rows = []

for position in range(3, 11):
    subset = validation[
        validation["position"] == position
    ]

    if subset.empty:
        continue

    observed_position = subset[
        "true_position_efficiency"
    ].to_numpy(dtype=float)

    prediction_sets = {
        "original_crispron": subset[
            "pred_position_efficiency"
        ].to_numpy(dtype=float),
        "head_only": subset[
            "head_only_pred_position_efficiency"
        ].to_numpy(dtype=float),
        "dense23_finetune": subset[
            "finetune_pred_position_efficiency"
        ].to_numpy(dtype=float),
    }

    for model_name, predicted_position in prediction_sets.items():
        row = calculate_metrics(
            model_name,
            observed_position,
            predicted_position,
        )

        row["position"] = position
        position_metric_rows.append(row)

position_metrics = pd.DataFrame(
    position_metric_rows
)


# --------------------------------------------------
# 13.编辑窗口比较
# --------------------------------------------------

window = (
    validation
    .groupby(
        "position",
        as_index=False,
    )
    .agg(
        n=(
            "true_position_efficiency",
            "size",
        ),
        observed_mean=(
            "true_position_efficiency",
            "mean",
        ),
        original_mean=(
            "pred_position_efficiency",
            "mean",
        ),
        head_only_mean=(
            "head_only_pred_position_efficiency",
            "mean",
        ),
        finetune_mean=(
            "finetune_pred_position_efficiency",
            "mean",
        ),
    )
)


# --------------------------------------------------
# 14.比较每条记录是否改善
# --------------------------------------------------

improvement_summary = pd.DataFrame([
    {
        "model": "head_only",
        "n": len(validation),
        "n_better_than_original": int(
            (
                validation["head_only_absolute_error"]
                < validation["original_absolute_error"]
            ).sum()
        ),
        "fraction_better_than_original": float(
            (
                validation["head_only_absolute_error"]
                < validation["original_absolute_error"]
            ).mean()
        ),
        "median_absolute_error": float(
            validation[
                "head_only_absolute_error"
            ].median()
        ),
        "median_original_absolute_error": float(
            validation[
                "original_absolute_error"
            ].median()
        ),
    },
    {
        "model": "dense23_finetune",
        "n": len(validation),
        "n_better_than_original": int(
            (
                validation["finetune_absolute_error"]
                < validation["original_absolute_error"]
            ).sum()
        ),
        "fraction_better_than_original": float(
            (
                validation["finetune_absolute_error"]
                < validation["original_absolute_error"]
            ).mean()
        ),
        "median_absolute_error": float(
            validation[
                "finetune_absolute_error"
            ].median()
        ),
        "median_original_absolute_error": float(
            validation[
                "original_absolute_error"
            ].median()
        ),
    },
])


# --------------------------------------------------
# 15.保存split
# --------------------------------------------------

split_table = (
    df[
        [
            "seq_id",
            "seq30",
        ]
    ]
    .drop_duplicates()
    .copy()
)

split_table["split"] = np.where(
    split_table["seq30"].isin(
        valid_targets
    ),
    "validation",
    "train",
)


# --------------------------------------------------
# 16.保存训练历史
# --------------------------------------------------

head_history_df = pd.DataFrame(
    head_history.history
)

head_history_df.insert(
    0,
    "epoch",
    np.arange(1, len(head_history_df) + 1),
)

finetune_history_df = pd.DataFrame(
    finetune_history.history
)

finetune_history_df.insert(
    0,
    "epoch",
    np.arange(1, len(finetune_history_df) + 1),
)


# --------------------------------------------------
# 17.保存结果
# --------------------------------------------------

validation.to_csv(
    outdir
    / "abe8e_position_transfer_validation.tsv",
    sep="\t",
    index=False,
)

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

position_metrics.to_csv(
    outdir
    / "abe8e_position_transfer_position_metrics.tsv",
    sep="\t",
    index=False,
)

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

improvement_summary.to_csv(
    outdir
    / "abe8e_position_transfer_improvement.tsv",
    sep="\t",
    index=False,
)

split_table.to_csv(
    outdir
    / "abe8e_position_transfer_split.tsv",
    sep="\t",
    index=False,
)

head_history_df.to_csv(
    outdir
    / "abe8e_position_transfer_head_history.tsv",
    sep="\t",
    index=False,
)

finetune_history_df.to_csv(
    outdir
    / "abe8e_position_transfer_finetune_history.tsv",
    sep="\t",
    index=False,
)


# --------------------------------------------------
# 18.画验证集平均编辑窗口
# --------------------------------------------------

plt.figure(figsize=(7, 5))

plt.plot(
    window["position"],
    window["observed_mean"],
    marker="o",
    label="Observed",
)

plt.plot(
    window["position"],
    window["original_mean"],
    marker="o",
    label="Original CRISPRon-ABE",
)

plt.plot(
    window["position"],
    window["head_only_mean"],
    marker="o",
    label="Head only",
)

plt.plot(
    window["position"],
    window["finetune_mean"],
    marker="o",
    label="Dense2/3 fine-tune",
)

plt.xlabel("Protospacer position")
plt.ylabel("Mean A-to-G efficiency (%)")
plt.title("ABE8e position transfer comparison")
plt.xticks(range(3, 11))
plt.ylim(0, 100)
plt.legend()
plt.tight_layout()

plt.savefig(
    outdir
    / "abe8e_position_transfer_comparison.png",
    dpi=300,
)

plt.close()


# --------------------------------------------------
# 19.画训练历史
# --------------------------------------------------

plt.figure(figsize=(7, 5))

plt.plot(
    head_history_df["epoch"],
    head_history_df["loss"],
    label="Train",
)

plt.plot(
    head_history_df["epoch"],
    head_history_df["val_loss"],
    label="Validation",
)

plt.xlabel("Epoch")
plt.ylabel("MAE loss")
plt.title("Head-only training")
plt.legend()
plt.tight_layout()

plt.savefig(
    outdir
    / "abe8e_position_transfer_head_history.png",
    dpi=300,
)

plt.close()


plt.figure(figsize=(7, 5))

plt.plot(
    finetune_history_df["epoch"],
    finetune_history_df["loss"],
    label="Train",
)

plt.plot(
    finetune_history_df["epoch"],
    finetune_history_df["val_loss"],
    label="Validation",
)

plt.xlabel("Epoch")
plt.ylabel("MAE loss")
plt.title("Dense2/3 fine-tuning")
plt.legend()
plt.tight_layout()

plt.savefig(
    outdir
    / "abe8e_position_transfer_finetune_history.png",
    dpi=300,
)

plt.close()


# --------------------------------------------------
# 20.打印总结
# --------------------------------------------------

print("\nDataset:")
print(
    f"all rows={len(df)}, "
    f"train rows={len(train_index)}, "
    f"validation rows={len(valid_index)}"
)

print(
    f"train 30mers={len(train_targets)}, "
    f"validation 30mers={len(valid_targets)}"
)

print("\nTrainable parameters:")
print(
    f"head_only={head_only_trainable}"
)

print(
    f"dense23_finetune={finetune_trainable}"
)

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

print("\nImprovement relative to original CRISPRon-ABE:")
print(
    improvement_summary.to_string(index=False)
)

print("\nWindow:")
print(
    window.to_string(index=False)
)

print(
    f"\nOutput directory: {outdir}"
)
