feat: add optimized Domino rollout to DFlash V2 (#36899)
Co-authored-by: Qiaolin-Yu <liin1211@outlook.com>
This commit is contained in:
@@ -73,6 +73,10 @@ class Spec:
|
|||||||
Optional[int],
|
Optional[int],
|
||||||
"DFLASH only. Block size (verify window length). Alias of --speculative-num-draft-tokens for DFLASH.",
|
"DFLASH only. Block size (verify window length). Alias of --speculative-num-draft-tokens for DFLASH.",
|
||||||
] = None
|
] = 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[
|
speculative_dspark_block_size: A[
|
||||||
Optional[int],
|
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.",
|
"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.",
|
||||||
|
|||||||
@@ -635,6 +635,35 @@ class DFlashDraftModel(nn.Module):
|
|||||||
)
|
)
|
||||||
self.hidden_norm = RMSNorm(hidden_size, eps=rms_norm_eps)
|
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:
|
def set_block_size(self, block_size: int) -> None:
|
||||||
"""Adopt the block size the worker resolved.
|
"""Adopt the block size the worker resolved.
|
||||||
|
|
||||||
@@ -728,6 +757,7 @@ class DFlashDraftModel(nn.Module):
|
|||||||
]
|
]
|
||||||
|
|
||||||
params_dict = dict(self.named_parameters())
|
params_dict = dict(self.named_parameters())
|
||||||
|
loaded_params = set()
|
||||||
|
|
||||||
# Alias the native export's "encoder." names.
|
# Alias the native export's "encoder." names.
|
||||||
_VENDOR_ENCODER_ALIASES = {
|
_VENDOR_ENCODER_ALIASES = {
|
||||||
@@ -752,6 +782,14 @@ class DFlashDraftModel(nn.Module):
|
|||||||
return None
|
return None
|
||||||
|
|
||||||
for name, loaded_weight in weights:
|
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:
|
for param_name, weight_name, shard_id in stacked_params_mapping:
|
||||||
if f".{weight_name}." not in name:
|
if f".{weight_name}." not in name:
|
||||||
continue
|
continue
|
||||||
@@ -762,6 +800,7 @@ class DFlashDraftModel(nn.Module):
|
|||||||
param = params_dict[resolved_name]
|
param = params_dict[resolved_name]
|
||||||
weight_loader = getattr(param, "weight_loader", default_weight_loader)
|
weight_loader = getattr(param, "weight_loader", default_weight_loader)
|
||||||
weight_loader(param, loaded_weight, shard_id)
|
weight_loader(param, loaded_weight, shard_id)
|
||||||
|
loaded_params.add(resolved_name)
|
||||||
break
|
break
|
||||||
else:
|
else:
|
||||||
resolved_name = resolve_param_name(name)
|
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"(num_context_features={self.num_context_features}, hidden_size={int(self.config.hidden_size)}), "
|
||||||
f"but got {loaded_shape} for weight '{name}'."
|
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 = getattr(param, "weight_loader", default_weight_loader)
|
||||||
weight_loader(param, loaded_weight)
|
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):
|
class DFlashLagunaAttention(DFlashAttention):
|
||||||
|
|||||||
@@ -538,6 +538,15 @@ class DFlashDraftConfig:
|
|||||||
target_layer_ids: Optional[List[int]]
|
target_layer_ids: Optional[List[int]]
|
||||||
mask_token: str
|
mask_token: str
|
||||||
mask_token_id: Optional[int]
|
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:
|
def require_num_layers(self) -> int:
|
||||||
if self.num_hidden_layers is None:
|
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}."
|
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(
|
return DFlashDraftConfig(
|
||||||
num_hidden_layers=num_hidden_layers,
|
num_hidden_layers=num_hidden_layers,
|
||||||
num_target_layers=num_target_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,
|
target_layer_ids=parsed_target_layer_ids,
|
||||||
mask_token=mask_token,
|
mask_token=mask_token,
|
||||||
mask_token_id=mask_token_id,
|
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,
|
is_dflash_sampling_verify_available,
|
||||||
parse_dflash_draft_config,
|
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 (
|
from sglang.srt.speculative.draft_worker_common import (
|
||||||
build_block_pos_offsets,
|
build_block_pos_offsets,
|
||||||
build_draft_tp_worker,
|
build_draft_tp_worker,
|
||||||
@@ -280,6 +284,55 @@ class _SelectorDraftSampler:
|
|||||||
self.q_out[:bs].copy_(q_rows)
|
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):
|
class DFlashWorkerV2(BaseSpecWorker):
|
||||||
"""DFLASH speculative decoding worker (spec-v2).
|
"""DFLASH speculative decoding worker (spec-v2).
|
||||||
|
|
||||||
@@ -333,6 +386,36 @@ class DFlashWorkerV2(BaseSpecWorker):
|
|||||||
draft_config = parse_dflash_draft_config(
|
draft_config = parse_dflash_draft_config(
|
||||||
draft_hf_config=self.draft_model_runner.model_config.hf_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:
|
if get_spec().speculative_num_draft_tokens is None:
|
||||||
# Should not happen (ServerArgs should have inferred it), but keep a fallback.
|
# Should not happen (ServerArgs should have inferred it), but keep a fallback.
|
||||||
self.block_size = int(draft_config.resolve_block_size(default=16))
|
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.draft_model.set_block_size(self.block_size)
|
||||||
self.speculative_num_draft_tokens = int(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 = draft_config.mask_token
|
||||||
self._mask_token_id_override = draft_config.mask_token_id
|
self._mask_token_id_override = draft_config.mask_token_id
|
||||||
@@ -373,6 +461,11 @@ class DFlashWorkerV2(BaseSpecWorker):
|
|||||||
self.draft_window_size,
|
self.draft_window_size,
|
||||||
self.use_compact_draft_cache,
|
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(
|
logger.info(
|
||||||
"DFLASH draft runner ready. mask_token=%s, mask_token_id=%s, mask_token_id_override=%s, noise_embed_scale=%s",
|
"DFLASH draft runner ready. mask_token=%s, mask_token_id=%s, mask_token_id_override=%s, noise_embed_scale=%s",
|
||||||
self._mask_token,
|
self._mask_token,
|
||||||
@@ -661,6 +754,28 @@ class DFlashWorkerV2(BaseSpecWorker):
|
|||||||
# Quantized lm_head (FP8/INT) would break the static matmul.
|
# Quantized lm_head (FP8/INT) would break the static matmul.
|
||||||
return _eager("quantized lm_head")
|
return _eager("quantized lm_head")
|
||||||
tp_group = get_tp_group()
|
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 not hasattr(lm_head, "shard_indices"):
|
||||||
if tp_group.world_size != 1:
|
if tp_group.world_size != 1:
|
||||||
# No shard metadata to recover per-rank vocab offsets from.
|
# 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_out = self.draft_model_runner.forward(forward_batch)
|
||||||
draft_logits_output = draft_out.logits_output
|
draft_logits_output = draft_out.logits_output
|
||||||
|
|
||||||
folded = self._draft_sampler is not None and draft_out.can_run_graph
|
if (
|
||||||
if folded:
|
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[
|
draft_next = self._draft_sampler.out[
|
||||||
: bs * (int(self.block_size) - 1)
|
: bs * (int(self.block_size) - 1)
|
||||||
].view(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()
|
||||||
Reference in New Issue
Block a user