nemo_automodel.components.models.common.mtp.mtp
nemo_automodel.components.models.common.mtp.mtp
Model-agnostic MTP scaffolding: depth iteration, token rolling, and loss.
Module Contents
Classes
Functions
API
Runtime configuration for the MTP block.
Bases: Module
Multi-Token Prediction block.
Holds a flat :class:nn.ModuleList of sublayers (length
num_physical_depths * pattern_length) where the first sublayer of
each physical depth carries the fusion modules (enorm, hnorm,
eh_proj) and the last sublayer of each physical depth carries
final_layernorm. This flat layout matches the HuggingFace export
format used by Nemotron-V3 (mtp.layers.{i}.*).
The model-specific sublayer construction (which decoder block to use, how
to handle MoE / attention / Mamba) is delegated to the caller via
sublayer_factory.
Parameters:
:class:MTPConfig describing depth and pattern.
List of block-type strings (one per inner
sublayer position), length must equal mtp_config.pattern_length.
Caller is responsible for parsing the model-specific symbol
convention; this module does not interpret symbols.
Callable
factory(global_idx, depth, sublayer_idx, block_type, has_fusion, has_final_norm) -> nn.Module
constructing one sublayer. The returned module must be callable
as sublayer(hidden_states, **kwargs) -> Tensor and, when
has_fusion=True, expose attributes enorm, hnorm,
eh_proj. When has_final_norm=True it must expose
final_layernorm.
Iterate over MTP depths and return per-depth hidden states.
Two mutually-exclusive input modes:
- Single-rank / first-stage PP (default): pass
input_idsplusembed_fn. The module rollsinput_idscumulatively left by 1 per depth and appliesembed_fnto produce the future-token embedding for that depth. - Final-stage PP / multimodal: pass
embed_inputs(a tuple of pre-rolled per-depth embeddings, lengthnum_depths). Used when the last PP stage no longer ownsembed_tokens, or for multimodal models (e.g. SALM) where some positions carry continuous audio embeddings that cannot be recovered by re-embedding an integer token id — the caller pre-rolls the fused embedding tensor and passes it here.
Parameters:
Output of the main model’s final norm (h_0);
shape matches the model’s residual stream.
Token ids [B, S] (or [T] in THD). Rolled
cumulatively left by 1 per depth. Mutually exclusive with
embed_inputs.
Callable applied to rolled input_ids to produce the
future-token embedding (typically the model’s input embedding
layer). Required when input_ids is supplied.
Optional tuple of num_depths pre-computed
future-token embeddings, one per depth in MTP order.
Mutually exclusive with input_ids/embed_fn.
Position ids matching input_ids. When supplied,
rolled cumulatively per depth in lockstep with input_ids
(so slot t carries the original position of the rolled
token) and forwarded to each sublayer via block_kwargs.
Required for RoPE-using sublayers; ignored by sublayers that
don’t consume it.
Forwarded to each sublayer’s __call__ (e.g.
attention_mask).
Returns: list[torch.Tensor]
List of length num_depths containing the hidden state
Return the model’s configured MTP auxiliary-loss scaling factor.
Roll a tensor along dim by shifts and zero the wrapped slice.
Used to shift input_ids / position_ids / labels left by one
position per MTP depth. Single-GPU path only (no CP / packed-sequence
handling).
Parameters:
Input tensor.
Number of positions to shift (negative = left shift).
Dimension to roll along.
Returns: torch.Tensor
New tensor with the trailing |shifts| positions along dim