使用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转换为非负的标准差,常用两种方法:
- 指数转换:用
torch.exp()将对数sigma转为正数,这是最直接的方式 - 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
相关产品推荐
相关产品推荐

