Optimize FP8 MLA KV cache writes with Triton kernel (#15522)
This commit is contained in:
@@ -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),
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user