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

使用torch.normal触发RuntimeError:要求所有std≥0.0,求解决方案

解决torch.normal的RuntimeError: std >= 0.0问题

问题代码

b=32
n_s = 10
dim = 64
slots_mu = nn.Parameter(torch.randn(1, 1, dim))
slots_log_sigma = nn.Parameter(torch.randn(1, 1, dim))

mu = slots_mu.expand(b, n_s, -1)
sigma = slots_log_sigma.expand(b, n_s, -1)

slots = torch.normal(mu, sigma)

错误信息

---> 10 slots = torch.normal(mu, sigma)

RuntimeError: normal expects all elements of std >= 0.0

问题原因

slots_log_sigma 使用 torch.randn() 初始化,会生成正负随机值。直接将其作为标准差(sigma)传入 torch.normal() 违反了标准差必须非负的要求,因此触发报错。

解决方案

需要将对数形式的sigma转换为非负的标准差,常用两种方法:

  1. 指数转换:用torch.exp()将对数sigma转为正数,这是最直接的方式
  2. Softplus转换:用torch.nn.functional.softplus(),相比指数更稳定,避免极端大值(当log_sigma很大时,exp会导致数值爆炸)

修改后的代码示例:

方法1:指数转换

b=32
n_s = 10
dim = 64
slots_mu = nn.Parameter(torch.randn(1, 1, dim))
slots_log_sigma = nn.Parameter(torch.randn(1, 1, dim))

mu = slots_mu.expand(b, n_s, -1)
# 用exp转换为非负标准差
sigma = torch.exp(slots_log_sigma.expand(b, n_s, -1))

slots = torch.normal(mu, sigma)

方法2:Softplus转换

import torch.nn.functional as F

b=32
n_s = 10
dim = 64
slots_mu = nn.Parameter(torch.randn(1, 1, dim))
slots_log_sigma = nn.Parameter(torch.randn(1, 1, dim))

mu = slots_mu.expand(b, n_s, -1)
# 用softplus转换为非负标准差,数值更稳定
sigma = F.softplus(slots_log_sigma.expand(b, n_s, -1))

slots = torch.normal(mu, sigma)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.20 11:01:30