#!/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.读取position-level数据
# --------------------------------------------------

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():
    examples = (
        df.loc[
            df["CRISPRoff_score"].isna(),
            "target_23mer",
        ]
        .drop_duplicates()
        .head(10)
        .tolist()
    )

    raise ValueError(
        "部分target缺少CRISPRoff score："
        f"{examples}"
    )


# --------------------------------------------------
# 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缺少列："
        f"{sorted(missing_on)}"
    )

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():
    examples = (
        df.loc[
            df["CRISPRon_score"].isna(),
            "seq30",
        ]
        .drop_duplicates()
        .head(10)
        .tolist()
    )

    raise ValueError(
        "部分target缺少CRISPRon score："
        f"{examples}"
    )


# --------------------------------------------------
# 4.构建输入
# --------------------------------------------------

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


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


# 与原ABE8e预测一致：
# 第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划分训练集和验证集
# --------------------------------------------------

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

if train_targets & valid_targets:
    raise ValueError(
        "train和validation存在重复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并截取sequence encoder
# --------------------------------------------------

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

sequence_encoder = Model(
    inputs=pretrained.inputs[0],
    outputs=pretrained.get_layer(
        "Dense1"
    ).output,
    name="CRISPRonABE_sequence_encoder",
)

sequence_encoder.trainable = False


print("\nSequence encoder:")
sequence_encoder.summary()


# --------------------------------------------------
# 7.构建新的position-level模型
# --------------------------------------------------

seq_input = keras.Input(
    shape=(30, 4),
    name="one_hot",
)

position_input = keras.Input(
    shape=(8,),
    name="position_properties",
)

energy_input = keras.Input(
    shape=(1,),
    name="energy_properties",
)

cas9_input = keras.Input(
    shape=(1,),
    name="cas9",
)

dataset_input = keras.Input(
    shape=(5,),
    name="dataset",
)


# training=False非常重要：
# 固定原encoder中的SpatialDropout行为
sequence_features = sequence_encoder(
    seq_input,
    training=False,
)


fusion = layers.Concatenate(
    name="PositionFusion"
)(
    [
        sequence_features,
        position_input,
        energy_input,
        cas9_input,
        dataset_input,
    ]
)


x = layers.Dense(
    64,
    activation="relu",
    kernel_regularizer=keras.regularizers.l2(
        1e-3
    ),
    name="PositionDense1",
)(fusion)

x = layers.Dropout(
    0.1,
    name="PositionDropout1",
)(x)

x = layers.Dense(
    16,
    activation="relu",
    kernel_regularizer=keras.regularizers.l2(
        1e-3
    ),
    name="PositionDense2",
)(x)

position_output = layers.Dense(
    1,
    name="PositionOutput",
)(x)


model = Model(
    inputs=[
        seq_input,
        position_input,
        energy_input,
        cas9_input,
        dataset_input,
    ],
    outputs=position_output,
)


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


print("\nNew position model:")
model.summary()

print(
    "\nTrainable parameters:",
    sum(
        np.prod(weight.shape)
        for weight in model.trainable_weights
    ),
)


# --------------------------------------------------
# 8.训练
# --------------------------------------------------

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

reduce_lr = (
    keras.callbacks.ReduceLROnPlateau(
        monitor="val_loss",
        factor=0.5,
        patience=10,
        min_lr=1e-6,
        verbose=1,
    )
)


history = model.fit(
    X_train,
    y_train,
    validation_data=(
        X_valid,
        y_valid,
    ),
    epochs=400,
    batch_size=32,
    verbose=2,
    callbacks=[
        early_stop,
        reduce_lr,
    ],
)


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

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[
    "encoder_position_prediction"
] = prediction

validation[
    "encoder_position_error"
] = (
    validation[
        "encoder_position_prediction"
    ]
    - validation[
        "true_position_efficiency"
    ]
)


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

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

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

encoder_prediction = (
    validation[
        "encoder_position_prediction"
    ]
    .to_numpy(dtype=float)
)


metric_rows = []

for model_name, predicted in [
    (
        "original_crispron",
        original_prediction,
    ),
    (
        "encoder_position",
        encoder_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
)


# --------------------------------------------------
# 11.逐position指标
# --------------------------------------------------

position_metric_rows = []

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

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

    for model_name, column in [
        (
            "original_crispron",
            "pred_position_efficiency",
        ),
        (
            "encoder_position",
            "encoder_position_prediction",
        ),
    ]:
        predicted_position = subset[
            column
        ].to_numpy(dtype=float)

        position_metric_rows.append({
            "model": model_name,
            "position": position,
            "n": len(subset),
            "pearson": np.corrcoef(
                observed_position,
                predicted_position,
            )[0, 1],
            "spearman": (
                pd.Series(
                    observed_position
                )
                .corr(
                    pd.Series(
                        predicted_position
                    ),
                    method="spearman",
                )
            ),
            "mae": mean_absolute_error(
                observed_position,
                predicted_position,
            ),
            "rmse": (
                mean_squared_error(
                    observed_position,
                    predicted_position,
                ) ** 0.5
            ),
        })


position_metrics = pd.DataFrame(
    position_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",
        ),
        encoder_mean=(
            "encoder_position_prediction",
            "mean",
        ),
    )
)


# --------------------------------------------------
# 13.训练history
# --------------------------------------------------

history_table = pd.DataFrame({
    "epoch": np.arange(
        1,
        len(history.history["loss"]) + 1,
    ),
    "loss": history.history["loss"],
    "val_loss": history.history["val_loss"],
})


# --------------------------------------------------
# 14.保存
# --------------------------------------------------

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

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

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

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

history_table.to_csv(
    outdir
    / "abe8e_encoder_position_history.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_encoder_position_split.tsv",
    sep="\t",
    index=False,
)


model.save(
    outdir
    / "model_1_1_encoder_position"
)


# --------------------------------------------------
# 15.编辑窗口图
# --------------------------------------------------

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["encoder_mean"],
    marker="o",
    label="Sequence encoder + new head",
)

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

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

plt.close()


# --------------------------------------------------
# 16.loss图
# --------------------------------------------------

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

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

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

plt.xlabel("Epoch")
plt.ylabel("MAE loss")
plt.title(
    "Sequence-encoder position training"
)
plt.legend()
plt.tight_layout()

plt.savefig(
    outdir
    / "abe8e_encoder_position_history.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("\nMetrics:")
print(
    metrics.to_string(
        index=False
    )
)

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

print(
    "\nBest validation MAE:",
    min(history.history["val_loss"]),
)

print(
    "\nOutput directory:",
    outdir,
)