#!/usr/bin/env python3

import argparse
import gc
from itertools import combinations
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 tensorflow import keras
from tensorflow.keras import layers, Model


parser = argparse.ArgumentParser()
parser.add_argument("--input", required=True)
parser.add_argument("--pretrained-model", required=True)
parser.add_argument("--outdir", required=True)
parser.add_argument("--seed", type=int, default=42)
parser.add_argument("--epochs", type=int, default=400)
parser.add_argument("--patience", type=int, default=40)
parser.add_argument("--learning-rate", type=float, default=1e-3)
parser.add_argument("--l2", type=float, default=1e-4)
args = parser.parse_args()

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

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

distance_columns = [f"distance_pos{i}" for i in range(3, 11)]
required = [
    "record_id",
    "target_id",
    "length_group_id",
    "reported_position",
    "target_30mer",
    "CRISPRoff_score",
    "CRISPRon_score",
    "sgRNA_length_numeric",
    "position_efficiency_mean",
    "original_pred_position_efficiency",
] + distance_columns

df = pd.read_csv(args.input, sep="\t")
missing = [column for column in required if column not in df.columns]
if missing:
    raise ValueError(f"输入表缺少列：{missing}")

df["target_id"] = df["target_id"].astype(str)
df["length_group_id"] = df["length_group_id"].astype(str)

for column in [
    "reported_position",
    "CRISPRoff_score",
    "CRISPRon_score",
    "sgRNA_length_numeric",
    "position_efficiency_mean",
    "original_pred_position_efficiency",
] + distance_columns:
    df[column] = pd.to_numeric(df[column], errors="raise")

if df[required].isna().any().any():
    missing_counts = df[required].isna().sum()
    raise ValueError(
        "建模字段存在缺失值：\n"
        + missing_counts[missing_counts > 0].to_string()
    )

if df["record_id"].duplicated().any():
    raise ValueError("record_id存在重复")

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

if not df["target_30mer"].str.fullmatch(r"[ACGT]{30}").all():
    bad = df.loc[
        ~df["target_30mer"].str.fullmatch(r"[ACGT]{30}"),
        "target_30mer",
    ].head(10).tolist()
    raise ValueError(f"存在非法30mer：{bad}")

if not df["reported_position"].between(3, 10).all():
    raise ValueError("reported_position必须位于3–10")

if df["target_id"].nunique() != 12:
    print(
        f"Warning: 当前target数量为{df['target_id'].nunique()}，"
        "不是预期的12"
    )

conditions = (
    df[
        [
            "target_id",
            "length_group_id",
            "target_30mer",
            "CRISPRoff_score",
            "CRISPRon_score",
            "sgRNA_length_numeric",
        ] + distance_columns
    ]
    .drop_duplicates(
        ["target_id", "length_group_id"]
    )
    .copy()
)

if len(conditions) != df[
    ["target_id", "length_group_id"]
].drop_duplicates().shape[0]:
    raise ValueError("target×length条件信息不唯一")

if (
    conditions.groupby("target_id")["target_30mer"].nunique() > 1
).any():
    raise ValueError("同一target在不同长度条件下出现多个target_30mer")


nt_index = {"A": 0, "T": 1, "G": 2, "C": 3}
dataset_weight = np.array([[0, 0, 0, 0, 1]], dtype=np.float32)


def one_hot_sequence(sequence):
    x = np.zeros((1, 30, 4), dtype=np.float32)
    for i, nt in enumerate(sequence):
        x[0, i, nt_index[nt]] = 1.0
    return x


def outcome_matrix(sequence):
    window = sequence[6:14]
    editable = [
        index
        for index, nt in enumerate(window)
        if nt == "A"
    ]
    if not editable:
        raise ValueError(
            f"{sequence}在position3–10没有可编辑A"
        )

    rows = []
    for size in range(1, len(editable) + 1):
        for subset in combinations(editable, size):
            row = np.zeros(8, dtype=np.float32)
            row[list(subset)] = 1.0
            rows.append(row)
    return np.asarray(rows, dtype=np.float32), editable


