optimize: optimize EmbeddingGemma prefill performance (#32383)
This commit is contained in:
@@ -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.
|
||||
@@ -1040,6 +1040,7 @@
|
||||
{
|
||||
"group": "Google",
|
||||
"pages": [
|
||||
"cookbook/autoregressive/Google/EmbeddingGemma",
|
||||
"cookbook/autoregressive/Google/Gemma4",
|
||||
"cookbook/autoregressive/Google/DiffusionGemma"
|
||||
]
|
||||
|
||||
@@ -1192,7 +1192,18 @@ class FlashAttentionBackend(AttentionBackend):
|
||||
is_swa_layer = (
|
||||
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
|
||||
# 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,
|
||||
@@ -1211,10 +1222,6 @@ class FlashAttentionBackend(AttentionBackend):
|
||||
q = q.to(self.kv_cache_dtype)
|
||||
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
|
||||
causal = True
|
||||
if layer.is_cross_attention or layer.attn_type == AttentionType.ENCODER_ONLY:
|
||||
causal = False
|
||||
|
||||
# Check if we should use local attention
|
||||
use_local_attn = (
|
||||
self.has_local_attention
|
||||
|
||||
@@ -1058,22 +1058,38 @@ class Gemma3RMSNorm(MultiPlatformOp):
|
||||
def _norm(self, x):
|
||||
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())
|
||||
# Llama does x.to(float16) * w whilst Gemma3 is (x * w).to(float16)
|
||||
# See https://github.com/huggingface/transformers/pull/29402
|
||||
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:
|
||||
return torch.ops.sgl_kernel.gemma3_rmsnorm_cpu(x, self.weight, self.eps)
|
||||
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)
|
||||
|
||||
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)
|
||||
return output
|
||||
|
||||
|
||||
@@ -868,6 +868,23 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
||||
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
|
||||
return self._extract_tokenizer_results(
|
||||
input_ids, token_type_ids, input_format, original_batch_size
|
||||
|
||||
@@ -14,10 +14,12 @@
|
||||
# limitations under the License.
|
||||
# ==============================================================================
|
||||
import copy
|
||||
import json
|
||||
from typing import Iterable, List, Optional, Set, Tuple
|
||||
|
||||
import einops
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from torch import nn
|
||||
from transformers import (
|
||||
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.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_amx_available = cpu_has_amx_support()
|
||||
@@ -374,12 +377,19 @@ class Gemma3DecoderLayer(nn.Module):
|
||||
position_embeddings_global: torch.Tensor,
|
||||
position_embeddings_local: torch.Tensor,
|
||||
forward_batch: ForwardBatch,
|
||||
residual: Optional[torch.Tensor] = None,
|
||||
**kwargs,
|
||||
) -> tuple[
|
||||
torch.FloatTensor, Optional[tuple[torch.FloatTensor, torch.FloatTensor]]
|
||||
]:
|
||||
residual = hidden_states
|
||||
hidden_states = self.input_layernorm(hidden_states)
|
||||
# Keep the residual live across layers so the add preceding the next
|
||||
# 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
|
||||
if self.self_attn.is_sliding:
|
||||
@@ -395,15 +405,13 @@ class Gemma3DecoderLayer(nn.Module):
|
||||
**kwargs,
|
||||
)
|
||||
hidden_states = self.post_attention_layernorm(hidden_states)
|
||||
hidden_states = residual + hidden_states
|
||||
|
||||
residual = hidden_states
|
||||
hidden_states = self.pre_feedforward_layernorm(hidden_states)
|
||||
hidden_states, residual = self.pre_feedforward_layernorm(
|
||||
hidden_states, residual
|
||||
)
|
||||
hidden_states = self.mlp(hidden_states)
|
||||
hidden_states = self.post_feedforward_layernorm(hidden_states)
|
||||
hidden_states = residual + hidden_states
|
||||
|
||||
outputs = (hidden_states,)
|
||||
outputs = (hidden_states, residual)
|
||||
|
||||
return outputs
|
||||
|
||||
@@ -626,21 +634,22 @@ class Gemma3TextModel(PreTrainedModel):
|
||||
hidden_states = input_embeds
|
||||
|
||||
aux_hidden_states = []
|
||||
residual = None
|
||||
|
||||
num_layers = len(self.layers)
|
||||
if _is_cpu and _is_cpu_amx_available:
|
||||
for i, layer in enumerate(self.layers):
|
||||
if i in self.layers_to_capture:
|
||||
aux_hidden_states.append(hidden_states)
|
||||
layer_outputs = layer(
|
||||
hidden_states, residual = layer(
|
||||
positions=positions,
|
||||
position_embeddings_global=None,
|
||||
position_embeddings_local=None,
|
||||
hidden_states=hidden_states,
|
||||
forward_batch=forward_batch,
|
||||
residual=residual,
|
||||
**kwargs,
|
||||
)
|
||||
hidden_states = layer_outputs[0]
|
||||
else:
|
||||
if positions.dim() == 1:
|
||||
positions = einops.rearrange(positions, "s -> 1 s")
|
||||
@@ -650,15 +659,15 @@ class Gemma3TextModel(PreTrainedModel):
|
||||
for i, layer in enumerate(self.layers):
|
||||
if i in self.layers_to_capture:
|
||||
aux_hidden_states.append(hidden_states)
|
||||
layer_outputs = layer(
|
||||
hidden_states, residual = layer(
|
||||
positions=positions,
|
||||
position_embeddings_global=position_embeddings_global,
|
||||
position_embeddings_local=position_embeddings_local,
|
||||
hidden_states=hidden_states,
|
||||
forward_batch=forward_batch,
|
||||
residual=residual,
|
||||
**kwargs,
|
||||
)
|
||||
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),
|
||||
@@ -666,7 +675,7 @@ class Gemma3TextModel(PreTrainedModel):
|
||||
if num_layers in self.layers_to_capture:
|
||||
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:
|
||||
return hidden_states
|
||||
@@ -960,9 +969,50 @@ class EmbeddingGemmaModel(Gemma3ForCausalLM):
|
||||
self.model = Gemma3TextModel(
|
||||
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
|
||||
|
||||
@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()
|
||||
def forward(
|
||||
self,
|
||||
@@ -977,7 +1027,10 @@ class EmbeddingGemmaModel(Gemma3ForCausalLM):
|
||||
hidden_states = self.model(
|
||||
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]]):
|
||||
"""Load both native Gemma3 and Sentence Transformers checkpoints.
|
||||
@@ -998,7 +1051,37 @@ class EmbeddingGemmaModel(Gemma3ForCausalLM):
|
||||
for name, weight in weights
|
||||
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]
|
||||
|
||||
@@ -3644,23 +3644,60 @@ class ServerArgs:
|
||||
# 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):
|
||||
# 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.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
|
||||
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
|
||||
)
|
||||
# CUDA-graph sizing has already run by this point and derives
|
||||
# its generic maximum from the 8K chunked-prefill default.
|
||||
# On the Hopper/Blackwell FA raw-K/V path, raise the unlocked
|
||||
# default to a full eight-way 2K embedding batch; callers can
|
||||
# still override this for larger aggregate prefills.
|
||||
prefill_config = self.cuda_graph_config.prefill
|
||||
# Unit-level capability tests may invoke this hook without
|
||||
# running the full CUDA-graph configuration parser, which is
|
||||
# where this internal lock set is normally initialized.
|
||||
# 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():
|
||||
# BCG is CUDA-only. Other graph backends do not support this
|
||||
# encoder-style prefill, so retain the eager Triton path.
|
||||
|
||||
Reference in New Issue
Block a user