Keras-rl2与TensorFlow版本兼容问题求助
Keras-rl2与TensorFlow兼容问题排查与解决
问题背景
使用Keras-rl2时遇到TensorFlow/Keras版本兼容问题,先后尝试最新版TensorFlow 2.16.1+Keras 3、降级至TensorFlow 2.13.0+Keras 2.13.0,均出现导入报错。
情况1:TensorFlow 2.16.1 + Keras 3
执行代码:
from rl.agents import DQNAgent
报错信息:
ModuleNotFoundError Traceback (most recent call last) Cell In[37], line 1 ----> 1 import rl.agents 3 print("RL agents library version:", rl.agents.__version__) File D:\Anaconda\Lib\site-packages\rl\agents\__init__.py:2 1 from .dqn import DQNAgent, NAFAgent, ContinuousDQNAgent ----> 2 from .ddpg import DDPGAgent 3 from .cem import CEMAgent 4 from .sarsa import SarsaAgent, SARSAAgent File D:\Anaconda\Lib\site-packages\rl\agents\dqn.py:8 5 from tensorflow.keras.layers import Lambda, Input, Layer, Dense 7 from rl.core import Agent ----> 8 from rl.policy import EpsGreedyQPolicy, GreedyQPolicy 9 from rl.util import * 12 def mean_q(y_true, y_pred): File D:\Anaconda\Lib\site-packages\rl\core.py:8 4 import numpy as np 5 from tensorflow.keras.callbacks import History 7 from rl.callbacks import ( ----> 8 CallbackList, 9 TestLogger, 10 TrainEpisodeLogger, 11 TrainIntervalLogger, 12 Visualizer 13 ) 16 class Agent: 17 """Abstract base class for all implemented agents. 18 19 Each agent interacts with the environment (as defined by the `Env` class) by first observing the (...) 37 processor (`Processor` instance): See [Processor](#processor) for details. 38 """ File D:\Anaconda\Lib\site-packages\rl\callbacks.py:12 9 from tensorflow.python.keras.callbacks import Callback as KerasCallback, CallbackList as KerasCallbackList 10 from tensorflow.python.keras.utils.generic_utils import Progbar ---> 12 class Callback(KerasCallback): 13 def _set_env(self, env): 14 self.env = env ModuleNotFoundError: No module named 'keras.utils.generic_utils'
情况2:TensorFlow 2.13.0 + Keras 2.13.0
执行同样导入代码,报错:
--------------------------------------------------------------------------- ImportError Traceback (most recent call last) Cell In[18], line 1 ----> 1 from rl.agents.dqn import DQNAgent File D:\Anaconda\envs\AI\Lib\site-packages\rl\agents\__init__.py:1 ----> 1 from .dqn import DQNAgent, NAFAgent, ContinuousDQNAgent 2 from .ddpg import DDPGAgent 3 from .cem import CEMAgent File D:\Anaconda\envs\AI\Lib\site-packages\rl\agents\dqn.py:7 4 from tensorflow.keras.models import Model 5 from tensorflow.keras.layers import Lambda, Input, Layer, Dense ----> 7 from rl.core import Agent 8 from rl.policy import EpsGreedyQPolicy, GreedyQPolicy 9 from rl.util import * File D:\Anaconda\envs\AI\Lib\site-packages\rl\core.py:7 4 import numpy as np 5 from tensorflow.keras.callbacks import History ----> 7 from rl.callbacks import ( 8 CallbackList, 9 TestLogger, 10 TrainEpisodeLogger, 11 TrainIntervalLogger, 12 Visualizer 13 ) 16 class Agent: 17 """Abstract base class for all implemented agents. 18 19 Each agent interacts with the environment (as defined by the `Env` class) by first observing the (...) 37 processor (`Processor` instance): See [Processor](#processor) for details. 38 """ File D:\Anaconda\envs\AI\Lib\site-packages\rl\callbacks.py:8 6 import numpy as np 7 import tensorflow as tf ----> 8 from tensorflow.keras import __version__ as KERAS_VERSION 9 from tensorflow.python.keras.callbacks import Callback as KerasCallback, CallbackList as KerasCallbackList 10 from tensorflow.python.keras.utils.generic_utils import Progbar ImportError: cannot import name '__version__' from 'tensorflow.keras' (D:\Anaconda\envs\AI\Lib\site-packages\keras\api\_v2\keras\__init__.py)
原因分析
Keras-rl2官方维护已停滞(最后更新于2021年),无法适配TensorFlow 2.10+版本的API变更:
- Keras 3/TensorFlow 2.16中,Keras模块结构大幅调整,
keras.utils.generic_utils这类旧路径被移除或重构,导致Keras-rl2导入代码失效。 - TensorFlow 2.13中,
tensorflow.keras.__version__的导入方式被修改,Keras-rl2的callbacks.py直接导入该属性的代码不再生效。
解决办法
方案1:使用社区维护分支
安装适配新版本TensorFlow的Keras-rl2社区分支:
pip install git+https://github.com/wau/keras-rl2.git
该分支修复了大部分新版本TensorFlow的API兼容问题。
方案2:锁定兼容版本
若不想使用第三方分支,可锁定到Keras-rl2官方支持的版本组合:
- TensorFlow 2.9.0 + Keras 2.9.0
- 安装命令:
pip install tensorflow==2.9.0 keras==2.9.0 keras-rl2
方案3:手动修改源码
针对报错点逐个修改Keras-rl2源码:
- 解决
ModuleNotFoundError: No module named 'keras.utils.generic_utils':
打开rl/callbacks.py,将from tensorflow.python.keras.utils.generic_utils import Progbar改为from keras.src.utils.generic_utils import Progbar - 解决
ImportError: cannot import name '__version__' from 'tensorflow.keras':
打开rl/callbacks.py,将from tensorflow.keras import __version__ as KERAS_VERSION改为import keras; KERAS_VERSION = keras.__version__
内容的提问来源于stack exchange,提问作者PACE
相关产品推荐
相关产品推荐

