Reuse draft MTP indexer topk only when model config allows

EAGLE V2 draft MTP can avoid recomputing NSA/DSA indexer topk across internal draft iterations when the model declares that those indices are shareable. The port follows the upstream narrow contract: store per-forward topk on ForwardBatch, enable reuse only for topk=1, and clear the transient state after draft_forward.

Constraint: topk > 1 reorders hidden rows in select_top_k_tokens, so carried topk indices would no longer match the hidden states.

Rejected: Reuse target-side topk for draft | broader semantic change not covered by the upstream fix or current tests.

Rejected: Skip loading draft indexer weights | separate memory optimization with correctness risk for models that do not enable MTP index sharing.

Confidence: high

Scope-risk: narrow

Directive: Do not enable index_share_for_mtp_iteration without preserving the topk==1 guard and per-draft-forward cleanup.

Tested: python -m pytest test/registered/spec/eagle/test_eagle_v2_draft_extend_contract.py -q

Tested: python -m py_compile python/sglang/srt/speculative/eagle_worker_v2.py python/sglang/srt/models/deepseek_nextn.py python/sglang/srt/model_executor/forward_batch_info.py test/registered/spec/eagle/test_eagle_v2_draft_extend_contract.py

Not-tested: full GLM5 EAGLE throughput/accuracy run
This commit is contained in:
laoyao0822
2026-06-28 00:41:04 +08:00
parent 648a33ab30
commit df3cc9abf8
4 changed files with 188 additions and 0 deletions

View File

@@ -401,6 +401,9 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
capture_hidden_mode: CaptureHiddenMode = None
capture_draft_hidden_states: bool = False
draft_hidden_states: Optional[torch.Tensor] = None
# For NSA/DSA topk_indices reuse across forward calls, e.g. EAGLE draft MTP.
topk_indices: Optional[torch.Tensor] = None
reuse_mtp_topk_indices: bool = False
# For padding
padded_static_len: int = -1 # -1 if not padded

View File

@@ -530,7 +530,14 @@ class DeepseekModelNextN(nn.Module):
forward_batch,
residual,
zero_allocator,
prev_topk_indices=(
forward_batch.topk_indices
if forward_batch.reuse_mtp_topk_indices
else None
),
)
if forward_batch.reuse_mtp_topk_indices:
forward_batch.topk_indices = topk_indices
if not forward_batch.forward_mode.is_idle():
if residual is not None:

View File

