如何修复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<128>[Dynamic<1>,]'. Position: 1 Declaration: forward(__torch__.models.preprocess.___torch_mangle_11.AugmentMelSTFT self, Tensor x) -> 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<128>[Dynamic<1>,]。 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
相关产品推荐
相关产品推荐

