PyTorch实现Kendall et al. 2018自步多任务加权损失时出现负损失,求问题排查
PyTorch实现Kendall et al. 2018自步多任务加权损失时出现负损失,求问题排查
嘿,我来帮你理一理这个问题!首先得明确:训练中出现负损失不一定是代码逻辑错误,但你的实现确实有可以优化的地方,我们可以一步步拆解原因和解决方案。
为什么会出现负损失?
你的代码里的加权损失逻辑是符合原论文公式的——对应 ( \frac{1}{2\sigma^2}L + \log\sigma )。但这里的 ( \log\sigma ) 项有个特点:当 ( \sigma < 1 ) 时,它是负数。如果这部分的绝对值超过了 ( \frac{1}{2\sigma^2}L ) 的值,整体损失就会变成负数。
举个实际例子:
- 假设模型学到的 ( \sigma = 0.5 )(小于1),那 ( \log(0.5) \approx -0.693 )
- 如果此时任务的原始损失 ( L = 0.1 ),那么 ( \frac{1}{2*(0.5)^2}0.1 = 20.1 = 0.2 )
- 加权损失就是 ( 0.2 + (-0.693) = -0.493 ),确实会是负数
这种情况在训练初期或者任务原始损失本身很小的时候特别容易出现——此时模型还没学到合适的σ值,可能让σ偏小,导致log项的负贡献盖过了损失项的正贡献。
你的实现可以优化的地方(同时缓解负损失)
原论文其实推荐了更稳定的参数化方式:不要直接学习σ,而是学习 ( \log\sigma )(即对数形式的不确定性),这样能避免用softplus做约束,数值稳定性更好,还能减少σ过小的情况。
给你调整后的代码,更贴近原论文的官方推荐实现:
import torch def apply_uncertainty_weights(log_sigma, loss): """ 基于Kendall et al. 2018的不确定性加权损失实现 参数: log_sigma: 模型学习的对数不确定性参数,初始化为torch.nn.Parameter(torch.zeros(1)) loss: 预测和目标之间的原始损失 返回: weighted_loss: 加权后的损失 """ # 从log_sigma推导σ²,完全避免数值不稳定问题 sigma_squared = torch.exp(2 * log_sigma) # log_sigma = logσ → 2logσ = logσ² → exp后得到σ² # 严格对应原论文公式计算加权损失 weighted_loss = (loss / (2 * sigma_squared)) + log_sigma return weighted_loss
这种实现的优势:
- 直接学习logσ,参数可以取任意实数值,不需要额外用softplus约束为正(因为σ=exp(logσ)天生为正)
- 当σ有变小的趋势时,( \frac{1}{2\sigma^2}L ) 会急剧增大,损失会变得很高,模型会主动避免σ过小的情况,自然减少logσ的负贡献盖过损失项的概率
- 去掉了手动加1e-8这类trick,代码更简洁也更贴近原论文逻辑
额外注意事项
- 如果是多任务学习,每个任务必须有独立的log_sigma参数,这是原论文的核心要求,模型才能为不同任务学习到专属的不确定性权重
- 训练中短暂出现负损失不一定是bug,只要模型最终能收敛到合理性能就没问题。但如果负损失持续存在且模型不收敛,要检查:
- 原始损失L的计算是否正确(比如是不是误用了负向损失函数?)
- 学习率是否过高,导致logσ的更新过于激进
- 多任务场景下是否把所有任务的加权损失加总后再反向传播,而不是单独看单个任务的负损失
备注:内容来源于stack exchange,提问作者bird
相关产品推荐
相关产品推荐

