#!/usr/bin/env python3
"""Run public synthetic diagnostics for advanced signal decompositions.

This benchmark is independent of MorphIQ Labs' private market-regime study.
It uses generated signals with known components to test canonical properties
of SWT, CWT/SSWT, the EMD family, VMD, and EWT. It does not use market data,
private feature recipes, production parameters, or product rankings.
"""

from __future__ import annotations

import csv
import json
import os
import shutil
from dataclasses import dataclass
from pathlib import Path

os.environ.setdefault("SSQ_PARALLEL", "0")

import matplotlib.pyplot as plt
import numpy as np
import pyewt
import pywt
from PyEMD import CEEMDAN, EEMD, EMD
from scipy.optimize import linear_sum_assignment
from scipy.signal import find_peaks
from sktime.libs.vmdpy import VMD
from ssqueezepy import Wavelet, ssq_cwt
from ssqueezepy.experimental import scale_to_freq


HERE = Path(__file__).resolve().parent
REPO_CANDIDATE = HERE.parents[1]
if (REPO_CANDIDATE / "package.json").exists():
    REPO_ROOT: Path | None = REPO_CANDIDATE
    RESULTS_DIR = HERE / "results"
    FIGURE_DIR = (
        REPO_ROOT / "public" / "images" / "posts" / "advanced-regime-methods"
    )
    PUBLIC_DATA_DIR: Path | None = (
        REPO_ROOT / "public" / "research" / "advanced-regime-methods"
    )
else:
    REPO_ROOT = None
    RESULTS_DIR = HERE / "results"
    FIGURE_DIR = HERE / "figures"
    PUBLIC_DATA_DIR = None

SEED = 20_260_809
BOOTSTRAP_SEED = 20_260_810
BOOTSTRAP_DRAWS = 4_000
SHIFT_COUNT = 32
TIME_FREQUENCY_TRIALS = 20
EMD_TRIALS = 20
EMD_ENSEMBLE_SIZE = 20
FILTER_BANK_TRIALS = 24

COLORS = {
    "ink": "#0B0D10",
    "graphite": "#3B4048",
    "mute": "#8E949E",
    "signal": "#E85A2C",
    "blue": "#3976A8",
    "green": "#39745B",
    "paper": "#F4F2ED",
    "rule": "#C7C3BA",
}


@dataclass(frozen=True)
class DiagnosticResult:
    experiment: str
    method: str
    metric: str
    estimate: float
    ci_low: float
    ci_high: float
    unit: str
    better: str


@dataclass(frozen=True)
class TimeFrequencyExample:
    signal: np.ndarray
    true_frequencies: np.ndarray
    cwt: np.ndarray
    cwt_frequencies: np.ndarray
    sswt: np.ndarray
    sswt_frequencies: np.ndarray


def bootstrap_ci(
    values: np.ndarray,
    rng: np.random.Generator,
) -> tuple[float, float, float]:
    clean = np.asarray(values, dtype=np.float64)
    clean = clean[np.isfinite(clean)]
    if clean.size == 0:
        return float("nan"), float("nan"), float("nan")
    estimate = float(np.mean(clean))
    samples = rng.choice(clean, size=(BOOTSTRAP_DRAWS, clean.size), replace=True)
    low, high = np.quantile(np.mean(samples, axis=1), [0.025, 0.975])
    return estimate, float(low), float(high)


def result_from_samples(
    experiment: str,
    method: str,
    metric: str,
    values: list[float] | np.ndarray,
    unit: str,
    better: str,
    rng: np.random.Generator,
) -> DiagnosticResult:
    estimate, low, high = bootstrap_ci(np.asarray(values, dtype=np.float64), rng)
    return DiagnosticResult(
        experiment=experiment,
        method=method,
        metric=metric,
        estimate=estimate,
        ci_low=low,
        ci_high=high,
        unit=unit,
        better=better,
    )


def coefficient_energy_shares(coefficients: list[np.ndarray]) -> np.ndarray:
    energies = np.asarray(
        [np.sum(np.abs(coefficient) ** 2) for coefficient in coefficients],
        dtype=np.float64,
    )
    return energies / np.maximum(np.sum(energies), np.finfo(float).tiny)


