diff --git a/python/sglang/srt/arg_groups/speculative_hook.py b/python/sglang/srt/arg_groups/speculative_hook.py index 77378509b5ca..3e970263f60e 100644 --- a/python/sglang/srt/arg_groups/speculative_hook.py +++ b/python/sglang/srt/arg_groups/speculative_hook.py @@ -10,21 +10,50 @@ logger = logging.getLogger(__name__) +_DSPARK_DRAFT_ARCHITECTURES = frozenset({"DSparkDraftModel", "Qwen3DSparkModel"}) + + +def _load_speculative_draft_config( + model_path: str, + *, + trust_remote_code: bool, + revision: Optional[str] = None, + model_override_args: Optional[dict] = None, + **kwargs, +): + from sglang.srt.utils.hf_transformers_utils import ( + get_config_allow_missing_model_type, + ) + + return get_config_allow_missing_model_type( + model_path, + trust_remote_code=trust_remote_code, + revision=revision, + model_override_args=model_override_args, + **kwargs, + ) + + +def _is_dspark_draft_config(config) -> bool: + draft_archs = getattr(config, "architectures", None) or [] + return any(arch in _DSPARK_DRAFT_ARCHITECTURES for arch in draft_archs) + def _resolve_speculative_algorithm_alias( speculative_algorithm: Optional[str], speculative_draft_model_path: Optional[str], trust_remote_code: bool = False, - kwargs: Optional[dict] = {}, + kwargs: Optional[dict] = None, ) -> Optional[str]: """Resolve CLI speculative algorithm; NEXTN/EAGLE may become FROZEN_KV_MTP for Gemma4 assistant drafts.""" + kwargs = kwargs or {} is_gemma4_draft = False if speculative_draft_model_path: - from sglang.srt.utils.hf_transformers_utils import get_config - - cfg = get_config( - speculative_draft_model_path, trust_remote_code=trust_remote_code, **kwargs + cfg = _load_speculative_draft_config( + speculative_draft_model_path, + trust_remote_code=trust_remote_code, + **kwargs, ) draft_archs = getattr(cfg, "architectures", None) or [] is_gemma4_draft = any( @@ -147,6 +176,26 @@ def _handle_dflash(server_args: ServerArgs) -> None: "DFLASH speculative decoding requires setting --speculative-draft-model-path." ) + draft_hf_config = None + draft_config_error = None + model_override_args = json.loads(server_args.json_model_override_args) + try: + draft_hf_config = _load_speculative_draft_config( + server_args.speculative_draft_model_path, + trust_remote_code=server_args.trust_remote_code, + revision=server_args.speculative_draft_model_revision, + model_override_args=model_override_args, + ) + except Exception as e: + draft_config_error = e + else: + if _is_dspark_draft_config(draft_hf_config): + raise ValueError( + "The draft checkpoint architecture is DSparkDraftModel, but " + "speculative_algorithm=DFLASH was requested. Use " + "--speculative-algorithm DSPARK for DSpark draft checkpoints." + ) + # DFLASH does not use EAGLE-style `num_steps`/`topk`, but those fields still # affect generic scheduler/KV-cache accounting (buffer sizing, KV freeing, # RoPE reservation). Force them to 1 to avoid surprising memory behavior. @@ -194,17 +243,10 @@ def _handle_dflash(server_args: ServerArgs) -> None: parse_dflash_draft_config, ) - model_override_args = json.loads(server_args.json_model_override_args) inferred_block_size = None try: - from sglang.srt.utils.hf_transformers_utils import get_config - - draft_hf_config = get_config( - server_args.speculative_draft_model_path, - trust_remote_code=server_args.trust_remote_code, - revision=server_args.speculative_draft_model_revision, - model_override_args=model_override_args, - ) + if draft_config_error is not None: + raise draft_config_error inferred_block_size = parse_dflash_draft_config( draft_hf_config=draft_hf_config ).resolve_block_size(default=None) @@ -340,9 +382,7 @@ def _handle_dspark(server_args: ServerArgs) -> None: model_override_args = json.loads(server_args.json_model_override_args) try: - from sglang.srt.utils.hf_transformers_utils import get_config - - draft_hf_config = get_config( + draft_hf_config = _load_speculative_draft_config( server_args.speculative_draft_model_path, trust_remote_code=server_args.trust_remote_code, revision=server_args.speculative_draft_model_revision, @@ -373,9 +413,7 @@ def _handle_dspark(server_args: ServerArgs) -> None: model_override_args = json.loads(server_args.json_model_override_args) config_gamma: Optional[int] = None try: - from sglang.srt.utils.hf_transformers_utils import get_config - - draft_hf_config = get_config( + draft_hf_config = _load_speculative_draft_config( server_args.speculative_draft_model_path, trust_remote_code=server_args.trust_remote_code, revision=server_args.speculative_draft_model_revision, @@ -396,8 +434,8 @@ def _handle_dspark(server_args: ServerArgs) -> None: if int(server_args.speculative_num_draft_tokens) != config_verify_window: raise ValueError( "DSpark speculative_num_draft_tokens must equal the draft " - "checkpoint block_size + 1 " - f"(= {config_verify_window} for block_size={config_gamma}), " + "checkpoint gamma + 1 " + f"(= {config_verify_window} for gamma={config_gamma}), " "but got speculative_num_draft_tokens=" f"{server_args.speculative_num_draft_tokens}." ) diff --git a/python/sglang/srt/configs/model_config.py b/python/sglang/srt/configs/model_config.py index 0a29c6fbc332..7dba49093aef 100644 --- a/python/sglang/srt/configs/model_config.py +++ b/python/sglang/srt/configs/model_config.py @@ -30,6 +30,7 @@ from sglang.srt.utils import is_hip, is_sm100_supported, retry from sglang.srt.utils.hf_transformers_utils import ( get_config, + get_config_allow_missing_model_type, get_context_length, get_generation_config, get_hf_text_config, @@ -100,6 +101,69 @@ def _hf_attr(config, name): return getattr(config, name, None) +def _restore_glm_moe_dsa_head_dims_from_raw_config( + hf_config: PretrainedConfig, + raw_config_dict: Optional[dict], + model_override_args: dict, + model_path: str, + trust_remote_code: bool, + revision: Optional[str], + kwargs: dict, +) -> None: + if _hf_arch(hf_config) != "GlmMoeDsaForCausalLM": + return + + if raw_config_dict is None: + try: + raw_config_dict, _ = PretrainedConfig.get_config_dict( + model_path, + trust_remote_code=trust_remote_code, + revision=revision, + **kwargs, + ) + except Exception as exc: + logger.warning( + "Failed to restore GLM DSA raw head dimensions from config.json: %s", + exc, + ) + return + + for name in ( + "qk_nope_head_dim", + "qk_rope_head_dim", + "qk_head_dim", + "v_head_dim", + ): + value = model_override_args.get(name, raw_config_dict.get(name)) + if value is not None: + setattr(hf_config, name, value) + + +def _normalize_nested_transformer_config(hf_config: PretrainedConfig) -> None: + """Materialize nested speculator transformer config fields on the HF config.""" + transformer_cfg = _hf_attr(hf_config, "transformer_layer_config") + if transformer_cfg is None: + return + + if isinstance(transformer_cfg, dict): + items = transformer_cfg.items() + else: + to_dict = getattr(transformer_cfg, "to_dict", None) + items = to_dict().items() if callable(to_dict) else vars(transformer_cfg).items() + + for key, value in items: + if key.startswith("_"): + continue + if _hf_attr(hf_config, key) is None: + setattr(hf_config, key, value) + + aux_layer_ids = _hf_attr(hf_config, "aux_hidden_state_layer_ids") + if _hf_attr(hf_config, "num_target_layers") is None and aux_layer_ids is not None: + parsed = [int(x) for x in aux_layer_ids] + if parsed: + setattr(hf_config, "num_target_layers", max(parsed) + 1) + + def is_deepseek_dsa(config) -> bool: return ( _hf_arch(config) @@ -271,15 +335,29 @@ def __init__( # get_config() is cached. ModelConfig mutates hf_config for draft-model # remapping and architecture-specific normalization, so each instance # must own an isolated copy. - self.hf_config = copy.deepcopy( - get_config( - self.model_path, - trust_remote_code=trust_remote_code, - revision=revision, - model_override_args=self.model_override_args, - model_config_parser=model_config_parser, - **kwargs, - ) + raw_config_dict = None + config_loader = ( + get_config_allow_missing_model_type if is_draft_model else get_config + ) + hf_config = config_loader( + self.model_path, + trust_remote_code=trust_remote_code, + revision=revision, + model_override_args=self.model_override_args, + model_config_parser=model_config_parser, + **kwargs, + ) + self.hf_config = copy.deepcopy(hf_config) + if is_draft_model: + _normalize_nested_transformer_config(self.hf_config) + _restore_glm_moe_dsa_head_dims_from_raw_config( + self.hf_config, + raw_config_dict, + self.model_override_args, + self.model_path, + trust_remote_code, + revision, + kwargs, ) self.hf_text_config = get_hf_text_config(self.hf_config) self.hf_generation_config = get_generation_config( diff --git a/python/sglang/srt/debug_utils/pr_fix_toggle.py b/python/sglang/srt/debug_utils/pr_fix_toggle.py index aecdca3e1822..8e8b1be5e142 100644 --- a/python/sglang/srt/debug_utils/pr_fix_toggle.py +++ b/python/sglang/srt/debug_utils/pr_fix_toggle.py @@ -73,12 +73,14 @@ - target: sglang.srt.mem_cache.common.get_req_to_token_extra_context_len edits: - match: | - if ( - server_args.speculative_algorithm is not None - and server_args.page_size > 1 - and (server_args.speculative_eagle_topk or 1) > 1 - ): - extra = max(extra, get_alloc_reserve_per_decode(server_args)) + if server_args.speculative_algorithm is not None: + from sglang.srt.speculative.spec_info import SpeculativeAlgorithm + spec_algo = SpeculativeAlgorithm.from_string(server_args.speculative_algorithm) + if ( + server_args.page_size > 1 + and (server_args.speculative_eagle_topk or 1) > 1 + ) or spec_algo.is_dflash_or_dspark(): + extra = max(extra, get_alloc_reserve_per_decode(server_args)) replacement: "" """ diff --git a/python/sglang/srt/layers/attention/dsa/dsa_indexer.py b/python/sglang/srt/layers/attention/dsa/dsa_indexer.py index b03deeb65829..c45677089f94 100644 --- a/python/sglang/srt/layers/attention/dsa/dsa_indexer.py +++ b/python/sglang/srt/layers/attention/dsa/dsa_indexer.py @@ -1556,7 +1556,12 @@ def _store_index_k_cache( layer_id=layer_id ) kv_cache = buf.view(-1, page_size, 132).view(fp8_dtype) - out_loc = forward_batch.out_cache_loc + out_loc = out_cache_loc + assert out_loc is not None, "DSA indexer cache store requires out_cache_loc" + assert out_loc.numel() == key.shape[0], ( + f"DSA indexer cache store got out_cache_loc tokens={out_loc.numel()} " + f"but key tokens={key.shape[0]}" + ) if not out_loc.is_contiguous(): out_loc = out_loc.contiguous() indexer_k_quant_and_cache( diff --git a/python/sglang/srt/layers/attention/dsa_backend.py b/python/sglang/srt/layers/attention/dsa_backend.py index 13f78f253238..cf98f1de17f2 100644 --- a/python/sglang/srt/layers/attention/dsa_backend.py +++ b/python/sglang/srt/layers/attention/dsa_backend.py @@ -50,6 +50,7 @@ seqlens_expand_triton, ) from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode +from sglang.srt.speculative.spec_info import SpeculativeAlgorithm from sglang.srt.utils import ( get_bool_env_var, is_cuda, @@ -407,6 +408,20 @@ def __init__( self.speculative_num_draft_tokens = ( model_runner.server_args.speculative_num_draft_tokens ) + if ( + self.speculative_num_draft_tokens is not None + and model_runner.is_draft_worker + ): + spec_algo = SpeculativeAlgorithm.from_string( + model_runner.server_args.speculative_algorithm + ) + if spec_algo.is_dspark(): + self.speculative_num_draft_tokens = ( + spec_algo.get_num_tokens_per_bs_for_target_verify( + int(self.speculative_num_draft_tokens), + is_draft_worker=True, + ) + ) self.speculative_step_id = speculative_step_id self.device_capability = torch.cuda.get_device_capability() @@ -647,7 +662,14 @@ def init_forward_metadata(self, forward_batch: ForwardBatch): cache_seqlens_int32 = (forward_batch.seq_lens + draft_token_num).to(torch.int32) cu_seqlens_k = compute_cu_seqlens(cache_seqlens_int32) - if forward_batch.seq_lens_cpu is not None: + if forward_batch.forward_mode.is_target_verify(): + if forward_batch.seq_lens_cpu is not None: + max_seqlen_k = int( + forward_batch.seq_lens_cpu.max().item() + draft_token_num + ) + else: + max_seqlen_k = int(cache_seqlens_int32.max().item()) + elif forward_batch.seq_lens_cpu is not None: max_seqlen_k = int( forward_batch.seq_lens_cpu.max().item() + draft_token_num ) @@ -657,6 +679,12 @@ def init_forward_metadata(self, forward_batch: ForwardBatch): # eager (e.g. over-capture-bs) fallback needs a length here. max_seqlen_k = int(forward_batch.seq_lens.max().item()) + draft_token_num # [b, max_seqlen_k] + row_width = self.req_to_token_pool.req_to_token.shape[1] + assert max_seqlen_k <= row_width, ( + f"DSA metadata max_seqlen_k={max_seqlen_k} exceeds req_to_token " + f"row width={row_width}; mode={forward_batch.forward_mode}, " + f"draft_token_num={draft_token_num}." + ) page_table = self.req_to_token_pool.req_to_token[ forward_batch.req_pool_indices, :max_seqlen_k ] @@ -705,8 +733,25 @@ def init_forward_metadata(self, forward_batch: ForwardBatch): extend_seq_lens_cpu = [self.speculative_num_draft_tokens] * batch_size forward_batch.extend_seq_lens_cpu = extend_seq_lens_cpu + layout = getattr(forward_batch.spec_info, "ragged_verify_layout", None) + if layout is None: + expected_out_tokens = batch_size * self.speculative_num_draft_tokens + assert ( + forward_batch.out_cache_loc is not None + and forward_batch.out_cache_loc.numel() >= expected_out_tokens + ), ( + f"DSA target verify expects at least {expected_out_tokens} " + f"out_cache_loc entries, got " + f"{None if forward_batch.out_cache_loc is None else forward_batch.out_cache_loc.numel()}." + ) + seqlens_expanded = seqlens_expand_triton( - torch.tensor(extend_seq_lens_cpu, dtype=torch.int32, device=device), + torch.full( + (batch_size,), + self.speculative_num_draft_tokens, + dtype=torch.int32, + device=device, + ), cache_seqlens_int32, self.speculative_num_draft_tokens * batch_size, self.speculative_num_draft_tokens, diff --git a/python/sglang/srt/layers/attention/triton_backend.py b/python/sglang/srt/layers/attention/triton_backend.py index fa765787a814..9a65745bd940 100644 --- a/python/sglang/srt/layers/attention/triton_backend.py +++ b/python/sglang/srt/layers/attention/triton_backend.py @@ -28,6 +28,7 @@ from sglang.srt.model_executor.cuda_graph_config import cuda_graph_fully_disabled from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode from sglang.srt.runtime_context import get_parallel +from sglang.srt.speculative.spec_info import SpeculativeAlgorithm from sglang.srt.speculative.spec_utils import ( draft_kv_indices_buffer_width, draft_kv_indices_used_len, @@ -148,6 +149,16 @@ def __init__( self.token_to_kv_pool_allocator = model_runner.token_to_kv_pool_allocator self.use_sliding_window_kv_pool = isinstance(self.token_to_kv_pool, SWAKVPool) self.num_draft_tokens = model_runner.server_args.speculative_num_draft_tokens + if self.num_draft_tokens is not None and model_runner.is_draft_worker: + spec_algo = SpeculativeAlgorithm.from_string( + model_runner.server_args.speculative_algorithm + ) + if spec_algo.is_dspark(): + self.num_draft_tokens = ( + spec_algo.get_num_tokens_per_bs_for_target_verify( + int(self.num_draft_tokens), is_draft_worker=True + ) + ) self.speculative_num_steps = model_runner.server_args.speculative_num_steps self.topk = model_runner.server_args.speculative_eagle_topk or 0 # Split-KV verify matches extend_attention_fwd only when the EAGLE tree diff --git a/python/sglang/srt/mem_cache/common.py b/python/sglang/srt/mem_cache/common.py index e50fcfc14bde..9b5f5c9294d9 100644 --- a/python/sglang/srt/mem_cache/common.py +++ b/python/sglang/srt/mem_cache/common.py @@ -257,18 +257,21 @@ def get_req_to_token_extra_context_len(server_args: ServerArgs) -> int: """req_to_token row headroom beyond the model context length. Sized to hold the decode over-allocation (kv_committed_len + - get_alloc_reserve_per_decode). The spec v2 page>1 topk>1 holey draft footprint - can outgrow the default num_draft_tokens headroom (PR #26972). + get_alloc_reserve_per_decode). Spec-v2 tree drafts and DFlash/DSpark draft + KV can outgrow the default num_draft_tokens headroom, especially when + overlap keeps a double buffer. """ # FIXME(lsyin): this is the temporary fix for the context length issue when # using speculative decoding extra = 4 + (server_args.max_speculative_num_draft_tokens or 0) - if ( - server_args.speculative_algorithm is not None - and server_args.page_size > 1 - and (server_args.speculative_eagle_topk or 1) > 1 - ): - extra = max(extra, get_alloc_reserve_per_decode(server_args)) + if server_args.speculative_algorithm is not None: + from sglang.srt.speculative.spec_info import SpeculativeAlgorithm + spec_algo = SpeculativeAlgorithm.from_string(server_args.speculative_algorithm) + if ( + server_args.page_size > 1 + and (server_args.speculative_eagle_topk or 1) > 1 + ) or spec_algo.is_dflash_or_dspark(): + extra = max(extra, get_alloc_reserve_per_decode(server_args)) return extra diff --git a/python/sglang/srt/models/dspark.py b/python/sglang/srt/models/dspark.py index bbc66a8eac0c..f62946fc9cfd 100644 --- a/python/sglang/srt/models/dspark.py +++ b/python/sglang/srt/models/dspark.py @@ -1,7 +1,7 @@ from __future__ import annotations import logging -from typing import Callable, Iterable, Optional, Tuple +from typing import Any, Callable, Iterable, Optional, Tuple import torch from torch import nn @@ -22,6 +22,49 @@ StepSampler = Callable[[torch.Tensor, int], torch.Tensor] +def _cfg_get(config: Any, key: str, default: Any = None) -> Any: + if isinstance(config, dict): + return config.get(key, default) + return getattr(config, key, default) + + +def _cfg_set(config: Any, key: str, value: Any) -> None: + if isinstance(config, dict): + config[key] = value + else: + setattr(config, key, value) + + +def _cfg_items(config: Any): + if isinstance(config, dict): + return config.items() + to_dict = getattr(config, "to_dict", None) + if callable(to_dict): + return to_dict().items() + return vars(config).items() + + +def normalize_dspark_draft_config(config: Any) -> Any: + """Expose nested speculator transformer fields as normal draft config attrs.""" + transformer_cfg = _cfg_get(config, "transformer_layer_config", None) + if transformer_cfg is None: + return config + + for key, value in _cfg_items(transformer_cfg): + if key.startswith("_"): + continue + if _cfg_get(config, key, None) is None: + _cfg_set(config, key, value) + + aux_layer_ids = _cfg_get(config, "aux_hidden_state_layer_ids", None) + if _cfg_get(config, "num_target_layers", None) is None and aux_layer_ids is not None: + parsed = [int(x) for x in aux_layer_ids] + if parsed: + _cfg_set(config, "num_target_layers", max(parsed) + 1) + + return config + + def gather_and_crop_vocab( local_logits: torch.Tensor, lm_head: nn.Module ) -> torch.Tensor: @@ -361,6 +404,7 @@ def build_confidence_head(config) -> Optional[nn.Module]: class DSparkDraftMixin: def __init__(self, config, quant_config=None, prefix: str = "") -> None: + config = normalize_dspark_draft_config(config) super().__init__(config=config, quant_config=quant_config, prefix=prefix) dspark_config = parse_dspark_draft_config(draft_hf_config=config) if not dspark_config.require_markov(): @@ -507,4 +551,4 @@ class Qwen3DSparkModel(DSparkDraftModel): pass -EntryClass = [Qwen3DSparkModel] +EntryClass = [Qwen3DSparkModel, DSparkDraftModel] diff --git a/python/sglang/srt/speculative/dflash_info_v2.py b/python/sglang/srt/speculative/dflash_info_v2.py index d19644074840..36c0d6dd2eb5 100644 --- a/python/sglang/srt/speculative/dflash_info_v2.py +++ b/python/sglang/srt/speculative/dflash_info_v2.py @@ -178,6 +178,14 @@ def prepare_for_decode(self, batch: ScheduleBatch): self.max_top_k = max(max_top_k, 1) self.uniform_top_k_value = uniform_top_k_value if uniform_top_k else None + row_width = batch.req_to_token_pool.req_to_token.shape[1] + max_reserved_len = int(nxt_kv_lens_cpu_t.max().item()) + assert max_reserved_len <= row_width, ( + f"DFLASH/DSPARK draft over-allocation ({max_reserved_len}) exceeds " + f"req_to_token row width ({row_width}); widen the row to hold committed " + f"+ get_alloc_reserve_per_decode()." + ) + caller_stream = None if plan_stream is not None: caller_stream = torch.get_device_module(batch.device).current_stream() diff --git a/python/sglang/srt/speculative/dflash_utils.py b/python/sglang/srt/speculative/dflash_utils.py index 1fb69ff56d86..db67cc3896a2 100644 --- a/python/sglang/srt/speculative/dflash_utils.py +++ b/python/sglang/srt/speculative/dflash_utils.py @@ -352,10 +352,16 @@ def _get_text_config(config: Any) -> Any: if config is None: return None if isinstance(config, dict): - return config.get("text_config", config) + text_config = config.get("text_config", None) + if text_config is not None: + return text_config + return config.get("transformer_layer_config", config) text_config = getattr(config, "text_config", None) if text_config is not None: return text_config + transformer_layer_config = getattr(config, "transformer_layer_config", None) + if transformer_layer_config is not None: + return transformer_layer_config get_text_config = getattr(config, "get_text_config", None) if callable(get_text_config): try: @@ -463,9 +469,16 @@ def parse_dflash_draft_config(*, draft_hf_config: Any) -> DFlashDraftConfig: field_name="DFLASH draft num_hidden_layers", min_value=1, ) + aux_layer_ids = _cfg_get(draft_hf_config, "aux_hidden_state_layer_ids", None) + inferred_num_target_layers = None + if aux_layer_ids is not None: + inferred_aux_layer_ids = [int(x) for x in aux_layer_ids] + if inferred_aux_layer_ids: + inferred_num_target_layers = max(inferred_aux_layer_ids) + 1 + raw_num_target_layers = dflash_cfg.get( "num_target_layers", - _cfg_get(draft_hf_config, "num_target_layers", None), + _cfg_get(draft_hf_config, "num_target_layers", inferred_num_target_layers), ) num_target_layers = _parse_optional_int( raw_num_target_layers, @@ -486,7 +499,7 @@ def parse_dflash_draft_config(*, draft_hf_config: Any) -> DFlashDraftConfig: layer_ids = dflash_cfg.get( "target_layer_ids", - _cfg_get(draft_hf_config, "target_layer_ids", None), + _cfg_get(draft_hf_config, "target_layer_ids", aux_layer_ids), ) parsed_target_layer_ids: Optional[List[int]] if layer_ids is None: diff --git a/python/sglang/srt/speculative/dflash_worker_v2.py b/python/sglang/srt/speculative/dflash_worker_v2.py index f63de41fd7ad..b4cbeac6c5e2 100644 --- a/python/sglang/srt/speculative/dflash_worker_v2.py +++ b/python/sglang/srt/speculative/dflash_worker_v2.py @@ -47,6 +47,21 @@ logger = logging.getLogger(__name__) + +def _copy_prefix_seq_lens_cpu( + dst: torch.Tensor, + prefix_lens: torch.Tensor, + seq_lens_cpu: Optional[torch.Tensor], +) -> int: + src = seq_lens_cpu + if src is None: + src = prefix_lens.detach().to(device=dst.device, dtype=dst.dtype) + elif src.dtype != dst.dtype: + src = src.to(dtype=dst.dtype) + dst.copy_(src) + return int(dst.sum().item()) + + _FusedKVMaterializeHelper = None @@ -1426,22 +1441,14 @@ def forward_batch_generation( draft_seq_lens = draft_prefix_lens draft_seq_lens_sum = int(seq_lens_cpu.sum().item()) else: - # Non-windowed path uses the shared overallocated mapping directly. - # Backend planning only needs a safe upper bound for the committed - # prefix lengths, not the full allocator reservation length. + # TARGET_VERIFY attention backends interpret seq_lens_cpu as the + # committed prefix and add the fixed draft width internally when + # planning metadata. Keep this mirror prefix-only; using allocator + # reservation or prefix + block double-counts the verify width. draft_seq_lens = prefix_lens - if batch.seq_lens_cpu is not None: - # Host bound = committed prefix + one verify block. - seq_lens_cpu.copy_(batch.seq_lens_cpu) - seq_lens_cpu.add_(block_size) - draft_seq_lens_sum = int(seq_lens_cpu.sum()) - elif draft_input.reserved_seq_lens_cpu is not None: - # GPU-only backend: reserved is a safe over-estimate. - seq_lens_cpu.copy_(draft_input.reserved_seq_lens_cpu) - draft_seq_lens_sum = int(draft_input.reserved_seq_lens_sum) - else: - seq_lens_cpu.copy_(prefix_lens.to("cpu", dtype=torch.int32)) - draft_seq_lens_sum = int(prefix_lens.sum().item()) + draft_seq_lens_sum = _copy_prefix_seq_lens_cpu( + seq_lens_cpu, prefix_lens, batch.seq_lens_cpu + ) forward_batch = ForwardBatch( forward_mode=ForwardMode.TARGET_VERIFY, @@ -1508,14 +1515,13 @@ def forward_batch_generation( ) seq_lens_cpu_backup = batch.seq_lens_cpu seq_lens_sum_backup = batch.seq_lens_sum - if seq_lens_cpu_backup is not None: - # Verify host bound = committed prefix + one verify block (matches draft). - verify_host_seq_lens = seq_lens_cpu_backup + block_size - batch.seq_lens_cpu = verify_host_seq_lens - batch.seq_lens_sum = int(verify_host_seq_lens.sum()) - elif draft_input.reserved_seq_lens_cpu is not None: - batch.seq_lens_cpu = draft_input.reserved_seq_lens_cpu - batch.seq_lens_sum = int(draft_input.reserved_seq_lens_sum) + base_seq_lens_cpu = ( + seq_lens_cpu_backup + if seq_lens_cpu_backup is not None + else batch.seq_lens.cpu() + ) + batch.seq_lens_cpu = base_seq_lens_cpu + batch.seq_lens_sum = int(base_seq_lens_cpu.sum()) verify_forward_batch, _ = verify_input.prepare_for_verify( batch, self.target_worker diff --git a/python/sglang/srt/speculative/draft_worker_common.py b/python/sglang/srt/speculative/draft_worker_common.py index b009436a642a..3eb282eadd11 100644 --- a/python/sglang/srt/speculative/draft_worker_common.py +++ b/python/sglang/srt/speculative/draft_worker_common.py @@ -19,7 +19,7 @@ ) from sglang.srt.speculative.dflash_info import DFlashVerifyInput from sglang.srt.speculative.dflash_info_v2 import DFlashDraftInputV2 -from sglang.srt.utils.hf_transformers_utils import get_config +from sglang.srt.utils.hf_transformers_utils import get_config_allow_missing_model_type logger = logging.getLogger(__name__) @@ -51,7 +51,7 @@ def _load_draft_hf_config(*, draft_server_args: ServerArgs) -> Optional[Any]: if not draft_model_path: return None model_override_args = json.loads(draft_server_args.json_model_override_args) - return get_config( + return get_config_allow_missing_model_type( draft_model_path, trust_remote_code=draft_server_args.trust_remote_code, revision=draft_server_args.speculative_draft_model_revision, diff --git a/python/sglang/srt/speculative/dspark_components/dspark_draft.py b/python/sglang/srt/speculative/dspark_components/dspark_draft.py index ebbd1ab5303f..24b1db443408 100644 --- a/python/sglang/srt/speculative/dspark_components/dspark_draft.py +++ b/python/sglang/srt/speculative/dspark_components/dspark_draft.py @@ -37,20 +37,34 @@ def __init__(self, *, model, gamma, max_bs, device, confidence_fn=None, out=None ) def __call__(self, hidden_states, input_ids): - bs = hidden_states.shape[0] // self.gamma - base_logits, confidence_tap = self.model.compute_base_logits(hidden_states) + draft_width = self.gamma + 1 + if hidden_states.shape[0] % draft_width != 0: + raise RuntimeError( + "DSpark folded draft sampler expects full blocks with " + f"anchor + gamma tokens, got {hidden_states.shape[0]} rows " + f"for gamma={self.gamma}." + ) + bs = hidden_states.shape[0] // draft_width + hidden_3d = hidden_states.view(bs, draft_width, -1) + ids_2d = input_ids.view(bs, draft_width) + anchor = ids_2d[:, 0] + # Slot 0 conditions the block. Slots 1..gamma are the draft tokens + # used by the verifier and confidence scheduler. + draft_hidden = hidden_3d[:, 1:, :].contiguous() + hidden_for_logits = draft_hidden.reshape(bs * self.gamma, -1) + + base_logits, confidence_tap = self.model.compute_base_logits(hidden_for_logits) base_logits = base_logits.view(bs, self.gamma, -1) - anchor = input_ids.view(bs, self.gamma)[:, 0] draft_tokens, _ = self.markov_head.sample_block( base_logits, first_prev_tokens=anchor, - hidden_states=hidden_states.view(bs, self.gamma, -1), + hidden_states=draft_hidden, sampler=greedy_step_sampler, ) self.out[: draft_tokens.numel()].copy_(draft_tokens.reshape(-1)) if self.confidence_out is not None: confidence = self.confidence_fn( - draft_hidden=hidden_states.view(bs, self.gamma, -1), + draft_hidden=draft_hidden, anchor_tokens=anchor, draft_tokens=draft_tokens, confidence_tap=confidence_tap, diff --git a/python/sglang/srt/speculative/dspark_components/dspark_draft_proposer.py b/python/sglang/srt/speculative/dspark_components/dspark_draft_proposer.py index 095c18756a53..a8fdfe776ff0 100644 --- a/python/sglang/srt/speculative/dspark_components/dspark_draft_proposer.py +++ b/python/sglang/srt/speculative/dspark_components/dspark_draft_proposer.py @@ -156,16 +156,20 @@ def _run_forward( embed_module, ) -> DraftForwardResult: gamma = self.gamma + draft_width = gamma + 1 prefix_lens = batch.seq_lens positions_2d = verify_window.positions_2d verify_cache_loc_2d = verify_window.verify_cache_loc_2d draft_block_ids = torch.full( - (bs, gamma), int(self._mask_token_id), dtype=torch.long, device=device + (bs, draft_width), + int(self._mask_token_id), + dtype=torch.long, + device=device, ) draft_block_ids[:, 0].copy_(draft_input.bonus_tokens.view(-1)) - draft_positions = positions_2d[:, :gamma].reshape(-1) - draft_cache_loc = verify_cache_loc_2d[:, :gamma].reshape(-1) + draft_positions = positions_2d[:, :draft_width].reshape(-1) + draft_cache_loc = verify_cache_loc_2d[:, :draft_width].reshape(-1) draft_owns_embed = hasattr(self.draft_model, "forward_embed") draft_input_embeds: Optional[torch.Tensor] = None @@ -174,13 +178,20 @@ def _run_forward( draft_input_embeds = noise_embedding.view(-1, noise_embedding.shape[-1]) if batch.seq_lens_cpu is not None: - draft_seq_lens_cpu = batch.seq_lens_cpu + gamma - draft_seq_lens_sum = int(draft_seq_lens_cpu.sum()) - elif draft_input.reserved_seq_lens_cpu is not None: - draft_seq_lens_cpu = draft_input.reserved_seq_lens_cpu - draft_seq_lens_sum = int(draft_input.reserved_seq_lens_sum) + # TARGET_VERIFY attention backends interpret seq_lens_cpu as the + # committed prefix and add the fixed draft width internally when + # sizing cache/page-table metadata. Passing prefix + gamma here + # makes DSA's indexer mirror longer than the real committed row. + draft_seq_lens_cpu = batch.seq_lens_cpu + seq_lens_sum = getattr(batch, "seq_lens_sum", None) + draft_seq_lens_sum = ( + int(seq_lens_sum) + if seq_lens_sum is not None + else int(batch.seq_lens_cpu.sum().item()) + ) else: - raise RuntimeError("DSpark decode expected batch.seq_lens_cpu, got None") + draft_seq_lens_cpu = prefix_lens.detach().to("cpu") + draft_seq_lens_sum = int(draft_seq_lens_cpu.sum().item()) draft_forward_batch = ForwardBatch( forward_mode=ForwardMode.TARGET_VERIFY, @@ -204,7 +215,11 @@ def _run_forward( raw_hidden = logits_output.hidden_states if raw_hidden is None: raise RuntimeError("DSpark draft model returned no hidden states.") - draft_hidden_3d = raw_hidden.view(bs, gamma, -1) + # DSpark/DFlash block slot 0 is the anchor token. The draft transformer + # must see it, but only slots 1..gamma are emitted as speculative tokens. + raw_hidden_3d = raw_hidden.view(bs, draft_width, -1) + draft_hidden_3d = raw_hidden_3d[:, 1:, :].contiguous() + raw_hidden = draft_hidden_3d.reshape(bs * gamma, -1) return DraftForwardResult( draft_block_ids=draft_block_ids, raw_hidden=raw_hidden, diff --git a/python/sglang/srt/speculative/dspark_components/dspark_target_verify.py b/python/sglang/srt/speculative/dspark_components/dspark_target_verify.py index d967f1ddf7af..5b8632fc1ef6 100644 --- a/python/sglang/srt/speculative/dspark_components/dspark_target_verify.py +++ b/python/sglang/srt/speculative/dspark_components/dspark_target_verify.py @@ -70,12 +70,13 @@ def run_non_compact( seq_lens_cpu_backup = batch.seq_lens_cpu seq_lens_sum_backup = batch.seq_lens_sum if not self._verify_backend_self_adds_seq_lens(): - if seq_lens_cpu_backup is not None: - batch.seq_lens_cpu = seq_lens_cpu_backup + verify_w - batch.seq_lens_sum = int(batch.seq_lens_cpu.sum()) - elif draft_input.reserved_seq_lens_cpu is not None: - batch.seq_lens_cpu = draft_input.reserved_seq_lens_cpu - batch.seq_lens_sum = int(draft_input.reserved_seq_lens_sum) + base_seq_lens_cpu = ( + seq_lens_cpu_backup + if seq_lens_cpu_backup is not None + else batch.seq_lens.cpu() + ) + batch.seq_lens_cpu = base_seq_lens_cpu + batch.seq_lens_sum = int(base_seq_lens_cpu.sum()) verify_forward_batch, _ = verify_input.prepare_for_verify( batch, self.target_worker @@ -156,16 +157,14 @@ def _run_ragged( batch.out_cache_loc = ragged_window.verify_cache_loc seq_lens_cpu_backup = batch.seq_lens_cpu seq_lens_sum_backup = batch.seq_lens_sum - if seq_lens_cpu_backup is not None: - verify_lens_cpu = ( - layout.verify_lens_cpu - if layout.verify_lens_cpu is not None - else layout.verify_lens.cpu().tolist() - ) - batch.seq_lens_cpu = seq_lens_cpu_backup + torch.tensor( - verify_lens_cpu, dtype=seq_lens_cpu_backup.dtype + if not self._verify_backend_self_adds_seq_lens(): + base_seq_lens_cpu = ( + seq_lens_cpu_backup + if seq_lens_cpu_backup is not None + else batch.seq_lens.cpu() ) - batch.seq_lens_sum = int(batch.seq_lens_cpu.sum()) + batch.seq_lens_cpu = base_seq_lens_cpu + batch.seq_lens_sum = int(base_seq_lens_cpu.sum()) verify_forward_batch, _ = verify_input.prepare_for_verify( batch, self.target_worker diff --git a/python/sglang/srt/speculative/dspark_components/dspark_utils.py b/python/sglang/srt/speculative/dspark_components/dspark_utils.py index dbfb56d2a7fc..e820cc8f7f75 100644 --- a/python/sglang/srt/speculative/dspark_components/dspark_utils.py +++ b/python/sglang/srt/speculative/dspark_components/dspark_utils.py @@ -77,10 +77,16 @@ def _get_text_config(config: Any) -> Any: if config is None: return None if isinstance(config, dict): - return config.get("text_config", config) + text_config = config.get("text_config", None) + if text_config is not None: + return text_config + return config.get("transformer_layer_config", config) text_config = getattr(config, "text_config", None) if text_config is not None: return text_config + transformer_layer_config = getattr(config, "transformer_layer_config", None) + if transformer_layer_config is not None: + return transformer_layer_config return config @@ -96,6 +102,47 @@ def _get_dspark_config(config: Any) -> dict: return {} +def _get_speculators_config(config: Any) -> dict: + cfg = _cfg_get(config, "speculators_config", None) + if cfg is None: + return {} + if isinstance(cfg, dict): + return cfg + try: + return dict(cfg) + except Exception: + return {} + + +def _resolve_speculators_proposal_gamma(config: Any) -> Optional[int]: + cfg = _get_speculators_config(config) + proposal_methods = cfg.get("proposal_methods") or [] + if not isinstance(proposal_methods, (list, tuple)) or not proposal_methods: + return None + + default_method = cfg.get("default_proposal_method") + selected = None + if default_method is not None: + for method in proposal_methods: + proposal_type = _cfg_get(method, "proposal_type", None) + if proposal_type == default_method: + selected = method + break + if selected is None: + selected = proposal_methods[0] + + speculative_tokens = _cfg_get(selected, "speculative_tokens", None) + if speculative_tokens is None: + return None + gamma = int(speculative_tokens) + if gamma < 1: + raise ValueError( + "DSpark speculators_config speculative_tokens must be positive, " + f"got {gamma}." + ) + return gamma + + def parse_dspark_draft_config(*, draft_hf_config: Any) -> DSparkDraftConfig: base = parse_dflash_draft_config(draft_hf_config=draft_hf_config) @@ -177,8 +224,11 @@ def parse_dspark_draft_config(*, draft_hf_config: Any) -> DSparkDraftConfig: f"DSpark mask_token_id must be non-negative, got {mask_token_id}." ) + speculators_gamma = _resolve_speculators_proposal_gamma(draft_hf_config) gamma = ( - int(prefixed_block_size) if prefixed_block_size is not None else base.block_size + int(prefixed_block_size) + if prefixed_block_size is not None + else speculators_gamma if speculators_gamma is not None else base.block_size ) if prefixed_target_layer_ids is not None: diff --git a/python/sglang/srt/speculative/dspark_components/dspark_worker_v2.py b/python/sglang/srt/speculative/dspark_components/dspark_worker_v2.py index 78b0dad3a3fe..9264bf9f9779 100644 --- a/python/sglang/srt/speculative/dspark_components/dspark_worker_v2.py +++ b/python/sglang/srt/speculative/dspark_components/dspark_worker_v2.py @@ -217,7 +217,7 @@ def __init__( length=self.verify_num_draft_tokens, device=self.device ) self._draft_block_spec_info = make_draft_block_spec_info( - draft_token_num=int(self.gamma), device=self.device + draft_token_num=int(self.verify_num_draft_tokens), device=self.device ) target_model = self.target_worker.model_runner.model diff --git a/python/sglang/srt/speculative/spec_info.py b/python/sglang/srt/speculative/spec_info.py index 124014517d7c..6a2e3c6e488c 100644 --- a/python/sglang/srt/speculative/spec_info.py +++ b/python/sglang/srt/speculative/spec_info.py @@ -205,8 +205,6 @@ def get_num_tokens_per_bs_for_target_verify( # graph support. We can use it for target verify, or we can use it for # other cases which is not target verify but fixed length prefill. # Here, we expose this interface to allow the other use cases. - if self.is_dspark() and is_draft_worker: - return num_draft_tokens - 1 return num_draft_tokens def create_worker( @@ -357,7 +355,7 @@ def create_dummy_verify_input( spec_info = DFlashVerifyInput( draft_token=None, positions=None, - draft_token_num=server_args.speculative_num_draft_tokens, + draft_token_num=num_tokens_per_bs, custom_mask=None, capture_hidden_mode=( CaptureHiddenMode.NULL if is_draft_worker else CaptureHiddenMode.FULL diff --git a/python/sglang/srt/utils/hf_transformers/__init__.py b/python/sglang/srt/utils/hf_transformers/__init__.py index 3e6b3fa78845..d1e4a9d8001e 100644 --- a/python/sglang/srt/utils/hf_transformers/__init__.py +++ b/python/sglang/srt/utils/hf_transformers/__init__.py @@ -36,7 +36,7 @@ get_sparse_attention_config, get_tokenizer_from_processor, ) -from .config import get_config +from .config import get_config, get_config_allow_missing_model_type from .processor import get_processor from .tokenizer import ( _fix_added_tokens_encoding, @@ -53,6 +53,7 @@ "check_gguf_file", "download_from_hf", "get_config", + "get_config_allow_missing_model_type", "get_context_length", "get_generation_config", "get_hf_text_config", diff --git a/python/sglang/srt/utils/hf_transformers/config.py b/python/sglang/srt/utils/hf_transformers/config.py index c1cb0b526015..e12f2c381255 100644 --- a/python/sglang/srt/utils/hf_transformers/config.py +++ b/python/sglang/srt/utils/hf_transformers/config.py @@ -58,6 +58,11 @@ def _is_legacy_glm_moe_dsa_layer_types_error(error: Exception) -> bool: ) +def is_missing_model_type_error(error: ValueError) -> bool: + error_msg = str(error) + return "Unrecognized model" in error_msg and "model_type" in error_msg + + def _load_glm_moe_dsa_config_without_legacy_layer_types( model, revision: Optional[str] = None, @@ -308,3 +313,46 @@ def get_config( _set_architectures(config, MODEL_FOR_CAUSAL_LM_MAPPING_NAMES[config.model_type]) return config + + +def get_config_allow_missing_model_type( + model: str, + trust_remote_code: bool, + revision: Optional[str] = None, + model_override_args: Optional[dict] = None, + model_config_parser: str = "auto", + **kwargs, +): + """Load a config, falling back to a generic config for draft-only checkpoints. + + Some speculative draft checkpoints are intentionally not target models and + omit ``model_type`` while still carrying SGLang-readable architecture and + draft metadata. ``AutoConfig`` rejects those before SGLang can inspect the + metadata, so fall back to a plain ``PretrainedConfig`` only for that exact + missing-model-type error. + """ + try: + return get_config( + model, + trust_remote_code=trust_remote_code, + revision=revision, + model_override_args=model_override_args, + model_config_parser=model_config_parser, + **kwargs, + ) + except ValueError as exc: + if not is_missing_model_type_error(exc): + raise + + from transformers import PretrainedConfig + + raw_config_dict, _ = PretrainedConfig.get_config_dict( + model, + trust_remote_code=trust_remote_code, + revision=revision, + **kwargs, + ) + config = PretrainedConfig.from_dict(raw_config_dict) + if model_override_args: + config.update(model_override_args) + return config diff --git a/test/registered/spec/dspark/test_dspark_draft_path_default.py b/test/registered/spec/dspark/test_dspark_draft_path_default.py index b09ae0651435..95e7cb1b96ee 100644 --- a/test/registered/spec/dspark/test_dspark_draft_path_default.py +++ b/test/registered/spec/dspark/test_dspark_draft_path_default.py @@ -1,7 +1,9 @@ import unittest from types import SimpleNamespace +from unittest.mock import patch from sglang.srt.arg_groups.speculative_hook import ( + _handle_dflash, _handle_dspark, _target_checkpoint_bundles_dspark_draft, ) @@ -29,6 +31,31 @@ def _plain_hf_config() -> SimpleNamespace: return SimpleNamespace(architectures=["DeepseekV4ForCausalLM"]) +def _redhat_glm52_dspark_config() -> SimpleNamespace: + return SimpleNamespace( + architectures=["DSparkDraftModel"], + aux_hidden_state_layer_ids=[8, 23, 39, 55, 70], + block_size=8, + markov_head_type="vanilla", + markov_rank=256, + mask_token_id=154856, + speculators_config={ + "default_proposal_method": "greedy", + "proposal_methods": [ + { + "proposal_type": "greedy", + "speculative_tokens": 7, + } + ], + }, + transformer_layer_config={ + "hidden_size": 6144, + "num_hidden_layers": 5, + "vocab_size": 154880, + }, + ) + + def _make_dspark_server_args( *, model_path: str, hf_config: SimpleNamespace ) -> ServerArgs: @@ -83,6 +110,72 @@ def test_explicit_draft_path_is_not_overwritten(self): "deepseek-ai/some-other-dspark-draft", ) + def test_external_config_infers_gamma_from_speculators_config(self): + server_args = _make_dspark_server_args( + model_path=_PLAIN_MODEL_PATH, hf_config=_plain_hf_config() + ) + server_args.speculative_draft_model_path = "RedHatAI/GLM-5.2-speculator.dspark" + server_args.speculative_dspark_block_size = None + + with patch( + "sglang.srt.arg_groups.speculative_hook._load_speculative_draft_config", + return_value=_redhat_glm52_dspark_config(), + ): + _handle_dspark(server_args) + + self.assertEqual(server_args.speculative_num_draft_tokens, 8) + + def test_explicit_num_draft_tokens_uses_external_config_loader(self): + server_args = _make_dspark_server_args( + model_path=_PLAIN_MODEL_PATH, hf_config=_plain_hf_config() + ) + server_args.speculative_draft_model_path = "RedHatAI/GLM-5.2-speculator.dspark" + server_args.speculative_dspark_block_size = None + server_args.speculative_num_draft_tokens = 8 + + with patch( + "sglang.srt.arg_groups.speculative_hook._load_speculative_draft_config", + return_value=_redhat_glm52_dspark_config(), + ) as load_config: + _handle_dspark(server_args) + + load_config.assert_called_once() + self.assertEqual(server_args.speculative_num_draft_tokens, 8) + + def test_explicit_num_draft_tokens_validates_gamma_from_speculators_config(self): + server_args = _make_dspark_server_args( + model_path=_PLAIN_MODEL_PATH, hf_config=_plain_hf_config() + ) + server_args.speculative_draft_model_path = "RedHatAI/GLM-5.2-speculator.dspark" + server_args.speculative_dspark_block_size = None + server_args.speculative_num_draft_tokens = 7 + + with ( + patch( + "sglang.srt.arg_groups.speculative_hook._load_speculative_draft_config", + return_value=_redhat_glm52_dspark_config(), + ), + self.assertRaisesRegex(ValueError, "gamma \\+ 1"), + ): + _handle_dspark(server_args) + + def test_dflash_rejects_dspark_draft_checkpoint(self): + server_args = _make_dspark_server_args( + model_path=_PLAIN_MODEL_PATH, hf_config=_plain_hf_config() + ) + server_args.speculative_algorithm = "DFLASH" + server_args.speculative_draft_model_path = "RedHatAI/GLM-5.2-speculator.dspark" + server_args.speculative_dspark_block_size = None + + with ( + patch( + "sglang.srt.arg_groups.speculative_hook._load_speculative_draft_config", + return_value=_redhat_glm52_dspark_config(), + ), + self.assertRaisesRegex(ValueError, "Use --speculative-algorithm DSPARK"), + ): + _handle_dflash(server_args) + if __name__ == "__main__": unittest.main() diff --git a/test/registered/spec/dspark/test_glm52_dspark_config.py b/test/registered/spec/dspark/test_glm52_dspark_config.py new file mode 100644 index 000000000000..693fc2188db4 --- /dev/null +++ b/test/registered/spec/dspark/test_glm52_dspark_config.py @@ -0,0 +1,213 @@ +import unittest +from types import SimpleNamespace +from unittest.mock import patch + +from transformers import PretrainedConfig + +from sglang.srt.configs.model_config import ( + _normalize_nested_transformer_config, + _restore_glm_moe_dsa_head_dims_from_raw_config, +) +from sglang.srt.models.dspark import EntryClass, normalize_dspark_draft_config +from sglang.srt.server_args import ServerArgs +from sglang.srt.speculative.draft_worker_common import ( + _resolve_draft_attention_backend, + draft_is_deepseek_v4, +) +from sglang.srt.speculative.dflash_utils import parse_dflash_draft_config +from sglang.srt.speculative.dspark_components.dspark_utils import ( + parse_dspark_draft_config, +) +from sglang.test.ci.ci_register import register_cpu_ci +from sglang.test.test_utils import CustomTestCase + +register_cpu_ci(est_time=10, suite="base-a-test-cpu") + + +def _glm52_redhat_dspark_config() -> SimpleNamespace: + return SimpleNamespace( + architectures=["DSparkDraftModel"], + aux_hidden_state_layer_ids=[8, 23, 39, 55, 70], + block_size=8, + confidence_head_with_markov=True, + enable_confidence_head=True, + markov_head_type="vanilla", + markov_rank=256, + mask_token_id=154856, + speculators_config={ + "algorithm": "dspark", + "default_proposal_method": "greedy", + "proposal_methods": [ + { + "proposal_type": "greedy", + "speculative_tokens": 7, + "verifier_accept_k": 1, + } + ], + }, + transformer_layer_config={ + "attention_bias": False, + "head_dim": 64, + "hidden_size": 6144, + "intermediate_size": 12288, + "layer_types": ["full_attention"] * 5, + "model_type": "qwen3", + "num_attention_heads": 64, + "num_hidden_layers": 5, + "num_key_value_heads": 64, + "rms_norm_eps": 1e-5, + "vocab_size": 154880, + }, + ) + + +def _glm52_redhat_dspark_raw_config() -> dict: + return dict(vars(_glm52_redhat_dspark_config())) + + +def _make_draft_server_args() -> ServerArgs: + server_args = ServerArgs(model_path="dummy") + server_args.speculative_draft_model_path = "RedHatAI/GLM-5.2-speculator.dspark" + server_args.speculative_draft_model_revision = None + server_args.speculative_draft_attention_backend = "triton" + server_args.trust_remote_code = False + server_args.json_model_override_args = "{}" + server_args.model_config_parser = "auto" + return server_args + + +class TestGLM52RedHatDSparkConfig(CustomTestCase): + def test_model_config_normalizer_materializes_nested_transformer_config(self): + config = PretrainedConfig( + architectures=["DSparkDraftModel"], + aux_hidden_state_layer_ids=[8, 23, 39, 55, 70], + transformer_layer_config={ + "hidden_size": 6144, + "num_hidden_layers": 5, + "vocab_size": 154880, + }, + ) + + _normalize_nested_transformer_config(config) + + self.assertEqual(config.hidden_size, 6144) + self.assertEqual(config.num_hidden_layers, 5) + self.assertEqual(config.vocab_size, 154880) + self.assertEqual(config.num_target_layers, 71) + + def test_dflash_parser_reads_nested_transformer_config_and_aux_layers(self): + parsed = parse_dflash_draft_config( + draft_hf_config=_glm52_redhat_dspark_config() + ) + + self.assertEqual(parsed.num_hidden_layers, 5) + self.assertEqual(parsed.block_size, 8) + self.assertEqual(parsed.target_layer_ids, [8, 23, 39, 55, 70]) + self.assertEqual(parsed.num_target_layers, 71) + + def test_dspark_parser_reads_redhat_fields(self): + parsed = parse_dspark_draft_config( + draft_hf_config=_glm52_redhat_dspark_config() + ) + + self.assertEqual(parsed.gamma, 7) + self.assertEqual(parsed.target_layer_ids, [8, 23, 39, 55, 70]) + self.assertEqual(parsed.markov_rank, 256) + self.assertEqual(parsed.markov_head_type, "vanilla") + self.assertEqual(parsed.mask_token_id, 154856) + + def test_dspark_draft_model_is_registered(self): + self.assertIn("DSparkDraftModel", {cls.__name__ for cls in EntryClass}) + + def test_normalize_dspark_draft_config_materializes_backbone_fields(self): + config = _glm52_redhat_dspark_config() + normalize_dspark_draft_config(config) + + self.assertEqual(config.hidden_size, 6144) + self.assertEqual(config.num_hidden_layers, 5) + self.assertEqual(config.vocab_size, 154880) + self.assertEqual(config.num_target_layers, 71) + + +class TestGLM52RedHatDSparkWorkerConfig(CustomTestCase): + def test_draft_worker_common_accepts_missing_model_type_config(self): + server_args = _make_draft_server_args() + missing_model_type = ValueError( + "Unrecognized model in RedHatAI/GLM-5.2-speculator.dspark. " + "Should have a `model_type` key in its config.json." + ) + + with ( + patch( + "sglang.srt.utils.hf_transformers.config.get_config", + side_effect=missing_model_type, + ), + patch( + "transformers.PretrainedConfig.get_config_dict", + return_value=(_glm52_redhat_dspark_raw_config(), {}), + ), + ): + self.assertFalse(draft_is_deepseek_v4(server_args=server_args)) + self.assertEqual( + _resolve_draft_attention_backend( + draft_server_args=server_args, algo_label="DSpark" + ), + "triton", + ) + + +class TestGLM52RawHeadDims(CustomTestCase): + def test_restore_raw_glm_moe_dsa_head_dims(self): + config = PretrainedConfig( + architectures=["GlmMoeDsaForCausalLM"], + qk_nope_head_dim=192, + qk_rope_head_dim=192, + qk_head_dim=384, + v_head_dim=256, + ) + raw_config = { + "qk_nope_head_dim": 192, + "qk_rope_head_dim": 64, + "qk_head_dim": 256, + "v_head_dim": 256, + } + + _restore_glm_moe_dsa_head_dims_from_raw_config( + config, + raw_config, + {}, + "unused", + False, + None, + {}, + ) + + self.assertEqual(config.qk_nope_head_dim, 192) + self.assertEqual(config.qk_rope_head_dim, 64) + self.assertEqual(config.qk_head_dim, 256) + self.assertEqual(config.v_head_dim, 256) + + def test_json_override_wins_over_raw_head_dims(self): + config = PretrainedConfig( + architectures=["GlmMoeDsaForCausalLM"], + qk_nope_head_dim=192, + qk_rope_head_dim=192, + qk_head_dim=384, + v_head_dim=256, + ) + + _restore_glm_moe_dsa_head_dims_from_raw_config( + config, + {"qk_rope_head_dim": 64}, + {"qk_rope_head_dim": 128}, + "unused", + False, + None, + {}, + ) + + self.assertEqual(config.qk_rope_head_dim, 128) + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/unit/spec/test_dflash_dspark_verify_lengths.py b/test/registered/unit/spec/test_dflash_dspark_verify_lengths.py new file mode 100644 index 000000000000..92dd46c7db9c --- /dev/null +++ b/test/registered/unit/spec/test_dflash_dspark_verify_lengths.py @@ -0,0 +1,346 @@ +import unittest +from types import SimpleNamespace +from unittest.mock import patch + +import torch + +from sglang.srt.environ import envs +from sglang.srt.mem_cache.common import ( + get_alloc_reserve_per_decode, + get_req_to_token_extra_context_len, +) +from sglang.srt.server_args import ServerArgs, set_global_server_args_for_scheduler +from sglang.srt.speculative.dflash_info import DFlashVerifyInput +from sglang.srt.speculative.dflash_info_v2 import DFlashDraftInputV2 +from sglang.srt.speculative.dflash_worker_v2 import _copy_prefix_seq_lens_cpu +from sglang.srt.speculative.dspark_components.dspark_draft_proposer import ( + DraftBlockProposer, +) +from sglang.srt.speculative.dspark_components.dspark_target_verify import ( + TargetVerifyExecutor, +) +from sglang.srt.speculative.spec_info import ( + SpeculativeAlgorithm, + create_dummy_verify_input, +) +from sglang.test.ci.ci_register import register_cpu_ci +from sglang.test.test_utils import CustomTestCase + +register_cpu_ci(est_time=3, suite="base-a-test-cpu") + + +def _spec_args(*, draft_tokens: int = 8, algorithm: str = "DSPARK") -> ServerArgs: + args = ServerArgs(model_path="dummy") + args.speculative_algorithm = algorithm + args.speculative_num_draft_tokens = draft_tokens + args.max_speculative_num_draft_tokens = draft_tokens + args.speculative_num_steps = None + args.speculative_eagle_topk = 1 + args.page_size = 1 + return args + + +class _FakeBatch: + def __init__(self, *, committed_lens, allocated_lens, row_width: int): + self.device = torch.device("cpu") + self.reqs = [ + SimpleNamespace( + kv_committed_len=committed_len, + kv_allocated_len=allocated_len, + sampling_params=SimpleNamespace(top_k=1), + ) + for committed_len, allocated_len in zip(committed_lens, allocated_lens) + ] + self.token_to_kv_pool_allocator = SimpleNamespace(page_size=1) + self.req_to_token_pool = SimpleNamespace( + req_to_token=torch.empty((len(self.reqs), row_width), dtype=torch.int32) + ) + + def batch_size(self): + return len(self.reqs) + + +class _FakeTargetWorker: + def __init__(self): + self.model_runner = SimpleNamespace(attn_backend=SimpleNamespace()) + + def forward_batch_generation(self, **kwargs): + return SimpleNamespace( + logits_output=SimpleNamespace(next_token_logits=torch.empty(0)), + can_run_cuda_graph=False, + ) + + +class TestDFlashDSparkVerifyLengths(CustomTestCase): + def test_req_to_token_headroom_covers_spec_v2_double_buffer(self): + args = _spec_args(draft_tokens=8) + self.assertEqual(get_alloc_reserve_per_decode(args), 16) + self.assertGreaterEqual(get_req_to_token_extra_context_len(args), 16) + + def test_req_to_token_headroom_keeps_non_tree_eagle_default(self): + args = _spec_args(draft_tokens=8, algorithm="EAGLE") + self.assertEqual(get_alloc_reserve_per_decode(args), 16) + self.assertEqual(get_req_to_token_extra_context_len(args), 12) + + def test_prepare_for_decode_keeps_committed_and_reserved_lengths_separate(self): + args = _spec_args(draft_tokens=4) + set_global_server_args_for_scheduler(args) + draft_input = DFlashDraftInputV2.create_idle_input(torch.device("cpu")) + batch = _FakeBatch( + committed_lens=[10, 20], + allocated_lens=[18, 28], + row_width=64, + ) + + with envs.SGLANG_ENABLE_OVERLAP_PLAN_STREAM.override(False): + draft_input.prepare_for_decode(batch) + + self.assertEqual(batch.seq_lens_cpu.tolist(), [10, 20]) + self.assertEqual(batch.seq_lens_sum, 30) + self.assertEqual(draft_input.reserved_seq_lens_cpu.tolist(), [18, 28]) + self.assertEqual(draft_input.reserved_seq_lens_sum, 46) + + def test_prepare_for_decode_fails_before_req_to_token_oob(self): + args = _spec_args(draft_tokens=8) + set_global_server_args_for_scheduler(args) + draft_input = DFlashDraftInputV2.create_idle_input(torch.device("cpu")) + batch = _FakeBatch(committed_lens=[8], allocated_lens=[24], row_width=23) + + with envs.SGLANG_ENABLE_OVERLAP_PLAN_STREAM.override(False): + with self.assertRaisesRegex(AssertionError, "over-allocation"): + draft_input.prepare_for_decode(batch) + + def test_dflash_draft_block_uses_prefix_seq_lens_cpu(self): + dst = torch.empty((2,), dtype=torch.int32) + prefix_lens = torch.tensor([10, 20], dtype=torch.int32) + host_lens = torch.tensor([10, 20], dtype=torch.int64) + + seq_lens_sum = _copy_prefix_seq_lens_cpu(dst, prefix_lens, host_lens) + self.assertEqual(dst.tolist(), [10, 20]) + self.assertEqual(seq_lens_sum, 30) + + dst.fill_(-1) + seq_lens_sum = _copy_prefix_seq_lens_cpu(dst, prefix_lens, None) + self.assertEqual(dst.tolist(), [10, 20]) + self.assertEqual(seq_lens_sum, 30) + + def test_dspark_target_verify_passes_prefix_seq_lens_cpu_not_reserved(self): + target_worker = _FakeTargetWorker() + executor = TargetVerifyExecutor( + target_worker=target_worker, + verify_num_draft_tokens=4, + model_runner=target_worker.model_runner, + kv_injector=SimpleNamespace(), + ) + batch = SimpleNamespace( + seq_lens=torch.tensor([10, 20], dtype=torch.int32), + seq_lens_cpu=None, + seq_lens_sum=None, + out_cache_loc=None, + ) + draft_input = SimpleNamespace( + reserved_seq_lens_cpu=torch.tensor([18, 28], dtype=torch.int32), + reserved_seq_lens_sum=46, + ) + verify_window = SimpleNamespace( + positions_2d=torch.arange(8, dtype=torch.int64).view(2, 4), + verify_cache_loc=torch.arange(8, dtype=torch.int64), + ) + seen = {} + + def capture_prepare(_verify_input, verify_batch, _target_worker): + seen["seq_lens_cpu"] = verify_batch.seq_lens_cpu.clone() + seen["seq_lens_sum"] = verify_batch.seq_lens_sum + return SimpleNamespace(), False + + with patch.object( + DFlashVerifyInput, "prepare_for_verify", new=capture_prepare + ): + executor.run_non_compact( + batch=batch, + draft_input=draft_input, + verify_ids_2d=torch.ones((2, 4), dtype=torch.int64), + verify_window=verify_window, + sampling_info=None, + ) + + self.assertEqual(seen["seq_lens_cpu"].tolist(), [10, 20]) + self.assertEqual(seen["seq_lens_sum"], 30) + self.assertIsNone(batch.seq_lens_cpu) + self.assertIsNone(batch.seq_lens_sum) + + def test_dspark_draft_proposer_passes_prefix_seq_lens_cpu_and_crops_anchor_hidden( + self, + ): + seen = {} + gamma = 4 + draft_width = gamma + 1 + bs = 2 + + class FakeDraftRunner: + device = "cpu" + + def forward(self, forward_batch): + seen["seq_lens"] = forward_batch.seq_lens.clone() + seen["seq_lens_cpu"] = forward_batch.seq_lens_cpu.clone() + seen["seq_lens_sum"] = forward_batch.seq_lens_sum + hidden = torch.arange( + bs * draft_width * 16, dtype=torch.float32 + ).view(bs * draft_width, 16) + return SimpleNamespace( + logits_output=SimpleNamespace(hidden_states=hidden), + can_run_graph=False, + ) + + proposer = DraftBlockProposer( + draft_model=SimpleNamespace(), + draft_model_runner=FakeDraftRunner(), + gamma=gamma, + mask_token_id=0, + draft_block_spec_info=SimpleNamespace(), + ) + batch = SimpleNamespace( + seq_lens=torch.tensor([10, 20], dtype=torch.int32), + seq_lens_cpu=torch.tensor([10, 20], dtype=torch.int32), + seq_lens_sum=30, + req_pool_indices=torch.tensor([0, 1], dtype=torch.int64), + ) + draft_input = SimpleNamespace( + bonus_tokens=torch.tensor([7, 8], dtype=torch.int64), + ) + verify_window = SimpleNamespace( + positions_2d=torch.arange(bs * draft_width, dtype=torch.int64).view( + bs, draft_width + ), + verify_cache_loc_2d=torch.arange( + bs * draft_width, dtype=torch.int64 + ).view(bs, draft_width), + ) + + out = proposer._run_forward( + batch=batch, + draft_input=draft_input, + verify_window=verify_window, + bs=bs, + device="cpu", + embed_module=torch.nn.Embedding(16, 16), + ) + + self.assertEqual(seen["seq_lens"].tolist(), [10, 20]) + self.assertEqual(seen["seq_lens_cpu"].tolist(), [10, 20]) + self.assertEqual(seen["seq_lens_sum"], 30) + self.assertEqual(tuple(out.draft_block_ids.shape), (bs, draft_width)) + self.assertEqual(tuple(out.draft_hidden_3d.shape), (bs, gamma, 16)) + self.assertEqual(out.raw_hidden[0].tolist(), list(range(16, 32))) + + def test_dspark_draft_proposer_derives_cpu_lens_from_gpu_only_batch(self): + seen = {} + gamma = 4 + draft_width = gamma + 1 + bs = 2 + + class FakeDraftRunner: + device = "cpu" + + def forward(self, forward_batch): + seen["seq_lens"] = forward_batch.seq_lens.clone() + seen["seq_lens_cpu"] = forward_batch.seq_lens_cpu.clone() + seen["seq_lens_sum"] = forward_batch.seq_lens_sum + return SimpleNamespace( + logits_output=SimpleNamespace( + hidden_states=torch.empty((bs * draft_width, 16)) + ), + can_run_graph=False, + ) + + proposer = DraftBlockProposer( + draft_model=SimpleNamespace(), + draft_model_runner=FakeDraftRunner(), + gamma=gamma, + mask_token_id=0, + draft_block_spec_info=SimpleNamespace(), + ) + batch = SimpleNamespace( + seq_lens=torch.tensor([10, 20], dtype=torch.int32), + seq_lens_cpu=None, + seq_lens_sum=None, + req_pool_indices=torch.tensor([0, 1], dtype=torch.int64), + ) + draft_input = SimpleNamespace( + bonus_tokens=torch.tensor([7, 8], dtype=torch.int64), + reserved_seq_lens_cpu=torch.tensor([18, 28], dtype=torch.int32), + reserved_seq_lens_sum=46, + ) + verify_window = SimpleNamespace( + positions_2d=torch.arange(bs * draft_width, dtype=torch.int64).view( + bs, draft_width + ), + verify_cache_loc_2d=torch.arange( + bs * draft_width, dtype=torch.int64 + ).view(bs, draft_width), + ) + + proposer._run_forward( + batch=batch, + draft_input=draft_input, + verify_window=verify_window, + bs=bs, + device="cpu", + embed_module=torch.nn.Embedding(16, 16), + ) + + self.assertEqual(seen["seq_lens"].tolist(), [10, 20]) + self.assertEqual(seen["seq_lens_cpu"].tolist(), [10, 20]) + self.assertEqual(seen["seq_lens_sum"], 30) + + def test_dspark_draft_dummy_verify_input_uses_verify_window(self): + args = _spec_args(draft_tokens=8) + spec_algorithm = SpeculativeAlgorithm.DSPARK + + draft_spec = create_dummy_verify_input( + spec_algorithm=spec_algorithm, + server_args=args, + custom_mask=torch.empty(0, dtype=torch.bool), + num_tokens_per_bs=8, + is_draft_worker=True, + ) + target_spec = create_dummy_verify_input( + spec_algorithm=spec_algorithm, + server_args=args, + custom_mask=torch.empty(0, dtype=torch.bool), + num_tokens_per_bs=8, + is_draft_worker=False, + ) + + self.assertEqual(draft_spec.draft_token_num, 8) + self.assertEqual(target_spec.draft_token_num, 8) + + def test_target_verify_width_adjustment_keeps_dspark_window(self): + self.assertEqual( + SpeculativeAlgorithm.DSPARK.get_num_tokens_per_bs_for_target_verify( + 8, is_draft_worker=True + ), + 8, + ) + self.assertEqual( + SpeculativeAlgorithm.DSPARK.get_num_tokens_per_bs_for_target_verify( + 8, is_draft_worker=False + ), + 8, + ) + self.assertEqual( + SpeculativeAlgorithm.EAGLE.get_num_tokens_per_bs_for_target_verify( + 8, is_draft_worker=True + ), + 8, + ) + self.assertEqual( + SpeculativeAlgorithm.NGRAM.get_num_tokens_per_bs_for_target_verify( + 8, is_draft_worker=True + ), + 8, + ) + + +if __name__ == "__main__": + unittest.main()