nemo_automodel.components.models.mimo_v25.model

View as Markdown

Module Contents

Classes

NameDescription
MiMoV2AttentionMiMoV2 hybrid attention (full or sliding-window).
MiMoV2DecoderLayer-
MiMoV2ForCausalLMNeMo AutoModel causal LM wrapper for MiMo-V2.5-Pro.
MiMoV2Model-
MiMoV2RotaryEmbedding-

Functions

Data

ModelClass

API

class nemo_automodel.components.models.mimo_v25.model.MiMoV2Attention(
is_swa: bool,
layer_idx: int,
projection_layout: str,
dtype: torch.dtype
)

Bases: Module

MiMoV2 hybrid attention (full or sliding-window).

attention_dropout
= getattr(config, 'attention_dropout', 0.0)
head_dim
k_proj
k_size
= self.num_key_value_heads * self.head_dim
num_attention_heads
num_key_value_groups
num_key_value_heads
o_hidden_size
= self.num_attention_heads * self.v_head_dim
o_proj
q_proj
q_size
= self.num_attention_heads * self.head_dim
qkv_proj
rope_dim
scaling
= self.head_dim ** -0.5
sliding_window
v_head_dim
v_proj
v_scale
= getattr(config, 'attention_value_scale', None)
v_size
= self.num_key_value_heads * self.v_head_dim
nemo_automodel.components.models.mimo_v25.model.MiMoV2Attention._forward_attention(
query_states: torch.Tensor,
key_states: torch.Tensor,
value_states: torch.Tensor,
input_shape: torch.Size,
position_embeddings: tuple[torch.Tensor, torch.Tensor],
attention_mask: torch.Tensor | None
) -> torch.Tensor
nemo_automodel.components.models.mimo_v25.model.MiMoV2Attention.forward(
hidden_states: torch.Tensor,
position_embeddings: tuple[torch.Tensor, torch.Tensor],
attention_mask: torch.Tensor | None,
kwargs: typing.Any = {}
) -> torch.Tensor
class nemo_automodel.components.models.mimo_v25.model.MiMoV2DecoderLayer(
layer_idx: int,
)

Bases: Module

attention_type
input_layernorm
mlp
post_attention_layernorm
self_attn
nemo_automodel.components.models.mimo_v25.model.MiMoV2DecoderLayer.forward(
hidden_states: torch.Tensor,
attention_mask: torch.Tensor | None = None,
position_embeddings: tuple[torch.Tensor, torch.Tensor],
padding_mask: torch.Tensor | None = None
) -> torch.Tensor
class nemo_automodel.components.models.mimo_v25.model.MiMoV2ForCausalLM(
kwargs = {}
)

Bases: HFCheckpointingMixin, Module, MoEFSDPSyncMixin

NeMo AutoModel causal LM wrapper for MiMo-V2.5-Pro.

_keep_in_fp32_modules_strict
backend
= backend or BackendConfig()
lm_head
model
state_dict_adapter
tie_word_embeddings_support
TieSupport = TieSupport.UNTIED_ONLY
nemo_automodel.components.models.mimo_v25.model.MiMoV2ForCausalLM.customize_pipeline_stage_modules(
module_names_per_stage: list[list[str]],
layers_prefix: str,
text_model: torch.nn.Module | None = None
) -> list[list[str]]

Keep the SWA rotary embedding on every PP stage.

nemo_automodel.components.models.mimo_v25.model.MiMoV2ForCausalLM.forward(
input_ids: torch.Tensor | None = None,
inputs_embeds: torch.FloatTensor | None = None,
position_ids: torch.LongTensor | None = None,
attention_mask: torch.Tensor | None = None,
padding_mask: torch.Tensor | None = None,
logits_to_keep: typing.Union[int, torch.Tensor] = 0,
output_hidden_states: bool | None = None,
kwargs: typing.Any = {}
) -> transformers.modeling_outputs.CausalLMOutputWithPast

Compute logits or pass hidden states to the next pipeline stage.

