You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.26 18:58:17