[Score API] Add SequenceClassification Model support (#22118)

This commit is contained in:
Sundara Raman Ramachandran
2026-04-08 01:30:58 -07:00
committed by GitHub
parent 213af1d4f7
commit 712c8c5051
11 changed files with 1287 additions and 131 deletions
@@ -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()
+442
View File
@@ -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()