Skip to content

vllm.model_executor.models.kimi_k25_vit

Vision tower implementation for Kimi-K2.5 model.

This module provides the vision encoder components for Kimi-K2.5, including 3D patch embedding, RoPE position embedding, and temporal pooling for video chunks.

Classes:

Functions:

KimiK25MultiModalProjector

Bases: Module

Multi-modal projector with patch merging for Kimi-K2.5.

Source code in vllm/model_executor/models/kimi_k25_vit.py
class KimiK25MultiModalProjector(nn.Module):
    """Multi-modal projector with patch merging for Kimi-K2.5."""

    def __init__(
        self,
        config: KimiK25VisionConfig,
        use_data_parallel: bool = False,
        quant_config: QuantizationConfig | None = None,
        prefix: str = "",
    ):
        super().__init__()
        self.use_data_parallel = use_data_parallel
        self.mm_projector_type = getattr(config, "mm_projector_type", "patchmerger")

        # Hidden size after patch merging
        merge_h, merge_w = config.merge_kernel_size
        self.hidden_size = config.hidden_size * merge_h * merge_w

        if self.mm_projector_type == "patchmergerv2":
            self.linear_1 = ReplicatedLinear(
                self.hidden_size,
                self.hidden_size,
                bias=False,
                quant_config=quant_config,
                prefix=f"{prefix}.linear_1",
            )
            self.linear_2 = ReplicatedLinear(
                self.hidden_size,
                getattr(config, "text_hidden_size", config.mm_hidden_size),
                bias=False,
                quant_config=quant_config,
                prefix=f"{prefix}.linear_2",
            )
            self.post_norm = torch.nn.RMSNorm(
                getattr(config, "text_hidden_size", config.mm_hidden_size),
                eps=config.projector_ln_eps,
            )
            self.act = GELUActivation()
            return

        self.pre_norm = torch.nn.LayerNorm(config.hidden_size, eps=1e-5)
        self.linear_1 = ReplicatedLinear(
            self.hidden_size,
            self.hidden_size,
            bias=True,
            quant_config=quant_config,
            prefix=f"{prefix}.linear_1",
        )
        self.linear_2 = ReplicatedLinear(
            self.hidden_size,
            config.mm_hidden_size,
            bias=True,
            quant_config=quant_config,
            prefix=f"{prefix}.linear_2",
        )
        self.act = GELUActivation()

    def forward(self, image_features: torch.Tensor) -> torch.Tensor:
        if self.mm_projector_type == "patchmergerv2":
            hidden_states = image_features.view(image_features.shape[0], -1)
            hidden_states, _ = self.linear_1(hidden_states)
            hidden_states = self.act(hidden_states)
            hidden_states, _ = self.linear_2(hidden_states)
            return self.post_norm(hidden_states)

        hidden_states = self.pre_norm(image_features).view(-1, self.hidden_size)
        hidden_states, _ = self.linear_1(hidden_states)
        hidden_states = self.act(hidden_states)
        hidden_states, _ = self.linear_2(hidden_states)
        return hidden_states

Learnable2DInterpPosEmbDivided_fixed

Bases: Module

2D learnable position embedding with temporal extension.

Source code in vllm/model_executor/models/kimi_k25_vit.py
class Learnable2DInterpPosEmbDivided_fixed(nn.Module):
    """2D learnable position embedding with temporal extension."""

    def __init__(
        self,
        height: int,
        width: int,
        num_frames: int,
        dim: int,
        interpolation_mode: str = "bicubic",
    ) -> None:
        super().__init__()
        self.height = height
        self.width = width
        self.num_frames = num_frames
        self.dim = dim
        self.interpolation_mode = interpolation_mode
        self.weight = nn.Parameter(torch.empty(height, width, dim))
        self.register_buffer(
            "time_weight",
            torch.from_numpy(get_1d_sincos_pos_embed(self.dim, self.num_frames))
            .float()
            .unsqueeze(1),
            persistent=False,
        )

        self.reset_parameters()

    def reset_parameters(self):
        nn.init.normal_(self.weight)

    def get_pos_embeds(self, grid_thws: torch.Tensor | list[list[int]]) -> torch.Tensor:
        pos_embs = []
        grid_thw_list = grid_thws if isinstance(grid_thws, list) else grid_thws.tolist()
        for t, h, w in grid_thw_list:
            assert t <= self.num_frames, f"t:{t} > self.num_frames:{self.num_frames}"
            if (h, w) == self.weight.shape[:-1]:
                pos_emb_2d = self.weight.flatten(end_dim=1)
            else:
                pos_emb_2d = get_rope_shape(
                    self.weight,
                    interpolation_mode=self.interpolation_mode,
                    shape=(h, w),
                )

            if t == 1:
                pos_emb_3d = pos_emb_2d
            else:
                pos_emb_3d = (
                    pos_emb_2d.unsqueeze(0).repeat(t, 1, 1) + self.time_weight[0:t]
                )

            pos_embs.append(pos_emb_3d.reshape(-1, pos_emb_3d.shape[-1]))

        return torch.cat(pos_embs)

    def forward(
        self, x: torch.Tensor, grid_thws: torch.Tensor | list[list[int]]
    ) -> torch.Tensor:
        return x + self.get_pos_embeds(grid_thws)

MLP2

Bases: Module

