#!/usr/bin/env python3

import argparse
from pathlib import Path

import pandas as pd


parser = argparse.ArgumentParser()
parser.add_argument("--model1-dir", required=True)
parser.add_argument("--robustness-dir", required=True)
parser.add_argument("--out", required=True)
args = parser.parse_args()

rows = []

model_dirs = {
    "model_1_1": Path(args.model1_dir),
}

robustness_dir = Path(args.robustness_dir)

for model_name in [
    "model_2_2",
    "model_3_3",
    "model_4_4",
    "model_5_5",
]:
    model_dirs[model_name] = (
        robustness_dir
        / model_name
        / "length_delta"
    )


for backbone, directory in model_dirs.items():

    metrics = pd.read_csv(
        directory
        / "abe8e_length_delta_metrics.tsv",
        sep="\t",
    )

    per_target = pd.read_csv(
        directory
        / "abe8e_length_delta_per_target.tsv",
        sep="\t",
    )

    for scope in [
        "all_lengths",
        "L20",
        "L18_19",
        "L16_17",
    ]:

        subset = (
            metrics[
                metrics["scope"] == scope
            ]
            .set_index("model")
        )

        length = subset.loc[
            "length_internal"
        ]

        delta = subset.loc[
            "length_delta_distance_internal"
        ]

        rows.append({
            "backbone": backbone,
            "scope": scope,
            "length_mae": length["mae"],
            "delta_mae": delta["mae"],
            "mae_improvement": (
                length["mae"]
                - delta["mae"]
            ),
            "length_rmse": length["rmse"],
            "delta_rmse": delta["rmse"],
            "length_pearson": length["pearson"],
            "delta_pearson": delta["pearson"],
        })


result = pd.DataFrame(rows)

result.to_csv(
    args.out,
    sep="\t",
    index=False,
)

print("\nBackbone robustness:")
print(
    result.to_string(
        index=False
    )
)

print("\nOverall summary:")

overall = result[
    result["scope"] == "all_lengths"
]

print(
    "Backbones with lower MAE after adding Δdistance:",
    (overall["mae_improvement"] > 0).sum(),
    "/",
    len(overall),
)

print(
    "Mean MAE improvement:",
    overall["mae_improvement"].mean(),
)

print(
    "Median MAE improvement:",
    overall["mae_improvement"].median(),
)