#!/usr/bin/env python3
# SPDX-License-Identifier: MIT
"""Optimize a frequency-scaled Rayleigh-Ritz basis for a quartic oscillator.

The dimensionless Hamiltonian is

    h = p**2 / 2 + q**2 / 2 + g * q**4,  g >= 0.

Only NumPy is required for the benchmark tables. Matplotlib is optional and is
used for the two-panel diagnostic plot. The calculation is deterministic.
"""

from __future__ import annotations

import argparse
import csv
import platform
from dataclasses import dataclass
from math import exp, log, sqrt
from pathlib import Path

import numpy as np


@dataclass(frozen=True)
class RitzResult:
    energy: float
    vector: np.ndarray
    virial_residual: float
    eigenpair_residual: float


@dataclass(frozen=True)
class Minimum:
    u: float
    energy: float

    @property
    def y(self) -> float:
        return exp(self.u)


def projected_matrices(
    k: int,
    y: float,
    g: float,
    *,
    padding: int = 6,
) -> tuple[np.ndarray, np.ndarray, np.ndarray, np.ndarray]:
    """Return H, T, q^2, and q^4 in {|0;y>, |2;y>, ..., |2K-2;y>}.

    q^4 changes oscillator quantum number by at most four. The ladder-operator
    matrices are therefore built above the retained boundary and projected only
    after the powers have been formed. This avoids a truncation-edge artifact.
    """

    if k < 1:
        raise ValueError("k must be at least one")
    if y <= 0:
        raise ValueError("y must be positive")
    if g < 0:
        raise ValueError("this benchmark assumes a stable coupling g >= 0")
    if padding < 4:
        raise ValueError("padding must be at least four for exact q^4 elements")

    n_max = 2 * (k - 1)
    full_dimension = n_max + padding + 1
    annihilation = np.zeros((full_dimension, full_dimension), dtype=float)
    n = np.arange(1, full_dimension)
    annihilation[n - 1, n] = np.sqrt(n)

    q_full = (annihilation + annihilation.T) / sqrt(2.0 * y)
    q2_full = q_full @ q_full
    q4_full = q2_full @ q2_full

    even = np.arange(0, 2 * k, 2)
    q2 = q2_full[np.ix_(even, even)]
    q4 = q4_full[np.ix_(even, even)]

    h_y = np.diag(y * (even + 0.5))
    kinetic = h_y - 0.5 * y**2 * q2
    hamiltonian = h_y + 0.5 * (1.0 - y**2) * q2 + g * q4
    return hamiltonian, kinetic, q2, q4


def ground_state(k: int, y: float, g: float, *, padding: int = 6) -> RitzResult:
    hamiltonian, kinetic, q2, q4 = projected_matrices(
        k, y, g, padding=padding
    )
    eigenvalues, eigenvectors = np.linalg.eigh(hamiltonian)
    energy = float(eigenvalues[0])
    vector = eigenvectors[:, 0]

    def expectation(operator: np.ndarray) -> float:
        return float(vector @ (operator @ vector))

    virial = (
        2.0 * expectation(kinetic)
        - expectation(q2)
        - 4.0 * g * expectation(q4)
    )
    residual = float(np.linalg.norm(hamiltonian @ vector - energy * vector))
    return RitzResult(energy, vector, virial, residual)


def energy(k: int, u: float, g: float, *, padding: int = 6) -> float:
    """Lowest Ritz value as a function of u = log(y)."""

    return ground_state(k, exp(u), g, padding=padding).energy


def gaussian_frequency(g: float) -> float:
    """Physical root of y^3 - y - 6g = 0 for the one-Gaussian ansatz."""

    lower = 1.0
    upper = max(2.0, (6.0 * g + 1.0) ** (1.0 / 3.0) + 1.0)
    for _ in range(100):
        middle = 0.5 * (lower + upper)
        if middle**3 - middle - 6.0 * g > 0:
            upper = middle
        else:
            lower = middle
    return 0.5 * (lower + upper)


def refine_basin(
    k: int,
    g: float,
    u_left: float,
    u_right: float,
    *,
    tolerance: float = 2.0e-12,
) -> Minimum:
    """Golden-section refinement inside one basin identified by a grid scan."""

    golden_ratio = 0.5 * (1.0 + sqrt(5.0))
    u_c = u_right - (u_right - u_left) / golden_ratio
    u_d = u_left + (u_right - u_left) / golden_ratio
    e_c = energy(k, u_c, g)
    e_d = energy(k, u_d, g)

    while u_right - u_left > tolerance:
        if e_c < e_d:
            u_right, u_d, e_d = u_d, u_c, e_c
            u_c = u_right - (u_right - u_left) / golden_ratio
            e_c = energy(k, u_c, g)
        else:
            u_left, u_c, e_c = u_c, u_d, e_d
            u_d = u_left + (u_right - u_left) / golden_ratio
            e_d = energy(k, u_d, g)

    u_min = 0.5 * (u_left + u_right)
    return Minimum(u_min, energy(k, u_min, g))


