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

如何修复PyTorch Android中forward()期望Tensor类型的CppException错误

问题:PyTorch Android加载模型时forward()参数类型不匹配错误

错误详情

Android端调用Module.load('mel_ptmobile_v2.pt')加载模型后,调用forward()抛出CppException,提示参数'x'期望Tensor类型,但实际传入Dynamic<128>[Dynamic<1>,]类型。

Android端日志:

mel, melInputTensor = org.pytorch.Tensor$Tensor_float32, [1, 840]
Caused by: com.facebook.jni.CppException: forward() Expected a value of type 'Tensor' for argument 'x' 
           but instead found type 'Dynamic&lt;128&gt;[Dynamic&lt;1&gt;,]'.
                 Position: 1
                 Declaration: forward(__torch__.models.preprocess.___torch_mangle_11.AugmentMelSTFT self, Tensor x) -&gt; Tensor
                 Exception raised from checkArg at /Users/huydo/Storage/mine/pytorch/aten/src/ATen/core/function_schema_inl.h:340 (most recent call first):
                 (no backtrace available)
                    at org.pytorch.NativePeer.forward(Native Method)
                    at org.pytorch.Module.forward(Module.java:52)

Android端调用代码:

val wavBatch = 1
val wavLength = 840
val dummyInput = dummyInput(wavBatch, wavLength, 0.0f)
val inputShape = longArrayOf(wavBatch.toLong(), wavLength.toLong())
melInputTensor = Tensor.fromBlob(dummyInput, inputShape)
if (DEBUG) Log.i(TAG,
    "mel, melInputTensor = " + melInputTensor?.javaClass?.name + ", " + melInputTensor?.shape().contentToString()
)

// 由com.facebook.jni.CppException引起:forward()期望参数'x'为Tensor类型,却发现类型为Dynamic&lt;128&gt;[Dynamic&lt;1&gt;,]。
val forward = melModel!!.forward(IValue.listFrom(melInputTensor))
// 由java.lang.IllegalStateException引起:期望IValue类型为Tuple,实际为Tensor
//val forward = melModel!!.forward(IValue.from(melInputTensor))

Python端正常运行情况

相同模型在Python中通过torch.load('mel_ptmobile_v2.pt')可正常运行:

Python端日志:

