Support fastsafetensors no-GDS loading and page-cache release (#31859)

This commit is contained in:
Nan Jiang
2026-07-31 23:12:32 +08:00
committed by GitHub
parent 5f9b0db18c
commit 89f4a80c1f
4 changed files with 131 additions and 7 deletions
+11
View File
@@ -399,6 +399,14 @@ class DefaultModelLoader(BaseModelLoader):
super().__init__(load_config)
extra_config = load_config.model_loader_extra_config
allowed_keys = {"enable_multithread_load", "num_threads"}
if load_config.load_format == LoadFormat.FASTSAFETENSORS:
allowed_keys.add("enable_gds")
if "enable_gds" in extra_config and not isinstance(
extra_config["enable_gds"], bool
):
raise ValueError(
"enable_gds in --model-loader-extra-config must be a boolean"
)
unexpected_keys = set(extra_config.keys()) - allowed_keys
if unexpected_keys:
@@ -609,8 +617,11 @@ class DefaultModelLoader(BaseModelLoader):
use_multithread = False
if self.load_config.load_format == LoadFormat.FASTSAFETENSORS:
enable_gds = extra_config.get("enable_gds", True)
weights_iterator = fastsafetensors_weights_iterator(
hf_weights_files,
enable_gds=enable_gds,
drop_cache_after_load=weight_loader_drop_cache_after_load,
)
elif use_multithread:
weights_iterator = buffered_multi_thread_safetensors_weights_iterator(
@@ -999,6 +999,8 @@ def safetensors_weights_iterator(
def fastsafetensors_weights_iterator(
hf_weights_files: List[str],
enable_gds: bool = True,
drop_cache_after_load: bool = False,
) -> Generator[Tuple[str, torch.Tensor], None, None]:
"""
Iterate over the weights in the model safetensor files
@@ -1036,7 +1038,7 @@ def fastsafetensors_weights_iterator(
disable=False,
bar_format=_BAR_FORMAT,
):
loader = SafeTensorsFileLoader(pg, device)
loader = SafeTensorsFileLoader(pg, device, nogds=not enable_gds)
rank_file_map = {i: [f] for i, f in enumerate(f_list)}
loader.add_filenames(rank_file_map)
try:
@@ -1050,6 +1052,9 @@ def fastsafetensors_weights_iterator(
pass
finally:
loader.close()
if drop_cache_after_load:
for loaded_file in rank_file_map.get(rank, []):
_drop_file_cache_after_load(loaded_file)
def multi_thread_safetensors_weights_iterator(