nemo_automodel.recipes.llm.train_dspark

View as Markdown

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

NameDescription
TrainDSparkRecipeRecipe for DSpark draft-model training on Qwen3, Gemma4, DeepSeek V4, GLM-5.2, and MiniMax M3 VL targets.
_DSparkMetricWindowMetric sums accumulated between two log points, reduced in one collective.
_DraftArgsDict with attribute access for the per-architecture draft-config builders.

Functions

NameDescription
_add_accept_rate_per_positionAdd measured per-position acceptance rates to a metrics dictionary.
_build_dspark_optimizerBuild the DSpark trainer’s optimizer from its optimizer: config.
_extract_mm_kwargsReturn only the multimodal keys present in batch, for generate_batch(**kwargs).
_init_dspark_wandbInitialize the rank-zero W&B run for a DSpark training job, or return None.
_packing_kwargsSequence-packing metadata from a dataloader batch (empty dict when unpacked).
_resolve_dspark_optimizer_specNormalize the recipe’s optimizer: config into a build_optimizer spec.
_resolve_wandb_kwargsConvert a wandb: config block into wandb.init kwargs, or None.
_resolve_warmup_stepsReturn the LR warmup length in optimizer steps.
_validate_cached_dspark_manifestValidate that a DSpark offline cache matches the configured target/draft run.
_validate_packing_gatesReject sequence-packing configs the DSpark path cannot honor (fail fast at setup).
mainEntrypoint for TrainDSparkRecipe.

Data

_DSPARK_MM_KEYS

_DSPARK_WINDOW_SCALARS

logger

API

class nemo_automodel.recipes.llm.train_dspark.TrainDSparkRecipe(
cfg
)

Bases: BaseRecipe

Recipe for DSpark draft-model training on Qwen3, Gemma4, DeepSeek V4, GLM-5.2, and MiniMax M3 VL targets.

nemo_automodel.recipes.llm.train_dspark.TrainDSparkRecipe._build_checkpointer(
target_path: str
) -> None

Build the checkpointer using the same plumbing as the EAGLE / DFlash recipes.

nemo_automodel.recipes.llm.train_dspark.TrainDSparkRecipe._finish_wandb() -> None
nemo_automodel.recipes.llm.train_dspark.TrainDSparkRecipe._forward_batch(
batch
)

Run one batch through live target capture or the offline cache.

nemo_automodel.recipes.llm.train_dspark.TrainDSparkRecipe._load_extra_state(
ckpt_dir: str
) -> None

Restore DSpark meta: global_step and epoch, and validate mask_token_id.

nemo_automodel.recipes.llm.train_dspark.TrainDSparkRecipe._log_saved_checkpoint(
kind: str,
epoch: int,
step: int
) -> None

Log a saved checkpoint on rank 0 when checkpointing is enabled.

nemo_automodel.recipes.llm.train_dspark.TrainDSparkRecipe._maybe_precompute_fp8_scales() -> None

Precompute float8 dynamic scales after an optimizer step (FSDP2 fp8 all-gather only).

nemo_automodel.recipes.llm.train_dspark.TrainDSparkRecipe._maybe_save_final_checkpoint(
completed_epochs: int
) -> bool

Always save the fully-trained model at the end, unless a cadence already saved the final step.

nemo_automodel.recipes.llm.train_dspark.TrainDSparkRecipe._maybe_save_step_checkpoint(
epoch: int
) -> bool

Save a checkpoint mid-epoch when ckpt_every_steps is configured.

nemo_automodel.recipes.llm.train_dspark.TrainDSparkRecipe._module()
nemo_automodel.recipes.llm.train_dspark.TrainDSparkRecipe._resolve_mask_token_id(
recipe_cfg,
vocab_size: int
) -> int
staticmethod

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.

nemo_automodel.recipes.llm.train_dspark.TrainDSparkRecipe._run_eval()

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.

nemo_automodel.recipes.llm.train_dspark.TrainDSparkRecipe._save_extra_state(
path: str,
epoch: int
) -> None

Persist DSpark meta: global_step, epoch, block_size, mask, and target layers.

