#!python
#
#    JKS - Measurement database system
#    Copyright (C) 2013-2026  Christoph Lehner (christoph.lehner@ur.de, https://github.com/lehner/jks)
#
#    This program is free software; you can redistribute it and/or modify
#    it under the terms of the GNU General Public License as published by
#    the Free Software Foundation; either version 2 of the License, or
#    (at your option) any later version.
#
#    This program is distributed in the hope that it will be useful,
#    but WITHOUT ANY WARRANTY; without even the implied warranty of
#    MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE.  See the
#    GNU General Public License for more details.
#
#    You should have received a copy of the GNU General Public License along
#    with this program; if not, write to the Free Software Foundation, Inc.,
#    51 Franklin Street, Fifth Floor, Boston, MA 02110-1301 USA.
#
import jks, sys, os, time, numpy as np

if len(sys.argv) not in (9, 10):
    print("%s database.jks tag_in list_of_weights_in list_of_weights_out omega_grid tag_lower tag_upper tag_out [dchi2]" % sys.argv[0])
    print("")
    print("- list_of_weights_X accepts a list of")
    print("  a) in integer t for which exp(-t*omega) will be used as the spectral weight")
    print("  b) a string tag to a weight")
    print("")
    print("- dchi2 (default 1) is the profile-likelihood threshold: the quoted band is")
    print("    { k.z : lower <= z <= upper, chi2(z) <= chi2_min(C) + dchi2 }")
    print("  i.e. the set of values of the output weight that are within dchi2 of the")
    print("  best description of the data within the bounds.  dchi2 = 1 is 68% for one")
    print("  parameter of interest; use 3.84 for 95%.")
    print("")
    print("- tag_lower / tag_upper (default 0 / +inf) are bounds  lower <= z <= upper  on the")
    print("  spectral weight at each omega_grid node, the tags are an array on omega_grid in the")
    print("  database; inf entries in upper leave that node unbounded.  With the defaults")
    print("  this is the positivity band of jks_plsa.")
    print("")
    print("- Note that an error on omega_grid is not propagated, use multiple omega_grid if needed")
    sys.exit(0)

db, tag_in, list_w_in, list_w_out, omega_grid, tag_lower, tag_upper, tag_out = sys.argv[1:9]
dchi2 = float(sys.argv[9]) if len(sys.argv) == 10 else 1.0
assert dchi2 > 0.0, "dchi2 must be positive"

res = jks.resamples(db)

omega_grid_arr = res.get(omega_grid).mean()

print(f"Computing bounds on the output weights using energy grid [{min(omega_grid_arr)},..,{max(omega_grid_arr)}] ({len(omega_grid_arr)} resolution)")

# Every weight is handed to the solver in rescaled grid units: column j of all
# input AND output weights is divided by d_j = max_i |w_in_i(omega_j)|, i.e.
# z_j -> d_j z_j.  k.z and the admissible set are invariant under this, so the
# result is unchanged, but the columns are O(1) instead of spanning
# exp(-t omega) over many orders of magnitude.  log d_j is formed in log space
# for integer weights, so exp(-t omega) is never evaluated where it would
# underflow.  A node no input weight sees keeps d_j = 1.
w_in, w_out = eval(list_w_in), eval(list_w_out)
assert isinstance(w_in, list) and isinstance(w_out, list)

def get_weight(e):
    assert isinstance(e, str), "weight not found"
    ww = np.asarray(res.get(e).mean(), float)
    assert ww.shape == omega_grid_arr.shape, f"{e} not found"
    return ww

with np.errstate(divide="ignore"):
    log_d = np.max([-e * omega_grid_arr if isinstance(e, int) else np.log(np.abs(get_weight(e)))
                    for e in w_in], axis=0)
log_d[~np.isfinite(log_d)] = 0.0

def process_weight_element(e, i):
    if isinstance(e, int):
        return np.exp(-e*omega_grid_arr - log_d), e
    return get_weight(e) * np.exp(-log_d), i

