From 71316fa8d02056a4d4cb365373191caaf13fded3 Mon Sep 17 00:00:00 2001 From: Ameer Haj Ali Date: Mon, 25 Nov 2019 00:11:24 -0800 Subject: [PATCH] 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 --- rllib/examples/custom_keras_model.py | 13 ++++++++++--- rllib/models/catalog.py | 9 ++++----- 2 files changed, 14 insertions(+), 8 deletions(-) diff --git a/rllib/examples/custom_keras_model.py b/rllib/examples/custom_keras_model.py index 741cd6d25..465af255d 100644 --- a/rllib/examples/custom_keras_model.py +++ b/rllib/examples/custom_keras_model.py @@ -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" diff --git a/rllib/models/catalog.py b/rllib/models/catalog.py index f79a016d4..a8adb05c8 100644 --- a/rllib/models/catalog.py +++ b/rllib/models/catalog.py @@ -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