Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
17 changes: 13 additions & 4 deletions python/sglang/srt/managers/tokenizer_manager_score_mixin.py
Original file line number Diff line number Diff line change
Expand Up @@ -225,7 +225,7 @@ def _process_multi_item_scoring_results(
# ------------------------------------------------------------------

def _multi_position_score_rows(
self, embedding: Any, apply_softmax: bool
self, embedding: Any, apply_softmax: bool, *, temperature: float
) -> List[List[float]]:
"""Validate a 2-D multi-position result embedding and return its per-token rows."""
embedding_tensor = torch.as_tensor(embedding)
Expand All @@ -240,7 +240,10 @@ def _multi_position_score_rows(
"pool per position do)."
)
if apply_softmax:
return torch.softmax(embedding_tensor, dim=-1).tolist()
scores_tensor = embedding_tensor.to(torch.float64)
# center before scaling so small temperatures do not overflow
centered = scores_tensor - scores_tensor.amax(dim=-1, keepdim=True)
return torch.softmax(centered / temperature, dim=-1).tolist()
return embedding if isinstance(embedding, list) else embedding_tensor.tolist()

def _multi_position_phs_matrix(self, phs: Any, expected_rows: int) -> torch.Tensor:
Expand All @@ -259,6 +262,7 @@ def _process_multi_item_extraction_results(
per_item_anchor_counts: List[int],
apply_softmax: bool,
return_pooled_hidden_states: bool = False,
temperature: float = 1.0,
) -> ScoreResult:
"""Process a fused multi-item score-extraction request (``--enable-mis``).

Expand All @@ -275,7 +279,9 @@ def _process_multi_item_extraction_results(
if embedding is None:
raise ValueError("Embedding not found in the result.")

rows = self._multi_position_score_rows(embedding, apply_softmax)
rows = self._multi_position_score_rows(
embedding, apply_softmax, temperature=temperature
)
total_anchors = sum(per_item_anchor_counts)
if len(rows) != total_anchors:
raise RuntimeError(
Expand Down Expand Up @@ -361,7 +367,9 @@ def _process_single_item_scoring_results(
prompt_tokens += result.get("meta_info", {}).get("prompt_tokens", 0)

if per_item_matrix:
rows = self._multi_position_score_rows(embedding, apply_softmax)
rows = self._multi_position_score_rows(
embedding, apply_softmax, temperature=temperature
)
scores.append(rows)
else:
if apply_softmax:
Expand Down Expand Up @@ -1044,6 +1052,7 @@ async def score_request(
per_item_anchor_counts,
apply_softmax,
return_pooled_hidden_states,
temperature=temperature,
)
elif use_multi_item_scoring:
# Multi-item scoring: extract scores from input_token_ids_logprobs or embedding
Expand Down
10 changes: 7 additions & 3 deletions test/registered/unit/managers/test_setwise_score_mixin.py
Original file line number Diff line number Diff line change
Expand Up @@ -409,20 +409,24 @@ def test_resolve_multi_position_pooling_mis_rejects_query_side_anchor(self):
def test_multi_position_score_rows_preserves_list_values_without_softmax(self):
# No softmax returns the original list as-is (no float round-trip).
emb = [[0.1, 0.2], [0.3, 0.4]]
rows = self.h._multi_position_score_rows(emb, apply_softmax=False)
rows = self.h._multi_position_score_rows(
emb, apply_softmax=False, temperature=1.0
)
self.assertIs(rows, emb)

def test_multi_position_score_rows_softmax_over_labels(self):
rows = self.h._multi_position_score_rows(
[[0.0, 0.0], [2.0, 2.0]], apply_softmax=True
[[0.0, 0.0], [2.0, 2.0]], apply_softmax=True, temperature=1.0
)
for row in rows:
self.assertAlmostEqual(sum(row), 1.0, places=5)
self.assertAlmostEqual(row[0], 0.5, places=5)

def test_multi_position_score_rows_rejects_non_matrix(self):
with self.assertRaisesRegex(ValueError, "expected a 2-D"):
self.h._multi_position_score_rows([0.1, 0.2], apply_softmax=False)
self.h._multi_position_score_rows(
[0.1, 0.2], apply_softmax=False, temperature=1.0
)

def test_multi_position_phs_matrix_rejects_row_count_mismatch(self):
with self.assertRaisesRegex(ValueError, "one row per score position"):
Expand Down
42 changes: 39 additions & 3 deletions test/registered/unit/test_token_scoring.py
Original file line number Diff line number Diff line change
Expand Up @@ -24,8 +24,10 @@
class ScoringManager(TokenizerManagerScoreMixin):
"""Replace only model execution; keep request construction and score extraction real."""

def __init__(self, enable_mis=False, generation=True):
self.server_args = ServerArgs(model_path="dummy", enable_mis=enable_mis)
def __init__(self, enable_mis=False, generation=True, **server_args):
self.server_args = ServerArgs(
model_path="dummy", enable_mis=enable_mis, **server_args
)
publish(self.server_args, role="test")
self.is_generation = generation
self.tokenizer = SimpleNamespace(vocab_size=8)
Expand All @@ -51,7 +53,10 @@ async def generate_request(self, request, raw_request):
results.append({"meta_info": meta})
else:
embedding = [1.0, 3.0]
if self.server_args.enable_mis:
if request.token_indices_to_pool is not None:
# setwise pools the head at every anchor: one row per position
embedding = [embedding] * len(request.token_indices_to_pool[index])
elif self.server_args.enable_mis:
count = len(request.multi_item_delimiter_indices[index])
embedding = [embedding] * count
results.append({"meta_info": meta, "embedding": embedding})
Expand Down Expand Up @@ -157,6 +162,37 @@ async def test_classification_temperature(self):
query=[], items=[[4]], return_token_logprobs=True
)

async def test_setwise_classification_temperature(self):
"""Setwise rows honor temperature in both the batched and --enable-mis paths."""
cases = {2.0: torch.softmax(torch.tensor([0.5, 1.5]), 0), 1e-300: [0.0, 1.0]}
for enable_mis in (False, True):
for temperature, expected in cases.items():
with self.subTest(enable_mis=enable_mis, temperature=temperature):
manager = ScoringManager(
enable_mis=enable_mis,
generation=False,
disable_radix_cache=True,
chunked_prefill_size=-1,
)
manager.model_config = SimpleNamespace(
hf_config=SimpleNamespace(
architectures=["Qwen3ForSequenceClassification"]
)
)
result = await manager.score_request(
query=[4],
items=[[5, 7], [6, 7, 7]],
apply_softmax=True,
score_extraction_token_id=7,
temperature=temperature,
)
self.assertEqual([len(rows) for rows in result.scores], [1, 2])
for rows in result.scores:
for row in rows:
torch.testing.assert_close(
torch.tensor(row), torch.as_tensor(expected)
)

async def test_invalid_candidates_and_temperature(self):
manager = ScoringManager()
for labels in ([], [[]], [[1]], [[1], []], [1, 1], [-1], [8], [True], "bad"):
Expand Down
Loading