This commit is contained in:
jiangzeyinzi
2024-10-21 11:36:02 +08:00
parent 0bba2c319d
commit cffd54a02a
4 changed files with 31 additions and 18 deletions
@@ -90,7 +90,10 @@ class DiffusionInference():
from safetensors.torch import load_file as load_safetensors
sd = load_safetensors(local_path)
else:
sd = torch.load(local_path, map_location='cpu', weights_only=True)
if 'weights_only' in torch.load.__code__.co_varnames:
sd = torch.load(local_path, map_location='cpu', weights_only=True)
else:
sd = torch.load(local_path, map_location='cpu')
first_stage_model_path = os.path.join(
os.path.dirname(local_path), 'first_stage_model.pth')
cond_stage_model_path = os.path.join(
@@ -40,8 +40,10 @@ class LargenInference(DiffusionInference):
from safetensors.torch import load_file as load_safetensors
sd = load_safetensors(local_path)
else:
sd = torch.load(local_path, map_location='cpu', weights_only=True)
if 'weights_only' in torch.load.__code__.co_varnames:
sd = torch.load(local_path, map_location='cpu', weights_only=True)
else:
sd = torch.load(local_path, map_location='cpu')
if 'model' in sd:
sd = sd['model']
+4 -1
View File
@@ -143,7 +143,10 @@ class TunerInference():
state_dict = {}
is_bin_file = True
if os.path.isfile(bin_file):
state_dict = torch.load(bin_file, weights_only=True)
if 'weights_only' in torch.load.__code__.co_varnames:
state_dict = torch.load(bin_file, weights_only=True)
else:
state_dict = torch.load(bin_file)
elif os.path.isfile(safe_file):
is_bin_file = False
from safetensors.torch import \