def run_shift_invariance(
    bootstrap_rng: np.random.Generator,
) -> tuple[list[DiagnosticResult], dict[str, np.ndarray]]:
    length = 512
    sample = np.arange(length, dtype=np.float64)
    low_packet = np.exp(-0.5 * ((sample - 116.0) / 28.0) ** 2) * np.sin(
        2.0 * np.pi * 0.045 * sample
    )
    high_packet = 0.8 * np.exp(-0.5 * ((sample - 318.0) / 17.0) ** 2) * np.sin(
        2.0 * np.pi * 0.205 * sample
    )
    step = 0.55 * (sample >= 224.0)
    signal = low_packet + high_packet + step

    def dwt_shares(values: np.ndarray) -> np.ndarray:
        coefficients = pywt.wavedec(
            values,
            wavelet="db4",
            mode="periodization",
            level=4,
        )
        return coefficient_energy_shares(coefficients)

    def swt_shares(values: np.ndarray) -> np.ndarray:
        coefficients = pywt.swt(
            values,
            wavelet="db4",
            level=4,
            trim_approx=True,
            norm=True,
        )
        return coefficient_energy_shares(list(coefficients))

    reference = {"Decimated DWT": dwt_shares(signal), "SWT": swt_shares(signal)}
    drifts: dict[str, list[float]] = {method: [] for method in reference}
    for shift in range(1, SHIFT_COUNT + 1):
        shifted = np.roll(signal, shift)
        current = {
            "Decimated DWT": dwt_shares(shifted),
            "SWT": swt_shares(shifted),
        }
        for method in reference:
            drifts[method].append(
                float(np.sum(np.abs(current[method] - reference[method])))
            )

    results: list[DiagnosticResult] = []
    for method, values in drifts.items():
        results.append(
            result_from_samples(
                "Shift invariance",
                method,
                "Mean scale-energy drift",
                values,
                "L1 distance",
                "lower",
                bootstrap_rng,
            )
        )
        maximum = float(np.max(values))
        results.append(
            DiagnosticResult(
                experiment="Shift invariance",
                method=method,
                metric="Maximum scale-energy drift",
                estimate=maximum,
                ci_low=maximum,
                ci_high=maximum,
                unit="L1 distance",
                better="lower",
            )
        )

    return results, {
        method: np.asarray(values, dtype=np.float64) for method, values in drifts.items()
    }


def generate_chirp_mixture(
    rng: np.random.Generator,
    length: int = 768,
) -> tuple[np.ndarray, np.ndarray]:
    position = np.linspace(0.0, 1.0, length, endpoint=False)
    low_frequency = 0.035 + 0.080 * position
    high_frequency = 0.260 - 0.065 * position + 0.008 * np.sin(2.0 * np.pi * position)
    frequencies = np.vstack((low_frequency, high_frequency))

    phases = rng.uniform(0.0, 2.0 * np.pi, size=2)
    low_phase = phases[0] + 2.0 * np.pi * np.cumsum(low_frequency)
    high_phase = phases[1] + 2.0 * np.pi * np.cumsum(high_frequency)
    low_amplitude = 0.85 + 0.15 * np.cos(2.0 * np.pi * position)
    high_amplitude = 0.72 + 0.12 * np.sin(4.0 * np.pi * position)
    clean = low_amplitude * np.sin(low_phase) + high_amplitude * np.sin(high_phase)
    return clean + 0.38 * rng.normal(size=length), frequencies


