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)
|
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
|
# Copied from original
|
||||||
def _quantize_k_cache_slow(
|
def _quantize_k_cache_slow(
|
||||||
input_k_cache: torch.Tensor, # (num_blocks, block_size, h_k, d)
|
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
|
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
|
@triton.jit
|
||||||
def _quantize_k_cache_fast_kernel(
|
def _quantize_k_cache_fast_kernel(
|
||||||
output_nope_q_ptr,
|
output_nope_q_ptr,
|
||||||
@@ -255,7 +410,50 @@ if __name__ == "__main__":
|
|||||||
)
|
)
|
||||||
|
|
||||||
print("Passed")
|
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 [
|
for num_blocks, block_size in [
|
||||||
(1, 64),
|
(1, 64),
|
||||||
|
|||||||
@@ -22,7 +22,10 @@ from typing import List
|
|||||||
from sglang.srt.configs.mamba_utils import BaseLinearStateParams
|
from sglang.srt.configs.mamba_utils import BaseLinearStateParams
|
||||||
from sglang.srt.environ import envs
|
from sglang.srt.environ import envs
|
||||||
from sglang.srt.layers.attention.nsa import index_buf_accessor
|
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
|
from sglang.srt.utils.torch_memory_saver_adapter import TorchMemorySaverAdapter
|
||||||
|
|
||||||
"""
|
"""
|
||||||
@@ -1597,12 +1600,22 @@ class MLATokenToKVPool(KVCache):
|
|||||||
layer_id = layer.layer_id
|
layer_id = layer.layer_id
|
||||||
|
|
||||||
if self.use_nsa and self.nsa_kv_cache_store_fp8:
|
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
|
# OPTIMIZATION: Quantize k_nope and k_rope separately to avoid concat overhead
|
||||||
# TODO no need to cat
|
# This also enables reuse of set_mla_kv_buffer_triton two-tensor write path
|
||||||
cache_k = torch.cat([cache_k_nope, cache_k_rope], dim=-1)
|
# quantize_k_cache_separate returns (nope_part, rope_part) as uint8 bytes
|
||||||
cache_k = quantize_k_cache(cache_k.unsqueeze(1)).squeeze(1)
|
cache_k_nope_fp8, cache_k_rope_fp8 = quantize_k_cache_separate(
|
||||||
cache_k = cache_k.view(self.store_dtype)
|
cache_k_nope, cache_k_rope
|
||||||
self.kv_buffer[layer_id - self.start_layer][loc] = cache_k
|
)
|
||||||
|
|
||||||
|
# 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:
|
else:
|
||||||
if cache_k_nope.dtype != self.dtype:
|
if cache_k_nope.dtype != self.dtype:
|
||||||
cache_k_nope = cache_k_nope.to(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)
|
loc = tl.load(loc_ptr + pid_loc).to(tl.int64)
|
||||||
dst_ptr = kv_buffer_ptr + loc * buffer_stride + offs
|
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:
|
if base + BLOCK <= nope_dim:
|
||||||
|
# Fast path: entire block is in nope region
|
||||||
src = tl.load(
|
src = tl.load(
|
||||||
cache_k_nope_ptr + pid_loc * nope_stride + offs,
|
cache_k_nope_ptr + pid_loc * nope_stride + offs,
|
||||||
mask=mask,
|
mask=mask,
|
||||||
)
|
)
|
||||||
else:
|
elif base >= nope_dim:
|
||||||
|
# Fast path: entire block is in rope region
|
||||||
offs_rope = offs - nope_dim
|
offs_rope = offs - nope_dim
|
||||||
src = tl.load(
|
src = tl.load(
|
||||||
cache_k_rope_ptr + pid_loc * rope_stride + offs_rope,
|
cache_k_rope_ptr + pid_loc * rope_stride + offs_rope,
|
||||||
mask=mask,
|
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)
|
tl.store(dst_ptr, src, mask=mask)
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user