[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.decode_backend = "fa2"
# Store multi-item scoring delimiter for efficient access
self.multi_item_scoring_delimiter = (
model_runner.server_args.multi_item_scoring_delimiter
)
# Store multi-item scoring flag for efficient access
self.enable_mis = model_runner.server_args.enable_mis
# FIXME: remove dllm workarounds from flashinfer
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)
"""
delimiter = self.multi_item_scoring_delimiter
if delimiter is None or forward_batch.forward_mode == ForwardMode.DECODE:
if not self.enable_mis or forward_batch.forward_mode == ForwardMode.DECODE:
return MultiItemScoringParams()
delimiter_mask = forward_batch.input_ids == delimiter
prefix_cache_lens = getattr(forward_batch, "extend_prefix_lens", None)
extend_seq_lens = getattr(forward_batch, "extend_seq_lens", None)
precomputed_indices = forward_batch.multi_item_delimiter_indices
if precomputed_indices is 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 = [], []
token_pos_in_items_len = 0
device = forward_batch.input_ids.device
# If no extend_seq_lens, treat whole batch as one sequence
if extend_seq_lens is None or len(extend_seq_lens) <= 1:
@@ -359,35 +360,44 @@ class FlashInferAttnBackend(AttentionBackend):
seq_start = 0
for i, seq_len in enumerate(extend_seq_lens):
seq_end = seq_start + seq_len
mask = delimiter_mask[seq_start:seq_end]
pos = forward_batch.positions[seq_start:seq_end]
delimiter_indices = torch.nonzero(mask, as_tuple=True)[0]
delimiter_indices_cpu = precomputed_indices[i]
if len(delimiter_indices_cpu) == 0:
seq_start = seq_end
continue
if len(delimiter_indices) > 0:
first_delim = delimiter_indices[0]
# Prefix length: store as scalar
prefix_len = first_delim + (
prefix_cache_lens[i] if prefix_cache_lens is not None else 0
)
prefix_len_ptr.append(
prefix_len.item() if torch.is_tensor(prefix_len) else prefix_len
)
first_delim = delimiter_indices_cpu[0].item() # CPU .item(), no GPU sync
delimiter_indices = delimiter_indices_cpu.to(device, non_blocking=True)
prefix_len = first_delim + (
prefix_cache_lens[i] if prefix_cache_lens is not None else 0
)
prefix_len_ptr.append(prefix_len)
# Compute relative positions within items after delimiters
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)
# 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
# Update forward_batch positions in-place
pos[first_delim:] = diff - 1
forward_batch.positions[seq_start:seq_end] = pos
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
token_pos_in_items_ptr.append(pos_within_item.to(torch.uint16))
forward_batch.positions[seq_start + first_delim : seq_end] = (
prefix_len + pos_within_item - 1
)
seq_start = seq_end
# Pad token_pos_in_items_ptr for batch processing
if 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 = [
torch.cat(
[
@@ -405,8 +415,6 @@ class FlashInferAttnBackend(AttentionBackend):
if not prefix_len_ptr or not token_pos_in_items_ptr:
return MultiItemScoringParams()
# Build final params
device = forward_batch.input_ids.device
return MultiItemScoringParams(
prefix_len_ptr=torch.tensor(
prefix_len_ptr, dtype=torch.uint32, device=device
@@ -470,7 +478,7 @@ class FlashInferAttnBackend(AttentionBackend):
prefix_lens = forward_batch.extend_prefix_lens
# 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:
# 1. Ragged wrapper doesn't support the specialized multi-item parameters
# (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
multi_item_params = MultiItemScoringParams()
if self.multi_item_scoring_delimiter is not None:
if self.enable_mis:
# Use new backend-specific implementation
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.return_full_logits = return_full_logits
self.multi_item_delimiter = (
get_global_server_args().multi_item_scoring_delimiter
)
self.enable_mis = get_global_server_args().enable_mis
# enable chunked logprobs processing
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,
hidden_states_before_norm: Optional[torch.Tensor] = None,
) -> LogitsProcessorOutput:
# Extract MIS indices before ForwardBatch → LogitsMetadata conversion
multi_item_delimiter_indices = None
if isinstance(logits_metadata, ForwardBatch):
multi_item_delimiter_indices = logits_metadata.multi_item_delimiter_indices
logits_metadata = LogitsMetadata.from_forward_batch(logits_metadata)
# Multi-item scoring only for prefill-only requests.
if self.multi_item_delimiter is not None and logits_metadata.is_prefill_only:
# Multi-item scoring only for prefill-only requests with pre-computed indices.
if multi_item_delimiter_indices is not None and logits_metadata.is_prefill_only:
return self.compute_logprobs_for_multi_item_scoring(
input_ids,
hidden_states,
lm_head,
logits_metadata,
self.multi_item_delimiter,
multi_item_delimiter_indices,
)
# Diffusion LLM only.
@@ -347,6 +348,9 @@ class LogitsProcessor(nn.Module):
return LogitsProcessorOutput(
next_token_logits=sampled_logits,
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,
)
@@ -1006,39 +1010,41 @@ class LogitsProcessor(nn.Module):
hidden_states,
lm_head: VocabParallelEmbedding,
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.
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.
Compute logprobs for multi-item scoring using pre-computed delimiter indices.
Sequence format: Query<delimiter>Item1<delimiter>Item2<delimiter>...
Scoring positions: Extracts logprobs at positions before each <delimiter>
Args:
input_ids (torch.Tensor): Input token IDs containing query and items separated by delimiters.
Shape: [total_sequence_length] for single request or [batch_total_length] for batch.
hidden_states (torch.Tensor): Hidden states from the model.
Shape: [sequence_length, hidden_dim].
lm_head (VocabParallelEmbedding): Language model head for computing logits.
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)
input_ids: Input token IDs. Shape: [total_sequence_length].
hidden_states: Hidden states from the model. Shape: [sequence_length, hidden_dim].
lm_head: Language model head for computing logits.
logits_metadata: Metadata containing batch info and logprob specs.
multi_item_delimiter_indices: Pre-computed delimiter positions per request (CPU tensors).
"""
multi_item_indices = (input_ids == delimiter_token).nonzero(as_tuple=True)[
0
] - 1
# Compute positions just before each delimiter.
# Build offset-adjusted indices on CPU, then do a single CPU→GPU transfer.
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
sliced_hidden = hidden_states[multi_item_indices]
@@ -1052,27 +1058,13 @@ class LogitsProcessor(nn.Module):
input_top_logprobs_idx = None
# 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 (
logits_metadata.token_ids_logprobs
or logits_metadata.extend_return_top_logprob
):
logits_metadata.extend_logprob_pruned_lens_cpu = []
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]
logits_metadata.extend_logprob_pruned_lens_cpu = [
len(t) for t in multi_item_delimiter_indices
]
# Get the logprobs of specified token ids
if logits_metadata.extend_token_ids_logprob:
@@ -1090,11 +1082,17 @@ class LogitsProcessor(nn.Module):
input_top_logprobs_idx,
) = get_top_logprobs_prefill(sliced_logprobs, logits_metadata)
# For input_token_logprobs, use delimiter token logprobs
input_token_logprobs = sliced_logprobs[:, delimiter_token]
# MIS scores come from input_token_ids_logprobs_val (label-token logprobs),
# 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(
next_token_logits=None, # Multi-item scoring doesn't need next token logits
next_token_logits=None,
input_token_logprobs=input_token_logprobs,
input_top_logprobs_val=input_top_logprobs_val,
input_top_logprobs_idx=input_top_logprobs_idx,
+65 -37
View File
@@ -12,7 +12,6 @@ import torch.nn as nn
from transformers import PretrainedConfig
from sglang.srt.layers.activation import get_cross_encoder_activation_function
from sglang.srt.server_args import get_global_server_args
if TYPE_CHECKING:
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}")
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(
score_head: nn.Module,
pooler: "Pooler",
@@ -75,47 +114,36 @@ def score_and_pool(
) -> EmbeddingPoolerOutput:
"""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``):
extract hidden states at positions just before each delimiter, apply the score head,
then split per-request.
MIS path (pre-computed delimiter indices on forward_batch): extract hidden
states at positions just before each delimiter, apply the score head, then
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
hidden states (before the score head) are included in the output.
"""
delimiter_token = get_global_server_args().multi_item_scoring_delimiter
if delimiter_token is not None and forward_batch.is_prefill_only:
delim_positions = (input_ids == delimiter_token).nonzero(as_tuple=True)[0]
# A delimiter at flat index 0 has no preceding hidden state to pool
delim_positions = delim_positions[delim_positions > 0]
if delim_positions.numel() > 0:
# Score only the tokens that precede a delimiter
pre_delim_hidden = hidden_states[delim_positions - 1]
scores = score_head(pre_delim_hidden)
# 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:
end = start + seq_len
mask = (delim_positions >= start) & (delim_positions < end)
per_request_scores.append(scores[mask])
if per_request_phs is not None:
per_request_phs.append(pre_delim_hidden[mask])
start = end
return EmbeddingPoolerOutput(
embeddings=per_request_scores,
pooled_hidden_states=per_request_phs,
)
if (
forward_batch.multi_item_delimiter_indices is not None
and forward_batch.is_prefill_only
):
# Pool hidden states at pre-delimiter positions, score only those —
# avoids wasting compute on tokens that never contribute to the output.
# pool_at_delimiter_positions returns one tensor per request; we concat
# to call score_head once, then split back per request.
per_request_phs = pool_at_delimiter_positions(
hidden_states, forward_batch, input_ids.device
)
phs_flat = torch.cat(per_request_phs, dim=0)
scores_flat = score_head(phs_flat)
delim_counts = [t.shape[0] for t in per_request_phs]
per_request_scores = list(scores_flat.split(delim_counts))
return EmbeddingPoolerOutput(
embeddings=per_request_scores,
pooled_hidden_states=(
per_request_phs if forward_batch.return_pooled_hidden_states else None
),
)
# Standard classification path: pool hidden states, then score.
pooled_hs = pool_hidden_states(pooler.pooling_type, hidden_states, forward_batch)
+28
View File
@@ -250,6 +250,10 @@ class GenerateReqInput(BaseReq):
image_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:
return (
has_valid_data(self.image_data)
@@ -685,6 +689,11 @@ class GenerateReqInput(BaseReq):
external_trace_header=self.external_trace_header,
http_worker_ipc=self.http_worker_ipc,
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
return sub
@@ -774,6 +783,9 @@ class TokenizedGenerateReqInput(BaseReq):
need_wait_for_mm_inputs: bool = False
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
time_stats: Optional[Union[APIServerReqTimeStats, DPControllerReqTimeStats]] = None
@@ -855,6 +867,10 @@ class EmbeddingReqInput(BaseReq):
# Whether to return pooled hidden states (pre-head transformer output)
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):
# 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:
@@ -957,6 +973,11 @@ class EmbeddingReqInput(BaseReq):
is_cross_encoder_request=True,
http_worker_ipc=self.http_worker_ipc,
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:
sub = EmbeddingReqInput(
@@ -981,6 +1002,11 @@ class EmbeddingReqInput(BaseReq):
http_worker_ipc=self.http_worker_ipc,
received_time=self.received_time,
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
return sub
@@ -1009,6 +1035,8 @@ class TokenizedEmbeddingReqInput(BaseReq):
# LoRA related
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
time_stats: Optional[Union[APIServerReqTimeStats, DPControllerReqTimeStats]] = None
@@ -597,6 +597,7 @@ class Req(ReqDllmMixin):
Union[APIServerReqTimeStats, DPControllerReqTimeStats]
] = None,
return_pooled_hidden_states: bool = False,
multi_item_delimiter_indices: Optional[List[int]] = None,
):
# Input and output info
self.rid = rid
@@ -614,6 +615,7 @@ class Req(ReqDllmMixin):
self.session = session
self.input_embeds = input_embeds
self.positional_embed_overrides = positional_embed_overrides
self.multi_item_delimiter_indices = multi_item_delimiter_indices
# For req-level memory management
self.kv_committed_len = 0
@@ -1441,6 +1443,9 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
# Whether this batch is prefill-only (no token generation needed)
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_consumer_index: int = -1
@@ -1817,6 +1822,23 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
self.token_type_ids = token_type_ids_tensor
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:
self.top_logprobs_nums = [r.top_logprobs_num 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,
is_prefill_only=self.is_prefill_only,
multi_item_delimiter_indices=self.multi_item_delimiter_indices,
dimensions=self.dimensions,
return_pooled_hidden_states=self.return_pooled_hidden_states,
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)
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
dllm_block_offsets: Optional[List[int]] = None
dllm_config: Optional[DllmConfig] = None
+2
View File
@@ -1881,6 +1881,7 @@ class Scheduler(
http_worker_ipc=recv_req.http_worker_ipc,
dllm_config=self.dllm_config,
time_stats=recv_req.time_stats,
multi_item_delimiter_indices=recv_req.multi_item_delimiter_indices,
)
req.tokenizer = self.tokenizer
@@ -2202,6 +2203,7 @@ class Scheduler(
http_worker_ipc=recv_req.http_worker_ipc,
time_stats=recv_req.time_stats,
return_pooled_hidden_states=recv_req.return_pooled_hidden_states,
multi_item_delimiter_indices=recv_req.multi_item_delimiter_indices,
)
req.tokenizer = self.tokenizer
@@ -21,7 +21,7 @@ from sglang.srt.managers.schedule_batch import (
ScheduleBatch,
)
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:
from sglang.srt.managers.scheduler import (
@@ -629,13 +629,12 @@ class SchedulerOutputProcessorMixin:
# Process logprob indices based on scoring type
if is_multi_item_scoring:
# Multi-item scoring: only include delimiter token positions
relevant_tokens = req.origin_input_ids[req.logprob_start_len :]
input_token_logprobs_idx = [
token_id
for token_id in relevant_tokens
if token_id == self.server_args.multi_item_scoring_delimiter
]
# MIS scores come from input_token_ids_logprobs, not input_token_logprobs.
# But the shared pipeline requires input_token_logprobs_idx to be the same
# length as input_token_logprobs_val (validated at line 816). We fill with
# MIS_DELIMITER_TOKEN_ID as a dummy — score_request() ignores this field.
delimiter_count = len(req.multi_item_delimiter_indices)
input_token_logprobs_idx = [MIS_DELIMITER_TOKEN_ID] * delimiter_count
else:
# Regular request: include all tokens from logprob_start_len onwards
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.
"""
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:
# Multi-item scoring: count delimiter tokens from logprob_start_len onwards
return sum(
1
for token_id in relevant_tokens
if token_id == self.server_args.multi_item_scoring_delimiter
)
return len(req.multi_item_delimiter_indices)
else:
# Regular request: all tokens from logprob_start_len onwards
return len(relevant_tokens)
return len(req.origin_input_ids[req.logprob_start_len :])
def _calculate_num_input_logprobs(
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)
if is_multi_item_scoring:
# Multi-item scoring: count delimiter tokens in the relevant portion
relevant_tokens = req.origin_input_ids[
extend_logprob_start_len:extend_input_len
]
# Count pre-computed delimiter indices within the extend range
return sum(
1
for token_id in relevant_tokens
if token_id == self.server_args.multi_item_scoring_delimiter
for idx in req.multi_item_delimiter_indices
if extend_logprob_start_len <= idx < extend_input_len
)
else:
# Regular request: all tokens in the range
@@ -758,7 +747,11 @@ class SchedulerOutputProcessorMixin:
token is configured. In this mode, only positions containing the
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(
self: Scheduler,
@@ -313,7 +313,6 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
self.processor = _processor
self.tokenizer = get_tokenizer_from_processor(self.processor)
os.environ["TOKENIZERS_PARALLELISM"] = "false"
self._initialize_multi_item_delimiter_text()
else:
self.mm_processor = self.processor = None
@@ -326,7 +325,6 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
trust_remote_code=server_args.trust_remote_code,
revision=server_args.revision,
)
self._initialize_multi_item_delimiter_text()
# Initialize async dynamic batch tokenizer if enabled (common for both multimodal and non-multimodal)
if (
@@ -1007,6 +1005,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
token_type_ids=token_type_ids,
need_wait_for_mm_inputs=obj.need_wait_for_mm_inputs,
num_items_assigned=obj.num_items_assigned,
multi_item_delimiter_indices=obj.multi_item_delimiter_indices,
)
elif isinstance(obj, EmbeddingReqInput):
# Resolve unresolved embed overrides now that input_ids are available
@@ -1033,6 +1032,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
lora_id=obj.lora_id,
http_worker_ipc=obj.http_worker_ipc,
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
@@ -8,6 +8,7 @@ import torch
from sglang.srt.configs.model_config import is_cross_encoding_pooler_model
from sglang.srt.managers.embed_types import PositionalEmbeds
from sglang.srt.managers.io_struct import EmbeddingReqInput, GenerateReqInput
from sglang.srt.server_args import MIS_DELIMITER_TOKEN_ID
logger = logging.getLogger(__name__)
@@ -76,27 +77,9 @@ class TokenizerManagerScoreMixin:
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(
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.
Format: query<delimiter>item1<delimiter>item2<delimiter>item3<delimiter>
@@ -107,18 +90,21 @@ class TokenizerManagerScoreMixin:
delimiter_token_id: Token ID to use as delimiter
Returns:
Combined token sequence
Tuple of (combined token sequence, delimiter indices)
"""
combined_sequence = query[:] # Start with query
delimiter_indices = []
for item in items:
delimiter_indices.append(len(combined_sequence))
combined_sequence.append(delimiter_token_id) # Add delimiter
combined_sequence.extend(item) # Add item tokens
# Add final delimiter after the last item for logprob extraction
delimiter_indices.append(len(combined_sequence))
combined_sequence.append(delimiter_token_id)
return combined_sequence
return combined_sequence, delimiter_indices
def _batch_tokenize_query_and_items(
self,
@@ -416,11 +402,14 @@ class TokenizerManagerScoreMixin:
embed_override_token_id: Optional[int],
query_embed_overrides: 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.
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.
Returns:
(text_prompts, input_ids, positional_embed_overrides, delimiter_indices)
"""
# Both query and items are token IDs
has_embeds = (
@@ -428,16 +417,17 @@ class TokenizerManagerScoreMixin:
)
if use_multi_item_scoring:
# Multi-item scoring: concatenate with delimiter token ID
# Format: query<delimiter_token_id>item1<delimiter_token_id>item2<delimiter_token_id>item3<delimiter_token_id>
delimiter_token_id = self.server_args.multi_item_scoring_delimiter
combined_input_ids = self._build_multi_item_token_sequence(
query, items, delimiter_token_id
# Multi-item scoring: concatenate with placeholder delimiter token.
# Positions are derived from item lengths (delimiter_indices), not
# by scanning for this token — it exists only for FlashInfer compat.
delimiter_token_id = MIS_DELIMITER_TOKEN_ID
combined_input_ids, delimiter_indices = (
self._build_multi_item_token_sequence(query, items, delimiter_token_id)
)
input_ids = [combined_input_ids]
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
all_embeds: List[torch.Tensor] = []
@@ -461,15 +451,15 @@ class TokenizerManagerScoreMixin:
current_offset += len(item) + 1 # +1 for delimiter
if all_embeds:
injection = [
positional_embed_overrides = [
PositionalEmbeds(
embeds=torch.cat(all_embeds, dim=0),
positions=all_positions,
)
]
else:
injection = None
return None, input_ids, injection
positional_embed_overrides = None
return None, input_ids, positional_embed_overrides, delimiter_indices
else:
# Single-item scoring: process each item separately
@@ -479,9 +469,9 @@ class TokenizerManagerScoreMixin:
input_ids = [query + item for item in items]
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):
item_embs = item_embed_overrides[i] if item_embed_overrides else None
pe = self._resolve_embed_overrides_for_request(
@@ -493,13 +483,14 @@ class TokenizerManagerScoreMixin:
item_position_offset=len(query),
item_label=f"items[{i}]",
)
injection.append(pe)
positional_embed_overrides.append(pe)
return (
None,
input_ids,
injection if any(pe is not None for pe in injection) else None,
positional_embed_overrides = (
positional_embed_overrides
if any(pe is not None for pe in positional_embed_overrides)
else None
)
return None, input_ids, positional_embed_overrides, None
# ------------------------------------------------------------------
# Main entry point
@@ -523,7 +514,7 @@ class TokenizerManagerScoreMixin:
This method supports two scoring approaches:
1. Single-Item scoring (default): Process each query+item pair independently
2. Multi-Item scoring: When multi_item_scoring_delimiter is set, combine query and
2. Multi-Item scoring: When --enable-mis is set, combine query and
multiple items into a single sequence using delimiter for efficient processing.
Note: item_first parameter is ignored in multi-item scoring mode since it uses
a fixed format: query<delimiter>item1<delimiter>item2<delimiter>item3<delimiter>
@@ -593,15 +584,13 @@ class TokenizerManagerScoreMixin:
f"Token ID {token_id} is out of vocabulary (vocab size: {vocab_size})"
)
# Check if multi-item scoring is enabled by presence of delimiter
use_multi_item_scoring = (
self.server_args.multi_item_scoring_delimiter is not None
and self.multi_item_delimiter_text is not None
)
# Check if multi-item scoring is enabled
use_multi_item_scoring = self.server_args.enable_mis
input_ids = None
text_prompts = None
positional_embed_overrides = None
delimiter_indices = None
use_text_prompts = isinstance(query, str) and not has_embeds
@@ -609,15 +598,17 @@ class TokenizerManagerScoreMixin:
# Both query and items are text
items_list = [items] if isinstance(items, str) else items
if use_multi_item_scoring:
# Multi-item scoring: tokenize separately then combine at token level
# to ensure the delimiter token ID is inserted exactly once per boundary
# (a text-level roundtrip through the tokenizer can alter boundary tokens)
delimiter_token_id = self.server_args.multi_item_scoring_delimiter
# Tokenize separately, then combine at token level with placeholder
# delimiter. Positions come from item lengths (delimiter_indices),
# not from scanning for this token — it's for FlashInfer compat only.
delimiter_token_id = MIS_DELIMITER_TOKEN_ID
query_ids, items_ids = self._batch_tokenize_query_and_items(
query, items_list
)
combined_input_ids = self._build_multi_item_token_sequence(
query_ids, items_ids, delimiter_token_id
combined_input_ids, delimiter_indices = (
self._build_multi_item_token_sequence(
query_ids, items_ids, delimiter_token_id
)
)
input_ids = [combined_input_ids]
else:
@@ -635,26 +626,30 @@ class TokenizerManagerScoreMixin:
):
# Both query and items are token IDs — tokenize text inputs if needed for embed overrides
query_ids, items_ids = query, items
_, input_ids, positional_embed_overrides = self._build_token_id_inputs(
query_ids,
items_ids,
item_first,
use_multi_item_scoring,
embed_override_token_id,
query_embed_overrides,
item_embed_overrides,
_, input_ids, positional_embed_overrides, delimiter_indices = (
self._build_token_id_inputs(
query_ids,
items_ids,
item_first,
use_multi_item_scoring,
embed_override_token_id,
query_embed_overrides,
item_embed_overrides,
)
)
elif has_embeds:
# Text inputs with embed overrides — need to tokenize first to resolve positions
query_ids, items_ids = self._batch_tokenize_query_and_items(query, items)
_, input_ids, positional_embed_overrides = self._build_token_id_inputs(
query_ids,
items_ids,
item_first,
use_multi_item_scoring,
embed_override_token_id,
query_embed_overrides,
item_embed_overrides,
_, input_ids, positional_embed_overrides, delimiter_indices = (
self._build_token_id_inputs(
query_ids,
items_ids,
item_first,
use_multi_item_scoring,
embed_override_token_id,
query_embed_overrides,
item_embed_overrides,
)
)
else:
raise ValueError(
@@ -679,6 +674,7 @@ class TokenizerManagerScoreMixin:
)
# Create the appropriate request type
mis_delimiter_indices = [delimiter_indices] if use_multi_item_scoring else None
if is_generation:
batch_request = GenerateReqInput(
text=text_prompts,
@@ -690,6 +686,7 @@ class TokenizerManagerScoreMixin:
stream=False,
sampling_params={"max_new_tokens": 0},
positional_embed_overrides=positional_embed_overrides,
multi_item_delimiter_indices=mis_delimiter_indices,
)
else:
batch_request = EmbeddingReqInput(
@@ -697,6 +694,7 @@ class TokenizerManagerScoreMixin:
input_ids=input_ids,
positional_embed_overrides=positional_embed_overrides,
return_pooled_hidden_states=return_pooled_hidden_states,
multi_item_delimiter_indices=mis_delimiter_indices,
)
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)
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
spec_info: Optional[SpecInput] = None
spec_algorithm: SpeculativeAlgorithm = None
@@ -468,6 +471,7 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
can_run_dp_cuda_graph=batch.can_run_dp_cuda_graph,
global_forward_mode=batch.global_forward_mode,
is_prefill_only=batch.is_prefill_only,
multi_item_delimiter_indices=batch.multi_item_delimiter_indices,
lora_ids=batch.lora_ids,
sampling_info=batch.sampling_info,
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"]
# 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 = [
"auto",
"deep_gemm",
@@ -601,10 +608,11 @@ class ServerArgs:
offload_mode: str = "cpu"
# Scoring configuration
# Delimiter token ID used to combine Query and Items into a single sequence for multi-item scoring.
# Format: Query<delimiter>Item1<delimiter>Item2<delimiter>...
# This enables efficient batch processing of multiple items against a single query.
multi_item_scoring_delimiter: Optional[Union[int]] = None
# Enable Multi-Item Scoring optimization. Combines query and multiple items
# into a single sequence for efficient batch processing. Item boundaries are
# determined by pre-computed delimiter indices (from item lengths), not by the
# placeholder token. See MIS_DELIMITER_TOKEN_ID for details.
enable_mis: bool = False
# Optimization/debug options
disable_radix_cache: bool = False
@@ -800,9 +808,6 @@ class ServerArgs:
# 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.
gpu_mem = get_device_memory_capacity(self.device)
@@ -823,6 +828,10 @@ class ServerArgs:
self._handle_nccl_pre_warm()
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.
self._handle_hicache()
@@ -1227,20 +1236,36 @@ class ServerArgs:
self.disable_piecewise_cuda_graph = True
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
spurious delimiter matches in score_and_pool's MIS path.
Auto-disables settings incompatible with MIS mechanics (CUDA graph,
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
if not self.disable_cuda_graph:
logger.warning(
"CUDA graph is disabled because --multi-item-scoring-delimiter is set."
)
logger.warning("CUDA graph is disabled because --enable-mis is set.")
self.disable_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):
"""
Configure GPU memory-dependent settings including
@@ -5739,10 +5764,13 @@ class ServerArgs:
# Args for multi-item-scoring
parser.add_argument(
"--multi-item-scoring-delimiter",
type=int,
default=ServerArgs.multi_item_scoring_delimiter,
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.",
"--enable-mis",
action="store_true",
default=ServerArgs.enable_mis,
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
@@ -6612,17 +6640,6 @@ class ServerArgs:
"--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
if self.enable_hisparse:
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 (
TokenizerManagerScoreMixin,
)
from sglang.srt.server_args import MIS_DELIMITER_TOKEN_ID
from sglang.test.ci.ci_register import register_cuda_ci
from sglang.test.test_utils import CustomTestCase
@@ -204,16 +205,15 @@ class TestEmbeddingReqInputEmbedOverride(CustomTestCase):
class _FakeServerArgs:
"""Minimal stub for server_args."""
def __init__(self, multi_item_scoring_delimiter=None):
self.multi_item_scoring_delimiter = multi_item_scoring_delimiter
def __init__(self, enable_mis=False):
self.enable_mis = enable_mis
class _FakeMixin(TokenizerManagerScoreMixin):
"""Minimal stub to call mixin methods without a full TokenizerManager."""
def __init__(self, delimiter=None):
self.server_args = _FakeServerArgs(delimiter)
self.multi_item_delimiter_text = None
def __init__(self, enable_mis=False):
self.server_args = _FakeServerArgs(enable_mis)
self.tokenizer = None
self.is_generation = True
@@ -334,17 +334,17 @@ class TestResolveEmbedOverridesForRequest(CustomTestCase):
# Score mixin: _build_token_id_inputs
# ========================================================================
DELIM_TOKEN = 99
DELIM_TOKEN = MIS_DELIMITER_TOKEN_ID
class TestBuildTokenIdInputs(CustomTestCase):
def setUp(self):
self.mixin = _FakeMixin(delimiter=DELIM_TOKEN)
self.mixin = _FakeMixin(enable_mis=True)
# --- single-item mode, no embeds ---
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],
items=[[3, 4], [5, 6]],
item_first=False,
@@ -354,10 +354,10 @@ class TestBuildTokenIdInputs(CustomTestCase):
item_embed_overrides=None,
)
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):
_, input_ids, injection = self.mixin._build_token_id_inputs(
_, input_ids, positional_embed_overrides, _ = self.mixin._build_token_id_inputs(
query=[1, 2],
items=[[3, 4]],
item_first=True,
@@ -367,12 +367,12 @@ class TestBuildTokenIdInputs(CustomTestCase):
item_embed_overrides=None,
)
self.assertEqual(input_ids, [[3, 4, 1, 2]])
self.assertIsNone(injection)
self.assertIsNone(positional_embed_overrides)
# --- multi-item mode, no embeds ---
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],
items=[[3, 4], [5, 6]],
item_first=False,
@@ -385,13 +385,13 @@ class TestBuildTokenIdInputs(CustomTestCase):
self.assertEqual(
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 ---
def test_single_item_query_embeds(self):
"""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],
items=[[20, 30], [40, 50]],
item_first=False,
@@ -401,15 +401,15 @@ class TestBuildTokenIdInputs(CustomTestCase):
item_embed_overrides=None,
)
self.assertEqual(input_ids, [[50, 10, 20, 30], [50, 10, 40, 50]])
self.assertIsNotNone(injection)
self.assertEqual(len(injection), 2)
self.assertIsNotNone(positional_embed_overrides)
self.assertEqual(len(positional_embed_overrides), 2)
# Each item gets its own PositionalEmbeds with query override at pos 0
self.assertEqual(injection[0].positions, [0])
self.assertEqual(injection[1].positions, [0])
self.assertEqual(positional_embed_overrides[0].positions, [0])
self.assertEqual(positional_embed_overrides[1].positions, [0])
def test_single_item_item_embeds(self):
"""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],
items=[[50, 30]],
item_first=False,
@@ -419,13 +419,13 @@ class TestBuildTokenIdInputs(CustomTestCase):
item_embed_overrides=[[_vec(2)]],
)
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
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):
"""When no items have placeholders, injection should be None."""
_, input_ids, injection = self.mixin._build_token_id_inputs(
"""When no items have placeholders, positional_embed_overrides should be None."""
_, input_ids, positional_embed_overrides, _ = self.mixin._build_token_id_inputs(
query=[10, 20],
items=[[30, 40]],
item_first=False,
@@ -434,11 +434,11 @@ class TestBuildTokenIdInputs(CustomTestCase):
query_embed_overrides=None,
item_embed_overrides=[None],
)
self.assertIsNone(injection)
self.assertIsNone(positional_embed_overrides)
def test_single_item_query_and_item_embeds(self):
"""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],
items=[[20, 50]],
item_first=False,
@@ -448,15 +448,15 @@ class TestBuildTokenIdInputs(CustomTestCase):
item_embed_overrides=[[_vec(2)]],
)
self.assertEqual(input_ids, [[50, 10, 20, 50]])
self.assertIsNotNone(injection)
pe = injection[0]
self.assertIsNotNone(positional_embed_overrides)
pe = positional_embed_overrides[0]
# query override at pos 0, item override at pos 3 (query_len=2 + idx=1)
self.assertEqual(pe.positions, [0, 3])
self.assertEqual(pe.embeds.shape, (2, HIDDEN_DIM))
def test_single_item_empty_query(self):
"""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=[],
items=[[50, 10]],
item_first=False,
@@ -466,15 +466,15 @@ class TestBuildTokenIdInputs(CustomTestCase):
item_embed_overrides=[[_vec(1)]],
)
self.assertEqual(input_ids, [[50, 10]])
self.assertIsNotNone(injection)
self.assertIsNotNone(positional_embed_overrides)
# 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 ---
def test_multi_item_with_query_and_item_embeds(self):
"""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],
items=[[20, 50], [30, 40]],
item_first=False,
@@ -483,13 +483,13 @@ class TestBuildTokenIdInputs(CustomTestCase):
query_embed_overrides=[_vec(1)],
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.assertIsNotNone(injection)
self.assertIsNotNone(positional_embed_overrides)
self.assertEqual(
len(injection), 1
len(positional_embed_overrides), 1
) # 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)
self.assertIn(0, 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")
_SEQCLS_MODEL = "Qwen/Qwen3-0.6B"
_QWEN3_EOT_TOKEN_ID = 151643
_CAUSAL_LM_MODEL = DEFAULT_SMALL_MODEL_NAME_FOR_TEST
_NUM_LABELS = 4
@@ -197,7 +196,7 @@ class TestPooledHiddenStatesMISEngine(CustomTestCase):
model_path=_SEQCLS_MODEL,
disable_radix_cache=True,
chunked_prefill_size=-1,
multi_item_scoring_delimiter=_QWEN3_EOT_TOKEN_ID,
enable_mis=True,
json_model_override_args=json.dumps(
{
"architectures": ["Qwen3ForSequenceClassification"],
+11 -11
View File
@@ -3,15 +3,16 @@
Two test classes, each with its own server instance:
TestCausalLMScoringHTTP basic endpoint: schema defaults, response
structure, error rejection (no MIS delimiter)
TestCausalLMMISScoringHTTP MIS mode: validates --multi-item-scoring-delimiter
CLI flag wiring and per-item output shape
structure, error rejection (no MIS)
TestCausalLMMISScoringHTTP MIS mode: validates --enable-mis CLI flag
wiring and per-item output shape
Engine-level correctness (numerical accuracy, batching, edge cases) lives in
test_score_engine.py. These tests focus on the HTTP integration seam:
Pydantic schema defaults, FastAPI routing, and server argument wiring.
"""
import os
import unittest
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")
_MODEL = DEFAULT_SMALL_MODEL_NAME_FOR_TEST # Llama-3.2-1B-Instruct
# <|eot_id|> for Llama-3.x Instruct — used as MIS delimiter
_LLAMA3_EOT_TOKEN_ID = 128009
_MODEL = os.environ.get("TEST_MODEL_NAME", DEFAULT_SMALL_MODEL_NAME_FOR_TEST)
# ---------------------------------------------------------------------------
@@ -41,7 +40,7 @@ _LLAMA3_EOT_TOKEN_ID = 128009
class TestCausalLMScoringHTTP(CustomTestCase):
"""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
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):
"""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
endpoint returns one probability vector per item when items are
@@ -163,8 +162,9 @@ class TestCausalLMMISScoringHTTP(CustomTestCase):
"--disable-radix-cache",
"--chunked-prefill-size",
"-1",
"--multi-item-scoring-delimiter",
str(_LLAMA3_EOT_TOKEN_ID),
"--enable-mis",
"--attention-backend",
"flashinfer",
],
)
@@ -4,14 +4,18 @@ Two model types, two scoring modes:
TestCausalLMScoring CausalLM, single-item and batched multi-item
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
exercises tokenization, forward pass, pooling, and score extraction without
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 os
import unittest
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")
_CAUSAL_LM_MODEL = DEFAULT_SMALL_MODEL_NAME_FOR_TEST # Llama-3.2-1B-Instruct
_SEQCLS_MODEL = "Qwen/Qwen3-0.6B" # backbone; arch overridden to SeqCls below
# <|endoftext|> for Qwen3 tokenizer — used as MIS delimiter
_QWEN3_EOT_TOKEN_ID = 151643
_CAUSAL_LM_MODEL = os.environ.get("TEST_MODEL_NAME", DEFAULT_SMALL_MODEL_NAME_FOR_TEST)
_SEQCLS_MODEL = os.environ.get("TEST_CLASSIFICATION_BASE_MODEL", "Qwen/Qwen3-0.6B")
# ---------------------------------------------------------------------------
@@ -368,10 +370,9 @@ class TestSeqClsScoring(CustomTestCase):
class TestSeqClsMISScoring(CustomTestCase):
"""SeqCls MIS: all items packed into one sequence separated by delimiter token.
score_and_pool() extracts per-item scores at delimiter positions.
Two sub-cases are tested:
- NUM_LABELS=2 standard binary classification head
- NUM_LABELS=12 stress-tests 2-D tensor indexing in score_and_pool()
Uses --enable-mis which hardcodes delimiter token ID 9999.
Basic pipeline correctness only thorough MIS tests (parity,
concurrency, advanced) live in test_multi_item_scoring.py.
"""
NUM_LABELS = 2
@@ -382,7 +383,8 @@ class TestSeqClsMISScoring(CustomTestCase):
model_path=_SEQCLS_MODEL,
disable_radix_cache=True,
chunked_prefill_size=-1,
multi_item_scoring_delimiter=_QWEN3_EOT_TOKEN_ID,
enable_mis=True,
attention_backend="flashinfer",
json_model_override_args=json.dumps(
{
"architectures": ["Qwen3ForSequenceClassification"],
@@ -432,32 +434,6 @@ class TestSeqClsMISScoring(CustomTestCase):
self.assertEqual(len(row), self.NUM_LABELS)
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)
@@ -479,7 +455,8 @@ class TestSeqClsMISAdvancedScoring(CustomTestCase):
model_path=_SEQCLS_MODEL,
disable_radix_cache=True,
chunked_prefill_size=-1,
multi_item_scoring_delimiter=_QWEN3_EOT_TOKEN_ID,
enable_mis=True,
attention_backend="flashinfer",
json_model_override_args=json.dumps(
{
"architectures": ["Qwen3ForSequenceClassification"],
@@ -506,17 +483,6 @@ class TestSeqClsMISAdvancedScoring(CustomTestCase):
self.assertEqual(len(row), self.NUM_LABELS)
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__":
unittest.main(verbosity=3)
@@ -1,12 +1,11 @@
"""Unit tests for score_and_pool in sglang.srt.layers.pooler.
All tests run on CPU no GPU required. The global server_args singleton
is mocked so the tests are hermetic.
All tests run on CPU no GPU required. MIS delimiter positions are passed
via forward_batch.multi_item_delimiter_indices (pre-computed by the caller).
"""
import unittest
from types import SimpleNamespace
from unittest.mock import patch
import torch
import torch.nn as nn
@@ -24,22 +23,22 @@ register_cpu_ci(est_time=9, suite="stage-a-test-cpu")
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."""
return SimpleNamespace(
extend_seq_lens=torch.tensor(extend_seq_lens, dtype=torch.long),
extend_seq_lens_cpu=extend_seq_lens,
is_prefill_only=is_prefill_only,
multi_item_delimiter_indices=multi_item_delimiter_indices,
dimensions=None,
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):
"""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.pooler = Pooler(pooling_type=PoolingType.LAST, normalize=False)
@patch("sglang.srt.layers.pooler.get_global_server_args")
def test_single_item_returns_scores(self, mock_get_args):
"""No delimiter -> single-item path returns [batch, num_labels]."""
mock_get_args.return_value = _mock_server_args(delimiter=None)
def test_single_item_returns_scores(self):
"""No delimiter indices -> single-item path returns [batch, num_labels]."""
hidden = torch.randn(8, self.hidden_dim)
fb = _make_forward_batch(extend_seq_lens=[5, 3])
input_ids = torch.arange(8)
@@ -64,30 +60,16 @@ class TestScoreAndPool(CustomTestCase):
self.assertIsInstance(out, EmbeddingPoolerOutput)
self.assertEqual(out.embeddings.shape, (2, self.num_labels))
@patch("sglang.srt.layers.pooler.get_global_server_args")
def test_mis_returns_per_request_list(self, mock_get_args):
"""Delimiter found -> returns a list with one tensor per request."""
delimiter_token = 99
mock_get_args.return_value = _mock_server_args(delimiter=delimiter_token)
input_ids = torch.tensor(
[
0,
1,
2,
delimiter_token,
3,
4,
5,
delimiter_token,
6,
7,
8,
delimiter_token,
]
)
def test_mis_returns_per_request_list(self):
"""Delimiter indices provided -> returns a list with one tensor per request."""
# Sequence: [0, 1, 2, D, 3, 4, 5, D, 6, 7, 8, D]
# Delimiters at positions 3, 7, 11 -> extract at 2, 6, 10
input_ids = torch.arange(12)
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)
@@ -95,20 +77,20 @@ class TestScoreAndPool(CustomTestCase):
self.assertEqual(len(out.embeddings), 1)
self.assertEqual(out.embeddings[0].shape, (3, self.num_labels))
@patch("sglang.srt.layers.pooler.get_global_server_args")
def test_mis_batched_splits_per_request(self, mock_get_args):
def test_mis_batched_splits_per_request(self):
"""Two batched MIS requests -> returns a list of length 2."""
delimiter_token = 99
mock_get_args.return_value = _mock_server_args(delimiter=delimiter_token)
# Request 1: [10, 11, delim, 12, 13, delim] -> 2 delimiters
# Request 2: [20, 21, 22, delim] -> 1 delimiter
req1 = [10, 11, delimiter_token, 12, 13, delimiter_token]
req2 = [20, 21, 22, delimiter_token]
# Request 1: [10, 11, D, 12, 13, D] -> delimiters at 2, 5
# Request 2: [20, 21, 22, D] -> delimiter at 3
req1 = [10, 11, 99, 12, 13, 99]
req2 = [20, 21, 22, 99]
input_ids = torch.tensor(req1 + req2)
hidden = torch.randn(len(input_ids), self.hidden_dim)
fb = _make_forward_batch(
extend_seq_lens=[len(req1), len(req2)], is_prefill_only=True
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)
@@ -118,42 +100,21 @@ class TestScoreAndPool(CustomTestCase):
self.assertEqual(out.embeddings[0].shape, (2, self.num_labels))
self.assertEqual(out.embeddings[1].shape, (1, self.num_labels))
@patch("sglang.srt.layers.pooler.get_global_server_args")
def test_mis_falls_back_when_no_delimiters_in_input(self, mock_get_args):
"""Delimiter configured but absent from input_ids -> single-item fallback."""
mock_get_args.return_value = _mock_server_args(delimiter=99)
def test_no_delimiter_indices_falls_back(self):
"""multi_item_delimiter_indices=None -> single-item fallback."""
input_ids = torch.tensor([0, 1, 2, 3, 4, 5, 6, 7])
hidden = torch.randn(8, self.hidden_dim)
fb = _make_forward_batch(extend_seq_lens=[5, 3], is_prefill_only=True)
fb = _make_forward_batch(extend_seq_lens=[5, 3])
out = score_and_pool(self.score_head, self.pooler, hidden, fb, input_ids)
self.assertIsInstance(out.embeddings, torch.Tensor)
self.assertEqual(out.embeddings.shape, (2, self.num_labels))
@patch("sglang.srt.layers.pooler.get_global_server_args")
def test_mis_falls_back_when_not_prefill_only(self, mock_get_args):
"""Delimiter configured, is_prefill_only=False -> single-item fallback."""
mock_get_args.return_value = _mock_server_args(delimiter=99)
input_ids = torch.tensor([0, 1, 2, 99, 3, 4, 5, 99])
hidden = torch.randn(8, self.hidden_dim)
fb = _make_forward_batch(extend_seq_lens=[5, 3], is_prefill_only=False)
out = score_and_pool(self.score_head, self.pooler, hidden, fb, input_ids)
self.assertIsInstance(out.embeddings, torch.Tensor)
self.assertEqual(out.embeddings.shape, (2, self.num_labels))
@patch("sglang.srt.layers.pooler.get_global_server_args")
def test_mis_extracts_positions_before_delimiter(self, mock_get_args):
def test_mis_extracts_positions_before_delimiter(self):
"""Verify MIS picks hidden states at index (delimiter_position - 1)."""
delimiter_token = 99
mock_get_args.return_value = _mock_server_args(delimiter=delimiter_token)
# Delimiters at indices 2 and 5 -> extract hidden at indices 1 and 4
input_ids = torch.tensor([10, 11, delimiter_token, 20, 21, delimiter_token])
input_ids = torch.tensor([10, 11, 99, 20, 21, 99])
hidden = (
torch.arange(len(input_ids))
.unsqueeze(1)
@@ -161,7 +122,10 @@ class TestScoreAndPool(CustomTestCase):
.expand(-1, self.hidden_dim)
.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)
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[1], hidden[4])
@patch("sglang.srt.layers.pooler.get_global_server_args")
def test_mis_ignores_delimiter_at_position_zero(self, mock_get_args):
"""A delimiter at flat index 0 has no preceding token and must be skipped."""
delimiter_token = 99
mock_get_args.return_value = _mock_server_args(delimiter=delimiter_token)
# Delimiter at index 0 should be ignored; only the one at index 3 counts
input_ids = torch.tensor([delimiter_token, 10, 11, delimiter_token])
def test_mis_delimiter_at_position_one(self):
"""Delimiters at positions 1 and 3 extract at indices 0 and 2."""
input_ids = torch.tensor([10, 99, 11, 99])
hidden = (
torch.arange(len(input_ids))
.unsqueeze(1)
@@ -187,7 +146,10 @@ class TestScoreAndPool(CustomTestCase):
.expand(-1, self.hidden_dim)
.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)
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)
self.assertEqual(len(out.embeddings), 1)
self.assertEqual(out.embeddings[0].shape[0], 1)
torch.testing.assert_close(out.embeddings[0][0], hidden[2])
@patch("sglang.srt.layers.pooler.get_global_server_args")
def test_single_item_scores_match_manual_computation(self, mock_get_args):
"""Single-item scores equal score_head applied to all tokens then pooled."""
mock_get_args.return_value = _mock_server_args(delimiter=None)
self.assertEqual(out.embeddings[0].shape[0], 2)
torch.testing.assert_close(out.embeddings[0][0], hidden[0])
torch.testing.assert_close(out.embeddings[0][1], hidden[2])
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)
fb = _make_forward_batch(extend_seq_lens=[5, 3])
input_ids = torch.arange(8)
out = score_and_pool(self.score_head, self.pooler, hidden, fb, input_ids)
# score-first-then-pool: matches the original Qwen3/Qwen2 classification forward
logits = self.score_head(hidden)
expected = self.pooler(logits, fb).embeddings
pooled = self.pooler(hidden, fb).embeddings
expected = self.score_head(pooled)
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__":
unittest.main()