nemo_automodel.components.models.common.mtp.mtp

View as Markdown

Model-agnostic MTP scaffolding: depth iteration, token rolling, and loss.

Module Contents

Classes

NameDescription
MTPConfigRuntime configuration for the MTP block.
MTPModuleMulti-Token Prediction block.

Functions

NameDescription
get_mtp_loss_scaling_factorReturn the model’s configured MTP auxiliary-loss scaling factor.
roll_tensorRoll a tensor along dim by shifts and zero the wrapped slice.

API

class nemo_automodel.components.models.common.mtp.mtp.MTPConfig(
num_layers: int = 0,
layer_pattern: str = '',
loss_scaling_factor: float = 0.1,
use_repeated_layer: bool = False
)
Dataclass

Runtime configuration for the MTP block.

enabled
bool
layer_pattern
str = ''
loss_scaling_factor
float = 0.1
num_layers
int = 0
num_physical_depths
int
pattern_length
int
total_sublayers
int
use_repeated_layer
bool = False
class nemo_automodel.components.models.common.mtp.mtp.MTPModule(
mtp_config: nemo_automodel.components.models.common.mtp.mtp.MTPConfig,
block_types_per_sublayer: list[str],
sublayer_factory: typing.Callable[..., torch.nn.Module]
)

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:

mtp_config
MTPConfig

:class:MTPConfig describing depth and pattern.

block_types_per_sublayer
list[str]

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.

sublayer_factory
Callable[..., nn.Module]

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.

layers
= nn.ModuleList(layers)
num_depths
int
pattern_length
int
nemo_automodel.components.models.common.mtp.mtp.MTPModule.forward(
hidden_states: torch.Tensor,
input_ids: torch.LongTensor | None = None,
embed_fn: typing.Callable[[torch.LongTensor], torch.Tensor] | None = None,
embed_inputs: tuple[torch.Tensor, ...] | None = None,
position_ids: torch.LongTensor | None = None,
block_kwargs = {}
) -> list[torch.Tensor]

Iterate over MTP depths and return per-depth hidden states.

Two mutually-exclusive input modes:

  • Single-rank / first-stage PP (default): pass input_ids plus embed_fn. The module rolls input_ids cumulatively left by 1 per depth and applies embed_fn to produce the future-token embedding for that depth.
  • Final-stage PP / multimodal: pass embed_inputs (a tuple of pre-rolled per-depth embeddings, length num_depths). Used when the last PP stage no longer owns embed_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:

hidden_states
torch.Tensor

Output of the main model’s final norm (h_0); shape matches the model’s residual stream.

input_ids
torch.LongTensor | NoneDefaults to None

Token ids [B, S] (or [T] in THD). Rolled cumulatively left by 1 per depth. Mutually exclusive with embed_inputs.

embed_fn
Callable[[torch.LongTensor], torch.Tensor] | NoneDefaults to None

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.

embed_inputs
tuple[torch.Tensor, ...] | NoneDefaults to None

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
torch.LongTensor | NoneDefaults to None

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.

**block_kwargs
Defaults to {}

Forwarded to each sublayer’s __call__ (e.g. attention_mask).

Returns: list[torch.Tensor]

List of length num_depths containing the hidden state

nemo_automodel.components.models.common.mtp.mtp.get_mtp_loss_scaling_factor(
model: torch.nn.Module,
default: float = 0.1
) -> float

Return the model’s configured MTP auxiliary-loss scaling factor.

nemo_automodel.components.models.common.mtp.mtp.roll_tensor(
t: torch.Tensor,
shifts: int = -1,
dim: int = -1
) -> torch.Tensor

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:

t
torch.Tensor

Input tensor.

shifts
intDefaults to -1

Number of positions to shift (negative = left shift).

dim
intDefaults to -1

Dimension to roll along.

Returns: torch.Tensor

New tensor with the trailing |shifts| positions along dim