[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
|
||||
|
||||
Reference in New Issue
Block a user