Skip to content

vllm.v1.attention.backends.mla.amx_mla

AMX-only, high-performance MLA backend for DeepSeek V2/V3/R1 on CPU.

Built on the AMX decode/extend/bmm kernels vendored under csrc/cpu/sgl-kernels/, plugged into vLLM's MLACommonBackend/ MLACommonImpl abstraction the same way every other concrete MLA backend (TritonMLAImpl, etc.) does.

This is a separate backend from the reference CPUMLABackend (vllm/v1/attention/backends/mla/cpu_mla.py): that one targets every CPU (any dtype, block_size=16, mla_decode_kvcache decode kernel + inherited SDPA-style prefill) as a functional/CI reference, not performance. This backend instead requires AMX (bf16 only, block_size a multiple of 32) and is selected by the platform layer in preference to the reference backend whenever the host supports it -- see CpuPlatform.get_attn_backend_cls.

Two points where this backend differs structurally from the GPU backends, both explained in the CPU MLA design plan:

  • forward_mha is fully overridden (not inherited from MLACommonBaseImpl): the GPU "compute-friendly" prefill path depends on a pluggable MLAPrefillBackend and CUDA-only chunked-context gather ops, neither of which exist on CPU. Instead, this attends directly in latent-MQA space via the extend_attention_cpu kernel, which handles cached-prefix continuation and fresh prefill in one causal pass (the KV cache already contains the new tokens by the time this runs, since do_kv_cache_update executes first).
  • Weight absorption for the prefill path is done with this impl's own VNNI-packed copies of W_UK/W_UV (computed in process_weights_after_loading, using bmm_cpu), rather than reading the generic layer.W_UK_T/layer.W_UV/layer._v_up_proj the way GPU backends do, since forward_mha's abstract signature has no layer parameter (only forward_mqa does) and that signature is left untouched. The decode path needs no such packing: MLAAttention.forward_impl already absorbs/de-absorbs Q around the forward_mqa call using the generic (unpacked) layer.W_UK_T/layer._v_up_proj, so forward_mqa here only has to invoke the decode kernel.

_block_quant_to_tensor_quant(weight, weight_scale, block_size)

Re-quantize FP8 block-quant weight to FP8 per-tensor quant.

Returns (fp8_weight, dequant_scale) where fp8_weight * dequant_scale approximates the original float weight. dequant_scale is a 0-dim scalar tensor as expected by bmm_cpu's scale argument.

NOTE: block->tensor re-quantization may slightly affect accuracy (equivalent to SGLang's block_quant_to_tensor_quant for CPU).

Source code in vllm/v1/attention/backends/mla/amx_mla.py
def _block_quant_to_tensor_quant(
    weight: torch.Tensor,  # FP8 [N, K] block-quantized weight
    weight_scale: torch.Tensor,  # [ceil(N/bn), ceil(K/bk)] block scales
    block_size: list[int],  # [block_n, block_k]
) -> tuple[torch.Tensor, torch.Tensor]:
    """Re-quantize FP8 block-quant weight to FP8 per-tensor quant.

    Returns (fp8_weight, dequant_scale) where fp8_weight * dequant_scale
    approximates the original float weight.  dequant_scale is a 0-dim
    scalar tensor as expected by bmm_cpu's ``scale`` argument.

    NOTE: block->tensor re-quantization may slightly affect accuracy
    (equivalent to SGLang's block_quant_to_tensor_quant for CPU).
    """
    block_n, block_k = block_size[0], block_size[1]
    n, k = weight.shape

    # Expand block scales [ceil(N/bn), ceil(K/bk)] -> [N, K]
    scale = weight_scale.to(torch.float32)
    if scale.shape[0] < n:
        scale = scale.repeat_interleave(block_n, dim=0)[:n, :]
    if scale.shape[1] < k:
        scale = scale.repeat_interleave(block_k, dim=1)[:, :k]

    # Dequantize to float32
    x_f32 = weight.to(torch.float32) * scale

    # Re-quantize per-tensor back to FP8
    fp8_max = torch.finfo(weight.dtype).max  # 448.0 for e4m3fn
    amax = x_f32.abs().amax().clamp(min=1e-12)
    quant_scale = fp8_max / amax
    x_fp8 = (x_f32 * quant_scale).clamp(-fp8_max, fp8_max).to(weight.dtype)

    # dequant_scale = 1 / quant_scale (0-dim tensor for bmm_cpu)
    dequant_scale = (amax / fp8_max).to(torch.float32)
    return x_fp8, dequant_scale

_compute_num_kv_splits(max_seq_len, num_threads)

Mirrors TritonMLAImpl's _compute_num_kv_splits, using the CPU thread count in place of SM count.

Source code in vllm/v1/attention/backends/mla/amx_mla.py
def _compute_num_kv_splits(max_seq_len: int, num_threads: int) -> int:
    """Mirrors TritonMLAImpl's _compute_num_kv_splits, using the CPU thread
    count in place of SM count."""
    ideal_splits = 1
    while ideal_splits < max(1, max_seq_len // _MIN_WORK_PER_SPLIT):
        ideal_splits *= 2
    max_splits = num_threads * _SPLIT_OCCUPANCY_MULTIPLIER
    return min(ideal_splits, max_splits)

_expand_block_table(block_table, block_size)

Adapter: vLLM's block-table paging -> the flat per-(request, position) physical row index the decode/extend kernels expect. Called once per step from AMXMLAMetadataBuilder.build, not per layer.

Source code in vllm/v1/attention/backends/mla/amx_mla.py
def _expand_block_table(block_table: torch.Tensor, block_size: int) -> torch.Tensor:
    """Adapter: vLLM's block-table paging -> the flat per-(request, position)
    physical row index the decode/extend kernels expect. Called once per step
    from ``AMXMLAMetadataBuilder.build``, not per layer.
    """
    offsets = torch.arange(
        block_size, device=block_table.device, dtype=block_table.dtype
    )
    flat = block_table.unsqueeze(-1) * block_size + offsets
    return flat.reshape(block_table.size(0), -1).contiguous()