Report per-token weight-version spans in generation meta info (#35926)
This commit is contained in:
@@ -578,6 +578,20 @@ Training workers gather weights (typically on TP rank 0), broadcast them to the
|
|||||||
- `engine.update_weights_from_distributed(names, dtypes, shapes, ...)`
|
- `engine.update_weights_from_distributed(names, dtypes, shapes, ...)`
|
||||||
- `engine.destroy_weights_update_group(group_name)`
|
- `engine.destroy_weights_update_group(group_name)`
|
||||||
|
|
||||||
|
### Per-Token Weight Version Attribution
|
||||||
|
|
||||||
|
A request can outlive a version change: the RL flow retracts it with `pause_generation`, refits, and continues it, or `POST /update_weight_version` relabels while it is still generating. `meta_info` then reports which tokens came from which weights:
|
||||||
|
|
||||||
|
```json Output
|
||||||
|
"weight_version": "42",
|
||||||
|
"weight_versions": [
|
||||||
|
{"version": "41", "start": 0, "end": 57},
|
||||||
|
{"version": "42", "start": 57, "end": 128}
|
||||||
|
]
|
||||||
|
```
|
||||||
|
|
||||||
|
- Ranges are half-open `[start, end)` over output-token indices, prompt excluded. The usual case is a single span.
|
||||||
|
|
||||||
## Easy To Postpone Generation
|
## Easy To Postpone Generation
|
||||||
|
|
||||||
Multi-turn RL rollouts often suffer from long-tail requests that block the entire batch. A small number of slow interactions can stall all GPUs, and the long-tail behavior makes profiling and monitoring difficult.
|
Multi-turn RL rollouts often suffer from long-tail requests that block the entire batch. A small number of slow interactions can stall all GPUs, and the long-tail behavior makes profiling and monitoring difficult.
|
||||||
|
|||||||
@@ -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.conversation import generate_chat_conv
|
||||||
from sglang.srt.parser.jinja_template_utils import process_content_for_template_format
|
from sglang.srt.parser.jinja_template_utils import process_content_for_template_format
|
||||||
from sglang.srt.parser.reasoning_parser import ReasoningParser
|
from sglang.srt.parser.reasoning_parser import ReasoningParser
|
||||||
|
from sglang.srt.utils.weight_versions import build_endpoint_weight_version_metadata
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from sglang.srt.managers.tokenizer_manager import TokenizerManager
|
from sglang.srt.managers.tokenizer_manager import TokenizerManager
|
||||||
@@ -2005,7 +2006,7 @@ class OpenAIServingChat(OpenAIServingBase):
|
|||||||
model=request.model,
|
model=request.model,
|
||||||
choices=choices,
|
choices=choices,
|
||||||
usage=usage,
|
usage=usage,
|
||||||
metadata={"weight_version": ret[0]["meta_info"]["weight_version"]},
|
metadata=build_endpoint_weight_version_metadata(ret[0]["meta_info"]),
|
||||||
sglext=response_sglext,
|
sglext=response_sglext,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -34,6 +34,7 @@ from sglang.srt.managers.io_struct import GenerateReqInput
|
|||||||
from sglang.srt.parser.code_completion_parser import (
|
from sglang.srt.parser.code_completion_parser import (
|
||||||
generate_completion_prompt_from_request,
|
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
|
from sglang.utils import convert_json_schema_to_str
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
@@ -634,7 +635,7 @@ class OpenAIServingCompletion(OpenAIServingBase):
|
|||||||
created=created,
|
created=created,
|
||||||
choices=choices,
|
choices=choices,
|
||||||
usage=usage,
|
usage=usage,
|
||||||
metadata={"weight_version": ret[0]["meta_info"]["weight_version"]},
|
metadata=build_endpoint_weight_version_metadata(ret[0]["meta_info"]),
|
||||||
sglext=response_sglext,
|
sglext=response_sglext,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -482,6 +482,7 @@ class DetokenizerManager(MultiHttpWorkerDetokenizerMixin):
|
|||||||
placeholder_tokens_idx=None,
|
placeholder_tokens_idx=None,
|
||||||
placeholder_tokens_val=None,
|
placeholder_tokens_val=None,
|
||||||
retraction_counts=recv_obj.retraction_counts,
|
retraction_counts=recv_obj.retraction_counts,
|
||||||
|
weight_versions=recv_obj.weight_versions,
|
||||||
token_steps=recv_obj.token_steps,
|
token_steps=recv_obj.token_steps,
|
||||||
dp_ranks=recv_obj.dp_ranks,
|
dp_ranks=recv_obj.dp_ranks,
|
||||||
time_stats=recv_obj.time_stats,
|
time_stats=recv_obj.time_stats,
|
||||||
|
|||||||
@@ -68,6 +68,7 @@ from sglang.srt.utils.msgspec_utils import (
|
|||||||
Base64Bytes,
|
Base64Bytes,
|
||||||
msgspec_struct_pydantic_core_schema,
|
msgspec_struct_pydantic_core_schema,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.utils.weight_versions import WeightVersionSpans
|
||||||
|
|
||||||
# Handle serialization of Image for pydantic
|
# Handle serialization of Image for pydantic
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
@@ -1455,6 +1456,8 @@ class BatchTokenIDOutput(BaseBatchReq, kw_only=True):
|
|||||||
# Number of times each request was retracted.
|
# Number of times each request was retracted.
|
||||||
retraction_counts: Optional[List[int]] = None
|
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.
|
# The trainer step id. Used to know which step's weights are used for sampling.
|
||||||
token_steps: Optional[List[List[int]]] = None
|
token_steps: Optional[List[List[int]]] = None
|
||||||
|
|
||||||
@@ -1546,6 +1549,8 @@ class BatchStrOutput(BaseBatchReq, kw_only=True):
|
|||||||
# Number of times each request was retracted.
|
# Number of times each request was retracted.
|
||||||
retraction_counts: Optional[List[int]] = None
|
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.
|
# The trainer step id. Used to know which step's weights are used for sampling.
|
||||||
token_steps: Optional[List[List[int]]] = None
|
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())
|
# The finished reason data (from BaseFinishReason.to_json())
|
||||||
finished_reason: Optional[FinishReasonDict] = None
|
finished_reason: Optional[FinishReasonDict] = None
|
||||||
abort_message: Optional[str] = None
|
abort_message: Optional[str] = None
|
||||||
|
weight_versions: Optional[WeightVersionSpans] = None
|
||||||
|
|
||||||
def __post_init__(self):
|
def __post_init__(self):
|
||||||
# FIXME: This is a hack to keep the same with the old code
|
# 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
|
output, "indexer_topk", i, check_length=False
|
||||||
),
|
),
|
||||||
retraction_counts=_extract_field_by_index(output, "retraction_counts", i),
|
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_idx=None,
|
||||||
placeholder_tokens_val=None,
|
placeholder_tokens_val=None,
|
||||||
token_steps=_extract_field_by_index(
|
token_steps=_extract_field_by_index(
|
||||||
@@ -383,6 +384,7 @@ def _handle_output_by_index(output, i):
|
|||||||
placeholder_tokens_idx=None,
|
placeholder_tokens_idx=None,
|
||||||
placeholder_tokens_val=None,
|
placeholder_tokens_val=None,
|
||||||
retraction_counts=_extract_field_by_index(output, "retraction_counts", i),
|
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(
|
token_steps=_extract_field_by_index(
|
||||||
output, "token_steps", i, check_length=False
|
output, "token_steps", i, check_length=False
|
||||||
),
|
),
|
||||||
|
|||||||
@@ -19,6 +19,10 @@ from sglang.srt.utils.common import (
|
|||||||
flatten_arrays_to_pinned_cpu,
|
flatten_arrays_to_pinned_cpu,
|
||||||
is_pin_memory_available,
|
is_pin_memory_available,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.utils.weight_versions import (
|
||||||
|
WeightVersionEvent,
|
||||||
|
truncate_weight_version_events,
|
||||||
|
)
|
||||||
|
|
||||||
# Copyright 2023-2024 SGLang Team
|
# Copyright 2023-2024 SGLang Team
|
||||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
# 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.
|
# Indicates if the req has ever been retracted.
|
||||||
self.retracted_stain = False
|
self.retracted_stain = False
|
||||||
|
|
||||||
|
self.weight_version_events: List[WeightVersionEvent] = []
|
||||||
|
|
||||||
# Incremental streamining
|
# Incremental streamining
|
||||||
self.send_token_offset: int = 0
|
self.send_token_offset: int = 0
|
||||||
self.send_decode_id_offset: int = 0
|
self.send_decode_id_offset: int = 0
|
||||||
@@ -1716,6 +1722,9 @@ class Req(ReqDllmMixin):
|
|||||||
# to ensure shape consistency in KV cache.
|
# to ensure shape consistency in KV cache.
|
||||||
if self.input_embeds is not None:
|
if self.input_embeds is not None:
|
||||||
self.output_ids = array("q")
|
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):
|
def offload_kv_cache(self, req_to_token_pool, token_to_kv_pool_allocator):
|
||||||
token_indices = req_to_token_pool.req_to_token[
|
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.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.nvtx_utils import scheduler_nvtx_method
|
||||||
from sglang.srt.utils.torch_memory_saver_adapter import TorchMemorySaverAdapter
|
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
|
from sglang.utils import TypeBasedDispatcher, get_exception_traceback
|
||||||
|
|
||||||
if is_mps():
|
if is_mps():
|
||||||
@@ -4586,7 +4590,20 @@ class Scheduler(
|
|||||||
|
|
||||||
old_version = get_serving().weight_version
|
old_version = get_serving().weight_version
|
||||||
get_context().override("scheduler.weight_version", weight_version=new_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]:
|
def collect_inflight_reqs(self) -> Set[Req]:
|
||||||
if self.ps.pp_size == 1:
|
if self.ps.pp_size == 1:
|
||||||
@@ -5258,4 +5275,12 @@ def run_scheduler_process(
|
|||||||
def _make_abort_req(
|
def _make_abort_req(
|
||||||
req: Req, finished_reason: Optional[FinishReasonDict] = None
|
req: Req, finished_reason: Optional[FinishReasonDict] = None
|
||||||
) -> AbortReq:
|
) -> 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.runtime_context import get_observability, get_serving
|
||||||
from sglang.srt.server_args import ServerArgs
|
from sglang.srt.server_args import ServerArgs
|
||||||
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
|
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
|
||||||
|
from sglang.srt.utils.weight_versions import compute_weight_version_spans
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from sglang.srt.managers.rust_server import RustServer
|
from sglang.srt.managers.rust_server import RustServer
|
||||||
@@ -168,6 +169,7 @@ class SchedulerOutputStreamer:
|
|||||||
default_force_stream_interval=DEFAULT_FORCE_STREAM_INTERVAL,
|
default_force_stream_interval=DEFAULT_FORCE_STREAM_INTERVAL,
|
||||||
get_cached_tokens_details=self.get_cached_tokens_details,
|
get_cached_tokens_details=self.get_cached_tokens_details,
|
||||||
rust_server_mode=self.rust_server is not None,
|
rust_server_mode=self.rust_server is not None,
|
||||||
|
current_weight_version=get_serving().weight_version,
|
||||||
)
|
)
|
||||||
for req in reqs:
|
for req in reqs:
|
||||||
if req is skip_req:
|
if req is skip_req:
|
||||||
@@ -316,6 +318,7 @@ class _GenerationStreamAccumulator:
|
|||||||
default_stream_interval: int
|
default_stream_interval: int
|
||||||
default_force_stream_interval: int
|
default_force_stream_interval: int
|
||||||
get_cached_tokens_details: Callable[[Req], Optional[CachedTokensDetails]]
|
get_cached_tokens_details: Callable[[Req], Optional[CachedTokensDetails]]
|
||||||
|
current_weight_version: Optional[str]
|
||||||
rids: list = field(default_factory=list)
|
rids: list = field(default_factory=list)
|
||||||
output_reqs: list[Req] = field(default_factory=list)
|
output_reqs: list[Req] = field(default_factory=list)
|
||||||
http_worker_ipcs: list = 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_correct_drafts_histogram: list = field(default_factory=list)
|
||||||
spec_cap_lens_histogram: list = field(default_factory=list)
|
spec_cap_lens_histogram: list = field(default_factory=list)
|
||||||
retraction_counts: 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
|
output_hidden_states: Optional[list] = None
|
||||||
routed_experts: Optional[list] = None
|
routed_experts: Optional[list] = None
|
||||||
indexer_topk: Optional[list] = None
|
indexer_topk: Optional[list] = None
|
||||||
@@ -489,6 +493,16 @@ class _GenerationStreamAccumulator:
|
|||||||
self.video_tokens.append(video_t)
|
self.video_tokens.append(video_t)
|
||||||
|
|
||||||
self.retraction_counts.append(req.retraction_count)
|
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)
|
self.time_stats.append(req.time_stats)
|
||||||
|
|
||||||
@@ -725,5 +739,8 @@ class _GenerationStreamAccumulator:
|
|||||||
placeholder_tokens_idx=None,
|
placeholder_tokens_idx=None,
|
||||||
placeholder_tokens_val=None,
|
placeholder_tokens_val=None,
|
||||||
retraction_counts=self.retraction_counts,
|
retraction_counts=self.retraction_counts,
|
||||||
|
weight_versions=(
|
||||||
|
self.weight_versions if any(self.weight_versions) else None
|
||||||
|
),
|
||||||
dp_ranks=dp_ranks,
|
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.network import get_zmq_socket
|
||||||
from sglang.srt.utils.request_logger import RequestLogger
|
from sglang.srt.utils.request_logger import RequestLogger
|
||||||
from sglang.srt.utils.watchdog import Watchdog
|
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
|
from sglang.utils import TypeBasedDispatcher, get_exception_traceback
|
||||||
|
|
||||||
asyncio.set_event_loop_policy(uvloop.EventLoopPolicy())
|
asyncio.set_event_loop_policy(uvloop.EventLoopPolicy())
|
||||||
@@ -2245,6 +2246,15 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
|||||||
"cached_tokens": recv_obj.cached_tokens[i],
|
"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
|
# Add detailed cache breakdown if available
|
||||||
if (
|
if (
|
||||||
hasattr(recv_obj, "cached_tokens_details")
|
hasattr(recv_obj, "cached_tokens_details")
|
||||||
@@ -3179,6 +3189,12 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
|||||||
"weight_version": self.config_value("weight_version"),
|
"weight_version": self.config_value("weight_version"),
|
||||||
"e2e_latency": state.time_stats.get_e2e_latency(),
|
"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)
|
is_stream = getattr(state.obj, "stream", False)
|
||||||
if getattr(state.obj, "return_logprob", False):
|
if getattr(state.obj, "return_logprob", False):
|
||||||
self.add_logprob_to_meta_info(
|
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
|
||||||
@@ -0,0 +1,590 @@
|
|||||||
|
import json
|
||||||
|
import time
|
||||||
|
import unittest
|
||||||
|
from concurrent.futures import ThreadPoolExecutor
|
||||||
|
|
||||||
|
import requests
|
||||||
|
|
||||||
|
from sglang.srt.utils import kill_process_tree
|
||||||
|
from sglang.test.ci.ci_register import register_cuda_ci
|
||||||
|
from sglang.test.test_utils import (
|
||||||
|
DEFAULT_MODEL_NAME_FOR_TEST_MLA,
|
||||||
|
DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
|
||||||
|
DEFAULT_URL_FOR_TEST,
|
||||||
|
CustomTestCase,
|
||||||
|
popen_launch_server,
|
||||||
|
)
|
||||||
|
|
||||||
|
register_cuda_ci(
|
||||||
|
est_time=180, stage="nightly", runner_config="2-gpu-large", nightly=True
|
||||||
|
)
|
||||||
|
|
||||||
|
_REQUEST_TIMEOUT = 180
|
||||||
|
|
||||||
|
|
||||||
|
def _assert_spans_contiguous(test, meta_info):
|
||||||
|
spans = meta_info["weight_versions"]
|
||||||
|
test.assertGreater(len(spans), 0)
|
||||||
|
test.assertEqual(spans[0]["start"], 0)
|
||||||
|
for prev, cur in zip(spans, spans[1:]):
|
||||||
|
test.assertEqual(prev["end"], cur["start"])
|
||||||
|
test.assertNotEqual(prev["version"], cur["version"])
|
||||||
|
test.assertEqual(meta_info["weight_version"], spans[-1]["version"])
|
||||||
|
return spans
|
||||||
|
|
||||||
|
|
||||||
|
class TestWeightVersionSpans(CustomTestCase):
|
||||||
|
@classmethod
|
||||||
|
def setUpClass(cls):
|
||||||
|
cls.model = DEFAULT_MODEL_NAME_FOR_TEST_MLA
|
||||||
|
cls.base_url = DEFAULT_URL_FOR_TEST
|
||||||
|
cls.process = popen_launch_server(
|
||||||
|
cls.model,
|
||||||
|
base_url=cls.base_url,
|
||||||
|
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
|
||||||
|
other_args=[
|
||||||
|
"--weight-version",
|
||||||
|
"base-v0",
|
||||||
|
"--trust-remote-code",
|
||||||
|
"--tp-size",
|
||||||
|
"2",
|
||||||
|
"--dp-size",
|
||||||
|
"2",
|
||||||
|
"--enable-dp-attention",
|
||||||
|
"--speculative-algorithm",
|
||||||
|
"EAGLE",
|
||||||
|
"--speculative-draft-model-path",
|
||||||
|
DEFAULT_MODEL_NAME_FOR_TEST_MLA,
|
||||||
|
"--speculative-num-steps",
|
||||||
|
"2",
|
||||||
|
"--speculative-eagle-topk",
|
||||||
|
"3",
|
||||||
|
"--speculative-num-draft-tokens",
|
||||||
|
"3",
|
||||||
|
"--cuda-graph-max-bs-decode",
|
||||||
|
"32",
|
||||||
|
"--max-running-requests",
|
||||||
|
"8",
|
||||||
|
],
|
||||||
|
)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def tearDownClass(cls):
|
||||||
|
kill_process_tree(cls.process.pid)
|
||||||
|
|
||||||
|
def _generate(self, max_new_tokens: int, prompt: str = "The capital of France is"):
|
||||||
|
response = requests.post(
|
||||||
|
f"{self.base_url}/generate",
|
||||||
|
json={
|
||||||
|
"text": prompt,
|
||||||
|
"sampling_params": {
|
||||||
|
"temperature": 0.8,
|
||||||
|
"max_new_tokens": max_new_tokens,
|
||||||
|
"ignore_eos": True,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
timeout=_REQUEST_TIMEOUT,
|
||||||
|
)
|
||||||
|
self.assertEqual(response.status_code, 200)
|
||||||
|
return response.json()
|
||||||
|
|
||||||
|
def _current_version(self):
|
||||||
|
response = requests.get(f"{self.base_url}/get_model_info", timeout=30)
|
||||||
|
self.assertEqual(response.status_code, 200)
|
||||||
|
return response.json()["weight_version"]
|
||||||
|
|
||||||
|
def _pause(self, mode: str):
|
||||||
|
requests.post(
|
||||||
|
f"{self.base_url}/pause_generation",
|
||||||
|
json={"mode": mode},
|
||||||
|
timeout=30,
|
||||||
|
).raise_for_status()
|
||||||
|
|
||||||
|
def _continue(self):
|
||||||
|
requests.post(
|
||||||
|
f"{self.base_url}/continue_generation",
|
||||||
|
json={},
|
||||||
|
timeout=30,
|
||||||
|
).raise_for_status()
|
||||||
|
|
||||||
|
def _set_weight_version(self, new_version: str, abort_all_requests: bool = False):
|
||||||
|
response = requests.post(
|
||||||
|
f"{self.base_url}/update_weight_version",
|
||||||
|
json={
|
||||||
|
"new_version": new_version,
|
||||||
|
"abort_all_requests": abort_all_requests,
|
||||||
|
},
|
||||||
|
timeout=30,
|
||||||
|
)
|
||||||
|
self.assertEqual(response.status_code, 200)
|
||||||
|
return response.json()
|
||||||
|
|
||||||
|
def _update_weights_from_disk(self, **fields) -> None:
|
||||||
|
response = requests.post(
|
||||||
|
f"{self.base_url}/update_weights_from_disk",
|
||||||
|
json={"model_path": self.model, "flush_cache": False, **fields},
|
||||||
|
timeout=_REQUEST_TIMEOUT,
|
||||||
|
)
|
||||||
|
self.assertEqual(response.status_code, 200)
|
||||||
|
self.assertTrue(response.json()["success"])
|
||||||
|
|
||||||
|
def _run_while_paused(
|
||||||
|
self,
|
||||||
|
num_requests: int,
|
||||||
|
while_paused,
|
||||||
|
mode: str = "retract",
|
||||||
|
max_new_tokens: int = 1024,
|
||||||
|
):
|
||||||
|
with ThreadPoolExecutor(max_workers=num_requests) as executor:
|
||||||
|
futures = [
|
||||||
|
executor.submit(
|
||||||
|
self._generate,
|
||||||
|
max_new_tokens=max_new_tokens - 16 * i,
|
||||||
|
prompt=f"Write a long story about the number {i}.",
|
||||||
|
)
|
||||||
|
for i in range(num_requests)
|
||||||
|
]
|
||||||
|
|
||||||
|
time.sleep(2)
|
||||||
|
self._pause(mode)
|
||||||
|
try:
|
||||||
|
while_paused()
|
||||||
|
finally:
|
||||||
|
self._continue()
|
||||||
|
|
||||||
|
return [future.result() for future in futures]
|
||||||
|
|
||||||
|
def test_01_single_span_without_update(self):
|
||||||
|
"""A request untouched by updates reports one span covering all output tokens."""
|
||||||
|
data = self._generate(max_new_tokens=8)
|
||||||
|
|
||||||
|
meta_info = data["meta_info"]
|
||||||
|
spans = _assert_spans_contiguous(self, meta_info)
|
||||||
|
self.assertEqual(len(spans), 1)
|
||||||
|
self.assertEqual(spans[0]["version"], "base-v0")
|
||||||
|
self.assertEqual(spans[0]["end"], meta_info["completion_tokens"])
|
||||||
|
self.assertEqual(meta_info["weight_version"], "base-v0")
|
||||||
|
|
||||||
|
def test_02_update_weight_version_endpoint_applies_to_new_requests(self):
|
||||||
|
"""The endpoint returns only once every scheduler stamps new requests with the new version."""
|
||||||
|
self._set_weight_version("endpoint-v1")
|
||||||
|
self.assertEqual(self._current_version(), "endpoint-v1")
|
||||||
|
|
||||||
|
with ThreadPoolExecutor(max_workers=4) as executor:
|
||||||
|
futures = [
|
||||||
|
executor.submit(self._generate, max_new_tokens=8) for _ in range(4)
|
||||||
|
]
|
||||||
|
results = [future.result() for future in futures]
|
||||||
|
|
||||||
|
for data in results:
|
||||||
|
meta_info = data["meta_info"]
|
||||||
|
spans = _assert_spans_contiguous(self, meta_info)
|
||||||
|
self.assertEqual(len(spans), 1)
|
||||||
|
self.assertEqual(spans[0]["version"], "endpoint-v1")
|
||||||
|
|
||||||
|
def test_03_spans_split_across_pause_update_continue(self):
|
||||||
|
"""Requests spanning pause -> update_weights_from_disk -> continue report one span per version."""
|
||||||
|
base_version = self._current_version()
|
||||||
|
|
||||||
|
results = self._run_while_paused(
|
||||||
|
num_requests=4,
|
||||||
|
while_paused=lambda: self._update_weights_from_disk(
|
||||||
|
weight_version="disk-v2"
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
multi_span_count = 0
|
||||||
|
for data in results:
|
||||||
|
meta_info = data["meta_info"]
|
||||||
|
spans = _assert_spans_contiguous(self, meta_info)
|
||||||
|
self.assertEqual(spans[-1]["end"], meta_info["completion_tokens"])
|
||||||
|
versions = [span["version"] for span in spans]
|
||||||
|
self.assertEqual(versions[0], base_version)
|
||||||
|
self.assertIn(versions[-1], (base_version, "disk-v2"))
|
||||||
|
if len(spans) > 1:
|
||||||
|
multi_span_count += 1
|
||||||
|
self.assertEqual(versions, [base_version, "disk-v2"])
|
||||||
|
self.assertGreater(spans[0]["end"], 0)
|
||||||
|
|
||||||
|
self.assertGreater(
|
||||||
|
multi_span_count,
|
||||||
|
0,
|
||||||
|
"No request spanned the weight update -- no boundary was recorded.",
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_04_openai_metadata_contains_weight_versions(self):
|
||||||
|
"""OpenAI-compatible responses surface the spans under response metadata."""
|
||||||
|
response = requests.post(
|
||||||
|
f"{self.base_url}/v1/chat/completions",
|
||||||
|
json={
|
||||||
|
"model": self.model,
|
||||||
|
"messages": [{"role": "user", "content": "Hello"}],
|
||||||
|
"max_tokens": 8,
|
||||||
|
"temperature": 0.0,
|
||||||
|
},
|
||||||
|
timeout=_REQUEST_TIMEOUT,
|
||||||
|
)
|
||||||
|
self.assertEqual(response.status_code, 200)
|
||||||
|
|
||||||
|
data = response.json()
|
||||||
|
metadata = data["metadata"]
|
||||||
|
self.assertIn("weight_versions", metadata)
|
||||||
|
spans = metadata["weight_versions"]
|
||||||
|
self.assertEqual(len(spans), 1)
|
||||||
|
self.assertEqual(spans[0]["version"], metadata["weight_version"])
|
||||||
|
self.assertEqual(metadata["weight_version"], self._current_version())
|
||||||
|
self.assertEqual(spans[0]["start"], 0)
|
||||||
|
self.assertEqual(spans[0]["end"], data["usage"]["completion_tokens"])
|
||||||
|
|
||||||
|
def test_04b_openai_metadata_reports_the_first_choice_when_n_is_greater_than_one(
|
||||||
|
self,
|
||||||
|
):
|
||||||
|
"""With n > 1 the single metadata block describes the first choice instead of going missing."""
|
||||||
|
response = requests.post(
|
||||||
|
f"{self.base_url}/v1/chat/completions",
|
||||||
|
json={
|
||||||
|
"model": self.model,
|
||||||
|
"messages": [{"role": "user", "content": "Hello"}],
|
||||||
|
"max_tokens": 8,
|
||||||
|
"temperature": 0.8,
|
||||||
|
"n": 2,
|
||||||
|
},
|
||||||
|
timeout=_REQUEST_TIMEOUT,
|
||||||
|
)
|
||||||
|
self.assertEqual(response.status_code, 200)
|
||||||
|
|
||||||
|
data = response.json()
|
||||||
|
self.assertEqual(len(data["choices"]), 2)
|
||||||
|
metadata = data["metadata"]
|
||||||
|
spans = metadata["weight_versions"]
|
||||||
|
self.assertEqual(len(spans), 1)
|
||||||
|
self.assertEqual(spans[0]["version"], metadata["weight_version"])
|
||||||
|
self.assertEqual(metadata["weight_version"], self._current_version())
|
||||||
|
self.assertEqual(spans[0]["start"], 0)
|
||||||
|
|
||||||
|
def test_05_aborted_retracted_requests_report_spans(self):
|
||||||
|
"""Requests aborted while retracted in the waiting queue still report their spans."""
|
||||||
|
|
||||||
|
def abort_all():
|
||||||
|
requests.post(
|
||||||
|
f"{self.base_url}/abort_request",
|
||||||
|
json={"abort_all": True},
|
||||||
|
timeout=30,
|
||||||
|
).raise_for_status()
|
||||||
|
|
||||||
|
results = self._run_while_paused(num_requests=4, while_paused=abort_all)
|
||||||
|
|
||||||
|
aborted_with_spans = 0
|
||||||
|
for data in results:
|
||||||
|
meta_info = data["meta_info"]
|
||||||
|
if meta_info["finish_reason"]["type"] != "abort":
|
||||||
|
continue
|
||||||
|
if "weight_versions" not in meta_info:
|
||||||
|
continue
|
||||||
|
spans = _assert_spans_contiguous(self, meta_info)
|
||||||
|
self.assertEqual(spans[-1]["end"], meta_info["completion_tokens"])
|
||||||
|
aborted_with_spans += 1
|
||||||
|
|
||||||
|
self.assertGreater(
|
||||||
|
aborted_with_spans,
|
||||||
|
0,
|
||||||
|
"No aborted request carried weight_versions -- the AbortReq path lost the spans.",
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_06_running_requests_split_without_abort(self):
|
||||||
|
"""A version bump with abort_all_requests=False splits requests still in the running batch."""
|
||||||
|
previous_version = self._current_version()
|
||||||
|
|
||||||
|
results = self._run_while_paused(
|
||||||
|
num_requests=4,
|
||||||
|
while_paused=lambda: self._set_weight_version("inplace-v3"),
|
||||||
|
mode="in_place",
|
||||||
|
max_new_tokens=1024,
|
||||||
|
)
|
||||||
|
self.assertEqual(self._current_version(), "inplace-v3")
|
||||||
|
|
||||||
|
split_count = 0
|
||||||
|
for data in results:
|
||||||
|
meta_info = data["meta_info"]
|
||||||
|
self.assertNotEqual(meta_info["finish_reason"]["type"], "abort")
|
||||||
|
spans = _assert_spans_contiguous(self, meta_info)
|
||||||
|
self.assertEqual(spans[-1]["end"], meta_info["completion_tokens"])
|
||||||
|
self.assertEqual(spans[0]["version"], previous_version)
|
||||||
|
if len(spans) > 1:
|
||||||
|
split_count += 1
|
||||||
|
self.assertEqual(
|
||||||
|
[span["version"] for span in spans],
|
||||||
|
[previous_version, "inplace-v3"],
|
||||||
|
)
|
||||||
|
self.assertGreater(spans[0]["end"], 0)
|
||||||
|
|
||||||
|
self.assertGreater(
|
||||||
|
split_count,
|
||||||
|
0,
|
||||||
|
"No running request was split -- the running batch was not visited.",
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_07_update_without_weight_version_does_not_split(self):
|
||||||
|
"""A refit that carries no weight_version leaves attribution untouched."""
|
||||||
|
version = self._current_version()
|
||||||
|
|
||||||
|
results = self._run_while_paused(
|
||||||
|
num_requests=4, while_paused=self._update_weights_from_disk
|
||||||
|
)
|
||||||
|
self.assertEqual(self._current_version(), version)
|
||||||
|
|
||||||
|
for data in results:
|
||||||
|
meta_info = data["meta_info"]
|
||||||
|
spans = _assert_spans_contiguous(self, meta_info)
|
||||||
|
self.assertEqual(len(spans), 1)
|
||||||
|
self.assertEqual(spans[0]["version"], version)
|
||||||
|
self.assertEqual(spans[0]["end"], meta_info["completion_tokens"])
|
||||||
|
|
||||||
|
def test_08_reannouncing_current_version_is_a_noop(self):
|
||||||
|
"""Re-announcing the version the server already has must not split anything."""
|
||||||
|
version = self._current_version()
|
||||||
|
|
||||||
|
results = self._run_while_paused(
|
||||||
|
num_requests=4,
|
||||||
|
while_paused=lambda: self._set_weight_version(version),
|
||||||
|
mode="in_place",
|
||||||
|
max_new_tokens=128,
|
||||||
|
)
|
||||||
|
self.assertEqual(self._current_version(), version)
|
||||||
|
|
||||||
|
for data in results:
|
||||||
|
spans = _assert_spans_contiguous(self, data["meta_info"])
|
||||||
|
self.assertEqual(len(spans), 1)
|
||||||
|
self.assertEqual(spans[0]["version"], version)
|
||||||
|
|
||||||
|
def test_09_three_spans_across_two_updates(self):
|
||||||
|
"""Two updates during one request produce three ordered, non-empty spans."""
|
||||||
|
first_version = self._current_version()
|
||||||
|
|
||||||
|
with ThreadPoolExecutor(max_workers=4) as executor:
|
||||||
|
futures = [
|
||||||
|
executor.submit(
|
||||||
|
self._generate,
|
||||||
|
max_new_tokens=2048,
|
||||||
|
prompt=f"Write a long story about the number {i}.",
|
||||||
|
)
|
||||||
|
for i in range(4)
|
||||||
|
]
|
||||||
|
|
||||||
|
for new_version in ("multi-a", "multi-b"):
|
||||||
|
time.sleep(2)
|
||||||
|
self._pause("in_place")
|
||||||
|
try:
|
||||||
|
self._set_weight_version(new_version)
|
||||||
|
finally:
|
||||||
|
self._continue()
|
||||||
|
|
||||||
|
results = [future.result() for future in futures]
|
||||||
|
|
||||||
|
expected = [first_version, "multi-a", "multi-b"]
|
||||||
|
three_span_count = 0
|
||||||
|
for data in results:
|
||||||
|
meta_info = data["meta_info"]
|
||||||
|
spans = _assert_spans_contiguous(self, meta_info)
|
||||||
|
versions = [span["version"] for span in spans]
|
||||||
|
self.assertEqual(versions, expected[: len(versions)])
|
||||||
|
for span in spans:
|
||||||
|
self.assertGreater(span["end"], span["start"])
|
||||||
|
self.assertEqual(spans[-1]["end"], meta_info["completion_tokens"])
|
||||||
|
if len(spans) == 3:
|
||||||
|
three_span_count += 1
|
||||||
|
|
||||||
|
self.assertGreater(
|
||||||
|
three_span_count,
|
||||||
|
0,
|
||||||
|
"No request spanned both updates -- multi-event accumulation was not exercised.",
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_10_abort_before_first_token_reports_empty_span(self):
|
||||||
|
"""A request aborted before producing a token still reports a well-formed span."""
|
||||||
|
version = self._current_version()
|
||||||
|
|
||||||
|
def abort_all():
|
||||||
|
requests.post(
|
||||||
|
f"{self.base_url}/abort_request",
|
||||||
|
json={"abort_all": True},
|
||||||
|
timeout=30,
|
||||||
|
).raise_for_status()
|
||||||
|
|
||||||
|
results = self._run_while_paused(
|
||||||
|
num_requests=16,
|
||||||
|
while_paused=abort_all,
|
||||||
|
)
|
||||||
|
|
||||||
|
empty_aborts = [
|
||||||
|
data["meta_info"]
|
||||||
|
for data in results
|
||||||
|
if data["meta_info"]["finish_reason"]["type"] == "abort"
|
||||||
|
and data["meta_info"]["completion_tokens"] == 0
|
||||||
|
]
|
||||||
|
self.assertGreater(len(empty_aborts), 0)
|
||||||
|
for meta_info in empty_aborts:
|
||||||
|
self.assertEqual(
|
||||||
|
meta_info["weight_versions"],
|
||||||
|
[{"version": version, "start": 0, "end": 0}],
|
||||||
|
)
|
||||||
|
self.assertEqual(meta_info["weight_version"], version)
|
||||||
|
|
||||||
|
def test_11_streaming_reports_spans_only_on_the_final_chunk(self):
|
||||||
|
"""Intermediate stream chunks carry no spans; the finishing chunk carries them all."""
|
||||||
|
max_new_tokens = 32
|
||||||
|
response = requests.post(
|
||||||
|
f"{self.base_url}/generate",
|
||||||
|
json={
|
||||||
|
"text": "The capital of France is",
|
||||||
|
"sampling_params": {
|
||||||
|
"temperature": 0.8,
|
||||||
|
"max_new_tokens": max_new_tokens,
|
||||||
|
"ignore_eos": True,
|
||||||
|
},
|
||||||
|
"stream": True,
|
||||||
|
},
|
||||||
|
stream=True,
|
||||||
|
timeout=_REQUEST_TIMEOUT,
|
||||||
|
)
|
||||||
|
self.assertEqual(response.status_code, 200)
|
||||||
|
|
||||||
|
chunks = []
|
||||||
|
for line in response.iter_lines(decode_unicode=True):
|
||||||
|
if not line or not line.startswith("data:"):
|
||||||
|
continue
|
||||||
|
payload = line[len("data:") :].strip()
|
||||||
|
if payload == "[DONE]":
|
||||||
|
break
|
||||||
|
chunks.append(json.loads(payload))
|
||||||
|
|
||||||
|
version = self._current_version()
|
||||||
|
self.assertGreater(len(chunks), 1)
|
||||||
|
for chunk in chunks[:-1]:
|
||||||
|
self.assertNotIn("weight_versions", chunk["meta_info"])
|
||||||
|
self.assertEqual(chunk["meta_info"]["weight_version"], version)
|
||||||
|
|
||||||
|
meta_info = chunks[-1]["meta_info"]
|
||||||
|
spans = _assert_spans_contiguous(self, meta_info)
|
||||||
|
self.assertEqual(spans[-1]["end"], meta_info["completion_tokens"])
|
||||||
|
self.assertEqual(spans[-1]["end"], max_new_tokens)
|
||||||
|
|
||||||
|
def test_12_completions_endpoint_reports_metadata(self):
|
||||||
|
"""/v1/completions surfaces the spans the same way /v1/chat/completions does."""
|
||||||
|
response = requests.post(
|
||||||
|
f"{self.base_url}/v1/completions",
|
||||||
|
json={
|
||||||
|
"model": self.model,
|
||||||
|
"prompt": "The capital of France is",
|
||||||
|
"max_tokens": 8,
|
||||||
|
"temperature": 0.0,
|
||||||
|
},
|
||||||
|
timeout=_REQUEST_TIMEOUT,
|
||||||
|
)
|
||||||
|
self.assertEqual(response.status_code, 200)
|
||||||
|
|
||||||
|
data = response.json()
|
||||||
|
metadata = data["metadata"]
|
||||||
|
self.assertEqual(metadata["weight_version"], self._current_version())
|
||||||
|
spans = metadata["weight_versions"]
|
||||||
|
self.assertEqual(len(spans), 1)
|
||||||
|
self.assertEqual(spans[0]["version"], metadata["weight_version"])
|
||||||
|
self.assertEqual(spans[0]["start"], 0)
|
||||||
|
self.assertEqual(spans[0]["end"], data["usage"]["completion_tokens"])
|
||||||
|
|
||||||
|
def test_13_new_requests_after_the_endpoint_returns_see_the_new_version(self):
|
||||||
|
"""Once the endpoint returns, every concurrently admitted request stamps the new version."""
|
||||||
|
self._set_weight_version("dp-v1")
|
||||||
|
self.assertEqual(self._current_version(), "dp-v1")
|
||||||
|
|
||||||
|
with ThreadPoolExecutor(max_workers=8) as executor:
|
||||||
|
futures = [
|
||||||
|
executor.submit(self._generate, max_new_tokens=8) for _ in range(8)
|
||||||
|
]
|
||||||
|
results = [future.result() for future in futures]
|
||||||
|
|
||||||
|
for data in results:
|
||||||
|
spans = _assert_spans_contiguous(self, data["meta_info"])
|
||||||
|
self.assertEqual(len(spans), 1)
|
||||||
|
self.assertEqual(spans[0]["version"], "dp-v1")
|
||||||
|
|
||||||
|
def test_14_all_inflight_requests_are_swept(self):
|
||||||
|
"""A version change must split every in-flight request, not just the one the sweep happens to reach first."""
|
||||||
|
previous_version = self._current_version()
|
||||||
|
|
||||||
|
with ThreadPoolExecutor(max_workers=8) as executor:
|
||||||
|
futures = [
|
||||||
|
executor.submit(
|
||||||
|
self._generate,
|
||||||
|
max_new_tokens=1024 - 16 * i,
|
||||||
|
prompt=f"Write a long story about the number {i}.",
|
||||||
|
)
|
||||||
|
for i in range(8)
|
||||||
|
]
|
||||||
|
time.sleep(2)
|
||||||
|
self._set_weight_version("dp-v2")
|
||||||
|
results = [future.result() for future in futures]
|
||||||
|
|
||||||
|
split_count = 0
|
||||||
|
for data in results:
|
||||||
|
meta_info = data["meta_info"]
|
||||||
|
spans = _assert_spans_contiguous(self, meta_info)
|
||||||
|
self.assertEqual(spans[-1]["end"], meta_info["completion_tokens"])
|
||||||
|
self.assertEqual(spans[0]["version"], previous_version)
|
||||||
|
if len(spans) > 1:
|
||||||
|
split_count += 1
|
||||||
|
self.assertEqual(
|
||||||
|
[span["version"] for span in spans],
|
||||||
|
[previous_version, "dp-v2"],
|
||||||
|
)
|
||||||
|
self.assertGreater(spans[0]["end"], 0)
|
||||||
|
|
||||||
|
self.assertGreater(
|
||||||
|
split_count,
|
||||||
|
1,
|
||||||
|
"At most one in-flight request was split -- an attention-DP rank was "
|
||||||
|
"likely missed by the weight version sweep.",
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_15_single_span_covers_all_accepted_tokens(self):
|
||||||
|
"""Draft-token overshoot must not leak past the reported completion length."""
|
||||||
|
data = self._generate(max_new_tokens=32)
|
||||||
|
|
||||||
|
meta_info = data["meta_info"]
|
||||||
|
spans = _assert_spans_contiguous(self, meta_info)
|
||||||
|
self.assertEqual(len(spans), 1)
|
||||||
|
self.assertEqual(spans[0]["version"], self._current_version())
|
||||||
|
self.assertEqual(spans[0]["end"], meta_info["completion_tokens"])
|
||||||
|
|
||||||
|
def test_16_retract_update_continue_keeps_exact_boundaries(self):
|
||||||
|
"""Spans stay contiguous and clamped when a retract-paused update lands mid-speculation."""
|
||||||
|
previous_version = self._current_version()
|
||||||
|
results = self._run_while_paused(
|
||||||
|
num_requests=4,
|
||||||
|
while_paused=lambda: self._set_weight_version("spec-v1"),
|
||||||
|
)
|
||||||
|
|
||||||
|
split_count = 0
|
||||||
|
for data in results:
|
||||||
|
meta_info = data["meta_info"]
|
||||||
|
spans = _assert_spans_contiguous(self, meta_info)
|
||||||
|
self.assertEqual(spans[-1]["end"], meta_info["completion_tokens"])
|
||||||
|
self.assertEqual(spans[0]["version"], previous_version)
|
||||||
|
if len(spans) > 1:
|
||||||
|
split_count += 1
|
||||||
|
self.assertEqual(
|
||||||
|
[span["version"] for span in spans],
|
||||||
|
[previous_version, "spec-v1"],
|
||||||
|
)
|
||||||
|
self.assertGreater(spans[0]["end"], 0)
|
||||||
|
|
||||||
|
self.assertGreater(
|
||||||
|
split_count,
|
||||||
|
0,
|
||||||
|
"No request spanned the update -- the retract boundary under "
|
||||||
|
"speculative decoding was not recorded.",
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
@@ -15,6 +15,7 @@ import msgspec
|
|||||||
|
|
||||||
from sglang.srt.managers import io_struct
|
from sglang.srt.managers import io_struct
|
||||||
from sglang.srt.managers.io_struct import (
|
from sglang.srt.managers.io_struct import (
|
||||||
|
AbortReq,
|
||||||
BackupDramReq,
|
BackupDramReq,
|
||||||
ChecksumInfo,
|
ChecksumInfo,
|
||||||
CheckWeightsReqOutput,
|
CheckWeightsReqOutput,
|
||||||
@@ -37,6 +38,7 @@ from sglang.srt.model_executor.cuda_graph_config import CudaGraphConfig
|
|||||||
from sglang.srt.utils.msgspec_utils import msgspec_to_builtins
|
from sglang.srt.utils.msgspec_utils import msgspec_to_builtins
|
||||||
from sglang.srt.utils.weight_checker import ChecksumInfo as PydanticChecksumInfo
|
from sglang.srt.utils.weight_checker import ChecksumInfo as PydanticChecksumInfo
|
||||||
from sglang.srt.utils.weight_checker import ParallelismInfo as PydanticParallelismInfo
|
from sglang.srt.utils.weight_checker import ParallelismInfo as PydanticParallelismInfo
|
||||||
|
from sglang.srt.utils.weight_versions import WeightVersionSpan
|
||||||
from sglang.test.ci.ci_register import register_cpu_ci
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
from sglang.test.test_utils import CustomTestCase
|
from sglang.test.test_utils import CustomTestCase
|
||||||
|
|
||||||
@@ -237,5 +239,34 @@ class TestMsgpackIpcRoundtrip(CustomTestCase):
|
|||||||
self.assertEqual(_round_trip(output), output)
|
self.assertEqual(_round_trip(output), output)
|
||||||
|
|
||||||
|
|
||||||
|
class TestWeightVersionSpansRoundTrip(CustomTestCase):
|
||||||
|
"""The per-request weight-version spans ride the same msgpack IPC path."""
|
||||||
|
|
||||||
|
def test_abort_req_carries_spans(self):
|
||||||
|
"""A scheduler-side abort keeps its spans across the wire."""
|
||||||
|
obj = AbortReq(
|
||||||
|
rid="r0",
|
||||||
|
weight_versions=[
|
||||||
|
WeightVersionSpan(version="v1", start=0, end=3),
|
||||||
|
WeightVersionSpan(version="v2", start=3, end=7),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
|
||||||
|
decoded = _double_hop(obj)
|
||||||
|
|
||||||
|
self.assertEqual(
|
||||||
|
decoded.weight_versions,
|
||||||
|
[
|
||||||
|
WeightVersionSpan(version="v1", start=0, end=3),
|
||||||
|
WeightVersionSpan(version="v2", start=3, end=7),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
self.assertIsInstance(decoded.weight_versions[0], WeightVersionSpan)
|
||||||
|
|
||||||
|
def test_abort_req_defaults_to_no_spans(self):
|
||||||
|
"""The field is optional on the wire, so an abort without spans decodes to None."""
|
||||||
|
self.assertIsNone(_round_trip(AbortReq(rid="r0")).weight_versions)
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|||||||
@@ -1,5 +1,6 @@
|
|||||||
import unittest
|
import unittest
|
||||||
|
|
||||||
|
from sglang.srt.utils.weight_versions import WeightVersionSpan
|
||||||
from sglang.test.ci.ci_register import register_cpu_ci
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
from sglang.test.test_utils import maybe_stub_sgl_kernel
|
from sglang.test.test_utils import maybe_stub_sgl_kernel
|
||||||
|
|
||||||
@@ -76,6 +77,13 @@ def _make_batch_str_output() -> BatchStrOutput:
|
|||||||
placeholder_tokens_idx=[None, None],
|
placeholder_tokens_idx=[None, None],
|
||||||
placeholder_tokens_val=[None, None],
|
placeholder_tokens_val=[None, None],
|
||||||
retraction_counts=[0, 0],
|
retraction_counts=[0, 0],
|
||||||
|
weight_versions=[
|
||||||
|
[
|
||||||
|
WeightVersionSpan(version="v1", start=0, end=3),
|
||||||
|
WeightVersionSpan(version="v2", start=3, end=5),
|
||||||
|
],
|
||||||
|
[WeightVersionSpan(version="v2", start=0, end=2)],
|
||||||
|
],
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
@@ -92,6 +100,31 @@ class TestMultiTokenizerMixin(unittest.TestCase):
|
|||||||
[{"device": 1, "host": 3}],
|
[{"device": 1, "host": 3}],
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def test_batch_str_output_keeps_weight_versions_nested_per_request(self):
|
||||||
|
"""Per-request segment lists stay one level nested after the split."""
|
||||||
|
output = _make_batch_str_output()
|
||||||
|
|
||||||
|
self.assertEqual(
|
||||||
|
_handle_output_by_index(output, 0).weight_versions,
|
||||||
|
[
|
||||||
|
[
|
||||||
|
WeightVersionSpan(version="v1", start=0, end=3),
|
||||||
|
WeightVersionSpan(version="v2", start=3, end=5),
|
||||||
|
]
|
||||||
|
],
|
||||||
|
)
|
||||||
|
self.assertEqual(
|
||||||
|
_handle_output_by_index(output, 1).weight_versions,
|
||||||
|
[[WeightVersionSpan(version="v2", start=0, end=2)]],
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_batch_str_output_without_weight_versions_stays_none(self):
|
||||||
|
"""An output from an older server without the field splits into None."""
|
||||||
|
output = _make_batch_str_output()
|
||||||
|
output.weight_versions = None
|
||||||
|
|
||||||
|
self.assertIsNone(_handle_output_by_index(output, 0).weight_versions)
|
||||||
|
|
||||||
def test_get_tokenizer_worker_class_uses_default(self):
|
def test_get_tokenizer_worker_class_uses_default(self):
|
||||||
self.assertIs(get_tokenizer_worker_class(DefaultServerArgs()), TokenizerWorker)
|
self.assertIs(get_tokenizer_worker_class(DefaultServerArgs()), TokenizerWorker)
|
||||||
|
|
||||||
|
|||||||
@@ -9,6 +9,10 @@ from sglang.srt.managers.scheduler_components.output_streamer import (
|
|||||||
_GenerationStreamAccumulator,
|
_GenerationStreamAccumulator,
|
||||||
)
|
)
|
||||||
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
|
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
|
||||||
|
from sglang.srt.utils.weight_versions import (
|
||||||
|
WeightVersionSpan,
|
||||||
|
record_weight_version_events,
|
||||||
|
)
|
||||||
from sglang.test.ci.ci_register import register_cpu_ci
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
|
|
||||||
register_cpu_ci(est_time=1, suite="base-a-test-cpu")
|
register_cpu_ci(est_time=1, suite="base-a-test-cpu")
|
||||||
@@ -58,6 +62,7 @@ class _FakeReq:
|
|||||||
self.mm_video_tokens = 0
|
self.mm_video_tokens = 0
|
||||||
self.multimodal_inputs = None
|
self.multimodal_inputs = None
|
||||||
self.customized_info = customized_info
|
self.customized_info = customized_info
|
||||||
|
self.weight_version_events = []
|
||||||
|
|
||||||
def finished(self):
|
def finished(self):
|
||||||
return self._finished
|
return self._finished
|
||||||
@@ -69,11 +74,26 @@ class _FakeReq:
|
|||||||
return False
|
return False
|
||||||
|
|
||||||
|
|
||||||
|
def _accumulator(current_weight_version="default"):
|
||||||
|
return _GenerationStreamAccumulator(
|
||||||
|
return_logprob=False,
|
||||||
|
return_hidden_states=False,
|
||||||
|
return_routed_experts=False,
|
||||||
|
return_indexer_topk=False,
|
||||||
|
spec_algorithm=SpeculativeAlgorithm.NONE,
|
||||||
|
disaggregation_mode=DisaggregationMode.NULL,
|
||||||
|
default_stream_interval=1,
|
||||||
|
default_force_stream_interval=1,
|
||||||
|
get_cached_tokens_details=lambda req: None,
|
||||||
|
current_weight_version=current_weight_version,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
class TestOutputStreamerCustomizedInfo(unittest.TestCase):
|
class TestOutputStreamerCustomizedInfo(unittest.TestCase):
|
||||||
def setUp(self):
|
def setUp(self):
|
||||||
serving_patch = patch(
|
serving_patch = patch(
|
||||||
"sglang.srt.managers.scheduler_components.output_streamer.get_serving",
|
"sglang.srt.managers.scheduler_components.output_streamer.get_serving",
|
||||||
return_value=SimpleNamespace(stream_interval=1),
|
return_value=SimpleNamespace(stream_interval=1, weight_version="default"),
|
||||||
)
|
)
|
||||||
observability_patch = patch(
|
observability_patch = patch(
|
||||||
"sglang.srt.managers.scheduler_components.output_streamer.get_observability",
|
"sglang.srt.managers.scheduler_components.output_streamer.get_observability",
|
||||||
@@ -84,22 +104,8 @@ class TestOutputStreamerCustomizedInfo(unittest.TestCase):
|
|||||||
self.addCleanup(serving_patch.stop)
|
self.addCleanup(serving_patch.stop)
|
||||||
self.addCleanup(observability_patch.stop)
|
self.addCleanup(observability_patch.stop)
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
def _accumulator():
|
|
||||||
return _GenerationStreamAccumulator(
|
|
||||||
return_logprob=False,
|
|
||||||
return_hidden_states=False,
|
|
||||||
return_routed_experts=False,
|
|
||||||
return_indexer_topk=False,
|
|
||||||
spec_algorithm=SpeculativeAlgorithm.NONE,
|
|
||||||
disaggregation_mode=DisaggregationMode.NULL,
|
|
||||||
default_stream_interval=1,
|
|
||||||
default_force_stream_interval=1,
|
|
||||||
get_cached_tokens_details=lambda req: None,
|
|
||||||
)
|
|
||||||
|
|
||||||
def test_customized_info_is_padded_for_mixed_batches(self):
|
def test_customized_info_is_padded_for_mixed_batches(self):
|
||||||
accumulator = self._accumulator()
|
accumulator = _accumulator()
|
||||||
|
|
||||||
accumulator.accept(req=_FakeReq("r0", [10, 11]))
|
accumulator.accept(req=_FakeReq("r0", [10, 11]))
|
||||||
accumulator.accept(
|
accumulator.accept(
|
||||||
@@ -334,5 +340,38 @@ class TestOutputStreamerCustomizedInfo(unittest.TestCase):
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class TestOutputStreamerWeightVersions(unittest.TestCase):
|
||||||
|
def test_payload_carries_spans_for_finished_requests(self):
|
||||||
|
"""Finished requests report their spans; still-generating ones report nothing."""
|
||||||
|
streaming_req = _FakeReq("r0", [10, 11])
|
||||||
|
finished_req = _FakeReq("r1", [20, 21, 22], finished=True)
|
||||||
|
record_weight_version_events([finished_req], old_version="v1")
|
||||||
|
finished_req.output_ids.extend([23, 24])
|
||||||
|
|
||||||
|
accumulator = _accumulator(current_weight_version="v2")
|
||||||
|
accumulator.accept(req=streaming_req)
|
||||||
|
accumulator.accept(req=finished_req)
|
||||||
|
payload = accumulator.to_payload(dp_rank=0, is_idle_batch=False)
|
||||||
|
|
||||||
|
self.assertEqual(
|
||||||
|
payload.weight_versions,
|
||||||
|
[
|
||||||
|
None,
|
||||||
|
[
|
||||||
|
WeightVersionSpan(version="v1", start=0, end=3),
|
||||||
|
WeightVersionSpan(version="v2", start=3, end=5),
|
||||||
|
],
|
||||||
|
],
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_payload_omits_spans_while_all_requests_stream(self):
|
||||||
|
"""A batch of unfinished requests puts nothing on the wire."""
|
||||||
|
accumulator = _accumulator(current_weight_version="v2")
|
||||||
|
accumulator.accept(req=_FakeReq("r0", [10]))
|
||||||
|
payload = accumulator.to_payload(dp_rank=0, is_idle_batch=False)
|
||||||
|
|
||||||
|
self.assertIsNone(payload.weight_versions)
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|||||||
@@ -73,6 +73,7 @@ def _make_accumulator() -> _GenerationStreamAccumulator:
|
|||||||
default_stream_interval=1,
|
default_stream_interval=1,
|
||||||
default_force_stream_interval=1,
|
default_force_stream_interval=1,
|
||||||
get_cached_tokens_details=lambda req: None,
|
get_cached_tokens_details=lambda req: None,
|
||||||
|
current_weight_version=None,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -41,6 +41,8 @@ class TestDisaggregationPriorityQueueing(unittest.TestCase):
|
|||||||
req = MagicMock()
|
req = MagicMock()
|
||||||
req.priority = priority
|
req.priority = priority
|
||||||
req.rid = "req"
|
req.rid = "req"
|
||||||
|
req.output_ids = []
|
||||||
|
req.weight_version_events = []
|
||||||
req.time_stats = MagicMock()
|
req.time_stats = MagicMock()
|
||||||
req.time_stats.trace_ctx = MagicMock()
|
req.time_stats.trace_ctx = MagicMock()
|
||||||
return req
|
return req
|
||||||
@@ -73,7 +75,11 @@ class TestDisaggregationPriorityQueueing(unittest.TestCase):
|
|||||||
scheduler.abort_on_priority_when_disabled = True
|
scheduler.abort_on_priority_when_disabled = True
|
||||||
req = self._new_req(priority=10)
|
req = self._new_req(priority=10)
|
||||||
|
|
||||||
scheduler._add_request_to_queue(req)
|
with patch(
|
||||||
|
"sglang.srt.managers.scheduler.get_serving",
|
||||||
|
return_value=SimpleNamespace(weight_version="v0"),
|
||||||
|
):
|
||||||
|
scheduler._add_request_to_queue(req)
|
||||||
|
|
||||||
scheduler.disagg_decode_prealloc_queue.add.assert_not_called()
|
scheduler.disagg_decode_prealloc_queue.add.assert_not_called()
|
||||||
scheduler.ipc_channels.send_to_tokenizer.send_output.assert_called_once()
|
scheduler.ipc_channels.send_to_tokenizer.send_output.assert_called_once()
|
||||||
|
|||||||
@@ -9,7 +9,7 @@ scheduler/test_scheduler_control.py.
|
|||||||
import time
|
import time
|
||||||
import unittest
|
import unittest
|
||||||
from types import SimpleNamespace
|
from types import SimpleNamespace
|
||||||
from unittest.mock import MagicMock
|
from unittest.mock import MagicMock, patch
|
||||||
|
|
||||||
from sglang.srt.environ import envs
|
from sglang.srt.environ import envs
|
||||||
from sglang.test.ci.ci_register import register_cpu_ci
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
@@ -29,6 +29,8 @@ class _FakeReq:
|
|||||||
self.rid = rid
|
self.rid = rid
|
||||||
self.to_finish = None
|
self.to_finish = None
|
||||||
self._finished = is_finished
|
self._finished = is_finished
|
||||||
|
self.output_ids = []
|
||||||
|
self.weight_version_events = []
|
||||||
self.time_stats = SimpleNamespace(
|
self.time_stats = SimpleNamespace(
|
||||||
wait_queue_entry_time=wait_entry,
|
wait_queue_entry_time=wait_entry,
|
||||||
forward_entry_time=forward_entry,
|
forward_entry_time=forward_entry,
|
||||||
@@ -54,6 +56,14 @@ def _scheduler(waiting_queue):
|
|||||||
|
|
||||||
|
|
||||||
class TestWaitingTimeout(CustomTestCase):
|
class TestWaitingTimeout(CustomTestCase):
|
||||||
|
def setUp(self):
|
||||||
|
patcher = patch(
|
||||||
|
"sglang.srt.managers.scheduler.get_serving",
|
||||||
|
return_value=SimpleNamespace(weight_version="v0"),
|
||||||
|
)
|
||||||
|
patcher.start()
|
||||||
|
self.addCleanup(patcher.stop)
|
||||||
|
|
||||||
def test_drops_only_reqs_past_the_deadline(self):
|
def test_drops_only_reqs_past_the_deadline(self):
|
||||||
now = time.perf_counter()
|
now = time.perf_counter()
|
||||||
stale = _req("stale", wait_entry=now - 10)
|
stale = _req("stale", wait_entry=now - 10)
|
||||||
|
|||||||
@@ -37,11 +37,27 @@ class TestSchedulerRecordWeightVersionChange(CustomTestCase):
|
|||||||
self.addCleanup(patcher.stop)
|
self.addCleanup(patcher.stop)
|
||||||
return serving
|
return serving
|
||||||
|
|
||||||
|
def _scheduler(
|
||||||
|
self, *, inflight=(), waiting=(), chunked=None, staging=()
|
||||||
|
) -> SimpleNamespace:
|
||||||
|
return SimpleNamespace(
|
||||||
|
collect_inflight_reqs=lambda: set(inflight),
|
||||||
|
waiting_queue=list(waiting),
|
||||||
|
chunked_req=chunked,
|
||||||
|
hisparse_coordinator=(
|
||||||
|
SimpleNamespace(
|
||||||
|
ack_staging_queue=[SimpleNamespace(req=req) for req in staging]
|
||||||
|
)
|
||||||
|
if staging
|
||||||
|
else None
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
def test_a_new_version_is_adopted(self):
|
def test_a_new_version_is_adopted(self):
|
||||||
"""The scheduler has to end up on the version it was told about, or nothing downstream can read it."""
|
"""The scheduler has to end up on the version it was told about, or nothing downstream can read it."""
|
||||||
serving = self._serving("v1")
|
serving = self._serving("v1")
|
||||||
|
|
||||||
Scheduler.record_weight_version_change(SimpleNamespace(), new_version="v2")
|
Scheduler.record_weight_version_change(self._scheduler(), new_version="v2")
|
||||||
|
|
||||||
self.assertEqual(serving.weight_version, "v2")
|
self.assertEqual(serving.weight_version, "v2")
|
||||||
|
|
||||||
@@ -49,7 +65,7 @@ class TestSchedulerRecordWeightVersionChange(CustomTestCase):
|
|||||||
"""Re-announcing the current version must not be treated as a change."""
|
"""Re-announcing the current version must not be treated as a change."""
|
||||||
serving = self._serving("v1")
|
serving = self._serving("v1")
|
||||||
|
|
||||||
Scheduler.record_weight_version_change(SimpleNamespace(), new_version="v1")
|
Scheduler.record_weight_version_change(self._scheduler(), new_version="v1")
|
||||||
|
|
||||||
self.assertEqual(serving.weight_version, "v1")
|
self.assertEqual(serving.weight_version, "v1")
|
||||||
|
|
||||||
@@ -57,10 +73,28 @@ class TestSchedulerRecordWeightVersionChange(CustomTestCase):
|
|||||||
"""An update that carries no version must leave the recorded one alone."""
|
"""An update that carries no version must leave the recorded one alone."""
|
||||||
serving = self._serving("v1")
|
serving = self._serving("v1")
|
||||||
|
|
||||||
Scheduler.record_weight_version_change(SimpleNamespace(), new_version=None)
|
Scheduler.record_weight_version_change(self._scheduler(), new_version=None)
|
||||||
|
|
||||||
self.assertEqual(serving.weight_version, "v1")
|
self.assertEqual(serving.weight_version, "v1")
|
||||||
|
|
||||||
|
def test_every_source_of_live_requests_is_stamped(self):
|
||||||
|
"""A request missed here keeps attributing its next tokens to the superseded version."""
|
||||||
|
self._serving("v1")
|
||||||
|
inflight, queued, chunked, staged = (object() for _ in range(4))
|
||||||
|
scheduler = self._scheduler(
|
||||||
|
inflight=[inflight], waiting=[queued], chunked=chunked, staging=[staged]
|
||||||
|
)
|
||||||
|
|
||||||
|
with patch(
|
||||||
|
"sglang.srt.managers.scheduler.record_weight_version_events",
|
||||||
|
return_value=0,
|
||||||
|
) as recorder:
|
||||||
|
Scheduler.record_weight_version_change(scheduler, new_version="v2")
|
||||||
|
|
||||||
|
self.assertEqual(
|
||||||
|
set(recorder.call_args.args[0]), {inflight, queued, chunked, staged}
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
class TestRecordWeightVersionAfterUpdate(CustomTestCase):
|
class TestRecordWeightVersionAfterUpdate(CustomTestCase):
|
||||||
def _updater(
|
def _updater(
|
||||||
|
|||||||
@@ -0,0 +1,629 @@
|
|||||||
|
import random
|
||||||
|
import unittest
|
||||||
|
from types import SimpleNamespace
|
||||||
|
from unittest.mock import patch
|
||||||
|
|
||||||
|
from sglang.srt.managers.io_struct import AbortReq
|
||||||
|
from sglang.srt.managers.scheduler import Scheduler, _make_abort_req
|
||||||
|
from sglang.srt.utils.weight_versions import (
|
||||||
|
WeightVersionEvent,
|
||||||
|
WeightVersionSpan,
|
||||||
|
add_weight_versions_to_meta_info,
|
||||||
|
build_endpoint_weight_version_metadata,
|
||||||
|
compute_weight_version_spans,
|
||||||
|
record_weight_version_events,
|
||||||
|
truncate_weight_version_events,
|
||||||
|
)
|
||||||
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
|
from sglang.test.test_utils import CustomTestCase
|
||||||
|
|
||||||
|
register_cpu_ci(est_time=4, suite="base-a-test-cpu")
|
||||||
|
|
||||||
|
|
||||||
|
class _ReqStub:
|
||||||
|
def __init__(self, output_len: int):
|
||||||
|
self.output_ids = [0] * output_len
|
||||||
|
self.weight_version_events = []
|
||||||
|
|
||||||
|
def record_weight_version_change(self, old_version):
|
||||||
|
record_weight_version_events([self], old_version=old_version)
|
||||||
|
|
||||||
|
def compute_weight_version_spans(self, current_version, num_output_tokens):
|
||||||
|
return compute_weight_version_spans(
|
||||||
|
self.weight_version_events,
|
||||||
|
current_version=current_version,
|
||||||
|
num_output_tokens=num_output_tokens,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _expected_spans(events, current_version, num_output_tokens):
|
||||||
|
"""Reference model: name the version that sampled each token, then run-length encode."""
|
||||||
|
per_token = []
|
||||||
|
for index in range(num_output_tokens):
|
||||||
|
owner = next(
|
||||||
|
(event.old_version for event in events if event.num_output_tokens > index),
|
||||||
|
current_version,
|
||||||
|
)
|
||||||
|
per_token.append(owner)
|
||||||
|
|
||||||
|
if not per_token:
|
||||||
|
first_event_end_at_zero = next(
|
||||||
|
(event for event in events if event.num_output_tokens >= 0), None
|
||||||
|
)
|
||||||
|
version = (
|
||||||
|
first_event_end_at_zero.old_version
|
||||||
|
if first_event_end_at_zero is not None
|
||||||
|
else current_version
|
||||||
|
)
|
||||||
|
return [WeightVersionSpan(version=version, start=0, end=0)]
|
||||||
|
|
||||||
|
spans = []
|
||||||
|
for index, version in enumerate(per_token):
|
||||||
|
if spans and spans[-1].version == version:
|
||||||
|
spans[-1].end = index + 1
|
||||||
|
else:
|
||||||
|
spans.append(WeightVersionSpan(version=version, start=index, end=index + 1))
|
||||||
|
return spans
|
||||||
|
|
||||||
|
|
||||||
|
class TestComputeWeightVersionSpans(CustomTestCase):
|
||||||
|
def test_no_events_returns_single_span(self):
|
||||||
|
"""A request untouched by updates gets one span covering all output tokens."""
|
||||||
|
self.assertEqual(
|
||||||
|
_ReqStub(5).compute_weight_version_spans(
|
||||||
|
current_version="v1", num_output_tokens=5
|
||||||
|
),
|
||||||
|
[WeightVersionSpan(version="v1", start=0, end=5)],
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_zero_output_tokens_returns_empty_span_span(self):
|
||||||
|
"""A request finishing with no output still reports the current version."""
|
||||||
|
self.assertEqual(
|
||||||
|
_ReqStub(0).compute_weight_version_spans(
|
||||||
|
current_version="v1", num_output_tokens=0
|
||||||
|
),
|
||||||
|
[WeightVersionSpan(version="v1", start=0, end=0)],
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_one_update_splits_into_two_spans(self):
|
||||||
|
"""An update at 3 output tokens attributes [0,3) to the old version and the rest to the new."""
|
||||||
|
req = _ReqStub(3)
|
||||||
|
req.record_weight_version_change(old_version="v1")
|
||||||
|
req.output_ids.extend([0] * 4)
|
||||||
|
self.assertEqual(
|
||||||
|
req.compute_weight_version_spans(current_version="v2", num_output_tokens=7),
|
||||||
|
[
|
||||||
|
WeightVersionSpan(version="v1", start=0, end=3),
|
||||||
|
WeightVersionSpan(version="v2", start=3, end=7),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_two_updates_split_into_three_spans(self):
|
||||||
|
"""Each update the request lives through adds one span."""
|
||||||
|
req = _ReqStub(2)
|
||||||
|
req.record_weight_version_change(old_version="v1")
|
||||||
|
req.output_ids.extend([0] * 3)
|
||||||
|
req.record_weight_version_change(old_version="v2")
|
||||||
|
req.output_ids.extend([0] * 1)
|
||||||
|
self.assertEqual(
|
||||||
|
req.compute_weight_version_spans(current_version="v3", num_output_tokens=6),
|
||||||
|
[
|
||||||
|
WeightVersionSpan(version="v1", start=0, end=2),
|
||||||
|
WeightVersionSpan(version="v2", start=2, end=5),
|
||||||
|
WeightVersionSpan(version="v3", start=5, end=6),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_update_before_first_token_records_no_event(self):
|
||||||
|
"""An update while the request has no output leaves the whole output on the new version."""
|
||||||
|
req = _ReqStub(0)
|
||||||
|
req.record_weight_version_change(old_version="v1")
|
||||||
|
self.assertEqual(req.weight_version_events, [])
|
||||||
|
req.output_ids.extend([0] * 4)
|
||||||
|
self.assertEqual(
|
||||||
|
req.compute_weight_version_spans(current_version="v2", num_output_tokens=4),
|
||||||
|
[WeightVersionSpan(version="v2", start=0, end=4)],
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_back_to_back_updates_skip_empty_span(self):
|
||||||
|
"""Two updates with no tokens in between produce no empty span."""
|
||||||
|
req = _ReqStub(2)
|
||||||
|
req.record_weight_version_change(old_version="v1")
|
||||||
|
req.record_weight_version_change(old_version="v2")
|
||||||
|
req.output_ids.extend([0] * 2)
|
||||||
|
self.assertEqual(
|
||||||
|
req.compute_weight_version_spans(current_version="v3", num_output_tokens=4),
|
||||||
|
[
|
||||||
|
WeightVersionSpan(version="v1", start=0, end=2),
|
||||||
|
WeightVersionSpan(version="v3", start=2, end=4),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_update_at_final_length_yields_no_trailing_span(self):
|
||||||
|
"""An event recorded exactly at the final output length adds no empty trailing span."""
|
||||||
|
req = _ReqStub(4)
|
||||||
|
req.record_weight_version_change(old_version="v1")
|
||||||
|
self.assertEqual(
|
||||||
|
req.compute_weight_version_spans(current_version="v2", num_output_tokens=4),
|
||||||
|
[WeightVersionSpan(version="v1", start=0, end=4)],
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_event_beyond_reported_length_is_clamped(self):
|
||||||
|
"""A spec-decode overshoot event is clamped to the reported output length."""
|
||||||
|
req = _ReqStub(6)
|
||||||
|
req.record_weight_version_change(old_version="v1")
|
||||||
|
self.assertEqual(
|
||||||
|
req.compute_weight_version_spans(current_version="v2", num_output_tokens=4),
|
||||||
|
[WeightVersionSpan(version="v1", start=0, end=4)],
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_clamp_to_zero_reports_the_pre_update_version(self):
|
||||||
|
"""With no visible tokens the single span keeps the version that sampled them."""
|
||||||
|
req = _ReqStub(3)
|
||||||
|
req.record_weight_version_change(old_version="v1")
|
||||||
|
self.assertEqual(
|
||||||
|
req.compute_weight_version_spans(current_version="v2", num_output_tokens=0),
|
||||||
|
[WeightVersionSpan(version="v1", start=0, end=0)],
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_clamped_output_can_end_before_the_current_version(self):
|
||||||
|
"""When the visible tokens stop early the newest version never appears."""
|
||||||
|
req = _ReqStub(3)
|
||||||
|
req.record_weight_version_change(old_version="v1")
|
||||||
|
req.output_ids.extend([0] * 2)
|
||||||
|
req.record_weight_version_change(old_version="v2")
|
||||||
|
self.assertEqual(
|
||||||
|
req.compute_weight_version_spans(current_version="v3", num_output_tokens=4),
|
||||||
|
[
|
||||||
|
WeightVersionSpan(version="v1", start=0, end=3),
|
||||||
|
WeightVersionSpan(version="v2", start=3, end=4),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_clamping_collapses_events_and_merges_into_one_span(self):
|
||||||
|
"""Events beyond the visible range collapse instead of leaving empty spans."""
|
||||||
|
req = _ReqStub(3)
|
||||||
|
req.record_weight_version_change(old_version="v1")
|
||||||
|
req.output_ids.extend([0] * 2)
|
||||||
|
req.record_weight_version_change(old_version="v2")
|
||||||
|
self.assertEqual(
|
||||||
|
req.compute_weight_version_spans(current_version="v1", num_output_tokens=2),
|
||||||
|
[WeightVersionSpan(version="v1", start=0, end=2)],
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_duplicate_event_at_the_same_length_is_idempotent(self):
|
||||||
|
"""Recording twice at the same token count matches recording once."""
|
||||||
|
req = _ReqStub(3)
|
||||||
|
req.record_weight_version_change(old_version="v1")
|
||||||
|
req.record_weight_version_change(old_version="v1")
|
||||||
|
req.output_ids.extend([0] * 2)
|
||||||
|
self.assertEqual(
|
||||||
|
req.compute_weight_version_spans(current_version="v2", num_output_tokens=5),
|
||||||
|
[
|
||||||
|
WeightVersionSpan(version="v1", start=0, end=3),
|
||||||
|
WeightVersionSpan(version="v2", start=3, end=5),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_version_returning_after_tokens_keeps_separate_spans(self):
|
||||||
|
"""A version that comes back after another one sampled tokens gets its own span."""
|
||||||
|
req = _ReqStub(2)
|
||||||
|
req.record_weight_version_change(old_version="v1")
|
||||||
|
req.output_ids.extend([0] * 2)
|
||||||
|
req.record_weight_version_change(old_version="v2")
|
||||||
|
req.output_ids.extend([0] * 2)
|
||||||
|
self.assertEqual(
|
||||||
|
req.compute_weight_version_spans(current_version="v1", num_output_tokens=6),
|
||||||
|
[
|
||||||
|
WeightVersionSpan(version="v1", start=0, end=2),
|
||||||
|
WeightVersionSpan(version="v2", start=2, end=4),
|
||||||
|
WeightVersionSpan(version="v1", start=4, end=6),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_spans_are_never_empty_even_when_everything_is_clamped_away(self):
|
||||||
|
"""add_weight_versions_to_meta_info reads the newest span by index, so an empty list would crash it."""
|
||||||
|
req = _ReqStub(3)
|
||||||
|
req.record_weight_version_change(old_version="v1")
|
||||||
|
req.record_weight_version_change(old_version="v2")
|
||||||
|
|
||||||
|
for current_version in ("v1", "v3"):
|
||||||
|
self.assertTrue(
|
||||||
|
req.compute_weight_version_spans(
|
||||||
|
current_version=current_version, num_output_tokens=0
|
||||||
|
)
|
||||||
|
)
|
||||||
|
self.assertTrue(
|
||||||
|
_ReqStub(0).compute_weight_version_spans(
|
||||||
|
current_version="v1", num_output_tokens=0
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_spans_satisfy_the_contract_for_random_event_sequences(self):
|
||||||
|
"""Randomized event sequences always yield ordered, contiguous, non-duplicated spans."""
|
||||||
|
rng = random.Random(0)
|
||||||
|
versions = ["v0", "v1", "v2"]
|
||||||
|
for _ in range(300):
|
||||||
|
counts = sorted(rng.randint(1, 8) for _ in range(rng.randint(0, 5)))
|
||||||
|
events = [
|
||||||
|
WeightVersionEvent(
|
||||||
|
old_version=rng.choice(versions), num_output_tokens=count
|
||||||
|
)
|
||||||
|
for count in counts
|
||||||
|
]
|
||||||
|
current_version = rng.choice(versions)
|
||||||
|
num_output_tokens = rng.randint(0, 10)
|
||||||
|
spans = compute_weight_version_spans(
|
||||||
|
events,
|
||||||
|
current_version=current_version,
|
||||||
|
num_output_tokens=num_output_tokens,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(
|
||||||
|
spans,
|
||||||
|
_expected_spans(events, current_version, num_output_tokens),
|
||||||
|
)
|
||||||
|
self.assertTrue(spans)
|
||||||
|
self.assertEqual(spans[0].start, 0)
|
||||||
|
self.assertEqual(spans[-1].end, num_output_tokens)
|
||||||
|
for previous, current in zip(spans, spans[1:]):
|
||||||
|
self.assertEqual(previous.end, current.start)
|
||||||
|
self.assertNotEqual(previous.version, current.version)
|
||||||
|
if num_output_tokens == 0:
|
||||||
|
self.assertEqual(len(spans), 1)
|
||||||
|
self.assertEqual(spans[0].end, 0)
|
||||||
|
else:
|
||||||
|
for span in spans:
|
||||||
|
self.assertLess(span.start, span.end)
|
||||||
|
|
||||||
|
|
||||||
|
class _ServingStub:
|
||||||
|
def __init__(self, weight_version):
|
||||||
|
self.weight_version = weight_version
|
||||||
|
|
||||||
|
|
||||||
|
class _ContextStub:
|
||||||
|
def __init__(self, serving: _ServingStub):
|
||||||
|
self.serving = serving
|
||||||
|
|
||||||
|
def override(self, source, **fields):
|
||||||
|
self.serving.weight_version = fields["weight_version"]
|
||||||
|
|
||||||
|
|
||||||
|
class _SchedulerStub:
|
||||||
|
collect_inflight_reqs = Scheduler.collect_inflight_reqs
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
version,
|
||||||
|
running,
|
||||||
|
waiting,
|
||||||
|
chunked=None,
|
||||||
|
last_batch=None,
|
||||||
|
pp_size=1,
|
||||||
|
hisparse=None,
|
||||||
|
):
|
||||||
|
self.serving = _ServingStub(version)
|
||||||
|
self.ps = SimpleNamespace(pp_size=pp_size)
|
||||||
|
self.running_batch = SimpleNamespace(reqs=running)
|
||||||
|
self.last_batch = last_batch
|
||||||
|
self.waiting_queue = waiting
|
||||||
|
self.chunked_req = chunked
|
||||||
|
self.hisparse_coordinator = hisparse
|
||||||
|
|
||||||
|
|
||||||
|
class TestSchedulerRecordWeightVersionChange(CustomTestCase):
|
||||||
|
def _scheduler(self, *args, **kwargs):
|
||||||
|
scheduler = _SchedulerStub(*args, **kwargs)
|
||||||
|
for name, value in (
|
||||||
|
("get_serving", scheduler.serving),
|
||||||
|
("get_context", _ContextStub(scheduler.serving)),
|
||||||
|
):
|
||||||
|
patcher = patch(f"sglang.srt.managers.scheduler.{name}", return_value=value)
|
||||||
|
patcher.start()
|
||||||
|
self.addCleanup(patcher.stop)
|
||||||
|
return scheduler
|
||||||
|
|
||||||
|
def test_records_on_running_waiting_and_chunked_requests(self):
|
||||||
|
"""A version change records an event on every live request holding output tokens."""
|
||||||
|
running_req = _ReqStub(3)
|
||||||
|
waiting_req = _ReqStub(5)
|
||||||
|
chunked_req = _ReqStub(2)
|
||||||
|
scheduler = self._scheduler(
|
||||||
|
"v1", [running_req], [waiting_req], chunked=chunked_req
|
||||||
|
)
|
||||||
|
|
||||||
|
Scheduler.record_weight_version_change(scheduler, new_version="v2")
|
||||||
|
|
||||||
|
self.assertEqual(scheduler.serving.weight_version, "v2")
|
||||||
|
for req, output_len in (
|
||||||
|
(running_req, 3),
|
||||||
|
(waiting_req, 5),
|
||||||
|
(chunked_req, 2),
|
||||||
|
):
|
||||||
|
self.assertEqual(len(req.weight_version_events), 1)
|
||||||
|
self.assertEqual(req.weight_version_events[0].old_version, "v1")
|
||||||
|
self.assertEqual(req.weight_version_events[0].num_output_tokens, output_len)
|
||||||
|
|
||||||
|
def test_records_on_last_batch_requests_exactly_once(self):
|
||||||
|
"""A request in both the running and last batch is recorded once, and last-batch-only requests are recorded too."""
|
||||||
|
shared_req = _ReqStub(3)
|
||||||
|
prefill_req = _ReqStub(1)
|
||||||
|
scheduler = self._scheduler(
|
||||||
|
"v1",
|
||||||
|
[shared_req],
|
||||||
|
[],
|
||||||
|
last_batch=SimpleNamespace(reqs=[shared_req, prefill_req]),
|
||||||
|
)
|
||||||
|
|
||||||
|
Scheduler.record_weight_version_change(scheduler, new_version="v2")
|
||||||
|
|
||||||
|
self.assertEqual(len(shared_req.weight_version_events), 1)
|
||||||
|
self.assertEqual(len(prefill_req.weight_version_events), 1)
|
||||||
|
|
||||||
|
def test_records_on_all_pp_microbatches(self):
|
||||||
|
"""With pipeline parallelism every microbatch is visited, not just the selected one."""
|
||||||
|
mb0_req = _ReqStub(2)
|
||||||
|
mb1_req = _ReqStub(4)
|
||||||
|
pending_req = _ReqStub(6)
|
||||||
|
scheduler = self._scheduler("v1", [], [], pp_size=2)
|
||||||
|
scheduler.running_mbs = [
|
||||||
|
SimpleNamespace(reqs=[mb0_req]),
|
||||||
|
SimpleNamespace(reqs=[mb1_req]),
|
||||||
|
]
|
||||||
|
scheduler.mbs = [None, SimpleNamespace(reqs=[pending_req])]
|
||||||
|
|
||||||
|
Scheduler.record_weight_version_change(scheduler, new_version="v2")
|
||||||
|
|
||||||
|
for req in (mb0_req, mb1_req, pending_req):
|
||||||
|
self.assertEqual(len(req.weight_version_events), 1)
|
||||||
|
|
||||||
|
def test_records_on_hisparse_staging_requests(self):
|
||||||
|
"""Requests parked in the HiSparse staging queue are swept like any live request."""
|
||||||
|
staging_req = _ReqStub(2)
|
||||||
|
coordinator = SimpleNamespace(
|
||||||
|
ack_staging_queue=[SimpleNamespace(req=staging_req)]
|
||||||
|
)
|
||||||
|
scheduler = self._scheduler("v1", [], [], hisparse=coordinator)
|
||||||
|
|
||||||
|
Scheduler.record_weight_version_change(scheduler, new_version="v2")
|
||||||
|
|
||||||
|
self.assertEqual(len(staging_req.weight_version_events), 1)
|
||||||
|
self.assertEqual(staging_req.weight_version_events[0].old_version, "v1")
|
||||||
|
self.assertEqual(staging_req.weight_version_events[0].num_output_tokens, 2)
|
||||||
|
|
||||||
|
def test_same_version_is_a_noop(self):
|
||||||
|
"""Re-announcing the current version must not record events."""
|
||||||
|
req = _ReqStub(3)
|
||||||
|
scheduler = self._scheduler("v1", [req], [])
|
||||||
|
|
||||||
|
Scheduler.record_weight_version_change(scheduler, new_version="v1")
|
||||||
|
|
||||||
|
self.assertEqual(scheduler.serving.weight_version, "v1")
|
||||||
|
self.assertEqual(req.weight_version_events, [])
|
||||||
|
|
||||||
|
def test_none_version_is_a_noop(self):
|
||||||
|
"""An update without a weight version must not disturb attribution."""
|
||||||
|
req = _ReqStub(3)
|
||||||
|
scheduler = self._scheduler("v1", [req], [])
|
||||||
|
|
||||||
|
Scheduler.record_weight_version_change(scheduler, new_version=None)
|
||||||
|
|
||||||
|
self.assertEqual(scheduler.serving.weight_version, "v1")
|
||||||
|
self.assertEqual(req.weight_version_events, [])
|
||||||
|
|
||||||
|
|
||||||
|
class TestRecordWeightVersionEvents(CustomTestCase):
|
||||||
|
def test_records_only_for_requests_that_have_output(self):
|
||||||
|
"""Requests without output tokens are skipped without stopping the sweep."""
|
||||||
|
empty_req = _ReqStub(0)
|
||||||
|
started_req = _ReqStub(4)
|
||||||
|
|
||||||
|
num_recorded = record_weight_version_events(
|
||||||
|
[empty_req, started_req, _ReqStub(2)], old_version="v1"
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(num_recorded, 2)
|
||||||
|
self.assertEqual(empty_req.weight_version_events, [])
|
||||||
|
self.assertEqual(len(started_req.weight_version_events), 1)
|
||||||
|
self.assertEqual(started_req.weight_version_events[0].old_version, "v1")
|
||||||
|
self.assertEqual(started_req.weight_version_events[0].num_output_tokens, 4)
|
||||||
|
|
||||||
|
|
||||||
|
class TestTruncateWeightVersionEvents(CustomTestCase):
|
||||||
|
def test_events_within_the_kept_prefix_are_preserved(self):
|
||||||
|
"""Events entirely inside the streamed prefix survive a retract untouched."""
|
||||||
|
events = [WeightVersionEvent(old_version="v1", num_output_tokens=3)]
|
||||||
|
self.assertEqual(
|
||||||
|
truncate_weight_version_events(events, num_kept_tokens=5), events
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_events_beyond_the_kept_prefix_are_clamped(self):
|
||||||
|
"""An event past the streamed prefix is clamped so re-generated tokens forget it."""
|
||||||
|
events = [
|
||||||
|
WeightVersionEvent(old_version="v1", num_output_tokens=3),
|
||||||
|
WeightVersionEvent(old_version="v2", num_output_tokens=9),
|
||||||
|
]
|
||||||
|
self.assertEqual(
|
||||||
|
truncate_weight_version_events(events, num_kept_tokens=5),
|
||||||
|
[
|
||||||
|
WeightVersionEvent(old_version="v1", num_output_tokens=3),
|
||||||
|
WeightVersionEvent(old_version="v2", num_output_tokens=5),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_zero_kept_tokens_drops_all_events(self):
|
||||||
|
"""With nothing streamed yet the whole history is discarded, as before the fix."""
|
||||||
|
events = [WeightVersionEvent(old_version="v1", num_output_tokens=3)]
|
||||||
|
self.assertEqual(truncate_weight_version_events(events, num_kept_tokens=0), [])
|
||||||
|
|
||||||
|
|
||||||
|
class TestSpanlessAbortPaths(CustomTestCase):
|
||||||
|
def test_priority_disabled_rejection_attaches_spans(self):
|
||||||
|
"""The pre-scheduler priority rejection reports spans like every other abort."""
|
||||||
|
req = _ReqStub(0)
|
||||||
|
req.rid = "r0"
|
||||||
|
req.priority = 5
|
||||||
|
req.time_stats = SimpleNamespace(
|
||||||
|
trace_ctx=SimpleNamespace(abort=lambda abort_info: None)
|
||||||
|
)
|
||||||
|
sent = []
|
||||||
|
scheduler = SimpleNamespace(
|
||||||
|
enable_priority_scheduling=False,
|
||||||
|
abort_on_priority_when_disabled=True,
|
||||||
|
ipc_channels=SimpleNamespace(
|
||||||
|
send_to_tokenizer=SimpleNamespace(
|
||||||
|
send_output=lambda obj, req_arg: sent.append(obj)
|
||||||
|
)
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
with patch(
|
||||||
|
"sglang.srt.managers.scheduler.get_serving",
|
||||||
|
return_value=SimpleNamespace(weight_version="v1"),
|
||||||
|
):
|
||||||
|
accepted = Scheduler._set_or_validate_priority(scheduler, req)
|
||||||
|
|
||||||
|
self.assertFalse(accepted)
|
||||||
|
self.assertEqual(
|
||||||
|
sent[0].weight_versions, [WeightVersionSpan(version="v1", start=0, end=0)]
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class TestMakeAbortReq(CustomTestCase):
|
||||||
|
def test_abort_req_carries_spans_clamped_to_sampled_tokens(self):
|
||||||
|
"""An aborted request reports the versions that sampled the tokens it produced."""
|
||||||
|
req = _ReqStub(4)
|
||||||
|
req.rid = "r0"
|
||||||
|
req.weight_version_events.append(
|
||||||
|
WeightVersionEvent(old_version="v1", num_output_tokens=3)
|
||||||
|
)
|
||||||
|
with patch(
|
||||||
|
"sglang.srt.managers.scheduler.get_serving",
|
||||||
|
return_value=SimpleNamespace(weight_version="v2"),
|
||||||
|
):
|
||||||
|
abort_req = _make_abort_req(req)
|
||||||
|
|
||||||
|
self.assertIsInstance(abort_req, AbortReq)
|
||||||
|
self.assertEqual(abort_req.rid, "r0")
|
||||||
|
self.assertEqual(
|
||||||
|
abort_req.weight_versions,
|
||||||
|
[
|
||||||
|
WeightVersionSpan(version="v1", start=0, end=3),
|
||||||
|
WeightVersionSpan(version="v2", start=3, end=4),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class TestBuildEndpointWeightVersionMetadata(CustomTestCase):
|
||||||
|
def test_metadata_projects_only_the_weight_fields(self):
|
||||||
|
"""Endpoint metadata exposes the version fields and nothing else from meta_info."""
|
||||||
|
spans = [
|
||||||
|
{"version": "v1", "start": 0, "end": 3},
|
||||||
|
{"version": "v2", "start": 3, "end": 7},
|
||||||
|
]
|
||||||
|
metadata = build_endpoint_weight_version_metadata(
|
||||||
|
{
|
||||||
|
"weight_version": "v2",
|
||||||
|
"weight_versions": spans,
|
||||||
|
"id": "r0",
|
||||||
|
"completion_tokens": 7,
|
||||||
|
}
|
||||||
|
)
|
||||||
|
self.assertEqual(metadata, {"weight_version": "v2", "weight_versions": spans})
|
||||||
|
|
||||||
|
def test_metadata_omits_spans_when_absent(self):
|
||||||
|
"""Responses without spans still report the legacy version alone."""
|
||||||
|
metadata = build_endpoint_weight_version_metadata({"weight_version": "v2"})
|
||||||
|
self.assertEqual(metadata, {"weight_version": "v2"})
|
||||||
|
|
||||||
|
|
||||||
|
class TestAddWeightVersionsToMetaInfo(CustomTestCase):
|
||||||
|
def test_meta_info_gets_dict_spans_and_legacy_field(self):
|
||||||
|
"""Spans become dicts and the legacy weight_version reports the newest span."""
|
||||||
|
meta_info = {"weight_version": "stale"}
|
||||||
|
add_weight_versions_to_meta_info(
|
||||||
|
meta_info,
|
||||||
|
[
|
||||||
|
WeightVersionSpan(version="v1", start=0, end=3),
|
||||||
|
WeightVersionSpan(version="v2", start=3, end=7),
|
||||||
|
],
|
||||||
|
num_output_tokens=7,
|
||||||
|
)
|
||||||
|
self.assertEqual(
|
||||||
|
meta_info["weight_versions"],
|
||||||
|
[
|
||||||
|
{"version": "v1", "start": 0, "end": 3},
|
||||||
|
{"version": "v2", "start": 3, "end": 7},
|
||||||
|
],
|
||||||
|
)
|
||||||
|
self.assertEqual(meta_info["weight_version"], "v2")
|
||||||
|
|
||||||
|
def test_spans_are_clamped_to_returned_tokens(self):
|
||||||
|
"""Aborts returning fewer tokens than sampled clamp spans to the visible range."""
|
||||||
|
meta_info = {}
|
||||||
|
add_weight_versions_to_meta_info(
|
||||||
|
meta_info,
|
||||||
|
[
|
||||||
|
WeightVersionSpan(version="v1", start=0, end=3),
|
||||||
|
WeightVersionSpan(version="v2", start=3, end=7),
|
||||||
|
],
|
||||||
|
num_output_tokens=5,
|
||||||
|
)
|
||||||
|
self.assertEqual(
|
||||||
|
meta_info["weight_versions"],
|
||||||
|
[
|
||||||
|
{"version": "v1", "start": 0, "end": 3},
|
||||||
|
{"version": "v2", "start": 3, "end": 5},
|
||||||
|
],
|
||||||
|
)
|
||||||
|
self.assertEqual(meta_info["weight_version"], "v2")
|
||||||
|
|
||||||
|
def test_clamp_never_extends_spans(self):
|
||||||
|
"""A response longer than the sampled range leaves the spans untouched."""
|
||||||
|
meta_info = {}
|
||||||
|
add_weight_versions_to_meta_info(
|
||||||
|
meta_info,
|
||||||
|
[WeightVersionSpan(version="v1", start=0, end=3)],
|
||||||
|
num_output_tokens=10,
|
||||||
|
)
|
||||||
|
self.assertEqual(
|
||||||
|
meta_info["weight_versions"], [{"version": "v1", "start": 0, "end": 3}]
|
||||||
|
)
|
||||||
|
self.assertEqual(meta_info["weight_version"], "v1")
|
||||||
|
|
||||||
|
def test_clamp_drops_trailing_spans_and_rewrites_the_legacy_field(self):
|
||||||
|
"""Dropping invisible spans also moves the legacy version back."""
|
||||||
|
meta_info = {}
|
||||||
|
add_weight_versions_to_meta_info(
|
||||||
|
meta_info,
|
||||||
|
[
|
||||||
|
WeightVersionSpan(version="v1", start=0, end=3),
|
||||||
|
WeightVersionSpan(version="v2", start=3, end=7),
|
||||||
|
],
|
||||||
|
num_output_tokens=3,
|
||||||
|
)
|
||||||
|
self.assertEqual(
|
||||||
|
meta_info["weight_versions"], [{"version": "v1", "start": 0, "end": 3}]
|
||||||
|
)
|
||||||
|
self.assertEqual(meta_info["weight_version"], "v1")
|
||||||
|
|
||||||
|
def test_clamp_to_zero_tokens_keeps_a_degenerate_span(self):
|
||||||
|
"""A response with no visible tokens still reports a well-formed empty span."""
|
||||||
|
meta_info = {}
|
||||||
|
add_weight_versions_to_meta_info(
|
||||||
|
meta_info,
|
||||||
|
[
|
||||||
|
WeightVersionSpan(version="v1", start=0, end=3),
|
||||||
|
WeightVersionSpan(version="v2", start=3, end=7),
|
||||||
|
],
|
||||||
|
num_output_tokens=0,
|
||||||
|
)
|
||||||
|
self.assertEqual(
|
||||||
|
meta_info["weight_versions"], [{"version": "v1", "start": 0, "end": 0}]
|
||||||
|
)
|
||||||
|
self.assertEqual(meta_info["weight_version"], "v1")
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
Reference in New Issue
Block a user