def build_condition(row, extra_values):
    sequence = row["target_30mer"]
    outcomes, editable = outcome_matrix(sequence)
    n_outcomes = len(outcomes)

    x_seq = np.repeat(
        one_hot_sequence(sequence),
        n_outcomes,
        axis=0,
    )
    x_energy = np.full(
        (n_outcomes, 1),
        float(row["CRISPRoff_score"]),
        dtype=np.float32,
    )
    x_cas9 = np.full(
        (n_outcomes, 1),
        float(row["CRISPRon_score"]),
        dtype=np.float32,
    )
    x_dataset = np.repeat(
        dataset_weight,
        n_outcomes,
        axis=0,
    )
    extra_values = np.asarray(
        extra_values,
        dtype=np.float32,
    ).reshape(1, -1)
    x_extra = np.repeat(
        extra_values,
        n_outcomes,
        axis=0,
    )

    group = df[
        (df["target_id"] == row["target_id"])
        & (df["length_group_id"] == row["length_group_id"])
    ].copy()

    label_index = (
        group["reported_position"].astype(int).to_numpy()
        - 3
    )

    if any(index not in editable for index in label_index):
        bad_positions = group.loc[
            ~group["reported_position"].astype(int).sub(3).isin(editable),
            "reported_position",
        ].tolist()
        raise ValueError(
            f"{row['target_id']} {row['length_group_id']}存在非A标签位置："
            f"{bad_positions}"
        )

    label = group[
        "position_efficiency_mean"
    ].to_numpy(dtype=np.float32)

    record_ids = group["record_id"].astype(str).tolist()

    return {
        "target_id": row["target_id"],
        "length_group_id": row["length_group_id"],
        "sequence": sequence,
        "outcomes": tf.constant(outcomes, dtype=tf.float32),
        "inputs": [
            tf.constant(x_seq),
            tf.constant(outcomes),
            tf.constant(x_energy),
            tf.constant(x_cas9),
            tf.constant(x_dataset),
            tf.constant(x_extra),
        ],
        "label_index": tf.constant(label_index, dtype=tf.int32),
        "label": tf.constant(label, dtype=tf.float32),
        "record_ids": record_ids,
    }


def aggregate_position_predictions(raw_prediction, outcomes):
    pred_eff = tf.reduce_mean(raw_prediction[:, 0])
    pred_freq = tf.nn.relu(raw_prediction[:, 1])

    total_freq = tf.reduce_sum(pred_freq)
    normalized_freq = tf.where(
        total_freq > 1e-8,
        pred_freq / total_freq * pred_eff,
        tf.zeros_like(pred_freq),
    )

    marginal = tf.linalg.matvec(
        tf.transpose(outcomes),
        normalized_freq,
    )
    return marginal


def activation_name(layer):
    return keras.activations.serialize(layer.activation)


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

dense2_layer = pretrained.get_layer("Dense2")
dense3_layer = pretrained.get_layer("Dense3")
output_layer = pretrained.get_layer("Output")

dense2_kernel, dense2_bias = dense2_layer.get_weights()
dense3_kernel, dense3_bias = dense3_layer.get_weights()
output_kernel, output_bias = output_layer.get_weights()


def build_internal_extra_model(extra_dim, model_name):
    seq_input = keras.Input(
        shape=(30, 4),
        name="one_hot",
    )
    outcome_input = keras.Input(
        shape=(8,),
        name="outcome_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",
    )
    extra_input = keras.Input(
        shape=(extra_dim,),
        name="extra_properties",
    )

    seq_features = sequence_encoder(
        seq_input,
        training=False,
    )

    base_fusion = layers.Concatenate(
        name="OriginalFusion",
    )(
        [
            seq_features,
            outcome_input,
            energy_input,
            cas9_input,
            dataset_input,
        ]
    )

    base_dense2_linear = layers.Dense(
        dense2_layer.units,
        activation=None,
        trainable=False,
        name="FrozenDense2Linear",
    )
    base_preactivation = base_dense2_linear(
        base_fusion
    )
    base_dense2_linear.set_weights(
        [dense2_kernel, dense2_bias]
    )

    extra_dense2 = layers.Dense(
        dense2_layer.units,
        activation=None,
        use_bias=True,
        kernel_initializer="zeros",
        bias_initializer="zeros",
        kernel_regularizer=keras.regularizers.l2(
            args.l2
        ),
        bias_regularizer=keras.regularizers.l2(
            args.l2
        ),
        name="ExtraToDense2",
    )
    extra_preactivation = extra_dense2(
        extra_input
    )

    dense2_preactivation = layers.Add(
        name="Dense2Preactivation",
    )(
        [
            base_preactivation,
            extra_preactivation,
        ]
    )

    dense2_output = layers.Activation(
        activation_name(dense2_layer),
        name="Dense2Activation",
    )(dense2_preactivation)

    frozen_dense3 = layers.Dense(
        dense3_layer.units,
        activation=activation_name(
            dense3_layer
        ),
        trainable=False,
        name="FrozenDense3",
    )
    dense3_output = frozen_dense3(
        dense2_output
    )
    frozen_dense3.set_weights(
        [dense3_kernel, dense3_bias]
    )

    frozen_output = layers.Dense(
        output_layer.units,
        activation=activation_name(
            output_layer
        ),
        trainable=False,
        name="FrozenOutput",
    )
    raw_output = frozen_output(
        dense3_output
    )
    frozen_output.set_weights(
        [output_kernel, output_bias]
    )

    model = Model(
        inputs=[
            seq_input,
            outcome_input,
            energy_input,
            cas9_input,
            dataset_input,
            extra_input,
        ],
        outputs=raw_output,
        name=model_name,
    )

    for layer in model.layers:
        layer.trainable = (
            layer.name == "ExtraToDense2"
        )
    sequence_encoder.trainable = False

    return model