@@ -152,6 +152,18 @@ class EagleDraftWorker(BaseDraftWorker):
# Alias for better readability
self.draft_runner = self.draft_worker.model_runner
# Reuse the first draft step's NSA/DSA indexer topk across the rest of
# the MTP iteration when the model config says it is safe. The reuse is
# only valid for topk == 1: select_top_k_tokens reorders rows for topk
# > 1, which would desynchronize carried indices from hidden states.
self.index_share_for_mtp_iteration = (
getattr(
self.draft_runner.model_config.hf_config,
"index_share_for_mtp_iteration",
False,
)
and self.topk == 1
)
self.eagle_use_aux_hidden_state = False
if self.speculative_algorithm.is_eagle3():
eagle_config = getattr(
@@ -441,6 +453,10 @@ class EagleDraftWorker(BaseDraftWorker):
token_list: List[torch.Tensor] = []
parents_list: List[torch.Tensor] = []
if self.index_share_for_mtp_iteration:
forward_batch.reuse_mtp_topk_indices = True
forward_batch.topk_indices = None
# Forward multiple steps
scores = None
for i in range(self.speculative_num_steps):
@@ -505,6 +521,10 @@ class EagleDraftWorker(BaseDraftWorker):
batch_size = parents_list[0].shape[0]
parent_list = torch.empty(batch_size, 0, device=parents_list[0].device)
if self.index_share_for_mtp_iteration:
forward_batch.topk_indices = None
forward_batch.reuse_mtp_topk_indices = False
return parent_list, top_scores_index, draft_tokens
def draft_extend(self):

View File

@@ -256,6 +256,164 @@ def _function_calls(func: ast.FunctionDef, name: str) -> int:
return count
def _find_class(tree: ast.AST, name: str) -> ast.ClassDef:
for node in ast.walk(tree):
if isinstance(node, ast.ClassDef) and node.name == name:
return node
raise AssertionError(f"class {name!r} not found")
def _find_method(class_node: ast.ClassDef, name: str) -> ast.FunctionDef:
for node in class_node.body:
if isinstance(node, ast.FunctionDef) and node.name == name:
return node
raise AssertionError(f"method {name!r} not found in class {class_node.name!r}")
def _assigned_self_attrs(func: ast.FunctionDef) -> set[str]:
attrs: set[str] = set()
for node in ast.walk(func):
if not isinstance(node, (ast.Assign, ast.AnnAssign, ast.AugAssign)):
continue
targets = []
if isinstance(node, ast.Assign):
targets.extend(node.targets)
else:
targets.append(node.target)
for target in targets:
if (
isinstance(target, ast.Attribute)
and isinstance(target.value, ast.Name)
and target.value.id == "self"
):
attrs.add(target.attr)
return attrs
def test_eagle_v2_draft_worker_enables_mtp_index_reuse_only_from_model_config():
"""MTP index reuse is model-config gated and only safe for topk=1.
select_top_k_tokens reorders rows when topk > 1, so reusing a previous
draft-step NSA/DSA index across later MTP iterations would desynchronize
hidden rows and indices unless topk is exactly one.
"""
tree = _parse_module("python/sglang/srt/speculative/eagle_worker_v2.py")
cls = _find_class(tree, "EagleDraftWorker")
init = _find_method(cls, "__init__")
assert "index_share_for_mtp_iteration" in _assigned_self_attrs(init)
assign = next(
node
for node in ast.walk(init)
if isinstance(node, ast.Assign)
and any(
isinstance(target, ast.Attribute)
and isinstance(target.value, ast.Name)
and target.value.id == "self"
and target.attr == "index_share_for_mtp_iteration"
for target in node.targets
)
)
text = ast.unparse(assign.value)
assert "index_share_for_mtp_iteration" in text
assert "self.topk == 1" in text
def test_eagle_v2_draft_forward_scopes_mtp_index_reuse_to_one_draft_forward():
"""The cached MTP topk indices are per draft_forward invocation.
The worker should clear stale topk indices before the first draft step and
after the final step, otherwise a later unrelated batch could reuse indices
from an older request.
"""
tree = _parse_module("python/sglang/srt/speculative/eagle_worker_v2.py")
cls = _find_class(tree, "EagleDraftWorker")
func = _find_method(cls, "draft_forward")
reuse_assigns = []
topk_clears = 0
for node in ast.walk(func):
if not isinstance(node, ast.Assign):
continue
for target in node.targets:
if (
isinstance(target, ast.Attribute)
and isinstance(target.value, ast.Name)
and target.value.id == "forward_batch"
):
if target.attr == "reuse_mtp_topk_indices":
reuse_assigns.append(node.value)
if (
target.attr == "topk_indices"
and isinstance(node.value, ast.Constant)
and node.value.value is None
):
topk_clears += 1
assert any(
isinstance(value, ast.Constant) and value.value is True
for value in reuse_assigns
)
assert any(
isinstance(value, ast.Constant) and value.value is False
for value in reuse_assigns
)
assert topk_clears >= 2
def test_deepseek_nextn_reuses_and_updates_mtp_topk_indices_when_requested():
"""NextN must pass cached MTP topk indices into the decoder and update them.
The attention layer already supports `prev_topk_indices`; the NextN wrapper
is responsible for threading the previous draft-step topk through
ForwardBatch only while EAGLE asks for reuse.
"""
tree = _parse_module("python/sglang/srt/models/deepseek_nextn.py")
cls = _find_class(tree, "DeepseekModelNextN")
func = _find_method(cls, "forward")
decoder_calls = []
stores_forward_batch_topk = False
for node in ast.walk(func):
if isinstance(node, ast.Call):
callee = node.func
if (
isinstance(callee, ast.Attribute)
and callee.attr == "decoder"
and isinstance(callee.value, ast.Name)
and callee.value.id == "self"
):
decoder_calls.append(node)
if not isinstance(node, ast.Assign):
continue
for target in node.targets:
if (
isinstance(target, ast.Attribute)
and target.attr == "topk_indices"
and isinstance(target.value, ast.Name)
and target.value.id == "forward_batch"
and isinstance(node.value, ast.Name)
and node.value.id == "topk_indices"
):
stores_forward_batch_topk = True
assert len(decoder_calls) == 1
prev_kwargs = [
kw.value
for kw in decoder_calls[0].keywords
if kw.arg == "prev_topk_indices"
]
assert len(prev_kwargs) == 1
assert "reuse_mtp_topk_indices" in ast.unparse(prev_kwargs[0])
assert "forward_batch.topk_indices" in ast.unparse(prev_kwargs[0])
assert stores_forward_batch_topk
def test_eagle_v2_verify_records_rebound_tensors_across_streams():
"""Spec v2 verify rebinds tensors while forward kernels still use them.