[NPU] Adapt DFlash2 speculative decoding to Ascend NPUs (#35629)

Signed-off-by: syd520zy <529477025@qq.com>
Co-authored-by: github-actions[bot] <github-actions[bot]@users.noreply.github.com>
This commit is contained in:
faceless void
2026-09-07 09:08:39 +08:00
committed by GitHub
co-authored by github-actions[bot]
parent 6252993afe
commit 30d0eb2ca9
4 changed files with 202 additions and 24 deletions
@@ -120,7 +120,9 @@ class TestDflashVerifyRunsMambaTrackHook(CustomTestCase):
calls.append("init_new")
return fake_forward_batch
# This test covers the generic/CUDA hook ordering, not NPU DSV4 bundle setup.
with (
mock.patch.object(dflash_info, "_is_npu", False),
mock.patch(
"sglang.srt.speculative.spec_utils.prepare_mamba_track_for_verify",
side_effect=fake_hook,
+142 -2
View File
@@ -221,10 +221,16 @@ def test_worker_folds_a_gate_admitted_quantized_selector_head(monkeypatch):
from sglang.srt.speculative import dflash_worker_v2 as worker_mod
built = {}
built_sampler = object()
def build_sampler(**kwargs):
built.update(kwargs)
return built_sampler
monkeypatch.setattr(
worker_mod,
"_SelectorDraftSampler",
lambda **kwargs: built.setdefault("sampler", object()),
build_sampler,
)
monkeypatch.setattr(
worker_mod,
@@ -245,13 +251,15 @@ def test_worker_folds_a_gate_admitted_quantized_selector_head(monkeypatch):
ps=SimpleNamespace(tp_rank=0),
draft_model=SimpleNamespace(lm_head=None),
device="cpu",
_selector_sampling_enabled=True,
_target_worker=SimpleNamespace(
model_runner=SimpleNamespace(model=SimpleNamespace(lm_head=quant_head))
),
)
sampler = worker_mod.DFlashWorkerV2._maybe_build_draft_sampler(worker)
assert sampler is built["sampler"]
assert sampler is built_sampler
assert built["sampling_enabled"] is True
assert worker.draft_model.lm_head is quant_head
# A packed head without an applicable quant method must stay eager.
@@ -263,6 +271,138 @@ def test_worker_folds_a_gate_admitted_quantized_selector_head(monkeypatch):
assert worker.draft_model.lm_head is None
def test_worker_warns_once_when_selector_sampling_is_disabled(monkeypatch):
from sglang.srt.speculative import dflash_worker_v2 as worker_mod
warnings = []
monkeypatch.setattr(
worker_mod.logger, "warning", lambda *args: warnings.append(args)
)
worker = SimpleNamespace(
selector=object(),
_selector_sampling_enabled=False,
_warned_sampling_fallback=False,
ps=SimpleNamespace(tp_rank=0),
)
batch = SimpleNamespace(sampling_info=SimpleNamespace(is_all_greedy=False))
worker_mod.DFlashWorkerV2._validate_phase1_sampling_support(worker, batch)
worker_mod.DFlashWorkerV2._validate_phase1_sampling_support(worker, batch)
assert worker._warned_sampling_fallback
assert len(warnings) == 1
assert "sampling distribution will not be preserved" in warnings[0][0]
worker._selector_sampling_enabled = True
worker._warned_sampling_fallback = False
worker_mod.DFlashWorkerV2._validate_phase1_sampling_support(worker, batch)
assert len(warnings) == 1
def test_disabled_selector_sampling_forces_greedy_draft():
from sglang.srt.speculative import dflash_worker_v2 as worker_mod
sampling_info = SimpleNamespace(
temperatures=torch.tensor([[0.7]]),
top_ks=torch.tensor([[8]]),
is_all_greedy=False,
)
sampler = worker_mod._SelectorDraftSampler.__new__(worker_mod._SelectorDraftSampler)
sampler.temperatures = torch.zeros(1)
sampler.greedy_mask = torch.zeros(1, dtype=torch.bool)
sampler.sampling_enabled = False
sampler.stage_sampling_params(bs=1, sampling_info=sampling_info)
torch.testing.assert_close(sampler.temperatures, torch.ones(1))
assert sampler.greedy_mask.tolist() == [True]
sampler.sampling_enabled = True
sampler.stage_sampling_params(bs=1, sampling_info=sampling_info)
torch.testing.assert_close(sampler.temperatures, torch.tensor([0.7]))
assert sampler.greedy_mask.tolist() == [False]
observed = {}
def sample_path(**kwargs):
observed.update(kwargs)
return torch.zeros((1, 1), dtype=torch.int64), torch.zeros((1, 1, 2))
selector = SimpleNamespace(
build_lattice=lambda **kwargs: torch.zeros((1, 1, 2, 2)),
sample_path=sample_path,
)
draft_model = SimpleNamespace(
lm_head=None,
candidate_selector=selector,
compute_candidates=lambda hidden: (
torch.zeros((1, 2), dtype=torch.int64),
torch.zeros((1, 2)),
),
)
worker = SimpleNamespace(
draft_model=draft_model,
selector=selector,
block_size=2,
_selector_sampling_enabled=False,
_selector_sample=None,
)
draft_logits_output = SimpleNamespace(hidden_states=torch.zeros((2, 4)))
worker_mod.DFlashWorkerV2._propose_selector_block(
worker,
draft_logits_output=draft_logits_output,
bs=1,
lm_head=object(),
anchor_token_ids=torch.zeros(1, dtype=torch.int64),
sampling_info=sampling_info,
)
torch.testing.assert_close(observed["temperatures"], torch.ones(1))
assert observed["greedy_mask"].tolist() == [True]
assert worker._selector_sample is None
def test_selector_accept_uses_greedy_fallback_without_staged_sample(monkeypatch):
from sglang.srt.speculative import dflash_worker_v2 as worker_mod
monkeypatch.setattr(
worker_mod, "is_dflash_sampling_verify_available", lambda: False
)
monkeypatch.setattr(
worker_mod,
"compute_dflash_correct_drafts_and_bonus",
lambda **kwargs: (torch.tensor([0]), torch.tensor([7])),
)
sync_sites = []
worker = SimpleNamespace(
_selector_sample=None,
_selector_sampling_accept=lambda **kwargs: pytest.fail(
"selector sampling must not run without a staged sample"
),
_tp_sync=SimpleNamespace(sync=lambda site, tensor: sync_sites.append(site)),
_use_triton_accept_bonus=False,
block_size=2,
)
result = worker_mod.DFlashWorkerV2._accept_block(
worker,
candidates=torch.tensor([[9, 1]]),
next_token_logits=torch.tensor([[[0.0, 1.0], [1.0, 0.0]]]),
sampling_info=SimpleNamespace(is_all_greedy=False),
draft_input=object(),
prefix_lens=torch.tensor([3]),
bs=1,
)
accept_len, commit_lens, bonus, out_tokens, _, target_predict = result
assert accept_len.tolist() == [0]
assert commit_lens.tolist() == [1]
assert bonus.tolist() == [7]
assert out_tokens.tolist() == [[7, 0]]
assert target_predict.tolist() == [[1, 0]]
assert sync_sites == [worker_mod.SpecTpSyncSite.DFLASH_ACCEPT_GREEDY]
def test_grouped_conv_supports_runtime_block_sizes():
"""The conv indexes a position inside the block, so it must follow whatever
block size the worker resolved -- including one that is not a power of two."""