Two-layer MLP with tensor parallel support.

Source code in vllm/model_executor/models/kimi_k25_vit.py
class MLP2(nn.Module):
    """Two-layer MLP with tensor parallel support."""

    def __init__(
        self,
        dims: list[int],
        activation,
        bias: bool = True,
        quant_config: QuantizationConfig | None = None,
        prefix: str = "",
        use_data_parallel: bool = False,
    ):
        super().__init__()
        assert len(dims) == 3
        self.use_data_parallel = use_data_parallel
        self.fc0 = ColumnParallelLinear(
            dims[0],
            dims[1],
            bias=bias,
            quant_config=quant_config,
            prefix=maybe_prefix(prefix, "fc0"),
            disable_tp=self.use_data_parallel,
        )
        self.fc1 = RowParallelLinear(
            dims[1],
            dims[2],
            bias=bias,
            quant_config=quant_config,
            prefix=maybe_prefix(prefix, "fc1"),
            disable_tp=self.use_data_parallel,
        )
        self.activation = activation

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        x, _ = self.fc0(x)
        x = self.activation(x)
        x, _ = self.fc1(x)
        return x

MoonViT3dEncoder

Bases: Module

Full encoder stack for MoonViT 3D.

Source code in vllm/model_executor/models/kimi_k25_vit.py
class MoonViT3dEncoder(nn.Module):
    """Full encoder stack for MoonViT 3D."""

    def __init__(
        self,
        hidden_dim: int,
        num_layers: int,
        block_cfg: dict,
        video_attn_type: str = "spatial_temporal",
        quant_config: QuantizationConfig | None = None,
        prefix: str = "",
    ) -> None:
        super().__init__()

        assert video_attn_type == "spatial_temporal", (
            f'video_attn_type must be "spatial_temporal", got {video_attn_type}'
        )
        self.video_attn_type = video_attn_type
        qkv_hidden_size = block_cfg.get("qkv_hidden_size") or block_cfg["hidden_dim"]
        self.rope_2d = Rope2DPosEmbRepeated(
            qkv_hidden_size // block_cfg["num_heads"], 512, 512
        )
        self.blocks = nn.ModuleList(
            [
                MoonViTEncoderLayer(
                    **block_cfg,
                    quant_config=quant_config,
                    prefix=f"{prefix}.blocks.{layer_idx}",
                )
                for layer_idx in range(num_layers)
            ]
        )
        self.final_layernorm = _make_vision_norm(
            block_cfg.get("norm_type", "layernorm"), hidden_dim
        )

    def prepare_encoder_metadata(
        self,
        grid_thw_list: list[list[int]],
        *,
        device: torch.device,
        max_batch_size: int | None = None,
        max_seqlen_override: int | None = None,
    ) -> dict[str, torch.Tensor | None]:
        metadata: dict[str, torch.Tensor | None] = {}
        metadata["rope_freqs_cis"] = self.rope_2d.get_freqs_cis(
            grid_thw_list, device=device
        )

        grid_thw_np = np.array(grid_thw_list, dtype=np.int32)
        lengths = grid_thw_np[:, 0] * grid_thw_np[:, 1] * grid_thw_np[:, 2]
        cu_seqlens = np.concatenate(
            [np.zeros(1, dtype=np.int32), lengths.cumsum(dtype=np.int32)]
        )
        if max_batch_size is not None:
            num_seqs = len(cu_seqlens) - 1
            if num_seqs < max_batch_size:
                cu_seqlens = np.concatenate(
                    [
                        cu_seqlens,
                        np.full(
                            max_batch_size - num_seqs,
                            cu_seqlens[-1],
                            dtype=np.int32,
                        ),
                    ]
                )

        attn_backend = self.blocks[0].attn.attn_backend
        metadata["sequence_lengths"] = MMEncoderAttention.maybe_compute_seq_lens(
            attn_backend, cu_seqlens, device
        )
        max_seqlen = (
            max_seqlen_override
            if max_seqlen_override is not None
            else MMEncoderAttention.compute_max_seqlen(attn_backend, cu_seqlens)
        )
        metadata["max_seqlen"] = torch.tensor(max_seqlen, dtype=torch.int32)
        metadata["cu_seqlens"] = MMEncoderAttention.maybe_recompute_cu_seqlens(
            attn_backend,
            cu_seqlens,
            self.blocks[0].hidden_dim,
            self.blocks[0].tp_size,
            device,
        )
        return metadata

    def forward(
        self,
        hidden_states: torch.Tensor,
        grid_thws: torch.Tensor | list[list[int]] | None,
        *,
        encoder_metadata: dict[str, torch.Tensor | None] | None = None,
    ) -> torch.Tensor:
        if encoder_metadata is None:
            assert grid_thws is not None
            grid_thw_list = (
                grid_thws if isinstance(grid_thws, list) else grid_thws.tolist()
            )
            encoder_metadata = self.prepare_encoder_metadata(
                grid_thw_list, device=hidden_states.device
            )

        rope_freqs_cis = encoder_metadata["rope_freqs_cis"]
        cu_seqlens = encoder_metadata["cu_seqlens"]
        max_seqlen = encoder_metadata["max_seqlen"]
        sequence_lengths = encoder_metadata.get("sequence_lengths")
        assert rope_freqs_cis is not None
        assert cu_seqlens is not None
        assert max_seqlen is not None

        for block in self.blocks:
            hidden_states = block(
                hidden_states,
                cu_seqlens,
                rope_freqs_cis=rope_freqs_cis,
                max_seqlen=max_seqlen,
                sequence_lengths=sequence_lengths,
            )

        hidden_states = self.final_layernorm(hidden_states)

        return hidden_states

