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
29 changes: 29 additions & 0 deletions docs/PROVENANCE_INVENTORY.md
Original file line number Diff line number Diff line change
@@ -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 |
15 changes: 15 additions & 0 deletions optimizers/self_consistent_trace_log_tracker/README.md
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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]]:
Expand Down Expand Up @@ -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 = []
Expand Down Expand Up @@ -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
Expand All @@ -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
Expand Down Expand Up @@ -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(
{
Expand Down Expand Up @@ -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
Expand All @@ -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),
}

Expand All @@ -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()
120 changes: 120 additions & 0 deletions optimizers/self_consistent_trace_log_tracker/tests/test_provenance.py
Original file line number Diff line number Diff line change
@@ -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()