def time_frequency_metrics(
    coefficients: np.ndarray,
    frequencies: np.ndarray,
    true_frequencies: np.ndarray,
) -> tuple[float, float]:
    energy = np.abs(np.asarray(coefficients)) ** 2
    frequency_axis = np.asarray(frequencies, dtype=np.float64).reshape(-1)
    central = slice(72, energy.shape[1] - 72)
    energy = energy[:, central]
    truth = true_frequencies[:, central]

    valid_band = (frequency_axis >= 0.015) & (frequency_axis <= 0.32)
    energy = energy[valid_band]
    frequency_axis = frequency_axis[valid_band]
    energy /= np.maximum(np.sum(energy, axis=0, keepdims=True), np.finfo(float).tiny)

    close_to_ridge = np.zeros_like(energy, dtype=bool)
    for ridge in truth:
        close_to_ridge |= np.abs(frequency_axis[:, np.newaxis] - ridge) <= 0.010
    ridge_share = float(np.mean(np.sum(np.where(close_to_ridge, energy, 0.0), axis=0)))

    entropy = -np.sum(energy * np.log(np.maximum(energy, np.finfo(float).tiny)), axis=0)
    effective_bins = float(np.mean(np.exp(entropy)))
    return ridge_share, effective_bins


def run_time_frequency_concentration(
    rng: np.random.Generator,
    bootstrap_rng: np.random.Generator,
) -> tuple[list[DiagnosticResult], TimeFrequencyExample]:
    wavelet = Wavelet(("gmw", {"beta": 20}))
    ridge_shares = {"CWT": [], "SSWT": []}
    effective_bins = {"CWT": [], "SSWT": []}
    example: TimeFrequencyExample | None = None

    for trial in range(TIME_FREQUENCY_TRIALS):
        signal, true_frequencies = generate_chirp_mixture(rng)
        sswt, cwt, sswt_frequencies, scales = ssq_cwt(
            signal,
            wavelet=wavelet,
            nv=24,
            fs=1.0,
            padtype="reflect",
            astensor=False,
        )
        cwt_frequencies = np.asarray(
            scale_to_freq(scales, wavelet, len(signal), fs=1.0), dtype=np.float64
        )
        sswt_frequencies = np.asarray(sswt_frequencies, dtype=np.float64)

        for method, coefficients, frequencies in (
            ("CWT", cwt, cwt_frequencies),
            ("SSWT", sswt, sswt_frequencies),
        ):
            ridge_share, bins = time_frequency_metrics(
                coefficients,
                frequencies,
                true_frequencies,
            )
            ridge_shares[method].append(ridge_share)
            effective_bins[method].append(bins)

        if trial == 0:
            example = TimeFrequencyExample(
                signal=signal,
                true_frequencies=true_frequencies,
                cwt=np.asarray(cwt),
                cwt_frequencies=cwt_frequencies,
                sswt=np.asarray(sswt),
                sswt_frequencies=sswt_frequencies,
            )

    results: list[DiagnosticResult] = []
    for method in ("CWT", "SSWT"):
        results.append(
            result_from_samples(
                "Time-frequency concentration",
                method,
                "Energy near true ridges",
                ridge_shares[method],
                "share",
                "higher",
                bootstrap_rng,
            )
        )
        results.append(
            result_from_samples(
                "Time-frequency concentration",
                method,
                "Effective occupied frequency bins",
                effective_bins[method],
                "bins",
                "lower",
                bootstrap_rng,
            )
        )

    assert example is not None
    return results, example


def generate_intermittent_modes(
    rng: np.random.Generator,
    length: int = 512,
) -> tuple[np.ndarray, np.ndarray]:
    sample = np.arange(length, dtype=np.float64)
    phases = rng.uniform(0.0, 2.0 * np.pi, size=2)
    low_mode = np.sin(2.0 * np.pi * 0.045 * sample + phases[0])
    gate = 0.5 * (
        np.tanh((sample - 0.22 * length) / 11.0)
        - np.tanh((sample - 0.78 * length) / 11.0)
    )
    high_mode = 0.78 * gate * np.sin(2.0 * np.pi * 0.180 * sample + phases[1])
    components = np.vstack((low_mode, high_mode))
    signal = np.sum(components, axis=0) + 0.34 * rng.normal(size=length)
    return signal, components