MoonViT3dPretrainedModel

Bases: Module

Main vision tower model.

Uses KimiK25VisionConfig directly from transformers_utils/configs/kimi_k25.py.

Methods:

Source code in vllm/model_executor/models/kimi_k25_vit.py
class MoonViT3dPretrainedModel(nn.Module):
    """Main vision tower model.

    Uses KimiK25VisionConfig directly from transformers_utils/configs/kimi_k25.py.
    """

    def __init__(
        self,
        config: KimiK25VisionConfig,
        quant_config: QuantizationConfig | None = None,
        prefix: str = "",
    ):
        super().__init__()
        config = deepcopy(config)
        self.config = config  # Required for run_dp_sharded_mrope_vision_model
        self.merge_kernel_size = config.merge_kernel_size
        self.patch_size = config.patch_size
        self.merge_type = config.merge_type

        self.patch_embed = MoonVision3dPatchEmbed(
            out_dim=config.hidden_size,
            patch_size=config.patch_size,
            pos_emb_height=config.init_pos_emb_height,
            pos_emb_width=config.init_pos_emb_width,
            pos_emb_time=config.init_pos_emb_time,
            pos_emb_type=config.pos_emb_type,
            patch_embed_proj_bias=getattr(config, "patch_embed_proj_bias", True),
            pos_emb_interpolation_mode=getattr(
                config, "pos_emb_interpolation_mode", "bicubic"
            ),
        )

        self.encoder = MoonViT3dEncoder(
            hidden_dim=config.hidden_size,
            num_layers=config.num_hidden_layers,
            block_cfg={
                "num_heads": config.num_attention_heads,
                "hidden_dim": config.hidden_size,
                "qkv_hidden_size": getattr(config, "qkv_hidden_size", None),
                "mlp_dim": config.intermediate_size,
                "activation": get_act_fn(
                    getattr(config, "activation_func", "gelu_pytorch_tanh")
                ),
                "attn_bias": getattr(config, "attn_bias", True),
                "norm_type": getattr(config, "norm_type", "layernorm"),
                "mlp_type": getattr(config, "mlp_type", "mlp2"),
                "linear_bias": getattr(config, "linear_bias", True),
            },
            video_attn_type=config.video_attn_type,
            quant_config=quant_config,
            prefix=maybe_prefix(prefix, "encoder"),
        )

    def forward(
        self,
        pixel_values: torch.Tensor,
        grid_thws: torch.Tensor | list[list[int]] | None,
        *,
        encoder_metadata: dict[str, torch.Tensor | None] | None = None,
    ) -> torch.Tensor:
        """
        Args:
            pixel_values (torch.Tensor): The input pixel values.
            grid_thws (torch.Tensor): Temporal, height and width.

        Returns:
            torch.Tensor: The output tokens.
        """
        if encoder_metadata is not None and "pos_embeds" in encoder_metadata:
            hidden_states = self.patch_embed(
                pixel_values,
                None,
                pos_embeds=encoder_metadata["pos_embeds"],
            )
            hidden_states = self.encoder(
                hidden_states,
                None,
                encoder_metadata=encoder_metadata,
            )
            merge_gather_idx = encoder_metadata["merge_gather_idx"]
            assert merge_gather_idx is not None
            return tpool_patch_merger_packed(hidden_states, merge_gather_idx)

        assert grid_thws is not None
        grid_thw_list = grid_thws if isinstance(grid_thws, list) else grid_thws.tolist()
        if encoder_metadata is None:
            encoder_metadata = self.encoder.prepare_encoder_metadata(
                grid_thw_list, device=pixel_values.device
            )

        hidden_states = self.patch_embed(pixel_values, grid_thw_list)
        hidden_states = self.encoder(
            hidden_states,
            grid_thw_list,
            encoder_metadata=encoder_metadata,
        )
        if (
            self.merge_type == "sd2_tpool"
        ):  # spatial downsampling 2x with temporal pooling all
            hidden_states = tpool_patch_merger(
                hidden_states, grid_thw_list, merge_kernel_size=self.merge_kernel_size
            )
        else:
            raise NotImplementedError(f"Not support {self.merge_type}")

        return hidden_states

    def prepare_encoder_cudagraph_metadata(
        self,
        grid_thw_list: list[list[int]],
        *,
        max_batch_size: int,
        max_seqlen_override: int | None = None,
        device: torch.device,
    ) -> dict[str, torch.Tensor | None]:
        """Precompute fixed-buffer metadata for image encoder CUDA graphs."""
        grid_thw_list = [list(map(int, grid)) for grid in grid_thw_list]
        metadata = self.encoder.prepare_encoder_metadata(
            grid_thw_list,
            device=device,
            max_batch_size=max_batch_size,
            max_seqlen_override=max_seqlen_override,
        )
        metadata["pos_embeds"] = self.patch_embed.pos_emb.get_pos_embeds(
            grid_thw_list
        ).to(device=device)
        merge_gather_idx = build_image_merge_gather_idx(
            grid_thw_list, self.merge_kernel_size
        )
        metadata["merge_gather_idx"] = torch.from_numpy(merge_gather_idx).to(
            device=device, non_blocking=True
        )
        return metadata

