diff --git a/rllib/BUILD b/rllib/BUILD index cc048db33..199cc5ad9 100644 --- a/rllib/BUILD +++ b/rllib/BUILD @@ -1089,13 +1089,6 @@ py_test( srcs = ["models/tests/test_distributions.py"] ) -py_test( - name = "test_models", - tags = ["models"], - size = "small", - srcs = ["models/tests/test_models.py"] -) - py_test( name = "test_preprocessors", tags = ["models"], diff --git a/rllib/agents/ddpg/ddpg_tf_model.py b/rllib/agents/ddpg/ddpg_tf_model.py index a5b3a0ecb..cddc38006 100644 --- a/rllib/agents/ddpg/ddpg_tf_model.py +++ b/rllib/agents/ddpg/ddpg_tf_model.py @@ -89,6 +89,7 @@ class DDPGTFModel(TFModelV2): actor_out = tf.keras.layers.Lambda(lambda_)(actor_out) self.policy_model = tf.keras.Model(self.model_out, actor_out) + self.register_variables(self.policy_model.variables) # Build the Q-model(s). self.actions_input = tf.keras.layers.Input( @@ -115,10 +116,12 @@ class DDPGTFModel(TFModelV2): return q_net self.q_model = build_q_net("q", self.model_out, self.actions_input) + self.register_variables(self.q_model.variables) if twin_q: self.twin_q_model = build_q_net("twin_q", self.model_out, self.actions_input) + self.register_variables(self.twin_q_model.variables) else: self.twin_q_model = None diff --git a/rllib/agents/dqn/distributional_q_tf_model.py b/rllib/agents/dqn/distributional_q_tf_model.py index d921f7b78..ed8064999 100644 --- a/rllib/agents/dqn/distributional_q_tf_model.py +++ b/rllib/agents/dqn/distributional_q_tf_model.py @@ -162,11 +162,13 @@ class DistributionalQTFModel(TFModelV2): q_out = build_action_value(name + "/action_value/", self.model_out) self.q_value_head = tf.keras.Model(self.model_out, q_out) + self.register_variables(self.q_value_head.variables) if dueling: state_out = build_state_score(name + "/state_value/", self.model_out) self.state_value_head = tf.keras.Model(self.model_out, state_out) + self.register_variables(self.state_value_head.variables) def get_q_value_distributions(self, model_out: TensorType) -> List[TensorType]: diff --git a/rllib/agents/sac/sac_tf_model.py b/rllib/agents/sac/sac_tf_model.py index 4c890385f..af6d1b539 100644 --- a/rllib/agents/sac/sac_tf_model.py +++ b/rllib/agents/sac/sac_tf_model.py @@ -91,6 +91,8 @@ class SACTFModel(TFModelV2): ]) self.shift_and_log_scale_diag = self.action_model(self.model_out) + self.register_variables(self.action_model.variables) + self.actions_input = None if not self.discrete: self.actions_input = tf.keras.layers.Input( @@ -121,10 +123,12 @@ class SACTFModel(TFModelV2): return q_net self.q_net = build_q_net("q", self.model_out, self.actions_input) + self.register_variables(self.q_net.variables) if twin_q: self.twin_q_net = build_q_net("twin_q", self.model_out, self.actions_input) + self.register_variables(self.twin_q_net.variables) else: self.twin_q_net = None @@ -143,6 +147,8 @@ class SACTFModel(TFModelV2): target_entropy = -np.prod(action_space.shape) self.target_entropy = target_entropy + self.register_variables([self.log_alpha]) + def get_q_values(self, model_out: TensorType, actions: Optional[TensorType] = None) -> TensorType: diff --git a/rllib/examples/batch_norm_model.py b/rllib/examples/batch_norm_model.py index 1a2a604ff..29f6dbe73 100644 --- a/rllib/examples/batch_norm_model.py +++ b/rllib/examples/batch_norm_model.py @@ -33,7 +33,6 @@ if __name__ == "__main__": "model": { "custom_model": "bn_model", }, - "lr": 0.0003, # Use GPUs iff `RLLIB_NUM_GPUS` env var set to > 0. "num_gpus": int(os.environ.get("RLLIB_NUM_GPUS", "0")), "num_workers": 0, @@ -46,7 +45,7 @@ if __name__ == "__main__": "episode_reward_mean": args.stop_reward, } - results = tune.run(args.run, stop=stop, config=config, verbose=2) + results = tune.run(args.run, stop=stop, config=config, verbose=1) if args.as_test: check_learning_achieved(results, args.stop_reward) diff --git a/rllib/examples/custom_env.py b/rllib/examples/custom_env.py index b9b060e43..25f5e61e6 100644 --- a/rllib/examples/custom_env.py +++ b/rllib/examples/custom_env.py @@ -71,6 +71,7 @@ class CustomModel(TFModelV2): model_config, name) self.model = FullyConnectedNetwork(obs_space, action_space, num_outputs, model_config, name) + self.register_variables(self.model.variables()) def forward(self, input_dict, state, seq_lens): return self.model.forward(input_dict, state, seq_lens) diff --git a/rllib/examples/custom_keras_model.py b/rllib/examples/custom_keras_model.py index 03ccb0859..a277ccd94 100644 --- a/rllib/examples/custom_keras_model.py +++ b/rllib/examples/custom_keras_model.py @@ -48,6 +48,7 @@ class MyKerasModel(TFModelV2): activation=None, kernel_initializer=normc_initializer(0.01))(layer_1) self.base_model = tf.keras.Model(self.inputs, [layer_out, value_out]) + self.register_variables(self.base_model.variables) def forward(self, input_dict, state, seq_lens): model_out, self._value_out = self.base_model(input_dict["obs"]) @@ -83,6 +84,7 @@ class MyKerasQModel(DistributionalQTFModel): activation=tf.nn.relu, kernel_initializer=normc_initializer(1.0))(layer_1) self.base_model = tf.keras.Model(self.inputs, layer_out) + self.register_variables(self.base_model.variables) # Implement the core forward method def forward(self, input_dict, state, seq_lens): diff --git a/rllib/examples/models/autoregressive_action_model.py b/rllib/examples/models/autoregressive_action_model.py index cb8af9ddc..5602f9b52 100644 --- a/rllib/examples/models/autoregressive_action_model.py +++ b/rllib/examples/models/autoregressive_action_model.py @@ -69,12 +69,14 @@ class AutoregressiveActionModel(TFModelV2): # Base layers self.base_model = tf.keras.Model(obs_input, [context, value_out]) + self.register_variables(self.base_model.variables) self.base_model.summary() # Autoregressive action sampler self.action_model = tf.keras.Model([ctx_input, a1_input], [a1_logits, a2_logits]) self.action_model.summary() + self.register_variables(self.action_model.variables) def forward(self, input_dict, state, seq_lens): context, self._value_out = self.base_model(input_dict["obs"]) diff --git a/rllib/examples/models/batch_norm_model.py b/rllib/examples/models/batch_norm_model.py index 0ad833c89..7d77ebc07 100644 --- a/rllib/examples/models/batch_norm_model.py +++ b/rllib/examples/models/batch_norm_model.py @@ -123,6 +123,7 @@ class KerasBatchNormModel(TFModelV2): self.base_model = tf.keras.models.Model( inputs=[inputs, is_training], outputs=[output, value_out]) + self.register_variables(self.base_model.variables) @override(ModelV2) def forward(self, input_dict, state, seq_lens): diff --git a/rllib/examples/models/centralized_critic_models.py b/rllib/examples/models/centralized_critic_models.py index 7f6e370bf..23f1e8b92 100644 --- a/rllib/examples/models/centralized_critic_models.py +++ b/rllib/examples/models/centralized_critic_models.py @@ -23,6 +23,7 @@ class CentralizedCriticModel(TFModelV2): # Base of the model self.model = FullyConnectedNetwork(obs_space, action_space, num_outputs, model_config, name) + self.register_variables(self.model.variables()) # Central VF maps (obs, opp_obs, opp_act) -> vf_pred obs = tf.keras.layers.Input(shape=(6, ), name="obs") @@ -36,6 +37,7 @@ class CentralizedCriticModel(TFModelV2): 1, activation=None, name="c_vf_out")(central_vf_dense) self.central_vf = tf.keras.Model( inputs=[obs, opp_obs, opp_act], outputs=central_vf_out) + self.register_variables(self.central_vf.variables) @override(ModelV2) def forward(self, input_dict, state, seq_lens): @@ -77,9 +79,11 @@ class YetAnotherCentralizedCriticModel(TFModelV2): num_outputs, model_config, name + "_action") + self.register_variables(self.action_model.variables()) self.value_model = FullyConnectedNetwork(obs_space, action_space, 1, model_config, name + "_vf") + self.register_variables(self.value_model.variables()) def forward(self, input_dict, state, seq_lens): self._value_out, _ = self.value_model({ diff --git a/rllib/examples/models/cnn_plus_fc_concat_model.py b/rllib/examples/models/cnn_plus_fc_concat_model.py index 6f8e3d85e..e8cae19dd 100644 --- a/rllib/examples/models/cnn_plus_fc_concat_model.py +++ b/rllib/examples/models/cnn_plus_fc_concat_model.py @@ -53,6 +53,7 @@ class CNNPlusFCConcatModel(TFModelV2): name="cnn_{}".format(i)) concat_size += cnn.num_outputs self.cnns[i] = cnn + self.register_variables(cnn.variables()) # Discrete inputs -> One-hot encode. elif isinstance(component, Discrete): concat_size += component.n @@ -81,6 +82,7 @@ class CNNPlusFCConcatModel(TFModelV2): kernel_initializer=normc_initializer(0.01))(concat_layer) self.logits_and_value_model = tf.keras.models.Model( concat_layer, [logits_layer, value_layer]) + self.register_variables(self.logits_and_value_model.variables) else: self.num_outputs = concat_size diff --git a/rllib/examples/models/custom_loss_model.py b/rllib/examples/models/custom_loss_model.py index 5e0d7b7c2..e7c19cb9e 100644 --- a/rllib/examples/models/custom_loss_model.py +++ b/rllib/examples/models/custom_loss_model.py @@ -29,6 +29,7 @@ class CustomLossModel(TFModelV2): num_outputs, model_config, name="fcnet") + self.register_variables(self.fcnet.variables()) @override(ModelV2) def forward(self, input_dict, state, seq_lens): diff --git a/rllib/examples/models/eager_model.py b/rllib/examples/models/eager_model.py index 3b3b190b9..a20236711 100644 --- a/rllib/examples/models/eager_model.py +++ b/rllib/examples/models/eager_model.py @@ -40,6 +40,7 @@ class EagerModel(TFModelV2): out = tf.keras.layers.Lambda(lambda_)(out) self.base_model = tf.keras.models.Model(inputs, [out, value_out]) + self.register_variables(self.base_model.variables) @override(ModelV2) def forward(self, input_dict, state, seq_lens): diff --git a/rllib/examples/models/mobilenet_v2_with_lstm_models.py b/rllib/examples/models/mobilenet_v2_with_lstm_models.py index c97778570..8afdbf188 100644 --- a/rllib/examples/models/mobilenet_v2_with_lstm_models.py +++ b/rllib/examples/models/mobilenet_v2_with_lstm_models.py @@ -68,6 +68,7 @@ class MobileV2PlusRNNModel(RecurrentNetwork): self.rnn_model = tf.keras.Model( inputs=[inputs, seq_in, state_in_h, state_in_c], outputs=[logits, values, state_h, state_c]) + self.register_variables(self.rnn_model.variables) self.rnn_model.summary() @override(RecurrentNetwork) diff --git a/rllib/examples/models/parametric_actions_model.py b/rllib/examples/models/parametric_actions_model.py index cbeda9645..abffbcafd 100644 --- a/rllib/examples/models/parametric_actions_model.py +++ b/rllib/examples/models/parametric_actions_model.py @@ -36,6 +36,7 @@ class ParametricActionsModel(DistributionalQTFModel): self.action_embed_model = FullyConnectedNetwork( Box(-1, 1, shape=true_obs_shape), action_space, action_embed_size, model_config, name + "_action_embed") + self.register_variables(self.action_embed_model.variables()) def forward(self, input_dict, state, seq_lens): # Extract the available actions tensor from the observation. diff --git a/rllib/examples/models/rnn_model.py b/rllib/examples/models/rnn_model.py index 5a661f8a6..84729ca3d 100644 --- a/rllib/examples/models/rnn_model.py +++ b/rllib/examples/models/rnn_model.py @@ -54,6 +54,7 @@ class RNNModel(RecurrentNetwork): self.rnn_model = tf.keras.Model( inputs=[input_layer, seq_in, state_in_h, state_in_c], outputs=[logits, values, state_h, state_c]) + self.register_variables(self.rnn_model.variables) self.rnn_model.summary() @override(RecurrentNetwork) diff --git a/rllib/examples/models/rnn_spy_model.py b/rllib/examples/models/rnn_spy_model.py index 19009d7ea..1b1d95f1e 100644 --- a/rllib/examples/models/rnn_spy_model.py +++ b/rllib/examples/models/rnn_spy_model.py @@ -107,6 +107,7 @@ class RNNSpyModel(RecurrentNetwork): [inputs, seq_lens, state_in_h, state_in_c], [logits, value_out, state_out_h, state_out_c]) self.base_model.summary() + self.register_variables(self.base_model.variables) @override(RecurrentNetwork) def forward_rnn(self, inputs, state, seq_lens): diff --git a/rllib/examples/models/shared_weights_model.py b/rllib/examples/models/shared_weights_model.py index 8f3bf7ea4..4e4c6c32a 100644 --- a/rllib/examples/models/shared_weights_model.py +++ b/rllib/examples/models/shared_weights_model.py @@ -39,6 +39,7 @@ class TF2SharedWeightsModel(TFModelV2): vf = tf.keras.layers.Dense( units=1, activation=None, name="value_out")(last_layer) self.base_model = tf.keras.models.Model(inputs, [output, vf]) + self.register_variables(self.base_model.variables) @override(ModelV2) def forward(self, input_dict, state, seq_lens): @@ -79,6 +80,7 @@ class SharedWeightsModel1(TFModelV2): vf = tf.keras.layers.Dense( units=1, activation=None, name="value_out")(last_layer) self.base_model = tf.keras.models.Model(inputs, [output, vf]) + self.register_variables(self.base_model.variables) @override(ModelV2) def forward(self, input_dict, state, seq_lens): @@ -112,6 +114,7 @@ class SharedWeightsModel2(TFModelV2): vf = tf.keras.layers.Dense( units=1, activation=None, name="value_out")(last_layer) self.base_model = tf.keras.models.Model(inputs, [output, vf]) + self.register_variables(self.base_model.variables) @override(ModelV2) def forward(self, input_dict, state, seq_lens): diff --git a/rllib/examples/models/simple_rpg_model.py b/rllib/examples/models/simple_rpg_model.py index 8615d8c30..072437e5a 100644 --- a/rllib/examples/models/simple_rpg_model.py +++ b/rllib/examples/models/simple_rpg_model.py @@ -47,6 +47,7 @@ class CustomTFRPGModel(TFModelV2): name) self.model = TFFCNet(obs_space, action_space, num_outputs, model_config, name) + self.register_variables(self.model.variables()) def forward(self, input_dict, state, seq_lens): # The unpacked input tensors, where M=MAX_PLAYERS, N=MAX_ITEMS: diff --git a/rllib/examples/models/trajectory_view_utilizing_models.py b/rllib/examples/models/trajectory_view_utilizing_models.py index 41f53d872..2360be025 100644 --- a/rllib/examples/models/trajectory_view_utilizing_models.py +++ b/rllib/examples/models/trajectory_view_utilizing_models.py @@ -36,6 +36,7 @@ class FrameStackingCartPoleModel(TFModelV2): out = tf.keras.layers.Dense(self.num_outputs)(layer1) values = tf.keras.layers.Dense(1)(layer1) self.base_model = tf.keras.models.Model([input_], [out, values]) + self.register_variables(self.base_model.variables) self._last_value = None diff --git a/rllib/models/catalog.py b/rllib/models/catalog.py index a6e7415d4..9638ed44b 100644 --- a/rllib/models/catalog.py +++ b/rllib/models/catalog.py @@ -378,9 +378,7 @@ class ModelCatalog: if model_config.get("use_lstm") else AttentionWrapper) model_cls._wrapped_forward = forward - # Obsolete: Track and warn if vars were created but not - # registered. Only still do this, if users do register their - # variables. If not (which they shouldn't), don't check here. + # Track and warn if vars were created but not registered. created = set() def track_var_creation(next_creator, **kw): @@ -409,27 +407,19 @@ class ModelCatalog: # Other error -> re-raise. else: raise e - - # User still registered TFModelV2's variables: Check, whether - # ok. - registered = set(instance.var_list) - if len(registered) > 0: - not_registered = set() - for var in created: - if var not in registered: - not_registered.add(var) - if not_registered: - raise ValueError( - "It looks like you are still using " - "`{}.register_variables()` to register your " - "model's weights. This is no longer required, but " - "if you are still calling this method at least " - "once, you must make sure to register all created " - "variables properly. The missing variables are {}," - " and you only registered {}. " - "Did you forget to call `register_variables()` on " - "some of the variables in question?".format( - instance, not_registered, registered)) + registered = set(instance.variables()) + not_registered = set() + for var in created: + if var not in registered: + not_registered.add(var) + if not_registered: + raise ValueError( + "It looks like variables {} were created as part " + "of {} but does not appear in model.variables() " + "({}). Did you forget to call " + "model.register_variables() on the variables in " + "question?".format(not_registered, instance, + registered)) elif framework == "torch": # Try wrapping custom model with LSTM/attention, if required. if model_config.get("use_lstm") or \ diff --git a/rllib/models/tests/test_distributions.py b/rllib/models/tests/test_distributions.py index 3dd14d0ae..987f76a56 100644 --- a/rllib/models/tests/test_distributions.py +++ b/rllib/models/tests/test_distributions.py @@ -94,7 +94,8 @@ class TestDistributions(unittest.TestCase): inputs = inputs_space.sample() - for fw, sess in framework_iterator(session=True): + for fw, sess in framework_iterator( + session=True, frameworks=("tf", "tf2", "torch")): # Create the correct distribution object. cls = JAXCategorical if fw == "jax" else Categorical if \ fw != "torch" else TorchCategorical @@ -217,7 +218,8 @@ class TestDistributions(unittest.TestCase): input_space = Box(-2.0, 2.0, shape=(2000, 10)) low, high = -2.0, 1.0 - for fw, sess in framework_iterator(session=True): + for fw, sess in framework_iterator( + frameworks=("torch", "tf", "tfe"), session=True): cls = SquashedGaussian if fw != "torch" else TorchSquashedGaussian # Do a stability test using extreme NN outputs to see whether @@ -308,7 +310,8 @@ class TestDistributions(unittest.TestCase): """Tests the DiagGaussian ActionDistribution for all frameworks.""" input_space = Box(-2.0, 1.0, shape=(2000, 10)) - for fw, sess in framework_iterator(session=True): + for fw, sess in framework_iterator( + frameworks=("torch", "tf", "tfe"), session=True): cls = DiagGaussian if fw != "torch" else TorchDiagGaussian # Do a stability test using extreme NN outputs to see whether diff --git a/rllib/models/tests/test_models.py b/rllib/models/tests/test_models.py deleted file mode 100644 index 424dea16c..000000000 --- a/rllib/models/tests/test_models.py +++ /dev/null @@ -1,59 +0,0 @@ -from gym.spaces import Box -import numpy as np -import unittest - -from ray.rllib.models.tf.tf_modelv2 import TFModelV2 -from ray.rllib.models.tf.fcnet import FullyConnectedNetwork -from ray.rllib.utils.framework import try_import_tf - -tf1, tf, tfv = try_import_tf() - - -class TestTFModel(TFModelV2): - def __init__(self, obs_space, action_space, num_outputs, model_config, - name): - super().__init__(obs_space, action_space, num_outputs, model_config, - name) - input_ = tf.keras.layers.Input(shape=(3, )) - output = tf.keras.layers.Dense(2)(input_) - # A keras model inside. - self.keras_model = tf.keras.models.Model([input_], [output]) - # A RLlib FullyConnectedNetwork (tf) inside (which is also a keras - # Model). - self.fc_net = FullyConnectedNetwork(obs_space, action_space, 3, {}, - "fc1") - - def forward(self, input_dict, state, seq_lens): - obs = input_dict["obs_flat"] - out1 = self.keras_model(obs) - out2, _ = self.fc_net({"obs": obs}) - return tf.concat([out1, out2], axis=-1), [] - - -class TestModels(unittest.TestCase): - """Tests ModelV2 classes and their modularization capabilities.""" - - def test_tf_modelv2(self): - obs_space = Box(-1.0, 1.0, (3, )) - action_space = Box(-1.0, 1.0, (2, )) - my_tf_model = TestTFModel(obs_space, action_space, 5, {}, - "my_tf_model") - # Call the model. - out, states = my_tf_model({"obs": np.array([obs_space.sample()])}) - self.assertTrue(out.shape == (1, 5)) - self.assertTrue(out.dtype == tf.float32) - self.assertTrue(states == []) - vars = my_tf_model.variables(as_dict=True) - self.assertTrue(len(vars) == 6) - self.assertTrue("keras_model.dense.kernel:0" in vars) - self.assertTrue("keras_model.dense.bias:0" in vars) - self.assertTrue("fc_net.base_model.fc_out.kernel:0" in vars) - self.assertTrue("fc_net.base_model.fc_out.bias:0" in vars) - self.assertTrue("fc_net.base_model.value_out.kernel:0" in vars) - self.assertTrue("fc_net.base_model.value_out.bias:0" in vars) - - -if __name__ == "__main__": - import pytest - import sys - sys.exit(pytest.main(["-v", __file__])) diff --git a/rllib/models/tf/attention_net.py b/rllib/models/tf/attention_net.py index fadd5ed89..4e79eb51a 100644 --- a/rllib/models/tf/attention_net.py +++ b/rllib/models/tf/attention_net.py @@ -117,6 +117,7 @@ class TrXLNet(RecurrentNetwork): name="logits")(E_out) self.base_model = tf.keras.models.Model([inputs], [logits]) + self.register_variables(self.base_model.variables) @override(RecurrentNetwork) def forward_rnn(self, inputs: TensorType, state: List[TensorType], @@ -286,6 +287,7 @@ class GTrXLNet(RecurrentNetwork): self.trxl_model = tf.keras.Model( inputs=[input_layer] + memory_ins, outputs=outs + memory_outs[:-1]) + self.register_variables(self.trxl_model.variables) self.trxl_model.summary() # __sphinx_doc_begin__ @@ -384,6 +386,7 @@ class AttentionWrapper(TFModelV2): position_wise_mlp_dim=cfg["attention_position_wise_mlp_dim"], init_gru_gate_bias=cfg["attention_init_gru_gate_bias"], ) + self.register_variables(self.gtrxl.variables()) # `self.num_outputs` right now is the number of nodes coming from the # attention net. @@ -396,9 +399,11 @@ class AttentionWrapper(TFModelV2): # values. out = tf.keras.layers.Dense(self.num_outputs, activation=None)(input_) self._logits_branch = tf.keras.models.Model([input_], [out]) + self.register_variables(self._logits_branch.variables) out = tf.keras.layers.Dense(1, activation=None)(input_) self._value_branch = tf.keras.models.Model([input_], [out]) + self.register_variables(self._value_branch.variables) self.view_requirements = self.gtrxl.view_requirements diff --git a/rllib/models/tf/fcnet.py b/rllib/models/tf/fcnet.py index eea01014d..e556741dd 100644 --- a/rllib/models/tf/fcnet.py +++ b/rllib/models/tf/fcnet.py @@ -33,6 +33,7 @@ class FullyConnectedNetwork(TFModelV2): num_outputs = num_outputs // 2 self.log_std_var = tf.Variable( [0.0] * num_outputs, dtype=tf.float32, name="log_std") + self.register_variables([self.log_std_var]) # We are using obs_flat, so take the flattened shape as input. inputs = tf.keras.layers.Input( @@ -114,6 +115,7 @@ class FullyConnectedNetwork(TFModelV2): self.base_model = tf.keras.Model( inputs, [(logits_out if logits_out is not None else last_layer), value_out]) + self.register_variables(self.base_model.variables) def forward(self, input_dict: Dict[str, TensorType], state: List[TensorType], diff --git a/rllib/models/tf/recurrent_net.py b/rllib/models/tf/recurrent_net.py index 8618cecc5..fa51b54d0 100644 --- a/rllib/models/tf/recurrent_net.py +++ b/rllib/models/tf/recurrent_net.py @@ -50,6 +50,7 @@ class RecurrentNetwork(TFModelV2): self.rnn_model = tf.keras.Model( inputs=[input_layer, seq_in, state_in_h, state_in_c], outputs=[output_layer, state_h, state_c]) + self.register_variables(self.rnn_model.variables) self.rnn_model.summary() """ @@ -178,6 +179,7 @@ class LSTMWrapper(RecurrentNetwork): self._rnn_model = tf.keras.Model( inputs=[input_layer, seq_in, state_in_h, state_in_c], outputs=[logits, values, state_h, state_c]) + self.register_variables(self._rnn_model.variables) self._rnn_model.summary() # Add prev-a/r to this model's view, if required. diff --git a/rllib/models/tf/tf_modelv2.py b/rllib/models/tf/tf_modelv2.py index 78e1e0276..09625781b 100644 --- a/rllib/models/tf/tf_modelv2.py +++ b/rllib/models/tf/tf_modelv2.py @@ -1,12 +1,9 @@ import contextlib import gym -import re from typing import List -from ray.util import log_once from ray.rllib.models.modelv2 import ModelV2 from ray.rllib.utils.annotations import override, PublicAPI -from ray.rllib.utils.deprecation import deprecation_warning from ray.rllib.utils.framework import try_import_tf from ray.rllib.utils.typing import ModelConfigDict, TensorType @@ -15,7 +12,7 @@ tf1, tf, tfv = try_import_tf() @PublicAPI class TFModelV2(ModelV2): - """TF version of ModelV2, which is always also a keras Model. + """TF version of ModelV2. Note that this class by itself is not a valid model unless you implement forward() in a subclass.""" @@ -36,18 +33,18 @@ class TFModelV2(ModelV2): value_layer = tf.keras.layers.Dense(...)(hidden_layer) self.base_model = tf.keras.Model( input_layer, [output_layer, value_layer]) + self.register_variables(self.base_model.variables) """ - super().__init__( + + ModelV2.__init__( + self, obs_space, action_space, num_outputs, model_config, name, framework="tf") - - # Deprecated: TFModelV2 now automatically track their variables. self.var_list = [] - if tf1.executing_eagerly(): self.graph = None else: @@ -68,41 +65,13 @@ class TFModelV2(ModelV2): def register_variables(self, variables: List[TensorType]) -> None: """Register the given list of variables with this model.""" - if log_once("deprecated_tfmodelv2_register_variables"): - deprecation_warning( - old="TFModelV2.register_variables", error=False) self.var_list.extend(variables) @override(ModelV2) def variables(self, as_dict: bool = False) -> List[TensorType]: if as_dict: - # Old way using `register_variables`. - if self.var_list: - return {v.name: v for v in self.var_list} - # New way: Automatically determine the var tree. - else: - ret = {} - for prop, value in self.__dict__.items(): - # Keras Model: key=k + "." + var-name (replace '/' by '.'). - if isinstance(value, tf.keras.models.Model): - for var in value.variables: - key = prop + "." + re.sub("/", ".", var.name) - ret[key] = var - # Other TFModelV2: Include its vars into ours. - elif isinstance(value, TFModelV2): - for key, var in value.variables(as_dict=True).items(): - ret[prop + "." + key] = var - # tf.Variable - elif isinstance(value, tf.Variable): - ret[prop] = value - return ret - - # Old way using `register_variables`. - if self.var_list: - return list(self.var_list) - # New way: Automatically determine the var tree. - else: - return list(self.variables(as_dict=True).values()) + return {v.name: v for v in self.var_list} + return list(self.var_list) @override(ModelV2) def trainable_variables(self, as_dict: bool = False) -> List[TensorType]: diff --git a/rllib/models/tf/visionnet.py b/rllib/models/tf/visionnet.py index 039ad4389..c2a8de5d2 100644 --- a/rllib/models/tf/visionnet.py +++ b/rllib/models/tf/visionnet.py @@ -140,6 +140,7 @@ class VisionNetwork(TFModelV2): lambda x: tf.squeeze(x, axis=[1, 2]))(last_layer) self.base_model = tf.keras.Model(inputs, [conv_out, value_out]) + self.register_variables(self.base_model.variables) def forward(self, input_dict: Dict[str, TensorType], state: List[TensorType], diff --git a/rllib/models/torch/torch_modelv2.py b/rllib/models/torch/torch_modelv2.py index f56cf9978..39bc336ed 100644 --- a/rllib/models/torch/torch_modelv2.py +++ b/rllib/models/torch/torch_modelv2.py @@ -11,7 +11,7 @@ _, nn = try_import_torch() @PublicAPI class TorchModelV2(ModelV2): - """Torch version of ModelV2, which is also always a torch.nn.Module. + """Torch version of ModelV2. Note that this class by itself is not a valid model unless you inherit from nn.Module and implement forward() in a subclass.""" diff --git a/rllib/tests/test_model_imports.py b/rllib/tests/test_model_imports.py index 405b96b90..b92f5d3a6 100644 --- a/rllib/tests/test_model_imports.py +++ b/rllib/tests/test_model_imports.py @@ -49,6 +49,8 @@ class MyKerasModel(TFModelV2): else: self.base_model = tf.keras.Model(self.inputs, layer_out) + self.register_variables(self.base_model.variables) + def forward(self, input_dict, state, seq_lens): if self.model_config["vf_share_layers"]: model_out, self._value_out = self.base_model(input_dict["obs"]) diff --git a/rllib/tests/test_nested_observation_spaces.py b/rllib/tests/test_nested_observation_spaces.py index 1a10e8c71..736e69780 100644 --- a/rllib/tests/test_nested_observation_spaces.py +++ b/rllib/tests/test_nested_observation_spaces.py @@ -240,6 +240,7 @@ class DictSpyModel(TFModelV2): self.num_outputs = num_outputs or 64 out = tf.keras.layers.Dense(self.num_outputs)(input_) self._main_layer = tf.keras.models.Model([input_], [out]) + self.register_variables(self._main_layer.variables) def forward(self, input_dict, state, seq_lens): def spy(pos, front_cam, task): @@ -281,6 +282,7 @@ class TupleSpyModel(TFModelV2): self.num_outputs = num_outputs or 64 out = tf.keras.layers.Dense(self.num_outputs)(input_) self._main_layer = tf.keras.models.Model([input_], [out]) + self.register_variables(self._main_layer.variables) def forward(self, input_dict, state, seq_lens): def spy(pos, cam, task): diff --git a/rllib/utils/deprecation.py b/rllib/utils/deprecation.py index 430d0a66a..05788059b 100644 --- a/rllib/utils/deprecation.py +++ b/rllib/utils/deprecation.py @@ -1,5 +1,4 @@ import logging -from typing import Optional, Union logger = logging.getLogger(__name__) @@ -9,18 +8,15 @@ logger = logging.getLogger(__name__) DEPRECATED_VALUE = -1 -def deprecation_warning( - old: str, - new: Optional[str] = None, - error: Optional[Union[bool, Exception]] = None) -> None: - """Warns (via the `logger` object) or throws a deprecation warning/error. +def deprecation_warning(old, new=None, error=None): + """ + Logs (via the `logger` object) or throws a deprecation warning/error. Args: old (str): A description of the "thing" that is to be deprecated. new (Optional[str]): A description of the new "thing" that replaces it. - error (Optional[Union[bool, Exception]]): Whether or which exception to - throw. If True, throw ValueError. If False, just warn. - If Exception, throw that Exception. + error (Optional[Union[bool,Exception]]): Whether or which exception to + throw. If True, throw ValueError. """ msg = "`{}` has been deprecated.{}".format( old, (" Use `{}` instead.".format(new) if new else "")) diff --git a/rllib/utils/exploration/curiosity.py b/rllib/utils/exploration/curiosity.py index 45b75ac22..ec91c53d3 100644 --- a/rllib/utils/exploration/curiosity.py +++ b/rllib/utils/exploration/curiosity.py @@ -208,6 +208,7 @@ class Curiosity(Exploration): self._curiosity_feature_net.base_model.variables + \ self._curiosity_inverse_fcnet.variables + \ self._curiosity_forward_fcnet.variables + self.model.register_variables(self._optimizer_var_list) self._optimizer = tf1.train.AdamOptimizer(learning_rate=self.lr) # Create placeholders and initialize the loss. if self.framework == "tf":