def correlation_matrix(targets: np.ndarray, modes: np.ndarray) -> np.ndarray:
    centered_targets = targets - np.mean(targets, axis=1, keepdims=True)
    centered_modes = modes - np.mean(modes, axis=1, keepdims=True)
    numerator = centered_targets @ centered_modes.T
    denominator = np.sqrt(
        np.sum(centered_targets**2, axis=1, keepdims=True)
        * np.sum(centered_modes**2, axis=1, keepdims=True).T
    )
    return numerator / np.maximum(denominator, np.finfo(float).tiny)


def mode_isolation_metrics(
    targets: np.ndarray,
    modes: np.ndarray,
) -> tuple[float, float]:
    correlations = np.abs(correlation_matrix(targets, modes))
    target_indices, mode_indices = linear_sum_assignment(-correlations)
    assigned_r2 = float(
        np.sum(correlations[target_indices, mode_indices] ** 2) / targets.shape[0]
    )
    best_modes = np.argmax(correlations, axis=1)
    collision = float(len(np.unique(best_modes)) < len(best_modes))
    return assigned_r2, collision


def run_emd_family(
    rng: np.random.Generator,
    bootstrap_rng: np.random.Generator,
) -> list[DiagnosticResult]:
    assigned_r2: dict[str, list[float]] = {
        "EMD": [],
        "EEMD": [],
        "CEEMDAN": [],
    }
    collisions: dict[str, list[float]] = {method: [] for method in assigned_r2}

    for trial in range(EMD_TRIALS):
        signal, targets = generate_intermittent_modes(rng)

        eemd = EEMD(trials=EMD_ENSEMBLE_SIZE, noise_width=0.05, parallel=False)
        eemd.noise_seed(SEED + 1_000 + trial)
        ceemdan = CEEMDAN(trials=EMD_ENSEMBLE_SIZE, epsilon=0.05, parallel=False)
        ceemdan.noise_seed(SEED + 2_000 + trial)

        decompositions = {
            "EMD": EMD().emd(signal, max_imf=6),
            "EEMD": eemd.eemd(signal, max_imf=6),
            "CEEMDAN": ceemdan.ceemdan(signal, max_imf=6),
        }
        for method, modes in decompositions.items():
            r2, collision = mode_isolation_metrics(targets, np.asarray(modes))
            assigned_r2[method].append(r2)
            collisions[method].append(collision)

    results: list[DiagnosticResult] = []
    for method in assigned_r2:
        results.append(
            result_from_samples(
                "Noise-assisted mode separation",
                method,
                "Distinct-mode assigned R-squared",
                assigned_r2[method],
                "R-squared",
                "higher",
                bootstrap_rng,
            )
        )
        results.append(
            result_from_samples(
                "Noise-assisted mode separation",
                method,
                "Mode-collision rate",
                collisions[method],
                "share",
                "lower",
                bootstrap_rng,
            )
        )
    return results


def generate_stationary_modes(
    rng: np.random.Generator,
    length: int = 512,
) -> tuple[np.ndarray, np.ndarray, np.ndarray]:
    sample = np.arange(length, dtype=np.float64)
    frequencies = np.asarray([0.035, 0.105, 0.235], dtype=np.float64)
    frequencies += rng.uniform(-0.008, 0.008, size=len(frequencies))
    amplitudes = np.asarray([1.0, 0.72, 0.52], dtype=np.float64)
    amplitudes *= rng.uniform(0.92, 1.08, size=len(amplitudes))
    phases = rng.uniform(0.0, 2.0 * np.pi, size=len(frequencies))
    components = np.vstack(
        [
            amplitude * np.sin(2.0 * np.pi * frequency * sample + phase)
            for amplitude, frequency, phase in zip(
                amplitudes, frequencies, phases, strict=True
            )
        ]
    )
    signal = np.sum(components, axis=0) + 0.28 * rng.normal(size=length)
    return signal, components, frequencies


def dominant_frequency(mode: np.ndarray) -> float:
    values = np.real(np.asarray(mode, dtype=np.complex128))
    values = values - np.mean(values)
    spectrum = np.abs(np.fft.rfft(values * np.hanning(len(values)))) ** 2
    frequencies = np.fft.rfftfreq(len(values))
    spectrum[0] = 0.0
    return float(frequencies[int(np.argmax(spectrum))])


