[Feature] Add token embedding overrides for sparse embedding replacement (#20960)

Co-authored-by: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
jsheng_Linkedin
2026-04-08 20:51:36 -07:00
committed by GitHub
co-authored by Claude Opus 4.6
parent a69be2e866
commit 6838a23226
18 changed files with 1190 additions and 25 deletions
+7
View File
@@ -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,
+8
View File
@@ -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,
)
@@ -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,
)
@@ -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
)
@@ -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
@@ -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,
)
@@ -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
]
+53
View File
@@ -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)})"
)
+52
View File
@@ -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
@@ -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
+2
View File
@@ -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,
@@ -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]]:
@@ -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: query<delimiter_token_id>item1<delimiter_token_id>item2<delimiter_token_id>item3<delimiter_token_id>
delimiter_token_id = self.server_args.multi_item_scoring_delimiter
combined_input_ids = self._build_multi_item_token_sequence(
query, items, delimiter_token_id
)
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: query<delimiter_token_id>item1<delimiter_token_id>item2<delimiter_token_id>item3<delimiter_token_id>
delimiter_token_id = self.server_args.multi_item_scoring_delimiter
combined_input_ids = self._build_multi_item_token_sequence(
query, items, delimiter_token_id
)
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__()
@@ -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
@@ -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,
@@ -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
@@ -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(
@@ -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,
)
# query<D>item1<D>item2<D>
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],
)
# query<D>item1<D>item2<D> = [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()