embedding: centralize capabilities and complete OpenAI compatibility (#32481)
This commit is contained in:
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user