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

PyTorch仅对超阈值分母做除法:torch.where导致模型不收敛问题

问题原因与解决方法

你的torch.where写法确实存在数值稳定性问题,导致模型不收敛。

两种实现的核心差异

  • 原代码是仅对符合条件的位置执行除法,min_mask为True的位置完全不参与除法计算,不会产生数值异常。
  • 改用torch.where后,所有位置都会先计算wav / wsq_ola,包括wsq_ola绝对值小于eps的位置——这时候分母接近0,会产生无穷大(inf)或非数值(nan)。这些异常值会在反向传播中污染梯度,最终导致模型无法收敛。

修复方案

要保证只有分母安全的位置才执行除法,同时兼容torchdynamo编译,可以先对wsq_ola做安全替换,再直接执行除法:

wsq_ola = wsq_ola.to(wav).expand_as(wav).clone()
# 将绝对值小于eps的分母替换为1,避免除零
safe_denominator = torch.where(wsq_ola.abs() < eps, torch.tensor(1.0, device=wsq_ola.device), wsq_ola)
wav = wav / safe_denominator

这个逻辑和原代码完全一致:分母小于阈值时,等价于wav除以1,保留原值;分母符合要求时,正常执行除法。而且所有计算都是批量完成的,没有原地修改,能通过torchdynamo编译,同时避免了数值异常。

补充说明

因为你的逻辑属于无训练权重的STFT层,数值稳定性对后续梯度传播影响极大——哪怕少量nan或inf出现,都会快速扩散到整个计算图,导致模型参数更新混乱。所以必须确保除法操作全程没有无效计算。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.29 03:33:15