nemo_automodel.recipes.llm.train_dspark.TrainDSparkRecipe._should_shard_dense_target(
recipe_cfg
) -> bool

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: if shard_dense_target is 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-pipelined model(...) call per rank (pp_size > 1 builds an AutoPipeline instead 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.
nemo_automodel.recipes.llm.train_dspark.TrainDSparkRecipe._wandb_log(
data: dict,
step: int
) -> None

Log rank-zero metrics when a W&B run is active.

nemo_automodel.recipes.llm.train_dspark.TrainDSparkRecipe.load_checkpoint(
restore_from: str | None = None
) -> None

Restore the DSpark draft model, optimizer, scheduler, RNG, and global_step.

nemo_automodel.recipes.llm.train_dspark.TrainDSparkRecipe.run_train_validation_loop()

Run the DSpark training loop.

nemo_automodel.recipes.llm.train_dspark.TrainDSparkRecipe.save_checkpoint(
epoch: int,
step: int,
train_loss: float | None = None,
val_loss: dict[str, float] | None = None,
best_metric_key: str = 'default',
is_final_checkpoint: bool = False
) -> None

Persist the DSpark draft model, optimizer, scheduler, RNG, and meta.

nemo_automodel.recipes.llm.train_dspark.TrainDSparkRecipe.setup()

Build the target model, DSpark draft, data, optimizer, and trainer module.

class nemo_automodel.recipes.llm.train_dspark._DSparkMetricWindow(
block_size: int,
device: torch.device | None = None,
loss: float = 0.0,
ce_loss: float = 0.0,
l1_loss: float = 0.0,
confidence_loss: float = 0.0,
tau_num: float = 0.0,
tau_den: float = 0.0,
confidence_abs_error_num: float = 0.0,
confidence_bias_num: float = 0.0,
confidence_cumprod_bias_num: float = 0.0,
confidence_diag_den: float = 0.0,
num_micro_batches: float = 0.0
)
Dataclass

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.

accept_den
Tensor = field(init=False)
accept_num
Tensor = field(init=False)
block_size
int
ce_loss
float = 0.0
confidence_abs_error_num
float = 0.0
confidence_bias_num
float = 0.0
confidence_cumprod_bias_num
float = 0.0
confidence_diag_den
float = 0.0
confidence_loss
float = 0.0
device
device | None = None
l1_loss
float = 0.0
loss
float = 0.0
num_micro_batches
float = 0.0
tau_den
float = 0.0
tau_num
float = 0.0
nemo_automodel.recipes.llm.train_dspark._DSparkMetricWindow.__post_init__() -> None
nemo_automodel.recipes.llm.train_dspark._DSparkMetricWindow.add(
) -> None

Accumulate one micro-batch’s outputs.

nemo_automodel.recipes.llm.train_dspark._DSparkMetricWindow.pack() -> torch.Tensor

Flatten the window into the 1-D tensor handed to the DP all-reduce.

nemo_automodel.recipes.llm.train_dspark._DSparkMetricWindow.reset() -> None

Zero every sum, starting a new window.

nemo_automodel.recipes.llm.train_dspark._DSparkMetricWindow.unpack(
reduced: torch.Tensor
) -> dict[str, float]

Turn the DP-reduced pack tensor into the metrics to log.

class nemo_automodel.recipes.llm.train_dspark._DraftArgs()

Bases: dict

Dict with attribute access for the per-architecture draft-config builders.

nemo_automodel.recipes.llm.train_dspark._DraftArgs.__getattr__(
key
)
nemo_automodel.recipes.llm.train_dspark._add_accept_rate_per_position(
metrics: dict[str, float],
accept_num: torch.Tensor,
accept_den: torch.Tensor
) -> None

Add measured per-position acceptance rates to a metrics dictionary.

nemo_automodel.recipes.llm.train_dspark._build_dspark_optimizer(
trainer_module,
opt_cfg,
device_mesh = None
) -> torch.optim.Optimizer

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).

nemo_automodel.recipes.llm.train_dspark._extract_mm_kwargs(
batch: dict
) -> dict

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.

nemo_automodel.recipes.llm.train_dspark._init_dspark_wandb(
is_main: bool,
wandb_cfg,
cfg_dict: dict,
default_name: str
)

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.

nemo_automodel.recipes.llm.train_dspark._packing_kwargs(
batch: dict
) -> dict

Sequence-packing metadata from a dataloader batch (empty dict when unpacked).

nemo_automodel.recipes.llm.train_dspark._resolve_dspark_optimizer_spec(
opt_cfg
) -> tuple[str, dict]

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).

nemo_automodel.recipes.llm.train_dspark._resolve_wandb_kwargs(
wandb_cfg: dict
) -> dict | None

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.

nemo_automodel.recipes.llm.train_dspark._resolve_warmup_steps(
warmup_ratio: float,
total_optim_steps: int,
min_warmup_steps: int = 20
) -> int

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).

nemo_automodel.recipes.llm.train_dspark._validate_cached_dspark_manifest(
cache_dir: str,
manifest: dict,
target_config,
target_layer_ids: list[int],
target_model: str,
target_model_type: str,
seq_length: int,
compute_dtype: torch.dtype
) -> None

Validate that a DSpark offline cache matches the configured target/draft run.

nemo_automodel.recipes.llm.train_dspark._validate_packing_gates(
cp_size: int,
target_attn_impl: str,
micro_batch_size: int
) -> None

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.

nemo_automodel.recipes.llm.train_dspark.main(
config_path: str | None = None
)

Entrypoint for TrainDSparkRecipe.

nemo_automodel.recipes.llm.train_dspark._DSPARK_MM_KEYS = tuple(k for k in VLM_INPUT_KEYS if k != 'input_ids')
nemo_automodel.recipes.llm.train_dspark._DSPARK_WINDOW_SCALARS = ('loss', 'ce_loss', 'l1_loss', 'confidence_loss', 'tau_num', 'tau_den', 'confide...
nemo_automodel.recipes.llm.train_dspark.logger = logging.getLogger(__name__)