forward(pixel_values, grid_thws, *, encoder_metadata=None)

Parameters:

  • pixel_values

    (Tensor) –

    The input pixel values.

  • grid_thws

    (Tensor) –

    Temporal, height and width.

Returns:

  • Tensor

    torch.Tensor: The output tokens.

Source code in vllm/model_executor/models/kimi_k25_vit.py
def forward(
    self,
    pixel_values: torch.Tensor,
    grid_thws: torch.Tensor | list[list[int]] | None,
    *,
    encoder_metadata: dict[str, torch.Tensor | None] | None = None,
) -> torch.Tensor:
    """
    Args:
        pixel_values (torch.Tensor): The input pixel values.
        grid_thws (torch.Tensor): Temporal, height and width.

    Returns:
        torch.Tensor: The output tokens.
    """
    if encoder_metadata is not None and "pos_embeds" in encoder_metadata:
        hidden_states = self.patch_embed(
            pixel_values,
            None,
            pos_embeds=encoder_metadata["pos_embeds"],
        )
        hidden_states = self.encoder(
            hidden_states,
            None,
            encoder_metadata=encoder_metadata,
        )
        merge_gather_idx = encoder_metadata["merge_gather_idx"]
        assert merge_gather_idx is not None
        return tpool_patch_merger_packed(hidden_states, merge_gather_idx)

    assert grid_thws is not None
    grid_thw_list = grid_thws if isinstance(grid_thws, list) else grid_thws.tolist()
    if encoder_metadata is None:
        encoder_metadata = self.encoder.prepare_encoder_metadata(
            grid_thw_list, device=pixel_values.device
        )

    hidden_states = self.patch_embed(pixel_values, grid_thw_list)
    hidden_states = self.encoder(
        hidden_states,
        grid_thw_list,
        encoder_metadata=encoder_metadata,
    )
    if (
        self.merge_type == "sd2_tpool"
    ):  # spatial downsampling 2x with temporal pooling all
        hidden_states = tpool_patch_merger(
            hidden_states, grid_thw_list, merge_kernel_size=self.merge_kernel_size
        )
    else:
        raise NotImplementedError(f"Not support {self.merge_type}")

    return hidden_states

prepare_encoder_cudagraph_metadata(grid_thw_list, *, max_batch_size, max_seqlen_override=None, device)

Precompute fixed-buffer metadata for image encoder CUDA graphs.

Source code in vllm/model_executor/models/kimi_k25_vit.py
def prepare_encoder_cudagraph_metadata(
    self,
    grid_thw_list: list[list[int]],
    *,
    max_batch_size: int,
    max_seqlen_override: int | None = None,
    device: torch.device,
) -> dict[str, torch.Tensor | None]:
    """Precompute fixed-buffer metadata for image encoder CUDA graphs."""
    grid_thw_list = [list(map(int, grid)) for grid in grid_thw_list]
    metadata = self.encoder.prepare_encoder_metadata(
        grid_thw_list,
        device=device,
        max_batch_size=max_batch_size,
        max_seqlen_override=max_seqlen_override,
    )
    metadata["pos_embeds"] = self.patch_embed.pos_emb.get_pos_embeds(
        grid_thw_list
    ).to(device=device)
    merge_gather_idx = build_image_merge_gather_idx(
        grid_thw_list, self.merge_kernel_size
    )
    metadata["merge_gather_idx"] = torch.from_numpy(merge_gather_idx).to(
        device=device, non_blocking=True
    )
    return metadata

MoonViTEncoderLayer

Bases: Module

Single encoder layer for MoonViT with TP/DP support.

Methods:

