#!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 list_of_tags_out lambda [alpha [p]]" % sys.argv[0])
    print("")
    print("Stores the kernels that jks_hlt would use (Hansen-Lupo-Tantalo,")
    print("arXiv:1903.06476) as weights on omega_grid, without evaluating them.")
    print("Arguments as for jks_hlt, except that list_of_tags_out holds one tag per")
    print("output 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).  kbar is stored under the corresponding tag of")
    print("  list_of_tags_out.  g depends on the data only through Cov and C_0, so kbar")
    print("  is the same in every jackknife block.")
    print("")
    print("- kbar lies exactly in the span of the input weights.  Used as an output")
    print("  weight of jks_plsa with the same input weights, the band of int rho kbar")
    print("  shows what spectral positivity adds to the HLT estimate of the same")
    print("  quantity (without positivity its error would be the HLT statistical error).")
    print("")
    print("- lambda in (0, 1), alpha (default 0), p (default 0): as for jks_hlt.")
    sys.exit(0)

db, tag_in, list_w_in, list_w_out, omega_grid, list_tags_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 kernels 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}"
tags_out = eval(list_tags_out)
assert isinstance(tags_out, list) and all(isinstance(t, str) for t in tags_out), \
    "list_of_tags_out must be a list of strings"
assert len(tags_out) == len(w_out), \
    f"list_of_tags_out has {len(tags_out)} entries for {len(w_out)} output weights"
assert len(set(tags_out)) == len(tags_out), "list_of_tags_out has duplicate tags"
exist = [t for t in tags_out if t in res.set]
assert not exist, f"tags already in the database: {exist}"

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 kernels in the database.")
    sys.exit(1)

dmp = [d for d in hlt.digits if d is not None]
if dmp:
    print(f"mpmath used for {len(dmp)} of {hlt.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    tag                   A/A0        B/C0^2    digits")
for j in range(hlt.p):
    s = f" {sel_w_out[j]!s:>6}  {tags_out[j]:20s}  {hlt.A_rel[j]: .3e}  {hlt.B_rel[j]: .3e}"
    s += f"  {hlt.digits[j]:6d}" if hlt.digits[j] is not None else "     dbl"
    print(s)
print("")
print(" A/A0 is the relative L2 mismatch of the stored kernel kbar from the output")
print(" weight k; B/C0^2 is the relative statistical error of int rho kbar.")
print(" digits: mpmath working precision of the solve (dbl: double precision,")
print(" used where the condition bound is <= 1e6); kbar is formed at that precision.")
print("")

for j, t in enumerate(tags_out):
    kbar = hlt.kernel(j)
    res.add(t, res.apply(lambda r: kbar))

res.save(db)

print("Done")
