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

This commit is contained in:
Sundara Raman Ramachandran
2026-04-08 01:30:58 -07:00
committed by GitHub
parent 213af1d4f7
commit 712c8c5051
11 changed files with 1287 additions and 131 deletions
@@ -33,36 +33,28 @@ class EngineScoreMixin:
item_first: bool = False,
) -> ScoreResult:
"""
Score the probability of specified token IDs appearing after the given (query + item) pair. For example:
query = "<|user|>Is the following city the capital of France? "
items = ["Paris <|assistant|>", "London <|assistant|>", "Berlin <|assistant|>"]
label_token_ids = [2332, 1223] # Token IDs for "Yes" and "No"
item_first = False
Score items against a query using the loaded model.
This would pass the following prompts to the model:
"<|user|>Is the following city the capital of France? Paris <|assistant|>"
"<|user|>Is the following city the capital of France? London <|assistant|>"
"<|user|>Is the following city the capital of France? Berlin <|assistant|>"
The api would then return the probabilities of the model producing "Yes" and "No" as the next token.
The output would look like:
[[0.9, 0.1], [0.2, 0.8], [0.1, 0.9]]
For generation (CausalLM) models, returns the probability of each label_token_id
being generated after the query+item prompt. Example:
query = "<|user|>Is the following city the capital of France? "
items = ["Paris <|assistant|>", "London <|assistant|>"]
label_token_ids = [2332, 1223] # "Yes" / "No"
# -> [[0.9, 0.1], [0.2, 0.8]]
For SequenceClassification models, returns the pooled class logits directly from
the classification head. label_token_ids is optional and ignored.
Args:
query: The query text or pre-tokenized query token IDs. Must be provided.
items: The item text(s) or pre-tokenized item token IDs. Must be provided.
label_token_ids: List of token IDs to compute probabilities for. If None, no token probabilities will be computed.
apply_softmax: Whether to normalize probabilities using softmax.
item_first: If True, prepend items to query. Otherwise append items to query.
query: The query text or pre-tokenized token IDs.
items: The item text(s) or pre-tokenized token IDs.
label_token_ids: Token IDs to score (required for CausalLM; ignored for
SequenceClassification).
apply_softmax: Whether to normalize scores using softmax.
item_first: If True, prepend items before query (single-item mode only).
Returns:
ScoreResult with:
scores: List of lists containing probabilities for each item and each label token
prompt_tokens: The number of prompt tokens processed.
Raises:
ValueError: If query is not provided, or if items is not provided,
or if token IDs are out of vocabulary, or if logprobs are not available for the specified tokens.
ScoreResult with scores (one list per item) and prompt token count.
"""
return self.loop.run_until_complete(
self.tokenizer_manager.score_request(
@@ -83,11 +75,7 @@ class EngineScoreMixin:
apply_softmax: bool = False,
item_first: bool = False,
) -> ScoreResult:
"""
Asynchronous version of score method.
See score() for detailed documentation.
"""
"""Asynchronous version of score(). See score() for full documentation."""
return await self.tokenizer_manager.score_request(
query=query,
items=items,
+1 -1
View File
@@ -1642,7 +1642,7 @@ async def retrieve_model(model: str):
@app.post("/v1/score", dependencies=[Depends(validate_json_request)])
async def v1_score_request(request: ScoringRequest, raw_request: Request):
"""Endpoint for the decoder-only scoring API. See Engine.score() for detailed documentation."""
"""Endpoint for the scoring API. Supports CausalLM (logprob-based) and SequenceClassification (class logit-based) models. See Engine.score() for documentation."""
return await raw_request.app.state.openai_serving_score.handle_request(
request, raw_request
)
+49
View File
@@ -11,6 +11,7 @@ from transformers import PretrainedConfig
from sglang.srt.layers.activation import get_cross_encoder_activation_function
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.server_args import get_global_server_args
class PoolingType(IntEnum):
@@ -25,6 +26,54 @@ class EmbeddingPoolerOutput:
embeddings: torch.Tensor | list[torch.Tensor]
def score_and_pool(
score_head: nn.Module,
pooler: "Pooler",
hidden_states: torch.Tensor,
forward_batch: ForwardBatch,
input_ids: torch.Tensor,
) -> EmbeddingPoolerOutput:
"""Apply a classification/score head with multi-item scoring (MIS) support.
When ``multi_item_scoring_delimiter`` is configured and found in
``input_ids``, takes the MIS path: extract hidden states at the positions
just before each delimiter, apply the score head only to those positions,
then split results per-request using ``forward_batch.extend_seq_lens``.
Otherwise, takes the normal single-item path: apply the score head to all
hidden states, then pool (matching the original classification model
forward logic).
"""
delimiter_token = get_global_server_args().multi_item_scoring_delimiter
if delimiter_token is not None and forward_batch.is_prefill_only:
delim_positions = (input_ids == delimiter_token).nonzero(as_tuple=True)[0]
# A delimiter at flat index 0 has no preceding hidden state to pool
delim_positions = delim_positions[delim_positions > 0]
if delim_positions.numel() > 0:
# Score only the tokens that precede a delimiter
scores = score_head(hidden_states[delim_positions - 1])
# Split per-request so the scheduler gets one tensor per request.
# Use CPU sequence lengths to avoid per-iteration GPU<->CPU sync
# from `.item()` calls on device tensors.
seq_lens = forward_batch.extend_seq_lens_cpu
start = 0
per_request = []
for seq_len in seq_lens:
end = start + seq_len
mask = (delim_positions >= start) & (delim_positions < end)
per_request.append(scores[mask])
start = end
return EmbeddingPoolerOutput(embeddings=per_request)
# Standard classification path: score all tokens, then pool.
logits = score_head(hidden_states)
pooled_logits = pooler(logits, forward_batch).embeddings
return EmbeddingPoolerOutput(embeddings=pooled_logits)
class Pooler(nn.Module):
"""A layer that pools specific information from hidden states.
This layer does the following:
@@ -1,9 +1,11 @@
import logging
import math
from dataclasses import dataclass
from typing import Any, Dict, List, Optional, Union
from typing import Any, Dict, List, Optional, Tuple, Union
from sglang.srt.managers.io_struct import GenerateReqInput
import torch
from sglang.srt.managers.io_struct import EmbeddingReqInput, GenerateReqInput
logger = logging.getLogger(__name__)
@@ -11,7 +13,7 @@ logger = logging.getLogger(__name__)
@dataclass(frozen=True, slots=True)
class ScoreResult:
scores: List[List[float]]
prompt_tokens: int
prompt_tokens: int = 0
class TokenizerManagerScoreMixin:
@@ -110,17 +112,52 @@ class TokenizerManagerScoreMixin:
return combined_sequence
def _batch_tokenize_query_and_items(
self,
query: Optional[Union[str, List[int]]],
items: Optional[Union[str, List[str], List[List[int]]]],
) -> Tuple[List[int], List[List[int]]]:
"""
Tokenize query and items into token IDs.
Args:
query: The query text (str) or pre-tokenized token IDs (List[int]).
items: Item texts or pre-tokenized token IDs.
Returns:
(query_ids, items_ids): query token IDs and list of per-item token IDs.
"""
if isinstance(query, str):
query_ids = self.tokenizer.encode(query)
else:
query_ids = list(query)
items_list = [items] if isinstance(items, str) else items
items_ids = []
for item in items_list:
if isinstance(item, str):
items_ids.append(self.tokenizer.encode(item))
else:
items_ids.append(list(item))
return query_ids, items_ids
def _process_multi_item_scoring_results(
self,
results: Any,
items: List,
label_token_ids: List[int],
label_token_ids: Optional[List[int]],
apply_softmax: bool,
batch_request=None,
) -> ScoreResult:
"""
Process results from multi-item scoring request.
Extracts logprobs at delimiter positions from input_token_ids_logprobs.
Extracts per-delimiter scores from whichever field the scheduler
populated (input_token_ids_logprobs for generation models,
embedding for classification models), then uniformly validates,
skips the query-boundary delimiter, and normalizes.
Args:
results: Results from generate_request
@@ -134,60 +171,74 @@ class TokenizerManagerScoreMixin:
scores: List of score lists, one for each prompt, each in the order of label_token_ids.
prompt_tokens: The number of prompt tokens processed.
"""
result = results[0] if isinstance(results, list) else results
meta_info = result.get("meta_info", {})
# For multi-item scoring, logprobs are in input_token_ids_logprobs
input_logprobs = meta_info.get("input_token_ids_logprobs", [])
prompt_tokens = meta_info.get("prompt_tokens", 0)
single_result = results[0] if isinstance(results, list) else results
meta_info = single_result.get("meta_info", {})
num_items = len(items) if isinstance(items, list) else 1
expected_count = num_items + 1
request_id = meta_info.get("id", "<unknown>")
prompt_tokens = meta_info.get("prompt_tokens", 0)
if not input_logprobs:
# Extract per-delimiter scores from whichever field has them
input_logprobs = meta_info.get("input_token_ids_logprobs", [])
embedding = single_result.get("embedding")
if input_logprobs:
# Generation model: extract label-token logprobs at each delimiter
per_delimiter_scores = []
for logprobs_data in input_logprobs:
logprobs = self._extract_logprobs_for_tokens(
logprobs_data, label_token_ids
)
score_list = self._convert_logprobs_to_scores(
logprobs, label_token_ids, apply_softmax
)
per_delimiter_scores.append(score_list)
elif embedding is not None:
# Classification model: scores are directly in 2D embedding.
if apply_softmax:
scores_tensor = (
torch.tensor(embedding)
if isinstance(embedding, list)
else embedding
)
scores_tensor = torch.nn.functional.softmax(scores_tensor, dim=-1)
per_delimiter_scores = scores_tensor.tolist()
else:
per_delimiter_scores = (
embedding if isinstance(embedding, list) else embedding.tolist()
)
else:
raise RuntimeError(
f"input_token_ids_logprobs is empty for multi-item scoring request {request_id}. "
"This indicates token_ids_logprobs were not computed properly for Multi-Item Scoring."
f"No scoring data found for multi-item scoring request {request_id}. "
"Expected either input_token_ids_logprobs or embedding."
)
scores = []
num_items = len(items) if isinstance(items, list) else 1
# Check if we have the expected number of logprobs
expected_logprobs_count = num_items + 1
if len(input_logprobs) != expected_logprobs_count:
# Validate delimiter count
if len(per_delimiter_scores) != expected_count:
raise RuntimeError(
f"Expected {expected_logprobs_count} input_token_ids_logprobs for multi-item scoring "
f"with {num_items} items, but got {len(input_logprobs)}. "
f"Expected {expected_count} delimiter entries for multi-item scoring "
f"with {num_items} items, but got {len(per_delimiter_scores)}. "
f"Request ID: {request_id}"
)
# Skip the first delimiter (between query and first item) and process remaining delimiter positions
# We want to exclude the first one since it represents the boundary between query and first item, not an item boundary
start_idx = 1 if len(input_logprobs) > 1 else 0
# Process logprobs for each item position (excluding first delimiter)
for item_idx in range(num_items):
logprob_idx = start_idx + item_idx
item_logprobs_data = input_logprobs[logprob_idx]
logprobs = self._extract_logprobs_for_tokens(
item_logprobs_data, label_token_ids
)
score_list = self._convert_logprobs_to_scores(
logprobs, label_token_ids, apply_softmax
)
scores.append(score_list)
# Skip the first delimiter (query-item boundary)
scores = per_delimiter_scores[1:]
return ScoreResult(scores=scores, prompt_tokens=prompt_tokens)
def _process_single_item_scoring_results(
self, results: Any, label_token_ids: List[int], apply_softmax: bool
self, results: Any, label_token_ids: Optional[List[int]], apply_softmax: bool
) -> ScoreResult:
"""
Process results from single-item scoring request.
Single-item scoring results are stored in output_token_ids_logprobs.
For generation (CausalLM) models: reads output_token_ids_logprobs.
For non-generation (SequenceClassification) models: reads the embedding field
which contains pooled class logits from the classification head.
Args:
results: Results from generate_request
label_token_ids: Token IDs to extract scores for
label_token_ids: Token IDs to extract scores for (generation models only)
apply_softmax: Whether to apply softmax normalization
Returns:
@@ -198,24 +249,48 @@ class TokenizerManagerScoreMixin:
scores = []
prompt_tokens = 0
for result in results:
# 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)
if not output_logprobs or len(output_logprobs) == 0:
raise RuntimeError(
f"output_logprobs is empty for request {result['meta_info'].get('id', '<unknown>')}."
is_generation = getattr(self, "is_generation", True)
if is_generation:
for result in results:
# 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)
# Extract logprobs for the first (and only) position
logprobs = self._extract_logprobs_for_tokens(
output_logprobs[0], label_token_ids
)
score_list = self._convert_logprobs_to_scores(
logprobs, label_token_ids, apply_softmax
)
scores.append(score_list)
if not output_logprobs or len(output_logprobs) == 0:
raise RuntimeError(
f"output_logprobs is empty for request "
f"{result['meta_info'].get('id', '<unknown>')}."
)
# Extract logprobs for the first (and only) position
logprobs = self._extract_logprobs_for_tokens(
output_logprobs[0], label_token_ids
)
score_list = self._convert_logprobs_to_scores(
logprobs, label_token_ids, apply_softmax
)
scores.append(score_list)
else:
for result in results:
embedding = result.get("embedding", None)
if embedding is None:
raise ValueError("Embedding not found in the result.")
prompt_tokens += result.get("meta_info", {}).get("prompt_tokens", 0)
if apply_softmax:
embedding = torch.softmax(
torch.as_tensor(embedding), dim=-1
).tolist()
# The classification head produces per-token logits, which the pooler reduces
# into a single vector per input. That vector is returned in the `.embeddings`
# field — not as semantic embeddings, but as pooled classification logits.
# The field name is reused for compatibility with the existing
# EmbeddingPoolerOutput API.
scores.append(embedding)
return ScoreResult(scores=scores, prompt_tokens=prompt_tokens)
@@ -242,6 +317,10 @@ class TokenizerManagerScoreMixin:
- Text: query<delimiter_text>item1<delimiter_text>item2<delimiter_text>item3<delimiter_text>
- Tokens: query<delimiter_token_id>item1<delimiter_token_id>item2<delimiter_token_id>item3<delimiter_token_id>
Supports two model types:
- Generation (CausalLM): Requires label_token_ids; returns logprob-based scores.
- SequenceClassification: label_token_ids is optional; returns pooled class logits.
Args:
query: The query text or pre-tokenized query token IDs
items: The item text(s) or pre-tokenized item token IDs
@@ -255,15 +334,19 @@ class TokenizerManagerScoreMixin:
scores: List of score lists, one for each prompt, each in the order of label_token_ids.
prompt_tokens: The number of prompt tokens processed.
"""
if label_token_ids is None:
raise ValueError("label_token_ids must be provided")
is_generation = getattr(self, "is_generation", True)
if is_generation and label_token_ids is None:
raise ValueError(
"label_token_ids is required for generation (CausalLM) models."
)
if items is None:
raise ValueError("items must be provided")
if not items:
return ScoreResult(scores=[], prompt_tokens=0)
if self.tokenizer is not None:
if self.tokenizer is not None and label_token_ids is not None:
vocab_size = self.tokenizer.vocab_size
for token_id in label_token_ids:
if token_id >= vocab_size:
@@ -277,14 +360,8 @@ class TokenizerManagerScoreMixin:
and self.multi_item_delimiter_text is not None
)
batch_request = GenerateReqInput(
token_ids_logprob=label_token_ids,
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,
stream=False,
sampling_params={"max_new_tokens": 0},
)
input_ids = None
text_prompts = None
# Handle string or tokenized query/items
if isinstance(query, str) and (
@@ -295,21 +372,23 @@ class TokenizerManagerScoreMixin:
items_list = [items] if isinstance(items, str) else items
if use_multi_item_scoring:
# Multi-item scoring: create single prompt with delimiter text
# Always use format: query<delimiter>item1<delimiter>item2<delimiter>item3<delimiter>
# (item_first is ignored for multi-item scoring)
delimiter = self.multi_item_delimiter_text
combined_items = delimiter.join(items_list)
# Add final delimiter after the last item for logprob extraction
single_prompt = f"{query}{delimiter}{combined_items}{delimiter}"
batch_request.text = [single_prompt]
# Multi-item scoring: tokenize separately then combine at token level
# to ensure the delimiter token ID is inserted exactly once per boundary
# (a text-level roundtrip through the tokenizer can alter boundary tokens)
delimiter_token_id = self.server_args.multi_item_scoring_delimiter
query_ids, items_ids = self._batch_tokenize_query_and_items(
query, items_list
)
combined_input_ids = self._build_multi_item_token_sequence(
query_ids, items_ids, delimiter_token_id
)
input_ids = [combined_input_ids]
else:
# Single-item scoring: create separate prompts for each item
if item_first:
prompts = [f"{item}{query}" for item in items_list]
text_prompts = [f"{item}{query}" for item in items_list]
else:
prompts = [f"{query}{item}" for item in items_list]
batch_request.text = prompts
text_prompts = [f"{query}{item}" for item in items_list]
elif (
isinstance(query, list)
@@ -325,23 +404,40 @@ class TokenizerManagerScoreMixin:
combined_input_ids = self._build_multi_item_token_sequence(
query, items, delimiter_token_id
)
batch_request.input_ids = [combined_input_ids]
input_ids = [combined_input_ids]
else:
# Single-item scoring: process each item separately
if item_first:
input_ids_list = [item + query for item in items]
input_ids = [item + query for item in items]
else:
input_ids_list = [query + item for item in items]
batch_request.input_ids = input_ids_list
input_ids = [query + item for item in items]
else:
raise ValueError(
"Invalid combination of query/items types for score_request."
)
# Create the appropriate request type
if is_generation:
batch_request = GenerateReqInput(
text=text_prompts,
input_ids=input_ids,
token_ids_logprob=label_token_ids,
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,
stream=False,
sampling_params={"max_new_tokens": 0},
)
else:
batch_request = EmbeddingReqInput(
text=text_prompts,
input_ids=input_ids,
)
results = await self.generate_request(batch_request, request).__anext__()
if use_multi_item_scoring:
# Multi-item scoring: extract scores from input_token_ids_logprobs
# Multi-item scoring: extract scores from input_token_ids_logprobs or embedding
return self._process_multi_item_scoring_results(
results, items, label_token_ids, apply_softmax, batch_request
)
@@ -368,8 +464,6 @@ class TokenizerManagerScoreMixin:
Returns:
List of scores in the same order as label_token_ids
"""
import torch
score_list = [
logprobs.get(token_id, float("-inf")) for token_id in label_token_ids
]
@@ -18,7 +18,12 @@ import torch
from torch import nn
from transformers import LlamaConfig
from sglang.srt.layers.pooler import EmbeddingPoolerOutput, Pooler, PoolingType
from sglang.srt.layers.pooler import (
EmbeddingPoolerOutput,
Pooler,
PoolingType,
score_and_pool,
)
from sglang.srt.layers.quantization.base_config import QuantizationConfig
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.model_loader.weight_utils import default_weight_loader
@@ -59,10 +64,13 @@ class LlamaForClassification(nn.Module):
), "LlamaForClassification is only used for embedding. Please add --is-embedding when you launch the server."
hidden_states = self.model(input_ids, positions, forward_batch, input_embeds)
last_token_hidden = self.pooler(hidden_states, forward_batch).embeddings
scores = self.classification_head(last_token_hidden)
return EmbeddingPoolerOutput(scores)
return score_and_pool(
self.classification_head,
self.pooler,
hidden_states,
forward_batch,
input_ids,
)
def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]):
params_dict = dict(self.named_parameters())
@@ -18,7 +18,12 @@ import torch
from torch import nn
from transformers import Qwen2Config
from sglang.srt.layers.pooler import EmbeddingPoolerOutput, Pooler, PoolingType
from sglang.srt.layers.pooler import (
EmbeddingPoolerOutput,
Pooler,
PoolingType,
score_and_pool,
)
from sglang.srt.layers.quantization.base_config import QuantizationConfig
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.models.qwen2 import Qwen2ForCausalLM, Qwen2Model
@@ -57,10 +62,9 @@ class Qwen2ForSequenceClassification(nn.Module):
), "Qwen2ForSequenceClassification is only used for embedding"
hidden_states = self.model(input_ids, positions, forward_batch, input_embeds)
logits = self.score(hidden_states)
pooled_logits = self.pooler(logits, forward_batch).embeddings
return EmbeddingPoolerOutput(pooled_logits)
return score_and_pool(
self.score, self.pooler, hidden_states, forward_batch, input_ids
)
def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]):
# Filter out lm_head weights of Qwen2ForCausalLM
@@ -19,7 +19,12 @@ import torch
from torch import nn
from transformers import Qwen2Config # Qwen3 uses Qwen2Config
from sglang.srt.layers.pooler import EmbeddingPoolerOutput, Pooler, PoolingType
from sglang.srt.layers.pooler import (
EmbeddingPoolerOutput,
Pooler,
PoolingType,
score_and_pool,
)
from sglang.srt.layers.quantization.base_config import QuantizationConfig
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.model_loader.weight_utils import default_weight_loader
@@ -62,10 +67,9 @@ class Qwen3ForPooledOutput(nn.Module):
assert get_embedding, f"{self.__class__.__name__} is only used for embedding"
hidden_states = self.model(input_ids, positions, forward_batch, input_embeds)
logits = self.score(hidden_states)
pooled_logits = self.pooler(logits, forward_batch).embeddings
return EmbeddingPoolerOutput(pooled_logits)
return score_and_pool(
self.score, self.pooler, hidden_states, forward_batch, input_ids
)
def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]):
stacked_params_mapping = [