From 3c8ccdf46b70edbf5c13daea682df4ab37ee9f41 Mon Sep 17 00:00:00 2001 From: harshal-96 Date: Tue, 1 Sep 2026 13:31:16 +0530 Subject: [PATCH 1/3] feat: support original z-lab Qwen3 DFlash checkpoints and mixed dtypes Validated dense Qwen3 end-to-end on the DFlash stack (RTX 4070 8 GB, target thewimo/Qwen3-4B-AWQ float16, draft z-lab/Qwen3-4B-DFlash-b16 bfloat16, greedy, 256 new tokens): 33.6-36.8 tok/s vs 25.1 baseline, mean 4.19 of 15 drafts accepted per block, coherent lossless output. Two gaps had to be fixed to get there: - parse_dflash_config required dflash_config['block_size'], but the original z-lab Qwen3 checkpoints (Qwen3-4B/8B-DFlash-b16) store block_size at the top level of config.json. Fall back to the top-level attribute (same resolution order as the reference implementation), for mask_token_id as well, with a regression test. - The draft ran the shared target embedding output and the target aux states in the target dtype through bfloat16 draft weights, which fails with a dtype mismatch for float16 targets (e.g. AWQ). Casting the whole engine to float16 instead is not an option: Qwen3 hidden state outliers overflow float16 during feature fusion and acceptance collapses to zero (measured 0.06 of 15). Cast to the draft dtype at the two ingestion boundaries; logits already cast via get_logits. Signed-off-by: harshal-96 --- lmdeploy/pytorch/models/qwen3_dflash.py | 7 +++++-- lmdeploy/pytorch/spec_decode/dflash_utils.py | 10 ++++++++-- tests/pytorch/spec_decode/test_dflash_utils.py | 15 +++++++++++++++ 3 files changed, 28 insertions(+), 4 deletions(-) diff --git a/lmdeploy/pytorch/models/qwen3_dflash.py b/lmdeploy/pytorch/models/qwen3_dflash.py index 75f42d8ed7..a0682c4d9b 100644 --- a/lmdeploy/pytorch/models/qwen3_dflash.py +++ b/lmdeploy/pytorch/models/qwen3_dflash.py @@ -352,7 +352,9 @@ def embed_input_ids(self, input_ids: torch.Tensor) -> torch.Tensor: if self.has_separate_mask_embedding and self.mask_token_id is not None: mask = (input_ids == int(self.mask_token_id)).unsqueeze(-1) embeds = torch.where(mask, self.mask_embedding.to(dtype=embeds.dtype), embeds) - return embeds + # the shared target embedding may run in a different dtype than the + # draft (e.g. float16 AWQ target with a bfloat16 draft checkpoint) + return embeds.to(self.dtype) def project_target_hidden(self, target_hidden: torch.Tensor) -> torch.Tensor: """Project concatenated target-layer hidden states into draft hidden @@ -362,7 +364,8 @@ def project_target_hidden(self, target_hidden: torch.Tensor) -> torch.Tensor: raise ValueError('DFlash target hidden feature dim mismatch. ' f'Expected shape [N, {expected}] from {self.num_context_features} target layers, ' f'got {tuple(target_hidden.shape)}.') - return self.hidden_norm(self.fc(target_hidden)) + # target aux states arrive in the target dtype; fuse in the draft dtype + return self.hidden_norm(self.fc(target_hidden.to(self.dtype))) def _rotary_pos_emb_for_context( self, diff --git a/lmdeploy/pytorch/spec_decode/dflash_utils.py b/lmdeploy/pytorch/spec_decode/dflash_utils.py index d5506d148a..e7deeb9a5c 100644 --- a/lmdeploy/pytorch/spec_decode/dflash_utils.py +++ b/lmdeploy/pytorch/spec_decode/dflash_utils.py @@ -139,14 +139,20 @@ def parse_dflash_config(draft_hf_config: Any, num_speculative_tokens: int, f'draft declares num_target_layers={num_target_layers}, ' f'but target ModelConfig has num_layers={target_num_layers}.') - max_query_length = dflash_config['block_size'] + # Original z-lab DFlash checkpoints (e.g. Qwen3-4B/8B-DFlash-b16) store + # block_size at the top level of config.json; newer checkpoints nest it + # in dflash_config. Mirror the reference implementation's fallback order + # (dflash_config first, then the top-level config attribute). + max_query_length = dflash_config.get('block_size', getattr(draft_hf_config, 'block_size', None)) + if max_query_length is None: + raise ValueError('DFlash checkpoint requires block_size in dflash_config or at the top level of its config.') query_length = num_speculative_tokens + 1 if query_length > max_query_length: raise ValueError('DFlash query length (1 + speculative_num_draft_tokens) must not exceed checkpoint ' 'dflash_config.block_size. ' f'Got block_size={max_query_length}, query_length={query_length}.') - mask_token_id = dflash_config.get('mask_token_id') + mask_token_id = dflash_config.get('mask_token_id', getattr(draft_hf_config, 'mask_token_id', None)) if mask_token_id is None: raise ValueError('DFlash checkpoint requires dflash_config.mask_token_id.') diff --git a/tests/pytorch/spec_decode/test_dflash_utils.py b/tests/pytorch/spec_decode/test_dflash_utils.py index 251487a3d7..595bea6130 100644 --- a/tests/pytorch/spec_decode/test_dflash_utils.py +++ b/tests/pytorch/spec_decode/test_dflash_utils.py @@ -71,6 +71,21 @@ def test_parse_dflash_config_valid(): assert mask_token_id == 32001 +def test_parse_dflash_config_top_level_checkpoint_layout(): + """Original z-lab DFlash checkpoints (e.g. Qwen3-4B-DFlash-b16) keep + block_size at the top level of config.json and nest only mask_token_id and + target_layer_ids inside dflash_config.""" + config = _draft_config(block_size=4, + dflash_config=dict( + mask_token_id=32001, + target_layer_ids=[1, 5, 9, 13], + )) + target_layer_ids, mask_token_id = _parse_dflash(config, num_speculative_tokens=3) + + assert target_layer_ids == (1, 5, 9, 13) + assert mask_token_id == 32001 + + def test_specdecode_config_stores_resolved_dflash_fields_directly(): target_layer_ids, mask_token_id = _parse_dflash(_draft_config(), num_speculative_tokens=3) cfg = SpecDecodeConfig(model='draft-model', From 5794d5a38fab13d698121db5fd71d2ef85bd4f46 Mon Sep 17 00:00:00 2001 From: harshal-96 Date: Sat, 19 Sep 2026 13:23:38 +0530 Subject: [PATCH 2/3] remove build_target_layer_ids fallback; require checkpoint target_layer_ids The evenly spaced fallback cannot be correct: the tap layers are a training-time choice baked into the checkpoint (the fc input width and which target layers the draft was trained against), so any load-time guess reads the wrong features. For 14 of the 20 published DFlash checkpoints the guess differs from the trained layers, and when the counts collide the load succeeds silently with wrong layers (z-lab/dflash#156). Every published DFlash checkpoint sets dflash_config.target_layer_ids, so a missing key now fails loudly. Matches the same removal in the reference repo (z-lab/dflash#157) and sglang (sgl-project/sglang#37476). Signed-off-by: harshal-96 --- lmdeploy/pytorch/spec_decode/dflash_utils.py | 31 ++++--------------- .../pytorch/spec_decode/test_dflash_utils.py | 20 ++++++------ 2 files changed, 17 insertions(+), 34 deletions(-) diff --git a/lmdeploy/pytorch/spec_decode/dflash_utils.py b/lmdeploy/pytorch/spec_decode/dflash_utils.py index e7deeb9a5c..435f3607a5 100644 --- a/lmdeploy/pytorch/spec_decode/dflash_utils.py +++ b/lmdeploy/pytorch/spec_decode/dflash_utils.py @@ -32,29 +32,6 @@ def _validate_target_layer_id_order(layer_ids: tuple[int, ...], field_name: str) prev_layer_id = layer_id -def build_target_layer_ids(num_target_layers: int, num_draft_layers: int) -> tuple[int, ...]: - """Select evenly spaced DFlash target layer ids. - - DFlash consumes hidden features sampled from the target model. This fallback matches the SGLang/vLLM convention used - when checkpoint metadata does not explicitly list target layers. - """ - if num_target_layers < 1: - raise ValueError(f'DFlash num_target_layers must be positive, got {num_target_layers!r}.') - if num_draft_layers < 1: - raise ValueError(f'DFlash num_hidden_layers must be positive, got {num_draft_layers!r}.') - if num_draft_layers == 1: - return (num_target_layers // 2,) - - start = 1 - end = num_target_layers - 3 - if end < start: - raise ValueError(f'DFlash target layer fallback requires at least 4 target layers, got {num_target_layers}.') - span = end - start - layer_ids = tuple(int(round(start + i * span / (num_draft_layers - 1))) for i in range(num_draft_layers)) - _validate_target_layer_id_order(layer_ids, 'DFlash fallback target_layer_ids') - return layer_ids - - def _normalize_sliding_window(sliding_window: Any) -> int | None: """Normalize the no-window values used by HF and ``ModelConfig``.""" if sliding_window in (None, 0, -1): @@ -130,7 +107,6 @@ def _validate_dflash_v1_supported(draft_hf_config: Any, dflash_config: Any) -> N def parse_dflash_config(draft_hf_config: Any, num_speculative_tokens: int, target_num_layers: int) -> tuple[tuple[int, ...], int]: """Return resolved ``(target_layer_ids, mask_token_id)`` metadata.""" - num_hidden_layers = draft_hf_config.num_hidden_layers num_target_layers = draft_hf_config.num_target_layers dflash_config = draft_hf_config.dflash_config _validate_dflash_v1_supported(draft_hf_config, dflash_config) @@ -158,7 +134,12 @@ def parse_dflash_config(draft_hf_config: Any, num_speculative_tokens: int, target_layer_ids = _parse_layer_ids(dflash_config.get('target_layer_ids')) if target_layer_ids is None: - target_layer_ids = build_target_layer_ids(num_target_layers, num_hidden_layers) + # The tap layers are a training-time choice baked into the checkpoint + # (fc input width and which target layers the draft was trained on), + # so no load-time guess can be correct; see z-lab/dflash#156. Every + # published DFlash checkpoint sets this key. + raise ValueError('DFlash checkpoint requires dflash_config.target_layer_ids; ' + 'a guessed layer list can silently read the wrong target layers.') for pos, layer_id in enumerate(target_layer_ids): if layer_id < 0 or layer_id >= num_target_layers: raise ValueError('DFlash target_layer_ids contains an out-of-range value: ' diff --git a/tests/pytorch/spec_decode/test_dflash_utils.py b/tests/pytorch/spec_decode/test_dflash_utils.py index 595bea6130..9b966fc153 100644 --- a/tests/pytorch/spec_decode/test_dflash_utils.py +++ b/tests/pytorch/spec_decode/test_dflash_utils.py @@ -23,7 +23,6 @@ _resolve_dflash_layer_attention, ) from lmdeploy.pytorch.spec_decode.dflash_utils import ( - build_target_layer_ids, parse_dflash_config, validate_dflash_cache_config, validate_dflash_dist_config, @@ -186,29 +185,32 @@ def test_parse_dflash_config_requires_mask_token_id(): @pytest.mark.parametrize('num_speculative_tokens', [3, 7, 15]) def test_parse_dflash_config_allows_runtime_query_up_to_checkpoint_block_size(num_speculative_tokens): - draft_config = _draft_config(dflash_config=dict(block_size=16, mask_token_id=32001)) + draft_config = _draft_config( + dflash_config=dict(block_size=16, mask_token_id=32001, target_layer_ids=[1, 5, 9, 13])) target_layer_ids, mask_token_id = _parse_dflash(draft_config, num_speculative_tokens=num_speculative_tokens) - assert target_layer_ids == build_target_layer_ids(16, 4) + assert target_layer_ids == (1, 5, 9, 13) assert mask_token_id == 32001 def test_parse_dflash_config_rejects_query_above_checkpoint_block_size(): - draft_config = _draft_config(dflash_config=dict(block_size=16, mask_token_id=32001)) + draft_config = _draft_config( + dflash_config=dict(block_size=16, mask_token_id=32001, target_layer_ids=[1, 5, 9, 13])) with pytest.raises(ValueError, match='must not exceed.*block_size'): _parse_dflash(draft_config, num_speculative_tokens=16) -def test_parse_dflash_config_resolves_target_layers_from_num_target_layers(): +def test_parse_dflash_config_requires_target_layer_ids(): + """The tap layers are a training-time choice baked into the checkpoint, so + a missing key must fail loudly instead of falling back to a guessed layer + list that can silently read the wrong target layers (z-lab/dflash#156).""" draft_config = _draft_config(dflash_config=dict(block_size=4, mask_token_id=32001)) - target_layer_ids, mask_token_id = _parse_dflash(draft_config, num_speculative_tokens=3) - - assert target_layer_ids == build_target_layer_ids(16, 4) - assert mask_token_id == 32001 + with pytest.raises(ValueError, match='target_layer_ids'): + _parse_dflash(draft_config, num_speculative_tokens=3) def test_parse_dflash_config_rejects_non_increasing_target_layers(): From 57883c2c9074129962a48f4bf814a14efec46f77 Mon Sep 17 00:00:00 2001 From: harshal-96 Date: Mon, 21 Sep 2026 00:06:52 +0530 Subject: [PATCH 3/3] address review: mixed-dtype tests, top-level mask_token_id test, clearer block_size error - Unit-test the two draft-dtype ingestion boundaries (shared target embeddings and projected target hidden states) with a float16 source and a bfloat16 draft; both fail on the uncast code. - Cover the top-level mask_token_id fallback, matching the block_size fallback coverage. - Mention both supported block_size locations in the query-length error. Signed-off-by: harshal-96 --- lmdeploy/pytorch/spec_decode/dflash_utils.py | 4 +- .../pytorch/spec_decode/test_dflash_utils.py | 52 +++++++++++++++++++ 2 files changed, 54 insertions(+), 2 deletions(-) diff --git a/lmdeploy/pytorch/spec_decode/dflash_utils.py b/lmdeploy/pytorch/spec_decode/dflash_utils.py index 435f3607a5..931725a233 100644 --- a/lmdeploy/pytorch/spec_decode/dflash_utils.py +++ b/lmdeploy/pytorch/spec_decode/dflash_utils.py @@ -124,8 +124,8 @@ def parse_dflash_config(draft_hf_config: Any, num_speculative_tokens: int, raise ValueError('DFlash checkpoint requires block_size in dflash_config or at the top level of its config.') query_length = num_speculative_tokens + 1 if query_length > max_query_length: - raise ValueError('DFlash query length (1 + speculative_num_draft_tokens) must not exceed checkpoint ' - 'dflash_config.block_size. ' + raise ValueError('DFlash query length (1 + speculative_num_draft_tokens) must not exceed the checkpoint ' + 'block_size (from dflash_config or the top level of its config). ' f'Got block_size={max_query_length}, query_length={query_length}.') mask_token_id = dflash_config.get('mask_token_id', getattr(draft_hf_config, 'mask_token_id', None)) diff --git a/tests/pytorch/spec_decode/test_dflash_utils.py b/tests/pytorch/spec_decode/test_dflash_utils.py index 9b966fc153..a87314efca 100644 --- a/tests/pytorch/spec_decode/test_dflash_utils.py +++ b/tests/pytorch/spec_decode/test_dflash_utils.py @@ -85,6 +85,18 @@ def test_parse_dflash_config_top_level_checkpoint_layout(): assert mask_token_id == 32001 +def test_parse_dflash_config_top_level_mask_token_id(): + """mask_token_id resolves from the top level of the draft config when + dflash_config does not carry it, same fallback order as block_size.""" + config = _draft_config(block_size=4, + mask_token_id=32001, + dflash_config=dict(target_layer_ids=[1, 5, 9, 13], )) + target_layer_ids, mask_token_id = _parse_dflash(config, num_speculative_tokens=3) + + assert target_layer_ids == (1, 5, 9, 13) + assert mask_token_id == 32001 + + def test_specdecode_config_stores_resolved_dflash_fields_directly(): target_layer_ids, mask_token_id = _parse_dflash(_draft_config(), num_speculative_tokens=3) cfg = SpecDecodeConfig(model='draft-model', @@ -904,3 +916,43 @@ def _materialize_context(context_inputs, target_hidden, cache_engine): assert captured['context_inputs'].input_ids.tolist() == [[10, 11, 12, 20, 21]] assert captured['context_inputs'].seq_length.tolist() == [3, 2] assert captured['target_hidden'].tolist() == extra_inputs.target_hidden_states.tolist() + + +class TestDFlashDraftModelMixedDtype: + """The draft ingests two target-side tensors (shared embeddings and aux + hidden states) that may arrive in the target dtype, e.g. a float16 AWQ + target with a bfloat16 draft checkpoint. + + Both ingestion boundaries must cast to the draft dtype; running everything in float16 instead is not an option + because Qwen3 hidden state outliers overflow float16 during feature fusion (acceptance collapses to near zero). + """ + + def _make_model(self): + from lmdeploy.pytorch.models.qwen3_dflash import DFlashDraftModel + + # bypass __init__: it requires a full engine build context, while + # the dtype contract lives entirely in these two methods + model = DFlashDraftModel.__new__(DFlashDraftModel) + torch.nn.Module.__init__(model) + model.dtype = torch.bfloat16 + return model + + def test_embed_input_ids_casts_to_draft_dtype(self): + model = self._make_model() + model.embed_tokens = torch.nn.Embedding(8, 4, dtype=torch.float16) + model.has_separate_mask_embedding = False + model.mask_token_id = None + + embeds = model.embed_input_ids(torch.tensor([[0, 1, 2]])) + + assert embeds.dtype == torch.bfloat16 + + def test_project_target_hidden_casts_to_draft_dtype(self): + model = self._make_model() + model.fc = torch.nn.Linear(8, 4, bias=False, dtype=torch.bfloat16) + model.hidden_norm = torch.nn.Identity() + model.num_context_features = 2 + + out = model.project_target_hidden(torch.randn(3, 8, dtype=torch.float16)) + + assert out.dtype == torch.bfloat16