使用TensorFlow 2.10.1与Keras-RL 0.4.2运行RL代码遇符号输入输出错误
问题描述
正在学习强化学习,从基础模型进阶到RL后,跟着旧教程复现代码时遇到兼容性错误。使用的库版本:TensorFlow 2.10.1、Keras-RL 0.4.2。
复现代码
import numpy as np import random import pygame import gym from rl.memory import SequentialMemory from rl.policy import BoltzmannQPolicy from rl.agents.dqn import DQNAgent from keras.layers import Dense, Flatten import tensorflow as tf env = gym.make("CartPole-v1", render_mode="rgb_array") states = env.observation_space.shape[0] actions = env.action_space.n def build_model(states, actions): model = tf.keras.Sequential() model.add(Dense(24, activation='relu', input_shape=(states,))) model.add(Dense(24, activation='relu')) model.add(Dense(actions, activation='linear')) model.build((None, states)) return model def buildAgent(model, actions): policy = BoltzmannQPolicy() memory = SequentialMemory(limit=50000, window_length=1) dqn = DQNAgent(model, memory=memory, policy=policy, nb_actions=actions, nb_steps_warmup=10, target_model_update=1e-2) return dqn model = build_model(states, actions) DQN = buildAgent(model, actions) DQN.compile(tf.keras.optimizers.Adam(learning_rate=1e-3), metrics=['mae']) DQN.fit(env, nb_steps=50000, visualize=False, verbose=1) scores = DQN.test(env, nb_episodes=100, visualize=True) print(np.mean(scores.history['episode_reward'])) model.save('model.h5')
报错信息
TypeError: Keras symbolic inputs/outputs do not implement `__len__`. You may be trying to pass Keras symbolic inputs/outputs to a TF API that does not register dispatching, preventing Keras from automatically converting the API call to a lambda layer in the Functional Model. This error will also get raised if you try asserting a symbolic input/output directly. 堆栈追踪: --------------------------------------------------------------------------- TypeError Traceback (most recent call last) ~\\AppData\\Local\\Temp\\ipykernel_9048\\467332075.py in <module> 30 model = build_model(states, actions) 31 ---> 32 DQN = buildAgent(model, actions) 33 34 DQN.compile(tf.keras.optimizers.Adam(learning_rate=1e-3), metrics=['mae']) ~\\AppData\\Local\\Temp\\ipykernel_9048\\467332075.py in buildAgent(model, actions) 25 memory = SequentialMemory(limit=50000, window_length=1) 26 dqn = DQNAgent(model, memory=memory, policy=policy, nb_actions=actions, nb_steps_warmup=10, ---> 27 target_model_update=1e-2) 28 return dqn 29 c:\\Users\\Cyril\\miniconda3\\envs\\dsl\\lib\\site-packages\\rl\\agents\\dqn.py in __init__(self, model, policy, test_policy, enable_double_dqn, enable_dueling_network, dueling_type, *args, **kwargs) 106 107 # Validate (important) input. ---> 108 if hasattr(model.output, '__len__') and len(model.output) > 1: 109 raise ValueError('Model \"{}\" has more than one output. DQN expects a model that has a single output.'.format(model)) 110 if model.output._keras_shape != (None, self.nb_actions): c:\\Users\\Cyril\\miniconda3\\envs\\dsl\\lib\\site-packages\\keras\\engine\\keras_tensor.py in __len__(self) 243 def __len__(self): 244 raise TypeError( ---> 245 "Keras symbolic inputs/outputs do not " 246 "implement `__len__`. You may be " 247 "trying to pass Keras symbolic inputs/outputs "
解决方案
这个错误源于Keras-RL 0.4.2与TensorFlow 2.x的API不兼容:旧版Keras-RL是为TensorFlow 1.x设计的,而TF2.x中Keras模型的输出是KerasTensor对象,不再支持__len__方法,导致DQNAgent初始化时的输出数量判断逻辑失效。
以下是三种可行的解决办法:
方法1:更换为支持TF2.x的Keras-RL2
卸载旧版Keras-RL,安装适配TF2.x的分支:
pip uninstall keras-rl -y pip install keras-rl2
替换后代码无需大幅修改,直接运行即可,API基本兼容旧版。
方法2:修改Keras-RL源码的判断逻辑
找到环境中rl/agents/dqn.py文件(路径类似你的conda环境路径/lib/site-packages/rl/agents/dqn.py),将第108行的判断代码替换为:
if isinstance(model.output, (list, tuple)) and len(model.output) > 1:
原代码通过hasattr(model.output, '__len__')判断是否为多输出,但TF2.x的KerasTensor虽然有__len__方法却会抛出错误,改为直接判断类型即可避免问题。
方法3:调整模型构建方式
修改build_model函数,改用函数式API构建模型,避免生成触发错误的KerasTensor:
def build_model(states, actions): inputs = tf.keras.Input(shape=(states,)) x = Dense(24, activation='relu')(inputs) x = Dense(24, activation='relu')(x) outputs = Dense(actions, activation='linear')(x) model = tf.keras.Model(inputs=inputs, outputs=outputs) return model
这种方式构建的模型,输出结构能被旧版Keras-RL正确识别。
内容的提问来源于stack exchange,提问作者bebel
相关产品推荐
相关产品推荐

