feat: add optimized Domino rollout to DFlash V2 (#36899)

Co-authored-by: Qiaolin-Yu <liin1211@outlook.com>
This commit is contained in:
Francis
2026-09-09 14:12:15 -07:00
committed by GitHub
co-authored by Qiaolin-Yu
parent 2948a62a6f
commit a84ffd1326
8 changed files with 861 additions and 2 deletions
@@ -73,6 +73,10 @@ class Spec:
Optional[int],
"DFLASH only. Block size (verify window length). Alias of --speculative-num-draft-tokens for DFLASH.",
] = None
speculative_domino_candidate_pool_size: A[
int,
"Domino only. Size of the approximate block-shared base-logit candidate pool. Set to 0 to score the full vocabulary.",
] = 2048
speculative_dspark_block_size: A[
Optional[int],
"DSPARK only. Draft block size gamma (number of proposed draft tokens). The verify window is gamma + 1, so this sets --speculative-num-draft-tokens = gamma + 1. Omit to auto-infer gamma from the draft checkpoint block_size.",
+62
View File
@@ -635,6 +635,35 @@ class DFlashDraftModel(nn.Module):
)
self.hidden_norm = RMSNorm(hidden_size, eps=rms_norm_eps)
# The model loader calls load_weights() before set_block_size(). Build
# Domino projector modules here so their parameters are present while
# checkpoint weights are loaded.
self.projector_type = draft_config.projector_type
self.shift_label = draft_config.shift_label
self.prefix_gru: Optional[nn.GRU] = None
self.embed_proj: Optional[nn.Sequential] = None
if draft_config.is_domino:
assert draft_config.gru_hidden_dim is not None
assert draft_config.emb_dim is not None
self.prefix_gru = nn.GRU(
input_size=hidden_size,
hidden_size=int(draft_config.gru_hidden_dim),
num_layers=1,
batch_first=True,
bias=False,
)
self.embed_proj = nn.Sequential(
nn.Linear(
hidden_size + int(draft_config.gru_hidden_dim),
int(draft_config.emb_dim),
bias=False,
),
nn.SiLU(),
nn.Linear(
int(draft_config.emb_dim), int(config.vocab_size), bias=False
),
)
def set_block_size(self, block_size: int) -> None:
"""Adopt the block size the worker resolved.
@@ -728,6 +757,7 @@ class DFlashDraftModel(nn.Module):
]
params_dict = dict(self.named_parameters())
loaded_params = set()
# Alias the native export's "encoder." names.
_VENDOR_ENCODER_ALIASES = {
@@ -752,6 +782,14 @@ class DFlashDraftModel(nn.Module):
return None
for name, loaded_weight in weights:
unprefixed_name = name.removeprefix("model.")
if self.projector_type != "domino" and unprefixed_name.startswith(
("prefix_gru.", "embed_proj.")
):
raise ValueError(
"DFLASH checkpoint contains Domino projector weights but "
f"projector_type={self.projector_type!r}."
)
for param_name, weight_name, shard_id in stacked_params_mapping:
if f".{weight_name}." not in name:
continue
@@ -762,6 +800,7 @@ class DFlashDraftModel(nn.Module):
param = params_dict[resolved_name]
weight_loader = getattr(param, "weight_loader", default_weight_loader)
weight_loader(param, loaded_weight, shard_id)
loaded_params.add(resolved_name)
break
else:
resolved_name = resolve_param_name(name)
@@ -796,8 +835,31 @@ class DFlashDraftModel(nn.Module):
f"(num_context_features={self.num_context_features}, hidden_size={int(self.config.hidden_size)}), "
f"but got {loaded_shape} for weight '{name}'."
)
if resolved_name.startswith(("prefix_gru.", "embed_proj.")) and tuple(
loaded_weight.shape
) != tuple(param.shape):
raise ValueError(
"DFLASH Domino projector weight shape mismatch: "
f"expected {resolved_name}{tuple(param.shape)}, got "
f"{tuple(loaded_weight.shape)} from {name!r}."
)
weight_loader = getattr(param, "weight_loader", default_weight_loader)
weight_loader(param, loaded_weight)
loaded_params.add(resolved_name)
if self.projector_type == "domino":
required = {
"prefix_gru.weight_ih_l0",
"prefix_gru.weight_hh_l0",
"embed_proj.0.weight",
"embed_proj.2.weight",
}
missing = required - loaded_params
if missing:
raise ValueError(
"DFLASH Domino checkpoint is missing required projector weights: "
f"{sorted(missing)}."
)
class DFlashLagunaAttention(DFlashAttention):
@@ -538,6 +538,15 @@ class DFlashDraftConfig:
target_layer_ids: Optional[List[int]]
mask_token: str
mask_token_id: Optional[int]
projector_type: Optional[str]
shift_label: Optional[bool]
pure_draft_prefix_len: Optional[int]
gru_hidden_dim: Optional[int]
emb_dim: Optional[int]
@property
def is_domino(self) -> bool:
return self.projector_type == "domino"
def require_num_layers(self) -> int:
if self.num_hidden_layers is None:
@@ -698,6 +707,72 @@ def parse_dflash_draft_config(*, draft_hf_config: Any) -> DFlashDraftConfig:
f"got {mask_token_id}."
)
projector_type = dflash_cfg.get(
"projector_type", _cfg_get(draft_hf_config, "projector_type", None)
)
shift_label = None
pure_draft_prefix_len = None
gru_hidden_dim = None
emb_dim = None
if projector_type == "domino":
shift_label = dflash_cfg.get(
"shift_label", _cfg_get(draft_hf_config, "shift_label", None)
)
pure_draft_prefix_len = _parse_optional_int(
dflash_cfg.get(
"pure_draft_prefix_len",
_cfg_get(draft_hf_config, "pure_draft_prefix_len", None),
),
field_name="DFLASH Domino pure_draft_prefix_len",
min_value=0,
)
gru_hidden_dim = _parse_optional_int(
dflash_cfg.get(
"gru_hidden_dim", _cfg_get(draft_hf_config, "gru_hidden_dim", None)
),
field_name="DFLASH Domino gru_hidden_dim",
min_value=1,
)
nested_emb_dim = _parse_optional_int(
dflash_cfg.get("emb_dim", None),
field_name="DFLASH Domino dflash_config.emb_dim",
min_value=1,
)
top_level_emb_dim = _parse_optional_int(
_cfg_get(draft_hf_config, "emb_dim", None),
field_name="DFLASH Domino top-level emb_dim",
min_value=1,
)
if (
nested_emb_dim is not None
and top_level_emb_dim is not None
and nested_emb_dim != top_level_emb_dim
):
raise ValueError(
"DFLASH Domino emb_dim differs between dflash_config and the "
f"top-level config: {nested_emb_dim} != {top_level_emb_dim}."
)
emb_dim = nested_emb_dim if nested_emb_dim is not None else top_level_emb_dim
if not isinstance(shift_label, bool):
raise ValueError(
"DFLASH Domino requires dflash_config.shift_label to be a bool, "
f"got {shift_label!r}."
)
if pure_draft_prefix_len != 1:
raise ValueError(
"DFLASH Domino currently requires pure_draft_prefix_len=1, "
f"got {pure_draft_prefix_len!r}."
)
if gru_hidden_dim is None:
raise ValueError("DFLASH Domino requires dflash_config.gru_hidden_dim.")
if emb_dim is None:
raise ValueError("DFLASH Domino requires dflash_config.emb_dim.")
if block_size is not None and block_size <= 1:
raise ValueError(
f"DFLASH Domino requires block_size > 1, got {block_size}."
)
return DFlashDraftConfig(
num_hidden_layers=num_hidden_layers,
num_target_layers=num_target_layers,
@@ -711,6 +786,11 @@ def parse_dflash_draft_config(*, draft_hf_config: Any) -> DFlashDraftConfig:
target_layer_ids=parsed_target_layer_ids,
mask_token=mask_token,
mask_token_id=mask_token_id,
projector_type=projector_type,
shift_label=shift_label,
pure_draft_prefix_len=pure_draft_prefix_len,
gru_hidden_dim=gru_hidden_dim,
emb_dim=emb_dim,
)
@@ -60,6 +60,10 @@ from sglang.srt.speculative.dflash_utils import (
is_dflash_sampling_verify_available,
parse_dflash_draft_config,
)
from sglang.srt.speculative.domino_utils import (
domino_greedy_rollout,
validate_domino_runtime,
)
from sglang.srt.speculative.draft_worker_common import (
build_block_pos_offsets,
build_draft_tp_worker,
@@ -280,6 +284,55 @@ class _SelectorDraftSampler:
self.q_out[:bs].copy_(q_rows)
class _DominoDraftSampler:
"""Capture-safe TP=1 Domino rollout over a fixed-size draft block."""
def __init__(
self,
*,
target_embedding,
lm_head_weight,
prefix_gru,
embed_proj,
vocab_size,
block_size,
shift_label,
max_bs,
candidate_pool_size,
):
self.target_embedding = target_embedding
self.lm_head_weight = lm_head_weight
self.prefix_gru = prefix_gru
self.embed_proj = embed_proj
self.vocab_size = int(vocab_size)
self.block_size = int(block_size)
self.shift_label = bool(shift_label)
self.candidate_pool_size = int(candidate_pool_size)
max_tokens = int(max_bs) * (self.block_size - 1)
self.out = torch.empty(
(max_tokens,), dtype=torch.int64, device=lm_head_weight.device
)
def __call__(self, hidden_states, input_ids=None):
if input_ids is None:
raise RuntimeError("Domino draft sampler requires block input_ids.")
bs = hidden_states.shape[0] // self.block_size
draft_hidden = hidden_states.view(bs, self.block_size, -1)
bonus_tokens = input_ids.view(bs, self.block_size)[:, 0]
proposals = domino_greedy_rollout(
draft_hidden=draft_hidden,
bonus_tokens=bonus_tokens,
target_embedding=self.target_embedding,
lm_head_weight=self.lm_head_weight,
prefix_gru=self.prefix_gru,
embed_proj=self.embed_proj,
vocab_size=self.vocab_size,
shift_label=self.shift_label,
candidate_pool_size=self.candidate_pool_size,
)
self.out[: bs * (self.block_size - 1)].copy_(proposals.reshape(-1))
class DFlashWorkerV2(BaseSpecWorker):
"""DFLASH speculative decoding worker (spec-v2).
@@ -333,6 +386,36 @@ class DFlashWorkerV2(BaseSpecWorker):
draft_config = parse_dflash_draft_config(
draft_hf_config=self.draft_model_runner.model_config.hf_config
)
self._is_domino = draft_config.is_domino
self.domino_candidate_pool_size = int(
get_spec().speculative_domino_candidate_pool_size
)
if self._is_domino:
if self.domino_candidate_pool_size < 0:
raise ValueError(
"--speculative-domino-candidate-pool-size must be non-negative, "
f"got {self.domino_candidate_pool_size}."
)
target_model = self.target_worker.model_runner.model
target_embedding = target_model.get_input_embeddings()
lm_head = getattr(target_model, "lm_head", None)
prefix_gru = getattr(self.draft_model, "prefix_gru", None)
embed_proj = getattr(self.draft_model, "embed_proj", None)
if lm_head is None or prefix_gru is None or embed_proj is None:
raise ValueError(
"DFLASH Domino requires target lm_head and loaded Domino projector modules."
)
validate_domino_runtime(
device=torch.device(self.device),
tp_size=int(get_tp_group().world_size),
target_vocab_size=int(self.model_runner.model_config.vocab_size),
draft_vocab_size=int(self.draft_model_runner.model_config.vocab_size),
hidden_size=int(self.draft_model.config.hidden_size),
target_embedding=target_embedding,
lm_head=lm_head,
prefix_gru=prefix_gru,
embed_proj=embed_proj,
)
if get_spec().speculative_num_draft_tokens is None:
# Should not happen (ServerArgs should have inferred it), but keep a fallback.
self.block_size = int(draft_config.resolve_block_size(default=16))
@@ -351,6 +434,11 @@ class DFlashWorkerV2(BaseSpecWorker):
)
self.draft_model.set_block_size(self.block_size)
self.speculative_num_draft_tokens = int(self.block_size)
if self._is_domino and self.block_size <= 1:
raise ValueError(
"DFLASH Domino requires speculative_num_draft_tokens > 1, "
f"got {self.block_size}."
)
self._mask_token = draft_config.mask_token
self._mask_token_id_override = draft_config.mask_token_id
@@ -373,6 +461,11 @@ class DFlashWorkerV2(BaseSpecWorker):
self.draft_window_size,
self.use_compact_draft_cache,
)
if self._is_domino:
logger.info(
"DFLASH Domino rollout enabled (BF16, TP=1, block-shared candidate pool size=%s).",
self.domino_candidate_pool_size,
)
logger.info(
"DFLASH draft runner ready. mask_token=%s, mask_token_id=%s, mask_token_id_override=%s, noise_embed_scale=%s",
self._mask_token,
@@ -661,6 +754,28 @@ class DFlashWorkerV2(BaseSpecWorker):
# Quantized lm_head (FP8/INT) would break the static matmul.
return _eager("quantized lm_head")
tp_group = get_tp_group()
if self._is_domino:
if tp_group.world_size != 1:
return _eager("Domino cuda graph currently requires tp=1")
prefix_gru = self.draft_model.prefix_gru
embed_proj = self.draft_model.embed_proj
if prefix_gru is None or embed_proj is None:
return _eager("Domino projector modules are unavailable")
if self.ps.tp_rank == 0:
logger.info(
"DFLASH Domino rollout folded into the draft cuda graph (tp=1)."
)
return _DominoDraftSampler(
target_embedding=target_model.get_input_embeddings(),
lm_head_weight=lm_head.weight,
prefix_gru=prefix_gru,
embed_proj=embed_proj,
vocab_size=int(self.model_runner.model_config.vocab_size),
block_size=self.block_size,
shift_label=self.draft_model.shift_label,
max_bs=max(get_exec().graph.cuda_graph_config.decode.bs),
candidate_pool_size=self.domino_candidate_pool_size,
)
if not hasattr(lm_head, "shard_indices"):
if tp_group.world_size != 1:
# No shard metadata to recover per-rank vocab offsets from.
@@ -2154,8 +2269,35 @@ class DFlashWorkerV2(BaseSpecWorker):
draft_out = self.draft_model_runner.forward(forward_batch)
draft_logits_output = draft_out.logits_output
folded = self._draft_sampler is not None and draft_out.can_run_graph
if folded:
if (
self._is_domino
and self._draft_sampler is not None
and draft_out.can_run_graph
):
draft_next = self._draft_sampler.out[
: bs * (int(self.block_size) - 1)
].view(bs, int(self.block_size) - 1)
elif self._is_domino:
draft_hidden = draft_logits_output.hidden_states
if draft_hidden is None:
raise RuntimeError("DFLASH draft model returned no hidden states.")
draft_hidden = draft_hidden.view(bs, int(self.block_size), -1)
prefix_gru = self.draft_model.prefix_gru
embed_proj = self.draft_model.embed_proj
if prefix_gru is None or embed_proj is None:
raise RuntimeError("DFLASH Domino projector modules are unavailable.")
draft_next = domino_greedy_rollout(
draft_hidden=draft_hidden,
bonus_tokens=block_ids[:, 0],
target_embedding=embed_module,
lm_head_weight=lm_head.weight,
prefix_gru=prefix_gru,
embed_proj=embed_proj,
vocab_size=int(self.model_runner.model_config.vocab_size),
shift_label=bool(self.draft_model.shift_label),
candidate_pool_size=self.domino_candidate_pool_size,
)
elif self._draft_sampler is not None and draft_out.can_run_graph:
draft_next = self._draft_sampler.out[
: bs * (int(self.block_size) - 1)
].view(bs, int(self.block_size) - 1)
@@ -0,0 +1,212 @@
from __future__ import annotations
import torch
import torch.nn.functional as F
from torch import nn
def _domino_gru_cell(
prefix_gru: nn.GRU, input: torch.Tensor, hidden: torch.Tensor
) -> torch.Tensor:
"""Run one feedback token without cuDNN's per-call weight packing."""
return torch.ops.aten.gru_cell.default(
input,
hidden,
prefix_gru.weight_ih_l0,
prefix_gru.weight_hh_l0,
prefix_gru.bias_ih_l0 if prefix_gru.bias else None,
prefix_gru.bias_hh_l0 if prefix_gru.bias else None,
)
def validate_domino_runtime(
*,
device: torch.device,
tp_size: int,
target_vocab_size: int,
draft_vocab_size: int,
hidden_size: int,
target_embedding: nn.Module,
lm_head: nn.Module,
prefix_gru: nn.GRU,
embed_proj: nn.Sequential,
) -> None:
"""Validate the deliberately narrow correctness-first Domino runtime."""
if device.type != "cuda":
raise ValueError(f"DFLASH Domino currently requires CUDA, got {device}.")
if int(tp_size) != 1:
raise ValueError(f"DFLASH Domino currently requires TP=1, got TP={tp_size}.")
if int(target_vocab_size) != int(draft_vocab_size):
raise ValueError(
"DFLASH Domino requires identical target and draft vocab sizes, "
f"got target={target_vocab_size}, draft={draft_vocab_size}."
)
embedding_weight = getattr(target_embedding, "weight", None)
lm_head_weight = getattr(lm_head, "weight", None)
if embedding_weight is None or lm_head_weight is None:
raise ValueError(
"DFLASH Domino requires target embedding and lm_head weight tensors."
)
shard = getattr(lm_head, "shard_indices", None)
if shard is not None:
if int(shard.num_added_elements) != 0:
raise ValueError(
"DFLASH Domino does not support added-vocab lm_head shards."
)
if int(shard.org_vocab_start_index) != 0 or int(shard.num_org_elements) != int(
target_vocab_size
):
raise ValueError(
"DFLASH Domino requires the complete target vocabulary on TP=1."
)
elif int(lm_head_weight.shape[0]) != int(target_vocab_size):
raise ValueError(
"DFLASH Domino lm_head row count must equal the target vocab size, "
f"got rows={int(lm_head_weight.shape[0])}, vocab={target_vocab_size}."
)
if int(embedding_weight.shape[0]) < int(target_vocab_size):
raise ValueError(
"DFLASH Domino target embedding has fewer rows than the target vocab "
f"size: rows={int(embedding_weight.shape[0])}, vocab={target_vocab_size}."
)
if int(embedding_weight.shape[-1]) != int(hidden_size) or int(
lm_head_weight.shape[-1]
) != int(hidden_size):
raise ValueError(
"DFLASH Domino target embedding/lm_head hidden size does not match "
f"the draft hidden size {hidden_size}."
)
if int(prefix_gru.input_size) != int(hidden_size):
raise ValueError(
"DFLASH Domino GRU input size does not match the draft hidden size."
)
if int(embed_proj[0].in_features) != int(hidden_size + prefix_gru.hidden_size):
raise ValueError("DFLASH Domino projector input shape is inconsistent.")
if int(embed_proj[2].out_features) != int(target_vocab_size):
raise ValueError("DFLASH Domino projector output vocab size is inconsistent.")
weights = (
embedding_weight,
lm_head_weight,
prefix_gru.weight_ih_l0,
prefix_gru.weight_hh_l0,
embed_proj[0].weight,
embed_proj[2].weight,
)
non_bf16 = [
str(weight.dtype) for weight in weights if weight.dtype != torch.bfloat16
]
if non_bf16:
raise ValueError(
"DFLASH Domino currently requires BF16 target and projector weights; "
f"found {non_bf16}."
)
@torch.no_grad()
def domino_greedy_rollout(
*,
draft_hidden: torch.Tensor,
bonus_tokens: torch.Tensor,
target_embedding: nn.Module,
lm_head_weight: torch.Tensor,
prefix_gru: nn.GRU,
embed_proj: nn.Sequential,
vocab_size: int,
shift_label: bool,
candidate_pool_size: int,
) -> torch.Tensor:
"""Generate a Domino chain using one block-shared base-logit candidate pool."""
if draft_hidden.ndim != 3:
raise ValueError(
f"draft_hidden must have shape [batch, block, hidden], got {tuple(draft_hidden.shape)}."
)
batch_size, block_size, hidden_size = draft_hidden.shape
if bonus_tokens.shape != (batch_size,):
raise ValueError(
f"bonus_tokens must have shape ({batch_size},), got {tuple(bonus_tokens.shape)}."
)
num_proposals = int(block_size) - 1
if num_proposals < 1:
raise ValueError(f"Domino requires block_size > 1, got {block_size}.")
candidate_pool_size = int(candidate_pool_size)
if candidate_pool_size < 0:
raise ValueError(
"Domino candidate_pool_size must be non-negative, "
f"got {candidate_pool_size}."
)
candidate_pool_size = min(candidate_pool_size, int(vocab_size))
start = 0 if shift_label else 1
z = draft_hidden[:, start : start + num_proposals, :]
if int(z.shape[1]) != num_proposals:
raise ValueError(
"Domino draft hidden states do not contain enough proposal positions."
)
weight = lm_head_weight[: int(vocab_size)]
z_for_logits = z.to(weight.dtype) if z.dtype != weight.dtype else z
logits_input = (
z_for_logits.transpose(0, 1)
.contiguous()
.view(num_proposals * batch_size, hidden_size)
)
base_logits = F.linear(logits_input, weight).view(num_proposals, batch_size, -1)
first_ids = torch.argmax(base_logits[0], dim=-1).to(torch.long)
proposals = [first_ids]
if num_proposals == 1:
return first_ids[:, None]
candidate_ids = None
candidate_base = None
candidate_weight = None
if 0 < candidate_pool_size < int(vocab_size):
feedback_logits = base_logits[1:]
candidate_ids = torch.topk(
feedback_logits.amax(dim=0),
k=candidate_pool_size,
dim=-1,
sorted=False,
).indices.contiguous()
candidate_base = torch.gather(
feedback_logits.transpose(0, 1),
2,
candidate_ids[:, None, :].expand(-1, num_proposals - 1, -1),
).transpose(0, 1)
candidate_weight = F.embedding(candidate_ids, embed_proj[2].weight)
prefix_ids = torch.stack((bonus_tokens, first_ids), dim=1)
_, gru_hidden = prefix_gru(target_embedding(prefix_ids))
for index in range(1, num_proposals):
step_hidden = z[:, index, :]
correction_hidden = embed_proj[1](
embed_proj[0](torch.cat((step_hidden, gru_hidden[0]), dim=-1))
)
if candidate_ids is None:
correction = embed_proj[2](correction_hidden)
next_ids = torch.argmax(base_logits[index] + correction, dim=-1).to(
torch.long
)
else:
correction = torch.bmm(
candidate_weight, correction_hidden.unsqueeze(-1)
).squeeze(-1)
candidate_position = torch.argmax(
candidate_base[index - 1] + correction, dim=-1
)
next_ids = torch.gather(
candidate_ids, 1, candidate_position[:, None]
).squeeze(1)
proposals.append(next_ids)
if index + 1 < num_proposals:
gru_hidden = _domino_gru_cell(
prefix_gru, target_embedding(next_ids), gru_hidden[0]
)[None]
return torch.stack(proposals, dim=1)
@@ -0,0 +1,104 @@
import unittest
from pathlib import Path
from tempfile import NamedTemporaryFile
import requests
from sglang.test.ci.ci_register import register_cuda_ci
from sglang.test.kits.eval_accuracy_kit import GSM8KMixin
from sglang.test.test_utils import (
DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
DEFAULT_URL_FOR_TEST,
CustomTestCase,
kill_process_tree,
popen_launch_server,
)
register_cuda_ci(est_time=600, stage="base-b", runner_config="1-gpu-small")
class TestDFlashDominoFullVocab(CustomTestCase):
model = "Qwen/Qwen3-8B"
draft_model = "Huang2020/Qwen3-8B-Domino-b16"
candidate_pool_size = 0
@classmethod
def setUpClass(cls):
cls.base_url = DEFAULT_URL_FOR_TEST
cls.server_log = NamedTemporaryFile(mode="w+", suffix="-domino.log")
cls.addClassCleanup(cls.server_log.close)
print(f"Domino server log: {cls.server_log.name}", flush=True)
cls.process = popen_launch_server(
cls.model,
cls.base_url,
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
return_stdout_stderr=(cls.server_log, cls.server_log),
other_args=[
"--trust-remote-code",
"--dtype",
"bfloat16",
"--tp-size",
"1",
"--attention-backend",
"triton",
"--speculative-algorithm",
"DFLASH",
"--speculative-draft-model-path",
cls.draft_model,
"--speculative-domino-candidate-pool-size",
str(cls.candidate_pool_size),
"--speculative-draft-attention-backend",
"triton",
"--cuda-graph-backend-decode",
"full",
"--cuda-graph-max-bs-decode",
"64",
"--max-running-requests",
"64",
"--mem-fraction-static",
"0.7",
],
)
def test_domino_runtime(self):
response = requests.get(self.base_url + "/server_info", timeout=10)
response.raise_for_status()
state = response.json()["internal_states"][0]
self.assertEqual(state["speculative_num_draft_tokens"], 16)
self.assertFalse(state["disable_overlap_schedule"])
self.assertEqual(
state["speculative_domino_candidate_pool_size"], self.candidate_pool_size
)
log = Path(self.server_log.name).read_text()
self.assertIn(
"DFLASH Domino rollout enabled (BF16, TP=1, "
f"block-shared candidate pool size={self.candidate_pool_size}).",
log,
)
self.assertIn("Domino rollout folded into the draft cuda graph", log)
self.assertIn(
"Capture draft verify CUDA graph begin. backend=full, num_tokens_per_req=16,",
log,
)
self.assertIn(
"Capture target verify CUDA graph begin. backend=full, num_tokens_per_req=16,",
log,
)
@classmethod
def tearDownClass(cls):
if hasattr(cls, "process") and cls.process:
kill_process_tree(cls.process.pid)
class TestDFlashDomino(TestDFlashDominoFullVocab, GSM8KMixin):
gsm8k_score_threshold = 0.90
gsm8k_num_examples = 200
gsm8k_accept_length_thres = 4.0
gsm8k_num_threads = 128
gsm8k_num_shots = 5
candidate_pool_size = 2048
if __name__ == "__main__":
unittest.main()
@@ -0,0 +1,125 @@
import unittest
import torch
from torch import nn
from sglang.srt.speculative.dflash_worker_v2 import _DominoDraftSampler
from sglang.srt.speculative.domino_utils import _domino_gru_cell, domino_greedy_rollout
from sglang.test.ci.ci_register import register_cuda_ci
from sglang.test.test_utils import CustomTestCase
register_cuda_ci(est_time=10, stage="base-b-kernel-unit", runner_config="1-gpu-large")
@unittest.skipUnless(torch.cuda.is_available(), "CUDA is required")
class TestDFlashDominoRollout(CustomTestCase):
def setUp(self):
torch.manual_seed(0)
self.embedding = nn.Embedding(31, 8, device="cuda", dtype=torch.bfloat16)
self.prefix_gru = nn.GRU(8, 4, batch_first=True, bias=False).cuda().bfloat16()
self.embed_proj = (
nn.Sequential(
nn.Linear(12, 5, bias=False), nn.SiLU(), nn.Linear(5, 31, bias=False)
)
.cuda()
.bfloat16()
)
self.lm_head_weight = torch.randn(31, 8, device="cuda", dtype=torch.bfloat16)
self.hidden = torch.randn(3, 16, 8, device="cuda", dtype=torch.bfloat16)
self.bonus_tokens = torch.tensor([1, 4, 9], device="cuda")
def rollout(self, hidden, bonus_tokens, pool_size=5, shift_label=True):
return domino_greedy_rollout(
draft_hidden=hidden,
bonus_tokens=bonus_tokens,
target_embedding=self.embedding,
lm_head_weight=self.lm_head_weight,
prefix_gru=self.prefix_gru,
embed_proj=self.embed_proj,
vocab_size=31,
shift_label=shift_label,
candidate_pool_size=pool_size,
)
def test_gru_feedback_matches_sequence(self):
embeddings = self.embedding(torch.tensor([[1, 2, 3], [4, 5, 6]], device="cuda"))
_, expected = self.prefix_gru(embeddings)
state = torch.zeros(2, 4, device="cuda", dtype=torch.bfloat16)
for step in embeddings.unbind(dim=1):
state = _domino_gru_cell(self.prefix_gru, step, state)
torch.testing.assert_close(state, expected[0], rtol=0.02, atol=0.002)
def test_candidate_pool_boundaries(self):
for shift_label in (True, False):
with self.subTest(shift_label=shift_label):
full = self.rollout(self.hidden, self.bonus_tokens, 0, shift_label)
for pool_size in (31, 32):
actual = self.rollout(
self.hidden, self.bonus_tokens, pool_size, shift_label
)
torch.testing.assert_close(actual, full, rtol=0, atol=0)
first_hidden = self.hidden[:, 0 if shift_label else 1]
expected_first = (first_hidden @ self.lm_head_weight.T).argmax(dim=-1)
for block_size in (2, 16):
limited = self.rollout(
self.hidden[:, :block_size], self.bonus_tokens, 1, shift_label
)
self.assertEqual(limited.shape, (3, block_size - 1))
torch.testing.assert_close(limited[:, 0], expected_first)
if block_size > 2:
torch.testing.assert_close(
limited[:, 1:], limited[:, 1:2].expand_as(limited[:, 1:])
)
def test_batch_matches_individual_requests(self):
for pool_size in (0, 5):
with self.subTest(pool_size=pool_size):
batched = self.rollout(self.hidden, self.bonus_tokens, pool_size)
individual = torch.cat(
[
self.rollout(hidden[None], bonus[None], pool_size)
for hidden, bonus in zip(self.hidden, self.bonus_tokens)
]
)
torch.testing.assert_close(batched, individual, rtol=0, atol=0)
def test_sampler_replays_with_new_inputs(self):
sampler = _DominoDraftSampler(
target_embedding=self.embedding,
lm_head_weight=self.lm_head_weight,
prefix_gru=self.prefix_gru,
embed_proj=self.embed_proj,
vocab_size=31,
block_size=16,
shift_label=True,
max_bs=3,
candidate_pool_size=5,
)
block_ids = torch.zeros(3, 16, device="cuda", dtype=torch.long)
block_ids[:, 0].copy_(self.bonus_tokens)
def sample():
sampler(self.hidden.flatten(0, 1), block_ids.flatten())
warmup = torch.cuda.Stream()
warmup.wait_stream(torch.cuda.current_stream())
with torch.cuda.stream(warmup):
sample()
torch.cuda.current_stream().wait_stream(warmup)
graph = torch.cuda.CUDAGraph()
with torch.cuda.graph(graph):
sample()
for _ in range(2):
self.hidden.copy_(torch.randn_like(self.hidden))
block_ids[:, 0].copy_(torch.randint(31, (3,), device="cuda"))
expected = self.rollout(self.hidden, block_ids[:, 0])
graph.replay()
torch.cuda.synchronize()
torch.testing.assert_close(
sampler.out.view(3, 15), expected, rtol=0, atol=0
)
if __name__ == "__main__":
unittest.main()
@@ -0,0 +1,130 @@
import unittest
from types import SimpleNamespace
import torch
from torch import nn
from sglang.srt.models.dflash import DFlashDraftModel
from sglang.srt.speculative.dflash_utils import parse_dflash_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 _domino_config(**overrides):
dflash_config = {
"projector_type": "domino",
"mask_token_id": 29,
"shift_label": True,
"target_layer_ids": [1, 3],
"pure_draft_prefix_len": 1,
"gru_hidden_dim": 4,
"emb_dim": 5,
}
dflash_config.update(overrides.pop("dflash_config", {}))
fields = {
"num_hidden_layers": 2,
"num_target_layers": 4,
"block_size": 16,
"hidden_size": 8,
"vocab_size": 31,
"emb_dim": 5,
"dflash_config": dflash_config,
}
fields.update(overrides)
return SimpleNamespace(**fields)
def _projector_model(projector_type="domino"):
model = DFlashDraftModel.__new__(DFlashDraftModel)
nn.Module.__init__(model)
model.projector_type = projector_type
model.config = SimpleNamespace(hidden_size=8)
if projector_type == "domino":
model.prefix_gru = nn.GRU(8, 4, batch_first=True, bias=False)
model.embed_proj = nn.Sequential(
nn.Linear(12, 5, bias=False),
nn.SiLU(),
nn.Linear(5, 31, bias=False),
)
else:
model.prefix_gru = None
model.embed_proj = None
return model
def _projector_weights(model):
return {
"prefix_gru.weight_ih_l0": torch.randn_like(model.prefix_gru.weight_ih_l0),
"prefix_gru.weight_hh_l0": torch.randn_like(model.prefix_gru.weight_hh_l0),
"embed_proj.0.weight": torch.randn_like(model.embed_proj[0].weight),
"embed_proj.2.weight": torch.randn_like(model.embed_proj[2].weight),
}
class TestDFlashDominoConfig(CustomTestCase):
def test_top_level_emb_dim_fallback(self):
config = _domino_config()
del config.dflash_config["emb_dim"]
self.assertEqual(parse_dflash_draft_config(draft_hf_config=config).emb_dim, 5)
def test_invalid_domino_config_fails_fast(self):
cases = {
"shift_label": {"shift_label": 1},
"pure_draft_prefix_len": {"pure_draft_prefix_len": 2},
"gru_hidden_dim": {"gru_hidden_dim": None},
}
for expected, updates in cases.items():
with self.subTest(expected=expected):
with self.assertRaisesRegex(ValueError, expected):
parse_dflash_draft_config(
draft_hf_config=_domino_config(dflash_config=updates)
)
config = _domino_config(dflash_config={"emb_dim": None}, emb_dim=None)
with self.assertRaisesRegex(ValueError, "emb_dim"):
parse_dflash_draft_config(draft_hf_config=config)
with self.assertRaisesRegex(ValueError, "block_size > 1"):
parse_dflash_draft_config(draft_hf_config=_domino_config(block_size=1))
def test_conflicting_emb_dim_fails(self):
with self.assertRaisesRegex(ValueError, "emb_dim differs"):
parse_dflash_draft_config(draft_hf_config=_domino_config(emb_dim=6))
class TestDFlashDominoWeights(CustomTestCase):
def test_projector_weights_load_exactly(self):
model = _projector_model()
weights = _projector_weights(model)
model.load_weights(weights.items())
for name, expected in weights.items():
torch.testing.assert_close(
dict(model.named_parameters())[name], expected, rtol=0, atol=0
)
def test_each_required_projector_weight_is_checked(self):
for missing_name in _projector_weights(_projector_model()):
with self.subTest(missing_name=missing_name):
model = _projector_model()
weights = _projector_weights(model)
del weights[missing_name]
with self.assertRaisesRegex(ValueError, missing_name):
model.load_weights(weights.items())
def test_projector_shape_mismatch_fails(self):
model = _projector_model()
weights = _projector_weights(model)
weights["embed_proj.2.weight"] = torch.empty(30, 5)
with self.assertRaisesRegex(ValueError, "shape mismatch"):
model.load_weights(weights.items())
def test_projector_weights_require_domino_config(self):
model = _projector_model(projector_type="domnio")
with self.assertRaisesRegex(ValueError, "projector_type"):
model.load_weights([("prefix_gru.weight_ih_l0", torch.empty(12, 8))])
if __name__ == "__main__":
unittest.main()