Skip to content
Merged
Show file tree
Hide file tree
Changes from 1 commit
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
306 changes: 4 additions & 302 deletions python/sglang/srt/managers/tokenizer_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -33,7 +33,6 @@

import fastapi
import orjson
import torch
import uvloop
import zmq
import zmq.asyncio
Expand Down Expand Up @@ -78,6 +77,9 @@
from sglang.srt.managers.scheduler import is_health_check_generate_req
from sglang.srt.managers.scheduler_input_blocker import input_blocker_guard_region
from sglang.srt.managers.tokenizer_communicator_mixin import TokenizerCommunicatorMixin
from sglang.srt.managers.tokenizer_manager_multiitem_mixin import (
TokenizerManagerMultiItemMixin,
)
from sglang.srt.metrics.collector import TokenizerMetricsCollector
from sglang.srt.sampling.sampling_params import SamplingParams
from sglang.srt.server_args import (
Expand Down Expand Up @@ -171,7 +173,7 @@ class ReqState:
output_token_ids_logprobs_idx: List = dataclasses.field(default_factory=list)


class TokenizerManager(TokenizerCommunicatorMixin):
class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixin):
"""TokenizerManager is a process that tokenizes the text."""

def __init__(
Expand Down Expand Up @@ -2088,50 +2090,6 @@ def _handle_update_weights_from_disk_req_output(self, recv_obj):
if len(self.model_update_tmp) == self.server_args.dp_size:
self.model_update_result.set_result(self.model_update_tmp)

def _initialize_multi_item_delimiter_text(self):
"""Initialize multi-item delimiter text from token ID after tokenizer is loaded."""
if (
hasattr(self.server_args, "multi_item_scoring_delimiter")
and self.server_args.multi_item_scoring_delimiter is not None
and self.tokenizer is not None
):
try:
self.multi_item_delimiter_text = self.tokenizer.decode(
[self.server_args.multi_item_scoring_delimiter],
skip_special_tokens=False,
)
except Exception as e:
logger.warning(
f"Failed to decode delimiter token {self.server_args.multi_item_scoring_delimiter}: {e}"
)
self.multi_item_delimiter_text = None

def _build_multi_item_token_sequence(
self, query: List[int], items: List[List[int]], delimiter_token_id: int
) -> List[int]:
"""
Build a single token sequence for multi-item scoring.
Format: query<delimiter>item1<delimiter>item2<delimiter>item3<delimiter>

Args:
query: Query token IDs
items: List of item token ID sequences
delimiter_token_id: Token ID to use as delimiter

Returns:
Combined token sequence
"""
combined_sequence = query[:] # Start with query

for item in items:
combined_sequence.append(delimiter_token_id) # Add delimiter
combined_sequence.extend(item) # Add item tokens

# Add final delimiter after the last item for logprob extraction
combined_sequence.append(delimiter_token_id)

return combined_sequence

def _extract_logprobs_for_tokens(
self, logprobs_data: List, label_token_ids: List[int]
) -> Dict[int, float]:
Expand All @@ -2152,262 +2110,6 @@ def _extract_logprobs_for_tokens(
logprobs[token_id] = logprob
return logprobs

def _convert_logprobs_to_scores(
self,
logprobs: Dict[int, float],
label_token_ids: List[int],
apply_softmax: bool,
) -> List[float]:
"""
Convert logprobs dictionary to ordered score list.

Args:
logprobs: Dictionary mapping token_id to logprob
label_token_ids: Token IDs in desired order
apply_softmax: Whether to apply softmax normalization

Returns:
List of scores in the same order as label_token_ids
"""
score_list = [
logprobs.get(token_id, float("-inf")) for token_id in label_token_ids
]

if apply_softmax:
score_list = torch.softmax(torch.tensor(score_list), dim=0).tolist()
else:
# Convert logprobs to probabilities if not using softmax
score_list = [
math.exp(x) if x != float("-inf") else 0.0 for x in score_list
]

return score_list

def _process_multi_item_scoring_results(
self,
results: Any,
items: List,
label_token_ids: List[int],
apply_softmax: bool,
batch_request=None,
) -> List[List[float]]:
"""
Process results from multi-item scoring request.
Extracts logprobs at delimiter positions from input_token_ids_logprobs.

Args:
results: Results from generate_request
items: List of items being scored
label_token_ids: Token IDs to extract scores for
apply_softmax: Whether to apply softmax normalization
batch_request: The original batch request containing input sequence

Returns:
List of score lists, one for each item
"""
single_result = results[0] if isinstance(results, list) else results

# For multi-item scoring, logprobs are in input_token_ids_logprobs
input_logprobs = single_result["meta_info"].get("input_token_ids_logprobs", [])

if not input_logprobs:
raise RuntimeError(
f"input_token_ids_logprobs is empty for multi-item scoring request {single_result['meta_info'].get('id', '<unknown>')}. "
"This indicates token_ids_logprobs were not computed properly for Mutil Item Scoring."
)

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:
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: {single_result['meta_info'].get('id', '<unknown>')}"
)

# 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)

return scores

def _process_single_item_scoring_results(
self, results: Any, label_token_ids: List[int], apply_softmax: bool
) -> List[List[float]]:
"""
Process results from single-item scoring request.
Single-item scoring results are stored in output_token_ids_logprobs.

Args:
results: Results from generate_request
label_token_ids: Token IDs to extract scores for
apply_softmax: Whether to apply softmax normalization

Returns:
List of score lists, one for each result
"""
scores = []

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", [])

if not output_logprobs or len(output_logprobs) == 0:
raise RuntimeError(
f"output_logprobs is empty for request {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)

return scores

async def score_request(
self,
query: Optional[Union[str, List[int]]] = None,
items: Optional[Union[str, List[str], List[List[int]]]] = None,
label_token_ids: Optional[List[int]] = None,
apply_softmax: bool = False,
item_first: bool = False,
request: Optional[Any] = None,
) -> List[List[float]]:
"""
Score the probability of specified token IDs appearing after the given (query + item) pair.

This method supports two scoring approaches:
1. Single-Item scoring (default): Process each query+item pair independently
2. Multi-Item scoring: When multi_item_scoring_delimiter is set, combine query and
multiple items into a single sequence using delimiter for efficient processing.
Note: item_first parameter is ignored in multi-item scoring mode since it uses
a fixed format: query<delimiter>item1<delimiter>item2<delimiter>item3<delimiter>

Multi-item scoring works with both text and pre-tokenized inputs:
- 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>

Args:
query: The query text or pre-tokenized query token IDs
items: The item text(s) or pre-tokenized item token IDs
label_token_ids: List of token IDs to compute probabilities for
apply_softmax: Whether to normalize probabilities using softmax
item_first: If True, prepend items to query. Ignored for multi-item scoring.
request: Optional FastAPI request object

Returns:
List of lists containing probabilities for each item and each label token
"""
if label_token_ids is None:
raise ValueError("label_token_ids must be provided")

if self.tokenizer is not None:
vocab_size = self.tokenizer.vocab_size
for token_id in label_token_ids:
if token_id >= vocab_size:
raise ValueError(
f"Token ID {token_id} is out of vocabulary (vocab size: {vocab_size})"
)

# Check if multi-item scoring is enabled by presence of delimiter
use_multi_item_scoring = (
self.server_args.multi_item_scoring_delimiter is not None
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},
)

# Handle string or tokenized query/items
if isinstance(query, str) and (
isinstance(items, str)
or (isinstance(items, list) and (not items or isinstance(items[0], str)))
):
# Both query and items are text
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]
else:
# Single-item scoring: create separate prompts for each item
if item_first:
prompts = [f"{item}{query}" for item in items_list]
else:
prompts = [f"{query}{item}" for item in items_list]
batch_request.text = prompts

elif (
isinstance(query, list)
and isinstance(items, list)
and items
and isinstance(items[0], list)
):
# Both query and items are token IDs
if use_multi_item_scoring:
# Multi-item scoring: concatenate with delimiter token ID
# Format: query<delimiter_token_id>item1<delimiter_token_id>item2<delimiter_token_id>item3<delimiter_token_id>
delimiter_token_id = self.server_args.multi_item_scoring_delimiter
combined_input_ids = self._build_multi_item_token_sequence(
query, items, delimiter_token_id
)
batch_request.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]
else:
input_ids_list = [query + item for item in items]
batch_request.input_ids = input_ids_list
else:
raise ValueError(
"Invalid combination of query/items types for score_request."
)

results = await self.generate_request(batch_request, request).__anext__()

if use_multi_item_scoring:
# Multi-item scoring: extract scores from input_token_ids_logprobs
return self._process_multi_item_scoring_results(
results, items, label_token_ids, apply_softmax, batch_request
)
else:
# Single-item scoring: process each result separately
return self._process_single_item_scoring_results(
results, label_token_ids, apply_softmax
)

async def watch_load_thread(self):
# Only for dp_controller when dp_size > 1
if (
Expand Down
Loading
Loading