Source code in vllm/model_executor/models/kimi_k25_vit.py
class MoonViTEncoderLayer(nn.Module):
    """Single encoder layer for MoonViT with TP/DP support."""

    def __init__(
        self,
        num_heads: int,
        hidden_dim: int,
        mlp_dim: int,
        quant_config: QuantizationConfig | None = None,
        prefix: str = "",
        *,
        activation=F.gelu,
        attn_bias: bool = False,
        qkv_hidden_size: int | None = None,
        norm_type: str = "layernorm",
        mlp_type: str = "mlp2",
        linear_bias: bool = True,
    ):
        super().__init__()
        self.use_data_parallel = is_vit_use_data_parallel(num_heads)

        self.num_heads = num_heads
        self.hidden_dim = hidden_dim
        self.qkv_hidden_size = (
            hidden_dim if qkv_hidden_size is None else qkv_hidden_size
        )
        self.hidden_size_per_attention_head = self.qkv_hidden_size // self.num_heads
        self.tp_size = (
            1 if self.use_data_parallel else get_tensor_model_parallel_world_size()
        )
        self.num_attention_heads_per_partition = divide(num_heads, self.tp_size)

        self.norm0 = _make_vision_norm(norm_type, hidden_dim)
        self.norm1 = _make_vision_norm(norm_type, hidden_dim)
        if mlp_type != "mlp2":
            raise NotImplementedError(f"Not support mlp_type: {mlp_type}")
        self.mlp = MLP2(
            [hidden_dim, mlp_dim, hidden_dim],
            activation,
            bias=linear_bias,
            quant_config=quant_config,
            prefix=f"{prefix}.mlp",
            use_data_parallel=self.use_data_parallel,
        )
        self.wqkv = QKVParallelLinear(
            hidden_size=hidden_dim,
            head_size=self.hidden_size_per_attention_head,
            total_num_heads=num_heads,
            total_num_kv_heads=num_heads,
            bias=attn_bias,
            quant_config=quant_config,
            prefix=f"{prefix}.wqkv",
            disable_tp=self.use_data_parallel,
        )
        self.wo = RowParallelLinear(
            self.qkv_hidden_size,
            hidden_dim,
            bias=attn_bias,
            quant_config=quant_config,
            prefix=f"{prefix}.wo",
            disable_tp=self.use_data_parallel,
        )
        self.attn = MMEncoderAttention(
            num_heads=self.num_attention_heads_per_partition,
            head_size=self.hidden_size_per_attention_head,
            scale=self.hidden_size_per_attention_head**-0.5,
            prefix=f"{prefix}.attn",
        )
        self.apply_rotary_emb = ApplyRotaryEmb(
            enforce_enable=True,
            is_neox_style=False,
            enable_fp32_compute=True,
        )

    def attention_qkvpacked(
        self,
        x: torch.Tensor,
        cu_seqlens: torch.Tensor,
        rope_freqs_cis: torch.Tensor,
        max_seqlen: torch.Tensor | None = None,
        sequence_lengths: torch.Tensor | None = None,
    ):
        """Compute self-attention with packed QKV.

        Args:
            x (torch.Tensor): (seqlen, hidden_dim)
            cu_seqlens (torch.Tensor): cumulative sequence lengths
        """
        seq_length = x.size(0)
        xqkv, _ = self.wqkv(x)

        qkv_shape = xqkv.size()[:-1] + (
            3,
            self.num_attention_heads_per_partition,
            self.hidden_size_per_attention_head,
        )
        # xqkv: (seqlen, 3, nheads, headdim)
        xqkv = xqkv.view(*qkv_shape)
        xq, xk, xv = torch.unbind(xqkv, dim=-3)

        _apply_rope_input_validation(xq, rope_freqs_cis)
        _apply_rope_input_validation(xk, rope_freqs_cis)
        rope_cos = rope_freqs_cis.real.contiguous()
        rope_sin = rope_freqs_cis.imag.contiguous()
        xq = self.apply_rotary_emb(xq, rope_cos, rope_sin)
        xk = self.apply_rotary_emb(xk, rope_cos, rope_sin)

        if max_seqlen is None:
            max_seqlen = (cu_seqlens[1:] - cu_seqlens[:-1]).max()
        attn_out = self.attn(
            xq.unsqueeze(0),
            xk.unsqueeze(0),
            xv.unsqueeze(0),
            cu_seqlens=cu_seqlens,
            max_seqlen=max_seqlen,
            sequence_lengths=sequence_lengths,
        )
        attn_out = attn_out.reshape(
            seq_length,
            self.num_attention_heads_per_partition
            * self.hidden_size_per_attention_head,
        )
        attn_out, _ = self.wo(attn_out)
        return attn_out

    def forward(
        self,
        hidden_states: torch.Tensor,
        cu_seqlens: torch.Tensor,
        rope_freqs_cis: torch.Tensor,
        max_seqlen: torch.Tensor | None = None,
        sequence_lengths: torch.Tensor | None = None,
    ):
        residual = hidden_states
        hidden_states = self.norm0(hidden_states)

        hidden_states = self.attention_qkvpacked(
            hidden_states,
            cu_seqlens,
            rope_freqs_cis,
            max_seqlen=max_seqlen,
            sequence_lengths=sequence_lengths,
        )
        hidden_states = residual + hidden_states

        residual = hidden_states
        hidden_states = self.norm1(hidden_states)
        hidden_states = self.mlp(hidden_states)
        hidden_states = residual + hidden_states

        return hidden_states

attention_qkvpacked(x, cu_seqlens, rope_freqs_cis, max_seqlen=None, sequence_lengths=None)

Compute self-attention with packed QKV.

Parameters:

  • x

    (Tensor) –

    (seqlen, hidden_dim)

  • cu_seqlens

    (Tensor) –

    cumulative sequence lengths

Source code in vllm/model_executor/models/kimi_k25_vit.py
def attention_qkvpacked(
    self,
    x: torch.Tensor,
    cu_seqlens: torch.Tensor,
    rope_freqs_cis: torch.Tensor,
    max_seqlen: torch.Tensor | None = None,
    sequence_lengths: torch.Tensor | None = None,
):
    """Compute self-attention with packed QKV.

    Args:
        x (torch.Tensor): (seqlen, hidden_dim)
        cu_seqlens (torch.Tensor): cumulative sequence lengths
    """
    seq_length = x.size(0)
    xqkv, _ = self.wqkv(x)

    qkv_shape = xqkv.size()[:-1] + (
        3,
        self.num_attention_heads_per_partition,
        self.hidden_size_per_attention_head,
    )
    # xqkv: (seqlen, 3, nheads, headdim)
    xqkv = xqkv.view(*qkv_shape)
    xq, xk, xv = torch.unbind(xqkv, dim=-3)

    _apply_rope_input_validation(xq, rope_freqs_cis)
    _apply_rope_input_validation(xk, rope_freqs_cis)
    rope_cos = rope_freqs_cis.real.contiguous()
    rope_sin = rope_freqs_cis.imag.contiguous()
    xq = self.apply_rotary_emb(xq, rope_cos, rope_sin)
    xk = self.apply_rotary_emb(xk, rope_cos, rope_sin)

    if max_seqlen is None:
        max_seqlen = (cu_seqlens[1:] - cu_seqlens[:-1]).max()
    attn_out = self.attn(
        xq.unsqueeze(0),
        xk.unsqueeze(0),
        xv.unsqueeze(0),
        cu_seqlens=cu_seqlens,
        max_seqlen=max_seqlen,
        sequence_lengths=sequence_lengths,
    )
    attn_out = attn_out.reshape(
        seq_length,
        self.num_attention_heads_per_partition
        * self.hidden_size_per_attention_head,
    )
    attn_out, _ = self.wo(attn_out)
    return attn_out

