[AMD] Fix Dspark accept length and reduce host bubble on DSV4 (#39116)
This commit is contained in:
@@ -10,7 +10,6 @@ from typing import (
|
|||||||
List,
|
List,
|
||||||
Literal,
|
Literal,
|
||||||
Optional,
|
Optional,
|
||||||
Tuple,
|
|
||||||
TypeVar,
|
TypeVar,
|
||||||
Union,
|
Union,
|
||||||
)
|
)
|
||||||
@@ -141,10 +140,11 @@ class UnifiedKvMetadata:
|
|||||||
"verify_store_state_slot",
|
"verify_store_state_slot",
|
||||||
"c4_out_loc",
|
"c4_out_loc",
|
||||||
"c128_out_loc",
|
"c128_out_loc",
|
||||||
|
# Captured store_cache reads swa_loc by address, and the eager
|
||||||
|
# target-verify path builds it outside the graph.
|
||||||
|
"swa_loc",
|
||||||
],
|
],
|
||||||
# swa_loc is recomputed each forward (recorded inside cuda graphs),
|
assign_fields=[],
|
||||||
# so it is rebound rather than copied across replays.
|
|
||||||
assign_fields=["swa_loc"],
|
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
@@ -482,9 +482,8 @@ class DeepseekV4HipRadixBackend(
|
|||||||
self.speculative_num_steps = speculative_num_steps
|
self.speculative_num_steps = speculative_num_steps
|
||||||
self.speculative_num_draft_tokens: int = get_spec().speculative_num_draft_tokens
|
self.speculative_num_draft_tokens: int = get_spec().speculative_num_draft_tokens
|
||||||
self.is_draft_worker = getattr(model_runner, "is_draft_worker", False)
|
self.is_draft_worker = getattr(model_runner, "is_draft_worker", False)
|
||||||
self.is_dspark_draft = (
|
self.is_dspark = model_runner.spec_algorithm.is_dspark()
|
||||||
self.is_draft_worker and model_runner.spec_algorithm.is_dspark()
|
self.is_dspark_draft = self.is_draft_worker and self.is_dspark
|
||||||
)
|
|
||||||
self.target_verify_num_draft_tokens = self.speculative_num_draft_tokens
|
self.target_verify_num_draft_tokens = self.speculative_num_draft_tokens
|
||||||
if self.is_dspark_draft:
|
if self.is_dspark_draft:
|
||||||
assert self.speculative_num_draft_tokens is not None
|
assert self.speculative_num_draft_tokens is not None
|
||||||
@@ -493,6 +492,15 @@ class DeepseekV4HipRadixBackend(
|
|||||||
# CUDA-side convention gamma + 1, so use an explicit effective value
|
# CUDA-side convention gamma + 1, so use an explicit effective value
|
||||||
# instead of mutating speculative_num_draft_tokens in place.
|
# instead of mutating speculative_num_draft_tokens in place.
|
||||||
self.target_verify_num_draft_tokens = self.speculative_num_draft_tokens - 1
|
self.target_verify_num_draft_tokens = self.speculative_num_draft_tokens - 1
|
||||||
|
# Past MAX_FUSED_ROWS the fp4 schedule falls back to AITER's preamble,
|
||||||
|
# which frees the scratch its kernels read -- not capture-safe.
|
||||||
|
self._fp4_graph_row_limit: Optional[int] = None
|
||||||
|
if self.enable_deepseek_v4_fp4_indexer and self.speculative_num_steps == 0:
|
||||||
|
from sglang.kernels.ops.attention.dsv4.fp4_indexer_schedule_hip import (
|
||||||
|
MAX_FUSED_ROWS,
|
||||||
|
)
|
||||||
|
|
||||||
|
self._fp4_graph_row_limit = MAX_FUSED_ROWS
|
||||||
self.speculative_step_id = speculative_step_id
|
self.speculative_step_id = speculative_step_id
|
||||||
self.forward_metadata: Union[
|
self.forward_metadata: Union[
|
||||||
DSV4Metadata,
|
DSV4Metadata,
|
||||||
@@ -545,32 +553,29 @@ class DeepseekV4HipRadixBackend(
|
|||||||
compress_gpu_plan: bool = False,
|
compress_gpu_plan: bool = False,
|
||||||
extend_start_loc: Optional[torch.Tensor] = None,
|
extend_start_loc: Optional[torch.Tensor] = None,
|
||||||
attach_decode_streams: bool = False,
|
attach_decode_streams: bool = False,
|
||||||
|
# Whether num_tokens == sum(extend_seq_lens) exactly, which lets the
|
||||||
|
# token map skip an implicit D2H.
|
||||||
|
exact_num_tokens: bool = True,
|
||||||
) -> DSV4Metadata:
|
) -> DSV4Metadata:
|
||||||
if extend_start_loc is not None:
|
|
||||||
from sglang.kernels.ops.attention.dsv4_attn_metadata_kernels import (
|
from sglang.kernels.ops.attention.dsv4_attn_metadata_kernels import (
|
||||||
ExpandPrefillCausally,
|
ExpandPrefillCausally,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# extend_start_loc and the CPU mirrors below only feed the torch
|
||||||
|
# fallback; the triton kernel cumsums extend_seq_lens on device, so
|
||||||
|
# every caller can share it instead of dropping to a per-request loop.
|
||||||
_expanded = ExpandPrefillCausally.execute(
|
_expanded = ExpandPrefillCausally.execute(
|
||||||
req_pool_indices=req_pool_indices,
|
req_pool_indices=req_pool_indices,
|
||||||
seq_lens=seq_lens,
|
seq_lens=seq_lens,
|
||||||
extend_seq_lens=extend_seq_lens,
|
extend_seq_lens=extend_seq_lens,
|
||||||
extend_start_loc=extend_start_loc,
|
extend_start_loc=extend_start_loc,
|
||||||
seq_lens_cpu=None,
|
seq_lens_cpu=seq_lens_cpu,
|
||||||
extend_seq_lens_cpu=None,
|
extend_seq_lens_cpu=extend_seq_lens_cpu,
|
||||||
num_tokens=num_tokens,
|
num_tokens=num_tokens,
|
||||||
padded_num_tokens=out_cache_loc.shape[0],
|
padded_num_tokens=out_cache_loc.shape[0],
|
||||||
)
|
)
|
||||||
seq_lens_casual = _expanded.seq_lens_casual
|
seq_lens_casual = _expanded.seq_lens_casual
|
||||||
req_pool_indices_repeated = _expanded.req_pool_indices_repeated
|
req_pool_indices_repeated = _expanded.req_pool_indices_repeated
|
||||||
else:
|
|
||||||
seq_lens_casual, req_pool_indices_repeated = self.expand_prefill_casually(
|
|
||||||
num_tokens=num_tokens,
|
|
||||||
seq_lens=seq_lens_cpu,
|
|
||||||
extend_seq_lens=extend_seq_lens_cpu,
|
|
||||||
req_pool_indices=req_pool_indices,
|
|
||||||
padded_num_tokens=out_cache_loc.shape[0],
|
|
||||||
)
|
|
||||||
core_attn_metadata = self.make_core_attn_metadata(
|
core_attn_metadata = self.make_core_attn_metadata(
|
||||||
req_to_token=self.req_to_token,
|
req_to_token=self.req_to_token,
|
||||||
req_pool_indices_repeated=req_pool_indices_repeated,
|
req_pool_indices_repeated=req_pool_indices_repeated,
|
||||||
@@ -586,7 +591,7 @@ class DeepseekV4HipRadixBackend(
|
|||||||
seq_lens,
|
seq_lens,
|
||||||
extend_seq_lens,
|
extend_seq_lens,
|
||||||
num_tokens,
|
num_tokens,
|
||||||
need_compress=need_compress,
|
exact_num_tokens=exact_num_tokens,
|
||||||
)
|
)
|
||||||
if attach_decode_streams:
|
if attach_decode_streams:
|
||||||
# Target-verify runs through the unified_kv DECODE kernel, so build
|
# Target-verify runs through the unified_kv DECODE kernel, so build
|
||||||
@@ -648,8 +653,29 @@ class DeepseekV4HipRadixBackend(
|
|||||||
seq_lens_cpu: Optional[List[int]] = None,
|
seq_lens_cpu: Optional[List[int]] = None,
|
||||||
ragged_layout=None,
|
ragged_layout=None,
|
||||||
) -> Union[DSV4Metadata, DSV4RawVerifyMetadata]:
|
) -> Union[DSV4Metadata, DSV4RawVerifyMetadata]:
|
||||||
# HIP path: build target-verify metadata eagerly. The raw/lazy-upgrade route can
|
# DSPARK verifies a uniform num_draft block, exactly what
|
||||||
# hit planner invariants during graph capture for DSV4+EAGLE.
|
# make_forward_metadata_from_raw_verify expands, so the build can be
|
||||||
|
# deferred into the graph. Graph path only: the upgrade sizes its page
|
||||||
|
# table by MAX_SEQ_LEN_FOR_CAPTURE, far wider than the live max_seq_len
|
||||||
|
# an eager caller passes. EAGLE and ragged layouts stay eager -- no raw
|
||||||
|
# expansion, and EAGLE's fixed-tier plan trips planner invariants.
|
||||||
|
if (
|
||||||
|
use_prefill_cuda_graph
|
||||||
|
and self.is_dspark
|
||||||
|
and ragged_layout is None
|
||||||
|
and out_cache_loc is not None
|
||||||
|
# Oversized batches keep the eager build; see _fp4_graph_row_limit.
|
||||||
|
and (
|
||||||
|
self._fp4_graph_row_limit is None
|
||||||
|
or self.target_verify_num_draft_tokens * len(seq_lens)
|
||||||
|
<= self._fp4_graph_row_limit
|
||||||
|
)
|
||||||
|
):
|
||||||
|
return DSV4RawVerifyMetadata(
|
||||||
|
req_pool_indices=req_pool_indices,
|
||||||
|
seq_lens=seq_lens,
|
||||||
|
out_cache_loc=out_cache_loc,
|
||||||
|
)
|
||||||
if seq_lens_cpu is None:
|
if seq_lens_cpu is None:
|
||||||
seq_lens_cpu = seq_lens.tolist()
|
seq_lens_cpu = seq_lens.tolist()
|
||||||
return self.init_forward_metadata_target_verify_old(
|
return self.init_forward_metadata_target_verify_old(
|
||||||
@@ -692,8 +718,11 @@ class DeepseekV4HipRadixBackend(
|
|||||||
num_tokens = ragged_layout.total_verify_tokens
|
num_tokens = ragged_layout.total_verify_tokens
|
||||||
if num_tokens is None:
|
if num_tokens is None:
|
||||||
num_tokens = int(verify_lens_dev.sum().item())
|
num_tokens = int(verify_lens_dev.sum().item())
|
||||||
|
exact_num_tokens = True
|
||||||
else:
|
else:
|
||||||
num_tokens = int(num_tokens)
|
num_tokens = int(num_tokens)
|
||||||
|
# Padded tier: num_tokens >= sum(verify_lens)
|
||||||
|
exact_num_tokens = False
|
||||||
extend_seq_lens_cpu = None
|
extend_seq_lens_cpu = None
|
||||||
seq_lens_cpu = None
|
seq_lens_cpu = None
|
||||||
else:
|
else:
|
||||||
@@ -704,6 +733,7 @@ class DeepseekV4HipRadixBackend(
|
|||||||
extend_seq_lens_cpu = [self.target_verify_num_draft_tokens] * batch_size
|
extend_seq_lens_cpu = [self.target_verify_num_draft_tokens] * batch_size
|
||||||
num_tokens = self.target_verify_num_draft_tokens * batch_size
|
num_tokens = self.target_verify_num_draft_tokens * batch_size
|
||||||
extend_seq_lens = self._move_to_device(extend_seq_lens_cpu)
|
extend_seq_lens = self._move_to_device(extend_seq_lens_cpu)
|
||||||
|
exact_num_tokens = True
|
||||||
if out_cache_loc is None:
|
if out_cache_loc is None:
|
||||||
out_cache_loc = seq_lens.new_zeros(num_tokens)
|
out_cache_loc = seq_lens.new_zeros(num_tokens)
|
||||||
return self.init_forward_metadata_prefill(
|
return self.init_forward_metadata_prefill(
|
||||||
@@ -720,6 +750,7 @@ class DeepseekV4HipRadixBackend(
|
|||||||
compress_gpu_plan=ragged_layout is not None,
|
compress_gpu_plan=ragged_layout is not None,
|
||||||
extend_start_loc=extend_start_loc,
|
extend_start_loc=extend_start_loc,
|
||||||
attach_decode_streams=True,
|
attach_decode_streams=True,
|
||||||
|
exact_num_tokens=exact_num_tokens,
|
||||||
)
|
)
|
||||||
|
|
||||||
def make_forward_metadata_from_raw_verify(
|
def make_forward_metadata_from_raw_verify(
|
||||||
@@ -753,6 +784,18 @@ class DeepseekV4HipRadixBackend(
|
|||||||
out_loc=out_cache_loc,
|
out_loc=out_cache_loc,
|
||||||
need_compress=True,
|
need_compress=True,
|
||||||
)
|
)
|
||||||
|
# extend_seq_lens is uniform here (seq_lens already carries the draft
|
||||||
|
# block, so the minimum above cannot trim it), hence an exact token count.
|
||||||
|
self._attach_unified_kv_prefill_meta(
|
||||||
|
core_attn_metadata,
|
||||||
|
req_pool_indices,
|
||||||
|
seq_lens,
|
||||||
|
extend_seq_lens,
|
||||||
|
num_draft_tokens * bs,
|
||||||
|
)
|
||||||
|
self._attach_unified_kv_decode_streams(
|
||||||
|
core_attn_metadata, req_pool_indices_repeated
|
||||||
|
)
|
||||||
indexer_metadata = self.init_forward_metadata_indexer(core_attn_metadata)
|
indexer_metadata = self.init_forward_metadata_indexer(core_attn_metadata)
|
||||||
create = functools.partial(
|
create = functools.partial(
|
||||||
create_paged_compressor_data,
|
create_paged_compressor_data,
|
||||||
@@ -843,7 +886,8 @@ class DeepseekV4HipRadixBackend(
|
|||||||
|
|
||||||
def init_forward_metadata_in_graph(self, forward_batch: ForwardBatch) -> None:
|
def init_forward_metadata_in_graph(self, forward_batch: ForwardBatch) -> None:
|
||||||
# Raw metadata must be materialized inside the graph to refresh on replay.
|
# Raw metadata must be materialized inside the graph to refresh on replay.
|
||||||
if isinstance(self.forward_metadata, DSV4RawVerifyMetadata):
|
upgraded_verify = isinstance(self.forward_metadata, DSV4RawVerifyMetadata)
|
||||||
|
if upgraded_verify:
|
||||||
self.forward_metadata = self.make_forward_metadata_from_raw_verify(
|
self.forward_metadata = self.make_forward_metadata_from_raw_verify(
|
||||||
raw_metadata=self.forward_metadata,
|
raw_metadata=self.forward_metadata,
|
||||||
)
|
)
|
||||||
@@ -894,6 +938,12 @@ class DeepseekV4HipRadixBackend(
|
|||||||
torch.int64
|
torch.int64
|
||||||
)
|
)
|
||||||
|
|
||||||
|
if upgraded_verify:
|
||||||
|
# The out-graph refresh saw raw metadata and skipped. Without this
|
||||||
|
# the logits kernel builds its own schedule -- the variant that
|
||||||
|
# frees the scratch it reads, which every replay would re-read.
|
||||||
|
self._refresh_fp4_prefill_workspace(forward_batch)
|
||||||
|
|
||||||
# Decode's schedule builder is capture-safe because the workspace pins
|
# Decode's schedule builder is capture-safe because the workspace pins
|
||||||
# the scratch it reads, so it can stay next to the metadata it consumes.
|
# the scratch it reads, so it can stay next to the metadata it consumes.
|
||||||
# Prefill/target-verify cannot; see _refresh_fp4_prefill_workspace.
|
# Prefill/target-verify cannot; see _refresh_fp4_prefill_workspace.
|
||||||
@@ -1159,6 +1209,7 @@ class DeepseekV4HipRadixBackend(
|
|||||||
extend_seq_lens=extend_seq_lens,
|
extend_seq_lens=extend_seq_lens,
|
||||||
extend_seq_lens_cpu=extend_seq_lens_cpu,
|
extend_seq_lens_cpu=extend_seq_lens_cpu,
|
||||||
need_compress=not is_draft,
|
need_compress=not is_draft,
|
||||||
|
exact_num_tokens=is_draft,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
raise NotImplementedError(f"unsupported mode {forward_batch.forward_mode=}")
|
raise NotImplementedError(f"unsupported mode {forward_batch.forward_mode=}")
|
||||||
@@ -1276,7 +1327,7 @@ class DeepseekV4HipRadixBackend(
|
|||||||
seq_lens: torch.Tensor,
|
seq_lens: torch.Tensor,
|
||||||
extend_seq_lens: torch.Tensor,
|
extend_seq_lens: torch.Tensor,
|
||||||
num_tokens: int,
|
num_tokens: int,
|
||||||
need_compress: bool = True,
|
exact_num_tokens: bool = True,
|
||||||
) -> None:
|
) -> None:
|
||||||
from sglang.kernels.ops.attention.dsv4.unified_kv_kernels.env_gate import (
|
from sglang.kernels.ops.attention.dsv4.unified_kv_kernels.env_gate import (
|
||||||
is_unified_kv_triton,
|
is_unified_kv_triton,
|
||||||
@@ -1289,18 +1340,12 @@ class DeepseekV4HipRadixBackend(
|
|||||||
seq_lens = seq_lens.to(torch.int64)
|
seq_lens = seq_lens.to(torch.int64)
|
||||||
extend_seq_lens = extend_seq_lens.to(torch.int64)
|
extend_seq_lens = extend_seq_lens.to(torch.int64)
|
||||||
# token -> req index (length L = sum(extend_seq_lens)).
|
# token -> req index (length L = sum(extend_seq_lens)).
|
||||||
# output_size skips the implicit sum() D2H on draft-extend. dropping it on the
|
# output_size skips the implicit sum() D2H, but it must equal L:
|
||||||
# target-extend path triggers a GPU memory access fault.
|
# exact_num_tokens tells whether num_tokens does.
|
||||||
if need_compress:
|
|
||||||
bid = torch.repeat_interleave(
|
bid = torch.repeat_interleave(
|
||||||
torch.arange(bs, device=device, dtype=torch.int64),
|
torch.arange(bs, device=device, dtype=torch.int64),
|
||||||
extend_seq_lens,
|
extend_seq_lens,
|
||||||
)
|
output_size=num_tokens if exact_num_tokens else None,
|
||||||
else:
|
|
||||||
bid = torch.repeat_interleave(
|
|
||||||
torch.arange(bs, device=device, dtype=torch.int64),
|
|
||||||
extend_seq_lens,
|
|
||||||
output_size=num_tokens,
|
|
||||||
)
|
)
|
||||||
if core.unified is None:
|
if core.unified is None:
|
||||||
core.unified = UnifiedKvMetadata()
|
core.unified = UnifiedKvMetadata()
|
||||||
@@ -1680,41 +1725,6 @@ class DeepseekV4HipRadixBackend(
|
|||||||
|
|
||||||
raise NotImplementedError("ragged attention")
|
raise NotImplementedError("ragged attention")
|
||||||
|
|
||||||
def expand_prefill_casually(
|
|
||||||
self,
|
|
||||||
num_tokens: int,
|
|
||||||
seq_lens: List[int],
|
|
||||||
extend_seq_lens: List[int],
|
|
||||||
req_pool_indices: torch.Tensor,
|
|
||||||
padded_num_tokens: Optional[int],
|
|
||||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
|
||||||
seq_lens_casual = torch.empty(num_tokens, **self.cuda_int32_kwargs)
|
|
||||||
idx_to_req_repeated = torch.empty(num_tokens, **self.cuda_int32_kwargs)
|
|
||||||
offset = 0
|
|
||||||
for i, (kv_len, qo_len) in enumerate(zip(seq_lens, extend_seq_lens)):
|
|
||||||
out = seq_lens_casual[offset : offset + qo_len]
|
|
||||||
offset += qo_len
|
|
||||||
torch.arange(kv_len - qo_len + 1, kv_len + 1, out=out)
|
|
||||||
idx_to_req_repeated[offset - qo_len : offset].fill_(i)
|
|
||||||
|
|
||||||
assert offset == num_tokens
|
|
||||||
req_pool_indices_repeated = req_pool_indices[idx_to_req_repeated]
|
|
||||||
|
|
||||||
if padded_num_tokens is not None and padded_num_tokens > num_tokens:
|
|
||||||
pad_size = padded_num_tokens - num_tokens
|
|
||||||
seq_lens_casual = torch.nn.functional.pad(
|
|
||||||
seq_lens_casual,
|
|
||||||
(0, pad_size),
|
|
||||||
value=1,
|
|
||||||
)
|
|
||||||
req_pool_indices_repeated = torch.nn.functional.pad(
|
|
||||||
req_pool_indices_repeated,
|
|
||||||
(0, pad_size),
|
|
||||||
value=req_pool_indices_repeated[-1].item(),
|
|
||||||
)
|
|
||||||
|
|
||||||
return seq_lens_casual, req_pool_indices_repeated
|
|
||||||
|
|
||||||
def expand_extend_with_same_length(
|
def expand_extend_with_same_length(
|
||||||
self,
|
self,
|
||||||
bs: int,
|
bs: int,
|
||||||
|
|||||||
@@ -11,12 +11,19 @@ Registry: nightly-amd-8-gpu-mi35x-deepseek-v4-pro-dspark suite
|
|||||||
import os
|
import os
|
||||||
import unittest
|
import unittest
|
||||||
from types import SimpleNamespace
|
from types import SimpleNamespace
|
||||||
|
from unittest import mock
|
||||||
|
|
||||||
import requests
|
import requests
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
|
from sglang.kernels.ops.attention.dsv4.fp4_indexer_schedule_hip import MAX_FUSED_ROWS
|
||||||
from sglang.kernels.ops.attention.dsv4.unified_kv_kernels import runtime
|
from sglang.kernels.ops.attention.dsv4.unified_kv_kernels import runtime
|
||||||
from sglang.kernels.ops.speculative.dspark import dspark_verify_window
|
from sglang.kernels.ops.speculative.dspark import dspark_verify_window
|
||||||
|
from sglang.srt.layers.attention.deepseek_v4_backend_hip_radix import (
|
||||||
|
DeepseekV4HipRadixBackend,
|
||||||
|
DSV4RawVerifyMetadata,
|
||||||
|
UnifiedKvMetadata,
|
||||||
|
)
|
||||||
from sglang.srt.utils import kill_process_tree
|
from sglang.srt.utils import kill_process_tree
|
||||||
from sglang.test.ci.ci_register import register_amd_ci
|
from sglang.test.ci.ci_register import register_amd_ci
|
||||||
from sglang.test.few_shot_gsm8k import run_eval as run_eval_few_shot_gsm8k
|
from sglang.test.few_shot_gsm8k import run_eval as run_eval_few_shot_gsm8k
|
||||||
@@ -60,6 +67,60 @@ FP4_ENV_VARS = {
|
|||||||
|
|
||||||
|
|
||||||
class TestDSparkUnifiedKVKernelsAMD(CustomTestCase):
|
class TestDSparkUnifiedKVKernelsAMD(CustomTestCase):
|
||||||
|
def test_unified_metadata_copy_updates_captured_swa_loc_in_place(self):
|
||||||
|
captured_swa_loc = torch.tensor([3, 5, 7], device=DEVICE, dtype=torch.int32)
|
||||||
|
replay_swa_loc = torch.tensor([11, 13, 17], device=DEVICE, dtype=torch.int32)
|
||||||
|
captured = UnifiedKvMetadata(swa_loc=captured_swa_loc)
|
||||||
|
replay = UnifiedKvMetadata(swa_loc=replay_swa_loc)
|
||||||
|
|
||||||
|
captured.copy_(replay)
|
||||||
|
|
||||||
|
self.assertIs(captured.swa_loc, captured_swa_loc)
|
||||||
|
self.assertTrue(torch.equal(captured.swa_loc, replay_swa_loc))
|
||||||
|
|
||||||
|
def test_dspark_verify_metadata_graph_routing(self):
|
||||||
|
num_draft_tokens = 7
|
||||||
|
max_fused_bs = MAX_FUSED_ROWS // num_draft_tokens
|
||||||
|
cases = (
|
||||||
|
("eligible", True, None, max_fused_bs, True),
|
||||||
|
("fp4_row_limit", True, None, max_fused_bs + 1, False),
|
||||||
|
("eagle", False, None, max_fused_bs, False),
|
||||||
|
("ragged", True, object(), max_fused_bs, False),
|
||||||
|
)
|
||||||
|
|
||||||
|
for name, is_dspark, ragged_layout, bs, expect_raw in cases:
|
||||||
|
with self.subTest(name=name):
|
||||||
|
backend = DeepseekV4HipRadixBackend.__new__(DeepseekV4HipRadixBackend)
|
||||||
|
backend.is_dspark = is_dspark
|
||||||
|
backend._fp4_graph_row_limit = MAX_FUSED_ROWS
|
||||||
|
backend.target_verify_num_draft_tokens = num_draft_tokens
|
||||||
|
eager_metadata = object()
|
||||||
|
backend.init_forward_metadata_target_verify_old = mock.Mock(
|
||||||
|
return_value=eager_metadata
|
||||||
|
)
|
||||||
|
|
||||||
|
req_pool_indices = torch.arange(bs, device=DEVICE, dtype=torch.int32)
|
||||||
|
seq_lens = torch.ones(bs, device=DEVICE, dtype=torch.int32)
|
||||||
|
out_cache_loc = torch.zeros(
|
||||||
|
bs * num_draft_tokens, device=DEVICE, dtype=torch.int32
|
||||||
|
)
|
||||||
|
result = backend.init_forward_metadata_target_verify(
|
||||||
|
max_seq_len=128,
|
||||||
|
req_pool_indices=req_pool_indices,
|
||||||
|
seq_lens=seq_lens,
|
||||||
|
out_cache_loc=out_cache_loc,
|
||||||
|
use_prefill_cuda_graph=True,
|
||||||
|
seq_lens_cpu=[1] * bs,
|
||||||
|
ragged_layout=ragged_layout,
|
||||||
|
)
|
||||||
|
|
||||||
|
if expect_raw:
|
||||||
|
self.assertIsInstance(result, DSV4RawVerifyMetadata)
|
||||||
|
backend.init_forward_metadata_target_verify_old.assert_not_called()
|
||||||
|
else:
|
||||||
|
self.assertIs(result, eager_metadata)
|
||||||
|
backend.init_forward_metadata_target_verify_old.assert_called_once()
|
||||||
|
|
||||||
def test_build_unified_commit_inject_layout(self):
|
def test_build_unified_commit_inject_layout(self):
|
||||||
stride, ring_stride = 7, 128
|
stride, ring_stride = 7, 128
|
||||||
req_pool_indices = torch.tensor([3, 0, 5, 1], device=DEVICE, dtype=torch.int32)
|
req_pool_indices = torch.tensor([3, 0, 5, 1], device=DEVICE, dtype=torch.int32)
|
||||||
|
|||||||
Reference in New Issue
Block a user