Skip to content

Commit 3c4bc69

Browse files
committed
feat(data): add forward data smoke
1 parent 3d17615 commit 3c4bc69

4 files changed

Lines changed: 203 additions & 1 deletion

File tree

docs/data_pipeline.md

Lines changed: 11 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -128,10 +128,21 @@ The intended Stream D scale gate is:
128128
--include-structure-ids
129129
```
130130

131+
## Local Forward Smoke
132+
133+
scripts/data_smoke.py remains ingress-only by default, but --forward-smoke
134+
adds a forward-only tiny HybridTinyLM check on the first full batch. The script
135+
threads LMTokenBatch.model_kwargs() side channels and token-aligned
136+
document_ids into the model, reports finite next-token loss and logits shape,
137+
and keeps training_wired=false with no GB10, distributed Megatron, or
138+
M4-vs-GB10 parity claim.
139+
131140
## Current Guardrails
132141

133142
- Packing is exported from cppmega_mlx.data for callers and tests.
134143
- The current training ingress still consumes dense LMTokenBatch rows.
144+
- scripts/data_smoke.py --forward-smoke verifies local batch -> model forward
145+
closure without training, benchmarking, or parity claims.
135146
- PyTorch DataLoader integration is explicit and optional; the MLX training hot
136147
path does not import torch unless the bridge is requested.
137148
- Mapping-batch and LMTokenBatch training can carry explicit packed document

scripts/data_smoke.py

Lines changed: 112 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,5 @@
11
#!/usr/bin/env python3
2-
"""Smoke local token dataset ingress without wiring training."""
2+
"""Smoke local token dataset ingress with an optional forward-only check."""
33

44
from __future__ import annotations
55

@@ -118,6 +118,32 @@ def build_parser() -> argparse.ArgumentParser:
118118
default=None,
119119
help="Optional dtype for raw .bin Megatron handoffs without .idx metadata.",
120120
)
121+
parser.add_argument(
122+
"--forward-smoke",
123+
action="store_true",
124+
help=(
125+
"Run a tiny HybridTinyLM forward/loss on the first opened batch. "
126+
"This is forward-only and does not wire training or parity claims."
127+
),
128+
)
129+
parser.add_argument(
130+
"--forward-seed",
131+
type=int,
132+
default=0,
133+
help="MLX random seed used when constructing the forward-smoke model.",
134+
)
135+
parser.add_argument(
136+
"--forward-hidden-size",
137+
type=int,
138+
default=16,
139+
help="Hidden size for the tiny forward-smoke HybridTinyLM.",
140+
)
141+
parser.add_argument(
142+
"--forward-attention-heads",
143+
type=int,
144+
default=4,
145+
help="Attention head count for the tiny forward-smoke HybridTinyLM.",
146+
)
121147
return parser
122148

123149

@@ -196,6 +222,11 @@ def run_smoke(args: argparse.Namespace) -> dict[str, Any]:
196222
payload["packing"] = _packing_receipt(first_batch, args)
197223
else:
198224
payload["packing"] = {"enabled": False}
225+
if args.forward_smoke:
226+
payload["forward"] = _forward_receipt(first_batch, dataset, args)
227+
payload["forward_wired"] = True
228+
else:
229+
payload["forward"] = {"enabled": False}
199230
return payload
200231

201232

@@ -311,6 +342,85 @@ def _packing_receipt(batch: LMTokenBatch, args: argparse.Namespace) -> dict[str,
311342
}
312343

313344

345+
def _forward_receipt(
346+
batch: LMTokenBatch,
347+
dataset: TokenBatchDataset,
348+
args: argparse.Namespace,
349+
) -> dict[str, Any]:
350+
if args.forward_hidden_size <= 0:
351+
raise SmokeError("--forward-hidden-size must be positive")
352+
if args.forward_attention_heads <= 0:
353+
raise SmokeError("--forward-attention-heads must be positive")
354+
if args.forward_hidden_size % args.forward_attention_heads != 0:
355+
raise SmokeError(
356+
"--forward-hidden-size must be divisible by --forward-attention-heads"
357+
)
358+
359+
from cppmega_mlx.models.hybrid_lm import HybridTinyConfig, HybridTinyLM
360+
import mlx.nn as nn
361+
362+
vocab_size = _forward_vocab_size(dataset)
363+
input_seq_len = int(batch.inputs.shape[1])
364+
mx.random.seed(int(args.forward_seed))
365+
config = HybridTinyConfig(
366+
vocab_size=vocab_size,
367+
hidden_size=int(args.forward_hidden_size),
368+
num_attention_heads=int(args.forward_attention_heads),
369+
depth=1,
370+
pattern="A",
371+
max_seq_length=max(2, input_seq_len),
372+
attention_sparse_topk=max(1, min(16, input_seq_len)),
373+
)
374+
model = HybridTinyLM(config)
375+
model_kwargs = batch.model_kwargs()
376+
document_ids = batch.input_document_ids
377+
if document_ids is not None:
378+
model_kwargs["document_ids"] = document_ids
379+
380+
logits = model(batch.inputs, **model_kwargs)
381+
targets = batch.targets
382+
token_losses = nn.losses.cross_entropy(
383+
logits.astype(mx.float32),
384+
targets,
385+
reduction="none",
386+
)
387+
mask = batch.target_mask
388+
ntokens = mask.sum()
389+
denom = mx.maximum(ntokens, mx.array(1.0, dtype=mx.float32))
390+
loss = (token_losses * mask).astype(mx.float32).sum() / denom
391+
mx.eval(logits, loss, ntokens)
392+
loss_value = float(loss.item())
393+
ntokens_value = float(ntokens.item())
394+
side_channel_kwargs = sorted(key for key in model_kwargs if key != "document_ids")
395+
return {
396+
"document_ids_used": document_ids is not None,
397+
"enabled": True,
398+
"finite_loss": bool(np.isfinite(loss_value)),
399+
"forward_only": True,
400+
"logits_shape": [int(dim) for dim in logits.shape],
401+
"loss": loss_value,
402+
"model": {
403+
"class": "HybridTinyLM",
404+
"hidden_size": int(args.forward_hidden_size),
405+
"num_attention_heads": int(args.forward_attention_heads),
406+
"pattern": "A",
407+
"seed": int(args.forward_seed),
408+
"vocab_size": vocab_size,
409+
},
410+
"ntokens": ntokens_value,
411+
"side_channel_model_kwargs": side_channel_kwargs,
412+
"target_shape": [int(dim) for dim in batch.targets.shape],
413+
"training_wired": False,
414+
}
415+
416+
417+
def _forward_vocab_size(dataset: TokenBatchDataset) -> int:
418+
token_min, token_max = dataset.token_id_range()
419+
if token_min < 0:
420+
raise SmokeError("forward smoke requires non-negative token IDs")
421+
return max(2, int(dataset.metadata.vocab_size), int(token_max) + 1)
422+
423+
314424
def _base_receipt(
315425
*,
316426
status: str,
@@ -326,6 +436,7 @@ def _base_receipt(
326436
"receipt_scope": "local_token_dataset_ingress_smoke",
327437
"status": status,
328438
"trainable_metal_kernel_adoption_claim": False,
439+
"forward_wired": False,
329440
"training_wired": False,
330441
}
331442
if dataset_format is not None:

tests/test_data_pipeline_doc.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -25,6 +25,7 @@ def test_data_pipeline_doc_states_ingress_packing_and_guardrails() -> None:
2525
"Multi-shard Megatron indexed directories batch across shard boundaries",
2626
"scripts/megatron_ingress_stress.py generates local Megatron Indexed fixtures",
2727
"explicitly makes no GB10, distributed Megatron, or M4-vs-GB10 parity claim",
28+
"scripts/data_smoke.py --forward-smoke verifies local batch -> model forward closure",
2829
"PyTorch DataLoader integration is explicit and optional",
2930
"M0.1 tokenizer parity is already closed",
3031
"schema drift fails closed",

tests/test_data_smoke_script.py

Lines changed: 79 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -114,6 +114,8 @@ def test_npz_smoke_reports_local_ingress_contract(tmp_path: Path) -> None:
114114
assert payload["m4_vs_gb10_parity_claim"] is False
115115
assert payload["distributed_megatron_parity_claim"] is False
116116
assert payload["trainable_metal_kernel_adoption_claim"] is False
117+
assert payload["forward_wired"] is False
118+
assert payload["forward"] == {"enabled": False}
117119
assert payload["training_wired"] is False
118120

119121

@@ -247,6 +249,83 @@ def test_megatron_multishard_smoke_reports_side_channels(tmp_path: Path) -> None
247249
assert payload["structure_side_channels"] == ["structure_ids", "dep_levels"]
248250
assert payload["structure_side_channels_present"] is True
249251
assert payload["distributed_megatron_parity_claim"] is False
252+
assert payload["forward_wired"] is False
253+
254+
255+
def test_npz_smoke_can_run_forward_only_model_path(tmp_path: Path) -> None:
256+
npz_path = tmp_path / "tokens.npz"
257+
write_npz(npz_path, include_structure=True)
258+
259+
result = run_script(
260+
str(npz_path),
261+
"--dataset-format",
262+
"npz",
263+
"--batch-size",
264+
"2",
265+
"--seq-len",
266+
"4",
267+
"--forward-smoke",
268+
)
269+
270+
assert result.returncode == 0, result.stderr
271+
payload = load_json(result)
272+
assert payload["status"] == "ok"
273+
assert payload["forward_wired"] is True
274+
assert payload["training_wired"] is False
275+
assert payload["forward"]["enabled"] is True
276+
assert payload["forward"]["forward_only"] is True
277+
assert payload["forward"]["finite_loss"] is True
278+
assert payload["forward"]["logits_shape"] == [2, 3, 32]
279+
assert payload["forward"]["target_shape"] == [2, 3]
280+
assert payload["forward"]["ntokens"] == 6.0
281+
assert payload["forward"]["document_ids_used"] is False
282+
assert payload["forward"]["side_channel_model_kwargs"] == [
283+
"dep_levels",
284+
"structure_ids",
285+
]
286+
assert payload["forward"]["training_wired"] is False
287+
assert payload["gb10_parity_claim"] is False
288+
289+
290+
def test_megatron_multishard_smoke_can_run_forward_with_document_ids(
291+
tmp_path: Path,
292+
) -> None:
293+
_write_structured_multishard_fixture(
294+
tmp_path,
295+
shard_docs=[
296+
[np.arange(8, dtype=np.int32)],
297+
[np.arange(100, 108, dtype=np.int32)],
298+
],
299+
)
300+
301+
result = run_script(
302+
str(tmp_path),
303+
"--batch-size",
304+
"2",
305+
"--seq-len",
306+
"4",
307+
"--require-structure-side-channels",
308+
"--forward-smoke",
309+
)
310+
311+
assert result.returncode == 0, result.stderr
312+
payload = load_json(result)
313+
assert payload["status"] == "ok"
314+
assert payload["dataset_format"] == "megatron"
315+
assert payload["forward_wired"] is True
316+
assert payload["training_wired"] is False
317+
assert payload["forward"]["enabled"] is True
318+
assert payload["forward"]["finite_loss"] is True
319+
assert payload["forward"]["logits_shape"] == [2, 3, 256]
320+
assert payload["forward"]["target_shape"] == [2, 3]
321+
assert payload["forward"]["ntokens"] == 6.0
322+
assert payload["forward"]["document_ids_used"] is True
323+
assert payload["forward"]["side_channel_model_kwargs"] == [
324+
"dep_levels",
325+
"structure_ids",
326+
]
327+
assert payload["forward"]["training_wired"] is False
328+
assert payload["distributed_megatron_parity_claim"] is False
250329

251330

252331
def test_unsupported_dataset_format_fails_closed_with_json(tmp_path: Path) -> None:

0 commit comments

Comments
 (0)