01

整体架构

MiniMaxM3TinyForConditionalGeneration(L2443)是 MiniMax M3-Tiny 多模态视觉语言模型在 vLLM 中的完整实现。它同时注册了 SupportsMultiModalSupportsPPHasInnerState 接口,能够处理文本、图像和视频三种模态输入。

四大核心组件
  • vision_towerMiniMaxVL7UVisionModelMiniMaxVL7UVisionTransformer,L777):ViT 视觉编码器,支持 2D/3D Conv + 2D/3D RoPE + Video Patchified Attention (VPA)。
  • multi_modal_projectorMiniMaxVL7UMultiModalProjector,L291):2 层 MLP(ColumnParallelLinear → act → RowParallelLinear),ViT 特征空间 → LLM 特征空间。
  • patch_merge_mlpMiniMaxVL7UPatchMerger,L324):空间 Token 压缩,将 spatial_merge_size² 个 patch token 合并为 1 个,实现 4× 压缩。
  • language_modelMiniMaxM2ForCausalLMMiniMaxM2SparseForCausalLM):根据 sparse_attention_config 切换 Dense/Sparse 后端。
minimax_m3_tiny.py — MiniMaxM3TinyForConditionalGeneration.__init__ L2461-L2528
class MiniMaxM3TinyForConditionalGeneration(nn.Module, SupportsMultiModal,
                                          SupportsPP, HasInnerState):

    def __init__(self, *, vllm_config: VllmConfig, prefix: str = "") -> None:
        super().__init__()
        config = vllm_config.model_config.hf_config
        quant_config = vllm_config.quant_config

        # 缓存图片/视频 token 信息,用于动态分割 multimodal_embeddings
        self._cached_image_token_counts = []
        self._cached_video_token_counts = []

        # 四大核心组件
        self.vision_tower = MiniMaxVL7UVisionModel(
            config=config.vision_config, quant_config=quant_config,
            require_post_norm=False,
            prefix=maybe_prefix(prefix, "vision_tower"))

        self.multi_modal_projector = MiniMaxVL7UMultiModalProjector(
            vision_hidden_size=config.vision_config.hidden_size,
            text_hidden_size=config.text_config.hidden_size,
            projector_hidden_act=config.projector_hidden_act,
            multimodal_projector_bias=True, ...)

        self.patch_merge_mlp = MiniMaxVL7UPatchMerger(
            spatial_merge_size=config.img_token_compression_config.spatial_merge_size,
            text_hidden_size=config.text_config.hidden_size, ...)

        # Dense vs Sparse 语言模型后端
        self.is_sparse_attention_model = (
            getattr(config.text_config, "sparse_attention_config", None) is not None
            and config.text_config.sparse_attention_config["use_sparse_attention"])
        if self.is_sparse_attention_model:
            self.language_model = MiniMaxM2SparseForCausalLM(...)
        else:
            self.language_model = MiniMaxM2ForCausalLM(...)
