diff --git a/python/sglang/srt/entrypoints/engine_score_mixin.py b/python/sglang/srt/entrypoints/engine_score_mixin.py index b2e305b6ffc3..d8f35318a297 100644 --- a/python/sglang/srt/entrypoints/engine_score_mixin.py +++ b/python/sglang/srt/entrypoints/engine_score_mixin.py @@ -56,8 +56,10 @@ def score( Setwise scoring is expressed via ``score_extraction_token_id``: pass a single item containing the whole candidate block with one extraction token per - candidate, and the head is pooled AT those positions, so ``scores`` becomes - the ``[N x num_labels]`` per-candidate matrix. + candidate, and the readout is taken AT those positions, so ``scores`` + becomes the ``[N x num_labels]`` per-candidate matrix. SequenceClassification + pools the head there; CausalLM reads label-token logprobs there (batched + path only). Args: query: The query text or pre-tokenized token IDs. @@ -69,10 +71,11 @@ def score( embed_override_token_id: Placeholder token ID used to locate override positions. query_embed_overrides: Embedding vectors replacing placeholder tokens in query. item_embed_overrides: Per-item embedding vectors replacing placeholder tokens in items. - score_extraction_token_id: SequenceClassification-only. When set, pool the - head at every occurrence of this token per sequence instead of the last - token; ``scores`` becomes nested — one ``[Ni x num_labels]`` matrix per - item (``len(scores) == len(items)``), where ``Ni`` is the number of + score_extraction_token_id: When set, read the score head (SeqCls) or + label-token logprobs (CausalLM) at every occurrence of + this token per sequence instead of the last token; ``scores`` becomes + nested — one ``[Ni x num_labels]`` matrix per item + (``len(scores) == len(items)``), where ``Ni`` is the number of extraction tokens (candidates) in item ``i``. return_pooled_hidden_states: Whether to include raw pooled transformer hidden states (before the task head) in the result. Only supported diff --git a/python/sglang/srt/entrypoints/openai/protocol.py b/python/sglang/srt/entrypoints/openai/protocol.py index 5e503c0b2907..de378b6e0218 100644 --- a/python/sglang/srt/entrypoints/openai/protocol.py +++ b/python/sglang/srt/entrypoints/openai/protocol.py @@ -1408,10 +1408,11 @@ class ScoringRequest(BaseModel): item_first: bool = False return_pooled_hidden_states: bool = False - # Setwise readout (SequenceClassification-only): when set, the head is pooled - # AT every occurrence of this token in each `query + item` sequence instead of - # the last token, and `scores` is returned nested (one `[Nᵢ x num_labels]` - # matrix per item). --enable-mis fuses items; otherwise each is scored alone. + # Setwise readout: when set, the readout is taken AT every occurrence of this + # token in each `query + item` sequence instead of the last token, and `scores` + # is returned nested (one `[Nᵢ x num_labels]` matrix per item). SeqCls pools the + # head there; CausalLM reads label-token logprobs there. Both support batched + # and --enable-mis. score_extraction_token: Optional[str] = None model: str = DEFAULT_MODEL_NAME diff --git a/python/sglang/srt/layers/logits_processor.py b/python/sglang/srt/layers/logits_processor.py index 5cb046c4e84b..d46227dc0d2f 100644 --- a/python/sglang/srt/layers/logits_processor.py +++ b/python/sglang/srt/layers/logits_processor.py @@ -505,10 +505,12 @@ def forward( aux_hidden_states: Optional[AuxHiddenStates] = None, hidden_states_before_norm: Optional[torch.Tensor] = None, ) -> LogitsProcessorOutput: - # Extract MIS indices before ForwardBatch → LogitsMetadata conversion + # Extract MIS / setwise indices before ForwardBatch → LogitsMetadata conversion multi_item_delimiter_indices = None + token_indices_to_pool = None if isinstance(logits_metadata, ForwardBatch): multi_item_delimiter_indices = logits_metadata.multi_item_delimiter_indices + token_indices_to_pool = logits_metadata.token_indices_to_pool logits_metadata = LogitsMetadata.from_forward_batch(logits_metadata) # Autotune dummy run discards this output. `is False` not `not`: None @@ -526,6 +528,18 @@ def forward( forward_mode=logits_metadata.forward_mode, ) + # Setwise scoring (CausalLM): read label-token logprobs AT each anchor + # position instead of the last token. Takes precedence over the MIS + # delimiter path; the two are mutually exclusive for generation. + if token_indices_to_pool is not None and logits_metadata.is_prefill_only: + return self.compute_logprobs_at_positions( + input_ids, + hidden_states, + lm_head, + logits_metadata, + token_indices_to_pool, + ) + # Multi-item scoring only for prefill-only requests with pre-computed indices. if multi_item_delimiter_indices is not None and logits_metadata.is_prefill_only: return self.compute_logprobs_for_multi_item_scoring( @@ -1252,6 +1266,89 @@ def compute_logprobs_for_multi_item_scoring( mm_input_embeds=logits_metadata.mm_input_embeds, ) + def compute_logprobs_at_positions( + self, + input_ids, + hidden_states, + lm_head: VocabParallelEmbedding, + logits_metadata: Union[LogitsMetadata, ForwardBatch], + token_indices_to_pool: List[torch.Tensor], + ): + """Compute label-token logprobs AT each requested position (setwise, CausalLM). + + Mirrors ``compute_logprobs_for_multi_item_scoring`` but reads the LM head + AT ``token_indices_to_pool`` (no delimiter - 1 shift, no discarded row): + each anchor's logprobs are P(next token | prefix up to and including the + anchor), one row per anchor. + """ + device = input_ids.device + all_tensors = [] + if logits_metadata.extend_seq_lens_cpu is not None: + offset = 0 + for req_seq_len, indices_tensor in zip( + logits_metadata.extend_seq_lens_cpu, token_indices_to_pool + ): + if len(indices_tensor) > 0: + all_tensors.append(indices_tensor + offset) + offset += req_seq_len + else: + all_tensors.append(token_indices_to_pool[0]) + pooled_indices = torch.cat(all_tensors).to(device, non_blocking=True) + + sliced_hidden = hidden_states[pooled_indices] + sliced_logits = self._get_logits(sliced_hidden, lm_head, logits_metadata) + sliced_logprobs = torch.nn.functional.log_softmax(sliced_logits, dim=-1) + + input_token_ids_logprobs_val = [] + input_token_ids_logprobs_idx = [] + input_top_logprobs_val = None + input_top_logprobs_idx = None + + if ( + logits_metadata.token_ids_logprobs + or logits_metadata.extend_return_top_logprob + ): + logits_metadata.extend_logprob_pruned_lens_cpu = [ + len(t) for t in token_indices_to_pool + ] + + if logits_metadata.extend_token_ids_logprob: + ( + input_token_ids_logprobs_val, + input_token_ids_logprobs_idx, + ) = get_token_ids_logprobs_raw( + sliced_logprobs, + logits_metadata.token_ids_logprobs, + stage=LogprobStage.PREFILL, + extend_logprob_pruned_lens_cpu=logits_metadata.extend_logprob_pruned_lens_cpu, + no_copy_to_cpu=True, + ) + + if logits_metadata.extend_return_top_logprob: + ( + input_top_logprobs_val, + input_top_logprobs_idx, + ) = get_top_logprobs_raw( + sliced_logprobs, + logits_metadata.top_logprobs_nums, + stage=LogprobStage.PREFILL, + extend_logprob_pruned_lens_cpu=logits_metadata.extend_logprob_pruned_lens_cpu, + ) + + # Zeros to satisfy the shared logprob pipeline's non-None / length asserts; + # score_request() reads only input_token_ids_logprobs_val (see the MIS path). + input_token_logprobs = torch.zeros(pooled_indices.shape[0], device=device) + + return LogitsProcessorOutput( + next_token_logits=None, + input_token_logprobs=input_token_logprobs, + input_top_logprobs_val=input_top_logprobs_val, + input_top_logprobs_idx=input_top_logprobs_idx, + input_token_ids_logprobs_val=input_token_ids_logprobs_val, + input_token_ids_logprobs_idx=input_token_ids_logprobs_idx, + mm_input_embeds=logits_metadata.mm_input_embeds, + ) + def _reassemble_tp_lm_head_all_to_all_output( all_to_all_output: torch.Tensor, diff --git a/python/sglang/srt/managers/io_struct.py b/python/sglang/srt/managers/io_struct.py index cdf56a2b7b85..9629b7aab621 100644 --- a/python/sglang/srt/managers/io_struct.py +++ b/python/sglang/srt/managers/io_struct.py @@ -357,6 +357,11 @@ class GenerateReqInput: # Batch-level: List[List[int]] (one per request). After __getitem__: List[int]. multi_item_delimiter_indices: Optional[Union[List[List[int]], List[int]]] = None + # Token positions for setwise pooling readout (CausalLM: label-token logprobs + # are read AT these positions instead of the last token). + # Batch-level: List[List[int]] (one per request). After __getitem__: List[int]. + token_indices_to_pool: Optional[Union[List[List[int]], List[int]]] = None + # Cache namespace used to isolate otherwise-identical prefixes. cache_salt: Optional[Union[List[str], str]] = None @@ -1023,6 +1028,11 @@ def __getitem__(self, i): if self.multi_item_delimiter_indices is not None else None ), + token_indices_to_pool=( + self.token_indices_to_pool[i] + if self.token_indices_to_pool is not None + else None + ), ) cache[i] = sub return sub @@ -1120,6 +1130,9 @@ class TokenizedGenerateReqInput(BaseReq, kw_only=True): # Pre-computed delimiter indices for multi-item scoring multi_item_delimiter_indices: Optional[List[int]] = None + # Token positions for setwise pooling readout (CausalLM) + token_indices_to_pool: Optional[List[int]] = None + # For observability # Pickled Optional[Union[APIServerReqTimeStats, DPControllerReqTimeStats]] time_stats: Optional[PickleWrapper] = None diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index f74ce7307d02..8f98be3c77b0 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -2798,6 +2798,7 @@ def handle_generate_request( dllm_config=self.dllm_config, time_stats=recv_req.time_stats, multi_item_delimiter_indices=recv_req.multi_item_delimiter_indices, + token_indices_to_pool=recv_req.token_indices_to_pool, ) req.tokenizer = self.tokenizer diff --git a/python/sglang/srt/managers/scheduler_components/logprob_result_processor.py b/python/sglang/srt/managers/scheduler_components/logprob_result_processor.py index c4eb98ab12eb..e55906c49ffd 100644 --- a/python/sglang/srt/managers/scheduler_components/logprob_result_processor.py +++ b/python/sglang/srt/managers/scheduler_components/logprob_result_processor.py @@ -27,24 +27,24 @@ def _process_input_token_logprobs( self, req: Req, input_token_logprobs: List ) -> None: """Process input token logprobs values and indices.""" - is_multi_item_scoring = self._is_multi_item_scoring(req) + uses_scoring_positions = self._uses_scoring_positions(req) - # Process logprob values - handle multi-item scoring vs regular requests - if is_multi_item_scoring: - # Multi-item scoring: use all logprobs as-is + # Process logprob values - handle position-based scoring vs regular requests + if uses_scoring_positions: + # Position-based scoring: use all logprobs as-is req.logprob.input_token_logprobs_val = input_token_logprobs else: # Regular request: add None at start, remove last (sampling token) req.logprob.input_token_logprobs_val = [None] + input_token_logprobs[:-1] # Process logprob indices based on scoring type - if is_multi_item_scoring: - # MIS scores come from input_token_ids_logprobs, not input_token_logprobs. + if uses_scoring_positions: + # Position-based scores come from input_token_ids_logprobs, not input_token_logprobs. # But the shared pipeline requires input_token_logprobs_idx to be the same # length as input_token_logprobs_val (validated at line 816). We fill with # MIS_DELIMITER_TOKEN_ID as a dummy — score_request() ignores this field. - delimiter_count = len(req.multi_item_delimiter_indices) - input_token_logprobs_idx = [MIS_DELIMITER_TOKEN_ID] * delimiter_count + position_count = len(self._scoring_positions(req)) + input_token_logprobs_idx = [MIS_DELIMITER_TOKEN_ID] * position_count else: # Regular request: include all tokens from logprob_start_len onwards input_token_logprobs_idx = req.origin_input_ids[req.logprob_start_len :] @@ -60,11 +60,11 @@ def _process_input_top_logprobs(self, req: Req) -> None: if req.logprob.top_logprobs_num <= 0: return - is_multi_item_scoring = self._is_multi_item_scoring(req) + uses_scoring_positions = self._uses_scoring_positions(req) - # Initialize arrays - multi-item scoring starts empty, others start with None - req.logprob.input_top_logprobs_val = [] if is_multi_item_scoring else [None] - req.logprob.input_top_logprobs_idx = [] if is_multi_item_scoring else [None] + # Initialize arrays - position-based scoring starts empty, others start with None + req.logprob.input_top_logprobs_val = [] if uses_scoring_positions else [None] + req.logprob.input_top_logprobs_idx = [] if uses_scoring_positions else [None] # Extend arrays with temp values for val, idx in zip( @@ -75,8 +75,8 @@ def _process_input_top_logprobs(self, req: Req) -> None: req.logprob.input_top_logprobs_val.extend(val) req.logprob.input_top_logprobs_idx.extend(idx) - # Remove last token (sampling token) for non multi-item scoring requests - if not is_multi_item_scoring: + # Remove last token (sampling token) for non position-based scoring requests + if not uses_scoring_positions: req.logprob.input_top_logprobs_val.pop() req.logprob.input_top_logprobs_idx.pop() @@ -118,14 +118,14 @@ def _process_input_token_ids_logprobs(self, req: Req) -> None: if req.logprob.token_ids_logprob is None: return - is_multi_item_scoring = self._is_multi_item_scoring(req) + uses_scoring_positions = self._uses_scoring_positions(req) - # Initialize arrays - multi-item scoring starts empty, others start with None + # Initialize arrays - position-based scoring starts empty, others start with None req.logprob.input_token_ids_logprobs_val = ( - [] if is_multi_item_scoring else [None] + [] if uses_scoring_positions else [None] ) req.logprob.input_token_ids_logprobs_idx = ( - [] if is_multi_item_scoring else [None] + [] if uses_scoring_positions else [None] ) # Process temp values - convert tensors to lists and extend arrays @@ -140,8 +140,8 @@ def _process_input_token_ids_logprobs(self, req: Req) -> None: ) req.logprob.input_token_ids_logprobs_idx.extend(idx) - # Remove last token (sampling token) for non multi-item scoring requests - if not is_multi_item_scoring: + # Remove last token (sampling token) for non position-based scoring requests + if not uses_scoring_positions: req.logprob.input_token_ids_logprobs_val.pop() req.logprob.input_token_ids_logprobs_idx.pop() @@ -150,15 +150,16 @@ def _process_input_token_ids_logprobs(self, req: Req) -> None: req.temp_input_token_ids_logprobs_val = None def _calculate_relevant_tokens_len(self, req: Req) -> int: - """Calculate the expected length of logprob arrays based on whether multi-item scoring is enabled. + """Calculate the expected length of logprob arrays based on whether position-based scoring is used. - For multi-item scoring, only delimiter positions have logprobs. - For regular requests, all positions from logprob_start_len onwards have logprobs. + For position-based scoring (MIS delimiters or setwise anchors), only those + positions have logprobs. For regular requests, all positions from + logprob_start_len onwards have logprobs. """ - is_multi_item_scoring = self._is_multi_item_scoring(req) + positions = self._scoring_positions(req) - if is_multi_item_scoring: - return len(req.multi_item_delimiter_indices) + if positions is not None: + return len(positions) else: return len(req.origin_input_ids[req.logprob_start_len :]) @@ -168,36 +169,53 @@ def calculate_num_input_logprobs( extend_input_len: int, extend_logprob_start_len: int, ) -> int: - """Calculate the number of input logprobs based on whether multi-item scoring is enabled. + """Calculate the number of input logprobs based on whether position-based scoring is used. - For multi-item scoring, only delimiter positions have logprobs. - For regular requests, all positions in the range have logprobs. + For position-based scoring (MIS delimiters or setwise anchors), only those + positions have logprobs. For regular requests, all positions in the range + have logprobs. """ - is_multi_item_scoring = self._is_multi_item_scoring(req) + positions = self._scoring_positions(req) - if is_multi_item_scoring: - # Count pre-computed delimiter indices within the extend range + if positions is not None: + # Count pre-computed scoring positions within the extend range return sum( 1 - for idx in req.multi_item_delimiter_indices + for idx in positions if extend_logprob_start_len <= idx < extend_input_len ) else: # Regular request: all tokens in the range return extend_input_len - extend_logprob_start_len - def _is_multi_item_scoring(self, req: Req) -> bool: - """Check if request uses multi-item scoring. + def _scoring_positions(self, req: Req): + """Position-based scoring readout indices for a prefill-only request, else None. - Multi-item scoring applies to prefill-only requests when a delimiter - token is configured. In this mode, only positions containing the - delimiter token receive logprobs. + Both MIS (delimiter positions, --enable-mis) and setwise scoring + (token_indices_to_pool, per-anchor readout) compute label-token logprobs + at a fixed set of positions instead of at every input token; the logprob + bookkeeping below is identical for the two. """ - return ( + if not req.is_prefill_only: + return None + # Setwise anchors take precedence over MIS delimiters: a fused setwise + # request carries both, and its logprobs are read at the anchors. + if req.token_indices_to_pool is not None: + return req.token_indices_to_pool + if ( get_exec().features.enable_mis - and req.is_prefill_only and req.multi_item_delimiter_indices is not None - ) + ): + return req.multi_item_delimiter_indices + return None + + def _uses_scoring_positions(self, req: Req) -> bool: + """Whether the request uses position-based scoring (MIS or setwise). + + In this mode only the scoring positions receive logprobs, so the shared + logprob pipeline skips the None prefix and last-token pop. + """ + return self._scoring_positions(req) is not None def add_input_logprob_return_values( self, diff --git a/python/sglang/srt/managers/tokenizer_manager.py b/python/sglang/srt/managers/tokenizer_manager.py index 7140fbbd5c64..954d59a6a84f 100644 --- a/python/sglang/srt/managers/tokenizer_manager.py +++ b/python/sglang/srt/managers/tokenizer_manager.py @@ -1533,6 +1533,7 @@ def _create_tokenized_object( need_wait_for_mm_inputs=obj.need_wait_for_mm_inputs, num_items_assigned=obj.num_items_assigned, multi_item_delimiter_indices=obj.multi_item_delimiter_indices, + token_indices_to_pool=obj.token_indices_to_pool, encoder_urls=obj.encoder_urls, ) elif isinstance(obj, EmbeddingReqInput): diff --git a/python/sglang/srt/managers/tokenizer_manager_score_mixin.py b/python/sglang/srt/managers/tokenizer_manager_score_mixin.py index 73e93c930539..14af8335a9e9 100644 --- a/python/sglang/srt/managers/tokenizer_manager_score_mixin.py +++ b/python/sglang/srt/managers/tokenizer_manager_score_mixin.py @@ -258,25 +258,61 @@ def _process_multi_item_extraction_results( results: Any, per_item_anchor_counts: List[int], apply_softmax: bool, + label_token_ids: Optional[List[List[int]]] = None, + temperature: float = 1.0, return_pooled_hidden_states: bool = False, ) -> ScoreResult: """Process a fused multi-item score-extraction request (``--enable-mis``). - The fused sequence returns one ``[ΣNᵢ, num_labels]`` matrix of every - anchor's logits in item order; split it back into one ``[Nᵢ, num_labels]`` - matrix per item using per_item_anchor_counts. + The fused sequence yields every anchor's scores in item order (a + ``[ΣNᵢ, num_labels]`` embedding for SequenceClassification, or per-anchor + label-token logprobs for CausalLM); split it back into one + ``[Nᵢ, num_labels]`` matrix per item using per_item_anchor_counts. """ single_result = results[0] if isinstance(results, list) else results meta_info = single_result.get("meta_info", {}) request_id = meta_info.get("id", "") prompt_tokens = meta_info.get("prompt_tokens", 0) + total_anchors = sum(per_item_anchor_counts) + + if self.is_generation: + # CausalLM: per-anchor label-token logprobs in item order; item i's + # anchors use item i's candidate labels (label_token_ids[i]). + anchor_logprobs = meta_info.get("input_token_ids_logprobs", []) + if not anchor_logprobs: + raise ValueError( + f"input_token_ids_logprobs not found in the result for " + f"request {request_id}." + ) + if len(anchor_logprobs) != total_anchors: + raise RuntimeError( + f"Expected {total_anchors} anchor rows across " + f"{len(per_item_anchor_counts)} items, but got " + f"{len(anchor_logprobs)}. Request ID: {request_id}" + ) + scores = [] + offset = 0 + for labels, count in zip(label_token_ids, per_item_anchor_counts): + scores.append( + [ + self._convert_logprobs_to_scores( + self._extract_logprobs_for_tokens( + anchor_logprobs[offset + j], labels + ), + labels, + apply_softmax, + temperature, + ) + for j in range(count) + ] + ) + offset += count + return ScoreResult(scores=scores, prompt_tokens=prompt_tokens) embedding = single_result.get("embedding") if embedding is None: raise ValueError("Embedding not found in the result.") - rows = self._multi_position_score_rows(embedding, apply_softmax) - total_anchors = sum(per_item_anchor_counts) if len(rows) != total_anchors: raise RuntimeError( f"Expected {total_anchors} anchor rows across " @@ -332,19 +368,38 @@ def _process_single_item_scoring_results( is_generation = self.is_generation if is_generation: for result, labels in zip(results, label_token_ids): - # For single-item scoring, logprobs are in output_token_ids_logprobs - output_logprobs = result["meta_info"].get( - "output_token_ids_logprobs", [] - ) - prompt_tokens += result["meta_info"].get("prompt_tokens", 0) + meta = result["meta_info"] + prompt_tokens += meta.get("prompt_tokens", 0) + + if per_item_matrix: + # Setwise: label-token logprobs were read AT each anchor + # position; one row per anchor, grouped as one item's matrix. + anchor_logprobs = meta.get("input_token_ids_logprobs", []) + if not anchor_logprobs: + raise RuntimeError( + f"input_token_ids_logprobs is empty for request " + f"{meta.get('id', '')}." + ) + item_matrix = [ + self._convert_logprobs_to_scores( + self._extract_logprobs_for_tokens(pos_logprobs, labels), + labels, + apply_softmax, + temperature, + ) + for pos_logprobs in anchor_logprobs + ] + scores.append(item_matrix) + continue + # Pointwise: logprobs are in output_token_ids_logprobs (last position) + output_logprobs = meta.get("output_token_ids_logprobs", []) if not output_logprobs or len(output_logprobs) == 0: raise RuntimeError( f"output_logprobs is empty for request " - f"{result['meta_info'].get('id', '')}." + f"{meta.get('id', '')}." ) - # Extract logprobs for the first (and only) position logprobs = self._extract_logprobs_for_tokens(output_logprobs[0], labels) score_list = self._convert_logprobs_to_scores( logprobs, labels, apply_softmax, temperature @@ -592,29 +647,26 @@ def _validate_score_extraction( is_generation: bool, item_first: bool, ) -> None: - """Validate that multi-position pooling readout is applicable. - - Supported only by SequenceClassification models whose forward routes the - head through ``score_and_pool`` (per-position pooling, in - ``is_score_and_pool_model``). Generation, cross-encoder, reward, and - embedding models pool a single vector and are rejected here before - inference. Requires radix cache and chunked prefill off and no auto-truncate, - since pooling positions are full-prompt coordinates. + """Validate that setwise pooling readout is applicable. + + SequenceClassification models must route the head through + ``score_and_pool`` (per-position pooling, in ``is_score_and_pool_model``); + cross-encoder, reward, and embedding models pool a single vector and are + rejected. CausalLM (generation) models read label-token logprobs at each + anchor via the LM head and are supported in both the batched and fused + (``--enable-mis``) paths. Requires radix cache and chunked prefill off and + no auto-truncate, since pooling positions are full-prompt coordinates. """ - if is_generation: - raise ValueError( - "score_extraction_token_id is only supported for " - "SequenceClassification models, not generation (CausalLM) models." - ) - architectures = self.model_config.hf_config.architectures or [] - if not is_score_and_pool_model(architectures): - raise ValueError( - "score_extraction_token_id is only supported for " - "SequenceClassification models that pool the head per position " - "(score_and_pool); model architecture(s) " - f"{architectures} do not (cross-encoder, reward, and embedding " - "models pool a single vector and ignore the readout positions)." - ) + if not is_generation: + architectures = self.model_config.hf_config.architectures or [] + if not is_score_and_pool_model(architectures): + raise ValueError( + "score_extraction_token_id is only supported for " + "SequenceClassification models that pool the head per position " + "(score_and_pool); model architecture(s) " + f"{architectures} do not (cross-encoder, reward, and embedding " + "models pool a single vector and ignore the readout positions)." + ) if item_first: raise ValueError( "item_first is not supported with score_extraction_token_id." @@ -786,12 +838,12 @@ async def score_request( - Generation (CausalLM): Requires label_token_ids; returns logprob-based scores. - SequenceClassification: label_token_ids is optional; returns pooled class logits. - Setwise scoring (SequenceClassification-only) is expressed via - score_extraction_token_id: when set, the head is pooled AT every occurrence - of this token in each ``query + item`` sequence instead of the last token, - and ``scores`` is returned nested (one ``[Nᵢ x num_labels]`` matrix per - item). With ``--enable-mis`` the items are fused into one multi-item - sequence; otherwise each item is scored independently in one batch. + Setwise scoring is expressed via score_extraction_token_id: when set, the + readout is taken AT every occurrence of this token in each ``query + item`` + sequence instead of the last token, and ``scores`` is returned nested (one + ``[Nᵢ x num_labels]`` matrix per item). SequenceClassification pools the + head at those positions; CausalLM reads label-token logprobs there. Both + support the batched and fused (``--enable-mis``) paths. return_pooled_hidden_states is only supported for non-generation models (SequenceClassification, RewardModel); raises ValueError for CausalLM. @@ -1016,12 +1068,16 @@ async def score_request( input_ids=input_ids, token_ids_logprob=request_labels, return_logprob=True, - # Set logprob_start_len=0 for multi-item scoring since we want logprobs at all delimiter positions - logprob_start_len=0 if use_multi_item_scoring else -1, + # logprob_start_len=0 so input-position logprobs are computed: + # multi-item scoring reads them at delimiters, setwise at anchors. + logprob_start_len=( + 0 if (use_multi_item_scoring or use_score_extraction) else -1 + ), stream=False, sampling_params={"max_new_tokens": 0}, positional_embed_overrides=positional_embed_overrides, multi_item_delimiter_indices=mis_delimiter_indices, + token_indices_to_pool=token_indices_to_pool, ) else: batch_request = EmbeddingReqInput( @@ -1037,13 +1093,15 @@ async def score_request( if use_multi_item_scoring and use_score_extraction: # Multi-item setwise: the items were fused into one sequence and the - # head pooled at every anchor. Split the flat [ΣNᵢ, num_labels] result - # back into one matrix per item (nested scores). + # readout taken at every anchor. Split the flat [ΣNᵢ, num_labels] + # result back into one matrix per item (nested scores). return self._process_multi_item_extraction_results( results, per_item_anchor_counts, apply_softmax, - return_pooled_hidden_states, + label_token_ids=label_token_ids, + temperature=temperature, + return_pooled_hidden_states=return_pooled_hidden_states, ) elif use_multi_item_scoring: # Multi-item scoring: extract scores from input_token_ids_logprobs or embedding diff --git a/test/registered/e2e/scoring/test_setwise_scoring.py b/test/registered/e2e/scoring/test_setwise_scoring.py index 669d2876ba92..e9c8087d17b8 100644 --- a/test/registered/e2e/scoring/test_setwise_scoring.py +++ b/test/registered/e2e/scoring/test_setwise_scoring.py @@ -14,11 +14,16 @@ """ import asyncio +import math import os import unittest import torch -from transformers import AutoModelForSequenceClassification, AutoTokenizer +from transformers import ( + AutoModelForCausalLM, + AutoModelForSequenceClassification, + AutoTokenizer, +) from sglang.srt.entrypoints.engine import Engine from sglang.test.ci.ci_register import register_cuda_ci @@ -30,6 +35,8 @@ "TEST_CLASSIFICATION_BASE_MODEL", "tomaarsen/Qwen3-Reranker-0.6B-seq-cls", ) +# CausalLM base (same tokenizer / anchor special token as the reranker above). +_CAUSAL_MODEL = os.environ.get("TEST_CAUSAL_LM_MODEL", "Qwen/Qwen3-0.6B") _ANCHOR_TOKEN = os.environ.get("TEST_SCORE_EXTRACTION_TOKEN", "<|object_ref_start|>") # float16 on flashinfer (no float32 prefill kernel); HF golden is float32 and the # tolerance below covers the fp16-vs-fp32 gap. All overridable via env. @@ -296,6 +303,185 @@ async def _run(): self.assertEqual(len(r.scores), 2) +class TestGenerationSetwiseScoringHFParity(CustomTestCase): + """CausalLM setwise scores must match an HF LM-head-at-anchor reference. + + At each anchor the engine reads label-token logprobs from the LM head + (P(next | prefix up to the anchor)); the HF reference gathers the same + log_softmax(logits) rows and exponentiates them (apply_softmax=False -> + probabilities), so a wrong pooling position changes the numbers. + """ + + @classmethod + def setUpClass(cls): + cls.tokenizer = AutoTokenizer.from_pretrained(_CAUSAL_MODEL) + cls.anchor_id = cls.tokenizer.convert_tokens_to_ids(_ANCHOR_TOKEN) + assert ( + cls.anchor_id is not None and cls.anchor_id != cls.tokenizer.unk_token_id + ), f"{_ANCHOR_TOKEN!r} did not resolve to a dedicated token id" + # Two arbitrary single-token labels to score at each anchor. + cls.label_token_ids = [ + cls.tokenizer.encode("yes", add_special_tokens=False)[-1], + cls.tokenizer.encode("no", add_special_tokens=False)[-1], + ] + cls.engine = Engine( + model_path=_CAUSAL_MODEL, + disable_radix_cache=True, + chunked_prefill_size=-1, + attention_backend="flashinfer", + dtype=_DTYPE, + mem_fraction_static=0.15, + ) + + @classmethod + def tearDownClass(cls): + if getattr(cls, "engine", None) is not None: + cls.engine.shutdown() + torch.cuda.empty_cache() + + def _build_prompt(self, n_anchors: int) -> str: + candidates = " ".join(f"Candidate {i}." for i in range(n_anchors)) + return f"Rank the candidates. {candidates} Scores:" + ( + _ANCHOR_TOKEN * n_anchors + ) + + def _hf_causal_setwise_reference(self, prompt: str): + """Reference: exp(log_softmax(LM-head logits)[anchor])[label] per anchor.""" + input_ids = self.tokenizer.encode(prompt) + anchor_positions = [i for i, t in enumerate(input_ids) if t == self.anchor_id] + self.assertGreater(len(anchor_positions), 0, "prompt has no anchor tokens") + + model = AutoModelForCausalLM.from_pretrained( + _CAUSAL_MODEL, torch_dtype=torch.float32 + ).eval() + try: + ids = torch.tensor([input_ids], dtype=torch.long) + with torch.no_grad(): + logits = model(input_ids=ids).logits[0] # [seq, vocab] + logprobs = torch.log_softmax(logits.float(), dim=-1) + ref = [ + [math.exp(logprobs[p, t].item()) for t in self.label_token_ids] + for p in anchor_positions + ] + return ref, anchor_positions + finally: + model.cpu() + del model + torch.cuda.empty_cache() + + def _assert_close(self, ref, sgl, atol=_ATOL, rtol=_RTOL): + self.assertEqual(len(ref), len(sgl), "row count mismatch") + for i, (rrow, srow) in enumerate(zip(ref, sgl)): + self.assertEqual(len(rrow), len(srow), f"row {i} width mismatch") + for r, s in zip(rrow, srow): + self.assertLessEqual(abs(r - s), atol + rtol * abs(r)) + + def test_generation_setwise_matches_hf_reference(self): + prompt = self._build_prompt(3) + ref, anchor_positions = self._hf_causal_setwise_reference(prompt) + + sgl = self.engine.score( + query="", + items=[prompt], + label_token_ids=self.label_token_ids, + apply_softmax=False, + score_extraction_token_id=self.anchor_id, + ).scores + + self.assertEqual(len(sgl), 1) # one item + self.assertEqual(len(sgl[0]), len(anchor_positions)) # one row per anchor + self._assert_close(ref, sgl[0]) + + def test_generation_setwise_multiple_items_match_hf(self): + prompt0 = self._build_prompt(3) + prompt1 = self._build_prompt(2) + ref0, anchors0 = self._hf_causal_setwise_reference(prompt0) + ref1, anchors1 = self._hf_causal_setwise_reference(prompt1) + + sgl = self.engine.score( + query="", + items=[prompt0, prompt1], + label_token_ids=self.label_token_ids, + apply_softmax=False, + score_extraction_token_id=self.anchor_id, + ).scores + + self.assertEqual(len(sgl), 2) # one matrix per item + self.assertEqual(len(sgl[0]), len(anchors0)) + self.assertEqual(len(sgl[1]), len(anchors1)) + self._assert_close(ref0, sgl[0]) + self._assert_close(ref1, sgl[1]) + + +class TestGenerationSetwiseMISScoring(CustomTestCase): + """Fused CausalLM setwise under ``--enable-mis``. + + Items are fused into one sequence; the LM head is read at each anchor with the + block-diagonal mask isolating sets. Asserts set isolation (the property only a + real fused forward exercises); shape/grouping is covered by the unit tests. + """ + + @classmethod + def setUpClass(cls): + cls.tokenizer = AutoTokenizer.from_pretrained(_CAUSAL_MODEL) + cls.anchor_id = cls.tokenizer.convert_tokens_to_ids(_ANCHOR_TOKEN) + cls.label_token_ids = [ + cls.tokenizer.encode("yes", add_special_tokens=False)[-1], + cls.tokenizer.encode("no", add_special_tokens=False)[-1], + ] + cls.engine = Engine( + model_path=_CAUSAL_MODEL, + disable_radix_cache=True, + chunked_prefill_size=-1, + enable_mis=True, + attention_backend="flashinfer", + dtype=_DTYPE, + mem_fraction_static=0.15, + ) + + @classmethod + def tearDownClass(cls): + if getattr(cls, "engine", None) is not None: + cls.engine.shutdown() + torch.cuda.empty_cache() + + @staticmethod + def _set_prompt(n_anchors: int) -> str: + candidates = " ".join(f"Candidate {i}." for i in range(n_anchors)) + return f"Rank the candidates. {candidates} Scores:" + ( + _ANCHOR_TOKEN * n_anchors + ) + + def _assert_matrix_close(self, a, b, atol=5e-2): + self.assertEqual(len(a), len(b), "row count mismatch") + for ra, rb in zip(a, b): + self.assertEqual(len(ra), len(rb), "width mismatch") + for x, y in zip(ra, rb): + self.assertLessEqual(abs(x - y), atol, f"{x} vs {y}") + + def test_mis_generation_set_isolation(self): + # The block-diagonal mask means set0's scores are unchanged by a second set. + set0 = self._set_prompt(3) + alone = self.engine.score( + query="Rank:", + items=[set0], + label_token_ids=self.label_token_ids, + apply_softmax=False, + score_extraction_token_id=self.anchor_id, + ).scores + fused = self.engine.score( + query="Rank:", + items=[set0, self._set_prompt(2)], + label_token_ids=self.label_token_ids, + apply_softmax=False, + score_extraction_token_id=self.anchor_id, + ).scores + + self.assertEqual(len(alone), 1) + self.assertEqual([len(m) for m in fused], [3, 2]) # nested per item + self._assert_matrix_close(alone[0], fused[0]) + + class TestSetwiseMultiItemMISScoring(CustomTestCase): """Multi-item setwise scoring under ``--enable-mis`` (fused sequence). diff --git a/test/registered/unit/managers/test_setwise_score_mixin.py b/test/registered/unit/managers/test_setwise_score_mixin.py index 52f409cd691e..0b9d96a7f0df 100644 --- a/test/registered/unit/managers/test_setwise_score_mixin.py +++ b/test/registered/unit/managers/test_setwise_score_mixin.py @@ -9,6 +9,7 @@ grouping, validation, anchor bucketing, and the serialization round-trip on CPU. """ +import math import unittest from types import SimpleNamespace @@ -43,6 +44,17 @@ def __init__(self, architecture="Qwen3ForSequenceClassification"): ) +class _GenHarness(TokenizerManagerScoreMixin): + """Bare CausalLM harness for the generation result parser.""" + + is_generation = True + + def __init__(self, architecture="LlamaForCausalLM"): + self.model_config = SimpleNamespace( + hf_config=SimpleNamespace(architectures=[architecture]) + ) + + class TestSingleItemScoringResults(CustomTestCase): @staticmethod def _result(embedding, phs=None, prompt_tokens=7): @@ -230,6 +242,12 @@ def test_validation_allows_multi_item(self): # both return one score matrix per item. self.h._validate_score_extraction(False, False) + def test_validation_allows_causal_lm(self): + # CausalLM setwise is supported (batched and --enable-mis); generation reads + # label-token logprobs at each anchor and skips the SequenceClassification + # score_and_pool allow-list check. + self.h._validate_score_extraction(True, False) + def test_anchor_counts_per_item_buckets_by_delimiter(self): # Fused: q item0(2 anchors) item1(1 anchor) . # delimiter_indices point at the delimiter tokens. @@ -433,6 +451,150 @@ def test_multi_position_phs_matrix_rejects_non_matrix(self): self.h._multi_position_phs_matrix(torch.randn(4), expected_rows=4) +class TestGenerationSetwiseResults(CustomTestCase): + """CausalLM setwise: label-token logprobs read at each anchor position. + + The head is the LM head; per-anchor logprobs arrive as the request's + ``input_token_ids_logprobs`` (one entry per anchor), and the parser groups + them into one ``[N x num_labels]`` matrix per item (nested), reusing the + pointwise logprob->score conversion. + """ + + def setUp(self): + super().setUp() + self.h = _GenHarness() + + @staticmethod + def _result(anchor_logprobs, prompt_tokens=7): + # anchor_logprobs: list (per anchor) of [(logprob, token_id, text), ...]. + return { + "meta_info": { + "id": "rid-1", + "prompt_tokens": prompt_tokens, + "input_token_ids_logprobs": anchor_logprobs, + } + } + + def test_generation_setwise_one_matrix_per_item(self): + result = self._result( + [ + [(-0.1, 10, "a"), (-2.0, 20, "b")], + [(-0.5, 10, "a"), (-1.0, 20, "b")], + ] + ) + res = self.h._process_single_item_scoring_results( + [result], + label_token_ids=[[10, 20]], + apply_softmax=False, + per_item_matrix=True, + ) + self.assertIsInstance(res, ScoreResult) + self.assertEqual(len(res.scores), 1) # one item + self.assertEqual(len(res.scores[0]), 2) # two anchors + self.assertAlmostEqual(res.scores[0][0][0], math.exp(-0.1), places=5) + self.assertAlmostEqual(res.scores[0][1][1], math.exp(-1.0), places=5) + self.assertEqual(res.prompt_tokens, 7) + + def test_generation_setwise_multiple_items_grouped_per_item(self): + r0 = self._result([[(-0.1, 10, "a")], [(-0.2, 10, "a")]]) + r1 = self._result([[(-0.3, 10, "a")]]) + res = self.h._process_single_item_scoring_results( + [r0, r1], + label_token_ids=[[10], [10]], + apply_softmax=False, + per_item_matrix=True, + ) + self.assertEqual(len(res.scores), 2) + self.assertEqual(len(res.scores[0]), 2) # item 0: two anchors + self.assertEqual(len(res.scores[1]), 1) # item 1: one anchor + + def test_generation_setwise_softmax_over_labels(self): + result = self._result([[(0.0, 10, "a"), (0.0, 20, "b")]]) + res = self.h._process_single_item_scoring_results( + [result], + label_token_ids=[[10, 20]], + apply_softmax=True, + per_item_matrix=True, + ) + row = res.scores[0][0] + self.assertAlmostEqual(sum(row), 1.0, places=5) + self.assertAlmostEqual(row[0], 0.5, places=5) + + def test_generation_setwise_missing_label_is_zero(self): + # A label token absent from an anchor's logprobs -> -inf -> 0.0 probability. + result = self._result([[(-0.1, 10, "a")]]) + res = self.h._process_single_item_scoring_results( + [result], + label_token_ids=[[10, 999]], + apply_softmax=False, + per_item_matrix=True, + ) + self.assertAlmostEqual(res.scores[0][0][0], math.exp(-0.1), places=5) + self.assertEqual(res.scores[0][0][1], 0.0) + + def test_generation_setwise_empty_logprobs_raises(self): + # No anchor logprobs -> hard error (a missing anchor would silently drop + # that item's scores). + with self.assertRaisesRegex(RuntimeError, "input_token_ids_logprobs"): + self.h._process_single_item_scoring_results( + [self._result([])], + label_token_ids=[[10]], + apply_softmax=False, + per_item_matrix=True, + ) + + def test_generation_pointwise_unchanged(self): + # Without per_item_matrix, the pointwise path still reads + # output_token_ids_logprobs (last position), one flat row per item. + result = { + "meta_info": { + "id": "rid-1", + "prompt_tokens": 5, + "output_token_ids_logprobs": [[(-0.1, 10, "a"), (-2.0, 20, "b")]], + } + } + res = self.h._process_single_item_scoring_results( + [result], + label_token_ids=[[10, 20]], + apply_softmax=False, + ) + self.assertEqual(len(res.scores), 1) + self.assertAlmostEqual(res.scores[0][0], math.exp(-0.1), places=5) + + def test_generation_mis_extraction_groups_scores_per_item(self): + # Fused --enable-mis request: one result with ΣNᵢ anchor logprobs, split + # back per item via per_item_anchor_counts (nested scores). + result = self._result( + [ + [(-0.1, 10, "a"), (-2.0, 20, "b")], # item 0, anchor 0 + [(-0.5, 10, "a"), (-1.0, 20, "b")], # item 0, anchor 1 + [(-0.3, 10, "a"), (-1.5, 20, "b")], # item 1, anchor 0 + ] + ) + res = self.h._process_multi_item_extraction_results( + [result], + per_item_anchor_counts=[2, 1], + apply_softmax=False, + label_token_ids=[[10, 20], [10, 20]], + ) + self.assertEqual(len(res.scores), 2) # one matrix per item + self.assertEqual(len(res.scores[0]), 2) # item 0: two anchors + self.assertEqual(len(res.scores[1]), 1) # item 1: one anchor + self.assertAlmostEqual(res.scores[0][0][0], math.exp(-0.1), places=5) + self.assertAlmostEqual(res.scores[1][0][1], math.exp(-1.5), places=5) + + def test_generation_mis_extraction_row_count_mismatch_raises(self): + # Anchor logprob count must equal sum(per_item_anchor_counts). + result = self._result([[(-0.1, 10, "a")], [(-0.2, 10, "a")]]) + with self.assertRaisesRegex(RuntimeError, "anchor rows"): + self.h._process_multi_item_extraction_results( + [result], + per_item_anchor_counts=[2, 1], + apply_softmax=False, + label_token_ids=[[10], [10]], + ) + + class TestSetwisePooledHiddenStatesRoundTrip(CustomTestCase): """Validate per-position pooled hidden states survive scheduler serialization.