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

RuntimeError:矩阵形状不匹配(摄氏度转华氏度模型训练报错)

问题原因分析

错误RuntimeError: mat1 and mat2 shapes cannot be multiplied (1x7 and 1x1)的核心是输入张量形状不匹配:

  • 你定义的nn.Linear(1, 1)层要求输入形状为(batch_size, 1)(每个样本对应1维特征)。
  • 但当前inputs = torch.from_numpy(celsius_arr)得到的是一维张量,形状为(7,)。PyTorch会自动将其视为(1, 7)的二维张量(把整个数组当作1个含7个特征的样本),这和线性层要求的输入维度完全不符,导致矩阵乘法失败。
修复方案

需要将输入和标签张量调整为(样本数, 特征数)的二维形状,具体有两种实现方式:

方式1:用unsqueeze()添加维度

在训练循环中修改输入和标签的转换代码:

inputs = torch.from_numpy(celsius_arr).unsqueeze(1)  # 形状变为(7, 1)
labels = torch.from_numpy(fahrenheit_arr).unsqueeze(1)  # 形状变为(7, 1)

方式2:用reshape()调整形状

同样在训练循环中修改:

inputs = torch.from_numpy(celsius_arr).reshape(-1, 1)  # -1自动计算样本数,最终形状(7,1)
labels = torch.from_numpy(fahrenheit_arr).reshape(-1, 1)

额外注意:推理阶段的输入形状

推理时输入的torch.tensor([100.0])也是一维张量,同样需要调整形状,否则会报相同错误:

print(model(torch.tensor([100.0]).unsqueeze(1)))
完整修复后的代码
# 用PyTorch训练摄氏度转华氏度预测模型的简单示例
import numpy as np
import torch
import torch.nn as nn

celsius_arr = [-40, -10, 0, 8, 15, 22, 38]
fahrenheit_arr = [-40, 14, 32, 46, 59, 72, 100]

# 转换为numpy数组
celsius_arr = np.array(celsius_arr, dtype=np.float32)
fahrenheit_arr = np.array(fahrenheit_arr, dtype=np.float32)

# 定义模型
class Net(nn.Module):
    def __init__(self):
        super(Net, self).__init__()
        self.linear = nn.Linear(1, 1)

    def forward(self, x):
        x = self.linear(x)
        return x

# 构建损失函数和优化器
model = Net()
criterion = nn.MSELoss()
optimizer = torch.optim.SGD(model.parameters(), lr=0.01)

# 训练循环
for epoch in range(1000):
    # 调整输入和标签的形状为(7,1)
    inputs = torch.from_numpy(celsius_arr).unsqueeze(1)
    labels = torch.from_numpy(fahrenheit_arr).unsqueeze(1)

    # 前向传播
    outputs = model(inputs)
    loss = criterion(outputs, labels)

    # 反向传播与优化
    optimizer.zero_grad()
    loss.backward()
    optimizer.step()

    if epoch % 100 == 0:
        print(f'epoch {epoch}, loss = {loss.item():.4f}')

# 推理:调整输入形状为(1,1)
print(model(torch.tensor([100.0]).unsqueeze(1)))

# 保存模型
torch.save(model.state_dict(), 'celsius_fahrenheit_model.pth')

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.17 01:55:20