自定义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
相关产品推荐
相关产品推荐

