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

如何用PyTorch建模并训练神经网络学习数学函数f(x)?

问题分析与正确实现方案

你的代码核心问题

  • 网络结构完全不匹配任务:f(x)=2x-1是纯线性函数,你的模型用了两层线性层+ReLU激活,不仅冗余,ReLU的非线性还会破坏线性关系;同时输入输出维度错误——单个x是1维特征,你却把batch_size=10作为输入层的in_features,混淆了特征维度和批量大小的概念。
  • 缺少完整训练流程:只定义了模型,没有数据生成、损失函数、优化器和训练循环,根本无法完成参数更新。

正确实现代码

import torch
import torch.nn as nn
import torch.optim as optim

# 1. 生成训练数据
# 生成10个[-5,5]范围内的随机x,形状为[10,1](批量10,每个样本1个特征)
x = torch.rand(10, 1) * 10 - 5
y = 2 * x - 1  # 直接生成对应标签

# 2. 定义匹配任务的模型
class LinearModel(nn.Module):
    def __init__(self):
        super().__init__()
        # 输入维度1(单个x值),输出维度1(预测的y值),单层线性层足够
        self.linear = nn.Linear(1, 1)
    
    def forward(self, x):
        return self.linear(x)  # 线性输出,无需激活函数

model = LinearModel()

# 3. 定义损失函数与优化器
criterion = nn.MSELoss()  # 回归任务用均方误差损失
optimizer = optim.SGD(model.parameters(), lr=0.01)  # 随机梯度下降优化器

# 4. 训练循环
epochs = 1000
for epoch in range(epochs):
    # 前向传播计算预测值
    y_pred = model(x)
    # 计算损失
    loss = criterion(y_pred, y)
    
    # 反向传播与参数更新
    optimizer.zero_grad()  # 清空上一轮梯度
    loss.backward()        # 反向传播计算梯度
    optimizer.step()       # 更新模型参数
    
    # 每100轮打印损失
    if (epoch + 1) % 100 == 0:
        print(f"Epoch [{epoch+1}/{epochs}], Loss: {loss.item():.4f}")

# 验证训练结果
print("\n训练后模型参数:")
print(f"权重w: {model.linear.weight.item():.4f}")
print(f"偏置b: {model.linear.bias.item():.4f}")

# 测试新样本
test_x = torch.tensor([[3.0], [-2.0]])
test_y_pred = model(test_x)
print("\n测试结果:")
for xi, yi_pred in zip(test_x, test_y_pred):
    print(f"x={xi.item()}, 预测y={yi_pred.item():.4f}, 真实y={2*xi.item()-1:.4f}")

关键说明

  • 模型设计:目标函数是纯线性的,单层nn.Linear(1,1)的数学表达式就是y = w*x + b,正好对应我们要学习的y=2x-1(w=2,b=-1),不需要任何激活函数。
  • 数据维度:要明确区分批量大小和特征维度——每个样本是单个数值(特征维度1),批量大小是10,所以数据形状为[10,1],而非把批量大小当作特征维度。
  • 训练流程:必须包含「前向传播→计算损失→清空梯度→反向传播→优化器更新」这几个核心步骤,才能让模型逐步学到正确的参数。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.31 04:39:34