This commit is contained in:
@@ -3,6 +3,12 @@ import unittest
|
||||
import torch
|
||||
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.managers.io_struct import (
|
||||
GetInternalStateReqOutput,
|
||||
PickleWrapper,
|
||||
msgpack_decode,
|
||||
msgpack_encode,
|
||||
)
|
||||
from sglang.srt.speculative.dspark_components.dspark_observability import (
|
||||
DecodeStepObservation,
|
||||
DsparkInfoDumper,
|
||||
@@ -12,6 +18,7 @@ from sglang.srt.speculative.dspark_components.dspark_observability import (
|
||||
resolve_components,
|
||||
resolve_enabled_components,
|
||||
)
|
||||
from sglang.srt.utils.msgspec_utils import msgspec_to_builtins
|
||||
from sglang.test.ci.ci_register import register_cpu_ci
|
||||
from sglang.test.test_utils import CustomTestCase
|
||||
|
||||
@@ -394,5 +401,39 @@ class TestReqsAndGpuTiming(CustomTestCase):
|
||||
self.assertGreater(record["target_verify_gpu_ms"], 0.0)
|
||||
|
||||
|
||||
class TestDumpCrossesMsgpackIpc(CustomTestCase):
|
||||
"""Guard the DSpark -> GetInternalStateReqOutput serialization contract.
|
||||
|
||||
`Scheduler.get_internal_state` stores `draft_worker.dump_info_records()` under
|
||||
`internal_state["dspark_info_record"]`, then ships the struct over the strict
|
||||
msgpack IPC path (issue #29465). That path has no PickleWrapper fallback: any
|
||||
value that is not msgpack-native (a numpy scalar, a torch tensor, an
|
||||
un-converted `msgspec.Struct`) raises at encode time. The dumper's own tests
|
||||
assert record *values* -- and `assertEqual(np.int64(3), 3)` passes -- so a
|
||||
scalar that silently became numpy would escape them but fail this round-trip.
|
||||
"""
|
||||
|
||||
def _real_dump(self):
|
||||
dumper, clock = make_dumper({"core"})
|
||||
dumper.observe_decode_step(make_obs(forward_ct=1))
|
||||
clock.advance(0.01)
|
||||
dumper.observe_decode_step(make_obs(forward_ct=2))
|
||||
dumped = dumper.dump()
|
||||
# DsparkObservability.dump_info_records appends this float onto the raw
|
||||
# dumper output before the scheduler reads it; mirror the full payload.
|
||||
dumped["simulate_acc_len"] = 4.0
|
||||
return dumped
|
||||
|
||||
def test_real_dump_output_round_trips_natively(self):
|
||||
internal_state = msgspec_to_builtins(
|
||||
{"dspark_info_record": self._real_dump(), "max_running_requests": 256}
|
||||
)
|
||||
output = GetInternalStateReqOutput(internal_state=internal_state)
|
||||
|
||||
decoded = msgpack_decode(msgpack_encode(output))
|
||||
self.assertNotIsInstance(decoded, PickleWrapper)
|
||||
self.assertEqual(decoded, output)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
||||
@@ -0,0 +1,241 @@
|
||||
"""Round-trip coverage for the IPC structs that used to be pickle-wrapped.
|
||||
|
||||
Issue #29465 Task 4 tightened the 13 types in `_REQ_TYPES_WITH_OPAQUE_FIELDS` to
|
||||
precise msgspec-native annotations and deleted the registry. This test proves
|
||||
each type now encodes natively over the msgpack IPC path (no `PickleWrapper`
|
||||
frame) by asserting `msgpack_decode(msgpack_encode(x)) == x`, and guards the
|
||||
type-specific decisions (the `ExpertWeightPointer` narrowing, the
|
||||
`CheckWeightsReqOutput` struct mirrors, and the internal-state sanitization).
|
||||
"""
|
||||
|
||||
import dataclasses
|
||||
import unittest
|
||||
|
||||
import msgspec
|
||||
|
||||
from sglang.srt.managers import io_struct
|
||||
from sglang.srt.managers.io_struct import (
|
||||
BackupDramReq,
|
||||
ChecksumInfo,
|
||||
CheckWeightsReqOutput,
|
||||
DumperControlReqInput,
|
||||
DumperControlReqOutput,
|
||||
ExpertWeightPointer,
|
||||
GetInternalStateReqOutput,
|
||||
GetWeightsByNameReqOutput,
|
||||
LoadLoRAAdapterFromTensorsReqInput,
|
||||
ParallelismInfo,
|
||||
RpcReqInput,
|
||||
SetInternalStateReq,
|
||||
SetInternalStateReqOutput,
|
||||
UpdateWeightFromDiskReqInput,
|
||||
VertexGenerateReqInput,
|
||||
msgpack_decode,
|
||||
msgpack_encode,
|
||||
)
|
||||
from sglang.srt.model_executor.cuda_graph_config import CudaGraphConfig
|
||||
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 ParallelismInfo as PydanticParallelismInfo
|
||||
from sglang.test.ci.ci_register import register_cpu_ci
|
||||
from sglang.test.test_utils import CustomTestCase
|
||||
|
||||
register_cpu_ci(est_time=10, suite="base-c-test-cpu")
|
||||
|
||||
|
||||
def _round_trip(obj):
|
||||
return msgpack_decode(msgpack_encode(obj))
|
||||
|
||||
|
||||
def _double_hop(obj):
|
||||
# MultiTokenizerRouter and the DP controller re-encode already-decoded
|
||||
# structs, so a single hop cannot catch re-encode bugs.
|
||||
return msgpack_decode(msgpack_encode(msgpack_decode(msgpack_encode(obj))))
|
||||
|
||||
|
||||
def _contains_dataclass(obj) -> bool:
|
||||
if dataclasses.is_dataclass(obj) and not isinstance(obj, type):
|
||||
return True
|
||||
if isinstance(obj, dict):
|
||||
return any(_contains_dataclass(v) for v in obj.values())
|
||||
if isinstance(obj, (list, tuple, set)):
|
||||
return any(_contains_dataclass(v) for v in obj)
|
||||
return False
|
||||
|
||||
|
||||
def _parallelism_info() -> ParallelismInfo:
|
||||
return ParallelismInfo(
|
||||
tp_rank=0, tp_size=2, dp_rank=0, dp_size=1, pp_rank=0, pp_size=1, rank=0, size=2
|
||||
)
|
||||
|
||||
|
||||
def _checksum_info(tag: str) -> ChecksumInfo:
|
||||
return ChecksumInfo(
|
||||
checksums={f"model.layers.{tag}": "deadbeef"},
|
||||
per_gpu_checksum="cafef00d",
|
||||
parallelism_info=_parallelism_info(),
|
||||
)
|
||||
|
||||
|
||||
# One representative instance per (now-tightened) ex-registry type. The 13th
|
||||
# entry, SetInjectDumpMetadataReqInput, was deleted as dead code, leaving 12.
|
||||
REGISTRY_TYPE_INSTANCES = {
|
||||
"UpdateWeightFromDiskReqInput": UpdateWeightFromDiskReqInput(
|
||||
model_path="dummy", manifest={"w": [1, 2], "meta": {"k": "v"}}
|
||||
),
|
||||
"BackupDramReq": BackupDramReq(
|
||||
rank=0,
|
||||
weight_pointer_map={
|
||||
"experts.0.gate_proj": ExpertWeightPointer(weight_ptr=8, byte_size=4),
|
||||
"experts.1.up_proj": ExpertWeightPointer(weight_ptr=16, byte_size=8),
|
||||
},
|
||||
session_id="session",
|
||||
buffer_size=1024,
|
||||
),
|
||||
"GetWeightsByNameReqOutput/flat": GetWeightsByNameReqOutput(
|
||||
parameter=[1.0, 2.5, 3.0]
|
||||
),
|
||||
"GetWeightsByNameReqOutput/nested": GetWeightsByNameReqOutput(
|
||||
parameter=[[1.0, 2.0], [3.0]]
|
||||
),
|
||||
"GetWeightsByNameReqOutput/none": GetWeightsByNameReqOutput(parameter=None),
|
||||
"CheckWeightsReqOutput": CheckWeightsReqOutput(
|
||||
success=True,
|
||||
message="Success.",
|
||||
payload=[_checksum_info("0"), _checksum_info("1")],
|
||||
),
|
||||
"GetInternalStateReqOutput": GetInternalStateReqOutput(
|
||||
internal_state={"a": 1, "b": [1, 2], "c": {"d": "e"}, "f": None}
|
||||
),
|
||||
"SetInternalStateReq": SetInternalStateReq(
|
||||
server_args={
|
||||
"pp_max_micro_batch_size": 4,
|
||||
"speculative_accept_threshold_acc": 0.5,
|
||||
}
|
||||
),
|
||||
"SetInternalStateReqOutput": SetInternalStateReqOutput(updated=True),
|
||||
"VertexGenerateReqInput": VertexGenerateReqInput(
|
||||
instances=[{"prompt": "hi"}], parameters={"max_tokens": 8}
|
||||
),
|
||||
"RpcReqInput/empty": RpcReqInput(method="collective_rpc", parameters={}),
|
||||
"RpcReqInput/scalars": RpcReqInput(
|
||||
method="collective_rpc",
|
||||
parameters={"flag": True, "n": 1, "ratio": 2.0, "name": "x", "opt": None},
|
||||
),
|
||||
"RpcReqInput/none": RpcReqInput(method="collective_rpc", parameters=None),
|
||||
"LoadLoRAAdapterFromTensorsReqInput": LoadLoRAAdapterFromTensorsReqInput(
|
||||
lora_name="adapter",
|
||||
config_dict={"r": 8, "lora_alpha": 16, "target_modules": ["q_proj", "v_proj"]},
|
||||
serialized_tensors="",
|
||||
added_tokens_config={"<extra>": 32000},
|
||||
),
|
||||
"DumperControlReqInput": DumperControlReqInput(method="start", body={"k": "v"}),
|
||||
"DumperControlReqOutput": DumperControlReqOutput(
|
||||
success=True, response=[{"worker": 0, "ok": True}]
|
||||
),
|
||||
}
|
||||
|
||||
NARROWED_BACKUP_KEYS = ("name", "shape", "numel", "dtype", "element_size")
|
||||
|
||||
|
||||
class TestMsgpackIpcRoundtrip(CustomTestCase):
|
||||
def test_registry_is_empty(self):
|
||||
# `getattr(..., ())` is deliberate: the acceptance criterion for Task 4 is
|
||||
# that the symbol is *deleted*, so this asserts its absence rather than
|
||||
# defensively reading a field. It must survive the symbol removal.
|
||||
self.assertEqual(
|
||||
getattr(io_struct, "_REQ_TYPES_WITH_OPAQUE_FIELDS", ()),
|
||||
(),
|
||||
)
|
||||
|
||||
def test_each_type_round_trips_natively(self):
|
||||
for name, instance in REGISTRY_TYPE_INSTANCES.items():
|
||||
with self.subTest(type=name):
|
||||
encoded = msgpack_encode(instance)
|
||||
# Natively encoded structs are never wrapped: a PickleWrapper
|
||||
# frame would decode back to a PickleWrapper, not the type.
|
||||
self.assertNotIsInstance(
|
||||
msgpack_decode(encoded), io_struct.PickleWrapper
|
||||
)
|
||||
self.assertEqual(_round_trip(instance), instance)
|
||||
self.assertEqual(_double_hop(instance), instance)
|
||||
|
||||
def test_backup_dram_req_is_narrowed(self):
|
||||
# ExpertWeightPointer carries only the two fields the consumer reads; the
|
||||
# five torch-metadata keys the producer used to send are gone.
|
||||
field_names = {f.name for f in msgspec.structs.fields(ExpertWeightPointer)}
|
||||
self.assertEqual(field_names, {"weight_ptr", "byte_size"})
|
||||
for dropped in NARROWED_BACKUP_KEYS:
|
||||
self.assertNotIn(dropped, field_names)
|
||||
|
||||
decoded = _round_trip(REGISTRY_TYPE_INSTANCES["BackupDramReq"])
|
||||
pointer = decoded.weight_pointer_map["experts.0.gate_proj"]
|
||||
self.assertEqual((pointer.weight_ptr, pointer.byte_size), (8, 4))
|
||||
|
||||
def test_check_weights_mirrors_match_pydantic_models(self):
|
||||
# Field-parity guard: the msgspec wire structs must not drift from the
|
||||
# pydantic source of truth in weight_checker.
|
||||
self.assertEqual(
|
||||
{f.name for f in msgspec.structs.fields(ParallelismInfo)},
|
||||
set(PydanticParallelismInfo.model_fields),
|
||||
)
|
||||
self.assertEqual(
|
||||
{f.name for f in msgspec.structs.fields(ChecksumInfo)},
|
||||
set(PydanticChecksumInfo.model_fields),
|
||||
)
|
||||
|
||||
def test_check_weights_multi_rank_payload(self):
|
||||
# tp>1 sends one ChecksumInfo per rank; the list round-trips and stays a
|
||||
# {field: value} dict once converted back to builtins for the HTTP body.
|
||||
instance = REGISTRY_TYPE_INSTANCES["CheckWeightsReqOutput"]
|
||||
decoded = _round_trip(instance)
|
||||
self.assertEqual(len(decoded.payload), 2)
|
||||
as_dict = msgspec_to_builtins(decoded.payload[0])
|
||||
self.assertEqual(as_dict["per_gpu_checksum"], "cafef00d")
|
||||
self.assertIn("tp_rank", as_dict["parallelism_info"])
|
||||
|
||||
def test_check_weights_producer_conversion(self):
|
||||
# Mirrors weight_updater.check_weights: WeightChecker returns
|
||||
# ChecksumInfo.model_dump() (a dict), converted to the msgspec struct via
|
||||
# msgspec.convert, and the result round-trips as the payload.
|
||||
pydantic_checksum = PydanticChecksumInfo(
|
||||
checksums={"model.layers.0": "deadbeef"},
|
||||
per_gpu_checksum="cafef00d",
|
||||
parallelism_info=PydanticParallelismInfo(
|
||||
tp_rank=0,
|
||||
tp_size=2,
|
||||
dp_rank=0,
|
||||
dp_size=1,
|
||||
pp_rank=0,
|
||||
pp_size=1,
|
||||
rank=0,
|
||||
size=2,
|
||||
),
|
||||
)
|
||||
converted = msgspec.convert(pydantic_checksum.model_dump(), ChecksumInfo)
|
||||
self.assertEqual(converted.per_gpu_checksum, "cafef00d")
|
||||
self.assertEqual(converted.parallelism_info.tp_rank, 0)
|
||||
output = CheckWeightsReqOutput(success=True, message="ok", payload=[converted])
|
||||
self.assertEqual(_round_trip(output), output)
|
||||
|
||||
def test_get_internal_state_sanitizes_dataclass(self):
|
||||
# A live vars(ServerArgs) dump holds a dataclass: cuda_graph_config is an
|
||||
# Optional[CudaGraphConfig]. The producer sanitizes via msgspec_to_builtins
|
||||
# so it does not survive onto the wire; materialize CudaGraphConfig
|
||||
# explicitly.
|
||||
raw = {
|
||||
"cuda_graph_config": CudaGraphConfig(),
|
||||
"max_running_requests": 256,
|
||||
}
|
||||
self.assertTrue(_contains_dataclass(raw))
|
||||
|
||||
sanitized = msgspec_to_builtins(raw)
|
||||
self.assertFalse(_contains_dataclass(sanitized))
|
||||
self.assertIsInstance(sanitized["cuda_graph_config"], dict)
|
||||
|
||||
output = GetInternalStateReqOutput(internal_state=sanitized)
|
||||
self.assertEqual(_round_trip(output), output)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
Reference in New Issue
Block a user