|
| 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