diff --git a/python/sglang/srt/managers/tokenizer_manager_score_mixin.py b/python/sglang/srt/managers/tokenizer_manager_score_mixin.py index 73e93c930539..6202e0d0a470 100644 --- a/python/sglang/srt/managers/tokenizer_manager_score_mixin.py +++ b/python/sglang/srt/managers/tokenizer_manager_score_mixin.py @@ -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) @@ -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: @@ -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``). @@ -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( @@ -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: @@ -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 diff --git a/test/registered/unit/managers/test_setwise_score_mixin.py b/test/registered/unit/managers/test_setwise_score_mixin.py index 52f409cd691e..78dc86748d59 100644 --- a/test/registered/unit/managers/test_setwise_score_mixin.py +++ b/test/registered/unit/managers/test_setwise_score_mixin.py @@ -409,12 +409,14 @@ 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) @@ -422,7 +424,9 @@ def test_multi_position_score_rows_softmax_over_labels(self): 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"): diff --git a/test/registered/unit/test_token_scoring.py b/test/registered/unit/test_token_scoring.py index 6c8e0994b7b7..be199172c864 100644 --- a/test/registered/unit/test_token_scoring.py +++ b/test/registered/unit/test_token_scoring.py @@ -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) @@ -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}) @@ -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"):