PyAntiGen 1.0.9__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- framework/AntimonyGen.py +48 -0
- framework/RxnDict_to_antimony.py +594 -0
- framework/TelluriumGen.py +16 -0
- framework/__init__.py +0 -0
- framework/antimony_utils.py +294 -0
- framework/cli.py +229 -0
- framework/data_interpolation.py +340 -0
- framework/isotopomer_tools.py +41 -0
- framework/model_generation.py +46 -0
- framework/models.py +189 -0
- framework/module_base.py +42 -0
- framework/pyantigen.py +51 -0
- framework/rate_laws.py +101 -0
- framework/reaction_creation.py +43 -0
- framework/template/Example/AntiGen_paths.py +23 -0
- framework/template/Example/Engine/Anchor_cache.py +193 -0
- framework/template/Example/Engine/Deadline.py +535 -0
- framework/template/Example/Engine/Evaluator.py +1176 -0
- framework/template/Example/Engine/Event_times.py +491 -0
- framework/template/Example/Engine/Fast_profile.py +701 -0
- framework/template/Example/Engine/Fit_cache.py +329 -0
- framework/template/Example/Engine/Identifiability.py +698 -0
- framework/template/Example/Engine/Model_optimize.py +1483 -0
- framework/template/Example/Engine/Model_simulate.py +124 -0
- framework/template/Example/Engine/Nuisance_sensitivity.py +298 -0
- framework/template/Example/Engine/Optimize.py +6862 -0
- framework/template/Example/Engine/Petab_export.py +398 -0
- framework/template/Example/Engine/Preequil_cache.py +361 -0
- framework/template/Example/Engine/Profile_checkpoint.py +399 -0
- framework/template/Example/Engine/Results.py +395 -0
- framework/template/Example/Engine/Sensitivity_analysis.py +320 -0
- framework/template/Example/Engine/Simulate.py +617 -0
- framework/template/Example/Flipflop_reference.py +401 -0
- framework/template/Example/Model_generate.py +37 -0
- framework/template/Example/Model_run.py +261 -0
- framework/template/Example/Modules/Data.py +63 -0
- framework/template/Example/Modules/Events.py +14 -0
- framework/template/Example/Modules/Experiment.py +194 -0
- framework/template/Example/Modules/Loss_config.py +61 -0
- framework/template/Example/Modules/Observed_species.py +3 -0
- framework/template/Example/Modules/Optimizer_settings.py +258 -0
- framework/template/Example/Modules/Plots.py +89 -0
- framework/template/Example/Modules/Solver_settings.py +16 -0
- framework/template/Example/Modules/Update_opt_parameters.py +24 -0
- framework/template/Example/Modules/Update_parameters.py +49 -0
- framework/template/data/ADneg.csv +27 -0
- framework/template/data/ADpos.csv +27 -0
- framework/template/data/Flipflop.csv +29 -0
- framework/template/data/make_flipflop_data.py +174 -0
- pyantigen-1.0.9.dist-info/METADATA +129 -0
- pyantigen-1.0.9.dist-info/RECORD +55 -0
- pyantigen-1.0.9.dist-info/WHEEL +5 -0
- pyantigen-1.0.9.dist-info/entry_points.txt +2 -0
- pyantigen-1.0.9.dist-info/licenses/LICENSE +21 -0
- pyantigen-1.0.9.dist-info/top_level.txt +1 -0
|
@@ -0,0 +1,701 @@
|
|
|
1
|
+
"""The fast profile: one capped profile point per side, at the slice crossing.
|
|
2
|
+
|
|
3
|
+
A full profile finds where dNLL crosses 1.9207 for every parameter, and on a
|
|
4
|
+
QSP model that is days. Most of that time is spent on the parameters that
|
|
5
|
+
turn out to be poorly determined, and for those the answer is already visible
|
|
6
|
+
much closer in. This pass asks a cheaper question at a single, well-chosen
|
|
7
|
+
point and reports what can be said from it.
|
|
8
|
+
|
|
9
|
+
The point
|
|
10
|
+
---------
|
|
11
|
+
The slice screen (:mod:`Engine.Identifiability`) has already walked each
|
|
12
|
+
side with the nuisance parameters held at their fitted values, so the ladder
|
|
13
|
+
point where the *slice* first sits above the threshold is known, together
|
|
14
|
+
with the slice's value there. That point, not an interpolated crossing, is
|
|
15
|
+
where the nuisance minimization is run: its slice value is a measured number,
|
|
16
|
+
so how much of it the other parameters absorb is measured too, and nothing
|
|
17
|
+
rests on an interpolation whose bias would otherwise leak into every verdict.
|
|
18
|
+
The profile crossing lies at or beyond the slice crossing, because the slice
|
|
19
|
+
is an upper bound on the profile. At that one point a nuisance minimization
|
|
20
|
+
is run, in a few short rounds, and the amount by which dNLL falls says how
|
|
21
|
+
much the other parameters can compensate:
|
|
22
|
+
|
|
23
|
+
* It stays at the threshold: nothing compensates. The profile and the slice
|
|
24
|
+
coincide here, the interval closes at the slice crossing, and the parameter
|
|
25
|
+
is as well determined as the slice suggests.
|
|
26
|
+
* It falls to near zero: the whole displacement is absorbed at no cost in
|
|
27
|
+
likelihood. There is a direction through parameter space along which the
|
|
28
|
+
data say nothing over at least this range, and the profile crossing, if it
|
|
29
|
+
exists at all, is far out.
|
|
30
|
+
* Somewhere in between: the interval is wider than the slice's by a factor
|
|
31
|
+
that can be estimated, and it is likely, not certainly, closed.
|
|
32
|
+
|
|
33
|
+
What is proof and what is not
|
|
34
|
+
-----------------------------
|
|
35
|
+
Every evaluation of the nuisance objective is an upper bound on the profile.
|
|
36
|
+
So a value *below* the threshold is a proof that the profile there is below
|
|
37
|
+
the threshold too, and a small value is a proof that the profile is at least
|
|
38
|
+
that small. Those are the one-directional statements this pass can make with
|
|
39
|
+
certainty, and they all widen intervals.
|
|
40
|
+
|
|
41
|
+
A value that *stays high* proves nothing on its own: Nelder-Mead in fourteen
|
|
42
|
+
nuisance dimensions can sit far above the minimum for a long time. "Proven
|
|
43
|
+
true" below therefore rests on the optimizer reporting convergence, which is
|
|
44
|
+
the same standard the full profile applies when it declares a crossing, and
|
|
45
|
+
no stronger. A parameter whose descent stalled without converging is reported
|
|
46
|
+
as likely, never proven.
|
|
47
|
+
|
|
48
|
+
Widths are extrapolated under a locally quadratic profile: if the profile
|
|
49
|
+
sits at dNLL d at the slice crossing, a distance s from the optimum, its own
|
|
50
|
+
crossing is at least s * sqrt(threshold / d) out. Because d is an upper bound
|
|
51
|
+
the estimate is a lower bound on the width, not a guess at it.
|
|
52
|
+
|
|
53
|
+
Verdicts
|
|
54
|
+
--------
|
|
55
|
+
Each side of each parameter is judged on the statement "the 95% interval is
|
|
56
|
+
closed on this side":
|
|
57
|
+
|
|
58
|
+
proven false the screen proved the slice, and so the profile, still
|
|
59
|
+
below the threshold decades out. The side is open.
|
|
60
|
+
unlikely true dNLL at the slice crossing fell to within a small fraction
|
|
61
|
+
of zero. The nuisance set compensates almost fully; the
|
|
62
|
+
extrapolated crossing is several times further out and
|
|
63
|
+
usually beyond where the screen walked.
|
|
64
|
+
likely true dNLL fell but stayed clear of zero, or stayed above the
|
|
65
|
+
threshold without the optimizer converging.
|
|
66
|
+
proven true the nuisance minimization converged above the threshold at
|
|
67
|
+
the slice crossing: the interval closes at or inside it.
|
|
68
|
+
no verdict the side could not be screened, the point could not be
|
|
69
|
+
evaluated, or a better optimum than the fit was found.
|
|
70
|
+
|
|
71
|
+
Every point computed here is a legitimate profile record and is written to
|
|
72
|
+
the same checkpoint the full profile reads, so a full profile run afterwards
|
|
73
|
+
starts with these points in hand and nothing is spent twice.
|
|
74
|
+
"""
|
|
75
|
+
|
|
76
|
+
import json
|
|
77
|
+
import os
|
|
78
|
+
import time
|
|
79
|
+
from datetime import datetime
|
|
80
|
+
|
|
81
|
+
import numpy as np
|
|
82
|
+
|
|
83
|
+
from Engine.Evaluator import FAILURE_VALUE
|
|
84
|
+
from Engine.Identifiability import (
|
|
85
|
+
MIN_REACH_DECADES,
|
|
86
|
+
SPAN_DECADES,
|
|
87
|
+
THRESHOLD,
|
|
88
|
+
_INCONCLUSIVE_STATES,
|
|
89
|
+
decades_from,
|
|
90
|
+
load_screen,
|
|
91
|
+
print_screen_report,
|
|
92
|
+
run_slice_screen,
|
|
93
|
+
save_screen,
|
|
94
|
+
)
|
|
95
|
+
|
|
96
|
+
REPORT_FILENAME = "fast_profile.json"
|
|
97
|
+
|
|
98
|
+
# Evaluations per round. Nelder-Mead needs n+1 to build its simplex, so a
|
|
99
|
+
# round has to be comfortably more than that to say anything; at ~21 s per
|
|
100
|
+
# evaluation on the SILK spec a round of 150 is about 50 minutes.
|
|
101
|
+
DEFAULT_ROUND_EVALS = 150
|
|
102
|
+
DEFAULT_ROUNDS = 3
|
|
103
|
+
|
|
104
|
+
# dNLL at or below this fraction of the threshold counts as "near zero".
|
|
105
|
+
# 0.1 puts the extrapolated crossing at least sqrt(10) ~ 3.2 times further out
|
|
106
|
+
# than the slice crossing.
|
|
107
|
+
DEFAULT_NEAR_ZERO_FRAC = 0.10
|
|
108
|
+
|
|
109
|
+
# A round that lowers dNLL by less than this fraction of the threshold has
|
|
110
|
+
# stalled. Used only to stop spending rounds on a point that is going nowhere;
|
|
111
|
+
# it never upgrades a verdict.
|
|
112
|
+
STALL_FRAC = 0.05
|
|
113
|
+
|
|
114
|
+
# dNLL below this is a better optimum than the fit, not a flat direction.
|
|
115
|
+
NEGATIVE_TOL = 1e-3
|
|
116
|
+
|
|
117
|
+
VERDICTS = ("proven true", "likely true", "unlikely true", "proven false",
|
|
118
|
+
"no verdict")
|
|
119
|
+
|
|
120
|
+
|
|
121
|
+
# ---------------------------------------------------------------------------
|
|
122
|
+
# Geometry
|
|
123
|
+
# ---------------------------------------------------------------------------
|
|
124
|
+
|
|
125
|
+
def slice_crossing(points, p_opt, threshold):
|
|
126
|
+
"""Where the slice crosses the threshold, from the screen's ladder.
|
|
127
|
+
|
|
128
|
+
Returns ``(x_cross, x_inner, x_outer)`` in optimizer space: the linear
|
|
129
|
+
interpolation of the crossing, and the two ladder points it lies between.
|
|
130
|
+
The optimum itself, at dNLL 0, is the inner point when the first ladder
|
|
131
|
+
point is already above the threshold. None if the slice never crossed.
|
|
132
|
+
|
|
133
|
+
Linear interpolation is chosen for its bias. A convex curve lies below its
|
|
134
|
+
chord, so the chord reaches the threshold first and the interpolated
|
|
135
|
+
crossing sits *inside* the true one. Every width this pass reports is a
|
|
136
|
+
lower bound, and placing the point inside keeps it one.
|
|
137
|
+
"""
|
|
138
|
+
prev_x, prev_d = float(p_opt), 0.0
|
|
139
|
+
for p in points:
|
|
140
|
+
d = p.get("dnll")
|
|
141
|
+
nll = p.get("nll")
|
|
142
|
+
if d is None or not np.isfinite(d) or nll is None or nll >= FAILURE_VALUE:
|
|
143
|
+
continue
|
|
144
|
+
x = float(p["x"])
|
|
145
|
+
if d > threshold:
|
|
146
|
+
if d == prev_d:
|
|
147
|
+
return x, prev_x, x
|
|
148
|
+
t = (threshold - prev_d) / (d - prev_d)
|
|
149
|
+
return prev_x + t * (x - prev_x), prev_x, x
|
|
150
|
+
prev_x, prev_d = x, float(d)
|
|
151
|
+
return None
|
|
152
|
+
|
|
153
|
+
|
|
154
|
+
def to_linear(x, is_log):
|
|
155
|
+
return float(10.0 ** x) if is_log else float(x)
|
|
156
|
+
|
|
157
|
+
|
|
158
|
+
def local_compensation(wald_cov, i):
|
|
159
|
+
"""Marginal over conditional standard error for parameter *i*, or None.
|
|
160
|
+
|
|
161
|
+
The conditional SE, 1/sqrt(H_ii), is the slice's own curvature; the
|
|
162
|
+
marginal SE, sqrt((H^-1)_ii), is the Wald interval's. Their ratio is the
|
|
163
|
+
square root of the variance inflation factor: how much wider the interval
|
|
164
|
+
is once the other parameters are free to move, in the quadratic
|
|
165
|
+
approximation. It is the local, linearised version of the question this
|
|
166
|
+
pass asks at the slice crossing, and it comes free with the Hessian.
|
|
167
|
+
"""
|
|
168
|
+
if wald_cov is None:
|
|
169
|
+
return None
|
|
170
|
+
try:
|
|
171
|
+
cov = np.asarray(wald_cov, dtype=float)
|
|
172
|
+
if cov.ndim != 2 or cov.shape[0] != cov.shape[1] or i >= cov.shape[0]:
|
|
173
|
+
return None
|
|
174
|
+
marginal = float(np.sqrt(cov[i, i]))
|
|
175
|
+
hess = np.linalg.inv(cov)
|
|
176
|
+
conditional = float(1.0 / np.sqrt(hess[i, i]))
|
|
177
|
+
except (np.linalg.LinAlgError, ValueError, FloatingPointError):
|
|
178
|
+
return None
|
|
179
|
+
if not (np.isfinite(marginal) and np.isfinite(conditional)) or conditional <= 0:
|
|
180
|
+
return None
|
|
181
|
+
return marginal / conditional
|
|
182
|
+
|
|
183
|
+
|
|
184
|
+
def width_factor(dnll, threshold):
|
|
185
|
+
"""Lower bound on (profile crossing distance) / (slice crossing distance).
|
|
186
|
+
|
|
187
|
+
Under a locally quadratic profile and with *dnll* an upper bound on the
|
|
188
|
+
profile at the slice crossing. Infinite when dnll is not positive: the
|
|
189
|
+
profile is flat to within resolution.
|
|
190
|
+
"""
|
|
191
|
+
if dnll is None or not np.isfinite(dnll):
|
|
192
|
+
return None
|
|
193
|
+
if dnll <= 0:
|
|
194
|
+
return float("inf")
|
|
195
|
+
return float(np.sqrt(threshold / dnll))
|
|
196
|
+
|
|
197
|
+
|
|
198
|
+
# ---------------------------------------------------------------------------
|
|
199
|
+
# Verdicts
|
|
200
|
+
# ---------------------------------------------------------------------------
|
|
201
|
+
|
|
202
|
+
def classify(side, threshold, near_zero_frac=DEFAULT_NEAR_ZERO_FRAC):
|
|
203
|
+
"""Verdict and one-line reason for one side's record. Pure."""
|
|
204
|
+
state = side.get("screen_state")
|
|
205
|
+
if state == "open":
|
|
206
|
+
reach = side.get("reach_decades")
|
|
207
|
+
how_far = (f"{reach:.2g} decade(s) out" if reach is not None
|
|
208
|
+
else f"out at {_g(side.get('reach_linear'))}")
|
|
209
|
+
return ("proven false",
|
|
210
|
+
f"the slice stays below {threshold:.4g} {how_far}, so the "
|
|
211
|
+
f"profile does too; the data do not exclude that value")
|
|
212
|
+
if state in _INCONCLUSIVE_STATES or state == "empty":
|
|
213
|
+
return ("no verdict", "the screen could not walk this side")
|
|
214
|
+
if side.get("x_point") is None:
|
|
215
|
+
return ("no verdict", "the slice never crossed within the walk")
|
|
216
|
+
if side.get("status") != "ok" or side.get("dnll") is None:
|
|
217
|
+
return ("no verdict",
|
|
218
|
+
f"the point just outside the slice crossing could not be "
|
|
219
|
+
f"evaluated ({side.get('status')})")
|
|
220
|
+
|
|
221
|
+
d = float(side["dnll"])
|
|
222
|
+
conv = bool(side.get("converged"))
|
|
223
|
+
nfev = side.get("nfev_total")
|
|
224
|
+
factor = width_factor(d, threshold)
|
|
225
|
+
absorbed = side.get("absorbed")
|
|
226
|
+
abs_txt = (f"{100 * absorbed:.0f}% of the slice's dNLL absorbed"
|
|
227
|
+
if absorbed is not None else "slice dNLL unknown")
|
|
228
|
+
at = f"{side.get('x_point_linear'):.4g}"
|
|
229
|
+
if d < -NEGATIVE_TOL:
|
|
230
|
+
return ("no verdict",
|
|
231
|
+
f"dNLL {d:.4g} at {at} is below zero: a better optimum than "
|
|
232
|
+
f"the fit exists there, so the anchor is wrong and nothing "
|
|
233
|
+
f"here can be read until the fit is redone from it")
|
|
234
|
+
if d <= near_zero_frac * threshold:
|
|
235
|
+
where = (f"; the crossing is at least {factor:.3g}x further out than "
|
|
236
|
+
f"{at}" if np.isfinite(factor) else
|
|
237
|
+
"; flat to within resolution")
|
|
238
|
+
beyond = (", beyond where the screen walked"
|
|
239
|
+
if side.get("est_beyond_reach") else "")
|
|
240
|
+
return ("unlikely true",
|
|
241
|
+
f"the other parameters absorb the displacement almost "
|
|
242
|
+
f"entirely at {at} ({abs_txt}, dNLL {d:.3g} of "
|
|
243
|
+
f"{threshold:.4g}){where}{beyond}")
|
|
244
|
+
if d < threshold:
|
|
245
|
+
tail = ("converged" if conv else
|
|
246
|
+
f"not converged after {nfev} evaluation(s), so it could fall "
|
|
247
|
+
f"further")
|
|
248
|
+
return ("likely true",
|
|
249
|
+
f"the profile is below the threshold at {at} ({abs_txt}, "
|
|
250
|
+
f"dNLL {d:.3g}), so the crossing is at least {factor:.3g}x "
|
|
251
|
+
f"further out; {tail}")
|
|
252
|
+
if conv:
|
|
253
|
+
return ("proven true",
|
|
254
|
+
f"the nuisance minimization converged above the threshold at "
|
|
255
|
+
f"{at} ({abs_txt}, dNLL {d:.3g}), so the interval closes at or "
|
|
256
|
+
f"inside it")
|
|
257
|
+
how = "stalled" if side.get("stalled") else "still descending"
|
|
258
|
+
return ("likely true",
|
|
259
|
+
f"dNLL {d:.3g} at {at} is still above the threshold after {nfev} "
|
|
260
|
+
f"evaluation(s) without converging ({how}, {abs_txt}); a flat "
|
|
261
|
+
f"direction has not been found, but the search is not finished")
|
|
262
|
+
|
|
263
|
+
|
|
264
|
+
# ---------------------------------------------------------------------------
|
|
265
|
+
# The pass
|
|
266
|
+
# ---------------------------------------------------------------------------
|
|
267
|
+
|
|
268
|
+
def _capped_kwargs(optimizer_kwargs, n_evals):
|
|
269
|
+
kw = dict(optimizer_kwargs or {})
|
|
270
|
+
opts = dict(kw.get("options") or {})
|
|
271
|
+
opts["maxfev"] = int(n_evals)
|
|
272
|
+
opts["maxiter"] = int(n_evals)
|
|
273
|
+
kw["options"] = opts
|
|
274
|
+
return kw
|
|
275
|
+
|
|
276
|
+
|
|
277
|
+
def run_fast_profile(batch, nll_batch, res_x, nll_at_optimum, param_names,
|
|
278
|
+
bounds, scales, method="Nelder-Mead", optimizer_kwargs=None,
|
|
279
|
+
wald_se=None, wald_cov=None, checkpoint=None, ckpt_dir=None,
|
|
280
|
+
threshold=THRESHOLD, round_evals=DEFAULT_ROUND_EVALS,
|
|
281
|
+
n_rounds=DEFAULT_ROUNDS,
|
|
282
|
+
near_zero_frac=DEFAULT_NEAR_ZERO_FRAC,
|
|
283
|
+
span_decades=SPAN_DECADES,
|
|
284
|
+
min_reach_decades=MIN_REACH_DECADES, verbose=True):
|
|
285
|
+
"""Screen, then one capped profile point per crossed side, in rounds.
|
|
286
|
+
|
|
287
|
+
*batch* has the profile_batch signature: ``batch(jobs, on_result, label)``.
|
|
288
|
+
*nll_batch* evaluates a list of full parameter vectors, for the screen.
|
|
289
|
+
Returns the report dict; :func:`print_fast_profile_report` renders it and
|
|
290
|
+
:func:`fast_profile_summary` shrinks it for the results snapshot.
|
|
291
|
+
"""
|
|
292
|
+
from Engine.Optimize import _cold_simplex, _param_bounds
|
|
293
|
+
|
|
294
|
+
res_x = np.asarray(res_x, dtype=float)
|
|
295
|
+
n = len(param_names)
|
|
296
|
+
scales = list(scales) if scales is not None else ["lin"] * n
|
|
297
|
+
t_start = time.time()
|
|
298
|
+
|
|
299
|
+
# ── The screen, reused when this fit already has one ──────────────────
|
|
300
|
+
screen = load_screen(ckpt_dir, param_names, res_x, threshold,
|
|
301
|
+
span_decades, min_reach_decades)
|
|
302
|
+
if screen is None:
|
|
303
|
+
screen = run_slice_screen(
|
|
304
|
+
nll_batch, res_x, nll_at_optimum, param_names, bounds,
|
|
305
|
+
scales=scales, wald_se=wald_se, threshold=threshold,
|
|
306
|
+
span_decades=span_decades, min_reach_decades=min_reach_decades,
|
|
307
|
+
verbose=verbose)
|
|
308
|
+
save_screen(screen, ckpt_dir)
|
|
309
|
+
elif verbose:
|
|
310
|
+
print(f"\n[fast profile] reusing the slice screen already run for this "
|
|
311
|
+
f"fit ({screen.get('n_evaluations', 0)} evaluation(s)).",
|
|
312
|
+
flush=True)
|
|
313
|
+
if verbose:
|
|
314
|
+
print_screen_report(screen)
|
|
315
|
+
|
|
316
|
+
# ── One side record per parameter per side ────────────────────────────
|
|
317
|
+
sides = {}
|
|
318
|
+
jobs = []
|
|
319
|
+
for i, name in enumerate(param_names):
|
|
320
|
+
is_log = scales[i] == "log10"
|
|
321
|
+
p_opt = float(res_x[i])
|
|
322
|
+
comp = local_compensation(wald_cov, i)
|
|
323
|
+
for sign, side_name in ((-1, "lower"), (1, "upper")):
|
|
324
|
+
rec = screen.get("parameters", {}).get(name, {}).get(side_name, {})
|
|
325
|
+
side = {
|
|
326
|
+
"param_idx": i, "name": name, "side": side_name, "sign": sign,
|
|
327
|
+
"is_log": is_log, "p_opt": p_opt, "p_opt_linear": to_linear(p_opt, is_log),
|
|
328
|
+
"screen_state": rec.get("state", "empty"),
|
|
329
|
+
"reach_decades": rec.get("reach_decades"),
|
|
330
|
+
"reach_linear": rec.get("reach"),
|
|
331
|
+
"certified_inner_linear": rec.get("inner_bracket"),
|
|
332
|
+
"local_compensation": comp,
|
|
333
|
+
# The interpolated slice crossing, for information; the point
|
|
334
|
+
# actually profiled is the ladder point just outside it.
|
|
335
|
+
"x_slice_linear": None,
|
|
336
|
+
"x_point": None, "x_point_linear": None, "point_decades": None,
|
|
337
|
+
"slice_dnll_at_point": None, "absorbed": None,
|
|
338
|
+
"dnll": None, "status": None, "converged": None,
|
|
339
|
+
"nfev_total": None, "rounds": [], "stalled": False,
|
|
340
|
+
"width_factor": None, "est_crossing_linear": None,
|
|
341
|
+
"est_decades": None, "est_beyond_reach": None,
|
|
342
|
+
}
|
|
343
|
+
sides[(i, side_name)] = side
|
|
344
|
+
if rec.get("state") != "crossed":
|
|
345
|
+
continue
|
|
346
|
+
cross = slice_crossing(rec.get("points", []), p_opt, threshold)
|
|
347
|
+
if cross is None:
|
|
348
|
+
continue
|
|
349
|
+
x_cross, _x_inner, x_point = (float(v) for v in cross)
|
|
350
|
+
side["x_slice_linear"] = to_linear(x_cross, is_log)
|
|
351
|
+
side["x_point"] = x_point
|
|
352
|
+
side["x_point_linear"] = to_linear(x_point, is_log)
|
|
353
|
+
d = decades_from(p_opt, x_point, is_log)
|
|
354
|
+
side["point_decades"] = float(d) if np.isfinite(d) else None
|
|
355
|
+
for p in rec.get("points", []):
|
|
356
|
+
if abs(float(p["x"]) - x_point) <= 1e-12 * max(1.0, abs(x_point)):
|
|
357
|
+
side["slice_dnll_at_point"] = float(p["dnll"])
|
|
358
|
+
break
|
|
359
|
+
x_slice = x_point
|
|
360
|
+
|
|
361
|
+
nb = None
|
|
362
|
+
if bounds is not None:
|
|
363
|
+
nb = [list(b) if b is not None else None
|
|
364
|
+
for b in (list(bounds[:i]) + list(bounds[i + 1:]))]
|
|
365
|
+
start = np.delete(res_x, i)
|
|
366
|
+
se_nuis = None
|
|
367
|
+
if wald_se is not None:
|
|
368
|
+
try:
|
|
369
|
+
se_arr = np.atleast_1d(np.asarray(wald_se, dtype=float))
|
|
370
|
+
if se_arr.size == n:
|
|
371
|
+
se_nuis = np.delete(se_arr, i)
|
|
372
|
+
except (TypeError, ValueError):
|
|
373
|
+
se_nuis = None
|
|
374
|
+
job = {
|
|
375
|
+
"param_idx": i,
|
|
376
|
+
"param_name": name,
|
|
377
|
+
"x_fixed": x_slice,
|
|
378
|
+
"x_fixed_linear": side["x_point_linear"],
|
|
379
|
+
"x_start": start.tolist(),
|
|
380
|
+
"warm_seeded": False,
|
|
381
|
+
"nuisance_bounds": nb,
|
|
382
|
+
"method": method,
|
|
383
|
+
"optimizer_kwargs": _capped_kwargs(optimizer_kwargs, round_evals),
|
|
384
|
+
# A cold grid point as far as the full profile is concerned:
|
|
385
|
+
# its warm pass owes this point a sweep like any other.
|
|
386
|
+
"phase": 1,
|
|
387
|
+
"direction": 0,
|
|
388
|
+
"x_step": abs(x_slice - p_opt),
|
|
389
|
+
"fast_profile": True,
|
|
390
|
+
"_side": side_name,
|
|
391
|
+
}
|
|
392
|
+
sim = _cold_simplex(start, se_nuis, bounds=nb)
|
|
393
|
+
if sim is not None:
|
|
394
|
+
job["initial_simplex"] = sim
|
|
395
|
+
jobs.append(job)
|
|
396
|
+
|
|
397
|
+
if verbose:
|
|
398
|
+
n_crossed = sum(1 for s in sides.values() if s["screen_state"] == "crossed")
|
|
399
|
+
print(f"\n[fast profile] {len(jobs)} side(s) crossed the slice threshold "
|
|
400
|
+
f"and get one capped profile point each at the slice crossing, "
|
|
401
|
+
f"up to {n_rounds} round(s) of {round_evals} evaluation(s); "
|
|
402
|
+
f"{2 * n - n_crossed} side(s) were settled or blocked by the "
|
|
403
|
+
f"screen.", flush=True)
|
|
404
|
+
|
|
405
|
+
# ── Rounds ────────────────────────────────────────────────────────────
|
|
406
|
+
def record(res):
|
|
407
|
+
i, side_name = int(res["param_idx"]), res.get("_side")
|
|
408
|
+
side = sides.get((i, side_name))
|
|
409
|
+
if side is None:
|
|
410
|
+
return
|
|
411
|
+
nll = res.get("nll")
|
|
412
|
+
ok = (res.get("status") == "ok" and nll is not None
|
|
413
|
+
and np.isfinite(nll) and nll < FAILURE_VALUE)
|
|
414
|
+
d = float(nll) - float(nll_at_optimum) if ok else None
|
|
415
|
+
prev = side["dnll"]
|
|
416
|
+
side["status"] = res.get("status")
|
|
417
|
+
side["rounds"].append({
|
|
418
|
+
"dnll": d, "nfev_total": res.get("nfev_total"),
|
|
419
|
+
"converged": bool(res.get("converged")),
|
|
420
|
+
"status": res.get("status"),
|
|
421
|
+
})
|
|
422
|
+
if ok and (prev is None or d <= prev):
|
|
423
|
+
side["dnll"] = d
|
|
424
|
+
side["converged"] = bool(res.get("converged"))
|
|
425
|
+
side["nfev_total"] = res.get("nfev_total")
|
|
426
|
+
side["_last"] = res
|
|
427
|
+
if checkpoint is not None:
|
|
428
|
+
keep = dict(res)
|
|
429
|
+
keep.pop("_side", None)
|
|
430
|
+
keep["dnll"] = d
|
|
431
|
+
keep.setdefault("warm_refined", 0)
|
|
432
|
+
checkpoint.append(keep)
|
|
433
|
+
elif ok:
|
|
434
|
+
# Not an improvement, but the spend and the simplex still move.
|
|
435
|
+
side["_last"] = dict(side.get("_last") or res, **{
|
|
436
|
+
k: res[k] for k in ("nfev_total", "nit_total", "nm_simplex",
|
|
437
|
+
"converged", "nuisance_x")
|
|
438
|
+
if k in res})
|
|
439
|
+
side["nfev_total"] = res.get("nfev_total")
|
|
440
|
+
side["converged"] = bool(res.get("converged"))
|
|
441
|
+
if prev is not None and d is not None:
|
|
442
|
+
side["stalled"] = (prev - d) < STALL_FRAC * threshold
|
|
443
|
+
|
|
444
|
+
deadline_hit = None
|
|
445
|
+
active = list(jobs)
|
|
446
|
+
for r in range(int(n_rounds)):
|
|
447
|
+
if not active:
|
|
448
|
+
break
|
|
449
|
+
label = f"fast-profile r{r + 1}"
|
|
450
|
+
if verbose and r:
|
|
451
|
+
print(f"\n[fast profile] round {r + 1}: {len(active)} point(s) "
|
|
452
|
+
f"continue.", flush=True)
|
|
453
|
+
try:
|
|
454
|
+
batch(active, on_result=record, label=label)
|
|
455
|
+
except Exception as exc: # DeadlineReached, or a pool failure
|
|
456
|
+
if type(exc).__name__ != "DeadlineReached":
|
|
457
|
+
raise
|
|
458
|
+
deadline_hit = str(exc)
|
|
459
|
+
if verbose:
|
|
460
|
+
print(f"[fast profile] {exc}; classifying what landed.",
|
|
461
|
+
flush=True)
|
|
462
|
+
break
|
|
463
|
+
|
|
464
|
+
nxt = []
|
|
465
|
+
for job in active:
|
|
466
|
+
side = sides[(job["param_idx"], job["_side"])]
|
|
467
|
+
d = side["dnll"]
|
|
468
|
+
last = side.get("_last")
|
|
469
|
+
if last is None or side["status"] != "ok" or d is None:
|
|
470
|
+
continue # failed: no more rounds
|
|
471
|
+
if d < -NEGATIVE_TOL or d <= near_zero_frac * threshold:
|
|
472
|
+
continue # verdict reached
|
|
473
|
+
if side["converged"]:
|
|
474
|
+
continue
|
|
475
|
+
if side["stalled"] and d >= threshold:
|
|
476
|
+
continue # going nowhere above the line
|
|
477
|
+
cont = dict(job)
|
|
478
|
+
cont["x_start"] = list(last.get("nuisance_x") or job["x_start"])
|
|
479
|
+
cont["optimizer_kwargs"] = _capped_kwargs(optimizer_kwargs,
|
|
480
|
+
round_evals * (r + 2))
|
|
481
|
+
cont["nfev_used"] = int(last.get("nfev_total") or 0)
|
|
482
|
+
cont["nit_used"] = int(last.get("nit_total") or 0)
|
|
483
|
+
cont["nll_so_far"] = last.get("nll")
|
|
484
|
+
cont["resumed"] = True
|
|
485
|
+
sim = last.get("nm_simplex")
|
|
486
|
+
if sim is not None:
|
|
487
|
+
cont["initial_simplex"] = sim
|
|
488
|
+
nxt.append(cont)
|
|
489
|
+
active = nxt
|
|
490
|
+
|
|
491
|
+
# ── Extrapolation and verdicts ────────────────────────────────────────
|
|
492
|
+
counts = {v: 0 for v in VERDICTS}
|
|
493
|
+
for side in sides.values():
|
|
494
|
+
side.pop("_last", None)
|
|
495
|
+
d = side["dnll"]
|
|
496
|
+
if d is not None and side["x_point"] is not None:
|
|
497
|
+
s_here = side.get("slice_dnll_at_point")
|
|
498
|
+
if s_here is not None and s_here > 0:
|
|
499
|
+
side["absorbed"] = float(np.clip(1.0 - d / s_here, 0.0, 1.0))
|
|
500
|
+
f = width_factor(d, threshold)
|
|
501
|
+
side["width_factor"] = f
|
|
502
|
+
if f is not None and np.isfinite(f):
|
|
503
|
+
x_est = side["p_opt"] + side["sign"] * f * abs(side["x_point"] - side["p_opt"])
|
|
504
|
+
side["est_crossing_linear"] = to_linear(x_est, side["is_log"])
|
|
505
|
+
dec = decades_from(side["p_opt"], x_est, side["is_log"])
|
|
506
|
+
side["est_decades"] = float(dec) if np.isfinite(dec) else None
|
|
507
|
+
reach = side.get("reach_decades")
|
|
508
|
+
if reach is not None and np.isfinite(dec):
|
|
509
|
+
side["est_beyond_reach"] = bool(dec > reach)
|
|
510
|
+
elif side.get("reach_linear") is not None:
|
|
511
|
+
# No decades for a linear parameter through zero; compare
|
|
512
|
+
# distances in linear units instead.
|
|
513
|
+
p_lin = side["p_opt_linear"]
|
|
514
|
+
side["est_beyond_reach"] = bool(
|
|
515
|
+
abs(side["est_crossing_linear"] - p_lin)
|
|
516
|
+
> abs(float(side["reach_linear"]) - p_lin))
|
|
517
|
+
else:
|
|
518
|
+
side["est_beyond_reach"] = None
|
|
519
|
+
elif f is not None:
|
|
520
|
+
side["est_beyond_reach"] = True
|
|
521
|
+
verdict, reason = classify(side, threshold, near_zero_frac)
|
|
522
|
+
side["verdict"] = verdict
|
|
523
|
+
side["reason"] = reason
|
|
524
|
+
counts[verdict] += 1
|
|
525
|
+
|
|
526
|
+
report = {
|
|
527
|
+
"threshold": float(threshold),
|
|
528
|
+
"anchor": float(nll_at_optimum),
|
|
529
|
+
"res_x": [float(v) for v in res_x],
|
|
530
|
+
"param_names": list(param_names),
|
|
531
|
+
"round_evals": int(round_evals),
|
|
532
|
+
"n_rounds": int(n_rounds),
|
|
533
|
+
"near_zero_frac": float(near_zero_frac),
|
|
534
|
+
"n_points": len(jobs),
|
|
535
|
+
"n_evaluations": int(sum(
|
|
536
|
+
(s.get("nfev_total") or 0) for s in sides.values())),
|
|
537
|
+
"wall_s": time.time() - t_start,
|
|
538
|
+
"deadline": deadline_hit,
|
|
539
|
+
"counts": counts,
|
|
540
|
+
"timestamp": datetime.now().isoformat(timespec="seconds"),
|
|
541
|
+
"parameters": {},
|
|
542
|
+
}
|
|
543
|
+
for (i, side_name), side in sides.items():
|
|
544
|
+
report["parameters"].setdefault(side["name"], {})[side_name] = {
|
|
545
|
+
k: v for k, v in side.items() if k not in ("param_idx", "name",
|
|
546
|
+
"side", "sign")}
|
|
547
|
+
return report
|
|
548
|
+
|
|
549
|
+
|
|
550
|
+
# ---------------------------------------------------------------------------
|
|
551
|
+
# Reporting
|
|
552
|
+
# ---------------------------------------------------------------------------
|
|
553
|
+
|
|
554
|
+
def save_report(report, ckpt_dir):
|
|
555
|
+
if not ckpt_dir:
|
|
556
|
+
return None
|
|
557
|
+
path = os.path.join(ckpt_dir, REPORT_FILENAME)
|
|
558
|
+
tmp = f"{path}.{os.getpid()}.tmp"
|
|
559
|
+
try:
|
|
560
|
+
os.makedirs(ckpt_dir, exist_ok=True)
|
|
561
|
+
with open(tmp, "w", encoding="utf-8") as fh:
|
|
562
|
+
json.dump(_jsonable(report), fh, indent=1)
|
|
563
|
+
os.replace(tmp, path)
|
|
564
|
+
except OSError:
|
|
565
|
+
try:
|
|
566
|
+
os.unlink(tmp)
|
|
567
|
+
except OSError:
|
|
568
|
+
pass
|
|
569
|
+
return None
|
|
570
|
+
return path
|
|
571
|
+
|
|
572
|
+
|
|
573
|
+
def _jsonable(obj):
|
|
574
|
+
if isinstance(obj, dict):
|
|
575
|
+
return {str(k): _jsonable(v) for k, v in obj.items()}
|
|
576
|
+
if isinstance(obj, (list, tuple)):
|
|
577
|
+
return [_jsonable(v) for v in obj]
|
|
578
|
+
if isinstance(obj, np.ndarray):
|
|
579
|
+
return _jsonable(obj.tolist())
|
|
580
|
+
if isinstance(obj, (np.integer,)):
|
|
581
|
+
return int(obj)
|
|
582
|
+
if isinstance(obj, (np.floating, float)):
|
|
583
|
+
v = float(obj)
|
|
584
|
+
return v if np.isfinite(v) else None
|
|
585
|
+
return obj
|
|
586
|
+
|
|
587
|
+
|
|
588
|
+
def _g(v, prec=4):
|
|
589
|
+
if v is None:
|
|
590
|
+
return "-"
|
|
591
|
+
try:
|
|
592
|
+
v = float(v)
|
|
593
|
+
except (TypeError, ValueError):
|
|
594
|
+
return str(v)
|
|
595
|
+
if not np.isfinite(v):
|
|
596
|
+
return "inf" if v > 0 else "-inf"
|
|
597
|
+
return f"{v:.{prec}g}"
|
|
598
|
+
|
|
599
|
+
|
|
600
|
+
def format_summary_table(report):
|
|
601
|
+
"""The end-of-run table, as text."""
|
|
602
|
+
thr = report["threshold"]
|
|
603
|
+
rows = []
|
|
604
|
+
for name, sides in report["parameters"].items():
|
|
605
|
+
for side_name in ("lower", "upper"):
|
|
606
|
+
s = sides.get(side_name)
|
|
607
|
+
if s is None:
|
|
608
|
+
continue
|
|
609
|
+
factor = s.get("width_factor")
|
|
610
|
+
factor_txt = ("flat" if factor is not None and not np.isfinite(factor)
|
|
611
|
+
else (f"x{factor:.2g}" if factor is not None else "-"))
|
|
612
|
+
absorbed = s.get("absorbed")
|
|
613
|
+
rows.append((
|
|
614
|
+
name, side_name,
|
|
615
|
+
_g(s.get("x_point_linear")),
|
|
616
|
+
_g(s.get("slice_dnll_at_point"), 3),
|
|
617
|
+
_g(s.get("dnll"), 3),
|
|
618
|
+
(f"{100 * absorbed:.0f}%" if absorbed is not None else "-"),
|
|
619
|
+
str(s.get("nfev_total") if s.get("nfev_total") is not None else "-"),
|
|
620
|
+
("yes" if s.get("converged") else
|
|
621
|
+
("no" if s.get("converged") is not None else "-")),
|
|
622
|
+
factor_txt,
|
|
623
|
+
_g(s.get("est_crossing_linear")),
|
|
624
|
+
_g(s.get("local_compensation"), 2),
|
|
625
|
+
s.get("verdict", "-"),
|
|
626
|
+
))
|
|
627
|
+
heads = ("parameter", "side", "point at", "slice dNLL", "profile dNLL",
|
|
628
|
+
"absorbed", "evals", "conv", "width >=", "crossing >=", "VIF^0.5",
|
|
629
|
+
"interval closed?")
|
|
630
|
+
widths = [max(len(h), *(len(r[c]) for r in rows)) if rows else len(h)
|
|
631
|
+
for c, h in enumerate(heads)]
|
|
632
|
+
line = " ".join(h.ljust(w) for h, w in zip(heads, widths))
|
|
633
|
+
out = ["", f"[fast profile] summary (threshold dNLL = {thr:.4g}):", "",
|
|
634
|
+
" " + line, " " + "-" * len(line)]
|
|
635
|
+
for r in rows:
|
|
636
|
+
out.append(" " + " ".join(v.ljust(w) for v, w in zip(r, widths)))
|
|
637
|
+
c = report["counts"]
|
|
638
|
+
out.append("")
|
|
639
|
+
out.append(" " + " ".join(f"{k}: {c.get(k, 0)}" for k in VERDICTS))
|
|
640
|
+
out.append("")
|
|
641
|
+
out.append(" Statement judged: 'the 95% interval is closed on this side'.")
|
|
642
|
+
out.append(" proven false the slice, and so the profile, stays below the "
|
|
643
|
+
"threshold decades out (a proof; the side is open).")
|
|
644
|
+
out.append(" unlikely true dNLL at the slice crossing fell to within "
|
|
645
|
+
f"{100 * report['near_zero_frac']:.0f}% of zero: the other "
|
|
646
|
+
"parameters absorb the displacement; 'crossing >=' is a lower "
|
|
647
|
+
"bound on where the interval could close.")
|
|
648
|
+
out.append(" likely true the profile fell below the threshold there "
|
|
649
|
+
"but stayed clear of zero, or stayed above it without converging.")
|
|
650
|
+
out.append(" proven true the nuisance minimization converged above the "
|
|
651
|
+
"threshold at the slice crossing (proof to optimizer tolerance, "
|
|
652
|
+
"the same standard the full profile uses).")
|
|
653
|
+
out.append(" no verdict unscreened, unevaluable, or a better optimum "
|
|
654
|
+
"than the fit was found.")
|
|
655
|
+
out.append(" 'width >=' is (profile crossing distance)/(slice crossing "
|
|
656
|
+
"distance) under a quadratic profile; since the point is an upper "
|
|
657
|
+
"bound it is a floor, not an estimate.")
|
|
658
|
+
out.append(" 'VIF^0.5' is Wald marginal SE over conditional SE from the "
|
|
659
|
+
"Hessian: the same compensation question, linearised at the "
|
|
660
|
+
"optimum.")
|
|
661
|
+
if report.get("deadline"):
|
|
662
|
+
out.append(f" Stopped on the wall clock: {report['deadline']}")
|
|
663
|
+
out.append("")
|
|
664
|
+
return "\n".join(out)
|
|
665
|
+
|
|
666
|
+
|
|
667
|
+
def print_fast_profile_report(report):
|
|
668
|
+
print(format_summary_table(report), flush=True)
|
|
669
|
+
for name, sides in report["parameters"].items():
|
|
670
|
+
for side_name in ("lower", "upper"):
|
|
671
|
+
s = sides.get(side_name)
|
|
672
|
+
if s is None:
|
|
673
|
+
continue
|
|
674
|
+
print(f" {name} ({side_name}): {s.get('verdict')} -- {s.get('reason')}")
|
|
675
|
+
print(flush=True)
|
|
676
|
+
|
|
677
|
+
|
|
678
|
+
def fast_profile_summary(report):
|
|
679
|
+
"""The report without its per-round detail, for the results snapshot."""
|
|
680
|
+
if not report:
|
|
681
|
+
return None
|
|
682
|
+
keep = ("verdict", "reason", "screen_state", "x_slice_linear",
|
|
683
|
+
"x_point_linear", "slice_dnll_at_point", "dnll", "absorbed",
|
|
684
|
+
"converged", "nfev_total", "width_factor", "est_crossing_linear",
|
|
685
|
+
"est_decades", "est_beyond_reach", "local_compensation",
|
|
686
|
+
"certified_inner_linear", "reach_decades")
|
|
687
|
+
return {
|
|
688
|
+
"threshold": report.get("threshold"),
|
|
689
|
+
"counts": report.get("counts"),
|
|
690
|
+
"n_points": report.get("n_points"),
|
|
691
|
+
"n_evaluations": report.get("n_evaluations"),
|
|
692
|
+
"round_evals": report.get("round_evals"),
|
|
693
|
+
"n_rounds": report.get("n_rounds"),
|
|
694
|
+
"near_zero_frac": report.get("near_zero_frac"),
|
|
695
|
+
"deadline": report.get("deadline"),
|
|
696
|
+
"parameters": {
|
|
697
|
+
name: {side: {k: _jsonable(rec.get(k)) for k in keep}
|
|
698
|
+
for side, rec in sides.items()}
|
|
699
|
+
for name, sides in report.get("parameters", {}).items()
|
|
700
|
+
},
|
|
701
|
+
}
|