def global_optimize(
    k: int,
    g: float,
    *,
    grid_points: int = 2401,
) -> tuple[Minimum, list[Minimum]]:
    """Scan globally in log-frequency, refine every basin, and choose the best."""

    if grid_points < 101:
        raise ValueError("grid_points must be at least 101")

    center = log(gaussian_frequency(g))
    u_min = min(-2.0, center - 2.5)
    u_max = max(4.0, center + 2.5)
    u_grid = np.linspace(u_min, u_max, grid_points)
    e_grid = np.array([energy(k, u, g) for u in u_grid])

    candidates: list[Minimum] = []
    for i in range(1, grid_points - 1):
        is_basin = e_grid[i] <= e_grid[i - 1] and e_grid[i] < e_grid[i + 1]
        if is_basin:
            candidates.append(
                refine_basin(k, g, float(u_grid[i - 1]), float(u_grid[i + 1]))
            )

    if not candidates:
        i = int(np.argmin(e_grid))
        if i == 0 or i == grid_points - 1:
            raise RuntimeError("the scan range does not bracket a minimum")
        candidates.append(
            refine_basin(k, g, float(u_grid[i - 1]), float(u_grid[i + 1]))
        )

    candidates.sort(key=lambda item: item.energy)
    return candidates[0], candidates


def curvature(k: int, g: float, u_min: float, *, step: float = 1.0e-2) -> float:
    return (
        energy(k, u_min + step, g)
        - 2.0 * energy(k, u_min, g)
        + energy(k, u_min - step, g)
    ) / step**2


def reference_energy(g: float, *, k_reference: int = 64) -> float:
    """Large-basis reference at the analytic Gaussian preconditioning scale."""

    return ground_state(k_reference, gaussian_frequency(g), g).energy


def parity_coupling_norm(*, dimension: int = 24, y: float = 3.0, g: float = 1.0) -> float:
    """Norm of the even-odd block in an independently built full matrix."""

    padded_dimension = dimension + 6
    annihilation = np.zeros((padded_dimension, padded_dimension), dtype=float)
    n = np.arange(1, padded_dimension)
    annihilation[n - 1, n] = np.sqrt(n)
    q = (annihilation + annihilation.T) / sqrt(2.0 * y)
    q2 = q @ q
    q4 = q2 @ q2
    numbers = np.arange(padded_dimension)
    hamiltonian = (
        np.diag(y * (numbers + 0.5))
        + 0.5 * (1.0 - y**2) * q2
        + g * q4
    )[:dimension, :dimension]
    even = np.arange(0, dimension, 2)
    odd = np.arange(1, dimension, 2)
    return float(np.linalg.norm(hamiltonian[np.ix_(even, odd)]))


def benchmark_rows(
    couplings: list[float],
    k_max: int,
    grid_points: int,
) -> list[dict[str, float | int]]:
    rows: list[dict[str, float | int]] = []
    for g in couplings:
        reference = reference_energy(g)
        for k in range(1, k_max + 1):
            optimum, basins = global_optimize(k, g, grid_points=grid_points)
            optimized = ground_state(k, optimum.y, g)
            fixed = ground_state(k, 1.0, g)
            rows.append(
                {
                    "g": g,
                    "K": k,
                    "local_minima": len(basins),
                    "y_opt": optimum.y,
                    "E_opt": optimized.energy,
                    "E_fixed_y_1": fixed.energy,
                    "E_reference": reference,
                    "error_opt": optimized.energy - reference,
                    "error_fixed_y_1": fixed.energy - reference,
                    "curvature_u": curvature(k, g, optimum.u),
                    "virial_residual": optimized.virial_residual,
                    "eigenpair_residual": optimized.eigenpair_residual,
                }
            )
    return rows


def write_benchmark_csv(rows: list[dict[str, float | int]], path: Path) -> None:
    with path.open("w", newline="", encoding="utf-8") as stream:
        writer = csv.DictWriter(stream, fieldnames=list(rows[0]))
        writer.writeheader()
        writer.writerows(rows)


def write_landscape_csv(path: Path, *, g: float = 1.0) -> None:
    reference = reference_energy(g)
    with path.open("w", newline="", encoding="utf-8") as stream:
        writer = csv.writer(stream)
        writer.writerow(["u", "y", "K", "E", "E_minus_reference"])
        for u in np.linspace(-0.5, 1.7, 221):
            for k in (1, 2, 4, 8):
                ritz_energy = energy(k, float(u), g)
                writer.writerow([u, exp(float(u)), k, ritz_energy, ritz_energy - reference])


