Support Gemma3/4 + Eagle3 (#23976)

This commit is contained in:
Charles Chen
2026-05-09 13:34:56 -07:00
committed by GitHub
parent 8087e07d52
commit 12f42f2e7e
5 changed files with 194 additions and 10 deletions
+68 -6
View File
@@ -12,7 +12,7 @@
# limitations under the License.
# ==============================================================================
import copy
from typing import Iterable, Optional, Set, Tuple
from typing import Iterable, List, Optional, Set, Tuple
import einops
import torch
@@ -24,7 +24,10 @@ from transformers import (
PreTrainedModel,
)
from sglang.srt.distributed import get_tensor_model_parallel_world_size
from sglang.srt.distributed import (
get_tensor_model_parallel_rank,
get_tensor_model_parallel_world_size,
)
from sglang.srt.layers.activation import GeluAndMul
from sglang.srt.layers.layernorm import Gemma3RMSNorm
from sglang.srt.layers.linear import (
@@ -599,6 +602,7 @@ class Gemma3TextModel(PreTrainedModel):
prefix=add_prefix("layers", prefix),
)
self.norm = Gemma3RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
self.layers_to_capture = []
self.post_init()
def forward(
@@ -614,8 +618,13 @@ class Gemma3TextModel(PreTrainedModel):
else:
hidden_states = input_embeds
aux_hidden_states = []
num_layers = len(self.layers)
if _is_cpu and _is_cpu_amx_available:
for layer in self.layers:
for i, layer in enumerate(self.layers):
if i in self.layers_to_capture:
aux_hidden_states.append(hidden_states)
layer_outputs = layer(
positions=positions,
position_embeddings_global=None,
@@ -631,7 +640,9 @@ class Gemma3TextModel(PreTrainedModel):
position_embeddings_global = self.rotary_emb(hidden_states, positions)
position_embeddings_local = self.rotary_emb_local(hidden_states, positions)
for layer in self.layers:
for i, layer in enumerate(self.layers):
if i in self.layers_to_capture:
aux_hidden_states.append(hidden_states)
layer_outputs = layer(
positions=positions,
position_embeddings_global=position_embeddings_global,
@@ -642,9 +653,18 @@ class Gemma3TextModel(PreTrainedModel):
)
hidden_states = layer_outputs[0]
# Capture the output of the last layer if requested.
# layers_to_capture uses +1 offset (captures input of layer i = output of i-1),
# so index num_layers means the output of the final layer.
if num_layers in self.layers_to_capture:
aux_hidden_states.append(hidden_states)
hidden_states = self.norm(hidden_states)
return hidden_states
if len(aux_hidden_states) == 0:
return hidden_states
return hidden_states, aux_hidden_states
class Gemma3ForCausalLM(PreTrainedModel):
@@ -722,6 +742,7 @@ class Gemma3ForCausalLM(PreTrainedModel):
quant_config=quant_config,
prefix=add_prefix("lm_head", prefix),
)
self.capture_aux_hidden_states = False
self.post_init()
def get_input_embeddings(self) -> nn.Embedding:
@@ -746,8 +767,16 @@ class Gemma3ForCausalLM(PreTrainedModel):
input_ids, positions, forward_batch, input_embeds, **kwargs
)
aux_hidden_states = None
if self.capture_aux_hidden_states:
hidden_states, aux_hidden_states = hidden_states
return self.logits_processor(
input_ids, hidden_states, self.model.embed_tokens, forward_batch
input_ids,
hidden_states,
self.model.embed_tokens,
forward_batch,
aux_hidden_states,
)
@torch.no_grad()
@@ -862,5 +891,38 @@ class Gemma3ForCausalLM(PreTrainedModel):
# )
return loaded_params
def set_eagle3_layers_to_capture(self, layer_ids: Optional[List[int]] = None):
if layer_ids is None:
self.capture_aux_hidden_states = True
num_layers = self.config.num_hidden_layers
self.model.layers_to_capture = [2, num_layers // 2, num_layers - 3]
else:
self.capture_aux_hidden_states = True
# we plus 1 here because in sglang, for the ith layer, it takes the output
# of the (i-1)th layer as aux hidden state
self.model.layers_to_capture = [val + 1 for val in layer_ids]
def _shard_weight(self, weight: torch.Tensor) -> torch.Tensor:
"""Shard a full embedding/lm_head weight along vocab dim for the current TP rank.
Gemma3 uses nn.Embedding (unsharded) but the Eagle3 draft model uses
VocabParallelEmbedding (sharded). This method extracts the correct
shard so the weights can be shared.
"""
tp_size = get_tensor_model_parallel_world_size()
if tp_size <= 1:
return weight
tp_rank = get_tensor_model_parallel_rank()
shard_size = (weight.shape[0] + tp_size - 1) // tp_size
return weight[tp_rank * shard_size : (tp_rank + 1) * shard_size]
def get_embed(self):
return self._shard_weight(self.model.embed_tokens.weight)
def get_embed_and_head(self):
embed = self._shard_weight(self.model.embed_tokens.weight)
head = self._shard_weight(self.lm_head.weight)
return embed, head
EntryClass = Gemma3ForCausalLM
+16
View File
@@ -480,5 +480,21 @@ class Gemma3ForConditionalGeneration(PreTrainedModel):
# f"Some weights are not initialized from checkpoints: {unloaded_params}")
return loaded_params
def get_embed_and_head(self):
# For EAGLE3, we delegate to the language model which should have this method
# If the language model doesn't have lm_head (like EAGLE3), we return None for head
embed = self.language_model.get_embed()
if hasattr(self.language_model, "get_embed_and_head"):
return self.language_model.get_embed_and_head()
elif hasattr(self.language_model, "lm_head"):
return embed, self.language_model.lm_head.weight
else:
# For EAGLE3, head might not be needed
return embed, None
def set_eagle3_layers_to_capture(self, layer_ids: Optional[List[int]] = None):
if hasattr(self.language_model, "set_eagle3_layers_to_capture"):
self.language_model.set_eagle3_layers_to_capture(layer_ids)
EntryClass = Gemma3ForConditionalGeneration
+60 -3
View File
@@ -14,7 +14,7 @@
import logging
import re
from typing import Iterable, Optional, Set, Tuple
from typing import Iterable, List, Optional, Set, Tuple
import torch
from torch import nn
@@ -25,6 +25,7 @@ from transformers import (
)
from sglang.srt.distributed import (
get_tensor_model_parallel_rank,
get_tensor_model_parallel_world_size,
)
from sglang.srt.layers.gemma4_fused_ops import (
@@ -689,6 +690,7 @@ class Gemma4TextModel(PreTrainedModel):
)
self.norm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
self.layers_to_capture = []
self.post_init()
def get_input_embeddings(self) -> nn.Embedding:
@@ -786,7 +788,13 @@ class Gemma4TextModel(PreTrainedModel):
hidden_states = input_embeds
aux_hidden_states = []
num_layers = len(self.layers)
for layer_idx, layer in enumerate(self.layers):
if layer_idx in self.layers_to_capture:
aux_hidden_states.append(hidden_states)
if per_layer_inputs is not None:
per_layer_input = per_layer_inputs[:, layer_idx, :]
else:
@@ -801,11 +809,21 @@ class Gemma4TextModel(PreTrainedModel):
hidden_states = layer_outputs[0]
residual = layer_outputs[1] if len(layer_outputs) > 1 else None
# Capture the output of the last layer if requested.
# layers_to_capture uses +1 offset, so num_layers means
# "output of the last layer" which is only available after the loop.
if num_layers in self.layers_to_capture:
aux_hidden_states.append(hidden_states)
if residual is None:
hidden_states = self.norm(hidden_states)
else:
hidden_states, _ = self.norm(hidden_states, residual)
return hidden_states
if len(aux_hidden_states) == 0:
return hidden_states
return hidden_states, aux_hidden_states
class Gemma4ForCausalLM(PreTrainedModel):
@@ -873,6 +891,7 @@ class Gemma4ForCausalLM(PreTrainedModel):
quant_config=quant_config,
prefix=add_prefix("lm_head", prefix),
)
self.capture_aux_hidden_states = False
self.post_init()
def get_input_embeddings(self) -> nn.Embedding:
@@ -905,8 +924,13 @@ class Gemma4ForCausalLM(PreTrainedModel):
per_layer_inputs,
**kwargs,
)
aux_hidden_states = None
if self.capture_aux_hidden_states:
hidden_states, aux_hidden_states = hidden_states
return self.logits_processor(
input_ids, hidden_states, self.lm_head, forward_batch
input_ids, hidden_states, self.lm_head, forward_batch, aux_hidden_states
)
def _get_k_eq_v_layers(self) -> set:
@@ -1035,5 +1059,38 @@ class Gemma4ForCausalLM(PreTrainedModel):
logger.log(level, "%s: %s", msg, names)
return loaded_params
def _shard_weight(self, weight: torch.Tensor) -> torch.Tensor:
"""Shard a full embedding/lm_head weight along vocab dim for the current TP rank.
Gemma4 uses nn.Embedding (unsharded) but the Eagle3 draft model uses
VocabParallelEmbedding (sharded). This method extracts the correct
shard so the weights can be shared.
"""
tp_size = get_tensor_model_parallel_world_size()
if tp_size <= 1:
return weight
tp_rank = get_tensor_model_parallel_rank()
shard_size = (weight.shape[0] + tp_size - 1) // tp_size
return weight[tp_rank * shard_size : (tp_rank + 1) * shard_size]
def get_embed(self):
return self._shard_weight(self.model.embed_tokens.weight)
def get_embed_and_head(self):
embed = self._shard_weight(self.model.embed_tokens.weight)
head = self._shard_weight(self.lm_head.weight)
return embed, head
def set_eagle3_layers_to_capture(self, layer_ids: Optional[List[int]] = None):
if layer_ids is None:
self.capture_aux_hidden_states = True
num_layers = self.config.num_hidden_layers
self.model.layers_to_capture = [2, num_layers // 2, num_layers - 3]
else:
self.capture_aux_hidden_states = True
# we plus 1 here because in sglang, for the ith layer, it takes the output
# of the (i-1)th layer as aux hidden state
self.model.layers_to_capture = [val + 1 for val in layer_ids]
EntryClass = Gemma4ForCausalLM
+34 -1
View File
@@ -221,6 +221,7 @@ class Gemma4ForConditionalGeneration(PreTrainedModel):
# Create logits processor for the multimodal model
self.logits_processor = LogitsProcessor(config.text_config)
self.capture_aux_hidden_states = False
self.post_init()
@@ -594,9 +595,18 @@ class Gemma4ForConditionalGeneration(PreTrainedModel):
**kwargs,
)
# Unpack aux_hidden_states if Eagle3 capture is active
aux_hidden_states = None
if self.capture_aux_hidden_states:
hidden_states, aux_hidden_states = hidden_states
# Process hidden states through logits processor
return self.logits_processor(
input_ids, hidden_states, self.language_model.embed_tokens, forward_batch
input_ids,
hidden_states,
self.language_model.embed_tokens,
forward_batch,
aux_hidden_states,
)
def tie_weights(self, recompute_mapping=False):
@@ -899,5 +909,28 @@ class Gemma4ForConditionalGeneration(PreTrainedModel):
else:
raise NotImplementedError()
def get_embed(self):
return self.language_model.embed_tokens.weight
def get_embed_and_head(self):
embed = self.language_model.embed_tokens.weight
# Gemma4 ties word embeddings, so embed_tokens serves as lm_head
return embed, embed
def set_eagle3_layers_to_capture(self, layer_ids: Optional[List[int]] = None):
self.capture_aux_hidden_states = True
text_config = self.config.text_config
if layer_ids is None:
num_layers = text_config.num_hidden_layers
self.language_model.layers_to_capture = [
2,
num_layers // 2,
num_layers - 3,
]
else:
# we plus 1 here because in sglang, for the ith layer, it takes the output
# of the (i-1)th layer as aux hidden state
self.language_model.layers_to_capture = [val + 1 for val in layer_ids]
EntryClass = Gemma4ForConditionalGeneration
+16
View File
@@ -135,6 +135,15 @@ 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)
self.fc = torch.nn.Linear(
self.hidden_size_in * 3,
config.hidden_size,
@@ -174,6 +183,13 @@ 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)
hidden_states = self.fc(hidden_states)
# idle batch