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

Co-authored-by: Chanh Nguyen <chanhnguyen@gmail.com>
Co-authored-by: Sundara Raman Ramachandran <sundar24295@gmail.com>
This commit is contained in:
jsheng_Linkedin
2026-04-20 22:50:40 -07:00
committed by GitHub
co-authored by Chanh Nguyen Sundara Raman Ramachandran
parent cfd49e233c
commit a8e3a534a4
17 changed files with 1071 additions and 431 deletions
@@ -126,10 +126,8 @@ class FlashInferAttnBackend(AttentionBackend):
self.prefill_backend = "fa2"
self.decode_backend = "fa2"
# Store multi-item scoring delimiter for efficient access
self.multi_item_scoring_delimiter = (
model_runner.server_args.multi_item_scoring_delimiter
)
# Store multi-item scoring flag for efficient access
self.enable_mis = model_runner.server_args.enable_mis
# FIXME: remove dllm workarounds from flashinfer
self.dllm_config = DllmConfig.from_server_args(model_runner.server_args)
@@ -342,15 +340,18 @@ class FlashInferAttnBackend(AttentionBackend):
- max_item_len_ptr: [2, 3] (max lengths per sequence)
"""
delimiter = self.multi_item_scoring_delimiter
if delimiter is None or forward_batch.forward_mode == ForwardMode.DECODE:
if not self.enable_mis or forward_batch.forward_mode == ForwardMode.DECODE:
return MultiItemScoringParams()
delimiter_mask = forward_batch.input_ids == delimiter
prefix_cache_lens = getattr(forward_batch, "extend_prefix_lens", None)
extend_seq_lens = getattr(forward_batch, "extend_seq_lens", None)
precomputed_indices = forward_batch.multi_item_delimiter_indices
if precomputed_indices is None:
return MultiItemScoringParams()
prefix_cache_lens = getattr(forward_batch, "extend_prefix_lens_cpu", None)
extend_seq_lens = getattr(forward_batch, "extend_seq_lens_cpu", None)
prefix_len_ptr, token_pos_in_items_ptr = [], []
token_pos_in_items_len = 0
device = forward_batch.input_ids.device
# If no extend_seq_lens, treat whole batch as one sequence
if extend_seq_lens is None or len(extend_seq_lens) <= 1:
@@ -359,35 +360,44 @@ class FlashInferAttnBackend(AttentionBackend):
seq_start = 0
for i, seq_len in enumerate(extend_seq_lens):
seq_end = seq_start + seq_len
mask = delimiter_mask[seq_start:seq_end]
pos = forward_batch.positions[seq_start:seq_end]
delimiter_indices = torch.nonzero(mask, as_tuple=True)[0]
delimiter_indices_cpu = precomputed_indices[i]
if len(delimiter_indices_cpu) == 0:
seq_start = seq_end
continue
if len(delimiter_indices) > 0:
first_delim = delimiter_indices[0]
# Prefix length: store as scalar
prefix_len = first_delim + (
prefix_cache_lens[i] if prefix_cache_lens is not None else 0
)
prefix_len_ptr.append(
prefix_len.item() if torch.is_tensor(prefix_len) else prefix_len
)
first_delim = delimiter_indices_cpu[0].item() # CPU .item(), no GPU sync
delimiter_indices = delimiter_indices_cpu.to(device, non_blocking=True)
prefix_len = first_delim + (
prefix_cache_lens[i] if prefix_cache_lens is not None else 0
)
prefix_len_ptr.append(prefix_len)
# Compute relative positions within items after delimiters
diff = pos[first_delim:] - torch.cummax(mask[first_delim:], 0)[1]
token_pos = (diff - pos[first_delim]).to(torch.uint16)
token_pos_in_items_ptr.append(token_pos)
# Compute relative positions within items using searchsorted (no GPU sync).
# suffix_range = [0, 1, 2, 3, 4, ...]
# searchsorted = bucket index for each position
# last_delim = delimiter offset at start of current bucket
# pos_within_item = suffix_range - last_delim
suffix_len = seq_len - first_delim
relative_positions = delimiter_indices - first_delim
# Update forward_batch positions in-place
pos[first_delim:] = diff - 1
forward_batch.positions[seq_start:seq_end] = pos
suffix_range = torch.arange(suffix_len, dtype=torch.int64, device=device)
bucket_idx = torch.searchsorted(
relative_positions, suffix_range, right=True
)
last_delim = relative_positions[torch.clamp(bucket_idx - 1, min=0)]
pos_within_item = suffix_range - last_delim
token_pos_in_items_ptr.append(pos_within_item.to(torch.uint16))
forward_batch.positions[seq_start + first_delim : seq_end] = (
prefix_len + pos_within_item - 1
)
seq_start = seq_end
# Pad token_pos_in_items_ptr for batch processing
if token_pos_in_items_ptr:
token_pos_in_items_len = max(t.numel() for t in token_pos_in_items_ptr)
device = forward_batch.input_ids.device
token_pos_in_items_ptr = [
torch.cat(
[
@@ -405,8 +415,6 @@ class FlashInferAttnBackend(AttentionBackend):
if not prefix_len_ptr or not token_pos_in_items_ptr:
return MultiItemScoringParams()
# Build final params
device = forward_batch.input_ids.device
return MultiItemScoringParams(
prefix_len_ptr=torch.tensor(
prefix_len_ptr, dtype=torch.uint32, device=device
@@ -470,7 +478,7 @@ class FlashInferAttnBackend(AttentionBackend):
prefix_lens = forward_batch.extend_prefix_lens
# Disable ragged wrapper and ensure prefix handling for multimodal and multi-item scoring
if self.is_multimodal or self.multi_item_scoring_delimiter is not None:
if self.is_multimodal or self.enable_mis:
# use_ragged = False: Multi-item scoring requires the paged wrapper because:
# 1. Ragged wrapper doesn't support the specialized multi-item parameters
# (prefix_len_ptr, token_pos_in_items_ptr, etc.)
@@ -487,7 +495,7 @@ class FlashInferAttnBackend(AttentionBackend):
# Process multi-item scoring in attention backend instead of ForwardBatch
multi_item_params = MultiItemScoringParams()
if self.multi_item_scoring_delimiter is not None:
if self.enable_mis:
# Use new backend-specific implementation
multi_item_params = self._process_multi_item_scoring(forward_batch)
+49 -51
View File
@@ -274,9 +274,7 @@ class LogitsProcessor(nn.Module):
self.final_logit_softcapping = None
self.return_full_logits = return_full_logits
self.multi_item_delimiter = (
get_global_server_args().multi_item_scoring_delimiter
)
self.enable_mis = get_global_server_args().enable_mis
# enable chunked logprobs processing
self.enable_logprobs_chunk = envs.SGLANG_ENABLE_LOGITS_PROCESSER_CHUNK.get()
@@ -292,17 +290,20 @@ class LogitsProcessor(nn.Module):
aux_hidden_states: Optional[torch.Tensor] = None,
hidden_states_before_norm: Optional[torch.Tensor] = None,
) -> LogitsProcessorOutput:
# Extract MIS indices before ForwardBatch → LogitsMetadata conversion
multi_item_delimiter_indices = None
if isinstance(logits_metadata, ForwardBatch):
multi_item_delimiter_indices = logits_metadata.multi_item_delimiter_indices
logits_metadata = LogitsMetadata.from_forward_batch(logits_metadata)
# Multi-item scoring only for prefill-only requests.
if self.multi_item_delimiter is not None and logits_metadata.is_prefill_only:
# Multi-item scoring only for prefill-only requests with pre-computed indices.
if multi_item_delimiter_indices is not None and logits_metadata.is_prefill_only:
return self.compute_logprobs_for_multi_item_scoring(
input_ids,
hidden_states,
lm_head,
logits_metadata,
self.multi_item_delimiter,
multi_item_delimiter_indices,
)
# Diffusion LLM only.
@@ -347,6 +348,9 @@ class LogitsProcessor(nn.Module):
return LogitsProcessorOutput(
next_token_logits=sampled_logits,
hidden_states=hidden_states_to_store,
# FIXME: These fields are not logits-related but are passed through here as a
# workaround since ForwardBatch is local to forward_batch_generation().
# They should be moved to GenerationBatchResult to keep this class clean.
mm_input_embeds=logits_metadata.mm_input_embeds,
)
@@ -1006,39 +1010,41 @@ class LogitsProcessor(nn.Module):
hidden_states,
lm_head: VocabParallelEmbedding,
logits_metadata: Union[LogitsMetadata, ForwardBatch],
delimiter_token: int,
multi_item_delimiter_indices: List[torch.Tensor],
):
"""
Compute logprobs for multi-item scoring using delimiter-based token extraction.
This method is designed for scenarios where you want to score multiple items/candidates
against a single query by combining them into one sequence separated by delimiters.
Compute logprobs for multi-item scoring using pre-computed delimiter indices.
Sequence format: Query<delimiter>Item1<delimiter>Item2<delimiter>...
Scoring positions: Extracts logprobs at positions before each <delimiter>
Args:
input_ids (torch.Tensor): Input token IDs containing query and items separated by delimiters.
Shape: [total_sequence_length] for single request or [batch_total_length] for batch.
hidden_states (torch.Tensor): Hidden states from the model.
Shape: [sequence_length, hidden_dim].
lm_head (VocabParallelEmbedding): Language model head for computing logits.
logits_metadata (Union[LogitsMetadata, ForwardBatch]): Metadata containing batch info
and token ID specifications for logprob extraction.
delimiter_token (int): Token ID used as delimiter between query and items.
Returns:
LogitsProcessorOutput: Contains:
- next_token_logits: None (not needed for scoring-only requests)
- input_token_logprobs: Logprobs of delimiter tokens at scoring positions
- input_top_logprobs_val: Top-k logprobs at delimiter positions (if requested)
- input_top_logprobs_idx: Top-k token indices at delimiter positions (if requested)
- input_token_ids_logprobs_val: Logprobs for user-requested token IDs (if any)
- input_token_ids_logprobs_idx: Indices for user-requested token IDs (if any)
input_ids: Input token IDs. Shape: [total_sequence_length].
hidden_states: Hidden states from the model. Shape: [sequence_length, hidden_dim].
lm_head: Language model head for computing logits.
logits_metadata: Metadata containing batch info and logprob specs.
multi_item_delimiter_indices: Pre-computed delimiter positions per request (CPU tensors).
"""
multi_item_indices = (input_ids == delimiter_token).nonzero(as_tuple=True)[
0
] - 1
# Compute positions just before each delimiter.
# Build offset-adjusted indices on CPU, then do a single CPU→GPU transfer.
device = input_ids.device
all_tensors = []
if logits_metadata.extend_seq_lens_cpu is not None:
offset = 0
for req_seq_len, indices_tensor in zip(
logits_metadata.extend_seq_lens_cpu, multi_item_delimiter_indices
):
if len(indices_tensor) > 0:
# Note: if the first delimiter is at position 0 (empty query),
# indices - 1 wraps to -1. This is harmless — the first
# delimiter entry is always discarded by
# _process_multi_item_scoring_results.
all_tensors.append(indices_tensor + (offset - 1))
offset += req_seq_len
else:
all_tensors.append(multi_item_delimiter_indices[0] - 1)
multi_item_indices = torch.cat(all_tensors).to(device, non_blocking=True)
# Extract hidden states at delimiter positions for multi-item scoring
sliced_hidden = hidden_states[multi_item_indices]
@@ -1052,27 +1058,13 @@ class LogitsProcessor(nn.Module):
input_top_logprobs_idx = None
# Recalculate extend_logprob_pruned_lens_cpu to match delimiter counts per request
# Original contains sequence lengths, but we need delimiter counts for sliced_logprobs
if (
logits_metadata.token_ids_logprobs
or logits_metadata.extend_return_top_logprob
):
logits_metadata.extend_logprob_pruned_lens_cpu = []
if logits_metadata.extend_seq_lens_cpu is not None:
# Multi-request batch: count delimiters per request
input_pt = 0
for req_seq_len in logits_metadata.extend_seq_lens_cpu:
req_input_ids = input_ids[input_pt : input_pt + req_seq_len]
delimiter_count = (req_input_ids == delimiter_token).sum().item()
logits_metadata.extend_logprob_pruned_lens_cpu.append(
delimiter_count
)
input_pt += req_seq_len
else:
# Single request case: one request gets all delimiters
total_delimiters = (input_ids == delimiter_token).sum().item()
logits_metadata.extend_logprob_pruned_lens_cpu = [total_delimiters]
logits_metadata.extend_logprob_pruned_lens_cpu = [
len(t) for t in multi_item_delimiter_indices
]
# Get the logprobs of specified token ids
if logits_metadata.extend_token_ids_logprob:
@@ -1090,11 +1082,17 @@ class LogitsProcessor(nn.Module):
input_top_logprobs_idx,
) = get_top_logprobs_prefill(sliced_logprobs, logits_metadata)
# For input_token_logprobs, use delimiter token logprobs
input_token_logprobs = sliced_logprobs[:, delimiter_token]
# MIS scores come from input_token_ids_logprobs_val (label-token logprobs),
# not from per-position input_token_logprobs. However, the shared logprob
# pipeline (add_input_logprob_return_values) asserts input_token_logprobs is
# non-None, converts it to a tuple, slices it, and validates its length —
# all before score_request() ever sees the result. We can't set it to None
# without changing those shared asserts, so we fill with zeros to satisfy
# the pipeline. score_request() ignores this field entirely.
input_token_logprobs = torch.zeros(multi_item_indices.shape[0], device=device)
return LogitsProcessorOutput(
next_token_logits=None, # Multi-item scoring doesn't need next token logits
next_token_logits=None,
input_token_logprobs=input_token_logprobs,
input_top_logprobs_val=input_top_logprobs_val,
input_top_logprobs_idx=input_top_logprobs_idx,
+65 -37
View File
@@ -12,7 +12,6 @@ import torch.nn as nn
from transformers import PretrainedConfig
from sglang.srt.layers.activation import get_cross_encoder_activation_function
from sglang.srt.server_args import get_global_server_args
if TYPE_CHECKING:
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
@@ -66,6 +65,46 @@ def pool_hidden_states(
raise ValueError(f"Unsupported pooling type: {pooling_type}")
def pool_at_delimiter_positions(
data: torch.Tensor,
forward_batch: ForwardBatch,
device: torch.device,
) -> List[torch.Tensor]:
"""Pool a tensor at the position before each MIS delimiter for every request.
Uses pre-computed delimiter indices from ForwardBatch (CPU tensors),
moves to GPU with non_blocking=True to avoid CUDA syncs.
Args:
data: 2-D tensor [total_tokens, dim] — hidden states or logits.
forward_batch: Forward batch with extend_seq_lens_cpu and
multi_item_delimiter_indices populated.
device: Device for the index tensor.
Returns:
One tensor per request, shaped [num_delimiters, dim].
"""
all_index_tensors: List[torch.Tensor] = []
delim_counts: List[int] = []
offset = 0
for req_idx, req_seq_len in enumerate(forward_batch.extend_seq_lens_cpu):
indices_tensor = forward_batch.multi_item_delimiter_indices[req_idx]
n = len(indices_tensor)
if n > 0:
# Note: if the first delimiter is at position 0 (empty query),
# indices - 1 wraps to -1. This is harmless — the first delimiter
# entry is always discarded by _process_multi_item_scoring_results.
all_index_tensors.append(indices_tensor + (offset - 1))
delim_counts.append(n)
offset += req_seq_len
if all_index_tensors:
index_tensor = torch.cat(all_index_tensors).to(device, non_blocking=True)
else:
index_tensor = torch.tensor([], dtype=torch.long, device=device)
return list(data[index_tensor].split(delim_counts))
def score_and_pool(
score_head: nn.Module,
pooler: "Pooler",
@@ -75,47 +114,36 @@ def score_and_pool(
) -> EmbeddingPoolerOutput:
"""Apply a classification/score head with MIS and pooled-hidden-states support.
MIS path (when ``multi_item_scoring_delimiter`` is set and found in ``input_ids``):
extract hidden states at positions just before each delimiter, apply the score head,
then split per-request.
MIS path (pre-computed delimiter indices on forward_batch): extract hidden
states at positions just before each delimiter, apply the score head, then
split per-request.
Standard path: apply the score head to all hidden states, then pool.
Standard path: pool hidden states, then apply the score head.
When ``forward_batch.return_pooled_hidden_states`` is True, the raw pooled
hidden states (before the score head) are included in the output.
"""
delimiter_token = get_global_server_args().multi_item_scoring_delimiter
if delimiter_token is not None and forward_batch.is_prefill_only:
delim_positions = (input_ids == delimiter_token).nonzero(as_tuple=True)[0]
# A delimiter at flat index 0 has no preceding hidden state to pool
delim_positions = delim_positions[delim_positions > 0]
if delim_positions.numel() > 0:
# Score only the tokens that precede a delimiter
pre_delim_hidden = hidden_states[delim_positions - 1]
scores = score_head(pre_delim_hidden)
# Split per-request so the scheduler gets one tensor per request.
# Use CPU sequence lengths to avoid per-iteration GPU<->CPU sync
# from `.item()` calls on device tensors.
seq_lens = forward_batch.extend_seq_lens_cpu
start = 0
per_request_scores: List[torch.Tensor] = []
per_request_phs: Optional[List[torch.Tensor]] = (
[] if forward_batch.return_pooled_hidden_states else None
)
for seq_len in seq_lens:
end = start + seq_len
mask = (delim_positions >= start) & (delim_positions < end)
per_request_scores.append(scores[mask])
if per_request_phs is not None:
per_request_phs.append(pre_delim_hidden[mask])
start = end
return EmbeddingPoolerOutput(
embeddings=per_request_scores,
pooled_hidden_states=per_request_phs,
)
if (
forward_batch.multi_item_delimiter_indices is not None
and forward_batch.is_prefill_only
):
# Pool hidden states at pre-delimiter positions, score only those —
# avoids wasting compute on tokens that never contribute to the output.
# pool_at_delimiter_positions returns one tensor per request; we concat
# to call score_head once, then split back per request.
per_request_phs = pool_at_delimiter_positions(
hidden_states, forward_batch, input_ids.device
)
phs_flat = torch.cat(per_request_phs, dim=0)
scores_flat = score_head(phs_flat)
delim_counts = [t.shape[0] for t in per_request_phs]
per_request_scores = list(scores_flat.split(delim_counts))
return EmbeddingPoolerOutput(
embeddings=per_request_scores,
pooled_hidden_states=(
per_request_phs if forward_batch.return_pooled_hidden_states else None
),
)
# Standard classification path: pool hidden states, then score.
pooled_hs = pool_hidden_states(pooler.pooling_type, hidden_states, forward_batch)
+28
View File
@@ -250,6 +250,10 @@ class GenerateReqInput(BaseReq):
image_max_dynamic_patch: Optional[int] = None
video_max_dynamic_patch: Optional[int] = None
# Pre-computed delimiter indices for multi-item scoring.
# Batch-level: List[List[int]] (one per request). After __getitem__: List[int].
multi_item_delimiter_indices: Optional[Union[List[List[int]], List[int]]] = None
def contains_mm_input(self) -> bool:
return (
has_valid_data(self.image_data)
@@ -685,6 +689,11 @@ class GenerateReqInput(BaseReq):
external_trace_header=self.external_trace_header,
http_worker_ipc=self.http_worker_ipc,
received_time=self.received_time,
multi_item_delimiter_indices=(
self.multi_item_delimiter_indices[i]
if self.multi_item_delimiter_indices is not None
else None
),
)
cache[i] = sub
return sub
@@ -774,6 +783,9 @@ class TokenizedGenerateReqInput(BaseReq):
need_wait_for_mm_inputs: bool = False
num_items_assigned: Optional[Dict[Modality, List[int]]] = None
# Pre-computed delimiter indices for multi-item scoring
multi_item_delimiter_indices: Optional[List[int]] = None
# For observability
time_stats: Optional[Union[APIServerReqTimeStats, DPControllerReqTimeStats]] = None
@@ -855,6 +867,10 @@ class EmbeddingReqInput(BaseReq):
# Whether to return pooled hidden states (pre-head transformer output)
return_pooled_hidden_states: bool = False
# Pre-computed delimiter indices for multi-item scoring.
# Batch-level: List[List[int]] (one per request). After __getitem__: List[int].
multi_item_delimiter_indices: Optional[Union[List[List[int]], List[int]]] = None
def normalize_batch_and_arguments(self):
# at least one of text, input_ids, or image should be provided
if self.text is None and self.input_ids is None and self.image_data is None:
@@ -957,6 +973,11 @@ class EmbeddingReqInput(BaseReq):
is_cross_encoder_request=True,
http_worker_ipc=self.http_worker_ipc,
return_pooled_hidden_states=self.return_pooled_hidden_states,
multi_item_delimiter_indices=(
self.multi_item_delimiter_indices[i]
if self.multi_item_delimiter_indices is not None
else None
),
)
else:
sub = EmbeddingReqInput(
@@ -981,6 +1002,11 @@ class EmbeddingReqInput(BaseReq):
http_worker_ipc=self.http_worker_ipc,
received_time=self.received_time,
return_pooled_hidden_states=self.return_pooled_hidden_states,
multi_item_delimiter_indices=(
self.multi_item_delimiter_indices[i]
if self.multi_item_delimiter_indices is not None
else None
),
)
cache[i] = sub
return sub
@@ -1009,6 +1035,8 @@ class TokenizedEmbeddingReqInput(BaseReq):
# LoRA related
lora_id: Optional[str] = None # None means just use the base model
# Pre-computed delimiter indices for multi-item scoring
multi_item_delimiter_indices: Optional[List[int]] = None
# For observability
time_stats: Optional[Union[APIServerReqTimeStats, DPControllerReqTimeStats]] = None
@@ -597,6 +597,7 @@ class Req(ReqDllmMixin):
Union[APIServerReqTimeStats, DPControllerReqTimeStats]
] = None,
return_pooled_hidden_states: bool = False,
multi_item_delimiter_indices: Optional[List[int]] = None,
):
# Input and output info
self.rid = rid
@@ -614,6 +615,7 @@ class Req(ReqDllmMixin):
self.session = session
self.input_embeds = input_embeds
self.positional_embed_overrides = positional_embed_overrides
self.multi_item_delimiter_indices = multi_item_delimiter_indices
# For req-level memory management
self.kv_committed_len = 0
@@ -1441,6 +1443,9 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
# Whether this batch is prefill-only (no token generation needed)
is_prefill_only: bool = False
# Multi-item scoring delimiter indices (set during prepare_for_extend)
multi_item_delimiter_indices: Optional[List[torch.Tensor]] = None
# hicache pointer for synchronizing data loading from CPU to GPU
hicache_consumer_index: int = -1
@@ -1817,6 +1822,23 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
self.token_type_ids = token_type_ids_tensor
self.seq_lens_sum = sum(seq_lens)
# Pre-compute delimiter indices as CPU tensors for MIS.
# When --enable-mis is on, every request in the batch is expected to
# carry delimiter indices (the score endpoint always produces MIS-structured
# requests). Consumers index this list without None-checking.
if get_global_server_args().enable_mis and any(
r.multi_item_delimiter_indices is not None for r in reqs
):
assert all(
r.multi_item_delimiter_indices is not None for r in reqs
), "MIS batch must have delimiter indices on every request"
self.multi_item_delimiter_indices = [
torch.tensor(r.multi_item_delimiter_indices, dtype=torch.int64)
for r in reqs
]
else:
self.multi_item_delimiter_indices = None
if self.return_logprob:
self.top_logprobs_nums = [r.top_logprobs_num for r in reqs]
self.token_ids_logprobs = [r.token_ids_logprob for r in reqs]
@@ -2464,6 +2486,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
),
extend_input_logprob_token_ids=self.extend_input_logprob_token_ids,
is_prefill_only=self.is_prefill_only,
multi_item_delimiter_indices=self.multi_item_delimiter_indices,
dimensions=self.dimensions,
return_pooled_hidden_states=self.return_pooled_hidden_states,
dllm_block_offsets=[req.dllm_block_offset for req in self.reqs],
@@ -2665,6 +2688,9 @@ class ModelWorkerBatch:
# Whether this batch is prefill-only (no token generation needed)
is_prefill_only: bool = False
# Pre-computed delimiter indices for multi-item scoring (CPU tensors, one per request)
multi_item_delimiter_indices: Optional[List[torch.Tensor]] = None
# Diffusion LLM
dllm_block_offsets: Optional[List[int]] = None
dllm_config: Optional[DllmConfig] = None
+2
View File
@@ -1881,6 +1881,7 @@ class Scheduler(
http_worker_ipc=recv_req.http_worker_ipc,
dllm_config=self.dllm_config,
time_stats=recv_req.time_stats,
multi_item_delimiter_indices=recv_req.multi_item_delimiter_indices,
)
req.tokenizer = self.tokenizer
@@ -2202,6 +2203,7 @@ class Scheduler(
http_worker_ipc=recv_req.http_worker_ipc,
time_stats=recv_req.time_stats,
return_pooled_hidden_states=recv_req.return_pooled_hidden_states,
multi_item_delimiter_indices=recv_req.multi_item_delimiter_indices,
)
req.tokenizer = self.tokenizer
@@ -21,7 +21,7 @@ from sglang.srt.managers.schedule_batch import (
ScheduleBatch,
)
from sglang.srt.mem_cache.common import release_kv_cache
from sglang.srt.server_args import get_global_server_args
from sglang.srt.server_args import MIS_DELIMITER_TOKEN_ID, get_global_server_args
if TYPE_CHECKING:
from sglang.srt.managers.scheduler import (
@@ -629,13 +629,12 @@ class SchedulerOutputProcessorMixin:
# Process logprob indices based on scoring type
if is_multi_item_scoring:
# Multi-item scoring: only include delimiter token positions
relevant_tokens = req.origin_input_ids[req.logprob_start_len :]
input_token_logprobs_idx = [
token_id
for token_id in relevant_tokens
if token_id == self.server_args.multi_item_scoring_delimiter
]
# MIS scores come from input_token_ids_logprobs, not input_token_logprobs.
# But the shared pipeline requires input_token_logprobs_idx to be the same
# length as input_token_logprobs_val (validated at line 816). We fill with
# MIS_DELIMITER_TOKEN_ID as a dummy — score_request() ignores this field.
delimiter_count = len(req.multi_item_delimiter_indices)
input_token_logprobs_idx = [MIS_DELIMITER_TOKEN_ID] * delimiter_count
else:
# Regular request: include all tokens from logprob_start_len onwards
input_token_logprobs_idx = req.origin_input_ids[req.logprob_start_len :]
@@ -714,18 +713,11 @@ class SchedulerOutputProcessorMixin:
For regular requests, all positions from logprob_start_len onwards have logprobs.
"""
is_multi_item_scoring = self._is_multi_item_scoring(req)
relevant_tokens = req.origin_input_ids[req.logprob_start_len :]
if is_multi_item_scoring:
# Multi-item scoring: count delimiter tokens from logprob_start_len onwards
return sum(
1
for token_id in relevant_tokens
if token_id == self.server_args.multi_item_scoring_delimiter
)
return len(req.multi_item_delimiter_indices)
else:
# Regular request: all tokens from logprob_start_len onwards
return len(relevant_tokens)
return len(req.origin_input_ids[req.logprob_start_len :])
def _calculate_num_input_logprobs(
self: Scheduler, req: Req, extend_input_len: int, extend_logprob_start_len: int
@@ -738,14 +730,11 @@ class SchedulerOutputProcessorMixin:
is_multi_item_scoring = self._is_multi_item_scoring(req)
if is_multi_item_scoring:
# Multi-item scoring: count delimiter tokens in the relevant portion
relevant_tokens = req.origin_input_ids[
extend_logprob_start_len:extend_input_len
]
# Count pre-computed delimiter indices within the extend range
return sum(
1
for token_id in relevant_tokens
if token_id == self.server_args.multi_item_scoring_delimiter
for idx in req.multi_item_delimiter_indices
if extend_logprob_start_len <= idx < extend_input_len
)
else:
# Regular request: all tokens in the range
@@ -758,7 +747,11 @@ class SchedulerOutputProcessorMixin:
token is configured. In this mode, only positions containing the
delimiter token receive logprobs.
"""
return req.is_prefill_only and self.server_args.multi_item_scoring_delimiter
return (
self.server_args.enable_mis
and req.is_prefill_only
and req.multi_item_delimiter_indices is not None
)
def add_input_logprob_return_values(
self: Scheduler,
@@ -313,7 +313,6 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
self.processor = _processor
self.tokenizer = get_tokenizer_from_processor(self.processor)
os.environ["TOKENIZERS_PARALLELISM"] = "false"
self._initialize_multi_item_delimiter_text()
else:
self.mm_processor = self.processor = None
@@ -326,7 +325,6 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
trust_remote_code=server_args.trust_remote_code,
revision=server_args.revision,
)
self._initialize_multi_item_delimiter_text()
# Initialize async dynamic batch tokenizer if enabled (common for both multimodal and non-multimodal)
if (
@@ -1007,6 +1005,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
token_type_ids=token_type_ids,
need_wait_for_mm_inputs=obj.need_wait_for_mm_inputs,
num_items_assigned=obj.num_items_assigned,
multi_item_delimiter_indices=obj.multi_item_delimiter_indices,
)
elif isinstance(obj, EmbeddingReqInput):
# Resolve unresolved embed overrides now that input_ids are available
@@ -1033,6 +1032,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
lora_id=obj.lora_id,
http_worker_ipc=obj.http_worker_ipc,
return_pooled_hidden_states=obj.return_pooled_hidden_states,
multi_item_delimiter_indices=obj.multi_item_delimiter_indices,
)
tokenized_obj.time_stats = self.rid_to_state[obj.rid].time_stats
@@ -8,6 +8,7 @@ import torch
from sglang.srt.configs.model_config import is_cross_encoding_pooler_model
from sglang.srt.managers.embed_types import PositionalEmbeds
from sglang.srt.managers.io_struct import EmbeddingReqInput, GenerateReqInput
from sglang.srt.server_args import MIS_DELIMITER_TOKEN_ID
logger = logging.getLogger(__name__)
@@ -76,27 +77,9 @@ class TokenizerManagerScoreMixin:
raise ValueError("Invalid prompts type for score_prompts.")
def _initialize_multi_item_delimiter_text(self):
"""Initialize multi-item delimiter text from token ID after tokenizer is loaded."""
if (
hasattr(self.server_args, "multi_item_scoring_delimiter")
and self.server_args.multi_item_scoring_delimiter is not None
and self.tokenizer is not None
):
try:
self.multi_item_delimiter_text = self.tokenizer.decode(
[self.server_args.multi_item_scoring_delimiter],
skip_special_tokens=False,
)
except Exception as e:
logger.warning(
f"Failed to decode delimiter token {self.server_args.multi_item_scoring_delimiter}: {e}"
)
self.multi_item_delimiter_text = None
def _build_multi_item_token_sequence(
self, query: List[int], items: List[List[int]], delimiter_token_id: int
) -> List[int]:
) -> Tuple[List[int], List[int]]:
"""
Build a single token sequence for multi-item scoring.
Format: query<delimiter>item1<delimiter>item2<delimiter>item3<delimiter>
@@ -107,18 +90,21 @@ class TokenizerManagerScoreMixin:
delimiter_token_id: Token ID to use as delimiter
Returns:
Combined token sequence
Tuple of (combined token sequence, delimiter indices)
"""
combined_sequence = query[:] # Start with query
delimiter_indices = []
for item in items:
delimiter_indices.append(len(combined_sequence))
combined_sequence.append(delimiter_token_id) # Add delimiter
combined_sequence.extend(item) # Add item tokens
# Add final delimiter after the last item for logprob extraction
delimiter_indices.append(len(combined_sequence))
combined_sequence.append(delimiter_token_id)
return combined_sequence
return combined_sequence, delimiter_indices
def _batch_tokenize_query_and_items(
self,
@@ -416,11 +402,14 @@ class TokenizerManagerScoreMixin:
embed_override_token_id: Optional[int],
query_embed_overrides: Optional[List[torch.Tensor]],
item_embed_overrides: Optional[List[Optional[List[torch.Tensor]]]],
) -> Tuple[None, List[List[int]], Optional[list]]:
) -> Tuple[None, List[List[int]], Optional[list], Optional[List[int]]]:
"""Build input_ids and resolve embed overrides for token-ID inputs.
Works identically for multi-item-scoring and single-item modes the only difference is
how input_ids are assembled and what position offset each item gets.
Returns:
(text_prompts, input_ids, positional_embed_overrides, delimiter_indices)
"""
# Both query and items are token IDs
has_embeds = (
@@ -428,16 +417,17 @@ class TokenizerManagerScoreMixin:
)
if use_multi_item_scoring:
# Multi-item scoring: concatenate with delimiter token ID
# Format: query<delimiter_token_id>item1<delimiter_token_id>item2<delimiter_token_id>item3<delimiter_token_id>
delimiter_token_id = self.server_args.multi_item_scoring_delimiter
combined_input_ids = self._build_multi_item_token_sequence(
query, items, delimiter_token_id
# Multi-item scoring: concatenate with placeholder delimiter token.
# Positions are derived from item lengths (delimiter_indices), not
# by scanning for this token — it exists only for FlashInfer compat.
delimiter_token_id = MIS_DELIMITER_TOKEN_ID
combined_input_ids, delimiter_indices = (
self._build_multi_item_token_sequence(query, items, delimiter_token_id)
)
input_ids = [combined_input_ids]
if not has_embeds:
return None, input_ids, None
return None, input_ids, None, delimiter_indices
# Resolve embed overrides across the combined multi-item-scoring sequence
all_embeds: List[torch.Tensor] = []
@@ -461,15 +451,15 @@ class TokenizerManagerScoreMixin:
current_offset += len(item) + 1 # +1 for delimiter
if all_embeds:
injection = [
positional_embed_overrides = [
PositionalEmbeds(
embeds=torch.cat(all_embeds, dim=0),
positions=all_positions,
)
]
else:
injection = None
return None, input_ids, injection
positional_embed_overrides = None
return None, input_ids, positional_embed_overrides, delimiter_indices
else:
# Single-item scoring: process each item separately
@@ -479,9 +469,9 @@ class TokenizerManagerScoreMixin:
input_ids = [query + item for item in items]
if not has_embeds:
return None, input_ids, None
return None, input_ids, None, None
injection = []
positional_embed_overrides = []
for i, item in enumerate(items):
item_embs = item_embed_overrides[i] if item_embed_overrides else None
pe = self._resolve_embed_overrides_for_request(
@@ -493,13 +483,14 @@ class TokenizerManagerScoreMixin:
item_position_offset=len(query),
item_label=f"items[{i}]",
)
injection.append(pe)
positional_embed_overrides.append(pe)
return (
None,
input_ids,
injection if any(pe is not None for pe in injection) else None,
positional_embed_overrides = (
positional_embed_overrides
if any(pe is not None for pe in positional_embed_overrides)
else None
)
return None, input_ids, positional_embed_overrides, None
# ------------------------------------------------------------------
# Main entry point
@@ -523,7 +514,7 @@ class TokenizerManagerScoreMixin:
This method supports two scoring approaches:
1. Single-Item scoring (default): Process each query+item pair independently
2. Multi-Item scoring: When multi_item_scoring_delimiter is set, combine query and
2. Multi-Item scoring: When --enable-mis is set, combine query and
multiple items into a single sequence using delimiter for efficient processing.
Note: item_first parameter is ignored in multi-item scoring mode since it uses
a fixed format: query<delimiter>item1<delimiter>item2<delimiter>item3<delimiter>
@@ -593,15 +584,13 @@ class TokenizerManagerScoreMixin:
f"Token ID {token_id} is out of vocabulary (vocab size: {vocab_size})"
)
# Check if multi-item scoring is enabled by presence of delimiter
use_multi_item_scoring = (
self.server_args.multi_item_scoring_delimiter is not None
and self.multi_item_delimiter_text is not None
)
# Check if multi-item scoring is enabled
use_multi_item_scoring = self.server_args.enable_mis
input_ids = None
text_prompts = None
positional_embed_overrides = None
delimiter_indices = None
use_text_prompts = isinstance(query, str) and not has_embeds
@@ -609,15 +598,17 @@ class TokenizerManagerScoreMixin:
# Both query and items are text
items_list = [items] if isinstance(items, str) else items
if use_multi_item_scoring:
# Multi-item scoring: tokenize separately then combine at token level
# to ensure the delimiter token ID is inserted exactly once per boundary
# (a text-level roundtrip through the tokenizer can alter boundary tokens)
delimiter_token_id = self.server_args.multi_item_scoring_delimiter
# Tokenize separately, then combine at token level with placeholder
# delimiter. Positions come from item lengths (delimiter_indices),
# not from scanning for this token — it's for FlashInfer compat only.
delimiter_token_id = MIS_DELIMITER_TOKEN_ID
query_ids, items_ids = self._batch_tokenize_query_and_items(
query, items_list
)
combined_input_ids = self._build_multi_item_token_sequence(
query_ids, items_ids, delimiter_token_id
combined_input_ids, delimiter_indices = (
self._build_multi_item_token_sequence(
query_ids, items_ids, delimiter_token_id
)
)
input_ids = [combined_input_ids]
else:
@@ -635,26 +626,30 @@ class TokenizerManagerScoreMixin:
):
# Both query and items are token IDs — tokenize text inputs if needed for embed overrides
query_ids, items_ids = query, items
_, input_ids, positional_embed_overrides = self._build_token_id_inputs(
query_ids,
items_ids,
item_first,
use_multi_item_scoring,
embed_override_token_id,
query_embed_overrides,
item_embed_overrides,
_, input_ids, positional_embed_overrides, delimiter_indices = (
self._build_token_id_inputs(
query_ids,
items_ids,
item_first,
use_multi_item_scoring,
embed_override_token_id,
query_embed_overrides,
item_embed_overrides,
)
)
elif has_embeds:
# Text inputs with embed overrides — need to tokenize first to resolve positions
query_ids, items_ids = self._batch_tokenize_query_and_items(query, items)
_, input_ids, positional_embed_overrides = self._build_token_id_inputs(
query_ids,
items_ids,
item_first,
use_multi_item_scoring,
embed_override_token_id,
query_embed_overrides,
item_embed_overrides,
_, input_ids, positional_embed_overrides, delimiter_indices = (
self._build_token_id_inputs(
query_ids,
items_ids,
item_first,
use_multi_item_scoring,
embed_override_token_id,
query_embed_overrides,
item_embed_overrides,
)
)
else:
raise ValueError(
@@ -679,6 +674,7 @@ class TokenizerManagerScoreMixin:
)
# Create the appropriate request type
mis_delimiter_indices = [delimiter_indices] if use_multi_item_scoring else None
if is_generation:
batch_request = GenerateReqInput(
text=text_prompts,
@@ -690,6 +686,7 @@ class TokenizerManagerScoreMixin:
stream=False,
sampling_params={"max_new_tokens": 0},
positional_embed_overrides=positional_embed_overrides,
multi_item_delimiter_indices=mis_delimiter_indices,
)
else:
batch_request = EmbeddingReqInput(
@@ -697,6 +694,7 @@ class TokenizerManagerScoreMixin:
input_ids=input_ids,
positional_embed_overrides=positional_embed_overrides,
return_pooled_hidden_states=return_pooled_hidden_states,
multi_item_delimiter_indices=mis_delimiter_indices,
)
results = await self.generate_request(batch_request, request).__anext__()
@@ -396,6 +396,9 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
# Whether this batch is prefill-only (no token generation needed)
is_prefill_only: bool = False
# Pre-computed delimiter indices for multi-item scoring (CPU tensors, one per request)
multi_item_delimiter_indices: Optional[List[torch.Tensor]] = None
# Speculative decoding
spec_info: Optional[SpecInput] = None
spec_algorithm: SpeculativeAlgorithm = None
@@ -468,6 +471,7 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
can_run_dp_cuda_graph=batch.can_run_dp_cuda_graph,
global_forward_mode=batch.global_forward_mode,
is_prefill_only=batch.is_prefill_only,
multi_item_delimiter_indices=batch.multi_item_delimiter_indices,
lora_ids=batch.lora_ids,
sampling_info=batch.sampling_info,
req_to_token_pool=model_runner.req_to_token_pool,
+46 -29
View File
@@ -159,6 +159,13 @@ DISAGG_TRANSFER_BACKEND_CHOICES = ["mooncake", "nixl", "ascend", "fake", "mori"]
GRAMMAR_BACKEND_CHOICES = ["xgrammar", "outlines", "llguidance", "none"]
# Placeholder token inserted between items in Multi-Item Scoring sequences:
# query<delim>item1<delim>item2<delim>... Positions are pre-computed from item
# lengths (multi_item_delimiter_indices); the token only exists for FlashInfer
# attention mask compat and logprob column indexing. Will be removed once the
# attention backend supports position-only MIS.
MIS_DELIMITER_TOKEN_ID = 9999
MOE_RUNNER_BACKEND_CHOICES = [
"auto",
"deep_gemm",
@@ -601,10 +608,11 @@ class ServerArgs:
offload_mode: str = "cpu"
# Scoring configuration
# Delimiter token ID used to combine Query and Items into a single sequence for multi-item scoring.
# Format: Query<delimiter>Item1<delimiter>Item2<delimiter>...
# This enables efficient batch processing of multiple items against a single query.
multi_item_scoring_delimiter: Optional[Union[int]] = None
# Enable Multi-Item Scoring optimization. Combines query and multiple items
# into a single sequence for efficient batch processing. Item boundaries are
# determined by pre-computed delimiter indices (from item lengths), not by the
# placeholder token. See MIS_DELIMITER_TOKEN_ID for details.
enable_mis: bool = False
# Optimization/debug options
disable_radix_cache: bool = False
@@ -800,9 +808,6 @@ class ServerArgs:
# Handle piecewise CUDA graph.
self._handle_piecewise_cuda_graph()
# Handle multi-item scoring constraints.
self._handle_multi_item_scoring()
# Get GPU memory capacity, which is a common dependency for several configuration steps.
gpu_mem = get_device_memory_capacity(self.device)
@@ -823,6 +828,10 @@ class ServerArgs:
self._handle_nccl_pre_warm()
self._handle_grammar_backend()
# Handle multi-item scoring constraints. Must run after the above so
# the final attention backend and chunked_prefill_size are in effect.
self._handle_multi_item_scoring()
# Handle Hicache settings.
self._handle_hicache()
@@ -1227,20 +1236,36 @@ class ServerArgs:
self.disable_piecewise_cuda_graph = True
def _handle_multi_item_scoring(self):
"""Disable CUDA graphs when multi-item scoring delimiter is set.
"""Setup and validate multi-item scoring constraints.
The padded static input_ids buffer used by CUDA graph replay causes
spurious delimiter matches in score_and_pool's MIS path.
Auto-disables settings incompatible with MIS mechanics (CUDA graph,
radix cache, chunked prefill). Asserts on attention backend since
changing it silently could surprise users who intentionally picked
a non-flashinfer backend.
"""
if self.multi_item_scoring_delimiter is None:
if not self.enable_mis:
return
if not self.disable_cuda_graph:
logger.warning(
"CUDA graph is disabled because --multi-item-scoring-delimiter is set."
)
logger.warning("CUDA graph is disabled because --enable-mis is set.")
self.disable_cuda_graph = True
self.disable_piecewise_cuda_graph = True
if not self.disable_radix_cache:
logger.warning("Radix cache is disabled because --enable-mis is set.")
self.disable_radix_cache = True
if self.chunked_prefill_size != -1:
logger.warning("Chunked prefill is disabled because --enable-mis is set.")
self.chunked_prefill_size = -1
prefill_backend, decode_backend = self.get_attention_backends()
assert prefill_backend == "flashinfer" and decode_backend == "flashinfer", (
"Multi-item scoring requires flashinfer attention backend for custom attention mask support. "
f"Please set --attention-backend flashinfer when using --enable-mis. "
f"Current backends: prefill={prefill_backend}, decode={decode_backend}"
)
def _handle_gpu_memory_settings(self, gpu_mem):
"""
Configure GPU memory-dependent settings including
@@ -5739,10 +5764,13 @@ class ServerArgs:
# Args for multi-item-scoring
parser.add_argument(
"--multi-item-scoring-delimiter",
type=int,
default=ServerArgs.multi_item_scoring_delimiter,
help="Delimiter token ID for multi-item scoring. Used to combine Query and Items into a single sequence: Query<delimiter>Item1<delimiter>Item2<delimiter>... This enables efficient batch processing of multiple items against a single query.",
"--enable-mis",
action="store_true",
default=ServerArgs.enable_mis,
help="Enable Multi-Item Scoring optimization. Combines query and multiple items "
"into a single sequence for efficient batch processing. "
"Requires --attention-backend flashinfer; auto-disables CUDA graph, "
"radix cache, and chunked prefill.",
)
# Optimization/debug options
@@ -6612,17 +6640,6 @@ class ServerArgs:
"--default-priority-value has no effect without --enable-priority-scheduling"
)
# Check multi-item scoring
if self.multi_item_scoring_delimiter is not None:
assert self.disable_radix_cache, (
"Multi-item scoring requires radix cache to be disabled. "
"Please set --disable-radix-cache when using --multi-item-scoring-delimiter."
)
assert self.chunked_prefill_size == -1, (
"Multi-item scoring requires chunked prefill to be disabled. "
"Please set --chunked-prefill-size -1 when using --multi-item-scoring-delimiter."
)
# Check hisparse
if self.enable_hisparse:
from sglang.srt.configs.model_config import is_deepseek_nsa