From 81e74340916a449941debc1aba1c4c858ca905b1 Mon Sep 17 00:00:00 2001 From: Sven Mika Date: Wed, 10 Feb 2021 15:21:46 +0100 Subject: [PATCH] [RLlib] TFPolicy.export_model: Add timestep placeholder to model's signature, if needed. (#13988) --- rllib/policy/tf_policy.py | 5 +++++ rllib/tests/test_export.py | 8 ++++++++ 2 files changed, 13 insertions(+) diff --git a/rllib/policy/tf_policy.py b/rllib/policy/tf_policy.py index f16f3f72a..e71cd2b44 100644 --- a/rllib/policy/tf_policy.py +++ b/rllib/policy/tf_policy.py @@ -709,9 +709,14 @@ class TFPolicy(Policy): input_signature["prev_reward"] = \ tf1.saved_model.utils.build_tensor_info( self._prev_reward_input) + input_signature["is_training"] = \ tf1.saved_model.utils.build_tensor_info(self._is_training) + if self._timestep is not None: + input_signature["timestep"] = \ + tf1.saved_model.utils.build_tensor_info(self._timestep) + for state_input in self._state_inputs: input_signature[state_input.name] = \ tf1.saved_model.utils.build_tensor_info(state_input) diff --git a/rllib/tests/test_export.py b/rllib/tests/test_export.py index 711cc85b5..bb8bde8e1 100644 --- a/rllib/tests/test_export.py +++ b/rllib/tests/test_export.py @@ -6,8 +6,11 @@ import unittest import ray from ray.rllib.agents.registry import get_trainer_class +from ray.rllib.utils.framework import try_import_tf from ray.tune.trial import ExportFormat +tf1, tf, tfv = try_import_tf() + CONFIGS = { "A3C": { "explore": False, @@ -105,6 +108,11 @@ def export_test(alg_name, failures): or not valid_tf_checkpoint(os.path.join(export_dir, ExportFormat.CHECKPOINT)): failures.append(alg_name) + + # Test loading the exported model. + model = tf.saved_model.load(os.path.join(export_dir, ExportFormat.MODEL)) + assert model + shutil.rmtree(export_dir)