MoonVision3dPatchEmbed

Bases: Module

3D patch embedding for vision tower.

Source code in vllm/model_executor/models/kimi_k25_vit.py
class MoonVision3dPatchEmbed(nn.Module):
    """3D patch embedding for vision tower."""

    def __init__(
        self,
        out_dim: int,
        in_dim: int = 3,
        patch_size: int | tuple[int, int] = (14, 14),
        pos_emb_height: int = 14,
        pos_emb_width: int = 14,
        pos_emb_time: int = 4,
        pos_emb_type: str = "divided_fixed",
        patch_embed_proj_bias: bool = True,
        pos_emb_interpolation_mode: str = "bicubic",
    ):
        super().__init__()
        assert isinstance(patch_size, int | Sequence), (
            f"Invalid patch_size type: {type(patch_size)}"
        )
        if isinstance(patch_size, int):
            patch_size = (patch_size, patch_size)
        assert len(patch_size) == 2, (
            f"Expected patch_size to be a tuple of 2, got {patch_size}"
        )
        self.patch_size = patch_size

        self.proj = nn.Conv2d(
            in_dim,
            out_dim,
            kernel_size=patch_size,
            stride=patch_size,
            bias=patch_embed_proj_bias,
        )

        if pos_emb_type == "divided_fixed":
            self.pos_emb = Learnable2DInterpPosEmbDivided_fixed(
                height=pos_emb_height,
                width=pos_emb_width,
                num_frames=pos_emb_time,
                dim=out_dim,
                interpolation_mode=pos_emb_interpolation_mode,
            )
        else:
            raise NotImplementedError(f"Not support pos_emb_type: {pos_emb_type}")

    def forward(
        self,
        x: torch.Tensor,
        grid_thws: torch.Tensor | list[list[int]] | None,
        *,
        pos_embeds: torch.Tensor | None = None,
    ) -> torch.Tensor:
        x = self._proj(x).view(x.size(0), -1)
        if pos_embeds is not None:
            return x + pos_embeds
        assert grid_thws is not None
        return self.pos_emb(x, grid_thws)

    def _proj(self, x: torch.Tensor) -> torch.Tensor:
        # MIOpen conv2d intermittently fails under load on ROCm; use aiter Triton.
        if current_platform.is_rocm() and x.dtype in (torch.float16, torch.bfloat16):
            from aiter.ops.triton.conv.conv2d import conv2d

            return conv2d(
                x,
                self.proj.weight,
                self.proj.bias,
                stride=self.patch_size,
                layout="nchw",
            )
        return self.proj(x)

Rope2DPosEmbRepeated

Bases: Module

2D rotary position embedding with multi-resolution support.

Methods:

