#!/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("--previous-outdir", 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=5e-4)
parser.add_argument("--l2", type=float, default=1e-4)
parser.add_argument("--bottleneck", type=int, default=4)
args = parser.parse_args()

outdir = Path(args.outdir)
previous_outdir = Path(args.previous_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}")

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

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

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存在重复")

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

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

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条件信息不唯一")

# 每个target的L20 distance作为结构基准
l20 = conditions[
    conditions["length_group_id"].eq("L20")
][["target_id"] + distance_columns].copy()

if l20["target_id"].duplicated().any():
    raise ValueError("同一target存在多个L20结构")

if set(l20["target_id"]) != set(conditions["target_id"]):
    missing_targets = sorted(
        set(conditions["target_id"]) - set(l20["target_id"])
    )
    raise ValueError(f"以下target缺少L20结构：{missing_targets}")

l20 = l20.rename(
    columns={
        column: f"{column}_L20"
        for column in distance_columns
    }
)

conditions = conditions.merge(
    l20,
    on="target_id",
    how="left",
    validate="many_to_one",
)

delta_columns = []
for column in distance_columns:
    delta_column = f"delta_{column}"
    conditions[delta_column] = (
        conditions[column]
        - conditions[f"{column}_L20"]
    )
    delta_columns.append(delta_column)

# L20的delta必须严格为0
l20_delta = conditions[
    conditions["length_group_id"].eq("L20")
][delta_columns].to_numpy(dtype=float)

