diff --git a/TTS/speaker_encoder/utils/prepare_voxceleb.py b/TTS/speaker_encoder/utils/prepare_voxceleb.py index 758e1cb3b..2c7208b53 100644 --- a/TTS/speaker_encoder/utils/prepare_voxceleb.py +++ b/TTS/speaker_encoder/utils/prepare_voxceleb.py @@ -100,7 +100,7 @@ def download_and_extract(directory, subset, urls): extract_path = zip_filepath.strip(".zip") # check zip file md5sum - md5 = hashlib.md5(open(zip_filepath, 'rb').read()).hexdigest() + md5 = hashlib.md5(open(zip_filepath, 'rb').read(), usedforsecurity=False).hexdigest() if md5 != MD5SUM[subset]: raise ValueError("md5sum of %s mismatch" % zip_filepath) diff --git a/TTS/tts/tf/utils/generic_utils.py b/TTS/tts/tf/utils/generic_utils.py index 7eba946b1..105e00816 100644 --- a/TTS/tts/tf/utils/generic_utils.py +++ b/TTS/tts/tf/utils/generic_utils.py @@ -18,8 +18,17 @@ def save_checkpoint(model, optimizer, current_step, epoch, r, output_path, **kwa pickle.dump(state, open(output_path, 'wb')) +class SafeUnpickler(pickle.Unpickler): + def find_class(self, module, name): + if module == "builtins" and name in {"dict", "list", "tuple", "set", "int", "float", "str", "bytes", "bool"}: + return super().find_class(module, name) + if "numpy" in module: + return super().find_class(module, name) + raise pickle.UnpicklingError(f"Unsafe global {module}.{name}") + + def load_checkpoint(model, checkpoint_path): - checkpoint = pickle.load(open(checkpoint_path, 'rb')) + checkpoint = SafeUnpickler(open(checkpoint_path, 'rb')).load() chkp_var_dict = {var.name: var.numpy() for var in checkpoint['model']} tf_vars = model.weights for tf_var in tf_vars: diff --git a/TTS/tts/tf/utils/io.py b/TTS/tts/tf/utils/io.py index 143422d27..a6879b3b6 100644 --- a/TTS/tts/tf/utils/io.py +++ b/TTS/tts/tf/utils/io.py @@ -16,8 +16,17 @@ def save_checkpoint(model, optimizer, current_step, epoch, r, output_path, **kwa pickle.dump(state, open(output_path, 'wb')) +class SafeUnpickler(pickle.Unpickler): + def find_class(self, module, name): + if module == "builtins" and name in {"dict", "list", "tuple", "set", "int", "float", "str", "bytes", "bool"}: + return super().find_class(module, name) + if "numpy" in module: + return super().find_class(module, name) + raise pickle.UnpicklingError(f"Unsafe global {module}.{name}") + + def load_checkpoint(model, checkpoint_path): - checkpoint = pickle.load(open(checkpoint_path, 'rb')) + checkpoint = SafeUnpickler(open(checkpoint_path, 'rb')).load() chkp_var_dict = {var.name: var.numpy() for var in checkpoint['model']} tf_vars = model.weights for tf_var in tf_vars: diff --git a/TTS/vocoder/tf/utils/io.py b/TTS/vocoder/tf/utils/io.py index c73c9cd86..81423c5ae 100644 --- a/TTS/vocoder/tf/utils/io.py +++ b/TTS/vocoder/tf/utils/io.py @@ -15,9 +15,18 @@ def save_checkpoint(model, current_step, epoch, output_path, **kwargs): pickle.dump(state, open(output_path, 'wb')) +class SafeUnpickler(pickle.Unpickler): + def find_class(self, module, name): + if module == "builtins" and name in {"dict", "list", "tuple", "set", "int", "float", "str", "bytes", "bool"}: + return super().find_class(module, name) + if "numpy" in module: + return super().find_class(module, name) + raise pickle.UnpicklingError(f"Unsafe global {module}.{name}") + + def load_checkpoint(model, checkpoint_path): """ Load TF Vocoder model """ - checkpoint = pickle.load(open(checkpoint_path, 'rb')) + checkpoint = SafeUnpickler(open(checkpoint_path, 'rb')).load() chkp_var_dict = {var.name: var.numpy() for var in checkpoint['model']} tf_vars = model.weights for tf_var in tf_vars: