Skip to content

vllm.model_executor.models.mimo_v2

Classes:

MiMoV2Model

Bases: Module, EagleModelMixin

Source code in vllm/model_executor/models/mimo_v2.py
@support_torch_compile
class MiMoV2Model(nn.Module, EagleModelMixin):
    def __init__(self, *, vllm_config: VllmConfig, prefix: str = ""):
        super().__init__()

        config = vllm_config.model_config.hf_config.get_text_config()
        quant_config = vllm_config.quant_config
        eplb_config = vllm_config.parallel_config.eplb_config

        self.config = config
        self.quant_config = quant_config
        self.vocab_size = config.vocab_size
        self.num_redundant_experts = eplb_config.num_redundant_experts

        if get_pp_group().is_first_rank or (
            config.tie_word_embeddings and get_pp_group().is_last_rank
        ):
            self.embed_tokens = VocabParallelEmbedding(
                config.vocab_size,
                config.hidden_size,
                quant_config=quant_config,
                prefix=f"{prefix}.embed_tokens",
            )
        else:
            self.embed_tokens = PPMissingLayer()

        self.start_layer, self.end_layer, self.layers = make_layers(
            config.num_hidden_layers,
            lambda prefix: MiMoV2FlashDecoderLayer(
                vllm_config=vllm_config,
                prefix=prefix,
            ),
            prefix=f"{prefix}.layers",
        )

        self.make_empty_intermediate_tensors = make_empty_intermediate_tensors_factory(
            ["hidden_states", "residual"], config.hidden_size
        )
        if get_pp_group().is_last_rank:
            self.norm = RMSNorm(config.hidden_size, eps=config.layernorm_epsilon)
        else:
            self.norm = PPMissingLayer()

    def embed_input_ids(self, input_ids: torch.Tensor) -> torch.Tensor:
        return self.embed_tokens(input_ids)

    def forward(
        self,
        input_ids: torch.Tensor | None,
        positions: torch.Tensor,
        intermediate_tensors: IntermediateTensors | None = None,
        inputs_embeds: torch.Tensor | None = None,
    ) -> torch.Tensor | IntermediateTensors:
        if get_pp_group().is_first_rank:
            if inputs_embeds is not None:
                hidden_states = inputs_embeds
            else:
                hidden_states = self.embed_input_ids(input_ids)
            residual = None
        else:
            assert intermediate_tensors is not None
            hidden_states = intermediate_tensors["hidden_states"]
            residual = intermediate_tensors["residual"]

        aux_hidden_states = self._maybe_add_hidden_state(
            [], self.start_layer, hidden_states, residual
        )
        for idx, layer in enumerate(
            islice(self.layers, self.start_layer, self.end_layer)
        ):
            hidden_states, residual = layer(positions, hidden_states, residual)
            self._maybe_add_hidden_state(
                aux_hidden_states, idx + 1, hidden_states, residual
            )

        if not get_pp_group().is_last_rank:
            return IntermediateTensors(
                {"hidden_states": hidden_states, "residual": residual}
            )

        hidden_states, _ = self.norm(hidden_states, residual)

        if len(aux_hidden_states) > 0:
            return hidden_states, aux_hidden_states
        return hidden_states

    def get_expert_mapping(self) -> list[tuple[str, str, int, str]]:
        # Params for weights, fp8 weight scales, fp8 activation scales
        # (param_name, weight_name, expert_id, shard_id)
        return fused_moe_make_expert_params_mapping(
            self,
            ckpt_gate_proj_name="gate_proj",
            ckpt_down_proj_name="down_proj",
            ckpt_up_proj_name="up_proj",
            num_experts=self.config.n_routed_experts,
            num_redundant_experts=self.num_redundant_experts,
        )

    def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]:
        stacked_params_mapping: list[tuple[str, str, str | int]] = [
            # (param_name, shard_name, shard_id)
            ("qkv_proj", "q_proj", "q"),
            ("qkv_proj", "k_proj", "k"),
            ("qkv_proj", "v_proj", "v"),
            ("gate_up_proj", "gate_proj", 0),
            ("gate_up_proj", "up_proj", 1),
        ]

        tp_rank = get_tensor_model_parallel_rank()
        tp_size = get_tensor_model_parallel_world_size()

        params_dict = dict(self.named_parameters(remove_duplicate=False))
        loaded_params: set[str] = set()
        expert_params_mapping = self.get_expert_mapping()
        # Pro-format fused qkv_proj arrives as two tensors (weight and
        # weight_scale_inv). Store them per-layer so that they can be
        # sharded together. The state must outlive this call: AutoWeightsLoader
        # delegates per contiguous group of names, so a pair can straddle two
        # calls and would otherwise be dropped silently.
        pending_fp8_qkv_proj = getattr(self, "_pending_fp8_qkv_proj", None)
        if pending_fp8_qkv_proj is None:
            self._pending_fp8_qkv_proj = pending_fp8_qkv_proj = {}
        for name, loaded_weight in weights:
            if "rotary_emb.inv_freq" in name:
                continue
            if "rotary_emb.cos_cached" in name or "rotary_emb.sin_cached" in name:
                continue
            if "mtp" in name:
                continue

            expert_matched = False
            for param_name, weight_name, expert_id, shard_id in expert_params_mapping:
                if weight_name not in name:
                    continue

                name_rewritten = name.replace(weight_name, param_name)

                if is_pp_missing_parameter(name_rewritten, self):
                    continue

                if (
                    name_rewritten.endswith(".bias") or name_rewritten.endswith("_bias")
                ) and name_rewritten not in params_dict:
                    continue

                if name_rewritten not in params_dict:
                    continue

                param = params_dict[name_rewritten]
                weight_loader = param.weight_loader

                weight_loader(
                    param,
                    loaded_weight,
                    name_rewritten,
                    shard_id=shard_id,
                    expert_id=expert_id,
                )
                loaded_params.add(name_rewritten)
                expert_matched = True
                break

            if expert_matched:
                continue
            # Support fused qkv_proj checkpoint (Pro format)
            if self._try_load_fp8_qkv_proj(
                name,
                loaded_weight,
                pending_fp8_qkv_proj,
                params_dict,
                loaded_params,
                tp_rank,
                tp_size,
            ):
                continue
            stacked_matched = False
            for param_name, weight_name, stacked_shard_id in stacked_params_mapping:
                if weight_name not in name:
                    continue
                name_rewritten = name.replace(weight_name, param_name)

                if (
                    name_rewritten.endswith(".bias")
                    and name_rewritten not in params_dict
                ):
                    continue

                if is_pp_missing_parameter(name_rewritten, self):
                    continue

                if name_rewritten not in params_dict:
                    continue

                param = params_dict[name_rewritten]
                weight_loader = getattr(param, "weight_loader", default_weight_loader)
                weight_loader(param, loaded_weight, stacked_shard_id)
                loaded_params.add(name_rewritten)

                stacked_matched = True
                break

            if stacked_matched:
                continue

            if name.endswith(".bias") and name not in params_dict:
                continue

            orig_name = name
            mapped_name = maybe_remap_kv_scale_name(name, params_dict)
            name = mapped_name if mapped_name is not None else orig_name

            if name not in params_dict:
                continue

            param = params_dict[name]

            if "attention_sink_bias" in name:
                total_heads = loaded_weight.shape[0]
                heads_per_rank = total_heads // tp_size
                head_start = tp_rank * heads_per_rank
                narrow_weight = loaded_weight.narrow(0, head_start, heads_per_rank)

                param.data.copy_(narrow_weight)
                loaded_params.add(name)
            else:
                weight_loader = getattr(param, "weight_loader", default_weight_loader)
                weight_loader(param, loaded_weight)
                loaded_params.add(name)

        return loaded_params

    def _try_load_fp8_qkv_proj(
        self,
        name: str,
        tensor: torch.Tensor,
        fp8_qkv_proj_dict: dict[str, dict[str, torch.Tensor]],
        params_dict: dict[str, torch.nn.Parameter],
        loaded_params: set[str],
        tp_rank: int,
        tp_size: int,
    ) -> bool:
        """The fused fp8 QKV projection weights and scale are stored separately.
        Special care must be taken while sharding these tensors across TP ranks.
        See _shard_fp8_qkv_proj for more details.

        Returns:
            True if ``tensor`` was an fp8 qkv_proj weight/scale and was consumed
            (caller should skip it); False otherwise, so the caller falls
            through to its normal loading path.

        """
        is_weight = (
            name.endswith("qkv_proj.weight") and tensor.dtype == torch.float8_e4m3fn
        )
        is_scale = name.endswith("qkv_proj.weight_scale_inv")
        if not is_weight and not is_scale:
            # Weight is not in FP8 format. Ignore.
            return False

        if is_pp_missing_parameter(name, self):
            # This qkv_proj is for a layer not on this PP rank.
            return True

        prefix, qkv_kind = name.rsplit(".", 1)
        entry = fp8_qkv_proj_dict.setdefault(prefix, {})
        entry[qkv_kind] = tensor
        if "weight" not in entry or "weight_scale_inv" not in entry:
            # Still waiting for the other param.
            return True
        del fp8_qkv_proj_dict[prefix]

        # Get self_attn module, which is a parent of qkv_proj.
        attn = self.get_submodule(prefix.rsplit(".", 1)[0])

        # Shard the qkv_proj per-rank.
        w_rank, s_rank = _shard_fp8_qkv_proj(
            entry["weight"],
            entry["weight_scale_inv"],
            num_heads=attn.total_num_heads,
            num_kv_heads=attn.total_num_kv_heads,
            head_dim=attn.head_dim,
            v_head_dim=attn.v_head_dim,
            tp_rank=tp_rank,
            tp_size=tp_size,
            # The fused qkv_proj is pre-sharded for this many ranks.
            ckpt_tp=self.config.num_key_value_heads,
        )
        sharded = {"weight": w_rank, "weight_scale_inv": s_rank}
        for kind, tensor in sharded.items():
            param_name = f"{prefix}.{kind}"
            param = params_dict[param_name]
            if tensor.shape[0] > param.shape[0]:
                tensor = tensor[: param.shape[0]]
            default_weight_loader(param, tensor)
            loaded_params.add(param_name)
        return True

