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

导入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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.18 03:07:24