Source code in vllm/model_executor/models/kimi_k25_vit.py
class Rope2DPosEmbRepeated(nn.Module):
    """2D rotary position embedding with multi-resolution support."""

    def __init__(self, dim: int, max_height: int, max_width: int, theta_base=10000):
        super().__init__()
        self.dim = dim
        assert self.dim % 4 == 0, "dim must be divisible by 4"
        self.max_height = max_height
        self.max_width = max_width
        self.theta_base = theta_base

    def extra_repr(self):
        return (
            f"dim={self.dim}, max_height={self.max_height}, "
            f"max_width={self.max_width}, theta_base={self.theta_base}"
        )

    def _precompute_freqs_cis(self, device: torch.device) -> torch.Tensor:
        """Calculate the cis(freqs) for each position in the 2D grid."""
        N = self.max_height * self.max_width
        flat_pos = torch.arange(0, N).float().to(device)
        x_pos = flat_pos % self.max_width
        y_pos = flat_pos // self.max_width
        dim_range = (
            torch.arange(0, self.dim, 4)[: (self.dim // 4)].float().to(device)
        )  # C/4
        freqs = 1.0 / (self.theta_base ** (dim_range / self.dim))
        x_freqs = torch.outer(x_pos, freqs).float()  # N, C/4
        y_freqs = torch.outer(y_pos, freqs).float()  # N, C/4
        x_cis = torch.polar(torch.ones_like(x_freqs), x_freqs)  # N, C/4
        y_cis = torch.polar(torch.ones_like(y_freqs), y_freqs)  # N, C/4
        # N, C/4, 2
        freqs_cis = torch.cat(
            [x_cis.unsqueeze(dim=-1), y_cis.unsqueeze(dim=-1)], dim=-1
        )
        # max_height, max_width, C/2
        freqs_cis = freqs_cis.reshape(self.max_height, self.max_width, -1)
        return freqs_cis

    def get_freqs_cis(
        self, grid_thws: torch.Tensor | list[list[int]], device: torch.device
    ) -> torch.Tensor:
        """
        Args:
            grid_thws (torch.Tensor): grid time, height and width

        Returns:
            freqs_cis: tensor of shape (sum(t * height * width), dim//2)
        """
        if not hasattr(self, "freqs_cis"):
            self.register_buffer(
                "freqs_cis", self._precompute_freqs_cis(device), persistent=False
            )

        shapes = grid_thws if isinstance(grid_thws, list) else grid_thws.tolist()
        assert all(
            1 <= h <= self.max_height and 1 <= w <= self.max_width for t, h, w in shapes
        ), (
            shapes,
            self.max_height,
            self.max_width,
        )
        freqs_cis = torch.cat(
            [
                self.freqs_cis[:h, :w].reshape(-1, self.dim // 2).repeat(t, 1)
                for t, h, w in shapes
            ],
            dim=0,
        )
        return freqs_cis

_precompute_freqs_cis(device)

Calculate the cis(freqs) for each position in the 2D grid.

Source code in vllm/model_executor/models/kimi_k25_vit.py
def _precompute_freqs_cis(self, device: torch.device) -> torch.Tensor:
    """Calculate the cis(freqs) for each position in the 2D grid."""
    N = self.max_height * self.max_width
    flat_pos = torch.arange(0, N).float().to(device)
    x_pos = flat_pos % self.max_width
    y_pos = flat_pos // self.max_width
    dim_range = (
        torch.arange(0, self.dim, 4)[: (self.dim // 4)].float().to(device)
    )  # C/4
    freqs = 1.0 / (self.theta_base ** (dim_range / self.dim))
    x_freqs = torch.outer(x_pos, freqs).float()  # N, C/4
    y_freqs = torch.outer(y_pos, freqs).float()  # N, C/4
    x_cis = torch.polar(torch.ones_like(x_freqs), x_freqs)  # N, C/4
    y_cis = torch.polar(torch.ones_like(y_freqs), y_freqs)  # N, C/4
    # N, C/4, 2
    freqs_cis = torch.cat(
        [x_cis.unsqueeze(dim=-1), y_cis.unsqueeze(dim=-1)], dim=-1
    )
    # max_height, max_width, C/2
    freqs_cis = freqs_cis.reshape(self.max_height, self.max_width, -1)
    return freqs_cis

get_freqs_cis(grid_thws, device)

Parameters:

  • grid_thws

    (Tensor) –

    grid time, height and width

Returns:

  • freqs_cis ( Tensor ) –

    tensor of shape (sum(t * height * width), dim//2)

Source code in vllm/model_executor/models/kimi_k25_vit.py
def get_freqs_cis(
    self, grid_thws: torch.Tensor | list[list[int]], device: torch.device
) -> torch.Tensor:
    """
    Args:
        grid_thws (torch.Tensor): grid time, height and width

    Returns:
        freqs_cis: tensor of shape (sum(t * height * width), dim//2)
    """
    if not hasattr(self, "freqs_cis"):
        self.register_buffer(
            "freqs_cis", self._precompute_freqs_cis(device), persistent=False
        )

    shapes = grid_thws if isinstance(grid_thws, list) else grid_thws.tolist()
    assert all(
        1 <= h <= self.max_height and 1 <= w <= self.max_width for t, h, w in shapes
    ), (
        shapes,
        self.max_height,
        self.max_width,
    )
    freqs_cis = torch.cat(
        [
            self.freqs_cis[:h, :w].reshape(-1, self.dim // 2).repeat(t, 1)
            for t, h, w in shapes
        ],
        dim=0,
    )
    return freqs_cis

build_image_merge_gather_idx(grid_thws, merge_kernel_size)

Build packed spatial-merge indices for image-only CUDA graphs.

Source code in vllm/model_executor/models/kimi_k25_vit.py
def build_image_merge_gather_idx(
    grid_thws: list[list[int]] | list[tuple[int, int, int]],
    merge_kernel_size: tuple[int, int],
) -> np.ndarray:
    """Build packed spatial-merge indices for image-only CUDA graphs."""
    kh, kw = merge_kernel_size
    parts: list[np.ndarray] = []
    offset = 0
    for t, h, w in grid_thws:
        if t != 1:
            raise ValueError("Image encoder CUDA graphs require grid T == 1")
        idx = np.arange(h * w, dtype=np.int64).reshape(h, w)
        idx = idx.reshape(h // kh, kh, w // kw, kw)
        parts.append(idx.transpose(0, 2, 1, 3).reshape(-1, kh * kw) + offset)
        offset += h * w
    if not parts:
        return np.empty((0, kh * kw), dtype=np.int64)
    return np.concatenate(parts)

get_1d_sincos_pos_embed(embed_dim, t_size, cls_token=False)

Generate 1D sincos positional embedding.

Source code in vllm/model_executor/models/kimi_k25_vit.py
def get_1d_sincos_pos_embed(embed_dim, t_size, cls_token=False):
    """Generate 1D sincos positional embedding."""
    grid_t = np.arange(t_size, dtype=np.float32)
    pos_embed = get_1d_sincos_pos_embed_from_grid(embed_dim, grid_t)
    if cls_token:
        pos_embed = np.concatenate([np.zeros([1, embed_dim]), pos_embed], axis=0)
    return pos_embed

get_1d_sincos_pos_embed_from_grid(embed_dim, pos)

Generate 1D sincos positional embedding from grid positions.

Source code in vllm/model_executor/models/kimi_k25_vit.py
def get_1d_sincos_pos_embed_from_grid(embed_dim, pos):
    """Generate 1D sincos positional embedding from grid positions."""
    assert embed_dim % 2 == 0
    omega = np.arange(embed_dim // 2, dtype=np.float32)
    omega /= embed_dim / 2.0
    omega = 1.0 / 10000**omega  # (D/2,)

    pos = pos.reshape(-1)  # (M,)
    out = np.einsum("m,d->md", pos, omega)  # (M, D/2), outer product

    emb_sin = np.sin(out)  # (M, D/2)
    emb_cos = np.cos(out)  # (M, D/2)

    emb = np.concatenate([emb_sin, emb_cos], axis=1)  # (M, D)
    return emb

mm_projector_forward(mm_projector, vt_output)

Apply MM projector to vision tower outputs.

Source code in vllm/model_executor/models/kimi_k25_vit.py
@torch.inference_mode()
def mm_projector_forward(mm_projector: torch.nn.Module, vt_output: list[torch.Tensor]):
    """Apply MM projector to vision tower outputs."""
    num_embedding_list = [x.shape[0] for x in vt_output]
    batched = torch.cat(vt_output, dim=0)
    projector_dtype = next(mm_projector.parameters()).dtype
    if batched.dtype != projector_dtype:
        batched = batched.to(projector_dtype)
    proj_out = mm_projector(batched)
    proj_out = proj_out.reshape(-1, proj_out.shape[-1])
    proj_out = torch.split(proj_out, num_embedding_list)
    return proj_out

tpool_patch_merger(x, grid_thws, merge_kernel_size=(2, 2))

Temporal pooling patch merger.

Source code in vllm/model_executor/models/kimi_k25_vit.py
def tpool_patch_merger(
    x: torch.Tensor,
    grid_thws: torch.Tensor | list[list[int]],
    merge_kernel_size: tuple[int, int] = (2, 2),
) -> list[torch.Tensor]:
    """Temporal pooling patch merger."""
    kh, kw = merge_kernel_size
    grid_thw_list = grid_thws if isinstance(grid_thws, list) else grid_thws.tolist()
    lengths = [t * h * w for t, h, w in grid_thw_list]
    seqs = x.split(lengths, dim=0)

    outputs = []
    for seq, (t, h, w) in zip(seqs, grid_thw_list):
        nh, nw = h // kh, w // kw
        # Reshape: (t*h*w, d) -> (t, nh, kh, nw, kw, d)
        v = seq.view(t, nh, kh, nw, kw, -1)
        # Temporal pooling first (reduces tensor size before permute)
        v = v.mean(dim=0)  # (nh, kh, nw, kw, d)
        # Spatial rearrangement: (nh, kh, nw, kw, d) -> (nh, nw, kh, kw, d)
        out = v.permute(0, 2, 1, 3, 4).reshape(nh * nw, kh * kw, -1)
        outputs.append(out)

    return outputs

tpool_patch_merger_packed(x, merge_gather_idx)

Apply the image-only spatial merge using precomputed tensor indices.

Source code in vllm/model_executor/models/kimi_k25_vit.py
def tpool_patch_merger_packed(
    x: torch.Tensor,
    merge_gather_idx: torch.Tensor,
) -> torch.Tensor:
    """Apply the image-only spatial merge using precomputed tensor indices."""
    return x[merge_gather_idx]

vision_tower_forward(vision_tower, pixel_values, grid_thw, mm_projector, use_data_parallel)

DP-sharded vision tower forward with mrope.

Uses vLLM's standard data parallelism utility to shard the batch across available GPUs, enabling parallel processing of vision features.

Source code in vllm/model_executor/models/kimi_k25_vit.py
@torch.inference_mode()
def vision_tower_forward(
    vision_tower: Any,
    pixel_values: torch.Tensor,
    grid_thw: torch.Tensor,
    mm_projector: Any,
    use_data_parallel: bool,
) -> list[torch.Tensor]:
    """DP-sharded vision tower forward with mrope.

    Uses vLLM's standard data parallelism utility to shard the batch
    across available GPUs, enabling parallel processing of vision features.
    """
    if use_data_parallel:
        grid_thw_list = grid_thw.tolist()
        vt_outputs = run_dp_sharded_mrope_vision_model(
            vision_model=vision_tower,
            pixel_values=pixel_values,
            grid_thw_list=grid_thw_list,
            rope_type="rope_2d",
        )
    else:
        grid_thw_list = grid_thw.tolist()
        encoder_metadata = vision_tower.encoder.prepare_encoder_metadata(
            grid_thw_list, device=pixel_values.device
        )
        vt_outputs = vision_tower(
            pixel_values,
            grid_thw_list,
            encoder_metadata=encoder_metadata,
        )
    tensors = mm_projector_forward(mm_projector, list(vt_outputs))
    return list(tensors)