#!/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",
)

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以外的位置")


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

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

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",
)

off_score = (
    off
    .groupby("guideSeq")["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"}

if required_on - set(on.columns):
    raise ValueError(
        f"crispron.csv缺少列："
        f"{sorted(required_on - set(on.columns))}"
    )

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

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

on_score = (
    on
    .drop_duplicates("30mer")
    .set_index("30mer")["CRISPRon"]
)

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 position_index, nucleotide in enumerate(sequence):
        x_seq[
            row_index,
            position_index,
            nt_index[nucleotide],
        ] = 1.0


# 用8维one-hot告诉模型“当前预测第几个位置”
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)
)


# Kissling ABE8e=第5个dataset
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,
)

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之后、原Output之前的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中outcome-frequency分支初始化新Output
old_kernel, old_bias = (
    pretrained
    .get_layer("Output")
    .get_weights()
)

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


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

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",
)

print("\nWarm-up:")
print(
    "trainable parameters:",
    sum(
        np.prod(weight.shape)
        for weight in model.trainable_weights
    ),
)

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

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


# --------------------------------------------------
# 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",
)

print("\nFine-tuning:")
print(
    "trainable parameters:",
    sum(
        np.prod(weight.shape)
        for weight in model.trainable_weights
    ),
)

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

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


# --------------------------------------------------
# 10.验证集预测
# --------------------------------------------------

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

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

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

validation[
    "transfer_pred_position_efficiency"
] = prediction

validation[
    "transfer_error"
] = (
    validation[
        "transfer_pred_position_efficiency"
    ]
    - validation[
        "true_position_efficiency"
    ]
)


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

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

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

transfer_prediction = (
    validation[
        "transfer_pred_position_efficiency"
    ]
    .to_numpy(dtype=float)
)

metric_rows = []

for model_name, predicted in [
    (
        "original_crispron",
        original_prediction,
    ),
    (
        "position_transfer",
        transfer_prediction,
    ),
]:
    metric_rows.append({
        "model": model_name,
        "n": len(observed),
        "pearson": np.corrcoef(
            observed,
            predicted,
        )[0, 1],
        "spearman": (
            pd.Series(observed)
            .corr(
                pd.Series(predicted),
                method="spearman",
            )
        ),
        "mae": mean_absolute_error(
            observed,
            predicted,
        ),
        "rmse": (
            mean_squared_error(
                observed,
                predicted,
            )
            ** 0.5
        ),
        "observed_mean": observed.mean(),
        "predicted_mean": predicted.mean(),
    })

metrics = pd.DataFrame(metric_rows)


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

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",
        ),
        transfer_mean=(
            "transfer_pred_position_efficiency",
            "mean",
        ),
    )
)


# --------------------------------------------------
# 13.保存结果
# --------------------------------------------------

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,
)

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

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

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

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

model.save(
    outdir
    / "model_1_1_position"
)


# --------------------------------------------------
# 14.画验证集窗口
# --------------------------------------------------

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["transfer_mean"],
    marker="o",
    label="Position transfer",
)

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

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

plt.close()


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("\nValidation metrics:")
print(metrics.to_string(index=False))

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

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