Skip to content

Commit 489fab9

Browse files
ci: add robust trace-log patch script
1 parent d359fe7 commit 489fab9

1 file changed

Lines changed: 307 additions & 0 deletions

File tree

Lines changed: 307 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,307 @@
1+
from pathlib import Path
2+
3+
diagnostics_path = Path("baseline/rg_baselines/diagnostics.py")
4+
diagnostics = diagnostics_path.read_text(encoding="utf-8")
5+
6+
new_clean = """\
7+
def clean_positive_eigenvalues(
8+
values: Any,
9+
*,
10+
expected_dimension: Optional[int] = None,
11+
) -> np.ndarray:
12+
\"\"\"Return positive eigenvalues in ascending order.
13+
14+
When ``expected_dimension`` is supplied, fail closed if the ESD is
15+
incomplete, non-finite, or rank deficient. This preserves
16+
WeightWatcher's full-M normalization instead of silently renormalizing a
17+
filtered positive-rank spectrum.
18+
\"\"\"
19+
20+
evals = np.asarray(values, dtype=float).reshape(-1)
21+
if expected_dimension is not None:
22+
expected = int(expected_dimension)
23+
if expected < 2:
24+
raise ValueError("expected spectral dimension must be at least two")
25+
if evals.size != expected:
26+
raise ValueError(
27+
"ESD dimension mismatch: "
28+
f"expected {expected} eigenvalues, received {evals.size}"
29+
)
30+
if not np.all(np.isfinite(evals)):
31+
raise ValueError("full ESD contains non-finite eigenvalues")
32+
if np.any(evals <= 0.0):
33+
positive = int(np.count_nonzero(evals > 0.0))
34+
raise ValueError(
35+
"rank-deficient ESD: "
36+
f"expected {expected} positive eigenvalues, found {positive}"
37+
)
38+
else:
39+
evals = evals[np.isfinite(evals) & (evals > 0.0)]
40+
41+
evals = np.sort(evals)
42+
if evals.size < 2:
43+
raise ValueError("fewer than two finite positive eigenvalues")
44+
return evals
45+
"""
46+
start = diagnostics.index("def clean_positive_eigenvalues(")
47+
end = diagnostics.index("\ndef _entropy_effective_rank", start)
48+
diagnostics = diagnostics[:start] + new_clean + diagnostics[end + 1 :]
49+
50+
new_metrics_prefix = """\
51+
def spectral_metrics_from_esd(
52+
raw_evals_ascending: Any,
53+
normalized_evals_ascending: Any,
54+
*,
55+
detx_num: int,
56+
num_pl_spikes: int,
57+
erg_gap: int,
58+
expected_dimension: Optional[int] = None,
59+
) -> dict[str, float | int]:
60+
\"\"\"Compute transparent metrics from one WeightWatcher ESD.
61+
62+
``normalized_evals_ascending`` must be produced by WeightWatcher's own
63+
``RMT_Util.rescale_eigenvalues``. The trace-log boundary and gap are not
64+
recomputed here: the supplied ``detx_num``, ``num_pl_spikes``, and
65+
``erg_gap`` must come from ``watcher.analyze(ERG=True)``.
66+
67+
``expected_dimension`` is the full spectral dimension
68+
``min(weight.shape)``. Strict baseline measurements require all of those
69+
eigenvalues to be finite and positive so WeightWatcher's normalization is
70+
not silently changed by positive-eigenvalue filtering.
71+
\"\"\"
72+
73+
raw = clean_positive_eigenvalues(
74+
raw_evals_ascending,
75+
expected_dimension=expected_dimension,
76+
)
77+
normalized = clean_positive_eigenvalues(
78+
normalized_evals_ascending,
79+
expected_dimension=expected_dimension,
80+
)
81+
if raw.size != normalized.size:
82+
raise ValueError("raw and normalized ESDs have different sizes")
83+
84+
count = int(raw.size)
85+
normalized_sum = float(np.sum(normalized))
86+
if not np.isclose(
87+
normalized_sum,
88+
float(count),
89+
rtol=1e-10,
90+
atol=1e-10 * max(count, 1),
91+
):
92+
raise ValueError(
93+
"WeightWatcher normalization audit failed: "
94+
f"sum={normalized_sum:.17g}, expected={count}"
95+
)
96+
97+
m_detx = int(detx_num)
98+
m_pl = int(num_pl_spikes)
99+
if not 1 <= m_detx <= count:
100+
raise ValueError(
101+
f"detX_num must lie in [1, {count}], received {m_detx}"
102+
)
103+
if not 1 <= m_pl <= count:
104+
raise ValueError(
105+
f"num_pl_spikes must lie in [1, {count}], received {m_pl}"
106+
)
107+
108+
expected_gap = m_detx - m_pl
109+
if int(erg_gap) != expected_gap:
110+
raise ValueError(
111+
f"WeightWatcher ERG_gap audit failed: {erg_gap} != {m_detx} - {m_pl}"
112+
)
113+
m_midpoint = int(math.floor((m_detx + m_pl) / 2.0))
114+
"""
115+
start = diagnostics.index("def spectral_metrics_from_esd(")
116+
body = diagnostics.index(" raw_desc = raw[::-1]", start)
117+
diagnostics = diagnostics[:start] + new_metrics_prefix + "\n" + diagnostics[body:]
118+
119+
old_sum = (
120+
' "rescaled_eigenvalue_sum": float(np.sum(normalized)),\n'
121+
' "rescale_sum_minus_num_eigenvalues": '
122+
'float(np.sum(normalized) - count),\n'
123+
)
124+
new_sum = (
125+
' "rescaled_eigenvalue_sum": normalized_sum,\n'
126+
' "rescale_sum_minus_num_eigenvalues": '
127+
'float(normalized_sum - count),\n'
128+
)
129+
if diagnostics.count(old_sum) != 1:
130+
raise RuntimeError("unexpected normalized-sum output source")
131+
diagnostics = diagnostics.replace(old_sum, new_sum, 1)
132+
133+
new_measure = """\
134+
parameter = parameter_map.get(parameter_name) if parameter_name else None
135+
if parameter is None:
136+
raise ValueError(
137+
"WeightWatcher layer could not be matched to a model matrix"
138+
)
139+
expected_dimension = int(min(parameter.shape))
140+
raw_esd = clean_positive_eigenvalues(
141+
_get_esd_compat(
142+
watcher,
143+
model=model_cpu,
144+
layer_id=int(layer_id),
145+
params=get_esd_params,
146+
),
147+
expected_dimension=expected_dimension,
148+
)
149+
normalized_esd, weight_scale = _rescale_with_weightwatcher(raw_esd)
150+
computed = spectral_metrics_from_esd(
151+
raw_esd,
152+
normalized_esd,
153+
detx_num=int(detx_num),
154+
num_pl_spikes=int(num_pl_spikes),
155+
erg_gap=erg_gap,
156+
expected_dimension=expected_dimension,
157+
)
158+
159+
"""
160+
start = diagnostics.index(" raw_esd = clean_positive_eigenvalues(")
161+
end = diagnostics.index(" record = {", start)
162+
diagnostics = diagnostics[:start] + new_measure + diagnostics[end:]
163+
164+
diagnostics = diagnostics.replace(
165+
' "layer_rows": int(parameter.shape[0]) '
166+
'if parameter is not None else np.nan,\n'
167+
' "layer_cols": int(parameter.shape[1]) '
168+
'if parameter is not None else np.nan,\n'
169+
' "layer_parameter_count": int(parameter.numel()) '
170+
'if parameter is not None else np.nan,\n',
171+
' "layer_rows": int(parameter.shape[0]),\n'
172+
' "layer_cols": int(parameter.shape[1]),\n'
173+
' "layer_parameter_count": int(parameter.numel()),\n',
174+
1,
175+
)
176+
diagnostics_path.write_text(diagnostics, encoding="utf-8")
177+
178+
tests_path = Path("baseline/tests/test_diagnostics.py")
179+
tests_path.write_text(
180+
"""\
181+
import unittest
182+
183+
import numpy as np
184+
185+
from rg_baselines.diagnostics import (
186+
clean_positive_eigenvalues,
187+
spectral_metrics_from_esd,
188+
)
189+
190+
191+
class SpectralMetricsTests(unittest.TestCase):
192+
def test_original_boundaries_and_midpoint(self) -> None:
193+
raw = np.arange(1.0, 11.0)
194+
normalized = raw * (len(raw) / raw.sum())
195+
metrics = spectral_metrics_from_esd(
196+
raw,
197+
normalized,
198+
detx_num=8,
199+
num_pl_spikes=4,
200+
erg_gap=4,
201+
expected_dimension=len(raw),
202+
)
203+
self.assertEqual(metrics["m_midpoint"], 6)
204+
self.assertEqual(metrics["ERG_gap"], 4)
205+
self.assertAlmostEqual(metrics["rescaled_eigenvalue_sum"], 10.0)
206+
self.assertAlmostEqual(
207+
metrics["rescale_sum_minus_num_eigenvalues"],
208+
0.0,
209+
)
210+
self.assertGreater(metrics["midpoint_energy_fraction"], 0.5)
211+
212+
def test_trace_log_matches_analytic_top_spectrum_value(self) -> None:
213+
raw = np.asarray([1.0, 2.0, 4.0, 8.0])
214+
normalized = raw * (len(raw) / raw.sum())
215+
metrics = spectral_metrics_from_esd(
216+
raw,
217+
normalized,
218+
detx_num=4,
219+
num_pl_spikes=2,
220+
erg_gap=2,
221+
expected_dimension=4,
222+
)
223+
224+
retained = normalized[::-1][:3]
225+
expected_total = float(np.sum(np.log(retained)))
226+
expected_per_eval = float(np.mean(np.log(retained)))
227+
self.assertEqual(metrics["m_midpoint"], 3)
228+
self.assertAlmostEqual(
229+
metrics["trace_log_midpoint_total"],
230+
expected_total,
231+
)
232+
self.assertAlmostEqual(
233+
metrics["trace_log_midpoint_per_eval"],
234+
expected_per_eval,
235+
)
236+
self.assertAlmostEqual(
237+
metrics["geometric_mean_midpoint"],
238+
float(np.exp(expected_per_eval)),
239+
)
240+
self.assertAlmostEqual(
241+
metrics["trace_log_midpoint_total"],
242+
3.0 * metrics["trace_log_midpoint_per_eval"],
243+
)
244+
245+
def test_gap_mismatch_is_rejected(self) -> None:
246+
raw = np.arange(1.0, 11.0)
247+
normalized = raw * (len(raw) / raw.sum())
248+
with self.assertRaisesRegex(ValueError, "ERG_gap audit failed"):
249+
spectral_metrics_from_esd(
250+
raw,
251+
normalized,
252+
detx_num=8,
253+
num_pl_spikes=4,
254+
erg_gap=3,
255+
expected_dimension=len(raw),
256+
)
257+
258+
def test_out_of_range_boundaries_are_rejected(self) -> None:
259+
raw = np.arange(1.0, 6.0)
260+
normalized = raw * (len(raw) / raw.sum())
261+
for field, detx_num, num_pl_spikes in (
262+
("detX_num", 6, 2),
263+
("num_pl_spikes", 4, 0),
264+
):
265+
with self.subTest(field=field):
266+
with self.assertRaisesRegex(ValueError, field):
267+
spectral_metrics_from_esd(
268+
raw,
269+
normalized,
270+
detx_num=detx_num,
271+
num_pl_spikes=num_pl_spikes,
272+
erg_gap=detx_num - num_pl_spikes,
273+
expected_dimension=len(raw),
274+
)
275+
276+
def test_rank_deficient_full_esd_is_rejected(self) -> None:
277+
with self.assertRaisesRegex(ValueError, "rank-deficient ESD"):
278+
clean_positive_eigenvalues(
279+
[0.0, 1.0, 2.0],
280+
expected_dimension=3,
281+
)
282+
283+
def test_incomplete_full_esd_is_rejected(self) -> None:
284+
with self.assertRaisesRegex(ValueError, "ESD dimension mismatch"):
285+
clean_positive_eigenvalues(
286+
[1.0, 2.0],
287+
expected_dimension=3,
288+
)
289+
290+
def test_incorrect_weightwatcher_normalization_is_rejected(self) -> None:
291+
raw = np.asarray([1.0, 2.0, 3.0, 4.0])
292+
with self.assertRaisesRegex(ValueError, "normalization audit failed"):
293+
spectral_metrics_from_esd(
294+
raw,
295+
raw,
296+
detx_num=4,
297+
num_pl_spikes=2,
298+
erg_gap=2,
299+
expected_dimension=4,
300+
)
301+
302+
303+
if __name__ == "__main__":
304+
unittest.main()
305+
""",
306+
encoding="utf-8",
307+
)

0 commit comments

Comments
 (0)