Kernel: optimize decoding metadata in NSA multi-spec backend with fused kernels (#17554)

This commit is contained in:
Johnsonms
2026-02-14 16:40:15 +08:00
committed by GitHub
parent 38473f8ee0
commit 34132d6da5
7 changed files with 2824 additions and 54 deletions
+2
View File
@@ -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"],
)
+307 -51
View File
@@ -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):