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
44from __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+
314424def _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 :
0 commit comments