def frequency_recovery_metrics(
    modes: np.ndarray | list[np.ndarray],
    true_frequencies: np.ndarray,
) -> tuple[float, float, float, float]:
    mode_array = np.asarray(modes)
    mode_frequencies = np.asarray(
        [dominant_frequency(mode) for mode in mode_array], dtype=np.float64
    )
    return frequency_candidate_metrics(mode_frequencies, true_frequencies)


def frequency_candidate_metrics(
    candidate_frequencies: np.ndarray,
    true_frequencies: np.ndarray,
) -> tuple[float, float, float, float]:
    mode_frequencies = np.asarray(candidate_frequencies, dtype=np.float64)
    costs = np.abs(true_frequencies[:, np.newaxis] - mode_frequencies[np.newaxis, :])
    targets, matched_modes = linear_sum_assignment(costs)
    matched_errors = costs[targets, matched_modes]
    recovered = int(np.count_nonzero(matched_errors <= 0.012))
    hit_rate = float(recovered / len(true_frequencies))
    mode_precision = float(recovered / len(mode_frequencies))
    missing = len(true_frequencies) - len(targets)
    mean_error = float(
        (np.sum(matched_errors) + 0.5 * missing) / len(true_frequencies)
    )
    extra_modes = float(max(0, len(mode_frequencies) - len(true_frequencies)))
    return hit_rate, mode_precision, mean_error, extra_modes


def periodogram_peaks(signal: np.ndarray, count: int = 3) -> np.ndarray:
    values = signal - np.mean(signal)
    spectrum = np.abs(np.fft.rfft(values * np.hanning(len(values)))) ** 2
    frequencies = np.fft.rfftfreq(len(values))
    peaks, _ = find_peaks(spectrum, distance=12)
    strongest = peaks[np.argsort(spectrum[peaks])[-count:]]
    return frequencies[strongest]


def ewt_modes(signal: np.ndarray, detector: str) -> list[np.ndarray]:
    params = pyewt.Default_Params()
    params.update(
        {
            "N": 3,
            "detect": detector,
            "Completion": False,
            "wavname": "meyer",
        }
    )
    modes, _, _ = pyewt.ewt1d(signal, params)
    return [np.real(np.asarray(mode)) for mode in modes]


def run_adaptive_filter_banks(
    rng: np.random.Generator,
    bootstrap_rng: np.random.Generator,
) -> list[DiagnosticResult]:
    methods = (
        "Periodogram top 3",
        "VMD K=2",
        "VMD K=3",
        "VMD K=4",
        "VMD K=5",
        "EWT local-max N=3",
        "EWT scale-space",
    )
    hit_rates: dict[str, list[float]] = {method: [] for method in methods}
    mode_precision: dict[str, list[float]] = {method: [] for method in methods}
    frequency_errors: dict[str, list[float]] = {method: [] for method in methods}
    extra_modes: dict[str, list[float]] = {method: [] for method in methods}

    for _ in range(FILTER_BANK_TRIALS):
        signal, _, frequencies = generate_stationary_modes(rng)
        decompositions: dict[str, np.ndarray | list[np.ndarray]] = {}
        baseline = frequency_candidate_metrics(periodogram_peaks(signal), frequencies)
        hit_rates["Periodogram top 3"].append(baseline[0])
        mode_precision["Periodogram top 3"].append(baseline[1])
        frequency_errors["Periodogram top 3"].append(baseline[2])
        extra_modes["Periodogram top 3"].append(baseline[3])
        for mode_count in (2, 3, 4, 5):
            modes, _, _ = VMD(
                signal,
                alpha=2_000,
                tau=0.0,
                K=mode_count,
                DC=0,
                init=1,
                tol=1e-7,
            )
            decompositions[f"VMD K={mode_count}"] = modes
        decompositions["EWT local-max N=3"] = ewt_modes(signal, "locmax")
        decompositions["EWT scale-space"] = ewt_modes(signal, "scalespace")

        for method, modes in decompositions.items():
            hit_rate, precision, frequency_error, extra = frequency_recovery_metrics(
                modes,
                frequencies,
            )
            hit_rates[method].append(hit_rate)
            mode_precision[method].append(precision)
            frequency_errors[method].append(frequency_error)
            extra_modes[method].append(extra)

    results: list[DiagnosticResult] = []
    for method in methods:
        results.append(
            result_from_samples(
                "Adaptive filter-bank contract",
                method,
                "Known frequencies recovered",
                hit_rates[method],
                "share",
                "higher",
                bootstrap_rng,
            )
        )
        results.append(
            result_from_samples(
                "Adaptive filter-bank contract",
                method,
                "Returned-mode precision",
                mode_precision[method],
                "share",
                "higher",
                bootstrap_rng,
            )
        )
        results.append(
            result_from_samples(
                "Adaptive filter-bank contract",
                method,
                "Mean assigned frequency error",
                frequency_errors[method],
                "cycles/sample",
                "lower",
                bootstrap_rng,
            )
        )
        results.append(
            result_from_samples(
                "Adaptive filter-bank contract",
                method,
                "Extra returned modes",
                extra_modes[method],
                "modes",
                "lower",
                bootstrap_rng,
            )
        )
    return results