def process_weight(w):
    t = [process_weight_element(e, i) for i, e in enumerate(w)]
    return [x[0] for x in t], [x[1] for x in t]

list_w_in, sel_w_in = process_weight(w_in)
list_w_out, sel_w_out = process_weight(w_out)

list_w_in = np.asarray(list_w_in, float)
list_w_out = np.asarray(list_w_out, float)

# create a temporary version of c_in projected
def get_c_in(r):
    raw_in = r[tag_in]
    return [ raw_in[s] for s in sel_w_in ]

c_in = res.apply(get_c_in)
c_in_mn = c_in.mean()
c_mn = np.asarray(c_in_mn, float)

# full input covariance.  The diagonal-only version mis-states the tension by
# O(1) on correlated data (lqcd: chi_min 1.09 diag vs 1.80 full) and it
# mis-states the propagated statistical error by up to a factor 1.6, because
# the jackknife blocks carry the correlations whatever metric is used here.
c_in_cv = np.array(c_in.cov())

# protect against problematic covariance matrices
if "JKS_CORRELATION_STRENGTH" in os.environ:
    lam = float(os.environ["JKS_CORRELATION_STRENGTH"])
    assert lam <= 1.0 and 0.0 <= lam, "JKS_CORRELATION_STRENGTH needs to be between 0 and 1"
    c_in_cv = lam * c_in_cv + (1 - lam) * np.diag(c_in_cv.diagonal())
evs, _ = np.linalg.eig(c_in_cv)
evs = sorted(evs.real)
kappa = abs(evs[-1] / evs[0])
print(f"Covariance condition number: {kappa:e}")
assert kappa < 1e10, f"Problematic condition number of covariance {kappa}"


# bounds on z, taken to the rescaled units of the weights (z -> d z, see above)
def get_bound(which):
    b = np.asarray(res.get(which).mean(), float)
    assert b.shape == omega_grid_arr.shape, f"{which} is not an array on the omega grid"
    return b, b * np.exp(log_d)
lower_z, lower_s = get_bound(tag_lower)
upper_z, upper_s = get_bound(tag_upper)

solver = jks.bounded_laplace.solver(list_w_in, c_in_cv, list_w_out, omega_grid_arr, lower=lower_s, upper=upper_s)

# --- tension between the data and the bounds ------------------------------------
# exact (NNLS), not the faceted lower bound of the old min_chi
z_ml, chi_min_in = solver.nnls(c_in_mn)
ml = list_w_out @ z_ml               # maximum-likelihood value of each output

# A nonzero chi_min is the NORMAL state, not a red flag: a correlator dominated by
# its ground state sits close to the boundary of the moment cone (a single
# exponential is an extreme ray of it), so an ordinary sub-sigma fluctuation of the
# central value lands outside.
#
# The record band is therefore taken at a RELATIVE threshold,
#     chi_r(C)^2 = chi_min(C)^2 + dchi2,
# so that  [lo, hi] = { y : min_{z>=0, k.z=y} chi2(z) <= chi2_min + dchi2 }
# is the Delta-chi^2 profile-likelihood interval of the functional k.z.  Unlike
# an absolute threshold it is never empty -- neither at the central value nor at
# any resample -- so no feasibility floor and no linearised fallback are needed,
# and for dchi2 = 1 it is calibrated: in closed-form Monte Carlo over the full
# analysis it covers the truth 66-70% of the time for outputs the data
# determine, rising (conservatively) to ~85% for extrapolated ones.
chi_r = (chi_min_in ** 2 + dchi2) ** 0.5
print(f"Tension with the bounds: chi_min = {chi_min_in:.4f}  (chi^2_min = {chi_min_in**2:.4f}, "
      f"{len(c_in_mn)} input weights)")
print(f"Record chi = {chi_r:.4f}  (chi^2 = chi^2_min + {dchi2:g}, profile-likelihood threshold)")
if chi_min_in == 0.0:
    print(" the data are reproduced exactly by a spectrum within the bounds on this grid")