def original_single_condition_prediction(row):
    outcomes, _ = outcome_matrix(
        row["target_30mer"]
    )
    n_outcomes = len(outcomes)

    x_seq = np.repeat(
        one_hot_sequence(
            row["target_30mer"]
        ),
        n_outcomes,
        axis=0,
    )
    x_energy = np.full(
        (n_outcomes, 1),
        float(row["CRISPRoff_score"]),
        dtype=np.float32,
    )
    x_cas9 = np.full(
        (n_outcomes, 1),
        float(row["CRISPRon_score"]),
        dtype=np.float32,
    )
    x_dataset = np.repeat(
        dataset_weight,
        n_outcomes,
        axis=0,
    )

    raw = pretrained(
        [
            x_seq,
            outcomes,
            x_energy,
            x_cas9,
            x_dataset,
        ],
        training=False,
    )

    marginal = aggregate_position_predictions(
        raw,
        tf.constant(
            outcomes,
            dtype=tf.float32,
        ),
    ).numpy()

    return marginal


single_baseline = {}
for row in conditions.itertuples(index=False):
    single_baseline[
        (row.target_id, row.length_group_id)
    ] = original_single_condition_prediction(
        pd.Series(row._asdict())
    )


first_row = conditions.iloc[0]
test_model = build_internal_extra_model(
    1,
    "equivalence_check",
)
test_condition = build_condition(
    first_row,
    [0.0],
)

old_raw = pretrained(
    [
        test_condition["inputs"][0],
        test_condition["inputs"][1],
        test_condition["inputs"][2],
        test_condition["inputs"][3],
        test_condition["inputs"][4],
    ],
    training=False,
).numpy()

new_raw = test_model(
    test_condition["inputs"],
    training=False,
).numpy()

max_difference = float(
    np.max(np.abs(old_raw - new_raw))
)

print(
    "Zero-extra equivalence check, max raw-output difference:",
    max_difference,
)

if max_difference > 1e-4:
    raise RuntimeError(
        "重建模型在extra=0时与原模型不一致，"
        "请停止训练并检查网络结构"
    )

del test_model
gc.collect()


def condition_errors(model, condition):
    raw = model(
        condition["inputs"],
        training=False,
    )
    marginal = aggregate_position_predictions(
        raw,
        condition["outcomes"],
    )
    predicted = tf.gather(
        marginal,
        condition["label_index"],
    )
    return predicted - condition["label"]


def evaluate_conditions(model, condition_list):
    errors = []
    for condition in condition_list:
        errors.append(
            condition_errors(
                model,
                condition,
            ).numpy()
        )
    errors = np.concatenate(errors)
    return float(np.mean(np.abs(errors)))


