[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:
co-authored by
Claude Opus 4.6
parent
a69be2e866
commit
6838a23226
@@ -510,6 +510,13 @@
|
||||
"cooldown_interval_minutes": 60,
|
||||
"reason": "custom override"
|
||||
},
|
||||
"fortunecookiee": {
|
||||
"can_tag_run_ci_label": true,
|
||||
"can_rerun_failed_ci": true,
|
||||
"can_rerun_stage": true,
|
||||
"cooldown_interval_minutes": 0,
|
||||
"reason": "custom override"
|
||||
},
|
||||
"fy1214": {
|
||||
"can_tag_run_ci_label": true,
|
||||
"can_rerun_failed_ci": true,
|
||||
|
||||
@@ -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
|
||||
]
|
||||
|
||||
@@ -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)})"
|
||||
)
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
# 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,
|
||||
)
|
||||
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:
|
||||
raise ValueError(
|
||||
"Invalid combination of query/items types for score_request."
|
||||
@@ -427,11 +629,13 @@ class TokenizerManagerScoreMixin:
|
||||
logprob_start_len=0 if use_multi_item_scoring else -1,
|
||||
stream=False,
|
||||
sampling_params={"max_new_tokens": 0},
|
||||
positional_embed_overrides=positional_embed_overrides,
|
||||
)
|
||||
else:
|
||||
batch_request = EmbeddingReqInput(
|
||||
text=text_prompts,
|
||||
input_ids=input_ids,
|
||||
positional_embed_overrides=positional_embed_overrides,
|
||||
)
|
||||
|
||||
results = await self.generate_request(batch_request, request).__anext__()
|
||||
|
||||
@@ -675,6 +675,9 @@ class CudaGraphRunner:
|
||||
return torch.int64
|
||||
|
||||
def can_run(self, forward_batch: ForwardBatch):
|
||||
# Disable for token embedding overrides (dynamic per-request)
|
||||
if forward_batch.replace_embeds is not None:
|
||||
return False
|
||||
if self.require_mlp_tp_gather:
|
||||
cuda_graph_bs = (
|
||||
max(forward_batch.global_num_tokens_cpu) // self.num_tokens_per_bs
|
||||
|
||||
@@ -359,6 +359,10 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
|
||||
# For input embeddings
|
||||
input_embeds: Optional[torch.Tensor] = None
|
||||
|
||||
# For token embedding overrides (sparse replacement at specific positions)
|
||||
replace_embeds: Optional[torch.Tensor] = None
|
||||
replace_positions: Optional[torch.Tensor] = None
|
||||
|
||||
# For cross-encoder model
|
||||
token_type_ids: Optional[torch.Tensor] = None
|
||||
|
||||
@@ -473,6 +477,8 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
|
||||
spec_info=batch.spec_info,
|
||||
capture_hidden_mode=batch.capture_hidden_mode,
|
||||
input_embeds=batch.input_embeds,
|
||||
replace_embeds=batch.replace_embeds,
|
||||
replace_positions=batch.replace_positions,
|
||||
token_type_ids=batch.token_type_ids,
|
||||
tbo_split_seq_index=batch.tbo_split_seq_index,
|
||||
dimensions=batch.dimensions,
|
||||
|
||||
@@ -2738,6 +2738,17 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
||||
kwargs["pp_proxy_tensors"] = pp_proxy_tensors
|
||||
if forward_batch.input_embeds is not None:
|
||||
kwargs["input_embeds"] = forward_batch.input_embeds.bfloat16()
|
||||
if (
|
||||
forward_batch.replace_embeds is not None
|
||||
and forward_batch.replace_positions is not None
|
||||
):
|
||||
# Token embedding overrides: get base embeddings, scatter replacements
|
||||
if "input_embeds" not in kwargs:
|
||||
embed_layer = self.model.get_input_embeddings()
|
||||
kwargs["input_embeds"] = embed_layer(forward_batch.input_ids)
|
||||
kwargs["input_embeds"][forward_batch.replace_positions] = (
|
||||
forward_batch.replace_embeds.to(kwargs["input_embeds"].dtype)
|
||||
)
|
||||
if not self.is_generation:
|
||||
kwargs["get_embedding"] = True
|
||||
|
||||
|
||||
@@ -417,6 +417,9 @@ class PiecewiseCudaGraphRunner:
|
||||
# TODO(yuwei): fix it
|
||||
if forward_batch.input_embeds is not None:
|
||||
return False
|
||||
# Disable for token embedding overrides (dynamic per-request)
|
||||
if forward_batch.replace_embeds is not None:
|
||||
return False
|
||||
num_tokens = len(forward_batch.input_ids)
|
||||
if forward_batch.return_logprob:
|
||||
for start_len, seq_len in zip(
|
||||
|
||||
@@ -0,0 +1,602 @@
|
||||
"""Unit tests for token embedding override support.
|
||||
|
||||
Covers:
|
||||
- PositionalEmbeds dataclass (embed_types.py)
|
||||
- convert_embeds_to_tensors (utils.py)
|
||||
- TokenizerManager._resolve_embed_overrides (tokenizer_manager.py)
|
||||
- positional_embed_overrides on GenerateReqInput/EmbeddingReqInput (io_struct.py)
|
||||
- Score mixin override resolution (tokenizer_manager_score_mixin.py)
|
||||
"""
|
||||
|
||||
import unittest
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.srt.entrypoints.openai.utils import convert_embeds_to_tensors
|
||||
from sglang.srt.managers.embed_types import PositionalEmbeds
|
||||
from sglang.srt.managers.io_struct import EmbeddingReqInput, GenerateReqInput
|
||||
from sglang.srt.managers.tokenizer_manager import TokenizerManager
|
||||
from sglang.srt.managers.tokenizer_manager_score_mixin import (
|
||||
TokenizerManagerScoreMixin,
|
||||
)
|
||||
from sglang.test.ci.ci_register import register_cuda_ci
|
||||
from sglang.test.test_utils import CustomTestCase
|
||||
|
||||
register_cuda_ci(est_time=5, suite="stage-b-test-1-gpu-small")
|
||||
|
||||
HIDDEN_DIM = 4
|
||||
|
||||
|
||||
def _vec(val: float = 1.0) -> torch.Tensor:
|
||||
"""Create a 1-D tensor of size HIDDEN_DIM."""
|
||||
return torch.full((HIDDEN_DIM,), val, dtype=torch.float32)
|
||||
|
||||
|
||||
def _vec2d(val: float = 1.0) -> torch.Tensor:
|
||||
"""Create a [1, HIDDEN_DIM] tensor."""
|
||||
return torch.full((1, HIDDEN_DIM), val, dtype=torch.float32)
|
||||
|
||||
|
||||
# ========================================================================
|
||||
# PositionalEmbeds
|
||||
# ========================================================================
|
||||
|
||||
|
||||
class TestPositionalEmbeds(CustomTestCase):
|
||||
def test_from_list_of_1d_tensors(self):
|
||||
pe = PositionalEmbeds(embeds=[_vec(1), _vec(2)], positions=[0, 5])
|
||||
self.assertEqual(pe.embeds.shape, (2, HIDDEN_DIM))
|
||||
self.assertAlmostEqual(pe.embeds[0, 0].item(), 1.0)
|
||||
self.assertAlmostEqual(pe.embeds[1, 0].item(), 2.0)
|
||||
|
||||
def test_from_list_of_2d_tensors(self):
|
||||
pe = PositionalEmbeds(embeds=[_vec2d(3), _vec2d(4)], positions=[1, 2])
|
||||
self.assertEqual(pe.embeds.shape, (2, HIDDEN_DIM))
|
||||
|
||||
def test_from_pre_stacked_tensor(self):
|
||||
stacked = torch.zeros(3, HIDDEN_DIM)
|
||||
pe = PositionalEmbeds(embeds=stacked, positions=[0, 1, 2])
|
||||
self.assertIs(pe.embeds, stacked)
|
||||
|
||||
def test_length_mismatch_raises(self):
|
||||
with self.assertRaises(ValueError):
|
||||
PositionalEmbeds(embeds=[_vec()], positions=[0, 1])
|
||||
|
||||
def test_empty(self):
|
||||
pe = PositionalEmbeds(embeds=torch.zeros(0, HIDDEN_DIM), positions=[])
|
||||
self.assertEqual(pe.embeds.shape[0], 0)
|
||||
|
||||
|
||||
# ========================================================================
|
||||
# convert_embeds_to_tensors
|
||||
# ========================================================================
|
||||
|
||||
|
||||
class TestConvertEmbedsToTensors(CustomTestCase):
|
||||
def test_none_returns_none(self):
|
||||
self.assertIsNone(convert_embeds_to_tensors(None))
|
||||
|
||||
def test_empty_list(self):
|
||||
self.assertEqual(convert_embeds_to_tensors([]), [])
|
||||
|
||||
def test_single_input(self):
|
||||
"""[num_replacements][hidden_size] -> [[tensor, ...]]"""
|
||||
result = convert_embeds_to_tensors([[1.0, 2.0], [3.0, 4.0]])
|
||||
self.assertEqual(len(result), 1) # wrapped in outer list
|
||||
self.assertEqual(len(result[0]), 2) # two replacement vectors
|
||||
self.assertTrue(torch.is_tensor(result[0][0]))
|
||||
self.assertEqual(result[0][0].tolist(), [1.0, 2.0])
|
||||
self.assertEqual(result[0][0].dtype, torch.float32)
|
||||
self.assertEqual(result[0][0].dim(), 1) # each vector is 1-D
|
||||
|
||||
def test_batch_input(self):
|
||||
"""[num_inputs][num_replacements][hidden_size] -> [[tensor, ...], ...]"""
|
||||
result = convert_embeds_to_tensors(
|
||||
[
|
||||
[[1.0, 2.0]],
|
||||
[[3.0, 4.0], [5.0, 6.0]],
|
||||
]
|
||||
)
|
||||
self.assertEqual(len(result), 2)
|
||||
self.assertEqual(len(result[0]), 1)
|
||||
self.assertEqual(len(result[1]), 2)
|
||||
|
||||
|
||||
# ========================================================================
|
||||
# TokenizerManager._resolve_embed_overrides
|
||||
# ========================================================================
|
||||
|
||||
|
||||
class TestResolveEmbedOverrides(CustomTestCase):
|
||||
def test_basic_resolution(self):
|
||||
embeds = [_vec(1), _vec(2)]
|
||||
pe = TokenizerManager._resolve_embed_overrides(
|
||||
input_ids=[10, 50, 20, 50, 30],
|
||||
token_id=50,
|
||||
embeds=embeds,
|
||||
)
|
||||
self.assertIsInstance(pe, PositionalEmbeds)
|
||||
self.assertEqual(pe.positions, [1, 3])
|
||||
self.assertEqual(pe.embeds.shape, (2, HIDDEN_DIM))
|
||||
|
||||
def test_no_placeholders_raises(self):
|
||||
with self.assertRaises(ValueError):
|
||||
TokenizerManager._resolve_embed_overrides(
|
||||
input_ids=[10, 20, 30],
|
||||
token_id=50,
|
||||
embeds=[_vec()],
|
||||
)
|
||||
|
||||
def test_count_mismatch_raises(self):
|
||||
with self.assertRaises(ValueError):
|
||||
TokenizerManager._resolve_embed_overrides(
|
||||
input_ids=[10, 50, 20],
|
||||
token_id=50,
|
||||
embeds=[_vec(1), _vec(2)],
|
||||
)
|
||||
|
||||
|
||||
# ========================================================================
|
||||
# io_struct: positional_embed_overrides on GenerateReqInput
|
||||
# ========================================================================
|
||||
|
||||
|
||||
class TestGenerateReqInputEmbedOverride(CustomTestCase):
|
||||
def test_single_override_in_getitem(self):
|
||||
"""Single PositionalEmbeds is shared across all items in __getitem__."""
|
||||
pe = PositionalEmbeds(embeds=[_vec()], positions=[0])
|
||||
req = GenerateReqInput(
|
||||
input_ids=[[1, 2], [3, 4]],
|
||||
sampling_params=[{}, {}],
|
||||
positional_embed_overrides=pe,
|
||||
)
|
||||
req.normalize_batch_and_arguments()
|
||||
item = req[0]
|
||||
self.assertIs(item.positional_embed_overrides, pe)
|
||||
|
||||
def test_batch_override_in_getitem(self):
|
||||
"""List[Optional[PositionalEmbeds]] is indexed per-item."""
|
||||
pe0 = PositionalEmbeds(embeds=[_vec(1)], positions=[0])
|
||||
pe1 = None
|
||||
req = GenerateReqInput(
|
||||
input_ids=[[1, 2], [3, 4]],
|
||||
sampling_params=[{}, {}],
|
||||
positional_embed_overrides=[pe0, pe1],
|
||||
)
|
||||
req.normalize_batch_and_arguments()
|
||||
self.assertEqual(req[0].positional_embed_overrides, pe0)
|
||||
self.assertIsNone(req[1].positional_embed_overrides)
|
||||
|
||||
|
||||
# ========================================================================
|
||||
# io_struct: embed override fields on EmbeddingReqInput
|
||||
# ========================================================================
|
||||
|
||||
|
||||
class TestEmbeddingReqInputEmbedOverride(CustomTestCase):
|
||||
def test_override_fields_in_getitem(self):
|
||||
"""embed_override_token_id, embed_overrides, and positional_embed_overrides
|
||||
are correctly sliced in __getitem__."""
|
||||
pe0 = PositionalEmbeds(embeds=[_vec(1)], positions=[0])
|
||||
pe1 = PositionalEmbeds(embeds=[_vec(2)], positions=[1])
|
||||
req = EmbeddingReqInput(
|
||||
input_ids=[[50, 10], [20, 50]],
|
||||
sampling_params=[{}, {}],
|
||||
embed_override_token_id=50,
|
||||
embed_overrides=[[_vec(1)], [_vec(2)]],
|
||||
positional_embed_overrides=[pe0, pe1],
|
||||
)
|
||||
req.normalize_batch_and_arguments()
|
||||
item0 = req[0]
|
||||
item1 = req[1]
|
||||
self.assertEqual(item0.embed_override_token_id, 50)
|
||||
self.assertEqual(len(item0.embed_overrides), 1)
|
||||
self.assertEqual(item0.positional_embed_overrides, pe0)
|
||||
self.assertEqual(item1.positional_embed_overrides, pe1)
|
||||
|
||||
|
||||
# ========================================================================
|
||||
# Score mixin: _resolve_overrides_for_sequence
|
||||
# ========================================================================
|
||||
|
||||
|
||||
class _FakeServerArgs:
|
||||
"""Minimal stub for server_args."""
|
||||
|
||||
def __init__(self, multi_item_scoring_delimiter=None):
|
||||
self.multi_item_scoring_delimiter = multi_item_scoring_delimiter
|
||||
|
||||
|
||||
class _FakeMixin(TokenizerManagerScoreMixin):
|
||||
"""Minimal stub to call mixin methods without a full TokenizerManager."""
|
||||
|
||||
def __init__(self, delimiter=None):
|
||||
self.server_args = _FakeServerArgs(delimiter)
|
||||
self.multi_item_delimiter_text = None
|
||||
self.tokenizer = None
|
||||
self.is_generation = True
|
||||
|
||||
|
||||
class TestResolveOverridesForSequence(CustomTestCase):
|
||||
def setUp(self):
|
||||
self.mixin = _FakeMixin()
|
||||
|
||||
def test_none_embeds_returns_empty(self):
|
||||
embeds, positions = self.mixin._resolve_overrides_for_sequence(
|
||||
token_ids=[10, 50, 20],
|
||||
embeds=None,
|
||||
embed_override_token_id=50,
|
||||
)
|
||||
self.assertEqual(embeds, [])
|
||||
self.assertEqual(positions, [])
|
||||
|
||||
def test_basic_resolution(self):
|
||||
e1, e2 = _vec(1), _vec(2)
|
||||
embeds, positions = self.mixin._resolve_overrides_for_sequence(
|
||||
token_ids=[50, 10, 50],
|
||||
embeds=[e1, e2],
|
||||
embed_override_token_id=50,
|
||||
)
|
||||
self.assertEqual(len(embeds), 2)
|
||||
self.assertEqual(positions, [0, 2])
|
||||
|
||||
def test_with_offset(self):
|
||||
embeds, positions = self.mixin._resolve_overrides_for_sequence(
|
||||
token_ids=[10, 50],
|
||||
embeds=[_vec()],
|
||||
embed_override_token_id=50,
|
||||
position_offset=100,
|
||||
)
|
||||
self.assertEqual(positions, [101])
|
||||
|
||||
def test_empty_embeds_list(self):
|
||||
"""Empty embeds list with no placeholders succeeds."""
|
||||
embeds, positions = self.mixin._resolve_overrides_for_sequence(
|
||||
token_ids=[10, 20],
|
||||
embeds=[],
|
||||
embed_override_token_id=50,
|
||||
)
|
||||
self.assertEqual(embeds, [])
|
||||
self.assertEqual(positions, [])
|
||||
|
||||
def test_count_mismatch_raises(self):
|
||||
with self.assertRaises(ValueError):
|
||||
self.mixin._resolve_overrides_for_sequence(
|
||||
token_ids=[50, 50],
|
||||
embeds=[_vec()],
|
||||
embed_override_token_id=50,
|
||||
)
|
||||
|
||||
|
||||
# ========================================================================
|
||||
# Score mixin: _resolve_embed_overrides_for_request
|
||||
# ========================================================================
|
||||
|
||||
|
||||
class TestResolveEmbedOverridesForRequest(CustomTestCase):
|
||||
def setUp(self):
|
||||
self.mixin = _FakeMixin()
|
||||
|
||||
def test_no_overrides_returns_none(self):
|
||||
result = self.mixin._resolve_embed_overrides_for_request(
|
||||
query=[10, 20],
|
||||
item=[30, 40],
|
||||
embed_override_token_id=50,
|
||||
query_embed_overrides=None,
|
||||
item_embeds=None,
|
||||
item_position_offset=2,
|
||||
item_label="items[0]",
|
||||
)
|
||||
self.assertIsNone(result)
|
||||
|
||||
def test_query_only_overrides(self):
|
||||
pe = self.mixin._resolve_embed_overrides_for_request(
|
||||
query=[50, 20],
|
||||
item=[30, 40],
|
||||
embed_override_token_id=50,
|
||||
query_embed_overrides=[_vec(1)],
|
||||
item_embeds=None,
|
||||
item_position_offset=2,
|
||||
item_label="items[0]",
|
||||
)
|
||||
self.assertIsInstance(pe, PositionalEmbeds)
|
||||
self.assertEqual(pe.positions, [0])
|
||||
self.assertEqual(pe.embeds.shape, (1, HIDDEN_DIM))
|
||||
|
||||
def test_item_only_overrides(self):
|
||||
pe = self.mixin._resolve_embed_overrides_for_request(
|
||||
query=[10, 20],
|
||||
item=[50, 40],
|
||||
embed_override_token_id=50,
|
||||
query_embed_overrides=None,
|
||||
item_embeds=[_vec(2)],
|
||||
item_position_offset=2,
|
||||
item_label="items[0]",
|
||||
)
|
||||
self.assertEqual(pe.positions, [2]) # offset applied
|
||||
|
||||
def test_query_and_item_overrides(self):
|
||||
pe = self.mixin._resolve_embed_overrides_for_request(
|
||||
query=[50, 20],
|
||||
item=[30, 50],
|
||||
embed_override_token_id=50,
|
||||
query_embed_overrides=[_vec(1)],
|
||||
item_embeds=[_vec(2)],
|
||||
item_position_offset=2,
|
||||
item_label="items[0]",
|
||||
)
|
||||
self.assertEqual(pe.positions, [0, 3]) # query pos 0, item pos 1+offset 2
|
||||
self.assertEqual(pe.embeds.shape, (2, HIDDEN_DIM))
|
||||
|
||||
|
||||
# ========================================================================
|
||||
# Score mixin: _build_token_id_inputs
|
||||
# ========================================================================
|
||||
|
||||
DELIM_TOKEN = 99
|
||||
|
||||
|
||||
class TestBuildTokenIdInputs(CustomTestCase):
|
||||
def setUp(self):
|
||||
self.mixin = _FakeMixin(delimiter=DELIM_TOKEN)
|
||||
|
||||
# --- single-item mode, no embeds ---
|
||||
|
||||
def test_single_item_no_embeds(self):
|
||||
_, input_ids, injection = self.mixin._build_token_id_inputs(
|
||||
query=[1, 2],
|
||||
items=[[3, 4], [5, 6]],
|
||||
item_first=False,
|
||||
use_multi_item_scoring=False,
|
||||
embed_override_token_id=None,
|
||||
query_embed_overrides=None,
|
||||
item_embed_overrides=None,
|
||||
)
|
||||
self.assertEqual(input_ids, [[1, 2, 3, 4], [1, 2, 5, 6]])
|
||||
self.assertIsNone(injection)
|
||||
|
||||
def test_single_item_no_embeds_item_first(self):
|
||||
_, input_ids, injection = self.mixin._build_token_id_inputs(
|
||||
query=[1, 2],
|
||||
items=[[3, 4]],
|
||||
item_first=True,
|
||||
use_multi_item_scoring=False,
|
||||
embed_override_token_id=None,
|
||||
query_embed_overrides=None,
|
||||
item_embed_overrides=None,
|
||||
)
|
||||
self.assertEqual(input_ids, [[3, 4, 1, 2]])
|
||||
self.assertIsNone(injection)
|
||||
|
||||
# --- multi-item mode, no embeds ---
|
||||
|
||||
def test_multi_item_no_embeds(self):
|
||||
_, input_ids, injection = self.mixin._build_token_id_inputs(
|
||||
query=[1, 2],
|
||||
items=[[3, 4], [5, 6]],
|
||||
item_first=False,
|
||||
use_multi_item_scoring=True,
|
||||
embed_override_token_id=None,
|
||||
query_embed_overrides=None,
|
||||
item_embed_overrides=None,
|
||||
)
|
||||
# query<D>item1<D>item2<D>
|
||||
self.assertEqual(
|
||||
input_ids, [[1, 2, DELIM_TOKEN, 3, 4, DELIM_TOKEN, 5, 6, DELIM_TOKEN]]
|
||||
)
|
||||
self.assertIsNone(injection)
|
||||
|
||||
# --- single-item mode, with embeds ---
|
||||
|
||||
def test_single_item_query_embeds(self):
|
||||
"""Query placeholder overrides are resolved per item."""
|
||||
_, input_ids, injection = self.mixin._build_token_id_inputs(
|
||||
query=[50, 10],
|
||||
items=[[20, 30], [40, 50]],
|
||||
item_first=False,
|
||||
use_multi_item_scoring=False,
|
||||
embed_override_token_id=50,
|
||||
query_embed_overrides=[_vec(1)],
|
||||
item_embed_overrides=None,
|
||||
)
|
||||
self.assertEqual(input_ids, [[50, 10, 20, 30], [50, 10, 40, 50]])
|
||||
self.assertIsNotNone(injection)
|
||||
self.assertEqual(len(injection), 2)
|
||||
# Each item gets its own PositionalEmbeds with query override at pos 0
|
||||
self.assertEqual(injection[0].positions, [0])
|
||||
self.assertEqual(injection[1].positions, [0])
|
||||
|
||||
def test_single_item_item_embeds(self):
|
||||
"""Per-item overrides with correct position offsets."""
|
||||
_, input_ids, injection = self.mixin._build_token_id_inputs(
|
||||
query=[10, 20],
|
||||
items=[[50, 30]],
|
||||
item_first=False,
|
||||
use_multi_item_scoring=False,
|
||||
embed_override_token_id=50,
|
||||
query_embed_overrides=None,
|
||||
item_embed_overrides=[[_vec(2)]],
|
||||
)
|
||||
self.assertEqual(input_ids, [[10, 20, 50, 30]])
|
||||
self.assertIsNotNone(injection)
|
||||
# item placeholder at index 0 of item, offset by query length 2
|
||||
self.assertEqual(injection[0].positions, [2])
|
||||
|
||||
def test_single_item_no_override_positions_returns_none_injection(self):
|
||||
"""When no items have placeholders, injection should be None."""
|
||||
_, input_ids, injection = self.mixin._build_token_id_inputs(
|
||||
query=[10, 20],
|
||||
items=[[30, 40]],
|
||||
item_first=False,
|
||||
use_multi_item_scoring=False,
|
||||
embed_override_token_id=50,
|
||||
query_embed_overrides=None,
|
||||
item_embed_overrides=[None],
|
||||
)
|
||||
self.assertIsNone(injection)
|
||||
|
||||
def test_single_item_query_and_item_embeds(self):
|
||||
"""Single-item mode with both query and item overrides in one request."""
|
||||
_, input_ids, injection = self.mixin._build_token_id_inputs(
|
||||
query=[50, 10],
|
||||
items=[[20, 50]],
|
||||
item_first=False,
|
||||
use_multi_item_scoring=False,
|
||||
embed_override_token_id=50,
|
||||
query_embed_overrides=[_vec(1)],
|
||||
item_embed_overrides=[[_vec(2)]],
|
||||
)
|
||||
self.assertEqual(input_ids, [[50, 10, 20, 50]])
|
||||
self.assertIsNotNone(injection)
|
||||
pe = injection[0]
|
||||
# query override at pos 0, item override at pos 3 (query_len=2 + idx=1)
|
||||
self.assertEqual(pe.positions, [0, 3])
|
||||
self.assertEqual(pe.embeds.shape, (2, HIDDEN_DIM))
|
||||
|
||||
def test_single_item_empty_query(self):
|
||||
"""Empty query with item-only overrides (valid from score_prompts)."""
|
||||
_, input_ids, injection = self.mixin._build_token_id_inputs(
|
||||
query=[],
|
||||
items=[[50, 10]],
|
||||
item_first=False,
|
||||
use_multi_item_scoring=False,
|
||||
embed_override_token_id=50,
|
||||
query_embed_overrides=None,
|
||||
item_embed_overrides=[[_vec(1)]],
|
||||
)
|
||||
self.assertEqual(input_ids, [[50, 10]])
|
||||
self.assertIsNotNone(injection)
|
||||
# item placeholder at absolute pos 0 (offset=len([])=0)
|
||||
self.assertEqual(injection[0].positions, [0])
|
||||
|
||||
# --- multi-item mode, with embeds ---
|
||||
|
||||
def test_multi_item_with_query_and_item_embeds(self):
|
||||
"""Multi-item mode resolves query overrides once and item overrides per item."""
|
||||
_, input_ids, injection = self.mixin._build_token_id_inputs(
|
||||
query=[50, 10],
|
||||
items=[[20, 50], [30, 40]],
|
||||
item_first=False,
|
||||
use_multi_item_scoring=True,
|
||||
embed_override_token_id=50,
|
||||
query_embed_overrides=[_vec(1)],
|
||||
item_embed_overrides=[[_vec(2)], None],
|
||||
)
|
||||
# query<D>item1<D>item2<D> = [50,10, 99, 20,50, 99, 30,40, 99]
|
||||
self.assertEqual(len(input_ids), 1)
|
||||
self.assertIsNotNone(injection)
|
||||
self.assertEqual(
|
||||
len(injection), 1
|
||||
) # single PositionalEmbeds for combined sequence
|
||||
pe = injection[0]
|
||||
# query override at pos 0, item[0] override at pos 4 (query_len=2 + delim=1 + idx=1)
|
||||
self.assertIn(0, pe.positions)
|
||||
self.assertIn(4, pe.positions)
|
||||
self.assertEqual(pe.embeds.shape[0], 2)
|
||||
|
||||
|
||||
# ========================================================================
|
||||
# Score mixin: score_request validation
|
||||
# ========================================================================
|
||||
|
||||
|
||||
class TestScoreRequestValidation(CustomTestCase):
|
||||
"""Test validation guards in score_request without running full pipeline."""
|
||||
|
||||
def setUp(self):
|
||||
self.mixin = _FakeMixin()
|
||||
|
||||
def _call(self, **kwargs):
|
||||
"""Wrapper to call score_request synchronously."""
|
||||
import asyncio
|
||||
|
||||
return asyncio.run(self.mixin.score_request(**kwargs))
|
||||
|
||||
def test_generation_requires_label_token_ids(self):
|
||||
self.mixin.is_generation = True
|
||||
with self.assertRaisesRegex(ValueError, "label_token_ids is required"):
|
||||
self._call(
|
||||
query=[1, 2],
|
||||
items=[[3, 4]],
|
||||
label_token_ids=None,
|
||||
)
|
||||
|
||||
def test_seq_classification_allows_none_label_token_ids(self):
|
||||
"""SequenceClassification models should not require label_token_ids.
|
||||
Verify it passes validation and reaches generate_request."""
|
||||
self.mixin.is_generation = False
|
||||
mock_result = AsyncMock()
|
||||
mock_result.__anext__ = AsyncMock(
|
||||
return_value=[{"embedding": [0.1, 0.9], "meta_info": {"prompt_tokens": 2}}]
|
||||
)
|
||||
self.mixin.generate_request = MagicMock(return_value=mock_result)
|
||||
result = self._call(
|
||||
query=[1, 2],
|
||||
items=[[3, 4]],
|
||||
label_token_ids=None,
|
||||
)
|
||||
self.mixin.generate_request.assert_called_once()
|
||||
self.assertEqual(len(result.scores), 1)
|
||||
|
||||
def test_items_none_raises(self):
|
||||
with self.assertRaisesRegex(ValueError, "items must be provided"):
|
||||
self._call(
|
||||
query=[1, 2],
|
||||
items=None,
|
||||
label_token_ids=[100],
|
||||
)
|
||||
|
||||
def test_empty_items_returns_empty(self):
|
||||
result = self._call(
|
||||
query=[1, 2],
|
||||
items=[],
|
||||
label_token_ids=[100],
|
||||
)
|
||||
self.assertEqual(result.scores, [])
|
||||
self.assertEqual(result.prompt_tokens, 0)
|
||||
|
||||
def test_embed_override_token_id_required_with_query_embeds(self):
|
||||
with self.assertRaisesRegex(ValueError, "embed_override_token_id is required"):
|
||||
self._call(
|
||||
query=[1, 2],
|
||||
items=[[3, 4]],
|
||||
label_token_ids=[100],
|
||||
query_embed_overrides=[_vec(1)],
|
||||
embed_override_token_id=None,
|
||||
)
|
||||
|
||||
def test_embed_override_token_id_required_with_item_embeds(self):
|
||||
with self.assertRaisesRegex(ValueError, "embed_override_token_id is required"):
|
||||
self._call(
|
||||
query=[1, 2],
|
||||
items=[[3, 4]],
|
||||
label_token_ids=[100],
|
||||
item_embed_overrides=[[_vec(1)]],
|
||||
embed_override_token_id=None,
|
||||
)
|
||||
|
||||
def test_item_first_with_embeds_raises(self):
|
||||
with self.assertRaisesRegex(ValueError, "item_first is not supported"):
|
||||
self._call(
|
||||
query=[1, 2],
|
||||
items=[[3, 4]],
|
||||
label_token_ids=[100],
|
||||
item_first=True,
|
||||
embed_override_token_id=50,
|
||||
query_embed_overrides=[_vec(1)],
|
||||
)
|
||||
|
||||
def test_item_embed_overrides_length_mismatch_raises(self):
|
||||
with self.assertRaisesRegex(ValueError, "must match items length"):
|
||||
self._call(
|
||||
query=[1, 2],
|
||||
items=[[3, 4], [5, 6]],
|
||||
label_token_ids=[100],
|
||||
embed_override_token_id=50,
|
||||
item_embed_overrides=[[_vec(1)]], # 1 override for 2 items
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
Reference in New Issue
Block a user