Optimize FP8 MLA KV cache writes with Triton kernel (#15522)

This commit is contained in:
Hudson Xing
2025-12-25 12:35:39 -08:00
committed by GitHub
parent 8087ef126f
commit 9d878c1f3e
3 changed files with 241 additions and 9 deletions

View File

@@ -13,6 +13,84 @@ def quantize_k_cache(cache_k):
return _quantize_k_cache_slow(cache_k)
def quantize_k_cache_separate(
k_nope: torch.Tensor,
k_rope: torch.Tensor,
tile_size: int = 128,
):
"""
Quantize k_nope and k_rope separately without concat, returns two tensors.
This avoids the concat operation and enables direct reuse of set_mla_kv_buffer_triton
by returning two separate byte tensors for the nope and rope parts.
Args:
k_nope: (num_tokens, dim_nope) or (num_tokens, 1, dim_nope)
Must have dim_nope=512 for FP8 MLA quantization
k_rope: (num_tokens, dim_rope) or (num_tokens, 1, dim_rope)
Must have dim_rope=64 for FP8 MLA quantization
tile_size: quantization tile size (default 128)
Returns:
Tuple of (nope_part, rope_part) where:
- nope_part: (num_tokens, 1, 528) as uint8 view, contains [nope_fp8(512) | scales(16)]
- rope_part: (num_tokens, 1, 128) as uint8 view, contains [rope_bf16_bytes(128)]
These two tensors can be directly passed to set_mla_kv_buffer_triton(kv_buffer, loc, nope_part, rope_part)
"""
# Squeeze middle dimension if present
k_nope_2d = k_nope.squeeze(1) if k_nope.ndim == 3 else k_nope
k_rope_2d = k_rope.squeeze(1) if k_rope.ndim == 3 else k_rope
num_tokens = k_nope_2d.shape[0]
dim_nope = k_nope_2d.shape[1]
dim_rope = k_rope_2d.shape[1]
# Validate dimensions for FP8 MLA
if dim_nope != 512:
raise ValueError(f"Expected dim_nope=512 for FP8 MLA, got {dim_nope}")
if dim_rope != 64:
raise ValueError(f"Expected dim_rope=64 for FP8 MLA, got {dim_rope}")
if k_rope_2d.shape[0] != num_tokens:
raise ValueError(
f"k_nope and k_rope must have same num_tokens, got {num_tokens} vs {k_rope_2d.shape[0]}"
)
# Call fast kernel that directly produces two separate outputs (single Triton kernel)
if NSA_QUANT_K_CACHE_FAST:
nope_part, rope_part = _quantize_k_cache_fast_separate(
k_nope=k_nope_2d, k_rope=k_rope_2d, group_size=tile_size
)
else:
# Fallback: use existing slow path with post-processing
cache_k_concat = torch.cat([k_nope_2d, k_rope_2d], dim=-1)
packed_output_4d = quantize_k_cache(cache_k_concat.unsqueeze(1).unsqueeze(1))
packed_output = packed_output_4d.squeeze(1).squeeze(1)
# Convert to uint8 bytes view
packed_bytes = packed_output.contiguous().view(torch.uint8)
# Strict byte-size validation
expected_total_bytes = 656 # 512 (nope_fp8) + 16 (scales) + 128 (rope_bf16)
if packed_bytes.shape[1] != expected_total_bytes:
raise ValueError(
f"Packed output has {packed_bytes.shape[1]} bytes, expected {expected_total_bytes}. "
f"Original dtype: {packed_output.dtype}, shape: {packed_output.shape}"
)
# Split into nope and rope parts
num_tiles = dim_nope // tile_size # 4
nope_part_bytes = dim_nope + num_tiles * 4 # 512 + 16 = 528
rope_part_bytes = 128
nope_part = packed_bytes[:, :nope_part_bytes].unsqueeze(1)
rope_part = packed_bytes[
:, nope_part_bytes : nope_part_bytes + rope_part_bytes
].unsqueeze(1)
return nope_part, rope_part
# Copied from original
def _quantize_k_cache_slow(
input_k_cache: torch.Tensor, # (num_blocks, block_size, h_k, d)
@@ -145,6 +223,83 @@ def _quantize_k_cache_fast(k_nope, k_rope, group_size: int = 128):
return output
def _quantize_k_cache_fast_separate(k_nope, k_rope, group_size: int = 128):
"""
Quantize k_nope and k_rope in a single Triton kernel, directly outputting two separate tensors.
This avoids packing/unpacking and enables direct use with set_mla_kv_buffer_triton.
:param k_nope: (num_tokens, dim_nope 512) bfloat16
:param k_rope: (num_tokens, dim_rope 64) bfloat16
:param group_size: quantization tile size (default 128, kernel is tuned for this value)
:return: Tuple of (nope_part_u8, rope_part_u8)
- nope_part_u8: (num_tokens, 1, nope_part_bytes) uint8, layout [nope_fp8(dim_nope) | scales(num_tiles*4)]
- rope_part_u8: (num_tokens, 1, rope_part_bytes) uint8, layout [rope_bf16_bytes(dim_rope*2)]
"""
num_tokens, dim_nope = k_nope.shape
num_tokens_, dim_rope = k_rope.shape
assert num_tokens == num_tokens_, f"k_nope and k_rope must have same num_tokens"
# Ensure contiguous tensors for kernel
k_nope = k_nope.contiguous()
k_rope = k_rope.contiguous()
num_tiles = dim_nope // group_size
# Calculate byte sizes based on validated dimensions
# nope_part: [FP8 quantized data (dim_nope bytes)] + [FP32 scales (num_tiles * 4 bytes)]
# rope_part: [BF16 raw data (dim_rope * 2 bytes)]
nope_part_bytes = (
dim_nope + num_tiles * 4
) # e.g., 512 + 4*4 = 528 for dim_nope=512, group_size=128
rope_part_bytes = (
dim_rope * k_rope.element_size()
) # e.g., 64 * 2 = 128 for dim_rope=64, BF16
# Allocate two separate output buffers (as uint8 for direct byte-level access)
nope_part_u8 = torch.empty(
(num_tokens, nope_part_bytes), dtype=torch.uint8, device=k_nope.device
)
rope_part_u8 = torch.empty(
(num_tokens, rope_part_bytes), dtype=torch.uint8, device=k_rope.device
)
# Create typed views for the kernel to write into
# Fixed byte layout for nope_part: [nope_fp8 (dim_nope bytes) | scales_fp32 (num_tiles*4 bytes)]
# Fixed byte layout for rope_part: [rope_bf16 (dim_rope*2 bytes)]
nope_q_view = nope_part_u8[:, :dim_nope].view(torch.float8_e4m3fn)
nope_s_view = nope_part_u8[:, dim_nope:].view(torch.float32)
rope_view = rope_part_u8.view(torch.bfloat16)
# Kernel launch parameters
num_blocks_per_token = triton.cdiv(dim_nope + dim_rope, group_size)
NUM_NOPE_BLOCKS = dim_nope // group_size
# Use the same kernel as _quantize_k_cache_fast (reuse existing implementation)
_quantize_k_cache_fast_kernel[(num_tokens, num_blocks_per_token)](
nope_q_view,
nope_s_view,
rope_view,
k_nope,
k_rope,
nope_q_view.stride(0),
nope_s_view.stride(0),
rope_view.stride(0),
k_nope.stride(0),
k_rope.stride(0),
NUM_NOPE_BLOCKS=NUM_NOPE_BLOCKS,
GROUP_SIZE=group_size,
DIM_NOPE=dim_nope,
DIM_ROPE=dim_rope,
FP8_MIN=torch.finfo(torch.float8_e4m3fn).min,
FP8_MAX=torch.finfo(torch.float8_e4m3fn).max,
)
# Add middle dimension for compatibility with set_mla_kv_buffer_triton
return nope_part_u8.unsqueeze(1), rope_part_u8.unsqueeze(1)
@triton.jit
def _quantize_k_cache_fast_kernel(
output_nope_q_ptr,
@@ -255,7 +410,50 @@ if __name__ == "__main__":
)
print("Passed")
print("Do benchmark...")
# Test quantize_k_cache_separate: verify output matches concat path
print("\nTesting quantize_k_cache_separate...")
for num_tokens in [64, 100]:
dim_nope = 512
dim_rope = 64
k_nope = torch.randn(
num_tokens, 1, dim_nope, dtype=torch.bfloat16, device="cuda"
)
k_rope = torch.randn(
num_tokens, 1, dim_rope, dtype=torch.bfloat16, device="cuda"
)
# Old path: concat then quantize
k_concat = torch.cat([k_nope, k_rope], dim=-1).squeeze(1) # (num_tokens, 576)
old_output = quantize_k_cache(k_concat.unsqueeze(1).unsqueeze(1)) # 4D input
old_output = old_output.squeeze(1).squeeze(1) # Back to (num_tokens, 656)
# New path: quantize separately
nope_part, rope_part = quantize_k_cache_separate(k_nope, k_rope)
new_bytes = torch.cat([nope_part.squeeze(1), rope_part.squeeze(1)], dim=-1)
# Compare byte-level equality
old_bytes = old_output.view(torch.uint8)
if old_bytes.shape != new_bytes.shape:
raise RuntimeError(
f"Shape mismatch: {old_bytes.shape} vs {new_bytes.shape}"
)
diff_bytes = (old_bytes != new_bytes).sum().item()
if diff_bytes > 0:
max_diff = (old_bytes.float() - new_bytes.float()).abs().max().item()
raise RuntimeError(
f"quantize_k_cache_separate output doesn't match concat path: "
f"{diff_bytes} differing bytes, max_diff={max_diff}"
)
print(f" num_tokens={num_tokens}: PASSED (outputs match byte-wise)")
print("quantize_k_cache_separate tests passed!")
print("\nDo benchmark...")
for num_blocks, block_size in [
(1, 64),

View File

@@ -22,7 +22,10 @@ from typing import List
from sglang.srt.configs.mamba_utils import BaseLinearStateParams
from sglang.srt.environ import envs
from sglang.srt.layers.attention.nsa import index_buf_accessor
from sglang.srt.layers.attention.nsa.quant_k_cache import quantize_k_cache
from sglang.srt.layers.attention.nsa.quant_k_cache import (
quantize_k_cache,
quantize_k_cache_separate,
)
from sglang.srt.utils.torch_memory_saver_adapter import TorchMemorySaverAdapter
"""
@@ -1597,12 +1600,22 @@ class MLATokenToKVPool(KVCache):
layer_id = layer.layer_id
if self.use_nsa and self.nsa_kv_cache_store_fp8:
# original cache_k: (num_tokens, num_heads 1, hidden 576); we unsqueeze the page_size=1 dim here
# TODO no need to cat
cache_k = torch.cat([cache_k_nope, cache_k_rope], dim=-1)
cache_k = quantize_k_cache(cache_k.unsqueeze(1)).squeeze(1)
cache_k = cache_k.view(self.store_dtype)
self.kv_buffer[layer_id - self.start_layer][loc] = cache_k
# OPTIMIZATION: Quantize k_nope and k_rope separately to avoid concat overhead
# This also enables reuse of set_mla_kv_buffer_triton two-tensor write path
# quantize_k_cache_separate returns (nope_part, rope_part) as uint8 bytes
cache_k_nope_fp8, cache_k_rope_fp8 = quantize_k_cache_separate(
cache_k_nope, cache_k_rope
)
# Reuse existing two-tensor write kernel (works with FP8 byte layout)
# cache_k_nope_fp8: (num_tokens, 1, 528) uint8 [nope_fp8(512) | scales(16)]
# cache_k_rope_fp8: (num_tokens, 1, 128) uint8 [rope_bf16_bytes(128)]
set_mla_kv_buffer_triton(
self.kv_buffer[layer_id - self.start_layer],
loc,
cache_k_nope_fp8,
cache_k_rope_fp8,
)
else:
if cache_k_nope.dtype != self.dtype:
cache_k_nope = cache_k_nope.to(self.dtype)

View File

@@ -46,17 +46,38 @@ def set_mla_kv_buffer_kernel(
loc = tl.load(loc_ptr + pid_loc).to(tl.int64)
dst_ptr = kv_buffer_ptr + loc * buffer_stride + offs
# Three-way branch to handle boundary correctly while preserving fast path
if base + BLOCK <= nope_dim:
# Fast path: entire block is in nope region
src = tl.load(
cache_k_nope_ptr + pid_loc * nope_stride + offs,
mask=mask,
)
else:
elif base >= nope_dim:
# Fast path: entire block is in rope region
offs_rope = offs - nope_dim
src = tl.load(
cache_k_rope_ptr + pid_loc * rope_stride + offs_rope,
mask=mask,
)
else:
# Boundary case: block spans nope/rope boundary (e.g., FP8 with nope_dim=528)
# Handle each offset individually to avoid negative indexing
is_nope = offs < nope_dim
is_rope = (offs >= nope_dim) & (offs < (nope_dim + rope_dim))
src_nope = tl.load(
cache_k_nope_ptr + pid_loc * nope_stride + offs,
mask=mask & is_nope,
other=0,
)
src_rope = tl.load(
cache_k_rope_ptr + pid_loc * rope_stride + (offs - nope_dim),
mask=mask & is_rope,
other=0,
)
src = tl.where(is_nope, src_nope, src_rope)
tl.store(dst_ptr, src, mask=mask)