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
@@ -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)