graph TD subgraph Input["输入层"] IDS["input_ids
文本 token 序列"] IMG["pixel_values
图像 [C,T,H,W]"] VID["pixel_values_videos
视频 [C,T,H,W]"] end subgraph Vision["视觉分支"] VT["vision_tower
MiniMaxVL7UVisionTransformer
2D/3D Conv + RoPE + VPA"] PROJ["multi_modal_projector
ColumnParallelLinear → act → RowParallelLinear"] MERGE["patch_merge_mlp
MiniMaxVL7UPatchMerger
spatial_merge_size² → 1"] end subgraph LM["语言模型"] EMB["embed_input_ids"] FUSE["merge_multimodal_embeddings
image_token_id=200025
video_token_id=200026"] LLM["language_model
Dense / Sparse 切换"] end IMG --> VT VID --> VT VT --> PROJ --> MERGE --> FUSE IDS --> EMB --> FUSE FUSE --> LLM

模型同时支持图像视频输入,使用不同的 token ID 进行占位:image_token_id=200025 对应 ]<]image[>[video_token_id=200026 对应 ]<]video[>[。图像和视频的 embedding 分别合并到文本序列中,而非混合处理。

02

CLIPVisionEmbeddings:2D/3D Conv Patch 编码

CLIPVisionEmbeddings(L434)是视觉编码器的入口层,根据 temporal_patch_size 配置自动选择 2D Conv 或 3D Conv:

minimax_m3_tiny.py — CLIPVisionEmbeddings.__init__ L434-L464
class CLIPVisionEmbeddings(nn.Module):
    def __init__(self, config: CLIPVisionConfig):
        super().__init__()
        self.patch_size = config.patch_size
        self.embed_dim = config.hidden_size
        self.temporal_patch_size = config.img_token_compression_config.get(
            "temporal_patch_size", 2)
        self.spatial_merge_size = config.img_token_compression_config.get(
            "spatial_merge_size", 2)

        if self.temporal_patch_size == 1:
            # 图像路径:标准 2D Conv
            self.patch_embedding = nn.Conv2d(
                in_channels=config.num_channels,
                out_channels=self.embed_dim,
                kernel_size=self.patch_size,
                stride=self.patch_size, bias=False)
        else:
            # 视频路径:3D Conv,时间维度步长 = temporal_patch_size
            self.patch_embedding = nn.Conv3d(
                in_channels=config.num_channels,
                out_channels=self.embed_dim,
                kernel_size=(self.temporal_patch_size,
                             self.patch_size, self.patch_size),
                stride=(self.temporal_patch_size,
                        self.patch_size, self.patch_size),
                bias=False)

forward 方法(L466)中最关键的处理是 Patch 区域重排——将空间维度按 spatial_merge_size 重新排列,使相邻 patch 在序列中连续,与后续 RoPE 位置编码和 PatchMerger 对齐:

minimax_m3_tiny.py — CLIPVisionEmbeddings.forward(3D Conv 路径) L513-L538
# 3D Conv 路径
# 输入: pixel_values [batch, channel, temporal, height, width]
if pixel_values.dim() == 4:
    pixel_values = pixel_values.unsqueeze(2)  # 添加时间维度

# 3D Conv: [B, C, T, H, W] -> [B, embed_dim, T', H', W']
# 其中 T' = T // temporal_patch_size, H' = H // patch_size, W' = W // patch_size
patch_embeds = self.patch_embedding(pixel_values)

_, embed_dim, n_temporal, n_height, n_width = patch_embeds.shape

# 关键:Patch 区域重排,与 3D RoPE 对齐
# [B, D, T', H', W'] -> [B, D, H'//m, W'//m, T', m, m]
patch_embeds = patch_embeds.reshape(
    batch_size, embed_dim, n_temporal,
    n_height // spatial_merge_size, spatial_merge_size,
    n_width // spatial_merge_size, spatial_merge_size
).permute(0, 1, 3, 5, 2, 4, 6)

# 展平为序列: [B, D, seq_len] -> [B, seq_len, D]
patch_embeds = patch_embeds.flatten(2).transpose(1, 2).contiguous()
Patch 区域重排的本质

重排操作将原始的逐行扫描顺序(H'×W')变为按 spatial_merge_size×spatial_merge_size 的块分组。例如 spatial_merge_size=2 时,原本相距 W' 个位置的上下两行 patch 变为在序列中连续排列。这样做是为了让后续 PatchMerger 能直接按连续的 4 个 token 分组合并,同时 RoPE 的高度/宽度位置 ID 也需要对应重排。

03

VisionFlashAttention2:变长注意力

VisionFlashAttention2(L563)是 ViT 的注意力层。与标准 FlashAttention 不同,它使用 flash_attn_varlen_func——变长序列的注意力实现,通过 cu_seq_len(cumulative sequence lengths)控制哪些 token 互相可见:

minimax_m3_tiny.py — VisionFlashAttention2.forward L628-L650
def forward(self, hidden_states, cu_seq_len, max_seqlen=None,
            rotary_pos_emb=None, **kwargs):
    seq_length = hidden_states.shape[0]

    # QKV 投影(支持 Tensor Parallel)
    qkv_states, _ = self.qkv_proj(hidden_states)
    q, k, v = self.split_qkv(qkv_states)  # [seq, 1, tp_head, head_dim]
    q, k, v = q.squeeze(1), k.squeeze(1), v.squeeze(1)

    # 应用 RoPE 位置编码
    if self.config.position_embedding_type == "rope":
        q = apply_rotary_pos_emb_vision(q.unsqueeze(0), rotary_pos_emb).squeeze(0)
        k = apply_rotary_pos_emb_vision(k.unsqueeze(0), rotary_pos_emb).squeeze(0)

    # 变长注意力:cu_seq_len 控制注意力边界
    attn_output = flash_attn_varlen_func(
        q, k, v,
        cu_seq_len, cu_seq_len,     # Q 和 KV 使用相同的边界
        max_seqlen, max_seqlen,
        causal=getattr(self.config, 'causal_attention', False))

    attn_output = attn_output.view(seq_length, -1)
    attn_output = self.out_proj(attn_output)[0].unsqueeze(1)
    return attn_output, None
cu_seq_len 的关键作用

cu_seq_len 是一个递增的整数数组,定义了每个"注意力组"的边界。例如 [0, 576, 1152, 1728] 表示 3 个独立的注意力组,每组 576 个 token,组内互相可见、组间互不可见。
video_segments=[4](所有帧在一个段)时,cu_seq_len = [0, 2304]——所有帧的 token 互相可见(Video Patchified Attention)。
video_segments=None 时,每帧独立,cu_seq_len = [0, 576, 1152, ...]——每帧只能看自己。

注意力层使用 QKVParallelLinear(L582)进行 Q/K/V 联合投影,支持 Tensor Parallel。在 TP 模式下,通过 all_gather_interleave(L541)收集各 rank 的 QKV 输出,再按 rank 切分分配:

minimax_m3_tiny.py — split_qkv(Tensor Parallel 处理) L597-L626
def split_qkv(self, qkv):
    seq_len, bs, _ = qkv.shape

    if self.tp_size > 1:
        # 跨 TP rank 收集完整 QKV,交错排列
        qkv = all_gather_interleave(qkv, self.qkv_proj.hidden_size, self.tp_size)

    q, k, v = qkv.chunk(3, dim=2)

    if self.tp_size > 1:
        # 各 rank 取自己负责的 head 切片
        splitter = partial(dist_utils.split_tensor_along_last_dim,
                          num_partitions=self.tp_size)
        q = splitter(q)[self.tp_rank]
        k = splitter(k)[self.tp_rank]
        v = splitter(v)[self.tp_rank]

    # reshape 为 [seq, batch, num_heads_per_partition, head_dim]
    new_shape = (seq_len, bs,
                 self.num_attention_heads_per_partition,
                 self.hidden_size_per_attention_head)
    q, k, v = (x.view(*new_shape) for x in (q, k, v))
    return q, k, v
04

2D/3D RoPE 位置编码

MiniMaxVL7UVisionTransformer 根据 rope_mode 配置使用 2D 或 3D RoPE。rot_pos_emb(L942)是统一入口:

4.1 2D RoPE(图像)

2D RoPE(L967)将 head_dim 平均分为 hw 两部分,使用 VisionRotaryEmbedding(L423)生成频率表,然后按 spatial_merge 重排后的位置 ID 查表拼接:

minimax_m3_tiny.py — _get_rope_embed_2d L967-L1013
def _get_rope_embed_2d(self, grid_thw, spatial_merge_size):
    pos_ids = []
    max_grid_size = 0

    for h, w in grid_thw:
        w = w // self.config.patch_size
        h = h // self.config.patch_size
        max_grid_size = max(max_grid_size, w, h)

        # 高度位置 ID:按 spatial_merge_size 重排
        hpos_ids = torch.arange(h).unsqueeze(1).expand(-1, w)
        hpos_ids = hpos_ids.reshape(
            h // spatial_merge_size, spatial_merge_size,
            w // spatial_merge_size, spatial_merge_size,
        ).permute(0, 2, 1, 3).flatten()

        # 宽度位置 ID:同样重排
        wpos_ids = torch.arange(w).unsqueeze(0).expand(h, -1)
        wpos_ids = wpos_ids.reshape(
            h // spatial_merge_size, spatial_merge_size,
            w // spatial_merge_size, spatial_merge_size,
        ).permute(0, 2, 1, 3).flatten()

        pos_ids.append(torch.stack([hpos_ids, wpos_ids], dim=-1))

    pos_ids = torch.cat(pos_ids, dim=0)
    # 查频率表:rotary_pos_emb_full[pos_ids] → [seq_len, rope_dim]
    rotary_pos_emb_full = self.rotary_pos_emb(max_grid_size)
    rotary_pos_emb = rotary_pos_emb_full[pos_ids].flatten(1)
    return rotary_pos_emb

4.2 3D RoPE(视频)

3D RoPE(L879)将 head_dim 三等分为 t_dimh_dimw_dim,使用三组独立的 inv_freq 缓冲区。每帧的 temporal 位置由 frame_index 决定(来自 _compute_frame_indices):

minimax_m3_tiny.py — _get_3d_rope_embed(单帧) L879-L940
def _get_3d_rope_embed(self, grid_h, grid_w, frame_index, spatial_merge_size):
    """单帧的 3D RoPE 位置嵌入"""
    tokens_per_frame = grid_h * grid_w

    # ===== 时间位置 ID =====
    # 当前帧的所有 token 使用相同的时间位置 frame_index
    tpos_ids = torch.full((tokens_per_frame,), frame_index, dtype=torch.long)

    # ===== 高度位置 ID(带 spatial_merge 重排)=====
    hpos_ids = torch.arange(grid_h).unsqueeze(1).expand(-1, grid_w)
    hpos_ids = hpos_ids.reshape(
        grid_h // spatial_merge_size, spatial_merge_size,
        grid_w // spatial_merge_size, spatial_merge_size,
    ).permute(0, 2, 1, 3).flatten()

    # ===== 宽度位置 ID(同样重排)=====
    wpos_ids = torch.arange(grid_w).unsqueeze(0).expand(grid_h, -1)
    wpos_ids = wpos_ids.reshape(
        grid_h // spatial_merge_size, spatial_merge_size,
        grid_w // spatial_merge_size, spatial_merge_size,
    ).permute(0, 2, 1, 3).flatten()

    # ===== 三维频率表 =====
    max_hw = max(grid_h, grid_w)
    freqs_t = torch.outer(torch.arange(max(frame_index + 1, 1)), self.inv_freq_t)
    freqs_h = torch.outer(torch.arange(max_hw), self.inv_freq_h)
    freqs_w = torch.outer(torch.arange(max_hw), self.inv_freq_w)

    # ===== 查表拼接 =====
    emb_t = freqs_t[tpos_ids]  # [seq_len, t_dim/2]
    emb_h = freqs_h[hpos_ids]  # [seq_len, h_dim/2]
    emb_w = freqs_w[wpos_ids]  # [seq_len, w_dim/2]

    rotary_pos_emb = torch.cat([emb_t, emb_h, emb_w], dim=-1)
    return rotary_pos_emb  # [seq_len, (t_dim + h_dim + w_dim) / 2]
3D RoPE 与 head_dim 的维度对齐

3D RoPE 的总维度 t_dim + h_dim + w_dim 可能不等于 head_dim(例如 head_dim=64,三等分各 20,总 60)。apply_rotary_pos_emb_vision(L382)采用分割-应用-拼接策略:只对前 rot_dim 维应用旋转,剩余维度直接传递(passthrough),与训练代码保持一致。

graph LR subgraph TwoD["2D RoPE(图像)"] H2["head_dim / 2
→ h_freq"] W2["head_dim / 2
→ w_freq"] R2["concat → [seq, head_dim/2]"] end subgraph ThreeD["3D RoPE(视频)"] T3["t_dim
→ t_freq"] H3["h_dim
→ h_freq"] W3["w_dim
→ w_freq"] R3["concat → [seq, (t+h+w)/2]"] end H2 --> R2 W2 --> R2 T3 --> R3 H3 --> R3 W3 --> R3
05

Video Patchified Attention (VPA)

VPA 是 M3-Tiny 视频理解的核心机制——通过 video_segments 控制哪些帧共享注意力。这涉及三个关键函数的协作:

5.1 _compute_cu_seq_len:注意力分组

minimax_m3_tiny.py — _compute_cu_seq_len L1183-L1270
def _compute_cu_seq_len(self, frame_token_counts, video_segments, device):
    """
    计算 flash_attn_varlen_func 的 cumulative sequence lengths。

    Args:
        frame_token_counts: 每个压缩帧的 token 数列表
        video_segments: 每个视频段的压缩帧数
            e.g., [8, 8] = 两段各 8 个压缩帧

    Examples:
        frame_token_counts = [576, 576, 576, 576]

        Case 1: video_segments = None(每帧独立)
            cu_seq_len = [0, 576, 1152, 1728, 2304]

        Case 2: video_segments = [4](所有帧一个段 → VPA)
            cu_seq_len = [0, 2304]
            所有帧互相可见。

        Case 3: video_segments = [2, 2](两段各 2 帧)
            cu_seq_len = [0, 1152, 2304]
            帧 0-1 互相可见,帧 2-3 互相可见。
    """
    num_frames = len(frame_token_counts)

    if video_segments is None:
        # 每帧独立
        cu_seq_len = [0] + frame_token_counts
        cu_seq_len = torch.tensor(cu_seq_len, device=device).to(torch.int32)
        return torch.cumsum(cu_seq_len, dim=0).to(torch.int32)

    # 应用 vision_segment_max_frames 限制
    effective_segments = self._apply_max_frames_limit(video_segments, num_frames)

    # 按段聚合 token 数
    cu_seq_len = [0]
    frame_idx = 0
    for segment_frame_count in effective_segments:
        segment_tokens = sum(frame_token_counts[frame_idx:frame_idx + segment_frame_count])
        cu_seq_len.append(segment_tokens)
        frame_idx += segment_frame_count

    cu_seq_len = torch.tensor(cu_seq_len, device=device).to(torch.int32)
    return torch.cumsum(cu_seq_len, dim=0).to(torch.int32)

5.2 _compute_frame_indices:帧内索引

_compute_frame_indices(L1063)为 3D RoPE 计算每帧在其所属视频段中的相对索引(frame_index),段间索引重置为 0:

minimax_m3_tiny.py — _compute_frame_indices L1063-L1103
def _compute_frame_indices(self, num_frames, video_segments):
    """
    根据 video_segments 计算每帧的 frame_index(段内索引)。

    video_segments=[3, 2] → frame_indices=[0, 1, 2, 0, 1]
    video_segments=None   → frame_indices=[0, 0, 0, ...](全部视为独立图像)

    若设置 vision_segment_max_frames=4:
      video_segments=[10] → effective=[4, 4, 2]
      → frame_indices=[0,1,2,3, 0,1,2,3, 0,1]
    """
    if video_segments is None:
        return [0] * num_frames

    # 应用 max_frames 限制(内含 _normalize_video_segments 处理 [0] 特殊值)
    effective_segments = self._apply_max_frames_limit(video_segments, num_frames)

    frame_indices = []
    for segment_size in effective_segments:
        for i in range(segment_size):
            frame_indices.append(i)

    return frame_indices

5.3 vision_segment_max_frames:超长段拆分

_apply_max_frames_limit(L1134)将超过 vision_segment_max_frames 的段自动拆分为多个小段。该限制针对压缩后的帧数,可通过环境变量 VISION_SEGMENT_MAX_FRAMES 覆盖配置值(L818):

minimax_m3_tiny.py — _apply_max_frames_limit L1134-L1181
def _apply_max_frames_limit(self, video_segments, num_frames):
    """
    将超长段拆分为多个小段。

    Example:
        temporal_patch_size=2, vision_segment_max_frames=4
        原始 20 帧 → 压缩 10 帧
        video_segments=[10] → effective_segments=[4, 4, 2]
    """
    normalized_segments = self._normalize_video_segments(video_segments, num_frames)

    if self.vision_segment_max_frames is None:
        return normalized_segments

    max_frames = self.vision_segment_max_frames
    effective_segments = []

    for segment_size in normalized_segments:
        if segment_size <= max_frames:
            effective_segments.append(segment_size)
        else:
            remaining = segment_size
            while remaining > 0:
                chunk_size = min(remaining, max_frames)
                effective_segments.append(chunk_size)
                remaining -= chunk_size

    return effective_segments
video_segments 中的特殊值 [0]

_normalize_video_segments(L1105)会将 [0][tensor(0)] 解释为"所有帧在一个 segment",即 [0][num_frames]。这在某些 API 调用场景中用作快捷写法。

06

VisionTransformer forward 流程

MiniMaxVL7UVisionTransformer.forward(L1272)将上述所有组件串联。输入是一个帧列表(每帧为 [C, T, H, W] 的 4D tensor),输出是所有帧拼接后的特征向量。

minimax_m3_tiny.py — VisionTransformer.forward L1272-L1382
def forward(self, pixel_values, image_sizes,
            feature_sample_layers=None, video_segments=None):
    inputs = []
    frame_token_counts = []
    max_seqlen = 0

    for pixel_value in pixel_values:
        # 1. Patch Embedding(2D/3D Conv + 区域重排)
        hidden_states = self.embeddings(pixel_value.unsqueeze(0))
        hidden_states = self.pre_layrnorm(hidden_states)
        # [B, S, C] → [S, B, C](flash-attn 输入格式)
        hidden_states = hidden_states.permute(1, 0, 2).contiguous()
        frame_token_counts.append(hidden_states.size(0))
        max_seqlen = max(max_seqlen, hidden_states.size(0))
        inputs.append(hidden_states)

    # 拼接所有帧的 token
    inputs = torch.cat(inputs, dim=0)  # [total_seq, 1, embed_dim]

    # 2. 计算 cu_seq_len(VPA 注意力分组)
    cu_seq_len = self._compute_cu_seq_len(
        frame_token_counts, video_segments, inputs.device)

    # 更新 max_seqlen 为段级别的最大值
    if video_segments is not None:
        effective_segments = self._apply_max_frames_limit(
            video_segments, len(frame_token_counts))
        frame_idx = 0
        segment_token_counts = []
        for seg_count in effective_segments:
            segment_tokens = sum(frame_token_counts[frame_idx:frame_idx + seg_count])
            segment_token_counts.append(segment_tokens)
            frame_idx += seg_count
        max_seqlen = max(segment_token_counts)

    # 3. 计算 RoPE 位置编码(2D 或 3D)
    rotary_pos_emb = self.rot_pos_emb(image_sizes, video_segments)

    # 4. N 层 CLIPEncoder(每层:LN → VisionFlashAttention2 → LN → CLIPMLP)
    encoder_outputs = self.encoder(
        inputs_embeds=inputs,
        cu_seq_len=cu_seq_len,
        max_seqlen=max_seqlen,
        rotary_pos_emb=rotary_pos_emb)

    # [S, B, C] → [total_seq, hidden_size]
    encoder_outputs = encoder_outputs.permute(1, 0, 2).contiguous()[0]
    return encoder_outputs
graph TD A["pixel_values
[C,T,H,W] × N帧"] -->|"逐帧"| B["CLIPVisionEmbeddings
2D/3D Conv + 区域重排"] B --> C["pre_layrnorm"] C --> D["torch.cat
拼接所有帧 token"] D --> E["_compute_cu_seq_len
VPA 注意力边界"] D --> F["rot_pos_emb
2D/3D RoPE"] E --> G["CLIPEncoder
N × (FlashAttn + CLIPMLP)"] F --> G D --> G G --> H["encoder_outputs
[total_seq, hidden_size]"]
07

Projector 与 PatchMerger

7.1 MiniMaxVL7UMultiModalProjector

MiniMaxVL7UMultiModalProjector(L291)是 ViT → LLM 的特征空间映射,2 层 MLP 结构使用 vLLM 的并行化 Linear:

minimax_m3_tiny.py — MiniMaxVL7UMultiModalProjector L291-L321
class MiniMaxVL7UMultiModalProjector(nn.Module):
    def __init__(self, vision_hidden_size, text_hidden_size,
                 projector_hidden_act, multimodal_projector_bias,
                 projector_hidden_size=None, quant_config=None, prefix=""):
        super().__init__()
        projector_intermediate = projector_hidden_size or text_hidden_size
        self.linear_1 = ColumnParallelLinear(
            vision_hidden_size, projector_intermediate,
            bias=multimodal_projector_bias, quant_config=quant_config)
        self.act = get_act_fn(projector_hidden_act)
        self.linear_2 = RowParallelLinear(
            projector_intermediate, text_hidden_size,
            bias=multimodal_projector_bias, quant_config=quant_config)

    def forward(self, image_features):
        hidden_states, _ = self.linear_1(image_features)
        hidden_states = self.act(hidden_states)
        hidden_states, _ = self.linear_2(hidden_states)
        return hidden_states

7.2 MiniMaxVL7UPatchMerger

MiniMaxVL7UPatchMerger(L324)的结构与 Projector 相同(ColumnParallelLinear → act → RowParallelLinear),但输入维度是 text_hidden_size × spatial_merge_size²——在调用前,spatial_merge_size² 个 token 已被拼接为一个长向量:

minimax_m3_tiny.py — MiniMaxVL7UPatchMerger L324-L352
class MiniMaxVL7UPatchMerger(nn.Module):
    def __init__(self, spatial_merge_size, text_hidden_size, ...):
        super().__init__()
        # 输入维度 = text_hidden_size × spatial_merge_size²
        # 输出维度 = text_hidden_size
        self.linear_1 = ColumnParallelLinear(
            text_hidden_size * spatial_merge_size ** 2,
            projector_intermediate, ...)
        self.act = get_act_fn(projector_hidden_act)
        self.linear_2 = RowParallelLinear(
            projector_intermediate, text_hidden_size, ...)

    def forward(self, image_features):
        hidden_states, _ = self.linear_1(image_features)
        hidden_states = self.act(hidden_states)
        hidden_states, _ = self.linear_2(hidden_states)
        return hidden_states
PatchMerger 的调用方式

注意 PatchMerger 的 forward 接收的是已拼接好的向量。拼接操作在 _process_image_input / _process_video_input 中完成:先按 chunk_sizes = [token_num × spatial_merge_size²] 分割 projector 输出,再 .view(-1, hidden_size × spatial_merge_size²) 重排为合并后的长向量输入 PatchMerger。

08

输入预处理与格式归一化

在进入 ViT 编码之前,图像和视频的原始输入需要经过预处理,统一转换为 3D Conv 所需的 4D 张量格式 [C, T, H, W]。这一步由 _parse_and_validate_image_input(L2786)和 _parse_and_validate_video_input(L2873)完成。

8.1 数据结构定义

图像和视频输入分别由两个 TypedDict 定义:

minimax_m3_tiny.py — 图像/视频输入类型定义 L229-L275
class MiniMaxVL7UImagePixelInputs(TypedDict, total=False):
    type: Literal["pixel_values"]
    pixel_values: torch.Tensor
    # Shape: (batch_size * num_images, num_channels, height, width)
    # 或 list(当不同图片尺寸不同时)
    image_sizes: torch.Tensor
    # Shape: (batch_size * num_images, 2)  → (height, width)
    video_segments: Optional[List[int]]
    # 例如 [16, 16] = 两个视频段,每段 16 帧共享注意力

class MiniMaxVL7UVideoPixelInputs(TypedDict, total=False):
    type: Literal["pixel_values_videos"]
    pixel_values_videos: torch.Tensor
    # Shape: (batch_size * num_videos, num_channels, height, width)
    video_sizes: torch.Tensor
    # Shape: (batch_size * num_videos, 3)  → (temporal, height, width)
    video_segments: Optional[List[int]]

8.2 图像输入归一化:_parse_and_validate_image_input

该函数(L2786)将来自不同来源的图像张量统一转换为 [C, T, H, W] 格式的 flat list,为 3D Conv 做准备:

minimax_m3_tiny.py — _parse_and_validate_image_input L2786-L2871
def _parse_and_validate_image_input(self, **kwargs):
    pixel_values = kwargs.pop("pixel_values", None)
    image_sizes = kwargs.pop("image_sizes", None)

    if pixel_values is not None:
        # 6D tensor: (batch, n_images, temporal_patch_size, C, H, W)
        if isinstance(pixel_values, torch.Tensor) and pixel_values.dim() == 6:
            new_pixel_values = []
            for i in range(pixel_values.shape[0]):
                for j in range(pixel_values.shape[1]):
                    # (T, C, H, W) → (C, T, H, W)  permute 通道到前面
                    new_pixel_values.append(
                        pixel_values[i, j].permute(1, 0, 2, 3))

        # 5D tensor: (batch, n_images, C, H, W) — 无时间维度
        elif isinstance(pixel_values, torch.Tensor) and pixel_values.dim() == 5:
            new_pixel_values = []
            for i in range(pixel_values.shape[0]):
                for j in range(pixel_values.shape[1]):
                    # (C, H, W) → (C, 1, H, W)  添加时间维度
                    new_pixel_values.append(
                        pixel_values[i, j].unsqueeze(1))

        # list of list:逐元素处理
        elif isinstance(pixel_values, list):
            new_pixel_values = []
            for i in range(len(pixel_values)):
                for j in range(len(pixel_values[i])):
                    item = pixel_values[i][j]
                    if item.dim() == 4:  # (T, C, H, W) → (C, T, H, W)
                        new_pixel_values.append(item.permute(1, 0, 2, 3))
                    elif item.dim() == 3:  # (C, H, W) → (C, 1, H, W)
                        new_pixel_values.append(item.unsqueeze(1))

    # 同步展平 image_sizes(3D tensor / nested list → flat list)
    return MiniMaxVL7UImagePixelInputs(
        type="pixel_values", pixel_values=new_pixel_values,
        image_sizes=new_image_sizes)
graph TD A["pixel_values 输入"] --> B{"检测维度"} B -->|"6D: [B, N, T, C, H, W]"| C["permute(1,0,2,3)
→ [C, T, H, W]"] B -->|"5D: [B, N, C, H, W]"| D["unsqueeze(1)
→ [C, 1, H, W]"] B -->|"list of list"| E["逐元素判断
4D→permute / 3D→unsqueeze"] C --> F["flat list of [C, T, H, W]"] D --> F E --> F F --> G["MiniMaxVL7UImagePixelInputs"]

8.3 视频输入归一化:_parse_and_validate_video_input

视频输入格式更复杂(可能是 7D tensor 或深度嵌套 list),因此使用递归展平策略。该函数(L2873)定义了三个嵌套辅助函数:

minimax_m3_tiny.py — _parse_and_validate_video_input 递归展平 L2896-L3024
# 辅助函数 1:处理单个 tensor → (C, T, H, W)
def process_single_tensor(tensor: torch.Tensor) -> torch.Tensor:
    # 先 squeeze 掉所有前导的 dim=1
    while tensor.dim() > 4 and tensor.shape[0] == 1:
        tensor = tensor.squeeze(0)
    if tensor.dim() == 4:
        if tensor.shape[0] == 3:    # 已经是 (C, T, H, W)
            return tensor
        if tensor.shape[1] == 3:    # 是 (T, C, H, W),需要 permute
            return tensor.permute(1, 0, 2, 3)
    if tensor.dim() == 3:           # (C, H, W) → 添加时间维度
        return tensor.unsqueeze(1)

# 辅助函数 2:递归展平嵌套的 tensor / list
def flatten_video_tensors(data) -> list[torch.Tensor]:
    if isinstance(data, torch.Tensor):
        if data.dim() <= 4:
            return [process_single_tensor(data)]
        flattened = []
        for child in data:
            flattened.extend(flatten_video_tensors(child))
        return flattened
    if isinstance(data, list):
        flattened = []
        for child in data:
            if child is None: continue
            flattened.extend(flatten_video_tensors(child))
        return flattened

# 辅助函数 3:递归展平 video_sizes
def flatten_video_sizes(data) -> list:
    # 1D tensor → [data], ≥2D tensor → 递归
    # tuple → [data], list of numbers → [data], nested list → 递归

# 辅助函数 4:递归展平 video_segments
def flatten_video_segments(data) -> Optional[list[int]]:
    # 0D tensor → [int(item)], ≥1D tensor → 递归
    # list/tuple → 递归, int/float → [int(data)]
递归展平的设计意图

vLLM 的输入可能来自不同的前端(HuggingFace processor、手动构造、batch 拼接等),每种来源的 tensor 嵌套层级不同。递归展平保证了无论输入嵌套多深,最终都能收敛到统一的 [C, T, H, W] flat list。同时 pixel_values_videosvideo_sizesvideo_segments 三者的展平结果数量必须一致,否则抛出 ValueError

通道维度自动检测

process_single_tensor 通过检查 shape[0]==3shape[1]==3 来自动判断通道维度的位置(channels_first vs temporal_first),这意味着 [C, T, H, W][T, C, H, W] 两种格式都能正确处理。

09

Token 计算与 Prompt 展开

在视觉特征编码之前,需要先计算每张图片/每帧视频对应多少 token,并在文本 prompt 中展开占位符。这由 ImageTokensCalculator(L1538)和 _get_prompt_updates(L2343)完成。

9.1 ImageTokensCalculator:三种 token 计算模式

minimax_m3_tiny.py — ImageTokensCalculator.calculate_tokens L1538-L1662
class ImageTokensCalculator:
    def __init__(self):
        self.patch_size = 14
        # 28 种预定义分辨率(336×336 到 2016×1008)
        self.default_grid_pinpoints = [
            [336, 336], [336, 672], ..., [2016, 1008]
        ]

    def calculate_tokens(self, height, width):
        mode = self.config['process_image_mode']
        if mode == 'resize':
            # 最简单:直接按 patch 面积计算
            return int(height * width / self.patch_size ** 2)

        elif mode == 'dynamic_res':
            return self._get_dynamic_res_num_token(
                height, width, max_img_h, max_img_w)

        else:  # anyres
            return self._get_anyres_num_token(height, width)
三种模式对比
  • resizeh × w / patch_size²,最简单直接
  • dynamic_res:先调用 get_hw_multiple_of 对齐到 patch_size × spatial_merge_size 的倍数,支持 interpolate(面积超阈值则缩放)和 patch_merge(额外除以 spatial_merge_size)两种压缩方法
  • anyres:从 28 种预定义分辨率中选最佳匹配(_select_best_resolution,最大化有效分辨率、最小化浪费),额外加上 336/14 的 thumbnail tokens

9.2 get_hw_multiple_of:尺寸对齐

get_hw_multiple_of(L355)将图像尺寸向上对齐到指定倍数,并可选地裁切到 max_size 限制内:

minimax_m3_tiny.py — get_hw_multiple_of L355-L372
def get_hw_multiple_of(image_size, multiple, max_size=None):
    w, h = image_size
    # 向上对齐到 multiple 的倍数
    new_w = w if w % multiple == 0 else w + (multiple - w % multiple)
    new_h = h if h % multiple == 0 else h + (multiple - h % multiple)

    if max_size is not None:
        max_w, max_h = max_size
        if new_w > max_w or new_h > max_h:
            # 按比例缩放到 max_size 内,再次对齐
            new_w_ = min((new_w * max_w) // new_w,
                         (new_w * max_h) // new_h)
            new_h_ = min((new_h * max_w) // new_w,
                         (new_h * max_h) // new_h)
            new_w = new_w_ + (multiple - new_w_ % multiple) % multiple
            new_h = new_h_ + (multiple - new_h_ % multiple) % multiple
    return new_w, new_h

9.3 _get_prompt_updates:占位符展开

_get_prompt_updates(L2343)定义了如何将 prompt 中的单个占位符 token 展开为 start + N × image/video token + end 的序列:

minimax_m3_tiny.py — _get_prompt_updates L2343-L2433
def _get_prompt_updates(self, mm_items, hf_processor_mm_kwargs, out_mm_kwargs):
    image_token_id = hf_config.image_token_index  # 200025
    video_token_id = hf_config.video_token_index  # 200026
    START_TOKEN_ID = 200029   # ]<]start of image[>[
    END_TOKEN_ID   = 200030   # ]<]end of image[>[

    def get_image_replacement(item_idx):
        num_tokens = get_num_tokens_for_item(item_idx)
        token_ids = ([START_TOKEN_ID]
                   + [image_token_id] * num_tokens
                   + [END_TOKEN_ID])
        # 只有 image_token_id 位置会被视觉 embedding 替换
        return PromptUpdateDetails.select_token_id(
            token_ids, image_token_id)

    def get_video_replacement(item_idx):
        num_tokens = get_num_tokens_for_item(item_idx, modality="video")
        token_ids = ([START_TOKEN_ID]
                   + [video_token_id] * num_tokens
                   + [END_TOKEN_ID])
        return PromptUpdateDetails.select_token_id(
            token_ids, video_token_id)

    return [
        PromptReplacement(modality="image",
            target="]<]image[>[",
            replacement=get_image_replacement),
        PromptReplacement(modality="video",
            target="]<]video[>[",
            replacement=get_video_replacement),
    ]
graph LR A["原始 prompt:
...文本... ]<]image[>[ ...文本..."] --> B["token 展开"] B --> C["[200029] + [200025]×N + [200030]"] C --> D["get_input_embeddings 中
200025 位置被视觉 embedding 替换
200029/200030 保持文本 embedding"]
视频帧的 token 数计算

视频使用 out_mm_kwargs["video_sizes"] 获取 resize 后的帧尺寸(而非原始帧尺寸),这是因为 mm_items 中包含的是原始帧(未经时间压缩),而 video_sizes 反映了 temporal 压缩后 resize 到 ViT 输入尺寸的实际大小,与 VisionTransformer 一致。

10

图像处理管线:_process_image_input

_process_image_input(L2720)处理图像输入,完整流程:ViT → Projector → 按图分割 → PatchMerger:

minimax_m3_tiny.py — _process_image_input L2720-L2782
def _process_image_input(self, image_input):
    self._cached_image_token_counts = []

    # Step 1: ViT 编码(无 video_segments → 每帧独立注意力)
    image_features = self._process_image_pixels(image_input)
    #   → self.vision_tower(pixel_values, image_sizes)

    # Step 2: 计算每张图片的 token 数量
    spatial_merge_size = config.img_token_compression_config.spatial_merge_size
    for image_size in image_input["image_sizes"]:
        height, width = image_size[0].item(), image_size[1].item()
        new_width, new_height = get_hw_multiple_of(
            (width, height), patch_size * spatial_merge_size, max_size=...)
        num_patches_w = new_width // patch_size // spatial_merge_size
        num_patches_h = new_height // patch_size // spatial_merge_size
        self._cached_image_token_counts.append(num_patches_w * num_patches_h)

    # Step 3: Projector(ViT dim → LLM dim)
    image_embeds = self.multi_modal_projector(image_features.unsqueeze(1))

    # Step 4: 按每张图片的 token 数分割,逐张 PatchMerger
    chunk_sizes = [t * spatial_merge_size**2
                   for t in self._cached_image_token_counts]
    image_embeds = torch.split(
        image_embeds.squeeze(1).unsqueeze(0), chunk_sizes, dim=1)

    final_image_features = []
    for chunk_size, image_feature in zip(chunk_sizes, image_embeds):
        # reshape: [1, chunk, hidden] → [chunk/s², hidden * s²]
        image_feature = self.patch_merge_mlp(
            image_feature.view(-1, hidden_size * spatial_merge_size**2))
        final_image_features.append(image_feature)

    return tuple(final_image_features)

关键点:

  • 所有图片先一起经过 ViT 编码,输出一个拼接的特征向量
  • 然后按 _cached_image_token_counts 分割为每张图片对应的 chunk
  • 每个 chunk 独立经过 PatchMerger 压缩(乘以 spatial_merge_size² 再除以 spatial_merge_size²
11

视频处理管线:_process_video_input

_process_video_input(L2637)是视频处理的核心,与图像管线最大的区别在于:video_segments 被传入 vision_tower,启用 VPA 跨帧注意力。

minimax_m3_tiny.py — _process_video_input L2637-L2707
def _process_video_input(self, video_input):
    self._cached_video_token_counts = []

    video_pixel_values = video_input["pixel_values_videos"]
    video_sizes = video_input["video_sizes"]
    video_segments = video_input.get("video_segments", None)

    # Step 1: ViT 编码 —— 传入 video_segments 启用 VPA
    video_features = self.vision_tower(
        video_pixel_values, video_sizes, video_segments=video_segments)

    # Step 2: 计算每个压缩帧的 token 数量
    for frame_size in video_sizes:
        height, width = frame_size[-2].item(), frame_size[-1].item()
        new_width, new_height = get_hw_multiple_of(
            (width, height), patch_size * spatial_merge_size, max_size=...)
        tokens = (new_width // patch_size // spatial_merge_size *
                  new_height // patch_size // spatial_merge_size)
        self._cached_video_token_counts.append(tokens)

    # Step 3: Projector
    video_embeds = self.multi_modal_projector(video_features.unsqueeze(1))

    # Step 4: 按帧分割 + PatchMerger
    chunk_sizes = [t * spatial_merge_size**2
                   for t in self._cached_video_token_counts]
    video_embeds = torch.split(
        video_embeds.squeeze(1).unsqueeze(0), chunk_sizes, dim=1)

    final_video_features = []
    for chunk_size, video_feature in zip(chunk_sizes, video_embeds):
        video_feature = self.patch_merge_mlp(
            video_feature.contiguous().view(
                -1, hidden_size * spatial_merge_size**2))
        final_video_features.append(video_feature)

    return tuple(final_video_features)
graph TD A["video_pixel_values
[C,T,H,W] × N 压缩帧"] --> B["vision_tower.forward
传入 video_segments"] B --> B1["CLIPVisionEmbeddings
3D Conv 时间压缩"] B1 --> B2["_compute_cu_seq_len
VPA 注意力分组"] B1 --> B3["_get_rope_embed_3d
3D RoPE"] B2 --> B4["CLIPEncoder
N 层 FlashAttn"] B3 --> B4 B4 --> C["multi_modal_projector
ViT dim → LLM dim"] C --> D["按帧分割
chunk_sizes = [t × s²]"] D --> E["patch_merge_mlp
view(-1, hidden × s²)
→ 4× 压缩"] E --> F["final_video_features
tuple of tensors"]
视频处理的完整数据维度追踪

以 8 帧视频(temporal_patch_size=2patch_size=14,图像 336×336spatial_merge_size=2)为例:

  • 原始输入:8 帧 → 经 3D Conv 压缩为 4 个压缩帧
  • 每帧 patch 数:336/14 = 24,经区域重排后 24×24 = 576 个 token
  • ViT 编码后:4 × 576 = 2304 个 token,维度 vit_hidden_dim
  • Projector 后:2304 个 token,维度 text_hidden_dim
  • PatchMerger 后:2304 / 4 = 576 个 token → 最终进入 LLM 的视频 token 数
12

多模态融合:get_input_embeddings

get_input_embeddings(L2529)负责将图像和视频的 embedding 分别合并到文本序列中。M3-Tiny 使用两个独立的 token ID进行占位:image_token_id=200025video_token_id=200026,因此图像和视频的 embedding 分别合并:

minimax_m3_tiny.py — get_input_embeddings L2529-L2600
def get_input_embeddings(self, input_ids, multimodal_embeddings=None):
    inputs_embeds = self.language_model.embed_input_ids(input_ids)

    if multimodal_embeddings is not None and len(multimodal_embeddings) != 0:
        image_token_id = self.config.image_token_index  # 200025
        video_token_id = self.config.video_token_index  # 200026

        # 分别合并 image 和 video embedding
        if "image_embeddings" in multimodal_embeddings:
            inputs_embeds = merge_multimodal_embeddings(
                input_ids, inputs_embeds,
                multimodal_embeddings["image_embeddings"],
                [image_token_id])

        if "video_embeddings" in multimodal_embeddings:
            inputs_embeds = merge_multimodal_embeddings(
                input_ids, inputs_embeds,
                multimodal_embeddings["video_embeddings"],
                [video_token_id])

    return inputs_embeds

merge_multimodal_embeddings(L114)是合并的核心实现。placeholder_token_id 可以是列表——此时使用 torch.isin 匹配所有占位 token,按序列中出现顺序进行切片替换:

minimax_m3_tiny.py — merge_multimodal_embeddings L114-L185
def merge_multimodal_embeddings(
    input_ids, inputs_embeds, multimodal_embeddings, placeholder_token_id):
    """
    文本序列结构示意:
    [text ... | 200029 | 200025×N | 200030 | text ... | 200029 | 200026×M | 200030 | ...]
              ↑ start   ↑ image tokens ↑ end          ↑ start   ↑ video tokens ↑ end

    placeholder_token_id 可以是列表(如 [200025]),
    通过 torch.isin 匹配所有占位位置,用视觉 embedding 逐个替换。
    """
    if isinstance(placeholder_token_id, list):
        placeholder_token_id = torch.tensor(placeholder_token_id, device=input_ids.device)
        return _merge_multimodal_embeddings(
            inputs_embeds,
            torch.isin(input_ids, placeholder_token_id),
            multimodal_embeddings)

def _merge_multimodal_embeddings(inputs_embeds, is_multimodal, multimodal_embeddings):
    flattened = _flatten_embeddings(multimodal_embeddings)
    # 直接原地替换
    inputs_embeds[is_multimodal] = flattened
    return inputs_embeds
Start/End Token 的作用

每个图像/视频 token 组由 200029]<]start of image[>[)和 200030]<]end of image[>[)包裹。这些标记 token 由 _get_prompt_updates(L2343)在 prompt 展开时注入,帮助语言模型识别视觉信息的边界。图像和视频使用相同的 start/end token ID。

13

MiniMaxVLProcessor:时间维度 Prompt 压缩

MiniMaxVLProcessor.__call__(L1844)在 HF processor 层面处理视频帧的时间压缩——当 temporal_patch_size=2 时,每 2 帧中只保留第 1 帧的 ]<]video[>[ token,第 2 帧的 token 和对应的时间戳被删除:

minimax_m3_tiny.py — MiniMaxVLProcessor 时间压缩逻辑 L1966-L2032
# 3D Conv 时间戳处理:跟踪当前 video_seg 中的帧计数器
frame_counter_in_current_video_seg = 0
current_video_seg_index = 0
last_timestamp_text = None

# 计算时间压缩后的 video_segments
video_segments_after_temporal_merge = [
    (seg + self.temporal_patch_size - 1) // self.temporal_patch_size
    for seg in video_segments
]

# 遍历 prompt 中的每个 token
for i, _sample in enumerate(split_text):
    if _sample == self.video_token:
        if has_video_item and video_segments is not None:
            # 只保留每 temporal_patch_size 帧中的第一帧
            if frame_counter_in_current_video_seg % self.temporal_patch_size == 0:
                # 保留这个 video token
                final_text += _sample
                video_index += 1
            else:
                # 跳过这个 video token(不添加到 final_text)
                # 同时删除被跳过帧对应的时间戳
                if last_timestamp_text is not None:
                    final_text = final_text[:-len(last_timestamp_text)]
                    last_timestamp_text = None

            frame_counter_in_current_video_seg += 1
            if frame_counter_in_current_video_seg == video_segments[current_video_seg_index]:
                current_video_seg_index += 1
                frame_counter_in_current_video_seg = 0

这段逻辑确保 prompt 中的 ]<]video[>[ 数量与 3D Conv 压缩后的帧数一致。例如原始 8 帧视频,temporal_patch_size=2,prompt 中只保留 4 个 video token,每个对应 2 帧的压缩表示。

video_segments 在 vLLM 中的传递链路

video_segmentsMiniMaxVLMultiModalDataParser._parse_video_data(L1752)中提取(需设置 USE_VIDEO_SEGMENTS=True 环境变量),经 MiniMaxVLProcessor.__call__ 写入 BatchFeature,再通过 _cached_apply_hf_processor(L2260)传入 mm_kwargs,最终在 _parse_and_validate_video_input(L2873)中被提取并传给 vision_tower

14

Forward 主流程与 Dense/Sparse 后端

forward(L3081)是模型的主入口,协调视觉编码和语言模型前向传播:

minimax_m3_tiny.py — forward L3081-L3112
def forward(self, input_ids, positions, intermediate_tensors=None,
            inputs_embeds=None, **kwargs):

    if intermediate_tensors is not None:
        inputs_embeds = None
    elif inputs_embeds is None:
        # Step 1: 获取视觉 embedding(分别处理 image 和 video)
        multimodal_embeddings = self.get_multimodal_embeddings(**kwargs)
        # Step 2: 合并到文本 embedding
        inputs_embeds = self.get_input_embeddings(
            input_ids, multimodal_embeddings)

    # Step 3: 语言模型前向传播
    hidden_states = self.language_model.model(
        input_ids, positions, intermediate_tensors,
        inputs_embeds=inputs_embeds, **kwargs)

    return hidden_states

get_multimodal_embeddings(L3029)按模态分别调用 _process_image_input_process_video_input,返回值是一个 dict:

minimax_m3_tiny.py — get_multimodal_embeddings L3029-L3079
def get_multimodal_embeddings(self, **kwargs):
    image_input = self._parse_and_validate_image_input(**kwargs)
    video_input = self._parse_and_validate_video_input(**kwargs)

    multimodal_embeddings = {}

    if image_input is not None:
        embeddings = self._process_image_input(image_input)
        multimodal_embeddings["image_embeddings"] = embeddings

    if video_input is not None:
        embeddings = self._process_video_input(video_input)
        multimodal_embeddings["video_embeddings"] = embeddings

    return multimodal_embeddings

14.1 Dense vs Sparse 语言模型后端

语言模型后端的选择在 __init__ 时根据 config.text_config.sparse_attention_config 完成:

graph TD CONFIG["text_config.sparse_attention_config"] NONE["None 或
use_sparse_attention=False"] SPARSE["use_sparse_attention=True"] DENSE["MiniMaxM2ForCausalLM
标准全注意力 + MoE FFN"] SPARSEM["MiniMaxM2SparseForCausalLM
稀疏注意力 + ConstCacheManager"] CONFIG --> NONE --> DENSE CONFIG --> SPARSE --> SPARSEM

Sparse 后端特有的 ConstCacheManager 支持固定大小 KV-cache,通过 get_seqlen_agnostic_capture_inputs(L3173)和 copy_inputs_before_cuda_graphs(L3180)与 vLLM 的 CUDA Graph 机制集成:

minimax_m3_tiny.py — Sparse 后端与 CUDA Graph 集成 L3173-L3185
def get_seqlen_agnostic_capture_inputs(self, batch_size):
    # 只有 Sparse 模型有 const_cache
    if self.is_sparse_attention_model:
        return self.language_model.model.const_cache \
            .get_seqlen_agnostic_capture_inputs(batch_size)
    return None

def copy_inputs_before_cuda_graphs(self, input_buffers, **kwargs):
    if self.is_sparse_attention_model:
        return self.language_model.model.const_cache \
            .copy_inputs_before_cuda_graphs(input_buffers, **kwargs)
    return None

14.2 完整数据流全景

flowchart TD subgraph INPUT["输入"] A["input_ids + pixel_values + pixel_values_videos"] end subgraph PARSE["解析 & 验证"] P1["_parse_and_validate_image_input
6D/5D/list → [C,T,H,W] 列表"] P2["_parse_and_validate_video_input
flatten + video_segments 提取"] end subgraph IMAGE["图像管线"] I1["vision_tower
无 VPA · 2D RoPE"] I2["multi_modal_projector"] I3["按图分割 + patch_merge_mlp"] end subgraph VIDEO["视频管线"] V1["vision_tower
VPA · 3D Conv · 3D RoPE"] V2["multi_modal_projector"] V3["按帧分割 + patch_merge_mlp"] end subgraph MERGE["融合"] M1["embed_input_ids
文本 embedding"] M2["merge image_embeddings
token_id=200025"] M3["merge video_embeddings
token_id=200026"] end subgraph LM["语言模型"] L1["MiniMaxM2ForCausalLM
或 MiniMaxM2SparseForCausalLM"] L2["compute_logits"] end A --> P1 --> I1 --> I2 --> I3 --> M2 A --> P2 --> V1 --> V2 --> V3 --> M3 A --> M1 --> M2 --> M3 --> L1 --> L2