Kernel: optimize decoding metadata in NSA multi-spec backend with fused kernels (#17554)
This commit is contained in:
@@ -362,6 +362,8 @@ class Envs:
|
||||
# NSA Backend
|
||||
SGLANG_NSA_FUSE_TOPK = EnvBool(True)
|
||||
SGLANG_NSA_ENABLE_MTP_PRECOMPUTE_METADATA = EnvBool(True)
|
||||
SGLANG_USE_FUSED_METADATA_COPY = EnvBool(True)
|
||||
SGLANG_VERIFY_FUSED_METADATA_COPY = EnvBool(False)
|
||||
SGLANG_NSA_FORCE_MLA = EnvBool(False)
|
||||
|
||||
# sgl-kernel
|
||||
|
||||
@@ -127,7 +127,7 @@ class NativeSparseAttnBackendMTPPrecomputeMixin:
|
||||
cu_seqlens_k = compute_cu_seqlens(cache_seqlens)
|
||||
|
||||
# Get page indices from cache
|
||||
page_indices = self.req_to_token[req_pool_indices, :max_len]
|
||||
page_indices = self.req_to_token[req_pool_indices, :max_len].contiguous()
|
||||
|
||||
# Compute NSA seqlens
|
||||
nsa_cache_seqlens = compute_nsa_seqlens(
|
||||
@@ -187,7 +187,7 @@ class NativeSparseAttnBackendMTPPrecomputeMixin:
|
||||
page_indices = self.req_to_token[req_pool_indices, :max_seqlen_k]
|
||||
page_indices = torch.repeat_interleave(
|
||||
page_indices, repeats=self.speculative_num_draft_tokens, dim=0
|
||||
)
|
||||
).contiguous()
|
||||
|
||||
# Generate expanded seqlens
|
||||
extend_seq_lens_cpu = [self.speculative_num_draft_tokens] * bs
|
||||
@@ -269,7 +269,7 @@ class NativeSparseAttnBackendMTPPrecomputeMixin:
|
||||
page_indices = self.req_to_token[req_pool_indices, :max_seqlen_k]
|
||||
page_indices = torch.repeat_interleave(
|
||||
page_indices, repeats=extend_seq_lens, dim=0
|
||||
)
|
||||
).contiguous()
|
||||
|
||||
# Generate expanded seqlens
|
||||
seqlens_expanded = torch.cat(
|
||||
|
||||
@@ -0,0 +1,407 @@
|
||||
"""
|
||||
Verification utilities for NSA backend fused metadata copy operations.
|
||||
|
||||
This module contains verification code to ensure that fused metadata copy kernels
|
||||
produce the same results as individual copy operations.
|
||||
"""
|
||||
|
||||
import torch
|
||||
|
||||
|
||||
def verify_single_backend_fused_metadata_copy(
|
||||
metadata,
|
||||
precomputed,
|
||||
forward_mode,
|
||||
bs,
|
||||
flashmla_num_splits_src=None,
|
||||
flashmla_metadata_src=None,
|
||||
flashmla_num_splits_dst=None,
|
||||
flashmla_metadata_dst=None,
|
||||
):
|
||||
"""
|
||||
Verify that the fused metadata copy kernel produces the same results as individual copies.
|
||||
|
||||
Args:
|
||||
metadata: The NSA metadata object containing destination tensors
|
||||
precomputed: The precomputed metadata containing source tensors
|
||||
forward_mode: The forward mode (decode, target_verify, or draft_extend)
|
||||
bs: Batch size
|
||||
flashmla_num_splits_src: Source FlashMLA num_splits tensor (optional)
|
||||
flashmla_metadata_src: Source FlashMLA metadata tensor (optional)
|
||||
flashmla_num_splits_dst: Destination FlashMLA num_splits tensor (optional)
|
||||
flashmla_metadata_dst: Destination FlashMLA metadata tensor (optional)
|
||||
|
||||
Raises:
|
||||
RuntimeError: If verification fails (tensors don't match)
|
||||
"""
|
||||
# Clone destination tensors to preserve fused kernel results
|
||||
fused_cache_seqlens = metadata.cache_seqlens_int32.clone()
|
||||
fused_cu_seqlens_k = metadata.cu_seqlens_k.clone()
|
||||
fused_page_table_1 = metadata.page_table_1.clone()
|
||||
fused_nsa_cache_seqlens = metadata.nsa_cache_seqlens_int32.clone()
|
||||
fused_nsa_seqlens_expanded = metadata.nsa_seqlens_expanded.clone()
|
||||
fused_nsa_cu_seqlens_k = metadata.nsa_cu_seqlens_k.clone()
|
||||
fused_real_page_table = (
|
||||
metadata.real_page_table.clone()
|
||||
if precomputed.real_page_table is not None
|
||||
else None
|
||||
)
|
||||
fused_flashmla_num_splits = None
|
||||
fused_flashmla_metadata = None
|
||||
if precomputed.flashmla_metadata is not None:
|
||||
fused_flashmla_num_splits = flashmla_num_splits_dst.clone()
|
||||
fused_flashmla_metadata = flashmla_metadata_dst.clone()
|
||||
|
||||
# Create reference tensors (zeroed out)
|
||||
ref_cache_seqlens = torch.zeros_like(metadata.cache_seqlens_int32)
|
||||
ref_cu_seqlens_k = torch.zeros_like(metadata.cu_seqlens_k)
|
||||
ref_page_table_1 = torch.zeros_like(metadata.page_table_1)
|
||||
ref_nsa_cache_seqlens = torch.zeros_like(metadata.nsa_cache_seqlens_int32)
|
||||
ref_nsa_seqlens_expanded = torch.zeros_like(metadata.nsa_seqlens_expanded)
|
||||
ref_nsa_cu_seqlens_k = torch.zeros_like(metadata.nsa_cu_seqlens_k)
|
||||
ref_real_page_table = (
|
||||
torch.zeros_like(metadata.real_page_table)
|
||||
if precomputed.real_page_table is not None
|
||||
else None
|
||||
)
|
||||
ref_flashmla_num_splits = None
|
||||
ref_flashmla_metadata = None
|
||||
if precomputed.flashmla_metadata is not None:
|
||||
ref_flashmla_num_splits = torch.zeros_like(flashmla_num_splits_dst)
|
||||
ref_flashmla_metadata = torch.zeros_like(flashmla_metadata_dst)
|
||||
|
||||
# Run individual copy operations (reference implementation)
|
||||
ref_cache_seqlens.copy_(precomputed.cache_seqlens)
|
||||
ref_cu_seqlens_k[1:].copy_(precomputed.cu_seqlens_k[1:])
|
||||
|
||||
if forward_mode.is_decode_or_idle():
|
||||
# Decode mode
|
||||
ref_page_table_1[:, : precomputed.max_len].copy_(precomputed.page_indices)
|
||||
ref_nsa_cache_seqlens.copy_(precomputed.nsa_cache_seqlens)
|
||||
elif forward_mode.is_target_verify():
|
||||
# Target verify mode
|
||||
ref_page_table_1[:, : precomputed.max_seqlen_k].copy_(precomputed.page_indices)
|
||||
ref_nsa_seqlens_expanded.copy_(precomputed.seqlens_expanded)
|
||||
ref_nsa_cache_seqlens.copy_(precomputed.nsa_cache_seqlens)
|
||||
elif forward_mode.is_draft_extend():
|
||||
# Draft extend mode
|
||||
rows = precomputed.page_indices.shape[0]
|
||||
cols = precomputed.max_seqlen_k
|
||||
ref_page_table_1[:rows, :cols].copy_(precomputed.page_indices)
|
||||
size = precomputed.seqlens_expanded_size
|
||||
ref_nsa_seqlens_expanded[:size].copy_(precomputed.seqlens_expanded)
|
||||
ref_nsa_cache_seqlens[:size].copy_(precomputed.nsa_cache_seqlens)
|
||||
|
||||
# Copy NSA cu_seqlens
|
||||
size = precomputed.seqlens_expanded_size
|
||||
ref_nsa_cu_seqlens_k[1 : 1 + size].copy_(precomputed.nsa_cu_seqlens_k[1 : 1 + size])
|
||||
|
||||
# Copy real page table
|
||||
if precomputed.real_page_table is not None:
|
||||
rows, cols = precomputed.real_page_table.shape
|
||||
ref_real_page_table[:rows, :cols].copy_(precomputed.real_page_table)
|
||||
|
||||
# Copy FlashMLA metadata
|
||||
if precomputed.flashmla_metadata is not None:
|
||||
size = precomputed.seqlens_expanded_size
|
||||
ref_flashmla_num_splits[: size + 1].copy_(flashmla_num_splits_src[: size + 1])
|
||||
ref_flashmla_metadata.copy_(flashmla_metadata_src)
|
||||
|
||||
# Compare results and crash if inconsistent
|
||||
def check_tensor_equal(name, fused, ref):
|
||||
if not torch.equal(fused, ref):
|
||||
max_diff = (fused.float() - ref.float()).abs().max().item()
|
||||
mismatched_elements = (fused != ref).sum().item()
|
||||
total_elements = fused.numel()
|
||||
raise RuntimeError(
|
||||
f"FUSED METADATA COPY VERIFICATION FAILED!\n"
|
||||
f"Tensor: {name}\n"
|
||||
f"Max difference: {max_diff}\n"
|
||||
f"Mismatched elements: {mismatched_elements}/{total_elements}\n"
|
||||
f"Fused shape: {fused.shape}, Ref shape: {ref.shape}\n"
|
||||
f"Forward mode: {forward_mode}, bs={bs}\n"
|
||||
f"The fused kernel produces different results than individual copies.\n"
|
||||
f"This indicates a bug in the fused metadata copy kernel."
|
||||
)
|
||||
|
||||
# Verify all tensors (only compare the slices that were actually updated)
|
||||
check_tensor_equal("cache_seqlens", fused_cache_seqlens, ref_cache_seqlens)
|
||||
check_tensor_equal("cu_seqlens_k", fused_cu_seqlens_k, ref_cu_seqlens_k)
|
||||
|
||||
# Compare page_table_1 only for the region that was updated
|
||||
if forward_mode.is_decode_or_idle():
|
||||
check_tensor_equal(
|
||||
"page_table_1",
|
||||
fused_page_table_1[:, : precomputed.max_len],
|
||||
ref_page_table_1[:, : precomputed.max_len],
|
||||
)
|
||||
elif forward_mode.is_target_verify():
|
||||
check_tensor_equal(
|
||||
"page_table_1",
|
||||
fused_page_table_1[:, : precomputed.max_seqlen_k],
|
||||
ref_page_table_1[:, : precomputed.max_seqlen_k],
|
||||
)
|
||||
elif forward_mode.is_draft_extend():
|
||||
rows = precomputed.page_indices.shape[0]
|
||||
cols = precomputed.max_seqlen_k
|
||||
check_tensor_equal(
|
||||
"page_table_1",
|
||||
fused_page_table_1[:rows, :cols],
|
||||
ref_page_table_1[:rows, :cols],
|
||||
)
|
||||
|
||||
# Compare nsa_cache_seqlens only for the region that was updated
|
||||
if forward_mode.is_decode_or_idle():
|
||||
check_tensor_equal(
|
||||
"nsa_cache_seqlens",
|
||||
fused_nsa_cache_seqlens,
|
||||
ref_nsa_cache_seqlens,
|
||||
)
|
||||
else: # TARGET_VERIFY or DRAFT_EXTEND
|
||||
size = precomputed.seqlens_expanded_size
|
||||
check_tensor_equal(
|
||||
"nsa_cache_seqlens",
|
||||
fused_nsa_cache_seqlens[:size],
|
||||
ref_nsa_cache_seqlens[:size],
|
||||
)
|
||||
|
||||
# Compare nsa_seqlens_expanded only for TARGET_VERIFY and DRAFT_EXTEND
|
||||
if forward_mode.is_target_verify() or forward_mode.is_draft_extend():
|
||||
size = precomputed.seqlens_expanded_size
|
||||
check_tensor_equal(
|
||||
"nsa_seqlens_expanded",
|
||||
fused_nsa_seqlens_expanded[:size],
|
||||
ref_nsa_seqlens_expanded[:size],
|
||||
)
|
||||
|
||||
# Compare nsa_cu_seqlens_k only for the region that was updated
|
||||
size = precomputed.seqlens_expanded_size
|
||||
check_tensor_equal(
|
||||
"nsa_cu_seqlens_k",
|
||||
fused_nsa_cu_seqlens_k[: 1 + size],
|
||||
ref_nsa_cu_seqlens_k[: 1 + size],
|
||||
)
|
||||
|
||||
if precomputed.real_page_table is not None:
|
||||
rows, cols = precomputed.real_page_table.shape
|
||||
check_tensor_equal(
|
||||
"real_page_table",
|
||||
fused_real_page_table[:rows, :cols],
|
||||
ref_real_page_table[:rows, :cols],
|
||||
)
|
||||
|
||||
if precomputed.flashmla_metadata is not None:
|
||||
size = precomputed.seqlens_expanded_size
|
||||
check_tensor_equal(
|
||||
"flashmla_num_splits",
|
||||
fused_flashmla_num_splits[: size + 1],
|
||||
ref_flashmla_num_splits[: size + 1],
|
||||
)
|
||||
check_tensor_equal(
|
||||
"flashmla_metadata",
|
||||
fused_flashmla_metadata,
|
||||
ref_flashmla_metadata,
|
||||
)
|
||||
|
||||
|
||||
def verify_multi_backend_fused_metadata_copy(
|
||||
metadata0,
|
||||
metadata1,
|
||||
metadata2,
|
||||
precomputed,
|
||||
bs,
|
||||
flashmla_num_splits_src=None,
|
||||
flashmla_metadata_src=None,
|
||||
):
|
||||
"""
|
||||
Verify that the multi-backend fused metadata copy kernel produces the same results
|
||||
as individual copies for all three backends.
|
||||
|
||||
Args:
|
||||
metadata0: The NSA metadata object for backend 0
|
||||
metadata1: The NSA metadata object for backend 1
|
||||
metadata2: The NSA metadata object for backend 2
|
||||
precomputed: The precomputed metadata containing source tensors
|
||||
bs: Batch size
|
||||
flashmla_num_splits_src: Source FlashMLA num_splits tensor (optional)
|
||||
flashmla_metadata_src: Source FlashMLA metadata tensor (optional)
|
||||
|
||||
Raises:
|
||||
RuntimeError: If verification fails (tensors don't match)
|
||||
"""
|
||||
# Clone destination tensors to preserve fused kernel results
|
||||
fused_results = []
|
||||
for idx, metadata in enumerate([metadata0, metadata1, metadata2]):
|
||||
fused_cache_seqlens = metadata.cache_seqlens_int32.clone()
|
||||
fused_cu_seqlens_k = metadata.cu_seqlens_k.clone()
|
||||
fused_page_table_1 = metadata.page_table_1.clone()
|
||||
fused_nsa_cache_seqlens = metadata.nsa_cache_seqlens_int32.clone()
|
||||
fused_nsa_cu_seqlens_k = metadata.nsa_cu_seqlens_k.clone()
|
||||
fused_real_page_table = (
|
||||
metadata.real_page_table.clone()
|
||||
if precomputed.real_page_table is not None
|
||||
else None
|
||||
)
|
||||
fused_flashmla_num_splits = None
|
||||
fused_flashmla_metadata = None
|
||||
if precomputed.flashmla_metadata is not None:
|
||||
fused_flashmla_num_splits = metadata.flashmla_metadata.num_splits.clone()
|
||||
fused_flashmla_metadata = (
|
||||
metadata.flashmla_metadata.flashmla_metadata.clone()
|
||||
)
|
||||
|
||||
fused_results.append(
|
||||
{
|
||||
"cache_seqlens": fused_cache_seqlens,
|
||||
"cu_seqlens_k": fused_cu_seqlens_k,
|
||||
"page_table_1": fused_page_table_1,
|
||||
"nsa_cache_seqlens": fused_nsa_cache_seqlens,
|
||||
"nsa_cu_seqlens_k": fused_nsa_cu_seqlens_k,
|
||||
"real_page_table": fused_real_page_table,
|
||||
"flashmla_num_splits": fused_flashmla_num_splits,
|
||||
"flashmla_metadata": fused_flashmla_metadata,
|
||||
}
|
||||
)
|
||||
|
||||
# Run individual copy operations for each backend (reference implementation)
|
||||
ref_results = []
|
||||
for idx in range(3):
|
||||
metadata = [metadata0, metadata1, metadata2][idx]
|
||||
|
||||
# Create reference tensors (zeroed out)
|
||||
ref_cache_seqlens = torch.zeros_like(metadata.cache_seqlens_int32)
|
||||
ref_cu_seqlens_k = torch.zeros_like(metadata.cu_seqlens_k)
|
||||
ref_page_table_1 = torch.zeros_like(metadata.page_table_1)
|
||||
ref_nsa_cache_seqlens = torch.zeros_like(metadata.nsa_cache_seqlens_int32)
|
||||
ref_nsa_cu_seqlens_k = torch.zeros_like(metadata.nsa_cu_seqlens_k)
|
||||
ref_real_page_table = (
|
||||
torch.zeros_like(metadata.real_page_table)
|
||||
if precomputed.real_page_table is not None
|
||||
else None
|
||||
)
|
||||
ref_flashmla_num_splits = None
|
||||
ref_flashmla_metadata = None
|
||||
if precomputed.flashmla_metadata is not None:
|
||||
ref_flashmla_num_splits = torch.zeros_like(
|
||||
metadata.flashmla_metadata.num_splits
|
||||
)
|
||||
ref_flashmla_metadata = torch.zeros_like(
|
||||
metadata.flashmla_metadata.flashmla_metadata
|
||||
)
|
||||
|
||||
# Copy operations (decode mode)
|
||||
ref_cache_seqlens.copy_(precomputed.cache_seqlens)
|
||||
ref_cu_seqlens_k[1:].copy_(precomputed.cu_seqlens_k[1:])
|
||||
ref_page_table_1[:, : precomputed.max_len].copy_(precomputed.page_indices)
|
||||
ref_nsa_cache_seqlens.copy_(precomputed.nsa_cache_seqlens)
|
||||
|
||||
# Copy NSA cu_seqlens
|
||||
size = precomputed.seqlens_expanded_size
|
||||
ref_nsa_cu_seqlens_k[1 : 1 + size].copy_(
|
||||
precomputed.nsa_cu_seqlens_k[1 : 1 + size]
|
||||
)
|
||||
|
||||
# Copy real page table
|
||||
if precomputed.real_page_table is not None:
|
||||
rows, cols = precomputed.real_page_table.shape
|
||||
ref_real_page_table[:rows, :cols].copy_(precomputed.real_page_table)
|
||||
|
||||
# Copy FlashMLA metadata
|
||||
if precomputed.flashmla_metadata is not None:
|
||||
ref_flashmla_num_splits[: size + 1].copy_(
|
||||
flashmla_num_splits_src[: size + 1]
|
||||
)
|
||||
ref_flashmla_metadata.copy_(flashmla_metadata_src)
|
||||
|
||||
ref_results.append(
|
||||
{
|
||||
"cache_seqlens": ref_cache_seqlens,
|
||||
"cu_seqlens_k": ref_cu_seqlens_k,
|
||||
"page_table_1": ref_page_table_1,
|
||||
"nsa_cache_seqlens": ref_nsa_cache_seqlens,
|
||||
"nsa_cu_seqlens_k": ref_nsa_cu_seqlens_k,
|
||||
"real_page_table": ref_real_page_table,
|
||||
"flashmla_num_splits": ref_flashmla_num_splits,
|
||||
"flashmla_metadata": ref_flashmla_metadata,
|
||||
}
|
||||
)
|
||||
|
||||
# Compare results for all 3 backends
|
||||
def check_tensor_equal(backend_idx, name, fused, ref):
|
||||
if not torch.equal(fused, ref):
|
||||
max_diff = (fused.float() - ref.float()).abs().max().item()
|
||||
mismatched_elements = (fused != ref).sum().item()
|
||||
total_elements = fused.numel()
|
||||
raise RuntimeError(
|
||||
f"MULTI-BACKEND FUSED METADATA COPY VERIFICATION FAILED!\n"
|
||||
f"Backend: {backend_idx}\n"
|
||||
f"Tensor: {name}\n"
|
||||
f"Max difference: {max_diff}\n"
|
||||
f"Mismatched elements: {mismatched_elements}/{total_elements}\n"
|
||||
f"Fused shape: {fused.shape}, Ref shape: {ref.shape}\n"
|
||||
f"Batch size: {bs}\n"
|
||||
f"The multi-backend fused kernel produces different results than individual copies.\n"
|
||||
f"This indicates a bug in the fused metadata copy kernel."
|
||||
)
|
||||
|
||||
# Verify all tensors for all 3 backends (multi-backend is DECODE mode only)
|
||||
for idx in range(3):
|
||||
fused = fused_results[idx]
|
||||
ref = ref_results[idx]
|
||||
|
||||
check_tensor_equal(
|
||||
idx,
|
||||
"cache_seqlens",
|
||||
fused["cache_seqlens"],
|
||||
ref["cache_seqlens"],
|
||||
)
|
||||
check_tensor_equal(
|
||||
idx,
|
||||
"cu_seqlens_k",
|
||||
fused["cu_seqlens_k"],
|
||||
ref["cu_seqlens_k"],
|
||||
)
|
||||
# Multi-backend is DECODE mode only, so compare only [:, :max_len]
|
||||
check_tensor_equal(
|
||||
idx,
|
||||
"page_table_1",
|
||||
fused["page_table_1"][:, : precomputed.max_len],
|
||||
ref["page_table_1"][:, : precomputed.max_len],
|
||||
)
|
||||
check_tensor_equal(
|
||||
idx,
|
||||
"nsa_cache_seqlens",
|
||||
fused["nsa_cache_seqlens"],
|
||||
ref["nsa_cache_seqlens"],
|
||||
)
|
||||
# DECODE mode uses bs for nsa_cu_seqlens_k size
|
||||
check_tensor_equal(
|
||||
idx,
|
||||
"nsa_cu_seqlens_k",
|
||||
fused["nsa_cu_seqlens_k"][: bs + 1],
|
||||
ref["nsa_cu_seqlens_k"][: bs + 1],
|
||||
)
|
||||
|
||||
if precomputed.real_page_table is not None:
|
||||
rows, cols = precomputed.real_page_table.shape
|
||||
check_tensor_equal(
|
||||
idx,
|
||||
"real_page_table",
|
||||
fused["real_page_table"][:rows, :cols],
|
||||
ref["real_page_table"][:rows, :cols],
|
||||
)
|
||||
|
||||
if precomputed.flashmla_metadata is not None:
|
||||
# DECODE mode uses bs + 1 for flashmla_num_splits
|
||||
check_tensor_equal(
|
||||
idx,
|
||||
"flashmla_num_splits",
|
||||
fused["flashmla_num_splits"][: bs + 1],
|
||||
ref["flashmla_num_splits"][: bs + 1],
|
||||
)
|
||||
check_tensor_equal(
|
||||
idx,
|
||||
"flashmla_metadata",
|
||||
fused["flashmla_metadata"],
|
||||
ref["flashmla_metadata"],
|
||||
)
|
||||
@@ -16,6 +16,10 @@ from sglang.srt.layers.attention.nsa.nsa_backend_mtp_precompute import (
|
||||
compute_cu_seqlens,
|
||||
)
|
||||
from sglang.srt.layers.attention.nsa.nsa_indexer import BaseIndexerMetadata
|
||||
from sglang.srt.layers.attention.nsa.nsa_mtp_verification import (
|
||||
verify_multi_backend_fused_metadata_copy,
|
||||
verify_single_backend_fused_metadata_copy,
|
||||
)
|
||||
from sglang.srt.layers.attention.nsa.quant_k_cache import quantize_k_cache
|
||||
from sglang.srt.layers.attention.nsa.transform_index import (
|
||||
transform_index_page_table_decode,
|
||||
@@ -63,6 +67,15 @@ else:
|
||||
# Reuse this workspace buffer across all NSA backend instances
|
||||
global_workspace_buffer = None
|
||||
|
||||
# Control whether to use fused metadata copy kernel (default: enabled)
|
||||
# Set SGLANG_USE_FUSED_METADATA_COPY=0 or false to disable
|
||||
_USE_FUSED_METADATA_COPY = envs.SGLANG_USE_FUSED_METADATA_COPY.get()
|
||||
|
||||
# Control whether to verify fused metadata copy against individual copies (default: disabled)
|
||||
# Set SGLANG_VERIFY_FUSED_METADATA_COPY=1 or true to enable verification
|
||||
# This will crash with detailed error message if any inconsistency is detected
|
||||
_VERIFY_FUSED_METADATA_COPY = envs.SGLANG_VERIFY_FUSED_METADATA_COPY.get()
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class NSAFlashMLAMetadata:
|
||||
@@ -1127,55 +1140,150 @@ class NativeSparseAttnBackend(
|
||||
|
||||
metadata = self.decode_cuda_graph_metadata[bs]
|
||||
|
||||
# Copy basic seqlens
|
||||
metadata.cache_seqlens_int32.copy_(precomputed.cache_seqlens)
|
||||
metadata.cu_seqlens_k[1:].copy_(precomputed.cu_seqlens_k[1:])
|
||||
# Track whether fused kernel succeeded
|
||||
fused_kernel_succeeded = False
|
||||
|
||||
# Mode-specific copy logic
|
||||
if forward_mode.is_decode_or_idle():
|
||||
# Decode mode
|
||||
metadata.page_table_1[:, : precomputed.max_len].copy_(
|
||||
precomputed.page_indices
|
||||
)
|
||||
metadata.nsa_cache_seqlens_int32.copy_(precomputed.nsa_cache_seqlens)
|
||||
# seqlens_expanded is same as cache_seqlens (already copied)
|
||||
# Use fused CUDA kernel for all copy operations
|
||||
if _USE_FUSED_METADATA_COPY:
|
||||
try:
|
||||
from sglang.jit_kernel.fused_metadata_copy import (
|
||||
fused_metadata_copy_cuda,
|
||||
)
|
||||
|
||||
elif forward_mode.is_target_verify():
|
||||
# Target verify mode
|
||||
metadata.page_table_1[:, : precomputed.max_seqlen_k].copy_(
|
||||
precomputed.page_indices
|
||||
)
|
||||
metadata.nsa_seqlens_expanded.copy_(precomputed.seqlens_expanded)
|
||||
metadata.nsa_cache_seqlens_int32.copy_(precomputed.nsa_cache_seqlens)
|
||||
# Map forward_mode to integer enum
|
||||
if forward_mode.is_decode_or_idle():
|
||||
mode_int = 0 # DECODE
|
||||
elif forward_mode.is_target_verify():
|
||||
mode_int = 1 # TARGET_VERIFY
|
||||
elif forward_mode.is_draft_extend():
|
||||
mode_int = 2 # DRAFT_EXTEND
|
||||
else:
|
||||
raise ValueError(f"Unsupported forward_mode: {forward_mode}")
|
||||
|
||||
elif forward_mode.is_draft_extend():
|
||||
# Draft extend mode
|
||||
rows = precomputed.page_indices.shape[0]
|
||||
cols = precomputed.max_seqlen_k
|
||||
metadata.page_table_1[:rows, :cols].copy_(precomputed.page_indices)
|
||||
# Prepare FlashMLA tensors if needed
|
||||
flashmla_num_splits_src = None
|
||||
flashmla_num_splits_dst = None
|
||||
flashmla_metadata_src = None
|
||||
flashmla_metadata_dst = None
|
||||
if precomputed.flashmla_metadata is not None:
|
||||
flashmla_num_splits_src = precomputed.flashmla_metadata.num_splits
|
||||
flashmla_num_splits_dst = metadata.flashmla_metadata.num_splits
|
||||
flashmla_metadata_src = (
|
||||
precomputed.flashmla_metadata.flashmla_metadata
|
||||
)
|
||||
flashmla_metadata_dst = metadata.flashmla_metadata.flashmla_metadata
|
||||
|
||||
# Call fused kernel
|
||||
fused_metadata_copy_cuda(
|
||||
# Source tensors
|
||||
precomputed.cache_seqlens,
|
||||
precomputed.cu_seqlens_k,
|
||||
precomputed.page_indices,
|
||||
precomputed.nsa_cache_seqlens,
|
||||
precomputed.seqlens_expanded,
|
||||
precomputed.nsa_cu_seqlens_k,
|
||||
precomputed.real_page_table,
|
||||
flashmla_num_splits_src,
|
||||
flashmla_metadata_src,
|
||||
# Destination tensors
|
||||
metadata.cache_seqlens_int32,
|
||||
metadata.cu_seqlens_k,
|
||||
metadata.page_table_1,
|
||||
metadata.nsa_cache_seqlens_int32,
|
||||
metadata.nsa_seqlens_expanded,
|
||||
metadata.nsa_cu_seqlens_k,
|
||||
(
|
||||
metadata.real_page_table
|
||||
if precomputed.real_page_table is not None
|
||||
else None
|
||||
),
|
||||
flashmla_num_splits_dst,
|
||||
flashmla_metadata_dst,
|
||||
# Parameters
|
||||
mode_int,
|
||||
bs,
|
||||
precomputed.max_len,
|
||||
precomputed.max_seqlen_k,
|
||||
precomputed.seqlens_expanded_size,
|
||||
)
|
||||
|
||||
# Successfully used fused kernel
|
||||
fused_kernel_succeeded = True
|
||||
|
||||
# Verification: compare fused kernel results against individual copies
|
||||
if _VERIFY_FUSED_METADATA_COPY:
|
||||
verify_single_backend_fused_metadata_copy(
|
||||
metadata=metadata,
|
||||
precomputed=precomputed,
|
||||
forward_mode=forward_mode,
|
||||
bs=bs,
|
||||
flashmla_num_splits_src=flashmla_num_splits_src,
|
||||
flashmla_metadata_src=flashmla_metadata_src,
|
||||
flashmla_num_splits_dst=flashmla_num_splits_dst,
|
||||
flashmla_metadata_dst=flashmla_metadata_dst,
|
||||
)
|
||||
except ImportError:
|
||||
print(
|
||||
"Warning: Fused metadata copy kernel not available, falling back to individual copies."
|
||||
)
|
||||
except Exception as e:
|
||||
print(
|
||||
f"Warning: Fused metadata copy kernel failed with error: {e}, falling back to individual copies."
|
||||
)
|
||||
|
||||
# Fallback to individual copy operations if fused kernel disabled or failed
|
||||
if not fused_kernel_succeeded:
|
||||
# Copy basic seqlens
|
||||
metadata.cache_seqlens_int32.copy_(precomputed.cache_seqlens)
|
||||
metadata.cu_seqlens_k[1:].copy_(precomputed.cu_seqlens_k[1:])
|
||||
|
||||
# Mode-specific copy logic
|
||||
if forward_mode.is_decode_or_idle():
|
||||
# Decode mode
|
||||
metadata.page_table_1[:, : precomputed.max_len].copy_(
|
||||
precomputed.page_indices
|
||||
)
|
||||
metadata.nsa_cache_seqlens_int32.copy_(precomputed.nsa_cache_seqlens)
|
||||
# seqlens_expanded is same as cache_seqlens (already copied)
|
||||
|
||||
elif forward_mode.is_target_verify():
|
||||
# Target verify mode
|
||||
metadata.page_table_1[:, : precomputed.max_seqlen_k].copy_(
|
||||
precomputed.page_indices
|
||||
)
|
||||
metadata.nsa_seqlens_expanded.copy_(precomputed.seqlens_expanded)
|
||||
metadata.nsa_cache_seqlens_int32.copy_(precomputed.nsa_cache_seqlens)
|
||||
|
||||
elif forward_mode.is_draft_extend():
|
||||
# Draft extend mode
|
||||
rows = precomputed.page_indices.shape[0]
|
||||
cols = precomputed.max_seqlen_k
|
||||
metadata.page_table_1[:rows, :cols].copy_(precomputed.page_indices)
|
||||
|
||||
size = precomputed.seqlens_expanded_size
|
||||
metadata.nsa_seqlens_expanded[:size].copy_(precomputed.seqlens_expanded)
|
||||
metadata.nsa_cache_seqlens_int32[:size].copy_(
|
||||
precomputed.nsa_cache_seqlens
|
||||
)
|
||||
|
||||
# Copy NSA cu_seqlens
|
||||
size = precomputed.seqlens_expanded_size
|
||||
metadata.nsa_seqlens_expanded[:size].copy_(precomputed.seqlens_expanded)
|
||||
metadata.nsa_cache_seqlens_int32[:size].copy_(precomputed.nsa_cache_seqlens)
|
||||
metadata.nsa_cu_seqlens_k[1 : 1 + size].copy_(
|
||||
precomputed.nsa_cu_seqlens_k[1 : 1 + size]
|
||||
)
|
||||
|
||||
# Copy NSA cu_seqlens
|
||||
size = precomputed.seqlens_expanded_size
|
||||
metadata.nsa_cu_seqlens_k[1 : 1 + size].copy_(
|
||||
precomputed.nsa_cu_seqlens_k[1 : 1 + size]
|
||||
)
|
||||
# Copy real page table
|
||||
if precomputed.real_page_table is not None:
|
||||
rows, cols = precomputed.real_page_table.shape
|
||||
metadata.real_page_table[:rows, :cols].copy_(
|
||||
precomputed.real_page_table
|
||||
)
|
||||
|
||||
# Copy real page table
|
||||
if precomputed.real_page_table is not None:
|
||||
rows, cols = precomputed.real_page_table.shape
|
||||
metadata.real_page_table[:rows, :cols].copy_(precomputed.real_page_table)
|
||||
else:
|
||||
# real_page_table is same as page_table_1 (already copied)
|
||||
pass
|
||||
|
||||
# Copy FlashMLA metadata
|
||||
if precomputed.flashmla_metadata is not None:
|
||||
flashmla_metadata = metadata.flashmla_metadata.slice(slice(0, size + 1))
|
||||
flashmla_metadata.copy_(precomputed.flashmla_metadata)
|
||||
# Copy FlashMLA metadata in fallback path
|
||||
if precomputed.flashmla_metadata is not None:
|
||||
size = precomputed.seqlens_expanded_size
|
||||
flashmla_metadata = metadata.flashmla_metadata.slice(slice(0, size + 1))
|
||||
flashmla_metadata.copy_(precomputed.flashmla_metadata)
|
||||
|
||||
self.forward_metadata = metadata
|
||||
|
||||
@@ -1958,15 +2066,163 @@ class NativeSparseAttnMultiStepBackend:
|
||||
spec_info=forward_batch.spec_info,
|
||||
)
|
||||
|
||||
# Fast copy to each backend (1-2x faster than computing N times)
|
||||
for i in range(self.speculative_num_steps):
|
||||
self.attn_backends[
|
||||
i
|
||||
].init_forward_metadata_replay_cuda_graph_from_precomputed(
|
||||
bs=bs,
|
||||
precomputed=precomputed,
|
||||
forward_mode=ForwardMode.DECODE,
|
||||
)
|
||||
# Use multi-backend fused copy when we have 3 or more backends
|
||||
# This is 3x faster than calling the single-backend copy 3 times
|
||||
if self.speculative_num_steps >= 3:
|
||||
try:
|
||||
from sglang.jit_kernel.fused_metadata_copy import (
|
||||
fused_metadata_copy_multi_cuda,
|
||||
)
|
||||
|
||||
metadata0 = self.attn_backends[0].decode_cuda_graph_metadata[bs]
|
||||
metadata1 = self.attn_backends[1].decode_cuda_graph_metadata[bs]
|
||||
metadata2 = self.attn_backends[2].decode_cuda_graph_metadata[bs]
|
||||
|
||||
# Set nsa_prefill_impl for first 3 backends (required by the method)
|
||||
for i in range(3):
|
||||
self.attn_backends[i].set_nsa_prefill_impl(forward_batch=None)
|
||||
|
||||
# Prepare FlashMLA tensors if needed
|
||||
flashmla_num_splits_src = None
|
||||
flashmla_metadata_src = None
|
||||
flashmla_num_splits_dst0 = None
|
||||
flashmla_num_splits_dst1 = None
|
||||
flashmla_num_splits_dst2 = None
|
||||
flashmla_metadata_dst0 = None
|
||||
flashmla_metadata_dst1 = None
|
||||
flashmla_metadata_dst2 = None
|
||||
|
||||
if precomputed.flashmla_metadata is not None:
|
||||
flashmla_num_splits_src = (
|
||||
precomputed.flashmla_metadata.num_splits
|
||||
)
|
||||
flashmla_metadata_src = (
|
||||
precomputed.flashmla_metadata.flashmla_metadata
|
||||
)
|
||||
flashmla_num_splits_dst0 = (
|
||||
metadata0.flashmla_metadata.num_splits
|
||||
)
|
||||
flashmla_num_splits_dst1 = (
|
||||
metadata1.flashmla_metadata.num_splits
|
||||
)
|
||||
flashmla_num_splits_dst2 = (
|
||||
metadata2.flashmla_metadata.num_splits
|
||||
)
|
||||
flashmla_metadata_dst0 = (
|
||||
metadata0.flashmla_metadata.flashmla_metadata
|
||||
)
|
||||
flashmla_metadata_dst1 = (
|
||||
metadata1.flashmla_metadata.flashmla_metadata
|
||||
)
|
||||
flashmla_metadata_dst2 = (
|
||||
metadata2.flashmla_metadata.flashmla_metadata
|
||||
)
|
||||
|
||||
# Call the multi-backend fused kernel for first 3 backends
|
||||
fused_metadata_copy_multi_cuda(
|
||||
# Source tensors
|
||||
precomputed.cache_seqlens,
|
||||
precomputed.cu_seqlens_k,
|
||||
precomputed.page_indices,
|
||||
precomputed.nsa_cache_seqlens,
|
||||
precomputed.nsa_cu_seqlens_k,
|
||||
precomputed.real_page_table,
|
||||
flashmla_num_splits_src,
|
||||
flashmla_metadata_src,
|
||||
# Destination tensors for backend 0
|
||||
metadata0.cache_seqlens_int32,
|
||||
metadata0.cu_seqlens_k,
|
||||
metadata0.page_table_1,
|
||||
metadata0.nsa_cache_seqlens_int32,
|
||||
metadata0.nsa_cu_seqlens_k,
|
||||
(
|
||||
metadata0.real_page_table
|
||||
if precomputed.real_page_table is not None
|
||||
else None
|
||||
),
|
||||
flashmla_num_splits_dst0,
|
||||
flashmla_metadata_dst0,
|
||||
# Destination tensors for backend 1
|
||||
metadata1.cache_seqlens_int32,
|
||||
metadata1.cu_seqlens_k,
|
||||
metadata1.page_table_1,
|
||||
metadata1.nsa_cache_seqlens_int32,
|
||||
metadata1.nsa_cu_seqlens_k,
|
||||
(
|
||||
metadata1.real_page_table
|
||||
if precomputed.real_page_table is not None
|
||||
else None
|
||||
),
|
||||
flashmla_num_splits_dst1,
|
||||
flashmla_metadata_dst1,
|
||||
# Destination tensors for backend 2
|
||||
metadata2.cache_seqlens_int32,
|
||||
metadata2.cu_seqlens_k,
|
||||
metadata2.page_table_1,
|
||||
metadata2.nsa_cache_seqlens_int32,
|
||||
metadata2.nsa_cu_seqlens_k,
|
||||
(
|
||||
metadata2.real_page_table
|
||||
if precomputed.real_page_table is not None
|
||||
else None
|
||||
),
|
||||
flashmla_num_splits_dst2,
|
||||
flashmla_metadata_dst2,
|
||||
# Parameters
|
||||
bs,
|
||||
precomputed.max_len,
|
||||
precomputed.seqlens_expanded_size,
|
||||
)
|
||||
|
||||
# Verification: compare fused kernel results against individual copies
|
||||
if _VERIFY_FUSED_METADATA_COPY:
|
||||
verify_multi_backend_fused_metadata_copy(
|
||||
metadata0=metadata0,
|
||||
metadata1=metadata1,
|
||||
metadata2=metadata2,
|
||||
precomputed=precomputed,
|
||||
bs=bs,
|
||||
flashmla_num_splits_src=flashmla_num_splits_src,
|
||||
flashmla_metadata_src=flashmla_metadata_src,
|
||||
)
|
||||
|
||||
# Copy remaining backends one by one (if > 3 backends)
|
||||
for i in range(3, self.speculative_num_steps):
|
||||
self.attn_backends[
|
||||
i
|
||||
].init_forward_metadata_replay_cuda_graph_from_precomputed(
|
||||
bs=bs,
|
||||
precomputed=precomputed,
|
||||
forward_mode=ForwardMode.DECODE,
|
||||
)
|
||||
except (ImportError, Exception) as e:
|
||||
# Fallback to loop if multi-backend kernel not available or fails
|
||||
if isinstance(e, ImportError):
|
||||
print(
|
||||
"Warning: Multi-backend fused metadata copy kernel not available, falling back to loop."
|
||||
)
|
||||
else:
|
||||
print(
|
||||
f"Warning: Multi-backend fused metadata copy kernel failed with error: {e}, falling back to loop."
|
||||
)
|
||||
for i in range(self.speculative_num_steps):
|
||||
self.attn_backends[
|
||||
i
|
||||
].init_forward_metadata_replay_cuda_graph_from_precomputed(
|
||||
bs=bs,
|
||||
precomputed=precomputed,
|
||||
forward_mode=ForwardMode.DECODE,
|
||||
)
|
||||
else:
|
||||
# Less than 3 backends: copy to each backend individually
|
||||
for i in range(self.speculative_num_steps):
|
||||
self.attn_backends[
|
||||
i
|
||||
].init_forward_metadata_replay_cuda_graph_from_precomputed(
|
||||
bs=bs,
|
||||
precomputed=precomputed,
|
||||
forward_mode=ForwardMode.DECODE,
|
||||
)
|
||||
else:
|
||||
# Fallback: compute metadata separately for each backend
|
||||
for i in range(self.speculative_num_steps):
|
||||
|
||||
Reference in New Issue
Block a user