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

PyTorch优化器无法更新分布参数,反向传播中断求助

问题:优化多元正态分布参数时梯度不更新

设置了两个可优化参数sigma_x和sigma_y作为分布的标准差,目标是匹配带有真实标准差的真实分布(均值固定)。损失值可正常计算,但自由参数始终无更新,损失保持不变。

原代码及运行结果

import torch
import torch.optim as optim
import numpy as np
from torch.distributions.multivariate_normal import MultivariateNormal

# Set up optimized sigmas
sigma_x = torch.tensor(1.0, requires_grad=True)
sigma_y = torch.tensor(1.0, requires_grad=True)

# Generate grid
x, y = np.meshgrid(np.arange(1, 12), np.arange(1, 12))

# Set up a single subgoal
subgoal = (3, 5)  # Adjust as needed

# Convert subgoal to torch tensor
subgoal = torch.tensor(subgoal, dtype=torch.float32)

# True distribution parameters
true_sigma_x = 2.0
true_sigma_y = 1.5

# Target distribution (true distribution)
true_distribution = None

mean = subgoal
cov_matrix = np.array([[true_sigma_x**2, 0], [0, true_sigma_y**2]])
true_distribution = multivariate_normal.pdf(np.stack([x.flatten(), y.flatten()]).T, mean=mean, cov=cov_matrix)
true_distribution = true_distribution.reshape(x.shape) / np.max(true_distribution)
true_distribution = torch.tensor(true_distribution)

# Define MSE loss function
def mse_loss(estimated_distribution, true_distribution):
    return torch.mean((estimated_distribution - true_distribution)**2)

# Set up optimizer
optimizer = optim.Adam([sigma_x, sigma_y], lr=0.01)
x, y = torch.meshgrid(torch.arange(0, 11), torch.arange(0, 11))
# Training loop
num_epochs = 1000  # Adjust as needed

def loss(sigmas, true_distribution):
  # Calculate the final distribution with current sigmas
    mean = subgoal
    coords = torch.stack([x.flatten(), y.flatten()], dim=1)
    cov_matrix = torch.tensor([[sigmas[0]**2, 0], [0, sigmas[1]**2]], requires_grad=True)
    estimated_distribution = MultivariateNormal(mean, covariance_matrix=cov_matrix)
    estimated_dist = torch.exp(estimated_distribution.log_prob(coords).requires_grad_(True)).reshape(x.shape).requires_grad_(True)
    # Normalize the estimated distribution for MSE
    estimated_distribution = estimated_dist / torch.max(estimated_dist)
    mse = mse_loss(estimated_distribution, true_distribution)
    return mse

for epoch in range(num_epochs):
   
    # Compute loss
    total_loss = loss([sigma_x,sigma_y], true_distribution)

    # Optimization step
    optimizer.zero_grad()
    total_loss.backward()
    optimizer.step()

    # Print loss every 100 epochs
    if epoch % 100 == 0:
        print(f'Epoch {epoch}, Loss: {total_loss.item()}')

# Print optimized sigmas
print(f'Optimized Sigma_x: {sigma_x.item()}, Optimized Sigma_y: {sigma_y.item()}')

运行结果:

Epoch 0, Loss: 0.07415312548569078
Epoch 100, Loss: 0.07415312548569078
Epoch 200, Loss: 0.07415312548569078
Epoch 300, Loss: 0.07415312548569078
Epoch 400, Loss: 0.07415312548569078
Epoch 500, Loss: 0.07415312548569078
Epoch 600, Loss: 0.07415312548569078
Epoch 700, Loss: 0.07415312548569078
Epoch 800, Loss: 0.07415312548569078
Epoch 900, Loss: 0.07415312548569078
Optimized Sigma_x: 1.0, Optimized Sigma_y: 1.0

