From a75b79e03bf867c775bc187667d16218c289fc10 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Do=C4=9Fa=C3=A7=20Eldenk?= Date: Mon, 11 May 2026 20:17:20 -0500 Subject: [PATCH] Feat: Support newer EAGLE-3 drafters (#24663) Co-authored-by: hnyls2002 --- python/sglang/srt/models/llama_eagle3.py | 105 +++++++++++++------- python/sglang/srt/speculative/eagle_info.py | 23 +++-- 2 files changed, 87 insertions(+), 41 deletions(-) diff --git a/python/sglang/srt/models/llama_eagle3.py b/python/sglang/srt/models/llama_eagle3.py index 0bff63788..9cafa8923 100644 --- a/python/sglang/srt/models/llama_eagle3.py +++ b/python/sglang/srt/models/llama_eagle3.py @@ -50,9 +50,13 @@ class LlamaDecoderLayer(LlamaDecoderLayer): ) -> None: super().__init__(config, layer_id, quant_config, prefix) + # Input layer concats embeds + target_hidden before qkv (input dim 2x). + self.is_input_layer = layer_id == 0 + hidden_size = 2 * self.hidden_size if self.is_input_layer else self.hidden_size + # override qkv self.self_attn.qkv_proj = QKVParallelLinear( - 2 * self.hidden_size, + hidden_size, self.self_attn.head_dim, self.self_attn.total_num_heads, self.self_attn.total_num_kv_heads, @@ -81,11 +85,16 @@ class LlamaDecoderLayer(LlamaDecoderLayer): residual: Optional[torch.Tensor], ) -> Tuple[torch.Tensor, torch.Tensor]: - residual = hidden_states - embeds = self.input_layernorm(embeds) - hidden_states = self.hidden_norm(hidden_states) + if self.is_input_layer: + # Input layer consumes target hidden states; no carried residual to fuse. + residual = hidden_states + hidden_states = self.hidden_norm(hidden_states) + embeds = self.input_layernorm(embeds) + hidden_states = torch.cat([embeds, hidden_states], dim=-1) + else: + # Fuse the previous layer's MLP residual add into hidden_norm. + hidden_states, residual = self.hidden_norm(hidden_states, residual) - hidden_states = torch.cat([embeds, hidden_states], dim=-1) # Self Attention hidden_states = self.self_attn( positions=positions, @@ -135,24 +144,42 @@ class LlamaModel(nn.Module): else: self.hidden_size_in = config.hidden_size - # Optional per-layer RMSNorm applied to each aux hidden state before - # concatenation, so that all three layers contribute equally regardless - # of their raw scale. Enabled via config "use_aux_norm": true. - self.use_aux_norm = getattr(config, "use_aux_norm", False) - if self.use_aux_norm: - self.aux_norm_low = RMSNorm(self.hidden_size_in, eps=config.rms_norm_eps) - self.aux_norm_mid = RMSNorm(self.hidden_size_in, eps=config.rms_norm_eps) - self.aux_norm_high = RMSNorm(self.hidden_size_in, eps=config.rms_norm_eps) + # num_aux resolution: explicit attr > eagle_config layer_ids > default 3. + self.num_aux_hidden_states = getattr(config, "num_aux_hidden_states", None) + if self.num_aux_hidden_states is None: + eagle_config = getattr(config, "eagle_config", None) or {} + layer_ids = eagle_config.get("eagle_aux_hidden_state_layer_ids") + self.num_aux_hidden_states = len(layer_ids) if layer_ids else 3 self.fc = torch.nn.Linear( - self.hidden_size_in * 3, + self.hidden_size_in * self.num_aux_hidden_states, config.hidden_size, bias=getattr(config, "bias", False), ) - self.midlayer = LlamaDecoderLayer(config, 0, quant_config, prefix) + # Per-aux RMSNorm before fc; enabled via `fc_norm` or legacy `use_aux_norm` flag. + use_fc_norm = getattr(config, "fc_norm", None) or getattr( + config, "use_aux_norm", False + ) + if use_fc_norm: + self.fc_norm = nn.ModuleList( + [ + RMSNorm(self.hidden_size_in, eps=config.rms_norm_eps) + for _ in range(self.num_aux_hidden_states) + ] + ) + else: + self.fc_norm = None + + self.layers = nn.ModuleList( + [ + LlamaDecoderLayer(config, i, quant_config, prefix) + for i in range(config.num_hidden_layers) + ] + ) self.norm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps) + self.norm_output = getattr(config, "norm_output", False) def forward( self, @@ -183,13 +210,12 @@ class LlamaModel(nn.Module): hidden_states = forward_batch.spec_info.hidden_states if hidden_states.shape[-1] != embeds.shape[-1]: - if self.use_aux_norm and hidden_states.shape[-1] == self.hidden_size_in * 3: - # Normalize each aux layer independently before fc projection. - h_low, h_mid, h_high = hidden_states.split(self.hidden_size_in, dim=-1) - h_low = self.aux_norm_low(h_low) - h_mid = self.aux_norm_mid(h_mid) - h_high = self.aux_norm_high(h_high) - hidden_states = torch.cat((h_low, h_mid, h_high), dim=-1) + if self.fc_norm is not None: + chunks = hidden_states.chunk(self.num_aux_hidden_states, dim=-1) + hidden_states = torch.cat( + [norm(chunk) for norm, chunk in zip(self.fc_norm, chunks)], + dim=-1, + ) hidden_states = self.fc(hidden_states) # idle batch @@ -197,20 +223,22 @@ class LlamaModel(nn.Module): return hidden_states, [hidden_states] residual = None - hidden_states, residual = self.midlayer( - positions, - embeds, - hidden_states, - forward_batch, - residual, - ) + for layer in self.layers: + hidden_states, residual = layer( + positions, + embeds, + hidden_states, + forward_batch, + residual, + ) hidden_states_to_logits, hidden_states_to_aux = self.norm( hidden_states, residual ) - # For draft decode, we capture the hidden state before norm - return hidden_states_to_logits, [hidden_states_to_aux] + # Draft decode captures pre-norm hidden by default; `norm_output` opts for normed. + aux = hidden_states_to_logits if self.norm_output else hidden_states_to_aux + return hidden_states_to_logits, [aux] class LlamaForCausalLMEagle3(LlamaForCausalLM): @@ -225,9 +253,6 @@ class LlamaForCausalLMEagle3(LlamaForCausalLM): self.quant_config = quant_config self.pp_group = get_pp_group() - if self.config.num_hidden_layers != 1: - raise ValueError("EAGLE3 currently only supports 1 layer") - self.model = LlamaModel( config, quant_config=quant_config, prefix=add_prefix("model", prefix) ) @@ -268,7 +293,19 @@ class LlamaForCausalLMEagle3(LlamaForCausalLM): (".gate_up_proj", ".up_proj", 1), ] + # Legacy weight names -> new module attribute names (backwards compat). + legacy_name_map = { + "midlayer": "layers.0", + "aux_norm_low": "fc_norm.0", + "aux_norm_mid": "fc_norm.1", + "aux_norm_high": "fc_norm.2", + } + for name, loaded_weight in weights: + for legacy, new in legacy_name_map.items(): + if legacy in name: + name = name.replace(legacy, new) + if "d2t" in name: # d2t stores diffs between draft id and target id self.hot_token_id = loaded_weight + torch.arange(loaded_weight.shape[0]) diff --git a/python/sglang/srt/speculative/eagle_info.py b/python/sglang/srt/speculative/eagle_info.py index 0cbe54483..d92a99dc3 100644 --- a/python/sglang/srt/speculative/eagle_info.py +++ b/python/sglang/srt/speculative/eagle_info.py @@ -846,17 +846,26 @@ class EagleDraftExtendInput(SpecInput): @classmethod def hidden_size_for(cls, worker) -> int: - """Extend-phase `hidden_states` width: target verify output (EAGLE - paper's "feature"). Widened to `target.hidden_size * 3` for EAGLE-3 - aux mode (low/mid/high features fused into a 3k-dim vector, reduced - by draft's FC).""" + """Extend-phase `hidden_states` width: target's `spec_hidden_size`, + widened to `num_aux * target_hidden` for EAGLE-3 aux mode.""" target_cfg = worker.target_worker.model_runner.model_config - if ( + if not ( worker.speculative_algorithm.is_eagle3() and worker.eagle_use_aux_hidden_state ): - return target_cfg.hidden_size * 3 - return target_cfg.spec_hidden_size + return target_cfg.spec_hidden_size + + hf_config = target_cfg.hf_config + + # `num_aux` resolution: explicit attr > eagle_config layer_ids > default 3. + num_aux = getattr(hf_config, "num_aux_hidden_states", None) + if num_aux is None: + eagle_config = getattr(hf_config, "eagle_config", None) or {} + layer_ids = eagle_config.get("eagle_aux_hidden_state_layer_ids") + num_aux = len(layer_ids) if layer_ids else 3 + + target_hidden = getattr(hf_config, "target_hidden_size", target_cfg.hidden_size) + return target_hidden * num_aux @classmethod def dtype_for(cls, worker) -> torch.dtype: