[Score API] Add Multi-Item Scoring with pre-computed delimiter indices (#22544)

Co-authored-by: Chanh Nguyen <chanhnguyen@gmail.com>
Co-authored-by: Sundara Raman Ramachandran <sundar24295@gmail.com>
This commit is contained in:
jsheng_Linkedin
2026-04-20 22:50:40 -07:00
committed by GitHub
co-authored by Chanh Nguyen Sundara Raman Ramachandran
parent cfd49e233c
commit a8e3a534a4
17 changed files with 1071 additions and 431 deletions
@@ -20,6 +20,7 @@ from sglang.srt.managers.tokenizer_manager import TokenizerManager
from sglang.srt.managers.tokenizer_manager_score_mixin import (
TokenizerManagerScoreMixin,
)
from sglang.srt.server_args import MIS_DELIMITER_TOKEN_ID
from sglang.test.ci.ci_register import register_cuda_ci
from sglang.test.test_utils import CustomTestCase
@@ -204,16 +205,15 @@ class TestEmbeddingReqInputEmbedOverride(CustomTestCase):
class _FakeServerArgs:
"""Minimal stub for server_args."""
def __init__(self, multi_item_scoring_delimiter=None):
self.multi_item_scoring_delimiter = multi_item_scoring_delimiter
def __init__(self, enable_mis=False):
self.enable_mis = enable_mis
class _FakeMixin(TokenizerManagerScoreMixin):
"""Minimal stub to call mixin methods without a full TokenizerManager."""
def __init__(self, delimiter=None):
self.server_args = _FakeServerArgs(delimiter)
self.multi_item_delimiter_text = None
def __init__(self, enable_mis=False):
self.server_args = _FakeServerArgs(enable_mis)
self.tokenizer = None
self.is_generation = True
@@ -334,17 +334,17 @@ class TestResolveEmbedOverridesForRequest(CustomTestCase):
# Score mixin: _build_token_id_inputs
# ========================================================================
DELIM_TOKEN = 99
DELIM_TOKEN = MIS_DELIMITER_TOKEN_ID
class TestBuildTokenIdInputs(CustomTestCase):
def setUp(self):
self.mixin = _FakeMixin(delimiter=DELIM_TOKEN)
self.mixin = _FakeMixin(enable_mis=True)
# --- single-item mode, no embeds ---
def test_single_item_no_embeds(self):
_, input_ids, injection = self.mixin._build_token_id_inputs(
_, input_ids, positional_embed_overrides, _ = self.mixin._build_token_id_inputs(
query=[1, 2],
items=[[3, 4], [5, 6]],
item_first=False,
@@ -354,10 +354,10 @@ class TestBuildTokenIdInputs(CustomTestCase):
item_embed_overrides=None,
)
self.assertEqual(input_ids, [[1, 2, 3, 4], [1, 2, 5, 6]])
self.assertIsNone(injection)
self.assertIsNone(positional_embed_overrides)
def test_single_item_no_embeds_item_first(self):
_, input_ids, injection = self.mixin._build_token_id_inputs(
_, input_ids, positional_embed_overrides, _ = self.mixin._build_token_id_inputs(
query=[1, 2],
items=[[3, 4]],
item_first=True,
@@ -367,12 +367,12 @@ class TestBuildTokenIdInputs(CustomTestCase):
item_embed_overrides=None,
)
self.assertEqual(input_ids, [[3, 4, 1, 2]])
self.assertIsNone(injection)
self.assertIsNone(positional_embed_overrides)
# --- multi-item mode, no embeds ---
def test_multi_item_no_embeds(self):
_, input_ids, injection = self.mixin._build_token_id_inputs(
_, input_ids, positional_embed_overrides, _ = self.mixin._build_token_id_inputs(
query=[1, 2],
items=[[3, 4], [5, 6]],
item_first=False,
@@ -385,13 +385,13 @@ class TestBuildTokenIdInputs(CustomTestCase):
self.assertEqual(
input_ids, [[1, 2, DELIM_TOKEN, 3, 4, DELIM_TOKEN, 5, 6, DELIM_TOKEN]]
)
self.assertIsNone(injection)
self.assertIsNone(positional_embed_overrides)
# --- single-item mode, with embeds ---
def test_single_item_query_embeds(self):
"""Query placeholder overrides are resolved per item."""
_, input_ids, injection = self.mixin._build_token_id_inputs(
_, input_ids, positional_embed_overrides, _ = self.mixin._build_token_id_inputs(
query=[50, 10],
items=[[20, 30], [40, 50]],
item_first=False,
@@ -401,15 +401,15 @@ class TestBuildTokenIdInputs(CustomTestCase):
item_embed_overrides=None,
)
self.assertEqual(input_ids, [[50, 10, 20, 30], [50, 10, 40, 50]])
self.assertIsNotNone(injection)
self.assertEqual(len(injection), 2)
self.assertIsNotNone(positional_embed_overrides)
self.assertEqual(len(positional_embed_overrides), 2)
# Each item gets its own PositionalEmbeds with query override at pos 0
self.assertEqual(injection[0].positions, [0])
self.assertEqual(injection[1].positions, [0])
self.assertEqual(positional_embed_overrides[0].positions, [0])
self.assertEqual(positional_embed_overrides[1].positions, [0])
def test_single_item_item_embeds(self):
"""Per-item overrides with correct position offsets."""
_, input_ids, injection = self.mixin._build_token_id_inputs(
_, input_ids, positional_embed_overrides, _ = self.mixin._build_token_id_inputs(
query=[10, 20],
items=[[50, 30]],
item_first=False,
@@ -419,13 +419,13 @@ class TestBuildTokenIdInputs(CustomTestCase):
item_embed_overrides=[[_vec(2)]],
)
self.assertEqual(input_ids, [[10, 20, 50, 30]])
self.assertIsNotNone(injection)
self.assertIsNotNone(positional_embed_overrides)
# item placeholder at index 0 of item, offset by query length 2
self.assertEqual(injection[0].positions, [2])
self.assertEqual(positional_embed_overrides[0].positions, [2])
def test_single_item_no_override_positions_returns_none_injection(self):
"""When no items have placeholders, injection should be None."""
_, input_ids, injection = self.mixin._build_token_id_inputs(
"""When no items have placeholders, positional_embed_overrides should be None."""
_, input_ids, positional_embed_overrides, _ = self.mixin._build_token_id_inputs(
query=[10, 20],
items=[[30, 40]],
item_first=False,
@@ -434,11 +434,11 @@ class TestBuildTokenIdInputs(CustomTestCase):
query_embed_overrides=None,
item_embed_overrides=[None],
)
self.assertIsNone(injection)
self.assertIsNone(positional_embed_overrides)
def test_single_item_query_and_item_embeds(self):
"""Single-item mode with both query and item overrides in one request."""
_, input_ids, injection = self.mixin._build_token_id_inputs(
_, input_ids, positional_embed_overrides, _ = self.mixin._build_token_id_inputs(
query=[50, 10],
items=[[20, 50]],
item_first=False,
@@ -448,15 +448,15 @@ class TestBuildTokenIdInputs(CustomTestCase):
item_embed_overrides=[[_vec(2)]],
)
self.assertEqual(input_ids, [[50, 10, 20, 50]])
self.assertIsNotNone(injection)
pe = injection[0]
self.assertIsNotNone(positional_embed_overrides)
pe = positional_embed_overrides[0]
# query override at pos 0, item override at pos 3 (query_len=2 + idx=1)
self.assertEqual(pe.positions, [0, 3])
self.assertEqual(pe.embeds.shape, (2, HIDDEN_DIM))
def test_single_item_empty_query(self):
"""Empty query with item-only overrides (valid from score_prompts)."""
_, input_ids, injection = self.mixin._build_token_id_inputs(
_, input_ids, positional_embed_overrides, _ = self.mixin._build_token_id_inputs(
query=[],
items=[[50, 10]],
item_first=False,
@@ -466,15 +466,15 @@ class TestBuildTokenIdInputs(CustomTestCase):
item_embed_overrides=[[_vec(1)]],
)
self.assertEqual(input_ids, [[50, 10]])
self.assertIsNotNone(injection)
self.assertIsNotNone(positional_embed_overrides)
# item placeholder at absolute pos 0 (offset=len([])=0)
self.assertEqual(injection[0].positions, [0])
self.assertEqual(positional_embed_overrides[0].positions, [0])
# --- multi-item mode, with embeds ---
def test_multi_item_with_query_and_item_embeds(self):
"""Multi-item mode resolves query overrides once and item overrides per item."""
_, input_ids, injection = self.mixin._build_token_id_inputs(
_, input_ids, positional_embed_overrides, _ = self.mixin._build_token_id_inputs(
query=[50, 10],
items=[[20, 50], [30, 40]],
item_first=False,
@@ -483,13 +483,13 @@ class TestBuildTokenIdInputs(CustomTestCase):
query_embed_overrides=[_vec(1)],
item_embed_overrides=[[_vec(2)], None],
)
# query<D>item1<D>item2<D> = [50,10, 99, 20,50, 99, 30,40, 99]
# query<D>item1<D>item2<D> = [50,10, DELIM, 20,50, DELIM, 30,40, DELIM]
self.assertEqual(len(input_ids), 1)
self.assertIsNotNone(injection)
self.assertIsNotNone(positional_embed_overrides)
self.assertEqual(
len(injection), 1
len(positional_embed_overrides), 1
) # single PositionalEmbeds for combined sequence
pe = injection[0]
pe = positional_embed_overrides[0]
# query override at pos 0, item[0] override at pos 4 (query_len=2 + delim=1 + idx=1)
self.assertIn(0, pe.positions)
self.assertIn(4, pe.positions)
@@ -0,0 +1,599 @@
"""Tests for the Multi-Item Scoring (MIS) optimization.
MIS is a server-side optimization enabled via --enable-mis that batches
multiple items into a single forward pass using delimiter tokens (token ID 9999).
This is different from batch scoring (multiple items in one API call) which
processes items as separate requests.
The key difference:
- Batch scoring: N items -> N separate forward passes
- MIS optimization: N items -> 1 forward pass with delimiter-separated items
These tests ensure the MIS optimization produces correct results and catches
bugs in tensor shape handling (e.g., 2D tensors [num_delimiters, num_label_tokens]).
"""
import asyncio
import os
import unittest
import torch
from transformers import AutoConfig, AutoTokenizer
from sglang.srt.entrypoints.engine import Engine
from sglang.test.ci.ci_register import register_cuda_ci
from sglang.test.test_utils import (
DEFAULT_SMALL_MODEL_NAME_FOR_TEST,
CustomTestCase,
)
register_cuda_ci(est_time=240, suite="stage-b-test-1-gpu-small")
TEST_MODEL_NAME = os.environ.get("TEST_MODEL_NAME", DEFAULT_SMALL_MODEL_NAME_FOR_TEST)
TEST_CLASSIFICATION_BASE_MODEL = os.environ.get(
"TEST_CLASSIFICATION_BASE_MODEL",
"tomaarsen/Qwen3-Reranker-0.6B-seq-cls",
)
_CLS_NUM_LABELS = AutoConfig.from_pretrained(TEST_CLASSIFICATION_BASE_MODEL).num_labels
class TestMISServerArgsValidation(unittest.TestCase):
"""Test ServerArgs defaults for MIS mode."""
def test_enable_mis_default(self):
"""Test that enable_mis defaults to False."""
from sglang.srt.server_args import ServerArgs
self.assertEqual(ServerArgs.enable_mis, False)
class TestMultiItemScoringOptimization(CustomTestCase):
"""Test the Multi-Item Scoring (MIS) optimization with generation models."""
@classmethod
def setUpClass(cls):
cls.engine = Engine(
model_path=TEST_MODEL_NAME,
disable_radix_cache=True,
chunked_prefill_size=-1,
enable_mis=True,
attention_backend="flashinfer",
mem_fraction_static=0.15,
)
cls.non_mis_engine = Engine(
model_path=TEST_MODEL_NAME,
disable_radix_cache=True,
chunked_prefill_size=-1,
mem_fraction_static=0.15,
)
@classmethod
def tearDownClass(cls):
if cls.engine is not None:
cls.engine.shutdown()
if cls.non_mis_engine is not None:
cls.non_mis_engine.shutdown()
torch.cuda.empty_cache()
def test_mis_basic(self):
"""Test basic MIS: correct shapes, valid probabilities."""
query = "Rate each option:"
items = ["Option A", "Option B", "Option C"]
label_token_ids = [9454, 2753] # "Yes" and "No" tokens
scores = self.engine.score(
query=query,
items=items,
label_token_ids=label_token_ids,
apply_softmax=True,
).scores
self.assertEqual(len(scores), len(items))
for i, score_list in enumerate(scores):
self.assertEqual(len(score_list), len(label_token_ids))
self.assertAlmostEqual(sum(score_list), 1.0, places=5)
for score in score_list:
self.assertGreaterEqual(score, 0)
self.assertLessEqual(score, 1)
def test_mis_consistency_with_single_item(self):
"""MIS with one item should match non-MIS scoring closely."""
query = "Is this a fact?\n"
items = [" The sun rises in the east"]
label_token_ids = [9454, 2753]
mis_scores = self.engine.score(
query=query,
items=items,
label_token_ids=label_token_ids,
apply_softmax=True,
).scores
non_mis_scores = self.non_mis_engine.score(
query=query,
items=items,
label_token_ids=label_token_ids,
apply_softmax=True,
).scores
self.assertEqual(len(mis_scores), 1)
self.assertEqual(len(non_mis_scores), 1)
for j, (m, n) in enumerate(zip(mis_scores[0], non_mis_scores[0])):
relative_diff = abs(m - n) / max(abs(n), 1e-6)
self.assertLess(
relative_diff,
0.08,
msg=f"label {j}: MIS={m} vs non-MIS={n} (diff: {relative_diff:.3f})",
)
def test_mis_empty_query(self):
"""MIS with empty query — delimiter indices start at position 0."""
items = ["alpha", "beta"]
label_token_ids = [9454, 2753]
scores = self.engine.score(
query="",
items=items,
label_token_ids=label_token_ids,
apply_softmax=True,
).scores
self.assertEqual(len(scores), len(items))
for score_list in scores:
self.assertEqual(len(score_list), len(label_token_ids))
self.assertAlmostEqual(sum(score_list), 1.0, places=5)
class TestMultiItemScoringClassification(CustomTestCase):
"""Test MIS with classification models.
Uses a pre-trained Qwen3ForSequenceClassification model so that the
classification head weights are deterministic across Engine instances.
"""
NUM_LABELS = _CLS_NUM_LABELS
def setUp(self):
self.engine = Engine(
model_path=TEST_CLASSIFICATION_BASE_MODEL,
disable_radix_cache=True,
chunked_prefill_size=-1,
enable_mis=True,
attention_backend="flashinfer",
mem_fraction_static=0.15,
)
def tearDown(self):
if self.engine is not None:
self.engine.shutdown()
torch.cuda.empty_cache()
def test_classification_mis_basic(self):
"""Classification MIS: correct shapes, valid softmax probabilities."""
query = "Rate each option:"
items = ["Option A", "Option B", "Option C"]
scores = self.engine.score(query=query, 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)
self.assertAlmostEqual(sum(score_list), 1.0, places=5)
for score in score_list:
self.assertGreaterEqual(score, 0)
self.assertLessEqual(score, 1)
def test_classification_mis_tokenized_input(self):
"""Classification MIS with pre-tokenized query and items."""
tokenizer = AutoTokenizer.from_pretrained(TEST_CLASSIFICATION_BASE_MODEL)
query_ids = tokenizer.encode("Rate each option:", add_special_tokens=False)
items_ids = [
tokenizer.encode(item, add_special_tokens=False)
for item in ["Option A", "Option B"]
]
scores = self.engine.score(
query=query_ids, items=items_ids, apply_softmax=True
).scores
self.assertEqual(len(scores), len(items_ids))
for score_list in scores:
self.assertEqual(len(score_list), self.NUM_LABELS)
self.assertAlmostEqual(sum(score_list), 1.0, places=5)
def test_classification_non_mis_fallback(self):
"""Classification model works correctly without --enable-mis."""
non_mis_engine = Engine(
model_path=TEST_CLASSIFICATION_BASE_MODEL,
disable_radix_cache=True,
mem_fraction_static=0.15,
)
try:
scores = non_mis_engine.score(
query="Test:", items=["A", "B"], apply_softmax=True
).scores
self.assertEqual(len(scores), 2)
for score_list in scores:
self.assertEqual(len(score_list), self.NUM_LABELS)
self.assertAlmostEqual(sum(score_list), 1.0, places=5)
finally:
non_mis_engine.shutdown()
torch.cuda.empty_cache()
class TestMultiItemScoringParity(CustomTestCase):
"""Test that MIS produces the same results as single-item scoring."""
@classmethod
def setUpClass(cls):
cls.engine_single = Engine(
model_path=TEST_MODEL_NAME,
disable_radix_cache=True,
log_level="error",
mem_fraction_static=0.15,
)
cls.engine_mis = Engine(
model_path=TEST_MODEL_NAME,
disable_radix_cache=True,
chunked_prefill_size=-1,
log_level="error",
enable_mis=True,
attention_backend="flashinfer",
mem_fraction_static=0.15,
)
@classmethod
def tearDownClass(cls):
if cls.engine_single is not None:
cls.engine_single.shutdown()
if cls.engine_mis is not None:
cls.engine_mis.shutdown()
torch.cuda.empty_cache()
def _compare_scores(
self, query, items, label_token_ids=None, apply_softmax=True, test_name=""
):
"""Compare MIS vs single-item scoring results."""
single_scores = self.engine_single.score(
query=query,
items=items,
label_token_ids=label_token_ids,
apply_softmax=apply_softmax,
).scores
mis_scores = self.engine_mis.score(
query=query,
items=items,
label_token_ids=label_token_ids,
apply_softmax=apply_softmax,
).scores
self.assertEqual(
len(mis_scores), len(single_scores), f"{test_name}: count mismatch"
)
for i, (ms, ss) in enumerate(zip(mis_scores, single_scores)):
self.assertEqual(len(ms), len(ss), f"{test_name}: item {i} length mismatch")
for j, (m, s) in enumerate(zip(ms, ss)):
self.assertAlmostEqual(
m,
s,
places=1,
msg=f"{test_name}: item {i} label {j}: MIS={m} vs single={s}",
)
def test_parity_basic(self):
tokenizer = AutoTokenizer.from_pretrained(TEST_MODEL_NAME)
query = "Rate this option:"
items = [" Option A", " Option B", " Option C"]
labels = [" good", " bad"]
label_ids = [tokenizer.encode(lb, add_special_tokens=False)[0] for lb in labels]
self._compare_scores(query, items, label_ids, test_name="basic")
def test_parity_tokenized_inputs(self):
tokenizer = AutoTokenizer.from_pretrained(TEST_MODEL_NAME)
query = "Rate this option:"
items = [" Option X", " Option Y"]
labels = [" good", " bad"]
query_ids = tokenizer.encode(query, add_special_tokens=False)
items_ids = [tokenizer.encode(i, add_special_tokens=False) for i in items]
label_ids = [tokenizer.encode(lb, add_special_tokens=False)[0] for lb in labels]
self._compare_scores(query_ids, items_ids, label_ids, test_name="tokenized")
def test_parity_without_softmax(self):
tokenizer = AutoTokenizer.from_pretrained(TEST_MODEL_NAME)
query = "The weather today is"
items = [" sunny", " cloudy", " rainy"]
labels = [" nice", " bad"]
label_ids = [tokenizer.encode(lb, add_special_tokens=False)[0] for lb in labels]
self._compare_scores(
query, items, label_ids, apply_softmax=False, test_name="no_softmax"
)
def test_parity_many_items(self):
tokenizer = AutoTokenizer.from_pretrained(TEST_MODEL_NAME)
query = "Rate this option from 1 to 5:"
items = [f" Option {i}" for i in range(10)]
labels = [" 1", " 2", " 3", " 4", " 5"]
label_ids = [tokenizer.encode(lb, add_special_tokens=False)[0] for lb in labels]
self._compare_scores(query, items, label_ids, test_name="many_items")
class TestMultiItemScoringClassificationParity(CustomTestCase):
"""Test that MIS multi-item batching matches single-item MIS scoring.
Both paths use the MIS engine (with delimiter tokens in the attention
context). The reference scores each item individually so each gets its
own forward pass; the batched path packs all items into one pass.
This isolates the MIS batching logic from the delimiter-presence effect.
"""
NUM_LABELS = _CLS_NUM_LABELS
@classmethod
def setUpClass(cls):
cls.engine = Engine(
model_path=TEST_CLASSIFICATION_BASE_MODEL,
disable_radix_cache=True,
chunked_prefill_size=-1,
enable_mis=True,
attention_backend="flashinfer",
mem_fraction_static=0.15,
)
@classmethod
def tearDownClass(cls):
if cls.engine is not None:
cls.engine.shutdown()
torch.cuda.empty_cache()
def _compare_scores(self, query, items, apply_softmax=True, test_name=""):
"""Compare MIS batched vs MIS single-item scoring results."""
single_scores = []
for item in items:
result = self.engine.score(
query=query,
items=[item],
apply_softmax=apply_softmax,
).scores
single_scores.append(result[0])
batched_scores = self.engine.score(
query=query,
items=items,
apply_softmax=apply_softmax,
).scores
self.assertEqual(
len(batched_scores), len(single_scores), f"{test_name}: count mismatch"
)
for i, (bs, ss) in enumerate(zip(batched_scores, single_scores)):
self.assertEqual(len(bs), len(ss), f"{test_name}: item {i} length mismatch")
for j, (b, s) in enumerate(zip(bs, ss)):
self.assertAlmostEqual(
b,
s,
places=1,
msg=f"{test_name}: item {i} label {j}: batched={b} vs single={s}",
)
def test_parity_basic(self):
query = "Rate this option:"
items = [" Option A", " Option B", " Option C"]
self._compare_scores(query, items, test_name="cls_basic")
def test_parity_tokenized_inputs(self):
tokenizer = AutoTokenizer.from_pretrained(TEST_CLASSIFICATION_BASE_MODEL)
query_ids = tokenizer.encode("Rate this option:", add_special_tokens=False)
items_ids = [
tokenizer.encode(item, add_special_tokens=False)
for item in [" Option X", " Option Y"]
]
self._compare_scores(query_ids, items_ids, test_name="cls_tokenized")
def test_parity_without_softmax(self):
query = "The weather today is"
items = [" sunny", " cloudy", " rainy"]
self._compare_scores(
query, items, apply_softmax=False, test_name="cls_no_softmax"
)
def test_parity_many_items(self):
query = "Classify this option:"
items = [f" Option {i}" for i in range(10)]
self._compare_scores(query, items, test_name="cls_many_items")
class TestMultiItemScoringClassificationMISvsNonMIS(CustomTestCase):
"""Test that MIS single-item approximates non-MIS single-item.
The MIS path inserts delimiter tokens into the attention context,
which slightly perturbs hidden states. After softmax the scores
should still be close. Uses places=1 (±0.05) tolerance.
Runs as a separate class so each engine is created and destroyed
independently to avoid GPU OOM.
"""
def test_mis_single_vs_non_mis(self):
non_mis_engine = Engine(
model_path=TEST_CLASSIFICATION_BASE_MODEL,
disable_radix_cache=True,
mem_fraction_static=0.15,
)
try:
query = "Rate this option:"
items = [" Option A", " Option B", " Option C"]
non_mis_scores = non_mis_engine.score(
query=query,
items=items,
apply_softmax=True,
).scores
finally:
non_mis_engine.shutdown()
torch.cuda.empty_cache()
mis_engine = Engine(
model_path=TEST_CLASSIFICATION_BASE_MODEL,
disable_radix_cache=True,
chunked_prefill_size=-1,
enable_mis=True,
attention_backend="flashinfer",
mem_fraction_static=0.15,
)
try:
mis_scores = mis_engine.score(
query=query,
items=items,
apply_softmax=True,
).scores
finally:
mis_engine.shutdown()
torch.cuda.empty_cache()
self.assertEqual(len(mis_scores), len(non_mis_scores))
for i, (ms, ns) in enumerate(zip(mis_scores, non_mis_scores)):
self.assertEqual(len(ms), len(ns))
for j, (m, n) in enumerate(zip(ms, ns)):
self.assertAlmostEqual(
m,
n,
places=1,
msg=f"item {i} label {j}: MIS={m} vs non-MIS={n}",
)
class TestMultiItemScoringClassificationAdvanced(CustomTestCase):
"""Advanced MIS tests for classification models: score distinctness,
determinism, and concurrent request handling."""
NUM_LABELS = _CLS_NUM_LABELS
@classmethod
def setUpClass(cls):
cls.engine = Engine(
model_path=TEST_CLASSIFICATION_BASE_MODEL,
disable_radix_cache=True,
chunked_prefill_size=-1,
enable_mis=True,
attention_backend="flashinfer",
mem_fraction_static=0.15,
)
@classmethod
def tearDownClass(cls):
if cls.engine is not None:
cls.engine.shutdown()
torch.cuda.empty_cache()
def test_items_produce_distinct_scores(self):
"""Different items must produce different score vectors.
Core regression test: before the delimiter-index fix, all items got
identical scores because the MIS attention mask only let delimiter
tokens attend to the query prefix.
"""
query = "Rate each option:"
items = [
"Option A is about cats",
"Option B is about dogs",
"Option C is about fish",
]
scores = self.engine.score(query=query, 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 broken. Scores: {scores[0]}",
)
def test_many_items_distinct(self):
"""Stress test: 15 items should not all produce identical scores."""
query = "Classify each city:"
items = [f"City {i}" for i in range(15)]
scores = self.engine.score(query=query, items=items).scores
self.assertEqual(len(scores), len(items))
for score_list in scores:
self.assertEqual(len(score_list), self.NUM_LABELS)
unique_count = len({tuple(s) for s in scores})
self.assertGreater(unique_count, 1, "All 15 items returned identical scores")
def test_deterministic(self):
"""Identical requests should return identical scores."""
query = "Evaluate:"
items = ["alpha", "beta", "gamma"]
scores1 = self.engine.score(query=query, items=items).scores
scores2 = self.engine.score(query=query, items=items).scores
self.assertEqual(
scores1, scores2, "Identical inputs must produce identical scores"
)
def test_concurrent_requests(self):
"""Concurrent MIS requests must produce the same scores as sequential.
Runs each request sequentially to get baseline scores, then runs all
concurrently and asserts the results match. This catches cross-request
contamination when multiple MIS requests share a GPU batch.
"""
test_cases = [
{"query": "Is this a fruit?", "items": ["apple", "car", "banana"]},
{"query": "Is this an animal?", "items": ["dog", "table"]},
{
"query": "Is this a country?",
"items": ["France", "pizza", "Japan", "chair"],
},
{"query": "Is this a color?", "items": ["red"]},
]
# Sequential baseline
sequential_scores = []
for tc in test_cases:
result = self.engine.score(query=tc["query"], items=tc["items"])
sequential_scores.append(result.scores)
# Concurrent execution
async def _gather():
return await asyncio.gather(
*(
self.engine.async_score(query=tc["query"], items=tc["items"])
for tc in test_cases
)
)
concurrent_results = self.engine.loop.run_until_complete(_gather())
for idx, (tc, seq_scores, conc_result) in enumerate(
zip(test_cases, sequential_scores, concurrent_results)
):
conc_scores = conc_result.scores
self.assertEqual(
len(conc_scores),
len(seq_scores),
f"Case {idx}: count mismatch",
)
for i, (cs, ss) in enumerate(zip(conc_scores, seq_scores)):
self.assertEqual(
len(cs),
len(ss),
f"Case {idx} item {i}: label count mismatch",
)
for j, (c, s) in enumerate(zip(cs, ss)):
self.assertAlmostEqual(
c,
s,
places=1,
msg=f"Case {idx} item {i} label {j}: "
f"concurrent={c} vs sequential={s}",
)
if __name__ == "__main__":
unittest.main()
@@ -30,7 +30,6 @@ from sglang.test.test_utils import (
register_cuda_ci(est_time=240, suite="stage-b-test-1-gpu-small")
_SEQCLS_MODEL = "Qwen/Qwen3-0.6B"
_QWEN3_EOT_TOKEN_ID = 151643
_CAUSAL_LM_MODEL = DEFAULT_SMALL_MODEL_NAME_FOR_TEST
_NUM_LABELS = 4
@@ -197,7 +196,7 @@ class TestPooledHiddenStatesMISEngine(CustomTestCase):
model_path=_SEQCLS_MODEL,
disable_radix_cache=True,
chunked_prefill_size=-1,
multi_item_scoring_delimiter=_QWEN3_EOT_TOKEN_ID,
enable_mis=True,
json_model_override_args=json.dumps(
{
"architectures": ["Qwen3ForSequenceClassification"],
+11 -11
View File
@@ -3,15 +3,16 @@
Two test classes, each with its own server instance:
TestCausalLMScoringHTTP — basic endpoint: schema defaults, response
structure, error rejection (no MIS delimiter)
TestCausalLMMISScoringHTTP — MIS mode: validates --multi-item-scoring-delimiter
CLI flag wiring and per-item output shape
structure, error rejection (no MIS)
TestCausalLMMISScoringHTTP — MIS mode: validates --enable-mis CLI flag
wiring and per-item output shape
Engine-level correctness (numerical accuracy, batching, edge cases) lives in
test_score_engine.py. These tests focus on the HTTP integration seam:
Pydantic schema defaults, FastAPI routing, and server argument wiring.
"""
import os
import unittest
import requests
@@ -28,9 +29,7 @@ from sglang.test.test_utils import (
register_cuda_ci(est_time=70, suite="stage-b-test-1-gpu-small")
_MODEL = DEFAULT_SMALL_MODEL_NAME_FOR_TEST # Llama-3.2-1B-Instruct
# <|eot_id|> for Llama-3.x Instruct — used as MIS delimiter
_LLAMA3_EOT_TOKEN_ID = 128009
_MODEL = os.environ.get("TEST_MODEL_NAME", DEFAULT_SMALL_MODEL_NAME_FOR_TEST)
# ---------------------------------------------------------------------------
@@ -41,7 +40,7 @@ _LLAMA3_EOT_TOKEN_ID = 128009
class TestCausalLMScoringHTTP(CustomTestCase):
"""Validates /v1/score HTTP integration — schema, defaults, and error handling.
Starts a plain CausalLM server (no --multi-item-scoring-delimiter) to test
Starts a plain CausalLM server (no --enable-mis) to test
the HTTP layer in isolation: response envelope shape, the apply_softmax
default (False), and Pydantic validation errors on malformed input.
"""
@@ -139,12 +138,12 @@ class TestCausalLMScoringHTTP(CustomTestCase):
# ---------------------------------------------------------------------------
# MIS scoring (with --multi-item-scoring-delimiter)
# MIS scoring (with --enable-mis)
# ---------------------------------------------------------------------------
class TestCausalLMMISScoringHTTP(CustomTestCase):
"""Validates /v1/score with --multi-item-scoring-delimiter.
"""Validates /v1/score with --enable-mis.
Confirms that the CLI flag is correctly wired into ServerArgs and that the
endpoint returns one probability vector per item when items are
@@ -163,8 +162,9 @@ class TestCausalLMMISScoringHTTP(CustomTestCase):
"--disable-radix-cache",
"--chunked-prefill-size",
"-1",
"--multi-item-scoring-delimiter",
str(_LLAMA3_EOT_TOKEN_ID),
"--enable-mis",
"--attention-backend",
"flashinfer",
],
)
@@ -4,14 +4,18 @@ Two model types, two scoring modes:
TestCausalLMScoring — CausalLM, single-item and batched multi-item
TestSeqClsScoring — SequenceClassification, single-item mode
TestSeqClsMISScoring — SequenceClassification, MIS delimiter mode
TestSeqClsMISScoring — SequenceClassification, MIS mode (--enable-mis)
TestSeqClsMISAdvancedScoring — SeqCls MIS with 12 labels (tensor shape stress)
The Engine (Python API) is the right layer for correctness testing: it
exercises tokenization, forward pass, pooling, and score extraction without
the HTTP serialization overhead. HTTP-layer tests live in test_score_api.py.
Thorough MIS tests (parity, concurrency, generation models) live in
test_multi_item_scoring.py.
"""
import json
import os
import unittest
from unittest.mock import patch
@@ -24,10 +28,8 @@ from sglang.test.test_utils import DEFAULT_SMALL_MODEL_NAME_FOR_TEST, CustomTest
register_cuda_ci(est_time=85, suite="stage-b-test-1-gpu-small")
_CAUSAL_LM_MODEL = DEFAULT_SMALL_MODEL_NAME_FOR_TEST # Llama-3.2-1B-Instruct
_SEQCLS_MODEL = "Qwen/Qwen3-0.6B" # backbone; arch overridden to SeqCls below
# <|endoftext|> for Qwen3 tokenizer — used as MIS delimiter
_QWEN3_EOT_TOKEN_ID = 151643
_CAUSAL_LM_MODEL = os.environ.get("TEST_MODEL_NAME", DEFAULT_SMALL_MODEL_NAME_FOR_TEST)
_SEQCLS_MODEL = os.environ.get("TEST_CLASSIFICATION_BASE_MODEL", "Qwen/Qwen3-0.6B")
# ---------------------------------------------------------------------------
@@ -368,10 +370,9 @@ class TestSeqClsScoring(CustomTestCase):
class TestSeqClsMISScoring(CustomTestCase):
"""SeqCls MIS: all items packed into one sequence separated by delimiter token.
score_and_pool() extracts per-item scores at delimiter positions.
Two sub-cases are tested:
- NUM_LABELS=2 — standard binary classification head
- NUM_LABELS=12 — stress-tests 2-D tensor indexing in score_and_pool()
Uses --enable-mis which hardcodes delimiter token ID 9999.
Basic pipeline correctness only — thorough MIS tests (parity,
concurrency, advanced) live in test_multi_item_scoring.py.
"""
NUM_LABELS = 2
@@ -382,7 +383,8 @@ class TestSeqClsMISScoring(CustomTestCase):
model_path=_SEQCLS_MODEL,
disable_radix_cache=True,
chunked_prefill_size=-1,
multi_item_scoring_delimiter=_QWEN3_EOT_TOKEN_ID,
enable_mis=True,
attention_backend="flashinfer",
json_model_override_args=json.dumps(
{
"architectures": ["Qwen3ForSequenceClassification"],
@@ -432,32 +434,6 @@ class TestSeqClsMISScoring(CustomTestCase):
self.assertEqual(len(row), self.NUM_LABELS)
self.assertAlmostEqual(sum(row), 1.0, places=5)
def test_mis_items_produce_distinct_scores(self):
"""Different items must yield different score vectors.
Catches bugs where all delimiter positions share the same pooled
hidden state (e.g. off-by-one in score_and_pool indexing).
"""
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))
self.assertFalse(
all(scores[0] == s for s in scores[1:]),
f"All items returned identical scores — delimiter indexing is likely broken. "
f"Scores: {scores[0]}",
)
def test_mis_deterministic(self):
"""Identical MIS requests return identical scores."""
kwargs = dict(query="Evaluate:", items=["alpha", "beta", "gamma"])
self.assertEqual(
self.engine.score(**kwargs).scores, self.engine.score(**kwargs).scores
)
# ---------------------------------------------------------------------------
# SequenceClassification — MIS with many labels (tensor shape stress test)
@@ -479,7 +455,8 @@ class TestSeqClsMISAdvancedScoring(CustomTestCase):
model_path=_SEQCLS_MODEL,
disable_radix_cache=True,
chunked_prefill_size=-1,
multi_item_scoring_delimiter=_QWEN3_EOT_TOKEN_ID,
enable_mis=True,
attention_backend="flashinfer",
json_model_override_args=json.dumps(
{
"architectures": ["Qwen3ForSequenceClassification"],
@@ -506,17 +483,6 @@ class TestSeqClsMISAdvancedScoring(CustomTestCase):
self.assertEqual(len(row), self.NUM_LABELS)
self.assertAlmostEqual(sum(row), 1.0, places=5)
def test_many_items_produce_distinct_scores(self):
"""15 items should not all return identical score vectors."""
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))
self.assertGreater(
len({tuple(s) for s in scores}),
1,
"All 15 items returned identical scores",
)
if __name__ == "__main__":
unittest.main(verbosity=3)
@@ -1,12 +1,11 @@
"""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.
All tests run on CPU — no GPU required. MIS delimiter positions are passed
via forward_batch.multi_item_delimiter_indices (pre-computed by the caller).
"""
import unittest
from types import SimpleNamespace
from unittest.mock import patch
import torch
import torch.nn as nn
@@ -24,22 +23,22 @@ register_cpu_ci(est_time=9, suite="stage-a-test-cpu")
def _make_forward_batch(
extend_seq_lens, is_prefill_only=False, return_pooled_hidden_states=False
extend_seq_lens,
multi_item_delimiter_indices=None,
return_pooled_hidden_states=False,
is_prefill_only=True,
):
"""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,
multi_item_delimiter_indices=multi_item_delimiter_indices,
dimensions=None,
return_pooled_hidden_states=return_pooled_hidden_states,
is_prefill_only=is_prefill_only,
)
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."""
@@ -50,11 +49,8 @@ class TestScoreAndPool(CustomTestCase):
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)
def test_single_item_returns_scores(self):
"""No delimiter indices -> single-item path returns [batch, num_labels]."""
hidden = torch.randn(8, self.hidden_dim)
fb = _make_forward_batch(extend_seq_lens=[5, 3])
input_ids = torch.arange(8)
@@ -64,30 +60,16 @@ class TestScoreAndPool(CustomTestCase):
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,
]
)
def test_mis_returns_per_request_list(self):
"""Delimiter indices provided -> returns a list with one tensor per request."""
# Sequence: [0, 1, 2, D, 3, 4, 5, D, 6, 7, 8, D]
# Delimiters at positions 3, 7, 11 -> extract at 2, 6, 10
input_ids = torch.arange(12)
hidden = torch.randn(len(input_ids), self.hidden_dim)
fb = _make_forward_batch(extend_seq_lens=[len(input_ids)], is_prefill_only=True)
fb = _make_forward_batch(
extend_seq_lens=[len(input_ids)],
multi_item_delimiter_indices=[torch.tensor([3, 7, 11])],
)
out = score_and_pool(self.score_head, self.pooler, hidden, fb, input_ids)
@@ -95,20 +77,20 @@ class TestScoreAndPool(CustomTestCase):
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):
def test_mis_batched_splits_per_request(self):
"""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]
# Request 1: [10, 11, D, 12, 13, D] -> delimiters at 2, 5
# Request 2: [20, 21, 22, D] -> delimiter at 3
req1 = [10, 11, 99, 12, 13, 99]
req2 = [20, 21, 22, 99]
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
extend_seq_lens=[len(req1), len(req2)],
multi_item_delimiter_indices=[
torch.tensor([2, 5]),
torch.tensor([3]),
],
)
out = score_and_pool(self.score_head, self.pooler, hidden, fb, input_ids)
@@ -118,42 +100,21 @@ class TestScoreAndPool(CustomTestCase):
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)
def test_no_delimiter_indices_falls_back(self):
"""multi_item_delimiter_indices=None -> single-item fallback."""
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)
fb = _make_forward_batch(extend_seq_lens=[5, 3])
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):
def test_mis_extracts_positions_before_delimiter(self):
"""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])
input_ids = torch.tensor([10, 11, 99, 20, 21, 99])
hidden = (
torch.arange(len(input_ids))
.unsqueeze(1)
@@ -161,7 +122,10 @@ class TestScoreAndPool(CustomTestCase):
.expand(-1, self.hidden_dim)
.clone()
)
fb = _make_forward_batch(extend_seq_lens=[len(input_ids)], is_prefill_only=True)
fb = _make_forward_batch(
extend_seq_lens=[len(input_ids)],
multi_item_delimiter_indices=[torch.tensor([2, 5])],
)
identity_head = nn.Linear(self.hidden_dim, self.hidden_dim, bias=False)
nn.init.eye_(identity_head.weight)
@@ -172,14 +136,9 @@ class TestScoreAndPool(CustomTestCase):
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])
def test_mis_delimiter_at_position_one(self):
"""Delimiters at positions 1 and 3 extract at indices 0 and 2."""
input_ids = torch.tensor([10, 99, 11, 99])
hidden = (
torch.arange(len(input_ids))
.unsqueeze(1)
@@ -187,7 +146,10 @@ class TestScoreAndPool(CustomTestCase):
.expand(-1, self.hidden_dim)
.clone()
)
fb = _make_forward_batch(extend_seq_lens=[len(input_ids)], is_prefill_only=True)
fb = _make_forward_batch(
extend_seq_lens=[len(input_ids)],
multi_item_delimiter_indices=[torch.tensor([1, 3])],
)
identity_head = nn.Linear(self.hidden_dim, self.hidden_dim, bias=False)
nn.init.eye_(identity_head.weight)
@@ -195,25 +157,37 @@ class TestScoreAndPool(CustomTestCase):
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)
self.assertEqual(out.embeddings[0].shape[0], 2)
torch.testing.assert_close(out.embeddings[0][0], hidden[0])
torch.testing.assert_close(out.embeddings[0][1], hidden[2])
def test_single_item_scores_match_manual_computation(self):
"""Single-item scores equal score_head applied to pooled hidden states."""
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
pooled = self.pooler(hidden, fb).embeddings
expected = self.score_head(pooled)
torch.testing.assert_close(out.embeddings, expected)
def test_empty_delimiter_indices(self):
"""Empty delimiter tensor per request -> returns list with empty tensor."""
input_ids = torch.arange(6)
hidden = torch.randn(6, self.hidden_dim)
fb = _make_forward_batch(
extend_seq_lens=[6],
multi_item_delimiter_indices=[torch.tensor([], dtype=torch.long)],
)
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, (0, self.num_labels))
if __name__ == "__main__":
unittest.main()