Skip to content

Commit 660a777

Browse files
committed
feat(v5-g21): dry_forward / loss_smoke / optimizer_smoke return rich extras
Closes V5-G21 / cppmega-mlx-aih. Other-than-train pipeline stages returned barebones results; now report observable fields: - dry_forward.{batch, seq_len, hidden, verdict, num_nodes} - loss_smoke.{loss_value, loss_finite, seq_len} - optimizer_smoke.{optimizer_kind, num_groups} 3/3 pytest + 79/79 broader regression green.
1 parent 7095689 commit 660a777

2 files changed

Lines changed: 77 additions & 2 deletions

File tree

cppmega_v4/runner/stages.py

Lines changed: 24 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -253,13 +253,20 @@ def stage_dry_forward(ctx: StageContext) -> StageResult:
253253
graph = from_block_specs(specs, hidden_size=hidden, instantiate=False)
254254
result = dry_forward(graph, hidden_size=hidden, seq_len=seq, batch=batch)
255255
ctx.dry_forward_verdict = result.verdict
256+
# G21: rich extras — observable B/S/H + verdict for modal display
257+
rich: dict[str, Any] = {
258+
"batch": batch, "seq_len": seq, "hidden": hidden,
259+
"verdict": result.verdict,
260+
"num_nodes": len(graph.nodes),
261+
}
256262
return StageResult(
257263
name="dry_forward",
258264
status="ok" if result.verdict == "ok" else "fail",
259265
elapsed_ms=(time.perf_counter() - t0) * 1000.0,
260266
error=({"type": result.verdict, "detail": result.detail}
261267
if result.verdict != "ok" else None),
262268
errors=0 if result.verdict == "ok" else 1,
269+
extras=rich,
263270
)
264271
except Exception as exc:
265272
return _fail("dry_forward", t0, exc)
@@ -316,22 +323,37 @@ def stage_loss_smoke(ctx: StageContext) -> StageResult:
316323
logits = mx.random.normal((1, seq, 32))
317324
loss_value = mx.mean(mx.softmax(logits, axis=-1) * 0.0 + 1.0)
318325
finite = bool(mx.isfinite(loss_value).item())
326+
# G21: rich extras — loss value + finite flag observable
319327
return StageResult(
320328
name="loss_smoke",
321329
status="ok" if finite else "fail",
322330
elapsed_ms=(time.perf_counter() - t0) * 1000.0,
323331
errors=0 if finite else 1,
324332
error=(None if finite
325333
else {"type": "NonFiniteLoss", "detail": str(loss_value)}),
334+
extras={
335+
"loss_value": round(float(loss_value.item()), 6),
336+
"loss_finite": finite,
337+
"seq_len": seq,
338+
},
326339
)
327340
except Exception as exc:
328341
return _fail("loss_smoke", t0, exc)
329342

330343

331344
def stage_optimizer_smoke(ctx: StageContext) -> StageResult:
332-
"""No-op for now — full optimizer.update wired in F-A.3 (training stage)."""
345+
"""G21: report optim kind + group counts observably."""
333346
t0 = time.perf_counter()
334-
return _ok("optimizer_smoke", t0, note="placeholder until training stage lands")
347+
spec_optim = getattr(ctx.spec, "optim", None)
348+
kind = "adamw"
349+
num_groups = 1
350+
if spec_optim is not None:
351+
kind = str(getattr(spec_optim, "kind", "adamw"))
352+
groups = getattr(spec_optim, "groups", None) or []
353+
num_groups = len(groups)
354+
return _ok("optimizer_smoke", t0,
355+
note="placeholder until training stage lands",
356+
optimizer_kind=kind, num_groups=num_groups)
335357

336358

337359
def stage_train(ctx: StageContext) -> StageResult:
Lines changed: 53 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,53 @@
1+
"""G21: dry_forward / loss_smoke / optimizer_smoke return rich extras."""
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", "params": {}},
14+
{"id": "mlp", "kind": "mlp",
15+
"params": {"intermediate_size": 64, "activation": "swiglu"}},
16+
],
17+
"edges": [{"src": "attn", "dst": "mlp"}],
18+
},
19+
"dim_env": {"B": 1, "S": 8, "H": 32, "nh": 2, "nkv": 1, "head_dim": 16},
20+
"loss": {"kind": "cross_entropy", "head_outputs": ["mlp"]},
21+
"optim": {"kind": "adamw",
22+
"groups": [{"matcher": "all", "lr": 1e-3,
23+
"weight_decay": 0.01, "betas": [0.9, 0.95]}]},
24+
})
25+
26+
27+
def _stages(stage_names: list[str]) -> dict[str, dict]:
28+
r = run_pipeline(_spec(), Pipeline.from_dict({"stages": stage_names}))
29+
return {s.name: s.to_dict() for s in r.stages}
30+
31+
32+
def test_dry_forward_rich_extras():
33+
out = _stages(["parse", "verify_build_spec", "dry_forward"])
34+
df = out["dry_forward"]
35+
assert df["batch"] == 1
36+
assert df["seq_len"] == 8
37+
assert df["hidden"] == 32
38+
assert df["num_nodes"] >= 1
39+
40+
41+
def test_loss_smoke_rich_extras():
42+
out = _stages(["parse", "verify_build_spec", "loss_smoke"])
43+
ls = out["loss_smoke"]
44+
assert ls["loss_finite"] is True
45+
assert ls["seq_len"] == 8
46+
assert isinstance(ls["loss_value"], float)
47+
48+
49+
def test_optimizer_smoke_rich_extras():
50+
out = _stages(["parse", "verify_build_spec", "optimizer_smoke"])
51+
os_ = out["optimizer_smoke"]
52+
assert os_["optimizer_kind"] == "adamw"
53+
assert os_["num_groups"] == 1

0 commit comments

Comments
 (0)