[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:
co-authored by
Chanh Nguyen
Sundara Raman Ramachandran
parent
cfd49e233c
commit
a8e3a534a4
@@ -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)
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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"],
|
||||
|
||||
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user