Eagle3 DP attention for Qwen3 MoE (#12002)

This commit is contained in:
Rain H
2025-10-29 20:25:17 +08:00
committed by GitHub
parent 42f8ea4030
commit 750940ae36
9 changed files with 219 additions and 27 deletions
+23 -1
View File
@@ -15,7 +15,7 @@
from dataclasses import dataclass
from enum import Enum, auto
from functools import partial
from typing import Dict, Optional
from typing import Dict, List, Optional
import torch
@@ -216,6 +216,28 @@ class LayerCommunicator:
get_global_server_args().speculative_algorithm
)
def prepare_attn_and_capture_last_layer_outputs(
self,
hidden_states: torch.Tensor,
residual: torch.Tensor,
forward_batch: ForwardBatch,
captured_last_layer_outputs: Optional[List[torch.Tensor]] = None,
):
hidden_states, residual = self.prepare_attn(
hidden_states, residual, forward_batch
)
if captured_last_layer_outputs is not None:
gathered_last_layer_output = self._communicate_simple_fn(
hidden_states=residual,
forward_batch=forward_batch,
context=self._context,
)
if gathered_last_layer_output is residual:
# Clone to avoid modifying the original residual by Custom RMSNorm inplace operation
gathered_last_layer_output = residual.clone()
captured_last_layer_outputs.append(gathered_last_layer_output)
return hidden_states, residual
def prepare_attn(
self,
hidden_states: torch.Tensor,
+11 -1
View File
@@ -19,6 +19,7 @@ from sglang.srt.utils import add_prefix
# https://github.com/SafeAILab/EAGLE/blob/main/eagle/model/cnets.py
"""Inference-only LLaMA-EAGLE model compatible with HuggingFace weights."""
import copy
from typing import Iterable, Optional, Tuple
import torch
@@ -161,6 +162,10 @@ class LlamaModel(nn.Module):
if hidden_states.shape[-1] != embeds.shape[-1]:
hidden_states = self.fc(hidden_states)
# idle batch
if hidden_states.shape[0] == 0:
return hidden_states, [hidden_states]
residual = None
hidden_states, residual = self.midlayer(
positions,
@@ -212,7 +217,12 @@ class LlamaForCausalLMEagle3(LlamaForCausalLM):
prefix=add_prefix("lm_head", prefix),
)
self.logits_processor = LogitsProcessor(config)
config_ = copy.deepcopy(config)
config_.vocab_size = (
config_.draft_vocab_size
) # draft logits processor has it's own vocab size
self.logits_processor = LogitsProcessor(config_)
self.capture_aux_hidden_states = True
self.hot_token_id = None
+30 -15
View File
@@ -473,10 +473,16 @@ class Qwen2MoeDecoderLayer(nn.Module):
hidden_states: torch.Tensor,
forward_batch: ForwardBatch,
residual: Optional[torch.Tensor],
captured_last_layer_outputs: Optional[List[torch.Tensor]] = None,
) -> Tuple[torch.Tensor, torch.Tensor]:
hidden_states, residual = self.layer_communicator.prepare_attn(
hidden_states, residual, forward_batch
hidden_states, residual = (
self.layer_communicator.prepare_attn_and_capture_last_layer_outputs(
hidden_states,
residual,
forward_batch,
captured_last_layer_outputs=captured_last_layer_outputs,
)
)
if hidden_states.shape[0] != 0:
@@ -553,6 +559,11 @@ class Qwen2MoeModel(nn.Module):
# For EAGLE3 support
self.layers_to_capture = []
def set_eagle3_layers_to_capture(self, layers_to_capture: List[int]):
self.layers_to_capture = layers_to_capture
for layer_id in self.layers_to_capture:
setattr(self.layers[layer_id], "_is_layer_to_capture", True)
def forward(
self,
input_ids: torch.Tensor,
@@ -585,12 +596,6 @@ class Qwen2MoeModel(nn.Module):
)
else:
for i in range(self.start_layer, self.end_layer):
if i in self.layers_to_capture:
aux_hidden_states.append(
hidden_states + residual
if residual is not None
else hidden_states
)
ctx = (
nullcontext()
if get_global_server_args().enable_piecewise_cuda_graph
@@ -599,7 +604,15 @@ class Qwen2MoeModel(nn.Module):
with ctx:
layer = self.layers[i]
hidden_states, residual = layer(
positions, hidden_states, forward_batch, residual
positions,
hidden_states,
forward_batch,
residual,
captured_last_layer_outputs=(
aux_hidden_states
if getattr(layer, "_is_layer_to_capture", False)
else None
),
)
if not self.pp_group.is_last_rank:
return PPProxyTensors(
@@ -830,13 +843,15 @@ class Qwen2MoeForCausalLM(nn.Module):
self.capture_aux_hidden_states = True
if layer_ids is None:
num_layers = self.config.num_hidden_layers
self.model.layers_to_capture = [
2,
num_layers // 2,
num_layers - 3,
] # Specific layers for EAGLE3 support
self.model.set_eagle3_layers_to_capture(
[
2,
num_layers // 2,
num_layers - 3,
]
) # Specific layers for EAGLE3 support
else:
self.model.layers_to_capture = [val + 1 for val in layer_ids]
self.model.set_eagle3_layers_to_capture([val + 1 for val in layer_ids])
EntryClass = Qwen2MoeForCausalLM
+16 -8
View File
@@ -537,10 +537,16 @@ class Qwen3MoeDecoderLayer(nn.Module):
hidden_states: torch.Tensor,
forward_batch: ForwardBatch,
residual: Optional[torch.Tensor],
captured_last_layer_outputs: Optional[List[torch.Tensor]] = None,
) -> Tuple[torch.Tensor, torch.Tensor]:
hidden_states, residual = self.layer_communicator.prepare_attn(
hidden_states, residual, forward_batch
hidden_states, residual = (
self.layer_communicator.prepare_attn_and_capture_last_layer_outputs(
hidden_states,
residual,
forward_batch,
captured_last_layer_outputs=captured_last_layer_outputs,
)
)
if hidden_states.shape[0] != 0:
@@ -772,13 +778,15 @@ class Qwen3MoeForCausalLM(nn.Module):
self.capture_aux_hidden_states = True
if layer_ids is None:
num_layers = self.config.num_hidden_layers
self.model.layers_to_capture = [
2,
num_layers // 2,
num_layers - 3,
] # Specific layers for EAGLE3 support
self.model.set_eagle3_layers_to_capture(
[
2,
num_layers // 2,
num_layers - 3,
]
) # Specific layers for EAGLE3 support
else:
self.model.layers_to_capture = [val + 1 for val in layer_ids]
self.model.set_eagle3_layers_to_capture([val + 1 for val in layer_ids])
def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]):
stacked_params_mapping = [
+1 -1
View File
@@ -822,7 +822,7 @@ class ServerArgs:
capture_bs = (
list(range(1, 9, 1))
+ list(range(10, 33, 2))
+ list(range(40, 64, 4))
+ list(range(40, 65, 4))
+ list(range(72, 257, 8))
+ list(range(272, self.cuda_graph_max_bs + 1, 16))
)
@@ -5,6 +5,7 @@ from typing import List, Optional, Tuple
import torch
from sglang.srt.distributed import get_tp_group
from sglang.srt.layers.dp_attention import get_attention_tp_group
from sglang.srt.layers.logits_processor import LogitsProcessorOutput
from sglang.srt.layers.sampler import get_token_ids_logprobs, get_top_logprobs
from sglang.srt.managers.schedule_batch import ScheduleBatch
@@ -117,7 +118,11 @@ class EAGLEWorker(TpModelWorker):
self.hot_token_id = None
# Init draft worker
with empty_context():
if server_args.enable_dp_attention and self.speculative_algorithm.is_eagle3():
ctx = draft_tp_context(get_attention_tp_group())
else:
ctx = empty_context()
with ctx:
super().__init__(
server_args=server_args,
gpu_id=gpu_id,
+2
View File
@@ -84,6 +84,8 @@ DEFAULT_MODEL_NAME_FOR_TEST_AWQ_INT4 = (
DEFAULT_EAGLE_TARGET_MODEL_FOR_TEST = "meta-llama/Llama-2-7b-chat-hf"
DEFAULT_EAGLE_DRAFT_MODEL_FOR_TEST = "lmsys/sglang-EAGLE-llama2-chat-7B"
DEFAULT_EAGLE_TARGET_MODEL_FOR_TEST_EAGLE3 = "meta-llama/Llama-3.1-8B-Instruct"
DEFAULT_EAGLE_DP_ATTENTION_TARGET_MODEL_FOR_TEST = "Qwen/Qwen3-30B-A3B"
DEFAULT_EAGLE_DP_ATTENTION_DRAFT_MODEL_FOR_TEST = "Tengyunw/qwen3_30b_moe_eagle3"
DEFAULT_MODEL_NAME_FOR_TEST_EAGLE3 = "lmsys/sglang-EAGLE3-LLaMA3.1-Instruct-8B"
DEFAULT_STANDALONE_SPECULATIVE_TARGET_MODEL_FOR_TEST = (
"meta-llama/Llama-3.1-8B-Instruct"