Report per-token weight-version spans in generation meta info (#35926)
This commit is contained in:
@@ -91,6 +91,7 @@ from sglang.srt.managers.io_struct import GenerateReqInput
|
||||
from sglang.srt.parser.conversation import generate_chat_conv
|
||||
from sglang.srt.parser.jinja_template_utils import process_content_for_template_format
|
||||
from sglang.srt.parser.reasoning_parser import ReasoningParser
|
||||
from sglang.srt.utils.weight_versions import build_endpoint_weight_version_metadata
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sglang.srt.managers.tokenizer_manager import TokenizerManager
|
||||
@@ -2005,7 +2006,7 @@ class OpenAIServingChat(OpenAIServingBase):
|
||||
model=request.model,
|
||||
choices=choices,
|
||||
usage=usage,
|
||||
metadata={"weight_version": ret[0]["meta_info"]["weight_version"]},
|
||||
metadata=build_endpoint_weight_version_metadata(ret[0]["meta_info"]),
|
||||
sglext=response_sglext,
|
||||
)
|
||||
|
||||
|
||||
@@ -34,6 +34,7 @@ from sglang.srt.managers.io_struct import GenerateReqInput
|
||||
from sglang.srt.parser.code_completion_parser import (
|
||||
generate_completion_prompt_from_request,
|
||||
)
|
||||
from sglang.srt.utils.weight_versions import build_endpoint_weight_version_metadata
|
||||
from sglang.utils import convert_json_schema_to_str
|
||||
|
||||
if TYPE_CHECKING:
|
||||
@@ -634,7 +635,7 @@ class OpenAIServingCompletion(OpenAIServingBase):
|
||||
created=created,
|
||||
choices=choices,
|
||||
usage=usage,
|
||||
metadata={"weight_version": ret[0]["meta_info"]["weight_version"]},
|
||||
metadata=build_endpoint_weight_version_metadata(ret[0]["meta_info"]),
|
||||
sglext=response_sglext,
|
||||
)
|
||||
|
||||
|
||||
@@ -482,6 +482,7 @@ class DetokenizerManager(MultiHttpWorkerDetokenizerMixin):
|
||||
placeholder_tokens_idx=None,
|
||||
placeholder_tokens_val=None,
|
||||
retraction_counts=recv_obj.retraction_counts,
|
||||
weight_versions=recv_obj.weight_versions,
|
||||
token_steps=recv_obj.token_steps,
|
||||
dp_ranks=recv_obj.dp_ranks,
|
||||
time_stats=recv_obj.time_stats,
|
||||
|
||||
@@ -68,6 +68,7 @@ from sglang.srt.utils.msgspec_utils import (
|
||||
Base64Bytes,
|
||||
msgspec_struct_pydantic_core_schema,
|
||||
)
|
||||
from sglang.srt.utils.weight_versions import WeightVersionSpans
|
||||
|
||||
# Handle serialization of Image for pydantic
|
||||
if TYPE_CHECKING:
|
||||
@@ -1455,6 +1456,8 @@ class BatchTokenIDOutput(BaseBatchReq, kw_only=True):
|
||||
# Number of times each request was retracted.
|
||||
retraction_counts: Optional[List[int]] = None
|
||||
|
||||
weight_versions: Optional[List[Optional[WeightVersionSpans]]] = None
|
||||
|
||||
# The trainer step id. Used to know which step's weights are used for sampling.
|
||||
token_steps: Optional[List[List[int]]] = None
|
||||
|
||||
@@ -1546,6 +1549,8 @@ class BatchStrOutput(BaseBatchReq, kw_only=True):
|
||||
# Number of times each request was retracted.
|
||||
retraction_counts: Optional[List[int]] = None
|
||||
|
||||
weight_versions: Optional[List[Optional[WeightVersionSpans]]] = None
|
||||
|
||||
# The trainer step id. Used to know which step's weights are used for sampling.
|
||||
token_steps: Optional[List[List[int]]] = None
|
||||
|
||||
@@ -2005,6 +2010,7 @@ class AbortReq(BaseReq, kw_only=True):
|
||||
# The finished reason data (from BaseFinishReason.to_json())
|
||||
finished_reason: Optional[FinishReasonDict] = None
|
||||
abort_message: Optional[str] = None
|
||||
weight_versions: Optional[WeightVersionSpans] = None
|
||||
|
||||
def __post_init__(self):
|
||||
# FIXME: This is a hack to keep the same with the old code
|
||||
|
||||
@@ -260,6 +260,7 @@ def _handle_output_by_index(output, i):
|
||||
output, "indexer_topk", i, check_length=False
|
||||
),
|
||||
retraction_counts=_extract_field_by_index(output, "retraction_counts", i),
|
||||
weight_versions=_extract_field_by_index(output, "weight_versions", i),
|
||||
placeholder_tokens_idx=None,
|
||||
placeholder_tokens_val=None,
|
||||
token_steps=_extract_field_by_index(
|
||||
@@ -383,6 +384,7 @@ def _handle_output_by_index(output, i):
|
||||
placeholder_tokens_idx=None,
|
||||
placeholder_tokens_val=None,
|
||||
retraction_counts=_extract_field_by_index(output, "retraction_counts", i),
|
||||
weight_versions=_extract_field_by_index(output, "weight_versions", i),
|
||||
token_steps=_extract_field_by_index(
|
||||
output, "token_steps", i, check_length=False
|
||||
),
|
||||
|
||||
@@ -19,6 +19,10 @@ from sglang.srt.utils.common import (
|
||||
flatten_arrays_to_pinned_cpu,
|
||||
is_pin_memory_available,
|
||||
)
|
||||
from sglang.srt.utils.weight_versions import (
|
||||
WeightVersionEvent,
|
||||
truncate_weight_version_events,
|
||||
)
|
||||
|
||||
# Copyright 2023-2024 SGLang Team
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
@@ -1038,6 +1042,8 @@ class Req(ReqDllmMixin):
|
||||
# Indicates if the req has ever been retracted.
|
||||
self.retracted_stain = False
|
||||
|
||||
self.weight_version_events: List[WeightVersionEvent] = []
|
||||
|
||||
# Incremental streamining
|
||||
self.send_token_offset: int = 0
|
||||
self.send_decode_id_offset: int = 0
|
||||
@@ -1716,6 +1722,9 @@ class Req(ReqDllmMixin):
|
||||
# to ensure shape consistency in KV cache.
|
||||
if self.input_embeds is not None:
|
||||
self.output_ids = array("q")
|
||||
self.weight_version_events = truncate_weight_version_events(
|
||||
self.weight_version_events, num_kept_tokens=self.send_token_offset
|
||||
)
|
||||
|
||||
def offload_kv_cache(self, req_to_token_pool, token_to_kv_pool_allocator):
|
||||
token_indices = req_to_token_pool.req_to_token[
|
||||
|
||||
@@ -332,6 +332,10 @@ from sglang.srt.utils.msgspec_utils import msgspec_to_builtins
|
||||
from sglang.srt.utils.numa_utils import get_numa_node_if_available, numa_bind_to_node
|
||||
from sglang.srt.utils.nvtx_utils import scheduler_nvtx_method
|
||||
from sglang.srt.utils.torch_memory_saver_adapter import TorchMemorySaverAdapter
|
||||
from sglang.srt.utils.weight_versions import (
|
||||
compute_weight_version_spans,
|
||||
record_weight_version_events,
|
||||
)
|
||||
from sglang.utils import TypeBasedDispatcher, get_exception_traceback
|
||||
|
||||
if is_mps():
|
||||
@@ -4586,7 +4590,20 @@ class Scheduler(
|
||||
|
||||
old_version = get_serving().weight_version
|
||||
get_context().override("scheduler.weight_version", weight_version=new_version)
|
||||
logger.info(f"Weight version changed. {old_version=} {new_version=}")
|
||||
|
||||
live_reqs = {
|
||||
*self.collect_inflight_reqs(),
|
||||
*self.waiting_queue,
|
||||
*([self.chunked_req] if self.chunked_req is not None else []),
|
||||
}
|
||||
if self.hisparse_coordinator is not None:
|
||||
live_reqs.update(
|
||||
act.req for act in self.hisparse_coordinator.ack_staging_queue
|
||||
)
|
||||
num_recorded = record_weight_version_events(live_reqs, old_version=old_version)
|
||||
logger.info(
|
||||
f"Weight version changed. {old_version=} {new_version=} {num_recorded=}"
|
||||
)
|
||||
|
||||
def collect_inflight_reqs(self) -> Set[Req]:
|
||||
if self.ps.pp_size == 1:
|
||||
@@ -5258,4 +5275,12 @@ def run_scheduler_process(
|
||||
def _make_abort_req(
|
||||
req: Req, finished_reason: Optional[FinishReasonDict] = None
|
||||
) -> AbortReq:
|
||||
return AbortReq(rid=req.rid, finished_reason=finished_reason)
|
||||
return AbortReq(
|
||||
rid=req.rid,
|
||||
finished_reason=finished_reason,
|
||||
weight_versions=compute_weight_version_spans(
|
||||
req.weight_version_events,
|
||||
current_version=get_serving().weight_version,
|
||||
num_output_tokens=len(req.output_ids),
|
||||
),
|
||||
)
|
||||
|
||||
@@ -31,6 +31,7 @@ from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache
|
||||
from sglang.srt.runtime_context import get_observability, get_serving
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
|
||||
from sglang.srt.utils.weight_versions import compute_weight_version_spans
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sglang.srt.managers.rust_server import RustServer
|
||||
@@ -168,6 +169,7 @@ class SchedulerOutputStreamer:
|
||||
default_force_stream_interval=DEFAULT_FORCE_STREAM_INTERVAL,
|
||||
get_cached_tokens_details=self.get_cached_tokens_details,
|
||||
rust_server_mode=self.rust_server is not None,
|
||||
current_weight_version=get_serving().weight_version,
|
||||
)
|
||||
for req in reqs:
|
||||
if req is skip_req:
|
||||
@@ -316,6 +318,7 @@ class _GenerationStreamAccumulator:
|
||||
default_stream_interval: int
|
||||
default_force_stream_interval: int
|
||||
get_cached_tokens_details: Callable[[Req], Optional[CachedTokensDetails]]
|
||||
current_weight_version: Optional[str]
|
||||
rids: list = field(default_factory=list)
|
||||
output_reqs: list[Req] = field(default_factory=list)
|
||||
http_worker_ipcs: list = field(default_factory=list)
|
||||
@@ -344,6 +347,7 @@ class _GenerationStreamAccumulator:
|
||||
spec_correct_drafts_histogram: list = field(default_factory=list)
|
||||
spec_cap_lens_histogram: list = field(default_factory=list)
|
||||
retraction_counts: list = field(default_factory=list)
|
||||
weight_versions: list = field(default_factory=list)
|
||||
output_hidden_states: Optional[list] = None
|
||||
routed_experts: Optional[list] = None
|
||||
indexer_topk: Optional[list] = None
|
||||
@@ -489,6 +493,16 @@ class _GenerationStreamAccumulator:
|
||||
self.video_tokens.append(video_t)
|
||||
|
||||
self.retraction_counts.append(req.retraction_count)
|
||||
if req.finished():
|
||||
self.weight_versions.append(
|
||||
compute_weight_version_spans(
|
||||
req.weight_version_events,
|
||||
current_version=self.current_weight_version,
|
||||
num_output_tokens=len(output_ids_),
|
||||
)
|
||||
)
|
||||
else:
|
||||
self.weight_versions.append(None)
|
||||
|
||||
self.time_stats.append(req.time_stats)
|
||||
|
||||
@@ -725,5 +739,8 @@ class _GenerationStreamAccumulator:
|
||||
placeholder_tokens_idx=None,
|
||||
placeholder_tokens_val=None,
|
||||
retraction_counts=self.retraction_counts,
|
||||
weight_versions=(
|
||||
self.weight_versions if any(self.weight_versions) else None
|
||||
),
|
||||
dp_ranks=dp_ranks,
|
||||
)
|
||||
|
||||
@@ -165,6 +165,7 @@ from sglang.srt.utils.hf_transformers_utils import (
|
||||
from sglang.srt.utils.network import get_zmq_socket
|
||||
from sglang.srt.utils.request_logger import RequestLogger
|
||||
from sglang.srt.utils.watchdog import Watchdog
|
||||
from sglang.srt.utils.weight_versions import add_weight_versions_to_meta_info
|
||||
from sglang.utils import TypeBasedDispatcher, get_exception_traceback
|
||||
|
||||
asyncio.set_event_loop_policy(uvloop.EventLoopPolicy())
|
||||
@@ -2245,6 +2246,15 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
||||
"cached_tokens": recv_obj.cached_tokens[i],
|
||||
}
|
||||
)
|
||||
if (
|
||||
recv_obj.weight_versions is not None
|
||||
and (spans := recv_obj.weight_versions[i]) is not None
|
||||
):
|
||||
add_weight_versions_to_meta_info(
|
||||
meta_info,
|
||||
spans,
|
||||
num_output_tokens=recv_obj.completion_tokens[i],
|
||||
)
|
||||
# Add detailed cache breakdown if available
|
||||
if (
|
||||
hasattr(recv_obj, "cached_tokens_details")
|
||||
@@ -3179,6 +3189,12 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
||||
"weight_version": self.config_value("weight_version"),
|
||||
"e2e_latency": state.time_stats.get_e2e_latency(),
|
||||
}
|
||||
if recv_obj.weight_versions is not None:
|
||||
add_weight_versions_to_meta_info(
|
||||
meta_info,
|
||||
recv_obj.weight_versions,
|
||||
num_output_tokens=len(state.output_ids),
|
||||
)
|
||||
is_stream = getattr(state.obj, "stream", False)
|
||||
if getattr(state.obj, "return_logprob", False):
|
||||
self.add_logprob_to_meta_info(
|
||||
|
||||
@@ -0,0 +1,117 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import dataclasses
|
||||
from typing import TYPE_CHECKING, Any, Dict, Iterable, List
|
||||
|
||||
import msgspec
|
||||
|
||||
from sglang.srt.utils.msgspec_utils import msgspec_struct_pydantic_core_schema
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sglang.srt.managers.schedule_batch import Req
|
||||
|
||||
|
||||
# ======================================================================
|
||||
# Shared types
|
||||
# ======================================================================
|
||||
class WeightVersionSpan(msgspec.Struct, kw_only=True, array_like=True):
|
||||
version: str
|
||||
start: int
|
||||
end: int
|
||||
|
||||
@classmethod
|
||||
def __get_pydantic_core_schema__(cls, source, handler):
|
||||
return msgspec_struct_pydantic_core_schema(cls, handler)
|
||||
|
||||
|
||||
WeightVersionSpans = List[WeightVersionSpan]
|
||||
|
||||
|
||||
@dataclasses.dataclass(frozen=True, slots=True)
|
||||
class WeightVersionEvent:
|
||||
old_version: str
|
||||
num_output_tokens: int
|
||||
|
||||
|
||||
# ======================================================================
|
||||
# Scheduler process
|
||||
# ======================================================================
|
||||
def record_weight_version_events(reqs: Iterable[Req], old_version: str) -> int:
|
||||
num_recorded = 0
|
||||
for req in reqs:
|
||||
if req.output_ids:
|
||||
req.weight_version_events.append(
|
||||
WeightVersionEvent(
|
||||
old_version=old_version,
|
||||
num_output_tokens=len(req.output_ids),
|
||||
)
|
||||
)
|
||||
num_recorded += 1
|
||||
return num_recorded
|
||||
|
||||
|
||||
def truncate_weight_version_events(
|
||||
events: List[WeightVersionEvent], num_kept_tokens: int
|
||||
) -> List[WeightVersionEvent]:
|
||||
truncated = [
|
||||
WeightVersionEvent(
|
||||
old_version=event.old_version,
|
||||
num_output_tokens=min(event.num_output_tokens, num_kept_tokens),
|
||||
)
|
||||
for event in events
|
||||
]
|
||||
return [event for event in truncated if event.num_output_tokens > 0]
|
||||
|
||||
|
||||
def compute_weight_version_spans(
|
||||
events: List[WeightVersionEvent],
|
||||
current_version: str,
|
||||
num_output_tokens: int,
|
||||
) -> WeightVersionSpans:
|
||||
changes = [(event.old_version, event.num_output_tokens) for event in events]
|
||||
changes.append((current_version, num_output_tokens))
|
||||
|
||||
spans: WeightVersionSpans = []
|
||||
for version, end in changes:
|
||||
end = min(end, num_output_tokens)
|
||||
if spans and end <= spans[-1].end:
|
||||
continue
|
||||
if spans and version == spans[-1].version:
|
||||
spans[-1].end = end
|
||||
continue
|
||||
start = spans[-1].end if spans else 0
|
||||
spans.append(WeightVersionSpan(version=version, start=start, end=end))
|
||||
return spans
|
||||
|
||||
|
||||
# ======================================================================
|
||||
# TokenizerManager
|
||||
# ======================================================================
|
||||
def add_weight_versions_to_meta_info(
|
||||
meta_info: Dict[str, Any],
|
||||
spans: WeightVersionSpans,
|
||||
num_output_tokens: int,
|
||||
) -> None:
|
||||
visible = [
|
||||
span for span in spans if span.start < num_output_tokens or span.start == 0
|
||||
]
|
||||
|
||||
meta_info["weight_versions"] = [
|
||||
{
|
||||
"version": span.version,
|
||||
"start": span.start,
|
||||
"end": min(span.end, num_output_tokens),
|
||||
}
|
||||
for span in visible
|
||||
]
|
||||
meta_info["weight_version"] = visible[-1].version
|
||||
|
||||
|
||||
# ======================================================================
|
||||
# OpenAI-compatible endpoints
|
||||
# ======================================================================
|
||||
def build_endpoint_weight_version_metadata(meta_info: Dict[str, Any]) -> Dict[str, Any]:
|
||||
metadata = {"weight_version": meta_info["weight_version"]}
|
||||
if "weight_versions" in meta_info:
|
||||
metadata["weight_versions"] = meta_info["weight_versions"]
|
||||
return metadata
|
||||
Reference in New Issue
Block a user