Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
7 changes: 5 additions & 2 deletions lmdeploy/pytorch/models/qwen3_dflash.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Comment on lines +355 to +357

def project_target_hidden(self, target_hidden: torch.Tensor) -> torch.Tensor:
"""Project concatenated target-layer hidden states into draft hidden
Expand All @@ -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,
Expand Down
45 changes: 16 additions & 29 deletions lmdeploy/pytorch/spec_decode/dflash_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down Expand Up @@ -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)
Expand All @@ -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: '
Expand Down
87 changes: 78 additions & 9 deletions tests/pytorch/spec_decode/test_dflash_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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',
Expand Down Expand Up @@ -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():
Expand Down Expand Up @@ -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
Loading