[GLM-5.3-Flash] Reduce KPool planning synchronization and overlap indexer preparation (#39695)
Co-authored-by: Xinyuan Tong <xinyuantong.cs@gmail.com>
This commit is contained in:
co-authored by
Xinyuan Tong
parent
c8eb54c41d
commit
9f21fbc34b
@@ -705,6 +705,7 @@ class TboForwardBatchPreparer:
|
|||||||
|
|
||||||
for key in [
|
for key in [
|
||||||
"req_pool_indices",
|
"req_pool_indices",
|
||||||
|
"req_pool_indices_cpu",
|
||||||
"seq_lens",
|
"seq_lens",
|
||||||
"seq_lens_cpu",
|
"seq_lens_cpu",
|
||||||
"extend_seq_lens",
|
"extend_seq_lens",
|
||||||
|
|||||||
@@ -14,6 +14,7 @@ from sglang.srt.layers.attention.dsa.dsa_indexer import (
|
|||||||
rotate_activation,
|
rotate_activation,
|
||||||
)
|
)
|
||||||
from sglang.srt.layers.attention.dsa.dsa_topk_backend import TopkTransformMethod
|
from sglang.srt.layers.attention.dsa.dsa_topk_backend import TopkTransformMethod
|
||||||
|
from sglang.srt.layers.attention.dsa.utils import dsa_use_prefill_cp
|
||||||
from sglang.srt.layers.layernorm import LayerNorm
|
from sglang.srt.layers.layernorm import LayerNorm
|
||||||
from sglang.srt.layers.utils import MultiPlatformOp
|
from sglang.srt.layers.utils import MultiPlatformOp
|
||||||
from sglang.srt.utils import add_prefix, ceil_align, is_cuda, is_hip, is_npu
|
from sglang.srt.utils import add_prefix, ceil_align, is_cuda, is_hip, is_npu
|
||||||
@@ -163,6 +164,36 @@ class IndexerKPool(MultiPlatformOp):
|
|||||||
weights = weights.unsqueeze(-1) * q_scale * self.softmax_scale
|
weights = weights.unsqueeze(-1) * q_scale * self.softmax_scale
|
||||||
return weights
|
return weights
|
||||||
|
|
||||||
|
@torch.compile(dynamic=True)
|
||||||
|
def _project_and_scale_head_gates(self, x: torch.Tensor) -> torch.Tensor:
|
||||||
|
weights, _ = self.weights_proj(x.float())
|
||||||
|
return weights * self.n_heads**-0.5
|
||||||
|
|
||||||
|
@torch.compile(dynamic=True)
|
||||||
|
def _apply_q_scale_and_softmax_scale(
|
||||||
|
self, weights: torch.Tensor, q_scale: torch.Tensor
|
||||||
|
) -> torch.Tensor:
|
||||||
|
return weights.unsqueeze(-1) * q_scale * self.softmax_scale
|
||||||
|
|
||||||
|
def _resolve_head_gate_weights(self, x, q_scale, head_weights):
|
||||||
|
if head_weights is not None:
|
||||||
|
return self._apply_q_scale_and_softmax_scale(head_weights, q_scale)
|
||||||
|
return self._get_logits_head_gate(x, q_scale)
|
||||||
|
|
||||||
|
def _can_overlap_prefill(
|
||||||
|
self,
|
||||||
|
forward_batch: ForwardBatch,
|
||||||
|
return_indices: bool,
|
||||||
|
) -> bool:
|
||||||
|
return (
|
||||||
|
self.alt_stream is not None
|
||||||
|
and return_indices
|
||||||
|
and forward_batch.forward_mode.is_extend_without_speculative()
|
||||||
|
and not get_is_capture_mode()
|
||||||
|
and not is_in_breakable_cuda_graph()
|
||||||
|
and not dsa_use_prefill_cp(forward_batch)
|
||||||
|
)
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _get_index_k_read_buffer(pool, layer_id: int) -> torch.Tensor:
|
def _get_index_k_read_buffer(pool, layer_id: int) -> torch.Tensor:
|
||||||
if hasattr(pool, "get_broadcastable_index_k_with_scale_buffer"):
|
if hasattr(pool, "get_broadcastable_index_k_with_scale_buffer"):
|
||||||
@@ -538,8 +569,11 @@ class IndexerKPool(MultiPlatformOp):
|
|||||||
enable_dual_stream: bool,
|
enable_dual_stream: bool,
|
||||||
forward_batch: ForwardBatch,
|
forward_batch: ForwardBatch,
|
||||||
precompute_compress_gate: bool = False,
|
precompute_compress_gate: bool = False,
|
||||||
|
precompute_head_gate: bool = False,
|
||||||
):
|
):
|
||||||
gate_score = None
|
gate_score = None
|
||||||
|
head_weights = None
|
||||||
|
apply_rope = not self.skip_rope and self.rope_head_dim > 0
|
||||||
if enable_dual_stream:
|
if enable_dual_stream:
|
||||||
current_stream = torch.cuda.current_stream()
|
current_stream = torch.cuda.current_stream()
|
||||||
self.alt_stream.wait_stream(current_stream)
|
self.alt_stream.wait_stream(current_stream)
|
||||||
@@ -557,6 +591,10 @@ class IndexerKPool(MultiPlatformOp):
|
|||||||
[self.rope_head_dim, self.head_dim - self.rope_head_dim],
|
[self.rope_head_dim, self.head_dim - self.rope_head_dim],
|
||||||
dim=-1,
|
dim=-1,
|
||||||
)
|
)
|
||||||
|
if precompute_head_gate:
|
||||||
|
head_weights = self._project_and_scale_head_gates(x)
|
||||||
|
if not apply_rope:
|
||||||
|
query = rotate_activation(query)
|
||||||
with torch.cuda.stream(self.alt_stream):
|
with torch.cuda.stream(self.alt_stream):
|
||||||
key, _ = self.wk(x)
|
key, _ = self.wk(x)
|
||||||
key = self.k_norm(key)
|
key = self.k_norm(key)
|
||||||
@@ -584,15 +622,16 @@ class IndexerKPool(MultiPlatformOp):
|
|||||||
key, [self.rope_head_dim, self.head_dim - self.rope_head_dim], dim=-1
|
key, [self.rope_head_dim, self.head_dim - self.rope_head_dim], dim=-1
|
||||||
)
|
)
|
||||||
|
|
||||||
if not self.skip_rope:
|
if apply_rope:
|
||||||
q_rope, k_rope = self.rotary_emb(positions, q_rope, k_rope)
|
q_rope, k_rope = self.rotary_emb(positions, q_rope, k_rope)
|
||||||
|
|
||||||
query[..., : self.rope_head_dim] = q_rope
|
query[..., : self.rope_head_dim] = q_rope
|
||||||
key[..., : self.rope_head_dim] = k_rope
|
key[..., : self.rope_head_dim] = k_rope
|
||||||
|
|
||||||
query = rotate_activation(query)
|
if apply_rope or not enable_dual_stream:
|
||||||
|
query = rotate_activation(query)
|
||||||
|
|
||||||
return query, key, gate_score
|
return query, key, gate_score, head_weights
|
||||||
|
|
||||||
def _get_k_bf16(
|
def _get_k_bf16(
|
||||||
self,
|
self,
|
||||||
@@ -1293,7 +1332,7 @@ class IndexerKPool(MultiPlatformOp):
|
|||||||
assert plan is not None, "DSA kpool target_verify requires kpool_write_plan"
|
assert plan is not None, "DSA kpool target_verify requires kpool_write_plan"
|
||||||
num_draft_tokens = plan.num_draft_tokens
|
num_draft_tokens = plan.num_draft_tokens
|
||||||
|
|
||||||
query, key, gate_score_maybe = self._get_q_k_bf16(
|
query, key, gate_score_maybe, head_weights = self._get_q_k_bf16(
|
||||||
q_lora,
|
q_lora,
|
||||||
x,
|
x,
|
||||||
positions,
|
positions,
|
||||||
@@ -1302,6 +1341,7 @@ class IndexerKPool(MultiPlatformOp):
|
|||||||
precompute_compress_gate=(
|
precompute_compress_gate=(
|
||||||
enable_dual_stream and self.compress_gate_stream is not None
|
enable_dual_stream and self.compress_gate_stream is not None
|
||||||
),
|
),
|
||||||
|
precompute_head_gate=enable_dual_stream and return_indices,
|
||||||
)
|
)
|
||||||
|
|
||||||
pool = get_token_to_kv_pool()
|
pool = get_token_to_kv_pool()
|
||||||
@@ -1340,7 +1380,7 @@ class IndexerKPool(MultiPlatformOp):
|
|||||||
self.alt_stream.wait_stream(self.compress_gate_stream)
|
self.alt_stream.wait_stream(self.compress_gate_stream)
|
||||||
if return_indices:
|
if return_indices:
|
||||||
q_fp8, q_scale = act_quant(query, self.block_size, self.scale_fmt)
|
q_fp8, q_scale = act_quant(query, self.block_size, self.scale_fmt)
|
||||||
weights = self._get_logits_head_gate(x, q_scale)
|
weights = self._resolve_head_gate_weights(x, q_scale, head_weights)
|
||||||
with torch.cuda.stream(self.alt_stream):
|
with torch.cuda.stream(self.alt_stream):
|
||||||
_compress_write()
|
_compress_write()
|
||||||
current_stream.wait_stream(self.alt_stream)
|
current_stream.wait_stream(self.alt_stream)
|
||||||
@@ -1348,7 +1388,7 @@ class IndexerKPool(MultiPlatformOp):
|
|||||||
_compress_write()
|
_compress_write()
|
||||||
if return_indices:
|
if return_indices:
|
||||||
q_fp8, q_scale = act_quant(query, self.block_size, self.scale_fmt)
|
q_fp8, q_scale = act_quant(query, self.block_size, self.scale_fmt)
|
||||||
weights = self._get_logits_head_gate(x, q_scale)
|
weights = self._resolve_head_gate_weights(x, q_scale, head_weights)
|
||||||
|
|
||||||
if not return_indices:
|
if not return_indices:
|
||||||
return None
|
return None
|
||||||
@@ -1476,44 +1516,24 @@ class IndexerKPool(MultiPlatformOp):
|
|||||||
and forward_batch.forward_mode.is_decode_or_idle()
|
and forward_batch.forward_mode.is_decode_or_idle()
|
||||||
and self.compress_gate_stream is not None
|
and self.compress_gate_stream is not None
|
||||||
)
|
)
|
||||||
query, key, gate_score = self._get_q_k_bf16(
|
query, key, gate_score, head_weights = self._get_q_k_bf16(
|
||||||
q_lora,
|
q_lora,
|
||||||
x,
|
x,
|
||||||
positions,
|
positions,
|
||||||
enable_dual_stream,
|
enable_dual_stream,
|
||||||
forward_batch=forward_batch,
|
forward_batch=forward_batch,
|
||||||
precompute_compress_gate=precompute_compress_gate,
|
precompute_compress_gate=precompute_compress_gate,
|
||||||
|
precompute_head_gate=enable_dual_stream and return_indices,
|
||||||
)
|
)
|
||||||
|
|
||||||
weights = None
|
has_kpool_extend_plan = metadata.attn_metadata.kpool_extend_plan is not None
|
||||||
kpool_extend_cache = None
|
is_prefill = forward_batch.forward_mode.is_extend_without_speculative()
|
||||||
if enable_dual_stream and forward_batch.forward_mode.is_decode_or_idle():
|
defer_kpool_cache_write = (
|
||||||
current_stream = torch.cuda.current_stream()
|
is_prefill and return_indices and not has_kpool_extend_plan
|
||||||
self.alt_stream.wait_stream(current_stream)
|
)
|
||||||
if gate_score is not None:
|
|
||||||
self.alt_stream.wait_stream(self.compress_gate_stream)
|
def compress_write():
|
||||||
with torch.cuda.stream(self.alt_stream):
|
return self._compress_write(
|
||||||
self._compress_write(
|
|
||||||
x=x,
|
|
||||||
key=key,
|
|
||||||
positions=positions,
|
|
||||||
forward_batch=forward_batch,
|
|
||||||
layer_id=layer_id,
|
|
||||||
metadata=metadata,
|
|
||||||
gate_score=gate_score,
|
|
||||||
)
|
|
||||||
q_fp8, q_scale = act_quant(query, self.block_size, self.scale_fmt)
|
|
||||||
weights = self._get_logits_head_gate(x, q_scale)
|
|
||||||
current_stream.wait_stream(self.alt_stream)
|
|
||||||
else:
|
|
||||||
q_fp8, q_scale = act_quant(query, self.block_size, self.scale_fmt)
|
|
||||||
has_kpool_extend_plan = metadata.attn_metadata.kpool_extend_plan is not None
|
|
||||||
defer_kpool_cache_write = (
|
|
||||||
forward_batch.forward_mode.is_extend_without_speculative()
|
|
||||||
and return_indices
|
|
||||||
and not has_kpool_extend_plan
|
|
||||||
)
|
|
||||||
kpool_extend_cache = self._compress_write(
|
|
||||||
x=x,
|
x=x,
|
||||||
key=key,
|
key=key,
|
||||||
positions=positions,
|
positions=positions,
|
||||||
@@ -1521,20 +1541,33 @@ class IndexerKPool(MultiPlatformOp):
|
|||||||
layer_id=layer_id,
|
layer_id=layer_id,
|
||||||
metadata=metadata,
|
metadata=metadata,
|
||||||
gate_score=gate_score,
|
gate_score=gate_score,
|
||||||
return_compressed=(
|
return_compressed=is_prefill and return_indices,
|
||||||
forward_batch.forward_mode.is_extend_without_speculative()
|
|
||||||
and return_indices
|
|
||||||
),
|
|
||||||
write_cache=not defer_kpool_cache_write,
|
write_cache=not defer_kpool_cache_write,
|
||||||
)
|
)
|
||||||
if (
|
|
||||||
forward_batch.forward_mode.is_extend_without_speculative()
|
|
||||||
and not return_indices
|
|
||||||
):
|
|
||||||
return None
|
|
||||||
|
|
||||||
if weights is None:
|
overlap_decode = (
|
||||||
weights = self._get_logits_head_gate(x, q_scale)
|
enable_dual_stream and forward_batch.forward_mode.is_decode_or_idle()
|
||||||
|
)
|
||||||
|
if overlap_decode or self._can_overlap_prefill(forward_batch, return_indices):
|
||||||
|
current_stream = torch.cuda.current_stream()
|
||||||
|
self.alt_stream.wait_stream(current_stream)
|
||||||
|
if gate_score is not None:
|
||||||
|
self.alt_stream.wait_stream(self.compress_gate_stream)
|
||||||
|
with torch.cuda.stream(self.alt_stream):
|
||||||
|
kpool_extend_cache = compress_write()
|
||||||
|
if return_indices:
|
||||||
|
q_fp8, q_scale = act_quant(query, self.block_size, self.scale_fmt)
|
||||||
|
weights = self._resolve_head_gate_weights(x, q_scale, head_weights)
|
||||||
|
current_stream.wait_stream(self.alt_stream)
|
||||||
|
else:
|
||||||
|
if return_indices:
|
||||||
|
q_fp8, q_scale = act_quant(query, self.block_size, self.scale_fmt)
|
||||||
|
kpool_extend_cache = compress_write()
|
||||||
|
if return_indices:
|
||||||
|
weights = self._resolve_head_gate_weights(x, q_scale, head_weights)
|
||||||
|
|
||||||
|
if not return_indices:
|
||||||
|
return None
|
||||||
|
|
||||||
if is_cuda():
|
if is_cuda():
|
||||||
if (
|
if (
|
||||||
|
|||||||
@@ -245,7 +245,10 @@ def _kpool_cpu_plan(
|
|||||||
if isinstance(extend_seq_lens_cpu, torch.Tensor):
|
if isinstance(extend_seq_lens_cpu, torch.Tensor):
|
||||||
extend_seq_lens_cpu = extend_seq_lens_cpu.tolist()
|
extend_seq_lens_cpu = extend_seq_lens_cpu.tolist()
|
||||||
seq_lens_cpu = forward_batch.seq_lens_cpu.tolist()
|
seq_lens_cpu = forward_batch.seq_lens_cpu.tolist()
|
||||||
req_pool_indices_cpu = forward_batch.req_pool_indices.tolist()
|
req_pool_indices_cpu = getattr(forward_batch, "req_pool_indices_cpu", None)
|
||||||
|
if req_pool_indices_cpu is None:
|
||||||
|
req_pool_indices_cpu = forward_batch.req_pool_indices
|
||||||
|
req_pool_indices_cpu = req_pool_indices_cpu.tolist()
|
||||||
|
|
||||||
_append_compress_rows(
|
_append_compress_rows(
|
||||||
plan,
|
plan,
|
||||||
@@ -411,7 +414,9 @@ def _kpool_plan_to_gpu(
|
|||||||
if need_paged:
|
if need_paged:
|
||||||
req_to_token = get_req_to_token_pool().req_to_token
|
req_to_token = get_req_to_token_pool().req_to_token
|
||||||
ragged_paged_page_table_row_index = torch.repeat_interleave(
|
ragged_paged_page_table_row_index = torch.repeat_interleave(
|
||||||
local_req_pool_indices.to(torch.int32), ragged_q_len_t
|
local_req_pool_indices.to(torch.int32),
|
||||||
|
ragged_q_len_t,
|
||||||
|
output_size=sum(cpu.ragged_q_len),
|
||||||
)
|
)
|
||||||
ragged_paged_page_table = req_to_token
|
ragged_paged_page_table = req_to_token
|
||||||
|
|
||||||
|
|||||||
@@ -2348,8 +2348,6 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
|||||||
# A full prefill has one result but can span several run_batch calls.
|
# A full prefill has one result but can span several run_batch calls.
|
||||||
split_prefill_start: Optional[Tuple[int, float]] = None
|
split_prefill_start: Optional[Tuple[int, float]] = None
|
||||||
|
|
||||||
# CPU mirror of req_pool_indices; schedule-path only (used in overlap_utils,
|
|
||||||
# not read by ForwardBatch), stale in spec draft window
|
|
||||||
req_pool_indices_cpu: torch.Tensor = None # shape: [b], int64
|
req_pool_indices_cpu: torch.Tensor = None # shape: [b], int64
|
||||||
|
|
||||||
# Forward-pass metrics
|
# Forward-pass metrics
|
||||||
@@ -3770,6 +3768,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
|||||||
prefix_lens=self.prefix_lens,
|
prefix_lens=self.prefix_lens,
|
||||||
req_to_token_pool=self.req_to_token_pool,
|
req_to_token_pool=self.req_to_token_pool,
|
||||||
req_pool_indices=self.req_pool_indices,
|
req_pool_indices=self.req_pool_indices,
|
||||||
|
req_pool_indices_cpu=self.req_pool_indices_cpu,
|
||||||
model_config=self.model_config,
|
model_config=self.model_config,
|
||||||
forward_mode=self.forward_mode,
|
forward_mode=self.forward_mode,
|
||||||
out_cache_loc=self.out_cache_loc,
|
out_cache_loc=self.out_cache_loc,
|
||||||
|
|||||||
@@ -547,6 +547,8 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
|
|||||||
# === Borrowed from ScheduleBatch: host metadata (CPU lists / mirrors) ===
|
# === Borrowed from ScheduleBatch: host metadata (CPU lists / mirrors) ===
|
||||||
# Optional seq_lens on cpu (CPU mirror of seq_lens)
|
# Optional seq_lens on cpu (CPU mirror of seq_lens)
|
||||||
seq_lens_cpu: Optional[torch.Tensor] = None
|
seq_lens_cpu: Optional[torch.Tensor] = None
|
||||||
|
# Fresh only for non-speculative extend; speculative modes use device slots.
|
||||||
|
req_pool_indices_cpu: Optional[torch.Tensor] = None
|
||||||
|
|
||||||
# For logprob
|
# For logprob
|
||||||
top_logprobs_nums: Optional[List[int]] = None
|
top_logprobs_nums: Optional[List[int]] = None
|
||||||
@@ -908,6 +910,11 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
|
|||||||
seq_lens_sum=batch.seq_lens_sum,
|
seq_lens_sum=batch.seq_lens_sum,
|
||||||
# Inputs aliased by reference from ScheduleBatch
|
# Inputs aliased by reference from ScheduleBatch
|
||||||
seq_lens_cpu=seq_lens_cpu,
|
seq_lens_cpu=seq_lens_cpu,
|
||||||
|
req_pool_indices_cpu=(
|
||||||
|
getattr(batch, "req_pool_indices_cpu", None)
|
||||||
|
if batch.forward_mode.is_extend_without_speculative()
|
||||||
|
else None
|
||||||
|
),
|
||||||
orig_seq_lens=batch.orig_seq_lens,
|
orig_seq_lens=batch.orig_seq_lens,
|
||||||
out_cache_loc_dsv4=batch.out_cache_loc_dsv4,
|
out_cache_loc_dsv4=batch.out_cache_loc_dsv4,
|
||||||
engram_history=batch.engram_history,
|
engram_history=batch.engram_history,
|
||||||
@@ -1671,6 +1678,10 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
|
|||||||
# Keep token-aligned inputs consistent after padding.
|
# Keep token-aligned inputs consistent after padding.
|
||||||
self.input_embeds = self._pad_tensor_to_size(self.input_embeds, num_tokens)
|
self.input_embeds = self._pad_tensor_to_size(self.input_embeds, num_tokens)
|
||||||
self.req_pool_indices = self._pad_tensor_to_size(self.req_pool_indices, bs)
|
self.req_pool_indices = self._pad_tensor_to_size(self.req_pool_indices, bs)
|
||||||
|
if self.req_pool_indices_cpu is not None:
|
||||||
|
self.req_pool_indices_cpu = self._pad_tensor_to_size(
|
||||||
|
self.req_pool_indices_cpu, bs
|
||||||
|
)
|
||||||
if self.lora_ids is not None:
|
if self.lora_ids is not None:
|
||||||
self.lora_ids.extend((bs - len(self.lora_ids)) * [None])
|
self.lora_ids.extend((bs - len(self.lora_ids)) * [None])
|
||||||
|
|
||||||
@@ -1836,6 +1847,8 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
|
|||||||
self.positions = self.positions[: self._original_num_tokens]
|
self.positions = self.positions[: self._original_num_tokens]
|
||||||
self.seq_lens = self.seq_lens[:bs]
|
self.seq_lens = self.seq_lens[:bs]
|
||||||
self.req_pool_indices = self.req_pool_indices[:bs]
|
self.req_pool_indices = self.req_pool_indices[:bs]
|
||||||
|
if self.req_pool_indices_cpu is not None:
|
||||||
|
self.req_pool_indices_cpu = self.req_pool_indices_cpu[:bs]
|
||||||
if self.seq_lens_cpu is not None:
|
if self.seq_lens_cpu is not None:
|
||||||
self.seq_lens_cpu = self.seq_lens_cpu[:bs]
|
self.seq_lens_cpu = self.seq_lens_cpu[:bs]
|
||||||
|
|
||||||
@@ -1845,6 +1858,8 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
|
|||||||
self.positions = self.positions[:num_tokens]
|
self.positions = self.positions[:num_tokens]
|
||||||
self.seq_lens = self.seq_lens[:bs]
|
self.seq_lens = self.seq_lens[:bs]
|
||||||
self.req_pool_indices = self.req_pool_indices[:bs]
|
self.req_pool_indices = self.req_pool_indices[:bs]
|
||||||
|
if self.req_pool_indices_cpu is not None:
|
||||||
|
self.req_pool_indices_cpu = self.req_pool_indices_cpu[:bs]
|
||||||
if self.seq_lens_cpu is not None:
|
if self.seq_lens_cpu is not None:
|
||||||
self.seq_lens_cpu = self.seq_lens_cpu[:bs]
|
self.seq_lens_cpu = self.seq_lens_cpu[:bs]
|
||||||
if logits_output.next_token_logits is not None:
|
if logits_output.next_token_logits is not None:
|
||||||
|
|||||||
@@ -1176,6 +1176,10 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
|
|||||||
padded_view.seq_lens = s["seq_lens"][:r]
|
padded_view.seq_lens = s["seq_lens"][:r]
|
||||||
padded_view.seq_lens_cpu = self._full_cg_seq_lens_cpu
|
padded_view.seq_lens_cpu = self._full_cg_seq_lens_cpu
|
||||||
padded_view.req_pool_indices = s["req_pool_indices"][:r]
|
padded_view.req_pool_indices = s["req_pool_indices"][:r]
|
||||||
|
if getattr(forward_batch, "req_pool_indices_cpu", None) is not None:
|
||||||
|
padded_view.req_pool_indices_cpu = padded_view._pad_tensor_to_size(
|
||||||
|
forward_batch.req_pool_indices_cpu, r
|
||||||
|
)
|
||||||
padded_view.extend_seq_lens = s["extend_seq_lens"][:r]
|
padded_view.extend_seq_lens = s["extend_seq_lens"][:r]
|
||||||
padded_view.extend_prefix_lens = s["extend_prefix_lens"][:r]
|
padded_view.extend_prefix_lens = s["extend_prefix_lens"][:r]
|
||||||
padded_view.max_seq_len_override = static_forward_batch.max_seq_len_override
|
padded_view.max_seq_len_override = static_forward_batch.max_seq_len_override
|
||||||
|
|||||||
@@ -0,0 +1,147 @@
|
|||||||
|
"""CPU coverage for KPool request-slot selection and paged query rows."""
|
||||||
|
|
||||||
|
import unittest
|
||||||
|
from types import SimpleNamespace
|
||||||
|
from unittest.mock import patch
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from sglang.srt.environ import envs
|
||||||
|
from sglang.srt.layers.attention.dsa import kpool_plan
|
||||||
|
from sglang.srt.layers.attention.dsa.dsa_topk_backend import TopkTransformMethod
|
||||||
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
|
|
||||||
|
register_cpu_ci(est_time=5, suite="base-a-test-cpu")
|
||||||
|
|
||||||
|
|
||||||
|
class _NoDeviceReadTensor(torch.Tensor):
|
||||||
|
def tolist(self):
|
||||||
|
raise AssertionError("Reading device request slots would synchronize")
|
||||||
|
|
||||||
|
|
||||||
|
def _batch(extend_lens=(3, 5), seq_lens=(5, 11), slots=(9, 2)):
|
||||||
|
indices = torch.tensor(slots, dtype=torch.int64)
|
||||||
|
return SimpleNamespace(
|
||||||
|
batch_size=len(slots),
|
||||||
|
extend_seq_lens_cpu=list(extend_lens),
|
||||||
|
seq_lens_cpu=torch.tensor(seq_lens, dtype=torch.int64),
|
||||||
|
seq_lens=torch.tensor(seq_lens, dtype=torch.int64),
|
||||||
|
req_pool_indices=indices.as_subclass(_NoDeviceReadTensor),
|
||||||
|
req_pool_indices_cpu=indices,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _expected_plan():
|
||||||
|
return kpool_plan._KPoolCpuPlan(
|
||||||
|
pool_batch_idx=[0, 1],
|
||||||
|
pool_req=[9, 2],
|
||||||
|
pool_pool_id=[0, 1],
|
||||||
|
pool_n_from_tail=[2, 2],
|
||||||
|
pool_chunk_src=[0, 3],
|
||||||
|
pool_tail_logical_base=[0, 4],
|
||||||
|
tail_req=[9, 2],
|
||||||
|
tail_dst_logical_start=[4, 8],
|
||||||
|
tail_chunk_src=[2, 5],
|
||||||
|
tail_n_write=[1, 3],
|
||||||
|
ragged_q_len=[3, 5],
|
||||||
|
ragged_pool_pages=[1, 1],
|
||||||
|
cu_pages_excl=[0, 1],
|
||||||
|
cu_q_len_excl=[0, 3],
|
||||||
|
total_pool_pages=2,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class TestKPoolPlannerCpuMirror(unittest.TestCase):
|
||||||
|
def test_mirror_avoids_device_read_and_preserves_compression_rows(self):
|
||||||
|
for tensor_lengths in (False, True):
|
||||||
|
with self.subTest(tensor_lengths=tensor_lengths):
|
||||||
|
batch = _batch()
|
||||||
|
if tensor_lengths:
|
||||||
|
batch.extend_seq_lens_cpu = torch.tensor(batch.extend_seq_lens_cpu)
|
||||||
|
plan = kpool_plan._kpool_cpu_plan(batch, 4, 64)
|
||||||
|
self.assertEqual(plan, _expected_plan())
|
||||||
|
|
||||||
|
def test_absent_or_none_mirror_uses_device_slots(self):
|
||||||
|
for absent in (False, True):
|
||||||
|
with self.subTest(absent=absent):
|
||||||
|
batch = _batch()
|
||||||
|
batch.req_pool_indices = batch.req_pool_indices_cpu
|
||||||
|
if absent:
|
||||||
|
del batch.req_pool_indices_cpu
|
||||||
|
else:
|
||||||
|
batch.req_pool_indices_cpu = None
|
||||||
|
plan = kpool_plan._kpool_cpu_plan(batch, 4, 64)
|
||||||
|
self.assertEqual(plan, _expected_plan())
|
||||||
|
|
||||||
|
def test_empty_mirror_needs_no_device_read(self):
|
||||||
|
plan = kpool_plan._kpool_cpu_plan(_batch((), (), ()), 4, 64)
|
||||||
|
self.assertEqual(plan, kpool_plan._KPoolCpuPlan())
|
||||||
|
|
||||||
|
def test_paged_query_rows_use_local_lengths_and_explicit_output_size(self):
|
||||||
|
batch = _batch()
|
||||||
|
cpu_plan = kpool_plan._kpool_cpu_plan(
|
||||||
|
batch,
|
||||||
|
4,
|
||||||
|
64,
|
||||||
|
local_extend_seq_lens_cpu=[1, 2],
|
||||||
|
local_seq_lens_cpu=[3, 8],
|
||||||
|
)
|
||||||
|
expected = _expected_plan()
|
||||||
|
expected.ragged_q_len = [1, 2]
|
||||||
|
expected.ragged_pool_pages = [0, 1]
|
||||||
|
expected.cu_pages_excl = [0, 0]
|
||||||
|
expected.cu_q_len_excl = [0, 1]
|
||||||
|
expected.total_pool_pages = 1
|
||||||
|
self.assertEqual(cpu_plan, expected)
|
||||||
|
|
||||||
|
original_tensor = torch.tensor
|
||||||
|
|
||||||
|
def cpu_tensor(*args, **kwargs):
|
||||||
|
kwargs.pop("pin_memory", None)
|
||||||
|
return original_tensor(*args, **kwargs)
|
||||||
|
|
||||||
|
page_table = torch.zeros(2, 4, dtype=torch.int32)
|
||||||
|
req_to_token = torch.zeros(10, 256, dtype=torch.int32)
|
||||||
|
seq_lens = torch.tensor([3, 7, 8], dtype=torch.int32)
|
||||||
|
local_slots = torch.tensor([9, 2], dtype=torch.int64)
|
||||||
|
with (
|
||||||
|
envs.SGLANG_DSA_FUSE_TOPK.override(True),
|
||||||
|
patch.object(torch, "tensor", side_effect=cpu_tensor),
|
||||||
|
patch.object(
|
||||||
|
torch, "repeat_interleave", wraps=torch.repeat_interleave
|
||||||
|
) as repeat,
|
||||||
|
patch.object(
|
||||||
|
kpool_plan,
|
||||||
|
"kpool_build_ragged_layout",
|
||||||
|
return_value=(torch.empty(0), torch.empty(0), torch.empty(0)),
|
||||||
|
),
|
||||||
|
patch.object(kpool_plan, "dsa_use_prefill_cp", return_value=False),
|
||||||
|
patch.object(kpool_plan, "_RAGGED_SCRATCH_K_U8", None),
|
||||||
|
patch.object(kpool_plan, "_RAGGED_SCRATCH_K_SCALE", None),
|
||||||
|
patch.object(
|
||||||
|
kpool_plan,
|
||||||
|
"get_req_to_token_pool",
|
||||||
|
return_value=SimpleNamespace(req_to_token=req_to_token),
|
||||||
|
),
|
||||||
|
):
|
||||||
|
plan = kpool_plan._kpool_plan_to_gpu(
|
||||||
|
cpu_plan,
|
||||||
|
batch,
|
||||||
|
page_table,
|
||||||
|
page_table,
|
||||||
|
seq_lens,
|
||||||
|
local_slots,
|
||||||
|
4,
|
||||||
|
64,
|
||||||
|
TopkTransformMethod.PAGED,
|
||||||
|
)
|
||||||
|
|
||||||
|
repeat.assert_called_once()
|
||||||
|
self.assertEqual(repeat.call_args.kwargs["output_size"], 3)
|
||||||
|
self.assertEqual(plan.ragged_paged_page_table_row_index.tolist(), [9, 2, 2])
|
||||||
|
self.assertEqual(plan.ragged_paged_page_table_row_index.dtype, torch.int32)
|
||||||
|
self.assertIs(plan.ragged_paged_page_table, req_to_token)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
@@ -0,0 +1,316 @@
|
|||||||
|
"""CPU scheduling checks: stream dependencies, cache contracts and gate math."""
|
||||||
|
|
||||||
|
import unittest
|
||||||
|
from contextlib import contextmanager, nullcontext
|
||||||
|
from types import MethodType, SimpleNamespace
|
||||||
|
from unittest.mock import Mock, patch
|
||||||
|
|
||||||
|
import torch
|
||||||
|
import torch.nn.functional as F
|
||||||
|
|
||||||
|
from sglang.kernels.ops.attention.dsa import triton_kernel
|
||||||
|
from sglang.srt.layers.attention.dsa import dsa_indexer_kpool as indexer_module
|
||||||
|
from sglang.srt.layers.attention.dsa.dsa_indexer_kpool import IndexerKPool
|
||||||
|
from sglang.srt.model_executor.forward_batch_info import ForwardMode
|
||||||
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
|
|
||||||
|
register_cpu_ci(est_time=5, suite="base-a-test-cpu")
|
||||||
|
|
||||||
|
|
||||||
|
def _eager(method):
|
||||||
|
return getattr(method, "_torchdynamo_orig_callable", method)
|
||||||
|
|
||||||
|
|
||||||
|
class _Stream:
|
||||||
|
def __init__(self, name, trace):
|
||||||
|
self.name = name
|
||||||
|
self.trace = trace
|
||||||
|
|
||||||
|
def wait_stream(self, other):
|
||||||
|
self.trace.append(("wait", self.name, other.name))
|
||||||
|
|
||||||
|
|
||||||
|
class _Streams:
|
||||||
|
def __init__(self):
|
||||||
|
self.trace = []
|
||||||
|
self.current = _Stream("main", self.trace)
|
||||||
|
self.alt = _Stream("alt", self.trace)
|
||||||
|
self.gate = _Stream("gate", self.trace)
|
||||||
|
|
||||||
|
@contextmanager
|
||||||
|
def use(self, stream):
|
||||||
|
previous = self.current
|
||||||
|
self.current = stream
|
||||||
|
try:
|
||||||
|
yield
|
||||||
|
finally:
|
||||||
|
self.current = previous
|
||||||
|
|
||||||
|
def record(self, operation):
|
||||||
|
self.trace.append((operation, self.current.name))
|
||||||
|
|
||||||
|
|
||||||
|
class TestKPoolStreamScheduling(unittest.TestCase):
|
||||||
|
def test_prefill_overlap_excludes_cp_and_graph_capture(self):
|
||||||
|
for mode in (ForwardMode.EXTEND, ForwardMode.DECODE, ForwardMode.TARGET_VERIFY):
|
||||||
|
for has_stream in (False, True):
|
||||||
|
for capture, breakable, cp in (
|
||||||
|
(False, False, False),
|
||||||
|
(True, False, False),
|
||||||
|
(False, True, False),
|
||||||
|
(False, False, True),
|
||||||
|
):
|
||||||
|
with (
|
||||||
|
self.subTest(
|
||||||
|
mode=mode,
|
||||||
|
stream=has_stream,
|
||||||
|
capture=capture,
|
||||||
|
breakable=breakable,
|
||||||
|
cp=cp,
|
||||||
|
),
|
||||||
|
patch.object(
|
||||||
|
indexer_module, "get_is_capture_mode", return_value=capture
|
||||||
|
),
|
||||||
|
patch.object(
|
||||||
|
indexer_module,
|
||||||
|
"is_in_breakable_cuda_graph",
|
||||||
|
return_value=breakable,
|
||||||
|
),
|
||||||
|
patch.object(
|
||||||
|
indexer_module, "dsa_use_prefill_cp", return_value=cp
|
||||||
|
),
|
||||||
|
):
|
||||||
|
indexer = SimpleNamespace(
|
||||||
|
alt_stream=object() if has_stream else None
|
||||||
|
)
|
||||||
|
actual = IndexerKPool._can_overlap_prefill(
|
||||||
|
indexer,
|
||||||
|
SimpleNamespace(forward_mode=mode),
|
||||||
|
return_indices=True,
|
||||||
|
)
|
||||||
|
self.assertEqual(
|
||||||
|
actual,
|
||||||
|
mode == ForwardMode.EXTEND
|
||||||
|
and has_stream
|
||||||
|
and not (capture or breakable or cp),
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_precomputed_head_gate_matches_original_math(self):
|
||||||
|
torch.manual_seed(19)
|
||||||
|
for dtype in (torch.bfloat16, torch.float16, torch.float32):
|
||||||
|
x = torch.randn(11, 32).to(dtype)
|
||||||
|
matrix = torch.randn(8, 32)
|
||||||
|
q_scale = torch.rand(11, 8, 1)
|
||||||
|
indexer = SimpleNamespace(
|
||||||
|
weights_proj=lambda value: (F.linear(value, matrix), None),
|
||||||
|
n_heads=8,
|
||||||
|
softmax_scale=128**-0.5,
|
||||||
|
)
|
||||||
|
expected = _eager(IndexerKPool._get_logits_head_gate)(indexer, x, q_scale)
|
||||||
|
projected = _eager(IndexerKPool._project_and_scale_head_gates)(indexer, x)
|
||||||
|
actual = _eager(IndexerKPool._apply_q_scale_and_softmax_scale)(
|
||||||
|
indexer, projected, q_scale
|
||||||
|
)
|
||||||
|
torch.testing.assert_close(actual, expected, rtol=0, atol=0)
|
||||||
|
|
||||||
|
def test_projection_reordering_preserves_rope_and_third_stream(self):
|
||||||
|
for skip_rope in (False, True):
|
||||||
|
streams = _Streams()
|
||||||
|
x = torch.arange(32, dtype=torch.float32).reshape(4, 8)
|
||||||
|
indexer = SimpleNamespace(
|
||||||
|
alt_stream=streams.alt,
|
||||||
|
compress_gate_stream=streams.gate,
|
||||||
|
half_device_sm_count=8,
|
||||||
|
head_dim=4,
|
||||||
|
rope_head_dim=2,
|
||||||
|
skip_rope=skip_rope,
|
||||||
|
index_kpool_compress_gate=torch.ones(4, 8),
|
||||||
|
)
|
||||||
|
|
||||||
|
def project_q(value):
|
||||||
|
streams.record("project_q")
|
||||||
|
return value.clone(), None
|
||||||
|
|
||||||
|
def project_k(value):
|
||||||
|
streams.record("project_k")
|
||||||
|
return value[:, :4].clone(), None
|
||||||
|
|
||||||
|
def head_gate(value):
|
||||||
|
streams.record("head_gate")
|
||||||
|
return value[:, :2].clone()
|
||||||
|
|
||||||
|
def rotate(value):
|
||||||
|
streams.record("rotate")
|
||||||
|
return value.flip(-1)
|
||||||
|
|
||||||
|
def rope(positions, q, k):
|
||||||
|
streams.record("rope")
|
||||||
|
return q + 1, k + 2
|
||||||
|
|
||||||
|
indexer.wq_b = project_q
|
||||||
|
indexer.wk = project_k
|
||||||
|
indexer.k_norm = lambda value: value
|
||||||
|
indexer.rotary_emb = rope
|
||||||
|
indexer._project_and_scale_head_gates = head_gate
|
||||||
|
with (
|
||||||
|
patch.object(
|
||||||
|
torch.cuda, "current_stream", side_effect=lambda: streams.current
|
||||||
|
),
|
||||||
|
patch.object(torch.cuda, "stream", side_effect=streams.use),
|
||||||
|
patch.object(
|
||||||
|
indexer_module.deep_gemm_wrapper,
|
||||||
|
"configure_deep_gemm_num_sms",
|
||||||
|
return_value=nullcontext(),
|
||||||
|
),
|
||||||
|
patch.object(indexer_module, "rotate_activation", side_effect=rotate),
|
||||||
|
):
|
||||||
|
actual = IndexerKPool._get_q_k_bf16(
|
||||||
|
indexer,
|
||||||
|
x,
|
||||||
|
x,
|
||||||
|
torch.arange(4),
|
||||||
|
True,
|
||||||
|
None,
|
||||||
|
precompute_compress_gate=True,
|
||||||
|
precompute_head_gate=True,
|
||||||
|
)
|
||||||
|
trace = streams.trace[:]
|
||||||
|
expected = IndexerKPool._get_q_k_bf16(
|
||||||
|
indexer, x, x, torch.arange(4), False, None
|
||||||
|
)
|
||||||
|
torch.testing.assert_close(actual[0], expected[0])
|
||||||
|
torch.testing.assert_close(actual[1], expected[1])
|
||||||
|
torch.testing.assert_close(
|
||||||
|
actual[2], F.linear(x, indexer.index_kpool_compress_gate)
|
||||||
|
)
|
||||||
|
self.assertIsNotNone(actual[3])
|
||||||
|
self.assertIn(("wait", "gate", "main"), trace)
|
||||||
|
join = trace.index(("wait", "main", "alt"))
|
||||||
|
self.assertLess(trace.index(("head_gate", "main")), join)
|
||||||
|
if skip_rope:
|
||||||
|
self.assertLess(trace.index(("rotate", "main")), join)
|
||||||
|
self.assertNotIn(("rope", "main"), trace)
|
||||||
|
else:
|
||||||
|
self.assertGreater(trace.index(("rope", "main")), join)
|
||||||
|
self.assertGreater(trace.index(("rotate", "main")), join)
|
||||||
|
|
||||||
|
def test_prefill_waits_before_topk_and_preserves_deferred_cache(self):
|
||||||
|
for has_plan in (False, True):
|
||||||
|
for return_indices in (False, True):
|
||||||
|
for cp in (False, True):
|
||||||
|
for num_tokens in (0, 128, 8192):
|
||||||
|
with self.subTest(
|
||||||
|
plan=has_plan,
|
||||||
|
indices=return_indices,
|
||||||
|
cp=cp,
|
||||||
|
tokens=num_tokens,
|
||||||
|
):
|
||||||
|
self._run_prefill(has_plan, return_indices, cp, num_tokens)
|
||||||
|
|
||||||
|
def _run_prefill(self, has_plan, return_indices, cp, num_tokens):
|
||||||
|
streams = _Streams()
|
||||||
|
x = torch.ones(num_tokens, 8)
|
||||||
|
compressed = object()
|
||||||
|
calls = []
|
||||||
|
metadata = SimpleNamespace(
|
||||||
|
attn_metadata=SimpleNamespace(
|
||||||
|
kpool_extend_plan=object() if has_plan else None
|
||||||
|
)
|
||||||
|
)
|
||||||
|
batch = SimpleNamespace(
|
||||||
|
forward_mode=ForwardMode.EXTEND,
|
||||||
|
seq_lens_cpu=torch.tensor([8192 + num_tokens]),
|
||||||
|
)
|
||||||
|
prepare_qk = Mock(return_value=(x, x, None, None))
|
||||||
|
indexer = SimpleNamespace(
|
||||||
|
alt_stream=streams.alt,
|
||||||
|
compress_gate_stream=streams.gate,
|
||||||
|
index_topk=16,
|
||||||
|
index_kpool=4,
|
||||||
|
index_kpool_compress=True,
|
||||||
|
block_size=128,
|
||||||
|
scale_fmt=None,
|
||||||
|
_get_q_k_bf16=prepare_qk,
|
||||||
|
)
|
||||||
|
indexer._can_overlap_prefill = MethodType(
|
||||||
|
IndexerKPool._can_overlap_prefill, indexer
|
||||||
|
)
|
||||||
|
|
||||||
|
def compress(**kwargs):
|
||||||
|
streams.record("compress")
|
||||||
|
calls.append(kwargs)
|
||||||
|
return compressed
|
||||||
|
|
||||||
|
def quant(*args):
|
||||||
|
streams.record("quant")
|
||||||
|
return x, torch.ones(num_tokens, 1, 1)
|
||||||
|
|
||||||
|
def head(*args):
|
||||||
|
streams.record("head")
|
||||||
|
return x
|
||||||
|
|
||||||
|
def topk(*args, **kwargs):
|
||||||
|
streams.record("topk")
|
||||||
|
if not has_plan:
|
||||||
|
self.assertIs(kwargs["kpool_extend_cache"], compressed)
|
||||||
|
return x
|
||||||
|
|
||||||
|
indexer._compress_write = compress
|
||||||
|
indexer._resolve_head_gate_weights = head
|
||||||
|
indexer._get_topk_ragged = topk
|
||||||
|
indexer._get_topk_ragged_kpool_plan = topk
|
||||||
|
with (
|
||||||
|
patch.object(indexer_module, "is_cuda", return_value=True),
|
||||||
|
patch.object(indexer_module, "is_hip", return_value=False),
|
||||||
|
patch.object(indexer_module, "is_npu", return_value=False),
|
||||||
|
patch.object(indexer_module, "get_is_capture_mode", return_value=False),
|
||||||
|
patch.object(
|
||||||
|
indexer_module, "is_in_breakable_cuda_graph", return_value=False
|
||||||
|
),
|
||||||
|
patch.object(indexer_module, "dsa_use_prefill_cp", return_value=cp),
|
||||||
|
patch.object(
|
||||||
|
indexer_module,
|
||||||
|
"get_attn_backend",
|
||||||
|
return_value=SimpleNamespace(
|
||||||
|
get_indexer_metadata=lambda *args: metadata
|
||||||
|
),
|
||||||
|
),
|
||||||
|
patch.object(
|
||||||
|
torch.cuda, "current_stream", side_effect=lambda: streams.current
|
||||||
|
),
|
||||||
|
patch.object(torch.cuda, "stream", side_effect=streams.use),
|
||||||
|
patch.object(triton_kernel, "act_quant", side_effect=quant),
|
||||||
|
):
|
||||||
|
actual = IndexerKPool._forward_cuda_impl(
|
||||||
|
indexer, x, x, torch.arange(num_tokens), batch, 0, return_indices
|
||||||
|
)
|
||||||
|
self.assertEqual(calls[0]["return_compressed"], return_indices)
|
||||||
|
self.assertEqual(calls[0]["write_cache"], has_plan or not return_indices)
|
||||||
|
self.assertFalse(prepare_qk.call_args.args[3])
|
||||||
|
self.assertFalse(prepare_qk.call_args.kwargs["precompute_head_gate"])
|
||||||
|
overlap = not cp and return_indices
|
||||||
|
self.assertIn(("compress", "alt" if overlap else "main"), streams.trace)
|
||||||
|
if overlap:
|
||||||
|
self.assertLess(
|
||||||
|
streams.trace.index(("wait", "alt", "main")),
|
||||||
|
streams.trace.index(("compress", "alt")),
|
||||||
|
)
|
||||||
|
join = streams.trace.index(("wait", "main", "alt"))
|
||||||
|
self.assertGreater(join, streams.trace.index(("compress", "alt")))
|
||||||
|
self.assertGreater(join, streams.trace.index(("quant", "main")))
|
||||||
|
self.assertGreater(join, streams.trace.index(("head", "main")))
|
||||||
|
self.assertLess(join, streams.trace.index(("topk", "main")))
|
||||||
|
else:
|
||||||
|
self.assertFalse(any(event[0] == "wait" for event in streams.trace))
|
||||||
|
if return_indices:
|
||||||
|
self.assertIs(actual, x)
|
||||||
|
self.assertIn(("quant", "main"), streams.trace)
|
||||||
|
else:
|
||||||
|
self.assertIsNone(actual)
|
||||||
|
self.assertNotIn(("quant", "main"), streams.trace)
|
||||||
|
self.assertNotIn(("head", "main"), streams.trace)
|
||||||
|
self.assertNotIn(("topk", "main"), streams.trace)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
@@ -10,6 +10,7 @@ from sglang.srt.model_executor import forward_batch_info
|
|||||||
from sglang.srt.model_executor.cuda_graph_config import Backend
|
from sglang.srt.model_executor.cuda_graph_config import Backend
|
||||||
from sglang.srt.model_executor.forward_batch_info import (
|
from sglang.srt.model_executor.forward_batch_info import (
|
||||||
CaptureHiddenMode,
|
CaptureHiddenMode,
|
||||||
|
ForwardBatch,
|
||||||
ForwardMode,
|
ForwardMode,
|
||||||
prefill_graph_tolerates_sum_len,
|
prefill_graph_tolerates_sum_len,
|
||||||
)
|
)
|
||||||
@@ -93,6 +94,62 @@ class TestPrefillCudaGraphPadding(CustomTestCase):
|
|||||||
forward_batch, num_qo_tokens=16
|
forward_batch, num_qo_tokens=16
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def test_full_replay_pads_request_slot_cpu_mirror(self):
|
||||||
|
for batch_size, has_mirror in ((2, True), (4, True), (2, False)):
|
||||||
|
with self.subTest(batch_size=batch_size, has_mirror=has_mirror):
|
||||||
|
runner = self._make_runner()
|
||||||
|
runner._is_full_backend = True
|
||||||
|
runner._capture_req_slots = 4
|
||||||
|
runner._full_cg_seq_lens_cpu = torch.full((4,), -1)
|
||||||
|
slots = torch.tensor([7, 2, 9, 5])[:batch_size]
|
||||||
|
seq_lens = torch.arange(1, batch_size + 1)
|
||||||
|
static_slots = torch.zeros(4, dtype=slots.dtype)
|
||||||
|
static_slots[:batch_size].copy_(slots)
|
||||||
|
static_lens = torch.zeros(4, dtype=seq_lens.dtype)
|
||||||
|
static_lens[:batch_size].copy_(seq_lens)
|
||||||
|
runner._prefill_static_buffers = {
|
||||||
|
"req_pool_indices": static_slots,
|
||||||
|
"seq_lens": static_lens,
|
||||||
|
"extend_seq_lens": static_lens.clone(),
|
||||||
|
"extend_prefix_lens": torch.zeros(4, dtype=torch.int64),
|
||||||
|
}
|
||||||
|
attn_backend = mock.Mock()
|
||||||
|
runner.model_runner = SimpleNamespace(attn_backend=attn_backend)
|
||||||
|
batch = ForwardBatch(
|
||||||
|
forward_mode=ForwardMode.EXTEND,
|
||||||
|
batch_size=batch_size,
|
||||||
|
input_ids=torch.arange(batch_size),
|
||||||
|
req_pool_indices=slots,
|
||||||
|
req_pool_indices_cpu=slots if has_mirror else None,
|
||||||
|
seq_lens=seq_lens,
|
||||||
|
seq_lens_cpu=seq_lens,
|
||||||
|
out_cache_loc=torch.arange(batch_size),
|
||||||
|
seq_lens_sum=int(seq_lens.sum()),
|
||||||
|
)
|
||||||
|
|
||||||
|
runner._prepare_forward_metadata_for_replay(
|
||||||
|
batch, batch, shape_key=ShapeKey(size=16)
|
||||||
|
)
|
||||||
|
|
||||||
|
attn_backend.init_forward_metadata_out_graph.assert_called_once()
|
||||||
|
padded = attn_backend.init_forward_metadata_out_graph.call_args.args[0]
|
||||||
|
self.assertIsInstance(padded, ForwardBatch)
|
||||||
|
self.assertIsNot(padded, batch)
|
||||||
|
self.assertEqual(padded.batch_size, 4)
|
||||||
|
torch.testing.assert_close(padded.seq_lens_cpu, static_lens)
|
||||||
|
torch.testing.assert_close(padded.req_pool_indices, static_slots)
|
||||||
|
if has_mirror:
|
||||||
|
torch.testing.assert_close(
|
||||||
|
padded.req_pool_indices_cpu, static_slots
|
||||||
|
)
|
||||||
|
self.assertEqual(padded.req_pool_indices_cpu.device.type, "cpu")
|
||||||
|
self.assertIs(batch.req_pool_indices_cpu, slots)
|
||||||
|
else:
|
||||||
|
self.assertIsNone(padded.req_pool_indices_cpu)
|
||||||
|
self.assertEqual(batch.batch_size, batch_size)
|
||||||
|
torch.testing.assert_close(batch.seq_lens_cpu, seq_lens)
|
||||||
|
attn_backend.init_forward_metadata.assert_not_called()
|
||||||
|
|
||||||
def _megamoe_no_prefill_cp(self, graph_has_dp_gather=False):
|
def _megamoe_no_prefill_cp(self, graph_has_dp_gather=False):
|
||||||
return (
|
return (
|
||||||
mock.patch.object(
|
mock.patch.object(
|
||||||
|
|||||||
Reference in New Issue
Block a user