[diffusion] feat: spill large tensors over shared memory like numpy arrays (#38656)
This commit is contained in:
@@ -1,5 +1,5 @@
|
|||||||
# SPDX-License-Identifier: Apache-2.0
|
# SPDX-License-Identifier: Apache-2.0
|
||||||
"""Helpers for transferring large numpy arrays between local scheduler processes."""
|
"""Helpers for transferring large arrays between local scheduler processes."""
|
||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
@@ -10,6 +10,11 @@ from pathlib import Path
|
|||||||
from typing import Any
|
from typing import Any
|
||||||
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||||
|
|
||||||
|
logger = init_logger(__name__)
|
||||||
|
|
||||||
_MIN_FILE_REF_BYTES = 32 << 20
|
_MIN_FILE_REF_BYTES = 32 << 20
|
||||||
|
|
||||||
@@ -28,6 +33,19 @@ class NumpyArrayFileRef:
|
|||||||
pass
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class TorchTensorFileRef:
|
||||||
|
"""A tensor spilled as raw bytes; dtype and shape are restored on the far side."""
|
||||||
|
|
||||||
|
ref: NumpyArrayFileRef
|
||||||
|
dtype: str
|
||||||
|
shape: tuple[int, ...]
|
||||||
|
|
||||||
|
def materialize(self) -> torch.Tensor:
|
||||||
|
flat = torch.from_numpy(self.ref.materialize())
|
||||||
|
return flat.view(getattr(torch, self.dtype)).reshape(self.shape)
|
||||||
|
|
||||||
|
|
||||||
def is_local_endpoint(endpoint: str) -> bool:
|
def is_local_endpoint(endpoint: str) -> bool:
|
||||||
return endpoint.startswith(
|
return endpoint.startswith(
|
||||||
("tcp://127.0.0.1:", "tcp://localhost:", "ipc://", "inproc://")
|
("tcp://127.0.0.1:", "tcp://localhost:", "ipc://", "inproc://")
|
||||||
@@ -42,9 +60,25 @@ def spill_large_arrays_to_file_refs(value: Any) -> Any:
|
|||||||
|
|
||||||
|
|
||||||
def _spill_large_arrays_to_file_refs(value: Any, directory: str) -> Any:
|
def _spill_large_arrays_to_file_refs(value: Any, directory: str) -> Any:
|
||||||
|
# A payload that does not fit in shared memory is still deliverable inline,
|
||||||
|
# so a full /dev/shm slows the reply down instead of failing the request.
|
||||||
if isinstance(value, np.ndarray) and value.nbytes >= _MIN_FILE_REF_BYTES:
|
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
|
# only spill if the array size is above the threshold. if not, it's not worth it
|
||||||
return _spill_array(value, directory)
|
try:
|
||||||
|
return _spill_array(value, directory)
|
||||||
|
except OSError:
|
||||||
|
logger.warning_once(
|
||||||
|
f"Spilling an array to {directory} failed; sending it inline."
|
||||||
|
)
|
||||||
|
return value
|
||||||
|
if _is_large_tensor(value):
|
||||||
|
try:
|
||||||
|
return _spill_tensor(value, directory)
|
||||||
|
except OSError:
|
||||||
|
logger.warning_once(
|
||||||
|
f"Spilling a tensor to {directory} failed; sending it inline."
|
||||||
|
)
|
||||||
|
return value
|
||||||
if isinstance(value, list):
|
if isinstance(value, list):
|
||||||
return [_spill_large_arrays_to_file_refs(item, directory) for item in value]
|
return [_spill_large_arrays_to_file_refs(item, directory) for item in value]
|
||||||
if isinstance(value, tuple):
|
if isinstance(value, tuple):
|
||||||
@@ -55,7 +89,7 @@ def _spill_large_arrays_to_file_refs(value: Any, directory: str) -> Any:
|
|||||||
|
|
||||||
|
|
||||||
def materialize_file_refs(value: Any) -> Any:
|
def materialize_file_refs(value: Any) -> Any:
|
||||||
if isinstance(value, NumpyArrayFileRef):
|
if isinstance(value, (NumpyArrayFileRef, TorchTensorFileRef)):
|
||||||
return value.materialize()
|
return value.materialize()
|
||||||
if isinstance(value, list):
|
if isinstance(value, list):
|
||||||
return [materialize_file_refs(item) for item in value]
|
return [materialize_file_refs(item) for item in value]
|
||||||
@@ -85,6 +119,25 @@ def _spill_array(array: np.ndarray, directory: str) -> NumpyArrayFileRef:
|
|||||||
return NumpyArrayFileRef(path=path)
|
return NumpyArrayFileRef(path=path)
|
||||||
|
|
||||||
|
|
||||||
|
def _is_large_tensor(value: Any) -> bool:
|
||||||
|
return (
|
||||||
|
isinstance(value, torch.Tensor)
|
||||||
|
and value.numel() * value.element_size() >= _MIN_FILE_REF_BYTES
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _spill_tensor(tensor: torch.Tensor, directory: str) -> TorchTensorFileRef:
|
||||||
|
host = tensor.detach().to("cpu", copy=False).contiguous()
|
||||||
|
# Spill the raw bytes: numpy has no bfloat16, and the byte view needs no
|
||||||
|
# dtype table to stay exact.
|
||||||
|
array = host.reshape(-1).view(torch.uint8).numpy()
|
||||||
|
return TorchTensorFileRef(
|
||||||
|
ref=_spill_array(array, directory),
|
||||||
|
dtype=str(host.dtype).removeprefix("torch."),
|
||||||
|
shape=tuple(host.shape),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def _array_ipc_dir() -> str | None:
|
def _array_ipc_dir() -> str | None:
|
||||||
shm_path = Path("/dev/shm")
|
shm_path = Path("/dev/shm")
|
||||||
if shm_path.is_dir() and os.access(shm_path, os.W_OK):
|
if shm_path.is_dir() and os.access(shm_path, os.W_OK):
|
||||||
|
|||||||
@@ -5,10 +5,12 @@ from pathlib import Path
|
|||||||
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import pytest
|
import pytest
|
||||||
|
import torch
|
||||||
|
|
||||||
from sglang.multimodal_gen.runtime import ipc_array
|
from sglang.multimodal_gen.runtime import ipc_array
|
||||||
from sglang.multimodal_gen.runtime.ipc_array import (
|
from sglang.multimodal_gen.runtime.ipc_array import (
|
||||||
NumpyArrayFileRef,
|
NumpyArrayFileRef,
|
||||||
|
TorchTensorFileRef,
|
||||||
is_local_endpoint,
|
is_local_endpoint,
|
||||||
materialize_file_refs,
|
materialize_file_refs,
|
||||||
spill_large_arrays_to_file_refs,
|
spill_large_arrays_to_file_refs,
|
||||||
@@ -66,8 +68,8 @@ def test_spill_removes_temp_file_when_save_fails(monkeypatch, tmp_path):
|
|||||||
monkeypatch.setattr(tempfile, "mkstemp", tracked_mkstemp)
|
monkeypatch.setattr(tempfile, "mkstemp", tracked_mkstemp)
|
||||||
monkeypatch.setattr(np, "save", fail_save)
|
monkeypatch.setattr(np, "save", fail_save)
|
||||||
|
|
||||||
with pytest.raises(OSError, match="simulated write failure"):
|
# A failed spill falls back to sending the payload inline.
|
||||||
spill_large_arrays_to_file_refs(array)
|
assert spill_large_arrays_to_file_refs(array) is array
|
||||||
|
|
||||||
assert created_paths
|
assert created_paths
|
||||||
assert not created_paths[0].exists()
|
assert not created_paths[0].exists()
|
||||||
@@ -79,3 +81,55 @@ def test_local_endpoint_detection():
|
|||||||
assert is_local_endpoint("ipc:///tmp/sgl.sock")
|
assert is_local_endpoint("ipc:///tmp/sgl.sock")
|
||||||
assert is_local_endpoint("inproc://scheduler")
|
assert is_local_endpoint("inproc://scheduler")
|
||||||
assert not is_local_endpoint("tcp://10.0.0.2:30000")
|
assert not is_local_endpoint("tcp://10.0.0.2:30000")
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize("dtype", [torch.float32, torch.bfloat16, torch.uint8])
|
||||||
|
def test_spill_large_tensors_round_trips(monkeypatch, tmp_path, dtype):
|
||||||
|
monkeypatch.setattr(ipc_array, "_array_ipc_dir", lambda: str(tmp_path))
|
||||||
|
elements = (
|
||||||
|
ipc_array._MIN_FILE_REF_BYTES // torch.empty((), dtype=dtype).element_size()
|
||||||
|
)
|
||||||
|
tensor = torch.arange(elements, dtype=torch.int64).to(dtype).reshape(2, -1)
|
||||||
|
|
||||||
|
spilled = spill_large_arrays_to_file_refs([tensor])
|
||||||
|
|
||||||
|
assert isinstance(spilled[0], TorchTensorFileRef)
|
||||||
|
spilled_path = Path(spilled[0].ref.path)
|
||||||
|
assert spilled_path.exists()
|
||||||
|
|
||||||
|
materialized = materialize_file_refs(spilled)[0]
|
||||||
|
|
||||||
|
assert materialized.dtype == tensor.dtype
|
||||||
|
assert materialized.shape == tensor.shape
|
||||||
|
assert torch.equal(materialized, tensor)
|
||||||
|
assert not spilled_path.exists()
|
||||||
|
|
||||||
|
|
||||||
|
def test_non_contiguous_tensor_round_trips(monkeypatch, tmp_path):
|
||||||
|
monkeypatch.setattr(ipc_array, "_array_ipc_dir", lambda: str(tmp_path))
|
||||||
|
elements = ipc_array._MIN_FILE_REF_BYTES // 2
|
||||||
|
tensor = torch.arange(elements, dtype=torch.float32).reshape(2, -1).T
|
||||||
|
|
||||||
|
materialized = materialize_file_refs(spill_large_arrays_to_file_refs(tensor))
|
||||||
|
|
||||||
|
assert torch.equal(materialized, tensor)
|
||||||
|
|
||||||
|
|
||||||
|
def test_small_tensors_are_kept_inline():
|
||||||
|
tensor = torch.zeros(16)
|
||||||
|
|
||||||
|
spilled = spill_large_arrays_to_file_refs((tensor,))
|
||||||
|
|
||||||
|
assert spilled[0] is tensor
|
||||||
|
|
||||||
|
|
||||||
|
def test_tensor_spill_falls_back_inline_when_shm_is_full(monkeypatch, tmp_path):
|
||||||
|
monkeypatch.setattr(ipc_array, "_array_ipc_dir", lambda: str(tmp_path))
|
||||||
|
tensor = torch.zeros(ipc_array._MIN_FILE_REF_BYTES, dtype=torch.uint8)
|
||||||
|
|
||||||
|
def fail_save(*args, **kwargs):
|
||||||
|
raise OSError("No space left on device")
|
||||||
|
|
||||||
|
monkeypatch.setattr(np, "save", fail_save)
|
||||||
|
|
||||||
|
assert spill_large_arrays_to_file_refs(tensor) is tensor
|
||||||
|
|||||||
Reference in New Issue
Block a user