Skip to content
Draft
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
82 changes: 60 additions & 22 deletions python/sglang/srt/arg_groups/speculative_hook.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down Expand Up @@ -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.
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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,
Expand All @@ -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}."
)
Expand Down
96 changes: 87 additions & 9 deletions python/sglang/srt/configs/model_config.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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(
Expand Down
14 changes: 8 additions & 6 deletions python/sglang/srt/debug_utils/pr_fix_toggle.py
Original file line number Diff line number Diff line change
Expand Up @@ -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: ""
"""

Expand Down
7 changes: 6 additions & 1 deletion python/sglang/srt/layers/attention/dsa/dsa_indexer.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down
49 changes: 47 additions & 2 deletions python/sglang/srt/layers/attention/dsa_backend.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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()
Expand Down Expand Up @@ -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
)
Expand All @@ -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
]
Expand Down Expand Up @@ -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,
Expand Down
11 changes: 11 additions & 0 deletions python/sglang/srt/layers/attention/triton_backend.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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
Expand Down
Loading