如何正确重塑数组?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
相关产品推荐
相关产品推荐

