[dflash] fa3/fa4: device-side page table; drop seq_lens_cpu D2H sync (#29343)
This commit is contained in:
@@ -14,6 +14,9 @@ from sglang.srt.layers.attention.triton_ops.metadata import (
|
|||||||
normal_decode_set_metadata,
|
normal_decode_set_metadata,
|
||||||
prepare_swa_spec_page_table_triton,
|
prepare_swa_spec_page_table_triton,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.layers.attention.triton_ops.trtllm_mha_page_table import (
|
||||||
|
build_trtllm_mha_page_table,
|
||||||
|
)
|
||||||
from sglang.srt.layers.attention.utils import assert_buffer_fits
|
from sglang.srt.layers.attention.utils import assert_buffer_fits
|
||||||
from sglang.srt.layers.cp.base import CPAttentionBackendKind, get_cp_strategy
|
from sglang.srt.layers.cp.base import CPAttentionBackendKind, get_cp_strategy
|
||||||
from sglang.srt.layers.cp.utils import is_cp_v2_active
|
from sglang.srt.layers.cp.utils import is_cp_v2_active
|
||||||
@@ -26,7 +29,7 @@ from sglang.srt.mem_cache.memory_pool import KVWriteLoc
|
|||||||
from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool
|
from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool
|
||||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
|
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
|
||||||
from sglang.srt.server_args import get_global_server_args
|
from sglang.srt.server_args import get_global_server_args
|
||||||
from sglang.srt.speculative.spec_info import SpecInput
|
from sglang.srt.speculative.spec_info import SpecInput, SpeculativeAlgorithm
|
||||||
from sglang.srt.utils import get_compiler_backend
|
from sglang.srt.utils import get_compiler_backend
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
@@ -241,6 +244,16 @@ class FlashAttentionBackend(AttentionBackend):
|
|||||||
self.kv_cache_dtype = model_runner.kv_cache_dtype
|
self.kv_cache_dtype = model_runner.kv_cache_dtype
|
||||||
self.kv_cache_dtype_str = model_runner.server_args.kv_cache_dtype
|
self.kv_cache_dtype_str = model_runner.server_args.kv_cache_dtype
|
||||||
self.page_size = model_runner.page_size
|
self.page_size = model_runner.page_size
|
||||||
|
# Static page-table width (upper bound). The device-side page-table build
|
||||||
|
# sizes to this constant, so no runtime host max is needed.
|
||||||
|
self.max_num_pages = (
|
||||||
|
self.max_context_len + self.page_size - 1
|
||||||
|
) // self.page_size
|
||||||
|
# Opt out of the seq_lens_cpu D2H only for dflash (the worker adapted to
|
||||||
|
# the GPU-only relay); EAGLE/MTP/standalone/non-spec keep the CPU mirror.
|
||||||
|
self.needs_cpu_seq_lens = not SpeculativeAlgorithm.from_string(
|
||||||
|
model_runner.server_args.speculative_algorithm
|
||||||
|
).is_dflash()
|
||||||
self.use_mla = model_runner.model_config.attention_arch == AttentionArch.MLA
|
self.use_mla = model_runner.model_config.attention_arch == AttentionArch.MLA
|
||||||
self.skip_prefill = skip_prefill
|
self.skip_prefill = skip_prefill
|
||||||
self.attn_cp_size = model_runner.attn_cp_size
|
self.attn_cp_size = model_runner.attn_cp_size
|
||||||
@@ -489,6 +502,13 @@ class FlashAttentionBackend(AttentionBackend):
|
|||||||
seqlens_in_batch = forward_batch.seq_lens
|
seqlens_in_batch = forward_batch.seq_lens
|
||||||
batch_size = forward_batch.batch_size
|
batch_size = forward_batch.batch_size
|
||||||
device = seqlens_in_batch.device
|
device = seqlens_in_batch.device
|
||||||
|
# Eager path needs a host int for dynamic page-table sizing: the CPU
|
||||||
|
# mirror when published, else a local D2H (not the overlap hot path).
|
||||||
|
seq_lens_cpu = (
|
||||||
|
forward_batch.seq_lens_cpu
|
||||||
|
if forward_batch.seq_lens_cpu is not None
|
||||||
|
else seqlens_in_batch.cpu()
|
||||||
|
)
|
||||||
|
|
||||||
if forward_batch.forward_mode.is_decode_or_idle():
|
if forward_batch.forward_mode.is_decode_or_idle():
|
||||||
# Draft Decode
|
# Draft Decode
|
||||||
@@ -497,7 +517,7 @@ class FlashAttentionBackend(AttentionBackend):
|
|||||||
metadata.cache_seqlens_int32 = (
|
metadata.cache_seqlens_int32 = (
|
||||||
seqlens_in_batch + (self.speculative_step_id + 1)
|
seqlens_in_batch + (self.speculative_step_id + 1)
|
||||||
).to(torch.int32)
|
).to(torch.int32)
|
||||||
metadata.max_seq_len_k = forward_batch.seq_lens_cpu.max().item() + (
|
metadata.max_seq_len_k = seq_lens_cpu.max().item() + (
|
||||||
self.speculative_step_id + 1
|
self.speculative_step_id + 1
|
||||||
)
|
)
|
||||||
metadata.cu_seqlens_q = torch.arange(
|
metadata.cu_seqlens_q = torch.arange(
|
||||||
@@ -516,7 +536,7 @@ class FlashAttentionBackend(AttentionBackend):
|
|||||||
# Draft-extend's idle batch (padded for DP MLP-sync) has no
|
# Draft-extend's idle batch (padded for DP MLP-sync) has no
|
||||||
# tree; build plain metadata (padded output is discarded).
|
# tree; build plain metadata (padded output is discarded).
|
||||||
metadata.cache_seqlens_int32 = seqlens_in_batch.to(torch.int32)
|
metadata.cache_seqlens_int32 = seqlens_in_batch.to(torch.int32)
|
||||||
metadata.max_seq_len_k = forward_batch.seq_lens_cpu.max().item()
|
metadata.max_seq_len_k = seq_lens_cpu.max().item()
|
||||||
metadata.cu_seqlens_q = torch.arange(
|
metadata.cu_seqlens_q = torch.arange(
|
||||||
0, batch_size + 1, dtype=torch.int32, device=device
|
0, batch_size + 1, dtype=torch.int32, device=device
|
||||||
)
|
)
|
||||||
@@ -529,7 +549,7 @@ class FlashAttentionBackend(AttentionBackend):
|
|||||||
else:
|
else:
|
||||||
metadata.cache_seqlens_int32 = (seqlens_in_batch).to(torch.int32)
|
metadata.cache_seqlens_int32 = (seqlens_in_batch).to(torch.int32)
|
||||||
metadata.max_seq_len_q = self.topk
|
metadata.max_seq_len_q = self.topk
|
||||||
metadata.max_seq_len_k = forward_batch.seq_lens_cpu.max().item()
|
metadata.max_seq_len_k = seq_lens_cpu.max().item()
|
||||||
metadata.cu_seqlens_q = torch.arange(
|
metadata.cu_seqlens_q = torch.arange(
|
||||||
0,
|
0,
|
||||||
batch_size * self.topk + 1,
|
batch_size * self.topk + 1,
|
||||||
@@ -579,7 +599,7 @@ class FlashAttentionBackend(AttentionBackend):
|
|||||||
else:
|
else:
|
||||||
# Normal Decode
|
# Normal Decode
|
||||||
metadata.cache_seqlens_int32 = seqlens_in_batch.to(torch.int32)
|
metadata.cache_seqlens_int32 = seqlens_in_batch.to(torch.int32)
|
||||||
metadata.max_seq_len_k = forward_batch.seq_lens_cpu.max().item()
|
metadata.max_seq_len_k = seq_lens_cpu.max().item()
|
||||||
metadata.cu_seqlens_q = torch.arange(
|
metadata.cu_seqlens_q = torch.arange(
|
||||||
0, batch_size + 1, dtype=torch.int32, device=device
|
0, batch_size + 1, dtype=torch.int32, device=device
|
||||||
)
|
)
|
||||||
@@ -625,8 +645,7 @@ class FlashAttentionBackend(AttentionBackend):
|
|||||||
).to(torch.int32)
|
).to(torch.int32)
|
||||||
metadata.max_seq_len_q = self.speculative_num_draft_tokens
|
metadata.max_seq_len_q = self.speculative_num_draft_tokens
|
||||||
metadata.max_seq_len_k = (
|
metadata.max_seq_len_k = (
|
||||||
forward_batch.seq_lens_cpu.max().item()
|
seq_lens_cpu.max().item() + self.speculative_num_draft_tokens
|
||||||
+ self.speculative_num_draft_tokens
|
|
||||||
)
|
)
|
||||||
metadata.cu_seqlens_q = torch.arange(
|
metadata.cu_seqlens_q = torch.arange(
|
||||||
0,
|
0,
|
||||||
@@ -649,7 +668,7 @@ class FlashAttentionBackend(AttentionBackend):
|
|||||||
else:
|
else:
|
||||||
metadata.cache_seqlens_int32 = forward_batch.seq_lens.to(torch.int32)
|
metadata.cache_seqlens_int32 = forward_batch.seq_lens.to(torch.int32)
|
||||||
metadata.max_seq_len_q = self.speculative_num_draft_tokens
|
metadata.max_seq_len_q = self.speculative_num_draft_tokens
|
||||||
metadata.max_seq_len_k = forward_batch.seq_lens_cpu.max().item()
|
metadata.max_seq_len_k = seq_lens_cpu.max().item()
|
||||||
metadata.cu_seqlens_q = torch.arange(
|
metadata.cu_seqlens_q = torch.arange(
|
||||||
0,
|
0,
|
||||||
batch_size * self.speculative_num_draft_tokens + 1,
|
batch_size * self.speculative_num_draft_tokens + 1,
|
||||||
@@ -749,7 +768,7 @@ class FlashAttentionBackend(AttentionBackend):
|
|||||||
include_draft_extend_v2=True
|
include_draft_extend_v2=True
|
||||||
):
|
):
|
||||||
metadata.cache_seqlens_int32 = seqlens_in_batch.to(torch.int32)
|
metadata.cache_seqlens_int32 = seqlens_in_batch.to(torch.int32)
|
||||||
metadata.max_seq_len_k = forward_batch.seq_lens_cpu.max().item()
|
metadata.max_seq_len_k = seq_lens_cpu.max().item()
|
||||||
metadata.cu_seqlens_k = torch.nn.functional.pad(
|
metadata.cu_seqlens_k = torch.nn.functional.pad(
|
||||||
torch.cumsum(seqlens_in_batch, dim=0, dtype=torch.int32), (1, 0)
|
torch.cumsum(seqlens_in_batch, dim=0, dtype=torch.int32), (1, 0)
|
||||||
)
|
)
|
||||||
@@ -797,7 +816,7 @@ class FlashAttentionBackend(AttentionBackend):
|
|||||||
self._pa_swa_prefill_lens[
|
self._pa_swa_prefill_lens[
|
||||||
forward_batch.req_pool_indices[:batch_size]
|
forward_batch.req_pool_indices[:batch_size]
|
||||||
] = forward_batch.seq_lens[:batch_size].to(torch.int32)
|
] = forward_batch.seq_lens[:batch_size].to(torch.int32)
|
||||||
max_pf = int(forward_batch.seq_lens_cpu[:batch_size].max().item())
|
max_pf = int(seq_lens_cpu[:batch_size].max().item())
|
||||||
if max_pf > self._pa_swa_max_prefill_len:
|
if max_pf > self._pa_swa_max_prefill_len:
|
||||||
self._pa_swa_max_prefill_len = max_pf
|
self._pa_swa_max_prefill_len = max_pf
|
||||||
|
|
||||||
@@ -2257,6 +2276,16 @@ class FlashAttentionBackend(AttentionBackend):
|
|||||||
|
|
||||||
return metadata, metadata_expand
|
return metadata, metadata_expand
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _host_max_seq_len(
|
||||||
|
seq_lens_cpu: Optional[torch.Tensor], seq_lens: torch.Tensor
|
||||||
|
) -> int:
|
||||||
|
"""Host-side max KV length: the CPU mirror when published, else a local
|
||||||
|
D2H. For cold paths (topk>1, draft-extend, eager) that need a host max --
|
||||||
|
not the dflash hot path (topk=1, device-side build)."""
|
||||||
|
src = seq_lens_cpu if seq_lens_cpu is not None else seq_lens.cpu()
|
||||||
|
return src.max().item()
|
||||||
|
|
||||||
def _apply_cuda_graph_metadata(
|
def _apply_cuda_graph_metadata(
|
||||||
self,
|
self,
|
||||||
bs: int,
|
bs: int,
|
||||||
@@ -2277,7 +2306,10 @@ class FlashAttentionBackend(AttentionBackend):
|
|||||||
are gone.
|
are gone.
|
||||||
"""
|
"""
|
||||||
seq_lens = seq_lens[:bs]
|
seq_lens = seq_lens[:bs]
|
||||||
seq_lens_cpu = seq_lens_cpu[:bs]
|
# The GPU-only path passes seq_lens_cpu=None; the topk>1 branches below
|
||||||
|
# still need a host max, so sync locally in that case (not the dflash
|
||||||
|
# overlap hot path, which uses topk=1 and the device-side build).
|
||||||
|
seq_lens_cpu = seq_lens_cpu[:bs] if seq_lens_cpu is not None else None
|
||||||
req_pool_indices = req_pool_indices[:bs]
|
req_pool_indices = req_pool_indices[:bs]
|
||||||
device = seq_lens.device
|
device = seq_lens.device
|
||||||
metadata = None
|
metadata = None
|
||||||
@@ -2298,17 +2330,9 @@ class FlashAttentionBackend(AttentionBackend):
|
|||||||
if self.topk <= 1:
|
if self.topk <= 1:
|
||||||
# When topk = 1, we use the normal decode metadata
|
# When topk = 1, we use the normal decode metadata
|
||||||
metadata = self.decode_cuda_graph_metadata[bs]
|
metadata = self.decode_cuda_graph_metadata[bs]
|
||||||
max_len = seq_lens_cpu.max().item()
|
# Page table built on-device (self-guards on cache_seqlens);
|
||||||
metadata.max_seq_len_k = max_len + self.speculative_step_id + 1
|
# max_seq_len_k left unset -- unread here (scheduler_metadata
|
||||||
max_seq_pages = (
|
# is normal-decode-only).
|
||||||
metadata.max_seq_len_k + self.page_size - 1
|
|
||||||
) // self.page_size
|
|
||||||
|
|
||||||
assert_buffer_fits(
|
|
||||||
max_seq_pages,
|
|
||||||
metadata.page_table.shape[1],
|
|
||||||
"FA3 draft-decode page_table",
|
|
||||||
)
|
|
||||||
normal_decode_set_metadata(
|
normal_decode_set_metadata(
|
||||||
metadata.cache_seqlens_int32,
|
metadata.cache_seqlens_int32,
|
||||||
metadata.cu_seqlens_k,
|
metadata.cu_seqlens_k,
|
||||||
@@ -2316,7 +2340,7 @@ class FlashAttentionBackend(AttentionBackend):
|
|||||||
self.req_to_token,
|
self.req_to_token,
|
||||||
req_pool_indices,
|
req_pool_indices,
|
||||||
self.decode_cuda_graph_metadata["strided_indices"],
|
self.decode_cuda_graph_metadata["strided_indices"],
|
||||||
max_seq_pages,
|
self.max_num_pages,
|
||||||
seq_lens,
|
seq_lens,
|
||||||
self.speculative_step_id + 1,
|
self.speculative_step_id + 1,
|
||||||
self.page_size,
|
self.page_size,
|
||||||
@@ -2341,7 +2365,9 @@ class FlashAttentionBackend(AttentionBackend):
|
|||||||
# metadata.cu_seqlens_q already set in capture
|
# metadata.cu_seqlens_q already set in capture
|
||||||
# metadata.cu_seqlens_k is not needed
|
# metadata.cu_seqlens_k is not needed
|
||||||
|
|
||||||
metadata.max_seq_len_k = seq_lens_cpu.max().item()
|
metadata.max_seq_len_k = self._host_max_seq_len(
|
||||||
|
seq_lens_cpu, seq_lens
|
||||||
|
)
|
||||||
max_seq_pages = (
|
max_seq_pages = (
|
||||||
metadata.max_seq_len_k + self.page_size - 1
|
metadata.max_seq_len_k + self.page_size - 1
|
||||||
) // self.page_size
|
) // self.page_size
|
||||||
@@ -2389,16 +2415,11 @@ class FlashAttentionBackend(AttentionBackend):
|
|||||||
else:
|
else:
|
||||||
# Normal Decode
|
# Normal Decode
|
||||||
metadata = self.decode_cuda_graph_metadata[bs]
|
metadata = self.decode_cuda_graph_metadata[bs]
|
||||||
max_len = seq_lens_cpu.max().item()
|
|
||||||
max_seq_pages = (max_len + self.page_size - 1) // self.page_size
|
|
||||||
metadata.max_seq_len_k = max_len
|
|
||||||
|
|
||||||
assert_buffer_fits(
|
|
||||||
max_seq_pages,
|
|
||||||
metadata.page_table.shape[1],
|
|
||||||
"FA3 decode page_table",
|
|
||||||
)
|
|
||||||
if self.is_prefill_aware_swa:
|
if self.is_prefill_aware_swa:
|
||||||
|
# Prefill-aware SWA still needs a host max to bound the
|
||||||
|
# per-batch page table built below.
|
||||||
|
max_len = self._host_max_seq_len(seq_lens_cpu, seq_lens)
|
||||||
|
metadata.max_seq_len_k = max_len
|
||||||
pa_max_len = min(
|
pa_max_len = min(
|
||||||
self._pa_swa_max_prefill_len + self.sliding_window_size,
|
self._pa_swa_max_prefill_len + self.sliding_window_size,
|
||||||
max_len,
|
max_len,
|
||||||
@@ -2417,6 +2438,14 @@ class FlashAttentionBackend(AttentionBackend):
|
|||||||
dst_kv_lens=metadata.cache_seqlens_int32,
|
dst_kv_lens=metadata.cache_seqlens_int32,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
|
# Page table uses the static max_num_pages bound (no D2H).
|
||||||
|
# max_seq_len_k only feeds scheduler_metadata below, so use
|
||||||
|
# the free CPU mirror for a tight split heuristic when present.
|
||||||
|
metadata.max_seq_len_k = (
|
||||||
|
seq_lens_cpu.max().item()
|
||||||
|
if seq_lens_cpu is not None
|
||||||
|
else self.max_context_len
|
||||||
|
)
|
||||||
normal_decode_set_metadata(
|
normal_decode_set_metadata(
|
||||||
metadata.cache_seqlens_int32,
|
metadata.cache_seqlens_int32,
|
||||||
metadata.cu_seqlens_k,
|
metadata.cu_seqlens_k,
|
||||||
@@ -2424,7 +2453,7 @@ class FlashAttentionBackend(AttentionBackend):
|
|||||||
self.req_to_token,
|
self.req_to_token,
|
||||||
req_pool_indices,
|
req_pool_indices,
|
||||||
self.decode_cuda_graph_metadata["strided_indices"],
|
self.decode_cuda_graph_metadata["strided_indices"],
|
||||||
max_seq_pages,
|
self.max_num_pages,
|
||||||
seq_lens,
|
seq_lens,
|
||||||
0,
|
0,
|
||||||
self.page_size,
|
self.page_size,
|
||||||
@@ -2464,40 +2493,33 @@ class FlashAttentionBackend(AttentionBackend):
|
|||||||
(seq_lens + self.speculative_num_draft_tokens)
|
(seq_lens + self.speculative_num_draft_tokens)
|
||||||
)
|
)
|
||||||
|
|
||||||
metadata.max_seq_len_k = (
|
# Page table built on-device (self-guards on cache_seqlens);
|
||||||
seq_lens_cpu.max().item() + self.speculative_num_draft_tokens
|
# max_seq_len_k left unset -- unread here (scheduler_metadata is
|
||||||
)
|
# normal-decode-only).
|
||||||
metadata.cu_seqlens_k[1:].copy_(
|
metadata.cu_seqlens_k[1:].copy_(
|
||||||
torch.cumsum(metadata.cache_seqlens_int32, dim=0, dtype=torch.int32)
|
torch.cumsum(metadata.cache_seqlens_int32, dim=0, dtype=torch.int32)
|
||||||
)
|
)
|
||||||
max_seq_pages = (
|
has_swa = self.use_sliding_window_kv_pool
|
||||||
metadata.max_seq_len_k + self.page_size - 1
|
build_trtllm_mha_page_table(
|
||||||
) // self.page_size
|
req_to_token=self.req_to_token,
|
||||||
page_indices = self.req_to_token[
|
req_pool_indices=req_pool_indices,
|
||||||
req_pool_indices[:, None],
|
cache_seqlens=metadata.cache_seqlens_int32,
|
||||||
self.decode_cuda_graph_metadata["strided_indices"][:max_seq_pages],
|
page_table=metadata.page_table,
|
||||||
]
|
page_size=self.page_size,
|
||||||
if (
|
swa_page_table=metadata.swa_page_table if has_swa else None,
|
||||||
self.use_sliding_window_kv_pool
|
full_to_swa=(
|
||||||
and metadata.swa_page_table is not None
|
self.token_to_kv_pool.full_to_swa_index_mapping
|
||||||
):
|
if has_swa
|
||||||
swa_page_indices = (
|
else None
|
||||||
self.token_to_kv_pool.translate_loc_from_full_to_swa(
|
),
|
||||||
page_indices
|
|
||||||
)
|
)
|
||||||
)
|
|
||||||
metadata.swa_page_table[:, :max_seq_pages].copy_(
|
|
||||||
swa_page_indices // self.page_size
|
|
||||||
)
|
|
||||||
page_indices //= self.page_size
|
|
||||||
metadata.page_table[:, :max_seq_pages].copy_(page_indices)
|
|
||||||
else:
|
else:
|
||||||
# When topk > 1, we need two specific target verify metadata, and then merge states
|
# When topk > 1, we need two specific target verify metadata, and then merge states
|
||||||
# 1. The first half of metadata for prefix tokens
|
# 1. The first half of metadata for prefix tokens
|
||||||
metadata = self.target_verify_metadata_topk_normal[bs]
|
metadata = self.target_verify_metadata_topk_normal[bs]
|
||||||
metadata.cache_seqlens_int32.copy_(seq_lens)
|
metadata.cache_seqlens_int32.copy_(seq_lens)
|
||||||
# metadata.max_seq_len_q = self.speculative_num_draft_tokens, already set in capture
|
# metadata.max_seq_len_q = self.speculative_num_draft_tokens, already set in capture
|
||||||
metadata.max_seq_len_k = seq_lens_cpu.max().item()
|
metadata.max_seq_len_k = self._host_max_seq_len(seq_lens_cpu, seq_lens)
|
||||||
# metadata.cu_seqlens_q already set in capture
|
# metadata.cu_seqlens_q already set in capture
|
||||||
metadata.cu_seqlens_k[1:].copy_(
|
metadata.cu_seqlens_k[1:].copy_(
|
||||||
torch.cumsum(metadata.cache_seqlens_int32, dim=0, dtype=torch.int32)
|
torch.cumsum(metadata.cache_seqlens_int32, dim=0, dtype=torch.int32)
|
||||||
@@ -2583,7 +2605,7 @@ class FlashAttentionBackend(AttentionBackend):
|
|||||||
metadata = self.draft_extend_metadata[bs]
|
metadata = self.draft_extend_metadata[bs]
|
||||||
metadata.cache_seqlens_int32.copy_(seq_lens)
|
metadata.cache_seqlens_int32.copy_(seq_lens)
|
||||||
|
|
||||||
metadata.max_seq_len_k = seq_lens_cpu.max().item()
|
metadata.max_seq_len_k = self._host_max_seq_len(seq_lens_cpu, seq_lens)
|
||||||
metadata.cu_seqlens_k[1:].copy_(
|
metadata.cu_seqlens_k[1:].copy_(
|
||||||
torch.cumsum(metadata.cache_seqlens_int32, dim=0, dtype=torch.int32)
|
torch.cumsum(metadata.cache_seqlens_int32, dim=0, dtype=torch.int32)
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -6,6 +6,11 @@ it never reads a runtime max (no D2H sync). This test checks the device build is
|
|||||||
bit-identical to the legacy gather for the columns each request uses, for both
|
bit-identical to the legacy gather for the columns each request uses, for both
|
||||||
the full page table and the SWA-translated page table, across context lengths,
|
the full page table and the SWA-translated page table, across context lengths,
|
||||||
page sizes, and batch sizes.
|
page sizes, and batch sizes.
|
||||||
|
|
||||||
|
It also pins the invariant that lets the no-host-max (GPU-only) path hand the
|
||||||
|
kernel a static ``max_num_pages``-wide buffer: every column past a request's
|
||||||
|
page count must be left untouched, i.e. the kernel bounds its writes by the
|
||||||
|
device-side ``cache_seqlens`` alone.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
import unittest
|
import unittest
|
||||||
@@ -20,7 +25,7 @@ from sglang.test.ci.ci_register import register_cuda_ci
|
|||||||
from sglang.test.test_utils import CustomTestCase
|
from sglang.test.test_utils import CustomTestCase
|
||||||
|
|
||||||
# Triton kernel unit test for the trtllm_mha device-side page-table build.
|
# Triton kernel unit test for the trtllm_mha device-side page-table build.
|
||||||
register_cuda_ci(est_time=10, stage="base-b", runner_config="1-gpu-small")
|
register_cuda_ci(est_time=14, stage="base-b", runner_config="1-gpu-small")
|
||||||
|
|
||||||
|
|
||||||
def _build_page_table_reference(
|
def _build_page_table_reference(
|
||||||
@@ -143,6 +148,82 @@ class TestTrtllmMhaPageTable(CustomTestCase):
|
|||||||
max_ctx, page_size, num_reqs=max(64, bs), bs=bs, swa=True
|
max_ctx, page_size, num_reqs=max(64, bs), bs=bs, swa=True
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def _run_self_guard_case(self, max_context_len, page_size, bs, swa=False):
|
||||||
|
"""Short sequences against a full static buffer -- the GPU-only shape.
|
||||||
|
|
||||||
|
Pre-fill the page table with a sentinel and run the kernel with the
|
||||||
|
static ``max_num_pages`` width (no host max to tighten it). Used columns
|
||||||
|
must hold the right block ids; every tail column must keep the sentinel,
|
||||||
|
proving the kernel never writes past the device-side ``cache_seqlens``.
|
||||||
|
"""
|
||||||
|
torch.manual_seed(0)
|
||||||
|
dev = "cuda"
|
||||||
|
num_reqs = max(64, bs)
|
||||||
|
max_num_pages = (max_context_len + page_size - 1) // page_size
|
||||||
|
n_slots = num_reqs * max_context_len
|
||||||
|
req_to_token = torch.randint(
|
||||||
|
0, n_slots, (num_reqs, max_context_len), dtype=torch.int32, device=dev
|
||||||
|
)
|
||||||
|
req_pool_indices = torch.randperm(num_reqs, device=dev)[:bs].to(torch.int32)
|
||||||
|
# Cap lengths well below max_context_len so most tail columns stay unused.
|
||||||
|
hi = max(2, max_context_len // 8)
|
||||||
|
cache_seqlens = torch.randint(1, hi + 1, (bs,), dtype=torch.int32, device=dev)
|
||||||
|
full_to_swa = (
|
||||||
|
torch.randint(0, n_slots, (n_slots,), dtype=torch.int32, device=dev)
|
||||||
|
if swa
|
||||||
|
else None
|
||||||
|
)
|
||||||
|
|
||||||
|
SENTINEL = -1
|
||||||
|
page_table = torch.full(
|
||||||
|
(bs, max_num_pages), SENTINEL, dtype=torch.int32, device=dev
|
||||||
|
)
|
||||||
|
swa_page_table = (
|
||||||
|
torch.full((bs, max_num_pages), SENTINEL, dtype=torch.int32, device=dev)
|
||||||
|
if swa
|
||||||
|
else None
|
||||||
|
)
|
||||||
|
build_trtllm_mha_page_table(
|
||||||
|
req_to_token=req_to_token,
|
||||||
|
req_pool_indices=req_pool_indices,
|
||||||
|
cache_seqlens=cache_seqlens,
|
||||||
|
page_table=page_table,
|
||||||
|
page_size=page_size,
|
||||||
|
swa_page_table=swa_page_table,
|
||||||
|
full_to_swa=full_to_swa,
|
||||||
|
)
|
||||||
|
pt_ref, swa_ref = _build_page_table_reference(
|
||||||
|
req_to_token, req_pool_indices, cache_seqlens, page_size, full_to_swa
|
||||||
|
)
|
||||||
|
|
||||||
|
tag = f"max_ctx={max_context_len} page_size={page_size} bs={bs} swa={swa}"
|
||||||
|
for i in range(bs):
|
||||||
|
npages = (int(cache_seqlens[i].item()) + page_size - 1) // page_size
|
||||||
|
self.assertTrue(
|
||||||
|
torch.equal(page_table[i, :npages], pt_ref[i, :npages]),
|
||||||
|
f"used-column mismatch req={i} {tag}",
|
||||||
|
)
|
||||||
|
self.assertTrue(
|
||||||
|
torch.all(page_table[i, npages:] == SENTINEL),
|
||||||
|
f"kernel wrote past cache_seqlens req={i} npages={npages} {tag}",
|
||||||
|
)
|
||||||
|
if swa:
|
||||||
|
self.assertTrue(
|
||||||
|
torch.equal(swa_page_table[i, :npages], swa_ref[i, :npages]),
|
||||||
|
f"swa used-column mismatch req={i} {tag}",
|
||||||
|
)
|
||||||
|
self.assertTrue(
|
||||||
|
torch.all(swa_page_table[i, npages:] == SENTINEL),
|
||||||
|
f"swa wrote past cache_seqlens req={i} npages={npages} {tag}",
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_writes_bounded_by_cache_seqlens(self):
|
||||||
|
for max_ctx in (4096, 131072):
|
||||||
|
for page_size in (1, 64, 256):
|
||||||
|
for bs in (1, 8):
|
||||||
|
self._run_self_guard_case(max_ctx, page_size, bs=bs)
|
||||||
|
self._run_self_guard_case(max_ctx, page_size, bs=bs, swa=True)
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|||||||
Reference in New Issue
Block a user