PyTorch A3C优化器报错:state_steps需为单例张量列表
解决PyTorch A3C中
opt.step()的RuntimeError问题 问题背景
基于MorvanZhou的A3C代码模板实现强化学习程序,运行到opt.step()时触发错误:
RuntimeError: API has changed, "state_steps" argument must contain a list of singleton tensors
程序基于torch.multiprocessing运行文本模拟,内存存储和模拟逻辑正常,但无法完成全局网络参数更新。
问题原因
新版PyTorch(2.0及以上版本)对优化器的状态存储做了变更:要求优化器state中的step必须是单例张量(singleton tensor),而非原来的整数类型。你自定义的SharedAdam类中,初始化时将state['step']设为整数0,不符合新版API要求。
修复方案
修改SharedAdam类的初始化逻辑,将state['step']从整数改为单例张量,并同步设置共享内存(适配多进程场景):
修改后的SharedAdam代码:
import torch class SharedAdam(torch.optim.Adam): def __init__(self, params, lr=1e-3, betas=(0.9, 0.99), eps=1e-8, weight_decay=0): super(SharedAdam, self).__init__(params, lr=lr, betas=betas, eps=eps, weight_decay=weight_decay) # State initialization for group in self.param_groups: for p in group['params']: state = self.state[p] # 将整数step替换为单例张量 state['step'] = torch.tensor(0, dtype=torch.long) state['exp_avg'] = torch.zeros_like(p.data) state['exp_avg_sq'] = torch.zeros_like(p.data) # 共享内存(多进程场景必需) state['step'].share_memory_() state['exp_avg'].share_memory_() state['exp_avg_sq'].share_memory_()
关键修改点说明
- 将
state['step'] = 0替换为state['step'] = torch.tensor(0, dtype=torch.long):用单例张量替代整数,符合新版PyTorch优化器API要求 - 添加
state['step'].share_memory_():因为A3C基于多进程运行,需保证step张量在进程间共享,避免状态不一致
内容的提问来源于stack exchange,提问作者Sebas Cubids
相关产品推荐
相关产品推荐