elif chi_min_in ** 2 <= len(c_in_mn):
    print(f" chi^2_min = {chi_min_in**2:.2f} for {len(c_in_mn)} input weights -- ordinary noise on")
    print(f" the central value, the data are compatible with a spectrum within the bounds")
else:
    print(f" the data lie {chi_min_in:.3f} sigma outside the admissible cone.  Check whether")
    print(f" the tension is structural: if chi_min jumps when an input weight is added,")
    print(f" the window cannot fit it rather than it being noise.  The band is still the")
    print(f" honest Delta-chi^2 interval, but it sits on a thin cap and its centre will")
    print(f" move a lot from resample to resample (watch stat/half below).")
print("")

# --- exact admissible range of every output weight (deterministic) -----------
t0 = time.time()
lo, hi, _, info = solver.bands(c_in_mn, chi_r)
t1 = time.time()
io = len(list_w_out)

bad = [j for j in range(io) if info["status"][j] not in ("ok", "zero_kernel")]
if bad:
    for j in bad:
        st = info["status"][j]
        if st in ("unbounded_hi", "unbounded_lo", "unbounded"):
            print(f"ERROR: output weight {sel_w_out[j]!r} is UNBOUNDED under the bounds "
                  f"(status {st}): the data weights leave a free direction in the "
                  f"spectrum (a grid node no input weight sees, or input weights that "
                  f"cancel on part of the grid) and this output weight is nonzero along it.")
        else:
            print(f"ERROR: output weight {sel_w_out[j]!r}: dual status {st} (no band) -- "
                  f"raise dchi2 (chi_min={chi_min_in:.3f}) or check the inputs.")
    print("  Refusing to store non-finite values in the database.")
    sys.exit(1)
central = 0.5 * (lo + hi)
half = 0.5 * (hi - lo)

print(f"Admissible range computed in {t1 - t0:.2f} s (max primal-dual gap "
      f"{max(v for v in info['gap_rel'] if v == v):.1e} relative)")
print("")

# --- store -------------------------------------------------------------------
# central = band midpoint; stat blocks = the EXACT band center at each jackknife
# resample, each at ITS OWN record chi (so the threshold stays a Delta-chi^2 and
# the band never collapses).
#
# The band is then NOT an independent systematic on top of the statistical
# error: for an output weight the data determine, the half-width of the
# Delta-chi^2 = 1 band IS the 1-sigma statistical error (in the unconstrained
# linear case max/min k.z over chi^2 <= 1 is exactly central +- sigma_k), and
# the jackknife wobble of the centre is the same sigma_k.  Adding them in
# quadrature double counts and was what made the output error at an input
# timeslice sqrt(2) times the input error.  Only the part of the band that the
# statistics does not already account for is stored as the '!band' variation,
#     sys = sqrt(max(half^2 - stat^2, 0)),
# so that the total from tcov() is exactly max(stat, half) -- the calibrated
# combination: 68% where statistics dominates, conservative where the
# band does.
#
# A resample whose band does not solve is an error, as at the central value.
# The exact solve has not failed on any data vector tried (up to 10 sigma from
# the mean, dchi2 0.1..3.84, up to 20 input weights, up to 4000 grid nodes), so
# a silent linearised stand-in would only hide a real problem.
chis = []

def get_c_out(r):
    c_r = np.array(get_c_in(r))
    if np.array_equal(c_r, c_mn):
        return central                      # the unresampled point, already solved
    chi_r_r = (solver.chi_min(c_r) ** 2 + dchi2) ** 0.5
    chis.append(chi_r_r)
    l_r, h_r, _, i_r = solver.bands(c_r, chi_r_r)
    bad = [(sel_w_out[j], s) for j, s in enumerate(i_r["status"]) if s not in ("ok", "zero_kernel")]
    if bad:
        print(f"ERROR: no band at a resample (record chi {chi_r_r:.3f}) for output "
              f"weight(s) {bad}.")
        print("  Refusing to store non-finite values in the database.")
        sys.exit(1)
    return 0.5 * (l_r + h_r)