def train_model(
    model,
    train_conditions,
    valid_conditions,
    fold_seed,
):
    optimizer = keras.optimizers.Adam(
        learning_rate=args.learning_rate
    )

    rng = np.random.default_rng(
        fold_seed
    )

    best_val = np.inf
    best_weights = model.get_layer(
        "ExtraToDense2"
    ).get_weights()
    wait = 0
    history_rows = []

    for epoch in range(
        1,
        args.epochs + 1,
    ):
        order = rng.permutation(
            len(train_conditions)
        )

        for condition_index in order:
            condition = train_conditions[
                condition_index
            ]

            with tf.GradientTape() as tape:
                raw = model(
                    condition["inputs"],
                    training=False,
                )
                marginal = (
                    aggregate_position_predictions(
                        raw,
                        condition["outcomes"],
                    )
                )
                predicted = tf.gather(
                    marginal,
                    condition["label_index"],
                )

                data_loss = tf.reduce_mean(
                    tf.abs(
                        predicted
                        - condition["label"]
                    )
                )

                regularization = (
                    tf.add_n(model.losses)
                    if model.losses
                    else tf.constant(
                        0.0,
                        dtype=tf.float32,
                    )
                )

                loss = (
                    data_loss
                    + regularization
                )

            variables = (
                model.trainable_variables
            )
            gradients = tape.gradient(
                loss,
                variables,
            )

            optimizer.apply_gradients(
                [
                    (gradient, variable)
                    for gradient, variable
                    in zip(
                        gradients,
                        variables,
                    )
                    if gradient is not None
                ]
            )

        train_mae = evaluate_conditions(
            model,
            train_conditions,
        )
        val_mae = evaluate_conditions(
            model,
            valid_conditions,
        )

        history_rows.append({
            "epoch": epoch,
            "train_mae": train_mae,
            "val_mae": val_mae,
        })

        if val_mae < best_val - 1e-4:
            best_val = val_mae
            best_weights = (
                model.get_layer(
                    "ExtraToDense2"
                ).get_weights()
            )
            wait = 0
        else:
            wait += 1

        if (
            epoch == 1
            or epoch % 25 == 0
        ):
            print(
                f"epoch={epoch:3d} "
                f"train_MAE={train_mae:.3f} "
                f"val_MAE={val_mae:.3f}"
            )

        if wait >= args.patience:
            break

    model.get_layer(
        "ExtraToDense2"
    ).set_weights(best_weights)

    return (
        pd.DataFrame(history_rows),
        best_val,
    )


def preprocessing_parameters(
    model_type,
    train_condition_table,
):
    if model_type == "calibration":
        return (
            np.array([0.0]),
            np.array([1.0]),
        )

    if model_type == "length":
        values = train_condition_table[
            ["sgRNA_length_numeric"]
        ].to_numpy(dtype=float)

    elif model_type == "distance":
        values = train_condition_table[
            distance_columns
        ].to_numpy(dtype=float)

    else:
        raise ValueError(
            f"未知model_type：{model_type}"
        )

    mean = values.mean(axis=0)
    sd = values.std(
        axis=0,
        ddof=0,
    )
    sd = np.where(
        sd < 1e-8,
        1.0,
        sd,
    )

    return mean, sd


def extra_values(
    model_type,
    row,
    mean,
    sd,
):
    if model_type == "calibration":
        return np.array([0.0])

    if model_type == "length":
        raw = np.array([
            float(
                row[
                    "sgRNA_length_numeric"
                ]
            )
        ])

    elif model_type == "distance":
        raw = row[
            distance_columns
        ].to_numpy(
            dtype=float
        )

    else:
        raise ValueError(model_type)

    return (raw - mean) / sd


def make_condition_list(
    condition_table,
    model_type,
    mean,
    sd,
):
    result = []

    for _, row in condition_table.iterrows():
        result.append(
            build_condition(
                row,
                extra_values(
                    model_type,
                    row,
                    mean,
                    sd,
                ),
            )
        )

    return result


targets = sorted(
    df["target_id"]
    .astype(str)
    .unique()
)

prediction_rows = []
history_rows = []
fold_rows = []

for row in df.itertuples(index=False):
    prediction_rows.append({
        "record_id": row.record_id,
        "target_id": row.target_id,
        "length_group_id": row.length_group_id,
        "reported_position": int(
            row.reported_position
        ),
        "observed": float(
            row.position_efficiency_mean
        ),
        "model": "original_ensemble",
        "prediction": float(
            row.original_pred_position_efficiency
        ),
        "fold": np.nan,
        "validation_target": "",
    })

