Skip to content

vllm.model_executor.layers.rotary_embedding.packed_qk_rope

In-place interleaved RoPE on the Q/K slices of a packed QKV buffer.

Functions:

packed_qk_rope_(xqkv, freqs_cis)

Rotate Q and K in place inside a packed QKV buffer.

Parameters:

  • xqkv

    (Tensor) –

    contiguous (seqlen, 3, nheads, headdim); only the Q and K slices are rotated.

  • freqs_cis

    (Tensor) –

    contiguous (seqlen, headdim // 2) complex64 rotary freqs, read through its interleaved re/im fp32 view.

Matches ApplyRotaryEmb(enable_fp32_compute=True) applied to Q and K separately to within one ULP, but in one kernel with ~5x less memory traffic. Both compute in fp32; the gap is which product the compiler keeps exact inside the fused multiply-add, which differs per backend. Requires triton: callers must check HAS_TRITON and use the unfused path otherwise. Contract violations raise.

Source code in vllm/model_executor/layers/rotary_embedding/packed_qk_rope.py
def packed_qk_rope_(xqkv: torch.Tensor, freqs_cis: torch.Tensor) -> None:
    """Rotate Q and K in place inside a packed QKV buffer.

    Args:
        xqkv: contiguous (seqlen, 3, nheads, headdim); only the Q and K
            slices are rotated.
        freqs_cis: contiguous (seqlen, headdim // 2) complex64 rotary freqs,
            read through its interleaved re/im fp32 view.

    Matches ``ApplyRotaryEmb(enable_fp32_compute=True)`` applied to Q and K
    separately to within one ULP, but in one kernel with ~5x less memory
    traffic. Both compute in fp32; the gap is which product the compiler
    keeps exact inside the fused multiply-add, which differs per backend.
    Requires triton: callers must check HAS_TRITON and use the unfused path
    otherwise. Contract violations raise.

    """
    assert xqkv.ndim == 4 and xqkv.size(1) == 3, xqkv.shape
    assert xqkv.is_contiguous(), xqkv.shape
    assert freqs_cis.ndim == 2 and freqs_cis.is_contiguous(), freqs_cis.shape
    assert freqs_cis.dtype == torch.complex64, freqs_cis.dtype

    seq_length, _, num_heads, head_dim = xqkv.shape
    rotary_dim = freqs_cis.size(-1) * 2
    assert freqs_cis.size(0) == seq_length, (freqs_cis.shape, xqkv.shape)
    assert rotary_dim == head_dim, (freqs_cis.shape, xqkv.shape)  # full rotation
    qk = xqkv.as_strided(
        (seq_length, 2 * num_heads, head_dim),
        (3 * num_heads * head_dim, head_dim, 1),
    )

    BLOCK_M = 4
    # seqlen on grid.x: HIP caps gridDim.y at 65535.
    grid = (triton.cdiv(seq_length, BLOCK_M), 2 * num_heads)
    _packed_qk_rope_kernel[grid](
        qk,
        freqs_cis.view(torch.float32),
        seq_length,
        qk.stride(0),
        qk.stride(1),
        rotary_dim=rotary_dim,
        BLOCK_M=BLOCK_M,
        num_warps=2 if rotary_dim <= 64 else 4,
    )