[DeepSeek-V4] Add an opt-in non-paged indexer for long-context prefill (#29619)
This commit is contained in:
@@ -866,6 +866,7 @@ class Envs:
|
||||
SGLANG_OPT_FUSE_MHC_POST_PRE = EnvBool(False)
|
||||
SGLANG_OPT_USE_TILELANG_INDEXER = EnvBool(False)
|
||||
SGLANG_OPT_USE_AITER_INDEXER = EnvBool(False)
|
||||
SGLANG_OPT_DSV4_NONPAGED_INDEXER = EnvBool(False)
|
||||
SGLANG_OPT_USE_JIT_INDEXER_METADATA = EnvBool(True)
|
||||
SGLANG_OPT_USE_ONLINE_COMPRESS = EnvBool(False)
|
||||
SGLANG_EXPERIMENTAL_ONLINE_C128_MTP = EnvBool(False)
|
||||
|
||||
@@ -543,11 +543,17 @@ class DeepseekV4AttnBackend(
|
||||
online_state_slot_offset=online_c128_state_slot_offset,
|
||||
)
|
||||
|
||||
def init_forward_metadata_indexer(self, core_attn_metadata: DSV4AttnMetadata):
|
||||
def init_forward_metadata_indexer(
|
||||
self,
|
||||
core_attn_metadata: DSV4AttnMetadata,
|
||||
*,
|
||||
use_prefill_cuda_graph: bool = False,
|
||||
):
|
||||
return PagedIndexerMetadata(
|
||||
page_size=self.page_size,
|
||||
page_table=core_attn_metadata.page_table,
|
||||
c4_seq_lens=core_attn_metadata.c4_topk_lengths_raw,
|
||||
use_prefill_cuda_graph=use_prefill_cuda_graph,
|
||||
)
|
||||
|
||||
def init_forward_metadata_decode(
|
||||
@@ -630,7 +636,10 @@ class DeepseekV4AttnBackend(
|
||||
is_prefill=True,
|
||||
)
|
||||
indexer_metadata = (
|
||||
self.init_forward_metadata_indexer(core_attn_metadata)
|
||||
self.init_forward_metadata_indexer(
|
||||
core_attn_metadata,
|
||||
use_prefill_cuda_graph=use_prefill_cuda_graph,
|
||||
)
|
||||
if need_compress
|
||||
else None
|
||||
)
|
||||
|
||||
@@ -17,10 +17,21 @@ from sglang.jit_kernel.dsv4 import (
|
||||
from sglang.srt.configs.deepseek_v4 import DeepSeekV4Config
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.layers.attention.dsv4.compressor import Compressor
|
||||
from sglang.srt.layers.attention.dsv4.metadata import PagedIndexerMetadata
|
||||
from sglang.srt.layers.attention.dsv4.metadata import (
|
||||
NonPagedIndexerPlan,
|
||||
PagedIndexerMetadata,
|
||||
)
|
||||
from sglang.srt.layers.dp_attention import get_attention_cp_size
|
||||
from sglang.srt.layers.linear import ReplicatedLinear
|
||||
from sglang.srt.model_executor.forward_batch_info import ForwardMode
|
||||
from sglang.srt.model_executor.runner_backend_utils.breakable_cuda_graph.context import (
|
||||
is_in_breakable_cuda_graph,
|
||||
)
|
||||
from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import (
|
||||
is_in_tc_piecewise_cuda_graph,
|
||||
)
|
||||
from sglang.srt.state_capturer.indexer_topk import get_global_indexer_capturer
|
||||
from sglang.srt.utils import add_prefix, is_hip
|
||||
from sglang.srt.utils import add_prefix, is_cuda, is_hip
|
||||
from sglang.srt.utils.common import is_sm120_supported
|
||||
|
||||
if TYPE_CHECKING:
|
||||
@@ -360,10 +371,9 @@ class C4IndexerBackendMixin:
|
||||
c4_indexer: C4Indexer,
|
||||
positions: torch.Tensor,
|
||||
forward_batch: ForwardBatch,
|
||||
token_to_kv_pool: DeepSeekV4TokenToKVPool,
|
||||
alt_streams: Optional[List[torch.cuda.Stream]] = None,
|
||||
q_lora_ready: Optional[torch.cuda.Event] = None,
|
||||
) -> Tuple[IndexerQuery, torch.Tensor, torch.Tensor]:
|
||||
) -> Tuple[IndexerQuery, torch.Tensor]:
|
||||
if TYPE_CHECKING:
|
||||
assert isinstance(self, CompressorBackendMixin)
|
||||
|
||||
@@ -382,9 +392,6 @@ class C4IndexerBackendMixin:
|
||||
layer_id=c4_indexer.layer_id,
|
||||
compressor=c4_indexer.compressor,
|
||||
)
|
||||
c4_indexer_kv_cache = token_to_kv_pool.get_index_k_with_scale_buffer(
|
||||
layer_id=c4_indexer.layer_id,
|
||||
)
|
||||
|
||||
# The weight projection is small and fast; compute it on its own
|
||||
# stream, then have the Q stream wait on it before launching the big
|
||||
@@ -401,7 +408,7 @@ class C4IndexerBackendMixin:
|
||||
q, weights = c4_indexer.compute_q(q_lora, positions, weights)
|
||||
|
||||
current_stream.wait_stream(stream_q)
|
||||
return q, weights, c4_indexer_kv_cache
|
||||
return q, weights
|
||||
|
||||
def _forward_prepare_normal(
|
||||
self,
|
||||
@@ -410,9 +417,8 @@ class C4IndexerBackendMixin:
|
||||
c4_indexer: C4Indexer,
|
||||
positions: torch.Tensor,
|
||||
forward_batch: ForwardBatch,
|
||||
token_to_kv_pool: DeepSeekV4TokenToKVPool,
|
||||
skip_compressor: bool = False,
|
||||
) -> Tuple[IndexerQuery, torch.Tensor, torch.Tensor]:
|
||||
) -> Tuple[IndexerQuery, torch.Tensor]:
|
||||
if TYPE_CHECKING:
|
||||
assert isinstance(self, CompressorBackendMixin)
|
||||
|
||||
@@ -425,10 +431,159 @@ class C4IndexerBackendMixin:
|
||||
layer_id=c4_indexer.layer_id,
|
||||
compressor=c4_indexer.compressor,
|
||||
)
|
||||
c4_indexer_kv_cache = token_to_kv_pool.get_index_k_with_scale_buffer(
|
||||
return q, weights
|
||||
|
||||
def _can_use_nonpaged_indexer(
|
||||
self,
|
||||
*,
|
||||
c4_indexer: C4Indexer,
|
||||
forward_batch: ForwardBatch,
|
||||
indexer_metadata: PagedIndexerMetadata,
|
||||
) -> bool:
|
||||
if not envs.SGLANG_OPT_DSV4_NONPAGED_INDEXER.get():
|
||||
return False
|
||||
# This path calls CUDA DeepGEMM and assumes the CUDA FP8+FP32 packed
|
||||
# indexer cache layout. Explicitly reject HIP, NPU, and other devices.
|
||||
if not is_cuda() or is_hip():
|
||||
return False
|
||||
# The gather plan is built from eager, child-local ForwardBatch metadata.
|
||||
# Rewritten, TBO-split, and graph-backed batches must use the paged path.
|
||||
if (
|
||||
forward_batch.forward_mode != ForwardMode.EXTEND
|
||||
or forward_batch._original_forward_mode is not None
|
||||
or forward_batch.tbo_parent_token_range is not None
|
||||
or forward_batch.batch_size != 1
|
||||
or indexer_metadata.use_prefill_cuda_graph
|
||||
):
|
||||
return False
|
||||
if (
|
||||
c4_indexer.use_fp4_indexer
|
||||
or envs.SGLANG_OPT_USE_TILELANG_INDEXER.get()
|
||||
or envs.SGLANG_OPT_USE_AITER_INDEXER.get()
|
||||
or envs.SGLANG_FP8_PAGED_MQA_LOGITS_TORCH.get()
|
||||
):
|
||||
return False
|
||||
if (
|
||||
get_attention_cp_size() != 1
|
||||
or self.hisparse_coordinator is not None
|
||||
or is_in_tc_piecewise_cuda_graph()
|
||||
or is_in_breakable_cuda_graph()
|
||||
):
|
||||
return False
|
||||
return not torch.cuda.is_current_stream_capturing()
|
||||
|
||||
def _get_nonpaged_indexer_plan(
|
||||
self,
|
||||
*,
|
||||
c4_indexer: C4Indexer,
|
||||
forward_batch: ForwardBatch,
|
||||
indexer_metadata: PagedIndexerMetadata,
|
||||
page_table: torch.Tensor,
|
||||
c4_seq_lens: torch.Tensor,
|
||||
query_rows: int,
|
||||
) -> Optional[NonPagedIndexerPlan]:
|
||||
if not self._can_use_nonpaged_indexer(
|
||||
c4_indexer=c4_indexer,
|
||||
forward_batch=forward_batch,
|
||||
indexer_metadata=indexer_metadata,
|
||||
):
|
||||
return None
|
||||
if indexer_metadata.nonpaged_plan is not None:
|
||||
return indexer_metadata.nonpaged_plan
|
||||
|
||||
if (
|
||||
forward_batch.seq_lens is None
|
||||
or forward_batch.seq_lens_cpu is None
|
||||
or forward_batch.extend_seq_lens_cpu is None
|
||||
or forward_batch.extend_seq_lens is None
|
||||
or forward_batch.extend_start_loc is None
|
||||
or forward_batch.extend_num_tokens is None
|
||||
):
|
||||
return None
|
||||
|
||||
def to_cpu_int_list(values) -> Optional[List[int]]:
|
||||
if isinstance(values, torch.Tensor):
|
||||
if values.device.type != "cpu":
|
||||
return None
|
||||
values = values.tolist()
|
||||
return [int(value) for value in values]
|
||||
|
||||
extend_lens_cpu = to_cpu_int_list(forward_batch.extend_seq_lens_cpu)
|
||||
seq_lens_cpu = to_cpu_int_list(forward_batch.seq_lens_cpu)
|
||||
if (
|
||||
extend_lens_cpu is None
|
||||
or seq_lens_cpu is None
|
||||
or len(extend_lens_cpu) != 1
|
||||
or len(seq_lens_cpu) != 1
|
||||
or extend_lens_cpu[0] <= 0
|
||||
):
|
||||
return None
|
||||
|
||||
actual_queries = extend_lens_cpu[0]
|
||||
if (
|
||||
actual_queries != query_rows
|
||||
or int(forward_batch.extend_num_tokens) != query_rows
|
||||
or forward_batch.seq_lens.numel() != 1
|
||||
or forward_batch.extend_seq_lens.numel() != 1
|
||||
or forward_batch.extend_start_loc.numel() != 1
|
||||
or page_table.dim() != 2
|
||||
or page_table.shape[0] < query_rows
|
||||
or c4_seq_lens.numel() < query_rows
|
||||
):
|
||||
return None
|
||||
|
||||
final_c4_len = seq_lens_cpu[0] // 4
|
||||
if final_c4_len <= 0:
|
||||
return None
|
||||
|
||||
request_page_table = page_table[:1].contiguous()
|
||||
ke = c4_seq_lens[:query_rows].reshape(-1).to(torch.int32).contiguous()
|
||||
gather_seq_lens = ke[-1:]
|
||||
ks = torch.zeros_like(ke)
|
||||
c4_page_size = indexer_metadata.c4_page_size
|
||||
max_seqlen_k = (final_c4_len + c4_page_size - 1) // c4_page_size * c4_page_size
|
||||
plan = NonPagedIndexerPlan(
|
||||
page_table=request_page_table,
|
||||
gather_seq_lens=gather_seq_lens,
|
||||
ks=ks,
|
||||
ke=ke,
|
||||
seq_len_sum=final_c4_len,
|
||||
max_seq_len=final_c4_len,
|
||||
max_seqlen_k=max_seqlen_k,
|
||||
query_rows=query_rows,
|
||||
)
|
||||
indexer_metadata.nonpaged_plan = plan
|
||||
return plan
|
||||
|
||||
@staticmethod
|
||||
def _forward_nonpaged_indexer(
|
||||
*,
|
||||
q_indexer: torch.Tensor,
|
||||
weights: torch.Tensor,
|
||||
c4_indexer: C4Indexer,
|
||||
token_to_kv_pool: DeepSeekV4TokenToKVPool,
|
||||
plan: NonPagedIndexerPlan,
|
||||
) -> torch.Tensor:
|
||||
import deep_gemm
|
||||
|
||||
k_u8, scale_u8 = token_to_kv_pool.get_index_k_scale_buffer(
|
||||
layer_id=c4_indexer.layer_id,
|
||||
seq_len_tensor=plan.gather_seq_lens,
|
||||
page_indices=plan.page_table,
|
||||
seq_len_sum=plan.seq_len_sum,
|
||||
max_seq_len=plan.max_seq_len,
|
||||
)
|
||||
k_fp8 = k_u8.view(FP8_DTYPE)
|
||||
k_scale = scale_u8.view(torch.float32).squeeze(-1)
|
||||
return deep_gemm.fp8_mqa_logits(
|
||||
q_indexer[: plan.query_rows],
|
||||
(k_fp8, k_scale),
|
||||
weights[: plan.query_rows],
|
||||
plan.ks,
|
||||
plan.ke,
|
||||
clean_logits=False,
|
||||
max_seqlen_k=plan.max_seqlen_k,
|
||||
)
|
||||
return q, weights, c4_indexer_kv_cache
|
||||
|
||||
def forward_c4_indexer(
|
||||
self,
|
||||
@@ -465,35 +620,27 @@ class C4IndexerBackendMixin:
|
||||
positions = positions[:num_queries]
|
||||
|
||||
if enable_multi_stream:
|
||||
q_indexer, weights, c4_indexer_kv_cache = (
|
||||
self._forward_prepare_multi_stream(
|
||||
x=x,
|
||||
q_lora=q_lora,
|
||||
c4_indexer=c4_indexer,
|
||||
positions=positions,
|
||||
forward_batch=forward_batch,
|
||||
token_to_kv_pool=token_to_kv_pool,
|
||||
alt_streams=alt_streams,
|
||||
q_lora_ready=q_lora_ready,
|
||||
)
|
||||
q_indexer, weights = self._forward_prepare_multi_stream(
|
||||
x=x,
|
||||
q_lora=q_lora,
|
||||
c4_indexer=c4_indexer,
|
||||
positions=positions,
|
||||
forward_batch=forward_batch,
|
||||
alt_streams=alt_streams,
|
||||
q_lora_ready=q_lora_ready,
|
||||
)
|
||||
else:
|
||||
assert q_lora_ready is None
|
||||
q_indexer, weights, c4_indexer_kv_cache = self._forward_prepare_normal(
|
||||
q_indexer, weights = self._forward_prepare_normal(
|
||||
x=x,
|
||||
q_lora=q_lora,
|
||||
c4_indexer=c4_indexer,
|
||||
positions=positions,
|
||||
forward_batch=forward_batch,
|
||||
token_to_kv_pool=token_to_kv_pool,
|
||||
skip_compressor=skip_compressor,
|
||||
)
|
||||
|
||||
assert len(c4_indexer_kv_cache.shape) == 2
|
||||
block_kv = 64
|
||||
num_heads_kv = 1
|
||||
use_fp4_indexer = c4_indexer.use_fp4_indexer
|
||||
head_dim_with_sf = 68 if use_fp4_indexer else 132
|
||||
|
||||
if use_fp4_indexer:
|
||||
q_fp4, q_sf = q_indexer
|
||||
@@ -504,9 +651,6 @@ class C4IndexerBackendMixin:
|
||||
assert len(q_indexer.shape) == 3
|
||||
q = q_indexer.unsqueeze(1)
|
||||
|
||||
c4_indexer_kv_cache = c4_indexer_kv_cache.view(
|
||||
c4_indexer_kv_cache.shape[0], block_kv, num_heads_kv, head_dim_with_sf
|
||||
)
|
||||
assert len(weights.shape) == 3
|
||||
weights = weights.squeeze(2)
|
||||
if use_fp4_indexer:
|
||||
@@ -550,16 +694,42 @@ class C4IndexerBackendMixin:
|
||||
_use_aiter = envs.SGLANG_OPT_USE_AITER_INDEXER.get() and not use_fp4_indexer
|
||||
if _c4sl.dim() == 1 and not _use_tilelang and not _use_aiter:
|
||||
_c4sl = _c4sl.unsqueeze(-1)
|
||||
logits = fn(
|
||||
q,
|
||||
c4_indexer_kv_cache,
|
||||
weights,
|
||||
_c4sl,
|
||||
page_table,
|
||||
indexer_metadata.deep_gemm_metadata,
|
||||
indexer_metadata.max_c4_seq_len,
|
||||
False,
|
||||
nonpaged_plan = self._get_nonpaged_indexer_plan(
|
||||
c4_indexer=c4_indexer,
|
||||
forward_batch=forward_batch,
|
||||
indexer_metadata=indexer_metadata,
|
||||
page_table=page_table,
|
||||
c4_seq_lens=c4_seq_lens,
|
||||
query_rows=query_rows,
|
||||
)
|
||||
if nonpaged_plan is not None:
|
||||
assert isinstance(q_indexer, torch.Tensor)
|
||||
logits = self._forward_nonpaged_indexer(
|
||||
q_indexer=q_indexer,
|
||||
weights=weights,
|
||||
c4_indexer=c4_indexer,
|
||||
token_to_kv_pool=token_to_kv_pool,
|
||||
plan=nonpaged_plan,
|
||||
)
|
||||
else:
|
||||
c4_indexer_kv_cache = token_to_kv_pool.get_index_k_with_scale_buffer(
|
||||
layer_id=c4_indexer.layer_id,
|
||||
)
|
||||
assert c4_indexer_kv_cache.dim() == 2
|
||||
head_dim_with_sf = 68 if use_fp4_indexer else 132
|
||||
c4_indexer_kv_cache = c4_indexer_kv_cache.view(
|
||||
c4_indexer_kv_cache.shape[0], 64, 1, head_dim_with_sf
|
||||
)
|
||||
logits = fn(
|
||||
q,
|
||||
c4_indexer_kv_cache,
|
||||
weights,
|
||||
_c4sl,
|
||||
page_table,
|
||||
indexer_metadata.deep_gemm_metadata,
|
||||
indexer_metadata.max_c4_seq_len,
|
||||
False,
|
||||
)
|
||||
|
||||
assert indexer_metadata.page_table is core_metadata.page_table
|
||||
if self.debug_use_external_c4_sparse_indices:
|
||||
|
||||
@@ -95,13 +95,29 @@ def copy_metadata(
|
||||
), f"{provided_fields - all_fields=}, {all_fields - provided_fields=}"
|
||||
|
||||
|
||||
@dataclass
|
||||
class NonPagedIndexerPlan:
|
||||
page_table: torch.Tensor
|
||||
gather_seq_lens: torch.Tensor
|
||||
ks: torch.Tensor
|
||||
ke: torch.Tensor
|
||||
seq_len_sum: int
|
||||
max_seq_len: int
|
||||
max_seqlen_k: int
|
||||
query_rows: int
|
||||
|
||||
|
||||
@dataclass
|
||||
class PagedIndexerMetadata:
|
||||
page_size: int
|
||||
page_table: torch.Tensor
|
||||
c4_seq_lens: torch.Tensor
|
||||
use_prefill_cuda_graph: bool = False
|
||||
deep_gemm_metadata: Any = field(init=False, repr=False)
|
||||
topk_metadata: torch.Tensor = field(init=False, repr=False)
|
||||
nonpaged_plan: Optional[NonPagedIndexerPlan] = field(
|
||||
init=False, repr=False, default=None
|
||||
)
|
||||
|
||||
def __post_init__(self):
|
||||
if (
|
||||
@@ -156,18 +172,19 @@ class PagedIndexerMetadata:
|
||||
def copy_(self, other: PagedIndexerMetadata):
|
||||
if is_hip():
|
||||
copy_fields = ["page_table", "c4_seq_lens"]
|
||||
assign_fields = ["deep_gemm_metadata"]
|
||||
assign_fields = ["deep_gemm_metadata", "nonpaged_plan"]
|
||||
else:
|
||||
copy_fields = ["page_table", "c4_seq_lens", "deep_gemm_metadata"]
|
||||
assign_fields = []
|
||||
assign_fields = ["nonpaged_plan"]
|
||||
copy_fields += ["topk_metadata"]
|
||||
copy_metadata(
|
||||
src=other,
|
||||
dst=self,
|
||||
check_eq_fields=["page_size"],
|
||||
check_eq_fields=["page_size", "use_prefill_cuda_graph"],
|
||||
copy_fields=copy_fields,
|
||||
assign_fields=assign_fields,
|
||||
)
|
||||
self.nonpaged_plan = None
|
||||
|
||||
|
||||
def maybe_copy_inplace(dst, *, src) -> None:
|
||||
|
||||
@@ -321,12 +321,19 @@ class DeepSeekV4IndexerPool(KVCache):
|
||||
def get_index_k_scale_buffer(
|
||||
self,
|
||||
layer_id: int,
|
||||
seq_len: int,
|
||||
seq_len_tensor: torch.Tensor,
|
||||
page_indices: torch.Tensor,
|
||||
seq_len_sum: int,
|
||||
max_seq_len: int,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
buf = self.index_k_with_scale_buffer[layer_id]
|
||||
return index_buf_accessor.GetKAndS.execute(
|
||||
self, buf, seq_len=seq_len, page_indices=page_indices
|
||||
self,
|
||||
buf,
|
||||
page_indices=page_indices,
|
||||
seq_len_tensor=seq_len_tensor,
|
||||
seq_len_sum=seq_len_sum,
|
||||
max_seq_len=max_seq_len,
|
||||
)
|
||||
|
||||
def set_index_k_scale_buffer(
|
||||
@@ -1055,14 +1062,20 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool):
|
||||
def get_index_k_scale_buffer(
|
||||
self,
|
||||
layer_id: int,
|
||||
seq_len: int,
|
||||
seq_len_tensor: torch.Tensor,
|
||||
page_indices: torch.Tensor,
|
||||
seq_len_sum: int,
|
||||
max_seq_len: int,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
self.wait_layer_transfer(layer_id)
|
||||
compress_ratio, compress_layer_id, _ = self.layer_mapping[layer_id]
|
||||
assert compress_ratio == 4, f"only c4 has indexer, got {compress_ratio = }"
|
||||
return self.c4_indexer_kv_pool.get_index_k_scale_buffer(
|
||||
compress_layer_id, seq_len, page_indices
|
||||
compress_layer_id,
|
||||
seq_len_tensor,
|
||||
page_indices,
|
||||
seq_len_sum,
|
||||
max_seq_len,
|
||||
)
|
||||
|
||||
def set_index_k_scale_buffer(
|
||||
|
||||
@@ -0,0 +1,208 @@
|
||||
import sys
|
||||
import unittest
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.layers.attention.dsv4.indexer import FP8_DTYPE, C4IndexerBackendMixin
|
||||
from sglang.srt.layers.attention.dsv4.metadata import NonPagedIndexerPlan
|
||||
from sglang.srt.model_executor.forward_batch_info import ForwardMode
|
||||
from sglang.test.ci.ci_register import register_cpu_ci
|
||||
from sglang.test.test_utils import CustomTestCase
|
||||
|
||||
register_cpu_ci(est_time=2, suite="base-a-test-cpu")
|
||||
|
||||
_INDEXER = "sglang.srt.layers.attention.dsv4.indexer"
|
||||
|
||||
|
||||
class TestDSV4NonPagedIndexer(CustomTestCase):
|
||||
def _is_eligible(self, **overrides):
|
||||
backend = SimpleNamespace(hisparse_coordinator=None)
|
||||
c4_indexer = SimpleNamespace(use_fp4_indexer=overrides.get("fp4", False))
|
||||
forward_batch = SimpleNamespace(
|
||||
forward_mode=overrides.get("mode", ForwardMode.EXTEND),
|
||||
_original_forward_mode=overrides.get("original_mode"),
|
||||
tbo_parent_token_range=overrides.get("tbo"),
|
||||
batch_size=overrides.get("batch_size", 1),
|
||||
)
|
||||
metadata = SimpleNamespace(
|
||||
use_prefill_cuda_graph=overrides.get("prefill_graph", False)
|
||||
)
|
||||
with (
|
||||
envs.SGLANG_OPT_DSV4_NONPAGED_INDEXER.override(
|
||||
overrides.get("enabled", True)
|
||||
),
|
||||
envs.SGLANG_OPT_USE_TILELANG_INDEXER.override(False),
|
||||
envs.SGLANG_OPT_USE_AITER_INDEXER.override(False),
|
||||
envs.SGLANG_FP8_PAGED_MQA_LOGITS_TORCH.override(False),
|
||||
patch(f"{_INDEXER}.is_cuda", return_value=True),
|
||||
patch(f"{_INDEXER}.is_hip", return_value=False),
|
||||
patch(f"{_INDEXER}.get_attention_cp_size", return_value=1),
|
||||
patch(
|
||||
f"{_INDEXER}.is_in_tc_piecewise_cuda_graph",
|
||||
return_value=overrides.get("piecewise_graph", False),
|
||||
),
|
||||
patch(f"{_INDEXER}.is_in_breakable_cuda_graph", return_value=False),
|
||||
patch("torch.cuda.is_current_stream_capturing", return_value=False),
|
||||
):
|
||||
return C4IndexerBackendMixin._can_use_nonpaged_indexer(
|
||||
backend,
|
||||
c4_indexer=c4_indexer,
|
||||
forward_batch=forward_batch,
|
||||
indexer_metadata=metadata,
|
||||
)
|
||||
|
||||
def test_eligibility_is_fail_closed(self):
|
||||
self.assertIs(envs.SGLANG_OPT_DSV4_NONPAGED_INDEXER.default, False)
|
||||
self.assertTrue(self._is_eligible())
|
||||
for case in (
|
||||
{"enabled": False},
|
||||
{"mode": ForwardMode.DECODE},
|
||||
{"original_mode": ForwardMode.DECODE},
|
||||
{"batch_size": 2},
|
||||
{"batch_size": 20_000},
|
||||
{"tbo": (1, 2)},
|
||||
{"prefill_graph": True},
|
||||
{"piecewise_graph": True},
|
||||
{"fp4": True},
|
||||
):
|
||||
with self.subTest(case=case):
|
||||
self.assertFalse(self._is_eligible(**case))
|
||||
|
||||
def test_single_request_plan_contract(self):
|
||||
backend = SimpleNamespace(_can_use_nonpaged_indexer=lambda **_: True)
|
||||
c4_indexer = SimpleNamespace(use_fp4_indexer=False)
|
||||
query_rows = 4
|
||||
batch = SimpleNamespace(
|
||||
seq_lens=torch.tensor([262], dtype=torch.int32),
|
||||
seq_lens_cpu=[262],
|
||||
extend_seq_lens_cpu=[query_rows],
|
||||
extend_seq_lens=torch.tensor([query_rows], dtype=torch.int32),
|
||||
extend_start_loc=torch.tensor([0], dtype=torch.int32),
|
||||
extend_num_tokens=query_rows,
|
||||
)
|
||||
metadata = SimpleNamespace(nonpaged_plan=None, c4_page_size=64)
|
||||
page_table = torch.tensor([[3, 1]], dtype=torch.int32).repeat(query_rows, 1)
|
||||
c4_seq_lens = torch.tensor([62, 63, 64, 65], dtype=torch.int32)
|
||||
|
||||
def build_plan():
|
||||
return C4IndexerBackendMixin._get_nonpaged_indexer_plan(
|
||||
backend,
|
||||
c4_indexer=c4_indexer,
|
||||
forward_batch=batch,
|
||||
indexer_metadata=metadata,
|
||||
page_table=page_table,
|
||||
c4_seq_lens=c4_seq_lens,
|
||||
query_rows=query_rows,
|
||||
)
|
||||
|
||||
plan = build_plan()
|
||||
self.assertEqual(
|
||||
(plan.seq_len_sum, plan.max_seqlen_k, plan.query_rows),
|
||||
(65, 128, query_rows),
|
||||
)
|
||||
torch.testing.assert_close(plan.page_table, page_table[:1])
|
||||
torch.testing.assert_close(plan.ke, c4_seq_lens)
|
||||
torch.testing.assert_close(plan.gather_seq_lens, c4_seq_lens[-1:])
|
||||
|
||||
metadata.nonpaged_plan = None
|
||||
batch.extend_seq_lens_cpu = [2, 2]
|
||||
self.assertIsNone(build_plan())
|
||||
|
||||
def test_extreme_plan_metadata_is_bounded_and_fail_closed(self):
|
||||
backend = SimpleNamespace(_can_use_nonpaged_indexer=lambda **_: True)
|
||||
c4_indexer = SimpleNamespace(use_fp4_indexer=False)
|
||||
query_rows = 4
|
||||
batch = SimpleNamespace(
|
||||
seq_lens=torch.tensor([500_000], dtype=torch.int32),
|
||||
seq_lens_cpu=[500_000],
|
||||
extend_seq_lens_cpu=[query_rows],
|
||||
extend_seq_lens=torch.tensor([query_rows], dtype=torch.int32),
|
||||
extend_start_loc=torch.tensor([0], dtype=torch.int32),
|
||||
extend_num_tokens=query_rows,
|
||||
)
|
||||
metadata = SimpleNamespace(nonpaged_plan=None, c4_page_size=64)
|
||||
page_table = torch.zeros((query_rows, 1), dtype=torch.int32)
|
||||
c4_seq_lens = torch.tensor(
|
||||
[124_997, 124_998, 124_999, 125_000], dtype=torch.int32
|
||||
)
|
||||
|
||||
def build_plan():
|
||||
return C4IndexerBackendMixin._get_nonpaged_indexer_plan(
|
||||
backend,
|
||||
c4_indexer=c4_indexer,
|
||||
forward_batch=batch,
|
||||
indexer_metadata=metadata,
|
||||
page_table=page_table,
|
||||
c4_seq_lens=c4_seq_lens,
|
||||
query_rows=query_rows,
|
||||
)
|
||||
|
||||
plan = build_plan()
|
||||
self.assertEqual(plan.seq_len_sum, 125_000)
|
||||
self.assertEqual(plan.max_seq_len, 125_000)
|
||||
self.assertEqual(plan.max_seqlen_k, 125_056)
|
||||
|
||||
metadata.nonpaged_plan = None
|
||||
batch.seq_lens = torch.tensor([500_000, 200], dtype=torch.int32)
|
||||
batch.seq_lens_cpu = [500_000, 200]
|
||||
batch.extend_seq_lens_cpu = [2, 2]
|
||||
batch.extend_seq_lens = torch.tensor([2, 2], dtype=torch.int32)
|
||||
batch.extend_start_loc = torch.tensor([0, 2], dtype=torch.int32)
|
||||
self.assertIsNone(build_plan())
|
||||
|
||||
def test_nonpaged_dispatch_uses_gathered_kv_contract(self):
|
||||
query_rows = 4
|
||||
plan = NonPagedIndexerPlan(
|
||||
page_table=torch.tensor([[3, 1]], dtype=torch.int32),
|
||||
gather_seq_lens=torch.tensor([65], dtype=torch.int32),
|
||||
ks=torch.zeros(query_rows, dtype=torch.int32),
|
||||
ke=torch.tensor([62, 63, 64, 65], dtype=torch.int32),
|
||||
seq_len_sum=65,
|
||||
max_seq_len=65,
|
||||
max_seqlen_k=128,
|
||||
query_rows=query_rows,
|
||||
)
|
||||
q_indexer = torch.zeros((6, 2, 128), dtype=torch.uint8).view(FP8_DTYPE)
|
||||
weights = torch.ones((6, 2), dtype=torch.float32)
|
||||
k_u8 = torch.zeros((65, 128), dtype=torch.uint8)
|
||||
scale_u8 = torch.zeros((65, 4), dtype=torch.uint8)
|
||||
token_to_kv_pool = MagicMock()
|
||||
token_to_kv_pool.get_index_k_scale_buffer.return_value = (k_u8, scale_u8)
|
||||
c4_indexer = SimpleNamespace(layer_id=17)
|
||||
expected = MagicMock(name="logits")
|
||||
deep_gemm = SimpleNamespace(fp8_mqa_logits=MagicMock(return_value=expected))
|
||||
|
||||
with patch.dict(sys.modules, {"deep_gemm": deep_gemm}):
|
||||
actual = C4IndexerBackendMixin._forward_nonpaged_indexer(
|
||||
q_indexer=q_indexer,
|
||||
weights=weights,
|
||||
c4_indexer=c4_indexer,
|
||||
token_to_kv_pool=token_to_kv_pool,
|
||||
plan=plan,
|
||||
)
|
||||
|
||||
self.assertIs(actual, expected)
|
||||
token_to_kv_pool.get_index_k_scale_buffer.assert_called_once_with(
|
||||
layer_id=17,
|
||||
seq_len_tensor=plan.gather_seq_lens,
|
||||
page_indices=plan.page_table,
|
||||
seq_len_sum=65,
|
||||
max_seq_len=65,
|
||||
)
|
||||
call = deep_gemm.fp8_mqa_logits.call_args
|
||||
torch.testing.assert_close(call.args[0], q_indexer[:query_rows])
|
||||
torch.testing.assert_close(call.args[1][0], k_u8.view(FP8_DTYPE))
|
||||
torch.testing.assert_close(
|
||||
call.args[1][1], scale_u8.view(torch.float32).squeeze(-1)
|
||||
)
|
||||
torch.testing.assert_close(call.args[2], weights[:query_rows])
|
||||
torch.testing.assert_close(call.args[3], plan.ks)
|
||||
torch.testing.assert_close(call.args[4], plan.ke)
|
||||
self.assertEqual(call.kwargs, {"clean_logits": False, "max_seqlen_k": 128})
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
Reference in New Issue
Block a user