wrap models with DistributionalQModel when running DQN (#6258)

* wrap models with DistributionalQModel when running DQN

* wrap only for tensorflow models

* Update custom_keras_model.py
This commit is contained in:
Ameer Haj Ali
2019-11-25 00:11:24 -08:00
committed by Eric Liang
parent 7917bbef78
commit 71316fa8d0
2 changed files with 14 additions and 8 deletions
+10 -3
View File
@@ -13,12 +13,14 @@ from ray.rllib.models.tf.misc import normc_initializer
from ray.rllib.models.tf.tf_modelv2 import TFModelV2
from ray.rllib.agents.dqn.distributional_q_model import DistributionalQModel
from ray.rllib.utils import try_import_tf
from ray.rllib.models.tf.visionnet_v2 import VisionNetwork as MyVisionNetwork
tf = try_import_tf()
parser = argparse.ArgumentParser()
parser.add_argument("--run", type=str, default="DQN") # Try PG, PPO, DQN
parser.add_argument("--stop", type=int, default=200)
parser.add_argument("--use_vision_network", action="store_true")
class MyKerasModel(TFModelV2):
@@ -90,13 +92,18 @@ class MyKerasQModel(DistributionalQModel):
if __name__ == "__main__":
ray.init()
args = parser.parse_args()
ModelCatalog.register_custom_model("keras_model", MyKerasModel)
ModelCatalog.register_custom_model("keras_q_model", MyKerasQModel)
ModelCatalog.register_custom_model(
"keras_model", MyVisionNetwork
if args.use_vision_network else MyKerasModel)
ModelCatalog.register_custom_model(
"keras_q_model", MyVisionNetwork
if args.use_vision_network else MyKerasQModel)
tune.run(
args.run,
stop={"episode_reward_mean": args.stop},
config={
"env": "CartPole-v0",
"env": "BreakoutNoFrameskip-v4"
if args.use_vision_network else "CartPole-v0",
"num_gpus": 0,
"model": {
"custom_model": "keras_q_model"
+4 -5
View File
@@ -257,12 +257,11 @@ class ModelCatalog(object):
model_cls = _global_registry.get(RLLIB_MODEL,
model_config["custom_model"])
if issubclass(model_cls, ModelV2):
if model_interface and not issubclass(model_cls,
model_interface):
raise ValueError("The given model must subclass",
model_interface)
if framework == "tf":
logger.info("Wrapping {} as {}".format(
model_cls, model_interface))
model_cls = ModelCatalog._wrap_if_needed(
model_cls, model_interface)
created = set()
# Track and warn if vars were created but not registered