[Score API] Add SequenceClassification Model support (#22118)
This commit is contained in:
+1
-1
@@ -1,3 +1,3 @@
|
|||||||
[codespell]
|
[codespell]
|
||||||
ignore-words-list = ans, als, hel, boostrap, childs, te, vas, hsa, ment, cann, thi, makro, wil, rouge, PRIS, ather
|
ignore-words-list = ans, als, hel, boostrap, childs, te, vas, hsa, ment, cann, thi, makro, wil, rouge, PRIS, ather, MIS
|
||||||
skip = *.json,*.jsonl,*.patch,*.txt
|
skip = *.json,*.jsonl,*.patch,*.txt
|
||||||
|
|||||||
@@ -33,36 +33,28 @@ class EngineScoreMixin:
|
|||||||
item_first: bool = False,
|
item_first: bool = False,
|
||||||
) -> ScoreResult:
|
) -> ScoreResult:
|
||||||
"""
|
"""
|
||||||
Score the probability of specified token IDs appearing after the given (query + item) pair. For example:
|
Score items against a query using the loaded model.
|
||||||
|
|
||||||
|
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? "
|
query = "<|user|>Is the following city the capital of France? "
|
||||||
items = ["Paris <|assistant|>", "London <|assistant|>", "Berlin <|assistant|>"]
|
items = ["Paris <|assistant|>", "London <|assistant|>"]
|
||||||
label_token_ids = [2332, 1223] # Token IDs for "Yes" and "No"
|
label_token_ids = [2332, 1223] # "Yes" / "No"
|
||||||
item_first = False
|
# -> [[0.9, 0.1], [0.2, 0.8]]
|
||||||
|
|
||||||
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 SequenceClassification models, returns the pooled class logits directly from
|
||||||
|
the classification head. label_token_ids is optional and ignored.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
query: The query text or pre-tokenized query token IDs. Must be provided.
|
query: The query text or pre-tokenized token IDs.
|
||||||
items: The item text(s) or pre-tokenized item token IDs. Must be provided.
|
items: The item text(s) or pre-tokenized token IDs.
|
||||||
label_token_ids: List of token IDs to compute probabilities for. If None, no token probabilities will be computed.
|
label_token_ids: Token IDs to score (required for CausalLM; ignored for
|
||||||
apply_softmax: Whether to normalize probabilities using softmax.
|
SequenceClassification).
|
||||||
item_first: If True, prepend items to query. Otherwise append items to query.
|
apply_softmax: Whether to normalize scores using softmax.
|
||||||
|
item_first: If True, prepend items before query (single-item mode only).
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
ScoreResult with:
|
ScoreResult with scores (one list per item) and prompt token count.
|
||||||
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.
|
|
||||||
"""
|
"""
|
||||||
return self.loop.run_until_complete(
|
return self.loop.run_until_complete(
|
||||||
self.tokenizer_manager.score_request(
|
self.tokenizer_manager.score_request(
|
||||||
@@ -83,11 +75,7 @@ class EngineScoreMixin:
|
|||||||
apply_softmax: bool = False,
|
apply_softmax: bool = False,
|
||||||
item_first: bool = False,
|
item_first: bool = False,
|
||||||
) -> ScoreResult:
|
) -> ScoreResult:
|
||||||
"""
|
"""Asynchronous version of score(). See score() for full documentation."""
|
||||||
Asynchronous version of score method.
|
|
||||||
|
|
||||||
See score() for detailed documentation.
|
|
||||||
"""
|
|
||||||
return await self.tokenizer_manager.score_request(
|
return await self.tokenizer_manager.score_request(
|
||||||
query=query,
|
query=query,
|
||||||
items=items,
|
items=items,
|
||||||
|
|||||||
@@ -1642,7 +1642,7 @@ async def retrieve_model(model: str):
|
|||||||
|
|
||||||
@app.post("/v1/score", dependencies=[Depends(validate_json_request)])
|
@app.post("/v1/score", dependencies=[Depends(validate_json_request)])
|
||||||
async def v1_score_request(request: ScoringRequest, raw_request: 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(
|
return await raw_request.app.state.openai_serving_score.handle_request(
|
||||||
request, raw_request
|
request, raw_request
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -11,6 +11,7 @@ from transformers import PretrainedConfig
|
|||||||
|
|
||||||
from sglang.srt.layers.activation import get_cross_encoder_activation_function
|
from sglang.srt.layers.activation import get_cross_encoder_activation_function
|
||||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
||||||
|
from sglang.srt.server_args import get_global_server_args
|
||||||
|
|
||||||
|
|
||||||
class PoolingType(IntEnum):
|
class PoolingType(IntEnum):
|
||||||
@@ -25,6 +26,54 @@ class EmbeddingPoolerOutput:
|
|||||||
embeddings: torch.Tensor | list[torch.Tensor]
|
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):
|
class Pooler(nn.Module):
|
||||||
"""A layer that pools specific information from hidden states.
|
"""A layer that pools specific information from hidden states.
|
||||||
This layer does the following:
|
This layer does the following:
|
||||||
|
|||||||
@@ -1,9 +1,11 @@
|
|||||||
import logging
|
import logging
|
||||||
import math
|
import math
|
||||||
from dataclasses import dataclass
|
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__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
@@ -11,7 +13,7 @@ logger = logging.getLogger(__name__)
|
|||||||
@dataclass(frozen=True, slots=True)
|
@dataclass(frozen=True, slots=True)
|
||||||
class ScoreResult:
|
class ScoreResult:
|
||||||
scores: List[List[float]]
|
scores: List[List[float]]
|
||||||
prompt_tokens: int
|
prompt_tokens: int = 0
|
||||||
|
|
||||||
|
|
||||||
class TokenizerManagerScoreMixin:
|
class TokenizerManagerScoreMixin:
|
||||||
@@ -110,17 +112,52 @@ class TokenizerManagerScoreMixin:
|
|||||||
|
|
||||||
return combined_sequence
|
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(
|
def _process_multi_item_scoring_results(
|
||||||
self,
|
self,
|
||||||
results: Any,
|
results: Any,
|
||||||
items: List,
|
items: List,
|
||||||
label_token_ids: List[int],
|
label_token_ids: Optional[List[int]],
|
||||||
apply_softmax: bool,
|
apply_softmax: bool,
|
||||||
batch_request=None,
|
batch_request=None,
|
||||||
) -> ScoreResult:
|
) -> ScoreResult:
|
||||||
"""
|
"""
|
||||||
Process results from multi-item scoring request.
|
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:
|
Args:
|
||||||
results: Results from generate_request
|
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.
|
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.
|
prompt_tokens: The number of prompt tokens processed.
|
||||||
"""
|
"""
|
||||||
result = results[0] if isinstance(results, list) else results
|
single_result = results[0] if isinstance(results, list) else results
|
||||||
meta_info = result.get("meta_info", {})
|
meta_info = single_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)
|
|
||||||
request_id = meta_info.get("id", "<unknown>")
|
|
||||||
|
|
||||||
if not input_logprobs:
|
|
||||||
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."
|
|
||||||
)
|
|
||||||
|
|
||||||
scores = []
|
|
||||||
num_items = len(items) if isinstance(items, list) else 1
|
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)
|
||||||
|
|
||||||
# Check if we have the expected number of logprobs
|
# Extract per-delimiter scores from whichever field has them
|
||||||
expected_logprobs_count = num_items + 1
|
input_logprobs = meta_info.get("input_token_ids_logprobs", [])
|
||||||
if len(input_logprobs) != expected_logprobs_count:
|
embedding = single_result.get("embedding")
|
||||||
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"Request ID: {request_id}"
|
|
||||||
)
|
|
||||||
|
|
||||||
# Skip the first delimiter (between query and first item) and process remaining delimiter positions
|
if input_logprobs:
|
||||||
# We want to exclude the first one since it represents the boundary between query and first item, not an item boundary
|
# Generation model: extract label-token logprobs at each delimiter
|
||||||
start_idx = 1 if len(input_logprobs) > 1 else 0
|
per_delimiter_scores = []
|
||||||
|
for logprobs_data in input_logprobs:
|
||||||
# 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(
|
logprobs = self._extract_logprobs_for_tokens(
|
||||||
item_logprobs_data, label_token_ids
|
logprobs_data, label_token_ids
|
||||||
)
|
)
|
||||||
score_list = self._convert_logprobs_to_scores(
|
score_list = self._convert_logprobs_to_scores(
|
||||||
logprobs, label_token_ids, apply_softmax
|
logprobs, label_token_ids, apply_softmax
|
||||||
)
|
)
|
||||||
scores.append(score_list)
|
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"No scoring data found for multi-item scoring request {request_id}. "
|
||||||
|
"Expected either input_token_ids_logprobs or embedding."
|
||||||
|
)
|
||||||
|
|
||||||
|
# Validate delimiter count
|
||||||
|
if len(per_delimiter_scores) != expected_count:
|
||||||
|
raise RuntimeError(
|
||||||
|
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 (query-item boundary)
|
||||||
|
scores = per_delimiter_scores[1:]
|
||||||
|
|
||||||
return ScoreResult(scores=scores, prompt_tokens=prompt_tokens)
|
return ScoreResult(scores=scores, prompt_tokens=prompt_tokens)
|
||||||
|
|
||||||
def _process_single_item_scoring_results(
|
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:
|
) -> ScoreResult:
|
||||||
"""
|
"""
|
||||||
Process results from single-item scoring request.
|
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:
|
Args:
|
||||||
results: Results from generate_request
|
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
|
apply_softmax: Whether to apply softmax normalization
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
@@ -198,14 +249,19 @@ class TokenizerManagerScoreMixin:
|
|||||||
scores = []
|
scores = []
|
||||||
prompt_tokens = 0
|
prompt_tokens = 0
|
||||||
|
|
||||||
|
is_generation = getattr(self, "is_generation", True)
|
||||||
|
if is_generation:
|
||||||
for result in results:
|
for result in results:
|
||||||
# For single-item scoring, logprobs are in output_token_ids_logprobs
|
# For single-item scoring, logprobs are in output_token_ids_logprobs
|
||||||
output_logprobs = result["meta_info"].get("output_token_ids_logprobs", [])
|
output_logprobs = result["meta_info"].get(
|
||||||
|
"output_token_ids_logprobs", []
|
||||||
|
)
|
||||||
prompt_tokens += result["meta_info"].get("prompt_tokens", 0)
|
prompt_tokens += result["meta_info"].get("prompt_tokens", 0)
|
||||||
|
|
||||||
if not output_logprobs or len(output_logprobs) == 0:
|
if not output_logprobs or len(output_logprobs) == 0:
|
||||||
raise RuntimeError(
|
raise RuntimeError(
|
||||||
f"output_logprobs is empty for request {result['meta_info'].get('id', '<unknown>')}."
|
f"output_logprobs is empty for request "
|
||||||
|
f"{result['meta_info'].get('id', '<unknown>')}."
|
||||||
)
|
)
|
||||||
|
|
||||||
# Extract logprobs for the first (and only) position
|
# Extract logprobs for the first (and only) position
|
||||||
@@ -216,6 +272,25 @@ class TokenizerManagerScoreMixin:
|
|||||||
logprobs, label_token_ids, apply_softmax
|
logprobs, label_token_ids, apply_softmax
|
||||||
)
|
)
|
||||||
scores.append(score_list)
|
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)
|
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>
|
- 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>
|
- 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:
|
Args:
|
||||||
query: The query text or pre-tokenized query token IDs
|
query: The query text or pre-tokenized query token IDs
|
||||||
items: The item text(s) or pre-tokenized item 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.
|
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.
|
prompt_tokens: The number of prompt tokens processed.
|
||||||
"""
|
"""
|
||||||
if label_token_ids is None:
|
is_generation = getattr(self, "is_generation", True)
|
||||||
raise ValueError("label_token_ids must be provided")
|
|
||||||
|
if is_generation and label_token_ids is None:
|
||||||
|
raise ValueError(
|
||||||
|
"label_token_ids is required for generation (CausalLM) models."
|
||||||
|
)
|
||||||
|
|
||||||
if items is None:
|
if items is None:
|
||||||
raise ValueError("items must be provided")
|
raise ValueError("items must be provided")
|
||||||
if not items:
|
if not items:
|
||||||
return ScoreResult(scores=[], prompt_tokens=0)
|
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
|
vocab_size = self.tokenizer.vocab_size
|
||||||
for token_id in label_token_ids:
|
for token_id in label_token_ids:
|
||||||
if token_id >= vocab_size:
|
if token_id >= vocab_size:
|
||||||
@@ -277,14 +360,8 @@ class TokenizerManagerScoreMixin:
|
|||||||
and self.multi_item_delimiter_text is not None
|
and self.multi_item_delimiter_text is not None
|
||||||
)
|
)
|
||||||
|
|
||||||
batch_request = GenerateReqInput(
|
input_ids = None
|
||||||
token_ids_logprob=label_token_ids,
|
text_prompts = None
|
||||||
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},
|
|
||||||
)
|
|
||||||
|
|
||||||
# Handle string or tokenized query/items
|
# Handle string or tokenized query/items
|
||||||
if isinstance(query, str) and (
|
if isinstance(query, str) and (
|
||||||
@@ -295,21 +372,23 @@ class TokenizerManagerScoreMixin:
|
|||||||
items_list = [items] if isinstance(items, str) else items
|
items_list = [items] if isinstance(items, str) else items
|
||||||
|
|
||||||
if use_multi_item_scoring:
|
if use_multi_item_scoring:
|
||||||
# Multi-item scoring: create single prompt with delimiter text
|
# Multi-item scoring: tokenize separately then combine at token level
|
||||||
# Always use format: query<delimiter>item1<delimiter>item2<delimiter>item3<delimiter>
|
# to ensure the delimiter token ID is inserted exactly once per boundary
|
||||||
# (item_first is ignored for multi-item scoring)
|
# (a text-level roundtrip through the tokenizer can alter boundary tokens)
|
||||||
delimiter = self.multi_item_delimiter_text
|
delimiter_token_id = self.server_args.multi_item_scoring_delimiter
|
||||||
combined_items = delimiter.join(items_list)
|
query_ids, items_ids = self._batch_tokenize_query_and_items(
|
||||||
# Add final delimiter after the last item for logprob extraction
|
query, items_list
|
||||||
single_prompt = f"{query}{delimiter}{combined_items}{delimiter}"
|
)
|
||||||
batch_request.text = [single_prompt]
|
combined_input_ids = self._build_multi_item_token_sequence(
|
||||||
|
query_ids, items_ids, delimiter_token_id
|
||||||
|
)
|
||||||
|
input_ids = [combined_input_ids]
|
||||||
else:
|
else:
|
||||||
# Single-item scoring: create separate prompts for each item
|
# Single-item scoring: create separate prompts for each item
|
||||||
if item_first:
|
if item_first:
|
||||||
prompts = [f"{item}{query}" for item in items_list]
|
text_prompts = [f"{item}{query}" for item in items_list]
|
||||||
else:
|
else:
|
||||||
prompts = [f"{query}{item}" for item in items_list]
|
text_prompts = [f"{query}{item}" for item in items_list]
|
||||||
batch_request.text = prompts
|
|
||||||
|
|
||||||
elif (
|
elif (
|
||||||
isinstance(query, list)
|
isinstance(query, list)
|
||||||
@@ -325,23 +404,40 @@ class TokenizerManagerScoreMixin:
|
|||||||
combined_input_ids = self._build_multi_item_token_sequence(
|
combined_input_ids = self._build_multi_item_token_sequence(
|
||||||
query, items, delimiter_token_id
|
query, items, delimiter_token_id
|
||||||
)
|
)
|
||||||
batch_request.input_ids = [combined_input_ids]
|
input_ids = [combined_input_ids]
|
||||||
else:
|
else:
|
||||||
# Single-item scoring: process each item separately
|
# Single-item scoring: process each item separately
|
||||||
if item_first:
|
if item_first:
|
||||||
input_ids_list = [item + query for item in items]
|
input_ids = [item + query for item in items]
|
||||||
else:
|
else:
|
||||||
input_ids_list = [query + item for item in items]
|
input_ids = [query + item for item in items]
|
||||||
batch_request.input_ids = input_ids_list
|
|
||||||
else:
|
else:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
"Invalid combination of query/items types for score_request."
|
"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__()
|
results = await self.generate_request(batch_request, request).__anext__()
|
||||||
|
|
||||||
if use_multi_item_scoring:
|
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(
|
return self._process_multi_item_scoring_results(
|
||||||
results, items, label_token_ids, apply_softmax, batch_request
|
results, items, label_token_ids, apply_softmax, batch_request
|
||||||
)
|
)
|
||||||
@@ -368,8 +464,6 @@ class TokenizerManagerScoreMixin:
|
|||||||
Returns:
|
Returns:
|
||||||
List of scores in the same order as label_token_ids
|
List of scores in the same order as label_token_ids
|
||||||
"""
|
"""
|
||||||
import torch
|
|
||||||
|
|
||||||
score_list = [
|
score_list = [
|
||||||
logprobs.get(token_id, float("-inf")) for token_id in label_token_ids
|
logprobs.get(token_id, float("-inf")) for token_id in label_token_ids
|
||||||
]
|
]
|
||||||
|
|||||||
@@ -18,7 +18,12 @@ import torch
|
|||||||
from torch import nn
|
from torch import nn
|
||||||
from transformers import LlamaConfig
|
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.layers.quantization.base_config import QuantizationConfig
|
||||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
||||||
from sglang.srt.model_loader.weight_utils import default_weight_loader
|
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."
|
), "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)
|
hidden_states = self.model(input_ids, positions, forward_batch, input_embeds)
|
||||||
last_token_hidden = self.pooler(hidden_states, forward_batch).embeddings
|
return score_and_pool(
|
||||||
scores = self.classification_head(last_token_hidden)
|
self.classification_head,
|
||||||
|
self.pooler,
|
||||||
return EmbeddingPoolerOutput(scores)
|
hidden_states,
|
||||||
|
forward_batch,
|
||||||
|
input_ids,
|
||||||
|
)
|
||||||
|
|
||||||
def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]):
|
def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]):
|
||||||
params_dict = dict(self.named_parameters())
|
params_dict = dict(self.named_parameters())
|
||||||
|
|||||||
@@ -18,7 +18,12 @@ import torch
|
|||||||
from torch import nn
|
from torch import nn
|
||||||
from transformers import Qwen2Config
|
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.layers.quantization.base_config import QuantizationConfig
|
||||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
||||||
from sglang.srt.models.qwen2 import Qwen2ForCausalLM, Qwen2Model
|
from sglang.srt.models.qwen2 import Qwen2ForCausalLM, Qwen2Model
|
||||||
@@ -57,10 +62,9 @@ class Qwen2ForSequenceClassification(nn.Module):
|
|||||||
), "Qwen2ForSequenceClassification is only used for embedding"
|
), "Qwen2ForSequenceClassification is only used for embedding"
|
||||||
|
|
||||||
hidden_states = self.model(input_ids, positions, forward_batch, input_embeds)
|
hidden_states = self.model(input_ids, positions, forward_batch, input_embeds)
|
||||||
logits = self.score(hidden_states)
|
return score_and_pool(
|
||||||
pooled_logits = self.pooler(logits, forward_batch).embeddings
|
self.score, self.pooler, hidden_states, forward_batch, input_ids
|
||||||
|
)
|
||||||
return EmbeddingPoolerOutput(pooled_logits)
|
|
||||||
|
|
||||||
def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]):
|
def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]):
|
||||||
# Filter out lm_head weights of Qwen2ForCausalLM
|
# Filter out lm_head weights of Qwen2ForCausalLM
|
||||||
|
|||||||
@@ -19,7 +19,12 @@ import torch
|
|||||||
from torch import nn
|
from torch import nn
|
||||||
from transformers import Qwen2Config # Qwen3 uses Qwen2Config
|
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.layers.quantization.base_config import QuantizationConfig
|
||||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
||||||
from sglang.srt.model_loader.weight_utils import default_weight_loader
|
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"
|
assert get_embedding, f"{self.__class__.__name__} is only used for embedding"
|
||||||
|
|
||||||
hidden_states = self.model(input_ids, positions, forward_batch, input_embeds)
|
hidden_states = self.model(input_ids, positions, forward_batch, input_embeds)
|
||||||
logits = self.score(hidden_states)
|
return score_and_pool(
|
||||||
pooled_logits = self.pooler(logits, forward_batch).embeddings
|
self.score, self.pooler, hidden_states, forward_batch, input_ids
|
||||||
|
)
|
||||||
return EmbeddingPoolerOutput(pooled_logits)
|
|
||||||
|
|
||||||
def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]):
|
def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]):
|
||||||
stacked_params_mapping = [
|
stacked_params_mapping = [
|
||||||
|
|||||||
@@ -0,0 +1,351 @@
|
|||||||
|
"""Tests for Scoring API with SequenceClassification models.
|
||||||
|
|
||||||
|
Covers both single-item and multi-item scoring (MIS) for classification
|
||||||
|
models. Uses json_model_override_args to override a regular Qwen3 model's
|
||||||
|
architecture to Qwen3ForSequenceClassification so we can validate the
|
||||||
|
scoring pipeline without needing a dedicated classification checkpoint
|
||||||
|
(the score head gets randomly initialised, which is fine for shape /
|
||||||
|
pipeline correctness).
|
||||||
|
"""
|
||||||
|
|
||||||
|
import json
|
||||||
|
import unittest
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from sglang.srt.entrypoints.engine import Engine
|
||||||
|
from sglang.test.ci.ci_register import register_cuda_ci
|
||||||
|
from sglang.test.test_utils import CustomTestCase
|
||||||
|
|
||||||
|
register_cuda_ci(est_time=120, suite="stage-b-test-1-gpu-small")
|
||||||
|
|
||||||
|
# A lightweight Qwen3 checkpoint whose backbone weights load cleanly into
|
||||||
|
# Qwen3ForSequenceClassification (the classification head is random).
|
||||||
|
TEST_BASE_MODEL = "Qwen/Qwen3-0.6B"
|
||||||
|
# <|endoftext|> for Qwen3 tokenizer — used as MIS delimiter
|
||||||
|
QWEN3_ENDOFTEXT_TOKEN_ID = 151643
|
||||||
|
|
||||||
|
|
||||||
|
class TestScoreClassification(CustomTestCase):
|
||||||
|
"""Single-item scoring with a SequenceClassification model."""
|
||||||
|
|
||||||
|
NUM_LABELS = 2
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def setUpClass(cls):
|
||||||
|
override_args = json.dumps(
|
||||||
|
{
|
||||||
|
"architectures": ["Qwen3ForSequenceClassification"],
|
||||||
|
"num_labels": cls.NUM_LABELS,
|
||||||
|
}
|
||||||
|
)
|
||||||
|
cls.engine = Engine(
|
||||||
|
model_path=TEST_BASE_MODEL,
|
||||||
|
disable_radix_cache=True,
|
||||||
|
json_model_override_args=override_args,
|
||||||
|
mem_fraction_static=0.15,
|
||||||
|
)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def tearDownClass(cls):
|
||||||
|
if hasattr(cls, "engine") and cls.engine is not None:
|
||||||
|
cls.engine.shutdown()
|
||||||
|
torch.cuda.empty_cache()
|
||||||
|
|
||||||
|
def test_basic_single_item(self):
|
||||||
|
"""Each item gets a score vector of length num_labels."""
|
||||||
|
scores = self.engine.score(
|
||||||
|
query="Rate each option:",
|
||||||
|
items=["Option A", "Option B"],
|
||||||
|
apply_softmax=True,
|
||||||
|
).scores
|
||||||
|
|
||||||
|
self.assertEqual(len(scores), 2)
|
||||||
|
for i, score_list in enumerate(scores):
|
||||||
|
self.assertEqual(len(score_list), self.NUM_LABELS)
|
||||||
|
self.assertAlmostEqual(
|
||||||
|
sum(score_list),
|
||||||
|
1.0,
|
||||||
|
places=5,
|
||||||
|
msg=f"Softmax scores for item {i} should sum to 1",
|
||||||
|
)
|
||||||
|
for val in score_list:
|
||||||
|
self.assertGreaterEqual(val, 0.0)
|
||||||
|
self.assertLessEqual(val, 1.0)
|
||||||
|
|
||||||
|
def test_single_item_edge_case(self):
|
||||||
|
"""Single item in the list."""
|
||||||
|
scores = self.engine.score(
|
||||||
|
query="Evaluate:",
|
||||||
|
items=["Only item"],
|
||||||
|
apply_softmax=True,
|
||||||
|
).scores
|
||||||
|
|
||||||
|
self.assertEqual(len(scores), 1)
|
||||||
|
self.assertEqual(len(scores[0]), self.NUM_LABELS)
|
||||||
|
self.assertAlmostEqual(sum(scores[0]), 1.0, places=5)
|
||||||
|
|
||||||
|
def test_raw_logits_without_softmax(self):
|
||||||
|
"""Without softmax, returns raw logits (no probability constraints)."""
|
||||||
|
scores = self.engine.score(
|
||||||
|
query="Evaluate:",
|
||||||
|
items=["Alpha", "Beta"],
|
||||||
|
apply_softmax=False,
|
||||||
|
).scores
|
||||||
|
|
||||||
|
self.assertEqual(len(scores), 2)
|
||||||
|
for score_list in scores:
|
||||||
|
self.assertEqual(len(score_list), self.NUM_LABELS)
|
||||||
|
for val in score_list:
|
||||||
|
self.assertTrue(
|
||||||
|
isinstance(val, (int, float)),
|
||||||
|
f"Expected numeric score, got {type(val)}",
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_deterministic(self):
|
||||||
|
"""Identical inputs yield near-identical scores (fp16 non-determinism allowed)."""
|
||||||
|
kwargs = dict(query="Evaluate:", items=["alpha", "beta", "gamma"])
|
||||||
|
|
||||||
|
scores1 = self.engine.score(**kwargs).scores
|
||||||
|
scores2 = self.engine.score(**kwargs).scores
|
||||||
|
|
||||||
|
self.assertEqual(len(scores1), len(scores2))
|
||||||
|
for s1, s2 in zip(scores1, scores2):
|
||||||
|
for v1, v2 in zip(s1, s2):
|
||||||
|
self.assertAlmostEqual(v1, v2, places=1)
|
||||||
|
|
||||||
|
def test_tokenized_inputs(self):
|
||||||
|
"""Pre-tokenized query and items work the same as text inputs."""
|
||||||
|
from transformers import AutoTokenizer
|
||||||
|
|
||||||
|
tokenizer = AutoTokenizer.from_pretrained(TEST_BASE_MODEL)
|
||||||
|
query_text = "Rate this:"
|
||||||
|
items_text = ["Good", "Bad"]
|
||||||
|
|
||||||
|
text_scores = self.engine.score(
|
||||||
|
query=query_text,
|
||||||
|
items=items_text,
|
||||||
|
apply_softmax=True,
|
||||||
|
).scores
|
||||||
|
|
||||||
|
query_ids = tokenizer.encode(query_text)
|
||||||
|
items_ids = [tokenizer.encode(item) for item in items_text]
|
||||||
|
token_scores = self.engine.score(
|
||||||
|
query=query_ids,
|
||||||
|
items=items_ids,
|
||||||
|
apply_softmax=True,
|
||||||
|
).scores
|
||||||
|
|
||||||
|
self.assertEqual(len(text_scores), len(token_scores))
|
||||||
|
for txt_s, tok_s in zip(text_scores, token_scores):
|
||||||
|
for t, k in zip(txt_s, tok_s):
|
||||||
|
self.assertAlmostEqual(t, k, places=4)
|
||||||
|
|
||||||
|
def test_label_token_ids_ignored(self):
|
||||||
|
"""SequenceClassification models ignore label_token_ids (no crash)."""
|
||||||
|
scores = self.engine.score(
|
||||||
|
query="Evaluate:",
|
||||||
|
items=["Test item"],
|
||||||
|
label_token_ids=[1, 2, 3],
|
||||||
|
apply_softmax=True,
|
||||||
|
).scores
|
||||||
|
|
||||||
|
self.assertEqual(len(scores), 1)
|
||||||
|
self.assertEqual(len(scores[0]), self.NUM_LABELS)
|
||||||
|
|
||||||
|
|
||||||
|
class TestScoreClassificationMIS(CustomTestCase):
|
||||||
|
"""Multi-item scoring (MIS) with a SequenceClassification model.
|
||||||
|
|
||||||
|
MIS packs all items into one sequence separated by a delimiter token.
|
||||||
|
The score_and_pool function extracts per-item scores at delimiter
|
||||||
|
positions.
|
||||||
|
"""
|
||||||
|
|
||||||
|
NUM_LABELS = 2
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def setUpClass(cls):
|
||||||
|
override_args = json.dumps(
|
||||||
|
{
|
||||||
|
"architectures": ["Qwen3ForSequenceClassification"],
|
||||||
|
"num_labels": cls.NUM_LABELS,
|
||||||
|
}
|
||||||
|
)
|
||||||
|
cls.engine = Engine(
|
||||||
|
model_path=TEST_BASE_MODEL,
|
||||||
|
disable_radix_cache=True,
|
||||||
|
chunked_prefill_size=-1,
|
||||||
|
multi_item_scoring_delimiter=QWEN3_ENDOFTEXT_TOKEN_ID,
|
||||||
|
json_model_override_args=override_args,
|
||||||
|
mem_fraction_static=0.15,
|
||||||
|
)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def tearDownClass(cls):
|
||||||
|
if hasattr(cls, "engine") and cls.engine is not None:
|
||||||
|
cls.engine.shutdown()
|
||||||
|
torch.cuda.empty_cache()
|
||||||
|
|
||||||
|
def test_mis_basic(self):
|
||||||
|
"""MIS produces one score vector per item."""
|
||||||
|
items = ["Option A", "Option B", "Option C"]
|
||||||
|
scores = self.engine.score(
|
||||||
|
query="Rate each option:",
|
||||||
|
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,
|
||||||
|
f"Item {i} should have {self.NUM_LABELS} scores",
|
||||||
|
)
|
||||||
|
self.assertAlmostEqual(
|
||||||
|
sum(score_list),
|
||||||
|
1.0,
|
||||||
|
places=5,
|
||||||
|
msg=f"Scores for item {i} should sum to 1",
|
||||||
|
)
|
||||||
|
for val in score_list:
|
||||||
|
self.assertGreaterEqual(val, 0.0)
|
||||||
|
self.assertLessEqual(val, 1.0)
|
||||||
|
|
||||||
|
def test_mis_many_items(self):
|
||||||
|
"""Stress test: 10 items."""
|
||||||
|
items = [f"Item {i}" for i in range(10)]
|
||||||
|
scores = self.engine.score(
|
||||||
|
query="Classify each:",
|
||||||
|
items=items,
|
||||||
|
apply_softmax=True,
|
||||||
|
).scores
|
||||||
|
|
||||||
|
self.assertEqual(len(scores), len(items))
|
||||||
|
for score_list in scores:
|
||||||
|
self.assertEqual(len(score_list), self.NUM_LABELS)
|
||||||
|
self.assertAlmostEqual(sum(score_list), 1.0, places=5)
|
||||||
|
|
||||||
|
def test_mis_single_item(self):
|
||||||
|
"""Edge case: single item through MIS path."""
|
||||||
|
scores = self.engine.score(
|
||||||
|
query="Evaluate:",
|
||||||
|
items=["Single item"],
|
||||||
|
apply_softmax=True,
|
||||||
|
).scores
|
||||||
|
|
||||||
|
self.assertEqual(len(scores), 1)
|
||||||
|
self.assertEqual(len(scores[0]), self.NUM_LABELS)
|
||||||
|
self.assertAlmostEqual(sum(scores[0]), 1.0, places=5)
|
||||||
|
|
||||||
|
def test_items_produce_distinct_scores(self):
|
||||||
|
"""Different items must produce different score vectors.
|
||||||
|
|
||||||
|
Even with a randomly initialised classification head, different
|
||||||
|
item texts produce different hidden states, so scores should
|
||||||
|
differ. This catches bugs where all delimiter tokens share the
|
||||||
|
same pooled representation.
|
||||||
|
"""
|
||||||
|
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))
|
||||||
|
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 likely broken. Scores: {scores[0]}",
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_deterministic(self):
|
||||||
|
"""Identical MIS requests return identical scores."""
|
||||||
|
kwargs = dict(
|
||||||
|
query="Evaluate:",
|
||||||
|
items=["alpha", "beta", "gamma"],
|
||||||
|
)
|
||||||
|
scores1 = self.engine.score(**kwargs).scores
|
||||||
|
scores2 = self.engine.score(**kwargs).scores
|
||||||
|
|
||||||
|
self.assertEqual(scores1, scores2)
|
||||||
|
|
||||||
|
def test_softmax_valid(self):
|
||||||
|
"""With softmax, each item's scores form a valid probability distribution."""
|
||||||
|
items = ["Option A", "Option B", "Option C"]
|
||||||
|
scores = self.engine.score(
|
||||||
|
query="Rate each option:",
|
||||||
|
items=items,
|
||||||
|
apply_softmax=True,
|
||||||
|
).scores
|
||||||
|
|
||||||
|
for i, score_list in enumerate(scores):
|
||||||
|
self.assertEqual(len(score_list), self.NUM_LABELS)
|
||||||
|
for val in score_list:
|
||||||
|
self.assertGreaterEqual(val, 0.0)
|
||||||
|
self.assertLessEqual(val, 1.0)
|
||||||
|
self.assertAlmostEqual(
|
||||||
|
sum(score_list),
|
||||||
|
1.0,
|
||||||
|
places=6,
|
||||||
|
msg=f"Softmax scores for item {i} don't sum to 1: {sum(score_list)}",
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class TestScoreClassificationMISAdvanced(CustomTestCase):
|
||||||
|
"""Advanced MIS tests with more labels to stress tensor shape handling."""
|
||||||
|
|
||||||
|
NUM_LABELS = 12
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def setUpClass(cls):
|
||||||
|
override_args = json.dumps(
|
||||||
|
{
|
||||||
|
"architectures": ["Qwen3ForSequenceClassification"],
|
||||||
|
"num_labels": cls.NUM_LABELS,
|
||||||
|
}
|
||||||
|
)
|
||||||
|
cls.engine = Engine(
|
||||||
|
model_path=TEST_BASE_MODEL,
|
||||||
|
disable_radix_cache=True,
|
||||||
|
chunked_prefill_size=-1,
|
||||||
|
multi_item_scoring_delimiter=QWEN3_ENDOFTEXT_TOKEN_ID,
|
||||||
|
json_model_override_args=override_args,
|
||||||
|
mem_fraction_static=0.15,
|
||||||
|
)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def tearDownClass(cls):
|
||||||
|
if hasattr(cls, "engine") and cls.engine is not None:
|
||||||
|
cls.engine.shutdown()
|
||||||
|
torch.cuda.empty_cache()
|
||||||
|
|
||||||
|
def test_many_labels_shape(self):
|
||||||
|
"""Verify correct shape with many labels (catches 2D tensor bugs)."""
|
||||||
|
items = [f"Item {i}" for i in range(5)]
|
||||||
|
scores = self.engine.score(
|
||||||
|
query="Classify:",
|
||||||
|
items=items,
|
||||||
|
apply_softmax=True,
|
||||||
|
).scores
|
||||||
|
|
||||||
|
self.assertEqual(len(scores), len(items))
|
||||||
|
for score_list in scores:
|
||||||
|
self.assertEqual(len(score_list), self.NUM_LABELS)
|
||||||
|
self.assertAlmostEqual(sum(score_list), 1.0, places=5)
|
||||||
|
|
||||||
|
def test_many_items_distinct(self):
|
||||||
|
"""15 items should not all produce identical scores."""
|
||||||
|
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))
|
||||||
|
unique_count = len({tuple(s) for s in scores})
|
||||||
|
self.assertGreater(unique_count, 1, "All 15 items returned identical scores")
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
@@ -0,0 +1,216 @@
|
|||||||
|
"""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.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import unittest
|
||||||
|
from types import SimpleNamespace
|
||||||
|
from unittest.mock import patch
|
||||||
|
|
||||||
|
import torch
|
||||||
|
import torch.nn as nn
|
||||||
|
|
||||||
|
from sglang.srt.layers.pooler import (
|
||||||
|
EmbeddingPoolerOutput,
|
||||||
|
Pooler,
|
||||||
|
PoolingType,
|
||||||
|
score_and_pool,
|
||||||
|
)
|
||||||
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
|
from sglang.test.test_utils import CustomTestCase
|
||||||
|
|
||||||
|
register_cpu_ci(est_time=5, suite="stage-a-test-cpu")
|
||||||
|
|
||||||
|
|
||||||
|
def _make_forward_batch(extend_seq_lens, is_prefill_only=False):
|
||||||
|
"""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,
|
||||||
|
dimensions=None,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
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."""
|
||||||
|
|
||||||
|
def setUp(self):
|
||||||
|
torch.manual_seed(42)
|
||||||
|
self.hidden_dim = 8
|
||||||
|
self.num_labels = 2
|
||||||
|
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)
|
||||||
|
|
||||||
|
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)
|
||||||
|
|
||||||
|
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,
|
||||||
|
]
|
||||||
|
)
|
||||||
|
hidden = torch.randn(len(input_ids), self.hidden_dim)
|
||||||
|
fb = _make_forward_batch(extend_seq_lens=[len(input_ids)], is_prefill_only=True)
|
||||||
|
|
||||||
|
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, (3, self.num_labels))
|
||||||
|
|
||||||
|
@patch("sglang.srt.layers.pooler.get_global_server_args")
|
||||||
|
def test_mis_batched_splits_per_request(self, mock_get_args):
|
||||||
|
"""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]
|
||||||
|
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
|
||||||
|
)
|
||||||
|
|
||||||
|
out = score_and_pool(self.score_head, self.pooler, hidden, fb, input_ids)
|
||||||
|
|
||||||
|
self.assertIsInstance(out.embeddings, list)
|
||||||
|
self.assertEqual(len(out.embeddings), 2)
|
||||||
|
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)
|
||||||
|
|
||||||
|
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)
|
||||||
|
|
||||||
|
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):
|
||||||
|
"""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])
|
||||||
|
hidden = (
|
||||||
|
torch.arange(len(input_ids))
|
||||||
|
.unsqueeze(1)
|
||||||
|
.float()
|
||||||
|
.expand(-1, self.hidden_dim)
|
||||||
|
.clone()
|
||||||
|
)
|
||||||
|
fb = _make_forward_batch(extend_seq_lens=[len(input_ids)], is_prefill_only=True)
|
||||||
|
|
||||||
|
identity_head = nn.Linear(self.hidden_dim, self.hidden_dim, bias=False)
|
||||||
|
nn.init.eye_(identity_head.weight)
|
||||||
|
|
||||||
|
out = score_and_pool(identity_head, self.pooler, hidden, fb, input_ids)
|
||||||
|
|
||||||
|
scores = out.embeddings[0]
|
||||||
|
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])
|
||||||
|
hidden = (
|
||||||
|
torch.arange(len(input_ids))
|
||||||
|
.unsqueeze(1)
|
||||||
|
.float()
|
||||||
|
.expand(-1, self.hidden_dim)
|
||||||
|
.clone()
|
||||||
|
)
|
||||||
|
fb = _make_forward_batch(extend_seq_lens=[len(input_ids)], is_prefill_only=True)
|
||||||
|
|
||||||
|
identity_head = nn.Linear(self.hidden_dim, self.hidden_dim, bias=False)
|
||||||
|
nn.init.eye_(identity_head.weight)
|
||||||
|
|
||||||
|
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)
|
||||||
|
|
||||||
|
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
|
||||||
|
torch.testing.assert_close(out.embeddings, expected)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
@@ -0,0 +1,442 @@
|
|||||||
|
import os
|
||||||
|
import signal
|
||||||
|
import subprocess
|
||||||
|
import tempfile
|
||||||
|
import time
|
||||||
|
import unittest
|
||||||
|
|
||||||
|
import requests
|
||||||
|
|
||||||
|
from sglang.test.test_utils import DEFAULT_SMALL_MODEL_NAME_FOR_TEST, CustomTestCase
|
||||||
|
|
||||||
|
|
||||||
|
class TestMultiItemScoringServer(CustomTestCase):
|
||||||
|
"""Test multi-item scoring functionality through the server API."""
|
||||||
|
|
||||||
|
def setUp(self):
|
||||||
|
"""Set up each test case."""
|
||||||
|
self.model_path = DEFAULT_SMALL_MODEL_NAME_FOR_TEST
|
||||||
|
self.port = 30001 # Use different port to avoid conflicts
|
||||||
|
self.host = "localhost"
|
||||||
|
self.base_url = f"http://{self.host}:{self.port}"
|
||||||
|
self.server_process = None
|
||||||
|
self.server_log_file = None
|
||||||
|
|
||||||
|
def tearDown(self):
|
||||||
|
"""Clean up after each test case."""
|
||||||
|
self.stop_server()
|
||||||
|
|
||||||
|
def start_server(self, multi_item_scoring_delimiter=None):
|
||||||
|
"""Start the SGLang server with multi-item scoring enabled."""
|
||||||
|
if self.server_process is not None:
|
||||||
|
self.stop_server()
|
||||||
|
|
||||||
|
# Create a temporary log file
|
||||||
|
self.server_log_file = tempfile.NamedTemporaryFile(mode="w+", delete=False)
|
||||||
|
|
||||||
|
cmd = [
|
||||||
|
"python",
|
||||||
|
"-m",
|
||||||
|
"sglang.launch_server",
|
||||||
|
"--model-path",
|
||||||
|
self.model_path,
|
||||||
|
"--port",
|
||||||
|
str(self.port),
|
||||||
|
"--host",
|
||||||
|
self.host,
|
||||||
|
"--chunked-prefill-size",
|
||||||
|
"-1",
|
||||||
|
"--dtype",
|
||||||
|
"float16",
|
||||||
|
"--max-prefill-tokens",
|
||||||
|
"30000",
|
||||||
|
"--mem-fraction-static",
|
||||||
|
"0.3",
|
||||||
|
"--disable-radix-cache",
|
||||||
|
"--disable-cuda-graph",
|
||||||
|
"--attention-backend",
|
||||||
|
"flashinfer",
|
||||||
|
]
|
||||||
|
|
||||||
|
if multi_item_scoring_delimiter is not None:
|
||||||
|
cmd.extend(
|
||||||
|
["--multi-item-scoring-delimiter", str(multi_item_scoring_delimiter)]
|
||||||
|
)
|
||||||
|
|
||||||
|
# Start server process
|
||||||
|
self.server_process = subprocess.Popen(
|
||||||
|
cmd,
|
||||||
|
stdout=self.server_log_file,
|
||||||
|
stderr=subprocess.STDOUT,
|
||||||
|
preexec_fn=os.setsid if hasattr(os, "setsid") else None,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Wait for server to start
|
||||||
|
max_wait_time = 60 # seconds
|
||||||
|
start_time = time.time()
|
||||||
|
|
||||||
|
while time.time() - start_time < max_wait_time:
|
||||||
|
try:
|
||||||
|
response = requests.get(f"{self.base_url}/get_model_info", timeout=5)
|
||||||
|
if response.status_code == 200:
|
||||||
|
return True
|
||||||
|
except requests.exceptions.RequestException:
|
||||||
|
pass
|
||||||
|
time.sleep(1)
|
||||||
|
|
||||||
|
# If we get here, server didn't start properly
|
||||||
|
self.stop_server()
|
||||||
|
raise RuntimeError("Failed to start SGLang server within timeout period")
|
||||||
|
|
||||||
|
def stop_server(self):
|
||||||
|
"""Stop the SGLang server."""
|
||||||
|
if self.server_process is not None:
|
||||||
|
try:
|
||||||
|
# Kill the process group
|
||||||
|
if hasattr(os, "killpg"):
|
||||||
|
os.killpg(os.getpgid(self.server_process.pid), signal.SIGTERM)
|
||||||
|
else:
|
||||||
|
self.server_process.terminate()
|
||||||
|
|
||||||
|
# Wait for graceful shutdown
|
||||||
|
try:
|
||||||
|
self.server_process.wait(timeout=10)
|
||||||
|
except subprocess.TimeoutExpired:
|
||||||
|
# Force kill if graceful shutdown fails
|
||||||
|
if hasattr(os, "killpg"):
|
||||||
|
os.killpg(os.getpgid(self.server_process.pid), signal.SIGKILL)
|
||||||
|
else:
|
||||||
|
self.server_process.kill()
|
||||||
|
self.server_process.wait()
|
||||||
|
except (ProcessLookupError, OSError):
|
||||||
|
# Process already terminated
|
||||||
|
pass
|
||||||
|
finally:
|
||||||
|
self.server_process = None
|
||||||
|
|
||||||
|
if self.server_log_file is not None:
|
||||||
|
try:
|
||||||
|
self.server_log_file.close()
|
||||||
|
os.unlink(self.server_log_file.name)
|
||||||
|
except (OSError, FileNotFoundError):
|
||||||
|
pass
|
||||||
|
finally:
|
||||||
|
self.server_log_file = None
|
||||||
|
|
||||||
|
def get_server_logs(self):
|
||||||
|
"""Get server logs for debugging."""
|
||||||
|
if self.server_log_file is not None:
|
||||||
|
try:
|
||||||
|
with open(self.server_log_file.name, "r") as f:
|
||||||
|
return f.read()
|
||||||
|
except (OSError, FileNotFoundError):
|
||||||
|
pass
|
||||||
|
return "No logs available"
|
||||||
|
|
||||||
|
def test_multi_item_scoring_server_basic(self):
|
||||||
|
"""Test basic multi-item scoring through server API."""
|
||||||
|
# Start server with multi-item scoring enabled
|
||||||
|
delimiter_token_id = 151655 # Example delimiter token ID
|
||||||
|
self.start_server(multi_item_scoring_delimiter=delimiter_token_id)
|
||||||
|
|
||||||
|
# Test data
|
||||||
|
query = "What is the capital of California? Answer Yes or No for each of the following options:"
|
||||||
|
items = ["Sacramento", "San Jose", "San Francisco"]
|
||||||
|
label_token_ids = [9454, 2753] # "Yes" and "No" tokens
|
||||||
|
|
||||||
|
# Make scoring request
|
||||||
|
payload = {
|
||||||
|
"query": query,
|
||||||
|
"items": items,
|
||||||
|
"label_token_ids": label_token_ids,
|
||||||
|
"model": self.model_path,
|
||||||
|
}
|
||||||
|
|
||||||
|
response = requests.post(
|
||||||
|
f"{self.base_url}/v1/score",
|
||||||
|
headers={"Content-Type": "application/json"},
|
||||||
|
json=payload,
|
||||||
|
timeout=30,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Check response
|
||||||
|
self.assertEqual(
|
||||||
|
response.status_code,
|
||||||
|
200,
|
||||||
|
f"Server returned {response.status_code}. Logs: {self.get_server_logs()}",
|
||||||
|
)
|
||||||
|
|
||||||
|
result = response.json()
|
||||||
|
|
||||||
|
# Verify response structure
|
||||||
|
self.assertIn("scores", result)
|
||||||
|
self.assertIn("model", result)
|
||||||
|
self.assertIn("object", result)
|
||||||
|
|
||||||
|
self.assertEqual(result["object"], "scoring")
|
||||||
|
self.assertEqual(result["model"], self.model_path)
|
||||||
|
|
||||||
|
# Verify scores
|
||||||
|
scores = result["scores"]
|
||||||
|
self.assertEqual(len(scores), len(items), "Should get one score list per item")
|
||||||
|
|
||||||
|
for i, score_list in enumerate(scores):
|
||||||
|
self.assertEqual(
|
||||||
|
len(score_list),
|
||||||
|
len(label_token_ids),
|
||||||
|
f"Item {i} should have {len(label_token_ids)} scores",
|
||||||
|
)
|
||||||
|
# Verify scores are probabilities (sum to 1)
|
||||||
|
self.assertAlmostEqual(
|
||||||
|
sum(score_list),
|
||||||
|
1.0,
|
||||||
|
places=6,
|
||||||
|
msg=f"Scores for item {i} should sum to 1",
|
||||||
|
)
|
||||||
|
# Verify all scores are non-negative
|
||||||
|
for j, score in enumerate(score_list):
|
||||||
|
self.assertGreaterEqual(
|
||||||
|
score, 0, f"Score {j} for item {i} should be non-negative"
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_multi_item_scoring_server_different_sizes(self):
|
||||||
|
"""Test multi-item scoring with different numbers of items through server."""
|
||||||
|
delimiter_token_id = 151655
|
||||||
|
self.start_server(multi_item_scoring_delimiter=delimiter_token_id)
|
||||||
|
|
||||||
|
query = "Rate each option:"
|
||||||
|
label_token_ids = [1, 2, 3, 4, 5]
|
||||||
|
|
||||||
|
test_cases = [
|
||||||
|
["Single item"],
|
||||||
|
["Item 1", "Item 2"],
|
||||||
|
["A", "B", "C", "D"],
|
||||||
|
["X", "Y", "Z", "W", "V", "U"],
|
||||||
|
]
|
||||||
|
|
||||||
|
for items in test_cases:
|
||||||
|
with self.subTest(items=items):
|
||||||
|
payload = {
|
||||||
|
"query": query,
|
||||||
|
"items": items,
|
||||||
|
"label_token_ids": label_token_ids,
|
||||||
|
"model": self.model_path,
|
||||||
|
}
|
||||||
|
|
||||||
|
response = requests.post(
|
||||||
|
f"{self.base_url}/v1/score",
|
||||||
|
headers={"Content-Type": "application/json"},
|
||||||
|
json=payload,
|
||||||
|
timeout=30,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(
|
||||||
|
response.status_code,
|
||||||
|
200,
|
||||||
|
f"Failed for items {items}. Logs: {self.get_server_logs()}",
|
||||||
|
)
|
||||||
|
|
||||||
|
result = response.json()
|
||||||
|
scores = result["scores"]
|
||||||
|
|
||||||
|
self.assertEqual(
|
||||||
|
len(scores), len(items), f"Should get {len(items)} score lists"
|
||||||
|
)
|
||||||
|
|
||||||
|
for i, score_list in enumerate(scores):
|
||||||
|
self.assertEqual(len(score_list), len(label_token_ids))
|
||||||
|
self.assertAlmostEqual(sum(score_list), 1.0, places=6)
|
||||||
|
|
||||||
|
def test_multi_item_scoring_server_empty_items(self):
|
||||||
|
"""Test multi-item scoring with empty items list through server."""
|
||||||
|
delimiter_token_id = 151655
|
||||||
|
self.start_server(multi_item_scoring_delimiter=delimiter_token_id)
|
||||||
|
|
||||||
|
payload = {
|
||||||
|
"query": "Test query",
|
||||||
|
"items": [],
|
||||||
|
"label_token_ids": [1, 2],
|
||||||
|
"model": self.model_path,
|
||||||
|
}
|
||||||
|
|
||||||
|
response = requests.post(
|
||||||
|
f"{self.base_url}/v1/score",
|
||||||
|
headers={"Content-Type": "application/json"},
|
||||||
|
json=payload,
|
||||||
|
timeout=30,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, 200)
|
||||||
|
result = response.json()
|
||||||
|
self.assertEqual(
|
||||||
|
len(result["scores"]), 0, "Should return empty list for empty items"
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_multi_item_scoring_server_without_delimiter(self):
|
||||||
|
"""Test that server works without multi-item scoring delimiter."""
|
||||||
|
# Start server without multi-item scoring delimiter
|
||||||
|
self.start_server(multi_item_scoring_delimiter=None)
|
||||||
|
|
||||||
|
payload = {
|
||||||
|
"query": "Test query",
|
||||||
|
"items": ["Item 1", "Item 2"],
|
||||||
|
"label_token_ids": [1, 2],
|
||||||
|
"model": self.model_path,
|
||||||
|
}
|
||||||
|
|
||||||
|
response = requests.post(
|
||||||
|
f"{self.base_url}/v1/score",
|
||||||
|
headers={"Content-Type": "application/json"},
|
||||||
|
json=payload,
|
||||||
|
timeout=30,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Should still work (falls back to regular scoring)
|
||||||
|
self.assertEqual(response.status_code, 200)
|
||||||
|
result = response.json()
|
||||||
|
self.assertIn("scores", result)
|
||||||
|
|
||||||
|
def test_multi_item_scoring_server_error_handling(self):
|
||||||
|
"""Test error handling in multi-item scoring server API."""
|
||||||
|
delimiter_token_id = 151655
|
||||||
|
self.start_server(multi_item_scoring_delimiter=delimiter_token_id)
|
||||||
|
|
||||||
|
# Test with invalid payload
|
||||||
|
invalid_payloads = [
|
||||||
|
{
|
||||||
|
"query": "Test",
|
||||||
|
"items": "not a list",
|
||||||
|
"label_token_ids": [1, 2],
|
||||||
|
"model": self.model_path,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"query": "Test",
|
||||||
|
"items": ["Item 1"],
|
||||||
|
"label_token_ids": "not a list",
|
||||||
|
"model": self.model_path,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"query": "Test",
|
||||||
|
"items": ["Item 1"],
|
||||||
|
"label_token_ids": [1, 2],
|
||||||
|
}, # Missing model
|
||||||
|
{
|
||||||
|
"items": ["Item 1"],
|
||||||
|
"label_token_ids": [1, 2],
|
||||||
|
"model": self.model_path,
|
||||||
|
}, # Missing query
|
||||||
|
]
|
||||||
|
|
||||||
|
for i, payload in enumerate(invalid_payloads):
|
||||||
|
with self.subTest(payload=i):
|
||||||
|
response = requests.post(
|
||||||
|
f"{self.base_url}/v1/score",
|
||||||
|
headers={"Content-Type": "application/json"},
|
||||||
|
json=payload,
|
||||||
|
timeout=30,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Should return error status
|
||||||
|
self.assertGreaterEqual(
|
||||||
|
response.status_code,
|
||||||
|
400,
|
||||||
|
f"Should return error for invalid payload {i}",
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_multi_item_scoring_server_consistency(self):
|
||||||
|
"""Test that multi-item scoring gives consistent results through server."""
|
||||||
|
delimiter_token_id = 151655
|
||||||
|
self.start_server(multi_item_scoring_delimiter=delimiter_token_id)
|
||||||
|
|
||||||
|
query = "Choose the best option:"
|
||||||
|
items = ["Option A", "Option B", "Option C"]
|
||||||
|
label_token_ids = [1, 2, 3]
|
||||||
|
|
||||||
|
payload = {
|
||||||
|
"query": query,
|
||||||
|
"items": items,
|
||||||
|
"label_token_ids": label_token_ids,
|
||||||
|
"model": self.model_path,
|
||||||
|
}
|
||||||
|
|
||||||
|
# Run the same test multiple times
|
||||||
|
scores1 = None
|
||||||
|
scores2 = None
|
||||||
|
|
||||||
|
for attempt in range(2):
|
||||||
|
response = requests.post(
|
||||||
|
f"{self.base_url}/v1/score",
|
||||||
|
headers={"Content-Type": "application/json"},
|
||||||
|
json=payload,
|
||||||
|
timeout=30,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(
|
||||||
|
response.status_code,
|
||||||
|
200,
|
||||||
|
f"Attempt {attempt + 1} failed. Logs: {self.get_server_logs()}",
|
||||||
|
)
|
||||||
|
|
||||||
|
result = response.json()
|
||||||
|
scores = result["scores"]
|
||||||
|
|
||||||
|
if attempt == 0:
|
||||||
|
scores1 = scores
|
||||||
|
else:
|
||||||
|
scores2 = scores
|
||||||
|
|
||||||
|
# Results should be identical (deterministic)
|
||||||
|
self.assertEqual(len(scores1), len(scores2), "Should get same number of items")
|
||||||
|
for i, (s1, s2) in enumerate(zip(scores1, scores2)):
|
||||||
|
self.assertEqual(
|
||||||
|
len(s1), len(s2), f"Item {i} should have same number of scores"
|
||||||
|
)
|
||||||
|
for j, (score1, score2) in enumerate(zip(s1, s2)):
|
||||||
|
self.assertAlmostEqual(
|
||||||
|
score1,
|
||||||
|
score2,
|
||||||
|
places=6,
|
||||||
|
msg=f"Score {j} for item {i} should be identical",
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_multi_item_scoring_server_large_batch(self):
|
||||||
|
"""Test multi-item scoring with large batch through server."""
|
||||||
|
delimiter_token_id = 151655
|
||||||
|
self.start_server(multi_item_scoring_delimiter=delimiter_token_id)
|
||||||
|
|
||||||
|
query = "Classify each item:"
|
||||||
|
items = [f"Item {i}" for i in range(10)] # 10 items (smaller for test)
|
||||||
|
label_token_ids = [1, 2, 3]
|
||||||
|
|
||||||
|
payload = {
|
||||||
|
"query": query,
|
||||||
|
"items": items,
|
||||||
|
"label_token_ids": label_token_ids,
|
||||||
|
"model": self.model_path,
|
||||||
|
}
|
||||||
|
|
||||||
|
response = requests.post(
|
||||||
|
f"{self.base_url}/v1/score",
|
||||||
|
headers={"Content-Type": "application/json"},
|
||||||
|
json=payload,
|
||||||
|
timeout=60, # Longer timeout for large batch
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(
|
||||||
|
response.status_code,
|
||||||
|
200,
|
||||||
|
f"Large batch failed. Logs: {self.get_server_logs()}",
|
||||||
|
)
|
||||||
|
|
||||||
|
result = response.json()
|
||||||
|
scores = result["scores"]
|
||||||
|
|
||||||
|
self.assertEqual(len(scores), len(items), "Should handle large batches")
|
||||||
|
|
||||||
|
for i, score_list in enumerate(scores):
|
||||||
|
self.assertEqual(len(score_list), len(label_token_ids))
|
||||||
|
self.assertAlmostEqual(sum(score_list), 1.0, places=6)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
Reference in New Issue
Block a user