feat: support TP>1 Domino rollout for DFlash V2 (#37069)

Co-authored-by: Qiaolin-Yu <liin1211@outlook.com>
This commit is contained in:
Francis
2026-09-11 18:51:11 -07:00
committed by GitHub
co-authored by Qiaolin-Yu
parent 0d1bea77da
commit e91c948057
4 changed files with 424 additions and 27 deletions
@@ -14,7 +14,7 @@ from sglang.test.test_utils import (
popen_launch_server,
)
register_cuda_ci(est_time=600, stage="base-b", runner_config="1-gpu-small")
register_cuda_ci(est_time=600, stage="base-b", runner_config="2-gpu-large")
class TestDFlashDominoFullVocab(CustomTestCase):
@@ -38,7 +38,7 @@ class TestDFlashDominoFullVocab(CustomTestCase):
"--dtype",
"bfloat16",
"--tp-size",
"1",
"2",
"--attention-backend",
"triton",
"--speculative-algorithm",
@@ -64,6 +64,7 @@ class TestDFlashDominoFullVocab(CustomTestCase):
response = requests.get(self.base_url + "/server_info", timeout=10)
response.raise_for_status()
state = response.json()["internal_states"][0]
self.assertEqual(state["tp_size"], 2)
self.assertEqual(state["speculative_num_draft_tokens"], 16)
self.assertFalse(state["disable_overlap_schedule"])
self.assertEqual(
@@ -71,11 +72,11 @@ class TestDFlashDominoFullVocab(CustomTestCase):
)
log = Path(self.server_log.name).read_text()
self.assertIn(
"DFLASH Domino rollout enabled (BF16, TP=1, "
"DFLASH Domino rollout enabled (BF16, TP=2, "
f"block-shared candidate pool size={self.candidate_pool_size}).",
log,
)
self.assertIn("Domino rollout folded into the draft cuda graph", log)
self.assertIn("Domino rollout folded into the draft cuda graph (tp=2)", log)
self.assertIn(
"Capture draft verify CUDA graph begin. backend=full, num_tokens_per_req=16,",
log,
@@ -6,6 +6,7 @@ from torch import nn
from sglang.srt.models.dflash import DFlashDraftModel
from sglang.srt.speculative.dflash_utils import parse_dflash_draft_config
from sglang.srt.speculative.domino_utils import validate_domino_runtime
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
@@ -126,5 +127,93 @@ class TestDFlashDominoWeights(CustomTestCase):
model.load_weights([("prefix_gru.weight_ih_l0", torch.empty(12, 8))])
class TestDFlashDominoRuntimeValidation(CustomTestCase):
def _modules(self, dtype=torch.bfloat16):
embedding = nn.Embedding(31, 8, dtype=dtype)
lm_head = nn.Linear(8, 31, bias=False, dtype=dtype)
prefix_gru = nn.GRU(8, 4, batch_first=True, bias=False, dtype=dtype)
embed_proj = nn.Sequential(
nn.Linear(12, 5, bias=False, dtype=dtype),
nn.SiLU(),
nn.Linear(5, 31, bias=False, dtype=dtype),
)
return embedding, lm_head, prefix_gru, embed_proj
def _tp2_modules(self):
embedding, lm_head, prefix_gru, embed_proj = self._modules()
embedding = nn.Embedding(16, 8, dtype=torch.bfloat16)
lm_head = nn.Linear(8, 16, bias=False, dtype=torch.bfloat16)
shard = SimpleNamespace(
num_added_elements=0,
org_vocab_start_index=0,
org_vocab_end_index=16,
num_org_elements=16,
num_org_elements_padded=16,
)
for module in (embedding, lm_head):
module.shard_indices = shard
module.org_vocab_size = 31
module.tp_size = 2
module.num_added_embeddings = 0
return embedding, lm_head, prefix_gru, embed_proj
def _validate(self, **overrides):
embedding, lm_head, prefix_gru, embed_proj = overrides.pop(
"modules", self._modules()
)
args = {
"device": torch.device("cuda"),
"tp_size": 1,
"tp_rank": 0,
"target_vocab_size": 31,
"draft_vocab_size": 31,
"hidden_size": 8,
"target_embedding": embedding,
"lm_head": lm_head,
"prefix_gru": prefix_gru,
"embed_proj": embed_proj,
}
args.update(overrides)
validate_domino_runtime(**args)
def test_tp_requires_vocab_shard_metadata(self):
with self.assertRaisesRegex(ValueError, "lm_head shard metadata"):
self._validate(tp_size=2)
def test_tp2_vocab_shards_supported(self):
self._validate(tp_size=2, modules=self._tp2_modules())
def test_tp2_incomplete_lm_head_shard_fails(self):
modules = self._tp2_modules()
modules[1].shard_indices = SimpleNamespace(
num_added_elements=0,
num_org_elements_padded=16,
)
with self.assertRaisesRegex(ValueError, "shard metadata is missing"):
self._validate(tp_size=2, modules=modules)
def test_tp_vocab_shard_must_match_rank(self):
modules = self._tp2_modules()
modules[1].shard_indices.org_vocab_start_index = 1
modules[1].shard_indices.org_vocab_end_index = 17
with self.assertRaisesRegex(ValueError, "does not match its TP rank"):
self._validate(tp_size=2, modules=modules)
def test_tp1_requires_complete_vocab_shard(self):
modules = self._modules()
modules[1].shard_indices = SimpleNamespace(
num_added_elements=0,
org_vocab_start_index=0,
org_vocab_end_index=30,
num_org_elements=30,
num_org_elements_padded=31,
)
modules[1].org_vocab_size = 31
modules[1].tp_size = 1
modules[1].num_added_embeddings = 0
with self.assertRaisesRegex(ValueError, "does not match its TP rank"):
self._validate(modules=modules)
if __name__ == "__main__":
unittest.main()