optimize: optimize EmbeddingGemma prefill performance (#32383)

This commit is contained in:
Mick
2026-07-27 17:34:29 +08:00
committed by GitHub
parent 9a0bd24bed
commit 08af5aea57
7 changed files with 285 additions and 37 deletions
@@ -0,0 +1,87 @@
---
title: EmbeddingGemma
description: Serve Google's EmbeddingGemma text embedding model with SGLang.
tag: NEW
---
## Overview
[EmbeddingGemma](https://huggingface.co/google/embeddinggemma-300m) is Google's 300M-parameter text embedding model. SGLang detects its bidirectional Gemma 3 encoder, applies normalized mean pooling, and serves embeddings through the OpenAI-compatible `/v1/embeddings` endpoint.
On NVIDIA CUDA, SGLang uses breakable CUDA graph (BCG) for its complete prefill by default. It also disables prefix caching and chunked prefill, which are incompatible with this bidirectional encoder.
## Prerequisites
- NVIDIA CUDA GPU.
- A Hugging Face account that has accepted the [EmbeddingGemma license](https://huggingface.co/google/embeddinggemma-300m).
- A Hugging Face access token. Export it before starting the server so it can download the gated checkpoint:
```bash
export HF_TOKEN=<your-hugging-face-token>
```
Install an SGLang build that includes EmbeddingGemma support:
```bash
pip install 'git+https://github.com/sgl-project/sglang.git#subdirectory=python'
```
## Start the server
The standard configuration detects EmbeddingGemma and enables embedding mode,
BCG, and the checkpoint's BF16 dtype automatically:
```bash
sglang serve \
--model-path google/embeddinggemma-300m \
--host 0.0.0.0
```
### Hopper performance defaults
On H100 and H200, SGLang automatically selects FA3 and captures BCG through
16,384 tokens, covering eight 2K embedding requests in one replay. No extra
performance flags are required for this workload.
To capture larger aggregate prefills, raise the BCG tier explicitly:
```bash
sglang serve \
--model-path google/embeddinggemma-300m \
--cuda-graph-max-bs-prefill 32768 \
--host 0.0.0.0
```
EmbeddingGemma automatically enables batch tokenization for list-valued
embedding requests, so do not add a separate tokenizer batching flag.
## Create embeddings
Send one string or a batch of strings to the OpenAI-compatible endpoint:
```bash
curl http://127.0.0.1:30000/v1/embeddings \
-H 'Content-Type: application/json' \
-d '{
"model": "google/embeddinggemma-300m",
"input": [
"A short guide to serving text embeddings.",
"Vector search retrieves semantically similar documents."
],
"encoding_format": "float"
}'
```
See [OpenAI-compatible embedding APIs](/docs/basic_usage/openai_api_embeddings) for Python and OpenAI client examples.
## Deployment behavior
EmbeddingGemma performs bidirectional attention over the complete input, so reusing a prefix KV cache or splitting the input into chunked prefills would produce incorrect attention states. SGLang applies the required settings automatically:
- disables RadixAttention prefix caching;
- disables chunked prefill;
- disables the decode CUDA graph because this is an embedding-only model;
- uses BCG for CUDA prefill;
- uses the FlashAttention raw-K/V path when the prefill backend is FA3 or FA4 on supported Hopper and Blackwell CUDA GPUs.
No prefill CUDA-graph override is required for this recipe. Keep BCG enabled to use the optimized EmbeddingGemma path.
+1
View File
@@ -1040,6 +1040,7 @@
{ {
"group": "Google", "group": "Google",
"pages": [ "pages": [
"cookbook/autoregressive/Google/EmbeddingGemma",
"cookbook/autoregressive/Google/Gemma4", "cookbook/autoregressive/Google/Gemma4",
"cookbook/autoregressive/Google/DiffusionGemma" "cookbook/autoregressive/Google/DiffusionGemma"
] ]
@@ -1192,7 +1192,18 @@ class FlashAttentionBackend(AttentionBackend):
is_swa_layer = ( is_swa_layer = (
layer.sliding_window_size is not None and layer.sliding_window_size > -1 layer.sliding_window_size is not None and layer.sliding_window_size > -1
) )
window_size = (layer.sliding_window_size, 0) if is_swa_layer else (-1, -1) causal = not (
layer.is_cross_attention
or layer.attn_type
in (AttentionType.ENCODER_ONLY, AttentionType.DECODER_BIDIRECTIONAL)
)
# FlashAttention's sliding-window tuple is (left, right). Bidirectional
# encoder layers must see the same local context on both sides.
window_size = (
(layer.sliding_window_size, 0 if causal else layer.sliding_window_size)
if is_swa_layer
else (-1, -1)
)
fa_k_descale, fa_v_descale = None, None fa_k_descale, fa_v_descale = None, None
# only use kv scaling if: 1) fp8 kv is explicitly enabled, 2) RadixAttention # only use kv scaling if: 1) fp8 kv is explicitly enabled, 2) RadixAttention
# has corresponding quantization method so that layer.k_scale is not None, # has corresponding quantization method so that layer.k_scale is not None,
@@ -1211,10 +1222,6 @@ class FlashAttentionBackend(AttentionBackend):
q = q.to(self.kv_cache_dtype) q = q.to(self.kv_cache_dtype)
q_rope = q_rope.to(self.kv_cache_dtype) if q_rope is not None else None q_rope = q_rope.to(self.kv_cache_dtype) if q_rope is not None else None
k_rope = k_rope.to(self.kv_cache_dtype) if k_rope is not None else None k_rope = k_rope.to(self.kv_cache_dtype) if k_rope is not None else None
causal = True
if layer.is_cross_attention or layer.attn_type == AttentionType.ENCODER_ONLY:
causal = False
# Check if we should use local attention # Check if we should use local attention
use_local_attn = ( use_local_attn = (
self.has_local_attention self.has_local_attention
+21 -5
View File
@@ -1058,22 +1058,38 @@ class Gemma3RMSNorm(MultiPlatformOp):
def _norm(self, x): def _norm(self, x):
return x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.eps) return x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.eps)
def forward_native(self, x): def forward_native(self, x, residual: Optional[torch.Tensor] = None):
if residual is not None:
residual = x + residual
x = residual
output = self._norm(x.float()) output = self._norm(x.float())
# Llama does x.to(float16) * w whilst Gemma3 is (x * w).to(float16) # Llama does x.to(float16) * w whilst Gemma3 is (x * w).to(float16)
# See https://github.com/huggingface/transformers/pull/29402 # See https://github.com/huggingface/transformers/pull/29402
output = output * (1.0 + self.weight.float()) output = output * (1.0 + self.weight.float())
return output.type_as(x) output = output.type_as(x)
return output if residual is None else (output, residual)
def forward_cpu(self, x): def forward_cpu(self, x, residual: Optional[torch.Tensor] = None):
if residual is not None:
return self.forward_native(x, residual)
if _is_cpu_amx_available and x.stride(-1) == 1: if _is_cpu_amx_available and x.stride(-1) == 1:
return torch.ops.sgl_kernel.gemma3_rmsnorm_cpu(x, self.weight, self.eps) return torch.ops.sgl_kernel.gemma3_rmsnorm_cpu(x, self.weight, self.eps)
return self.forward_native(x) return self.forward_native(x)
def forward_cuda(self, x): def forward_cuda(self, x, residual: Optional[torch.Tensor] = None):
if residual is not None:
# The decoder residual is token-major and contiguous. The fused
# kernel updates both tensors in place: x becomes the normalized
# output and residual becomes x + residual for the next layer.
gemma_fused_add_rmsnorm(x, residual, self.weight.data, self.eps)
return x, residual
if x.dim() == 2:
return gemma_rmsnorm(x, self.weight.data, self.eps)
return self.forward_native(x) return self.forward_native(x)
def forward_npu(self, x): def forward_npu(self, x, residual: Optional[torch.Tensor] = None):
if residual is not None:
return self.forward_native(x, residual)
output, _ = torch_npu.npu_gemma_rms_norm(x, self.weight, self.eps) output, _ = torch_npu.npu_gemma_rms_norm(x, self.weight, self.eps)
return output return output
@@ -868,6 +868,23 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
encoded.get("token_type_ids") if is_cross_encoder else None encoded.get("token_type_ids") if is_cross_encoder else None
) )
# vLLM's OpenAI embeddings endpoint includes special tokens for
# encoder models. EmbeddingGemma's restored Gemma tokenizer adds BOS
# but, by its checkpoint default, omits EOS. Add EOS explicitly here
# rather than mutating tokenizer-global post-processing state.
if (
self.model_config.is_embedding_gemma
and self.tokenizer.eos_token_id is not None
):
input_ids = [
(
ids
if ids and ids[-1] == self.tokenizer.eos_token_id
else [*ids, self.tokenizer.eos_token_id]
)
for ids in input_ids
]
# Step 4: Extract results based on input format # Step 4: Extract results based on input format
return self._extract_tokenizer_results( return self._extract_tokenizer_results(
input_ids, token_type_ids, input_format, original_batch_size input_ids, token_type_ids, input_format, original_batch_size
+99 -16
View File
@@ -14,10 +14,12 @@
# limitations under the License. # limitations under the License.
# ============================================================================== # ==============================================================================
import copy import copy
import json
from typing import Iterable, List, Optional, Set, Tuple from typing import Iterable, List, Optional, Set, Tuple
import einops import einops
import torch import torch
import torch.nn.functional as F
from torch import nn from torch import nn
from transformers import ( from transformers import (
ROPE_INIT_FUNCTIONS, ROPE_INIT_FUNCTIONS,
@@ -46,6 +48,7 @@ from sglang.srt.model_loader.weight_utils import (
) )
from sglang.srt.runtime_context import get_parallel from sglang.srt.runtime_context import get_parallel
from sglang.srt.utils import add_prefix, cpu_has_amx_support, is_cpu, make_layers from sglang.srt.utils import add_prefix, cpu_has_amx_support, is_cpu, make_layers
from sglang.srt.utils.hf_transformers.common import _resolve_local_or_cached_file
_is_cpu = is_cpu() _is_cpu = is_cpu()
_is_cpu_amx_available = cpu_has_amx_support() _is_cpu_amx_available = cpu_has_amx_support()
@@ -374,12 +377,19 @@ class Gemma3DecoderLayer(nn.Module):
position_embeddings_global: torch.Tensor, position_embeddings_global: torch.Tensor,
position_embeddings_local: torch.Tensor, position_embeddings_local: torch.Tensor,
forward_batch: ForwardBatch, forward_batch: ForwardBatch,
residual: Optional[torch.Tensor] = None,
**kwargs, **kwargs,
) -> tuple[ ) -> tuple[
torch.FloatTensor, Optional[tuple[torch.FloatTensor, torch.FloatTensor]] torch.FloatTensor, Optional[tuple[torch.FloatTensor, torch.FloatTensor]]
]: ]:
residual = hidden_states # Keep the residual live across layers so the add preceding the next
hidden_states = self.input_layernorm(hidden_states) # RMSNorm is fused by Gemma3RMSNorm. This matches the upstream Gemma3
# residual layout and is safe to capture in a breakable CUDA graph.
if residual is None:
residual = hidden_states
hidden_states = self.input_layernorm(hidden_states)
else:
hidden_states, residual = self.input_layernorm(hidden_states, residual)
# apply global RoPE to non-sliding layer only # apply global RoPE to non-sliding layer only
if self.self_attn.is_sliding: if self.self_attn.is_sliding:
@@ -395,15 +405,13 @@ class Gemma3DecoderLayer(nn.Module):
**kwargs, **kwargs,
) )
hidden_states = self.post_attention_layernorm(hidden_states) hidden_states = self.post_attention_layernorm(hidden_states)
hidden_states = residual + hidden_states hidden_states, residual = self.pre_feedforward_layernorm(
hidden_states, residual
residual = hidden_states )
hidden_states = self.pre_feedforward_layernorm(hidden_states)
hidden_states = self.mlp(hidden_states) hidden_states = self.mlp(hidden_states)
hidden_states = self.post_feedforward_layernorm(hidden_states) hidden_states = self.post_feedforward_layernorm(hidden_states)
hidden_states = residual + hidden_states
outputs = (hidden_states,) outputs = (hidden_states, residual)
return outputs return outputs
@@ -626,21 +634,22 @@ class Gemma3TextModel(PreTrainedModel):
hidden_states = input_embeds hidden_states = input_embeds
aux_hidden_states = [] aux_hidden_states = []
residual = None
num_layers = len(self.layers) num_layers = len(self.layers)
if _is_cpu and _is_cpu_amx_available: if _is_cpu and _is_cpu_amx_available:
for i, layer in enumerate(self.layers): for i, layer in enumerate(self.layers):
if i in self.layers_to_capture: if i in self.layers_to_capture:
aux_hidden_states.append(hidden_states) aux_hidden_states.append(hidden_states)
layer_outputs = layer( hidden_states, residual = layer(
positions=positions, positions=positions,
position_embeddings_global=None, position_embeddings_global=None,
position_embeddings_local=None, position_embeddings_local=None,
hidden_states=hidden_states, hidden_states=hidden_states,
forward_batch=forward_batch, forward_batch=forward_batch,
residual=residual,
**kwargs, **kwargs,
) )
hidden_states = layer_outputs[0]
else: else:
if positions.dim() == 1: if positions.dim() == 1:
positions = einops.rearrange(positions, "s -> 1 s") positions = einops.rearrange(positions, "s -> 1 s")
@@ -650,15 +659,15 @@ class Gemma3TextModel(PreTrainedModel):
for i, layer in enumerate(self.layers): for i, layer in enumerate(self.layers):
if i in self.layers_to_capture: if i in self.layers_to_capture:
aux_hidden_states.append(hidden_states) aux_hidden_states.append(hidden_states)
layer_outputs = layer( hidden_states, residual = layer(
positions=positions, positions=positions,
position_embeddings_global=position_embeddings_global, position_embeddings_global=position_embeddings_global,
position_embeddings_local=position_embeddings_local, position_embeddings_local=position_embeddings_local,
hidden_states=hidden_states, hidden_states=hidden_states,
forward_batch=forward_batch, forward_batch=forward_batch,
residual=residual,
**kwargs, **kwargs,
) )
hidden_states = layer_outputs[0]
# Capture the output of the last layer if requested. # Capture the output of the last layer if requested.
# layers_to_capture uses +1 offset (captures input of layer i = output of i-1), # layers_to_capture uses +1 offset (captures input of layer i = output of i-1),
@@ -666,7 +675,7 @@ class Gemma3TextModel(PreTrainedModel):
if num_layers in self.layers_to_capture: if num_layers in self.layers_to_capture:
aux_hidden_states.append(hidden_states) aux_hidden_states.append(hidden_states)
hidden_states = self.norm(hidden_states) hidden_states, _ = self.norm(hidden_states, residual)
if len(aux_hidden_states) == 0: if len(aux_hidden_states) == 0:
return hidden_states return hidden_states
@@ -960,9 +969,50 @@ class EmbeddingGemmaModel(Gemma3ForCausalLM):
self.model = Gemma3TextModel( self.model = Gemma3TextModel(
config, quant_config, prefix=add_prefix("model", prefix) config, quant_config, prefix=add_prefix("model", prefix)
) )
self.pooler = Pooler(pooling_type=PoolingType.MEAN, normalize=True) # SentenceTransformers applies mean pooling, then its optional Dense
# projector modules, then L2 normalization. Keep normalization outside
# Pooler so this ordering is preserved.
self.pooler = Pooler(pooling_type=PoolingType.MEAN, normalize=False)
self.projector = self._build_sentence_transformer_projector(config)
self.capture_aux_hidden_states = False self.capture_aux_hidden_states = False
@staticmethod
def _build_sentence_transformer_projector(config: Gemma3TextConfig):
"""Create the checkpoint's SentenceTransformers Dense tail, if present."""
model_path = getattr(config, "_name_or_path", "")
try:
modules_path = _resolve_local_or_cached_file(model_path, "modules.json")
with open(modules_path) as f:
module_specs = json.load(f)
layers = []
for spec in module_specs:
if spec.get("type") != "sentence_transformers.models.Dense":
continue
dense_config_path = _resolve_local_or_cached_file(
model_path, f"{spec['path']}/config.json"
)
with open(dense_config_path) as f:
dense_config = json.load(f)
if dense_config.get("activation_function") not in (
None,
"torch.nn.modules.linear.Identity",
):
raise ValueError(
"EmbeddingGemma only supports identity SentenceTransformers "
"Dense activations"
)
layers.append(
nn.Linear(
dense_config["in_features"],
dense_config["out_features"],
bias=dense_config.get("bias", True),
)
)
return nn.Sequential(*layers) if layers else None
except (FileNotFoundError, OSError, ValueError, KeyError, json.JSONDecodeError):
return None
@torch.no_grad() @torch.no_grad()
def forward( def forward(
self, self,
@@ -977,7 +1027,10 @@ class EmbeddingGemmaModel(Gemma3ForCausalLM):
hidden_states = self.model( hidden_states = self.model(
input_ids, positions, forward_batch, input_embeds, **kwargs input_ids, positions, forward_batch, input_embeds, **kwargs
) )
return self.pooler(hidden_states, forward_batch) pooled = self.pooler(hidden_states, forward_batch).embeddings
if self.projector is not None:
pooled = self.projector(pooled)
return EmbeddingPoolerOutput(embeddings=F.normalize(pooled, p=2, dim=-1))
def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]): def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]):
"""Load both native Gemma3 and Sentence Transformers checkpoints. """Load both native Gemma3 and Sentence Transformers checkpoints.
@@ -998,7 +1051,37 @@ class EmbeddingGemmaModel(Gemma3ForCausalLM):
for name, weight in weights for name, weight in weights
if name.startswith("model.") or name.startswith(backbone_prefixes) if name.startswith("model.") or name.startswith(backbone_prefixes)
) )
return super().load_weights(remapped_weights) loaded_params = super().load_weights(remapped_weights)
if self.projector is not None:
model_path = getattr(self.config, "_name_or_path", "")
modules_path = _resolve_local_or_cached_file(model_path, "modules.json")
with open(modules_path) as f:
module_specs = json.load(f)
dense_specs = [
spec
for spec in module_specs
if spec.get("type") == "sentence_transformers.models.Dense"
]
from safetensors.torch import load_file
for layer, spec in zip(self.projector, dense_specs):
weights_path = _resolve_local_or_cached_file(
model_path, f"{spec['path']}/model.safetensors"
)
weights = load_file(weights_path, device="cpu")
weight_key = next(
key
for key in ("weight", "linear.weight", "dense.weight")
if key in weights
)
layer.weight.data.copy_(weights[weight_key].to(layer.weight.device))
if layer.bias is not None:
layer.bias.data.copy_(
weights[weight_key.replace("weight", "bias")].to(
layer.bias.device
)
)
return loaded_params
EntryClass = [Gemma3ForCausalLM, EmbeddingGemmaModel] EntryClass = [Gemma3ForCausalLM, EmbeddingGemmaModel]
+48 -11
View File
@@ -3644,23 +3644,60 @@ class ServerArgs:
# Breakable CUDA Graph captures one complete prefill and is the graph # Breakable CUDA Graph captures one complete prefill and is the graph
# mode validated for this encoder-style attention. # mode validated for this encoder-style attention.
if getattr(model_config, "is_embedding_gemma", False): if getattr(model_config, "is_embedding_gemma", False):
# This is an encoder-only model even though its HF architecture is
# named Gemma3TextModel. Marking it as embedding mode enables the
# FlashAttention raw-K/V fast path, which does not write or read
# the paged KV cache during its single prefill forward.
self.is_embedding = True
self.disable_radix_cache = True self.disable_radix_cache = True
self.chunked_prefill_size = -1 self.chunked_prefill_size = -1
# Submit a list-valued embeddings request atomically so BCG can
# replay its full prefill batch instead of starting item zero
# while the remaining texts are still being tokenized.
self.enable_tokenizer_batch_encode = True
requested_prefill_backend = (
self.prefill_attention_backend or self.attention_backend
)
if (
is_cuda()
and (is_sm90_supported() or is_sm100_supported())
and requested_prefill_backend in (None, "fa3", "fa4")
):
# Hopper/Blackwell's default FA backend can consume raw K/V
# tensors for a single embedding prefill. Enable its no-KV
# pool path before memory-pool sizing; an explicit non-FA
# backend retains the existing paged-KV behavior.
self.prefill_only_disable_kv_cache = True
self._validate_prefill_only_disable_kv_cache_args()
self.cuda_graph_config.decode.backend = Backend.DISABLED self.cuda_graph_config.decode.backend = Backend.DISABLED
if is_cuda() and self.cuda_graph_config.prefill.backend != Backend.DISABLED: if is_cuda() and self.cuda_graph_config.prefill.backend != Backend.DISABLED:
self.cuda_graph_config.prefill.backend = Backend.BREAKABLE self.cuda_graph_config.prefill.backend = Backend.BREAKABLE
# CUDA-graph sizing has already run by this point. With # CUDA-graph sizing has already run by this point and derives
# chunked prefill disabled its generic default is -1, which # its generic maximum from the 8K chunked-prefill default.
# otherwise leaves BCG with no shapes to capture. Use the # On the Hopper/Blackwell FA raw-K/V path, raise the unlocked
# model's maximum request length as the safe default; callers # default to a full eight-way 2K embedding batch; callers can
# can still raise it for larger aggregate prefill batches. # still override this for larger aggregate prefills.
if (self.cuda_graph_config.prefill.max_bs or 0) <= 0: prefill_config = self.cuda_graph_config.prefill
self.cuda_graph_config.prefill.max_bs = model_config.context_len # Unit-level capability tests may invoke this hook without
self.cuda_graph_config.prefill.bs = ( # running the full CUDA-graph configuration parser, which is
self._generate_prefill_cuda_graph_batch_sizes( # where this internal lock set is normally initialized.
model_config.context_len # Treat that minimal construction as having no user-locked
) # graph settings.
cuda_graph_config_locked = getattr(
self, "_cuda_graph_config_locked", set()
)
if (Phase.PREFILL, "max_bs") not in cuda_graph_config_locked:
prefill_config.max_bs = max(
prefill_config.max_bs or 0,
model_config.context_len,
16384,
) )
if (Phase.PREFILL, "bs") not in cuda_graph_config_locked:
prefill_config.bs = (
self._generate_prefill_cuda_graph_batch_sizes(
prefill_config.max_bs
)
)
elif not is_cuda(): elif not is_cuda():
# BCG is CUDA-only. Other graph backends do not support this # BCG is CUDA-only. Other graph backends do not support this
# encoder-style prefill, so retain the eager Triton path. # encoder-style prefill, so retain the eager Triton path.