_try_load_fp8_qkv_proj(name, tensor, fp8_qkv_proj_dict, params_dict, loaded_params, tp_rank, tp_size)

The fused fp8 QKV projection weights and scale are stored separately. Special care must be taken while sharding these tensors across TP ranks. See _shard_fp8_qkv_proj for more details.

Returns:

  • bool –

    True if tensor was an fp8 qkv_proj weight/scale and was consumed

  • bool –

    (caller should skip it); False otherwise, so the caller falls

  • bool –

    through to its normal loading path.

Source code in vllm/model_executor/models/mimo_v2.py
def _try_load_fp8_qkv_proj(
    self,
    name: str,
    tensor: torch.Tensor,
    fp8_qkv_proj_dict: dict[str, dict[str, torch.Tensor]],
    params_dict: dict[str, torch.nn.Parameter],
    loaded_params: set[str],
    tp_rank: int,
    tp_size: int,
) -> bool:
    """The fused fp8 QKV projection weights and scale are stored separately.
    Special care must be taken while sharding these tensors across TP ranks.
    See _shard_fp8_qkv_proj for more details.

    Returns:
        True if ``tensor`` was an fp8 qkv_proj weight/scale and was consumed
        (caller should skip it); False otherwise, so the caller falls
        through to its normal loading path.

    """
    is_weight = (
        name.endswith("qkv_proj.weight") and tensor.dtype == torch.float8_e4m3fn
    )
    is_scale = name.endswith("qkv_proj.weight_scale_inv")
    if not is_weight and not is_scale:
        # Weight is not in FP8 format. Ignore.
        return False

    if is_pp_missing_parameter(name, self):
        # This qkv_proj is for a layer not on this PP rank.
        return True

    prefix, qkv_kind = name.rsplit(".", 1)
    entry = fp8_qkv_proj_dict.setdefault(prefix, {})
    entry[qkv_kind] = tensor
    if "weight" not in entry or "weight_scale_inv" not in entry:
        # Still waiting for the other param.
        return True
    del fp8_qkv_proj_dict[prefix]

    # Get self_attn module, which is a parent of qkv_proj.
    attn = self.get_submodule(prefix.rsplit(".", 1)[0])

    # Shard the qkv_proj per-rank.
    w_rank, s_rank = _shard_fp8_qkv_proj(
        entry["weight"],
        entry["weight_scale_inv"],
        num_heads=attn.total_num_heads,
        num_kv_heads=attn.total_num_kv_heads,
        head_dim=attn.head_dim,
        v_head_dim=attn.v_head_dim,
        tp_rank=tp_rank,
        tp_size=tp_size,
        # The fused qkv_proj is pre-sharded for this many ranks.
        ckpt_tp=self.config.num_key_value_heads,
    )
    sharded = {"weight": w_rank, "weight_scale_inv": s_rank}
    for kind, tensor in sharded.items():
        param_name = f"{prefix}.{kind}"
        param = params_dict[param_name]
        if tensor.shape[0] > param.shape[0]:
            tensor = tensor[: param.shape[0]]
        default_weight_loader(param, tensor)
        loaded_params.add(param_name)
    return True

