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
相关产品推荐
相关产品推荐

