diff --git a/.github/CI_PERMISSIONS.json b/.github/CI_PERMISSIONS.json index b5ecc9876..735f4c2ab 100644 --- a/.github/CI_PERMISSIONS.json +++ b/.github/CI_PERMISSIONS.json @@ -510,6 +510,13 @@ "cooldown_interval_minutes": 60, "reason": "custom override" }, + "fortunecookiee": { + "can_tag_run_ci_label": true, + "can_rerun_failed_ci": true, + "can_rerun_stage": true, + "cooldown_interval_minutes": 0, + "reason": "custom override" + }, "fy1214": { "can_tag_run_ci_label": true, "can_rerun_failed_ci": true, diff --git a/python/sglang/srt/entrypoints/engine.py b/python/sglang/srt/entrypoints/engine.py index d864e4aba..a012fde79 100644 --- a/python/sglang/srt/entrypoints/engine.py +++ b/python/sglang/srt/entrypoints/engine.py @@ -446,6 +446,8 @@ class Engine(EngineScoreMixin, EngineBase): video_data: Optional[MultimodalDataInputFormat] = None, dimensions: Optional[int] = None, lora_path: Optional[Union[List[Optional[str]], Optional[str]]] = None, + embed_override_token_id: Optional[int] = None, + embed_overrides: Optional[List[List[torch.Tensor]]] = None, external_trace_header: Optional[Dict] = None, rid: Optional[Union[List[str], str]] = None, ) -> Dict: @@ -460,6 +462,8 @@ class Engine(EngineScoreMixin, EngineBase): video_data=video_data, dimensions=dimensions, lora_path=lora_path, + embed_override_token_id=embed_override_token_id, + embed_overrides=embed_overrides, external_trace_header=external_trace_header, rid=rid, ) @@ -475,6 +479,8 @@ class Engine(EngineScoreMixin, EngineBase): video_data: Optional[MultimodalDataInputFormat] = None, dimensions: Optional[int] = None, lora_path: Optional[Union[List[Optional[str]], Optional[str]]] = None, + embed_override_token_id: Optional[int] = None, + embed_overrides: Optional[List[List[torch.Tensor]]] = None, external_trace_header: Optional[Dict] = None, rid: Optional[Union[List[str], str]] = None, ) -> Dict: @@ -491,6 +497,8 @@ class Engine(EngineScoreMixin, EngineBase): video_data=video_data, dimensions=dimensions, lora_path=lora_path, + embed_override_token_id=embed_override_token_id, + embed_overrides=embed_overrides, external_trace_header=external_trace_header, rid=rid, ) diff --git a/python/sglang/srt/entrypoints/engine_score_mixin.py b/python/sglang/srt/entrypoints/engine_score_mixin.py index 7e3c622a5..089693e80 100644 --- a/python/sglang/srt/entrypoints/engine_score_mixin.py +++ b/python/sglang/srt/entrypoints/engine_score_mixin.py @@ -20,6 +20,8 @@ by TokenizerManagerScoreMixin. from typing import List, Optional, Union +import torch + from sglang.srt.managers.tokenizer_manager_score_mixin import ScoreResult @@ -31,6 +33,12 @@ class EngineScoreMixin: label_token_ids: Optional[List[int]] = None, apply_softmax: bool = False, item_first: bool = False, + # Placeholder token id in query/items that indicates override locations. + embed_override_token_id: Optional[int] = None, + # Query embedding overrides. + query_embed_overrides: Optional[List[torch.Tensor]] = None, + # Item embedding overrides: per-item list of tensors. + item_embed_overrides: Optional[List[Optional[List[torch.Tensor]]]] = None, ) -> ScoreResult: """ Score items against a query using the loaded model. @@ -52,6 +60,9 @@ class EngineScoreMixin: SequenceClassification). apply_softmax: Whether to normalize scores using softmax. item_first: If True, prepend items before query (single-item mode only). + embed_override_token_id: Placeholder token ID used to locate override positions. + query_embed_overrides: Embedding vectors replacing placeholder tokens in query. + item_embed_overrides: Per-item embedding vectors replacing placeholder tokens in items. Returns: ScoreResult with scores (one list per item) and prompt token count. @@ -63,6 +74,9 @@ class EngineScoreMixin: label_token_ids=label_token_ids, apply_softmax=apply_softmax, item_first=item_first, + embed_override_token_id=embed_override_token_id, + query_embed_overrides=query_embed_overrides, + item_embed_overrides=item_embed_overrides, request=None, ) ) @@ -74,6 +88,9 @@ class EngineScoreMixin: label_token_ids: Optional[List[int]] = None, apply_softmax: bool = False, item_first: bool = False, + embed_override_token_id: Optional[int] = None, + query_embed_overrides: Optional[List[torch.Tensor]] = None, + item_embed_overrides: Optional[List[Optional[List[torch.Tensor]]]] = None, ) -> ScoreResult: """Asynchronous version of score(). See score() for full documentation.""" return await self.tokenizer_manager.score_request( @@ -82,5 +99,8 @@ class EngineScoreMixin: label_token_ids=label_token_ids, apply_softmax=apply_softmax, item_first=item_first, + embed_override_token_id=embed_override_token_id, + query_embed_overrides=query_embed_overrides, + item_embed_overrides=item_embed_overrides, request=None, ) diff --git a/python/sglang/srt/entrypoints/openai/protocol.py b/python/sglang/srt/entrypoints/openai/protocol.py index c40bd37d9..165850075 100644 --- a/python/sglang/srt/entrypoints/openai/protocol.py +++ b/python/sglang/srt/entrypoints/openai/protocol.py @@ -942,6 +942,11 @@ class EmbeddingRequest(BaseModel): priority: Optional[int] = None # LoRA adapter path(s) lora_path: Optional[Union[List[Optional[str]], Optional[str]]] = None + # Placeholder token id used to locate embedding override positions in input token IDs. + embed_override_token_id: Optional[int] = None + # Per-input embedding overrides (null entries skip that input). + # Shape: [num_inputs][num_replacements][hidden_size] + embed_overrides: Optional[List[Optional[List[List[float]]]]] = None class EmbeddingObject(BaseModel): @@ -995,6 +1000,16 @@ class ScoringRequest(BaseModel): items: Optional[Union[str, List[str], List[List[int]]]] = ( None # Item text(s) or pre-tokenized token IDs ) + # Placeholder token id used to locate embedding override positions in query/items. + embed_override_token_id: Optional[int] = None + # Query embedding overrides. + query_embed_overrides: Optional[List[List[float]]] = ( + None # [num_query_embed_overrides][hidden_size] + ) + # Per-item embedding overrides (null entries skip that item). + item_embed_overrides: Optional[List[Optional[List[List[float]]]]] = ( + None # [num_items][num_item_embed_overrides][hidden_size] + ) label_token_ids: Optional[List[int]] = ( None # Token IDs to compute probabilities for ) diff --git a/python/sglang/srt/entrypoints/openai/serving_embedding.py b/python/sglang/srt/entrypoints/openai/serving_embedding.py index 7555cc476..5f2f3658c 100644 --- a/python/sglang/srt/entrypoints/openai/serving_embedding.py +++ b/python/sglang/srt/entrypoints/openai/serving_embedding.py @@ -14,6 +14,7 @@ from sglang.srt.entrypoints.openai.protocol import ( UsageInfo, ) from sglang.srt.entrypoints.openai.serving_base import OpenAIServingBase +from sglang.srt.entrypoints.openai.utils import convert_embeds_to_tensors from sglang.srt.managers.io_struct import EmbeddingReqInput from sglang.srt.parser.conversation import generate_embedding_convs @@ -59,14 +60,14 @@ class OpenAIServingEmbedding(OpenAIServingBase): # List of strings for i, item in enumerate(input): if not isinstance(item, str): - return f"All items in input list must be strings" + return "All items in input list must be strings" if not item.strip(): return f"Input at index {i} cannot be empty or whitespace only" elif isinstance(first_item, int): # List of integers (token IDs) for i, item in enumerate(input): if not isinstance(item, int): - return f"All items in input list must be integers" + return "All items in input list must be integers" if item < 0: return f"Token ID at index {i} must be non-negative" return None @@ -129,6 +130,26 @@ class OpenAIServingEmbedding(OpenAIServingBase): # Resolve LoRA adapter from model parameter or explicit lora_path lora_path = self._resolve_lora_path(request.model, request.lora_path) + # Validate pairing: both or neither must be provided + if ( + request.embed_overrides is not None + and request.embed_override_token_id is None + ): + raise ValueError( + "embed_override_token_id is required when embed_overrides is provided" + ) + if ( + request.embed_override_token_id is not None + and request.embed_overrides is None + ): + raise ValueError( + "embed_override_token_id requires embed_overrides to be provided" + ) + + # Convert float lists to tensors; position resolution is deferred + # to the tokenizer manager (after tokenization for text inputs). + embed_overrides = convert_embeds_to_tensors(request.embed_overrides) + adapted_request = EmbeddingReqInput( **prompt_kwargs, rid=request.rid, @@ -136,6 +157,8 @@ class OpenAIServingEmbedding(OpenAIServingBase): routing_key=self.extract_routing_key(raw_request), dimensions=request.dimensions, lora_path=lora_path, + embed_override_token_id=request.embed_override_token_id, + embed_overrides=embed_overrides, ) return adapted_request, request diff --git a/python/sglang/srt/entrypoints/openai/serving_score.py b/python/sglang/srt/entrypoints/openai/serving_score.py index e9fb5f8c0..ff480b632 100644 --- a/python/sglang/srt/entrypoints/openai/serving_score.py +++ b/python/sglang/srt/entrypoints/openai/serving_score.py @@ -1,6 +1,7 @@ import logging from typing import Union +import torch from fastapi import Request from sglang.srt.entrypoints.openai.protocol import ( @@ -42,13 +43,38 @@ class OpenAIServingScore(OpenAIServingBase): ) -> Union[ScoringResponse, ErrorResponse]: """Handle the scoring request""" try: - # Use tokenizer_manager's score_request method directly + # query_embed_overrides is [num_replacements][hidden_size] -> List[Tensor] + query_embed_overrides = ( + [ + torch.tensor(v, dtype=torch.float32) + for v in request.query_embed_overrides + ] + if request.query_embed_overrides is not None + else None + ) + # item_embed_overrides is [num_items][num_replacements][hidden_size] -> List[Optional[List[Tensor]]] + item_embed_overrides = ( + [ + ( + [torch.tensor(v, dtype=torch.float32) for v in per_item] + if per_item is not None + else None + ) + for per_item in request.item_embed_overrides + ] + if request.item_embed_overrides is not None + else None + ) + result = await self.tokenizer_manager.score_request( query=request.query, items=request.items, label_token_ids=request.label_token_ids, apply_softmax=request.apply_softmax, item_first=request.item_first, + embed_override_token_id=request.embed_override_token_id, + query_embed_overrides=query_embed_overrides, + item_embed_overrides=item_embed_overrides, request=raw_request, ) diff --git a/python/sglang/srt/entrypoints/openai/utils.py b/python/sglang/srt/entrypoints/openai/utils.py index 2e08a2717..0756994aa 100644 --- a/python/sglang/srt/entrypoints/openai/utils.py +++ b/python/sglang/srt/entrypoints/openai/utils.py @@ -1,6 +1,8 @@ import logging from typing import Any, Dict, List, Optional, Union +import torch + from sglang.srt.entrypoints.openai.protocol import ( CachedTokensDetails, ChatCompletionRequest, @@ -132,3 +134,41 @@ def process_cached_tokens_details_from_ret( device=details.get("device", 0), host=details.get("host", 0), ) + + +def convert_embeds_to_tensors( + embeds: Optional[Union[List[Optional[List[List[float]]]], List[List[float]]]], +) -> Optional[List[Optional[List[torch.Tensor]]]]: + """Convert nested float lists from the HTTP API to lists of tensors. + + Accepts either: + - None -> returns None + - List[List[float]] (single input) -> [[tensor, ...]] + - List[Optional[List[List[float]]]] (batch) -> [Optional[List[tensor]], ...] + Each innermost List[float] becomes a 1-D torch.Tensor. + Per-input None entries are preserved (no overrides for that input). + """ + if embeds is None: + return None + if len(embeds) == 0: + return [] + # Find first non-None entry to detect nesting depth + first_non_none = next((e for e in embeds if e is not None), None) + if first_non_none is None: + # All entries are None + return [None] * len(embeds) + # Detect nesting depth by checking the first non-None entry: + # - Single input [num_replacements][hidden_size]: first element is List[float] + # - Batch [num_inputs][num_replacements][hidden_size]: first element is List[List[float]] + if not first_non_none or not isinstance(first_non_none[0], list): + # Single input: each entry is a float vector + return [[torch.tensor(vec, dtype=torch.float32) for vec in embeds]] + # Otherwise it's batch: [num_inputs][num_replacements][hidden_size] + return [ + ( + [torch.tensor(vec, dtype=torch.float32) for vec in per_input] + if per_input is not None + else None + ) + for per_input in embeds + ] diff --git a/python/sglang/srt/managers/embed_types.py b/python/sglang/srt/managers/embed_types.py new file mode 100644 index 000000000..987102aed --- /dev/null +++ b/python/sglang/srt/managers/embed_types.py @@ -0,0 +1,53 @@ +# Copyright 2023-2024 SGLang Team +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================== +""" +Dataclasses for embedding injection. + +These are placed in a separate module to avoid circular imports between +io_struct.py and schedule_batch.py. +""" + +from dataclasses import dataclass +from typing import List, Union + +import torch + + +@dataclass +class PositionalEmbeds: + """Embeddings to place at specific token positions. + + Accepts either a list of [1, hidden_dim] tensors or a pre-stacked [N, hidden_dim] tensor. + In both cases, __post_init__ stacks into a single [N, hidden_dim] tensor to reduce + ZMQ serialization overhead. + + Attributes: + embeds: Stacked tensor of shape [N, hidden_dim] after __post_init__. + positions: List of positions where embeddings should be injected. + """ + + embeds: Union[List[torch.Tensor], torch.Tensor] + positions: List[int] + + def __post_init__(self): + # Stack list of tensors into a single [N, hidden_dim] tensor + if isinstance(self.embeds, list): + self.embeds = torch.cat( + [e if e.dim() == 2 else e.unsqueeze(0) for e in self.embeds], dim=0 + ) + if self.embeds.shape[0] != len(self.positions): + raise ValueError( + f"embeds length ({self.embeds.shape[0]}) != " + f"positions length ({len(self.positions)})" + ) diff --git a/python/sglang/srt/managers/io_struct.py b/python/sglang/srt/managers/io_struct.py index 8c2e0bf7b..5d48b1538 100644 --- a/python/sglang/srt/managers/io_struct.py +++ b/python/sglang/srt/managers/io_struct.py @@ -29,6 +29,7 @@ from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Union import torch from sglang.srt.lora.lora_registry import LoRARef +from sglang.srt.managers.embed_types import PositionalEmbeds from sglang.srt.managers.schedule_batch import BaseFinishReason, Modality from sglang.srt.multimodal.mm_utils import has_valid_data from sglang.srt.observability.req_time_stats import ( @@ -137,6 +138,10 @@ class GenerateReqInput(BaseReq): input_ids: Optional[Union[List[List[int]], List[int]]] = None # The embeddings for input_ids; one can specify either text or input_ids or input_embeds. input_embeds: Optional[Union[List[List[List[float]]], List[List[float]]]] = None + # Embedding overrides to place at specific token positions. + # Runtime type: Optional[Union[PositionalEmbeds, List[Optional[PositionalEmbeds]]]] + # Typed as Any to avoid Pydantic/FastAPI schema errors (PositionalEmbeds contains torch.Tensor). + positional_embed_overrides: Any = None # The image input. It can be an image instance, file name, URL, or base64 encoded string. # Can be formatted as: # - Single image for a single request @@ -600,6 +605,16 @@ class GenerateReqInput(BaseReq): ): raise ValueError("Session params must be a dict or a list of dicts.") + def _get_positional_embed_overrides_item( + self, i: int + ) -> Optional[PositionalEmbeds]: + """Extract the i-th item from positional_embed_overrides.""" + if self.positional_embed_overrides is None: + return None + if isinstance(self.positional_embed_overrides, PositionalEmbeds): + return self.positional_embed_overrides + return self.positional_embed_overrides[i] + def __getitem__(self, i): # Cache sub-objects so that repeated obj[i] calls return the same instance. # This avoids subtle bugs where different call sites get divergent objects. @@ -612,6 +627,7 @@ class GenerateReqInput(BaseReq): input_embeds=( self.input_embeds[i] if self.input_embeds is not None else None ), + positional_embed_overrides=self._get_positional_embed_overrides_item(i), image_data=self.image_data[i], video_data=self.video_data[i], audio_data=self.audio_data[i], @@ -706,6 +722,9 @@ class TokenizedGenerateReqInput(BaseReq): # The input embeds input_embeds: Optional[Union[List[List[List[float]]], List[List[float]]]] = None + # Embedding overrides to place at specific token positions. + positional_embed_overrides: Optional[PositionalEmbeds] = None + # Session info for continual prompting session_params: Optional[SessionParams] = None @@ -791,6 +810,18 @@ class EmbeddingReqInput(BaseReq): audio_data: Optional[MultimodalDataInputFormat] = None # The token ids for text; one can either specify text or input_ids. input_ids: Optional[Union[List[List[int]], List[int]]] = None + # Placeholder token ID used to locate embedding override positions in input token IDs. + embed_override_token_id: Optional[int] = None + # Unresolved embedding overrides: per-input list of tensors. + # Position resolution happens in the tokenizer manager after tokenization. + # Shape: [num_inputs][num_replacements] where each entry is a torch.Tensor of [hidden_size]. + # Per-input entry may be None when only some inputs in a batch need overrides. + # Runtime type: Optional[List[Optional[List[torch.Tensor]]]] + # Typed as Any to avoid Pydantic/FastAPI schema errors (contains torch.Tensor). + embed_overrides: Any = None + # Resolved embedding overrides with positions (set by tokenizer manager or score mixin). + # Runtime type: Optional[Union[PositionalEmbeds, List[Optional[PositionalEmbeds]]]] + positional_embed_overrides: Any = None # Dummy sampling params for compatibility sampling_params: Optional[Union[List[Dict], Dict]] = None # Dummy input embeds for compatibility @@ -896,6 +927,16 @@ class EmbeddingReqInput(BaseReq): or has_valid_data(self.audio_data) ) + def _get_positional_embed_overrides_item( + self, i: int + ) -> Optional[PositionalEmbeds]: + """Extract the i-th item from positional_embed_overrides.""" + if self.positional_embed_overrides is None: + return None + if isinstance(self.positional_embed_overrides, PositionalEmbeds): + return self.positional_embed_overrides + return self.positional_embed_overrides[i] + def __getitem__(self, i): # Cache sub-objects so that repeated obj[i] calls return the same instance. cache = self.__dict__.setdefault("_sub_obj_cache", {}) @@ -905,6 +946,7 @@ class EmbeddingReqInput(BaseReq): if self.is_cross_encoder_request: sub = EmbeddingReqInput( text=[self.text[i]] if self.text is not None else None, + positional_embed_overrides=self._get_positional_embed_overrides_item(i), sampling_params=self.sampling_params[i], rid=self.rid[i], lora_path=self.lora_path[i] if self.lora_path is not None else None, @@ -916,6 +958,13 @@ class EmbeddingReqInput(BaseReq): sub = EmbeddingReqInput( text=self.text[i] if self.text is not None else None, input_ids=self.input_ids[i] if self.input_ids is not None else None, + embed_override_token_id=self.embed_override_token_id, + embed_overrides=( + self.embed_overrides[i] + if self.embed_overrides is not None + else None + ), + positional_embed_overrides=self._get_positional_embed_overrides_item(i), image_data=self.image_data[i] if self.image_data is not None else None, audio_data=self.audio_data[i] if self.audio_data is not None else None, video_data=self.video_data[i] if self.video_data is not None else None, @@ -944,12 +993,15 @@ class TokenizedEmbeddingReqInput(BaseReq): token_type_ids: List[int] # Dummy sampling params for compatibility sampling_params: SamplingParams + # Embedding overrides to place at specific token positions. + positional_embed_overrides: Optional[PositionalEmbeds] = None # For DP routing routed_dp_rank: Optional[int] = None # Priority for the request priority: Optional[int] = None # The number of dimensions the resulting output embeddings should have. It is applicable for Matryoshka Embeddings. dimensions: Optional[int] = None + # LoRA related lora_id: Optional[str] = None # None means just use the base model # For observability diff --git a/python/sglang/srt/managers/schedule_batch.py b/python/sglang/srt/managers/schedule_batch.py index 9eeab2572..078d6ab91 100644 --- a/python/sglang/srt/managers/schedule_batch.py +++ b/python/sglang/srt/managers/schedule_batch.py @@ -59,6 +59,7 @@ from sglang.srt.distributed.parallel_state import get_tensor_model_parallel_rank from sglang.srt.dllm.mixin.req import ReqDllmMixin from sglang.srt.environ import envs from sglang.srt.layers.attention.fla.chunk_delta_h import CHUNK_SIZE as FLA_CHUNK_SIZE +from sglang.srt.managers.embed_types import PositionalEmbeds from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache, MatchPrefixParams from sglang.srt.mem_cache.common import ( @@ -571,6 +572,7 @@ class Req(ReqDllmMixin): origin_input_ids_unpadded: Optional[Tuple[int]] = None, lora_id: Optional[str] = None, input_embeds: Optional[List[List[float]]] = None, + positional_embed_overrides: Optional[PositionalEmbeds] = None, token_type_ids: List[int] = None, session: Optional[Session] = None, custom_logit_processor: Optional[str] = None, @@ -610,6 +612,7 @@ class Req(ReqDllmMixin): self.fill_ids = [] self.session = session self.input_embeds = input_embeds + self.positional_embed_overrides = positional_embed_overrides # For req-level memory management self.kv_committed_len = 0 @@ -973,6 +976,12 @@ class Req(ReqDllmMixin): max_prefix_len = max(max_prefix_len, 0) token_ids = self.fill_ids[:max_prefix_len] + # Disable prefix caching when embed overrides are present: same token IDs + # with different override vectors must not share cached KV values. + if self.positional_embed_overrides is not None: + max_prefix_len = 0 + token_ids = [] + if tree_cache is not None: if cow_mamba is None: cow_mamba = tree_cache.supports_mamba() @@ -1332,6 +1341,9 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin): # Batched arguments to model runner input_ids: torch.Tensor = None # shape: [b], int64 input_embeds: torch.Tensor = None # shape: [b, hidden_size], float32 + # Token replacement embeddings and absolute positions (optional). + replace_embeds: Optional[torch.Tensor] = None + replace_positions: Optional[torch.Tensor] = None ne_token_table: torch.Tensor = None token_type_ids: torch.Tensor = None # shape: [b], int64 req_pool_indices: torch.Tensor = None # shape: [b], int64 @@ -1618,6 +1630,11 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin): # Set fields input_embeds = [] + all_replace_embeds: List[torch.Tensor] = [] + all_replace_positions: List[int] = [] + has_replace_embeds = False + input_id_pointer = 0 + input_id_lens = [len(input_id) for input_id in input_ids] extend_input_logprob_token_ids = [] multimodal_inputs = [] mamba_track_mask_cpu = [] @@ -1642,6 +1659,27 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin): req.input_embeds[pre_len : pre_len + req.extend_input_len] ) + if req.positional_embed_overrides is not None: + # Override positions are absolute in the full sequence. + # Convert to extend-tensor coordinates by subtracting pre_len, + # then skip any that fall within the cached prefix. + embeds_to_add = [] + for embed_idx, pos in enumerate( + req.positional_embed_overrides.positions + ): + extend_pos = pos - pre_len + if extend_pos < 0 or extend_pos >= req.extend_input_len: + continue # Outside current extend chunk, skip + embeds_to_add.append((embed_idx, input_id_pointer + extend_pos)) + if embeds_to_add: + has_replace_embeds = True + indices, positions = zip(*embeds_to_add) + all_replace_embeds.append( + req.positional_embed_overrides.embeds[list(indices)] + ) + all_replace_positions.extend(positions) + input_id_pointer += input_id_lens[i] + multimodal_inputs.append(req.multimodal_inputs) # Only calculate cached_tokens once. Once retracted, the 'retracted_stain' @@ -1737,6 +1775,17 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin): else: extend_input_logprob_token_ids = None + if has_replace_embeds: + replace_embeds_tensor = torch.cat(all_replace_embeds, dim=0).to( + self.device, non_blocking=True + ) + replace_positions_tensor = torch.tensor( + all_replace_positions, dtype=torch.long, device=self.device + ) + else: + replace_embeds_tensor = None + replace_positions_tensor = None + self.input_ids = input_ids_tensor self.req_pool_indices = req_pool_indices_tensor self.orig_seq_lens = orig_seq_lens_tensor @@ -1748,6 +1797,8 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin): if input_embeds else None ) + self.replace_embeds = replace_embeds_tensor + self.replace_positions = replace_positions_tensor self.multimodal_inputs = multimodal_inputs self.token_type_ids = token_type_ids_tensor self.seq_lens_sum = sum(seq_lens) @@ -2367,6 +2418,8 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin): lora_ids=[req.lora_id for req in self.reqs], sampling_info=self.sampling_info, input_embeds=self.input_embeds, + replace_embeds=self.replace_embeds, + replace_positions=self.replace_positions, ne_token_table=self.ne_token_table, token_type_ids=self.token_type_ids, spec_algorithm=self.spec_algorithm, @@ -2542,6 +2595,8 @@ class ModelWorkerBatch: # The input Embeds input_embeds: Optional[torch.Tensor] = None + replace_embeds: Optional[torch.Tensor] = None + replace_positions: Optional[torch.Tensor] = None # token table for ngram embedding ne_token_table: Optional[torch.Tensor] = None diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index 377a7ec74..5cfb32c68 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -1803,6 +1803,7 @@ class Scheduler( stream=recv_req.stream, lora_id=recv_req.lora_id, input_embeds=recv_req.input_embeds, + positional_embed_overrides=recv_req.positional_embed_overrides, token_type_ids=recv_req.token_type_ids, custom_logit_processor=recv_req.custom_logit_processor, require_reasoning=recv_req.require_reasoning, @@ -2128,6 +2129,7 @@ class Scheduler( recv_req.input_text, recv_req.input_ids, recv_req.sampling_params, + positional_embed_overrides=recv_req.positional_embed_overrides, token_type_ids=recv_req.token_type_ids, routed_dp_rank=recv_req.routed_dp_rank, priority=recv_req.priority, diff --git a/python/sglang/srt/managers/tokenizer_manager.py b/python/sglang/srt/managers/tokenizer_manager.py index 2f2876a8d..8e1166c77 100644 --- a/python/sglang/srt/managers/tokenizer_manager.py +++ b/python/sglang/srt/managers/tokenizer_manager.py @@ -33,6 +33,7 @@ from typing import Any, Awaitable, Dict, List, Optional, Tuple, Union import fastapi import pybase64 +import torch import uvloop import zmq import zmq.asyncio @@ -45,6 +46,7 @@ from sglang.srt.environ import envs from sglang.srt.lora.lora_registry import LoRARef, LoRARegistry from sglang.srt.managers.async_dynamic_batch_tokenizer import AsyncDynamicbatchTokenizer from sglang.srt.managers.disagg_service import start_disagg_service +from sglang.srt.managers.embed_types import PositionalEmbeds from sglang.srt.managers.io_struct import ( AbortReq, ActiveRanksOutput, @@ -1000,6 +1002,7 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerScoreMixin): bootstrap_room=obj.bootstrap_room, lora_id=obj.lora_id, input_embeds=input_embeds, + positional_embed_overrides=obj.positional_embed_overrides, session_params=session_params, custom_logit_processor=obj.custom_logit_processor, require_reasoning=obj.require_reasoning, @@ -1015,12 +1018,24 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerScoreMixin): num_items_assigned=obj.num_items_assigned, ) elif isinstance(obj, EmbeddingReqInput): + # Resolve unresolved embed overrides now that input_ids are available + positional_embed_overrides = obj.positional_embed_overrides + if ( + positional_embed_overrides is None + and obj.embed_overrides is not None + and obj.embed_override_token_id is not None + ): + positional_embed_overrides = self._resolve_embed_overrides( + input_ids, obj.embed_override_token_id, obj.embed_overrides + ) + tokenized_obj = TokenizedEmbeddingReqInput( input_text, input_ids, mm_inputs, token_type_ids, sampling_params, + positional_embed_overrides=positional_embed_overrides, rid=obj.rid, priority=obj.priority, dimensions=obj.dimensions, @@ -1033,6 +1048,26 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerScoreMixin): return tokenized_obj + @staticmethod + def _resolve_embed_overrides( + input_ids: List[int], + token_id: int, + embeds: List[torch.Tensor], + ) -> PositionalEmbeds: + """Resolve placeholder positions in input_ids and create PositionalEmbeds. + + Scans input_ids for occurrences of token_id and pairs them with the + provided embedding tensors. + """ + positions = [idx for idx, tok in enumerate(input_ids) if tok == token_id] + if len(positions) != len(embeds): + raise ValueError( + f"input contains {len(positions)} occurrences of " + f"embed_override_token_id={token_id}, " + f"but embed_overrides has {len(embeds)} entries." + ) + return PositionalEmbeds(embeds=embeds, positions=positions) + async def _batch_tokenize_and_process( self, batch_size: int, obj: Union[GenerateReqInput, EmbeddingReqInput] ) -> List[Union[TokenizedGenerateReqInput, TokenizedEmbeddingReqInput]]: diff --git a/python/sglang/srt/managers/tokenizer_manager_score_mixin.py b/python/sglang/srt/managers/tokenizer_manager_score_mixin.py index 80d6171b5..e6180430e 100644 --- a/python/sglang/srt/managers/tokenizer_manager_score_mixin.py +++ b/python/sglang/srt/managers/tokenizer_manager_score_mixin.py @@ -5,6 +5,7 @@ from typing import Any, Dict, List, Optional, Tuple, Union import torch +from sglang.srt.managers.embed_types import PositionalEmbeds from sglang.srt.managers.io_struct import EmbeddingReqInput, GenerateReqInput logger = logging.getLogger(__name__) @@ -294,6 +295,181 @@ class TokenizerManagerScoreMixin: return ScoreResult(scores=scores, prompt_tokens=prompt_tokens) + # ------------------------------------------------------------------ + # Embed override position resolution + # ------------------------------------------------------------------ + + def _resolve_overrides_for_sequence( + self, + token_ids: List[int], + embeds: Optional[List[torch.Tensor]], + embed_override_token_id: int, + position_offset: int = 0, + label: str = "input", + ) -> Tuple[List[torch.Tensor], List[int]]: + """Scan token_ids for placeholder occurrences and pair with embeddings. + + Args: + token_ids: The token sequence to scan. + embeds: Embedding tensors to place at placeholder positions (None = skip). + embed_override_token_id: The placeholder token ID. + position_offset: Added to each found position (for absolute coordinates). + label: Label for error messages (e.g. "query", "items[2]"). + + Returns: + (embeds, positions) lists. Empty lists if embeds is None. + """ + if embeds is None: + return [], [] + positions = [ + idx + position_offset + for idx, tok in enumerate(token_ids) + if tok == embed_override_token_id + ] + if len(positions) != len(embeds): + raise ValueError( + f"{label} contains {len(positions)} occurrences of " + f"embed_override_token_id={embed_override_token_id}, " + f"but {len(embeds)} override embeddings were provided." + ) + return embeds, positions + + def _resolve_embed_overrides_for_request( + self, + query: List[int], + item: List[int], + embed_override_token_id: int, + query_embed_overrides: Optional[List[torch.Tensor]], + item_embeds: Optional[List[torch.Tensor]], + item_position_offset: int, + item_label: str, + ) -> Optional[PositionalEmbeds]: + """Resolve embed overrides for a single query+item pair. + + Returns PositionalEmbeds if any overrides exist, None otherwise. + """ + q_embeds, q_positions = self._resolve_overrides_for_sequence( + query, + query_embed_overrides, + embed_override_token_id, + position_offset=0, + label="query", + ) + i_embeds, i_positions = self._resolve_overrides_for_sequence( + item, + item_embeds, + embed_override_token_id, + position_offset=item_position_offset, + label=item_label, + ) + all_embeds = q_embeds + i_embeds + all_positions = q_positions + i_positions + if not all_embeds: + return None + return PositionalEmbeds(embeds=all_embeds, positions=all_positions) + + # ------------------------------------------------------------------ + # Input preparation (tokenization + input_ids construction) + # ------------------------------------------------------------------ + + def _build_token_id_inputs( + self, + query: List[int], + items: List[List[int]], + item_first: bool, + use_multi_item_scoring: bool, + 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]]: + """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. + """ + # Both query and items are token IDs + has_embeds = ( + query_embed_overrides is not None or item_embed_overrides is not None + ) + + 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 + ) + input_ids = [combined_input_ids] + + if not has_embeds: + return None, input_ids, None + + # Resolve embed overrides across the combined multi-item-scoring sequence + all_embeds: List[torch.Tensor] = [] + all_positions: List[int] = [] + current_offset = len(query) + 1 # +1 for first delimiter + 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( + query if i == 0 else [], # only resolve query overrides once + item, + embed_override_token_id, + query_embed_overrides if i == 0 else None, + item_embs, + current_offset, + f"items[{i}]", + ) + if pe is not None: + # pe.embeds is a stacked tensor after PositionalEmbeds.__post_init__ + all_embeds.append(pe.embeds) + all_positions.extend(pe.positions) + current_offset += len(item) + 1 # +1 for delimiter + + if all_embeds: + injection = [ + PositionalEmbeds( + embeds=torch.cat(all_embeds, dim=0), + positions=all_positions, + ) + ] + else: + injection = None + return None, input_ids, injection + + else: + # Single-item scoring: process each item separately + if item_first: + input_ids = [item + query for item in items] + else: + input_ids = [query + item for item in items] + + if not has_embeds: + return None, input_ids, None + + injection = [] + 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( + query, + item, + embed_override_token_id, + query_embed_overrides, + item_embs, + item_position_offset=len(query), + item_label=f"items[{i}]", + ) + injection.append(pe) + + return ( + None, + input_ids, + injection if any(pe is not None for pe in injection) else None, + ) + + # ------------------------------------------------------------------ + # Main entry point + # ------------------------------------------------------------------ + async def score_request( self, query: Optional[Union[str, List[int]]] = None, @@ -301,6 +477,9 @@ class TokenizerManagerScoreMixin: label_token_ids: Optional[List[int]] = None, apply_softmax: bool = False, item_first: bool = False, + embed_override_token_id: Optional[int] = None, + query_embed_overrides: Optional[List[torch.Tensor]] = None, + item_embed_overrides: Optional[List[Optional[List[torch.Tensor]]]] = None, request: Optional[Any] = None, ) -> ScoreResult: """ @@ -327,6 +506,9 @@ class TokenizerManagerScoreMixin: label_token_ids: List of token IDs to compute probabilities for apply_softmax: Whether to normalize probabilities using softmax item_first: If True, prepend items to query. Ignored for multi-item scoring. + embed_override_token_id: Placeholder token ID for embedding override positions. + query_embed_overrides: Embedding vectors replacing placeholder tokens in query. + item_embed_overrides: Per-item embedding vectors replacing placeholder tokens in items. request: Optional FastAPI request object Returns: @@ -340,12 +522,26 @@ class TokenizerManagerScoreMixin: raise ValueError( "label_token_ids is required for generation (CausalLM) models." ) - if items is None: raise ValueError("items must be provided") if not items: return ScoreResult(scores=[], prompt_tokens=0) + has_embeds = ( + query_embed_overrides is not None or item_embed_overrides is not None + ) + if has_embeds and embed_override_token_id is None: + raise ValueError( + "embed_override_token_id is required when query_embed_overrides " + "or item_embed_overrides are supplied." + ) + if item_first and has_embeds: + raise ValueError("item_first is not supported when embeddings are supplied") + if item_embed_overrides is not None and len(item_embed_overrides) != len(items): + raise ValueError( + f"item_embed_overrides length ({len(item_embed_overrides)}) " + f"must match items length ({len(items)})." + ) if self.tokenizer is not None and label_token_ids is not None: vocab_size = self.tokenizer.vocab_size for token_id in label_token_ids: @@ -362,15 +558,13 @@ class TokenizerManagerScoreMixin: input_ids = None text_prompts = None + positional_embed_overrides = None - # Handle string or tokenized query/items - if isinstance(query, str) and ( - isinstance(items, str) - or (isinstance(items, list) and (not items or isinstance(items[0], str))) - ): + use_text_prompts = isinstance(query, str) and not has_embeds + + if use_text_prompts: # 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 @@ -396,21 +590,29 @@ class TokenizerManagerScoreMixin: and items and isinstance(items[0], list) ): - # Both query and items are token IDs - 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 - ) - input_ids = [combined_input_ids] - else: - # Single-item scoring: process each item separately - if item_first: - input_ids = [item + query for item in items] - else: - input_ids = [query + item for item in items] + # 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, + ) + 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, + ) else: raise ValueError( "Invalid combination of query/items types for score_request." @@ -427,11 +629,13 @@ class TokenizerManagerScoreMixin: logprob_start_len=0 if use_multi_item_scoring else -1, stream=False, sampling_params={"max_new_tokens": 0}, + positional_embed_overrides=positional_embed_overrides, ) else: batch_request = EmbeddingReqInput( text=text_prompts, input_ids=input_ids, + positional_embed_overrides=positional_embed_overrides, ) results = await self.generate_request(batch_request, request).__anext__() diff --git a/python/sglang/srt/model_executor/cuda_graph_runner.py b/python/sglang/srt/model_executor/cuda_graph_runner.py index 69cb176ef..e906de2b7 100644 --- a/python/sglang/srt/model_executor/cuda_graph_runner.py +++ b/python/sglang/srt/model_executor/cuda_graph_runner.py @@ -675,6 +675,9 @@ class CudaGraphRunner: return torch.int64 def can_run(self, forward_batch: ForwardBatch): + # Disable for token embedding overrides (dynamic per-request) + if forward_batch.replace_embeds is not None: + return False if self.require_mlp_tp_gather: cuda_graph_bs = ( max(forward_batch.global_num_tokens_cpu) // self.num_tokens_per_bs diff --git a/python/sglang/srt/model_executor/forward_batch_info.py b/python/sglang/srt/model_executor/forward_batch_info.py index eaecdc54b..831b3b6a0 100644 --- a/python/sglang/srt/model_executor/forward_batch_info.py +++ b/python/sglang/srt/model_executor/forward_batch_info.py @@ -359,6 +359,10 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin): # For input embeddings input_embeds: Optional[torch.Tensor] = None + # For token embedding overrides (sparse replacement at specific positions) + replace_embeds: Optional[torch.Tensor] = None + replace_positions: Optional[torch.Tensor] = None + # For cross-encoder model token_type_ids: Optional[torch.Tensor] = None @@ -473,6 +477,8 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin): spec_info=batch.spec_info, capture_hidden_mode=batch.capture_hidden_mode, input_embeds=batch.input_embeds, + replace_embeds=batch.replace_embeds, + replace_positions=batch.replace_positions, token_type_ids=batch.token_type_ids, tbo_split_seq_index=batch.tbo_split_seq_index, dimensions=batch.dimensions, diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py index a00d0f989..6ddb53ff1 100644 --- a/python/sglang/srt/model_executor/model_runner.py +++ b/python/sglang/srt/model_executor/model_runner.py @@ -2738,6 +2738,17 @@ class ModelRunner(ModelRunnerKVCacheMixin): kwargs["pp_proxy_tensors"] = pp_proxy_tensors if forward_batch.input_embeds is not None: kwargs["input_embeds"] = forward_batch.input_embeds.bfloat16() + if ( + forward_batch.replace_embeds is not None + and forward_batch.replace_positions is not None + ): + # Token embedding overrides: get base embeddings, scatter replacements + if "input_embeds" not in kwargs: + embed_layer = self.model.get_input_embeddings() + kwargs["input_embeds"] = embed_layer(forward_batch.input_ids) + kwargs["input_embeds"][forward_batch.replace_positions] = ( + forward_batch.replace_embeds.to(kwargs["input_embeds"].dtype) + ) if not self.is_generation: kwargs["get_embedding"] = True diff --git a/python/sglang/srt/model_executor/piecewise_cuda_graph_runner.py b/python/sglang/srt/model_executor/piecewise_cuda_graph_runner.py index 6f0ba95a4..efec70dc3 100644 --- a/python/sglang/srt/model_executor/piecewise_cuda_graph_runner.py +++ b/python/sglang/srt/model_executor/piecewise_cuda_graph_runner.py @@ -417,6 +417,9 @@ class PiecewiseCudaGraphRunner: # TODO(yuwei): fix it if forward_batch.input_embeds is not None: return False + # Disable for token embedding overrides (dynamic per-request) + if forward_batch.replace_embeds is not None: + return False num_tokens = len(forward_batch.input_ids) if forward_batch.return_logprob: for start_len, seq_len in zip( diff --git a/test/registered/unit/managers/test_embed_overrides.py b/test/registered/unit/managers/test_embed_overrides.py new file mode 100644 index 000000000..42a77c6a9 --- /dev/null +++ b/test/registered/unit/managers/test_embed_overrides.py @@ -0,0 +1,602 @@ +"""Unit tests for token embedding override support. + +Covers: +- PositionalEmbeds dataclass (embed_types.py) +- convert_embeds_to_tensors (utils.py) +- TokenizerManager._resolve_embed_overrides (tokenizer_manager.py) +- positional_embed_overrides on GenerateReqInput/EmbeddingReqInput (io_struct.py) +- Score mixin override resolution (tokenizer_manager_score_mixin.py) +""" + +import unittest +from unittest.mock import AsyncMock, MagicMock + +import torch + +from sglang.srt.entrypoints.openai.utils import convert_embeds_to_tensors +from sglang.srt.managers.embed_types import PositionalEmbeds +from sglang.srt.managers.io_struct import EmbeddingReqInput, GenerateReqInput +from sglang.srt.managers.tokenizer_manager import TokenizerManager +from sglang.srt.managers.tokenizer_manager_score_mixin import ( + TokenizerManagerScoreMixin, +) +from sglang.test.ci.ci_register import register_cuda_ci +from sglang.test.test_utils import CustomTestCase + +register_cuda_ci(est_time=5, suite="stage-b-test-1-gpu-small") + +HIDDEN_DIM = 4 + + +def _vec(val: float = 1.0) -> torch.Tensor: + """Create a 1-D tensor of size HIDDEN_DIM.""" + return torch.full((HIDDEN_DIM,), val, dtype=torch.float32) + + +def _vec2d(val: float = 1.0) -> torch.Tensor: + """Create a [1, HIDDEN_DIM] tensor.""" + return torch.full((1, HIDDEN_DIM), val, dtype=torch.float32) + + +# ======================================================================== +# PositionalEmbeds +# ======================================================================== + + +class TestPositionalEmbeds(CustomTestCase): + def test_from_list_of_1d_tensors(self): + pe = PositionalEmbeds(embeds=[_vec(1), _vec(2)], positions=[0, 5]) + self.assertEqual(pe.embeds.shape, (2, HIDDEN_DIM)) + self.assertAlmostEqual(pe.embeds[0, 0].item(), 1.0) + self.assertAlmostEqual(pe.embeds[1, 0].item(), 2.0) + + def test_from_list_of_2d_tensors(self): + pe = PositionalEmbeds(embeds=[_vec2d(3), _vec2d(4)], positions=[1, 2]) + self.assertEqual(pe.embeds.shape, (2, HIDDEN_DIM)) + + def test_from_pre_stacked_tensor(self): + stacked = torch.zeros(3, HIDDEN_DIM) + pe = PositionalEmbeds(embeds=stacked, positions=[0, 1, 2]) + self.assertIs(pe.embeds, stacked) + + def test_length_mismatch_raises(self): + with self.assertRaises(ValueError): + PositionalEmbeds(embeds=[_vec()], positions=[0, 1]) + + def test_empty(self): + pe = PositionalEmbeds(embeds=torch.zeros(0, HIDDEN_DIM), positions=[]) + self.assertEqual(pe.embeds.shape[0], 0) + + +# ======================================================================== +# convert_embeds_to_tensors +# ======================================================================== + + +class TestConvertEmbedsToTensors(CustomTestCase): + def test_none_returns_none(self): + self.assertIsNone(convert_embeds_to_tensors(None)) + + def test_empty_list(self): + self.assertEqual(convert_embeds_to_tensors([]), []) + + def test_single_input(self): + """[num_replacements][hidden_size] -> [[tensor, ...]]""" + result = convert_embeds_to_tensors([[1.0, 2.0], [3.0, 4.0]]) + self.assertEqual(len(result), 1) # wrapped in outer list + self.assertEqual(len(result[0]), 2) # two replacement vectors + self.assertTrue(torch.is_tensor(result[0][0])) + self.assertEqual(result[0][0].tolist(), [1.0, 2.0]) + self.assertEqual(result[0][0].dtype, torch.float32) + self.assertEqual(result[0][0].dim(), 1) # each vector is 1-D + + def test_batch_input(self): + """[num_inputs][num_replacements][hidden_size] -> [[tensor, ...], ...]""" + result = convert_embeds_to_tensors( + [ + [[1.0, 2.0]], + [[3.0, 4.0], [5.0, 6.0]], + ] + ) + self.assertEqual(len(result), 2) + self.assertEqual(len(result[0]), 1) + self.assertEqual(len(result[1]), 2) + + +# ======================================================================== +# TokenizerManager._resolve_embed_overrides +# ======================================================================== + + +class TestResolveEmbedOverrides(CustomTestCase): + def test_basic_resolution(self): + embeds = [_vec(1), _vec(2)] + pe = TokenizerManager._resolve_embed_overrides( + input_ids=[10, 50, 20, 50, 30], + token_id=50, + embeds=embeds, + ) + self.assertIsInstance(pe, PositionalEmbeds) + self.assertEqual(pe.positions, [1, 3]) + self.assertEqual(pe.embeds.shape, (2, HIDDEN_DIM)) + + def test_no_placeholders_raises(self): + with self.assertRaises(ValueError): + TokenizerManager._resolve_embed_overrides( + input_ids=[10, 20, 30], + token_id=50, + embeds=[_vec()], + ) + + def test_count_mismatch_raises(self): + with self.assertRaises(ValueError): + TokenizerManager._resolve_embed_overrides( + input_ids=[10, 50, 20], + token_id=50, + embeds=[_vec(1), _vec(2)], + ) + + +# ======================================================================== +# io_struct: positional_embed_overrides on GenerateReqInput +# ======================================================================== + + +class TestGenerateReqInputEmbedOverride(CustomTestCase): + def test_single_override_in_getitem(self): + """Single PositionalEmbeds is shared across all items in __getitem__.""" + pe = PositionalEmbeds(embeds=[_vec()], positions=[0]) + req = GenerateReqInput( + input_ids=[[1, 2], [3, 4]], + sampling_params=[{}, {}], + positional_embed_overrides=pe, + ) + req.normalize_batch_and_arguments() + item = req[0] + self.assertIs(item.positional_embed_overrides, pe) + + def test_batch_override_in_getitem(self): + """List[Optional[PositionalEmbeds]] is indexed per-item.""" + pe0 = PositionalEmbeds(embeds=[_vec(1)], positions=[0]) + pe1 = None + req = GenerateReqInput( + input_ids=[[1, 2], [3, 4]], + sampling_params=[{}, {}], + positional_embed_overrides=[pe0, pe1], + ) + req.normalize_batch_and_arguments() + self.assertEqual(req[0].positional_embed_overrides, pe0) + self.assertIsNone(req[1].positional_embed_overrides) + + +# ======================================================================== +# io_struct: embed override fields on EmbeddingReqInput +# ======================================================================== + + +class TestEmbeddingReqInputEmbedOverride(CustomTestCase): + def test_override_fields_in_getitem(self): + """embed_override_token_id, embed_overrides, and positional_embed_overrides + are correctly sliced in __getitem__.""" + pe0 = PositionalEmbeds(embeds=[_vec(1)], positions=[0]) + pe1 = PositionalEmbeds(embeds=[_vec(2)], positions=[1]) + req = EmbeddingReqInput( + input_ids=[[50, 10], [20, 50]], + sampling_params=[{}, {}], + embed_override_token_id=50, + embed_overrides=[[_vec(1)], [_vec(2)]], + positional_embed_overrides=[pe0, pe1], + ) + req.normalize_batch_and_arguments() + item0 = req[0] + item1 = req[1] + self.assertEqual(item0.embed_override_token_id, 50) + self.assertEqual(len(item0.embed_overrides), 1) + self.assertEqual(item0.positional_embed_overrides, pe0) + self.assertEqual(item1.positional_embed_overrides, pe1) + + +# ======================================================================== +# Score mixin: _resolve_overrides_for_sequence +# ======================================================================== + + +class _FakeServerArgs: + """Minimal stub for server_args.""" + + def __init__(self, multi_item_scoring_delimiter=None): + self.multi_item_scoring_delimiter = multi_item_scoring_delimiter + + +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 + self.tokenizer = None + self.is_generation = True + + +class TestResolveOverridesForSequence(CustomTestCase): + def setUp(self): + self.mixin = _FakeMixin() + + def test_none_embeds_returns_empty(self): + embeds, positions = self.mixin._resolve_overrides_for_sequence( + token_ids=[10, 50, 20], + embeds=None, + embed_override_token_id=50, + ) + self.assertEqual(embeds, []) + self.assertEqual(positions, []) + + def test_basic_resolution(self): + e1, e2 = _vec(1), _vec(2) + embeds, positions = self.mixin._resolve_overrides_for_sequence( + token_ids=[50, 10, 50], + embeds=[e1, e2], + embed_override_token_id=50, + ) + self.assertEqual(len(embeds), 2) + self.assertEqual(positions, [0, 2]) + + def test_with_offset(self): + embeds, positions = self.mixin._resolve_overrides_for_sequence( + token_ids=[10, 50], + embeds=[_vec()], + embed_override_token_id=50, + position_offset=100, + ) + self.assertEqual(positions, [101]) + + def test_empty_embeds_list(self): + """Empty embeds list with no placeholders succeeds.""" + embeds, positions = self.mixin._resolve_overrides_for_sequence( + token_ids=[10, 20], + embeds=[], + embed_override_token_id=50, + ) + self.assertEqual(embeds, []) + self.assertEqual(positions, []) + + def test_count_mismatch_raises(self): + with self.assertRaises(ValueError): + self.mixin._resolve_overrides_for_sequence( + token_ids=[50, 50], + embeds=[_vec()], + embed_override_token_id=50, + ) + + +# ======================================================================== +# Score mixin: _resolve_embed_overrides_for_request +# ======================================================================== + + +class TestResolveEmbedOverridesForRequest(CustomTestCase): + def setUp(self): + self.mixin = _FakeMixin() + + def test_no_overrides_returns_none(self): + result = self.mixin._resolve_embed_overrides_for_request( + query=[10, 20], + item=[30, 40], + embed_override_token_id=50, + query_embed_overrides=None, + item_embeds=None, + item_position_offset=2, + item_label="items[0]", + ) + self.assertIsNone(result) + + def test_query_only_overrides(self): + pe = self.mixin._resolve_embed_overrides_for_request( + query=[50, 20], + item=[30, 40], + embed_override_token_id=50, + query_embed_overrides=[_vec(1)], + item_embeds=None, + item_position_offset=2, + item_label="items[0]", + ) + self.assertIsInstance(pe, PositionalEmbeds) + self.assertEqual(pe.positions, [0]) + self.assertEqual(pe.embeds.shape, (1, HIDDEN_DIM)) + + def test_item_only_overrides(self): + pe = self.mixin._resolve_embed_overrides_for_request( + query=[10, 20], + item=[50, 40], + embed_override_token_id=50, + query_embed_overrides=None, + item_embeds=[_vec(2)], + item_position_offset=2, + item_label="items[0]", + ) + self.assertEqual(pe.positions, [2]) # offset applied + + def test_query_and_item_overrides(self): + pe = self.mixin._resolve_embed_overrides_for_request( + query=[50, 20], + item=[30, 50], + embed_override_token_id=50, + query_embed_overrides=[_vec(1)], + item_embeds=[_vec(2)], + item_position_offset=2, + item_label="items[0]", + ) + self.assertEqual(pe.positions, [0, 3]) # query pos 0, item pos 1+offset 2 + self.assertEqual(pe.embeds.shape, (2, HIDDEN_DIM)) + + +# ======================================================================== +# Score mixin: _build_token_id_inputs +# ======================================================================== + +DELIM_TOKEN = 99 + + +class TestBuildTokenIdInputs(CustomTestCase): + def setUp(self): + self.mixin = _FakeMixin(delimiter=DELIM_TOKEN) + + # --- single-item mode, no embeds --- + + def test_single_item_no_embeds(self): + _, input_ids, injection = self.mixin._build_token_id_inputs( + query=[1, 2], + items=[[3, 4], [5, 6]], + item_first=False, + use_multi_item_scoring=False, + embed_override_token_id=None, + query_embed_overrides=None, + item_embed_overrides=None, + ) + self.assertEqual(input_ids, [[1, 2, 3, 4], [1, 2, 5, 6]]) + self.assertIsNone(injection) + + def test_single_item_no_embeds_item_first(self): + _, input_ids, injection = self.mixin._build_token_id_inputs( + query=[1, 2], + items=[[3, 4]], + item_first=True, + use_multi_item_scoring=False, + embed_override_token_id=None, + query_embed_overrides=None, + item_embed_overrides=None, + ) + self.assertEqual(input_ids, [[3, 4, 1, 2]]) + self.assertIsNone(injection) + + # --- multi-item mode, no embeds --- + + def test_multi_item_no_embeds(self): + _, input_ids, injection = self.mixin._build_token_id_inputs( + query=[1, 2], + items=[[3, 4], [5, 6]], + item_first=False, + use_multi_item_scoring=True, + embed_override_token_id=None, + query_embed_overrides=None, + item_embed_overrides=None, + ) + # queryitem1item2 + self.assertEqual( + input_ids, [[1, 2, DELIM_TOKEN, 3, 4, DELIM_TOKEN, 5, 6, DELIM_TOKEN]] + ) + self.assertIsNone(injection) + + # --- 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( + query=[50, 10], + items=[[20, 30], [40, 50]], + item_first=False, + use_multi_item_scoring=False, + embed_override_token_id=50, + query_embed_overrides=[_vec(1)], + item_embed_overrides=None, + ) + self.assertEqual(input_ids, [[50, 10, 20, 30], [50, 10, 40, 50]]) + self.assertIsNotNone(injection) + self.assertEqual(len(injection), 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]) + + def test_single_item_item_embeds(self): + """Per-item overrides with correct position offsets.""" + _, input_ids, injection = self.mixin._build_token_id_inputs( + query=[10, 20], + items=[[50, 30]], + item_first=False, + use_multi_item_scoring=False, + embed_override_token_id=50, + query_embed_overrides=None, + item_embed_overrides=[[_vec(2)]], + ) + self.assertEqual(input_ids, [[10, 20, 50, 30]]) + self.assertIsNotNone(injection) + # item placeholder at index 0 of item, offset by query length 2 + self.assertEqual(injection[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( + query=[10, 20], + items=[[30, 40]], + item_first=False, + use_multi_item_scoring=False, + embed_override_token_id=50, + query_embed_overrides=None, + item_embed_overrides=[None], + ) + self.assertIsNone(injection) + + 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( + query=[50, 10], + items=[[20, 50]], + item_first=False, + use_multi_item_scoring=False, + embed_override_token_id=50, + query_embed_overrides=[_vec(1)], + item_embed_overrides=[[_vec(2)]], + ) + self.assertEqual(input_ids, [[50, 10, 20, 50]]) + self.assertIsNotNone(injection) + pe = injection[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( + query=[], + items=[[50, 10]], + item_first=False, + use_multi_item_scoring=False, + embed_override_token_id=50, + query_embed_overrides=None, + item_embed_overrides=[[_vec(1)]], + ) + self.assertEqual(input_ids, [[50, 10]]) + self.assertIsNotNone(injection) + # item placeholder at absolute pos 0 (offset=len([])=0) + self.assertEqual(injection[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( + query=[50, 10], + items=[[20, 50], [30, 40]], + item_first=False, + use_multi_item_scoring=True, + embed_override_token_id=50, + query_embed_overrides=[_vec(1)], + item_embed_overrides=[[_vec(2)], None], + ) + # queryitem1item2 = [50,10, 99, 20,50, 99, 30,40, 99] + self.assertEqual(len(input_ids), 1) + self.assertIsNotNone(injection) + self.assertEqual( + len(injection), 1 + ) # single PositionalEmbeds for combined sequence + pe = injection[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) + self.assertEqual(pe.embeds.shape[0], 2) + + +# ======================================================================== +# Score mixin: score_request validation +# ======================================================================== + + +class TestScoreRequestValidation(CustomTestCase): + """Test validation guards in score_request without running full pipeline.""" + + def setUp(self): + self.mixin = _FakeMixin() + + def _call(self, **kwargs): + """Wrapper to call score_request synchronously.""" + import asyncio + + return asyncio.run(self.mixin.score_request(**kwargs)) + + def test_generation_requires_label_token_ids(self): + self.mixin.is_generation = True + with self.assertRaisesRegex(ValueError, "label_token_ids is required"): + self._call( + query=[1, 2], + items=[[3, 4]], + label_token_ids=None, + ) + + def test_seq_classification_allows_none_label_token_ids(self): + """SequenceClassification models should not require label_token_ids. + Verify it passes validation and reaches generate_request.""" + self.mixin.is_generation = False + mock_result = AsyncMock() + mock_result.__anext__ = AsyncMock( + return_value=[{"embedding": [0.1, 0.9], "meta_info": {"prompt_tokens": 2}}] + ) + self.mixin.generate_request = MagicMock(return_value=mock_result) + result = self._call( + query=[1, 2], + items=[[3, 4]], + label_token_ids=None, + ) + self.mixin.generate_request.assert_called_once() + self.assertEqual(len(result.scores), 1) + + def test_items_none_raises(self): + with self.assertRaisesRegex(ValueError, "items must be provided"): + self._call( + query=[1, 2], + items=None, + label_token_ids=[100], + ) + + def test_empty_items_returns_empty(self): + result = self._call( + query=[1, 2], + items=[], + label_token_ids=[100], + ) + self.assertEqual(result.scores, []) + self.assertEqual(result.prompt_tokens, 0) + + def test_embed_override_token_id_required_with_query_embeds(self): + with self.assertRaisesRegex(ValueError, "embed_override_token_id is required"): + self._call( + query=[1, 2], + items=[[3, 4]], + label_token_ids=[100], + query_embed_overrides=[_vec(1)], + embed_override_token_id=None, + ) + + def test_embed_override_token_id_required_with_item_embeds(self): + with self.assertRaisesRegex(ValueError, "embed_override_token_id is required"): + self._call( + query=[1, 2], + items=[[3, 4]], + label_token_ids=[100], + item_embed_overrides=[[_vec(1)]], + embed_override_token_id=None, + ) + + def test_item_first_with_embeds_raises(self): + with self.assertRaisesRegex(ValueError, "item_first is not supported"): + self._call( + query=[1, 2], + items=[[3, 4]], + label_token_ids=[100], + item_first=True, + embed_override_token_id=50, + query_embed_overrides=[_vec(1)], + ) + + def test_item_embed_overrides_length_mismatch_raises(self): + with self.assertRaisesRegex(ValueError, "must match items length"): + self._call( + query=[1, 2], + items=[[3, 4], [5, 6]], + label_token_ids=[100], + embed_override_token_id=50, + item_embed_overrides=[[_vec(1)]], # 1 override for 2 items + ) + + +if __name__ == "__main__": + unittest.main()