feat: add coordinated checkpoint prefetch for network filesystem loading (#20843)
This commit is contained in:
@@ -510,6 +510,8 @@ Please consult the documentation below and [server_args.py](https://github.com/s
|
|||||||
| --- | --- | --- | --- |
|
| --- | --- | --- | --- |
|
||||||
| `--custom-weight-loader` | The custom dataloader which used to update the model. Should be set with a valid import path, such as my_package.weight_load_func | `None` | List[str] |
|
| `--custom-weight-loader` | The custom dataloader which used to update the model. Should be set with a valid import path, such as my_package.weight_load_func | `None` | List[str] |
|
||||||
| `--weight-loader-disable-mmap` | Disable mmap while loading weight using safetensors. | `False` | bool flag (set to enable) |
|
| `--weight-loader-disable-mmap` | Disable mmap while loading weight using safetensors. | `False` | bool flag (set to enable) |
|
||||||
|
| `--weight-loader-prefetch-checkpoints` | Prefetch checkpoint files into OS page cache before loading. Each rank prefetches a fraction of the shards in a background thread, reducing total network I/O on shared filesystems (NFS/Lustre) from N\*checkpoint to 1\*checkpoint. Recommended for models on network storage. | `False` | bool flag (set to enable) |
|
||||||
|
| `--weight-loader-prefetch-num-threads` | Number of threads per rank for checkpoint prefetching. | `4` | Type: int |
|
||||||
| `--remote-instance-weight-loader-seed-instance-ip` | The ip of the seed instance for loading weights from remote instance. | `None` | Type: str |
|
| `--remote-instance-weight-loader-seed-instance-ip` | The ip of the seed instance for loading weights from remote instance. | `None` | Type: str |
|
||||||
| `--remote-instance-weight-loader-seed-instance-service-port` | The service port of the seed instance for loading weights from remote instance. | `None` | Type: int |
|
| `--remote-instance-weight-loader-seed-instance-service-port` | The service port of the seed instance for loading weights from remote instance. | `None` | Type: int |
|
||||||
| `--remote-instance-weight-loader-send-weights-group-ports` | The communication group ports for loading weights from remote instance. | `None` | Type: JSON list |
|
| `--remote-instance-weight-loader-send-weights-group-ports` | The communication group ports for loading weights from remote instance. | `None` | Type: JSON list |
|
||||||
|
|||||||
@@ -20,6 +20,7 @@ SGLang supports various environment variables that can be used to configure its
|
|||||||
| `SGLANG_REQ_WAITING_TIMEOUT` | Timeout (in seconds) for requests waiting in the queue before being scheduled | `-1` |
|
| `SGLANG_REQ_WAITING_TIMEOUT` | Timeout (in seconds) for requests waiting in the queue before being scheduled | `-1` |
|
||||||
| `SGLANG_REQ_RUNNING_TIMEOUT` | Timeout (in seconds) for requests running in the decode batch | `-1` |
|
| `SGLANG_REQ_RUNNING_TIMEOUT` | Timeout (in seconds) for requests running in the decode batch | `-1` |
|
||||||
| `SGLANG_CACHE_DIR` | Cache directory for model weights and other data | `~/.cache/sglang` |
|
| `SGLANG_CACHE_DIR` | Cache directory for model weights and other data | `~/.cache/sglang` |
|
||||||
|
| `SGLANG_PREFETCH_BLOCK_SIZE_MB` | Block size (in MB) for sequential checkpoint prefetch reads that warm the OS page cache before workers load weights via mmap | `16` |
|
||||||
|
|
||||||
## Performance Tuning
|
## Performance Tuning
|
||||||
|
|
||||||
|
|||||||
@@ -163,6 +163,7 @@ class Envs:
|
|||||||
SGLANG_USE_MODELSCOPE = EnvBool(False)
|
SGLANG_USE_MODELSCOPE = EnvBool(False)
|
||||||
SGLANG_SORT_WEIGHT_FILES = EnvBool(False)
|
SGLANG_SORT_WEIGHT_FILES = EnvBool(False)
|
||||||
SGLANG_DISABLED_MODEL_ARCHS = EnvTuple(tuple())
|
SGLANG_DISABLED_MODEL_ARCHS = EnvTuple(tuple())
|
||||||
|
SGLANG_PREFETCH_BLOCK_SIZE_MB = EnvInt(16)
|
||||||
|
|
||||||
# Logging Options
|
# Logging Options
|
||||||
SGLANG_LOG_GC = EnvBool(False)
|
SGLANG_LOG_GC = EnvBool(False)
|
||||||
|
|||||||
@@ -506,9 +506,10 @@ class DefaultModelLoader(BaseModelLoader):
|
|||||||
hf_weights_files,
|
hf_weights_files,
|
||||||
)
|
)
|
||||||
elif use_safetensors:
|
elif use_safetensors:
|
||||||
weight_loader_disable_mmap = (
|
server_args = get_global_server_args()
|
||||||
get_global_server_args().weight_loader_disable_mmap
|
weight_loader_disable_mmap = server_args.weight_loader_disable_mmap
|
||||||
)
|
weight_loader_prefetch = server_args.weight_loader_prefetch_checkpoints
|
||||||
|
prefetch_num_threads = server_args.weight_loader_prefetch_num_threads
|
||||||
|
|
||||||
if self.load_config.load_format == LoadFormat.FASTSAFETENSORS:
|
if self.load_config.load_format == LoadFormat.FASTSAFETENSORS:
|
||||||
weights_iterator = fastsafetensors_weights_iterator(
|
weights_iterator = fastsafetensors_weights_iterator(
|
||||||
@@ -521,10 +522,15 @@ class DefaultModelLoader(BaseModelLoader):
|
|||||||
"num_threads", self.DEFAULT_NUM_THREADS
|
"num_threads", self.DEFAULT_NUM_THREADS
|
||||||
),
|
),
|
||||||
disable_mmap=weight_loader_disable_mmap,
|
disable_mmap=weight_loader_disable_mmap,
|
||||||
|
prefetch=weight_loader_prefetch,
|
||||||
|
prefetch_num_threads=prefetch_num_threads,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
weights_iterator = safetensors_weights_iterator(
|
weights_iterator = safetensors_weights_iterator(
|
||||||
hf_weights_files, disable_mmap=weight_loader_disable_mmap
|
hf_weights_files,
|
||||||
|
disable_mmap=weight_loader_disable_mmap,
|
||||||
|
prefetch=weight_loader_prefetch,
|
||||||
|
prefetch_num_threads=prefetch_num_threads,
|
||||||
)
|
)
|
||||||
|
|
||||||
else:
|
else:
|
||||||
|
|||||||
@@ -39,6 +39,7 @@ from sglang.srt.configs.model_config import ModelConfig
|
|||||||
from sglang.srt.distributed import (
|
from sglang.srt.distributed import (
|
||||||
get_tensor_model_parallel_rank,
|
get_tensor_model_parallel_rank,
|
||||||
get_tensor_model_parallel_world_size,
|
get_tensor_model_parallel_world_size,
|
||||||
|
get_world_group,
|
||||||
)
|
)
|
||||||
from sglang.srt.layers.dp_attention import get_attention_tp_rank
|
from sglang.srt.layers.dp_attention import get_attention_tp_rank
|
||||||
from sglang.srt.layers.quantization import QuantizationConfig, get_quantization_config
|
from sglang.srt.layers.quantization import QuantizationConfig, get_quantization_config
|
||||||
@@ -68,6 +69,18 @@ except ImportError as e:
|
|||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
# Block size for sequential checkpoint prefetch reads (page cache warming).
|
||||||
|
_PREFETCH_BLOCK_SIZE = None
|
||||||
|
|
||||||
|
|
||||||
|
def _get_prefetch_block_size() -> int:
|
||||||
|
global _PREFETCH_BLOCK_SIZE
|
||||||
|
if _PREFETCH_BLOCK_SIZE is None:
|
||||||
|
from sglang.srt.environ import envs
|
||||||
|
|
||||||
|
_PREFETCH_BLOCK_SIZE = envs.SGLANG_PREFETCH_BLOCK_SIZE_MB.get() * 1024 * 1024
|
||||||
|
return _PREFETCH_BLOCK_SIZE
|
||||||
|
|
||||||
|
|
||||||
# use system-level temp directory for file locks, so that multiple users
|
# use system-level temp directory for file locks, so that multiple users
|
||||||
# can share the same lock without error.
|
# can share the same lock without error.
|
||||||
@@ -700,16 +713,127 @@ def np_cache_weights_iterator(
|
|||||||
yield name, torch.from_numpy(param)
|
yield name, torch.from_numpy(param)
|
||||||
|
|
||||||
|
|
||||||
|
def _prefetch_checkpoint_file(file_path: str) -> None:
|
||||||
|
"""Prefetch a checkpoint file into the OS page cache.
|
||||||
|
|
||||||
|
Reads the file sequentially in 16 MB blocks so the kernel caches its pages
|
||||||
|
before workers load the same file via mmap.
|
||||||
|
"""
|
||||||
|
with open(file_path, "rb") as f:
|
||||||
|
while f.read(_get_prefetch_block_size()):
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
def _prefetch_all_checkpoints(
|
||||||
|
sorted_files: List[str],
|
||||||
|
num_threads: int = 4,
|
||||||
|
) -> None:
|
||||||
|
"""Start prefetching checkpoint files into page cache in a background thread.
|
||||||
|
|
||||||
|
When multiple ranks on the same node load the same checkpoint (e.g.
|
||||||
|
DP-attention), each rank independently mmaps the same files, causing
|
||||||
|
redundant NFS/Lustre reads. By distributing the prefetch across ranks
|
||||||
|
(each rank reads 1/Nth of the shards), the total network I/O is reduced
|
||||||
|
from N * checkpoint_size to 1 * checkpoint_size, with subsequent
|
||||||
|
mmap accesses hitting the shared OS page cache.
|
||||||
|
|
||||||
|
The prefetch runs in a background thread so that loading can start
|
||||||
|
immediately and benefit from pages that have already been cached,
|
||||||
|
rather than blocking until all files are prefetched. This pipelining
|
||||||
|
naturally adapts to any RAM size — even if the full checkpoint does
|
||||||
|
not fit in page cache, the prefetch thread stays ahead of the loader.
|
||||||
|
"""
|
||||||
|
import asyncio
|
||||||
|
import threading
|
||||||
|
import time
|
||||||
|
|
||||||
|
# Use node-local rank so that each node independently prefetches the
|
||||||
|
# full checkpoint into its own page cache. Global rank would split files
|
||||||
|
# across nodes, but page cache is not shared across nodes.
|
||||||
|
if torch.distributed.is_initialized():
|
||||||
|
world_group = get_world_group()
|
||||||
|
local_rank = world_group.local_rank
|
||||||
|
local_world_size = world_group.local_size or world_group.world_size
|
||||||
|
else:
|
||||||
|
local_rank = 0
|
||||||
|
local_world_size = 1
|
||||||
|
|
||||||
|
my_files = sorted_files[local_rank::local_world_size]
|
||||||
|
total_for_rank = len(my_files)
|
||||||
|
|
||||||
|
logger.info(
|
||||||
|
"Rank %d: prefetching %d/%d checkpoint shards into page cache "
|
||||||
|
"(background, %d local ranks sharing the work, %d threads per rank)...",
|
||||||
|
local_rank,
|
||||||
|
total_for_rank,
|
||||||
|
len(sorted_files),
|
||||||
|
local_world_size,
|
||||||
|
num_threads,
|
||||||
|
)
|
||||||
|
|
||||||
|
async def _prefetch_all() -> None:
|
||||||
|
semaphore = asyncio.Semaphore(num_threads)
|
||||||
|
completed = 0
|
||||||
|
next_log_pct = 10
|
||||||
|
|
||||||
|
async def prefetch_one(path: str) -> None:
|
||||||
|
nonlocal completed, next_log_pct
|
||||||
|
try:
|
||||||
|
async with semaphore:
|
||||||
|
await asyncio.to_thread(_prefetch_checkpoint_file, path)
|
||||||
|
completed += 1
|
||||||
|
if total_for_rank > 0 and next_log_pct <= 100:
|
||||||
|
pct = 100 * completed / total_for_rank
|
||||||
|
if pct >= next_log_pct:
|
||||||
|
logger.info(
|
||||||
|
"Rank %d: prefetching checkpoint files: %d%% (%d/%d)",
|
||||||
|
local_rank,
|
||||||
|
next_log_pct,
|
||||||
|
completed,
|
||||||
|
total_for_rank,
|
||||||
|
)
|
||||||
|
next_log_pct += 10
|
||||||
|
except Exception:
|
||||||
|
logger.warning(
|
||||||
|
"Failed to prefetch checkpoint file %r.",
|
||||||
|
path,
|
||||||
|
exc_info=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
await asyncio.gather(*(prefetch_one(p) for p in my_files))
|
||||||
|
|
||||||
|
def _run_prefetch() -> None:
|
||||||
|
start = time.perf_counter()
|
||||||
|
asyncio.run(_prefetch_all())
|
||||||
|
elapsed = time.perf_counter() - start
|
||||||
|
logger.info(
|
||||||
|
"Rank %d: prefetching checkpoint files into page cache "
|
||||||
|
"finished in %.2fs",
|
||||||
|
local_rank,
|
||||||
|
elapsed,
|
||||||
|
)
|
||||||
|
|
||||||
|
threading.Thread(target=_run_prefetch, daemon=True).start()
|
||||||
|
|
||||||
|
|
||||||
def safetensors_weights_iterator(
|
def safetensors_weights_iterator(
|
||||||
hf_weights_files: List[str],
|
hf_weights_files: List[str],
|
||||||
disable_mmap: bool = False,
|
disable_mmap: bool = False,
|
||||||
|
prefetch: bool = False,
|
||||||
|
prefetch_num_threads: int = 4,
|
||||||
) -> Generator[Tuple[str, torch.Tensor], None, None]:
|
) -> Generator[Tuple[str, torch.Tensor], None, None]:
|
||||||
"""Iterate over the weights in the model safetensor files."""
|
"""Iterate over the weights in the model safetensor files."""
|
||||||
enable_tqdm = (
|
enable_tqdm = (
|
||||||
not torch.distributed.is_initialized() or torch.distributed.get_rank() == 0
|
not torch.distributed.is_initialized() or torch.distributed.get_rank() == 0
|
||||||
)
|
)
|
||||||
|
|
||||||
|
sorted_files = sorted(hf_weights_files)
|
||||||
|
|
||||||
|
if prefetch and not disable_mmap:
|
||||||
|
_prefetch_all_checkpoints(sorted_files, num_threads=prefetch_num_threads)
|
||||||
|
|
||||||
for st_file in tqdm(
|
for st_file in tqdm(
|
||||||
hf_weights_files,
|
sorted_files,
|
||||||
desc="Loading safetensors checkpoint shards",
|
desc="Loading safetensors checkpoint shards",
|
||||||
disable=not enable_tqdm,
|
disable=not enable_tqdm,
|
||||||
bar_format=BAR_FORMAT,
|
bar_format=BAR_FORMAT,
|
||||||
@@ -821,6 +945,8 @@ def buffered_multi_thread_safetensors_weights_iterator(
|
|||||||
hf_weights_files: List[str],
|
hf_weights_files: List[str],
|
||||||
max_workers: int,
|
max_workers: int,
|
||||||
disable_mmap: bool = False,
|
disable_mmap: bool = False,
|
||||||
|
prefetch: bool = False,
|
||||||
|
prefetch_num_threads: int = 4,
|
||||||
) -> Generator[Tuple[str, torch.Tensor], None, None]:
|
) -> Generator[Tuple[str, torch.Tensor], None, None]:
|
||||||
"""Multi-threaded safetensor loader with bounded memory via a sliding window.
|
"""Multi-threaded safetensor loader with bounded memory via a sliding window.
|
||||||
|
|
||||||
@@ -828,6 +954,9 @@ def buffered_multi_thread_safetensors_weights_iterator(
|
|||||||
max_workers loading concurrently + 1 prefetched and ready to yield.
|
max_workers loading concurrently + 1 prefetched and ready to yield.
|
||||||
Peak CPU RAM ≈ (max_workers + 2) × shard_file_size.
|
Peak CPU RAM ≈ (max_workers + 2) × shard_file_size.
|
||||||
"""
|
"""
|
||||||
|
sorted_files = sorted(hf_weights_files)
|
||||||
|
if prefetch and not disable_mmap:
|
||||||
|
_prefetch_all_checkpoints(sorted_files, num_threads=prefetch_num_threads)
|
||||||
enable_tqdm = (
|
enable_tqdm = (
|
||||||
not torch.distributed.is_initialized() or torch.distributed.get_rank() == 0
|
not torch.distributed.is_initialized() or torch.distributed.get_rank() == 0
|
||||||
)
|
)
|
||||||
@@ -845,7 +974,7 @@ def buffered_multi_thread_safetensors_weights_iterator(
|
|||||||
buffer_size = max_workers + 1
|
buffer_size = max_workers + 1
|
||||||
|
|
||||||
with concurrent.futures.ThreadPoolExecutor(max_workers=max_workers) as executor:
|
with concurrent.futures.ThreadPoolExecutor(max_workers=max_workers) as executor:
|
||||||
file_iter = iter(hf_weights_files)
|
file_iter = iter(sorted_files)
|
||||||
pending: collections.deque = collections.deque()
|
pending: collections.deque = collections.deque()
|
||||||
|
|
||||||
# Seed the buffer.
|
# Seed the buffer.
|
||||||
|
|||||||
@@ -722,6 +722,8 @@ class ServerArgs:
|
|||||||
# For model weight update and weight loading
|
# For model weight update and weight loading
|
||||||
custom_weight_loader: Optional[List[str]] = None
|
custom_weight_loader: Optional[List[str]] = None
|
||||||
weight_loader_disable_mmap: bool = False
|
weight_loader_disable_mmap: bool = False
|
||||||
|
weight_loader_prefetch_checkpoints: bool = False
|
||||||
|
weight_loader_prefetch_num_threads: int = 4
|
||||||
remote_instance_weight_loader_seed_instance_ip: Optional[str] = None
|
remote_instance_weight_loader_seed_instance_ip: Optional[str] = None
|
||||||
remote_instance_weight_loader_seed_instance_service_port: Optional[int] = None
|
remote_instance_weight_loader_seed_instance_service_port: Optional[int] = None
|
||||||
remote_instance_weight_loader_send_weights_group_ports: Optional[List[int]] = None
|
remote_instance_weight_loader_send_weights_group_ports: Optional[List[int]] = None
|
||||||
@@ -6234,6 +6236,20 @@ class ServerArgs:
|
|||||||
action="store_true",
|
action="store_true",
|
||||||
help="Disable mmap while loading weight using safetensors.",
|
help="Disable mmap while loading weight using safetensors.",
|
||||||
)
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--weight-loader-prefetch-checkpoints",
|
||||||
|
action="store_true",
|
||||||
|
help="Prefetch checkpoint files into OS page cache before loading. "
|
||||||
|
"Each rank prefetches a fraction of the shards, reducing total "
|
||||||
|
"network I/O on shared filesystems (NFS/Lustre) from N*checkpoint "
|
||||||
|
"to 1*checkpoint. Recommended for models on network storage.",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--weight-loader-prefetch-num-threads",
|
||||||
|
type=int,
|
||||||
|
default=ServerArgs.weight_loader_prefetch_num_threads,
|
||||||
|
help="Number of threads per rank for checkpoint prefetching (default: 4).",
|
||||||
|
)
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--remote-instance-weight-loader-seed-instance-ip",
|
"--remote-instance-weight-loader-seed-instance-ip",
|
||||||
type=str,
|
type=str,
|
||||||
|
|||||||
@@ -0,0 +1,49 @@
|
|||||||
|
import unittest
|
||||||
|
|
||||||
|
import sglang as sgl
|
||||||
|
from sglang.test.ci.ci_register import register_cuda_ci
|
||||||
|
from sglang.test.test_utils import CustomTestCase
|
||||||
|
|
||||||
|
register_cuda_ci(est_time=300, suite="nightly-4-gpu")
|
||||||
|
|
||||||
|
PROMPTS = [
|
||||||
|
"Hello, my name is",
|
||||||
|
"The president of the United States is",
|
||||||
|
"The capital of France is",
|
||||||
|
"The future of AI is",
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
class TestPrefetchCheckpointsMultiGPU(CustomTestCase):
|
||||||
|
"""Verify that --weight-loader-prefetch-checkpoints works with DP attention."""
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def setUpClass(cls):
|
||||||
|
cls.engine = sgl.Engine(
|
||||||
|
model_path="Qwen/Qwen1.5-MoE-A2.7B-Chat",
|
||||||
|
tp_size=4,
|
||||||
|
dp_size=4,
|
||||||
|
enable_dp_attention=True,
|
||||||
|
disable_radix_cache=True,
|
||||||
|
weight_loader_prefetch_checkpoints=True,
|
||||||
|
cuda_graph_max_bs=1,
|
||||||
|
max_total_tokens=256,
|
||||||
|
)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def tearDownClass(cls):
|
||||||
|
if hasattr(cls, "engine") and cls.engine:
|
||||||
|
cls.engine.shutdown()
|
||||||
|
|
||||||
|
def test_generate_with_prefetch(self):
|
||||||
|
"""Server launched with prefetch must produce valid output."""
|
||||||
|
outputs = self.engine.generate(PROMPTS)
|
||||||
|
self.assertEqual(len(outputs), len(PROMPTS))
|
||||||
|
for i, output in enumerate(outputs):
|
||||||
|
text = output["text"]
|
||||||
|
self.assertIsInstance(text, str)
|
||||||
|
self.assertGreater(len(text), 0, f"Prompt {i} produced empty output")
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
@@ -0,0 +1,55 @@
|
|||||||
|
"""
|
||||||
|
Unit tests for coordinated checkpoint prefetch.
|
||||||
|
|
||||||
|
Verifies that weights loaded with prefetch enabled are bit-identical
|
||||||
|
to weights loaded without prefetch.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import os
|
||||||
|
import tempfile
|
||||||
|
import unittest
|
||||||
|
from unittest.mock import patch
|
||||||
|
|
||||||
|
import safetensors.torch
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from sglang.srt.model_loader.weight_utils import (
|
||||||
|
safetensors_weights_iterator,
|
||||||
|
)
|
||||||
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
|
|
||||||
|
register_cpu_ci(est_time=5, suite="stage-a-test-cpu")
|
||||||
|
|
||||||
|
|
||||||
|
class TestPrefetchWeightsIdentical(unittest.TestCase):
|
||||||
|
"""Verify that loading with prefetch yields identical weights to without."""
|
||||||
|
|
||||||
|
def _create_safetensors_files(self, tmpdir, num_shards=3):
|
||||||
|
"""Create real safetensors files with known tensor content."""
|
||||||
|
paths = []
|
||||||
|
for i in range(num_shards):
|
||||||
|
tensors = {
|
||||||
|
f"layer{i}.weight": torch.randn(32, 32),
|
||||||
|
f"layer{i}.bias": torch.randn(32),
|
||||||
|
}
|
||||||
|
path = os.path.join(tmpdir, f"model-{i:05d}.safetensors")
|
||||||
|
safetensors.torch.save_file(tensors, path)
|
||||||
|
paths.append(path)
|
||||||
|
return paths
|
||||||
|
|
||||||
|
@patch("torch.distributed.is_initialized", return_value=False)
|
||||||
|
def test_weights_match_with_and_without_prefetch(self, _):
|
||||||
|
"""Tensors yielded must be bit-identical regardless of prefetch flag."""
|
||||||
|
with tempfile.TemporaryDirectory() as tmpdir:
|
||||||
|
paths = self._create_safetensors_files(tmpdir)
|
||||||
|
|
||||||
|
without = dict(safetensors_weights_iterator(paths, prefetch=False))
|
||||||
|
with_pf = dict(safetensors_weights_iterator(paths, prefetch=True))
|
||||||
|
|
||||||
|
self.assertEqual(set(without.keys()), set(with_pf.keys()))
|
||||||
|
for name in without:
|
||||||
|
torch.testing.assert_close(without[name], with_pf[name])
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
Reference in New Issue
Block a user