Parameters:

input_ids
torch.Tensor | NoneDefaults to None

Token IDs of shape [batch, sequence] on the first stage, or hidden states of shape [batch, sequence, hidden] thereafter.

inputs_embeds
torch.FloatTensor | NoneDefaults to None

Optional embeddings of shape [batch, sequence, hidden].

position_ids
torch.LongTensor | NoneDefaults to None

Optional positions of shape [batch, sequence] or [1, sequence] for broadcasting across the batch.

attention_mask
torch.Tensor | NoneDefaults to None

Optional padding mask of shape [batch, sequence] or additive attention mask of shape [batch, 1, sequence, sequence].

padding_mask
torch.Tensor | NoneDefaults to None

Optional padding indicators of shape [batch, sequence].

logits_to_keep
Union[int, torch.Tensor]Defaults to 0

Number of trailing token positions to retain, or indices of shape [retained_sequence]. Zero retains all positions.

output_hidden_states
bool | NoneDefaults to None

Whether to include the decoder hidden states.

**kwargs
AnyDefaults to {}

Additional decoder arguments, including optional cache_position of shape [sequence].

Returns: CausalLMOutputWithPast

A model output containing logits of shape [batch, retained_sequence,

classmethod
pretrained_model_name_or_path: str,
model_args = (),
kwargs = {}
classmethod
nemo_automodel.components.models.mimo_v25.model.MiMoV2ForCausalLM.get_input_embeddings() -> torch.nn.Embedding
nemo_automodel.components.models.mimo_v25.model.MiMoV2ForCausalLM.get_output_embeddings() -> torch.nn.Linear
nemo_automodel.components.models.mimo_v25.model.MiMoV2ForCausalLM.initialize_weights(
buffer_device: torch.device | None = None,
dtype: torch.dtype = torch.bfloat16
) -> None
nemo_automodel.components.models.mimo_v25.model.MiMoV2ForCausalLM.set_input_embeddings(
value: torch.nn.Embedding
) -> None
nemo_automodel.components.models.mimo_v25.model.MiMoV2ForCausalLM.set_output_embeddings(
new_embeddings: torch.nn.Linear
) -> None

Bases: Module

embed_tokens
has_sliding_layers
layers
norm
rotary_emb
= MiMoV2RotaryEmbedding(config=config, is_swa=False)
swa_rotary_emb
= MiMoV2RotaryEmbedding(config=config, is_swa=True)
nemo_automodel.components.models.mimo_v25.model.MiMoV2Model._build_causal_mask_mapping(
inputs_embeds: torch.Tensor,
attention_mask: torch.Tensor | dict[str, torch.Tensor] | None,
position_ids: torch.Tensor
) -> dict[str, torch.Tensor]

Build full and sliding masks for the uncached training sequence.

Parameters:

inputs_embeds
torch.Tensor

Embeddings of shape [batch, sequence, hidden].

attention_mask
torch.Tensor | dict[str, torch.Tensor] | None

Padding mask of shape [batch, sequence], an additive mask of shape [batch, 1, sequence, sequence], or a mapping of attention types to masks with that four-dimensional layout.

position_ids
torch.Tensor

Positions of shape [batch, sequence] or [1, sequence].

Returns: dict[str, torch.Tensor]

Full and sliding additive masks of shape [batch, 1, sequence, sequence].

nemo_automodel.components.models.mimo_v25.model.MiMoV2Model.forward(
input_ids: torch.Tensor | None = None,
inputs_embeds: torch.FloatTensor | None = None,
position_ids: torch.LongTensor | None = None,
attention_mask: torch.Tensor | None = None,
padding_mask: torch.Tensor | None = None,
cache_position: torch.LongTensor | None = None,
kwargs: typing.Any = {}
) -> torch.Tensor

Run the decoder layers owned by this pipeline stage.

Parameters:

input_ids
torch.Tensor | NoneDefaults to None

Token IDs of shape [batch, sequence] on the embedding stage, or hidden states of shape [batch, sequence, hidden] on later stages.

inputs_embeds
torch.FloatTensor | NoneDefaults to None

Optional embeddings of shape [batch, sequence, hidden].

position_ids
torch.LongTensor | NoneDefaults to None

Optional positions of shape [batch, sequence] or [1, sequence] for broadcasting across the batch.

attention_mask
torch.Tensor | NoneDefaults to None

Optional padding mask of shape [batch, sequence] or additive attention mask of shape [batch, 1, sequence, sequence].

padding_mask
torch.Tensor | NoneDefaults to None

Optional padding indicators of shape [batch, sequence].

cache_position
torch.LongTensor | NoneDefaults to None

Optional token positions of shape [sequence].

**kwargs
AnyDefaults to {}

Unused compatibility arguments.

Returns: torch.Tensor

Hidden states of shape [batch, sequence, hidden], normalized only

nemo_automodel.components.models.mimo_v25.model.MiMoV2Model.init_weights(
buffer_device: torch.device | None = None
) -> None
class nemo_automodel.components.models.mimo_v25.model.MiMoV2RotaryEmbedding(
is_swa: bool,
device: torch.device | None = None
)

Bases: Module

config
= copy(config)
inv_freq
Tensor
max_seq_len_cached
= config.max_position_embeddings
original_inv_freq
= self.inv_freq
original_max_seq_len
= config.max_position_embeddings
rope_init_fn
rope_type
nemo_automodel.components.models.mimo_v25.model.MiMoV2RotaryEmbedding.compute_default_rope_parameters(
device: torch.device | None = None,
seq_len: int | None = None,
layer_type: str | None = None
) -> tuple[torch.Tensor, float]
staticmethod
nemo_automodel.components.models.mimo_v25.model.MiMoV2RotaryEmbedding.forward(
x: torch.Tensor,
position_ids: torch.Tensor
) -> tuple[torch.Tensor, torch.Tensor]
nemo_automodel.components.models.mimo_v25.model._convert_bool_4d_mask_to_additive(
mask: torch.Tensor,
dtype: torch.dtype
) -> torch.Tensor
nemo_automodel.components.models.mimo_v25.model._derive_padding_mask(
attention_mask: torch.Tensor
) -> torch.Tensor
nemo_automodel.components.models.mimo_v25.model._ensure_additive_mask(
mask: torch.Tensor | None,
batch_size: int,
seq_len: int,
dtype: torch.dtype,
device: torch.device,
attention_mask: torch.Tensor | None,
sliding_window: int | None
) -> torch.Tensor
nemo_automodel.components.models.mimo_v25.model._fallback_additive_mask(
batch_size: int,
seq_len: int,
dtype: torch.dtype,
device: torch.device,
attention_mask: torch.Tensor | None = None,
sliding_window: int | None = None
) -> torch.Tensor
nemo_automodel.components.models.mimo_v25.model.apply_rotary_pos_emb(
q: torch.Tensor,
k: torch.Tensor,
cos: torch.Tensor,
sin: torch.Tensor,
position_ids: torch.Tensor | None = None,
unsqueeze_dim: int = 1
) -> tuple[torch.Tensor, torch.Tensor]
nemo_automodel.components.models.mimo_v25.model.eager_attention_forward(
module: torch.nn.Module,
query: torch.Tensor,
key: torch.Tensor,
value: torch.Tensor,
attention_mask: torch.Tensor | None,
scaling: float,
dropout: float = 0.0,
sinks: torch.Tensor | None = None,
kwargs: typing.Any = {}
) -> tuple[torch.Tensor, torch.Tensor]
nemo_automodel.components.models.mimo_v25.model.repeat_kv(
hidden_states: torch.Tensor,
n_rep: int
) -> torch.Tensor
nemo_automodel.components.models.mimo_v25.model.rotate_half(
x: torch.Tensor
) -> torch.Tensor
nemo_automodel.components.models.mimo_v25.model.ModelClass = MiMoV2ForCausalLM