embedding: centralize capabilities and complete OpenAI compatibility (#32481)

This commit is contained in:
Mick
2026-07-30 10:28:52 +08:00
committed by GitHub
parent 313a518bee
commit 22faf9fef8
16 changed files with 728 additions and 28 deletions
@@ -11,7 +11,7 @@ import unittest
from collections import Counter
from pathlib import Path
from types import SimpleNamespace
from unittest.mock import patch
from unittest.mock import MagicMock, patch
import numpy as np
from PIL import Image
@@ -42,6 +42,13 @@ from sglang.benchmark.datasets.mooncake import get_mooncake_request_over_time
from sglang.benchmark.datasets.openai_dataset import sample_openai_requests
from sglang.benchmark.datasets.random import sample_random_requests
from sglang.benchmark.datasets.sharegpt import sample_sharegpt_requests
from sglang.benchmark.serving import (
_BACKEND_API_PATHS,
_EMBEDDING_BACKENDS,
ASYNC_REQUEST_FUNCS,
async_request_openai_embeddings,
flush_server_cache,
)
from sglang.test.ci.ci_register import register_cpu_ci
register_cpu_ci(est_time=40, suite="base-a-test-cpu")
@@ -87,6 +94,33 @@ def create_lightweight_tokenizer() -> PreTrainedTokenizerFast:
return hf_tokenizer
class TestEmbeddingBenchmarkBackends(unittest.TestCase):
def test_vllm_embedding_reuses_the_openai_embedding_request_path(self):
self.assertIn("vllm-embedding", _EMBEDDING_BACKENDS)
self.assertIs(
ASYNC_REQUEST_FUNCS["vllm-embedding"], async_request_openai_embeddings
)
self.assertEqual(_BACKEND_API_PATHS["vllm-embedding"], "/v1/embeddings")
def test_embedding_cache_flush_uses_the_engine_specific_endpoint(self):
with (
patch("sglang.benchmark.serving.get_auth_headers", return_value={}),
patch("sglang.benchmark.serving.requests.post") as post,
):
post.return_value = MagicMock()
flush_server_cache("http://127.0.0.1:8000", "vllm-embedding")
post.assert_called_once_with(
"http://127.0.0.1:8000/reset_prefix_cache", headers={}
)
post.reset_mock()
flush_server_cache("http://127.0.0.1:30000", "sglang-embedding")
post.assert_called_once_with(
"http://127.0.0.1:30000/flush_cache", headers={}
)
class DummyProcessor:
def __init__(self, tokenizer: PreTrainedTokenizerFast):
self.tokenizer = tokenizer
@@ -0,0 +1,122 @@
import unittest
from types import SimpleNamespace
from sglang.srt.configs.embedding_model_spec import (
AttentionPattern,
BCGEligibility,
BCGPrefillPolicy,
EmbeddingExecution,
EmbeddingTask,
PoolingStrategy,
embedding_support_matrix,
resolve_embedding_model_spec,
resolved_embedding_plan,
)
from sglang.test.ci.ci_register import register_cpu_ci
register_cpu_ci(est_time=5, suite="base-a-test-cpu")
class TestEmbeddingModelSpec(unittest.TestCase):
def test_embedding_gemma_declares_full_encoder_bcg_contract(self):
spec = resolve_embedding_model_spec(
["Gemma3TextModel"],
is_embedding_requested=False,
is_embedding_gemma=True,
)
self.assertEqual(spec.family, "embeddinggemma")
self.assertEqual(spec.task, EmbeddingTask.EMBED)
self.assertEqual(spec.pooling, PoolingStrategy.MEAN)
self.assertTrue(spec.auto_enable_embedding)
self.assertTrue(spec.safe_disable_kv_cache)
self.assertEqual(spec.bcg_prefill_policy, BCGPrefillPolicy.FULL_ENCODER)
self.assertEqual(spec.execution, EmbeddingExecution.ENCODER_ONLY)
self.assertEqual(spec.attention, AttentionPattern.BIDIRECTIONAL)
self.assertEqual(spec.bcg_eligibility, BCGEligibility.FULL_ENCODER)
def test_encoder_embedding_models_enable_embedding_mode_automatically(self):
spec = resolve_embedding_model_spec(
["BertModel"],
is_embedding_requested=False,
is_embedding_gemma=False,
)
self.assertEqual(spec.family, "bert")
self.assertEqual(spec.task, EmbeddingTask.EMBED)
self.assertEqual(spec.pooling, PoolingStrategy.CLS)
self.assertFalse(spec.requires_embedding_flag)
self.assertTrue(spec.auto_enable_embedding)
def test_support_matrix_is_derived_from_the_same_registry(self):
matrix = embedding_support_matrix()
by_architecture = {row["architecture"]: row for row in matrix}
self.assertEqual(len(matrix), 7)
self.assertEqual(by_architecture["BertModel"]["family"], "bert")
self.assertEqual(by_architecture["BertModel"]["attention"], "bidirectional")
self.assertTrue(by_architecture["CLIPModel"]["supports_multimodal"])
self.assertEqual(
by_architecture["Gemma3TextModel (use_bidirectional_attention=true)"][
"bcg_eligibility"
],
"full_encoder",
)
def test_resolved_plan_reports_effective_runtime_knobs(self):
spec = resolve_embedding_model_spec(
["Gemma3TextModel"],
is_embedding_requested=False,
is_embedding_gemma=True,
)
plan = resolved_embedding_plan(
spec,
server_args=SimpleNamespace(
is_embedding=True,
cuda_graph_config=SimpleNamespace(
prefill=SimpleNamespace(
backend="breakable", max_bs=16384, bs=[1024, 16384]
)
),
prefill_only_disable_kv_cache=True,
disable_radix_cache=True,
chunked_prefill_size=-1,
),
model_config=SimpleNamespace(
is_matryoshka=False, matryoshka_dimensions=None
),
)
self.assertTrue(plan["enabled"])
self.assertTrue(plan["bcg"]["enabled"])
self.assertEqual(plan["bcg"]["capture_token_budget"], 16384)
self.assertEqual(plan["bcg"]["capture_batch_sizes"], [1024, 16384])
self.assertTrue(plan["cache"]["kv_cache_disabled"])
self.assertTrue(plan["cache"]["radix_cache_disabled"])
self.assertTrue(plan["cache"]["chunked_prefill_disabled"])
def test_decoder_embedding_intent_does_not_assume_encoder_fast_path(self):
spec = resolve_embedding_model_spec(
["Qwen3ForCausalLM"],
is_embedding_requested=True,
is_embedding_gemma=False,
)
self.assertEqual(spec.family, "explicit_decoder_embedding")
self.assertEqual(spec.task, EmbeddingTask.EMBED)
self.assertFalse(spec.safe_disable_kv_cache)
self.assertEqual(spec.bcg_prefill_policy, BCGPrefillPolicy.DEFAULT)
def test_unknown_generation_model_has_no_embedding_contract_without_intent(self):
spec = resolve_embedding_model_spec(
["Qwen3ForCausalLM"],
is_embedding_requested=False,
is_embedding_gemma=False,
)
self.assertEqual(spec.task, EmbeddingTask.NONE)
self.assertEqual(spec.family, "none")
if __name__ == "__main__":
unittest.main()
@@ -4,6 +4,7 @@ import unittest
from types import SimpleNamespace
from unittest.mock import patch
from sglang.srt.configs.embedding_model_spec import resolve_embedding_model_spec
from sglang.srt.configs.model_config import (
is_multimodal_piecewise_cuda_graph_supported,
)
@@ -128,6 +129,24 @@ class TestMultimodalPiecewiseCudaGraph(CustomTestCase):
self.assertEqual(args.cuda_graph_config.decode.backend, Backend.DISABLED)
self.assertEqual(args.cuda_graph_config.prefill.backend, Backend.BREAKABLE)
def test_encoder_embedding_model_enables_embedding_mode_without_flag(self):
args = ServerArgs(model_path="dummy")
args.is_embedding = False
args.model_config = SimpleNamespace(
embedding_model_spec=resolve_embedding_model_spec(
["BertModel"],
is_embedding_requested=False,
is_embedding_gemma=False,
),
is_multimodal=False,
hf_config=SimpleNamespace(architectures=["BertModel"]),
)
with patch.object(args, "get_model_config", return_value=args.model_config):
args._handle_model_capability_adjustments()
self.assertTrue(args.is_embedding)
if __name__ == "__main__":
unittest.main()
@@ -2,9 +2,11 @@
Unit tests for the OpenAIServingEmbedding class from serving_embedding.py.
"""
import base64
import importlib
import importlib.abc
import importlib.machinery
import struct
import sys
import types
import unittest
@@ -339,6 +341,29 @@ class ServingEmbeddingTestCase(unittest.TestCase):
self.image_only_multimodal_req
)
def test_base64_embedding_response_uses_little_endian_float32(self):
response = self.serving_embedding._build_embedding_response(
[{"embedding": [0.25, -1.5], "meta_info": {"prompt_tokens": 2}}],
encoding_format="base64",
)
encoded_embedding = response.data[0].embedding
self.assertIsInstance(encoded_embedding, str)
self.assertEqual(
struct.unpack("<2f", base64.b64decode(encoded_embedding)), (0.25, -1.5)
)
self.assertEqual(response.usage.prompt_tokens, 2)
def test_rejects_unknown_embedding_encoding_format(self):
invalid_request = EmbeddingRequest(
model="test-model", input="hello", encoding_format="binary"
)
self.assertIn(
"encoding_format must be either",
self.serving_embedding._validate_request(invalid_request),
)
if __name__ == "__main__":
unittest.main(verbosity=2)