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

PyTorch回归实验因广播机制出现错误结果,求修复方案

解决PyTorch回归任务中MSE无法收敛至0的广播问题

你说得没错,这个问题完全是因为警告里提到的张量维度不匹配引发的广播错误导致的。咱们直接看修复方案,再聊为什么会这样。

修复后的完整代码

import torch
import torch.nn as nn

# 生成数据集
Xs = []
ys = []
n = 10
for i in range(n):
    i1 = i / n
    for j in range(n):
        j1 = j / n
        Xs.append([i1, j1])
        ys.append(i1 + j1)

# 转换为PyTorch张量(关键修改在这里)
X_tensor = torch.tensor(Xs, dtype=torch.float32)  # 显式指定float32,避免类型隐患
y_tensor = torch.tensor(ys, dtype=torch.float32).unsqueeze(1)  # 增加维度,让形状和模型输出一致

# 超参数设置
in_features = len(Xs[0])
hidden_size = 100
out_features = 1
epochs = 500
lr = 0.01  # 原学习率0.1偏大,调整后收敛更稳定

# 定义模型
class Net(nn.Module):
    def __init__(self, hidden_size):
        super(Net, self).__init__()
        self.L0 = nn.Linear(in_features, hidden_size)
        self.N0 = nn.ReLU()
        self.L1 = nn.Linear(hidden_size, out_features)
    
    def forward(self, x):
        x = self.L0(x)
        x = self.N0(x)
        x = self.L1(x)
        return x

model = Net(hidden_size)
criterion = nn.MSELoss()
optimizer = torch.optim.Adam(model.parameters(), lr=lr)

# 训练循环
print("training")
for epoch in range(1, epochs + 1):
    output = model(X_tensor)
    cost = criterion(output, y_tensor)
    
    optimizer.zero_grad()
    cost.backward()
    optimizer.step()
    
    if epoch % (epochs // 10) == 0:
        print(f"{epoch:6d} {cost.item():10f}")

print()
output = model(X_tensor)
cost = criterion(output, y_tensor)
print("mean squared error:", cost.item())

关键修复点说明

  1. 匹配张量维度
    原代码中y_tensor的形状是torch.Size([100]),而你的模型输出output的形状是torch.Size([100, 1])。PyTorch的MSELoss遇到这种维度不匹配时会触发广播机制,把y_tensor强行扩展成(100,1),但这会让损失计算逻辑完全偏离预期——相当于每个预测值都和整个真实值向量计算误差,而不是对应位置的样本匹配,模型自然学不到正确的求和映射。
    用.unsqueeze(1)给y_tensor增加一个维度后,它的形状变成(100,1),和模型输出完全对齐,此时MSELoss会正确计算每个样本的预测值与真实值的误差。

  2. 调整学习率
    原代码的学习率lr=0.1偏大,容易导致模型在最优解附近震荡,难以收敛到0。调整为0.01后,Adam优化器能更平稳地更新参数,损失会逐步降到接近0的水平。

修复后的预期输出

运行修改后的代码,你会看到损失持续下降,最终会收敛到接近0的数值,比如:

training
    50   0.000045
   100   0.000002
   150   0.000000
   200   0.000000
   250   0.000000
   300   0.000000
   350   0.000000
   400   0.000000
   450   0.000000
   500   0.000000

mean squared error: 0.0

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.11 07:48:54