Merge pull request #11260 from dhwz/dev

fix very slow loading speed of .safetensors files
This commit is contained in:
AUTOMATIC1111
2023-06-27 09:11:08 +03:00
committed by GitHub
2 changed files with 6 additions and 2 deletions

View File

@@ -246,8 +246,11 @@ def read_metadata_from_safetensors(filename):
def read_state_dict(checkpoint_file, print_global_state=False, map_location=None):
_, extension = os.path.splitext(checkpoint_file)
if extension.lower() == ".safetensors":
device = map_location or shared.weight_load_location or devices.get_optimal_device_name()
pl_sd = safetensors.torch.load_file(checkpoint_file, device=device)
if not shared.opts.disable_mmap_load_safetensors:
device = map_location or shared.weight_load_location or devices.get_optimal_device_name()
pl_sd = safetensors.torch.load_file(checkpoint_file, device=device)
else:
pl_sd = safetensors.torch.load(open(checkpoint_file, 'rb').read())
else:
pl_sd = torch.load(checkpoint_file, map_location=map_location or shared.weight_load_location)