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..931725a233 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) @@ -139,20 +115,31 @@ 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. ' + 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') + 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.') 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 251487a3d7..a87314efca 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, @@ -71,6 +70,33 @@ 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_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', @@ -171,29 +197,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(): @@ -887,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