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

强化学习DQN训练CartPole-v1遇输入形状不匹配ValueError求助

问题分析与解决

错误原因

  1. 缺少gym库导入:代码中使用gym.make('CartPole-v1')但未导入gym模块,这是基础依赖缺失问题。
  2. 无效的观测调用:print(env.observation())是错误写法——Gym环境没有observation()方法,正确获取初始观测的方式是env.reset()。这行错误代码可能导致环境内部状态异常,进而让训练时传入模型的观测形状变成(1,2),与模型期望的(1,4)不匹配。

修正后的代码

import random
import numpy as np
import gym  # 新增gym导入
from tensorflow.keras.models import Sequential
from tensorflow.keras.layers import Dense, Flatten
from tensorflow.keras.optimizers import Adam
from rl.agents import DQNAgent
from rl.policy import BoltzmannQPolicy
from rl.memory import SequentialMemory

def build_model(states, actions):
    model = Sequential()
    model.add(Flatten(input_shape=(1, states)))
    model.add(Dense(24, activation='relu'))
    model.add(Dense(24, activation='relu'))
    model.add(Dense(actions, activation='linear'))
    return model

def build_agent(model, actions):
    policy = BoltzmannQPolicy()
    memory = SequentialMemory(limit=50000, window_length=1)
    dqn = DQNAgent(model=model, memory=memory, policy=policy, 
                  nb_actions=actions, nb_steps_warmup=10, target_model_update=1e-2)
    return dqn

def main():
    env = gym.make('CartPole-v1')
    states = env.observation_space.shape[0]
    actions = env.action_space.n
    
    # 修正观测获取方式
    initial_obs = env.reset()
    print(f"初始观测形状: {initial_obs.shape}, 观测值: {initial_obs}")

    model = build_model(states, actions)
    model.summary()  # 可选:查看模型结构确认输入输出形状

    dqn = build_agent(model, actions)
    dqn.compile(Adam(learning_rate=1e-3), metrics=['mae'])
    dqn.fit(env, nb_steps=50000, visualize=False, verbose=1)

main()

额外说明

  • CartPole-v1的观测空间是4维(小车位置、小车速度、杆角度、杆角速度),因此states值为4,模型输入形状(1,4)完全符合要求。
  • 运行前请确保依赖齐全:pip install gym tensorflow keras-rl

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.17 14:41:04