[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, "cooldown_interval_minutes": 60,
"reason": "custom override" "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": { "fy1214": {
"can_tag_run_ci_label": true, "can_tag_run_ci_label": true,
"can_rerun_failed_ci": true, "can_rerun_failed_ci": true,
+8
View File
@@ -446,6 +446,8 @@ class Engine(EngineScoreMixin, EngineBase):
video_data: Optional[MultimodalDataInputFormat] = None, video_data: Optional[MultimodalDataInputFormat] = None,
dimensions: Optional[int] = None, dimensions: Optional[int] = None,
lora_path: Optional[Union[List[Optional[str]], Optional[str]]] = 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, external_trace_header: Optional[Dict] = None,
rid: Optional[Union[List[str], str]] = None, rid: Optional[Union[List[str], str]] = None,
) -> Dict: ) -> Dict:
@@ -460,6 +462,8 @@ class Engine(EngineScoreMixin, EngineBase):
video_data=video_data, video_data=video_data,
dimensions=dimensions, dimensions=dimensions,
lora_path=lora_path, lora_path=lora_path,
embed_override_token_id=embed_override_token_id,
embed_overrides=embed_overrides,
external_trace_header=external_trace_header, external_trace_header=external_trace_header,
rid=rid, rid=rid,
) )
@@ -475,6 +479,8 @@ class Engine(EngineScoreMixin, EngineBase):
video_data: Optional[MultimodalDataInputFormat] = None, video_data: Optional[MultimodalDataInputFormat] = None,
dimensions: Optional[int] = None, dimensions: Optional[int] = None,
lora_path: Optional[Union[List[Optional[str]], Optional[str]]] = 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, external_trace_header: Optional[Dict] = None,
rid: Optional[Union[List[str], str]] = None, rid: Optional[Union[List[str], str]] = None,
) -> Dict: ) -> Dict:
@@ -491,6 +497,8 @@ class Engine(EngineScoreMixin, EngineBase):
video_data=video_data, video_data=video_data,
dimensions=dimensions, dimensions=dimensions,
lora_path=lora_path, lora_path=lora_path,
embed_override_token_id=embed_override_token_id,
embed_overrides=embed_overrides,
external_trace_header=external_trace_header, external_trace_header=external_trace_header,
rid=rid, rid=rid,
) )
@@ -20,6 +20,8 @@ by TokenizerManagerScoreMixin.
from typing import List, Optional, Union from typing import List, Optional, Union
import torch
from sglang.srt.managers.tokenizer_manager_score_mixin import ScoreResult from sglang.srt.managers.tokenizer_manager_score_mixin import ScoreResult
@@ -31,6 +33,12 @@ class EngineScoreMixin:
label_token_ids: Optional[List[int]] = None, label_token_ids: Optional[List[int]] = None,
apply_softmax: bool = False, apply_softmax: bool = False,
item_first: 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: ) -> ScoreResult:
""" """
Score items against a query using the loaded model. Score items against a query using the loaded model.
@@ -52,6 +60,9 @@ class EngineScoreMixin:
SequenceClassification). SequenceClassification).
apply_softmax: Whether to normalize scores using softmax. apply_softmax: Whether to normalize scores using softmax.
item_first: If True, prepend items before query (single-item mode only). 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: Returns:
ScoreResult with scores (one list per item) and prompt token count. ScoreResult with scores (one list per item) and prompt token count.
@@ -63,6 +74,9 @@ class EngineScoreMixin:
label_token_ids=label_token_ids, label_token_ids=label_token_ids,
apply_softmax=apply_softmax, apply_softmax=apply_softmax,
item_first=item_first, 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, request=None,
) )
) )
@@ -74,6 +88,9 @@ class EngineScoreMixin:
label_token_ids: Optional[List[int]] = None, label_token_ids: Optional[List[int]] = None,
apply_softmax: bool = False, apply_softmax: bool = False,
item_first: 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: ) -> ScoreResult:
"""Asynchronous version of score(). See score() for full documentation.""" """Asynchronous version of score(). See score() for full documentation."""
return await self.tokenizer_manager.score_request( return await self.tokenizer_manager.score_request(
@@ -82,5 +99,8 @@ class EngineScoreMixin:
label_token_ids=label_token_ids, label_token_ids=label_token_ids,
apply_softmax=apply_softmax, apply_softmax=apply_softmax,
item_first=item_first, 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, request=None,
) )
@@ -942,6 +942,11 @@ class EmbeddingRequest(BaseModel):
priority: Optional[int] = None priority: Optional[int] = None
# LoRA adapter path(s) # LoRA adapter path(s)
lora_path: Optional[Union[List[Optional[str]], Optional[str]]] = None 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): class EmbeddingObject(BaseModel):
@@ -995,6 +1000,16 @@ class ScoringRequest(BaseModel):
items: Optional[Union[str, List[str], List[List[int]]]] = ( items: Optional[Union[str, List[str], List[List[int]]]] = (
None # Item text(s) or pre-tokenized token IDs 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]] = ( label_token_ids: Optional[List[int]] = (
None # Token IDs to compute probabilities for None # Token IDs to compute probabilities for
) )
@@ -14,6 +14,7 @@ from sglang.srt.entrypoints.openai.protocol import (
UsageInfo, UsageInfo,
) )
from sglang.srt.entrypoints.openai.serving_base import OpenAIServingBase 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.managers.io_struct import EmbeddingReqInput
from sglang.srt.parser.conversation import generate_embedding_convs from sglang.srt.parser.conversation import generate_embedding_convs
@@ -59,14 +60,14 @@ class OpenAIServingEmbedding(OpenAIServingBase):
# List of strings # List of strings
for i, item in enumerate(input): for i, item in enumerate(input):
if not isinstance(item, str): 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(): if not item.strip():
return f"Input at index {i} cannot be empty or whitespace only" return f"Input at index {i} cannot be empty or whitespace only"
elif isinstance(first_item, int): elif isinstance(first_item, int):
# List of integers (token IDs) # List of integers (token IDs)
for i, item in enumerate(input): for i, item in enumerate(input):
if not isinstance(item, int): 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: if item < 0:
return f"Token ID at index {i} must be non-negative" return f"Token ID at index {i} must be non-negative"
return None return None
@@ -129,6 +130,26 @@ class OpenAIServingEmbedding(OpenAIServingBase):
# Resolve LoRA adapter from model parameter or explicit lora_path # Resolve LoRA adapter from model parameter or explicit lora_path
lora_path = self._resolve_lora_path(request.model, request.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( adapted_request = EmbeddingReqInput(
**prompt_kwargs, **prompt_kwargs,
rid=request.rid, rid=request.rid,
@@ -136,6 +157,8 @@ class OpenAIServingEmbedding(OpenAIServingBase):
routing_key=self.extract_routing_key(raw_request), routing_key=self.extract_routing_key(raw_request),
dimensions=request.dimensions, dimensions=request.dimensions,
lora_path=lora_path, lora_path=lora_path,
embed_override_token_id=request.embed_override_token_id,
embed_overrides=embed_overrides,
) )
return adapted_request, request return adapted_request, request
@@ -1,6 +1,7 @@
import logging import logging
from typing import Union from typing import Union
import torch
from fastapi import Request from fastapi import Request
from sglang.srt.entrypoints.openai.protocol import ( from sglang.srt.entrypoints.openai.protocol import (
@@ -42,13 +43,38 @@ class OpenAIServingScore(OpenAIServingBase):
) -> Union[ScoringResponse, ErrorResponse]: ) -> Union[ScoringResponse, ErrorResponse]:
"""Handle the scoring request""" """Handle the scoring request"""
try: 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( result = await self.tokenizer_manager.score_request(
query=request.query, query=request.query,
items=request.items, items=request.items,
label_token_ids=request.label_token_ids, label_token_ids=request.label_token_ids,
apply_softmax=request.apply_softmax, apply_softmax=request.apply_softmax,
item_first=request.item_first, 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, request=raw_request,
) )
@@ -1,6 +1,8 @@
import logging import logging
from typing import Any, Dict, List, Optional, Union from typing import Any, Dict, List, Optional, Union
import torch
from sglang.srt.entrypoints.openai.protocol import ( from sglang.srt.entrypoints.openai.protocol import (
CachedTokensDetails, CachedTokensDetails,
ChatCompletionRequest, ChatCompletionRequest,
@@ -132,3 +134,41 @@ def process_cached_tokens_details_from_ret(
device=details.get("device", 0), device=details.get("device", 0),
host=details.get("host", 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 import torch
from sglang.srt.lora.lora_registry import LoRARef 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.managers.schedule_batch import BaseFinishReason, Modality
from sglang.srt.multimodal.mm_utils import has_valid_data from sglang.srt.multimodal.mm_utils import has_valid_data
from sglang.srt.observability.req_time_stats import ( 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 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. # 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 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. # The image input. It can be an image instance, file name, URL, or base64 encoded string.
# Can be formatted as: # Can be formatted as:
# - Single image for a single request # - 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.") 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): def __getitem__(self, i):
# Cache sub-objects so that repeated obj[i] calls return the same instance. # Cache sub-objects so that repeated obj[i] calls return the same instance.
# This avoids subtle bugs where different call sites get divergent objects. # This avoids subtle bugs where different call sites get divergent objects.
@@ -612,6 +627,7 @@ class GenerateReqInput(BaseReq):
input_embeds=( input_embeds=(
self.input_embeds[i] if self.input_embeds is not None else None 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], image_data=self.image_data[i],
video_data=self.video_data[i], video_data=self.video_data[i],
audio_data=self.audio_data[i], audio_data=self.audio_data[i],
@@ -706,6 +722,9 @@ class TokenizedGenerateReqInput(BaseReq):
# The input embeds # The input embeds
input_embeds: Optional[Union[List[List[List[float]]], List[List[float]]]] = None 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 info for continual prompting
session_params: Optional[SessionParams] = None session_params: Optional[SessionParams] = None
@@ -791,6 +810,18 @@ class EmbeddingReqInput(BaseReq):
audio_data: Optional[MultimodalDataInputFormat] = None audio_data: Optional[MultimodalDataInputFormat] = None
# The token ids for text; one can either specify text or input_ids. # The token ids for text; one can either specify text or input_ids.
input_ids: Optional[Union[List[List[int]], List[int]]] = None 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 # Dummy sampling params for compatibility
sampling_params: Optional[Union[List[Dict], Dict]] = None sampling_params: Optional[Union[List[Dict], Dict]] = None
# Dummy input embeds for compatibility # Dummy input embeds for compatibility
@@ -896,6 +927,16 @@ class EmbeddingReqInput(BaseReq):
or has_valid_data(self.audio_data) 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): def __getitem__(self, i):
# Cache sub-objects so that repeated obj[i] calls return the same instance. # Cache sub-objects so that repeated obj[i] calls return the same instance.
cache = self.__dict__.setdefault("_sub_obj_cache", {}) cache = self.__dict__.setdefault("_sub_obj_cache", {})
@@ -905,6 +946,7 @@ class EmbeddingReqInput(BaseReq):
if self.is_cross_encoder_request: if self.is_cross_encoder_request:
sub = EmbeddingReqInput( sub = EmbeddingReqInput(
text=[self.text[i]] if self.text is not None else None, 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], sampling_params=self.sampling_params[i],
rid=self.rid[i], rid=self.rid[i],
lora_path=self.lora_path[i] if self.lora_path is not None else None, lora_path=self.lora_path[i] if self.lora_path is not None else None,
@@ -916,6 +958,13 @@ class EmbeddingReqInput(BaseReq):
sub = EmbeddingReqInput( sub = EmbeddingReqInput(
text=self.text[i] if self.text is not None else None, 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, 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, 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, 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, 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] token_type_ids: List[int]
# Dummy sampling params for compatibility # Dummy sampling params for compatibility
sampling_params: SamplingParams sampling_params: SamplingParams
# Embedding overrides to place at specific token positions.
positional_embed_overrides: Optional[PositionalEmbeds] = None
# For DP routing # For DP routing
routed_dp_rank: Optional[int] = None routed_dp_rank: Optional[int] = None
# Priority for the request # Priority for the request
priority: Optional[int] = None priority: Optional[int] = None
# The number of dimensions the resulting output embeddings should have. It is applicable for Matryoshka Embeddings. # The number of dimensions the resulting output embeddings should have. It is applicable for Matryoshka Embeddings.
dimensions: Optional[int] = None dimensions: Optional[int] = None
# LoRA related # LoRA related
lora_id: Optional[str] = None # None means just use the base model lora_id: Optional[str] = None # None means just use the base model
# For observability # 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.dllm.mixin.req import ReqDllmMixin
from sglang.srt.environ import envs 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.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.allocator import BaseTokenToKVPoolAllocator
from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache, MatchPrefixParams from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache, MatchPrefixParams
from sglang.srt.mem_cache.common import ( from sglang.srt.mem_cache.common import (
@@ -571,6 +572,7 @@ class Req(ReqDllmMixin):
origin_input_ids_unpadded: Optional[Tuple[int]] = None, origin_input_ids_unpadded: Optional[Tuple[int]] = None,
lora_id: Optional[str] = None, lora_id: Optional[str] = None,
input_embeds: Optional[List[List[float]]] = None, input_embeds: Optional[List[List[float]]] = None,
positional_embed_overrides: Optional[PositionalEmbeds] = None,
token_type_ids: List[int] = None, token_type_ids: List[int] = None,
session: Optional[Session] = None, session: Optional[Session] = None,
custom_logit_processor: Optional[str] = None, custom_logit_processor: Optional[str] = None,
@@ -610,6 +612,7 @@ class Req(ReqDllmMixin):
self.fill_ids = [] self.fill_ids = []
self.session = session self.session = session
self.input_embeds = input_embeds self.input_embeds = input_embeds
self.positional_embed_overrides = positional_embed_overrides
# For req-level memory management # For req-level memory management
self.kv_committed_len = 0 self.kv_committed_len = 0
@@ -973,6 +976,12 @@ class Req(ReqDllmMixin):
max_prefix_len = max(max_prefix_len, 0) max_prefix_len = max(max_prefix_len, 0)
token_ids = self.fill_ids[:max_prefix_len] 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 tree_cache is not None:
if cow_mamba is None: if cow_mamba is None:
cow_mamba = tree_cache.supports_mamba() cow_mamba = tree_cache.supports_mamba()
@@ -1332,6 +1341,9 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
# Batched arguments to model runner # Batched arguments to model runner
input_ids: torch.Tensor = None # shape: [b], int64 input_ids: torch.Tensor = None # shape: [b], int64
input_embeds: torch.Tensor = None # shape: [b, hidden_size], float32 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 ne_token_table: torch.Tensor = None
token_type_ids: torch.Tensor = None # shape: [b], int64 token_type_ids: torch.Tensor = None # shape: [b], int64
req_pool_indices: torch.Tensor = None # shape: [b], int64 req_pool_indices: torch.Tensor = None # shape: [b], int64
@@ -1618,6 +1630,11 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
# Set fields # Set fields
input_embeds = [] 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 = [] extend_input_logprob_token_ids = []
multimodal_inputs = [] multimodal_inputs = []
mamba_track_mask_cpu = [] mamba_track_mask_cpu = []
@@ -1642,6 +1659,27 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
req.input_embeds[pre_len : pre_len + req.extend_input_len] 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) multimodal_inputs.append(req.multimodal_inputs)
# Only calculate cached_tokens once. Once retracted, the 'retracted_stain' # Only calculate cached_tokens once. Once retracted, the 'retracted_stain'
@@ -1737,6 +1775,17 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
else: else:
extend_input_logprob_token_ids = None 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.input_ids = input_ids_tensor
self.req_pool_indices = req_pool_indices_tensor self.req_pool_indices = req_pool_indices_tensor
self.orig_seq_lens = orig_seq_lens_tensor self.orig_seq_lens = orig_seq_lens_tensor
@@ -1748,6 +1797,8 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
if input_embeds if input_embeds
else None else None
) )
self.replace_embeds = replace_embeds_tensor
self.replace_positions = replace_positions_tensor
self.multimodal_inputs = multimodal_inputs self.multimodal_inputs = multimodal_inputs
self.token_type_ids = token_type_ids_tensor self.token_type_ids = token_type_ids_tensor
self.seq_lens_sum = sum(seq_lens) self.seq_lens_sum = sum(seq_lens)
@@ -2367,6 +2418,8 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
lora_ids=[req.lora_id for req in self.reqs], lora_ids=[req.lora_id for req in self.reqs],
sampling_info=self.sampling_info, sampling_info=self.sampling_info,
input_embeds=self.input_embeds, input_embeds=self.input_embeds,
replace_embeds=self.replace_embeds,
replace_positions=self.replace_positions,
ne_token_table=self.ne_token_table, ne_token_table=self.ne_token_table,
token_type_ids=self.token_type_ids, token_type_ids=self.token_type_ids,
spec_algorithm=self.spec_algorithm, spec_algorithm=self.spec_algorithm,
@@ -2542,6 +2595,8 @@ class ModelWorkerBatch:
# The input Embeds # The input Embeds
input_embeds: Optional[torch.Tensor] = None input_embeds: Optional[torch.Tensor] = None
replace_embeds: Optional[torch.Tensor] = None
replace_positions: Optional[torch.Tensor] = None
# token table for ngram embedding # token table for ngram embedding
ne_token_table: Optional[torch.Tensor] = None ne_token_table: Optional[torch.Tensor] = None
+2
View File
@@ -1803,6 +1803,7 @@ class Scheduler(
stream=recv_req.stream, stream=recv_req.stream,
lora_id=recv_req.lora_id, lora_id=recv_req.lora_id,
input_embeds=recv_req.input_embeds, input_embeds=recv_req.input_embeds,
positional_embed_overrides=recv_req.positional_embed_overrides,
token_type_ids=recv_req.token_type_ids, token_type_ids=recv_req.token_type_ids,
custom_logit_processor=recv_req.custom_logit_processor, custom_logit_processor=recv_req.custom_logit_processor,
require_reasoning=recv_req.require_reasoning, require_reasoning=recv_req.require_reasoning,
@@ -2128,6 +2129,7 @@ class Scheduler(
recv_req.input_text, recv_req.input_text,
recv_req.input_ids, recv_req.input_ids,
recv_req.sampling_params, recv_req.sampling_params,
positional_embed_overrides=recv_req.positional_embed_overrides,
token_type_ids=recv_req.token_type_ids, token_type_ids=recv_req.token_type_ids,
routed_dp_rank=recv_req.routed_dp_rank, routed_dp_rank=recv_req.routed_dp_rank,
priority=recv_req.priority, priority=recv_req.priority,
@@ -33,6 +33,7 @@ from typing import Any, Awaitable, Dict, List, Optional, Tuple, Union
import fastapi import fastapi
import pybase64 import pybase64
import torch
import uvloop import uvloop
import zmq import zmq
import zmq.asyncio 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.lora.lora_registry import LoRARef, LoRARegistry
from sglang.srt.managers.async_dynamic_batch_tokenizer import AsyncDynamicbatchTokenizer from sglang.srt.managers.async_dynamic_batch_tokenizer import AsyncDynamicbatchTokenizer
from sglang.srt.managers.disagg_service import start_disagg_service 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 ( from sglang.srt.managers.io_struct import (
AbortReq, AbortReq,
ActiveRanksOutput, ActiveRanksOutput,
@@ -1000,6 +1002,7 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerScoreMixin):
bootstrap_room=obj.bootstrap_room, bootstrap_room=obj.bootstrap_room,
lora_id=obj.lora_id, lora_id=obj.lora_id,
input_embeds=input_embeds, input_embeds=input_embeds,
positional_embed_overrides=obj.positional_embed_overrides,
session_params=session_params, session_params=session_params,
custom_logit_processor=obj.custom_logit_processor, custom_logit_processor=obj.custom_logit_processor,
require_reasoning=obj.require_reasoning, require_reasoning=obj.require_reasoning,
@@ -1015,12 +1018,24 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerScoreMixin):
num_items_assigned=obj.num_items_assigned, num_items_assigned=obj.num_items_assigned,
) )
elif isinstance(obj, EmbeddingReqInput): 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( tokenized_obj = TokenizedEmbeddingReqInput(
input_text, input_text,
input_ids, input_ids,
mm_inputs, mm_inputs,
token_type_ids, token_type_ids,
sampling_params, sampling_params,
positional_embed_overrides=positional_embed_overrides,
rid=obj.rid, rid=obj.rid,
priority=obj.priority, priority=obj.priority,
dimensions=obj.dimensions, dimensions=obj.dimensions,
@@ -1033,6 +1048,26 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerScoreMixin):
return tokenized_obj 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( async def _batch_tokenize_and_process(
self, batch_size: int, obj: Union[GenerateReqInput, EmbeddingReqInput] self, batch_size: int, obj: Union[GenerateReqInput, EmbeddingReqInput]
) -> List[Union[TokenizedGenerateReqInput, TokenizedEmbeddingReqInput]]: ) -> List[Union[TokenizedGenerateReqInput, TokenizedEmbeddingReqInput]]:
@@ -5,6 +5,7 @@ from typing import Any, Dict, List, Optional, Tuple, Union
import torch import torch
from sglang.srt.managers.embed_types import PositionalEmbeds
from sglang.srt.managers.io_struct import EmbeddingReqInput, GenerateReqInput from sglang.srt.managers.io_struct import EmbeddingReqInput, GenerateReqInput
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -294,6 +295,181 @@ class TokenizerManagerScoreMixin:
return ScoreResult(scores=scores, prompt_tokens=prompt_tokens) 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( async def score_request(
self, self,
query: Optional[Union[str, List[int]]] = None, query: Optional[Union[str, List[int]]] = None,
@@ -301,6 +477,9 @@ class TokenizerManagerScoreMixin:
label_token_ids: Optional[List[int]] = None, label_token_ids: Optional[List[int]] = None,
apply_softmax: bool = False, apply_softmax: bool = False,
item_first: 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, request: Optional[Any] = None,
) -> ScoreResult: ) -> ScoreResult:
""" """
@@ -327,6 +506,9 @@ class TokenizerManagerScoreMixin:
label_token_ids: List of token IDs to compute probabilities for label_token_ids: List of token IDs to compute probabilities for
apply_softmax: Whether to normalize probabilities using softmax apply_softmax: Whether to normalize probabilities using softmax
item_first: If True, prepend items to query. Ignored for multi-item scoring. 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 request: Optional FastAPI request object
Returns: Returns:
@@ -340,12 +522,26 @@ class TokenizerManagerScoreMixin:
raise ValueError( raise ValueError(
"label_token_ids is required for generation (CausalLM) models." "label_token_ids is required for generation (CausalLM) models."
) )
if items is None: if items is None:
raise ValueError("items must be provided") raise ValueError("items must be provided")
if not items: if not items:
return ScoreResult(scores=[], prompt_tokens=0) 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: if self.tokenizer is not None and label_token_ids is not None:
vocab_size = self.tokenizer.vocab_size vocab_size = self.tokenizer.vocab_size
for token_id in label_token_ids: for token_id in label_token_ids:
@@ -362,15 +558,13 @@ class TokenizerManagerScoreMixin:
input_ids = None input_ids = None
text_prompts = None text_prompts = None
positional_embed_overrides = None
# Handle string or tokenized query/items use_text_prompts = isinstance(query, str) and not has_embeds
if isinstance(query, str) and (
isinstance(items, str) if use_text_prompts:
or (isinstance(items, list) and (not items or isinstance(items[0], str)))
):
# Both query and items are text # Both query and items are text
items_list = [items] if isinstance(items, str) else items items_list = [items] if isinstance(items, str) else items
if use_multi_item_scoring: if use_multi_item_scoring:
# Multi-item scoring: tokenize separately then combine at token level # Multi-item scoring: tokenize separately then combine at token level
# to ensure the delimiter token ID is inserted exactly once per boundary # to ensure the delimiter token ID is inserted exactly once per boundary
@@ -396,21 +590,29 @@ class TokenizerManagerScoreMixin:
and items and items
and isinstance(items[0], list) and isinstance(items[0], list)
): ):
# Both query and items are token IDs # Both query and items are token IDs — tokenize text inputs if needed for embed overrides
if use_multi_item_scoring: query_ids, items_ids = query, items
# Multi-item scoring: concatenate with delimiter token ID _, input_ids, positional_embed_overrides = self._build_token_id_inputs(
# Format: query<delimiter_token_id>item1<delimiter_token_id>item2<delimiter_token_id>item3<delimiter_token_id> query_ids,
delimiter_token_id = self.server_args.multi_item_scoring_delimiter items_ids,
combined_input_ids = self._build_multi_item_token_sequence( item_first,
query, items, delimiter_token_id use_multi_item_scoring,
embed_override_token_id,
query_embed_overrides,
item_embed_overrides,
)
elif has_embeds:
# Text inputs with embed overrides — need to tokenize first to resolve positions
query_ids, items_ids = self._batch_tokenize_query_and_items(query, items)
_, input_ids, positional_embed_overrides = self._build_token_id_inputs(
query_ids,
items_ids,
item_first,
use_multi_item_scoring,
embed_override_token_id,
query_embed_overrides,
item_embed_overrides,
) )
input_ids = [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]
else: else:
raise ValueError( raise ValueError(
"Invalid combination of query/items types for score_request." "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, logprob_start_len=0 if use_multi_item_scoring else -1,
stream=False, stream=False,
sampling_params={"max_new_tokens": 0}, sampling_params={"max_new_tokens": 0},
positional_embed_overrides=positional_embed_overrides,
) )
else: else:
batch_request = EmbeddingReqInput( batch_request = EmbeddingReqInput(
text=text_prompts, text=text_prompts,
input_ids=input_ids, input_ids=input_ids,
positional_embed_overrides=positional_embed_overrides,
) )
results = await self.generate_request(batch_request, request).__anext__() results = await self.generate_request(batch_request, request).__anext__()
@@ -675,6 +675,9 @@ class CudaGraphRunner:
return torch.int64 return torch.int64
def can_run(self, forward_batch: ForwardBatch): 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: if self.require_mlp_tp_gather:
cuda_graph_bs = ( cuda_graph_bs = (
max(forward_batch.global_num_tokens_cpu) // self.num_tokens_per_bs max(forward_batch.global_num_tokens_cpu) // self.num_tokens_per_bs
@@ -359,6 +359,10 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
# For input embeddings # For input embeddings
input_embeds: Optional[torch.Tensor] = None 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 # For cross-encoder model
token_type_ids: Optional[torch.Tensor] = None token_type_ids: Optional[torch.Tensor] = None
@@ -473,6 +477,8 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
spec_info=batch.spec_info, spec_info=batch.spec_info,
capture_hidden_mode=batch.capture_hidden_mode, capture_hidden_mode=batch.capture_hidden_mode,
input_embeds=batch.input_embeds, input_embeds=batch.input_embeds,
replace_embeds=batch.replace_embeds,
replace_positions=batch.replace_positions,
token_type_ids=batch.token_type_ids, token_type_ids=batch.token_type_ids,
tbo_split_seq_index=batch.tbo_split_seq_index, tbo_split_seq_index=batch.tbo_split_seq_index,
dimensions=batch.dimensions, dimensions=batch.dimensions,
@@ -2738,6 +2738,17 @@ class ModelRunner(ModelRunnerKVCacheMixin):
kwargs["pp_proxy_tensors"] = pp_proxy_tensors kwargs["pp_proxy_tensors"] = pp_proxy_tensors
if forward_batch.input_embeds is not None: if forward_batch.input_embeds is not None:
kwargs["input_embeds"] = forward_batch.input_embeds.bfloat16() 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: if not self.is_generation:
kwargs["get_embedding"] = True kwargs["get_embedding"] = True
@@ -417,6 +417,9 @@ class PiecewiseCudaGraphRunner:
# TODO(yuwei): fix it # TODO(yuwei): fix it
if forward_batch.input_embeds is not None: if forward_batch.input_embeds is not None:
return False 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) num_tokens = len(forward_batch.input_ids)
if forward_batch.return_logprob: if forward_batch.return_logprob:
for start_len, seq_len in zip( 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()