def validate(rows: list[dict[str, float | int]]) -> None:
    harmonic = ground_state(1, 1.0, 0.0)
    if abs(harmonic.energy - 0.5) > 5.0e-15:
        raise AssertionError("harmonic-limit check failed")

    hamiltonian, _, _, _ = projected_matrices(8, 3.0, 1.0)
    if np.linalg.norm(hamiltonian - hamiltonian.T) > 1.0e-13:
        raise AssertionError("Hermiticity check failed")
    if parity_coupling_norm() > 1.0e-13:
        raise AssertionError("parity-block check failed")

    e_padding_6 = ground_state(8, 3.0, 1.0, padding=6).energy
    e_padding_10 = ground_state(8, 3.0, 1.0, padding=10).energy
    if abs(e_padding_6 - e_padding_10) > 2.0e-14:
        raise AssertionError("operator-padding check failed")

    grouped: dict[float, list[dict[str, float | int]]] = {}
    for row in rows:
        grouped.setdefault(float(row["g"]), []).append(row)
        if float(row["error_opt"]) < -5.0e-12:
            raise AssertionError("variational upper-bound check failed")
        if abs(float(row["virial_residual"])) > 2.0e-7:
            raise AssertionError("virial stationarity check failed")
        if float(row["eigenpair_residual"]) > 2.0e-12:
            raise AssertionError("eigenpair residual check failed")

    for coupling_rows in grouped.values():
        energies = [float(row["E_opt"]) for row in coupling_rows]
        if any(right > left + 5.0e-13 for left, right in zip(energies, energies[1:])):
            raise AssertionError("nested-subspace monotonicity check failed")


def plot_diagnostics(
    rows: list[dict[str, float | int]],
    output_path: Path,
    *,
    g: float = 1.0,
) -> None:
    try:
        import matplotlib

        matplotlib.use("Agg")
        import matplotlib.pyplot as plt
    except ImportError as error:
        raise RuntimeError("Matplotlib is required unless --skip-plot is used") from error

    reference = reference_energy(g)
    u_grid = np.linspace(-0.5, 1.7, 221)
    figure, axes = plt.subplots(2, 1, figsize=(7.2, 8.0), constrained_layout=True)

    for k, style in zip((1, 2, 4, 8), ("-", "--", "-.", ":")):
        errors = np.array([energy(k, float(u), g) - reference for u in u_grid])
        axes[0].semilogy(u_grid, np.maximum(errors, np.finfo(float).eps), style, label=f"K={k}")
    axes[0].set_xlabel("log-frequency u = ln y")
    axes[0].set_ylabel("Ritz excess E_K(u) - E_ref")
    axes[0].grid(True, which="both", alpha=0.25)
    axes[0].legend()

    selected = [row for row in rows if float(row["g"]) == g]
    k_values = np.array([int(row["K"]) for row in selected])
    optimized_errors = np.array([max(float(row["error_opt"]), np.finfo(float).eps) for row in selected])
    fixed_errors = np.array([max(float(row["error_fixed_y_1"]), np.finfo(float).eps) for row in selected])
    axes[1].semilogy(k_values, fixed_errors, "o-", label="fixed y=1")
    axes[1].semilogy(k_values, optimized_errors, "s--", label="globally optimized y")
    axes[1].set_xlabel("number K of even basis states")
    axes[1].set_ylabel("energy excess above E_ref")
    axes[1].grid(True, which="both", alpha=0.25)
    axes[1].legend()
    figure.savefig(output_path, dpi=180)
    plt.close(figure)


def main() -> None:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--output-dir", type=Path, default=Path("variational-output"))
    parser.add_argument("--k-max", type=int, default=10)
    parser.add_argument("--grid-points", type=int, default=2401)
    parser.add_argument("--skip-plot", action="store_true")
    args = parser.parse_args()

    args.output_dir.mkdir(parents=True, exist_ok=True)
    couplings = [0.1, 1.0, 10.0]
    rows = benchmark_rows(couplings, args.k_max, args.grid_points)
    validate(rows)
    write_benchmark_csv(rows, args.output_dir / "variational-benchmark.csv")
    write_landscape_csv(args.output_dir / "variational-landscape.csv")
    if not args.skip_plot:
        try:
            plot_diagnostics(rows, args.output_dir / "variational-optimization.png")
        except RuntimeError as error:
            print(f"Plot skipped: {error}")

    print(f"Python {platform.python_version()}")
    print(f"NumPy {np.__version__}")
    print("Validation: PASS")
    print(f"Output: {args.output_dir.resolve()}")


if __name__ == "__main__":
    main()
