diff --git a/docs/cookbook/autoregressive/Qwen/Qwen3.8-27B.mdx b/docs/cookbook/autoregressive/Qwen/Qwen3.8-27B.mdx index 2baa0a63e..927522010 100644 --- a/docs/cookbook/autoregressive/Qwen/Qwen3.8-27B.mdx +++ b/docs/cookbook/autoregressive/Qwen/Qwen3.8-27B.mdx @@ -247,9 +247,24 @@ checkpoint's calibration scales automatically. `--speculative-algorithm DFLASH --speculative-draft-model-path incoai/Qwen3.8-27B-DFlash2 --speculative-num-draft-tokens 8` (8 is the draft's block size, and it is the `D` term in the ratio, same value as - DSpark's). The selector projects candidates through the target `lm_head`, + DSpark's). Runs on Ascend NPUs as well + ([#35629](https://github.com/sgl-project/sglang/pull/35629)): the selector + verify falls back to argmax there, matching what EAGLE and first-generation + DFlash drafts (such as `z-lab/Qwen3-8B-DFlash-b16`) already do on NPU. + Therefore, NPU currently guarantees lossless verification only for greedy + requests; use `temperature=0` and `top_k=1`. Non-greedy requests log a warning; + both draft proposal and target verification fall back to greedy, so their + requested sampling distribution is not preserved. The selector projects + candidates through the target `lm_head`, including quantized heads, so it runs on the NVFP4 checkpoint (whose head is NVFP4-packed; the BF16 and FP8 checkpoints keep a dense head). + The Ascend comparison in #35629 used a 910C with BF16 target weights, + `--tp-size 2 --attention-backend ascend --mamba-ssm-dtype bfloat16 + --mamba-scheduler-strategy extra_buffer`, and disabled RadixCache for both + baseline and DFlash2 to exclude cache warm-up and prefix reuse. The DFlash2 + run added the three flags shown above. + Accuracy used zero-shot GSM8K with greedy sampling, `max_new_tokens=2048`, + 128 examples, and concurrency levels 1, 2, 4, 8, and 16. Validation: NVFP4 measured end-to-end on RTX PRO 6000 and RTX 5090; the RTX PRO 6000 BF16/FP8 cells boot and serve; all 12 DGX Spark DFLASH2 cells boot and serve on `1cf2b8c` with the selector folded into the draft CUDA diff --git a/python/sglang/srt/speculative/dflash_worker_v2.py b/python/sglang/srt/speculative/dflash_worker_v2.py index 7df89d5a8..a008f7a0b 100644 --- a/python/sglang/srt/speculative/dflash_worker_v2.py +++ b/python/sglang/srt/speculative/dflash_worker_v2.py @@ -217,16 +217,19 @@ def _selector_lattice(draft_model, pred_hidden, anchor_token_ids): class _SelectorDraftSampler: - """Selector decode folded into the draft cuda graph, greedy and T>0 alike. + """Selector decode folded into the draft cuda graph. - One captured graph serves both: it always walks the sampling path, and a static - greedy_mask selects the argmax per row. + On sampling-enabled backends, one captured graph serves both: it always walks + the sampling path, and a static greedy_mask selects the argmax per row. """ - def __init__(self, *, draft_model, block_size, max_bs, device): + def __init__( + self, *, draft_model, block_size, max_bs, device, sampling_enabled: bool + ): self.draft_model = draft_model self.selector = draft_model.candidate_selector self.block_size = int(block_size) + self.sampling_enabled = sampling_enabled max_bs, gamma, top_k = int(max_bs), self.block_size - 1, self.selector.top_k self.out = torch.empty((max_bs * gamma,), dtype=torch.int64, device=device) # Written by the host before replay, or read after it; the addresses are @@ -244,7 +247,7 @@ class _SelectorDraftSampler: def stage_sampling_params(self, *, bs: int, sampling_info) -> None: """Host-side refresh of the static sampling params; must run before the draft graph replay that consumes them.""" - if sampling_info is None: + if sampling_info is None or not self.sampling_enabled: self.temperatures[:bs].fill_(1.0) self.greedy_mask[:bs].fill_(True) return @@ -325,6 +328,8 @@ class DFlashWorkerV2(BaseSpecWorker): self._draft_sampler = None self.draft_model = bundle.draft_model self.selector = self.draft_model.candidate_selector + # Ascend keeps selector proposal aligned with its greedy-only verify path. + self._selector_sampling_enabled = not _is_npu draft_config = parse_dflash_draft_config( draft_hf_config=self.draft_model_runner.model_config.hf_config ) @@ -639,14 +644,16 @@ class DFlashWorkerV2(BaseSpecWorker): self.draft_model.lm_head = lm_head if self.ps.tp_rank == 0: logger.info( - "DFLASH selector decode (greedy + sampling) folded into the " - "draft cuda graph." + "DFLASH selector decode folded into the draft cuda graph " + "(sampling_enabled=%s).", + self._selector_sampling_enabled, ) return _SelectorDraftSampler( draft_model=self.draft_model, block_size=self.block_size, max_bs=max(get_exec().graph.cuda_graph_config.decode.bs), device=self.device, + sampling_enabled=self._selector_sampling_enabled, ) if not hasattr(lm_head, "weight"): return _eager("quantized lm_head has no dense weight") @@ -1106,19 +1113,22 @@ class DFlashWorkerV2(BaseSpecWorker): # Clamped like DSpark so greedy rows don't divide by zero. temperatures = ( torch.ones(bs, dtype=torch.float32, device=device) - if sampling_info is None + if sampling_info is None or not self._selector_sampling_enabled else sampling_info.temperatures.view(-1).float().clamp_min(1e-5) ) + greedy_mask = ( + torch.ones(bs, dtype=torch.bool, device=device) + if not self._selector_sampling_enabled + else resolve_greedy_mask(bs=bs, sampling_info=sampling_info, device=device) + ) tokens, q_rows = self.selector.sample_path( candidate_ids=candidate_ids, scores=scores, uniforms=torch.rand(bs, num_pred, dtype=torch.float32, device=device), temperatures=temperatures, - greedy_mask=resolve_greedy_mask( - bs=bs, sampling_info=sampling_info, device=device - ), + greedy_mask=greedy_mask, ) - if not _is_all_greedy(sampling_info): + if self._selector_sampling_enabled and not _is_all_greedy(sampling_info): self._selector_sample = (candidate_ids, q_rows) return tokens.view(bs, num_pred) @@ -1833,13 +1843,20 @@ class DFlashWorkerV2(BaseSpecWorker): def _validate_phase1_sampling_support(self, batch: ScheduleBatch) -> None: sampling_info = batch.sampling_info - # A selector draft carries its own q and verifies through accept_sampling, so - # it never falls back to greedy argmax however this build was compiled. - if ( - sampling_info is None - or sampling_info.is_all_greedy - or self.selector is not None - ): + if sampling_info is None or sampling_info.is_all_greedy: + return + + if self.selector is not None: + if self._selector_sampling_enabled: + return + if not self._warned_sampling_fallback and self.ps.tp_rank == 0: + logger.warning( + "DFLASH non-greedy verification is unavailable on this " + "build/device; falling back to greedy argmax verification. " + "The requested sampling distribution will not be preserved; " + "use temperature=0 and top_k=1 for lossless greedy decoding." + ) + self._warned_sampling_fallback = True return if ( @@ -2142,7 +2159,11 @@ class DFlashWorkerV2(BaseSpecWorker): draft_next = self._draft_sampler.out[ : bs * (int(self.block_size) - 1) ].view(bs, int(self.block_size) - 1) - if self.selector is not None and not _is_all_greedy(batch.sampling_info): + if ( + self.selector is not None + and not _is_all_greedy(batch.sampling_info) + and self._selector_sampling_enabled + ): self._selector_sample = ( self._draft_sampler.candidate_out[:bs], self._draft_sampler.q_out[:bs], @@ -2216,7 +2237,7 @@ class DFlashWorkerV2(BaseSpecWorker): batch=None, forward_batch=verify_forward_batch, is_verify=True, - skip_attn_backend_init=True, + skip_attn_backend_init=True if not _is_npu else None, ) logits_output = target_out.logits_output can_run_cuda_graph = target_out.can_run_cuda_graph diff --git a/test/registered/unit/spec/test_dflash_extra_buffer_lazy.py b/test/registered/unit/spec/test_dflash_extra_buffer_lazy.py index 2a7d73f71..fa2a3dca7 100644 --- a/test/registered/unit/spec/test_dflash_extra_buffer_lazy.py +++ b/test/registered/unit/spec/test_dflash_extra_buffer_lazy.py @@ -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, diff --git a/test/registered/unit/spec/test_dflash_logits.py b/test/registered/unit/spec/test_dflash_logits.py index 821375e07..df343070f 100644 --- a/test/registered/unit/spec/test_dflash_logits.py +++ b/test/registered/unit/spec/test_dflash_logits.py @@ -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."""