Report per-token weight-version spans in generation meta info (#35926)

This commit is contained in:
fzyzcjy
2026-08-24 20:18:52 +08:00
committed by GitHub
parent 981dfa2b83
commit 3b24d8981b
20 changed files with 1607 additions and 25 deletions
@@ -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,
+6
View File
@@ -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[
+27 -2
View File
@@ -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(
+117
View File
@@ -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()