for _, condition_row in conditions.iterrows():
    marginal = single_baseline[
        (
            condition_row["target_id"],
            condition_row[
                "length_group_id"
            ],
        )
    ]

    group = df[
        (df["target_id"] == condition_row["target_id"])
        & (
            df["length_group_id"]
            == condition_row[
                "length_group_id"
            ]
        )
    ]

    for row in group.itertuples(index=False):
        prediction_rows.append({
            "record_id": row.record_id,
            "target_id": row.target_id,
            "length_group_id": row.length_group_id,
            "reported_position": int(
                row.reported_position
            ),
            "observed": float(
                row.position_efficiency_mean
            ),
            "model": "original_model_1_1",
            "prediction": float(
                np.clip(
                    marginal[
                        int(
                            row.reported_position
                        ) - 3
                    ],
                    0,
                    100,
                )
            ),
            "fold": np.nan,
            "validation_target": "",
        })


model_dimensions = {
    "calibration": 1,
    "length": 1,
    "distance": 8,
}

for fold, test_target in enumerate(
    targets,
    start=1,
):
    remaining = [
        target
        for target in targets
        if target != test_target
    ]

    validation_target = remaining[
        (fold - 1) % len(remaining)
    ]
    train_targets = [
        target
        for target in remaining
        if target != validation_target
    ]

    train_table = conditions[
        conditions["target_id"].isin(
            train_targets
        )
    ]
    validation_table = conditions[
        conditions["target_id"].eq(
            validation_target
        )
    ]
    test_table = conditions[
        conditions["target_id"].eq(
            test_target
        )
    ]

    print(
        "\n"
        + "=" * 70
        + f"\nFold {fold}: "
        f"test={test_target}, "
        f"validation={validation_target}, "
        f"train_targets={len(train_targets)}"
    )

    for model_type in [
        "calibration",
        "length",
        "distance",
    ]:
        print(
            f"\n--- {model_type} ---"
        )

        mean, sd = preprocessing_parameters(
            model_type,
            train_table,
        )

        train_condition_list = (
            make_condition_list(
                train_table,
                model_type,
                mean,
                sd,
            )
        )
        valid_condition_list = (
            make_condition_list(
                validation_table,
                model_type,
                mean,
                sd,
            )
        )
        test_condition_list = (
            make_condition_list(
                test_table,
                model_type,
                mean,
                sd,
            )
        )

        model = build_internal_extra_model(
            model_dimensions[
                model_type
            ],
            f"{model_type}_fold{fold}",
        )

        history, best_val = train_model(
            model,
            train_condition_list,
            valid_condition_list,
            args.seed
            + fold * 100
            + model_dimensions[
                model_type
            ],
        )

        history.insert(
            0,
            "model",
            model_type,
        )
        history.insert(
            1,
            "fold",
            fold,
        )
        history.insert(
            2,
            "test_target",
            test_target,
        )
        history.insert(
            3,
            "validation_target",
            validation_target,
        )

        history_rows.append(history)

        branch_kernel, branch_bias = (
            model.get_layer(
                "ExtraToDense2"
            ).get_weights()
        )

        np.savez_compressed(
            outdir
            / "branch_weights"
            / (
                f"{model_type}_fold"
                f"{fold:02d}.npz"
            ),
            kernel=branch_kernel,
            bias=branch_bias,
            mean=mean,
            sd=sd,
            test_target=np.array(
                [test_target]
            ),
            validation_target=np.array(
                [validation_target]
            ),
        )

        test_errors = []

        for condition in test_condition_list:
            raw = model(
                condition["inputs"],
                training=False,
            )
            marginal = (
                aggregate_position_predictions(
                    raw,
                    condition["outcomes"],
                )
                .numpy()
            )

            for (
                record_id,
                position_index,
                observed_value,
            ) in zip(
                condition["record_ids"],
                condition[
                    "label_index"
                ].numpy(),
                condition[
                    "label"
                ].numpy(),
            ):
                prediction = float(
                    np.clip(
                        marginal[
                            position_index
                        ],
                        0,
                        100,
                    )
                )

                prediction_rows.append({
                    "record_id": record_id,
                    "target_id": condition[
                        "target_id"
                    ],
                    "length_group_id": (
                        condition[
                            "length_group_id"
                        ]
                    ),
                    "reported_position": int(
                        position_index + 3
                    ),
                    "observed": float(
                        observed_value
                    ),
                    "model": (
                        model_type
                        + "_internal"
                    ),
                    "prediction": prediction,
                    "fold": fold,
                    "validation_target": (
                        validation_target
                    ),
                })

                test_errors.append(
                    prediction
                    - float(
                        observed_value
                    )
                )

        test_errors = np.asarray(
            test_errors
        )

        fold_rows.append({
            "model": (
                model_type
                + "_internal"
            ),
            "fold": fold,
            "test_target": test_target,
            "validation_target": (
                validation_target
            ),
            "train_targets": len(
                train_targets
            ),
            "best_validation_mae": (
                best_val
            ),
            "test_mae": float(
                np.mean(
                    np.abs(
                        test_errors
                    )
                )
            ),
            "test_rmse": float(
                np.sqrt(
                    np.mean(
                        test_errors ** 2
                    )
                )
            ),
            "epochs_ran": len(history),
        })

        del model
        gc.collect()