inputs 0 134400 [-1.4953613e-03 -1.6479492e-03 -1.4648438e-03 ... 
inputs 1 torch.Size([1, 134400]) tensor([[-1.4954e-03, -1.6479e-03, -1.4648e-03,  . <class 'torch.Tensor'>
inputs 2 torch.Size([1, 128, 420]) tensor([[[-0.7647, -0.5746, -0.6255,  ..., -1.3985
outputs 0 torch.Size([1, 527]) tensor([[ -3.1250,  -6.5625,  -7.0312,  -7.7500, 

Python端调用代码:

# 将模型用于将波形预处理为梅尔频谱
mel = load_model_from_uri(mel_ptmobile_name)

(waveform, _) = librosa.core.load(audio_path, sr=sample_rate, mono=True)
if DEBUG: print('inputs 0', len(waveform), str(waveform)[:50])
waveform = torch.from_numpy(waveform[None, :]).to(device)
if DEBUG: print('inputs 1', waveform.shape, str(waveform)[:50], type(waveform))

# 我们的模型以半精度模式(torch.float16)训练
# 在cuda上运行torch.float16可获得最佳性能
# 在cpu上运行torch.float32性能相近,使用torch.bfloat16性能更差
with torch.no_grad(), autocast(device_type=device.type) if cuda else nullcontext():
    spec = mel(waveform)
    if DEBUG: print('inputs 2', spec.shape, str(spec)[:50])

Python端模型定义:

class AugmentMelSTFT(nn.Module):
    def __init__(self, n_mels=128, sr=32000, win_length=800, hopsize=320, n_fft=1024, freqm=48, timem=192,
                 fmin=0.0, fmax=None, fmin_aug_range=10, fmax_aug_range=2000):
        torch.nn.Module.__init__(self)
        # adapted from: https://github.com/CPJKU/kagglebirds2020/commit/70f8308b39011b09d41eb0f4ace5aa7d2b0e806e

        self.win_length = win_length
        self.n_mels = n_mels
        self.n_fft = n_fft
        self.sr = sr
        self.fmin = fmin
        if fmax is None:
            fmax = sr // 2 - fmax_aug_range // 2
            if DEBUG: print(f"Warning: FMAX is None setting to {fmax} ")
        self.fmax = fmax
        self.hopsize = hopsize
        self.register_buffer('window',
                             torch.hann_window(win_length, periodic=False),
                             persistent=False)
        assert fmin_aug_range >= 1, f"fmin_aug_range={fmin_aug_range} should be >=1; 1 means no augmentation"
        assert fmax_aug_range >= 1, f"fmax_aug_range={fmax_aug_range} should be >=1; 1 means no augmentation"
        self.fmin_aug_range = fmin_aug_range
        self.fmax_aug_range = fmax_aug_range

        self.register_buffer("preemphasis_coefficient", torch.as_tensor([[[-.97, 1]]]), persistent=False)
        if freqm == 0:
            self.freqm = torch.nn.Identity()
        else:
            self.freqm = torchaudio.transforms.FrequencyMasking(freqm, iid_masks=True)
        if timem == 0:
            self.timem = torch.nn.Identity()
        else:
            self.timem = torchaudio.transforms.TimeMasking(timem, iid_masks=True)

    def forward(self, x):
        if onnx_conf.DEBUG: print('mel.forward,', x.shape, x[0][0].dtype, type(x))
        x = nn.functional.conv1d(x.unsqueeze(1), self.preemphasis_coefficient).squeeze(1)
        x = torch.stft(x, self.n_fft, hop_length=self.hopsize, win_length=self.win_length,
                       center=True, normalized=False, window=self.window, return_complex=False)
        # x = stft(x, self.n_fft, hop_length=self.hopsize, win_length=self.win_length,
        #          center=True, normalized=False, window=self.window, return_complex=False)
        x = (x ** 2).sum(dim=-1)  # power mag
        fmin = self.fmin + torch.randint(self.fmin_aug_range, (1,)).item()
        fmax = self.fmax + self.fmax_aug_range // 2 - torch.randint(self.fmax_aug_range, (1,)).item()
        # 不对评估数据做增强
        if not self.training:
            fmin = self.fmin
            fmax = self.fmax

        mel_basis, _ = torchaudio.compliance.kaldi.get_mel_banks(self.n_mels, self.n_fft, self.sr,
                                                                 fmin, fmax, vtln_low=100.0, vtln_high=-500.,
                                                                 vtln_warp_factor=1.0)
        mel_basis = torch.as_tensor(torch.nn.functional.pad(mel_basis, (0, 1), mode='constant', value=0),
                                    device=x.device)
        with torch.cuda.amp.autocast(enabled=False):
            melspec = torch.matmul(mel_basis, x)

        melspec = (melspec + 0.00001).log()

        if self.training:
            melspec = self.freqm(melspec)
            melspec = self.timem(melspec)

        melspec = (melspec + 4.5) / 5.  # 快速归一化

        return melspec

解决方案

1. 重新导出TorchScript格式模型

PyTorch Android仅支持TorchScript格式的模型,直接用torch.save保存的普通nn.Module无法在Android上正确运行,必须通过torch.jit.trace或torch.jit.script导出:

import torch
from your_model_module import AugmentMelSTFT

# 初始化模型(参数要和训练时一致)
model = AugmentMelSTFT(n_mels=128, sr=32000, win_length=800, hopsize=320, n_fft=1024, freqm=48, timem=192)
# 加载训练好的权重
model.load_state_dict(torch.load('mel_ptmobile_v2.pt'))
# 设置为评估模式
model.eval()

# 创建一个和Android端输入形状一致的dummy tensor
dummy_input = torch.randn(1, 840)
# 追踪模型生成TorchScript
traced_model = torch.jit.trace(model, dummy_input)
# 导出模型
traced_model.save('mel_ptmobile_v2_ts.pt')

注意:如果模型中有动态控制流(比如if self.training),建议用torch.jit.script代替torch.jit.trace,确保所有分支都被正确捕获。

2. 修正Android端调用代码

模型导出后,Android端需传入单个Tensor对应的IValue,而非列表:

// 替换原来的forward调用
val outputIValue = melModel!!.forward(IValue.from(melInputTensor))
// 转换为Tensor
val outputTensor = outputIValue.toTensor()

额外注意事项

  • 确保Android端输入Tensor的形状、数据类型和Python端一致(比如都是float32)。
  • 导出模型时需将模型设置为eval()模式,避免训练时的增强逻辑(比如频率/时间掩码)在推理阶段执行。

内容的提问来源于stack exchange,提问作者thecr0w

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.29 07:00:54