diff --git a/python/sglang/srt/layers/attention/flashinfer_backend.py b/python/sglang/srt/layers/attention/flashinfer_backend.py index c8128058e..282aa443a 100644 --- a/python/sglang/srt/layers/attention/flashinfer_backend.py +++ b/python/sglang/srt/layers/attention/flashinfer_backend.py @@ -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) diff --git a/python/sglang/srt/layers/logits_processor.py b/python/sglang/srt/layers/logits_processor.py index 1980e5520..7151ef1f3 100644 --- a/python/sglang/srt/layers/logits_processor.py +++ b/python/sglang/srt/layers/logits_processor.py @@ -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: QueryItem1Item2... Scoring positions: Extracts logprobs at positions before each 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, diff --git a/python/sglang/srt/layers/pooler.py b/python/sglang/srt/layers/pooler.py index 582bf0a97..a31e60dfd 100644 --- a/python/sglang/srt/layers/pooler.py +++ b/python/sglang/srt/layers/pooler.py @@ -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) diff --git a/python/sglang/srt/managers/io_struct.py b/python/sglang/srt/managers/io_struct.py index a3f6e4222..c27f6cbad 100644 --- a/python/sglang/srt/managers/io_struct.py +++ b/python/sglang/srt/managers/io_struct.py @@ -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 diff --git a/python/sglang/srt/managers/schedule_batch.py b/python/sglang/srt/managers/schedule_batch.py index 31d01a695..771584e44 100644 --- a/python/sglang/srt/managers/schedule_batch.py +++ b/python/sglang/srt/managers/schedule_batch.py @@ -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 diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index 3be81753b..8f39cad01 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -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 diff --git a/python/sglang/srt/managers/scheduler_output_processor_mixin.py b/python/sglang/srt/managers/scheduler_output_processor_mixin.py index 27cd17025..7b131595e 100644 --- a/python/sglang/srt/managers/scheduler_output_processor_mixin.py +++ b/python/sglang/srt/managers/scheduler_output_processor_mixin.py @@ -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, diff --git a/python/sglang/srt/managers/tokenizer_manager.py b/python/sglang/srt/managers/tokenizer_manager.py index 8a49717c5..aa9567d07 100644 --- a/python/sglang/srt/managers/tokenizer_manager.py +++ b/python/sglang/srt/managers/tokenizer_manager.py @@ -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 diff --git a/python/sglang/srt/managers/tokenizer_manager_score_mixin.py b/python/sglang/srt/managers/tokenizer_manager_score_mixin.py index d520cf782..bb05a57c0 100644 --- a/python/sglang/srt/managers/tokenizer_manager_score_mixin.py +++ b/python/sglang/srt/managers/tokenizer_manager_score_mixin.py @@ -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: queryitem1item2item3 @@ -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: queryitem1item2item3 - 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: queryitem1item2item3 @@ -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__() diff --git a/python/sglang/srt/model_executor/forward_batch_info.py b/python/sglang/srt/model_executor/forward_batch_info.py index f8704040b..ff66b4099 100644 --- a/python/sglang/srt/model_executor/forward_batch_info.py +++ b/python/sglang/srt/model_executor/forward_batch_info.py @@ -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, diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index 0d4fd94a7..51e408dc7 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -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: +# queryitem1item2... 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: QueryItem1Item2... - # 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: QueryItem1Item2... 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 diff --git a/test/registered/prefill_only/test_embed_overrides.py b/test/registered/prefill_only/test_embed_overrides.py index 7075472e6..01349dbd1 100644 --- a/test/registered/prefill_only/test_embed_overrides.py +++ b/test/registered/prefill_only/test_embed_overrides.py @@ -20,6 +20,7 @@ from sglang.srt.managers.tokenizer_manager import TokenizerManager from sglang.srt.managers.tokenizer_manager_score_mixin import ( TokenizerManagerScoreMixin, ) +from sglang.srt.server_args import MIS_DELIMITER_TOKEN_ID from sglang.test.ci.ci_register import register_cuda_ci from sglang.test.test_utils import CustomTestCase @@ -204,16 +205,15 @@ class TestEmbeddingReqInputEmbedOverride(CustomTestCase): class _FakeServerArgs: """Minimal stub for server_args.""" - def __init__(self, multi_item_scoring_delimiter=None): - self.multi_item_scoring_delimiter = multi_item_scoring_delimiter + def __init__(self, enable_mis=False): + self.enable_mis = enable_mis class _FakeMixin(TokenizerManagerScoreMixin): """Minimal stub to call mixin methods without a full TokenizerManager.""" - def __init__(self, delimiter=None): - self.server_args = _FakeServerArgs(delimiter) - self.multi_item_delimiter_text = None + def __init__(self, enable_mis=False): + self.server_args = _FakeServerArgs(enable_mis) self.tokenizer = None self.is_generation = True @@ -334,17 +334,17 @@ class TestResolveEmbedOverridesForRequest(CustomTestCase): # Score mixin: _build_token_id_inputs # ======================================================================== -DELIM_TOKEN = 99 +DELIM_TOKEN = MIS_DELIMITER_TOKEN_ID class TestBuildTokenIdInputs(CustomTestCase): def setUp(self): - self.mixin = _FakeMixin(delimiter=DELIM_TOKEN) + self.mixin = _FakeMixin(enable_mis=True) # --- single-item mode, no embeds --- def test_single_item_no_embeds(self): - _, input_ids, injection = self.mixin._build_token_id_inputs( + _, input_ids, positional_embed_overrides, _ = self.mixin._build_token_id_inputs( query=[1, 2], items=[[3, 4], [5, 6]], item_first=False, @@ -354,10 +354,10 @@ class TestBuildTokenIdInputs(CustomTestCase): item_embed_overrides=None, ) self.assertEqual(input_ids, [[1, 2, 3, 4], [1, 2, 5, 6]]) - self.assertIsNone(injection) + self.assertIsNone(positional_embed_overrides) def test_single_item_no_embeds_item_first(self): - _, input_ids, injection = self.mixin._build_token_id_inputs( + _, input_ids, positional_embed_overrides, _ = self.mixin._build_token_id_inputs( query=[1, 2], items=[[3, 4]], item_first=True, @@ -367,12 +367,12 @@ class TestBuildTokenIdInputs(CustomTestCase): item_embed_overrides=None, ) self.assertEqual(input_ids, [[3, 4, 1, 2]]) - self.assertIsNone(injection) + self.assertIsNone(positional_embed_overrides) # --- multi-item mode, no embeds --- def test_multi_item_no_embeds(self): - _, input_ids, injection = self.mixin._build_token_id_inputs( + _, input_ids, positional_embed_overrides, _ = self.mixin._build_token_id_inputs( query=[1, 2], items=[[3, 4], [5, 6]], item_first=False, @@ -385,13 +385,13 @@ class TestBuildTokenIdInputs(CustomTestCase): self.assertEqual( input_ids, [[1, 2, DELIM_TOKEN, 3, 4, DELIM_TOKEN, 5, 6, DELIM_TOKEN]] ) - self.assertIsNone(injection) + self.assertIsNone(positional_embed_overrides) # --- single-item mode, with embeds --- def test_single_item_query_embeds(self): """Query placeholder overrides are resolved per item.""" - _, input_ids, injection = self.mixin._build_token_id_inputs( + _, input_ids, positional_embed_overrides, _ = self.mixin._build_token_id_inputs( query=[50, 10], items=[[20, 30], [40, 50]], item_first=False, @@ -401,15 +401,15 @@ class TestBuildTokenIdInputs(CustomTestCase): item_embed_overrides=None, ) self.assertEqual(input_ids, [[50, 10, 20, 30], [50, 10, 40, 50]]) - self.assertIsNotNone(injection) - self.assertEqual(len(injection), 2) + self.assertIsNotNone(positional_embed_overrides) + self.assertEqual(len(positional_embed_overrides), 2) # Each item gets its own PositionalEmbeds with query override at pos 0 - self.assertEqual(injection[0].positions, [0]) - self.assertEqual(injection[1].positions, [0]) + self.assertEqual(positional_embed_overrides[0].positions, [0]) + self.assertEqual(positional_embed_overrides[1].positions, [0]) def test_single_item_item_embeds(self): """Per-item overrides with correct position offsets.""" - _, input_ids, injection = self.mixin._build_token_id_inputs( + _, input_ids, positional_embed_overrides, _ = self.mixin._build_token_id_inputs( query=[10, 20], items=[[50, 30]], item_first=False, @@ -419,13 +419,13 @@ class TestBuildTokenIdInputs(CustomTestCase): item_embed_overrides=[[_vec(2)]], ) self.assertEqual(input_ids, [[10, 20, 50, 30]]) - self.assertIsNotNone(injection) + self.assertIsNotNone(positional_embed_overrides) # item placeholder at index 0 of item, offset by query length 2 - self.assertEqual(injection[0].positions, [2]) + self.assertEqual(positional_embed_overrides[0].positions, [2]) def test_single_item_no_override_positions_returns_none_injection(self): - """When no items have placeholders, injection should be None.""" - _, input_ids, injection = self.mixin._build_token_id_inputs( + """When no items have placeholders, positional_embed_overrides should be None.""" + _, input_ids, positional_embed_overrides, _ = self.mixin._build_token_id_inputs( query=[10, 20], items=[[30, 40]], item_first=False, @@ -434,11 +434,11 @@ class TestBuildTokenIdInputs(CustomTestCase): query_embed_overrides=None, item_embed_overrides=[None], ) - self.assertIsNone(injection) + self.assertIsNone(positional_embed_overrides) def test_single_item_query_and_item_embeds(self): """Single-item mode with both query and item overrides in one request.""" - _, input_ids, injection = self.mixin._build_token_id_inputs( + _, input_ids, positional_embed_overrides, _ = self.mixin._build_token_id_inputs( query=[50, 10], items=[[20, 50]], item_first=False, @@ -448,15 +448,15 @@ class TestBuildTokenIdInputs(CustomTestCase): item_embed_overrides=[[_vec(2)]], ) self.assertEqual(input_ids, [[50, 10, 20, 50]]) - self.assertIsNotNone(injection) - pe = injection[0] + self.assertIsNotNone(positional_embed_overrides) + pe = positional_embed_overrides[0] # query override at pos 0, item override at pos 3 (query_len=2 + idx=1) self.assertEqual(pe.positions, [0, 3]) self.assertEqual(pe.embeds.shape, (2, HIDDEN_DIM)) def test_single_item_empty_query(self): """Empty query with item-only overrides (valid from score_prompts).""" - _, input_ids, injection = self.mixin._build_token_id_inputs( + _, input_ids, positional_embed_overrides, _ = self.mixin._build_token_id_inputs( query=[], items=[[50, 10]], item_first=False, @@ -466,15 +466,15 @@ class TestBuildTokenIdInputs(CustomTestCase): item_embed_overrides=[[_vec(1)]], ) self.assertEqual(input_ids, [[50, 10]]) - self.assertIsNotNone(injection) + self.assertIsNotNone(positional_embed_overrides) # item placeholder at absolute pos 0 (offset=len([])=0) - self.assertEqual(injection[0].positions, [0]) + self.assertEqual(positional_embed_overrides[0].positions, [0]) # --- multi-item mode, with embeds --- def test_multi_item_with_query_and_item_embeds(self): """Multi-item mode resolves query overrides once and item overrides per item.""" - _, input_ids, injection = self.mixin._build_token_id_inputs( + _, input_ids, positional_embed_overrides, _ = self.mixin._build_token_id_inputs( query=[50, 10], items=[[20, 50], [30, 40]], item_first=False, @@ -483,13 +483,13 @@ class TestBuildTokenIdInputs(CustomTestCase): query_embed_overrides=[_vec(1)], item_embed_overrides=[[_vec(2)], None], ) - # queryitem1item2 = [50,10, 99, 20,50, 99, 30,40, 99] + # queryitem1item2 = [50,10, DELIM, 20,50, DELIM, 30,40, DELIM] self.assertEqual(len(input_ids), 1) - self.assertIsNotNone(injection) + self.assertIsNotNone(positional_embed_overrides) self.assertEqual( - len(injection), 1 + len(positional_embed_overrides), 1 ) # single PositionalEmbeds for combined sequence - pe = injection[0] + pe = positional_embed_overrides[0] # query override at pos 0, item[0] override at pos 4 (query_len=2 + delim=1 + idx=1) self.assertIn(0, pe.positions) self.assertIn(4, pe.positions) diff --git a/test/registered/prefill_only/test_multi_item_scoring.py b/test/registered/prefill_only/test_multi_item_scoring.py new file mode 100644 index 000000000..7ec80e98c --- /dev/null +++ b/test/registered/prefill_only/test_multi_item_scoring.py @@ -0,0 +1,599 @@ +"""Tests for the Multi-Item Scoring (MIS) optimization. + +MIS is a server-side optimization enabled via --enable-mis that batches +multiple items into a single forward pass using delimiter tokens (token ID 9999). +This is different from batch scoring (multiple items in one API call) which +processes items as separate requests. + +The key difference: +- Batch scoring: N items -> N separate forward passes +- MIS optimization: N items -> 1 forward pass with delimiter-separated items + +These tests ensure the MIS optimization produces correct results and catches +bugs in tensor shape handling (e.g., 2D tensors [num_delimiters, num_label_tokens]). +""" + +import asyncio +import os +import unittest + +import torch +from transformers import AutoConfig, AutoTokenizer + +from sglang.srt.entrypoints.engine import Engine +from sglang.test.ci.ci_register import register_cuda_ci +from sglang.test.test_utils import ( + DEFAULT_SMALL_MODEL_NAME_FOR_TEST, + CustomTestCase, +) + +register_cuda_ci(est_time=240, suite="stage-b-test-1-gpu-small") + +TEST_MODEL_NAME = os.environ.get("TEST_MODEL_NAME", DEFAULT_SMALL_MODEL_NAME_FOR_TEST) +TEST_CLASSIFICATION_BASE_MODEL = os.environ.get( + "TEST_CLASSIFICATION_BASE_MODEL", + "tomaarsen/Qwen3-Reranker-0.6B-seq-cls", +) +_CLS_NUM_LABELS = AutoConfig.from_pretrained(TEST_CLASSIFICATION_BASE_MODEL).num_labels + + +class TestMISServerArgsValidation(unittest.TestCase): + """Test ServerArgs defaults for MIS mode.""" + + def test_enable_mis_default(self): + """Test that enable_mis defaults to False.""" + from sglang.srt.server_args import ServerArgs + + self.assertEqual(ServerArgs.enable_mis, False) + + +class TestMultiItemScoringOptimization(CustomTestCase): + """Test the Multi-Item Scoring (MIS) optimization with generation models.""" + + @classmethod + def setUpClass(cls): + cls.engine = Engine( + model_path=TEST_MODEL_NAME, + disable_radix_cache=True, + chunked_prefill_size=-1, + enable_mis=True, + attention_backend="flashinfer", + mem_fraction_static=0.15, + ) + cls.non_mis_engine = Engine( + model_path=TEST_MODEL_NAME, + disable_radix_cache=True, + chunked_prefill_size=-1, + mem_fraction_static=0.15, + ) + + @classmethod + def tearDownClass(cls): + if cls.engine is not None: + cls.engine.shutdown() + if cls.non_mis_engine is not None: + cls.non_mis_engine.shutdown() + torch.cuda.empty_cache() + + def test_mis_basic(self): + """Test basic MIS: correct shapes, valid probabilities.""" + query = "Rate each option:" + items = ["Option A", "Option B", "Option C"] + label_token_ids = [9454, 2753] # "Yes" and "No" tokens + + scores = self.engine.score( + query=query, + items=items, + label_token_ids=label_token_ids, + apply_softmax=True, + ).scores + + self.assertEqual(len(scores), len(items)) + for i, score_list in enumerate(scores): + self.assertEqual(len(score_list), len(label_token_ids)) + self.assertAlmostEqual(sum(score_list), 1.0, places=5) + for score in score_list: + self.assertGreaterEqual(score, 0) + self.assertLessEqual(score, 1) + + def test_mis_consistency_with_single_item(self): + """MIS with one item should match non-MIS scoring closely.""" + query = "Is this a fact?\n" + items = [" The sun rises in the east"] + label_token_ids = [9454, 2753] + + mis_scores = self.engine.score( + query=query, + items=items, + label_token_ids=label_token_ids, + apply_softmax=True, + ).scores + + non_mis_scores = self.non_mis_engine.score( + query=query, + items=items, + label_token_ids=label_token_ids, + apply_softmax=True, + ).scores + + self.assertEqual(len(mis_scores), 1) + self.assertEqual(len(non_mis_scores), 1) + for j, (m, n) in enumerate(zip(mis_scores[0], non_mis_scores[0])): + relative_diff = abs(m - n) / max(abs(n), 1e-6) + self.assertLess( + relative_diff, + 0.08, + msg=f"label {j}: MIS={m} vs non-MIS={n} (diff: {relative_diff:.3f})", + ) + + def test_mis_empty_query(self): + """MIS with empty query — delimiter indices start at position 0.""" + items = ["alpha", "beta"] + label_token_ids = [9454, 2753] + + scores = self.engine.score( + query="", + items=items, + label_token_ids=label_token_ids, + apply_softmax=True, + ).scores + + self.assertEqual(len(scores), len(items)) + for score_list in scores: + self.assertEqual(len(score_list), len(label_token_ids)) + self.assertAlmostEqual(sum(score_list), 1.0, places=5) + + +class TestMultiItemScoringClassification(CustomTestCase): + """Test MIS with classification models. + + Uses a pre-trained Qwen3ForSequenceClassification model so that the + classification head weights are deterministic across Engine instances. + """ + + NUM_LABELS = _CLS_NUM_LABELS + + def setUp(self): + self.engine = Engine( + model_path=TEST_CLASSIFICATION_BASE_MODEL, + disable_radix_cache=True, + chunked_prefill_size=-1, + enable_mis=True, + attention_backend="flashinfer", + mem_fraction_static=0.15, + ) + + def tearDown(self): + if self.engine is not None: + self.engine.shutdown() + torch.cuda.empty_cache() + + def test_classification_mis_basic(self): + """Classification MIS: correct shapes, valid softmax probabilities.""" + query = "Rate each option:" + items = ["Option A", "Option B", "Option C"] + + scores = self.engine.score(query=query, items=items, apply_softmax=True).scores + + self.assertEqual(len(scores), len(items)) + for i, score_list in enumerate(scores): + self.assertEqual(len(score_list), self.NUM_LABELS) + self.assertAlmostEqual(sum(score_list), 1.0, places=5) + for score in score_list: + self.assertGreaterEqual(score, 0) + self.assertLessEqual(score, 1) + + def test_classification_mis_tokenized_input(self): + """Classification MIS with pre-tokenized query and items.""" + tokenizer = AutoTokenizer.from_pretrained(TEST_CLASSIFICATION_BASE_MODEL) + query_ids = tokenizer.encode("Rate each option:", add_special_tokens=False) + items_ids = [ + tokenizer.encode(item, add_special_tokens=False) + for item in ["Option A", "Option B"] + ] + + scores = self.engine.score( + query=query_ids, items=items_ids, apply_softmax=True + ).scores + + self.assertEqual(len(scores), len(items_ids)) + for score_list in scores: + self.assertEqual(len(score_list), self.NUM_LABELS) + self.assertAlmostEqual(sum(score_list), 1.0, places=5) + + def test_classification_non_mis_fallback(self): + """Classification model works correctly without --enable-mis.""" + non_mis_engine = Engine( + model_path=TEST_CLASSIFICATION_BASE_MODEL, + disable_radix_cache=True, + mem_fraction_static=0.15, + ) + try: + scores = non_mis_engine.score( + query="Test:", items=["A", "B"], apply_softmax=True + ).scores + + self.assertEqual(len(scores), 2) + for score_list in scores: + self.assertEqual(len(score_list), self.NUM_LABELS) + self.assertAlmostEqual(sum(score_list), 1.0, places=5) + finally: + non_mis_engine.shutdown() + torch.cuda.empty_cache() + + +class TestMultiItemScoringParity(CustomTestCase): + """Test that MIS produces the same results as single-item scoring.""" + + @classmethod + def setUpClass(cls): + cls.engine_single = Engine( + model_path=TEST_MODEL_NAME, + disable_radix_cache=True, + log_level="error", + mem_fraction_static=0.15, + ) + cls.engine_mis = Engine( + model_path=TEST_MODEL_NAME, + disable_radix_cache=True, + chunked_prefill_size=-1, + log_level="error", + enable_mis=True, + attention_backend="flashinfer", + mem_fraction_static=0.15, + ) + + @classmethod + def tearDownClass(cls): + if cls.engine_single is not None: + cls.engine_single.shutdown() + if cls.engine_mis is not None: + cls.engine_mis.shutdown() + torch.cuda.empty_cache() + + def _compare_scores( + self, query, items, label_token_ids=None, apply_softmax=True, test_name="" + ): + """Compare MIS vs single-item scoring results.""" + single_scores = self.engine_single.score( + query=query, + items=items, + label_token_ids=label_token_ids, + apply_softmax=apply_softmax, + ).scores + + mis_scores = self.engine_mis.score( + query=query, + items=items, + label_token_ids=label_token_ids, + apply_softmax=apply_softmax, + ).scores + + self.assertEqual( + len(mis_scores), len(single_scores), f"{test_name}: count mismatch" + ) + for i, (ms, ss) in enumerate(zip(mis_scores, single_scores)): + self.assertEqual(len(ms), len(ss), f"{test_name}: item {i} length mismatch") + for j, (m, s) in enumerate(zip(ms, ss)): + self.assertAlmostEqual( + m, + s, + places=1, + msg=f"{test_name}: item {i} label {j}: MIS={m} vs single={s}", + ) + + def test_parity_basic(self): + tokenizer = AutoTokenizer.from_pretrained(TEST_MODEL_NAME) + query = "Rate this option:" + items = [" Option A", " Option B", " Option C"] + labels = [" good", " bad"] + label_ids = [tokenizer.encode(lb, add_special_tokens=False)[0] for lb in labels] + self._compare_scores(query, items, label_ids, test_name="basic") + + def test_parity_tokenized_inputs(self): + tokenizer = AutoTokenizer.from_pretrained(TEST_MODEL_NAME) + query = "Rate this option:" + items = [" Option X", " Option Y"] + labels = [" good", " bad"] + query_ids = tokenizer.encode(query, add_special_tokens=False) + items_ids = [tokenizer.encode(i, add_special_tokens=False) for i in items] + label_ids = [tokenizer.encode(lb, add_special_tokens=False)[0] for lb in labels] + self._compare_scores(query_ids, items_ids, label_ids, test_name="tokenized") + + def test_parity_without_softmax(self): + tokenizer = AutoTokenizer.from_pretrained(TEST_MODEL_NAME) + query = "The weather today is" + items = [" sunny", " cloudy", " rainy"] + labels = [" nice", " bad"] + label_ids = [tokenizer.encode(lb, add_special_tokens=False)[0] for lb in labels] + self._compare_scores( + query, items, label_ids, apply_softmax=False, test_name="no_softmax" + ) + + def test_parity_many_items(self): + tokenizer = AutoTokenizer.from_pretrained(TEST_MODEL_NAME) + query = "Rate this option from 1 to 5:" + items = [f" Option {i}" for i in range(10)] + labels = [" 1", " 2", " 3", " 4", " 5"] + label_ids = [tokenizer.encode(lb, add_special_tokens=False)[0] for lb in labels] + self._compare_scores(query, items, label_ids, test_name="many_items") + + +class TestMultiItemScoringClassificationParity(CustomTestCase): + """Test that MIS multi-item batching matches single-item MIS scoring. + + Both paths use the MIS engine (with delimiter tokens in the attention + context). The reference scores each item individually so each gets its + own forward pass; the batched path packs all items into one pass. + This isolates the MIS batching logic from the delimiter-presence effect. + """ + + NUM_LABELS = _CLS_NUM_LABELS + + @classmethod + def setUpClass(cls): + cls.engine = Engine( + model_path=TEST_CLASSIFICATION_BASE_MODEL, + disable_radix_cache=True, + chunked_prefill_size=-1, + enable_mis=True, + attention_backend="flashinfer", + mem_fraction_static=0.15, + ) + + @classmethod + def tearDownClass(cls): + if cls.engine is not None: + cls.engine.shutdown() + torch.cuda.empty_cache() + + def _compare_scores(self, query, items, apply_softmax=True, test_name=""): + """Compare MIS batched vs MIS single-item scoring results.""" + single_scores = [] + for item in items: + result = self.engine.score( + query=query, + items=[item], + apply_softmax=apply_softmax, + ).scores + single_scores.append(result[0]) + + batched_scores = self.engine.score( + query=query, + items=items, + apply_softmax=apply_softmax, + ).scores + + self.assertEqual( + len(batched_scores), len(single_scores), f"{test_name}: count mismatch" + ) + for i, (bs, ss) in enumerate(zip(batched_scores, single_scores)): + self.assertEqual(len(bs), len(ss), f"{test_name}: item {i} length mismatch") + for j, (b, s) in enumerate(zip(bs, ss)): + self.assertAlmostEqual( + b, + s, + places=1, + msg=f"{test_name}: item {i} label {j}: batched={b} vs single={s}", + ) + + def test_parity_basic(self): + query = "Rate this option:" + items = [" Option A", " Option B", " Option C"] + self._compare_scores(query, items, test_name="cls_basic") + + def test_parity_tokenized_inputs(self): + tokenizer = AutoTokenizer.from_pretrained(TEST_CLASSIFICATION_BASE_MODEL) + query_ids = tokenizer.encode("Rate this option:", add_special_tokens=False) + items_ids = [ + tokenizer.encode(item, add_special_tokens=False) + for item in [" Option X", " Option Y"] + ] + self._compare_scores(query_ids, items_ids, test_name="cls_tokenized") + + def test_parity_without_softmax(self): + query = "The weather today is" + items = [" sunny", " cloudy", " rainy"] + self._compare_scores( + query, items, apply_softmax=False, test_name="cls_no_softmax" + ) + + def test_parity_many_items(self): + query = "Classify this option:" + items = [f" Option {i}" for i in range(10)] + self._compare_scores(query, items, test_name="cls_many_items") + + +class TestMultiItemScoringClassificationMISvsNonMIS(CustomTestCase): + """Test that MIS single-item approximates non-MIS single-item. + + The MIS path inserts delimiter tokens into the attention context, + which slightly perturbs hidden states. After softmax the scores + should still be close. Uses places=1 (±0.05) tolerance. + + Runs as a separate class so each engine is created and destroyed + independently to avoid GPU OOM. + """ + + def test_mis_single_vs_non_mis(self): + non_mis_engine = Engine( + model_path=TEST_CLASSIFICATION_BASE_MODEL, + disable_radix_cache=True, + mem_fraction_static=0.15, + ) + try: + query = "Rate this option:" + items = [" Option A", " Option B", " Option C"] + non_mis_scores = non_mis_engine.score( + query=query, + items=items, + apply_softmax=True, + ).scores + finally: + non_mis_engine.shutdown() + torch.cuda.empty_cache() + + mis_engine = Engine( + model_path=TEST_CLASSIFICATION_BASE_MODEL, + disable_radix_cache=True, + chunked_prefill_size=-1, + enable_mis=True, + attention_backend="flashinfer", + mem_fraction_static=0.15, + ) + try: + mis_scores = mis_engine.score( + query=query, + items=items, + apply_softmax=True, + ).scores + finally: + mis_engine.shutdown() + torch.cuda.empty_cache() + + self.assertEqual(len(mis_scores), len(non_mis_scores)) + for i, (ms, ns) in enumerate(zip(mis_scores, non_mis_scores)): + self.assertEqual(len(ms), len(ns)) + for j, (m, n) in enumerate(zip(ms, ns)): + self.assertAlmostEqual( + m, + n, + places=1, + msg=f"item {i} label {j}: MIS={m} vs non-MIS={n}", + ) + + +class TestMultiItemScoringClassificationAdvanced(CustomTestCase): + """Advanced MIS tests for classification models: score distinctness, + determinism, and concurrent request handling.""" + + NUM_LABELS = _CLS_NUM_LABELS + + @classmethod + def setUpClass(cls): + cls.engine = Engine( + model_path=TEST_CLASSIFICATION_BASE_MODEL, + disable_radix_cache=True, + chunked_prefill_size=-1, + enable_mis=True, + attention_backend="flashinfer", + mem_fraction_static=0.15, + ) + + @classmethod + def tearDownClass(cls): + if cls.engine is not None: + cls.engine.shutdown() + torch.cuda.empty_cache() + + def test_items_produce_distinct_scores(self): + """Different items must produce different score vectors. + + Core regression test: before the delimiter-index fix, all items got + identical scores because the MIS attention mask only let delimiter + tokens attend to the query prefix. + """ + query = "Rate each option:" + items = [ + "Option A is about cats", + "Option B is about dogs", + "Option C is about fish", + ] + + scores = self.engine.score(query=query, items=items).scores + + self.assertEqual(len(scores), len(items)) + all_identical = all(scores[0] == s for s in scores[1:]) + self.assertFalse( + all_identical, + f"All {len(items)} items returned identical scores — " + f"MIS delimiter indexing is broken. Scores: {scores[0]}", + ) + + def test_many_items_distinct(self): + """Stress test: 15 items should not all produce identical scores.""" + query = "Classify each city:" + items = [f"City {i}" for i in range(15)] + + scores = self.engine.score(query=query, items=items).scores + + self.assertEqual(len(scores), len(items)) + for score_list in scores: + self.assertEqual(len(score_list), self.NUM_LABELS) + + unique_count = len({tuple(s) for s in scores}) + self.assertGreater(unique_count, 1, "All 15 items returned identical scores") + + def test_deterministic(self): + """Identical requests should return identical scores.""" + query = "Evaluate:" + items = ["alpha", "beta", "gamma"] + + scores1 = self.engine.score(query=query, items=items).scores + scores2 = self.engine.score(query=query, items=items).scores + + self.assertEqual( + scores1, scores2, "Identical inputs must produce identical scores" + ) + + def test_concurrent_requests(self): + """Concurrent MIS requests must produce the same scores as sequential. + + Runs each request sequentially to get baseline scores, then runs all + concurrently and asserts the results match. This catches cross-request + contamination when multiple MIS requests share a GPU batch. + """ + test_cases = [ + {"query": "Is this a fruit?", "items": ["apple", "car", "banana"]}, + {"query": "Is this an animal?", "items": ["dog", "table"]}, + { + "query": "Is this a country?", + "items": ["France", "pizza", "Japan", "chair"], + }, + {"query": "Is this a color?", "items": ["red"]}, + ] + + # Sequential baseline + sequential_scores = [] + for tc in test_cases: + result = self.engine.score(query=tc["query"], items=tc["items"]) + sequential_scores.append(result.scores) + + # Concurrent execution + async def _gather(): + return await asyncio.gather( + *( + self.engine.async_score(query=tc["query"], items=tc["items"]) + for tc in test_cases + ) + ) + + concurrent_results = self.engine.loop.run_until_complete(_gather()) + + for idx, (tc, seq_scores, conc_result) in enumerate( + zip(test_cases, sequential_scores, concurrent_results) + ): + conc_scores = conc_result.scores + self.assertEqual( + len(conc_scores), + len(seq_scores), + f"Case {idx}: count mismatch", + ) + for i, (cs, ss) in enumerate(zip(conc_scores, seq_scores)): + self.assertEqual( + len(cs), + len(ss), + f"Case {idx} item {i}: label count mismatch", + ) + for j, (c, s) in enumerate(zip(cs, ss)): + self.assertAlmostEqual( + c, + s, + places=1, + msg=f"Case {idx} item {i} label {j}: " + f"concurrent={c} vs sequential={s}", + ) + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/prefill_only/test_pooled_hidden_states.py b/test/registered/prefill_only/test_pooled_hidden_states.py index 66582a773..dd1757302 100644 --- a/test/registered/prefill_only/test_pooled_hidden_states.py +++ b/test/registered/prefill_only/test_pooled_hidden_states.py @@ -30,7 +30,6 @@ from sglang.test.test_utils import ( register_cuda_ci(est_time=240, suite="stage-b-test-1-gpu-small") _SEQCLS_MODEL = "Qwen/Qwen3-0.6B" -_QWEN3_EOT_TOKEN_ID = 151643 _CAUSAL_LM_MODEL = DEFAULT_SMALL_MODEL_NAME_FOR_TEST _NUM_LABELS = 4 @@ -197,7 +196,7 @@ class TestPooledHiddenStatesMISEngine(CustomTestCase): model_path=_SEQCLS_MODEL, disable_radix_cache=True, chunked_prefill_size=-1, - multi_item_scoring_delimiter=_QWEN3_EOT_TOKEN_ID, + enable_mis=True, json_model_override_args=json.dumps( { "architectures": ["Qwen3ForSequenceClassification"], diff --git a/test/registered/prefill_only/test_score_api.py b/test/registered/prefill_only/test_score_api.py index c096290d6..4def12765 100644 --- a/test/registered/prefill_only/test_score_api.py +++ b/test/registered/prefill_only/test_score_api.py @@ -3,15 +3,16 @@ Two test classes, each with its own server instance: TestCausalLMScoringHTTP — basic endpoint: schema defaults, response - structure, error rejection (no MIS delimiter) - TestCausalLMMISScoringHTTP — MIS mode: validates --multi-item-scoring-delimiter - CLI flag wiring and per-item output shape + structure, error rejection (no MIS) + TestCausalLMMISScoringHTTP — MIS mode: validates --enable-mis CLI flag + wiring and per-item output shape Engine-level correctness (numerical accuracy, batching, edge cases) lives in test_score_engine.py. These tests focus on the HTTP integration seam: Pydantic schema defaults, FastAPI routing, and server argument wiring. """ +import os import unittest import requests @@ -28,9 +29,7 @@ from sglang.test.test_utils import ( register_cuda_ci(est_time=70, suite="stage-b-test-1-gpu-small") -_MODEL = DEFAULT_SMALL_MODEL_NAME_FOR_TEST # Llama-3.2-1B-Instruct -# <|eot_id|> for Llama-3.x Instruct — used as MIS delimiter -_LLAMA3_EOT_TOKEN_ID = 128009 +_MODEL = os.environ.get("TEST_MODEL_NAME", DEFAULT_SMALL_MODEL_NAME_FOR_TEST) # --------------------------------------------------------------------------- @@ -41,7 +40,7 @@ _LLAMA3_EOT_TOKEN_ID = 128009 class TestCausalLMScoringHTTP(CustomTestCase): """Validates /v1/score HTTP integration — schema, defaults, and error handling. - Starts a plain CausalLM server (no --multi-item-scoring-delimiter) to test + Starts a plain CausalLM server (no --enable-mis) to test the HTTP layer in isolation: response envelope shape, the apply_softmax default (False), and Pydantic validation errors on malformed input. """ @@ -139,12 +138,12 @@ class TestCausalLMScoringHTTP(CustomTestCase): # --------------------------------------------------------------------------- -# MIS scoring (with --multi-item-scoring-delimiter) +# MIS scoring (with --enable-mis) # --------------------------------------------------------------------------- class TestCausalLMMISScoringHTTP(CustomTestCase): - """Validates /v1/score with --multi-item-scoring-delimiter. + """Validates /v1/score with --enable-mis. Confirms that the CLI flag is correctly wired into ServerArgs and that the endpoint returns one probability vector per item when items are @@ -163,8 +162,9 @@ class TestCausalLMMISScoringHTTP(CustomTestCase): "--disable-radix-cache", "--chunked-prefill-size", "-1", - "--multi-item-scoring-delimiter", - str(_LLAMA3_EOT_TOKEN_ID), + "--enable-mis", + "--attention-backend", + "flashinfer", ], ) diff --git a/test/registered/prefill_only/test_score_engine.py b/test/registered/prefill_only/test_score_engine.py index 767e8a91e..19d580f6d 100644 --- a/test/registered/prefill_only/test_score_engine.py +++ b/test/registered/prefill_only/test_score_engine.py @@ -4,14 +4,18 @@ Two model types, two scoring modes: TestCausalLMScoring — CausalLM, single-item and batched multi-item TestSeqClsScoring — SequenceClassification, single-item mode - TestSeqClsMISScoring — SequenceClassification, MIS delimiter mode + TestSeqClsMISScoring — SequenceClassification, MIS mode (--enable-mis) + TestSeqClsMISAdvancedScoring — SeqCls MIS with 12 labels (tensor shape stress) The Engine (Python API) is the right layer for correctness testing: it exercises tokenization, forward pass, pooling, and score extraction without the HTTP serialization overhead. HTTP-layer tests live in test_score_api.py. +Thorough MIS tests (parity, concurrency, generation models) live in +test_multi_item_scoring.py. """ import json +import os import unittest from unittest.mock import patch @@ -24,10 +28,8 @@ from sglang.test.test_utils import DEFAULT_SMALL_MODEL_NAME_FOR_TEST, CustomTest register_cuda_ci(est_time=85, suite="stage-b-test-1-gpu-small") -_CAUSAL_LM_MODEL = DEFAULT_SMALL_MODEL_NAME_FOR_TEST # Llama-3.2-1B-Instruct -_SEQCLS_MODEL = "Qwen/Qwen3-0.6B" # backbone; arch overridden to SeqCls below -# <|endoftext|> for Qwen3 tokenizer — used as MIS delimiter -_QWEN3_EOT_TOKEN_ID = 151643 +_CAUSAL_LM_MODEL = os.environ.get("TEST_MODEL_NAME", DEFAULT_SMALL_MODEL_NAME_FOR_TEST) +_SEQCLS_MODEL = os.environ.get("TEST_CLASSIFICATION_BASE_MODEL", "Qwen/Qwen3-0.6B") # --------------------------------------------------------------------------- @@ -368,10 +370,9 @@ class TestSeqClsScoring(CustomTestCase): class TestSeqClsMISScoring(CustomTestCase): """SeqCls MIS: all items packed into one sequence separated by delimiter token. - score_and_pool() extracts per-item scores at delimiter positions. - Two sub-cases are tested: - - NUM_LABELS=2 — standard binary classification head - - NUM_LABELS=12 — stress-tests 2-D tensor indexing in score_and_pool() + Uses --enable-mis which hardcodes delimiter token ID 9999. + Basic pipeline correctness only — thorough MIS tests (parity, + concurrency, advanced) live in test_multi_item_scoring.py. """ NUM_LABELS = 2 @@ -382,7 +383,8 @@ class TestSeqClsMISScoring(CustomTestCase): model_path=_SEQCLS_MODEL, disable_radix_cache=True, chunked_prefill_size=-1, - multi_item_scoring_delimiter=_QWEN3_EOT_TOKEN_ID, + enable_mis=True, + attention_backend="flashinfer", json_model_override_args=json.dumps( { "architectures": ["Qwen3ForSequenceClassification"], @@ -432,32 +434,6 @@ class TestSeqClsMISScoring(CustomTestCase): self.assertEqual(len(row), self.NUM_LABELS) self.assertAlmostEqual(sum(row), 1.0, places=5) - def test_mis_items_produce_distinct_scores(self): - """Different items must yield different score vectors. - - Catches bugs where all delimiter positions share the same pooled - hidden state (e.g. off-by-one in score_and_pool indexing). - """ - items = [ - "Option A is about cats", - "Option B is about dogs", - "Option C is about fish", - ] - scores = self.engine.score(query="Rate each option:", items=items).scores - self.assertEqual(len(scores), len(items)) - self.assertFalse( - all(scores[0] == s for s in scores[1:]), - f"All items returned identical scores — delimiter indexing is likely broken. " - f"Scores: {scores[0]}", - ) - - def test_mis_deterministic(self): - """Identical MIS requests return identical scores.""" - kwargs = dict(query="Evaluate:", items=["alpha", "beta", "gamma"]) - self.assertEqual( - self.engine.score(**kwargs).scores, self.engine.score(**kwargs).scores - ) - # --------------------------------------------------------------------------- # SequenceClassification — MIS with many labels (tensor shape stress test) @@ -479,7 +455,8 @@ class TestSeqClsMISAdvancedScoring(CustomTestCase): model_path=_SEQCLS_MODEL, disable_radix_cache=True, chunked_prefill_size=-1, - multi_item_scoring_delimiter=_QWEN3_EOT_TOKEN_ID, + enable_mis=True, + attention_backend="flashinfer", json_model_override_args=json.dumps( { "architectures": ["Qwen3ForSequenceClassification"], @@ -506,17 +483,6 @@ class TestSeqClsMISAdvancedScoring(CustomTestCase): self.assertEqual(len(row), self.NUM_LABELS) self.assertAlmostEqual(sum(row), 1.0, places=5) - def test_many_items_produce_distinct_scores(self): - """15 items should not all return identical score vectors.""" - items = [f"City {i}" for i in range(15)] - scores = self.engine.score(query="Classify each city:", items=items).scores - self.assertEqual(len(scores), len(items)) - self.assertGreater( - len({tuple(s) for s in scores}), - 1, - "All 15 items returned identical scores", - ) - if __name__ == "__main__": unittest.main(verbosity=3) diff --git a/test/registered/unit/layers/test_pooler_score_and_pool.py b/test/registered/unit/layers/test_pooler_score_and_pool.py index 2f0c4a1b9..829b228ba 100644 --- a/test/registered/unit/layers/test_pooler_score_and_pool.py +++ b/test/registered/unit/layers/test_pooler_score_and_pool.py @@ -1,12 +1,11 @@ """Unit tests for score_and_pool in sglang.srt.layers.pooler. -All tests run on CPU — no GPU required. The global server_args singleton -is mocked so the tests are hermetic. +All tests run on CPU — no GPU required. MIS delimiter positions are passed +via forward_batch.multi_item_delimiter_indices (pre-computed by the caller). """ import unittest from types import SimpleNamespace -from unittest.mock import patch import torch import torch.nn as nn @@ -24,22 +23,22 @@ register_cpu_ci(est_time=9, suite="stage-a-test-cpu") def _make_forward_batch( - extend_seq_lens, is_prefill_only=False, return_pooled_hidden_states=False + extend_seq_lens, + multi_item_delimiter_indices=None, + return_pooled_hidden_states=False, + is_prefill_only=True, ): """Build a minimal ForwardBatch stub for pooler unit tests.""" return SimpleNamespace( extend_seq_lens=torch.tensor(extend_seq_lens, dtype=torch.long), extend_seq_lens_cpu=extend_seq_lens, - is_prefill_only=is_prefill_only, + multi_item_delimiter_indices=multi_item_delimiter_indices, dimensions=None, return_pooled_hidden_states=return_pooled_hidden_states, + is_prefill_only=is_prefill_only, ) -def _mock_server_args(delimiter=None): - return SimpleNamespace(multi_item_scoring_delimiter=delimiter) - - class TestScoreAndPool(CustomTestCase): """Unit tests for the score_and_pool helper function.""" @@ -50,11 +49,8 @@ class TestScoreAndPool(CustomTestCase): self.score_head = nn.Linear(self.hidden_dim, self.num_labels, bias=False) self.pooler = Pooler(pooling_type=PoolingType.LAST, normalize=False) - @patch("sglang.srt.layers.pooler.get_global_server_args") - def test_single_item_returns_scores(self, mock_get_args): - """No delimiter -> single-item path returns [batch, num_labels].""" - mock_get_args.return_value = _mock_server_args(delimiter=None) - + def test_single_item_returns_scores(self): + """No delimiter indices -> single-item path returns [batch, num_labels].""" hidden = torch.randn(8, self.hidden_dim) fb = _make_forward_batch(extend_seq_lens=[5, 3]) input_ids = torch.arange(8) @@ -64,30 +60,16 @@ class TestScoreAndPool(CustomTestCase): self.assertIsInstance(out, EmbeddingPoolerOutput) self.assertEqual(out.embeddings.shape, (2, self.num_labels)) - @patch("sglang.srt.layers.pooler.get_global_server_args") - def test_mis_returns_per_request_list(self, mock_get_args): - """Delimiter found -> returns a list with one tensor per request.""" - delimiter_token = 99 - mock_get_args.return_value = _mock_server_args(delimiter=delimiter_token) - - input_ids = torch.tensor( - [ - 0, - 1, - 2, - delimiter_token, - 3, - 4, - 5, - delimiter_token, - 6, - 7, - 8, - delimiter_token, - ] - ) + def test_mis_returns_per_request_list(self): + """Delimiter indices provided -> returns a list with one tensor per request.""" + # Sequence: [0, 1, 2, D, 3, 4, 5, D, 6, 7, 8, D] + # Delimiters at positions 3, 7, 11 -> extract at 2, 6, 10 + input_ids = torch.arange(12) hidden = torch.randn(len(input_ids), self.hidden_dim) - fb = _make_forward_batch(extend_seq_lens=[len(input_ids)], is_prefill_only=True) + fb = _make_forward_batch( + extend_seq_lens=[len(input_ids)], + multi_item_delimiter_indices=[torch.tensor([3, 7, 11])], + ) out = score_and_pool(self.score_head, self.pooler, hidden, fb, input_ids) @@ -95,20 +77,20 @@ class TestScoreAndPool(CustomTestCase): self.assertEqual(len(out.embeddings), 1) self.assertEqual(out.embeddings[0].shape, (3, self.num_labels)) - @patch("sglang.srt.layers.pooler.get_global_server_args") - def test_mis_batched_splits_per_request(self, mock_get_args): + def test_mis_batched_splits_per_request(self): """Two batched MIS requests -> returns a list of length 2.""" - delimiter_token = 99 - mock_get_args.return_value = _mock_server_args(delimiter=delimiter_token) - - # Request 1: [10, 11, delim, 12, 13, delim] -> 2 delimiters - # Request 2: [20, 21, 22, delim] -> 1 delimiter - req1 = [10, 11, delimiter_token, 12, 13, delimiter_token] - req2 = [20, 21, 22, delimiter_token] + # Request 1: [10, 11, D, 12, 13, D] -> delimiters at 2, 5 + # Request 2: [20, 21, 22, D] -> delimiter at 3 + req1 = [10, 11, 99, 12, 13, 99] + req2 = [20, 21, 22, 99] input_ids = torch.tensor(req1 + req2) hidden = torch.randn(len(input_ids), self.hidden_dim) fb = _make_forward_batch( - extend_seq_lens=[len(req1), len(req2)], is_prefill_only=True + extend_seq_lens=[len(req1), len(req2)], + multi_item_delimiter_indices=[ + torch.tensor([2, 5]), + torch.tensor([3]), + ], ) out = score_and_pool(self.score_head, self.pooler, hidden, fb, input_ids) @@ -118,42 +100,21 @@ class TestScoreAndPool(CustomTestCase): self.assertEqual(out.embeddings[0].shape, (2, self.num_labels)) self.assertEqual(out.embeddings[1].shape, (1, self.num_labels)) - @patch("sglang.srt.layers.pooler.get_global_server_args") - def test_mis_falls_back_when_no_delimiters_in_input(self, mock_get_args): - """Delimiter configured but absent from input_ids -> single-item fallback.""" - mock_get_args.return_value = _mock_server_args(delimiter=99) - + def test_no_delimiter_indices_falls_back(self): + """multi_item_delimiter_indices=None -> single-item fallback.""" input_ids = torch.tensor([0, 1, 2, 3, 4, 5, 6, 7]) hidden = torch.randn(8, self.hidden_dim) - fb = _make_forward_batch(extend_seq_lens=[5, 3], is_prefill_only=True) + fb = _make_forward_batch(extend_seq_lens=[5, 3]) out = score_and_pool(self.score_head, self.pooler, hidden, fb, input_ids) self.assertIsInstance(out.embeddings, torch.Tensor) self.assertEqual(out.embeddings.shape, (2, self.num_labels)) - @patch("sglang.srt.layers.pooler.get_global_server_args") - def test_mis_falls_back_when_not_prefill_only(self, mock_get_args): - """Delimiter configured, is_prefill_only=False -> single-item fallback.""" - mock_get_args.return_value = _mock_server_args(delimiter=99) - - input_ids = torch.tensor([0, 1, 2, 99, 3, 4, 5, 99]) - hidden = torch.randn(8, self.hidden_dim) - fb = _make_forward_batch(extend_seq_lens=[5, 3], is_prefill_only=False) - - out = score_and_pool(self.score_head, self.pooler, hidden, fb, input_ids) - - self.assertIsInstance(out.embeddings, torch.Tensor) - self.assertEqual(out.embeddings.shape, (2, self.num_labels)) - - @patch("sglang.srt.layers.pooler.get_global_server_args") - def test_mis_extracts_positions_before_delimiter(self, mock_get_args): + def test_mis_extracts_positions_before_delimiter(self): """Verify MIS picks hidden states at index (delimiter_position - 1).""" - delimiter_token = 99 - mock_get_args.return_value = _mock_server_args(delimiter=delimiter_token) - # Delimiters at indices 2 and 5 -> extract hidden at indices 1 and 4 - input_ids = torch.tensor([10, 11, delimiter_token, 20, 21, delimiter_token]) + input_ids = torch.tensor([10, 11, 99, 20, 21, 99]) hidden = ( torch.arange(len(input_ids)) .unsqueeze(1) @@ -161,7 +122,10 @@ class TestScoreAndPool(CustomTestCase): .expand(-1, self.hidden_dim) .clone() ) - fb = _make_forward_batch(extend_seq_lens=[len(input_ids)], is_prefill_only=True) + fb = _make_forward_batch( + extend_seq_lens=[len(input_ids)], + multi_item_delimiter_indices=[torch.tensor([2, 5])], + ) identity_head = nn.Linear(self.hidden_dim, self.hidden_dim, bias=False) nn.init.eye_(identity_head.weight) @@ -172,14 +136,9 @@ class TestScoreAndPool(CustomTestCase): torch.testing.assert_close(scores[0], hidden[1]) torch.testing.assert_close(scores[1], hidden[4]) - @patch("sglang.srt.layers.pooler.get_global_server_args") - def test_mis_ignores_delimiter_at_position_zero(self, mock_get_args): - """A delimiter at flat index 0 has no preceding token and must be skipped.""" - delimiter_token = 99 - mock_get_args.return_value = _mock_server_args(delimiter=delimiter_token) - - # Delimiter at index 0 should be ignored; only the one at index 3 counts - input_ids = torch.tensor([delimiter_token, 10, 11, delimiter_token]) + def test_mis_delimiter_at_position_one(self): + """Delimiters at positions 1 and 3 extract at indices 0 and 2.""" + input_ids = torch.tensor([10, 99, 11, 99]) hidden = ( torch.arange(len(input_ids)) .unsqueeze(1) @@ -187,7 +146,10 @@ class TestScoreAndPool(CustomTestCase): .expand(-1, self.hidden_dim) .clone() ) - fb = _make_forward_batch(extend_seq_lens=[len(input_ids)], is_prefill_only=True) + fb = _make_forward_batch( + extend_seq_lens=[len(input_ids)], + multi_item_delimiter_indices=[torch.tensor([1, 3])], + ) identity_head = nn.Linear(self.hidden_dim, self.hidden_dim, bias=False) nn.init.eye_(identity_head.weight) @@ -195,25 +157,37 @@ class TestScoreAndPool(CustomTestCase): out = score_and_pool(identity_head, self.pooler, hidden, fb, input_ids) self.assertEqual(len(out.embeddings), 1) - self.assertEqual(out.embeddings[0].shape[0], 1) - torch.testing.assert_close(out.embeddings[0][0], hidden[2]) - - @patch("sglang.srt.layers.pooler.get_global_server_args") - def test_single_item_scores_match_manual_computation(self, mock_get_args): - """Single-item scores equal score_head applied to all tokens then pooled.""" - mock_get_args.return_value = _mock_server_args(delimiter=None) + self.assertEqual(out.embeddings[0].shape[0], 2) + torch.testing.assert_close(out.embeddings[0][0], hidden[0]) + torch.testing.assert_close(out.embeddings[0][1], hidden[2]) + def test_single_item_scores_match_manual_computation(self): + """Single-item scores equal score_head applied to pooled hidden states.""" hidden = torch.randn(8, self.hidden_dim) fb = _make_forward_batch(extend_seq_lens=[5, 3]) input_ids = torch.arange(8) out = score_and_pool(self.score_head, self.pooler, hidden, fb, input_ids) - # score-first-then-pool: matches the original Qwen3/Qwen2 classification forward - logits = self.score_head(hidden) - expected = self.pooler(logits, fb).embeddings + pooled = self.pooler(hidden, fb).embeddings + expected = self.score_head(pooled) torch.testing.assert_close(out.embeddings, expected) + def test_empty_delimiter_indices(self): + """Empty delimiter tensor per request -> returns list with empty tensor.""" + input_ids = torch.arange(6) + hidden = torch.randn(6, self.hidden_dim) + fb = _make_forward_batch( + extend_seq_lens=[6], + multi_item_delimiter_indices=[torch.tensor([], dtype=torch.long)], + ) + + out = score_and_pool(self.score_head, self.pooler, hidden, fb, input_ids) + + self.assertIsInstance(out.embeddings, list) + self.assertEqual(len(out.embeddings), 1) + self.assertEqual(out.embeddings[0].shape, (0, self.num_labels)) + if __name__ == "__main__": unittest.main()