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

如何正确重塑数组?CartPole-v1强化学习代码报错求助

强化学习CartPole-v1环境数组重塑报错解决

我参考旧教程构建DQN强化学习智能体,已经修复了gym step函数的兼容问题,但卡在了numpy数组重塑的报错上。查过Numpy文档、视频评论和同类教程都没解决。

我的代码(DQNAgent代码未包含,需要可提供)

env = gym.make('CartPole-v1')
state_size = env.observation_space.shape[0]
action_size = env.action_space.n

batch_size = 32
n_episodes = 1001

# 存储模型
output_dir = './model_output/cartpole'
if not os.path.exists(output_dir):
    os.makedirs(output_dir)
agent = DQNAgent(state_size, action_size)

done = False
for e in range(n_episodes):
    state = env.reset()  # 重置环境状态
    print(state)
    print([1, state_size])
    state = np.reshape(state, [1, state_size])

    for time in range(5000):  # 游戏最长持续5000步
        # env.render()

        action = agent.act(state)  # 获取动作(0-左,1-右)
        # 前期动作随机,后期会逐渐偏向最优策略
        next_state, reward, done, truncated, _ = env.step(action)
        reward = reward if not done else -10  # 对失败动作施加惩罚
        next_state = np.reshape(next_state, [1, state_size])

        agent.remember(state, action, reward, next_state, done)
        state = next_state
        if done:  # 打印当前回合表现
            print(f"回合: {e}/{n_episodes}, 得分: {time}, 探索率: {agent.epsilon:.2}")
            break

        if len(agent.memory) > batch_size:  # 当经验池足够大时,开始训练
            agent.replay(batch_size)

报错信息

回溯(最近的调用最后):
  文件 "C:\Users\josep\VS Code Files\Python Projects\RL_Gym\main.py", 第120行, in <module>
    state = np.reshape(state, [1, state_size])
  文件 "<__array_function__ internals>", 第200行, in reshape
  文件 "C:\Users\josep\VS Code Files\Python Projects\RL_Gym\venv\lib\site-packages\numpy\core\fromnumeric.py", 第298行, in reshape
    return _wrapfunc(a, 'reshape', newshape, order=order)
  文件 "C:\Users\josep\VS Code Files\Python Projects\RL_Gym\venv\lib\site-packages\numpy\core\fromnumeric.py", 第54行, in _wrapfunc
    return _wrapit(obj, method, *args, **kwds)
  文件 "C:\Users\josep\VS Code Files\Python Projects\RL_Gym\venv\lib\site-packages\numpy\core\fromnumeric.py", 第43行, in _wrapit
    result = getattr(asarray(obj), method)(*args, **kwds)
ValueError: 用序列设置数组元素。请求的数组在1维后具有非均匀形状。检测到的形状是(2,) + 非均匀部分。

解决方法

这个报错的核心原因是Gym 0.26版本之后,env.reset()的返回值格式改变:旧版本只返回状态数组,新版本返回的是(状态数组, 额外信息字典)的元组。你打印的state其实是包含两个元素的元组,不是单纯的状态数组,所以reshape时才会报“非均匀形状”的错误。

修复只要改一行代码:

# 原来的代码
state = env.reset()
# 修改为(用下划线接收不需要的额外信息)
state, _ = env.reset()

修改后state就会是正常的4维状态数组,执行np.reshape(state, [1, state_size])就不会报错了。另外你代码里step函数的返回值处理(next_state, reward, done, truncated, _)是正确的,符合新版本Gym的要求,保持即可。

优质强化学习学习资源推荐

  • 斯坦福CS234强化学习课程:理论与实践结合,从基础到进阶的经典教程,配套作业和代码实现
  • 《Reinforcement Learning: An Introduction》(Sutton&Barto):强化学习领域经典教材,系统梳理核心理论体系
  • OpenAI Spinning Up:面向初学者的实践导向教程,包含基础算法的极简实现,适合快速上手
  • Hugging Face RL4Docs:提供多种主流强化学习算法的代码示例和详细文档,适合参考工程实现

内容的提问来源于stack exchange,提问作者Jslater2001

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.31 09:31:05