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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.23 13:47:18