GPU上使用torchaudio.transforms.MelSpectrogram报错的解决咨询
问题
在GPU上使用torchaudio计算MelSpectrogram时,CPU运行正常,但切换到CUDA时出现RuntimeError,提示stft input and window must be on the same device but got self on cuda:0 and window on cpu。当前torch版本为2.2.2+cu121,测试代码及错误栈如下:
测试代码
from typing import Optional import torch import torchaudio import numpy as np from tests.__init__ import ( __target_clock__ as TARGET_CLOCK, __number_of_test_data_vals__ as NUMBER_OF_TEST_DATA_VALS, ) # Set general parameters: TARGET_DEVICE = "CUDA" TARGET_FREQUENCY: int = 440 NUMBER_OF_FFT_SLOTS: int = 1024 HOP_LENGTH: Optional[int] = None NUMBER_OF_MEL_SLOTS: int = 128 if __name__ == "__main__": target_device = torch.device( "cuda" if (TARGET_DEVICE == "CUDA" and torch.cuda.is_available()) else "cpu" ) print(f"Using device {target_device}") sampling_vec: np.ndarray = np.arange(NUMBER_OF_TEST_DATA_VALS) / TARGET_CLOCK frequency_vec: np.ndarray = np.sin( 2 * np.pi * TARGET_FREQUENCY * sampling_vec ).astype("float32") frequency_tensor: torch.Tensor = torch.Tensor(frequency_vec).to(target_device) mel_spectrogram: torch.Tensor = torchaudio.transforms.MelSpectrogram( sample_rate=TARGET_CLOCK, n_fft=NUMBER_OF_FFT_SLOTS, hop_length=HOP_LENGTH, n_mels=NUMBER_OF_MEL_SLOTS, )(frequency_tensor) print(f"Obtained MEL-Spectrogram: {mel_spectrogram}")
错误信息
Traceback (most recent call last): File "<frozen runpy>", line 198, in _run_module_as_main File "<frozen runpy>", line 88, in _run_code File "~\testing_modules\test_melspectrogram_GPU.py", line 39, in <module> mel_spectrogram: torch.Tensor = torchaudio.transforms.MelSpectrogram( ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ File "~\AppData\Local\pypoetry\Cache\virtualenvs\testbed-rg5q6nje-py3.11\Lib\site-packages\torch\nn\modules\module.py", line 1511, in _wrapped_call_impl return self._call_impl(*args, **kwargs) ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ File "~\AppData\Local\pypoetry\Cache\virtualenvs\testbed-rg5q6nje-py3.11\Lib\site-packages\torch\nn\modules\module.py", line 1520, in _call_impl return forward_call(*args, **kwargs) ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ File "~\AppData\Local\pypoetry\Cache\virtualenvs\testbed-rg5q6nje-py3.11\Lib\site-packages\torchaudio\transforms\_transforms.py", line 619, in forward specgram = self.spectrogram(waveform) ^^^^^^^^^^^^^^^^^^^^^^^^^^ File "~\AppData\Local\pypoetry\Cache\virtualenvs\testbed-rg5q6nje-py3.11\Lib\site-packages\torch\nn\modules\module.py", line 1511, in _wrapped_call_impl return self._call_impl(*args, **kwargs) ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ File "~\AppData\Local\pypoetry\Cache\virtualenvs\testbed-rg5q6nje-py3.11\Lib\site-packages\torch\nn\modules\module.py", line 1520, in _call_impl return forward_call(*args, **kwargs) ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ File "~\AppData\Local\pypoetry\Cache\virtualenvs\testbed-rg5q6nje-py3.11\Lib\site-packages\torchaudio\transforms\_transforms.py", line 110, in forward return F.spectrogram( ^^^^^^^^^^^^^^ File "~\AppData\Local\pypoetry\Cache\virtualenvs\testbed-rg5q6nje-py3.11\Lib\site-packages\torchaudio\functional\functional.py", line 126, in spectrogram spec_f = torch.stft( ^^^^^^^^^^^ File "~\AppData\Local\pypoetry\Cache\virtualenvs\testbed-rg5q6nje-py3.11\Lib\site-packages\torch\functional.py", line 660, in stft return _VF.stft(input, n_fft, hop_length, win_length, window, # type: ignore[attr-defined] ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ RuntimeError: stft input and window must be on the same device but got self on cuda:0 and window on cpu
原因分析
torchaudio.transforms.MelSpectrogram是PyTorch的nn.Module子类,初始化时会默认在CPU上创建STFT所需的窗口张量(如汉明窗)。当输入张量被移到CUDA后直接调用该变换时,变换内部的窗口张量仍停留在CPU,导致设备不匹配,触发RuntimeError。
解决方法
有两种可行的修复方式:
方法一:将MelSpectrogram模块移到目标设备
创建MelSpectrogram实例后,调用.to(target_device)将模块及其内部所有参数(包括窗口张量)移到CUDA设备,再处理输入:
# 创建变换实例并移到目标设备 mel_transform = torchaudio.transforms.MelSpectrogram( sample_rate=TARGET_CLOCK, n_fft=NUMBER_OF_FFT_SLOTS, hop_length=HOP_LENGTH, n_mels=NUMBER_OF_MEL_SLOTS, ).to(target_device) # 处理输入 mel_spectrogram: torch.Tensor = mel_transform(frequency_tensor)
方法二:显式指定窗口并移到目标设备
手动创建窗口张量并移到目标设备,在初始化MelSpectrogram时传入该窗口:
# 创建窗口并移到目标设备 window = torch.hann_window(NUMBER_OF_FFT_SLOTS).to(target_device) # 初始化变换时传入窗口 mel_spectrogram: torch.Tensor = torchaudio.transforms.MelSpectrogram( sample_rate=TARGET_CLOCK, n_fft=NUMBER_OF_FFT_SLOTS, hop_length=HOP_LENGTH, n_mels=NUMBER_OF_MEL_SLOTS, window=window, )(frequency_tensor)
内容的提问来源于stack exchange,提问作者arc_lupus
相关产品推荐
相关产品推荐