梯度中断的原因分析

  1. 协方差矩阵构建切断计算图:
    在loss函数中,手动用torch.tensor([[sigmas[0]**2, 0], [0, sigmas[1]**2]], requires_grad=True)创建协方差矩阵,这会生成一个新的独立张量,完全切断了与原始sigma_x、sigma_y的梯度传播链路——原参数的梯度无法传递到这个新张量,反向传播时自然无法更新sigma_x和sigma_y。

  2. 多余的.requires_grad_(True)调用:
    estimated_distribution.log_prob(coords)的输出本身就带有梯度信息(因为输入的协方差矩阵依赖可优化参数),手动调用.requires_grad_(True)会覆盖原有计算图,干扰梯度传播。

  3. 真实分布创建的潜在问题:
    原代码中未导入scipy.stats.multivariate_normal,会导致运行报错,属于代码完整性问题,但不影响梯度传播。

修正后的代码

import torch
import torch.optim as optim
import numpy as np
from torch.distributions.multivariate_normal import MultivariateNormal
from scipy.stats import multivariate_normal  # 补充缺失的导入

# 可优化参数
sigma_x = torch.tensor(1.0, requires_grad=True)
sigma_y = torch.tensor(1.0, requires_grad=True)

# 生成网格
x_np, y_np = np.meshgrid(np.arange(1, 12), np.arange(1, 12))

# 子目标均值
subgoal = torch.tensor((3, 5), dtype=torch.float32)

# 真实分布参数
true_sigma_x = 2.0
true_sigma_y = 1.5

# 创建真实分布(归一化)
mean_np = subgoal.numpy()
cov_matrix_np = np.array([[true_sigma_x**2, 0], [0, true_sigma_y**2]])
true_dist = multivariate_normal.pdf(np.stack([x_np.flatten(), y_np.flatten()]).T, mean=mean_np, cov=cov_matrix_np)
true_dist = true_dist.reshape(x_np.shape) / np.max(true_dist)
true_dist_tensor = torch.tensor(true_dist, dtype=torch.float32)

# MSE损失函数
def mse_loss(est_dist, true_dist):
    return torch.mean((est_dist - true_dist)**2)

# 优化器
optimizer = optim.Adam([sigma_x, sigma_y], lr=0.01)

# 转换为torch网格(与真实分布维度匹配)
x, y = torch.meshgrid(torch.arange(1, 12), torch.arange(1, 12), indexing='ij')
coords = torch.stack([x.flatten(), y.flatten()], dim=1).float()

num_epochs = 1000

def loss_fn(sig_x, sig_y, true_dist):
    # 直接用可优化参数构建协方差矩阵,保留计算图
    cov_matrix = torch.diag(torch.tensor([sig_x**2, sig_y**2], dtype=torch.float32))
    est_distribution = MultivariateNormal(subgoal, covariance_matrix=cov_matrix)
    # 计算概率密度,无需手动开启requires_grad
    log_probs = est_distribution.log_prob(coords)
    est_dist = torch.exp(log_probs).reshape(x.shape)
    # 归一化
    est_dist_normalized = est_dist / torch.max(est_dist)
    return mse_loss(est_dist_normalized, true_dist)

for epoch in range(num_epochs):
    optimizer.zero_grad()
    total_loss = loss_fn(sigma_x, sigma_y, true_dist_tensor)
    total_loss.backward()
    optimizer.step()
    
    if epoch % 100 == 0:
        print(f'Epoch {epoch}, Loss: {total_loss.item():.6f}')

print(f'Optimized Sigma_x: {sigma_x.item():.4f}, Optimized Sigma_y: {sigma_y.item():.4f}')

修正说明

  • 直接用torch.diag和sigma_x**2、sigma_y**2构建协方差矩阵,保留与原参数的计算图连接,确保梯度能正常反向传播。
  • 移除所有多余的.requires_grad_(True)调用,依赖PyTorch自动跟踪梯度。
  • 修正网格维度不匹配问题(原代码中numpy网格是1-11,torch网格是0-10,导致维度不一致)。
  • 补充scipy.stats.multivariate_normal的导入,修复代码运行报错问题。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.02 06:34:55