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

PyTorch神经网络参数更新及参数整体乘以0.9的实现方法

嘿,我来帮你拆解这两个PyTorch相关的问题:

1. 请问如何在PyTorch框架中更新神经网络的参数?

在PyTorch里,参数更新主要有两种常见方式,适配不同的使用场景:

方式一:使用优化器自动更新(最常用)

这是模型训练的标准流程,PyTorch的优化器会帮你处理梯度计算和参数更新的细节,步骤非常清晰:

  1. 先定义好模型、损失函数,再初始化优化器(需要传入模型的所有可训练参数)
  2. 前向传播计算预测结果和损失值
  3. 清空之前累积的梯度(避免梯度叠加影响更新)
  4. 反向传播计算每个参数的梯度
  5. 调用优化器的step()方法完成参数更新

示例代码:

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

# 定义一个简单的全连接网络
class MyNet(nn.Module):
    def __init__(self):
        super().__init__()
        self.fc1 = nn.Linear(20, 10)
        self.fc2 = nn.Linear(10, 2)
    
    def forward(self, x):
        x = torch.relu(self.fc1(x))
        return self.fc2(x)

# 初始化核心组件
model = MyNet()
criterion = nn.CrossEntropyLoss()
optimizer = optim.Adam(model.parameters(), lr=0.001)

# 模拟输入数据和标签
input_data = torch.randn(16, 20)
labels = torch.randint(0, 2, (16,))

# 前向传播计算损失
output = model(input_data)
loss = criterion(output, labels)

# 执行参数更新
optimizer.zero_grad()  # 清空梯度缓存
loss.backward()        # 反向传播计算梯度
optimizer.step()       # 优化器自动完成参数更新

方式二:手动更新参数(自定义逻辑)

如果需要实现特殊的更新规则(比如自定义学习率调度、梯度裁剪),可以手动遍历模型参数进行修改。注意要在torch.no_grad()上下文里操作,避免PyTorch跟踪这个修改的梯度。

示例代码(以基础梯度下降为例):

learning_rate = 0.001

# 先完成梯度计算
optimizer.zero_grad()
loss.backward()

# 手动更新参数
with torch.no_grad():
    for param in model.parameters():
        param -= learning_rate * param.grad
2. 假设我需要将PyTorch中继承自torch.nn.Module类的神经网络实例的所有参数都乘以0.9,该如何实现?

这个需求很容易实现,只需要遍历模型的所有参数,直接对参数值进行缩放操作即可,同样要在torch.no_grad()上下文里执行,避免不必要的梯度跟踪:

示例代码:

# 假设model是你的神经网络实例
with torch.no_grad():
    for param in model.parameters():
        param.data *= 0.9

补充说明:

  • 使用param.data是为了直接修改参数的底层张量数据,不会触发梯度计算的跟踪;如果在no_grad()上下文里,直接写param *= 0.9也是有效的,但param.data的写法更明确,能清晰表达我们只是修改参数值的意图。
  • 不管参数是否设置了requires_grad=False(比如冻结的层),这个操作都能生效,因为我们是直接修改参数的存储值。

内容的提问来源于stack exchange,提问作者the-bass

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.21 03:47:43