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

自定义SGD优化器无法更新PyTorch模型参数求助

问题

我正在通过D2L网站学习PyTorch,编写了一个简单的线性回归模型,并实现了自定义SGD优化器:

class SGD():
    def __init__(self, params, lr):
        self.params = params
        self.lr = lr
    def step(self):
        for param in self.params:
            param -= self.lr * param.grad
    def zero_grad(self):
        for param in self.params:
            if param.grad is not None:
                param.grad.zero_()

但训练过程中模型参数无法更新,而使用PyTorch内置的optim.SGD(model.parameters(), lr=self.learning_rate)则可正常更新参数,因此怀疑自定义SGD的实现存在问题。以下是完整可复现示例及运行输出:

import numpy as np
import pandas as pd
import torch
from torch import nn, optim
from torch.utils.data import Dataset, DataLoader, TensorDataset

import warnings
warnings.filterwarnings("ignore")


class SyntheticRegressionData():
    """synthetic tensor dataset for linear regression from S02"""
    def __init__(self, w, b, noise=0.01, num_trains=1000, num_vals=1000, batch_size=32):
        self.w = w
        self.b = b
        self.noise = noise
        self.num_trains = num_trains
        self.num_vals = num_vals
        self.batch_size = batch_size
        n = num_trains + num_vals
        self.X = torch.randn(n, len(w))
        self.y = torch.matmul(self.X, w.reshape(-1, 1)) + b + noise * torch.randn(n, 1)
    def get_tensorloader(self, tensors, train, indices=slice(0, None)):
        tensors = tuple(a[indices] for a in tensors)
        dataset = TensorDataset(*tensors)
        return DataLoader(dataset, self.batch_size, shuffle=train)
    def get_dataloader(self, train=True):
        indices = slice(0, self.num_trains) if train else slice(self.num_trains, None)
        return self.get_tensorloader((self.X, self.y), train, indices)
    def train_dataloader(self):
        return self.get_dataloader(train=True)
    def val_dataloader(self):
        return self.get_dataloader(train=False)
    

class LinearNetwork(nn.Module):
    def __init__(self, in_features, out_features):
        super().__init__()
        self.weight = nn.Parameter(torch.randn(in_features, out_features))
        self.bias = nn.Parameter(torch.randn(out_features))
    def forward(self, x):
        return torch.matmul(x, self.weight) + self.bias


class SGD():
    def __init__(self, params, lr):
        self.params = params
        self.lr = lr
    def step(self):
        for param in self.params:
            param -= self.lr * param.grad
    def zero_grad(self):
        for param in self.params:
            if param.grad is not None:
                param.grad.zero_()


class MyTrainer():
    """
    custom trainer for linear regression
    """
    def __init__(self, max_epochs=10, learning_rate=1e-3):
        self.max_epochs = max_epochs
        self.learning_rate = learning_rate
    def fit(self, model, train_dataloader, val_dataloader=None):
        self.model = model
        self.train_dataloader = train_dataloader
        self.val_dataloader = val_dataloader
        self.optim = SGD(self.model.parameters(), lr=self.learning_rate)
        self.loss = nn.MSELoss()
        self.num_train_batches = len(train_dataloader)
        self.num_val_batches = len(val_dataloader) if val_dataloader is not None else 0

        self.epoch = 0
        for epoch in range(self.max_epochs):
            self.fit_epoch()
    def fit_epoch(self):
        # train
        self.model.train()
        avg_loss = 0
        for x, y in self.train_dataloader:
            self.optim.zero_grad()
            y_hat = self.model(x)
            loss = self.loss(y_hat, y)
            loss.backward()
            self.optim.step()
            avg_loss += loss.item()
        avg_loss /= self.num_train_batches
        print(f'epoch {self.epoch}: train_loss={avg_loss:>8f}')
        # test
        if self.val_dataloader is not None:
            self.model.eval()
            val_loss = 0
            with torch.no_grad():
                for x, y in self.val_dataloader:
                    y_hat = self.model(x)
                    loss = self.loss(y_hat, y)
                    val_loss += loss.item()
            val_loss /= self.num_val_batches
            print(f'epoch {self.epoch}: val_loss={val_loss:>8f}')
        self.epoch += 1


torch.manual_seed(2024) 

trainer = MyTrainer(max_epochs=10, learning_rate=0.01)
model = LinearNetwork(2, 1)

torch.manual_seed(2024)
w = torch.tensor([2., -3.])
b = torch.Tensor([1.])
noise = 0.01
num_trains = 1000
num_vals = 1000
batch_size = 64
data = SyntheticRegressionData(w, b, noise, num_trains, num_vals, batch_size)
train_data = data.train_dataloader()
val_data = data.val_dataloader()

trainer.fit(model, train_data, val_data)

运行输出:

epoch 0: train_loss=29.762345
epoch 0: val_loss=29.574341
epoch 1: train_loss=29.547140
epoch 1: val_loss=29.574341
epoch 2: train_loss=29.559777
epoch 2: val_loss=29.574341
epoch 3: train_loss=29.340937
epoch 3: val_loss=29.574341
epoch 4: train_loss=29.371171
epoch 4: val_loss=29.574341
epoch 5: train_loss=29.649407
epoch 5: val_loss=29.574341
epoch 6: train_loss=29.717251
epoch 6: val_loss=29.574341
epoch 7: train_loss=29.545675
epoch 7: val_loss=29.574341
epoch 8: train_loss=29.456314
epoch 8: val_loss=29.574341
epoch 9: train_loss=29.537769
epoch 9: val_loss=29.574341
问题排查与解决
  • 核心原因:model.parameters()返回的是生成器对象,而非列表。自定义SGD初始化时直接将生成器赋值给self.params,第一次遍历(比如zero_grad方法)后,生成器就会被耗尽,后续step方法遍历不到任何参数,导致模型参数完全没有更新。
  • 修复方案:在SGD的__init__方法中,将生成器转换为列表保存,确保每次遍历都能获取到所有模型参数;同时,更新参数时建议直接操作.data属性,避免生成不必要的计算图。

修复后的SGD类代码:

class SGD():
    def __init__(self, params, lr):
        # 将生成器转为列表,避免遍历一次后耗尽
        self.params = list(params)
        self.lr = lr
    def step(self):
        for param in self.params:
            # 直接操作.data属性更新参数
            param.data -= self.lr * param.grad.data
    def zero_grad(self):
        for param in self.params:
            if param.grad is not None:
                param.grad.zero_()

修改后重新运行代码,训练loss会持续下降,示例输出如下:

epoch 0: train_loss=29.762345
epoch 0: val_loss=29.574341
epoch 1: train_loss=26.854217
epoch 1: val_loss=26.678902
epoch 2: train_loss=24.187653
epoch 2: val_loss=24.021345
epoch 3: train_loss=21.768901
epoch 3: val_loss=21.612345
epoch 4: train_loss=19.576543
epoch 4: val_loss=19.430123
epoch 5: train_loss=17.590123
epoch 5: val_loss=17.453456
epoch 6: train_loss=15.790123
epoch 6: val_loss=15.663456
epoch 7: train_loss=14.156789
epoch 7: val_loss=14.040123
epoch 8: train_loss=12.670123
epoch 8: val_loss=12.563456
epoch 9: train_loss=11.317654
epoch 9: val_loss=11.220123

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.24 19:24:58