#!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, numpy as np
if len(sys.argv) not in (8, 9, 10):
    print("%s database.jks tag_in list_of_weights_in list_of_weights_out omega_grid tag_out lambda [alpha [p]]" % sys.argv[0])
    print("")
    print("Linear (Hansen-Lupo-Tantalo, arXiv:1903.06476) estimate of the output weights.")
    print("Same calling convention and storage as jks_plsa; no positivity is used.")
    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("- each output weight k is approximated by  kbar(omega) = sum_i g_i w_in_i(omega),")
    print("  g minimizing")
    print("    W[g] = (1 - lambda) A[g]/A[0] + lambda B[g]/C_0^2")
    print("    A[g] = int_grid d omega exp(alpha omega) omega^(2p) (k(omega) - kbar(omega))^2")
    print("    B[g] = g^T Cov g")
    print("  (integral by the trapezoidal rule on omega_grid, C_0 the mean of the first")
    print("  input weight).  The output is g.C, its jackknife blocks g.C^(b).")
    print("")
    print("- lambda in (0, 1): small lambda = faithful kernel, large statistical error.")
    print("  Only the statistical error is stored; the lambda dependence (and the")
    print("  omega_grid dependence) is left to be estimated by running at several values.")
    print("- alpha (default 0): exponential weight of the kernel norm A.")
    print("- p (default 0): power weight of the kernel norm A.  This is the same as")
    print("  scaling all input AND output weights by omega^p (the data and the")
    print("  estimated quantity are unchanged, only the norm of k - kbar is).")
    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_out = sys.argv[1:7]
lam = float(sys.argv[7])
alpha = float(sys.argv[8]) if len(sys.argv) >= 9 else 0.0
pw = float(sys.argv[9]) if len(sys.argv) == 10 else 0.0
assert 0.0 < lam < 1.0, "lambda must be in (0, 1)"

res = jks.resamples(db)

omega_grid_arr = np.asarray(res.get(omega_grid).mean(), float)
K = len(omega_grid_arr)

print(f"HLT linear estimate on energy grid [{min(omega_grid_arr)},..,{max(omega_grid_arr)}] "
      f"({K} resolution), lambda = {lam:g}, alpha = {alpha:g}, p = {pw:g}")

# Unlike jks_plsa, no column rescaling: the kernel norm A[g] is a statement about
# the functions on the grid in physical units.  exp(-t omega) underflowing to 0
# is harmless here.
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

def process_weight_element(e, i):
    if isinstance(e, int):
        return np.exp(-e * omega_grid_arr), e
    return get_weight(e), 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)      # (m, K)
list_w_out = np.asarray(list_w_out, float)    # (p, K)

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_mn = np.asarray(c_in.mean(), float)
c_in_cv = np.array(c_in.cov())

if "JKS_CORRELATION_STRENGTH" in os.environ:
    lcs = float(os.environ["JKS_CORRELATION_STRENGTH"])
    assert lcs <= 1.0 and 0.0 <= lcs, "JKS_CORRELATION_STRENGTH needs to be between 0 and 1"
    c_in_cv = lcs * c_in_cv + (1 - lcs) * 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}"
try:
    hlt = jks.hlt.solver(w_in, list_w_in, w_out, list_w_out, c_in_cv, c_mn, omega_grid_arr,
                         lam, alpha, pw)
except RuntimeError as e:
    print(f"ERROR: {e}.")
    print("  Refusing to store non-converged values in the database.")
    sys.exit(1)
A_rel, B_rel, digits = hlt.A_rel, hlt.B_rel, hlt.digits
p = hlt.p

central = hlt.estimate(c_mn)

# stat blocks: the estimate is linear in the data, so the jackknife blocks are g.C^(b)
jk = res.apply(lambda r: hlt.estimate(get_c_in(r)))
stat = np.sqrt(np.clip(np.array(jk.cov()).diagonal(), 0.0, None))

dmp = [d for d in digits if d is not None]
if dmp:
    print(f"mpmath used for {len(dmp)} of {p} output weights ({min(dmp)}..{max(dmp)} digits, "
          f"each verified against a solve with 20 fewer digits)")
else:
    print("double precision sufficient for all output weights")
print("")
print(" out    central       stat        A/A0        B/C0^2    digits")
for j in range(p):
    s = f" {sel_w_out[j]!s:>6}  {central[j]: .5e}  {stat[j]: .3e}"
    s += f"  {A_rel[j]: .3e}  {B_rel[j]: .3e}"
    s += f"  {digits[j]:6d}" if digits[j] is not None else "     dbl"
    print(s)
print("")
print(" central = g.C with g the HLT coefficients at lambda; stat = jackknife error of")
print(" g.C (exactly the linearly propagated error, the estimate is linear in C).")
print(" No lambda systematic is included: estimate it from runs at several lambda.")
print(" A/A0 is the relative L2 mismatch of the reconstructed kernel: the estimate")
print(" is of int rho kbar, not of int rho k.  Nothing bounds int rho (k - kbar)")
print(" here -- compare with jks_plsa, whose band contains that mismatch rigorously.")
print(" digits: mpmath working precision of the solve (dbl: double precision,")
print(" used where the condition bound is <= 1e6).")
print("")

res.add(tag_out, jk)

res.save(db)

print("Done")