_requantize_fp8(grouped, rows_rank, block, dtype)

Block-quantize a rank's [Q | K | V] rows back to fp8.

A rank's rows need not end on a block boundary (a single 1856-row slice is 14.5 blocks) while scaled_quantize requires both dims to be multiples of the block size, so pad the tail with zeros: zero rows cannot raise a block's amax, so every scale -- and the number of scale rows -- is unchanged. The padding is dropped again here.

Source code in vllm/model_executor/models/mimo_v2.py
def _requantize_fp8(
    grouped: torch.Tensor, rows_rank: int, block: int, dtype: torch.dtype
) -> tuple[torch.Tensor, torch.Tensor]:
    """Block-quantize a rank's ``[Q | K | V]`` rows back to fp8.

    A rank's rows need not end on a block boundary (a single 1856-row slice is
    14.5 blocks) while ``scaled_quantize`` requires both dims to be multiples of
    the block size, so pad the tail with zeros: zero rows cannot raise a block's
    amax, so every scale -- and the number of scale rows -- is unchanged. The
    padding is dropped again here.
    """
    padded = cdiv(rows_rank, block) * block
    if padded != rows_rank:
        grouped = torch.cat(
            [grouped, grouped.new_zeros(padded - rows_rank, grouped.shape[1])], dim=0
        )
    w_rank, s_rank = scaled_quantize(
        grouped, GroupShape(block, block), dtype, compute_dtype=torch.float32
    )
    return w_rank[:rows_rank], s_rank

