[Score API] Add SequenceClassification Model support (#22118)
This commit is contained in:
@@ -0,0 +1,351 @@
|
||||
"""Tests for Scoring API with SequenceClassification models.
|
||||
|
||||
Covers both single-item and multi-item scoring (MIS) for classification
|
||||
models. Uses json_model_override_args to override a regular Qwen3 model's
|
||||
architecture to Qwen3ForSequenceClassification so we can validate the
|
||||
scoring pipeline without needing a dedicated classification checkpoint
|
||||
(the score head gets randomly initialised, which is fine for shape /
|
||||
pipeline correctness).
|
||||
"""
|
||||
|
||||
import json
|
||||
import unittest
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.srt.entrypoints.engine import Engine
|
||||
from sglang.test.ci.ci_register import register_cuda_ci
|
||||
from sglang.test.test_utils import CustomTestCase
|
||||
|
||||
register_cuda_ci(est_time=120, suite="stage-b-test-1-gpu-small")
|
||||
|
||||
# A lightweight Qwen3 checkpoint whose backbone weights load cleanly into
|
||||
# Qwen3ForSequenceClassification (the classification head is random).
|
||||
TEST_BASE_MODEL = "Qwen/Qwen3-0.6B"
|
||||
# <|endoftext|> for Qwen3 tokenizer — used as MIS delimiter
|
||||
QWEN3_ENDOFTEXT_TOKEN_ID = 151643
|
||||
|
||||
|
||||
class TestScoreClassification(CustomTestCase):
|
||||
"""Single-item scoring with a SequenceClassification model."""
|
||||
|
||||
NUM_LABELS = 2
|
||||
|
||||
@classmethod
|
||||
def setUpClass(cls):
|
||||
override_args = json.dumps(
|
||||
{
|
||||
"architectures": ["Qwen3ForSequenceClassification"],
|
||||
"num_labels": cls.NUM_LABELS,
|
||||
}
|
||||
)
|
||||
cls.engine = Engine(
|
||||
model_path=TEST_BASE_MODEL,
|
||||
disable_radix_cache=True,
|
||||
json_model_override_args=override_args,
|
||||
mem_fraction_static=0.15,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def tearDownClass(cls):
|
||||
if hasattr(cls, "engine") and cls.engine is not None:
|
||||
cls.engine.shutdown()
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
def test_basic_single_item(self):
|
||||
"""Each item gets a score vector of length num_labels."""
|
||||
scores = self.engine.score(
|
||||
query="Rate each option:",
|
||||
items=["Option A", "Option B"],
|
||||
apply_softmax=True,
|
||||
).scores
|
||||
|
||||
self.assertEqual(len(scores), 2)
|
||||
for i, score_list in enumerate(scores):
|
||||
self.assertEqual(len(score_list), self.NUM_LABELS)
|
||||
self.assertAlmostEqual(
|
||||
sum(score_list),
|
||||
1.0,
|
||||
places=5,
|
||||
msg=f"Softmax scores for item {i} should sum to 1",
|
||||
)
|
||||
for val in score_list:
|
||||
self.assertGreaterEqual(val, 0.0)
|
||||
self.assertLessEqual(val, 1.0)
|
||||
|
||||
def test_single_item_edge_case(self):
|
||||
"""Single item in the list."""
|
||||
scores = self.engine.score(
|
||||
query="Evaluate:",
|
||||
items=["Only item"],
|
||||
apply_softmax=True,
|
||||
).scores
|
||||
|
||||
self.assertEqual(len(scores), 1)
|
||||
self.assertEqual(len(scores[0]), self.NUM_LABELS)
|
||||
self.assertAlmostEqual(sum(scores[0]), 1.0, places=5)
|
||||
|
||||
def test_raw_logits_without_softmax(self):
|
||||
"""Without softmax, returns raw logits (no probability constraints)."""
|
||||
scores = self.engine.score(
|
||||
query="Evaluate:",
|
||||
items=["Alpha", "Beta"],
|
||||
apply_softmax=False,
|
||||
).scores
|
||||
|
||||
self.assertEqual(len(scores), 2)
|
||||
for score_list in scores:
|
||||
self.assertEqual(len(score_list), self.NUM_LABELS)
|
||||
for val in score_list:
|
||||
self.assertTrue(
|
||||
isinstance(val, (int, float)),
|
||||
f"Expected numeric score, got {type(val)}",
|
||||
)
|
||||
|
||||
def test_deterministic(self):
|
||||
"""Identical inputs yield near-identical scores (fp16 non-determinism allowed)."""
|
||||
kwargs = dict(query="Evaluate:", items=["alpha", "beta", "gamma"])
|
||||
|
||||
scores1 = self.engine.score(**kwargs).scores
|
||||
scores2 = self.engine.score(**kwargs).scores
|
||||
|
||||
self.assertEqual(len(scores1), len(scores2))
|
||||
for s1, s2 in zip(scores1, scores2):
|
||||
for v1, v2 in zip(s1, s2):
|
||||
self.assertAlmostEqual(v1, v2, places=1)
|
||||
|
||||
def test_tokenized_inputs(self):
|
||||
"""Pre-tokenized query and items work the same as text inputs."""
|
||||
from transformers import AutoTokenizer
|
||||
|
||||
tokenizer = AutoTokenizer.from_pretrained(TEST_BASE_MODEL)
|
||||
query_text = "Rate this:"
|
||||
items_text = ["Good", "Bad"]
|
||||
|
||||
text_scores = self.engine.score(
|
||||
query=query_text,
|
||||
items=items_text,
|
||||
apply_softmax=True,
|
||||
).scores
|
||||
|
||||
query_ids = tokenizer.encode(query_text)
|
||||
items_ids = [tokenizer.encode(item) for item in items_text]
|
||||
token_scores = self.engine.score(
|
||||
query=query_ids,
|
||||
items=items_ids,
|
||||
apply_softmax=True,
|
||||
).scores
|
||||
|
||||
self.assertEqual(len(text_scores), len(token_scores))
|
||||
for txt_s, tok_s in zip(text_scores, token_scores):
|
||||
for t, k in zip(txt_s, tok_s):
|
||||
self.assertAlmostEqual(t, k, places=4)
|
||||
|
||||
def test_label_token_ids_ignored(self):
|
||||
"""SequenceClassification models ignore label_token_ids (no crash)."""
|
||||
scores = self.engine.score(
|
||||
query="Evaluate:",
|
||||
items=["Test item"],
|
||||
label_token_ids=[1, 2, 3],
|
||||
apply_softmax=True,
|
||||
).scores
|
||||
|
||||
self.assertEqual(len(scores), 1)
|
||||
self.assertEqual(len(scores[0]), self.NUM_LABELS)
|
||||
|
||||
|
||||
class TestScoreClassificationMIS(CustomTestCase):
|
||||
"""Multi-item scoring (MIS) with a SequenceClassification model.
|
||||
|
||||
MIS packs all items into one sequence separated by a delimiter token.
|
||||
The score_and_pool function extracts per-item scores at delimiter
|
||||
positions.
|
||||
"""
|
||||
|
||||
NUM_LABELS = 2
|
||||
|
||||
@classmethod
|
||||
def setUpClass(cls):
|
||||
override_args = json.dumps(
|
||||
{
|
||||
"architectures": ["Qwen3ForSequenceClassification"],
|
||||
"num_labels": cls.NUM_LABELS,
|
||||
}
|
||||
)
|
||||
cls.engine = Engine(
|
||||
model_path=TEST_BASE_MODEL,
|
||||
disable_radix_cache=True,
|
||||
chunked_prefill_size=-1,
|
||||
multi_item_scoring_delimiter=QWEN3_ENDOFTEXT_TOKEN_ID,
|
||||
json_model_override_args=override_args,
|
||||
mem_fraction_static=0.15,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def tearDownClass(cls):
|
||||
if hasattr(cls, "engine") and cls.engine is not None:
|
||||
cls.engine.shutdown()
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
def test_mis_basic(self):
|
||||
"""MIS produces one score vector per item."""
|
||||
items = ["Option A", "Option B", "Option C"]
|
||||
scores = self.engine.score(
|
||||
query="Rate each option:",
|
||||
items=items,
|
||||
apply_softmax=True,
|
||||
).scores
|
||||
|
||||
self.assertEqual(len(scores), len(items))
|
||||
for i, score_list in enumerate(scores):
|
||||
self.assertEqual(
|
||||
len(score_list),
|
||||
self.NUM_LABELS,
|
||||
f"Item {i} should have {self.NUM_LABELS} scores",
|
||||
)
|
||||
self.assertAlmostEqual(
|
||||
sum(score_list),
|
||||
1.0,
|
||||
places=5,
|
||||
msg=f"Scores for item {i} should sum to 1",
|
||||
)
|
||||
for val in score_list:
|
||||
self.assertGreaterEqual(val, 0.0)
|
||||
self.assertLessEqual(val, 1.0)
|
||||
|
||||
def test_mis_many_items(self):
|
||||
"""Stress test: 10 items."""
|
||||
items = [f"Item {i}" for i in range(10)]
|
||||
scores = self.engine.score(
|
||||
query="Classify each:",
|
||||
items=items,
|
||||
apply_softmax=True,
|
||||
).scores
|
||||
|
||||
self.assertEqual(len(scores), len(items))
|
||||
for score_list in scores:
|
||||
self.assertEqual(len(score_list), self.NUM_LABELS)
|
||||
self.assertAlmostEqual(sum(score_list), 1.0, places=5)
|
||||
|
||||
def test_mis_single_item(self):
|
||||
"""Edge case: single item through MIS path."""
|
||||
scores = self.engine.score(
|
||||
query="Evaluate:",
|
||||
items=["Single item"],
|
||||
apply_softmax=True,
|
||||
).scores
|
||||
|
||||
self.assertEqual(len(scores), 1)
|
||||
self.assertEqual(len(scores[0]), self.NUM_LABELS)
|
||||
self.assertAlmostEqual(sum(scores[0]), 1.0, places=5)
|
||||
|
||||
def test_items_produce_distinct_scores(self):
|
||||
"""Different items must produce different score vectors.
|
||||
|
||||
Even with a randomly initialised classification head, different
|
||||
item texts produce different hidden states, so scores should
|
||||
differ. This catches bugs where all delimiter tokens share the
|
||||
same pooled representation.
|
||||
"""
|
||||
items = [
|
||||
"Option A is about cats",
|
||||
"Option B is about dogs",
|
||||
"Option C is about fish",
|
||||
]
|
||||
scores = self.engine.score(query="Rate each option:", items=items).scores
|
||||
|
||||
self.assertEqual(len(scores), len(items))
|
||||
all_identical = all(scores[0] == s for s in scores[1:])
|
||||
self.assertFalse(
|
||||
all_identical,
|
||||
f"All {len(items)} items returned identical scores — "
|
||||
f"MIS delimiter indexing is likely broken. Scores: {scores[0]}",
|
||||
)
|
||||
|
||||
def test_deterministic(self):
|
||||
"""Identical MIS requests return identical scores."""
|
||||
kwargs = dict(
|
||||
query="Evaluate:",
|
||||
items=["alpha", "beta", "gamma"],
|
||||
)
|
||||
scores1 = self.engine.score(**kwargs).scores
|
||||
scores2 = self.engine.score(**kwargs).scores
|
||||
|
||||
self.assertEqual(scores1, scores2)
|
||||
|
||||
def test_softmax_valid(self):
|
||||
"""With softmax, each item's scores form a valid probability distribution."""
|
||||
items = ["Option A", "Option B", "Option C"]
|
||||
scores = self.engine.score(
|
||||
query="Rate each option:",
|
||||
items=items,
|
||||
apply_softmax=True,
|
||||
).scores
|
||||
|
||||
for i, score_list in enumerate(scores):
|
||||
self.assertEqual(len(score_list), self.NUM_LABELS)
|
||||
for val in score_list:
|
||||
self.assertGreaterEqual(val, 0.0)
|
||||
self.assertLessEqual(val, 1.0)
|
||||
self.assertAlmostEqual(
|
||||
sum(score_list),
|
||||
1.0,
|
||||
places=6,
|
||||
msg=f"Softmax scores for item {i} don't sum to 1: {sum(score_list)}",
|
||||
)
|
||||
|
||||
|
||||
class TestScoreClassificationMISAdvanced(CustomTestCase):
|
||||
"""Advanced MIS tests with more labels to stress tensor shape handling."""
|
||||
|
||||
NUM_LABELS = 12
|
||||
|
||||
@classmethod
|
||||
def setUpClass(cls):
|
||||
override_args = json.dumps(
|
||||
{
|
||||
"architectures": ["Qwen3ForSequenceClassification"],
|
||||
"num_labels": cls.NUM_LABELS,
|
||||
}
|
||||
)
|
||||
cls.engine = Engine(
|
||||
model_path=TEST_BASE_MODEL,
|
||||
disable_radix_cache=True,
|
||||
chunked_prefill_size=-1,
|
||||
multi_item_scoring_delimiter=QWEN3_ENDOFTEXT_TOKEN_ID,
|
||||
json_model_override_args=override_args,
|
||||
mem_fraction_static=0.15,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def tearDownClass(cls):
|
||||
if hasattr(cls, "engine") and cls.engine is not None:
|
||||
cls.engine.shutdown()
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
def test_many_labels_shape(self):
|
||||
"""Verify correct shape with many labels (catches 2D tensor bugs)."""
|
||||
items = [f"Item {i}" for i in range(5)]
|
||||
scores = self.engine.score(
|
||||
query="Classify:",
|
||||
items=items,
|
||||
apply_softmax=True,
|
||||
).scores
|
||||
|
||||
self.assertEqual(len(scores), len(items))
|
||||
for score_list in scores:
|
||||
self.assertEqual(len(score_list), self.NUM_LABELS)
|
||||
self.assertAlmostEqual(sum(score_list), 1.0, places=5)
|
||||
|
||||
def test_many_items_distinct(self):
|
||||
"""15 items should not all produce identical scores."""
|
||||
items = [f"City {i}" for i in range(15)]
|
||||
scores = self.engine.score(query="Classify each city:", items=items).scores
|
||||
|
||||
self.assertEqual(len(scores), len(items))
|
||||
unique_count = len({tuple(s) for s in scores})
|
||||
self.assertGreater(unique_count, 1, "All 15 items returned identical scores")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,216 @@
|
||||
"""Unit tests for score_and_pool in sglang.srt.layers.pooler.
|
||||
|
||||
All tests run on CPU — no GPU required. The global server_args singleton
|
||||
is mocked so the tests are hermetic.
|
||||
"""
|
||||
|
||||
import unittest
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import patch
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from sglang.srt.layers.pooler import (
|
||||
EmbeddingPoolerOutput,
|
||||
Pooler,
|
||||
PoolingType,
|
||||
score_and_pool,
|
||||
)
|
||||
from sglang.test.ci.ci_register import register_cpu_ci
|
||||
from sglang.test.test_utils import CustomTestCase
|
||||
|
||||
register_cpu_ci(est_time=5, suite="stage-a-test-cpu")
|
||||
|
||||
|
||||
def _make_forward_batch(extend_seq_lens, is_prefill_only=False):
|
||||
"""Build a minimal ForwardBatch stub for pooler unit tests."""
|
||||
return SimpleNamespace(
|
||||
extend_seq_lens=torch.tensor(extend_seq_lens, dtype=torch.long),
|
||||
extend_seq_lens_cpu=extend_seq_lens,
|
||||
is_prefill_only=is_prefill_only,
|
||||
dimensions=None,
|
||||
)
|
||||
|
||||
|
||||
def _mock_server_args(delimiter=None):
|
||||
return SimpleNamespace(multi_item_scoring_delimiter=delimiter)
|
||||
|
||||
|
||||
class TestScoreAndPool(CustomTestCase):
|
||||
"""Unit tests for the score_and_pool helper function."""
|
||||
|
||||
def setUp(self):
|
||||
torch.manual_seed(42)
|
||||
self.hidden_dim = 8
|
||||
self.num_labels = 2
|
||||
self.score_head = nn.Linear(self.hidden_dim, self.num_labels, bias=False)
|
||||
self.pooler = Pooler(pooling_type=PoolingType.LAST, normalize=False)
|
||||
|
||||
@patch("sglang.srt.layers.pooler.get_global_server_args")
|
||||
def test_single_item_returns_scores(self, mock_get_args):
|
||||
"""No delimiter -> single-item path returns [batch, num_labels]."""
|
||||
mock_get_args.return_value = _mock_server_args(delimiter=None)
|
||||
|
||||
hidden = torch.randn(8, self.hidden_dim)
|
||||
fb = _make_forward_batch(extend_seq_lens=[5, 3])
|
||||
input_ids = torch.arange(8)
|
||||
|
||||
out = score_and_pool(self.score_head, self.pooler, hidden, fb, input_ids)
|
||||
|
||||
self.assertIsInstance(out, EmbeddingPoolerOutput)
|
||||
self.assertEqual(out.embeddings.shape, (2, self.num_labels))
|
||||
|
||||
@patch("sglang.srt.layers.pooler.get_global_server_args")
|
||||
def test_mis_returns_per_request_list(self, mock_get_args):
|
||||
"""Delimiter found -> returns a list with one tensor per request."""
|
||||
delimiter_token = 99
|
||||
mock_get_args.return_value = _mock_server_args(delimiter=delimiter_token)
|
||||
|
||||
input_ids = torch.tensor(
|
||||
[
|
||||
0,
|
||||
1,
|
||||
2,
|
||||
delimiter_token,
|
||||
3,
|
||||
4,
|
||||
5,
|
||||
delimiter_token,
|
||||
6,
|
||||
7,
|
||||
8,
|
||||
delimiter_token,
|
||||
]
|
||||
)
|
||||
hidden = torch.randn(len(input_ids), self.hidden_dim)
|
||||
fb = _make_forward_batch(extend_seq_lens=[len(input_ids)], is_prefill_only=True)
|
||||
|
||||
out = score_and_pool(self.score_head, self.pooler, hidden, fb, input_ids)
|
||||
|
||||
self.assertIsInstance(out.embeddings, list)
|
||||
self.assertEqual(len(out.embeddings), 1)
|
||||
self.assertEqual(out.embeddings[0].shape, (3, self.num_labels))
|
||||
|
||||
@patch("sglang.srt.layers.pooler.get_global_server_args")
|
||||
def test_mis_batched_splits_per_request(self, mock_get_args):
|
||||
"""Two batched MIS requests -> returns a list of length 2."""
|
||||
delimiter_token = 99
|
||||
mock_get_args.return_value = _mock_server_args(delimiter=delimiter_token)
|
||||
|
||||
# Request 1: [10, 11, delim, 12, 13, delim] -> 2 delimiters
|
||||
# Request 2: [20, 21, 22, delim] -> 1 delimiter
|
||||
req1 = [10, 11, delimiter_token, 12, 13, delimiter_token]
|
||||
req2 = [20, 21, 22, delimiter_token]
|
||||
input_ids = torch.tensor(req1 + req2)
|
||||
hidden = torch.randn(len(input_ids), self.hidden_dim)
|
||||
fb = _make_forward_batch(
|
||||
extend_seq_lens=[len(req1), len(req2)], is_prefill_only=True
|
||||
)
|
||||
|
||||
out = score_and_pool(self.score_head, self.pooler, hidden, fb, input_ids)
|
||||
|
||||
self.assertIsInstance(out.embeddings, list)
|
||||
self.assertEqual(len(out.embeddings), 2)
|
||||
self.assertEqual(out.embeddings[0].shape, (2, self.num_labels))
|
||||
self.assertEqual(out.embeddings[1].shape, (1, self.num_labels))
|
||||
|
||||
@patch("sglang.srt.layers.pooler.get_global_server_args")
|
||||
def test_mis_falls_back_when_no_delimiters_in_input(self, mock_get_args):
|
||||
"""Delimiter configured but absent from input_ids -> single-item fallback."""
|
||||
mock_get_args.return_value = _mock_server_args(delimiter=99)
|
||||
|
||||
input_ids = torch.tensor([0, 1, 2, 3, 4, 5, 6, 7])
|
||||
hidden = torch.randn(8, self.hidden_dim)
|
||||
fb = _make_forward_batch(extend_seq_lens=[5, 3], is_prefill_only=True)
|
||||
|
||||
out = score_and_pool(self.score_head, self.pooler, hidden, fb, input_ids)
|
||||
|
||||
self.assertIsInstance(out.embeddings, torch.Tensor)
|
||||
self.assertEqual(out.embeddings.shape, (2, self.num_labels))
|
||||
|
||||
@patch("sglang.srt.layers.pooler.get_global_server_args")
|
||||
def test_mis_falls_back_when_not_prefill_only(self, mock_get_args):
|
||||
"""Delimiter configured, is_prefill_only=False -> single-item fallback."""
|
||||
mock_get_args.return_value = _mock_server_args(delimiter=99)
|
||||
|
||||
input_ids = torch.tensor([0, 1, 2, 99, 3, 4, 5, 99])
|
||||
hidden = torch.randn(8, self.hidden_dim)
|
||||
fb = _make_forward_batch(extend_seq_lens=[5, 3], is_prefill_only=False)
|
||||
|
||||
out = score_and_pool(self.score_head, self.pooler, hidden, fb, input_ids)
|
||||
|
||||
self.assertIsInstance(out.embeddings, torch.Tensor)
|
||||
self.assertEqual(out.embeddings.shape, (2, self.num_labels))
|
||||
|
||||
@patch("sglang.srt.layers.pooler.get_global_server_args")
|
||||
def test_mis_extracts_positions_before_delimiter(self, mock_get_args):
|
||||
"""Verify MIS picks hidden states at index (delimiter_position - 1)."""
|
||||
delimiter_token = 99
|
||||
mock_get_args.return_value = _mock_server_args(delimiter=delimiter_token)
|
||||
|
||||
# Delimiters at indices 2 and 5 -> extract hidden at indices 1 and 4
|
||||
input_ids = torch.tensor([10, 11, delimiter_token, 20, 21, delimiter_token])
|
||||
hidden = (
|
||||
torch.arange(len(input_ids))
|
||||
.unsqueeze(1)
|
||||
.float()
|
||||
.expand(-1, self.hidden_dim)
|
||||
.clone()
|
||||
)
|
||||
fb = _make_forward_batch(extend_seq_lens=[len(input_ids)], is_prefill_only=True)
|
||||
|
||||
identity_head = nn.Linear(self.hidden_dim, self.hidden_dim, bias=False)
|
||||
nn.init.eye_(identity_head.weight)
|
||||
|
||||
out = score_and_pool(identity_head, self.pooler, hidden, fb, input_ids)
|
||||
|
||||
scores = out.embeddings[0]
|
||||
torch.testing.assert_close(scores[0], hidden[1])
|
||||
torch.testing.assert_close(scores[1], hidden[4])
|
||||
|
||||
@patch("sglang.srt.layers.pooler.get_global_server_args")
|
||||
def test_mis_ignores_delimiter_at_position_zero(self, mock_get_args):
|
||||
"""A delimiter at flat index 0 has no preceding token and must be skipped."""
|
||||
delimiter_token = 99
|
||||
mock_get_args.return_value = _mock_server_args(delimiter=delimiter_token)
|
||||
|
||||
# Delimiter at index 0 should be ignored; only the one at index 3 counts
|
||||
input_ids = torch.tensor([delimiter_token, 10, 11, delimiter_token])
|
||||
hidden = (
|
||||
torch.arange(len(input_ids))
|
||||
.unsqueeze(1)
|
||||
.float()
|
||||
.expand(-1, self.hidden_dim)
|
||||
.clone()
|
||||
)
|
||||
fb = _make_forward_batch(extend_seq_lens=[len(input_ids)], is_prefill_only=True)
|
||||
|
||||
identity_head = nn.Linear(self.hidden_dim, self.hidden_dim, bias=False)
|
||||
nn.init.eye_(identity_head.weight)
|
||||
|
||||
out = score_and_pool(identity_head, self.pooler, hidden, fb, input_ids)
|
||||
|
||||
self.assertEqual(len(out.embeddings), 1)
|
||||
self.assertEqual(out.embeddings[0].shape[0], 1)
|
||||
torch.testing.assert_close(out.embeddings[0][0], hidden[2])
|
||||
|
||||
@patch("sglang.srt.layers.pooler.get_global_server_args")
|
||||
def test_single_item_scores_match_manual_computation(self, mock_get_args):
|
||||
"""Single-item scores equal score_head applied to all tokens then pooled."""
|
||||
mock_get_args.return_value = _mock_server_args(delimiter=None)
|
||||
|
||||
hidden = torch.randn(8, self.hidden_dim)
|
||||
fb = _make_forward_batch(extend_seq_lens=[5, 3])
|
||||
input_ids = torch.arange(8)
|
||||
|
||||
out = score_and_pool(self.score_head, self.pooler, hidden, fb, input_ids)
|
||||
|
||||
# score-first-then-pool: matches the original Qwen3/Qwen2 classification forward
|
||||
logits = self.score_head(hidden)
|
||||
expected = self.pooler(logits, fb).embeddings
|
||||
torch.testing.assert_close(out.embeddings, expected)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,442 @@
|
||||
import os
|
||||
import signal
|
||||
import subprocess
|
||||
import tempfile
|
||||
import time
|
||||
import unittest
|
||||
|
||||
import requests
|
||||
|
||||
from sglang.test.test_utils import DEFAULT_SMALL_MODEL_NAME_FOR_TEST, CustomTestCase
|
||||
|
||||
|
||||
class TestMultiItemScoringServer(CustomTestCase):
|
||||
"""Test multi-item scoring functionality through the server API."""
|
||||
|
||||
def setUp(self):
|
||||
"""Set up each test case."""
|
||||
self.model_path = DEFAULT_SMALL_MODEL_NAME_FOR_TEST
|
||||
self.port = 30001 # Use different port to avoid conflicts
|
||||
self.host = "localhost"
|
||||
self.base_url = f"http://{self.host}:{self.port}"
|
||||
self.server_process = None
|
||||
self.server_log_file = None
|
||||
|
||||
def tearDown(self):
|
||||
"""Clean up after each test case."""
|
||||
self.stop_server()
|
||||
|
||||
def start_server(self, multi_item_scoring_delimiter=None):
|
||||
"""Start the SGLang server with multi-item scoring enabled."""
|
||||
if self.server_process is not None:
|
||||
self.stop_server()
|
||||
|
||||
# Create a temporary log file
|
||||
self.server_log_file = tempfile.NamedTemporaryFile(mode="w+", delete=False)
|
||||
|
||||
cmd = [
|
||||
"python",
|
||||
"-m",
|
||||
"sglang.launch_server",
|
||||
"--model-path",
|
||||
self.model_path,
|
||||
"--port",
|
||||
str(self.port),
|
||||
"--host",
|
||||
self.host,
|
||||
"--chunked-prefill-size",
|
||||
"-1",
|
||||
"--dtype",
|
||||
"float16",
|
||||
"--max-prefill-tokens",
|
||||
"30000",
|
||||
"--mem-fraction-static",
|
||||
"0.3",
|
||||
"--disable-radix-cache",
|
||||
"--disable-cuda-graph",
|
||||
"--attention-backend",
|
||||
"flashinfer",
|
||||
]
|
||||
|
||||
if multi_item_scoring_delimiter is not None:
|
||||
cmd.extend(
|
||||
["--multi-item-scoring-delimiter", str(multi_item_scoring_delimiter)]
|
||||
)
|
||||
|
||||
# Start server process
|
||||
self.server_process = subprocess.Popen(
|
||||
cmd,
|
||||
stdout=self.server_log_file,
|
||||
stderr=subprocess.STDOUT,
|
||||
preexec_fn=os.setsid if hasattr(os, "setsid") else None,
|
||||
)
|
||||
|
||||
# Wait for server to start
|
||||
max_wait_time = 60 # seconds
|
||||
start_time = time.time()
|
||||
|
||||
while time.time() - start_time < max_wait_time:
|
||||
try:
|
||||
response = requests.get(f"{self.base_url}/get_model_info", timeout=5)
|
||||
if response.status_code == 200:
|
||||
return True
|
||||
except requests.exceptions.RequestException:
|
||||
pass
|
||||
time.sleep(1)
|
||||
|
||||
# If we get here, server didn't start properly
|
||||
self.stop_server()
|
||||
raise RuntimeError("Failed to start SGLang server within timeout period")
|
||||
|
||||
def stop_server(self):
|
||||
"""Stop the SGLang server."""
|
||||
if self.server_process is not None:
|
||||
try:
|
||||
# Kill the process group
|
||||
if hasattr(os, "killpg"):
|
||||
os.killpg(os.getpgid(self.server_process.pid), signal.SIGTERM)
|
||||
else:
|
||||
self.server_process.terminate()
|
||||
|
||||
# Wait for graceful shutdown
|
||||
try:
|
||||
self.server_process.wait(timeout=10)
|
||||
except subprocess.TimeoutExpired:
|
||||
# Force kill if graceful shutdown fails
|
||||
if hasattr(os, "killpg"):
|
||||
os.killpg(os.getpgid(self.server_process.pid), signal.SIGKILL)
|
||||
else:
|
||||
self.server_process.kill()
|
||||
self.server_process.wait()
|
||||
except (ProcessLookupError, OSError):
|
||||
# Process already terminated
|
||||
pass
|
||||
finally:
|
||||
self.server_process = None
|
||||
|
||||
if self.server_log_file is not None:
|
||||
try:
|
||||
self.server_log_file.close()
|
||||
os.unlink(self.server_log_file.name)
|
||||
except (OSError, FileNotFoundError):
|
||||
pass
|
||||
finally:
|
||||
self.server_log_file = None
|
||||
|
||||
def get_server_logs(self):
|
||||
"""Get server logs for debugging."""
|
||||
if self.server_log_file is not None:
|
||||
try:
|
||||
with open(self.server_log_file.name, "r") as f:
|
||||
return f.read()
|
||||
except (OSError, FileNotFoundError):
|
||||
pass
|
||||
return "No logs available"
|
||||
|
||||
def test_multi_item_scoring_server_basic(self):
|
||||
"""Test basic multi-item scoring through server API."""
|
||||
# Start server with multi-item scoring enabled
|
||||
delimiter_token_id = 151655 # Example delimiter token ID
|
||||
self.start_server(multi_item_scoring_delimiter=delimiter_token_id)
|
||||
|
||||
# Test data
|
||||
query = "What is the capital of California? Answer Yes or No for each of the following options:"
|
||||
items = ["Sacramento", "San Jose", "San Francisco"]
|
||||
label_token_ids = [9454, 2753] # "Yes" and "No" tokens
|
||||
|
||||
# Make scoring request
|
||||
payload = {
|
||||
"query": query,
|
||||
"items": items,
|
||||
"label_token_ids": label_token_ids,
|
||||
"model": self.model_path,
|
||||
}
|
||||
|
||||
response = requests.post(
|
||||
f"{self.base_url}/v1/score",
|
||||
headers={"Content-Type": "application/json"},
|
||||
json=payload,
|
||||
timeout=30,
|
||||
)
|
||||
|
||||
# Check response
|
||||
self.assertEqual(
|
||||
response.status_code,
|
||||
200,
|
||||
f"Server returned {response.status_code}. Logs: {self.get_server_logs()}",
|
||||
)
|
||||
|
||||
result = response.json()
|
||||
|
||||
# Verify response structure
|
||||
self.assertIn("scores", result)
|
||||
self.assertIn("model", result)
|
||||
self.assertIn("object", result)
|
||||
|
||||
self.assertEqual(result["object"], "scoring")
|
||||
self.assertEqual(result["model"], self.model_path)
|
||||
|
||||
# Verify scores
|
||||
scores = result["scores"]
|
||||
self.assertEqual(len(scores), len(items), "Should get one score list per item")
|
||||
|
||||
for i, score_list in enumerate(scores):
|
||||
self.assertEqual(
|
||||
len(score_list),
|
||||
len(label_token_ids),
|
||||
f"Item {i} should have {len(label_token_ids)} scores",
|
||||
)
|
||||
# Verify scores are probabilities (sum to 1)
|
||||
self.assertAlmostEqual(
|
||||
sum(score_list),
|
||||
1.0,
|
||||
places=6,
|
||||
msg=f"Scores for item {i} should sum to 1",
|
||||
)
|
||||
# Verify all scores are non-negative
|
||||
for j, score in enumerate(score_list):
|
||||
self.assertGreaterEqual(
|
||||
score, 0, f"Score {j} for item {i} should be non-negative"
|
||||
)
|
||||
|
||||
def test_multi_item_scoring_server_different_sizes(self):
|
||||
"""Test multi-item scoring with different numbers of items through server."""
|
||||
delimiter_token_id = 151655
|
||||
self.start_server(multi_item_scoring_delimiter=delimiter_token_id)
|
||||
|
||||
query = "Rate each option:"
|
||||
label_token_ids = [1, 2, 3, 4, 5]
|
||||
|
||||
test_cases = [
|
||||
["Single item"],
|
||||
["Item 1", "Item 2"],
|
||||
["A", "B", "C", "D"],
|
||||
["X", "Y", "Z", "W", "V", "U"],
|
||||
]
|
||||
|
||||
for items in test_cases:
|
||||
with self.subTest(items=items):
|
||||
payload = {
|
||||
"query": query,
|
||||
"items": items,
|
||||
"label_token_ids": label_token_ids,
|
||||
"model": self.model_path,
|
||||
}
|
||||
|
||||
response = requests.post(
|
||||
f"{self.base_url}/v1/score",
|
||||
headers={"Content-Type": "application/json"},
|
||||
json=payload,
|
||||
timeout=30,
|
||||
)
|
||||
|
||||
self.assertEqual(
|
||||
response.status_code,
|
||||
200,
|
||||
f"Failed for items {items}. Logs: {self.get_server_logs()}",
|
||||
)
|
||||
|
||||
result = response.json()
|
||||
scores = result["scores"]
|
||||
|
||||
self.assertEqual(
|
||||
len(scores), len(items), f"Should get {len(items)} score lists"
|
||||
)
|
||||
|
||||
for i, score_list in enumerate(scores):
|
||||
self.assertEqual(len(score_list), len(label_token_ids))
|
||||
self.assertAlmostEqual(sum(score_list), 1.0, places=6)
|
||||
|
||||
def test_multi_item_scoring_server_empty_items(self):
|
||||
"""Test multi-item scoring with empty items list through server."""
|
||||
delimiter_token_id = 151655
|
||||
self.start_server(multi_item_scoring_delimiter=delimiter_token_id)
|
||||
|
||||
payload = {
|
||||
"query": "Test query",
|
||||
"items": [],
|
||||
"label_token_ids": [1, 2],
|
||||
"model": self.model_path,
|
||||
}
|
||||
|
||||
response = requests.post(
|
||||
f"{self.base_url}/v1/score",
|
||||
headers={"Content-Type": "application/json"},
|
||||
json=payload,
|
||||
timeout=30,
|
||||
)
|
||||
|
||||
self.assertEqual(response.status_code, 200)
|
||||
result = response.json()
|
||||
self.assertEqual(
|
||||
len(result["scores"]), 0, "Should return empty list for empty items"
|
||||
)
|
||||
|
||||
def test_multi_item_scoring_server_without_delimiter(self):
|
||||
"""Test that server works without multi-item scoring delimiter."""
|
||||
# Start server without multi-item scoring delimiter
|
||||
self.start_server(multi_item_scoring_delimiter=None)
|
||||
|
||||
payload = {
|
||||
"query": "Test query",
|
||||
"items": ["Item 1", "Item 2"],
|
||||
"label_token_ids": [1, 2],
|
||||
"model": self.model_path,
|
||||
}
|
||||
|
||||
response = requests.post(
|
||||
f"{self.base_url}/v1/score",
|
||||
headers={"Content-Type": "application/json"},
|
||||
json=payload,
|
||||
timeout=30,
|
||||
)
|
||||
|
||||
# Should still work (falls back to regular scoring)
|
||||
self.assertEqual(response.status_code, 200)
|
||||
result = response.json()
|
||||
self.assertIn("scores", result)
|
||||
|
||||
def test_multi_item_scoring_server_error_handling(self):
|
||||
"""Test error handling in multi-item scoring server API."""
|
||||
delimiter_token_id = 151655
|
||||
self.start_server(multi_item_scoring_delimiter=delimiter_token_id)
|
||||
|
||||
# Test with invalid payload
|
||||
invalid_payloads = [
|
||||
{
|
||||
"query": "Test",
|
||||
"items": "not a list",
|
||||
"label_token_ids": [1, 2],
|
||||
"model": self.model_path,
|
||||
},
|
||||
{
|
||||
"query": "Test",
|
||||
"items": ["Item 1"],
|
||||
"label_token_ids": "not a list",
|
||||
"model": self.model_path,
|
||||
},
|
||||
{
|
||||
"query": "Test",
|
||||
"items": ["Item 1"],
|
||||
"label_token_ids": [1, 2],
|
||||
}, # Missing model
|
||||
{
|
||||
"items": ["Item 1"],
|
||||
"label_token_ids": [1, 2],
|
||||
"model": self.model_path,
|
||||
}, # Missing query
|
||||
]
|
||||
|
||||
for i, payload in enumerate(invalid_payloads):
|
||||
with self.subTest(payload=i):
|
||||
response = requests.post(
|
||||
f"{self.base_url}/v1/score",
|
||||
headers={"Content-Type": "application/json"},
|
||||
json=payload,
|
||||
timeout=30,
|
||||
)
|
||||
|
||||
# Should return error status
|
||||
self.assertGreaterEqual(
|
||||
response.status_code,
|
||||
400,
|
||||
f"Should return error for invalid payload {i}",
|
||||
)
|
||||
|
||||
def test_multi_item_scoring_server_consistency(self):
|
||||
"""Test that multi-item scoring gives consistent results through server."""
|
||||
delimiter_token_id = 151655
|
||||
self.start_server(multi_item_scoring_delimiter=delimiter_token_id)
|
||||
|
||||
query = "Choose the best option:"
|
||||
items = ["Option A", "Option B", "Option C"]
|
||||
label_token_ids = [1, 2, 3]
|
||||
|
||||
payload = {
|
||||
"query": query,
|
||||
"items": items,
|
||||
"label_token_ids": label_token_ids,
|
||||
"model": self.model_path,
|
||||
}
|
||||
|
||||
# Run the same test multiple times
|
||||
scores1 = None
|
||||
scores2 = None
|
||||
|
||||
for attempt in range(2):
|
||||
response = requests.post(
|
||||
f"{self.base_url}/v1/score",
|
||||
headers={"Content-Type": "application/json"},
|
||||
json=payload,
|
||||
timeout=30,
|
||||
)
|
||||
|
||||
self.assertEqual(
|
||||
response.status_code,
|
||||
200,
|
||||
f"Attempt {attempt + 1} failed. Logs: {self.get_server_logs()}",
|
||||
)
|
||||
|
||||
result = response.json()
|
||||
scores = result["scores"]
|
||||
|
||||
if attempt == 0:
|
||||
scores1 = scores
|
||||
else:
|
||||
scores2 = scores
|
||||
|
||||
# Results should be identical (deterministic)
|
||||
self.assertEqual(len(scores1), len(scores2), "Should get same number of items")
|
||||
for i, (s1, s2) in enumerate(zip(scores1, scores2)):
|
||||
self.assertEqual(
|
||||
len(s1), len(s2), f"Item {i} should have same number of scores"
|
||||
)
|
||||
for j, (score1, score2) in enumerate(zip(s1, s2)):
|
||||
self.assertAlmostEqual(
|
||||
score1,
|
||||
score2,
|
||||
places=6,
|
||||
msg=f"Score {j} for item {i} should be identical",
|
||||
)
|
||||
|
||||
def test_multi_item_scoring_server_large_batch(self):
|
||||
"""Test multi-item scoring with large batch through server."""
|
||||
delimiter_token_id = 151655
|
||||
self.start_server(multi_item_scoring_delimiter=delimiter_token_id)
|
||||
|
||||
query = "Classify each item:"
|
||||
items = [f"Item {i}" for i in range(10)] # 10 items (smaller for test)
|
||||
label_token_ids = [1, 2, 3]
|
||||
|
||||
payload = {
|
||||
"query": query,
|
||||
"items": items,
|
||||
"label_token_ids": label_token_ids,
|
||||
"model": self.model_path,
|
||||
}
|
||||
|
||||
response = requests.post(
|
||||
f"{self.base_url}/v1/score",
|
||||
headers={"Content-Type": "application/json"},
|
||||
json=payload,
|
||||
timeout=60, # Longer timeout for large batch
|
||||
)
|
||||
|
||||
self.assertEqual(
|
||||
response.status_code,
|
||||
200,
|
||||
f"Large batch failed. Logs: {self.get_server_logs()}",
|
||||
)
|
||||
|
||||
result = response.json()
|
||||
scores = result["scores"]
|
||||
|
||||
self.assertEqual(len(scores), len(items), "Should handle large batches")
|
||||
|
||||
for i, score_list in enumerate(scores):
|
||||
self.assertEqual(len(score_list), len(label_token_ids))
|
||||
self.assertAlmostEqual(sum(score_list), 1.0, places=6)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
Reference in New Issue
Block a user