diff --git a/docs/PROVENANCE_INVENTORY.md b/docs/PROVENANCE_INVENTORY.md new file mode 100644 index 0000000..a83267f --- /dev/null +++ b/docs/PROVENANCE_INVENTORY.md @@ -0,0 +1,29 @@ +# Optimizer provenance field inventory + +**Class B/C companion.** Snapshot of whether packages emit local-delta-style +provenance fields (`actuator_id`, `ecs_backend`, `dose_definition`, `dose_value`, +`is_first_apply`, and where present `is_first_due`). Update when packages gain +logging. + +| Package | Provenance fields | Notes | +|---|---|---| +| `wwpgd_local_delta` | **yes** | SoT grammar (#34): null dose on no-op; schedule-aware first apply | +| `trace_log_tracker` | **pending / open PR** | Midpoint sibling: open collab PR **#49** (F1 `is_first_apply` vs `is_first_due`) when not yet on main | +| `self_consistent_trace_log_tracker` | **yes** (step stats) | This package: `ecs_backend=self_consistent_F_m`; first **successful** apply per param (F1); null dose on no-op | +| `adaptive_spectral_guard` | no | Different cadence model | +| `ecs_probe_loss_trace_wall` | no | | +| `spectral_rg_flow_projector` | no | | +| `full_matrix_log_rg` | no | | + +**Rule:** extend one package at a time; do not invent dose when no correction ran. + +**Field meanings (joint tables):** + +| Field | Meaning | +|---|---| +| `actuator_id` | Package / actuator name | +| `ecs_backend` | Support lineage label (`midpoint_pl_detx`, `self_consistent_F_m`, …) | +| `dose_definition` | Named ratio definition string | +| `dose_value` | Realized dose or **null** if no correction committed | +| `is_first_apply` | First **successful** correction for that parameter (not first schedule tick) | +| `is_first_due` | First schedule-due step (clock); may have null dose | diff --git a/optimizers/self_consistent_trace_log_tracker/README.md b/optimizers/self_consistent_trace_log_tracker/README.md index b9d3783..f8c6a70 100644 --- a/optimizers/self_consistent_trace_log_tracker/README.md +++ b/optimizers/self_consistent_trace_log_tracker/README.md @@ -187,6 +187,21 @@ optimizer.set_support_states(checkpoint.supports, replace=True) After each later outer checkpoint, call `set_support_states(...)` again. +## Provenance logging (step stats) + +`pop_step_stats()` rows include logging-only provenance fields aligned with the +local-delta package grammar and the midpoint `trace_log_tracker` extension: + +- `actuator_id = self_consistent_trace_log_tracker` +- `ecs_backend = self_consistent_F_m` (bulk-effective participation-ratio / `F(m)` lineage — **not** a free-fit MP edge label) +- `dose_definition = correction_frobenius_over_base_step_delta_frobenius` +- `dose_value` (null when no correction applied) +- `is_first_apply` (first **successful** correction per parameter) +- `is_first_due` (first schedule-due clock step; may have null dose) + +These fields do **not** change correction mathematics. See +[`docs/PROVENANCE_INVENTORY.md`](../../docs/PROVENANCE_INVENTORY.md). + ## Run the tests ```bash diff --git a/optimizers/self_consistent_trace_log_tracker/rg_sc_trace_log/wrapper.py b/optimizers/self_consistent_trace_log_tracker/rg_sc_trace_log/wrapper.py index 905f690..cdf92ce 100644 --- a/optimizers/self_consistent_trace_log_tracker/rg_sc_trace_log/wrapper.py +++ b/optimizers/self_consistent_trace_log_tracker/rg_sc_trace_log/wrapper.py @@ -136,6 +136,8 @@ def __init__( self.support_states: dict[str, AdaptiveSupportState] = {} self.global_step = 0 self._last_step_stats: list[dict[str, Any]] = [] + # First *successful* correction per parameter (dose not null); not first schedule-due. + self._applied_parameters: set[str] = set() @property def param_groups(self) -> list[MutableMapping[str, Any]]: @@ -210,6 +212,42 @@ def get_support_states(self) -> dict[str, AdaptiveSupportState]: def get_supports(self) -> dict[str, int]: return {name: int(state.ecs_rank) for name, state in self.support_states.items()} + def _first_due_step(self) -> int: + """First global_step index at which a correction is schedule-due.""" + warmup = int(self.config.warmup_steps) + every = int(self.config.apply_every_steps) + step = warmup + 1 + while step % every != 0: + step += 1 + return int(step) + + def _provenance_fields( + self, + *, + global_step: int, + dose_value: Optional[float], + parameter: str = "", + ) -> dict[str, Any]: + """Logging-only fields (local_delta #34 grammar + F1 scheduled≠applied). + + ``is_first_apply`` is true only on the first successful correction for + that parameter (``dose_value`` not null). Schedule-due steps with null + dose are never first-apply. ``is_first_due`` marks the first + schedule-due step for clock analysis (may be true with null dose). + """ + applied = dose_value is not None + is_first_apply = bool(applied and parameter not in self._applied_parameters) + if applied and parameter: + self._applied_parameters.add(parameter) + return { + "actuator_id": "self_consistent_trace_log_tracker", + "ecs_backend": "self_consistent_F_m", + "dose_definition": "correction_frobenius_over_base_step_delta_frobenius", + "dose_value": None if dose_value is None else float(dose_value), + "is_first_apply": is_first_apply, + "is_first_due": int(global_step) == self._first_due_step(), + } + def pop_step_stats(self) -> list[dict[str, Any]]: stats = self._last_step_stats self._last_step_stats = [] @@ -287,6 +325,11 @@ def _prepare_geometries( "parameter": name, "status": "geometry_failed", "reason": str(exc), + **self._provenance_fields( + global_step=next_step, + dose_value=None, + parameter=name, + ), } ) continue @@ -306,6 +349,11 @@ def _prepare_geometries( "coordinate_residual": float( geometry.residual.detach().float().cpu() ), + **self._provenance_fields( + global_step=next_step, + dose_value=None, + parameter=name, + ), } ) continue @@ -361,6 +409,7 @@ def step(self, closure: Optional[Any] = None) -> Any: ) parameter.copy_(before + result.corrected_delta) self.support_states[name] = new_state + dose = float(result.correction_ratio) if result.applied else None self._last_step_stats.append( { @@ -411,6 +460,11 @@ def step(self, closure: Optional[Any] = None) -> Any: "largest_retained_singular_value": ( geometry.largest_retained_singular_value ), + **self._provenance_fields( + global_step=self.global_step, + dose_value=dose, + parameter=name, + ), } ) return loss @@ -422,6 +476,7 @@ def state_dict(self) -> dict[str, Any]: name: state.to_dict() for name, state in self.support_states.items() }, "global_step": int(self.global_step), + "applied_parameters": sorted(self._applied_parameters), "config": asdict(self.config), } @@ -430,3 +485,8 @@ def load_state_dict(self, state_dict: Mapping[str, Any]) -> None: self.support_states = {} self.set_support_states(state_dict.get("support_states", {}), replace=True) self.global_step = int(state_dict.get("global_step", 0)) + applied = state_dict.get("applied_parameters") + if isinstance(applied, (list, set, tuple)): + self._applied_parameters = set(str(x) for x in applied) + else: + self._applied_parameters = set() diff --git a/optimizers/self_consistent_trace_log_tracker/tests/test_provenance.py b/optimizers/self_consistent_trace_log_tracker/tests/test_provenance.py new file mode 100644 index 0000000..9e2c79d --- /dev/null +++ b/optimizers/self_consistent_trace_log_tracker/tests/test_provenance.py @@ -0,0 +1,120 @@ +"""Provenance logging fields on SC trace-log step stats (logging only).""" + +from __future__ import annotations + +import unittest + +import torch +import torch.nn as nn + +from rg_sc_trace_log.ecs import AdaptiveSupportState +from rg_sc_trace_log.wrapper import ( + SelfConsistentTraceLogConfig, + SelfConsistentTraceLogRGWrapper, +) + + +class SCProvenanceTests(unittest.TestCase): + def _make_wrapper(self, **config_kwargs): + torch.manual_seed(0) + model = nn.Linear(6, 9, bias=False).double() + base = torch.optim.SGD(model.parameters(), lr=0.05) + cfg = SelfConsistentTraceLogConfig( + mode="one_sided", + min_retained=2, + min_ecs_size=2, + correction_scale=1.0, + max_correction_ratio=None, + ridge_relative=0.0, + bootstrap_without_weightwatcher=False, + refresh_ecs_every_steps=0, + **config_kwargs, + ) + wrapper = SelfConsistentTraceLogRGWrapper( + base, + model.named_parameters(), + config=cfg, + ) + state = AdaptiveSupportState( + ecs_rank=4, + normalization_dimension=5.0, + bulk_effective_count=1.0, + trace_log_per_eval=0.0, + status="test", + pl_rank=4, + working_rank=4, + ) + wrapper.set_support_states({"weight": state}, replace=True) + return model, wrapper + + def test_ok_rows_carry_provenance_and_dose(self): + model, wrapper = self._make_wrapper(warmup_steps=0, apply_every_steps=1) + model.zero_grad(set_to_none=True) + (model.weight ** 2).sum().backward() + wrapper.step() + stats = wrapper.pop_step_stats() + self.assertTrue(stats) + for row in stats: + self.assertEqual(row["actuator_id"], "self_consistent_trace_log_tracker") + self.assertEqual(row["ecs_backend"], "self_consistent_F_m") + self.assertEqual( + row["dose_definition"], + "correction_frobenius_over_base_step_delta_frobenius", + ) + self.assertIn( + row["status"], + {"ok", "skipped", "geometry_failed", "geometry_skipped"}, + ) + self.assertIn("is_first_due", row) + if row["status"] == "ok": + self.assertIsNotNone(row["dose_value"]) + self.assertGreaterEqual(float(row["dose_value"]), 0.0) + self.assertIs(row["is_first_apply"], True) + elif row["status"] in { + "skipped", + "geometry_failed", + "geometry_skipped", + }: + self.assertIsNone(row["dose_value"]) + self.assertIs(row["is_first_apply"], False) + + def test_first_apply_respects_warmup_and_cadence(self): + model, wrapper = self._make_wrapper(warmup_steps=2, apply_every_steps=2) + self.assertEqual(wrapper._first_due_step(), 4) + for _ in range(4): + model.zero_grad(set_to_none=True) + (model.weight ** 2).sum().backward() + wrapper.step() + stats = wrapper.pop_step_stats() + self.assertTrue(stats) + for row in stats: + self.assertTrue(row["is_first_due"]) + if row["dose_value"] is not None: + self.assertTrue(row["is_first_apply"]) + else: + self.assertFalse(row["is_first_apply"]) + for _ in range(2): + model.zero_grad(set_to_none=True) + (model.weight ** 2).sum().backward() + wrapper.step() + stats2 = wrapper.pop_step_stats() + self.assertTrue(stats2) + self.assertTrue(all(not row["is_first_due"] for row in stats2)) + self.assertTrue(all(not row["is_first_apply"] for row in stats2)) + + def test_state_dict_preserves_applied_parameters(self): + model, wrapper = self._make_wrapper(warmup_steps=0, apply_every_steps=1) + model.zero_grad(set_to_none=True) + (model.weight ** 2).sum().backward() + wrapper.step() + stats = wrapper.pop_step_stats() + if any(r.get("dose_value") is not None for r in stats): + state = wrapper.state_dict() + self.assertIn("applied_parameters", state) + _, second = self._make_wrapper(warmup_steps=0, apply_every_steps=1) + second.load_state_dict(state) + self.assertTrue(second._applied_parameters) + + +if __name__ == "__main__": + unittest.main()