Skip to content

Commit f3e8a6c

Browse files
committed
feat(v7-h07): per_brick_grad_norms in train extras + 1/1 pytest
1 parent 8bdde83 commit f3e8a6c

2 files changed

Lines changed: 56 additions & 0 deletions

File tree

cppmega_v4/runner/stages.py

Lines changed: 13 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1747,6 +1747,8 @@ def _base(k: str) -> str | None:
17471747
"top1_token_drift": top1_token_drift,
17481748
},
17491749
"side_channels_observed": side_channels_observed,
1750+
"per_brick_grad_norms": _safe_per_brick_grads(
1751+
all_modules, grads if 'grads' in dir() else None),
17501752
"side_channels_forward_effect": {
17511753
"doc_ids_mask_density": sc_doc_ids_mask_density,
17521754
"doc_mask_applied": sc_doc_mask_applied,
@@ -2174,6 +2176,17 @@ def read_ckpt_metadata(path: str) -> dict | None:
21742176
return out or None
21752177

21762178

2179+
def _safe_per_brick_grads(model: Any, grads: Any) -> dict[str, float]:
2180+
"""V7-H07: extras.per_brick_grad_norms — best-effort wrapper."""
2181+
if grads is None:
2182+
return {}
2183+
try:
2184+
from cppmega_v4.runtime.per_brick_probes import per_brick_grad_norms
2185+
return per_brick_grad_norms(model, grads)
2186+
except Exception:
2187+
return {}
2188+
2189+
21772190
def _ema_smooth(values: list[float], window: int = 10) -> list[float]:
21782191
"""G15: simple windowed-mean smoothing (cheaper than true EMA, same
21792192
job for visualising convergence)."""
Lines changed: 43 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,43 @@
1+
"""V7-H07: extras.per_brick_grad_norms surfaces real per-layer grads."""
2+
3+
from __future__ import annotations
4+
5+
from cppmega_v4.jsonrpc.schema import VerifyParams
6+
from cppmega_v4.runner import Pipeline, run_pipeline
7+
8+
9+
def _spec() -> VerifyParams:
10+
return VerifyParams.model_validate({
11+
"graph": {
12+
"nodes": [
13+
{"id": "attn", "kind": "attention",
14+
"params": {"num_heads": 4, "head_dim": 64}},
15+
{"id": "mlp", "kind": "mlp", "params": {}},
16+
],
17+
"edges": [{"src": "attn", "dst": "mlp"}],
18+
},
19+
"dim_env": {"B": 1, "S": 8, "H": 128,
20+
"nh": 2, "nkv": 1, "head_dim": 64},
21+
"loss": {"kind": "cross_entropy", "head_outputs": ["mlp"]},
22+
"optim": {"kind": "adamw",
23+
"groups": [{"matcher": "all", "lr": 1e-3,
24+
"weight_decay": 0.01,
25+
"betas": [0.9, 0.95]}]},
26+
})
27+
28+
29+
def test_v7_h07_extras_carries_per_brick_grad_norms():
30+
rep = run_pipeline(_spec(), Pipeline.from_dict({
31+
"stages": ["parse", "verify_build_spec", "build_model", "train"],
32+
"stage_options": {"train": {"num_steps": 2}},
33+
}))
34+
tr = next(s for s in rep.stages if s.name == "train")
35+
pbg = tr.extras["per_brick_grad_norms"]
36+
assert isinstance(pbg, dict)
37+
assert len(pbg) >= 1
38+
for k, v in pbg.items():
39+
assert k.startswith("layers.") or k in (
40+
"shared_expert", "train_token_embedding",
41+
"side_channel_token_embedding",
42+
)
43+
assert v > 0

0 commit comments

Comments
 (0)