nemo_automodel.components.speculative.dspark.target

View as Markdown

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

NameDescription
DSparkTargetBatchTarget-model features needed by the DSpark trainer.
HFDSparkTargetModelCapture intermediate + final hidden states from a frozen HF causal LM.
_DSparkTargetFeatureProviderOptional model-owned selection of modules carrying DSpark features.

Data

__all__

API

class nemo_automodel.components.speculative.dspark.target.DSparkTargetBatch(
target_hidden_states: torch.Tensor,
target_last_hidden_states: torch.Tensor,
input_ids: torch.Tensor,
loss_mask: torch.Tensor,
position_ids: torch.Tensor | None = None,
seq_lens: torch.Tensor | None = None,
doc_remaining: torch.Tensor | None = None
)
Dataclass

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.

doc_remaining
Tensor | None = None
input_ids
Tensor
loss_mask
Tensor
position_ids
Tensor | None = None
seq_lens
Tensor | None = None
target_hidden_states
Tensor
target_last_hidden_states
Tensor
class nemo_automodel.components.speculative.dspark.target.HFDSparkTargetModel(
model: torch.nn.Module,
target_layer_ids: typing.Sequence[int],
cp_mesh = None
)

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.

_cp_size
= cp_mesh.size() if cp_mesh is not None else 1
_num_layers
= len(self._get_transformer_layers())
_target_feature_modules
model
= model.eval()
target_layer_ids
nemo_automodel.components.speculative.dspark.target.HFDSparkTargetModel._collapse_hc_streams(
tensor: torch.Tensor
) -> torch.Tensor

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.

nemo_automodel.components.speculative.dspark.target.HFDSparkTargetModel._get_final_norm() -> torch.nn.Module

Return the final norm module whose output feeds lm_head.

nemo_automodel.components.speculative.dspark.target.HFDSparkTargetModel._get_transformer_layers() -> list[torch.nn.Module]

Return decoder layers as an ordered, integer-indexable list.

nemo_automodel.components.speculative.dspark.target.HFDSparkTargetModel._inner_model() -> torch.nn.Module

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

nemo_automodel.components.speculative.dspark.target.HFDSparkTargetModel.generate_batch(
input_ids: torch.Tensor,
attention_mask: torch.Tensor,
loss_mask: torch.Tensor,
position_ids: torch.Tensor | None = None,
seq_lens: torch.Tensor | None = None,
doc_remaining: torch.Tensor | None = None,
mm_kwargs: torch.Tensor = {}

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:

input_ids
torch.Tensor

Long tensor of shape [batch, sequence] containing target tokens.

attention_mask
torch.Tensor

Binary tensor of shape [batch, sequence], with one for valid tokens and zero for padding.

loss_mask
torch.Tensor

Tensor of shape [batch, sequence] selecting supervised tokens.

position_ids
torch.Tensor | NoneDefaults to None

Optional long tensor of shape [batch, sequence] containing per-document reset positions for packed input.

seq_lens
torch.Tensor | NoneDefaults to None

Optional long tensor of shape [batch, max_documents] containing packed document lengths.

doc_remaining
torch.Tensor | NoneDefaults to None

Optional long tensor of shape [batch, sequence] containing the valid tokens remaining in each position’s packed document.

**mm_kwargs
torch.TensorDefaults to {}

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,

nemo_automodel.components.speculative.dspark.target.HFDSparkTargetModel.get_input_embeddings() -> torch.nn.Embedding

Return the target model input embeddings.

nemo_automodel.components.speculative.dspark.target.HFDSparkTargetModel.get_output_embeddings() -> torch.nn.Module

Return the target model output embeddings (lm_head).

class nemo_automodel.components.speculative.dspark.target._DSparkTargetFeatureProvider()
Protocol

Optional model-owned selection of modules carrying DSpark features.

nemo_automodel.components.speculative.dspark.target._DSparkTargetFeatureProvider.get_dspark_target_feature_modules(
layer_ids: list[int]
) -> tuple[torch.nn.Module, ...]

Return modules whose first forward input is each requested feature.

nemo_automodel.components.speculative.dspark.target.__all__ = ['HFDSparkTargetModel', 'DSparkTargetBatch']