"""Reproduce the diagnostic values in Tables 1 and 2 and Fig. 1.

Article title: Cost-sensitive spectral sampling algorithms for randomized block Kaczmarz methods
Journal: Numerical Algorithms
Author: Shreyhaan Sarkar
Affiliation: Cornell University, Ithaca, NY, United States
Corresponding author email: sms736@cornell.edu

Tests A and B are closed-form coordinate-block tests. Test C is a
non-diagonal sparse block-catalogue test. The randomized simulation in Test C
uses fixed random seeds and reports median work over 20 trials.
"""

import numpy as np
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
from scipy.optimize import minimize_scalar

TOL = 1e-6

def cert_work(mu, cbar, tol=TOL):
    return cbar * np.ceil(np.log(tol) / np.log(1 - mu))


def print_closed_form_tests():
    rows = [
        ("A", "uniform/norm", 0.00870, 1.00),
        ("A", "spectral design", 0.05000, 1.00),
        ("B", "uniform", 0.16667, 34.17),
        ("B", "rank/norm", 0.57500, 150.25),
        ("B", "cost-blind spectral", 0.75000, 200.00),
        ("B", "cost-sensitive spectral", 0.05000, 1.00),
    ]
    for test, rule, mu, cbar in rows:
        eta = mu / cbar
        print(f"{test:1s} {rule:24s} mu={mu:8.5f} cbar={cbar:8.2f} eta={eta:10.7f} Wcert={cert_work(mu,cbar):8.0f}")


def sparse_random_A(m=250, n=30, nnz_per_row=5, seed=7):
    rng = np.random.default_rng(seed)
    A = np.zeros((m, n))
    for i in range(m):
        idx = rng.choice(n, size=nnz_per_row, replace=False)
        vals = rng.normal(size=nnz_per_row)
        A[i, idx] = vals
        A[i] /= np.linalg.norm(A[i])
    return A


def make_catalogue(A, seed=17):
    rng = np.random.default_rng(seed)
    m = A.shape[0]
    blocks = []
    for row in rng.choice(m, size=70, replace=False):
        blocks.append(np.array([row], dtype=int))
    for k, count in [(5, 45), (10, 25), (25, 10)]:
        for _ in range(count):
            if rng.random() < 0.5 and m - k > 0:
                start = rng.integers(0, m - k + 1)
                rows = np.arange(start, start + k)
            else:
                rows = rng.choice(m, size=k, replace=False)
            blocks.append(np.array(rows, dtype=int))
    return blocks


def projector_for_block(A, rows, tol=1e-12):
    M = A[rows, :]
    _, s, Vt = np.linalg.svd(M, full_matrices=False)
    if len(s) == 0:
        return np.zeros((A.shape[1], A.shape[1]))
    r = np.sum(s > tol * max(M.shape) * s[0])
    Q = Vt[:r, :].T
    return Q @ Q.T


def costs(A, blocks, gamma=0.25):
    out = []
    for rows in blocks:
        k = len(rows)
        out.append(np.count_nonzero(A[rows, :]) + gamma * k**3)
    return np.array(out, dtype=float)


def mu_of_p(Ps, p):
    H = np.einsum("i,ijk->jk", p, Ps)
    H = (H + H.T) / 2
    return np.linalg.eigvalsh(H)[0]


def frank_wolfe_e_design(Ks, max_iter=300, tol=1e-6):
    M = Ks.shape[0]
    q = np.ones(M) / M
    H = Ks.mean(axis=0)
    for it in range(max_iter):
        H = (H + H.T) / 2
        vals, vecs = np.linalg.eigh(H)
        lam = vals[0]
        V = vecs[:, vals <= lam + 1e-8 * max(1, abs(lam))]
        Z = V @ V.T / V.shape[1]
        g = np.einsum("ijk,jk->i", Ks, Z)
        j = int(np.argmax(g))
        if g[j] - lam <= tol:
            break
        D = Ks[j] - H
        def neg_lam(a):
            X = H + a * D
            return -np.linalg.eigvalsh((X + X.T) / 2)[0]
        res = minimize_scalar(neg_lam, bounds=(0, 1), method="bounded", options={"xatol": 1e-5})
        a = float(res.x)
        if a < 1e-8:
            a = 2.0 / (it + 2.0)
        q *= (1 - a)
        q[j] += a
        H += a * D
    return q