if np.max(np.abs(l20_delta)) > 1e-8:
    raise ValueError("L20 delta distance没有归零")


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, length_value, delta_value):
    outcomes, editable = 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,
    )
    x_length = np.full(
        (n_outcomes, 1),
        float(length_value),
        dtype=np.float32,
    )
    x_delta = np.repeat(
        np.asarray(delta_value, dtype=np.float32).reshape(1, -1),
        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):
        raise ValueError(
            f"{row['target_id']} {row['length_group_id']}存在非A标签位置"
        )

    return {
        "target_id": row["target_id"],
        "length_group_id": row["length_group_id"],
        "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_length),
            tf.constant(x_delta),
        ],
        "label_index": tf.constant(label_index, dtype=tf.int32),
        "label": tf.constant(
            group["position_efficiency_mean"].to_numpy(dtype=np.float32)
        ),
        "record_ids": group["record_id"].astype(str).tolist(),
    }


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

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


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_length_delta_model(length_kernel, length_bias, fold):
    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")
    length_input = keras.Input(shape=(1,), name="length_properties")
    delta_input = keras.Input(shape=(8,), name="delta_distance_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 = layers.Dense(
        dense2_layer.units,
        activation=None,
        trainable=False,
        name="FrozenDense2Linear",
    )
    base_preactivation = base_dense2(base_fusion)
    base_dense2.set_weights([dense2_kernel, dense2_bias])

    length_layer = layers.Dense(
        dense2_layer.units,
        activation=None,
        use_bias=True,
        trainable=False,
        name="FrozenLengthToDense2",
    )
    length_preactivation = length_layer(length_input)
    length_layer.set_weights([length_kernel, length_bias])

    # 8维结构变化先压成少量latent factors，再映射到Dense2。
    # 两层都不使用bias，因此L20的delta=0时结构修正严格为0。
    delta_latent = layers.Dense(
        args.bottleneck,
        activation="tanh",
        use_bias=False,
        kernel_regularizer=keras.regularizers.l2(args.l2),
        name="DeltaDistanceBottleneck",
    )(delta_input)

    delta_preactivation = layers.Dense(
        dense2_layer.units,
        activation=None,
        use_bias=False,
        kernel_initializer="zeros",
        kernel_regularizer=keras.regularizers.l2(args.l2),
        name="DeltaDistanceToDense2",
    )(delta_latent)

    dense2_preactivation = layers.Add(name="Dense2Preactivation")(
        [
            base_preactivation,
            length_preactivation,
            delta_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,
            length_input,
            delta_input,
        ],
        outputs=raw_output,
        name=f"length_delta_distance_fold{fold}",
    )

    for layer in model.layers:
        layer.trainable = layer.name in {
            "DeltaDistanceBottleneck",
            "DeltaDistanceToDense2",
        }

    sequence_encoder.trainable = False
    return model


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 = [
        condition_errors(model, condition).numpy()
        for condition in condition_list
    ]
    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 = {
        "bottleneck": model.get_layer(
            "DeltaDistanceBottleneck"
        ).get_weights(),
        "projection": model.get_layer(
            "DeltaDistanceToDense2"
        ).get_weights(),
    }
    wait = 0
    history_rows = []

    for epoch in range(1, args.epochs + 1):
        for condition_index in rng.permutation(len(train_conditions)):
            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

            gradients = tape.gradient(
                loss,
                model.trainable_variables,
            )
            optimizer.apply_gradients(
                [
                    (gradient, variable)
                    for gradient, variable
                    in zip(gradients, model.trainable_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 = {
                "bottleneck": model.get_layer(
                    "DeltaDistanceBottleneck"
                ).get_weights(),
                "projection": model.get_layer(
                    "DeltaDistanceToDense2"
                ).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("DeltaDistanceBottleneck").set_weights(
        best_weights["bottleneck"]
    )
    model.get_layer("DeltaDistanceToDense2").set_weights(
        best_weights["projection"]
    )

    return pd.DataFrame(history_rows), best_val


previous_predictions = pd.read_csv(
    previous_outdir / "abe8e_internal_transfer_predictions.tsv",
    sep="\t",
)

targets = sorted(df["target_id"].unique())
new_prediction_rows = []
history_rows = []
fold_rows = []

for fold, test_target in enumerate(targets, start=1):
    length_weights_file = (
        previous_outdir
        / "branch_weights"
        / f"length_fold{fold:02d}.npz"
    )

    if not length_weights_file.exists():
        raise FileNotFoundError(
            f"找不到上一阶段length权重：{length_weights_file}"
        )

    saved = np.load(length_weights_file, allow_pickle=True)
    saved_test_target = str(saved["test_target"][0])
    validation_target = str(saved["validation_target"][0])

    if saved_test_target != test_target:
        raise ValueError(
            f"fold{fold} target不一致："
            f"当前={test_target}, saved={saved_test_target}"
        )

    length_kernel = saved["kernel"]
    length_bias = saved["bias"]
    length_mean = np.asarray(saved["mean"], dtype=float)
    length_sd = np.asarray(saved["sd"], dtype=float)

    train_targets = [
        target
        for target in targets
        if target not in {test_target, validation_target}
    ]

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

    # delta只做scale，不做中心化：
    # 这样L20的delta=0经过预处理后仍然严格为0。
    delta_sd = (
        train_table[delta_columns]
        .to_numpy(dtype=float)
        .std(axis=0, ddof=0)
    )
    delta_sd = np.where(delta_sd < 1e-8, 1.0, delta_sd)

    def make_condition_list(table):
        result = []
        for _, row in table.iterrows():
            length_value = (
                float(row["sgRNA_length_numeric"]) - length_mean[0]
            ) / length_sd[0]

            delta_value = (
                row[delta_columns].to_numpy(dtype=float)
                / delta_sd
            )

            result.append(
                build_condition(
                    row,
                    length_value,
                    delta_value,
                )
            )
        return result

    train_conditions = make_condition_list(train_table)
    valid_conditions = make_condition_list(valid_table)
    test_conditions = make_condition_list(test_table)

    model = build_length_delta_model(
        length_kernel,
        length_bias,
        fold,
    )

    # delta投影初始化为0，因此训练前必须与上一阶段length_internal一致
    previous_test = previous_predictions[
        (previous_predictions["model"] == "length_internal")
        & (previous_predictions["target_id"] == test_target)
    ].set_index("record_id")["prediction"]

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

        for record_id, position_index in zip(
            condition["record_ids"],
            condition["label_index"].numpy(),
        ):
            if record_id not in previous_test.index:
                raise ValueError(
                    f"上一阶段缺少length_internal预测：{record_id}"
                )
            initial_prediction = float(marginal[position_index])
            initial_differences.append(
                abs(initial_prediction - float(previous_test.loc[record_id]))
            )

    max_initial_difference = max(initial_differences)
    print(
        "\n"
        + "=" * 70
        + f"\nFold {fold}: test={test_target}, "
        f"validation={validation_target}, "
        f"length-equivalence max diff={max_initial_difference:.6g}"
    )

    if max_initial_difference > 1e-3:
        raise RuntimeError(
            "length+delta模型在delta分支为0时与上一阶段length_internal不一致"
        )

    history, best_val = train_model(
        model,
        train_conditions,
        valid_conditions,
        args.seed + fold * 100,
    )

    history.insert(0, "fold", fold)
    history.insert(1, "test_target", test_target)
    history.insert(2, "validation_target", validation_target)
    history_rows.append(history)

    bottleneck_weights = model.get_layer(
        "DeltaDistanceBottleneck"
    ).get_weights()[0]
    projection_weights = model.get_layer(
        "DeltaDistanceToDense2"
    ).get_weights()[0]

    np.savez_compressed(
        outdir
        / "branch_weights"
        / f"length_delta_distance_fold{fold:02d}.npz",
        bottleneck_kernel=bottleneck_weights,
        projection_kernel=projection_weights,
        delta_sd=delta_sd,
        length_mean=length_mean,
        length_sd=length_sd,
        test_target=np.array([test_target]),
        validation_target=np.array([validation_target]),
    )

    test_errors = []

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

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

            new_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),
                "model": "length_delta_distance_internal",
                "prediction": prediction,
                "fold": fold,
                "validation_target": validation_target,
            })

            test_errors.append(prediction - float(observed))

    test_errors = np.asarray(test_errors)

    fold_rows.append({
        "model": "length_delta_distance_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),
        "max_initial_length_difference": max_initial_difference,
    })

    del model
    gc.collect()


new_predictions = pd.DataFrame(new_prediction_rows)
histories = pd.concat(history_rows, ignore_index=True)
fold_metrics = pd.DataFrame(fold_rows)

comparison_predictions = pd.concat(
    [
        previous_predictions[
            previous_predictions["model"].isin(
                [
                    "original_ensemble",
                    "original_model_1_1",
                    "calibration_internal",
                    "length_internal",
                    "distance_internal",
                ]
            )
        ],
        new_predictions,
    ],
    ignore_index=True,
)


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 comparison_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 comparison_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 = (
    comparison_predictions
    .groupby(
        ["model", "length_group_id", "reported_position"],
        as_index=False,
    )
    .agg(
        n=("observed", "size"),
        observed_mean=("observed", "mean"),
        predicted_mean=("prediction", "mean"),
    )
)

new_predictions.to_csv(
    outdir / "abe8e_length_delta_predictions.tsv",
    sep="\t",
    index=False,
)

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

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

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

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

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

plot_models = [
    "length_internal",
    "distance_internal",
    "length_delta_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")

        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"Length + structural-change model: {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 / f"abe8e_length_delta_{safe_name}.png",
        dpi=300,
    )
    plt.close()


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

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

print("\nOutput directory:", outdir)