t0 = time.time()
jk = res.apply(get_c_out)
t1 = time.time()
stat = np.sqrt(np.clip(np.array(jk.cov()).diagonal(), 0.0, None))   # exact, from the blocks
sys_ = np.sqrt(np.clip(half ** 2 - stat ** 2, 0.0, None))
total = np.sqrt(stat ** 2 + sys_ ** 2)     # == max(stat, half), evaluated the
                                           # same way tcov() will evaluate it

nres = len(chis)
cc = np.array(chis) if nres else np.array([chi_r])
print(f"{nres} of {res.N} blocks re-solved in {t1 - t0:.2f} s (the rest leave the input "
      f"weights unchanged); record chi over them: "
      f"min {cc.min():.3f}  median {np.median(cc):.3f}  max {cc.max():.3f}")
print("")

band_dom = half > 1.5 * stat                    # the bounds set the answer
ill_posed = (half > 0) & (stat > 1.5 * half)    # the centre is not an estimate
print(" out    central       stat        sys         total     stat/half   ML k.z*      (c-ML)/tot")
for j in range(io):
    s = f" {sel_w_out[j]!s:>6}  {central[j]: .5e}  {stat[j]: .3e}  {sys_[j]: .3e}  {total[j]: .3e}  "
    s += f"{stat[j] / half[j]:8.2f}" if half[j] > 0 else "       -"
    s += f"  {ml[j]: .5e}"
    s += f"  {(central[j] - ml[j]) / total[j]: 8.2f}" if total[j] > 0 else "         -"
    if band_dom[j]:
        s += "  [band-dominated]"
    elif ill_posed[j]:
        s += "  [ill-posed]"
    print(s)
print("")
print(" total = max(stat, half-width of the Delta-chi^2 band); stat is the exact")
print(" per-resample wobble of the band center (jackknife, gaussian where it dominates),")
print(" sys = sqrt(half^2 - stat^2) is stored as the '!band' variation so that")
print(" tcov() reproduces the total.")
print(" stat/half ~ 1 is the normal, data-dominated case: for an output the data")
print(" determine, the Delta-chi^2 = 1 half-width IS the 1-sigma statistical error.")
print(" [band-dominated] (half > 1.5 stat): the bounds, not the statistics, set the")
print(" answer -- the error is a range, not a gaussian sigma.")
print(" [ill-posed]      (stat > 1.5 half): the band center moves further than the band")
print(" is wide, so the center is not an estimate of anything stable -- quote the")
print(" interval (central +- total) and do not use the central value.")
print(" ML k.z* is the maximum-likelihood spectrum within the bounds evaluated on the output")
print(" weight; unlike the band center it does not move with dchi2.  A large")
print(" |c-ML|/total flags an output where the band is strongly asymmetric.")
nbd = int(ill_posed.sum())
if nbd:
    print("")
    print(f"WARNING: {nbd} of {io} output weights are ill-posed (stat > 1.5 half):")
    print("         their central values are not meaningful, only the intervals are.")
print("")

desc = ("Bounded-spectrum band beyond the statistical error (range = central +- shift), "
        "dual certificate, profile threshold chi^2 = chi^2_min + %g (chi_min=%.3f, "
        "record chi=%.3f), omega grid tag '%s' [%g..%g] x %d nodes, full input "
        "covariance, deterministic.  stat blocks = exact band center per resample at "
        "its own record chi.  shift = "
        "sqrt(max(half^2-stat^2,0)) so that stat (+) shift = max(stat, half).  "
        "Input weights: %s.  Box prior: lower %s, upper %s"
        % (dchi2, chi_min_in, chi_r, omega_grid, min(omega_grid_arr), max(omega_grid_arr),
           len(omega_grid_arr), sel_w_in, tag_lower, tag_upper))
jk = jk.addvar("band", central + sys_, desc)
res.add(tag_out, jk)

res.save(db)

print("Done")
