[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
+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(