【Pytorch深度学习代码错误案例】参数和优化器没有对应导致参数无法更新错误案例(附查看梯度和参数的方法)
·
问题描述
采用深度学习和强化学习方法训练一个能够下五子棋的AI,深度学习的训练的网络是强化学习将会用到的策略网络。在深度学习训练部分出现了有梯度但是参数不更新的现象。
原因分析:
经过问题排查,发现是在装载一个新的模型训练的时候,用的优化器绑定的却是旧的模型参数。导致每次计算出的梯度其实是在更新旧模型参数,并非新装载的模型。
下面将有问题的部分代码列出来。
class DeepLearning:
"""深度学习训练策略网络
"""
def __init__(self):
# ROOT为当前文件所在根路径, policy.pth是最优的模型参数
self.policy = self.load_model(ROOT / "model/policy/policy.pth")
self.optimizer = torch.optim.Adam(self.policy.parameters(), lr=1E-2)
# 其他属性
# ...
def load_model(self, policy_path):
"""根据路径装载模型
"""
# ...
def train(self, data_name):
"""训练模型
"""
# 在训练时装载成最新的训练参数, 就是这里出了问题
self.policy = self.load_model(ROOT / "model/policy/last.pth")
# 其他参数
# ...
# 开始训练
# ...
# 更新参数
self.optimizer.step()
记__init__()函数中的self.policy为P1,train()函数中的self.policy为P2,__init__()函数中的self.optimizer为O1。
就可以将问题描述为P1和O1是相互对应的,但是P2却没有对应的O2,在train()函数中调用的self.optimizer.step()其实更新的是P1而非P2,这就导致了虽然计算出了梯度,但是却没有让参数更新的问题。
解决方案:
解决方案很简单。可以直接删掉train()函数中的装载模型语句,直接在__init__()当中就装载好last.pth,即删除P2;也可以在train()函数中添加一个优化器,绑定的参数是新的模型参数,即添加O2。
第一种修改方案修改后的代码:
class DeepLearning:
"""深度学习训练策略网络
"""
def __init__(self):
# ROOT为当前文件所在根路径, policy.pth是最优的模型参数
self.policy = self.load_model(ROOT / "model/policy/last.pth")
self.optimizer = torch.optim.Adam(self.policy.parameters(), lr=1E-2)
# 其他属性
# ...
def load_model(self, policy_path):
"""根据路径装载模型
"""
# ...
def train(self, data_name):
"""训练模型
"""
# 其他参数
# ...
# 开始训练
# ...
# 更新参数
self.optimizer.step()
第二种修改方案修改后的代码:
class DeepLearning:
"""深度学习训练策略网络
"""
def __init__(self):
# ROOT为当前文件所在根路径, policy.pth是最优的模型参数
self.policy = self.load_model(ROOT / "model/policy/policy.pth")
self.optimizer = torch.optim.Adam(self.policy.parameters(), lr=1E-2)
# 其他属性
# ...
def load_model(self, policy_path):
"""根据路径装载模型
"""
# ...
def train(self, data_name):
"""训练模型
"""
# 在训练时装载成最新的训练参数, 就是这里出了问题
self.policy = self.load_model(ROOT / "model/policy/last.pth")
# 新的优化器
self.optimizer = torch.optim.Adam(self.policy.parameters(), lr=1E-2)
# 其他参数
# ...
# 开始训练
# ...
# 更新参数
self.optimizer.step()
附录
查看梯度和参数的方法
查看参数和梯度的方法
# 模型参数字典, 键是网络每层的名字, 值是对应参数张量
dict(model.named_parameters())
# 查看参数字典都有哪些名字
dict(model.named_parameters()).keys()
# 查找到名字后, 就可以查看参数和对应梯度了
dict(model.named_parameters())[name]
dict(model.named_parameters())[name].grad()
另一种查看参数的方法
# 查看模型的所有参数, 这将返回一个OrderDict, 键是网络每层的名字, 值是对应参数张量
model.state_dict()
# 查看所有名字
model.state_dict().keys()
# 查找到想要看的参数对应的名字后查找对应参数
model.state_dict()[name]
DAMO开发者矩阵,由阿里巴巴达摩院和中国互联网协会联合发起,致力于探讨最前沿的技术趋势与应用成果,搭建高质量的交流与分享平台,推动技术创新与产业应用链接,围绕“人工智能与新型计算”构建开放共享的开发者生态。
更多推荐


所有评论(0)