def configure_matplotlib() -> None:
    plt.rcParams.update(
        {
            "figure.facecolor": COLORS["paper"],
            "axes.facecolor": COLORS["paper"],
            "savefig.facecolor": COLORS["paper"],
            "text.color": COLORS["ink"],
            "axes.labelcolor": COLORS["graphite"],
            "axes.edgecolor": COLORS["rule"],
            "xtick.color": COLORS["graphite"],
            "ytick.color": COLORS["graphite"],
            "font.family": "sans-serif",
            "font.size": 10,
            "axes.titleweight": "bold",
            "axes.titlelocation": "left",
        }
    )


def plot_time_frequency(example: TimeFrequencyExample) -> None:
    configure_matplotlib()
    figure, axes = plt.subplots(3, 1, figsize=(12, 9.2), sharex=True)
    sample = np.arange(len(example.signal))

    axes[0].plot(sample, example.signal, color=COLORS["ink"], linewidth=0.8)
    axes[0].set_title("Noisy two-component AM-FM signal")
    axes[0].set_ylabel("Amplitude")
    axes[0].grid(axis="y", color=COLORS["rule"], linewidth=0.5, alpha=0.65)

    for axis, coefficients, frequencies, title in (
        (axes[1], example.cwt, example.cwt_frequencies, "CWT magnitude"),
        (axes[2], example.sswt, example.sswt_frequencies, "Synchrosqueezed CWT magnitude"),
    ):
        magnitude = np.log1p(6.0 * np.abs(coefficients))
        axis.pcolormesh(
            sample,
            frequencies,
            magnitude,
            shading="auto",
            cmap="magma",
            rasterized=True,
        )
        axis.plot(sample, example.true_frequencies[0], color="white", linewidth=0.8)
        axis.plot(sample, example.true_frequencies[1], color="white", linewidth=0.8)
        axis.set_ylim(0.01, 0.31)
        axis.set_ylabel("Cycles / sample")
        axis.set_title(title)

    axes[-1].set_xlabel("Observation")
    for axis in axes:
        axis.spines[["top", "right"]].set_visible(False)
    figure.suptitle(
        "Same wavelet, different concentration",
        x=0.06,
        ha="left",
        fontsize=15,
        fontweight="bold",
    )
    figure.tight_layout(rect=(0, 0, 1, 0.965), h_pad=1.5)
    FIGURE_DIR.mkdir(parents=True, exist_ok=True)
    figure.savefig(FIGURE_DIR / "cwt-vs-sswt.png", dpi=180, bbox_inches="tight")
    plt.close(figure)


def select_result(
    results: list[DiagnosticResult],
    experiment: str,
    method: str,
    metric: str,
) -> DiagnosticResult:
    return next(
        result
        for result in results
        if result.experiment == experiment
        and result.method == method
        and result.metric == metric
    )


