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.
Globally shifted MTP tensors prepared before context-parallel sharding.
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.
Three 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. - Context parallel: pass
input_ids_per_depthplusembed_fn. Each tensor is globally shifted and then sharded into the local CP token layout, so this module embeds it directly without a rank-local roll. - 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);
tensor of shape [batch, sequence, hidden] or
[tokens, hidden] for THD.
Token ids of shape [batch, sequence] (or
[tokens] in THD). Rolled
cumulatively left by 1 per depth. Mutually exclusive with
input_ids_per_depth and embed_inputs.
Optional tuple of num_depths pre-shifted
token-ID tensors. Each has local CP shape
[batch, sequence] or [tokens] and is embedded directly
with embed_fn. Requires position_ids_per_depth.
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, each of shape
[batch, sequence, hidden] or [tokens, hidden].
Mutually exclusive with input_ids/
input_ids_per_depth/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.
Optional tuple of num_depths
pre-computed future-token position tensors. Each tensor has
shape [batch, sequence], [axes, batch, sequence] for
multi-axis RoPE, or [tokens] for THD. When supplied, these
tensors are forwarded directly instead of rolling rank-local
position_ids. Use this when context parallelism has already sharded the sequence.
Required with input_ids_per_depth and incompatible with
rank-local input_ids rolling.
Forwarded to each sublayer’s __call__ (e.g.
attention_mask).
Returns: list[torch.Tensor]
List of length num_depths containing hidden states of shape
Normalize supported packed-boundary metadata to token-aligned IDs.
Parameters:
Unsharded batch. Optional seq_idx or _packed_seq_ids tensors have shape [batch, sequence] (or [sequence] when batch is one). seq_lens_padded has shape [batch, num_sequences]. cu_seqlens_padded or cu_seqlens contains flattened cumulative boundaries of shape [num_sequences + 1] with optional negative sentinels.
Global token-ID tensor of shape [batch, sequence] whose materialized token layout defines the expected output shape.
Returns: torch.LongTensor | None
Sequence-ID tensor of shape [batch, sequence] on the input device,
Expand padded packed-sequence lengths into token-aligned sequence IDs.
Parameters:
Tensor of shape [batch, num_sequences] or [num_sequences] for a single batch row. Negative entries are unused sentinels; nonnegative entries are materialized sequence lengths including padding.
Expected batch dimension of the returned tensor.
Expected materialized token count in each batch row.
Device for the returned token-aligned IDs.
Returns: torch.LongTensor
Sequence-ID tensor of shape [batch, sequence] on device.
Return the model’s configured MTP auxiliary-loss scaling factor.
Prepare global future-token tensors before context-parallel sharding.
Each MTP depth is shifted in global sequence order before CP partitions the token axis. Packed boundaries are preserved, so no future token, position, or target crosses from one document into another. Missing or shared position IDs are materialized in batch so the main model and MTP heads are subsequently sharded from the same global source.
Parameters:
Mutable unsharded batch. input_ids and labels are tensors of shape [batch, sequence]. Optional position_ids has shape [batch, sequence], shared shape [1, sequence], or multi-axis RoPE shape [axes, batch, sequence]. Packed boundaries may use the tensor layouts documented by _packed_seq_ids_from_batch.
Number of MTP future-token depths; must be positive.
Fill value for invalid targets at trailing and packed boundary positions.
Returns: MTPContextParallelInputs
Per-depth input IDs, position IDs, targets, and validity masks. Token
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
Shift a token-aligned tensor left without crossing sequence boundaries.
Parameters:
Token-aligned tensor in global sequence order. Its batch and
sequence axes are selected by batch_dim and seq_dim.
Number of future-token positions to shift; must be positive.
Optional sequence IDs of shape [batch, sequence]. Tokens
whose shifted source has a different ID are filled.
Scalar used for trailing and cross-sequence positions.
Batch dimension in tensor.
Sequence dimension in tensor.
Returns: torch.Tensor
Tensor with the same shape, dtype, and device as tensor. The output