def simulate(Ps, p, c, ntrials=20, seed=107):
    rng = np.random.default_rng(seed)
    works = []
    for _ in range(ntrials):
        e = rng.normal(size=Ps.shape[1])
        e /= np.linalg.norm(e)
        work = 0.0
        while np.dot(e, e) > TOL:
            idx = rng.choice(Ps.shape[0], p=p)
            e -= Ps[idx] @ e
            work += c[idx]
        works.append(work)
    return float(np.median(works))


def simulate_error_curve(Ps, p, c, work_grid, ntrials=20, seed=307):
    """Return median relative squared error on a common work grid."""
    rng = np.random.default_rng(seed)
    curves = []
    max_work = float(work_grid[-1])
    for _ in range(ntrials):
        e = rng.normal(size=Ps.shape[1])
        e /= np.linalg.norm(e)
        works = [0.0]
        errs = [1.0]
        work = 0.0
        # Continue long enough for slow rules to appear on the plot.
        while work < max_work and errs[-1] > 1e-12:
            idx = rng.choice(Ps.shape[0], p=p)
            e -= Ps[idx] @ e
            work += c[idx]
            works.append(work)
            errs.append(float(np.dot(e, e)))
        # Step-function interpolation: value after all updates whose work is <= grid point.
        vals = []
        j = 0
        for W in work_grid:
            while j + 1 < len(works) and works[j + 1] <= W:
                j += 1
            vals.append(errs[j])
        curves.append(vals)
    return np.median(np.array(curves), axis=0)


def run_test_c(make_plot=True):
    A = sparse_random_A()
    blocks = make_catalogue(A)
    c = costs(A, blocks)
    Ps = np.array([projector_for_block(A, rows) for rows in blocks])
    M = Ps.shape[0]
    p_unif = np.ones(M) / M
    ranks = np.array([np.trace(P) for P in Ps])
    p_rank = ranks / ranks.sum()
    p_cost_blind = frank_wolfe_e_design(Ps, max_iter=200, tol=1e-5)
    q_cost = frank_wolfe_e_design(Ps / c[:, None, None], max_iter=300, tol=1e-6)
    w = q_cost / c
    p_cost = w / w.sum()

    rows = [
        ("C", "uniform", p_unif),
        ("C", "norm/rank", p_rank),
        ("C", "cost-blind spectral", p_cost_blind),
        ("C", "cost-sensitive spectral", p_cost),
    ]
    for test, rule, p in rows:
        mu = mu_of_p(Ps, p)
        cbar = float(np.dot(p, c))
        eta = mu / cbar
        med = simulate(Ps, p, c)
        print(f"{test:1s} {rule:24s} mu={mu:8.5f} cbar={cbar:8.2f} eta={eta:10.7f} Wcert={cert_work(mu,cbar):8.0f} Wemp={med:8.0f}")

    if make_plot:
        work_grid = np.linspace(0, 50000, 401)
        plt.figure(figsize=(6.2, 4.0))
        styles = ["-", "--", "-.", ":"]
        markers = [None, "o", "s", "^"]
        for idx, (_, rule, p) in enumerate(rows):
            curve = simulate_error_curve(Ps, p, c, work_grid, ntrials=20)
            plt.semilogy(work_grid, curve, linestyle=styles[idx], marker=markers[idx],
                         markevery=55, linewidth=1.4, markersize=3.0, label=rule)
        plt.axhline(TOL, linestyle="--", linewidth=1)
        plt.xlabel("cumulative work")
        plt.ylabel(r"median $\|e_k\|^2/\|e_0\|^2$")
        plt.ylim(1e-7, 1.2)
        plt.xlim(0, 50000)
        plt.legend(fontsize=8)
        plt.tight_layout()
        plt.savefig("Fig1.pdf")
        plt.savefig("Fig1.png", dpi=200)



if __name__ == "__main__":
    print_closed_form_tests()
    run_test_c()
