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
梯度中断的原因分析
协方差矩阵构建切断计算图:
在loss函数中,手动用torch.tensor([[sigmas[0]**2, 0], [0, sigmas[1]**2]], requires_grad=True)创建协方差矩阵,这会生成一个新的独立张量,完全切断了与原始sigma_x、sigma_y的梯度传播链路——原参数的梯度无法传递到这个新张量,反向传播时自然无法更新sigma_x和sigma_y。多余的
.requires_grad_(True)调用:estimated_distribution.log_prob(coords)的输出本身就带有梯度信息(因为输入的协方差矩阵依赖可优化参数),手动调用.requires_grad_(True)会覆盖原有计算图,干扰梯度传播。真实分布创建的潜在问题:
原代码中未导入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
相关产品推荐
相关产品推荐