predictions = pd.DataFrame(
    prediction_rows
)

histories = pd.concat(
    history_rows,
    ignore_index=True,
)

fold_metrics = pd.DataFrame(
    fold_rows
)


def metric_row(model_name, scope, data):
    observed = data[
        "observed"
    ].to_numpy(dtype=float)
    predicted = data[
        "prediction"
    ].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(),
    }


metric_rows = []

for model_name, model_data in predictions.groupby(
    "model",
    sort=False,
):
    metric_rows.append(
        metric_row(
            model_name,
            "all_lengths",
            model_data,
        )
    )

    for (
        length_group,
        length_data,
    ) in model_data.groupby(
        "length_group_id",
        sort=False,
    ):
        metric_rows.append(
            metric_row(
                model_name,
                str(
                    length_group
                ),
                length_data,
            )
        )

metrics = pd.DataFrame(
    metric_rows
)


per_target_rows = []

for (
    model_name,
    target_id,
), group in predictions.groupby(
    [
        "model",
        "target_id",
    ]
):
    per_target_rows.append(
        metric_row(
            model_name,
            target_id,
            group,
        )
    )

per_target = pd.DataFrame(
    per_target_rows
).rename(
    columns={"scope": "target_id"}
)


window = (
    predictions
    .groupby(
        [
            "model",
            "length_group_id",
            "reported_position",
        ],
        as_index=False,
    )
    .agg(
        n=("observed", "size"),
        observed_mean=(
            "observed",
            "mean",
        ),
        predicted_mean=(
            "prediction",
            "mean",
        ),
    )
)


predictions.to_csv(
    outdir
    / "abe8e_internal_transfer_predictions.tsv",
    sep="\t",
    index=False,
)

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

fold_metrics.to_csv(
    outdir
    / "abe8e_internal_transfer_folds.tsv",
    sep="\t",
    index=False,
)

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

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

histories.to_csv(
    outdir
    / "abe8e_internal_transfer_history.tsv",
    sep="\t",
    index=False,
)


plot_models = [
    "original_ensemble",
    "original_model_1_1",
    "calibration_internal",
    "length_internal",
    "distance_internal",
]

for length_group in df[
    "length_group_id"
].drop_duplicates():
    subset = window[
        window["length_group_id"]
        == length_group
    ]

    observed_window = (
        subset[
            [
                "reported_position",
                "observed_mean",
            ]
        ]
        .drop_duplicates()
        .sort_values(
            "reported_position"
        )
    )

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

    plt.plot(
        observed_window[
            "reported_position"
        ],
        observed_window[
            "observed_mean"
        ],
        marker="o",
        label="Observed",
    )

    for model_name in plot_models:
        model_data = subset[
            subset["model"]
            == model_name
        ].sort_values(
            "reported_position"
        )

        if model_data.empty:
            continue

        plt.plot(
            model_data[
                "reported_position"
            ],
            model_data[
                "predicted_mean"
            ],
            marker="o",
            label=model_name,
        )

    plt.xlabel(
        "Protospacer position"
    )
    plt.ylabel(
        "A-to-G efficiency (%)"
    )
    plt.title(
        f"Internal transfer: "
        f"{length_group}"
    )
    plt.xticks(
        range(3, 11)
    )
    plt.ylim(
        0,
        100,
    )
    plt.legend(
        fontsize=8
    )
    plt.tight_layout()

    safe_name = (
        str(length_group)
        .replace("/", "_")
        .replace(" ", "_")
    )

    plt.savefig(
        outdir
        / (
            "abe8e_internal_transfer_"
            f"{safe_name}.png"
        ),
        dpi=300,
    )
    plt.close()


print("\nDataset summary:")
print(
    f"rows={len(df)}, "
    f"targets={df['target_id'].nunique()}, "
    f"conditions={len(conditions)}, "
    f"length_groups={df['length_group_id'].nunique()}"
)

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

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

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