Skip to content

vllm.models.deepseek_v41.common.ops.fused_layout

Weight permutations for FlashMLA's mega-attention kernel.

The kernel reads Q with 16-element head-dim chunks interleaved across heads and writes O with 32-element chunks interleaved across the 8 heads of a wo_a group. Every permutation here satisfies fused = standard[perm], and is applied once to wq_b rows and wo_a columns at load time so the surrounding GEMMs produce and consume the kernel's layouts directly -- no per-step shuffle.

Functions:

_bytes_view(t)

Byte view of a 1-byte-element tensor, so fp8/ue8m0 can be gathered.

Source code in vllm/models/deepseek_v41/common/ops/fused_layout.py
def _bytes_view(t: torch.Tensor) -> torch.Tensor:
    """Byte view of a 1-byte-element tensor, so fp8/ue8m0 can be gathered."""
    return t.view(torch.uint8) if t.element_size() == 1 else t

o_fused_chunk_permutation(heads_per_group=WV_GROUP_SIZE, head_dim=HEAD_DIM)

Per-32-element-chunk form of :func:o_fused_permutation, for scales.

Source code in vllm/models/deepseek_v41/common/ops/fused_layout.py
def o_fused_chunk_permutation(
    heads_per_group: int = WV_GROUP_SIZE, head_dim: int = HEAD_DIM
) -> torch.Tensor:
    """Per-32-element-chunk form of :func:`o_fused_permutation`, for scales."""
    return o_fused_permutation(heads_per_group, head_dim)[::O_CHUNK] // O_CHUNK

o_fused_permutation(heads_per_group=WV_GROUP_SIZE, head_dim=HEAD_DIM)

fused[(c * G + h) * 32 + j] = standard[h * D + c * 32 + j] per group.

Source code in vllm/models/deepseek_v41/common/ops/fused_layout.py
def o_fused_permutation(
    heads_per_group: int = WV_GROUP_SIZE, head_dim: int = HEAD_DIM
) -> torch.Tensor:
    """``fused[(c * G + h) * 32 + j] = standard[h * D + c * 32 + j]`` per group."""
    h = torch.arange(heads_per_group).view(-1, 1, 1)
    c = torch.arange(head_dim // O_CHUNK).view(1, -1, 1)
    j = torch.arange(O_CHUNK).view(1, 1, -1)
    fused_index = (c * heads_per_group + h) * O_CHUNK + j
    perm = torch.empty(heads_per_group * head_dim, dtype=torch.long)
    perm[fused_index.reshape(-1)] = torch.arange(heads_per_group * head_dim)
    return perm

permute_on_load(param, perm, dim)

Make param's weight loader gather the loaded shard by perm.

Source code in vllm/models/deepseek_v41/common/ops/fused_layout.py
def permute_on_load(param: torch.Tensor, perm: torch.Tensor, dim: int) -> None:
    """Make ``param``'s weight loader gather the loaded shard by ``perm``."""

    def gather(t: torch.Tensor) -> torch.Tensor:
        return _bytes_view(t).index_select(dim, perm.to(t.device)).view(t.dtype)

    param.weight_loader = composed_weight_loader(  # type: ignore[attr-defined]
        param.weight_loader,  # type: ignore[attr-defined]
        gather,
    )

permute_q_to_fused(q)

[N, H, D] standard layout -> the same shape in the fused layout.

Source code in vllm/models/deepseek_v41/common/ops/fused_layout.py
def permute_q_to_fused(q: torch.Tensor) -> torch.Tensor:
    """``[N, H, D]`` standard layout -> the same shape in the fused layout."""
    n, h, d = q.shape
    perm = q_fused_permutation(h, d).to(q.device)
    return q.reshape(n, h * d)[:, perm].view(n, h, d)

q_fused_permutation(num_heads, head_dim=HEAD_DIM)

fused[(d // 16) * (H * 16) + h * 16 + d % 16] = standard[h * D + d].

Source code in vllm/models/deepseek_v41/common/ops/fused_layout.py
def q_fused_permutation(num_heads: int, head_dim: int = HEAD_DIM) -> torch.Tensor:
    """``fused[(d // 16) * (H * 16) + h * 16 + d % 16] = standard[h * D + d]``."""
    h = torch.arange(num_heads).view(num_heads, 1)
    d = torch.arange(head_dim).view(1, head_dim)
    fused_index = (d // Q_CHUNK) * (num_heads * Q_CHUNK) + h * Q_CHUNK + d % Q_CHUNK
    perm = torch.empty(num_heads * head_dim, dtype=torch.long)
    perm[fused_index.reshape(-1)] = torch.arange(num_heads * head_dim)
    return perm