model: support EmbeddingGemma (#32375)

This commit is contained in:
Mick
2026-07-27 10:40:47 +08:00
committed by GitHub
parent a358374ae9
commit abb8f4b5e3
12 changed files with 209 additions and 16 deletions
+14 -1
View File
@@ -292,6 +292,7 @@ class ModelConfig:
)
)
self.hf_text_config = get_hf_text_config(self.hf_config)
self.is_embedding_gemma = is_embedding_gemma(self.hf_text_config)
rope_scaling = getattr(self.hf_text_config, "rope_parameters", None) or getattr(
self.hf_text_config, "rope_scaling", {}
@@ -399,7 +400,7 @@ class ModelConfig:
self.hf_text_config, "attention_chunk_size", None
)
self.sliding_window_size = self._get_sliding_window_size()
self.is_generation = is_generation_model(
self.is_generation = not self.is_embedding_gemma and is_generation_model(
self.hf_config.architectures, is_embedding
)
# The vision_config/audio_config attribute heuristic is only applied when
@@ -1668,6 +1669,18 @@ def _get_and_verify_dtype(
return torch_dtype
def is_embedding_gemma(config) -> bool:
"""Whether ``config`` is Google's bidirectional EmbeddingGemma checkpoint.
EmbeddingGemma uses the otherwise generative ``Gemma3TextModel``
architecture, so its model type alone is insufficient for dispatch. The
upstream ``use_bidirectional_attention`` flag is the defining distinction.
"""
return getattr(config, "model_type", None) == "gemma3_text" and getattr(
config, "use_bidirectional_attention", False
)
def is_generation_model(model_architectures: List[str], is_embedding: bool = False):
# We have two ways to determine whether a model is a generative model.
# 1. Check the model architecture
@@ -29,7 +29,12 @@ from sglang.srt.layers.dcp import (
from sglang.srt.layers.radix_attention import AttentionType
from sglang.srt.mem_cache.memory_pool import KVWriteLoc
from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool
from sglang.srt.model_executor.cuda_graph_config import cuda_graph_fully_disabled
from sglang.srt.model_executor.cuda_graph_config import (
Backend,
Phase,
check_cuda_graph_backend,
cuda_graph_fully_disabled,
)
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
from sglang.srt.runtime_context import get_parallel
from sglang.srt.speculative.spec_utils import (
@@ -223,7 +228,14 @@ class TritonAttnBackend(AttentionBackend):
self.use_pdl = False
self.allow_bidirectional_attention_in_extend = (
cuda_graph_fully_disabled()
# BCG captures one complete prefill forward. It is therefore safe
# for encoder-style attention, unlike the other CUDA graph modes
# that can split or pad requests. Eager prefill remains supported
# as before.
(
cuda_graph_fully_disabled()
or check_cuda_graph_backend(Phase.PREFILL, Backend.BREAKABLE)
)
and model_runner.server_args.chunked_prefill_size == -1
)
+11 -2
View File
@@ -20,6 +20,7 @@ if TYPE_CHECKING:
class PoolingType(IntEnum):
LAST = 0
CLS = 1
MEAN = 2
@dataclass
@@ -48,7 +49,7 @@ def pool_hidden_states(
hidden_states: torch.Tensor,
forward_batch: ForwardBatch,
) -> torch.Tensor:
"""Pool hidden_states by PoolingType (LAST/CLS).
"""Pool hidden_states by PoolingType (LAST/CLS/MEAN).
Raw pooling only — no normalize, no dim truncation.
Returns shape (batch_size, hidden_size).
@@ -61,6 +62,14 @@ def pool_hidden_states(
first_token_flat_indices = torch.zeros_like(prompt_lens)
first_token_flat_indices[1:] += torch.cumsum(prompt_lens, dim=0)[:-1]
return hidden_states[first_token_flat_indices]
elif pooling_type == PoolingType.MEAN:
prompt_lens = forward_batch.extend_seq_lens
end_indices = torch.cumsum(prompt_lens, dim=0) - 1
cumulative_hidden_states = torch.cumsum(hidden_states, dim=0)
sums = cumulative_hidden_states[end_indices]
preceding_sums = torch.zeros_like(sums)
preceding_sums[1:] = cumulative_hidden_states[end_indices[:-1]]
return (sums - preceding_sums) / prompt_lens.unsqueeze(-1)
else:
raise ValueError(f"Unsupported pooling type: {pooling_type}")
@@ -163,7 +172,7 @@ class Pooler(nn.Module):
2. Normalizes output if specified.
3. Returns structured results as `PoolerOutput`.
Attributes:
pooling_type: The type of pooling to use (LAST, AVERAGE, MAX).
pooling_type: The type of pooling to use (LAST, CLS, MEAN).
normalize: Whether to normalize the pooled data.
"""
+8 -2
View File
@@ -3529,18 +3529,24 @@ class Scheduler(
with self.forward_stream_ctx:
self.forward_stream.wait_stream(self.schedule_stream)
resolve_forward_inputs(batch, self.future_map)
pooler_output = self.tp_worker.forward_batch_embedding(batch)
pooler_output, can_run_cuda_graph = (
self.tp_worker.forward_batch_embedding(batch)
)
ret = EmbeddingBatchResult(
embeddings=pooler_output.embeddings,
pooled_hidden_states=pooler_output.pooled_hidden_states,
can_run_cuda_graph=can_run_cuda_graph,
)
ret.copy_to_cpu()
else:
resolve_forward_inputs(batch, self.future_map)
pooler_output = self.tp_worker.forward_batch_embedding(batch)
pooler_output, can_run_cuda_graph = (
self.tp_worker.forward_batch_embedding(batch)
)
ret = EmbeddingBatchResult(
embeddings=pooler_output.embeddings,
pooled_hidden_states=pooler_output.pooled_hidden_states,
can_run_cuda_graph=can_run_cuda_graph,
)
self._maybe_report_active_ranks()
+2 -2
View File
@@ -266,8 +266,8 @@ class BaseTpWorker(ABC):
self.model_runner,
return_hidden_states_before_norm=False,
)
output = self.model_runner.forward(forward_batch).logits_output
return output # Returns EmbeddingPoolerOutput
output = self.model_runner.forward(forward_batch)
return output.logits_output, output.can_run_graph
class TpModelWorker(BaseTpWorker):
+1 -4
View File
@@ -288,10 +288,7 @@ class EmbeddingBatchResult:
embeddings: torch.Tensor
pooled_hidden_states: Optional[torch.Tensor] = None
copy_done: Optional[torch.cuda.Event] = None
@property
def can_run_cuda_graph(self) -> bool:
return False
can_run_cuda_graph: bool = False
@torch.profiler.record_function("copy_embedding_to_cpu")
def copy_to_cpu(self):
+5
View File
@@ -196,6 +196,11 @@ def get_model_architecture(model_config: ModelConfig) -> Tuple[Type[nn.Module],
from sglang.srt.models.registry import ModelRegistry
architectures = getattr(model_config.hf_config, "architectures", [])
# EmbeddingGemma is serialized as Gemma3TextModel, which is also the name
# of the HF backbone. Route the bidirectional variant to SGLang's pooled
# embedding wrapper instead of falling back to the generic HF backend.
if getattr(model_config, "is_embedding_gemma", False):
architectures = ["EmbeddingGemmaModel"]
# Special handling for quantized Mixtral.
# FIXME(woosuk): This is a temporary hack.
mixtral_supported = [
+69 -2
View File
@@ -34,6 +34,7 @@ from sglang.srt.layers.linear import (
RowParallelLinear,
)
from sglang.srt.layers.logits_processor import LogitsProcessor
from sglang.srt.layers.pooler import EmbeddingPoolerOutput, Pooler, PoolingType
from sglang.srt.layers.quantization.base_config import QuantizationConfig
from sglang.srt.layers.radix_attention import AttentionType, RadixAttention
from sglang.srt.layers.rotary_embedding import apply_rotary_pos_emb, get_rope
@@ -224,7 +225,14 @@ class Gemma3Attention(nn.Module):
sliding_window_size=self.sliding_window,
quant_config=quant_config,
prefix=add_prefix("attn", prefix),
attn_type=AttentionType.DECODER_BIDIRECTIONAL,
# Gemma3 uses this attention implementation for both its causal
# LMs and EmbeddingGemma. Only the latter enables bidirectional
# prompt attention in its upstream config.
attn_type=(
AttentionType.DECODER_BIDIRECTIONAL
if getattr(config, "use_bidirectional_attention", False)
else AttentionType.DECODER
),
)
# Gemma3 adds normalization for q and k
@@ -934,4 +942,63 @@ class Gemma3ForCausalLM(PreTrainedModel):
return embed, head
EntryClass = Gemma3ForCausalLM
class EmbeddingGemmaModel(Gemma3ForCausalLM):
"""EmbeddingGemma's Gemma3 encoder with normalized mean pooling."""
def __init__(
self,
config: Gemma3TextConfig,
quant_config: Optional[QuantizationConfig] = None,
prefix: str = "",
) -> None:
# Do not initialize Gemma3ForCausalLM's unused LM head. Keeping the
# backbone under ``model`` also lets BCG capture only the transformer
# body and run this pooler as the eager tail.
PreTrainedModel.__init__(self, config=config)
self.config = config
self.quant_config = quant_config
self.model = Gemma3TextModel(
config, quant_config, prefix=add_prefix("model", prefix)
)
self.pooler = Pooler(pooling_type=PoolingType.MEAN, normalize=True)
self.capture_aux_hidden_states = False
@torch.no_grad()
def forward(
self,
input_ids: torch.Tensor,
positions: torch.Tensor,
forward_batch: ForwardBatch,
input_embeds: torch.Tensor = None,
get_embedding: bool = True,
**kwargs,
) -> EmbeddingPoolerOutput:
assert get_embedding, "EmbeddingGemmaModel is only used for embeddings"
hidden_states = self.model(
input_ids, positions, forward_batch, input_embeds, **kwargs
)
return self.pooler(hidden_states, forward_batch)
def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]):
"""Load both native Gemma3 and Sentence Transformers checkpoints.
The official-style Gemma3 checkpoints prefix backbone parameters with
``model.``, while the Sentence Transformers packaging used by
EmbeddingGemma stores the same backbone at the checkpoint root (and
includes unrelated ``*_Dense`` modules). Normalize the latter form
before delegating to the Gemma3 loader.
"""
backbone_prefixes = ("embed_tokens.", "layers.", "norm.")
remapped_weights = (
(
f"model.{name}" if name.startswith(backbone_prefixes) else name,
weight,
)
for name, weight in weights
if name.startswith("model.") or name.startswith(backbone_prefixes)
)
return super().load_weights(remapped_weights)
EntryClass = [Gemma3ForCausalLM, EmbeddingGemmaModel]
+32
View File
@@ -3638,6 +3638,38 @@ class ServerArgs:
"prompt attention."
)
# EmbeddingGemma is a Gemma3TextModel with bidirectional prompt
# attention. Prefix reuse and split prefills would reuse K/V states
# whose values depend on later prompt tokens, so both are invalid.
# Breakable CUDA Graph captures one complete prefill and is the graph
# mode validated for this encoder-style attention.
if getattr(model_config, "is_embedding_gemma", False):
self.disable_radix_cache = True
self.chunked_prefill_size = -1
self.cuda_graph_config.decode.backend = Backend.DISABLED
if is_cuda() and self.cuda_graph_config.prefill.backend != Backend.DISABLED:
self.cuda_graph_config.prefill.backend = Backend.BREAKABLE
# CUDA-graph sizing has already run by this point. With
# chunked prefill disabled its generic default is -1, which
# otherwise leaves BCG with no shapes to capture. Use the
# model's maximum request length as the safe default; callers
# can still raise it for larger aggregate prefill batches.
if (self.cuda_graph_config.prefill.max_bs or 0) <= 0:
self.cuda_graph_config.prefill.max_bs = model_config.context_len
self.cuda_graph_config.prefill.bs = (
self._generate_prefill_cuda_graph_batch_sizes(
model_config.context_len
)
)
elif not is_cuda():
# BCG is CUDA-only. Other graph backends do not support this
# encoder-style prefill, so retain the eager Triton path.
self.cuda_graph_config.prefill.backend = Backend.DISABLED
logger.info(
"EmbeddingGemma detected: disabling radix cache and chunked "
"prefill; using breakable CUDA graph for CUDA prefill."
)
if (
model_config.is_multimodal
and not model_config.is_multimodal_chunked_prefill_supported
@@ -3,7 +3,10 @@
import unittest
from types import SimpleNamespace
from sglang.srt.configs.model_config import get_hybrid_layer_ids
from sglang.srt.configs.model_config import (
get_hybrid_layer_ids,
is_embedding_gemma,
)
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
@@ -35,5 +38,19 @@ class TestHybridLayerIds(CustomTestCase):
)
class TestEmbeddingGemmaConfig(CustomTestCase):
def test_detects_bidirectional_gemma3_text_config(self):
config = SimpleNamespace(
model_type="gemma3_text", use_bidirectional_attention=True
)
self.assertTrue(is_embedding_gemma(config))
def test_does_not_misclassify_causal_gemma3(self):
config = SimpleNamespace(
model_type="gemma3_text", use_bidirectional_attention=False
)
self.assertFalse(is_embedding_gemma(config))
if __name__ == "__main__":
unittest.main()
@@ -100,6 +100,32 @@ class TestMultimodalPiecewiseCudaGraph(CustomTestCase):
self.assertFalse(runner.can_run_graph(forward_batch))
def test_embedding_gemma_forces_breakable_prefill(self):
args = ServerArgs(model_path="dummy")
args.model_config = SimpleNamespace(
is_embedding_gemma=True,
is_multimodal=False,
context_len=2048,
hf_config=SimpleNamespace(architectures=["Gemma3TextModel"]),
)
args.cuda_graph_config = CudaGraphConfig(
decode=PhaseConfig(backend=Backend.FULL),
prefill=PhaseConfig(backend=Backend.TC_PIECEWISE),
)
args.disable_radix_cache = False
args.chunked_prefill_size = 2048
with (
patch.object(args, "get_model_config", return_value=args.model_config),
patch("sglang.srt.server_args.is_cuda", return_value=True),
):
args._handle_model_capability_adjustments()
self.assertTrue(args.disable_radix_cache)
self.assertEqual(args.chunked_prefill_size, -1)
self.assertEqual(args.cuda_graph_config.decode.backend, Backend.DISABLED)
self.assertEqual(args.cuda_graph_config.prefill.backend, Backend.BREAKABLE)
if __name__ == "__main__":
unittest.main()
@@ -162,6 +162,15 @@ class TestScoreAndPool(CustomTestCase):
expected = self.score_head(pooled)
torch.testing.assert_close(out.embeddings, expected)
def test_mean_pooling_respects_packed_sequence_boundaries(self):
hidden = torch.tensor([[1.0], [3.0], [7.0], [9.0], [11.0]])
fb = _make_forward_batch(extend_seq_lens=[2, 3])
pooler = Pooler(pooling_type=PoolingType.MEAN, normalize=False)
pooled = pooler(hidden, fb).embeddings
torch.testing.assert_close(pooled, torch.tensor([[2.0], [9.0]]))
def test_empty_delimiter_indices(self):
"""Empty delimiter tensor per request -> returns list with empty tensor."""
input_ids = torch.arange(6)