运行LJSpeech Hugging Face示例报错:张量设备不匹配求助
问题:FastSpeech2-LJSpeech运行时设备不匹配错误
环境配置:CUDA 11.7、PyTorch 1.13.1、Fairseq 0.12.2
运行facebook/fastspeech2-en-ljspeech示例时,即使将模型移至唯一GPU,仍出现设备不匹配错误:
RuntimeError: Expected all tensors to be on the same device, but found at least two devices, cuda:0 and cpu! (when checking argument for argument index in method wrapper__index_select)
使用的代码:
from fairseq.checkpoint_utils import load_model_ensemble_and_task_from_hf_hub from fairseq.models.text_to_speech.hub_interface import TTSHubInterface import IPython.display as ipd import torch models, cfg, task = load_model_ensemble_and_task_from_hf_hub( "facebook/fastspeech2-en-ljspeech", arg_overrides={"vocoder": "hifigan", "fp16": False} ) model = models[0].to(torch.device('cuda')) models[0] = model TTSHubInterface.update_cfg_with_data_cfg(cfg, task.data_cfg) generator = task.build_generator(models, cfg) text = "Hello, this is a test run." sample = TTSHubInterface.get_model_input(task, text) wav, rate = TTSHubInterface.get_prediction(task, model, generator, sample) ipd.Audio(wav, rate=rate)
解决方案
问题核心是输入样本sample、任务字典相关张量仍留在CPU,与GPU上的模型张量设备不匹配。以下是修正步骤和代码:
关键修改点
- 将生成的输入样本
sample移至GPU - 将任务的源字典(词表)移至GPU(Fairseq TTS任务中,字典嵌入张量会参与计算)
- 统一设置默认设备,减少手动迁移的遗漏
修正后的代码
from fairseq.checkpoint_utils import load_model_ensemble_and_task_from_hf_hub from fairseq.models.text_to_speech.hub_interface import TTSHubInterface import IPython.display as ipd import torch # 统一设备配置 device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') models, cfg, task = load_model_ensemble_and_task_from_hf_hub( "facebook/fastspeech2-en-ljspeech", arg_overrides={"vocoder": "hifigan", "fp16": False} ) # 迁移模型到目标设备 model = models[0].to(device) models[0] = model # 迁移任务字典到目标设备(关键步骤) task.src_dict = task.src_dict.to(device) TTSHubInterface.update_cfg_with_data_cfg(cfg, task.data_cfg) generator = task.build_generator(models, cfg) text = "Hello, this is a test run." sample = TTSHubInterface.get_model_input(task, text) # 迁移输入样本到目标设备 sample = sample.to(device) wav, rate = TTSHubInterface.get_prediction(task, model, generator, sample) ipd.Audio(wav, rate=rate)
额外排查建议
- 若仍报错,检查HiFiGAN vocoder是否被正确迁移到GPU:部分场景下vocoder会单独加载,需确保其设备与模型一致
- 确认Fairseq版本兼容性:0.12.2与PyTorch 1.13.1基本兼容,若问题持续可尝试升级Fairseq至较新版本
内容的提问来源于stack exchange,提问作者carlos
相关产品推荐
相关产品推荐

