强化学习DQN训练CartPole-v1遇输入形状不匹配ValueError求助
问题分析与解决
错误原因
- 缺少gym库导入:代码中使用
gym.make('CartPole-v1')但未导入gym模块,这是基础依赖缺失问题。 - 无效的观测调用:
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
相关产品推荐
相关产品推荐

