From 672ec3b35e5afffb25b68a4328e94f0c912f04fa Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Eren=20G=C3=B6lge?= Date: Sat, 8 Jul 2023 11:40:44 +0200 Subject: [PATCH] Fix #2749 (#2750) --- TTS/tts/layers/bark/hubert/tokenizer.py | 2 +- TTS/tts/layers/bark/inference_funcs.py | 4 +--- 2 files changed, 2 insertions(+), 4 deletions(-) diff --git a/TTS/tts/layers/bark/hubert/tokenizer.py b/TTS/tts/layers/bark/hubert/tokenizer.py index be9a50f8..eb6dd0a8 100644 --- a/TTS/tts/layers/bark/hubert/tokenizer.py +++ b/TTS/tts/layers/bark/hubert/tokenizer.py @@ -121,7 +121,7 @@ class HubertTokenizer(nn.Module): data_from_model.output_size, data_from_model.version, ) - model.load_state_dict(torch.load(path)) + model.load_state_dict(torch.load(path, map_location=map_location)) if map_location: model = model.to(map_location) return model diff --git a/TTS/tts/layers/bark/inference_funcs.py b/TTS/tts/layers/bark/inference_funcs.py index da962ab1..3a73e6d4 100644 --- a/TTS/tts/layers/bark/inference_funcs.py +++ b/TTS/tts/layers/bark/inference_funcs.py @@ -136,9 +136,7 @@ def generate_voice( hubert_model = CustomHubert(checkpoint_path=model.config.LOCAL_MODEL_PATHS["hubert"]).to(model.device) # Load the CustomTokenizer model - tokenizer = HubertTokenizer.load_from_checkpoint(model.config.LOCAL_MODEL_PATHS["hubert_tokenizer"]).to( - model.device - ) # Automatically uses + tokenizer = HubertTokenizer.load_from_checkpoint(model.config.LOCAL_MODEL_PATHS["hubert_tokenizer"], map_location=model.device) # semantic_tokens = model.text_to_semantic( # text, max_gen_duration_s=seconds, top_k=50, top_p=0.95, temp=0.7 # ) # not 100%