Support fastsafetensors no-GDS loading and page-cache release (#31859)
This commit is contained in:
@@ -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(
|
||||
|
||||
Reference in New Issue
Block a user