Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
17 changes: 17 additions & 0 deletions docs/PROVENANCE_INVENTORY.md
Original file line number Diff line number Diff line change
@@ -0,0 +1,17 @@
# 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`). 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` | **yes** (step stats) | #34-style fields + `is_first_due`; `is_first_apply` = first **successful** apply per param (F1) |
| `self_consistent_trace_log_tracker` | no | Candidate for same grammar |
| `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.
9 changes: 9 additions & 0 deletions optimizers/trace_log_tracker/README.md
Original file line number Diff line number Diff line change
Expand Up @@ -135,3 +135,12 @@ The WeightWatcher analysis updates the RG model's cached midpoint ranks for the
cd optimizers/trace_log_tracker
PYTHONPATH=. python -m unittest discover -s tests -v
```

## Provenance fields on step stats

`pop_step_stats()` rows include logging-only provenance fields aligned with the
local-delta package grammar: `actuator_id`, `ecs_backend`, `dose_definition`,
`dose_value` (null when no correction applied), `is_first_apply` (first
**successful** correction per parameter — not first schedule-due step), and
`is_first_due` (first schedule-due clock step; may have null dose). These
fields do **not** change correction mathematics.
58 changes: 57 additions & 1 deletion optimizers/trace_log_tracker/rg_trace_log/wrapper.py
Original file line number Diff line number Diff line change
Expand Up @@ -87,6 +87,8 @@ def __init__(
self.supports: dict[str, int] = {}
self.global_step = 0
self._last_step_stats: list[dict[str, Any]] = []
# First *successful* correction per parameter (dose not null); not first schedule-due step.
self._applied_parameters: set[str] = set()

@property
def param_groups(self) -> list[MutableMapping[str, Any]]:
Expand Down Expand Up @@ -129,6 +131,43 @@ def set_supports(self, supports: Mapping[str, int]) -> None:
def get_supports(self) -> dict[str, int]:
return dict(self.supports)


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": "trace_log_tracker",
"ecs_backend": "midpoint_pl_detx",
"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 = []
Expand All @@ -154,11 +193,13 @@ def _prepare_geometries(self) -> dict[str, tuple[torch.Tensor, TraceLogGeometry]
eps=self.config.eps,
)
except (RuntimeError, ValueError) as exc:
step_idx = self.global_step + 1
self._last_step_stats.append({
"global_step": self.global_step + 1,
"global_step": step_idx,
"parameter": name,
"status": "geometry_failed",
"reason": str(exc),
**self._provenance_fields(global_step=step_idx, dose_value=None, parameter=name),
})
continue
prepared[name] = (before, geometry)
Expand Down Expand Up @@ -190,6 +231,7 @@ def step(self, closure: Optional[Any] = None) -> Any:
eps=self.config.eps,
)
parameter.copy_(before + result.corrected_delta)
dose = float(result.correction_ratio) if result.applied else None
self._last_step_stats.append({
"global_step": self.global_step,
"parameter": name,
Expand All @@ -209,6 +251,11 @@ def step(self, closure: Optional[Any] = None) -> Any:
"gradient_radial_inner_product": geometry.radial_inner_product,
"smallest_retained_singular_value": geometry.smallest_retained_singular_value,
"largest_retained_singular_value": geometry.largest_retained_singular_value,
**self._provenance_fields(
global_step=self.global_step,
dose_value=dose,
parameter=name,
),
})
return loss

Expand All @@ -217,10 +264,19 @@ def state_dict(self) -> dict[str, Any]:
"base_optimizer": self.base_optimizer.state_dict(),
"supports": dict(self.supports),
"global_step": int(self.global_step),
"applied_parameters": sorted(self._applied_parameters),
"config": asdict(self.config),
}

def load_state_dict(self, state_dict: Mapping[str, Any]) -> None:
self.base_optimizer.load_state_dict(state_dict["base_optimizer"])
self.set_supports(state_dict.get("supports", {}))
self.global_step = int(state_dict.get("global_step", 0))
applied = state_dict.get("applied_parameters", state_dict.get("has_applied_correction"))
if isinstance(applied, (list, set, tuple)):
self._applied_parameters = set(str(x) for x in applied)
elif applied:
# legacy bool: mark all current supports as already applied
self._applied_parameters = set(self.supports)
else:
self._applied_parameters = set()
92 changes: 92 additions & 0 deletions optimizers/trace_log_tracker/tests/test_provenance.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,92 @@
"""Provenance logging fields on trace-log step stats (logging only)."""

from __future__ import annotations

import unittest

import torch
import torch.nn as nn

from rg_trace_log.wrapper import TraceLogConfig, TraceLogRGWrapper


class TraceLogProvenanceTests(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 = TraceLogConfig(
mode="one_sided",
normalization="raw",
min_retained=2,
correction_scale=1.0,
max_correction_ratio=None,
**config_kwargs,
)
wrapper = TraceLogRGWrapper(
base,
model.named_parameters(),
config=cfg,
)
# retain most modes so geometry succeeds
wrapper.set_supports({"weight": 4})
return model, wrapper

def test_ok_rows_carry_provenance_and_dose(self):
model, wrapper = self._make_wrapper(warmup_steps=0, apply_every_steps=1)
# force a contracting step: set grad opposite a random weight move
model.zero_grad(set_to_none=True)
loss = (model.weight ** 2).sum()
loss.backward()
wrapper.step()
stats = wrapper.pop_step_stats()
self.assertTrue(stats)
for row in stats:
self.assertEqual(row["actuator_id"], "trace_log_tracker")
self.assertEqual(row["ecs_backend"], "midpoint_pl_detx")
self.assertEqual(
row["dose_definition"],
"correction_frobenius_over_base_step_delta_frobenius",
)
self.assertIn(row["status"], {"ok", "skipped", "geometry_failed"})
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"}:
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)
# first due step is 4 ( >2 and multiple of 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()
# steps 1-3: no correction stats (not due); step 4: stats with first_apply
# pop after each? step clears; only last step has stats
stats = wrapper.pop_step_stats()
self.assertTrue(stats)
# first due may skip with null dose; first_apply only if applied
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"])
# one more due step (6): not first due; not first apply if already applied
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))


if __name__ == "__main__":
unittest.main()