[diffusion] optimize: optimize frame returns path (#24616)
This commit is contained in:
@@ -0,0 +1,92 @@
|
|||||||
|
# SPDX-License-Identifier: Apache-2.0
|
||||||
|
"""Helpers for transferring large numpy arrays between local scheduler processes."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import os
|
||||||
|
import tempfile
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
import numpy as np
|
||||||
|
|
||||||
|
_MIN_FILE_REF_BYTES = 32 << 20
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class NumpyArrayFileRef:
|
||||||
|
path: str
|
||||||
|
|
||||||
|
def materialize(self) -> np.ndarray:
|
||||||
|
try:
|
||||||
|
return np.load(self.path, allow_pickle=False)
|
||||||
|
finally:
|
||||||
|
try:
|
||||||
|
os.unlink(self.path)
|
||||||
|
except FileNotFoundError:
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
def is_local_endpoint(endpoint: str) -> bool:
|
||||||
|
return endpoint.startswith(
|
||||||
|
("tcp://127.0.0.1:", "tcp://localhost:", "ipc://", "inproc://")
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def spill_large_arrays_to_file_refs(value: Any) -> Any:
|
||||||
|
directory = _array_ipc_dir()
|
||||||
|
if directory is None:
|
||||||
|
return value
|
||||||
|
return _spill_large_arrays_to_file_refs(value, directory)
|
||||||
|
|
||||||
|
|
||||||
|
def _spill_large_arrays_to_file_refs(value: Any, directory: str) -> Any:
|
||||||
|
if isinstance(value, np.ndarray) and value.nbytes >= _MIN_FILE_REF_BYTES:
|
||||||
|
# only spill if the array size is above the threshold. if not, it's not worth it
|
||||||
|
return _spill_array(value, directory)
|
||||||
|
if isinstance(value, list):
|
||||||
|
return [_spill_large_arrays_to_file_refs(item, directory) for item in value]
|
||||||
|
if isinstance(value, tuple):
|
||||||
|
return tuple(
|
||||||
|
_spill_large_arrays_to_file_refs(item, directory) for item in value
|
||||||
|
)
|
||||||
|
return value
|
||||||
|
|
||||||
|
|
||||||
|
def materialize_file_refs(value: Any) -> Any:
|
||||||
|
if isinstance(value, NumpyArrayFileRef):
|
||||||
|
return value.materialize()
|
||||||
|
if isinstance(value, list):
|
||||||
|
return [materialize_file_refs(item) for item in value]
|
||||||
|
if isinstance(value, tuple):
|
||||||
|
return tuple(materialize_file_refs(item) for item in value)
|
||||||
|
return value
|
||||||
|
|
||||||
|
|
||||||
|
def _spill_array(array: np.ndarray, directory: str) -> NumpyArrayFileRef:
|
||||||
|
if not array.flags.c_contiguous:
|
||||||
|
array = np.ascontiguousarray(array)
|
||||||
|
|
||||||
|
fd, path = tempfile.mkstemp(
|
||||||
|
prefix="sgldiffusion-array-",
|
||||||
|
suffix=".npy",
|
||||||
|
dir=directory,
|
||||||
|
)
|
||||||
|
try:
|
||||||
|
with os.fdopen(fd, "wb") as f:
|
||||||
|
np.save(f, array, allow_pickle=False)
|
||||||
|
except Exception:
|
||||||
|
try:
|
||||||
|
os.unlink(path)
|
||||||
|
except FileNotFoundError:
|
||||||
|
pass
|
||||||
|
raise
|
||||||
|
return NumpyArrayFileRef(path=path)
|
||||||
|
|
||||||
|
|
||||||
|
def _array_ipc_dir() -> str | None:
|
||||||
|
shm_path = Path("/dev/shm")
|
||||||
|
if shm_path.is_dir() and os.access(shm_path, os.W_OK):
|
||||||
|
return str(shm_path)
|
||||||
|
return None
|
||||||
@@ -10,6 +10,7 @@ from contextlib import ExitStack
|
|||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass, field
|
||||||
from typing import Any, Callable, List, Union
|
from typing import Any, Callable, List, Union
|
||||||
|
|
||||||
|
import numpy as np
|
||||||
import torch
|
import torch
|
||||||
from setproctitle import setproctitle
|
from setproctitle import setproctitle
|
||||||
|
|
||||||
@@ -31,7 +32,10 @@ from sglang.multimodal_gen.runtime.distributed.parallel_state import (
|
|||||||
get_ulysses_parallel_rank,
|
get_ulysses_parallel_rank,
|
||||||
get_ulysses_parallel_world_size,
|
get_ulysses_parallel_world_size,
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.entrypoints.utils import save_outputs
|
from sglang.multimodal_gen.runtime.entrypoints.utils import (
|
||||||
|
post_process_sample,
|
||||||
|
save_outputs,
|
||||||
|
)
|
||||||
from sglang.multimodal_gen.runtime.loader.weight_utils import compute_weights_checksum
|
from sglang.multimodal_gen.runtime.loader.weight_utils import compute_weights_checksum
|
||||||
from sglang.multimodal_gen.runtime.loader.weights_updater import (
|
from sglang.multimodal_gen.runtime.loader.weights_updater import (
|
||||||
WeightsUpdater,
|
WeightsUpdater,
|
||||||
@@ -366,6 +370,9 @@ class GPUWorker:
|
|||||||
if torch.cuda.is_initialized():
|
if torch.cuda.is_initialized():
|
||||||
torch.cuda.empty_cache()
|
torch.cuda.empty_cache()
|
||||||
|
|
||||||
|
# Keep return_frames payloads off the scheduler's tensor ZMQ path.
|
||||||
|
self._materialize_frame_outputs_for_return(output_batch, req)
|
||||||
|
|
||||||
if torch.cuda.is_initialized() and output_batch.output is None:
|
if torch.cuda.is_initialized() and output_batch.output is None:
|
||||||
torch.cuda.empty_cache()
|
torch.cuda.empty_cache()
|
||||||
|
|
||||||
@@ -401,6 +408,72 @@ class GPUWorker:
|
|||||||
torch.cuda.empty_cache()
|
torch.cuda.empty_cache()
|
||||||
return output_batch
|
return output_batch
|
||||||
|
|
||||||
|
def _materialize_frame_outputs_for_return(
|
||||||
|
self, output_batch: OutputBatch, req: Req
|
||||||
|
) -> None:
|
||||||
|
if self.rank != 0 or output_batch.output is None or not req.return_frames:
|
||||||
|
return
|
||||||
|
|
||||||
|
if (
|
||||||
|
os.environ.get("SGLANG_DIFFUSION_SYNC_STAGE_PROFILING", "0") == "1"
|
||||||
|
and torch.cuda.is_initialized()
|
||||||
|
):
|
||||||
|
torch.cuda.synchronize()
|
||||||
|
start_time = time.perf_counter()
|
||||||
|
output_batch.output = [
|
||||||
|
self._materialize_frame_output(output, output_batch, req)
|
||||||
|
for output in output_batch.output
|
||||||
|
]
|
||||||
|
if output_batch.metrics is not None:
|
||||||
|
if (
|
||||||
|
os.environ.get("SGLANG_DIFFUSION_SYNC_STAGE_PROFILING", "0") == "1"
|
||||||
|
and torch.cuda.is_initialized()
|
||||||
|
):
|
||||||
|
torch.cuda.synchronize()
|
||||||
|
output_batch.metrics.record_stage(
|
||||||
|
"GPUWorker.frame_materialize_for_return",
|
||||||
|
time.perf_counter() - start_time,
|
||||||
|
)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _materialize_frame_output(
|
||||||
|
output: Any, output_batch: OutputBatch, req: Req
|
||||||
|
) -> np.ndarray:
|
||||||
|
if (
|
||||||
|
isinstance(output, torch.Tensor)
|
||||||
|
and not req.enable_frame_interpolation
|
||||||
|
and not req.enable_upscaling
|
||||||
|
):
|
||||||
|
if output.dim() == 3:
|
||||||
|
output = output.unsqueeze(1)
|
||||||
|
output = (output * 255).clamp(0, 255).to(torch.uint8)
|
||||||
|
return output.permute(1, 2, 3, 0).cpu().numpy()
|
||||||
|
|
||||||
|
if (
|
||||||
|
isinstance(output, np.ndarray)
|
||||||
|
and output.dtype == np.uint8
|
||||||
|
and output.ndim == 4
|
||||||
|
and output.shape[-1] in (1, 3, 4)
|
||||||
|
):
|
||||||
|
return output
|
||||||
|
|
||||||
|
frames = post_process_sample(
|
||||||
|
output,
|
||||||
|
req.data_type,
|
||||||
|
req.fps,
|
||||||
|
save_output=False,
|
||||||
|
audio_sample_rate=output_batch.audio_sample_rate,
|
||||||
|
output_compression=req.output_compression,
|
||||||
|
enable_frame_interpolation=req.enable_frame_interpolation,
|
||||||
|
frame_interpolation_exp=req.frame_interpolation_exp,
|
||||||
|
frame_interpolation_scale=req.frame_interpolation_scale,
|
||||||
|
frame_interpolation_model_path=req.frame_interpolation_model_path,
|
||||||
|
enable_upscaling=req.enable_upscaling,
|
||||||
|
upscaling_model_path=req.upscaling_model_path,
|
||||||
|
upscaling_scale=req.upscaling_scale,
|
||||||
|
)
|
||||||
|
return np.asarray(frames)
|
||||||
|
|
||||||
def _record_output_peak_memory(self, output_batch: OutputBatch) -> None:
|
def _record_output_peak_memory(self, output_batch: OutputBatch) -> None:
|
||||||
if self.rank != 0 or current_platform.is_cpu():
|
if self.rank != 0 or current_platform.is_cpu():
|
||||||
return
|
return
|
||||||
@@ -492,6 +565,7 @@ class GPUWorker:
|
|||||||
first_req = reqs[0]
|
first_req = reqs[0]
|
||||||
shared_output_fields = (
|
shared_output_fields = (
|
||||||
"save_output",
|
"save_output",
|
||||||
|
"return_frames",
|
||||||
"return_file_paths_only",
|
"return_file_paths_only",
|
||||||
"data_type",
|
"data_type",
|
||||||
"fps",
|
"fps",
|
||||||
|
|||||||
@@ -8,9 +8,10 @@ import pickle
|
|||||||
import tempfile
|
import tempfile
|
||||||
import time
|
import time
|
||||||
from collections import deque
|
from collections import deque
|
||||||
|
from contextlib import contextmanager
|
||||||
from copy import deepcopy
|
from copy import deepcopy
|
||||||
from enum import Enum
|
from enum import Enum
|
||||||
from typing import Any, List
|
from typing import Any, Iterator, List
|
||||||
|
|
||||||
import zmq
|
import zmq
|
||||||
|
|
||||||
@@ -35,6 +36,10 @@ from sglang.multimodal_gen.runtime.entrypoints.utils import (
|
|||||||
ShutdownReq,
|
ShutdownReq,
|
||||||
UnmergeLoraWeightsReq,
|
UnmergeLoraWeightsReq,
|
||||||
)
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.ipc_array import (
|
||||||
|
is_local_endpoint,
|
||||||
|
spill_large_arrays_to_file_refs,
|
||||||
|
)
|
||||||
from sglang.multimodal_gen.runtime.managers.cpu_worker import CPUWorker
|
from sglang.multimodal_gen.runtime.managers.cpu_worker import CPUWorker
|
||||||
from sglang.multimodal_gen.runtime.managers.dynamic_batch_admission import (
|
from sglang.multimodal_gen.runtime.managers.dynamic_batch_admission import (
|
||||||
BatchAdmissionController,
|
BatchAdmissionController,
|
||||||
@@ -608,7 +613,37 @@ class Scheduler(SchedulerDisaggMixin):
|
|||||||
replies to client, only on rank 0
|
replies to client, only on rank 0
|
||||||
"""
|
"""
|
||||||
if not is_warmup and self.receiver is not None and identity is not None:
|
if not is_warmup and self.receiver is not None and identity is not None:
|
||||||
self.receiver.send_multipart([identity, b"", pickle.dumps(output_batch)])
|
# if the server is local, use temp file to spill the frame array instead of
|
||||||
|
# leaving it in OutputBatch to be pickled later
|
||||||
|
if is_local_endpoint(self.server_args.scheduler_endpoint):
|
||||||
|
with self._record_return_stage(
|
||||||
|
output_batch, "Scheduler.return_result.spill_arrays"
|
||||||
|
):
|
||||||
|
output_batch.output = spill_large_arrays_to_file_refs(
|
||||||
|
output_batch.output
|
||||||
|
)
|
||||||
|
|
||||||
|
with self._record_return_stage(
|
||||||
|
output_batch, "Scheduler.return_result.pickle"
|
||||||
|
):
|
||||||
|
payload = pickle.dumps(output_batch)
|
||||||
|
|
||||||
|
with self._record_return_stage(
|
||||||
|
output_batch, "Scheduler.return_result.send"
|
||||||
|
):
|
||||||
|
self.receiver.send_multipart([identity, b"", payload])
|
||||||
|
|
||||||
|
@contextmanager
|
||||||
|
def _record_return_stage(
|
||||||
|
self, output_batch: OutputBatch, stage_name: str
|
||||||
|
) -> Iterator[None]:
|
||||||
|
"""helper function to record a stage metric"""
|
||||||
|
start_time = time.perf_counter()
|
||||||
|
yield
|
||||||
|
if output_batch.metrics is not None:
|
||||||
|
output_batch.metrics.record_stage(
|
||||||
|
stage_name, time.perf_counter() - start_time
|
||||||
|
)
|
||||||
|
|
||||||
def _try_merge_generation_reqs(self, reqs: List[Req]) -> Req | None:
|
def _try_merge_generation_reqs(self, reqs: List[Req]) -> Req | None:
|
||||||
"""Create a batched generation request from compatible requests.
|
"""Create a batched generation request from compatible requests.
|
||||||
|
|||||||
@@ -383,7 +383,7 @@ class OutputBatch:
|
|||||||
Final output (after pipeline completion)
|
Final output (after pipeline completion)
|
||||||
"""
|
"""
|
||||||
|
|
||||||
output: torch.Tensor | None = None
|
output: Any | None = None
|
||||||
audio: torch.Tensor | None = None
|
audio: torch.Tensor | None = None
|
||||||
audio_sample_rate: int | None = None
|
audio_sample_rate: int | None = None
|
||||||
trajectory_timesteps: torch.Tensor | None = None
|
trajectory_timesteps: torch.Tensor | None = None
|
||||||
|
|||||||
@@ -1,9 +1,12 @@
|
|||||||
import pickle
|
import pickle
|
||||||
|
import time
|
||||||
from typing import Any
|
from typing import Any
|
||||||
|
|
||||||
import zmq
|
import zmq
|
||||||
import zmq.asyncio
|
import zmq.asyncio
|
||||||
|
|
||||||
|
from sglang.multimodal_gen.runtime.ipc_array import materialize_file_refs
|
||||||
|
from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import OutputBatch
|
||||||
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
||||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||||
|
|
||||||
@@ -80,6 +83,7 @@ class SchedulerClient:
|
|||||||
try:
|
try:
|
||||||
self.scheduler_socket.send_pyobj(batch)
|
self.scheduler_socket.send_pyobj(batch)
|
||||||
output_batch = self.scheduler_socket.recv_pyobj()
|
output_batch = self.scheduler_socket.recv_pyobj()
|
||||||
|
_materialize_output_batch_file_refs(output_batch)
|
||||||
return output_batch
|
return output_batch
|
||||||
except zmq.error.Again:
|
except zmq.error.Again:
|
||||||
logger.error("Timeout waiting for response from scheduler.")
|
logger.error("Timeout waiting for response from scheduler.")
|
||||||
@@ -162,7 +166,9 @@ class AsyncSchedulerClient:
|
|||||||
try:
|
try:
|
||||||
await socket.send(pickle.dumps(batch))
|
await socket.send(pickle.dumps(batch))
|
||||||
payload = await socket.recv()
|
payload = await socket.recv()
|
||||||
return pickle.loads(payload)
|
output_batch = pickle.loads(payload)
|
||||||
|
_materialize_output_batch_file_refs(output_batch)
|
||||||
|
return output_batch
|
||||||
except zmq.error.Again:
|
except zmq.error.Again:
|
||||||
logger.error("Timeout waiting for response from scheduler.")
|
logger.error("Timeout waiting for response from scheduler.")
|
||||||
raise TimeoutError("Scheduler did not respond in time.")
|
raise TimeoutError("Scheduler did not respond in time.")
|
||||||
@@ -203,3 +209,16 @@ class AsyncSchedulerClient:
|
|||||||
# Singleton instances for easy access
|
# Singleton instances for easy access
|
||||||
async_scheduler_client = AsyncSchedulerClient()
|
async_scheduler_client = AsyncSchedulerClient()
|
||||||
sync_scheduler_client = SchedulerClient()
|
sync_scheduler_client = SchedulerClient()
|
||||||
|
|
||||||
|
|
||||||
|
def _materialize_output_batch_file_refs(output_batch: Any) -> None:
|
||||||
|
if not isinstance(output_batch, OutputBatch):
|
||||||
|
return
|
||||||
|
|
||||||
|
start_time = time.perf_counter()
|
||||||
|
output_batch.output = materialize_file_refs(output_batch.output)
|
||||||
|
if output_batch.metrics is not None:
|
||||||
|
output_batch.metrics.record_stage(
|
||||||
|
"SchedulerClient.materialize_file_refs",
|
||||||
|
time.perf_counter() - start_time,
|
||||||
|
)
|
||||||
|
|||||||
@@ -0,0 +1,81 @@
|
|||||||
|
# SPDX-License-Identifier: Apache-2.0
|
||||||
|
|
||||||
|
import tempfile
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
import numpy as np
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from sglang.multimodal_gen.runtime import ipc_array
|
||||||
|
from sglang.multimodal_gen.runtime.ipc_array import (
|
||||||
|
NumpyArrayFileRef,
|
||||||
|
is_local_endpoint,
|
||||||
|
materialize_file_refs,
|
||||||
|
spill_large_arrays_to_file_refs,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_spill_large_arrays_round_trips_and_removes_file(monkeypatch, tmp_path):
|
||||||
|
monkeypatch.setattr(ipc_array, "_array_ipc_dir", lambda: str(tmp_path))
|
||||||
|
array = np.arange(ipc_array._MIN_FILE_REF_BYTES, dtype=np.uint8)
|
||||||
|
|
||||||
|
spilled = spill_large_arrays_to_file_refs([array])
|
||||||
|
|
||||||
|
assert isinstance(spilled[0], NumpyArrayFileRef)
|
||||||
|
spilled_path = Path(spilled[0].path)
|
||||||
|
assert spilled_path.exists()
|
||||||
|
|
||||||
|
materialized = materialize_file_refs(spilled)
|
||||||
|
|
||||||
|
assert np.array_equal(materialized[0], array)
|
||||||
|
assert not spilled_path.exists()
|
||||||
|
|
||||||
|
|
||||||
|
def test_small_arrays_are_kept_inline():
|
||||||
|
array = np.arange(16, dtype=np.uint8)
|
||||||
|
|
||||||
|
spilled = spill_large_arrays_to_file_refs((array,))
|
||||||
|
|
||||||
|
assert spilled[0] is array
|
||||||
|
|
||||||
|
|
||||||
|
def test_large_arrays_are_kept_inline_without_shm(monkeypatch):
|
||||||
|
monkeypatch.setattr(ipc_array, "_array_ipc_dir", lambda: None)
|
||||||
|
array = np.arange(ipc_array._MIN_FILE_REF_BYTES, dtype=np.uint8)
|
||||||
|
|
||||||
|
spilled = spill_large_arrays_to_file_refs(array)
|
||||||
|
|
||||||
|
assert spilled is array
|
||||||
|
|
||||||
|
|
||||||
|
def test_spill_removes_temp_file_when_save_fails(monkeypatch, tmp_path):
|
||||||
|
monkeypatch.setattr(ipc_array, "_array_ipc_dir", lambda: str(tmp_path))
|
||||||
|
array = np.arange(ipc_array._MIN_FILE_REF_BYTES, dtype=np.uint8)
|
||||||
|
created_paths = []
|
||||||
|
|
||||||
|
def fail_save(*args, **kwargs):
|
||||||
|
raise OSError("simulated write failure")
|
||||||
|
|
||||||
|
original_mkstemp = tempfile.mkstemp
|
||||||
|
|
||||||
|
def tracked_mkstemp(*args, **kwargs):
|
||||||
|
fd, path = original_mkstemp(*args, **kwargs)
|
||||||
|
created_paths.append(Path(path))
|
||||||
|
return fd, path
|
||||||
|
|
||||||
|
monkeypatch.setattr(tempfile, "mkstemp", tracked_mkstemp)
|
||||||
|
monkeypatch.setattr(np, "save", fail_save)
|
||||||
|
|
||||||
|
with pytest.raises(OSError, match="simulated write failure"):
|
||||||
|
spill_large_arrays_to_file_refs(array)
|
||||||
|
|
||||||
|
assert created_paths
|
||||||
|
assert not created_paths[0].exists()
|
||||||
|
|
||||||
|
|
||||||
|
def test_local_endpoint_detection():
|
||||||
|
assert is_local_endpoint("tcp://127.0.0.1:30000")
|
||||||
|
assert is_local_endpoint("tcp://localhost:30000")
|
||||||
|
assert is_local_endpoint("ipc:///tmp/sgl.sock")
|
||||||
|
assert is_local_endpoint("inproc://scheduler")
|
||||||
|
assert not is_local_endpoint("tcp://10.0.0.2:30000")
|
||||||
Reference in New Issue
Block a user