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
@@ -285,7 +285,7 @@ class _SelectorDraftSampler:
class _DominoDraftSampler:
"""Capture-safe TP=1 Domino rollout over a fixed-size draft block."""
"""Capture-safe Domino rollout over a fixed-size draft block."""
def __init__(
self,
@@ -299,6 +299,10 @@ class _DominoDraftSampler:
shift_label,
max_bs,
candidate_pool_size,
tp_group=None,
lm_head_org_vocab_start=0,
lm_head_num_org=None,
lm_head_num_org_padded=None,
):
self.target_embedding = target_embedding
self.lm_head_weight = lm_head_weight
@@ -308,6 +312,10 @@ class _DominoDraftSampler:
self.block_size = int(block_size)
self.shift_label = bool(shift_label)
self.candidate_pool_size = int(candidate_pool_size)
self.tp_group = tp_group
self.lm_head_org_vocab_start = int(lm_head_org_vocab_start)
self.lm_head_num_org = lm_head_num_org
self.lm_head_num_org_padded = lm_head_num_org_padded
max_tokens = int(max_bs) * (self.block_size - 1)
self.out = torch.empty(
(max_tokens,), dtype=torch.int64, device=lm_head_weight.device
@@ -329,6 +337,11 @@ class _DominoDraftSampler:
vocab_size=self.vocab_size,
shift_label=self.shift_label,
candidate_pool_size=self.candidate_pool_size,
tp_group=self.tp_group,
lm_head_org_vocab_start=self.lm_head_org_vocab_start,
lm_head_num_org=self.lm_head_num_org,
lm_head_num_org_padded=self.lm_head_num_org_padded,
prefer_tp_candidate_pool=bs > 1,
)
self.out[: bs * (self.block_size - 1)].copy_(proposals.reshape(-1))
@@ -408,6 +421,7 @@ class DFlashWorkerV2(BaseSpecWorker):
validate_domino_runtime(
device=torch.device(self.device),
tp_size=int(get_tp_group().world_size),
tp_rank=int(self.ps.tp_rank),
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),
@@ -463,7 +477,8 @@ class DFlashWorkerV2(BaseSpecWorker):
)
if self._is_domino:
logger.info(
"DFLASH Domino rollout enabled (BF16, TP=1, block-shared candidate pool size=%s).",
"DFLASH Domino rollout enabled (BF16, TP=%s, block-shared candidate pool size=%s).",
int(get_tp_group().world_size),
self.domino_candidate_pool_size,
)
logger.info(
@@ -755,16 +770,16 @@ class DFlashWorkerV2(BaseSpecWorker):
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)."
"DFLASH Domino rollout folded into the draft cuda graph (tp=%s).",
int(tp_group.world_size),
)
shard = getattr(lm_head, "shard_indices", None)
return _DominoDraftSampler(
target_embedding=target_model.get_input_embeddings(),
lm_head_weight=lm_head.weight,
@@ -775,6 +790,16 @@ class DFlashWorkerV2(BaseSpecWorker):
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,
tp_group=tp_group,
lm_head_org_vocab_start=(
int(shard.org_vocab_start_index) if shard is not None else 0
),
lm_head_num_org=(
int(shard.num_org_elements) if shard is not None else None
),
lm_head_num_org_padded=(
int(shard.num_org_elements_padded) if shard is not None else None
),
)
if not hasattr(lm_head, "shard_indices"):
if tp_group.world_size != 1:
@@ -2287,6 +2312,8 @@ class DFlashWorkerV2(BaseSpecWorker):
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.")
tp_group = get_tp_group()
shard = getattr(lm_head, "shard_indices", None)
draft_next = domino_greedy_rollout(
draft_hidden=draft_hidden,
bonus_tokens=block_ids[:, 0],
@@ -2297,6 +2324,16 @@ class DFlashWorkerV2(BaseSpecWorker):
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,
tp_group=tp_group,
lm_head_org_vocab_start=(
int(shard.org_vocab_start_index) if shard is not None else 0
),
lm_head_num_org=(
int(shard.num_org_elements) if shard is not None else None
),
lm_head_num_org_padded=(
int(shard.num_org_elements_padded) if shard is not None else None
),
)
elif self._draft_sampler is not None and draft_out.can_run_graph:
draft_next = self._draft_sampler.out[
+288 -18
View File
@@ -19,10 +19,113 @@ def _domino_gru_cell(
)
# Throughput policy for the logical full-vocabulary tensor. TP>1 uses the
# compact candidate pool when this gathered representation would be large.
_DOMINO_TP_FULL_BASE_LOGITS_MAX_BYTES = 96 * 1024 * 1024
def _domino_tp_first_ids(
local_logits: torch.Tensor,
*,
org_vocab_start: int,
num_org: int,
tp_group,
) -> torch.Tensor:
"""Select the full-vocab first argmax from contiguous vocab shards."""
local_max, local_arg = torch.max(local_logits[:, :num_org], dim=-1)
local_ids = local_arg.to(torch.int64) + int(org_vocab_start)
tp_size = int(tp_group.world_size)
batch_size = int(local_logits.shape[0])
gathered_max = torch.empty(
(tp_size * batch_size,), dtype=local_max.dtype, device=local_max.device
)
gathered_ids = torch.empty(
(tp_size * batch_size,), dtype=torch.int64, device=local_max.device
)
tp_group.all_gather_into_tensor(gathered_max, local_max.contiguous())
tp_group.all_gather_into_tensor(gathered_ids, local_ids.contiguous())
gathered_max = gathered_max.view(tp_size, batch_size)
gathered_ids = gathered_ids.view(tp_size, batch_size)
best_rank = torch.argmax(gathered_max, dim=0, keepdim=True)
return torch.gather(gathered_ids, 0, best_rank).squeeze(0)
def _domino_tp_candidate_state(
local_feedback_logits: torch.Tensor,
*,
candidate_pool_size: int,
org_vocab_start: int,
num_org: int,
tp_group,
) -> tuple[torch.Tensor, torch.Tensor]:
"""Build one strict global candidate pool from vocab-sharded base logits."""
num_steps, batch_size, local_vocab_size = local_feedback_logits.shape
tp_size = int(tp_group.world_size)
local_k = min(int(candidate_pool_size), int(local_vocab_size))
if int(candidate_pool_size) > tp_size * local_k:
raise ValueError("Domino TP shards do not cover the requested candidate pool.")
pool_scores = local_feedback_logits.amax(dim=0)
if num_org < local_vocab_size:
pool_scores[:, num_org:].fill_(float("-inf"))
local_scores, local_positions = torch.topk(
pool_scores, k=local_k, dim=-1, sorted=False
)
local_ids = local_positions.to(torch.int64) + int(org_vocab_start)
gathered_scores = torch.empty(
(tp_size * batch_size, local_k),
dtype=local_scores.dtype,
device=local_scores.device,
)
gathered_ids = torch.empty(
(tp_size * batch_size, local_k),
dtype=torch.int64,
device=local_scores.device,
)
tp_group.all_gather_into_tensor(gathered_scores, local_scores.contiguous())
tp_group.all_gather_into_tensor(gathered_ids, local_ids.contiguous())
gathered_scores = (
gathered_scores.view(tp_size, batch_size, local_k)
.permute(1, 0, 2)
.reshape(batch_size, tp_size * local_k)
)
gathered_ids = (
gathered_ids.view(tp_size, batch_size, local_k)
.permute(1, 0, 2)
.reshape(batch_size, tp_size * local_k)
)
global_positions = torch.topk(
gathered_scores,
k=int(candidate_pool_size),
dim=-1,
sorted=False,
).indices
candidate_ids = torch.gather(gathered_ids, 1, global_positions).contiguous()
owned = (candidate_ids >= int(org_vocab_start)) & (
candidate_ids < int(org_vocab_start + num_org)
)
local_positions = (candidate_ids - int(org_vocab_start)).clamp(
0, max(num_org - 1, 0)
)
candidate_base = torch.gather(
local_feedback_logits.transpose(0, 1),
2,
local_positions[:, None, :].expand(-1, num_steps, -1),
)
# Each global candidate is owned by exactly one vocab shard, so SUM
# reconstructs its base logit without gathering the full vocabulary.
candidate_base.masked_fill_(~owned[:, None, :], 0)
candidate_base = tp_group.all_reduce(candidate_base.contiguous())
return candidate_ids, candidate_base.transpose(0, 1)
def validate_domino_runtime(
*,
device: torch.device,
tp_size: int,
tp_rank: int,
target_vocab_size: int,
draft_vocab_size: int,
hidden_size: int,
@@ -34,8 +137,14 @@ def validate_domino_runtime(
"""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}.")
tp_size = int(tp_size)
if tp_size < 1:
raise ValueError(f"DFLASH Domino requires TP>=1, got TP={tp_size}.")
tp_rank = int(tp_rank)
if not 0 <= tp_rank < tp_size:
raise ValueError(
f"DFLASH Domino requires 0<=TP rank<TP size, got rank={tp_rank}, size={tp_size}."
)
if int(target_vocab_size) != int(draft_vocab_size):
raise ValueError(
"DFLASH Domino requires identical target and draft vocab sizes, "
@@ -49,25 +158,108 @@ def validate_domino_runtime(
"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:
lm_head_shard = getattr(lm_head, "shard_indices", None)
if lm_head_shard is not None:
if (
int(getattr(lm_head, "num_added_embeddings", 0)) != 0
or int(lm_head_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(
if int(getattr(lm_head, "org_vocab_size", target_vocab_size)) != int(
target_vocab_size
):
raise ValueError(
"DFLASH Domino requires the complete target vocabulary on TP=1."
"DFLASH Domino lm_head original vocab size does not match the target."
)
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(getattr(lm_head, "tp_size", tp_size)) != tp_size:
raise ValueError(
"DFLASH Domino lm_head TP size does not match the runtime TP size."
)
required_shard_fields = (
"org_vocab_start_index",
"org_vocab_end_index",
"num_org_elements",
"num_org_elements_padded",
)
missing_shard_fields = [
name for name in required_shard_fields if not hasattr(lm_head_shard, name)
]
if missing_shard_fields:
raise ValueError(
"DFLASH Domino lm_head shard metadata is missing: "
+ ", ".join(missing_shard_fields)
)
org_vocab_start = int(lm_head_shard.org_vocab_start_index)
org_vocab_end = int(lm_head_shard.org_vocab_end_index)
num_org = int(lm_head_shard.num_org_elements)
num_org_padded = int(lm_head_shard.num_org_elements_padded)
if (
num_org <= 0
or org_vocab_start < 0
or org_vocab_end != org_vocab_start + num_org
or org_vocab_end > int(target_vocab_size)
or num_org_padded < num_org
):
raise ValueError("DFLASH Domino lm_head original-vocab shard is invalid.")
if int(lm_head_weight.shape[0]) < num_org_padded:
raise ValueError(
"DFLASH Domino lm_head weight is smaller than its padded vocab shard."
)
expected_start = tp_rank * num_org_padded
expected_end = min(expected_start + num_org_padded, int(target_vocab_size))
if (
org_vocab_start != expected_start
or org_vocab_end != expected_end
or num_org != expected_end - expected_start
):
raise ValueError(
"DFLASH Domino lm_head vocab shard does not match its TP rank."
)
else:
if tp_size != 1:
raise ValueError("DFLASH Domino requires lm_head shard metadata for TP>1.")
if 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):
embedding_shard = getattr(target_embedding, "shard_indices", None)
if embedding_shard is not None:
if (
int(getattr(target_embedding, "num_added_embeddings", 0)) != 0
or int(embedding_shard.num_added_elements) != 0
):
raise ValueError(
"DFLASH Domino does not support added-vocab embedding shards."
)
if int(getattr(target_embedding, "org_vocab_size", target_vocab_size)) != int(
target_vocab_size
):
raise ValueError(
"DFLASH Domino embedding original vocab size does not match the target."
)
embedding_tp_size = int(getattr(target_embedding, "tp_size", tp_size))
if embedding_tp_size not in (1, tp_size):
raise ValueError(
"DFLASH Domino embedding TP size does not match the runtime TP size."
)
required_embedding_rows = (
int(target_vocab_size)
if embedding_tp_size == 1
else int(embedding_shard.num_org_elements_padded)
)
if int(embedding_weight.shape[0]) < required_embedding_rows:
raise ValueError(
"DFLASH Domino embedding weight is smaller than its padded vocab shard."
)
elif int(embedding_weight.shape[0]) < int(target_vocab_size):
if tp_size != 1:
raise ValueError(
"DFLASH Domino requires embedding shard metadata for TP>1."
)
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}."
@@ -119,6 +311,11 @@ def domino_greedy_rollout(
vocab_size: int,
shift_label: bool,
candidate_pool_size: int,
tp_group=None,
lm_head_org_vocab_start: int = 0,
lm_head_num_org: int | None = None,
lm_head_num_org_padded: int | None = None,
prefer_tp_candidate_pool: bool | None = None,
) -> torch.Tensor:
"""Generate a Domino chain using one block-shared base-logit candidate pool."""
if draft_hidden.ndim != 3:
@@ -148,24 +345,97 @@ def domino_greedy_rollout(
"Domino draft hidden states do not contain enough proposal positions."
)
weight = lm_head_weight[: int(vocab_size)]
tp_size = int(tp_group.world_size) if tp_group is not None else 1
if tp_size > 1 and (lm_head_num_org is None or lm_head_num_org_padded is None):
raise ValueError(
"Domino TP rollout requires local lm_head vocab shard metadata."
)
local_vocab_size = int(lm_head_num_org_padded) if tp_size > 1 else int(vocab_size)
if tp_size > 1:
num_org = int(lm_head_num_org)
org_vocab_start = int(lm_head_org_vocab_start)
if (
num_org <= 0
or num_org > local_vocab_size
or org_vocab_start < 0
or org_vocab_start + num_org > int(vocab_size)
):
raise ValueError(
"Domino TP rollout received an invalid lm_head vocab shard."
)
if int(lm_head_weight.shape[0]) < local_vocab_size:
raise ValueError(
"Domino lm_head weight is smaller than its padded vocab shard."
)
weight = lm_head_weight[:local_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)
local_logits = F.linear(logits_input, weight).view(
num_proposals, batch_size, local_vocab_size
)
if prefer_tp_candidate_pool is None:
full_base_logits_bytes = (
num_proposals * batch_size * int(vocab_size) * local_logits.element_size()
)
prefer_tp_candidate_pool = (
full_base_logits_bytes > _DOMINO_TP_FULL_BASE_LOGITS_MAX_BYTES
)
use_tp_candidate_pool = (
tp_size > 1
and prefer_tp_candidate_pool
and 0 < candidate_pool_size < int(vocab_size)
)
first_ids = None
candidate_ids = None
candidate_base = None
if tp_size == 1:
base_logits = local_logits[:, :, : int(vocab_size)]
elif use_tp_candidate_pool:
first_ids = _domino_tp_first_ids(
local_logits[0],
org_vocab_start=int(lm_head_org_vocab_start),
num_org=int(lm_head_num_org),
tp_group=tp_group,
)
if num_proposals > 1:
candidate_ids, candidate_base = _domino_tp_candidate_state(
local_logits[1:],
candidate_pool_size=candidate_pool_size,
org_vocab_start=int(lm_head_org_vocab_start),
num_org=int(lm_head_num_org),
tp_group=tp_group,
)
base_logits = None
else:
local_logits_t = local_logits.view(
num_proposals * batch_size, local_vocab_size
).T.contiguous()
gathered_logits = torch.empty(
(tp_size * local_vocab_size, num_proposals * batch_size),
dtype=local_logits.dtype,
device=local_logits.device,
)
tp_group.all_gather_into_tensor(gathered_logits, local_logits_t)
base_logits = (
gathered_logits.T[:, : int(vocab_size)]
.contiguous()
.view(num_proposals, batch_size, int(vocab_size))
)
first_ids = torch.argmax(base_logits[0], dim=-1).to(torch.long)
if first_ids is None:
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):
if candidate_ids is not None:
candidate_weight = F.embedding(candidate_ids, embed_proj[2].weight)
elif 0 < candidate_pool_size < int(vocab_size):
feedback_logits = base_logits[1:]
candidate_ids = torch.topk(
feedback_logits.amax(dim=0),