导入DQNAgent致Keras模型compile报错,如何指定使用training.py?
解决keras-rl2导入DQNAgent后模型compile不支持jit_compile的问题
问题概述
导入rl.agents.dqn.DQNAgent后,Keras Sequential模型调用compile时传入jit_compile=True会触发错误,原因是keras-rl2强制使用了training_v1.py中的旧版compile方法,该方法不支持jit_compile参数;而原生Keras的training.py中的compile方法是支持该参数的。
复现代码
from keras.models import Sequential from keras.layers import Dense, Flatten, LeakyReLU from keras.regularizers import l1 from rl.agents.dqn import DQNAgent reg = l1(1e-5) relu_alpha = 0.01 model = Sequential() model.add(Flatten(input_shape=[128,10,20])) model.add(Dense(16, kernel_regularizer = reg)) model.add(LeakyReLU(alpha = relu_alpha)) model.add(Dense(3, activation = "linear", kernel_regularizer = reg)) model.compile(loss='mse', jit_compile=True)
错误信息
File "C:\Users\lbosc\anaconda3\envs\ml_env3\lib\site-packages\keras\engine\training_v1.py", line 306, in compile raise TypeError( TypeError: Invalid keyword argument(s) in `compile`: {'jit_compile'}
问题原因
rl.agents.dqn.DQNAgent在导入过程中,会将Keras模型的compile方法替换为training_v1.py中的旧版本,该版本未实现jit_compile参数支持;而未导入DQNAgent时,模型使用的是training.py中的原生compile方法,支持该参数。
解决方案
方案1:先编译模型,再导入DQNAgent
利用DQNAgent在导入时才替换compile方法的特性,先完成模型的编译操作,再导入DQNAgent:
from keras.models import Sequential from keras.layers import Dense, Flatten, LeakyReLU from keras.regularizers import l1 reg = l1(1e-5) relu_alpha = 0.01 # 构建模型 model = Sequential() model.add(Flatten(input_shape=[128,10,20])) model.add(Dense(16, kernel_regularizer = reg)) model.add(LeakyReLU(alpha = relu_alpha)) model.add(Dense(3, activation = "linear", kernel_regularizer = reg)) # 先使用原生compile方法完成编译 model.compile(loss='mse', jit_compile=True) # 之后再导入DQNAgent,不会影响已编译的模型 from rl.agents.dqn import DQNAgent # 后续正常初始化并使用DQNAgent # dqn_agent = DQNAgent(model=model, ...)
方案2:手动替换回原生compile方法
如果必须先导入DQNAgent,可以在导入后手动将模型的compile方法替换为Keras原生版本:
from keras.models import Sequential from keras.layers import Dense, Flatten, LeakyReLU from keras.regularizers import l1 from rl.agents.dqn import DQNAgent # 导入原生compile方法所在的Model类 from keras.engine.training import Model as NativeModel reg = l1(1e-5) relu_alpha = 0.01 # 构建模型 model = Sequential() model.add(Flatten(input_shape=[128,10,20])) model.add(Dense(16, kernel_regularizer = reg)) model.add(LeakyReLU(alpha = relu_alpha)) model.add(Dense(3, activation = "linear", kernel_regularizer = reg)) # 将模型的compile方法替换为原生版本 model.compile = NativeModel.compile.__get__(model, Sequential) # 现在可以正常使用jit_compile参数编译 model.compile(loss='mse', jit_compile=True)
方案3:修改keras-rl2源码(不推荐)
找到keras-rl2库中导入training_v1的位置(通常在rl/agents/dqn.py或相关核心模块中),将类似from keras.engine.training_v1 import Model的语句替换为from keras.engine.training import Model,强制库使用新版compile方法。
注意:该方法会影响全局环境中的keras-rl2使用,后续升级库时修改会被覆盖,仅在无其他方案时考虑。
内容的提问来源于stack exchange,提问作者Luca
相关产品推荐
相关产品推荐

