#!/usr/bin/env python3
"""Exploratory held-out extension from Nahor through Moses."""

from __future__ import annotations

import math
from pathlib import Path

import numpy as np
import pandas as pd

from run_analysis import (
    AMBIENTS,
    N_GRID,
    LIFE_ALT,
    LIFE_PRIMARY,
    emission,
    fit_training,
    transition_matrix,
)


ROOT = Path(__file__).resolve().parent
RESULTS = ROOT / "results"
LATER_NAMES = ["Nahor", "Terah", "Abraham", "Isaac", "Jacob", "Levi", "Kohath", "Amram", "Moses"]
LATER_LIVES = np.array([208, 205, 175, 180, 147, 137, 133, 137, 120], dtype=float)


def score_extension(training_y, ambient, n, fit, model):
    d = training_y[0] - ambient
    means = ambient + d * np.arange(n + 1) / n
    alpha = fit["posterior_k7"].copy()
    logscore = 0.0
    rows = []
    for j, (name, obs) in enumerate(zip(LATER_NAMES, LATER_LIVES)):
        if model == "continuation":
            h = fit["retention"]
        elif model == "babel_pulse_0.5":
            h = 0.5 if j == 0 else fit["retention"]
        elif model == "babel_persistent_0.5":
            h = 0.5
        else:
            raise ValueError(model)
        prior = alpha @ transition_matrix(n, h)
        pred_mean = float(prior @ means)
        pred_sd = float(np.sqrt(prior @ ((means - pred_mean) ** 2) + fit["sigma"] ** 2))
        e = emission(obs, means, fit["sigma"])
        density = float(prior @ e)
        logscore += math.log(max(density, 1e-300))
        alpha = prior * e
        alpha /= max(alpha.sum(), 1e-300)
        rows.append(
            {
                "name": name,
                "observed": obs,
                "transition_retention": h,
                "predicted_mean": pred_mean,
                "predicted_sd": pred_sd,
                "sequential_logscore": math.log(max(density, 1e-300)),
                "cumulative_logscore": logscore,
            }
        )
    return logscore, rows


def main():
    raw_rows = []
    trajectory_rows = []
    profiles = []
    for series, y in [("Shelah_533_primary", LIFE_PRIMARY), ("Shelah_460_sensitivity", LIFE_ALT)]:
        for ambient in AMBIENTS:
            for n in N_GRID:
                fit = fit_training(y, ambient, n)
                profiles.append(
                    {
                        "series": series,
                        "ambient": ambient,
                        "N": n,
                        "train_loglik": fit["train_loglik"],
                    }
                )
                for model in ["continuation", "babel_pulse_0.5", "babel_persistent_0.5"]:
                    score, trajectory = score_extension(y, ambient, n, fit, model)
                    raw_rows.append(
                        {
                            "series": series,
                            "ambient": ambient,
                            "N": n,
                            "model": model,
                            "extension_logscore": score,
                        }
                    )
                    for row in trajectory:
                        trajectory_rows.append(
                            {
                                "series": series,
                                "ambient": ambient,
                                "N": n,
                                "model": model,
                                **row,
                            }
                        )

    raw = pd.DataFrame(raw_rows)
    profiles = pd.DataFrame(profiles)
    trajectories = pd.DataFrame(trajectory_rows)
    summaries = []
    for (series, ambient), p in profiles.groupby(["series", "ambient"]):
        p = p.set_index("N")
        weights = np.exp(p.train_loglik - p.train_loglik.max())
        weights /= weights.sum()
        model_scores = {}
        for model in ["continuation", "babel_pulse_0.5", "babel_persistent_0.5"]:
            q = raw[(raw.series == series) & (raw.ambient == ambient) & (raw.model == model)].set_index("N")
            model_scores[model] = math.log(float(np.sum(weights * np.exp(q.extension_logscore))))
        summaries.append(
            {
                "series": series,
                "ambient": ambient,
                "logscore_continuation": model_scores["continuation"],
                "logscore_babel_pulse": model_scores["babel_pulse_0.5"],
                "logscore_babel_persistent": model_scores["babel_persistent_0.5"],
                "predictive_ratio_pulse_vs_continuation": math.exp(model_scores["babel_pulse_0.5"] - model_scores["continuation"]),
                "predictive_ratio_persistent_vs_continuation": math.exp(model_scores["babel_persistent_0.5"] - model_scores["continuation"]),
            }
        )

    summary = pd.DataFrame(summaries)
    raw.to_csv(RESULTS / "extended_validation_raw.csv", index=False)
    trajectories.to_csv(RESULTS / "extended_validation_trajectories.csv", index=False)
    summary.to_csv(RESULTS / "extended_validation_summary.csv", index=False)
    print(summary.to_string(index=False))


if __name__ == "__main__":
    main()