def draw_summary_panels(
    axes: np.ndarray,
    results: list[DiagnosticResult],
) -> None:
    shift_axis, time_frequency_axis, emd_axis, filter_bank_axis = list(
        np.asarray(axes, dtype=object).reshape(-1)
    )
    shift_methods = ("Decimated DWT", "SWT")
    shift_rows = [
        select_result(
            results,
            "Shift invariance",
            method,
            "Mean scale-energy drift",
        )
        for method in shift_methods
    ]
    shift_axis.bar(
        shift_methods,
        [max(row.estimate, 1e-16) for row in shift_rows],
        color=(COLORS["graphite"], COLORS["signal"]),
    )
    shift_axis.set_yscale("log")
    shift_axis.set_ylabel("Mean L1 drift, log scale")
    shift_axis.set_title("Shift the same signal")

    tf_methods = ("CWT", "SSWT")
    tf_rows = [
        select_result(
            results,
            "Time-frequency concentration",
            method,
            "Energy near true ridges",
        )
        for method in tf_methods
    ]
    time_frequency_axis.bar(
        tf_methods,
        [100.0 * row.estimate for row in tf_rows],
        color=(COLORS["graphite"], COLORS["signal"]),
    )
    time_frequency_axis.set_ylim(0, 100)
    time_frequency_axis.set_ylabel("Energy within tolerance (%)")
    time_frequency_axis.set_title("Track known instantaneous frequencies")

    emd_methods = ("EMD", "EEMD", "CEEMDAN")
    emd_rows = [
        select_result(
            results,
            "Noise-assisted mode separation",
            method,
            "Distinct-mode assigned R-squared",
        )
        for method in emd_methods
    ]
    emd_axis.bar(
        emd_methods,
        [row.estimate for row in emd_rows],
        color=(COLORS["graphite"], COLORS["blue"], COLORS["signal"]),
    )
    emd_axis.set_ylim(0, 1)
    emd_axis.set_ylabel("Assigned component R-squared")
    emd_axis.set_title("Separate intermittent modes under noise")

    bank_methods = (
        "Periodogram top 3",
        "VMD K=2",
        "VMD K=3",
        "VMD K=4",
        "VMD K=5",
        "EWT local-max N=3",
        "EWT scale-space",
    )
    bank_rows = [
        select_result(
            results,
            "Adaptive filter-bank contract",
            method,
            "Known frequencies recovered",
        )
        for method in bank_methods
    ]
    precision_rows = [
        select_result(
            results,
            "Adaptive filter-bank contract",
            method,
            "Returned-mode precision",
        )
        for method in bank_methods
    ]
    labels = ("FFT 3", "VMD 2", "VMD 3", "VMD 4", "VMD 5", "EWT 3", "EWT auto")
    positions = np.arange(len(labels))
    width = 0.38
    filter_bank_axis.bar(
        positions - width / 2,
        [100.0 * row.estimate for row in bank_rows],
        width=width,
        color=COLORS["signal"],
        label="Known-frequency recall",
    )
    filter_bank_axis.bar(
        positions + width / 2,
        [100.0 * row.estimate for row in precision_rows],
        width=width,
        color=COLORS["graphite"],
        label="Returned-mode precision",
    )
    filter_bank_axis.set_xticks(positions, labels)
    filter_bank_axis.set_ylim(0, 100)
    filter_bank_axis.set_ylabel("Share (%)")
    filter_bank_axis.set_title("Specify or infer the filter-bank contract")
    filter_bank_axis.legend(frameon=False, fontsize=8, loc="lower left")

    for axis in (shift_axis, time_frequency_axis, emd_axis, filter_bank_axis):
        axis.grid(axis="y", color=COLORS["rule"], linewidth=0.5, alpha=0.65)
        axis.set_axisbelow(True)
        axis.spines[["top", "right"]].set_visible(False)
        axis.tick_params(axis="x", labelrotation=12)


