Refactor: separate module for utilities related to loading/saving checkpoints and safetensors
This commit is contained in:
@@ -80,9 +80,9 @@ def get_net_g(model_path: str, version: str, device: str, hps: HyperParameters):
|
||||
net_g.state_dict()
|
||||
_ = net_g.eval()
|
||||
if model_path.endswith(".pth") or model_path.endswith(".pt"):
|
||||
_ = utils.load_checkpoint(model_path, net_g, None, skip_optimizer=True)
|
||||
_ = utils.checkpoints.load_checkpoint(model_path, net_g, None, skip_optimizer=True)
|
||||
elif model_path.endswith(".safetensors"):
|
||||
_ = utils.load_safetensors(model_path, net_g, True)
|
||||
_ = utils.safetensors.load_safetensors(model_path, net_g, True)
|
||||
else:
|
||||
raise ValueError(f"Unknown model format: {model_path}")
|
||||
return net_g
|
||||
|
||||
Reference in New Issue
Block a user