Skip to content

vllm.v1.attention.ops.rocm_aiter_mla_merge

Merge AITER segmented MLA split-K partials into output plus natural-log LSE.

This mirrors AITER's own segment reduction, including its tiles_per_segment = cdiv(seq_len, NUM_SEGMENTS * TILE_SIZE) partitioning, so it has to move together with the skip_reduce=True call in the AITER MLA backend. It is a rank-local split-K merge and unrelated to any collective reduce; the natural-log LSE it returns is what the cross-rank DCP merge consumes.

Functions:

merge_mla_segments_triton(segm_output, segm_max, segm_expsum, seq_lens, tile_size, out_dtype)

Merge AITER base-2 segment partials into output and natural-log LSE.

Source code in vllm/v1/attention/ops/rocm_aiter_mla_merge.py
def merge_mla_segments_triton(
    segm_output: torch.Tensor,
    segm_max: torch.Tensor,
    segm_expsum: torch.Tensor,
    seq_lens: torch.Tensor,
    tile_size: int,
    out_dtype: torch.dtype,
) -> tuple[torch.Tensor, torch.Tensor]:
    """Merge AITER base-2 segment partials into output and natural-log LSE."""
    num_tokens, num_heads, num_segments, kv_lora_rank = segm_output.shape
    output = torch.empty(
        (num_tokens, num_heads, kv_lora_rank),
        dtype=out_dtype,
        device=segm_output.device,
    )
    lse = torch.empty(
        (num_tokens, num_heads),
        dtype=torch.float32,
        device=segm_output.device,
    )
    _merge_mla_segments_kernel[(num_tokens, num_heads)](
        output,
        lse,
        segm_output,
        segm_max,
        segm_expsum,
        seq_lens,
        num_query_heads=num_heads,
        out_stride0=output.stride(0),
        out_stride1=output.stride(1),
        lse_stride0=lse.stride(0),
        TILE_SIZE=tile_size,
        KV_LORA_RANK=kv_lora_rank,
        NUM_SEGMENTS_PER_SEQ=num_segments,
        LOGE2=LOGE2,
    )
    return output, lse