[diffusion] feat: add weight source reader (#35668)
This commit is contained in:
@@ -0,0 +1,75 @@
|
|||||||
|
# SPDX-License-Identifier: Apache-2.0
|
||||||
|
"""Ways of reading checkpoint weights, and the rule for picking one."""
|
||||||
|
|
||||||
|
from sglang.multimodal_gen import envs
|
||||||
|
from sglang.multimodal_gen.runtime.loader.weight_readers.base import (
|
||||||
|
WeightReader,
|
||||||
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.loader.weight_readers.runai_streamer import (
|
||||||
|
RunaiStreamerReader,
|
||||||
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.loader.weight_readers.safetensors_mmap import (
|
||||||
|
SafetensorsMmapReader,
|
||||||
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||||
|
|
||||||
|
logger = init_logger(__name__)
|
||||||
|
|
||||||
|
# The fallback is last and always available, so selection cannot come up empty.
|
||||||
|
_READERS: tuple[type, ...] = (RunaiStreamerReader, SafetensorsMmapReader)
|
||||||
|
FALLBACK_READER = SafetensorsMmapReader
|
||||||
|
|
||||||
|
|
||||||
|
def available_reader_names() -> list[str]:
|
||||||
|
return [b.name for b in _READERS if b.is_available()]
|
||||||
|
|
||||||
|
|
||||||
|
def select_weight_reader(
|
||||||
|
*,
|
||||||
|
requested: str | None = None,
|
||||||
|
needs_key_filter: bool = False,
|
||||||
|
) -> WeightReader:
|
||||||
|
"""Pick a reader, honouring an explicit request where it can be honoured.
|
||||||
|
|
||||||
|
`requested` names a reader; None means take the environment's preference.
|
||||||
|
A reader that cannot skip keys is passed over when the caller needs to,
|
||||||
|
because reading the whole checkpoint to discard most of it is worse than
|
||||||
|
reading the part that was asked for more slowly.
|
||||||
|
"""
|
||||||
|
if requested is not None:
|
||||||
|
chosen = next((b for b in _READERS if b.name == requested), None)
|
||||||
|
if chosen is None:
|
||||||
|
raise ValueError(
|
||||||
|
f"unknown weight reader {requested!r}; "
|
||||||
|
f"available: {available_reader_names()}"
|
||||||
|
)
|
||||||
|
elif envs.SGLANG_USE_RUNAI_MODEL_STREAMER and RunaiStreamerReader.is_available():
|
||||||
|
chosen = RunaiStreamerReader
|
||||||
|
else:
|
||||||
|
chosen = FALLBACK_READER
|
||||||
|
|
||||||
|
if not chosen.is_available():
|
||||||
|
logger.info(
|
||||||
|
"Weight reader %s is not installed; using %s",
|
||||||
|
chosen.name,
|
||||||
|
FALLBACK_READER.name,
|
||||||
|
)
|
||||||
|
chosen = FALLBACK_READER
|
||||||
|
if needs_key_filter and not chosen.supports_key_filter:
|
||||||
|
logger.debug(
|
||||||
|
"Weight reader %s cannot skip keys at load time; using %s",
|
||||||
|
chosen.name,
|
||||||
|
FALLBACK_READER.name,
|
||||||
|
)
|
||||||
|
chosen = FALLBACK_READER
|
||||||
|
return chosen()
|
||||||
|
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
"FALLBACK_READER",
|
||||||
|
"RunaiStreamerReader",
|
||||||
|
"SafetensorsMmapReader",
|
||||||
|
"WeightReader",
|
||||||
|
"available_reader_names",
|
||||||
|
"select_weight_reader",
|
||||||
|
]
|
||||||
@@ -0,0 +1,47 @@
|
|||||||
|
# SPDX-License-Identifier: Apache-2.0
|
||||||
|
"""What a weight source has to provide, and what distinguishes one from another.
|
||||||
|
|
||||||
|
Reading a checkpoint used to be a boolean: Run:ai streamer, or `safe_open`. The
|
||||||
|
two differ in more than speed, and the differences decide correctness and memory
|
||||||
|
behaviour rather than taste, so they are stated here as capabilities:
|
||||||
|
|
||||||
|
``supports_key_filter``
|
||||||
|
Whether the reader can skip keys while reading. The streamer materializes
|
||||||
|
every tensor before handing any of them over, so a caller that only wants
|
||||||
|
part of a checkpoint cannot save anything by asking it.
|
||||||
|
|
||||||
|
``retains_file_mapping``
|
||||||
|
Whether the tensors it yields are views into the checkpoint file. Those
|
||||||
|
pages are file-backed, so the kernel can drop them under pressure without
|
||||||
|
swap; a reader that copies into anonymous memory gives the kernel nothing
|
||||||
|
to reclaim.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from typing import Callable, ClassVar, Iterator, Protocol, runtime_checkable
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
|
||||||
|
@runtime_checkable
|
||||||
|
class WeightReader(Protocol):
|
||||||
|
"""Yields ``(name, tensor)`` for every weight in a set of checkpoint files."""
|
||||||
|
|
||||||
|
name: ClassVar[str]
|
||||||
|
supports_key_filter: ClassVar[bool]
|
||||||
|
retains_file_mapping: ClassVar[bool]
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def is_available(cls) -> bool:
|
||||||
|
"""Whether this reader can run at all in this install."""
|
||||||
|
|
||||||
|
def iter_weights(
|
||||||
|
self,
|
||||||
|
files: list[str],
|
||||||
|
*,
|
||||||
|
device: str,
|
||||||
|
to_cpu: bool,
|
||||||
|
key_filter: Callable[[str], bool] | None = None,
|
||||||
|
clone_tensors: bool = True,
|
||||||
|
show_progress: bool = True,
|
||||||
|
) -> Iterator[tuple[str, torch.Tensor]]:
|
||||||
|
"""Iterate the weights, in whatever order the reader finds them."""
|
||||||
@@ -0,0 +1,56 @@
|
|||||||
|
# SPDX-License-Identifier: Apache-2.0
|
||||||
|
"""Run:ai Model Streamer: fastest to read, but it copies into anonymous memory."""
|
||||||
|
|
||||||
|
from typing import Callable, ClassVar, Iterator
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||||
|
|
||||||
|
logger = init_logger(__name__)
|
||||||
|
|
||||||
|
try:
|
||||||
|
from runai_model_streamer import SafetensorsStreamer
|
||||||
|
|
||||||
|
HAS_RUNAI_MODEL_STREAMER = True
|
||||||
|
except ImportError:
|
||||||
|
SafetensorsStreamer = None
|
||||||
|
HAS_RUNAI_MODEL_STREAMER = False
|
||||||
|
|
||||||
|
|
||||||
|
class RunaiStreamerReader:
|
||||||
|
name: ClassVar[str] = "runai_streamer"
|
||||||
|
# it materializes every tensor before yielding, so filtering saves nothing
|
||||||
|
supports_key_filter: ClassVar[bool] = False
|
||||||
|
retains_file_mapping: ClassVar[bool] = False
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def is_available(cls) -> bool:
|
||||||
|
return HAS_RUNAI_MODEL_STREAMER
|
||||||
|
|
||||||
|
def iter_weights(
|
||||||
|
self,
|
||||||
|
files: list[str],
|
||||||
|
*,
|
||||||
|
device: str,
|
||||||
|
to_cpu: bool,
|
||||||
|
key_filter: Callable[[str], bool] | None = None,
|
||||||
|
clone_tensors: bool = True,
|
||||||
|
show_progress: bool = True,
|
||||||
|
) -> Iterator[tuple[str, torch.Tensor]]:
|
||||||
|
logger.info(
|
||||||
|
"Loading safetensors with Run:ai Model Streamer to %s",
|
||||||
|
"cpu" if to_cpu else device,
|
||||||
|
)
|
||||||
|
with SafetensorsStreamer() as streamer:
|
||||||
|
if to_cpu:
|
||||||
|
streamer.stream_files(files)
|
||||||
|
else:
|
||||||
|
streamer.stream_files(files, device=device)
|
||||||
|
for name, tensor in streamer.get_tensors():
|
||||||
|
if key_filter is not None and not key_filter(name):
|
||||||
|
continue
|
||||||
|
if to_cpu or clone_tensors:
|
||||||
|
yield name, tensor.clone().detach()
|
||||||
|
else:
|
||||||
|
yield name, tensor
|
||||||
@@ -0,0 +1,47 @@
|
|||||||
|
# SPDX-License-Identifier: Apache-2.0
|
||||||
|
"""`safe_open`: slower to read, and the only source whose pages stay reclaimable.
|
||||||
|
|
||||||
|
`safe_open` maps the file, so a CPU tensor it yields is a view into the
|
||||||
|
checkpoint rather than a copy. Those pages are file-backed, which is what lets
|
||||||
|
the kernel drop them under memory pressure even on a host with no swap.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from typing import Callable, ClassVar, Iterator
|
||||||
|
|
||||||
|
import torch
|
||||||
|
from safetensors.torch import safe_open
|
||||||
|
from tqdm.auto import tqdm
|
||||||
|
|
||||||
|
_BAR_FORMAT = "{desc}: {percentage:.0f}%|{bar}| {n_fmt}/{total_fmt}"
|
||||||
|
|
||||||
|
|
||||||
|
class SafetensorsMmapReader:
|
||||||
|
name: ClassVar[str] = "safetensors"
|
||||||
|
supports_key_filter: ClassVar[bool] = True
|
||||||
|
retains_file_mapping: ClassVar[bool] = True
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def is_available(cls) -> bool:
|
||||||
|
return True
|
||||||
|
|
||||||
|
def iter_weights(
|
||||||
|
self,
|
||||||
|
files: list[str],
|
||||||
|
*,
|
||||||
|
device: str,
|
||||||
|
to_cpu: bool,
|
||||||
|
key_filter: Callable[[str], bool] | None = None,
|
||||||
|
clone_tensors: bool = True,
|
||||||
|
show_progress: bool = True,
|
||||||
|
) -> Iterator[tuple[str, torch.Tensor]]:
|
||||||
|
for path in tqdm(
|
||||||
|
files,
|
||||||
|
desc="Loading safetensors checkpoint shards",
|
||||||
|
disable=not show_progress,
|
||||||
|
bar_format=_BAR_FORMAT,
|
||||||
|
):
|
||||||
|
with safe_open(path, framework="pt", device=device) as handle:
|
||||||
|
for name in handle.keys(): # noqa: SIM118
|
||||||
|
if key_filter is not None and not key_filter(name):
|
||||||
|
continue
|
||||||
|
yield name, handle.get_tensor(name)
|
||||||
@@ -18,16 +18,16 @@ from safetensors.torch import safe_open
|
|||||||
from torch.distributed.tensor import DTensor
|
from torch.distributed.tensor import DTensor
|
||||||
from tqdm.auto import tqdm
|
from tqdm.auto import tqdm
|
||||||
|
|
||||||
try:
|
|
||||||
from runai_model_streamer import SafetensorsStreamer
|
|
||||||
|
|
||||||
HAS_RUNAI_MODEL_STREAMER = True
|
|
||||||
except ImportError:
|
|
||||||
HAS_RUNAI_MODEL_STREAMER = False
|
|
||||||
|
|
||||||
from sglang.multimodal_gen import envs
|
|
||||||
from sglang.multimodal_gen.runtime.distributed import get_local_torch_device
|
from sglang.multimodal_gen.runtime.distributed import get_local_torch_device
|
||||||
from sglang.multimodal_gen.runtime.loader.weight_load_plan import WeightLoadPlan
|
from sglang.multimodal_gen.runtime.loader.weight_load_plan import WeightLoadPlan
|
||||||
|
from sglang.multimodal_gen.runtime.loader.weight_readers import (
|
||||||
|
FALLBACK_READER,
|
||||||
|
RunaiStreamerReader,
|
||||||
|
select_weight_reader,
|
||||||
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.loader.weight_readers.runai_streamer import (
|
||||||
|
HAS_RUNAI_MODEL_STREAMER,
|
||||||
|
)
|
||||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||||
|
|
||||||
logger = init_logger(__name__)
|
logger = init_logger(__name__)
|
||||||
@@ -230,14 +230,19 @@ def safetensors_weights_iterator(
|
|||||||
device = str(checkpoint_device)
|
device = str(checkpoint_device)
|
||||||
else:
|
else:
|
||||||
device = "cpu" if to_cpu else str(get_local_torch_device())
|
device = "cpu" if to_cpu else str(get_local_torch_device())
|
||||||
if use_runai_model_streamer is None:
|
# The caller may still pass the old boolean; map it onto a backend name so
|
||||||
use_runai_model_streamer = (
|
# there is one place that decides, and it is the place that knows which
|
||||||
HAS_RUNAI_MODEL_STREAMER and envs.SGLANG_USE_RUNAI_MODEL_STREAMER
|
# backends can skip keys.
|
||||||
|
requested = None
|
||||||
|
if use_runai_model_streamer is not None:
|
||||||
|
requested = (
|
||||||
|
RunaiStreamerReader.name
|
||||||
|
if use_runai_model_streamer
|
||||||
|
else FALLBACK_READER.name
|
||||||
)
|
)
|
||||||
if key_filter is not None:
|
backend = select_weight_reader(
|
||||||
# streamer filters after materializing all tensors, so it cannot skip
|
requested=requested, needs_key_filter=key_filter is not None
|
||||||
# a checkpoint partition at load time
|
)
|
||||||
use_runai_model_streamer = False
|
|
||||||
|
|
||||||
# Validate files before loading
|
# Validate files before loading
|
||||||
corrupted_files, duplicate_files_by_key = _scan_safetensors_files(hf_weights_files)
|
corrupted_files, duplicate_files_by_key = _scan_safetensors_files(hf_weights_files)
|
||||||
@@ -273,38 +278,14 @@ def safetensors_weights_iterator(
|
|||||||
|
|
||||||
_raise_if_duplicate_safetensors_keys(duplicate_files_by_key)
|
_raise_if_duplicate_safetensors_keys(duplicate_files_by_key)
|
||||||
|
|
||||||
if use_runai_model_streamer:
|
yield from backend.iter_weights(
|
||||||
logger.info(
|
hf_weights_files,
|
||||||
"Loading safetensors with Run:ai Model Streamer to %s",
|
device=device,
|
||||||
"cpu" if to_cpu else device,
|
to_cpu=to_cpu,
|
||||||
)
|
key_filter=key_filter,
|
||||||
with SafetensorsStreamer() as streamer:
|
clone_tensors=clone_streamed_tensors,
|
||||||
if to_cpu:
|
show_progress=enable_tqdm,
|
||||||
streamer.stream_files(hf_weights_files)
|
)
|
||||||
else:
|
|
||||||
streamer.stream_files(hf_weights_files, device=device)
|
|
||||||
for name, tensor in streamer.get_tensors():
|
|
||||||
if key_filter is not None and not key_filter(name):
|
|
||||||
continue
|
|
||||||
if to_cpu:
|
|
||||||
yield name, tensor.clone().detach()
|
|
||||||
elif clone_streamed_tensors:
|
|
||||||
yield name, tensor.clone().detach()
|
|
||||||
else:
|
|
||||||
yield name, tensor
|
|
||||||
else:
|
|
||||||
for st_file in tqdm(
|
|
||||||
hf_weights_files,
|
|
||||||
desc="Loading safetensors checkpoint shards",
|
|
||||||
disable=not enable_tqdm,
|
|
||||||
bar_format=_BAR_FORMAT,
|
|
||||||
):
|
|
||||||
with safe_open(st_file, framework="pt", device=device) as f:
|
|
||||||
for name in f.keys(): # noqa: SIM118
|
|
||||||
if key_filter is not None and not key_filter(name):
|
|
||||||
continue
|
|
||||||
param = f.get_tensor(name)
|
|
||||||
yield name, param
|
|
||||||
|
|
||||||
|
|
||||||
def _load_pt_file(bin_file: str, device: str) -> dict:
|
def _load_pt_file(bin_file: str, device: str) -> dict:
|
||||||
|
|||||||
@@ -0,0 +1,68 @@
|
|||||||
|
"""Which backend reads the checkpoint, and why that choice is not a boolean."""
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from sglang.multimodal_gen.runtime.loader import weight_readers
|
||||||
|
from sglang.multimodal_gen.runtime.loader.weight_readers import (
|
||||||
|
FALLBACK_READER,
|
||||||
|
RunaiStreamerReader,
|
||||||
|
SafetensorsMmapReader,
|
||||||
|
available_reader_names,
|
||||||
|
select_weight_reader,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class TestCapabilities:
|
||||||
|
def test_the_streamer_cannot_skip_keys(self):
|
||||||
|
# it materializes every tensor before yielding any of them
|
||||||
|
assert not RunaiStreamerReader.supports_key_filter
|
||||||
|
|
||||||
|
def test_only_the_mapping_backend_leaves_pages_reclaimable(self):
|
||||||
|
assert SafetensorsMmapReader.retains_file_mapping
|
||||||
|
assert not RunaiStreamerReader.retains_file_mapping
|
||||||
|
|
||||||
|
def test_the_fallback_is_always_available(self):
|
||||||
|
assert FALLBACK_READER.is_available()
|
||||||
|
assert FALLBACK_READER.name in available_reader_names()
|
||||||
|
|
||||||
|
|
||||||
|
class TestSelection:
|
||||||
|
def test_an_explicit_request_is_honoured(self):
|
||||||
|
assert select_weight_reader(requested="safetensors").name == "safetensors"
|
||||||
|
|
||||||
|
def test_an_unknown_name_is_an_error_not_a_silent_fallback(self):
|
||||||
|
with pytest.raises(ValueError, match="unknown weight reader"):
|
||||||
|
select_weight_reader(requested="does_not_exist")
|
||||||
|
|
||||||
|
def test_a_key_filter_passes_over_a_backend_that_cannot_filter(self):
|
||||||
|
# reading the whole checkpoint to discard most of it is worse than
|
||||||
|
# reading the requested part more slowly
|
||||||
|
chosen = select_weight_reader(requested="runai_streamer", needs_key_filter=True)
|
||||||
|
assert chosen.name == FALLBACK_READER.name
|
||||||
|
|
||||||
|
def test_a_key_filter_leaves_a_capable_backend_alone(self):
|
||||||
|
chosen = select_weight_reader(requested="safetensors", needs_key_filter=True)
|
||||||
|
assert chosen.name == "safetensors"
|
||||||
|
|
||||||
|
def test_an_unavailable_backend_falls_back(self, monkeypatch):
|
||||||
|
monkeypatch.setattr(
|
||||||
|
RunaiStreamerReader, "is_available", classmethod(lambda cls: False)
|
||||||
|
)
|
||||||
|
assert (
|
||||||
|
select_weight_reader(requested="runai_streamer").name
|
||||||
|
== FALLBACK_READER.name
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_the_environment_decides_when_nothing_is_requested(self, monkeypatch):
|
||||||
|
monkeypatch.setattr(
|
||||||
|
weight_readers.envs, "SGLANG_USE_RUNAI_MODEL_STREAMER", False
|
||||||
|
)
|
||||||
|
assert select_weight_reader().name == FALLBACK_READER.name
|
||||||
|
|
||||||
|
def test_the_environment_can_ask_for_the_streamer(self, monkeypatch):
|
||||||
|
if not RunaiStreamerReader.is_available():
|
||||||
|
pytest.skip("run:ai model streamer is not installed")
|
||||||
|
monkeypatch.setattr(
|
||||||
|
weight_readers.envs, "SGLANG_USE_RUNAI_MODEL_STREAMER", True
|
||||||
|
)
|
||||||
|
assert select_weight_reader().name == "runai_streamer"
|
||||||
Reference in New Issue
Block a user