[Score API] Add Multi-Item Scoring with pre-computed delimiter indices (#22544)

Co-authored-by: Chanh Nguyen <chanhnguyen@gmail.com>
Co-authored-by: Sundara Raman Ramachandran <sundar24295@gmail.com>
This commit is contained in:
jsheng_Linkedin
2026-04-20 22:50:40 -07:00
committed by GitHub
co-authored by Chanh Nguyen Sundara Raman Ramachandran
parent cfd49e233c
commit a8e3a534a4
17 changed files with 1071 additions and 431 deletions
@@ -126,10 +126,8 @@ class FlashInferAttnBackend(AttentionBackend):
self.prefill_backend = "fa2" self.prefill_backend = "fa2"
self.decode_backend = "fa2" self.decode_backend = "fa2"
# Store multi-item scoring delimiter for efficient access # Store multi-item scoring flag for efficient access
self.multi_item_scoring_delimiter = ( self.enable_mis = model_runner.server_args.enable_mis
model_runner.server_args.multi_item_scoring_delimiter
)
# FIXME: remove dllm workarounds from flashinfer # FIXME: remove dllm workarounds from flashinfer
self.dllm_config = DllmConfig.from_server_args(model_runner.server_args) self.dllm_config = DllmConfig.from_server_args(model_runner.server_args)
@@ -342,15 +340,18 @@ class FlashInferAttnBackend(AttentionBackend):
- max_item_len_ptr: [2, 3] (max lengths per sequence) - max_item_len_ptr: [2, 3] (max lengths per sequence)
""" """
delimiter = self.multi_item_scoring_delimiter if not self.enable_mis or forward_batch.forward_mode == ForwardMode.DECODE:
if delimiter is None or forward_batch.forward_mode == ForwardMode.DECODE:
return MultiItemScoringParams() return MultiItemScoringParams()
delimiter_mask = forward_batch.input_ids == delimiter precomputed_indices = forward_batch.multi_item_delimiter_indices
prefix_cache_lens = getattr(forward_batch, "extend_prefix_lens", None) if precomputed_indices is None:
extend_seq_lens = getattr(forward_batch, "extend_seq_lens", None) return MultiItemScoringParams()
prefix_cache_lens = getattr(forward_batch, "extend_prefix_lens_cpu", None)
extend_seq_lens = getattr(forward_batch, "extend_seq_lens_cpu", None)
prefix_len_ptr, token_pos_in_items_ptr = [], [] prefix_len_ptr, token_pos_in_items_ptr = [], []
token_pos_in_items_len = 0 token_pos_in_items_len = 0
device = forward_batch.input_ids.device
# If no extend_seq_lens, treat whole batch as one sequence # If no extend_seq_lens, treat whole batch as one sequence
if extend_seq_lens is None or len(extend_seq_lens) <= 1: if extend_seq_lens is None or len(extend_seq_lens) <= 1:
@@ -359,35 +360,44 @@ class FlashInferAttnBackend(AttentionBackend):
seq_start = 0 seq_start = 0
for i, seq_len in enumerate(extend_seq_lens): for i, seq_len in enumerate(extend_seq_lens):
seq_end = seq_start + seq_len seq_end = seq_start + seq_len
mask = delimiter_mask[seq_start:seq_end] delimiter_indices_cpu = precomputed_indices[i]
pos = forward_batch.positions[seq_start:seq_end] if len(delimiter_indices_cpu) == 0:
delimiter_indices = torch.nonzero(mask, as_tuple=True)[0] seq_start = seq_end
continue
if len(delimiter_indices) > 0: first_delim = delimiter_indices_cpu[0].item() # CPU .item(), no GPU sync
first_delim = delimiter_indices[0] delimiter_indices = delimiter_indices_cpu.to(device, non_blocking=True)
# Prefix length: store as scalar
prefix_len = first_delim + ( prefix_len = first_delim + (
prefix_cache_lens[i] if prefix_cache_lens is not None else 0 prefix_cache_lens[i] if prefix_cache_lens is not None else 0
) )
prefix_len_ptr.append( prefix_len_ptr.append(prefix_len)
prefix_len.item() if torch.is_tensor(prefix_len) else prefix_len
# Compute relative positions within items using searchsorted (no GPU sync).
# suffix_range = [0, 1, 2, 3, 4, ...]
# searchsorted = bucket index for each position
# last_delim = delimiter offset at start of current bucket
# pos_within_item = suffix_range - last_delim
suffix_len = seq_len - first_delim
relative_positions = delimiter_indices - first_delim
suffix_range = torch.arange(suffix_len, dtype=torch.int64, device=device)
bucket_idx = torch.searchsorted(
relative_positions, suffix_range, right=True
) )
last_delim = relative_positions[torch.clamp(bucket_idx - 1, min=0)]
pos_within_item = suffix_range - last_delim
# Compute relative positions within items after delimiters token_pos_in_items_ptr.append(pos_within_item.to(torch.uint16))
diff = pos[first_delim:] - torch.cummax(mask[first_delim:], 0)[1]
token_pos = (diff - pos[first_delim]).to(torch.uint16)
token_pos_in_items_ptr.append(token_pos)
# Update forward_batch positions in-place forward_batch.positions[seq_start + first_delim : seq_end] = (
pos[first_delim:] = diff - 1 prefix_len + pos_within_item - 1
forward_batch.positions[seq_start:seq_end] = pos )
seq_start = seq_end seq_start = seq_end
# Pad token_pos_in_items_ptr for batch processing # Pad token_pos_in_items_ptr for batch processing
if token_pos_in_items_ptr: if token_pos_in_items_ptr:
token_pos_in_items_len = max(t.numel() for t in token_pos_in_items_ptr) token_pos_in_items_len = max(t.numel() for t in token_pos_in_items_ptr)
device = forward_batch.input_ids.device
token_pos_in_items_ptr = [ token_pos_in_items_ptr = [
torch.cat( torch.cat(
[ [
@@ -405,8 +415,6 @@ class FlashInferAttnBackend(AttentionBackend):
if not prefix_len_ptr or not token_pos_in_items_ptr: if not prefix_len_ptr or not token_pos_in_items_ptr:
return MultiItemScoringParams() return MultiItemScoringParams()
# Build final params
device = forward_batch.input_ids.device
return MultiItemScoringParams( return MultiItemScoringParams(
prefix_len_ptr=torch.tensor( prefix_len_ptr=torch.tensor(
prefix_len_ptr, dtype=torch.uint32, device=device prefix_len_ptr, dtype=torch.uint32, device=device
@@ -470,7 +478,7 @@ class FlashInferAttnBackend(AttentionBackend):
prefix_lens = forward_batch.extend_prefix_lens prefix_lens = forward_batch.extend_prefix_lens
# Disable ragged wrapper and ensure prefix handling for multimodal and multi-item scoring # Disable ragged wrapper and ensure prefix handling for multimodal and multi-item scoring
if self.is_multimodal or self.multi_item_scoring_delimiter is not None: if self.is_multimodal or self.enable_mis:
# use_ragged = False: Multi-item scoring requires the paged wrapper because: # use_ragged = False: Multi-item scoring requires the paged wrapper because:
# 1. Ragged wrapper doesn't support the specialized multi-item parameters # 1. Ragged wrapper doesn't support the specialized multi-item parameters
# (prefix_len_ptr, token_pos_in_items_ptr, etc.) # (prefix_len_ptr, token_pos_in_items_ptr, etc.)
@@ -487,7 +495,7 @@ class FlashInferAttnBackend(AttentionBackend):
# Process multi-item scoring in attention backend instead of ForwardBatch # Process multi-item scoring in attention backend instead of ForwardBatch
multi_item_params = MultiItemScoringParams() multi_item_params = MultiItemScoringParams()
if self.multi_item_scoring_delimiter is not None: if self.enable_mis:
# Use new backend-specific implementation # Use new backend-specific implementation
multi_item_params = self._process_multi_item_scoring(forward_batch) multi_item_params = self._process_multi_item_scoring(forward_batch)
+49 -51
View File
@@ -274,9 +274,7 @@ class LogitsProcessor(nn.Module):
self.final_logit_softcapping = None self.final_logit_softcapping = None
self.return_full_logits = return_full_logits self.return_full_logits = return_full_logits
self.multi_item_delimiter = ( self.enable_mis = get_global_server_args().enable_mis
get_global_server_args().multi_item_scoring_delimiter
)
# enable chunked logprobs processing # enable chunked logprobs processing
self.enable_logprobs_chunk = envs.SGLANG_ENABLE_LOGITS_PROCESSER_CHUNK.get() self.enable_logprobs_chunk = envs.SGLANG_ENABLE_LOGITS_PROCESSER_CHUNK.get()
@@ -292,17 +290,20 @@ class LogitsProcessor(nn.Module):
aux_hidden_states: Optional[torch.Tensor] = None, aux_hidden_states: Optional[torch.Tensor] = None,
hidden_states_before_norm: Optional[torch.Tensor] = None, hidden_states_before_norm: Optional[torch.Tensor] = None,
) -> LogitsProcessorOutput: ) -> LogitsProcessorOutput:
# Extract MIS indices before ForwardBatch → LogitsMetadata conversion
multi_item_delimiter_indices = None
if isinstance(logits_metadata, ForwardBatch): if isinstance(logits_metadata, ForwardBatch):
multi_item_delimiter_indices = logits_metadata.multi_item_delimiter_indices
logits_metadata = LogitsMetadata.from_forward_batch(logits_metadata) logits_metadata = LogitsMetadata.from_forward_batch(logits_metadata)
# Multi-item scoring only for prefill-only requests. # Multi-item scoring only for prefill-only requests with pre-computed indices.
if self.multi_item_delimiter is not None and logits_metadata.is_prefill_only: if multi_item_delimiter_indices is not None and logits_metadata.is_prefill_only:
return self.compute_logprobs_for_multi_item_scoring( return self.compute_logprobs_for_multi_item_scoring(
input_ids, input_ids,
hidden_states, hidden_states,
lm_head, lm_head,
logits_metadata, logits_metadata,
self.multi_item_delimiter, multi_item_delimiter_indices,
) )
# Diffusion LLM only. # Diffusion LLM only.
@@ -347,6 +348,9 @@ class LogitsProcessor(nn.Module):
return LogitsProcessorOutput( return LogitsProcessorOutput(
next_token_logits=sampled_logits, next_token_logits=sampled_logits,
hidden_states=hidden_states_to_store, hidden_states=hidden_states_to_store,
# FIXME: These fields are not logits-related but are passed through here as a
# workaround since ForwardBatch is local to forward_batch_generation().
# They should be moved to GenerationBatchResult to keep this class clean.
mm_input_embeds=logits_metadata.mm_input_embeds, mm_input_embeds=logits_metadata.mm_input_embeds,
) )
@@ -1006,39 +1010,41 @@ class LogitsProcessor(nn.Module):
hidden_states, hidden_states,
lm_head: VocabParallelEmbedding, lm_head: VocabParallelEmbedding,
logits_metadata: Union[LogitsMetadata, ForwardBatch], logits_metadata: Union[LogitsMetadata, ForwardBatch],
delimiter_token: int, multi_item_delimiter_indices: List[torch.Tensor],
): ):
""" """
Compute logprobs for multi-item scoring using delimiter-based token extraction. Compute logprobs for multi-item scoring using pre-computed delimiter indices.
This method is designed for scenarios where you want to score multiple items/candidates
against a single query by combining them into one sequence separated by delimiters.
Sequence format: Query<delimiter>Item1<delimiter>Item2<delimiter>... Sequence format: Query<delimiter>Item1<delimiter>Item2<delimiter>...
Scoring positions: Extracts logprobs at positions before each <delimiter> Scoring positions: Extracts logprobs at positions before each <delimiter>
Args: Args:
input_ids (torch.Tensor): Input token IDs containing query and items separated by delimiters. input_ids: Input token IDs. Shape: [total_sequence_length].
Shape: [total_sequence_length] for single request or [batch_total_length] for batch. hidden_states: Hidden states from the model. Shape: [sequence_length, hidden_dim].
hidden_states (torch.Tensor): Hidden states from the model. lm_head: Language model head for computing logits.
Shape: [sequence_length, hidden_dim]. logits_metadata: Metadata containing batch info and logprob specs.
lm_head (VocabParallelEmbedding): Language model head for computing logits. multi_item_delimiter_indices: Pre-computed delimiter positions per request (CPU tensors).
logits_metadata (Union[LogitsMetadata, ForwardBatch]): Metadata containing batch info
and token ID specifications for logprob extraction.
delimiter_token (int): Token ID used as delimiter between query and items.
Returns:
LogitsProcessorOutput: Contains:
- next_token_logits: None (not needed for scoring-only requests)
- input_token_logprobs: Logprobs of delimiter tokens at scoring positions
- input_top_logprobs_val: Top-k logprobs at delimiter positions (if requested)
- input_top_logprobs_idx: Top-k token indices at delimiter positions (if requested)
- input_token_ids_logprobs_val: Logprobs for user-requested token IDs (if any)
- input_token_ids_logprobs_idx: Indices for user-requested token IDs (if any)
""" """
multi_item_indices = (input_ids == delimiter_token).nonzero(as_tuple=True)[ # Compute positions just before each delimiter.
0 # Build offset-adjusted indices on CPU, then do a single CPU→GPU transfer.
] - 1 device = input_ids.device
all_tensors = []
if logits_metadata.extend_seq_lens_cpu is not None:
offset = 0
for req_seq_len, indices_tensor in zip(
logits_metadata.extend_seq_lens_cpu, multi_item_delimiter_indices
):
if len(indices_tensor) > 0:
# Note: if the first delimiter is at position 0 (empty query),
# indices - 1 wraps to -1. This is harmless — the first
# delimiter entry is always discarded by
# _process_multi_item_scoring_results.
all_tensors.append(indices_tensor + (offset - 1))
offset += req_seq_len
else:
all_tensors.append(multi_item_delimiter_indices[0] - 1)
multi_item_indices = torch.cat(all_tensors).to(device, non_blocking=True)
# Extract hidden states at delimiter positions for multi-item scoring # Extract hidden states at delimiter positions for multi-item scoring
sliced_hidden = hidden_states[multi_item_indices] sliced_hidden = hidden_states[multi_item_indices]
@@ -1052,27 +1058,13 @@ class LogitsProcessor(nn.Module):
input_top_logprobs_idx = None input_top_logprobs_idx = None
# Recalculate extend_logprob_pruned_lens_cpu to match delimiter counts per request # Recalculate extend_logprob_pruned_lens_cpu to match delimiter counts per request
# Original contains sequence lengths, but we need delimiter counts for sliced_logprobs
if ( if (
logits_metadata.token_ids_logprobs logits_metadata.token_ids_logprobs
or logits_metadata.extend_return_top_logprob or logits_metadata.extend_return_top_logprob
): ):
logits_metadata.extend_logprob_pruned_lens_cpu = [] logits_metadata.extend_logprob_pruned_lens_cpu = [
len(t) for t in multi_item_delimiter_indices
if logits_metadata.extend_seq_lens_cpu is not None: ]
# Multi-request batch: count delimiters per request
input_pt = 0
for req_seq_len in logits_metadata.extend_seq_lens_cpu:
req_input_ids = input_ids[input_pt : input_pt + req_seq_len]
delimiter_count = (req_input_ids == delimiter_token).sum().item()
logits_metadata.extend_logprob_pruned_lens_cpu.append(
delimiter_count
)
input_pt += req_seq_len
else:
# Single request case: one request gets all delimiters
total_delimiters = (input_ids == delimiter_token).sum().item()
logits_metadata.extend_logprob_pruned_lens_cpu = [total_delimiters]
# Get the logprobs of specified token ids # Get the logprobs of specified token ids
if logits_metadata.extend_token_ids_logprob: if logits_metadata.extend_token_ids_logprob:
@@ -1090,11 +1082,17 @@ class LogitsProcessor(nn.Module):
input_top_logprobs_idx, input_top_logprobs_idx,
) = get_top_logprobs_prefill(sliced_logprobs, logits_metadata) ) = get_top_logprobs_prefill(sliced_logprobs, logits_metadata)
# For input_token_logprobs, use delimiter token logprobs # MIS scores come from input_token_ids_logprobs_val (label-token logprobs),
input_token_logprobs = sliced_logprobs[:, delimiter_token] # not from per-position input_token_logprobs. However, the shared logprob
# pipeline (add_input_logprob_return_values) asserts input_token_logprobs is
# non-None, converts it to a tuple, slices it, and validates its length —
# all before score_request() ever sees the result. We can't set it to None
# without changing those shared asserts, so we fill with zeros to satisfy
# the pipeline. score_request() ignores this field entirely.
input_token_logprobs = torch.zeros(multi_item_indices.shape[0], device=device)
return LogitsProcessorOutput( return LogitsProcessorOutput(
next_token_logits=None, # Multi-item scoring doesn't need next token logits next_token_logits=None,
input_token_logprobs=input_token_logprobs, input_token_logprobs=input_token_logprobs,
input_top_logprobs_val=input_top_logprobs_val, input_top_logprobs_val=input_top_logprobs_val,
input_top_logprobs_idx=input_top_logprobs_idx, input_top_logprobs_idx=input_top_logprobs_idx,
+61 -33
View File
@@ -12,7 +12,6 @@ import torch.nn as nn
from transformers import PretrainedConfig 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.server_args import get_global_server_args
if TYPE_CHECKING: if TYPE_CHECKING:
from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.model_executor.forward_batch_info import ForwardBatch
@@ -66,6 +65,46 @@ def pool_hidden_states(
raise ValueError(f"Unsupported pooling type: {pooling_type}") raise ValueError(f"Unsupported pooling type: {pooling_type}")
def pool_at_delimiter_positions(
data: torch.Tensor,
forward_batch: ForwardBatch,
device: torch.device,
) -> List[torch.Tensor]:
"""Pool a tensor at the position before each MIS delimiter for every request.
Uses pre-computed delimiter indices from ForwardBatch (CPU tensors),
moves to GPU with non_blocking=True to avoid CUDA syncs.
Args:
data: 2-D tensor [total_tokens, dim] — hidden states or logits.
forward_batch: Forward batch with extend_seq_lens_cpu and
multi_item_delimiter_indices populated.
device: Device for the index tensor.
Returns:
One tensor per request, shaped [num_delimiters, dim].
"""
all_index_tensors: List[torch.Tensor] = []
delim_counts: List[int] = []
offset = 0
for req_idx, req_seq_len in enumerate(forward_batch.extend_seq_lens_cpu):
indices_tensor = forward_batch.multi_item_delimiter_indices[req_idx]
n = len(indices_tensor)
if n > 0:
# Note: if the first delimiter is at position 0 (empty query),
# indices - 1 wraps to -1. This is harmless — the first delimiter
# entry is always discarded by _process_multi_item_scoring_results.
all_index_tensors.append(indices_tensor + (offset - 1))
delim_counts.append(n)
offset += req_seq_len
if all_index_tensors:
index_tensor = torch.cat(all_index_tensors).to(device, non_blocking=True)
else:
index_tensor = torch.tensor([], dtype=torch.long, device=device)
return list(data[index_tensor].split(delim_counts))
def score_and_pool( def score_and_pool(
score_head: nn.Module, score_head: nn.Module,
pooler: "Pooler", pooler: "Pooler",
@@ -75,46 +114,35 @@ def score_and_pool(
) -> EmbeddingPoolerOutput: ) -> EmbeddingPoolerOutput:
"""Apply a classification/score head with MIS and pooled-hidden-states support. """Apply a classification/score head with MIS and pooled-hidden-states support.
MIS path (when ``multi_item_scoring_delimiter`` is set and found in ``input_ids``): MIS path (pre-computed delimiter indices on forward_batch): extract hidden
extract hidden states at positions just before each delimiter, apply the score head, states at positions just before each delimiter, apply the score head, then
then split per-request. split per-request.
Standard path: apply the score head to all hidden states, then pool. Standard path: pool hidden states, then apply the score head.
When ``forward_batch.return_pooled_hidden_states`` is True, the raw pooled When ``forward_batch.return_pooled_hidden_states`` is True, the raw pooled
hidden states (before the score head) are included in the output. hidden states (before the score head) are included in the output.
""" """
delimiter_token = get_global_server_args().multi_item_scoring_delimiter if (
if delimiter_token is not None and forward_batch.is_prefill_only: forward_batch.multi_item_delimiter_indices is not None
delim_positions = (input_ids == delimiter_token).nonzero(as_tuple=True)[0] and forward_batch.is_prefill_only
# A delimiter at flat index 0 has no preceding hidden state to pool ):
delim_positions = delim_positions[delim_positions > 0] # Pool hidden states at pre-delimiter positions, score only those —
# avoids wasting compute on tokens that never contribute to the output.
if delim_positions.numel() > 0: # pool_at_delimiter_positions returns one tensor per request; we concat
# Score only the tokens that precede a delimiter # to call score_head once, then split back per request.
pre_delim_hidden = hidden_states[delim_positions - 1] per_request_phs = pool_at_delimiter_positions(
scores = score_head(pre_delim_hidden) hidden_states, forward_batch, input_ids.device
# 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_scores: List[torch.Tensor] = []
per_request_phs: Optional[List[torch.Tensor]] = (
[] if forward_batch.return_pooled_hidden_states else None
) )
for seq_len in seq_lens: phs_flat = torch.cat(per_request_phs, dim=0)
end = start + seq_len scores_flat = score_head(phs_flat)
mask = (delim_positions >= start) & (delim_positions < end) delim_counts = [t.shape[0] for t in per_request_phs]
per_request_scores.append(scores[mask]) per_request_scores = list(scores_flat.split(delim_counts))
if per_request_phs is not None:
per_request_phs.append(pre_delim_hidden[mask])
start = end
return EmbeddingPoolerOutput( return EmbeddingPoolerOutput(
embeddings=per_request_scores, embeddings=per_request_scores,
pooled_hidden_states=per_request_phs, pooled_hidden_states=(
per_request_phs if forward_batch.return_pooled_hidden_states else None
),
) )
# Standard classification path: pool hidden states, then score. # Standard classification path: pool hidden states, then score.
+28
View File
@@ -250,6 +250,10 @@ class GenerateReqInput(BaseReq):
image_max_dynamic_patch: Optional[int] = None image_max_dynamic_patch: Optional[int] = None
video_max_dynamic_patch: Optional[int] = None video_max_dynamic_patch: Optional[int] = None
# Pre-computed delimiter indices for multi-item scoring.
# Batch-level: List[List[int]] (one per request). After __getitem__: List[int].
multi_item_delimiter_indices: Optional[Union[List[List[int]], List[int]]] = None
def contains_mm_input(self) -> bool: def contains_mm_input(self) -> bool:
return ( return (
has_valid_data(self.image_data) has_valid_data(self.image_data)
@@ -685,6 +689,11 @@ class GenerateReqInput(BaseReq):
external_trace_header=self.external_trace_header, external_trace_header=self.external_trace_header,
http_worker_ipc=self.http_worker_ipc, http_worker_ipc=self.http_worker_ipc,
received_time=self.received_time, received_time=self.received_time,
multi_item_delimiter_indices=(
self.multi_item_delimiter_indices[i]
if self.multi_item_delimiter_indices is not None
else None
),
) )
cache[i] = sub cache[i] = sub
return sub return sub
@@ -774,6 +783,9 @@ class TokenizedGenerateReqInput(BaseReq):
need_wait_for_mm_inputs: bool = False need_wait_for_mm_inputs: bool = False
num_items_assigned: Optional[Dict[Modality, List[int]]] = None num_items_assigned: Optional[Dict[Modality, List[int]]] = None
# Pre-computed delimiter indices for multi-item scoring
multi_item_delimiter_indices: Optional[List[int]] = None
# For observability # For observability
time_stats: Optional[Union[APIServerReqTimeStats, DPControllerReqTimeStats]] = None time_stats: Optional[Union[APIServerReqTimeStats, DPControllerReqTimeStats]] = None
@@ -855,6 +867,10 @@ class EmbeddingReqInput(BaseReq):
# Whether to return pooled hidden states (pre-head transformer output) # Whether to return pooled hidden states (pre-head transformer output)
return_pooled_hidden_states: bool = False return_pooled_hidden_states: bool = False
# Pre-computed delimiter indices for multi-item scoring.
# Batch-level: List[List[int]] (one per request). After __getitem__: List[int].
multi_item_delimiter_indices: Optional[Union[List[List[int]], List[int]]] = None
def normalize_batch_and_arguments(self): def normalize_batch_and_arguments(self):
# at least one of text, input_ids, or image should be provided # at least one of text, input_ids, or image should be provided
if self.text is None and self.input_ids is None and self.image_data is None: if self.text is None and self.input_ids is None and self.image_data is None:
@@ -957,6 +973,11 @@ class EmbeddingReqInput(BaseReq):
is_cross_encoder_request=True, is_cross_encoder_request=True,
http_worker_ipc=self.http_worker_ipc, http_worker_ipc=self.http_worker_ipc,
return_pooled_hidden_states=self.return_pooled_hidden_states, return_pooled_hidden_states=self.return_pooled_hidden_states,
multi_item_delimiter_indices=(
self.multi_item_delimiter_indices[i]
if self.multi_item_delimiter_indices is not None
else None
),
) )
else: else:
sub = EmbeddingReqInput( sub = EmbeddingReqInput(
@@ -981,6 +1002,11 @@ class EmbeddingReqInput(BaseReq):
http_worker_ipc=self.http_worker_ipc, http_worker_ipc=self.http_worker_ipc,
received_time=self.received_time, received_time=self.received_time,
return_pooled_hidden_states=self.return_pooled_hidden_states, return_pooled_hidden_states=self.return_pooled_hidden_states,
multi_item_delimiter_indices=(
self.multi_item_delimiter_indices[i]
if self.multi_item_delimiter_indices is not None
else None
),
) )
cache[i] = sub cache[i] = sub
return sub return sub
@@ -1009,6 +1035,8 @@ class TokenizedEmbeddingReqInput(BaseReq):
# LoRA related # LoRA related
lora_id: Optional[str] = None # None means just use the base model lora_id: Optional[str] = None # None means just use the base model
# Pre-computed delimiter indices for multi-item scoring
multi_item_delimiter_indices: Optional[List[int]] = None
# For observability # For observability
time_stats: Optional[Union[APIServerReqTimeStats, DPControllerReqTimeStats]] = None time_stats: Optional[Union[APIServerReqTimeStats, DPControllerReqTimeStats]] = None
@@ -597,6 +597,7 @@ class Req(ReqDllmMixin):
Union[APIServerReqTimeStats, DPControllerReqTimeStats] Union[APIServerReqTimeStats, DPControllerReqTimeStats]
] = None, ] = None,
return_pooled_hidden_states: bool = False, return_pooled_hidden_states: bool = False,
multi_item_delimiter_indices: Optional[List[int]] = None,
): ):
# Input and output info # Input and output info
self.rid = rid self.rid = rid
@@ -614,6 +615,7 @@ class Req(ReqDllmMixin):
self.session = session self.session = session
self.input_embeds = input_embeds self.input_embeds = input_embeds
self.positional_embed_overrides = positional_embed_overrides self.positional_embed_overrides = positional_embed_overrides
self.multi_item_delimiter_indices = multi_item_delimiter_indices
# For req-level memory management # For req-level memory management
self.kv_committed_len = 0 self.kv_committed_len = 0
@@ -1441,6 +1443,9 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
# Whether this batch is prefill-only (no token generation needed) # Whether this batch is prefill-only (no token generation needed)
is_prefill_only: bool = False is_prefill_only: bool = False
# Multi-item scoring delimiter indices (set during prepare_for_extend)
multi_item_delimiter_indices: Optional[List[torch.Tensor]] = None
# hicache pointer for synchronizing data loading from CPU to GPU # hicache pointer for synchronizing data loading from CPU to GPU
hicache_consumer_index: int = -1 hicache_consumer_index: int = -1
@@ -1817,6 +1822,23 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
self.token_type_ids = token_type_ids_tensor self.token_type_ids = token_type_ids_tensor
self.seq_lens_sum = sum(seq_lens) self.seq_lens_sum = sum(seq_lens)
# Pre-compute delimiter indices as CPU tensors for MIS.
# When --enable-mis is on, every request in the batch is expected to
# carry delimiter indices (the score endpoint always produces MIS-structured
# requests). Consumers index this list without None-checking.
if get_global_server_args().enable_mis and any(
r.multi_item_delimiter_indices is not None for r in reqs
):
assert all(
r.multi_item_delimiter_indices is not None for r in reqs
), "MIS batch must have delimiter indices on every request"
self.multi_item_delimiter_indices = [
torch.tensor(r.multi_item_delimiter_indices, dtype=torch.int64)
for r in reqs
]
else:
self.multi_item_delimiter_indices = None
if self.return_logprob: if self.return_logprob:
self.top_logprobs_nums = [r.top_logprobs_num for r in reqs] self.top_logprobs_nums = [r.top_logprobs_num for r in reqs]
self.token_ids_logprobs = [r.token_ids_logprob for r in reqs] self.token_ids_logprobs = [r.token_ids_logprob for r in reqs]
@@ -2464,6 +2486,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
), ),
extend_input_logprob_token_ids=self.extend_input_logprob_token_ids, extend_input_logprob_token_ids=self.extend_input_logprob_token_ids,
is_prefill_only=self.is_prefill_only, is_prefill_only=self.is_prefill_only,
multi_item_delimiter_indices=self.multi_item_delimiter_indices,
dimensions=self.dimensions, dimensions=self.dimensions,
return_pooled_hidden_states=self.return_pooled_hidden_states, return_pooled_hidden_states=self.return_pooled_hidden_states,
dllm_block_offsets=[req.dllm_block_offset for req in self.reqs], dllm_block_offsets=[req.dllm_block_offset for req in self.reqs],
@@ -2665,6 +2688,9 @@ class ModelWorkerBatch:
# Whether this batch is prefill-only (no token generation needed) # Whether this batch is prefill-only (no token generation needed)
is_prefill_only: bool = False is_prefill_only: bool = False
# Pre-computed delimiter indices for multi-item scoring (CPU tensors, one per request)
multi_item_delimiter_indices: Optional[List[torch.Tensor]] = None
# Diffusion LLM # Diffusion LLM
dllm_block_offsets: Optional[List[int]] = None dllm_block_offsets: Optional[List[int]] = None
dllm_config: Optional[DllmConfig] = None dllm_config: Optional[DllmConfig] = None
+2
View File
@@ -1881,6 +1881,7 @@ class Scheduler(
http_worker_ipc=recv_req.http_worker_ipc, http_worker_ipc=recv_req.http_worker_ipc,
dllm_config=self.dllm_config, dllm_config=self.dllm_config,
time_stats=recv_req.time_stats, time_stats=recv_req.time_stats,
multi_item_delimiter_indices=recv_req.multi_item_delimiter_indices,
) )
req.tokenizer = self.tokenizer req.tokenizer = self.tokenizer
@@ -2202,6 +2203,7 @@ class Scheduler(
http_worker_ipc=recv_req.http_worker_ipc, http_worker_ipc=recv_req.http_worker_ipc,
time_stats=recv_req.time_stats, time_stats=recv_req.time_stats,
return_pooled_hidden_states=recv_req.return_pooled_hidden_states, return_pooled_hidden_states=recv_req.return_pooled_hidden_states,
multi_item_delimiter_indices=recv_req.multi_item_delimiter_indices,
) )
req.tokenizer = self.tokenizer req.tokenizer = self.tokenizer
@@ -21,7 +21,7 @@ from sglang.srt.managers.schedule_batch import (
ScheduleBatch, ScheduleBatch,
) )
from sglang.srt.mem_cache.common import release_kv_cache from sglang.srt.mem_cache.common import release_kv_cache
from sglang.srt.server_args import get_global_server_args from sglang.srt.server_args import MIS_DELIMITER_TOKEN_ID, get_global_server_args
if TYPE_CHECKING: if TYPE_CHECKING:
from sglang.srt.managers.scheduler import ( from sglang.srt.managers.scheduler import (
@@ -629,13 +629,12 @@ class SchedulerOutputProcessorMixin:
# Process logprob indices based on scoring type # Process logprob indices based on scoring type
if is_multi_item_scoring: if is_multi_item_scoring:
# Multi-item scoring: only include delimiter token positions # MIS scores come from input_token_ids_logprobs, not input_token_logprobs.
relevant_tokens = req.origin_input_ids[req.logprob_start_len :] # But the shared pipeline requires input_token_logprobs_idx to be the same
input_token_logprobs_idx = [ # length as input_token_logprobs_val (validated at line 816). We fill with
token_id # MIS_DELIMITER_TOKEN_ID as a dummy — score_request() ignores this field.
for token_id in relevant_tokens delimiter_count = len(req.multi_item_delimiter_indices)
if token_id == self.server_args.multi_item_scoring_delimiter input_token_logprobs_idx = [MIS_DELIMITER_TOKEN_ID] * delimiter_count
]
else: else:
# Regular request: include all tokens from logprob_start_len onwards # Regular request: include all tokens from logprob_start_len onwards
input_token_logprobs_idx = req.origin_input_ids[req.logprob_start_len :] input_token_logprobs_idx = req.origin_input_ids[req.logprob_start_len :]
@@ -714,18 +713,11 @@ class SchedulerOutputProcessorMixin:
For regular requests, all positions from logprob_start_len onwards have logprobs. For regular requests, all positions from logprob_start_len onwards have logprobs.
""" """
is_multi_item_scoring = self._is_multi_item_scoring(req) is_multi_item_scoring = self._is_multi_item_scoring(req)
relevant_tokens = req.origin_input_ids[req.logprob_start_len :]
if is_multi_item_scoring: if is_multi_item_scoring:
# Multi-item scoring: count delimiter tokens from logprob_start_len onwards return len(req.multi_item_delimiter_indices)
return sum(
1
for token_id in relevant_tokens
if token_id == self.server_args.multi_item_scoring_delimiter
)
else: else:
# Regular request: all tokens from logprob_start_len onwards return len(req.origin_input_ids[req.logprob_start_len :])
return len(relevant_tokens)
def _calculate_num_input_logprobs( def _calculate_num_input_logprobs(
self: Scheduler, req: Req, extend_input_len: int, extend_logprob_start_len: int self: Scheduler, req: Req, extend_input_len: int, extend_logprob_start_len: int
@@ -738,14 +730,11 @@ class SchedulerOutputProcessorMixin:
is_multi_item_scoring = self._is_multi_item_scoring(req) is_multi_item_scoring = self._is_multi_item_scoring(req)
if is_multi_item_scoring: if is_multi_item_scoring:
# Multi-item scoring: count delimiter tokens in the relevant portion # Count pre-computed delimiter indices within the extend range
relevant_tokens = req.origin_input_ids[
extend_logprob_start_len:extend_input_len
]
return sum( return sum(
1 1
for token_id in relevant_tokens for idx in req.multi_item_delimiter_indices
if token_id == self.server_args.multi_item_scoring_delimiter if extend_logprob_start_len <= idx < extend_input_len
) )
else: else:
# Regular request: all tokens in the range # Regular request: all tokens in the range
@@ -758,7 +747,11 @@ class SchedulerOutputProcessorMixin:
token is configured. In this mode, only positions containing the token is configured. In this mode, only positions containing the
delimiter token receive logprobs. delimiter token receive logprobs.
""" """
return req.is_prefill_only and self.server_args.multi_item_scoring_delimiter return (
self.server_args.enable_mis
and req.is_prefill_only
and req.multi_item_delimiter_indices is not None
)
def add_input_logprob_return_values( def add_input_logprob_return_values(
self: Scheduler, self: Scheduler,
@@ -313,7 +313,6 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
self.processor = _processor self.processor = _processor
self.tokenizer = get_tokenizer_from_processor(self.processor) self.tokenizer = get_tokenizer_from_processor(self.processor)
os.environ["TOKENIZERS_PARALLELISM"] = "false" os.environ["TOKENIZERS_PARALLELISM"] = "false"
self._initialize_multi_item_delimiter_text()
else: else:
self.mm_processor = self.processor = None self.mm_processor = self.processor = None
@@ -326,7 +325,6 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
trust_remote_code=server_args.trust_remote_code, trust_remote_code=server_args.trust_remote_code,
revision=server_args.revision, revision=server_args.revision,
) )
self._initialize_multi_item_delimiter_text()
# Initialize async dynamic batch tokenizer if enabled (common for both multimodal and non-multimodal) # Initialize async dynamic batch tokenizer if enabled (common for both multimodal and non-multimodal)
if ( if (
@@ -1007,6 +1005,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
token_type_ids=token_type_ids, token_type_ids=token_type_ids,
need_wait_for_mm_inputs=obj.need_wait_for_mm_inputs, need_wait_for_mm_inputs=obj.need_wait_for_mm_inputs,
num_items_assigned=obj.num_items_assigned, num_items_assigned=obj.num_items_assigned,
multi_item_delimiter_indices=obj.multi_item_delimiter_indices,
) )
elif isinstance(obj, EmbeddingReqInput): elif isinstance(obj, EmbeddingReqInput):
# Resolve unresolved embed overrides now that input_ids are available # Resolve unresolved embed overrides now that input_ids are available
@@ -1033,6 +1032,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
lora_id=obj.lora_id, lora_id=obj.lora_id,
http_worker_ipc=obj.http_worker_ipc, http_worker_ipc=obj.http_worker_ipc,
return_pooled_hidden_states=obj.return_pooled_hidden_states, return_pooled_hidden_states=obj.return_pooled_hidden_states,
multi_item_delimiter_indices=obj.multi_item_delimiter_indices,
) )
tokenized_obj.time_stats = self.rid_to_state[obj.rid].time_stats tokenized_obj.time_stats = self.rid_to_state[obj.rid].time_stats
@@ -8,6 +8,7 @@ import torch
from sglang.srt.configs.model_config import is_cross_encoding_pooler_model from sglang.srt.configs.model_config import is_cross_encoding_pooler_model
from sglang.srt.managers.embed_types import PositionalEmbeds from sglang.srt.managers.embed_types import PositionalEmbeds
from sglang.srt.managers.io_struct import EmbeddingReqInput, GenerateReqInput from sglang.srt.managers.io_struct import EmbeddingReqInput, GenerateReqInput
from sglang.srt.server_args import MIS_DELIMITER_TOKEN_ID
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -76,27 +77,9 @@ class TokenizerManagerScoreMixin:
raise ValueError("Invalid prompts type for score_prompts.") raise ValueError("Invalid prompts type for score_prompts.")
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( def _build_multi_item_token_sequence(
self, query: List[int], items: List[List[int]], delimiter_token_id: int self, query: List[int], items: List[List[int]], delimiter_token_id: int
) -> List[int]: ) -> Tuple[List[int], List[int]]:
""" """
Build a single token sequence for multi-item scoring. Build a single token sequence for multi-item scoring.
Format: query<delimiter>item1<delimiter>item2<delimiter>item3<delimiter> Format: query<delimiter>item1<delimiter>item2<delimiter>item3<delimiter>
@@ -107,18 +90,21 @@ class TokenizerManagerScoreMixin:
delimiter_token_id: Token ID to use as delimiter delimiter_token_id: Token ID to use as delimiter
Returns: Returns:
Combined token sequence Tuple of (combined token sequence, delimiter indices)
""" """
combined_sequence = query[:] # Start with query combined_sequence = query[:] # Start with query
delimiter_indices = []
for item in items: for item in items:
delimiter_indices.append(len(combined_sequence))
combined_sequence.append(delimiter_token_id) # Add delimiter combined_sequence.append(delimiter_token_id) # Add delimiter
combined_sequence.extend(item) # Add item tokens combined_sequence.extend(item) # Add item tokens
# Add final delimiter after the last item for logprob extraction # Add final delimiter after the last item for logprob extraction
delimiter_indices.append(len(combined_sequence))
combined_sequence.append(delimiter_token_id) combined_sequence.append(delimiter_token_id)
return combined_sequence return combined_sequence, delimiter_indices
def _batch_tokenize_query_and_items( def _batch_tokenize_query_and_items(
self, self,
@@ -416,11 +402,14 @@ class TokenizerManagerScoreMixin:
embed_override_token_id: Optional[int], embed_override_token_id: Optional[int],
query_embed_overrides: Optional[List[torch.Tensor]], query_embed_overrides: Optional[List[torch.Tensor]],
item_embed_overrides: Optional[List[Optional[List[torch.Tensor]]]], item_embed_overrides: Optional[List[Optional[List[torch.Tensor]]]],
) -> Tuple[None, List[List[int]], Optional[list]]: ) -> Tuple[None, List[List[int]], Optional[list], Optional[List[int]]]:
"""Build input_ids and resolve embed overrides for token-ID inputs. """Build input_ids and resolve embed overrides for token-ID inputs.
Works identically for multi-item-scoring and single-item modes — the only difference is Works identically for multi-item-scoring and single-item modes — the only difference is
how input_ids are assembled and what position offset each item gets. how input_ids are assembled and what position offset each item gets.
Returns:
(text_prompts, input_ids, positional_embed_overrides, delimiter_indices)
""" """
# Both query and items are token IDs # Both query and items are token IDs
has_embeds = ( has_embeds = (
@@ -428,16 +417,17 @@ class TokenizerManagerScoreMixin:
) )
if use_multi_item_scoring: if use_multi_item_scoring:
# Multi-item scoring: concatenate with delimiter token ID # Multi-item scoring: concatenate with placeholder delimiter token.
# Format: query<delimiter_token_id>item1<delimiter_token_id>item2<delimiter_token_id>item3<delimiter_token_id> # Positions are derived from item lengths (delimiter_indices), not
delimiter_token_id = self.server_args.multi_item_scoring_delimiter # by scanning for this token — it exists only for FlashInfer compat.
combined_input_ids = self._build_multi_item_token_sequence( delimiter_token_id = MIS_DELIMITER_TOKEN_ID
query, items, delimiter_token_id combined_input_ids, delimiter_indices = (
self._build_multi_item_token_sequence(query, items, delimiter_token_id)
) )
input_ids = [combined_input_ids] input_ids = [combined_input_ids]
if not has_embeds: if not has_embeds:
return None, input_ids, None return None, input_ids, None, delimiter_indices
# Resolve embed overrides across the combined multi-item-scoring sequence # Resolve embed overrides across the combined multi-item-scoring sequence
all_embeds: List[torch.Tensor] = [] all_embeds: List[torch.Tensor] = []
@@ -461,15 +451,15 @@ class TokenizerManagerScoreMixin:
current_offset += len(item) + 1 # +1 for delimiter current_offset += len(item) + 1 # +1 for delimiter
if all_embeds: if all_embeds:
injection = [ positional_embed_overrides = [
PositionalEmbeds( PositionalEmbeds(
embeds=torch.cat(all_embeds, dim=0), embeds=torch.cat(all_embeds, dim=0),
positions=all_positions, positions=all_positions,
) )
] ]
else: else:
injection = None positional_embed_overrides = None
return None, input_ids, injection return None, input_ids, positional_embed_overrides, delimiter_indices
else: else:
# Single-item scoring: process each item separately # Single-item scoring: process each item separately
@@ -479,9 +469,9 @@ class TokenizerManagerScoreMixin:
input_ids = [query + item for item in items] input_ids = [query + item for item in items]
if not has_embeds: if not has_embeds:
return None, input_ids, None return None, input_ids, None, None
injection = [] positional_embed_overrides = []
for i, item in enumerate(items): for i, item in enumerate(items):
item_embs = item_embed_overrides[i] if item_embed_overrides else None item_embs = item_embed_overrides[i] if item_embed_overrides else None
pe = self._resolve_embed_overrides_for_request( pe = self._resolve_embed_overrides_for_request(
@@ -493,13 +483,14 @@ class TokenizerManagerScoreMixin:
item_position_offset=len(query), item_position_offset=len(query),
item_label=f"items[{i}]", item_label=f"items[{i}]",
) )
injection.append(pe) positional_embed_overrides.append(pe)
return ( positional_embed_overrides = (
None, positional_embed_overrides
input_ids, if any(pe is not None for pe in positional_embed_overrides)
injection if any(pe is not None for pe in injection) else None, else None
) )
return None, input_ids, positional_embed_overrides, None
# ------------------------------------------------------------------ # ------------------------------------------------------------------
# Main entry point # Main entry point
@@ -523,7 +514,7 @@ class TokenizerManagerScoreMixin:
This method supports two scoring approaches: This method supports two scoring approaches:
1. Single-Item scoring (default): Process each query+item pair independently 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 2. Multi-Item scoring: When --enable-mis is set, combine query and
multiple items into a single sequence using delimiter for efficient processing. 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 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> a fixed format: query<delimiter>item1<delimiter>item2<delimiter>item3<delimiter>
@@ -593,15 +584,13 @@ class TokenizerManagerScoreMixin:
f"Token ID {token_id} is out of vocabulary (vocab size: {vocab_size})" f"Token ID {token_id} is out of vocabulary (vocab size: {vocab_size})"
) )
# Check if multi-item scoring is enabled by presence of delimiter # Check if multi-item scoring is enabled
use_multi_item_scoring = ( use_multi_item_scoring = self.server_args.enable_mis
self.server_args.multi_item_scoring_delimiter is not None
and self.multi_item_delimiter_text is not None
)
input_ids = None input_ids = None
text_prompts = None text_prompts = None
positional_embed_overrides = None positional_embed_overrides = None
delimiter_indices = None
use_text_prompts = isinstance(query, str) and not has_embeds use_text_prompts = isinstance(query, str) and not has_embeds
@@ -609,16 +598,18 @@ class TokenizerManagerScoreMixin:
# Both query and items are text # Both query and items are text
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: tokenize separately then combine at token level # Tokenize separately, then combine at token level with placeholder
# to ensure the delimiter token ID is inserted exactly once per boundary # delimiter. Positions come from item lengths (delimiter_indices),
# (a text-level roundtrip through the tokenizer can alter boundary tokens) # not from scanning for this token — it's for FlashInfer compat only.
delimiter_token_id = self.server_args.multi_item_scoring_delimiter delimiter_token_id = MIS_DELIMITER_TOKEN_ID
query_ids, items_ids = self._batch_tokenize_query_and_items( query_ids, items_ids = self._batch_tokenize_query_and_items(
query, items_list query, items_list
) )
combined_input_ids = self._build_multi_item_token_sequence( combined_input_ids, delimiter_indices = (
self._build_multi_item_token_sequence(
query_ids, items_ids, delimiter_token_id query_ids, items_ids, delimiter_token_id
) )
)
input_ids = [combined_input_ids] 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
@@ -635,7 +626,8 @@ class TokenizerManagerScoreMixin:
): ):
# Both query and items are token IDs — tokenize text inputs if needed for embed overrides # Both query and items are token IDs — tokenize text inputs if needed for embed overrides
query_ids, items_ids = query, items query_ids, items_ids = query, items
_, input_ids, positional_embed_overrides = self._build_token_id_inputs( _, input_ids, positional_embed_overrides, delimiter_indices = (
self._build_token_id_inputs(
query_ids, query_ids,
items_ids, items_ids,
item_first, item_first,
@@ -644,10 +636,12 @@ class TokenizerManagerScoreMixin:
query_embed_overrides, query_embed_overrides,
item_embed_overrides, item_embed_overrides,
) )
)
elif has_embeds: elif has_embeds:
# Text inputs with embed overrides — need to tokenize first to resolve positions # Text inputs with embed overrides — need to tokenize first to resolve positions
query_ids, items_ids = self._batch_tokenize_query_and_items(query, items) query_ids, items_ids = self._batch_tokenize_query_and_items(query, items)
_, input_ids, positional_embed_overrides = self._build_token_id_inputs( _, input_ids, positional_embed_overrides, delimiter_indices = (
self._build_token_id_inputs(
query_ids, query_ids,
items_ids, items_ids,
item_first, item_first,
@@ -656,6 +650,7 @@ class TokenizerManagerScoreMixin:
query_embed_overrides, query_embed_overrides,
item_embed_overrides, item_embed_overrides,
) )
)
else: else:
raise ValueError( raise ValueError(
"Invalid combination of query/items types for score_request." "Invalid combination of query/items types for score_request."
@@ -679,6 +674,7 @@ class TokenizerManagerScoreMixin:
) )
# Create the appropriate request type # Create the appropriate request type
mis_delimiter_indices = [delimiter_indices] if use_multi_item_scoring else None
if is_generation: if is_generation:
batch_request = GenerateReqInput( batch_request = GenerateReqInput(
text=text_prompts, text=text_prompts,
@@ -690,6 +686,7 @@ class TokenizerManagerScoreMixin:
stream=False, stream=False,
sampling_params={"max_new_tokens": 0}, sampling_params={"max_new_tokens": 0},
positional_embed_overrides=positional_embed_overrides, positional_embed_overrides=positional_embed_overrides,
multi_item_delimiter_indices=mis_delimiter_indices,
) )
else: else:
batch_request = EmbeddingReqInput( batch_request = EmbeddingReqInput(
@@ -697,6 +694,7 @@ class TokenizerManagerScoreMixin:
input_ids=input_ids, input_ids=input_ids,
positional_embed_overrides=positional_embed_overrides, positional_embed_overrides=positional_embed_overrides,
return_pooled_hidden_states=return_pooled_hidden_states, return_pooled_hidden_states=return_pooled_hidden_states,
multi_item_delimiter_indices=mis_delimiter_indices,
) )
results = await self.generate_request(batch_request, request).__anext__() results = await self.generate_request(batch_request, request).__anext__()
@@ -396,6 +396,9 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
# Whether this batch is prefill-only (no token generation needed) # Whether this batch is prefill-only (no token generation needed)
is_prefill_only: bool = False is_prefill_only: bool = False
# Pre-computed delimiter indices for multi-item scoring (CPU tensors, one per request)
multi_item_delimiter_indices: Optional[List[torch.Tensor]] = None
# Speculative decoding # Speculative decoding
spec_info: Optional[SpecInput] = None spec_info: Optional[SpecInput] = None
spec_algorithm: SpeculativeAlgorithm = None spec_algorithm: SpeculativeAlgorithm = None
@@ -468,6 +471,7 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
can_run_dp_cuda_graph=batch.can_run_dp_cuda_graph, can_run_dp_cuda_graph=batch.can_run_dp_cuda_graph,
global_forward_mode=batch.global_forward_mode, global_forward_mode=batch.global_forward_mode,
is_prefill_only=batch.is_prefill_only, is_prefill_only=batch.is_prefill_only,
multi_item_delimiter_indices=batch.multi_item_delimiter_indices,
lora_ids=batch.lora_ids, lora_ids=batch.lora_ids,
sampling_info=batch.sampling_info, sampling_info=batch.sampling_info,
req_to_token_pool=model_runner.req_to_token_pool, req_to_token_pool=model_runner.req_to_token_pool,
+46 -29
View File
@@ -159,6 +159,13 @@ DISAGG_TRANSFER_BACKEND_CHOICES = ["mooncake", "nixl", "ascend", "fake", "mori"]
GRAMMAR_BACKEND_CHOICES = ["xgrammar", "outlines", "llguidance", "none"] GRAMMAR_BACKEND_CHOICES = ["xgrammar", "outlines", "llguidance", "none"]
# Placeholder token inserted between items in Multi-Item Scoring sequences:
# query<delim>item1<delim>item2<delim>... Positions are pre-computed from item
# lengths (multi_item_delimiter_indices); the token only exists for FlashInfer
# attention mask compat and logprob column indexing. Will be removed once the
# attention backend supports position-only MIS.
MIS_DELIMITER_TOKEN_ID = 9999
MOE_RUNNER_BACKEND_CHOICES = [ MOE_RUNNER_BACKEND_CHOICES = [
"auto", "auto",
"deep_gemm", "deep_gemm",
@@ -601,10 +608,11 @@ class ServerArgs:
offload_mode: str = "cpu" offload_mode: str = "cpu"
# Scoring configuration # Scoring configuration
# Delimiter token ID used to combine Query and Items into a single sequence for multi-item scoring. # Enable Multi-Item Scoring optimization. Combines query and multiple items
# Format: Query<delimiter>Item1<delimiter>Item2<delimiter>... # into a single sequence for efficient batch processing. Item boundaries are
# This enables efficient batch processing of multiple items against a single query. # determined by pre-computed delimiter indices (from item lengths), not by the
multi_item_scoring_delimiter: Optional[Union[int]] = None # placeholder token. See MIS_DELIMITER_TOKEN_ID for details.
enable_mis: bool = False
# Optimization/debug options # Optimization/debug options
disable_radix_cache: bool = False disable_radix_cache: bool = False
@@ -800,9 +808,6 @@ class ServerArgs:
# Handle piecewise CUDA graph. # Handle piecewise CUDA graph.
self._handle_piecewise_cuda_graph() self._handle_piecewise_cuda_graph()
# Handle multi-item scoring constraints.
self._handle_multi_item_scoring()
# Get GPU memory capacity, which is a common dependency for several configuration steps. # Get GPU memory capacity, which is a common dependency for several configuration steps.
gpu_mem = get_device_memory_capacity(self.device) gpu_mem = get_device_memory_capacity(self.device)
@@ -823,6 +828,10 @@ class ServerArgs:
self._handle_nccl_pre_warm() self._handle_nccl_pre_warm()
self._handle_grammar_backend() self._handle_grammar_backend()
# Handle multi-item scoring constraints. Must run after the above so
# the final attention backend and chunked_prefill_size are in effect.
self._handle_multi_item_scoring()
# Handle Hicache settings. # Handle Hicache settings.
self._handle_hicache() self._handle_hicache()
@@ -1227,20 +1236,36 @@ class ServerArgs:
self.disable_piecewise_cuda_graph = True self.disable_piecewise_cuda_graph = True
def _handle_multi_item_scoring(self): def _handle_multi_item_scoring(self):
"""Disable CUDA graphs when multi-item scoring delimiter is set. """Setup and validate multi-item scoring constraints.
The padded static input_ids buffer used by CUDA graph replay causes Auto-disables settings incompatible with MIS mechanics (CUDA graph,
spurious delimiter matches in score_and_pool's MIS path. radix cache, chunked prefill). Asserts on attention backend since
changing it silently could surprise users who intentionally picked
a non-flashinfer backend.
""" """
if self.multi_item_scoring_delimiter is None: if not self.enable_mis:
return return
if not self.disable_cuda_graph: if not self.disable_cuda_graph:
logger.warning( logger.warning("CUDA graph is disabled because --enable-mis is set.")
"CUDA graph is disabled because --multi-item-scoring-delimiter is set."
)
self.disable_cuda_graph = True self.disable_cuda_graph = True
self.disable_piecewise_cuda_graph = True self.disable_piecewise_cuda_graph = True
if not self.disable_radix_cache:
logger.warning("Radix cache is disabled because --enable-mis is set.")
self.disable_radix_cache = True
if self.chunked_prefill_size != -1:
logger.warning("Chunked prefill is disabled because --enable-mis is set.")
self.chunked_prefill_size = -1
prefill_backend, decode_backend = self.get_attention_backends()
assert prefill_backend == "flashinfer" and decode_backend == "flashinfer", (
"Multi-item scoring requires flashinfer attention backend for custom attention mask support. "
f"Please set --attention-backend flashinfer when using --enable-mis. "
f"Current backends: prefill={prefill_backend}, decode={decode_backend}"
)
def _handle_gpu_memory_settings(self, gpu_mem): def _handle_gpu_memory_settings(self, gpu_mem):
""" """
Configure GPU memory-dependent settings including Configure GPU memory-dependent settings including
@@ -5739,10 +5764,13 @@ class ServerArgs:
# Args for multi-item-scoring # Args for multi-item-scoring
parser.add_argument( parser.add_argument(
"--multi-item-scoring-delimiter", "--enable-mis",
type=int, action="store_true",
default=ServerArgs.multi_item_scoring_delimiter, default=ServerArgs.enable_mis,
help="Delimiter token ID for multi-item scoring. Used to combine Query and Items into a single sequence: Query<delimiter>Item1<delimiter>Item2<delimiter>... This enables efficient batch processing of multiple items against a single query.", help="Enable Multi-Item Scoring optimization. Combines query and multiple items "
"into a single sequence for efficient batch processing. "
"Requires --attention-backend flashinfer; auto-disables CUDA graph, "
"radix cache, and chunked prefill.",
) )
# Optimization/debug options # Optimization/debug options
@@ -6612,17 +6640,6 @@ class ServerArgs:
"--default-priority-value has no effect without --enable-priority-scheduling" "--default-priority-value has no effect without --enable-priority-scheduling"
) )
# Check multi-item scoring
if self.multi_item_scoring_delimiter is not None:
assert self.disable_radix_cache, (
"Multi-item scoring requires radix cache to be disabled. "
"Please set --disable-radix-cache when using --multi-item-scoring-delimiter."
)
assert self.chunked_prefill_size == -1, (
"Multi-item scoring requires chunked prefill to be disabled. "
"Please set --chunked-prefill-size -1 when using --multi-item-scoring-delimiter."
)
# Check hisparse # Check hisparse
if self.enable_hisparse: if self.enable_hisparse:
from sglang.srt.configs.model_config import is_deepseek_nsa from sglang.srt.configs.model_config import is_deepseek_nsa
@@ -20,6 +20,7 @@ from sglang.srt.managers.tokenizer_manager import TokenizerManager
from sglang.srt.managers.tokenizer_manager_score_mixin import ( from sglang.srt.managers.tokenizer_manager_score_mixin import (
TokenizerManagerScoreMixin, TokenizerManagerScoreMixin,
) )
from sglang.srt.server_args import MIS_DELIMITER_TOKEN_ID
from sglang.test.ci.ci_register import register_cuda_ci from sglang.test.ci.ci_register import register_cuda_ci
from sglang.test.test_utils import CustomTestCase from sglang.test.test_utils import CustomTestCase
@@ -204,16 +205,15 @@ class TestEmbeddingReqInputEmbedOverride(CustomTestCase):
class _FakeServerArgs: class _FakeServerArgs:
"""Minimal stub for server_args.""" """Minimal stub for server_args."""
def __init__(self, multi_item_scoring_delimiter=None): def __init__(self, enable_mis=False):
self.multi_item_scoring_delimiter = multi_item_scoring_delimiter self.enable_mis = enable_mis
class _FakeMixin(TokenizerManagerScoreMixin): class _FakeMixin(TokenizerManagerScoreMixin):
"""Minimal stub to call mixin methods without a full TokenizerManager.""" """Minimal stub to call mixin methods without a full TokenizerManager."""
def __init__(self, delimiter=None): def __init__(self, enable_mis=False):
self.server_args = _FakeServerArgs(delimiter) self.server_args = _FakeServerArgs(enable_mis)
self.multi_item_delimiter_text = None
self.tokenizer = None self.tokenizer = None
self.is_generation = True self.is_generation = True
@@ -334,17 +334,17 @@ class TestResolveEmbedOverridesForRequest(CustomTestCase):
# Score mixin: _build_token_id_inputs # Score mixin: _build_token_id_inputs
# ======================================================================== # ========================================================================
DELIM_TOKEN = 99 DELIM_TOKEN = MIS_DELIMITER_TOKEN_ID
class TestBuildTokenIdInputs(CustomTestCase): class TestBuildTokenIdInputs(CustomTestCase):
def setUp(self): def setUp(self):
self.mixin = _FakeMixin(delimiter=DELIM_TOKEN) self.mixin = _FakeMixin(enable_mis=True)
# --- single-item mode, no embeds --- # --- single-item mode, no embeds ---
def test_single_item_no_embeds(self): def test_single_item_no_embeds(self):
_, input_ids, injection = self.mixin._build_token_id_inputs( _, input_ids, positional_embed_overrides, _ = self.mixin._build_token_id_inputs(
query=[1, 2], query=[1, 2],
items=[[3, 4], [5, 6]], items=[[3, 4], [5, 6]],
item_first=False, item_first=False,
@@ -354,10 +354,10 @@ class TestBuildTokenIdInputs(CustomTestCase):
item_embed_overrides=None, item_embed_overrides=None,
) )
self.assertEqual(input_ids, [[1, 2, 3, 4], [1, 2, 5, 6]]) self.assertEqual(input_ids, [[1, 2, 3, 4], [1, 2, 5, 6]])
self.assertIsNone(injection) self.assertIsNone(positional_embed_overrides)
def test_single_item_no_embeds_item_first(self): def test_single_item_no_embeds_item_first(self):
_, input_ids, injection = self.mixin._build_token_id_inputs( _, input_ids, positional_embed_overrides, _ = self.mixin._build_token_id_inputs(
query=[1, 2], query=[1, 2],
items=[[3, 4]], items=[[3, 4]],
item_first=True, item_first=True,
@@ -367,12 +367,12 @@ class TestBuildTokenIdInputs(CustomTestCase):
item_embed_overrides=None, item_embed_overrides=None,
) )
self.assertEqual(input_ids, [[3, 4, 1, 2]]) self.assertEqual(input_ids, [[3, 4, 1, 2]])
self.assertIsNone(injection) self.assertIsNone(positional_embed_overrides)
# --- multi-item mode, no embeds --- # --- multi-item mode, no embeds ---
def test_multi_item_no_embeds(self): def test_multi_item_no_embeds(self):
_, input_ids, injection = self.mixin._build_token_id_inputs( _, input_ids, positional_embed_overrides, _ = self.mixin._build_token_id_inputs(
query=[1, 2], query=[1, 2],
items=[[3, 4], [5, 6]], items=[[3, 4], [5, 6]],
item_first=False, item_first=False,
@@ -385,13 +385,13 @@ class TestBuildTokenIdInputs(CustomTestCase):
self.assertEqual( self.assertEqual(
input_ids, [[1, 2, DELIM_TOKEN, 3, 4, DELIM_TOKEN, 5, 6, DELIM_TOKEN]] input_ids, [[1, 2, DELIM_TOKEN, 3, 4, DELIM_TOKEN, 5, 6, DELIM_TOKEN]]
) )
self.assertIsNone(injection) self.assertIsNone(positional_embed_overrides)
# --- single-item mode, with embeds --- # --- single-item mode, with embeds ---
def test_single_item_query_embeds(self): def test_single_item_query_embeds(self):
"""Query placeholder overrides are resolved per item.""" """Query placeholder overrides are resolved per item."""
_, input_ids, injection = self.mixin._build_token_id_inputs( _, input_ids, positional_embed_overrides, _ = self.mixin._build_token_id_inputs(
query=[50, 10], query=[50, 10],
items=[[20, 30], [40, 50]], items=[[20, 30], [40, 50]],
item_first=False, item_first=False,
@@ -401,15 +401,15 @@ class TestBuildTokenIdInputs(CustomTestCase):
item_embed_overrides=None, item_embed_overrides=None,
) )
self.assertEqual(input_ids, [[50, 10, 20, 30], [50, 10, 40, 50]]) self.assertEqual(input_ids, [[50, 10, 20, 30], [50, 10, 40, 50]])
self.assertIsNotNone(injection) self.assertIsNotNone(positional_embed_overrides)
self.assertEqual(len(injection), 2) self.assertEqual(len(positional_embed_overrides), 2)
# Each item gets its own PositionalEmbeds with query override at pos 0 # Each item gets its own PositionalEmbeds with query override at pos 0
self.assertEqual(injection[0].positions, [0]) self.assertEqual(positional_embed_overrides[0].positions, [0])
self.assertEqual(injection[1].positions, [0]) self.assertEqual(positional_embed_overrides[1].positions, [0])
def test_single_item_item_embeds(self): def test_single_item_item_embeds(self):
"""Per-item overrides with correct position offsets.""" """Per-item overrides with correct position offsets."""
_, input_ids, injection = self.mixin._build_token_id_inputs( _, input_ids, positional_embed_overrides, _ = self.mixin._build_token_id_inputs(
query=[10, 20], query=[10, 20],
items=[[50, 30]], items=[[50, 30]],
item_first=False, item_first=False,
@@ -419,13 +419,13 @@ class TestBuildTokenIdInputs(CustomTestCase):
item_embed_overrides=[[_vec(2)]], item_embed_overrides=[[_vec(2)]],
) )
self.assertEqual(input_ids, [[10, 20, 50, 30]]) self.assertEqual(input_ids, [[10, 20, 50, 30]])
self.assertIsNotNone(injection) self.assertIsNotNone(positional_embed_overrides)
# item placeholder at index 0 of item, offset by query length 2 # item placeholder at index 0 of item, offset by query length 2
self.assertEqual(injection[0].positions, [2]) self.assertEqual(positional_embed_overrides[0].positions, [2])
def test_single_item_no_override_positions_returns_none_injection(self): def test_single_item_no_override_positions_returns_none_injection(self):
"""When no items have placeholders, injection should be None.""" """When no items have placeholders, positional_embed_overrides should be None."""
_, input_ids, injection = self.mixin._build_token_id_inputs( _, input_ids, positional_embed_overrides, _ = self.mixin._build_token_id_inputs(
query=[10, 20], query=[10, 20],
items=[[30, 40]], items=[[30, 40]],
item_first=False, item_first=False,
@@ -434,11 +434,11 @@ class TestBuildTokenIdInputs(CustomTestCase):
query_embed_overrides=None, query_embed_overrides=None,
item_embed_overrides=[None], item_embed_overrides=[None],
) )
self.assertIsNone(injection) self.assertIsNone(positional_embed_overrides)
def test_single_item_query_and_item_embeds(self): def test_single_item_query_and_item_embeds(self):
"""Single-item mode with both query and item overrides in one request.""" """Single-item mode with both query and item overrides in one request."""
_, input_ids, injection = self.mixin._build_token_id_inputs( _, input_ids, positional_embed_overrides, _ = self.mixin._build_token_id_inputs(
query=[50, 10], query=[50, 10],
items=[[20, 50]], items=[[20, 50]],
item_first=False, item_first=False,
@@ -448,15 +448,15 @@ class TestBuildTokenIdInputs(CustomTestCase):
item_embed_overrides=[[_vec(2)]], item_embed_overrides=[[_vec(2)]],
) )
self.assertEqual(input_ids, [[50, 10, 20, 50]]) self.assertEqual(input_ids, [[50, 10, 20, 50]])
self.assertIsNotNone(injection) self.assertIsNotNone(positional_embed_overrides)
pe = injection[0] pe = positional_embed_overrides[0]
# query override at pos 0, item override at pos 3 (query_len=2 + idx=1) # query override at pos 0, item override at pos 3 (query_len=2 + idx=1)
self.assertEqual(pe.positions, [0, 3]) self.assertEqual(pe.positions, [0, 3])
self.assertEqual(pe.embeds.shape, (2, HIDDEN_DIM)) self.assertEqual(pe.embeds.shape, (2, HIDDEN_DIM))
def test_single_item_empty_query(self): def test_single_item_empty_query(self):
"""Empty query with item-only overrides (valid from score_prompts).""" """Empty query with item-only overrides (valid from score_prompts)."""
_, input_ids, injection = self.mixin._build_token_id_inputs( _, input_ids, positional_embed_overrides, _ = self.mixin._build_token_id_inputs(
query=[], query=[],
items=[[50, 10]], items=[[50, 10]],
item_first=False, item_first=False,
@@ -466,15 +466,15 @@ class TestBuildTokenIdInputs(CustomTestCase):
item_embed_overrides=[[_vec(1)]], item_embed_overrides=[[_vec(1)]],
) )
self.assertEqual(input_ids, [[50, 10]]) self.assertEqual(input_ids, [[50, 10]])
self.assertIsNotNone(injection) self.assertIsNotNone(positional_embed_overrides)
# item placeholder at absolute pos 0 (offset=len([])=0) # item placeholder at absolute pos 0 (offset=len([])=0)
self.assertEqual(injection[0].positions, [0]) self.assertEqual(positional_embed_overrides[0].positions, [0])
# --- multi-item mode, with embeds --- # --- multi-item mode, with embeds ---
def test_multi_item_with_query_and_item_embeds(self): def test_multi_item_with_query_and_item_embeds(self):
"""Multi-item mode resolves query overrides once and item overrides per item.""" """Multi-item mode resolves query overrides once and item overrides per item."""
_, input_ids, injection = self.mixin._build_token_id_inputs( _, input_ids, positional_embed_overrides, _ = self.mixin._build_token_id_inputs(
query=[50, 10], query=[50, 10],
items=[[20, 50], [30, 40]], items=[[20, 50], [30, 40]],
item_first=False, item_first=False,
@@ -483,13 +483,13 @@ class TestBuildTokenIdInputs(CustomTestCase):
query_embed_overrides=[_vec(1)], query_embed_overrides=[_vec(1)],
item_embed_overrides=[[_vec(2)], None], item_embed_overrides=[[_vec(2)], None],
) )
# query<D>item1<D>item2<D> = [50,10, 99, 20,50, 99, 30,40, 99] # query<D>item1<D>item2<D> = [50,10, DELIM, 20,50, DELIM, 30,40, DELIM]
self.assertEqual(len(input_ids), 1) self.assertEqual(len(input_ids), 1)
self.assertIsNotNone(injection) self.assertIsNotNone(positional_embed_overrides)
self.assertEqual( self.assertEqual(
len(injection), 1 len(positional_embed_overrides), 1
) # single PositionalEmbeds for combined sequence ) # single PositionalEmbeds for combined sequence
pe = injection[0] pe = positional_embed_overrides[0]
# query override at pos 0, item[0] override at pos 4 (query_len=2 + delim=1 + idx=1) # query override at pos 0, item[0] override at pos 4 (query_len=2 + delim=1 + idx=1)
self.assertIn(0, pe.positions) self.assertIn(0, pe.positions)
self.assertIn(4, pe.positions) self.assertIn(4, pe.positions)
@@ -0,0 +1,599 @@
"""Tests for the Multi-Item Scoring (MIS) optimization.
MIS is a server-side optimization enabled via --enable-mis that batches
multiple items into a single forward pass using delimiter tokens (token ID 9999).
This is different from batch scoring (multiple items in one API call) which
processes items as separate requests.
The key difference:
- Batch scoring: N items -> N separate forward passes
- MIS optimization: N items -> 1 forward pass with delimiter-separated items
These tests ensure the MIS optimization produces correct results and catches
bugs in tensor shape handling (e.g., 2D tensors [num_delimiters, num_label_tokens]).
"""
import asyncio
import os
import unittest
import torch
from transformers import AutoConfig, AutoTokenizer
from sglang.srt.entrypoints.engine import Engine
from sglang.test.ci.ci_register import register_cuda_ci
from sglang.test.test_utils import (
DEFAULT_SMALL_MODEL_NAME_FOR_TEST,
CustomTestCase,
)
register_cuda_ci(est_time=240, suite="stage-b-test-1-gpu-small")
TEST_MODEL_NAME = os.environ.get("TEST_MODEL_NAME", DEFAULT_SMALL_MODEL_NAME_FOR_TEST)
TEST_CLASSIFICATION_BASE_MODEL = os.environ.get(
"TEST_CLASSIFICATION_BASE_MODEL",
"tomaarsen/Qwen3-Reranker-0.6B-seq-cls",
)
_CLS_NUM_LABELS = AutoConfig.from_pretrained(TEST_CLASSIFICATION_BASE_MODEL).num_labels
class TestMISServerArgsValidation(unittest.TestCase):
"""Test ServerArgs defaults for MIS mode."""
def test_enable_mis_default(self):
"""Test that enable_mis defaults to False."""
from sglang.srt.server_args import ServerArgs
self.assertEqual(ServerArgs.enable_mis, False)
class TestMultiItemScoringOptimization(CustomTestCase):
"""Test the Multi-Item Scoring (MIS) optimization with generation models."""
@classmethod
def setUpClass(cls):
cls.engine = Engine(
model_path=TEST_MODEL_NAME,
disable_radix_cache=True,
chunked_prefill_size=-1,
enable_mis=True,
attention_backend="flashinfer",
mem_fraction_static=0.15,
)
cls.non_mis_engine = Engine(
model_path=TEST_MODEL_NAME,
disable_radix_cache=True,
chunked_prefill_size=-1,
mem_fraction_static=0.15,
)
@classmethod
def tearDownClass(cls):
if cls.engine is not None:
cls.engine.shutdown()
if cls.non_mis_engine is not None:
cls.non_mis_engine.shutdown()
torch.cuda.empty_cache()
def test_mis_basic(self):
"""Test basic MIS: correct shapes, valid probabilities."""
query = "Rate each option:"
items = ["Option A", "Option B", "Option C"]
label_token_ids = [9454, 2753] # "Yes" and "No" tokens
scores = self.engine.score(
query=query,
items=items,
label_token_ids=label_token_ids,
apply_softmax=True,
).scores
self.assertEqual(len(scores), len(items))
for i, score_list in enumerate(scores):
self.assertEqual(len(score_list), len(label_token_ids))
self.assertAlmostEqual(sum(score_list), 1.0, places=5)
for score in score_list:
self.assertGreaterEqual(score, 0)
self.assertLessEqual(score, 1)
def test_mis_consistency_with_single_item(self):
"""MIS with one item should match non-MIS scoring closely."""
query = "Is this a fact?\n"
items = [" The sun rises in the east"]
label_token_ids = [9454, 2753]
mis_scores = self.engine.score(
query=query,
items=items,
label_token_ids=label_token_ids,
apply_softmax=True,
).scores
non_mis_scores = self.non_mis_engine.score(
query=query,
items=items,
label_token_ids=label_token_ids,
apply_softmax=True,
).scores
self.assertEqual(len(mis_scores), 1)
self.assertEqual(len(non_mis_scores), 1)
for j, (m, n) in enumerate(zip(mis_scores[0], non_mis_scores[0])):
relative_diff = abs(m - n) / max(abs(n), 1e-6)
self.assertLess(
relative_diff,
0.08,
msg=f"label {j}: MIS={m} vs non-MIS={n} (diff: {relative_diff:.3f})",
)
def test_mis_empty_query(self):
"""MIS with empty query — delimiter indices start at position 0."""
items = ["alpha", "beta"]
label_token_ids = [9454, 2753]
scores = self.engine.score(
query="",
items=items,
label_token_ids=label_token_ids,
apply_softmax=True,
).scores
self.assertEqual(len(scores), len(items))
for score_list in scores:
self.assertEqual(len(score_list), len(label_token_ids))
self.assertAlmostEqual(sum(score_list), 1.0, places=5)
class TestMultiItemScoringClassification(CustomTestCase):
"""Test MIS with classification models.
Uses a pre-trained Qwen3ForSequenceClassification model so that the
classification head weights are deterministic across Engine instances.
"""
NUM_LABELS = _CLS_NUM_LABELS
def setUp(self):
self.engine = Engine(
model_path=TEST_CLASSIFICATION_BASE_MODEL,
disable_radix_cache=True,
chunked_prefill_size=-1,
enable_mis=True,
attention_backend="flashinfer",
mem_fraction_static=0.15,
)
def tearDown(self):
if self.engine is not None:
self.engine.shutdown()
torch.cuda.empty_cache()
def test_classification_mis_basic(self):
"""Classification MIS: correct shapes, valid softmax probabilities."""
query = "Rate each option:"
items = ["Option A", "Option B", "Option C"]
scores = self.engine.score(query=query, 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)
self.assertAlmostEqual(sum(score_list), 1.0, places=5)
for score in score_list:
self.assertGreaterEqual(score, 0)
self.assertLessEqual(score, 1)
def test_classification_mis_tokenized_input(self):
"""Classification MIS with pre-tokenized query and items."""
tokenizer = AutoTokenizer.from_pretrained(TEST_CLASSIFICATION_BASE_MODEL)
query_ids = tokenizer.encode("Rate each option:", add_special_tokens=False)
items_ids = [
tokenizer.encode(item, add_special_tokens=False)
for item in ["Option A", "Option B"]
]
scores = self.engine.score(
query=query_ids, items=items_ids, apply_softmax=True
).scores
self.assertEqual(len(scores), len(items_ids))
for score_list in scores:
self.assertEqual(len(score_list), self.NUM_LABELS)
self.assertAlmostEqual(sum(score_list), 1.0, places=5)
def test_classification_non_mis_fallback(self):
"""Classification model works correctly without --enable-mis."""
non_mis_engine = Engine(
model_path=TEST_CLASSIFICATION_BASE_MODEL,
disable_radix_cache=True,
mem_fraction_static=0.15,
)
try:
scores = non_mis_engine.score(
query="Test:", items=["A", "B"], apply_softmax=True
).scores
self.assertEqual(len(scores), 2)
for score_list in scores:
self.assertEqual(len(score_list), self.NUM_LABELS)
self.assertAlmostEqual(sum(score_list), 1.0, places=5)
finally:
non_mis_engine.shutdown()
torch.cuda.empty_cache()
class TestMultiItemScoringParity(CustomTestCase):
"""Test that MIS produces the same results as single-item scoring."""
@classmethod
def setUpClass(cls):
cls.engine_single = Engine(
model_path=TEST_MODEL_NAME,
disable_radix_cache=True,
log_level="error",
mem_fraction_static=0.15,
)
cls.engine_mis = Engine(
model_path=TEST_MODEL_NAME,
disable_radix_cache=True,
chunked_prefill_size=-1,
log_level="error",
enable_mis=True,
attention_backend="flashinfer",
mem_fraction_static=0.15,
)
@classmethod
def tearDownClass(cls):
if cls.engine_single is not None:
cls.engine_single.shutdown()
if cls.engine_mis is not None:
cls.engine_mis.shutdown()
torch.cuda.empty_cache()
def _compare_scores(
self, query, items, label_token_ids=None, apply_softmax=True, test_name=""
):
"""Compare MIS vs single-item scoring results."""
single_scores = self.engine_single.score(
query=query,
items=items,
label_token_ids=label_token_ids,
apply_softmax=apply_softmax,
).scores
mis_scores = self.engine_mis.score(
query=query,
items=items,
label_token_ids=label_token_ids,
apply_softmax=apply_softmax,
).scores
self.assertEqual(
len(mis_scores), len(single_scores), f"{test_name}: count mismatch"
)
for i, (ms, ss) in enumerate(zip(mis_scores, single_scores)):
self.assertEqual(len(ms), len(ss), f"{test_name}: item {i} length mismatch")
for j, (m, s) in enumerate(zip(ms, ss)):
self.assertAlmostEqual(
m,
s,
places=1,
msg=f"{test_name}: item {i} label {j}: MIS={m} vs single={s}",
)
def test_parity_basic(self):
tokenizer = AutoTokenizer.from_pretrained(TEST_MODEL_NAME)
query = "Rate this option:"
items = [" Option A", " Option B", " Option C"]
labels = [" good", " bad"]
label_ids = [tokenizer.encode(lb, add_special_tokens=False)[0] for lb in labels]
self._compare_scores(query, items, label_ids, test_name="basic")
def test_parity_tokenized_inputs(self):
tokenizer = AutoTokenizer.from_pretrained(TEST_MODEL_NAME)
query = "Rate this option:"
items = [" Option X", " Option Y"]
labels = [" good", " bad"]
query_ids = tokenizer.encode(query, add_special_tokens=False)
items_ids = [tokenizer.encode(i, add_special_tokens=False) for i in items]
label_ids = [tokenizer.encode(lb, add_special_tokens=False)[0] for lb in labels]
self._compare_scores(query_ids, items_ids, label_ids, test_name="tokenized")
def test_parity_without_softmax(self):
tokenizer = AutoTokenizer.from_pretrained(TEST_MODEL_NAME)
query = "The weather today is"
items = [" sunny", " cloudy", " rainy"]
labels = [" nice", " bad"]
label_ids = [tokenizer.encode(lb, add_special_tokens=False)[0] for lb in labels]
self._compare_scores(
query, items, label_ids, apply_softmax=False, test_name="no_softmax"
)
def test_parity_many_items(self):
tokenizer = AutoTokenizer.from_pretrained(TEST_MODEL_NAME)
query = "Rate this option from 1 to 5:"
items = [f" Option {i}" for i in range(10)]
labels = [" 1", " 2", " 3", " 4", " 5"]
label_ids = [tokenizer.encode(lb, add_special_tokens=False)[0] for lb in labels]
self._compare_scores(query, items, label_ids, test_name="many_items")
class TestMultiItemScoringClassificationParity(CustomTestCase):
"""Test that MIS multi-item batching matches single-item MIS scoring.
Both paths use the MIS engine (with delimiter tokens in the attention
context). The reference scores each item individually so each gets its
own forward pass; the batched path packs all items into one pass.
This isolates the MIS batching logic from the delimiter-presence effect.
"""
NUM_LABELS = _CLS_NUM_LABELS
@classmethod
def setUpClass(cls):
cls.engine = Engine(
model_path=TEST_CLASSIFICATION_BASE_MODEL,
disable_radix_cache=True,
chunked_prefill_size=-1,
enable_mis=True,
attention_backend="flashinfer",
mem_fraction_static=0.15,
)
@classmethod
def tearDownClass(cls):
if cls.engine is not None:
cls.engine.shutdown()
torch.cuda.empty_cache()
def _compare_scores(self, query, items, apply_softmax=True, test_name=""):
"""Compare MIS batched vs MIS single-item scoring results."""
single_scores = []
for item in items:
result = self.engine.score(
query=query,
items=[item],
apply_softmax=apply_softmax,
).scores
single_scores.append(result[0])
batched_scores = self.engine.score(
query=query,
items=items,
apply_softmax=apply_softmax,
).scores
self.assertEqual(
len(batched_scores), len(single_scores), f"{test_name}: count mismatch"
)
for i, (bs, ss) in enumerate(zip(batched_scores, single_scores)):
self.assertEqual(len(bs), len(ss), f"{test_name}: item {i} length mismatch")
for j, (b, s) in enumerate(zip(bs, ss)):
self.assertAlmostEqual(
b,
s,
places=1,
msg=f"{test_name}: item {i} label {j}: batched={b} vs single={s}",
)
def test_parity_basic(self):
query = "Rate this option:"
items = [" Option A", " Option B", " Option C"]
self._compare_scores(query, items, test_name="cls_basic")
def test_parity_tokenized_inputs(self):
tokenizer = AutoTokenizer.from_pretrained(TEST_CLASSIFICATION_BASE_MODEL)
query_ids = tokenizer.encode("Rate this option:", add_special_tokens=False)
items_ids = [
tokenizer.encode(item, add_special_tokens=False)
for item in [" Option X", " Option Y"]
]
self._compare_scores(query_ids, items_ids, test_name="cls_tokenized")
def test_parity_without_softmax(self):
query = "The weather today is"
items = [" sunny", " cloudy", " rainy"]
self._compare_scores(
query, items, apply_softmax=False, test_name="cls_no_softmax"
)
def test_parity_many_items(self):
query = "Classify this option:"
items = [f" Option {i}" for i in range(10)]
self._compare_scores(query, items, test_name="cls_many_items")
class TestMultiItemScoringClassificationMISvsNonMIS(CustomTestCase):
"""Test that MIS single-item approximates non-MIS single-item.
The MIS path inserts delimiter tokens into the attention context,
which slightly perturbs hidden states. After softmax the scores
should still be close. Uses places=1 (±0.05) tolerance.
Runs as a separate class so each engine is created and destroyed
independently to avoid GPU OOM.
"""
def test_mis_single_vs_non_mis(self):
non_mis_engine = Engine(
model_path=TEST_CLASSIFICATION_BASE_MODEL,
disable_radix_cache=True,
mem_fraction_static=0.15,
)
try:
query = "Rate this option:"
items = [" Option A", " Option B", " Option C"]
non_mis_scores = non_mis_engine.score(
query=query,
items=items,
apply_softmax=True,
).scores
finally:
non_mis_engine.shutdown()
torch.cuda.empty_cache()
mis_engine = Engine(
model_path=TEST_CLASSIFICATION_BASE_MODEL,
disable_radix_cache=True,
chunked_prefill_size=-1,
enable_mis=True,
attention_backend="flashinfer",
mem_fraction_static=0.15,
)
try:
mis_scores = mis_engine.score(
query=query,
items=items,
apply_softmax=True,
).scores
finally:
mis_engine.shutdown()
torch.cuda.empty_cache()
self.assertEqual(len(mis_scores), len(non_mis_scores))
for i, (ms, ns) in enumerate(zip(mis_scores, non_mis_scores)):
self.assertEqual(len(ms), len(ns))
for j, (m, n) in enumerate(zip(ms, ns)):
self.assertAlmostEqual(
m,
n,
places=1,
msg=f"item {i} label {j}: MIS={m} vs non-MIS={n}",
)
class TestMultiItemScoringClassificationAdvanced(CustomTestCase):
"""Advanced MIS tests for classification models: score distinctness,
determinism, and concurrent request handling."""
NUM_LABELS = _CLS_NUM_LABELS
@classmethod
def setUpClass(cls):
cls.engine = Engine(
model_path=TEST_CLASSIFICATION_BASE_MODEL,
disable_radix_cache=True,
chunked_prefill_size=-1,
enable_mis=True,
attention_backend="flashinfer",
mem_fraction_static=0.15,
)
@classmethod
def tearDownClass(cls):
if cls.engine is not None:
cls.engine.shutdown()
torch.cuda.empty_cache()
def test_items_produce_distinct_scores(self):
"""Different items must produce different score vectors.
Core regression test: before the delimiter-index fix, all items got
identical scores because the MIS attention mask only let delimiter
tokens attend to the query prefix.
"""
query = "Rate each option:"
items = [
"Option A is about cats",
"Option B is about dogs",
"Option C is about fish",
]
scores = self.engine.score(query=query, 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 broken. Scores: {scores[0]}",
)
def test_many_items_distinct(self):
"""Stress test: 15 items should not all produce identical scores."""
query = "Classify each city:"
items = [f"City {i}" for i in range(15)]
scores = self.engine.score(query=query, items=items).scores
self.assertEqual(len(scores), len(items))
for score_list in scores:
self.assertEqual(len(score_list), self.NUM_LABELS)
unique_count = len({tuple(s) for s in scores})
self.assertGreater(unique_count, 1, "All 15 items returned identical scores")
def test_deterministic(self):
"""Identical requests should return identical scores."""
query = "Evaluate:"
items = ["alpha", "beta", "gamma"]
scores1 = self.engine.score(query=query, items=items).scores
scores2 = self.engine.score(query=query, items=items).scores
self.assertEqual(
scores1, scores2, "Identical inputs must produce identical scores"
)
def test_concurrent_requests(self):
"""Concurrent MIS requests must produce the same scores as sequential.
Runs each request sequentially to get baseline scores, then runs all
concurrently and asserts the results match. This catches cross-request
contamination when multiple MIS requests share a GPU batch.
"""
test_cases = [
{"query": "Is this a fruit?", "items": ["apple", "car", "banana"]},
{"query": "Is this an animal?", "items": ["dog", "table"]},
{
"query": "Is this a country?",
"items": ["France", "pizza", "Japan", "chair"],
},
{"query": "Is this a color?", "items": ["red"]},
]
# Sequential baseline
sequential_scores = []
for tc in test_cases:
result = self.engine.score(query=tc["query"], items=tc["items"])
sequential_scores.append(result.scores)
# Concurrent execution
async def _gather():
return await asyncio.gather(
*(
self.engine.async_score(query=tc["query"], items=tc["items"])
for tc in test_cases
)
)
concurrent_results = self.engine.loop.run_until_complete(_gather())
for idx, (tc, seq_scores, conc_result) in enumerate(
zip(test_cases, sequential_scores, concurrent_results)
):
conc_scores = conc_result.scores
self.assertEqual(
len(conc_scores),
len(seq_scores),
f"Case {idx}: count mismatch",
)
for i, (cs, ss) in enumerate(zip(conc_scores, seq_scores)):
self.assertEqual(
len(cs),
len(ss),
f"Case {idx} item {i}: label count mismatch",
)
for j, (c, s) in enumerate(zip(cs, ss)):
self.assertAlmostEqual(
c,
s,
places=1,
msg=f"Case {idx} item {i} label {j}: "
f"concurrent={c} vs sequential={s}",
)
if __name__ == "__main__":
unittest.main()
@@ -30,7 +30,6 @@ from sglang.test.test_utils import (
register_cuda_ci(est_time=240, suite="stage-b-test-1-gpu-small") register_cuda_ci(est_time=240, suite="stage-b-test-1-gpu-small")
_SEQCLS_MODEL = "Qwen/Qwen3-0.6B" _SEQCLS_MODEL = "Qwen/Qwen3-0.6B"
_QWEN3_EOT_TOKEN_ID = 151643
_CAUSAL_LM_MODEL = DEFAULT_SMALL_MODEL_NAME_FOR_TEST _CAUSAL_LM_MODEL = DEFAULT_SMALL_MODEL_NAME_FOR_TEST
_NUM_LABELS = 4 _NUM_LABELS = 4
@@ -197,7 +196,7 @@ class TestPooledHiddenStatesMISEngine(CustomTestCase):
model_path=_SEQCLS_MODEL, model_path=_SEQCLS_MODEL,
disable_radix_cache=True, disable_radix_cache=True,
chunked_prefill_size=-1, chunked_prefill_size=-1,
multi_item_scoring_delimiter=_QWEN3_EOT_TOKEN_ID, enable_mis=True,
json_model_override_args=json.dumps( json_model_override_args=json.dumps(
{ {
"architectures": ["Qwen3ForSequenceClassification"], "architectures": ["Qwen3ForSequenceClassification"],
+11 -11
View File
@@ -3,15 +3,16 @@
Two test classes, each with its own server instance: Two test classes, each with its own server instance:
TestCausalLMScoringHTTP — basic endpoint: schema defaults, response TestCausalLMScoringHTTP — basic endpoint: schema defaults, response
structure, error rejection (no MIS delimiter) structure, error rejection (no MIS)
TestCausalLMMISScoringHTTP — MIS mode: validates --multi-item-scoring-delimiter TestCausalLMMISScoringHTTP — MIS mode: validates --enable-mis CLI flag
CLI flag wiring and per-item output shape wiring and per-item output shape
Engine-level correctness (numerical accuracy, batching, edge cases) lives in Engine-level correctness (numerical accuracy, batching, edge cases) lives in
test_score_engine.py. These tests focus on the HTTP integration seam: test_score_engine.py. These tests focus on the HTTP integration seam:
Pydantic schema defaults, FastAPI routing, and server argument wiring. Pydantic schema defaults, FastAPI routing, and server argument wiring.
""" """
import os
import unittest import unittest
import requests import requests
@@ -28,9 +29,7 @@ from sglang.test.test_utils import (
register_cuda_ci(est_time=70, suite="stage-b-test-1-gpu-small") register_cuda_ci(est_time=70, suite="stage-b-test-1-gpu-small")
_MODEL = DEFAULT_SMALL_MODEL_NAME_FOR_TEST # Llama-3.2-1B-Instruct _MODEL = os.environ.get("TEST_MODEL_NAME", DEFAULT_SMALL_MODEL_NAME_FOR_TEST)
# <|eot_id|> for Llama-3.x Instruct — used as MIS delimiter
_LLAMA3_EOT_TOKEN_ID = 128009
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
@@ -41,7 +40,7 @@ _LLAMA3_EOT_TOKEN_ID = 128009
class TestCausalLMScoringHTTP(CustomTestCase): class TestCausalLMScoringHTTP(CustomTestCase):
"""Validates /v1/score HTTP integration — schema, defaults, and error handling. """Validates /v1/score HTTP integration — schema, defaults, and error handling.
Starts a plain CausalLM server (no --multi-item-scoring-delimiter) to test Starts a plain CausalLM server (no --enable-mis) to test
the HTTP layer in isolation: response envelope shape, the apply_softmax the HTTP layer in isolation: response envelope shape, the apply_softmax
default (False), and Pydantic validation errors on malformed input. default (False), and Pydantic validation errors on malformed input.
""" """
@@ -139,12 +138,12 @@ class TestCausalLMScoringHTTP(CustomTestCase):
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
# MIS scoring (with --multi-item-scoring-delimiter) # MIS scoring (with --enable-mis)
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
class TestCausalLMMISScoringHTTP(CustomTestCase): class TestCausalLMMISScoringHTTP(CustomTestCase):
"""Validates /v1/score with --multi-item-scoring-delimiter. """Validates /v1/score with --enable-mis.
Confirms that the CLI flag is correctly wired into ServerArgs and that the Confirms that the CLI flag is correctly wired into ServerArgs and that the
endpoint returns one probability vector per item when items are endpoint returns one probability vector per item when items are
@@ -163,8 +162,9 @@ class TestCausalLMMISScoringHTTP(CustomTestCase):
"--disable-radix-cache", "--disable-radix-cache",
"--chunked-prefill-size", "--chunked-prefill-size",
"-1", "-1",
"--multi-item-scoring-delimiter", "--enable-mis",
str(_LLAMA3_EOT_TOKEN_ID), "--attention-backend",
"flashinfer",
], ],
) )
@@ -4,14 +4,18 @@ Two model types, two scoring modes:
TestCausalLMScoring — CausalLM, single-item and batched multi-item TestCausalLMScoring — CausalLM, single-item and batched multi-item
TestSeqClsScoring — SequenceClassification, single-item mode TestSeqClsScoring — SequenceClassification, single-item mode
TestSeqClsMISScoring — SequenceClassification, MIS delimiter mode TestSeqClsMISScoring — SequenceClassification, MIS mode (--enable-mis)
TestSeqClsMISAdvancedScoring — SeqCls MIS with 12 labels (tensor shape stress)
The Engine (Python API) is the right layer for correctness testing: it The Engine (Python API) is the right layer for correctness testing: it
exercises tokenization, forward pass, pooling, and score extraction without exercises tokenization, forward pass, pooling, and score extraction without
the HTTP serialization overhead. HTTP-layer tests live in test_score_api.py. the HTTP serialization overhead. HTTP-layer tests live in test_score_api.py.
Thorough MIS tests (parity, concurrency, generation models) live in
test_multi_item_scoring.py.
""" """
import json import json
import os
import unittest import unittest
from unittest.mock import patch from unittest.mock import patch
@@ -24,10 +28,8 @@ from sglang.test.test_utils import DEFAULT_SMALL_MODEL_NAME_FOR_TEST, CustomTest
register_cuda_ci(est_time=85, suite="stage-b-test-1-gpu-small") register_cuda_ci(est_time=85, suite="stage-b-test-1-gpu-small")
_CAUSAL_LM_MODEL = DEFAULT_SMALL_MODEL_NAME_FOR_TEST # Llama-3.2-1B-Instruct _CAUSAL_LM_MODEL = os.environ.get("TEST_MODEL_NAME", DEFAULT_SMALL_MODEL_NAME_FOR_TEST)
_SEQCLS_MODEL = "Qwen/Qwen3-0.6B" # backbone; arch overridden to SeqCls below _SEQCLS_MODEL = os.environ.get("TEST_CLASSIFICATION_BASE_MODEL", "Qwen/Qwen3-0.6B")
# <|endoftext|> for Qwen3 tokenizer — used as MIS delimiter
_QWEN3_EOT_TOKEN_ID = 151643
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
@@ -368,10 +370,9 @@ class TestSeqClsScoring(CustomTestCase):
class TestSeqClsMISScoring(CustomTestCase): class TestSeqClsMISScoring(CustomTestCase):
"""SeqCls MIS: all items packed into one sequence separated by delimiter token. """SeqCls MIS: all items packed into one sequence separated by delimiter token.
score_and_pool() extracts per-item scores at delimiter positions. Uses --enable-mis which hardcodes delimiter token ID 9999.
Two sub-cases are tested: Basic pipeline correctness only — thorough MIS tests (parity,
- NUM_LABELS=2 — standard binary classification head concurrency, advanced) live in test_multi_item_scoring.py.
- NUM_LABELS=12 — stress-tests 2-D tensor indexing in score_and_pool()
""" """
NUM_LABELS = 2 NUM_LABELS = 2
@@ -382,7 +383,8 @@ class TestSeqClsMISScoring(CustomTestCase):
model_path=_SEQCLS_MODEL, model_path=_SEQCLS_MODEL,
disable_radix_cache=True, disable_radix_cache=True,
chunked_prefill_size=-1, chunked_prefill_size=-1,
multi_item_scoring_delimiter=_QWEN3_EOT_TOKEN_ID, enable_mis=True,
attention_backend="flashinfer",
json_model_override_args=json.dumps( json_model_override_args=json.dumps(
{ {
"architectures": ["Qwen3ForSequenceClassification"], "architectures": ["Qwen3ForSequenceClassification"],
@@ -432,32 +434,6 @@ class TestSeqClsMISScoring(CustomTestCase):
self.assertEqual(len(row), self.NUM_LABELS) self.assertEqual(len(row), self.NUM_LABELS)
self.assertAlmostEqual(sum(row), 1.0, places=5) self.assertAlmostEqual(sum(row), 1.0, places=5)
def test_mis_items_produce_distinct_scores(self):
"""Different items must yield different score vectors.
Catches bugs where all delimiter positions share the same pooled
hidden state (e.g. off-by-one in score_and_pool indexing).
"""
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))
self.assertFalse(
all(scores[0] == s for s in scores[1:]),
f"All items returned identical scores — delimiter indexing is likely broken. "
f"Scores: {scores[0]}",
)
def test_mis_deterministic(self):
"""Identical MIS requests return identical scores."""
kwargs = dict(query="Evaluate:", items=["alpha", "beta", "gamma"])
self.assertEqual(
self.engine.score(**kwargs).scores, self.engine.score(**kwargs).scores
)
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
# SequenceClassification — MIS with many labels (tensor shape stress test) # SequenceClassification — MIS with many labels (tensor shape stress test)
@@ -479,7 +455,8 @@ class TestSeqClsMISAdvancedScoring(CustomTestCase):
model_path=_SEQCLS_MODEL, model_path=_SEQCLS_MODEL,
disable_radix_cache=True, disable_radix_cache=True,
chunked_prefill_size=-1, chunked_prefill_size=-1,
multi_item_scoring_delimiter=_QWEN3_EOT_TOKEN_ID, enable_mis=True,
attention_backend="flashinfer",
json_model_override_args=json.dumps( json_model_override_args=json.dumps(
{ {
"architectures": ["Qwen3ForSequenceClassification"], "architectures": ["Qwen3ForSequenceClassification"],
@@ -506,17 +483,6 @@ class TestSeqClsMISAdvancedScoring(CustomTestCase):
self.assertEqual(len(row), self.NUM_LABELS) self.assertEqual(len(row), self.NUM_LABELS)
self.assertAlmostEqual(sum(row), 1.0, places=5) self.assertAlmostEqual(sum(row), 1.0, places=5)
def test_many_items_produce_distinct_scores(self):
"""15 items should not all return identical score vectors."""
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))
self.assertGreater(
len({tuple(s) for s in scores}),
1,
"All 15 items returned identical scores",
)
if __name__ == "__main__": if __name__ == "__main__":
unittest.main(verbosity=3) unittest.main(verbosity=3)
@@ -1,12 +1,11 @@
"""Unit tests for score_and_pool in sglang.srt.layers.pooler. """Unit tests for score_and_pool in sglang.srt.layers.pooler.
All tests run on CPU — no GPU required. The global server_args singleton All tests run on CPU — no GPU required. MIS delimiter positions are passed
is mocked so the tests are hermetic. via forward_batch.multi_item_delimiter_indices (pre-computed by the caller).
""" """
import unittest import unittest
from types import SimpleNamespace from types import SimpleNamespace
from unittest.mock import patch
import torch import torch
import torch.nn as nn import torch.nn as nn
@@ -24,22 +23,22 @@ register_cpu_ci(est_time=9, suite="stage-a-test-cpu")
def _make_forward_batch( def _make_forward_batch(
extend_seq_lens, is_prefill_only=False, return_pooled_hidden_states=False extend_seq_lens,
multi_item_delimiter_indices=None,
return_pooled_hidden_states=False,
is_prefill_only=True,
): ):
"""Build a minimal ForwardBatch stub for pooler unit tests.""" """Build a minimal ForwardBatch stub for pooler unit tests."""
return SimpleNamespace( return SimpleNamespace(
extend_seq_lens=torch.tensor(extend_seq_lens, dtype=torch.long), extend_seq_lens=torch.tensor(extend_seq_lens, dtype=torch.long),
extend_seq_lens_cpu=extend_seq_lens, extend_seq_lens_cpu=extend_seq_lens,
is_prefill_only=is_prefill_only, multi_item_delimiter_indices=multi_item_delimiter_indices,
dimensions=None, dimensions=None,
return_pooled_hidden_states=return_pooled_hidden_states, return_pooled_hidden_states=return_pooled_hidden_states,
is_prefill_only=is_prefill_only,
) )
def _mock_server_args(delimiter=None):
return SimpleNamespace(multi_item_scoring_delimiter=delimiter)
class TestScoreAndPool(CustomTestCase): class TestScoreAndPool(CustomTestCase):
"""Unit tests for the score_and_pool helper function.""" """Unit tests for the score_and_pool helper function."""
@@ -50,11 +49,8 @@ class TestScoreAndPool(CustomTestCase):
self.score_head = nn.Linear(self.hidden_dim, self.num_labels, bias=False) self.score_head = nn.Linear(self.hidden_dim, self.num_labels, bias=False)
self.pooler = Pooler(pooling_type=PoolingType.LAST, normalize=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):
def test_single_item_returns_scores(self, mock_get_args): """No delimiter indices -> single-item path returns [batch, num_labels]."""
"""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) hidden = torch.randn(8, self.hidden_dim)
fb = _make_forward_batch(extend_seq_lens=[5, 3]) fb = _make_forward_batch(extend_seq_lens=[5, 3])
input_ids = torch.arange(8) input_ids = torch.arange(8)
@@ -64,30 +60,16 @@ class TestScoreAndPool(CustomTestCase):
self.assertIsInstance(out, EmbeddingPoolerOutput) self.assertIsInstance(out, EmbeddingPoolerOutput)
self.assertEqual(out.embeddings.shape, (2, self.num_labels)) 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):
def test_mis_returns_per_request_list(self, mock_get_args): """Delimiter indices provided -> returns a list with one tensor per request."""
"""Delimiter found -> returns a list with one tensor per request.""" # Sequence: [0, 1, 2, D, 3, 4, 5, D, 6, 7, 8, D]
delimiter_token = 99 # Delimiters at positions 3, 7, 11 -> extract at 2, 6, 10
mock_get_args.return_value = _mock_server_args(delimiter=delimiter_token) input_ids = torch.arange(12)
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) hidden = torch.randn(len(input_ids), self.hidden_dim)
fb = _make_forward_batch(extend_seq_lens=[len(input_ids)], is_prefill_only=True) fb = _make_forward_batch(
extend_seq_lens=[len(input_ids)],
multi_item_delimiter_indices=[torch.tensor([3, 7, 11])],
)
out = score_and_pool(self.score_head, self.pooler, hidden, fb, input_ids) out = score_and_pool(self.score_head, self.pooler, hidden, fb, input_ids)
@@ -95,20 +77,20 @@ class TestScoreAndPool(CustomTestCase):
self.assertEqual(len(out.embeddings), 1) self.assertEqual(len(out.embeddings), 1)
self.assertEqual(out.embeddings[0].shape, (3, self.num_labels)) 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):
def test_mis_batched_splits_per_request(self, mock_get_args):
"""Two batched MIS requests -> returns a list of length 2.""" """Two batched MIS requests -> returns a list of length 2."""
delimiter_token = 99 # Request 1: [10, 11, D, 12, 13, D] -> delimiters at 2, 5
mock_get_args.return_value = _mock_server_args(delimiter=delimiter_token) # Request 2: [20, 21, 22, D] -> delimiter at 3
req1 = [10, 11, 99, 12, 13, 99]
# Request 1: [10, 11, delim, 12, 13, delim] -> 2 delimiters req2 = [20, 21, 22, 99]
# 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) input_ids = torch.tensor(req1 + req2)
hidden = torch.randn(len(input_ids), self.hidden_dim) hidden = torch.randn(len(input_ids), self.hidden_dim)
fb = _make_forward_batch( fb = _make_forward_batch(
extend_seq_lens=[len(req1), len(req2)], is_prefill_only=True extend_seq_lens=[len(req1), len(req2)],
multi_item_delimiter_indices=[
torch.tensor([2, 5]),
torch.tensor([3]),
],
) )
out = score_and_pool(self.score_head, self.pooler, hidden, fb, input_ids) out = score_and_pool(self.score_head, self.pooler, hidden, fb, input_ids)
@@ -118,42 +100,21 @@ class TestScoreAndPool(CustomTestCase):
self.assertEqual(out.embeddings[0].shape, (2, self.num_labels)) self.assertEqual(out.embeddings[0].shape, (2, self.num_labels))
self.assertEqual(out.embeddings[1].shape, (1, self.num_labels)) self.assertEqual(out.embeddings[1].shape, (1, self.num_labels))
@patch("sglang.srt.layers.pooler.get_global_server_args") def test_no_delimiter_indices_falls_back(self):
def test_mis_falls_back_when_no_delimiters_in_input(self, mock_get_args): """multi_item_delimiter_indices=None -> single-item fallback."""
"""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]) input_ids = torch.tensor([0, 1, 2, 3, 4, 5, 6, 7])
hidden = torch.randn(8, self.hidden_dim) hidden = torch.randn(8, self.hidden_dim)
fb = _make_forward_batch(extend_seq_lens=[5, 3], is_prefill_only=True) fb = _make_forward_batch(extend_seq_lens=[5, 3])
out = score_and_pool(self.score_head, self.pooler, hidden, fb, input_ids) out = score_and_pool(self.score_head, self.pooler, hidden, fb, input_ids)
self.assertIsInstance(out.embeddings, torch.Tensor) self.assertIsInstance(out.embeddings, torch.Tensor)
self.assertEqual(out.embeddings.shape, (2, self.num_labels)) 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):
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).""" """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 # 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]) input_ids = torch.tensor([10, 11, 99, 20, 21, 99])
hidden = ( hidden = (
torch.arange(len(input_ids)) torch.arange(len(input_ids))
.unsqueeze(1) .unsqueeze(1)
@@ -161,7 +122,10 @@ class TestScoreAndPool(CustomTestCase):
.expand(-1, self.hidden_dim) .expand(-1, self.hidden_dim)
.clone() .clone()
) )
fb = _make_forward_batch(extend_seq_lens=[len(input_ids)], is_prefill_only=True) fb = _make_forward_batch(
extend_seq_lens=[len(input_ids)],
multi_item_delimiter_indices=[torch.tensor([2, 5])],
)
identity_head = nn.Linear(self.hidden_dim, self.hidden_dim, bias=False) identity_head = nn.Linear(self.hidden_dim, self.hidden_dim, bias=False)
nn.init.eye_(identity_head.weight) nn.init.eye_(identity_head.weight)
@@ -172,14 +136,9 @@ class TestScoreAndPool(CustomTestCase):
torch.testing.assert_close(scores[0], hidden[1]) torch.testing.assert_close(scores[0], hidden[1])
torch.testing.assert_close(scores[1], hidden[4]) torch.testing.assert_close(scores[1], hidden[4])
@patch("sglang.srt.layers.pooler.get_global_server_args") def test_mis_delimiter_at_position_one(self):
def test_mis_ignores_delimiter_at_position_zero(self, mock_get_args): """Delimiters at positions 1 and 3 extract at indices 0 and 2."""
"""A delimiter at flat index 0 has no preceding token and must be skipped.""" input_ids = torch.tensor([10, 99, 11, 99])
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 = ( hidden = (
torch.arange(len(input_ids)) torch.arange(len(input_ids))
.unsqueeze(1) .unsqueeze(1)
@@ -187,7 +146,10 @@ class TestScoreAndPool(CustomTestCase):
.expand(-1, self.hidden_dim) .expand(-1, self.hidden_dim)
.clone() .clone()
) )
fb = _make_forward_batch(extend_seq_lens=[len(input_ids)], is_prefill_only=True) fb = _make_forward_batch(
extend_seq_lens=[len(input_ids)],
multi_item_delimiter_indices=[torch.tensor([1, 3])],
)
identity_head = nn.Linear(self.hidden_dim, self.hidden_dim, bias=False) identity_head = nn.Linear(self.hidden_dim, self.hidden_dim, bias=False)
nn.init.eye_(identity_head.weight) nn.init.eye_(identity_head.weight)
@@ -195,25 +157,37 @@ class TestScoreAndPool(CustomTestCase):
out = score_and_pool(identity_head, self.pooler, hidden, fb, input_ids) out = score_and_pool(identity_head, self.pooler, hidden, fb, input_ids)
self.assertEqual(len(out.embeddings), 1) self.assertEqual(len(out.embeddings), 1)
self.assertEqual(out.embeddings[0].shape[0], 1) self.assertEqual(out.embeddings[0].shape[0], 2)
torch.testing.assert_close(out.embeddings[0][0], hidden[2]) torch.testing.assert_close(out.embeddings[0][0], hidden[0])
torch.testing.assert_close(out.embeddings[0][1], 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)
def test_single_item_scores_match_manual_computation(self):
"""Single-item scores equal score_head applied to pooled hidden states."""
hidden = torch.randn(8, self.hidden_dim) hidden = torch.randn(8, self.hidden_dim)
fb = _make_forward_batch(extend_seq_lens=[5, 3]) fb = _make_forward_batch(extend_seq_lens=[5, 3])
input_ids = torch.arange(8) input_ids = torch.arange(8)
out = score_and_pool(self.score_head, self.pooler, hidden, fb, input_ids) out = score_and_pool(self.score_head, self.pooler, hidden, fb, input_ids)
# score-first-then-pool: matches the original Qwen3/Qwen2 classification forward pooled = self.pooler(hidden, fb).embeddings
logits = self.score_head(hidden) expected = self.score_head(pooled)
expected = self.pooler(logits, fb).embeddings
torch.testing.assert_close(out.embeddings, expected) torch.testing.assert_close(out.embeddings, expected)
def test_empty_delimiter_indices(self):
"""Empty delimiter tensor per request -> returns list with empty tensor."""
input_ids = torch.arange(6)
hidden = torch.randn(6, self.hidden_dim)
fb = _make_forward_batch(
extend_seq_lens=[6],
multi_item_delimiter_indices=[torch.tensor([], dtype=torch.long)],
)
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, (0, self.num_labels))
if __name__ == "__main__": if __name__ == "__main__":
unittest.main() unittest.main()