nemo_automodel.recipes.llm.train_dspark
nemo_automodel.recipes.llm.train_dspark
DSpark draft-model training recipe (Qwen3, Gemma4, DeepSeek V4, GLM-5.2, and MiniMax M3 VL targets).
DSpark is a semi-autoregressive parallel drafter: a parallel backbone produces a block of tokens per anchor in one pass, a serial Markov head injects intra-block dependency, and a confidence head predicts per-position acceptance. This recipe mirrors the EAGLE / DFlash scaffolding — online target hidden-state capture, gradient accumulation with a trailing-window flush, and the shared checkpointer plumbing — and trains the draft with the three-term DSpark objective.
Module Contents
Classes
Functions
Data
API
Bases: BaseRecipe
Recipe for DSpark draft-model training on Qwen3, Gemma4, DeepSeek V4, GLM-5.2, and MiniMax M3 VL targets.
Build the checkpointer using the same plumbing as the EAGLE / DFlash recipes.
Run one batch through live target capture or the offline cache.
Restore DSpark meta: global_step and epoch, and validate mask_token_id.
Log a saved checkpoint on rank 0 when checkpointing is enabled.
Precompute float8 dynamic scales after an optimizer step (FSDP2 fp8 all-gather only).
Always save the fully-trained model at the end, unless a cadence already saved the final step.
Save a checkpoint mid-epoch when ckpt_every_steps is configured.
Resolve and validate the MASK token id filling non-anchor block positions.
The draft’s embed_tokens row at this id is the learned “predict here”
signal. It must be a deliberately chosen reserved / unused token id (never a
silent fallback to pad, which is commonly aliased to eos), and the
inference runtime must fill block slots with the same id.
Evaluate the draft on the validation stream.
Reports the loss and the acceptance diagnostics that decide whether the
draft is worth serving: the per-position accept_rate@k, its aggregate,
the expected accepted block length tau, and the confidence head’s
calibration against the measured acceptance. Every batch already computes
these (DSparkStepMetrics); training reduces them over a log
window and validation over the whole split, both as unreduced
numerator/denominator sums so the ratio is formed once, after the
data-parallel reduction, rather than averaged over per-rank ratios.
Returns:
The metric dict, or None when no validation dataloader is configured.
Persist DSpark meta: global_step, epoch, block_size, mask, and target layers.
Whether to load a frozen dense target FSDP2-sharded via the standard distributed setup.
Opt-in (recipe_args.shard_dense_target, default False). A dense target
(Qwen3 / Gemma4) is otherwise loaded whole and replicated on every rank. For a large
dense target (e.g. Gemma4-31B) the frozen target is ~62 GiB, leaving no room for the
draft’s training activations, so training OOMs at the first backward on 80 GiB GPUs.
Loading it through create_distributed_setup_from_config +
NeMoAutoModelForCausalLM.from_pretrained(distributed_setup=...) FSDP2-shards it
across the mesh, the same path the MoE / VL targets already use.
A small target (e.g. Qwen3-0.6B) stays replicated by default, since sharding a target
that already fits is pure all-gather overhead. Requires distributed.strategy='fsdp2'
on more than one rank; otherwise the request is ignored with a warning and the target
stays replicated.
Raises:
ValueError: ifshard_dense_targetis requested together with a model-parallel or replication axis (tp_size/pp_size/cp_size/ep_size/dp_replicate_size> 1). DSpark’s forward-hook hidden-state capture needs one non-pipelinedmodel(...)call per rank (pp_size > 1builds anAutoPipelineinstead of a module), the other model-parallel axes are untested for the frozen dense target here, and HSDP replication re-replicates the target across the replicate dimension, defeating the sharding.
Log rank-zero metrics when a W&B run is active.
Restore the DSpark draft model, optimizer, scheduler, RNG, and global_step.
Run the DSpark training loop.
Persist the DSpark draft model, optimizer, scheduler, RNG, and meta.
Build the target model, DSpark draft, data, optimizer, and trainer module.
Metric sums accumulated between two log points, reduced in one collective.
The scalar sums and the two [block_size] per-position accept vectors are
concatenated into a single tensor by pack so one all-reduce covers the
whole window, and unpack turns the reduced tensor into the metrics to log.
The losses are window means of already normalized per-micro-batch values, so they
divide by the micro-batch count. The acceptance diagnostics accumulate as
(num, den) sums and divide once after the reduction, which gives the exact
global ratio regardless of per-rank token imbalance. A diagnostic whose denominator
is zero was not measured this window (e.g. an ablation without the confidence head)
and is omitted, so it shows no curve rather than a flat zero that reads like
collapsed acceptance.
Accumulate one micro-batch’s outputs.
Flatten the window into the 1-D tensor handed to the DP all-reduce.
Zero every sum, starting a new window.
Turn the DP-reduced pack tensor into the metrics to log.
Bases: dict
Dict with attribute access for the per-architecture draft-config builders.
Add measured per-position acceptance rates to a metrics dictionary.
Build the DSpark trainer’s optimizer from its optimizer: config.
Thin wrapper around build_optimizer so TrainDSparkRecipe.setup has a
single, unit-testable call site (build_optimizer itself needs no
distributed environment for a non-pipelined single-part model like the
DSpark draft, so this is testable with a plain CPU module).
Return only the multimodal keys present in batch, for generate_batch(**kwargs).
Empty for a text-only batch (Qwen3, Gemma4, or MiniMax M3 without
multimodal: true), so the generate_batch call is unchanged in that case.
Initialize the rank-zero W&B run for a DSpark training job, or return None.
Centralizes the is_main / block-presence / enable gating that
TrainDSparkRecipe.setup previously inlined, so it is unit-testable
without a distributed environment.
Sequence-packing metadata from a dataloader batch (empty dict when unpacked).
Normalize the recipe’s optimizer: config into a build_optimizer spec.
Reads an optional _target_ (a registry short name such as "fused_adam"
or a dotted import path, e.g. transformer_engine.pytorch.optimizers.FusedAdam)
plus whatever other fields the config carries — lr/betas/weight_decay
and any optimizer-specific kwargs (master_weights, master_weight_dtype,
exp_avg_dtype, exp_avg_sq_dtype, store_param_remainders, …) — and
returns the (target, kwargs) tuple that build_optimizer resolves via its
registry / dotted-import-path / OptimizerFromFactoryConfig escape hatch.
Absent an explicit _target_, this defaults to plain torch.optim.AdamW
with its prior betas/weight_decay defaults (matching the previous
hardcoded behavior, so existing DSpark configs are unaffected). Those two
AdamW-shaped defaults are only injected in that no-_target_ case: forcing
them onto an arbitrary explicit _target_ would break optimizers that do
not accept a betas kwarg (e.g. plain SGD).
Convert a wandb: config block into wandb.init kwargs, or None.
enable is the examples’ documentation-only opt-in flag (W&B logging is
opt-in: example configs ship the block with enable: false so users start
logging by flipping it to true instead of commenting the block in/out);
it is not a real wandb.init kwarg, so strip it before forwarding the rest
— passing it through raises TypeError: init() got an unexpected keyword argument 'enable'. Returns None when enable is explicitly False.
Return the LR warmup length in optimizer steps.
warmup_ratio * total_optim_steps collapses to a handful of steps (or fewer)
on short / small-dataset runs, dropping a freshly-initialized draft (random
attention layers, Markov head, confidence head) to near-peak LR within the
first few optimizer steps — a reliable way to trigger an early loss spike.
Floor the ratio-derived step count at min_warmup_steps unless the caller
explicitly opts out of warmup with warmup_ratio<=0 (e.g. the smoke config).
Validate that a DSpark offline cache matches the configured target/draft run.
Reject sequence-packing configs the DSpark path cannot honor (fail fast at setup).
Context parallelism shards the sequence and strips the block-causal mask packing
relies on, and a FlashAttention target packs documents from per-document
position_ids only at batch size 1.
Entrypoint for TrainDSparkRecipe.