nemo_automodel.components.speculative.dspark.target
nemo_automodel.components.speculative.dspark.target
Target-model wrapper for DSpark training (online hidden-state capture).
DSpark feeds the draft two things from the frozen target: the concatenation of a
configured set of decoder-layer hidden states (the draft fc context), and the
final post-norm hidden state (the input the target’s lm_head consumes, used
by the TV / confidence losses). Both are captured in a single forward pass via
forward hooks, mirroring the DFlash target wrapper.
Module Contents
Classes
Data
API
Target-model features needed by the DSpark trainer.
position_ids / seq_lens / doc_remaining are None off the
packing path and carry the (unshifted) packing metadata to the trainer on it.
Capture intermediate + final hidden states from a frozen HF causal LM.
A forward hook on decoder layer i captures hidden_states[i + 1] (the
HuggingFace output_hidden_states offset-1 convention); a hook on the final
norm captures the post-norm last hidden state. Models with a nonstandard
feature contract can expose get_dspark_target_feature_modules; the wrapper
then captures the first input of each returned module instead.
Collapse a 4D Hyper-Connection stream [B, S, hc_mult, H] to [B, S, H].
DeepSeek V4 decoder layers emit hc_mult parallel residual copies; only the
final-norm output is already collapsed. For an intermediate target-feature
layer we reduce the streams with their mean: a simple, in-distribution
reduction that the draft’s learnable fc then reprojects. We deliberately
avoid the model’s final hc_head here, since it is trained for the
last-layer stream distribution, not the intermediate ones. Non-HC targets
emit 3D states and pass through unchanged.
Return the final norm module whose output feeds lm_head.
Return decoder layers as an ordered, integer-indexable list.
Return the base transformer module that owns layers and norm.
Handles the common nestings: a plain causal LM (model.model), a
decoder-only base (model), and multimodal targets whose text stack is
under language_model (e.g. Gemma4: model.model.language_model).
Run the target model once and capture the DSpark context + last hidden state.
Features normally follow common.extract_context_feature: -1 is
the embedding output, the final layer is the post-norm hidden state, and
any other id is that decoder layer’s output. A model-owned feature provider
can replace those intermediate capture points. The final-norm output is
always returned separately for the TV and confidence losses.
Parameters:
Long tensor of shape [batch, sequence] containing target tokens.
Binary tensor of shape [batch, sequence], with one for valid tokens and zero for padding.
Tensor of shape [batch, sequence] selecting supervised tokens.
Optional long tensor of shape [batch, sequence] containing per-document reset positions for packed input.
Optional long tensor of shape [batch, max_documents] containing packed document lengths.
Optional long tensor of shape [batch, sequence] containing the valid tokens remaining in each position’s packed document.
Multimodal tensors accepted by the wrapped target. Their layouts follow that target model’s forward contract.
Returns: DSparkTargetBatch
DSparkTargetBatch with target_hidden_states of shape [batch, sequence,
Return the target model input embeddings.
Return the target model output embeddings (lm_head).
Optional model-owned selection of modules carrying DSpark features.
Return modules whose first forward input is each requested feature.