_shard_fp8_qkv_proj(w_full, s_full, num_heads, num_kv_heads, head_dim, v_head_dim, tp_rank, tp_size, ckpt_tp, block=128)

Shard the fp8 qkv_proj weights for tp_rank.

The checkpoint stores the fused QKV pre-sharded for ckpt_tp ranks (the model config's num_key_value_heads), each chunk holding that slice's Q, K and V rows:

[Q_0 | K_0 | V_0 | Q_1 | K_1 | V_1 | ... | Q_n | K_n | V_n]   (n = ckpt_tp)

Per chunk, Q has (num_heads / ckpt_tp) * head_dim rows, K has (num_kv_heads / ckpt_tp) * head_dim rows, and V has (num_kv_heads / ckpt_tp) * v_head_dim rows, and the fp8 block scales are tiled per chunk too (ceil(rows_per_chunk / block) rows each).

ckpt_tp is not the layer's KV-head count: a MiMo-V2.5 SWA layer has 8 KV heads over 4 chunks, so each chunk carries two KV heads (3712 rows = 29 blocks) while a GA layer has 4 KV heads over 4 chunks (3392 rows each).

The forward expects each rank's slice de-interleaved:

[Q_1 | Q_2 | ... | Q_g | K_1 | K_2 | ... | K_g | V_1 | V_2 | ... | V_g]

When tp_size == ckpt_tp the checkpoint chunk is that layout, so a plain chunk of both weight and scale suffices. Otherwise each rank's Q, K and V rows are gathered from the chunks that hold them, dequantized with the chunk's own scales, reordered, and re-quantized to fp8.

Source code in vllm/model_executor/models/mimo_v2.py
def _shard_fp8_qkv_proj(
    w_full: torch.Tensor,
    s_full: torch.Tensor,
    num_heads: int,
    num_kv_heads: int,
    head_dim: int,
    v_head_dim: int,
    tp_rank: int,
    tp_size: int,
    ckpt_tp: int,
    block: int = 128,
) -> tuple[torch.Tensor, torch.Tensor]:
    """Shard the fp8 qkv_proj weights for ``tp_rank``.

    The checkpoint stores the fused QKV pre-sharded for ``ckpt_tp`` ranks (the
    model config's ``num_key_value_heads``), each chunk holding that slice's
    Q, K and V rows:

        [Q_0 | K_0 | V_0 | Q_1 | K_1 | V_1 | ... | Q_n | K_n | V_n]   (n = ckpt_tp)

    Per chunk, Q has ``(num_heads / ckpt_tp) * head_dim`` rows, K has
    ``(num_kv_heads / ckpt_tp) * head_dim`` rows, and V has
    ``(num_kv_heads / ckpt_tp) * v_head_dim`` rows, and the fp8 block scales
    are tiled per chunk too (``ceil(rows_per_chunk / block)`` rows each).

    ``ckpt_tp`` is not the layer's KV-head count: a MiMo-V2.5 SWA layer has 8
    KV heads over 4 chunks, so each chunk carries two KV heads (3712 rows = 29
    blocks) while a GA layer has 4 KV heads over 4 chunks (3392 rows each).

    The forward expects each rank's slice de-interleaved:

        [Q_1 | Q_2 | ... | Q_g | K_1 | K_2 | ... | K_g | V_1 | V_2 | ... | V_g]

    When ``tp_size == ckpt_tp`` the checkpoint chunk *is* that layout, so a
    plain chunk of both weight and scale suffices. Otherwise each rank's Q, K
    and V rows are gathered from the chunks that hold them, dequantized with
    the chunk's own scales, reordered, and re-quantized to fp8.
    """
    assert num_heads % tp_size == 0, (
        f"num_heads={num_heads} must be divisible by tp_size={tp_size}."
    )
    if ckpt_tp <= 0 or num_heads % ckpt_tp or num_kv_heads % ckpt_tp:
        raise ValueError(
            f"fused qkv_proj is pre-sharded at {ckpt_tp} chunks, which do not "
            f"divide num_heads={num_heads} / num_kv_heads={num_kv_heads}."
        )
    # When there are fewer KV heads than ranks, vLLM replicates them
    # (`num_kv_head_replicas`) and rank r owns KV head r // replicas, which
    # keeps every Q head grouped with the KV head it attends to.
    if tp_size <= num_kv_heads:
        assert num_kv_heads % tp_size == 0, (
            f"num_kv_heads={num_kv_heads} must be divisible by tp_size={tp_size}."
        )
        kv_head_ids = list(
            range(
                tp_rank * (num_kv_heads // tp_size),
                (tp_rank + 1) * (num_kv_heads // tp_size),
            )
        )
    else:
        assert tp_size % num_kv_heads == 0, (
            f"tp_size={tp_size} must be divisible by num_kv_heads={num_kv_heads}."
        )
        kv_head_ids = [tp_rank // (tp_size // num_kv_heads)]
    q_head_ids = list(
        range(tp_rank * (num_heads // tp_size), (tp_rank + 1) * (num_heads // tp_size))
    )

    rows_per_chunk = w_full.shape[0] // ckpt_tp
    q_per_chunk = (num_heads // ckpt_tp) * head_dim
    k_per_chunk = (num_kv_heads // ckpt_tp) * head_dim
    v_per_chunk = (num_kv_heads // ckpt_tp) * v_head_dim
    if q_per_chunk + k_per_chunk + v_per_chunk != rows_per_chunk:
        raise ValueError(
            f"fused qkv_proj has {w_full.shape[0]} rows, not {ckpt_tp} chunks of "
            f"{q_per_chunk} Q + {k_per_chunk} K + {v_per_chunk} V rows."
        )

    # The scales are tiled per chunk; they collapse to one continuous grid when
    # a chunk is a whole number of blocks (the SWA chunks are: 3712 = 29 * 128).
    chunk_scale_rows = cdiv(rows_per_chunk, block)
    rows = torch.arange(w_full.shape[0])
    per_chunk_scales = s_full.shape[0] == ckpt_tp * chunk_scale_rows
    if per_chunk_scales:
        scale_index = (rows // rows_per_chunk) * chunk_scale_rows + (
            rows % rows_per_chunk
        ) // block
    elif s_full.shape[0] == cdiv(w_full.shape[0], block):
        scale_index = rows // block
    else:
        raise ValueError(
            f"fused qkv_proj scale has {s_full.shape[0]} rows, expected either "
            f"{ckpt_tp * chunk_scale_rows} ({ckpt_tp} chunks of "
            f"{chunk_scale_rows} rows) or {cdiv(w_full.shape[0], block)} "
            f"(one continuous grid)"
        )

    if tp_size == ckpt_tp and per_chunk_scales:
        # One checkpoint chunk per rank: already [Q | K | V] for that rank, and
        # its scale rows line up with ceil(rows_per_chunk / block).
        return (
            w_full.chunk(ckpt_tp, dim=0)[tp_rank],
            s_full.chunk(ckpt_tp, dim=0)[tp_rank],
        )

    # Gather this rank's Q, K and V rows from the chunks that hold them.
    q_heads_per_chunk = num_heads // ckpt_tp
    kv_heads_per_chunk = num_kv_heads // ckpt_tp
    head_rows = torch.arange(head_dim)
    v_head_rows = torch.arange(v_head_dim)
    row_index: list[torch.Tensor] = []
    for head in q_head_ids:
        chunk = head // q_heads_per_chunk
        row_index.append(
            chunk * rows_per_chunk + (head % q_heads_per_chunk) * head_dim + head_rows
        )
    for head in kv_head_ids:
        chunk = head // kv_heads_per_chunk
        row_index.append(
            chunk * rows_per_chunk
            + q_per_chunk
            + (head % kv_heads_per_chunk) * head_dim
            + head_rows
        )
    for head in kv_head_ids:
        chunk = head // kv_heads_per_chunk
        row_index.append(
            chunk * rows_per_chunk
            + q_per_chunk
            + k_per_chunk
            + (head % kv_heads_per_chunk) * v_head_dim
            + v_head_rows
        )
    index = torch.cat(row_index)
    grouped = w_full[index].to(torch.float32) * s_full[
        scale_index[index]
    ].repeat_interleave(block, dim=1)
    return _requantize_fp8(grouped, index.numel(), block, w_full.dtype)