File indexing completed on 2026-08-30 08:17:27
0001
0002 """Export scale maps to TH2D for ROOT plotting."""
0003 import argparse, numpy as np, torch, uproot
0004 from ml_momentum_calibration_reso_v2 import (KappaNet, ETA_RANGE, semi_even_pt_grid, load_stage1)
0005
0006 DEFAULT_PT_VALUES = (0.3, 0.5, 1.0, 2.0, 3.0)
0007
0008 ap = argparse.ArgumentParser()
0009 ap.add_argument("--model", default="calib_out_upgrade20260731/model.pt")
0010 ap.add_argument("--out", default="calib_out_upgrade20260731/kappa_maps.root")
0011 ap.add_argument("--hidden", type=int, default=48)
0012 ap.add_argument("--pt_mode", choices=("fixed", "even"), default="fixed",
0013 help="fixed: use --pt_values; even: use --slices/--pt_min/--pt_max")
0014 ap.add_argument("--pt_values", nargs="+", default=DEFAULT_PT_VALUES,
0015 help="fixed pT slices in GeV; accepts spaces or commas")
0016 ap.add_argument("--slices", type=int, default=9)
0017 ap.add_argument("--pt_min", type=float, default=0.10)
0018 ap.add_argument("--pt_max", type=float, default=5.0)
0019 ap.add_argument("--n_eta", type=int, default=200)
0020 ap.add_argument("--n_phi", type=int, default=200)
0021 a = ap.parse_args()
0022
0023 def parse_pt_values(values):
0024 pts = []
0025 for value in values:
0026 if isinstance(value, (int, float)):
0027 pts.append(float(value))
0028 continue
0029 pts.extend(float(item) for item in value.split(",") if item.strip())
0030 return np.asarray(pts, dtype=float)
0031
0032 def pt_slices(args):
0033 if args.pt_mode == "even":
0034 return semi_even_pt_grid(args.pt_min, args.pt_max, args.slices)
0035
0036 pts = parse_pt_values(args.pt_values)
0037 if pts.size == 0:
0038 raise ValueError("--pt_values must contain at least one pT slice")
0039 if np.any(~np.isfinite(pts)) or np.any(pts <= 0.0):
0040 raise ValueError("--pt_values must be finite positive values")
0041 return pts
0042
0043 model = load_stage1(a.model, hidden=a.hidden).eval()
0044
0045 eta_e = np.linspace(*ETA_RANGE, a.n_eta + 1)
0046 phi_e = np.linspace(-np.pi, np.pi, a.n_phi + 1)
0047 E, P = np.meshgrid(0.5 * (eta_e[1:] + eta_e[:-1]),
0048 0.5 * (phi_e[1:] + phi_e[:-1]), indexing="ij")
0049 t = lambda x: torch.tensor(np.ascontiguousarray(x.ravel()), dtype=torch.float64)
0050 r = lambda v: v.numpy().reshape(E.shape)
0051
0052 pts = pt_slices(a)
0053 tag = lambda pt: f"pt{pt:.2f}".replace(".", "p")
0054
0055 with uproot.recreate(a.out) as f, torch.no_grad():
0056 for pt in pts:
0057 PT = t(np.full(E.shape, pt))
0058 eps, dlt = model.eps_delta(PT, t(E), t(P))
0059 f[f"eps_{tag(pt)}"] = (r(eps), eta_e, phi_e)
0060 f[f"delta_{tag(pt)}"] = (r(dlt), eta_e, phi_e)
0061 for q, lab in ((+1.0, "qplus"), (-1.0, "qminus")):
0062 k, _, _ = model.kappa(PT, t(E), t(P), torch.full_like(PT, q))
0063 f[f"kappa_{lab}_{tag(pt)}"] = (r(k), eta_e, phi_e)
0064 f["slices"] = {"pt": pts, "curv": 1.0 / pts}
0065
0066 print(f"wrote {len(pts)} pT slices ({a.n_eta}x{a.n_phi}) to {a.out}")