def plot_summary(results: list[DiagnosticResult]) -> None:
    configure_matplotlib()
    figure, axes = plt.subplots(2, 2, figsize=(12, 8.8))
    draw_summary_panels(axes, results)

    figure.suptitle(
        "Four diagnostics, four different questions",
        x=0.06,
        ha="left",
        fontsize=15,
        fontweight="bold",
    )
    figure.tight_layout(rect=(0, 0, 1, 0.96), h_pad=2.5, w_pad=2.5)
    FIGURE_DIR.mkdir(parents=True, exist_ok=True)
    figure.savefig(FIGURE_DIR / "advanced-method-summary.png", dpi=180, bbox_inches="tight")
    plt.close(figure)

    mobile_figure, mobile_axes = plt.subplots(4, 1, figsize=(8.2, 14.5))
    draw_summary_panels(mobile_axes, results)
    mobile_figure.suptitle(
        "Four diagnostics, four different questions",
        x=0.08,
        ha="left",
        fontsize=15,
        fontweight="bold",
    )
    mobile_figure.tight_layout(rect=(0, 0, 1, 0.97), h_pad=2.8)
    mobile_figure.savefig(
        FIGURE_DIR / "advanced-method-summary-mobile.png",
        dpi=180,
        bbox_inches="tight",
    )
    plt.close(mobile_figure)


def write_results(results: list[DiagnosticResult]) -> None:
    rows = [
        {
            "experiment": result.experiment,
            "method": result.method,
            "metric": result.metric,
            "estimate": result.estimate,
            "ci_low": result.ci_low,
            "ci_high": result.ci_high,
            "unit": result.unit,
            "better": result.better,
        }
        for result in results
    ]
    metadata = {
        "experiment": "Synthetic diagnostics for advanced signal decompositions",
        "disclosure": "public",
        "market_data_used": False,
        "seed": SEED,
        "bootstrap_seed": BOOTSTRAP_SEED,
        "bootstrap_draws": BOOTSTRAP_DRAWS,
        "replicates": {
            "shift_count": SHIFT_COUNT,
            "time_frequency_trials": TIME_FREQUENCY_TRIALS,
            "emd_trials": EMD_TRIALS,
            "emd_ensemble_size": EMD_ENSEMBLE_SIZE,
            "filter_bank_trials": FILTER_BANK_TRIALS,
        },
        "results": rows,
    }

    destinations = [RESULTS_DIR]
    if PUBLIC_DATA_DIR is not None:
        destinations.append(PUBLIC_DATA_DIR)
    for destination in destinations:
        destination.mkdir(parents=True, exist_ok=True)
        with (destination / "summary.csv").open("w", newline="", encoding="utf-8") as handle:
            writer = csv.DictWriter(
                handle,
                fieldnames=list(rows[0].keys()),
                lineterminator="\n",
            )
            writer.writeheader()
            writer.writerows(rows)
        (destination / "summary.json").write_text(
            json.dumps(metadata, indent=2) + "\n",
            encoding="utf-8",
        )


def main() -> None:
    rng = np.random.default_rng(SEED)
    bootstrap_rng = np.random.default_rng(BOOTSTRAP_SEED)

    shift_results, _ = run_shift_invariance(bootstrap_rng)
    time_frequency_results, example = run_time_frequency_concentration(
        rng,
        bootstrap_rng,
    )
    emd_results = run_emd_family(rng, bootstrap_rng)
    filter_bank_results = run_adaptive_filter_banks(rng, bootstrap_rng)
    results = [
        *shift_results,
        *time_frequency_results,
        *emd_results,
        *filter_bank_results,
    ]

    write_results(results)
    plot_time_frequency(example)
    plot_summary(results)
    if PUBLIC_DATA_DIR is not None:
        PUBLIC_DATA_DIR.mkdir(parents=True, exist_ok=True)
        shutil.copy2(__file__, PUBLIC_DATA_DIR / "benchmark.py")
        shutil.copy2(HERE / "requirements.txt", PUBLIC_DATA_DIR / "requirements.txt")

    for result in results:
        print(
            f"{result.experiment} | {result.method} | {result.metric}: "
            f"{result.estimate:.6f} [{result.ci_low:.6f}, {result.ci_high:.6f}]"
        )


if __name__ == "__main__":
    main()
