Skip to content

Commit ef0ec38

Browse files
committed
feat(v7-h05): per-step train event helper from extras
Closes V7-H05 (cppmega-mlx-wvzv) building block: synthesises a sequence of {step, loss, lr, grad_norm, mem_mb, throughput_tok_s, ts} events from finalised stage_train extras — same payload shape a real /ws/train/{job_id} WS push will eventually use. The UI can render the per-step table + sparkline as soon as the result modal opens, without waiting for a deeper stage_train callback rewrite (deferred follow-up). Tests (tests/v4/test_train_events.py): 4/4 — documented keys, throughput + mem populated, missing lr_trajectory tolerated, empty losses yields no events.
1 parent 1bf185c commit ef0ec38

2 files changed

Lines changed: 100 additions & 0 deletions

File tree

cppmega_v4/runtime/train_events.py

Lines changed: 50 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,50 @@
1+
"""V7-H05: per-step training event extraction from stage_train extras.
2+
3+
Produces a stream of {step, loss, lr, grad_norm, mem_mb,
4+
throughput_tok_s, ts} events derived from finalised extras. Until
5+
stage_train emits live callbacks (deeper rewrite), the UI can render
6+
post-hoc per-step events the moment the modal opens — same payload
7+
shape a real /ws/train/{job_id} stream will eventually push.
8+
9+
Usage:
10+
events = list(train_events_from_extras(extras, batch=1, seq=16))
11+
for e in events:
12+
ui.append_row(e)
13+
"""
14+
15+
from __future__ import annotations
16+
17+
import time
18+
from typing import Iterator
19+
20+
21+
def train_events_from_extras(extras: dict, *,
22+
batch: int = 1, seq: int = 16,
23+
start_ts: float | None = None,
24+
) -> Iterator[dict]:
25+
"""Yield per-step training events from stage_train extras."""
26+
losses = list(extras.get("losses", []))
27+
lrs = list(extras.get("lr_trajectory", []))
28+
elapsed_ms = float(extras.get("elapsed_ms", 0.0)) or 1.0
29+
n = max(1, len(losses))
30+
per_step_ms = elapsed_ms / n
31+
tokens_per_step = max(1, batch * seq)
32+
throughput = tokens_per_step / max(per_step_ms / 1000.0, 1e-6)
33+
base_ts = start_ts if start_ts is not None else time.time()
34+
mem_peak = extras.get("memory_peak_bytes")
35+
mem_mb = (round(int(mem_peak) / (1024 * 1024), 4)
36+
if mem_peak else None)
37+
for i, loss in enumerate(losses):
38+
lr = float(lrs[i]) if i < len(lrs) else None
39+
yield {
40+
"step": i,
41+
"loss": float(loss),
42+
"lr": lr,
43+
"grad_norm": None, # not snapshotted per-step today
44+
"mem_mb": mem_mb,
45+
"throughput_tok_s": round(throughput, 4),
46+
"ts": round(base_ts + (i + 1) * per_step_ms / 1000.0, 6),
47+
}
48+
49+
50+
__all__ = ["train_events_from_extras"]

tests/v4/test_train_events.py

Lines changed: 50 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,50 @@
1+
"""V7-H05: per-step train event generator from extras."""
2+
3+
from __future__ import annotations
4+
5+
from cppmega_v4.runtime.train_events import train_events_from_extras
6+
7+
8+
def _extras() -> dict:
9+
return {
10+
"losses": [5.0, 4.5, 4.1],
11+
"lr_trajectory": [1e-4, 3e-4, 1e-3],
12+
"elapsed_ms": 30.0,
13+
"memory_peak_bytes": 4 * 1024 * 1024,
14+
}
15+
16+
17+
def test_v7_h05_events_per_step_with_documented_keys():
18+
events = list(train_events_from_extras(_extras(),
19+
batch=1, seq=16))
20+
assert len(events) == 3
21+
keys = {"step", "loss", "lr", "grad_norm", "mem_mb",
22+
"throughput_tok_s", "ts"}
23+
for e in events:
24+
assert keys.issubset(e.keys()), f"missing keys in {e}"
25+
assert [e["step"] for e in events] == [0, 1, 2]
26+
assert [e["loss"] for e in events] == [5.0, 4.5, 4.1]
27+
assert [e["lr"] for e in events] == [1e-4, 3e-4, 1e-3]
28+
29+
30+
def test_v7_h05_throughput_positive_and_mem_present():
31+
events = list(train_events_from_extras(_extras(),
32+
batch=1, seq=16))
33+
for e in events:
34+
assert e["throughput_tok_s"] > 0
35+
assert e["mem_mb"] == 4.0
36+
assert e["ts"] > 0
37+
38+
39+
def test_v7_h05_handles_missing_lr_trajectory():
40+
e = {"losses": [1.0, 2.0]}
41+
events = list(train_events_from_extras(e, batch=1, seq=8))
42+
assert len(events) == 2
43+
for ev in events:
44+
assert ev["lr"] is None
45+
assert ev["mem_mb"] is None
46+
47+
48+
def test_v7_h05_empty_losses_yields_no_events():
49+
events = list(train_events_from_extras({"losses": []}))
50+
assert events == []

0 commit comments

Comments
 (0)