Skip to content

ci: stage MNIST trace-log robustness patch #1

ci: stage MNIST trace-log robustness patch

ci: stage MNIST trace-log robustness patch #1

name: Harden MNIST trace-log diagnostics
on:
push:
branches:
- "agent/harden-mnist-trace-log"
paths:
- ".github/workflows/harden-mnist-trace-log.yml"
permissions:
contents: write
jobs:
patch-test-and-publish:
runs-on: ubuntu-latest
timeout-minutes: 45
steps:
- uses: actions/checkout@v4
with:
ref: agent/harden-mnist-trace-log
fetch-depth: 0
- uses: actions/setup-python@v5
with:
python-version: "3.11"
cache: pip
cache-dependency-path: |
baseline/pyproject.toml
baseline/requirements.txt
- name: Apply trace-log robustness patch
shell: bash
run: |
python - <<'PY'
from pathlib import Path
diagnostics_path = Path("baseline/rg_baselines/diagnostics.py")
diagnostics = diagnostics_path.read_text(encoding="utf-8")
old_clean = '''\
def clean_positive_eigenvalues(values: Any) -> np.ndarray:

Check failure on line 41 in .github/workflows/harden-mnist-trace-log.yml

View workflow run for this annotation

GitHub Actions / .github/workflows/harden-mnist-trace-log.yml

Invalid workflow file

You have an error in your yaml syntax on line 41
"""Return finite positive eigenvalues in ascending order."""
evals = np.asarray(values, dtype=float).reshape(-1)
evals = evals[np.isfinite(evals) & (evals > 0.0)]
evals = np.sort(evals)
if evals.size < 2:
raise ValueError("fewer than two finite positive eigenvalues")
return evals
'''
new_clean = '''\
def clean_positive_eigenvalues(
values: Any,
*,
expected_dimension: Optional[int] = None,
) -> np.ndarray:
"""Return positive eigenvalues in ascending order.
When ``expected_dimension`` is supplied, fail closed if the ESD is
incomplete, non-finite, or rank deficient. This preserves
WeightWatcher's full-M normalization instead of silently renormalizing a
filtered positive-rank spectrum.
"""
evals = np.asarray(values, dtype=float).reshape(-1)
if expected_dimension is not None:
expected = int(expected_dimension)
if expected < 2:
raise ValueError("expected spectral dimension must be at least two")
if evals.size != expected:
raise ValueError(
"ESD dimension mismatch: "
f"expected {expected} eigenvalues, received {evals.size}"
)
if not np.all(np.isfinite(evals)):
raise ValueError("full ESD contains non-finite eigenvalues")
if np.any(evals <= 0.0):
positive = int(np.count_nonzero(evals > 0.0))
raise ValueError(
"rank-deficient ESD: "
f"expected {expected} positive eigenvalues, found {positive}"
)
else:
evals = evals[np.isfinite(evals) & (evals > 0.0)]
evals = np.sort(evals)
if evals.size < 2:
raise ValueError("fewer than two finite positive eigenvalues")
return evals
'''
if diagnostics.count(old_clean) != 1:
raise RuntimeError("unexpected clean_positive_eigenvalues source")
diagnostics = diagnostics.replace(old_clean, new_clean, 1)
old_metrics = '''\
def spectral_metrics_from_esd(
raw_evals_ascending: Any,
normalized_evals_ascending: Any,
*,
detx_num: int,
num_pl_spikes: int,
erg_gap: int,
) -> dict[str, float | int]:
"""Compute transparent metrics from one WeightWatcher ESD.
``normalized_evals_ascending`` must be produced by WeightWatcher's own
``RMT_Util.rescale_eigenvalues``. The trace-log boundary and gap are not
recomputed here: the supplied ``detx_num``, ``num_pl_spikes``, and
``erg_gap`` must come from ``watcher.analyze(ERG=True)``.
"""
raw = clean_positive_eigenvalues(raw_evals_ascending)
normalized = clean_positive_eigenvalues(normalized_evals_ascending)
if raw.size != normalized.size:
raise ValueError("raw and normalized ESDs have different sizes")
count = int(raw.size)
m_detx = int(np.clip(int(detx_num), 1, count))
m_pl = int(np.clip(int(num_pl_spikes), 1, count))
expected_gap = m_detx - m_pl
if int(erg_gap) != expected_gap:
raise ValueError(
f"WeightWatcher ERG_gap audit failed: {erg_gap} != {m_detx} - {m_pl}"
)
m_midpoint = int(np.clip(math.floor((m_detx + m_pl) / 2.0), 1, count))
'''
new_metrics = '''\
def spectral_metrics_from_esd(
raw_evals_ascending: Any,
normalized_evals_ascending: Any,
*,
detx_num: int,
num_pl_spikes: int,
erg_gap: int,
expected_dimension: Optional[int] = None,
) -> dict[str, float | int]:
"""Compute transparent metrics from one WeightWatcher ESD.
``normalized_evals_ascending`` must be produced by WeightWatcher's own
``RMT_Util.rescale_eigenvalues``. The trace-log boundary and gap are not
recomputed here: the supplied ``detx_num``, ``num_pl_spikes``, and
``erg_gap`` must come from ``watcher.analyze(ERG=True)``.
``expected_dimension`` is the full spectral dimension
``min(weight.shape)``. Strict baseline measurements require all of those
eigenvalues to be finite and positive so WeightWatcher's normalization is
not silently changed by positive-eigenvalue filtering.
"""
raw = clean_positive_eigenvalues(
raw_evals_ascending,
expected_dimension=expected_dimension,
)
normalized = clean_positive_eigenvalues(
normalized_evals_ascending,
expected_dimension=expected_dimension,
)
if raw.size != normalized.size:
raise ValueError("raw and normalized ESDs have different sizes")
count = int(raw.size)
normalized_sum = float(np.sum(normalized))
if not np.isclose(
normalized_sum,
float(count),
rtol=1e-10,
atol=1e-10 * max(count, 1),
):
raise ValueError(
"WeightWatcher normalization audit failed: "
f"sum={normalized_sum:.17g}, expected={count}"
)
m_detx = int(detx_num)
m_pl = int(num_pl_spikes)
if not 1 <= m_detx <= count:
raise ValueError(
f"detX_num must lie in [1, {count}], received {m_detx}"
)
if not 1 <= m_pl <= count:
raise ValueError(
f"num_pl_spikes must lie in [1, {count}], received {m_pl}"
)
expected_gap = m_detx - m_pl
if int(erg_gap) != expected_gap:
raise ValueError(
f"WeightWatcher ERG_gap audit failed: {erg_gap} != {m_detx} - {m_pl}"
)
m_midpoint = int(math.floor((m_detx + m_pl) / 2.0))
'''
if diagnostics.count(old_metrics) != 1:
raise RuntimeError("unexpected spectral_metrics_from_esd source")
diagnostics = diagnostics.replace(old_metrics, new_metrics, 1)
old_sum = '''\
"rescaled_eigenvalue_sum": float(np.sum(normalized)),
"rescale_sum_minus_num_eigenvalues": float(np.sum(normalized) - count),
'''
new_sum = '''\
"rescaled_eigenvalue_sum": normalized_sum,
"rescale_sum_minus_num_eigenvalues": float(normalized_sum - count),
'''
if diagnostics.count(old_sum) != 1:
raise RuntimeError("unexpected normalized-sum output source")
diagnostics = diagnostics.replace(old_sum, new_sum, 1)
old_measure = '''\
raw_esd = clean_positive_eigenvalues(
_get_esd_compat(
watcher,
model=model_cpu,
layer_id=int(layer_id),
params=get_esd_params,
)
)
normalized_esd, weight_scale = _rescale_with_weightwatcher(raw_esd)
computed = spectral_metrics_from_esd(
raw_esd,
normalized_esd,
detx_num=int(detx_num),
num_pl_spikes=int(num_pl_spikes),
erg_gap=erg_gap,
)
parameter = parameter_map.get(parameter_name) if parameter_name else None
'''
new_measure = '''\
parameter = parameter_map.get(parameter_name) if parameter_name else None
if parameter is None:
raise ValueError(
"WeightWatcher layer could not be matched to a model matrix"
)
expected_dimension = int(min(parameter.shape))
raw_esd = clean_positive_eigenvalues(
_get_esd_compat(
watcher,
model=model_cpu,
layer_id=int(layer_id),
params=get_esd_params,
),
expected_dimension=expected_dimension,
)
normalized_esd, weight_scale = _rescale_with_weightwatcher(raw_esd)
computed = spectral_metrics_from_esd(
raw_esd,
normalized_esd,
detx_num=int(detx_num),
num_pl_spikes=int(num_pl_spikes),
erg_gap=erg_gap,
expected_dimension=expected_dimension,
)
'''
if diagnostics.count(old_measure) != 1:
raise RuntimeError("unexpected WeightWatcher measurement source")
diagnostics = diagnostics.replace(old_measure, new_measure, 1)
old_shape = '''\
"layer_rows": int(parameter.shape[0]) if parameter is not None else np.nan,
"layer_cols": int(parameter.shape[1]) if parameter is not None else np.nan,
"layer_parameter_count": int(parameter.numel()) if parameter is not None else np.nan,
'''
new_shape = '''\
"layer_rows": int(parameter.shape[0]),
"layer_cols": int(parameter.shape[1]),
"layer_parameter_count": int(parameter.numel()),
'''
if diagnostics.count(old_shape) != 1:
raise RuntimeError("unexpected layer-shape source")
diagnostics = diagnostics.replace(old_shape, new_shape, 1)
diagnostics_path.write_text(diagnostics, encoding="utf-8")
tests_path = Path("baseline/tests/test_diagnostics.py")
tests_path.write_text(
'''\
import unittest
import numpy as np
from rg_baselines.diagnostics import (
clean_positive_eigenvalues,
spectral_metrics_from_esd,
)
class SpectralMetricsTests(unittest.TestCase):
def test_original_boundaries_and_midpoint(self) -> None:
raw = np.arange(1.0, 11.0)
normalized = raw * (len(raw) / raw.sum())
metrics = spectral_metrics_from_esd(
raw,
normalized,
detx_num=8,
num_pl_spikes=4,
erg_gap=4,
expected_dimension=len(raw),
)
self.assertEqual(metrics["m_midpoint"], 6)
self.assertEqual(metrics["ERG_gap"], 4)
self.assertAlmostEqual(metrics["rescaled_eigenvalue_sum"], 10.0)
self.assertAlmostEqual(
metrics["rescale_sum_minus_num_eigenvalues"],
0.0,
)
self.assertGreater(metrics["midpoint_energy_fraction"], 0.5)
def test_trace_log_matches_analytic_top_spectrum_value(self) -> None:
raw = np.asarray([1.0, 2.0, 4.0, 8.0])
normalized = raw * (len(raw) / raw.sum())
metrics = spectral_metrics_from_esd(
raw,
normalized,
detx_num=4,
num_pl_spikes=2,
erg_gap=2,
expected_dimension=4,
)
retained = normalized[::-1][:3]
expected_total = float(np.sum(np.log(retained)))
expected_per_eval = float(np.mean(np.log(retained)))
self.assertEqual(metrics["m_midpoint"], 3)
self.assertAlmostEqual(
metrics["trace_log_midpoint_total"],
expected_total,
)
self.assertAlmostEqual(
metrics["trace_log_midpoint_per_eval"],
expected_per_eval,
)
self.assertAlmostEqual(
metrics["geometric_mean_midpoint"],
float(np.exp(expected_per_eval)),
)
self.assertAlmostEqual(
metrics["trace_log_midpoint_total"],
3.0 * metrics["trace_log_midpoint_per_eval"],
)
def test_gap_mismatch_is_rejected(self) -> None:
raw = np.arange(1.0, 11.0)
normalized = raw * (len(raw) / raw.sum())
with self.assertRaisesRegex(ValueError, "ERG_gap audit failed"):
spectral_metrics_from_esd(
raw,
normalized,
detx_num=8,
num_pl_spikes=4,
erg_gap=3,
expected_dimension=len(raw),
)
def test_out_of_range_boundaries_are_rejected(self) -> None:
raw = np.arange(1.0, 6.0)
normalized = raw * (len(raw) / raw.sum())
for field, detx_num, num_pl_spikes in (
("detX_num", 6, 2),
("num_pl_spikes", 4, 0),
):
with self.subTest(field=field):
with self.assertRaisesRegex(ValueError, field):
spectral_metrics_from_esd(
raw,
normalized,
detx_num=detx_num,
num_pl_spikes=num_pl_spikes,
erg_gap=detx_num - num_pl_spikes,
expected_dimension=len(raw),
)
def test_rank_deficient_full_esd_is_rejected(self) -> None:
with self.assertRaisesRegex(ValueError, "rank-deficient ESD"):
clean_positive_eigenvalues(
[0.0, 1.0, 2.0],
expected_dimension=3,
)
def test_incomplete_full_esd_is_rejected(self) -> None:
with self.assertRaisesRegex(ValueError, "ESD dimension mismatch"):
clean_positive_eigenvalues(
[1.0, 2.0],
expected_dimension=3,
)
def test_incorrect_weightwatcher_normalization_is_rejected(self) -> None:
raw = np.asarray([1.0, 2.0, 3.0, 4.0])
with self.assertRaisesRegex(ValueError, "normalization audit failed"):
spectral_metrics_from_esd(
raw,
raw,
detx_num=4,
num_pl_spikes=2,
erg_gap=2,
expected_dimension=4,
)
if __name__ == "__main__":
unittest.main()
''',
encoding="utf-8",
)
PY
- name: Install complete baseline test environment
run: |
python -m pip install --upgrade pip
python -m pip install torch torchvision --index-url https://download.pytorch.org/whl/cpu
python -m pip install -e './baseline[experiment]'
python -m pip install nbformat
python -m pip check
- name: Compile and test baseline suite
env:
PYTHONPATH: baseline
MPLBACKEND: Agg
run: |
git diff --check
python -m compileall -q baseline/rg_baselines baseline/tests
python -m unittest baseline.tests.test_diagnostics -v
python -m unittest discover -s baseline/tests -v
- name: Commit tested patch
shell: bash
run: |
set -euo pipefail
git config user.name "github-actions[bot]"
git config user.email \
"41898282+github-actions[bot]@users.noreply.github.com"
git add \
baseline/rg_baselines/diagnostics.py \
baseline/tests/test_diagnostics.py
git commit -m "Harden MNIST trace-log diagnostics"
git push origin HEAD:agent/harden-mnist-trace-log