Convert IPC dataclasses to msgspec.Struct with opt-in msgpack transport (#28688)
Co-authored-by: Lianmin Zheng <lianminzheng@gmail.com>
This commit is contained in:
co-authored by
Lianmin Zheng
parent
714011a40f
commit
be1930133a
@@ -22,7 +22,7 @@ import torch
|
||||
import torch.distributed as dist
|
||||
import zmq
|
||||
|
||||
from sglang.srt.managers.io_struct import sock_recv, sock_send
|
||||
from sglang.srt.managers.io_struct import sock_recv, sock_send, wrap_as_pickle
|
||||
|
||||
# -------------------------------------- config base ------------------------------------------
|
||||
|
||||
@@ -1440,7 +1440,7 @@ def _create_zmq_rpc_broadcast(
|
||||
except Exception as e:
|
||||
_log(f"[ZmqRpc] error inside handler: {e}")
|
||||
resp = {"result": None, "error": str(e)}
|
||||
sock_send(sock, resp)
|
||||
sock_send(sock, wrap_as_pickle(resp))
|
||||
|
||||
thread = threading.Thread(target=serve_loop, daemon=True)
|
||||
thread.start()
|
||||
@@ -1479,11 +1479,13 @@ class _ZmqRpcHandle:
|
||||
def call(*args, **kwargs):
|
||||
sock_send(
|
||||
self._socket,
|
||||
{
|
||||
"method": method_name,
|
||||
"args": args,
|
||||
"kwargs": kwargs,
|
||||
},
|
||||
wrap_as_pickle(
|
||||
{
|
||||
"method": method_name,
|
||||
"args": args,
|
||||
"kwargs": kwargs,
|
||||
}
|
||||
),
|
||||
)
|
||||
response = sock_recv(self._socket)
|
||||
if response["error"]:
|
||||
|
||||
@@ -26,7 +26,7 @@ from sglang.srt.disaggregation.encode_server import (
|
||||
handle_scheduler_receive_url_request,
|
||||
launch_encoder,
|
||||
)
|
||||
from sglang.srt.managers.io_struct import async_sock_send
|
||||
from sglang.srt.managers.io_struct import async_sock_send, wrap_as_pickle
|
||||
from sglang.srt.managers.schedule_batch import Modality
|
||||
from sglang.srt.server_args import PortArgs, ServerArgs
|
||||
from sglang.srt.utils import random_uuid
|
||||
@@ -96,7 +96,7 @@ class SGLangEncoderServer(SGLangEncoderServicer):
|
||||
"part_idx": request.part_idx,
|
||||
}
|
||||
for socket in self.send_sockets:
|
||||
await async_sock_send(socket, request_dict)
|
||||
await async_sock_send(socket, wrap_as_pickle(request_dict))
|
||||
|
||||
# gRPC encode is image-only; encoder.encode() requires modality
|
||||
(
|
||||
|
||||
@@ -50,6 +50,7 @@ from sglang.srt.managers.io_struct import (
|
||||
async_sock_recv,
|
||||
async_sock_send,
|
||||
sock_send,
|
||||
wrap_as_pickle,
|
||||
)
|
||||
from sglang.srt.managers.schedule_batch import Modality, MultimodalDataItem
|
||||
from sglang.srt.mem_cache.multimodal_cache import EmbeddingResult, MultiModalStaticCache
|
||||
@@ -2398,12 +2399,14 @@ class EncoderScheduler:
|
||||
for sock in self.send_sockets:
|
||||
sock_send(
|
||||
sock,
|
||||
{
|
||||
"type": "batch_encode",
|
||||
"modality": modality.name,
|
||||
"requests": requests,
|
||||
"enter_time": start,
|
||||
},
|
||||
wrap_as_pickle(
|
||||
{
|
||||
"type": "batch_encode",
|
||||
"modality": modality.name,
|
||||
"requests": requests,
|
||||
"enter_time": start,
|
||||
}
|
||||
),
|
||||
)
|
||||
|
||||
logger.info(f"Dispatching batch of {len(group)} {modality.name} requests")
|
||||
@@ -2449,7 +2452,7 @@ class EncoderScheduler:
|
||||
req = p.request
|
||||
try:
|
||||
for sock in self.send_sockets:
|
||||
sock_send(sock, req)
|
||||
sock_send(sock, wrap_as_pickle(req))
|
||||
result = await self.encoder.encode_request(req, modality)
|
||||
if not p.future.done():
|
||||
p.future.set_result(result)
|
||||
@@ -2738,7 +2741,7 @@ class DPDispatcher:
|
||||
)
|
||||
|
||||
try:
|
||||
await async_sock_send(self.dispatch_sockets[rank], request)
|
||||
await async_sock_send(self.dispatch_sockets[rank], wrap_as_pickle(request))
|
||||
# An alive-but-stuck worker (NCCL deadlock etc.) wouldn't trip
|
||||
# the watchdog, so bound the wait explicitly.
|
||||
return await asyncio.wait_for(future, timeout=ENCODER_REQ_TIMEOUT)
|
||||
@@ -2787,7 +2790,7 @@ class DPDispatcher:
|
||||
f"dp_rank={rank}, pending={self.pending_counts}"
|
||||
)
|
||||
try:
|
||||
await async_sock_send(self.dispatch_sockets[rank], request)
|
||||
await async_sock_send(self.dispatch_sockets[rank], wrap_as_pickle(request))
|
||||
return await asyncio.wait_for(future, timeout=ENCODER_REQ_TIMEOUT)
|
||||
except asyncio.TimeoutError:
|
||||
self.pending_futures[rank].pop(key, None)
|
||||
@@ -2828,7 +2831,9 @@ class DPDispatcher:
|
||||
self.req_id_to_rank[req_id] = rank
|
||||
rank_keys.append((rank, req_id))
|
||||
request_copy = {**request, "req_id": req_id}
|
||||
await async_sock_send(self.dispatch_sockets[rank], request_copy)
|
||||
await async_sock_send(
|
||||
self.dispatch_sockets[rank], wrap_as_pickle(request_copy)
|
||||
)
|
||||
futures.append(future)
|
||||
# Concurrent wait → total bounded by eff_timeout, not
|
||||
# dp_size × eff_timeout.
|
||||
@@ -3056,7 +3061,7 @@ async def _dp_worker_handle_request(
|
||||
# pyzmq async send isn't safe for concurrent senders.
|
||||
try:
|
||||
async with send_lock:
|
||||
await async_sock_send(send_sock, envelope)
|
||||
await async_sock_send(send_sock, wrap_as_pickle(envelope))
|
||||
except Exception:
|
||||
logger.error(
|
||||
f"DP worker {dp_rank} failed to send envelope for "
|
||||
@@ -3603,7 +3608,7 @@ async def handle_encode_request(request: dict):
|
||||
)
|
||||
else:
|
||||
for socket in send_sockets:
|
||||
sock_send(socket, request)
|
||||
sock_send(socket, wrap_as_pickle(request))
|
||||
nbytes, embedding_len, embedding_dim, error_msg, error_code = (
|
||||
await encoder.encode_request(request, modality)
|
||||
)
|
||||
@@ -3824,7 +3829,7 @@ async def health_generate():
|
||||
|
||||
# Broadcast to other TP ranks so distributed ops stay in sync
|
||||
for socket in send_sockets:
|
||||
sock_send(socket, dummy_request)
|
||||
sock_send(socket, wrap_as_pickle(dummy_request))
|
||||
|
||||
# Run encode on rank 0 with timeout
|
||||
_, _, _, error_msg, _ = await asyncio.wait_for(
|
||||
|
||||
@@ -109,6 +109,7 @@ from sglang.srt.utils import (
|
||||
set_prometheus_multiproc_dir,
|
||||
set_ulimit,
|
||||
)
|
||||
from sglang.srt.utils.msgspec_utils import msgspec_to_builtins
|
||||
from sglang.srt.utils.network import get_zmq_socket, is_port_available
|
||||
from sglang.srt.utils.torch_memory_saver_adapter import TorchMemorySaverAdapter
|
||||
from sglang.srt.utils.watchdog import SubprocessWatchdog
|
||||
@@ -994,12 +995,14 @@ class Engine(EngineScoreMixin, EngineBase):
|
||||
internal_states = self.loop.run_until_complete(
|
||||
self.tokenizer_manager.get_internal_state()
|
||||
)
|
||||
return {
|
||||
**dataclasses.asdict(self.tokenizer_manager.server_args),
|
||||
**self._scheduler_init_result.scheduler_infos[0],
|
||||
"internal_states": internal_states,
|
||||
"version": __version__,
|
||||
}
|
||||
return msgspec_to_builtins(
|
||||
{
|
||||
**dataclasses.asdict(self.tokenizer_manager.server_args),
|
||||
**self._scheduler_init_result.scheduler_infos[0],
|
||||
"internal_states": internal_states,
|
||||
"version": __version__,
|
||||
}
|
||||
)
|
||||
|
||||
def init_weights_update_group(
|
||||
self,
|
||||
|
||||
@@ -15,6 +15,8 @@ from typing import Any, Awaitable, Callable, Dict, List, Optional
|
||||
|
||||
from pydantic import ValidationError
|
||||
|
||||
from sglang.srt.utils.msgspec_utils import msgspec_to_builtins
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@@ -383,9 +385,9 @@ class RuntimeHandle:
|
||||
return json.dumps(result, default=str)
|
||||
|
||||
def get_server_info(self) -> str:
|
||||
result: Dict[str, Any] = dict(dataclasses.asdict(self.server_args))
|
||||
result: Dict[str, Any] = dataclasses.asdict(self.server_args)
|
||||
result.update(self.scheduler_info)
|
||||
return json.dumps(result, default=str)
|
||||
return json.dumps(msgspec_to_builtins(result), default=str)
|
||||
|
||||
def health_check(self) -> bool:
|
||||
from sglang.srt.managers.tokenizer_manager import ServerStatus
|
||||
|
||||
@@ -173,6 +173,7 @@ from sglang.srt.utils.json_response import (
|
||||
dumps_json,
|
||||
orjson_response,
|
||||
)
|
||||
from sglang.srt.utils.msgspec_utils import msgspec_to_builtins
|
||||
from sglang.srt.utils.watchdog import SubprocessWatchdog
|
||||
from sglang.utils import get_exception_traceback
|
||||
from sglang.version import __version__
|
||||
@@ -698,16 +699,18 @@ async def server_info():
|
||||
server_args = _global_state.tokenizer_manager.server_args
|
||||
|
||||
# server_args.model_config is not serializable but should be excluded by asdict.
|
||||
return {
|
||||
**dataclasses.asdict(server_args),
|
||||
**_global_state.scheduler_info,
|
||||
"internal_states": internal_states,
|
||||
"version": __version__,
|
||||
# Structured KV-event publisher descriptor for KV-aware routers.
|
||||
# `None` when publishing is disabled or misconfigured; see
|
||||
# `ServerArgs.describe_kv_events_publisher` for the precise contract.
|
||||
"kv_events": server_args.describe_kv_events_publisher(),
|
||||
}
|
||||
return msgspec_to_builtins(
|
||||
{
|
||||
**dataclasses.asdict(server_args),
|
||||
**_global_state.scheduler_info,
|
||||
"internal_states": internal_states,
|
||||
"version": __version__,
|
||||
# Structured KV-event publisher descriptor for KV-aware routers.
|
||||
# `None` when publishing is disabled or misconfigured; see
|
||||
# `ServerArgs.describe_kv_events_publisher` for the precise contract.
|
||||
"kv_events": server_args.describe_kv_events_publisher(),
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
@app.get("/get_load")
|
||||
@@ -1417,7 +1420,7 @@ async def load_lora_adapter(
|
||||
"""Load a new LoRA adapter without re-launching the server."""
|
||||
result = await _global_state.tokenizer_manager.load_lora_adapter(obj, request)
|
||||
status_code = HTTPStatus.OK if result.success else HTTPStatus.BAD_REQUEST
|
||||
return ORJSONResponse(result, status_code=status_code)
|
||||
return ORJSONResponse(msgspec_to_builtins(result), status_code=status_code)
|
||||
|
||||
|
||||
@app.api_route("/load_lora_adapter_from_tensors", methods=["POST"])
|
||||
@@ -1429,7 +1432,7 @@ async def load_lora_adapter_from_tensors(
|
||||
obj, request
|
||||
)
|
||||
status_code = HTTPStatus.OK if result.success else HTTPStatus.BAD_REQUEST
|
||||
return ORJSONResponse(result, status_code=status_code)
|
||||
return ORJSONResponse(msgspec_to_builtins(result), status_code=status_code)
|
||||
|
||||
|
||||
@app.api_route("/unload_lora_adapter", methods=["POST"])
|
||||
@@ -1440,7 +1443,7 @@ async def unload_lora_adapter(
|
||||
"""Load a new LoRA adapter without re-launching the server."""
|
||||
result = await _global_state.tokenizer_manager.unload_lora_adapter(obj, request)
|
||||
status_code = HTTPStatus.OK if result.success else HTTPStatus.BAD_REQUEST
|
||||
return ORJSONResponse(result, status_code=status_code)
|
||||
return ORJSONResponse(msgspec_to_builtins(result), status_code=status_code)
|
||||
|
||||
|
||||
@app.api_route("/open_session", methods=["GET", "POST"])
|
||||
|
||||
@@ -241,6 +241,10 @@ class Envs:
|
||||
SGLANG_LOG_SCHEDULER_STATUS_TARGET = EnvStr("")
|
||||
SGLANG_LOG_SCHEDULER_STATUS_INTERVAL = EnvFloat(60.0)
|
||||
|
||||
# IPC
|
||||
SGLANG_USE_PICKLE_IPC = EnvBool(True)
|
||||
SGLANG_LOG_PICKLE_IPC_OBJECTS = EnvBool(False)
|
||||
|
||||
# SGLang CI
|
||||
SGLANG_IS_IN_CI = EnvBool(False)
|
||||
SGLANG_IS_IN_CI_AMD = EnvBool(False)
|
||||
|
||||
@@ -15,16 +15,17 @@
|
||||
|
||||
import asyncio
|
||||
from collections import OrderedDict
|
||||
from dataclasses import dataclass, field, fields
|
||||
from typing import Dict, List, Optional, Union
|
||||
from uuid import NAMESPACE_URL, uuid4, uuid5
|
||||
|
||||
import msgspec
|
||||
from msgspec.structs import fields
|
||||
|
||||
from sglang.srt.utils import ConcurrentCounter
|
||||
from sglang.srt.utils.aio_rwlock import RWLock
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class LoRARef:
|
||||
class LoRARef(msgspec.Struct, frozen=True, array_like=True):
|
||||
"""
|
||||
Reference record for a LoRA model.
|
||||
|
||||
@@ -33,7 +34,7 @@ class LoRARef:
|
||||
keys (e.g., radix cache).
|
||||
"""
|
||||
|
||||
lora_id: str = field(default_factory=lambda: uuid4().hex)
|
||||
lora_id: str = msgspec.field(default_factory=lambda: uuid4().hex)
|
||||
lora_name: Optional[str] = None
|
||||
lora_path: Optional[str] = None
|
||||
pinned: Optional[bool] = None
|
||||
|
||||
@@ -38,6 +38,8 @@ from sglang.srt.managers.io_struct import (
|
||||
TokenizedGenerateReqInput,
|
||||
sock_recv,
|
||||
sock_send,
|
||||
unwrap_from_pickle,
|
||||
wrap_as_pickle,
|
||||
)
|
||||
from sglang.srt.managers.load_snapshot import create_load_snapshot_reader
|
||||
from sglang.srt.managers.schedule_batch import Req
|
||||
@@ -238,10 +240,14 @@ class DataParallelController:
|
||||
if refresh_load_budget and self.refresh_load_budget_on_dispatch:
|
||||
self.refresh_load_budget()
|
||||
|
||||
req.time_stats = DPControllerReqTimeStats.new_from_obj(req.time_stats)
|
||||
time_stats = DPControllerReqTimeStats.new_from_obj(
|
||||
unwrap_from_pickle(req.time_stats)
|
||||
)
|
||||
|
||||
req.time_stats.set_dp_dispatch_time()
|
||||
time_stats.set_dp_dispatch_time()
|
||||
req.time_stats = wrap_as_pickle(time_stats)
|
||||
self.dispatching(req)
|
||||
req.time_stats = time_stats
|
||||
req.time_stats.set_dp_dispatch_finish_time()
|
||||
|
||||
def dispatch_batch_generate(self, batch_req: BatchTokenizedGenerateReqInput):
|
||||
@@ -382,11 +388,11 @@ class DataParallelController:
|
||||
connected_clients = 0
|
||||
while connected_clients < expected_clients:
|
||||
# Wait for client handshake
|
||||
client_rank = sock_recv(rep_socket).decode()
|
||||
client_rank = sock_recv(rep_socket)
|
||||
logger.debug(f"Received handshake from node {client_rank}")
|
||||
|
||||
# Send worker ports to client
|
||||
sock_send(rep_socket, worker_ports)
|
||||
sock_send(rep_socket, wrap_as_pickle(worker_ports))
|
||||
connected_clients += 1
|
||||
logger.debug(
|
||||
f"Sent worker ports to {connected_clients}/{expected_clients} nodes"
|
||||
@@ -411,7 +417,7 @@ class DataParallelController:
|
||||
while True:
|
||||
# Wait for client handshake
|
||||
try:
|
||||
client_rank = sock_recv(rep_socket).decode()
|
||||
client_rank = sock_recv(rep_socket)
|
||||
except Exception:
|
||||
logger.exception(
|
||||
"Failed to recv/decode handshake in reply thread; continue"
|
||||
@@ -420,7 +426,7 @@ class DataParallelController:
|
||||
logger.debug(f"Received handshake from node {client_rank}")
|
||||
|
||||
# Send worker ports to client
|
||||
sock_send(rep_socket, worker_ports)
|
||||
sock_send(rep_socket, wrap_as_pickle(worker_ports))
|
||||
logger.debug(f"Sent worker ports to node {client_rank}")
|
||||
|
||||
def _receive_ports_as_client(self, endpoint: str, node_rank: int) -> List[int]:
|
||||
@@ -433,7 +439,7 @@ class DataParallelController:
|
||||
|
||||
try:
|
||||
# Send handshake with our node rank
|
||||
sock_send(req_socket, str(node_rank).encode())
|
||||
sock_send(req_socket, wrap_as_pickle(str(node_rank)))
|
||||
|
||||
# Receive worker ports
|
||||
worker_ports = sock_recv(req_socket)
|
||||
|
||||
@@ -12,20 +12,19 @@
|
||||
# limitations under the License.
|
||||
# ==============================================================================
|
||||
"""
|
||||
Dataclasses for embedding injection.
|
||||
Structs for embedding injection.
|
||||
|
||||
These are placed in a separate module to avoid circular imports between
|
||||
io_struct.py and schedule_batch.py.
|
||||
"""
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import List
|
||||
|
||||
import msgspec
|
||||
import torch
|
||||
|
||||
|
||||
@dataclass
|
||||
class PositionalEmbeds:
|
||||
class PositionalEmbeds(msgspec.Struct, array_like=True):
|
||||
"""Embeddings to place at specific token positions.
|
||||
|
||||
Accepts either a list of [1, hidden_dim] tensors or a pre-stacked [N, hidden_dim] tensor.
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -53,6 +53,8 @@ from sglang.srt.managers.io_struct import (
|
||||
async_sock_send,
|
||||
sock_recv,
|
||||
sock_send,
|
||||
unwrap_from_pickle,
|
||||
wrap_as_pickle,
|
||||
)
|
||||
from sglang.srt.managers.load_snapshot import (
|
||||
create_load_snapshot_reader,
|
||||
@@ -122,17 +124,29 @@ def _extract_field_by_index(
|
||||
if field is None:
|
||||
return None
|
||||
|
||||
should_wrap_result = field_name in ("customized_info", "time_stats")
|
||||
if should_wrap_result:
|
||||
field = unwrap_from_pickle(field)
|
||||
if field is None:
|
||||
return None
|
||||
|
||||
if isinstance(field, dict):
|
||||
new_field = {}
|
||||
for k, v in field.items():
|
||||
new_field[k] = v[index] if len(v) > index else None
|
||||
if len(v) > index:
|
||||
new_field[k] = [v[index]] if should_wrap_result else v[index]
|
||||
else:
|
||||
new_field[k] = [None] if should_wrap_result else None
|
||||
if should_wrap_result:
|
||||
return wrap_as_pickle(new_field) if new_field else None
|
||||
return new_field
|
||||
|
||||
if check_length:
|
||||
if len(field) <= index:
|
||||
return None
|
||||
|
||||
return [field[index]]
|
||||
new_field = [field[index]]
|
||||
return wrap_as_pickle(new_field) if should_wrap_result else new_field
|
||||
|
||||
|
||||
def _handle_output_by_index(output, i):
|
||||
|
||||
@@ -266,6 +266,7 @@ from sglang.srt.utils.hf_transformers_utils import (
|
||||
get_tokenizer,
|
||||
get_tokenizer_from_processor,
|
||||
)
|
||||
from sglang.srt.utils.msgspec_utils import msgspec_to_builtins
|
||||
from sglang.srt.utils.numa_utils import get_numa_node_if_available, numa_bind_to_node
|
||||
from sglang.srt.utils.nvtx_utils import scheduler_nvtx_method
|
||||
from sglang.srt.utils.tensor_bridge import use_mlx
|
||||
@@ -3752,7 +3753,7 @@ class Scheduler(
|
||||
# This field is not serializable.
|
||||
ret.pop("model_config", None)
|
||||
|
||||
return GetInternalStateReqOutput(internal_state=ret)
|
||||
return GetInternalStateReqOutput(internal_state=msgspec_to_builtins(ret))
|
||||
|
||||
def set_internal_state(self, recv_req: SetInternalStateReq):
|
||||
server_args_dict = recv_req.server_args
|
||||
@@ -3801,7 +3802,7 @@ class Scheduler(
|
||||
server_args.pop("model_config", None)
|
||||
return SetInternalStateReqOutput(
|
||||
updated=if_success,
|
||||
server_args=server_args,
|
||||
server_args=msgspec_to_builtins(server_args),
|
||||
)
|
||||
|
||||
def save_remote_model(self, **kwargs):
|
||||
|
||||
@@ -1,6 +1,5 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import dataclasses
|
||||
import time
|
||||
from dataclasses import dataclass
|
||||
from typing import (
|
||||
@@ -10,13 +9,14 @@ from typing import (
|
||||
Optional,
|
||||
)
|
||||
|
||||
import msgspec
|
||||
import zmq
|
||||
|
||||
from sglang.srt.disaggregation.kv_events import (
|
||||
EventPublisherFactory,
|
||||
KVEventBatch,
|
||||
)
|
||||
from sglang.srt.managers.io_struct import sock_send
|
||||
from sglang.srt.managers.io_struct import hook_custom_types, sock_send
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
|
||||
@@ -26,8 +26,7 @@ if TYPE_CHECKING:
|
||||
class SchedulerStats: ... # type: ignore[no-redef]
|
||||
|
||||
|
||||
@dataclasses.dataclass
|
||||
class KvMetrics:
|
||||
class KvMetrics(msgspec.Struct, tag=True, kw_only=True, array_like=True):
|
||||
request_active_slots: int = 0
|
||||
request_total_slots: int = 0
|
||||
kv_active_blocks: int = 0
|
||||
@@ -38,6 +37,9 @@ class KvMetrics:
|
||||
data_parallel_rank: int = 0
|
||||
|
||||
|
||||
hook_custom_types(KvMetrics)
|
||||
|
||||
|
||||
@dataclass(kw_only=True, slots=True)
|
||||
class SchedulerKvEventsPublisher:
|
||||
kv_events_config: Optional[str]
|
||||
|
||||
@@ -19,6 +19,7 @@ from sglang.srt.managers.io_struct import (
|
||||
BatchEmbeddingOutput,
|
||||
BatchTokenIDOutput,
|
||||
CachedTokensDetails,
|
||||
wrap_as_pickle,
|
||||
)
|
||||
from sglang.srt.managers.schedule_batch import (
|
||||
BaseFinishReason,
|
||||
@@ -229,7 +230,7 @@ class SchedulerOutputStreamer:
|
||||
BatchEmbeddingOutput(
|
||||
rids=rids,
|
||||
http_worker_ipcs=http_worker_ipcs,
|
||||
time_stats=time_stats,
|
||||
time_stats=wrap_as_pickle(time_stats),
|
||||
finished_reasons=finished_reasons,
|
||||
embeddings=embeddings,
|
||||
prompt_tokens=prompt_tokens,
|
||||
@@ -516,7 +517,7 @@ class _GenerationStreamAccumulator:
|
||||
spec_verify_ct=self.spec_verify_ct,
|
||||
spec_num_correct_drafts=self.spec_num_correct_drafts,
|
||||
spec_correct_drafts_histogram=self.spec_correct_drafts_histogram,
|
||||
time_stats=self.time_stats,
|
||||
time_stats=wrap_as_pickle(self.time_stats),
|
||||
finished_reasons=self.finished_reasons,
|
||||
decoded_texts=self.decoded_texts,
|
||||
decode_ids=self.decode_ids_list,
|
||||
@@ -549,7 +550,9 @@ class _GenerationStreamAccumulator:
|
||||
output_hidden_states=self.output_hidden_states,
|
||||
routed_experts=self.routed_experts,
|
||||
indexer_topk=self.indexer_topk,
|
||||
customized_info=self.customized_info,
|
||||
customized_info=(
|
||||
wrap_as_pickle(self.customized_info) if self.customized_info else None
|
||||
),
|
||||
placeholder_tokens_idx=None,
|
||||
placeholder_tokens_val=None,
|
||||
retraction_counts=self.retraction_counts,
|
||||
|
||||
@@ -89,6 +89,9 @@ class SchedulerRequestReceiver:
|
||||
|
||||
recv_reqs = self._broadcast_reqs_across_ranks(recv_reqs)
|
||||
|
||||
if self.ps.pp_rank == 0:
|
||||
self.unwrap_pickle_wrapper(recv_reqs)
|
||||
|
||||
recv_reqs = self._apply_mm_receiver(recv_reqs)
|
||||
|
||||
self._finalize_shm_features(recv_reqs)
|
||||
@@ -195,6 +198,19 @@ class SchedulerRequestReceiver:
|
||||
)
|
||||
return recv_reqs
|
||||
|
||||
def unwrap_pickle_wrapper(self, recv_reqs: Optional[List]) -> None:
|
||||
if not recv_reqs:
|
||||
return
|
||||
|
||||
for req in recv_reqs:
|
||||
if isinstance(req, (TokenizedGenerateReqInput, TokenizedEmbeddingReqInput)):
|
||||
req.unwrap_pickle_fields()
|
||||
elif isinstance(
|
||||
req, (BatchTokenizedGenerateReqInput, BatchTokenizedEmbeddingReqInput)
|
||||
):
|
||||
for sub_req in req:
|
||||
sub_req.unwrap_pickle_fields()
|
||||
|
||||
def _apply_mm_receiver(self, recv_reqs: List) -> List:
|
||||
# Process MM requests under EPD-disaggregation mode
|
||||
if (
|
||||
|
||||
@@ -80,6 +80,7 @@ from sglang.srt.managers.io_struct import (
|
||||
async_sock_recv,
|
||||
async_sock_send,
|
||||
sock_send,
|
||||
unwrap_from_pickle,
|
||||
)
|
||||
from sglang.srt.managers.load_snapshot import create_load_snapshot_reader
|
||||
from sglang.srt.managers.mm_utils import TensorTransportMode, wrap_shm_features
|
||||
@@ -1332,7 +1333,10 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
||||
):
|
||||
tokenized_obj.time_stats.set_api_server_dispatch_time()
|
||||
tokenized_obj = wrap_shm_features(tokenized_obj)
|
||||
time_stats = tokenized_obj.time_stats
|
||||
tokenized_obj.wrap_pickle_fields()
|
||||
self._dispatch_to_scheduler(tokenized_obj)
|
||||
tokenized_obj.time_stats = time_stats
|
||||
tokenized_obj.time_stats.set_api_server_dispatch_finish_time()
|
||||
|
||||
def _send_batch_request(
|
||||
@@ -1342,13 +1346,19 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
||||
],
|
||||
):
|
||||
"""Send a batch of tokenized requests as a single batched request to the scheduler."""
|
||||
set_time_batch(tokenized_objs, "set_api_server_dispatch_time")
|
||||
time_stats = [tokenized_obj.time_stats for tokenized_obj in tokenized_objs]
|
||||
for tokenized_obj in tokenized_objs:
|
||||
tokenized_obj.wrap_pickle_fields()
|
||||
|
||||
if isinstance(tokenized_objs[0], TokenizedGenerateReqInput):
|
||||
batch_req = BatchTokenizedGenerateReqInput(batch=tokenized_objs)
|
||||
else:
|
||||
batch_req = BatchTokenizedEmbeddingReqInput(batch=tokenized_objs)
|
||||
|
||||
set_time_batch(tokenized_objs, "set_api_server_dispatch_time")
|
||||
self._dispatch_to_scheduler(batch_req)
|
||||
for tokenized_obj, time_stat in zip(tokenized_objs, time_stats):
|
||||
tokenized_obj.time_stats = time_stat
|
||||
set_time_batch(tokenized_objs, "set_api_server_dispatch_finish_time")
|
||||
|
||||
def _coalesce_streaming_chunks(
|
||||
@@ -1856,6 +1866,11 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
||||
BatchTokenIDOutput,
|
||||
],
|
||||
):
|
||||
recv_obj.time_stats = unwrap_from_pickle(recv_obj.time_stats)
|
||||
if isinstance(recv_obj, (BatchStrOutput, BatchTokenIDOutput)):
|
||||
customized_info = unwrap_from_pickle(recv_obj.customized_info)
|
||||
else:
|
||||
customized_info = None
|
||||
pending_notify: dict[str, ReqState] = {}
|
||||
batch_notify_size = self.server_args.batch_notify_size
|
||||
for i, rid in enumerate(recv_obj.rids):
|
||||
@@ -1911,8 +1926,8 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
||||
meta_info["cached_tokens_details"] = recv_obj.cached_tokens_details[
|
||||
i
|
||||
]
|
||||
if recv_obj.customized_info is not None:
|
||||
for k, v in recv_obj.customized_info.items():
|
||||
if customized_info is not None:
|
||||
for k, v in customized_info.items():
|
||||
if k not in state.customized_info_accumulated:
|
||||
state.customized_info_accumulated[k] = []
|
||||
state.customized_info_accumulated[k].extend(v[i])
|
||||
|
||||
@@ -15,7 +15,7 @@
|
||||
|
||||
import logging
|
||||
import math
|
||||
from typing import Any, Dict, List, Optional, Set, Union
|
||||
from typing import Dict, List, Optional, Set, Union
|
||||
|
||||
import msgspec
|
||||
|
||||
@@ -25,6 +25,17 @@ try:
|
||||
except ImportError:
|
||||
import sre_parse # Python < 3.11
|
||||
|
||||
# JSON-safe value types for custom_params. Must survive msgpack IPC
|
||||
# without PickleWrapper. After deserialization on the scheduler side,
|
||||
# Req.__init__ injects "__req__" (a Req object) into the dict in-process;
|
||||
# that augmented dict is never re-serialized.
|
||||
_JsonScalar = Union[None, bool, int, float, str]
|
||||
CustomParamValue = Union[
|
||||
_JsonScalar,
|
||||
List[_JsonScalar],
|
||||
Dict[str, _JsonScalar],
|
||||
]
|
||||
|
||||
_SAMPLING_EPS = 1e-6
|
||||
TOP_K_ALL = 1 << 30
|
||||
|
||||
@@ -96,7 +107,7 @@ class SamplingParams(msgspec.Struct, kw_only=True, omit_defaults=True):
|
||||
skip_special_tokens: bool = True
|
||||
spaces_between_special_tokens: bool = True
|
||||
no_stop_trim: bool = False
|
||||
custom_params: Optional[Dict[str, Any]] = None
|
||||
custom_params: Optional[Dict[str, CustomParamValue]] = None
|
||||
stream_interval: Optional[int] = None
|
||||
logit_bias: Optional[Dict[str, float]] = None
|
||||
sampling_seed: Optional[int] = None
|
||||
|
||||
@@ -0,0 +1,108 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import base64
|
||||
import binascii
|
||||
from typing import Any
|
||||
|
||||
import msgspec
|
||||
from pydantic_core import core_schema
|
||||
|
||||
|
||||
class Base64Bytes:
|
||||
"""Pydantic marker for HTTP JSON base64-encoded bytes fields."""
|
||||
|
||||
def __get_pydantic_core_schema__(self, source_type: Any, handler):
|
||||
return core_schema.no_info_before_validator_function(
|
||||
self._decode_value,
|
||||
handler(source_type),
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def _decode_value(cls, value: Any) -> Any:
|
||||
if isinstance(value, str):
|
||||
try:
|
||||
return base64.b64decode(value, validate=True)
|
||||
except binascii.Error as exc:
|
||||
raise ValueError("Expected base64-encoded bytes") from exc
|
||||
|
||||
if isinstance(value, list):
|
||||
return [cls._decode_value(item) for item in value]
|
||||
|
||||
if isinstance(value, tuple):
|
||||
return tuple(cls._decode_value(item) for item in value)
|
||||
|
||||
return value
|
||||
|
||||
|
||||
def msgspec_to_builtins(obj: Any) -> Any:
|
||||
"""Recursively convert msgspec structs to dict/list Python builtins."""
|
||||
if isinstance(obj, msgspec.Struct):
|
||||
return {
|
||||
field.name: msgspec_to_builtins(getattr(obj, field.name))
|
||||
for field in msgspec.structs.fields(type(obj))
|
||||
}
|
||||
|
||||
if isinstance(obj, dict):
|
||||
return {key: msgspec_to_builtins(value) for key, value in obj.items()}
|
||||
|
||||
if isinstance(obj, list):
|
||||
return [msgspec_to_builtins(item) for item in obj]
|
||||
|
||||
if isinstance(obj, tuple):
|
||||
return tuple(msgspec_to_builtins(item) for item in obj)
|
||||
|
||||
if isinstance(obj, set):
|
||||
return [msgspec_to_builtins(item) for item in obj]
|
||||
|
||||
return obj
|
||||
|
||||
|
||||
def msgspec_struct_pydantic_core_schema(cls: type[msgspec.Struct], handler):
|
||||
fields = {}
|
||||
for struct_field in msgspec.structs.fields(cls):
|
||||
field_schema = handler.generate_schema(struct_field.type)
|
||||
required = (
|
||||
struct_field.default is msgspec.NODEFAULT
|
||||
and struct_field.default_factory is msgspec.NODEFAULT
|
||||
)
|
||||
|
||||
if struct_field.default is not msgspec.NODEFAULT:
|
||||
field_schema = core_schema.with_default_schema(
|
||||
field_schema,
|
||||
default=struct_field.default,
|
||||
)
|
||||
elif struct_field.default_factory is not msgspec.NODEFAULT:
|
||||
field_schema = core_schema.with_default_schema(
|
||||
field_schema,
|
||||
default_factory=struct_field.default_factory,
|
||||
)
|
||||
|
||||
fields[struct_field.name] = core_schema.typed_dict_field(
|
||||
field_schema,
|
||||
required=required,
|
||||
)
|
||||
|
||||
typed_dict_schema = core_schema.typed_dict_schema(
|
||||
fields,
|
||||
cls_name=cls.__name__,
|
||||
extra_behavior="ignore",
|
||||
ref=cls.__name__,
|
||||
)
|
||||
|
||||
def build_struct(value):
|
||||
return value if isinstance(value, cls) else cls(**value)
|
||||
|
||||
dict_to_struct_schema = core_schema.no_info_after_validator_function(
|
||||
build_struct,
|
||||
typed_dict_schema,
|
||||
)
|
||||
return core_schema.json_or_python_schema(
|
||||
json_schema=dict_to_struct_schema,
|
||||
python_schema=core_schema.union_schema(
|
||||
[
|
||||
core_schema.is_instance_schema(cls),
|
||||
dict_to_struct_schema,
|
||||
],
|
||||
mode="left_to_right",
|
||||
),
|
||||
)
|
||||
@@ -12,7 +12,7 @@ import zmq
|
||||
|
||||
from sglang.srt.entrypoints.http_server import launch_server
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.managers.io_struct import sock_recv, sock_send
|
||||
from sglang.srt.managers.io_struct import sock_recv, sock_send, wrap_as_pickle
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
from sglang.srt.utils.network import get_free_port, get_zmq_socket_on_host
|
||||
from sglang.test.scripted_runtime.io_struct import (
|
||||
@@ -90,7 +90,7 @@ class ScriptedHttpServer:
|
||||
raise RuntimeError(f"ScriptedHttpServer is dirty: {self._dirty}")
|
||||
|
||||
fn_path = f"{script_fn.__module__}:{script_fn.__qualname__}"
|
||||
sock_send(self._socket, RunScript(fn_path=fn_path, args=args))
|
||||
sock_send(self._socket, wrap_as_pickle(RunScript(fn_path=fn_path, args=args)))
|
||||
|
||||
if not self._socket.poll(int(timeout_s * 1000)):
|
||||
if not self._server_process.is_alive():
|
||||
@@ -117,7 +117,7 @@ class ScriptedHttpServer:
|
||||
fatal_error: Optional[OutOfBandError] = None
|
||||
try:
|
||||
try:
|
||||
sock_send(self._socket, Shutdown())
|
||||
sock_send(self._socket, wrap_as_pickle(Shutdown()))
|
||||
except zmq.ZMQError:
|
||||
pass
|
||||
|
||||
|
||||
@@ -11,7 +11,7 @@ from typing import TYPE_CHECKING, Generator, List, Optional, Tuple
|
||||
import zmq
|
||||
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.managers.io_struct import sock_recv, sock_send
|
||||
from sglang.srt.managers.io_struct import sock_recv, sock_send, wrap_as_pickle
|
||||
from sglang.srt.utils.network import get_zmq_socket
|
||||
from sglang.test.scripted_runtime.background_http_poster import BackgroundHttpPoster
|
||||
from sglang.test.scripted_runtime.context import ScriptedContext
|
||||
@@ -153,7 +153,7 @@ class ScriptedSchedulerHook:
|
||||
socket = get_zmq_socket(ctx_zmq, zmq.PAIR, endpoint, bind=False)
|
||||
try:
|
||||
yield from _drive_engine_through_warmup(self._context)
|
||||
sock_send(socket, HookReady())
|
||||
sock_send(socket, wrap_as_pickle(HookReady()))
|
||||
while True:
|
||||
msg = sock_recv(socket)
|
||||
match msg:
|
||||
@@ -170,10 +170,12 @@ class ScriptedSchedulerHook:
|
||||
except Exception:
|
||||
sock_send(
|
||||
socket,
|
||||
ScriptFailed(traceback=traceback.format_exc()),
|
||||
wrap_as_pickle(
|
||||
ScriptFailed(traceback=traceback.format_exc())
|
||||
),
|
||||
)
|
||||
else:
|
||||
sock_send(socket, ScriptSucceeded())
|
||||
sock_send(socket, wrap_as_pickle(ScriptSucceeded()))
|
||||
case _:
|
||||
raise ValueError(f"dispatch loop: unknown command {msg!r}")
|
||||
finally:
|
||||
|
||||
@@ -41,6 +41,11 @@ class ScriptedTokenizerRecvProxy:
|
||||
"ScriptedTokenizerRecvProxy.recv_pyobj: blocking recv is not supported"
|
||||
)
|
||||
|
||||
def recv(self, flags: int = 0) -> bytes:
|
||||
raise NotImplementedError(
|
||||
"TODO: support ScriptedTokenizerRecvProxy.recv for msgpack IPC"
|
||||
)
|
||||
|
||||
def wait_until_arrived(
|
||||
self,
|
||||
predicate: Callable[[Any], bool],
|
||||
|
||||
Reference in New Issue
Block a user