mirror of
https://github.com/wassname/ray.git
synced 2026-08-12 12:20:11 +08:00
[RLlib] Examples folder restructuring (Model examples; final part). (#8278)
- This PR completes any previously missing PyTorch Model counterparts to TFModels in examples/models. - It also makes sure, all example scripts in the rllib/examples folder are tested for both frameworks and learn the given task (this is often currently not checked) using a --as-test flag in connection with a --stop-reward.
This commit is contained in:
@@ -4,70 +4,20 @@ import random
|
||||
import ray
|
||||
from ray import tune
|
||||
from ray.rllib.agents.trainer_template import build_trainer
|
||||
from ray.rllib.examples.models.eager_model import EagerModel
|
||||
from ray.rllib.models import ModelCatalog
|
||||
from ray.rllib.models.modelv2 import ModelV2
|
||||
from ray.rllib.models.tf.fcnet_v2 import FullyConnectedNetwork
|
||||
from ray.rllib.models.tf.tf_modelv2 import TFModelV2
|
||||
from ray.rllib.policy.sample_batch import SampleBatch
|
||||
from ray.rllib.policy.tf_policy_template import build_tf_policy
|
||||
from ray.rllib.utils import try_import_tf
|
||||
from ray.rllib.utils.annotations import override
|
||||
from ray.rllib.utils.framework import try_import_tf
|
||||
from ray.rllib.utils.test_utils import check_learning_achieved
|
||||
|
||||
tf = try_import_tf()
|
||||
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--iters", type=int, default=200)
|
||||
|
||||
|
||||
class EagerModel(TFModelV2):
|
||||
"""Example of using embedded eager execution in a custom model.
|
||||
|
||||
This shows how to use tf.py_function() to execute a snippet of TF code
|
||||
in eager mode. Here the `self.forward_eager` method just prints out
|
||||
the intermediate tensor for debug purposes, but you can in general
|
||||
perform any TF eager operation in tf.py_function().
|
||||
"""
|
||||
|
||||
def __init__(self, observation_space, action_space, num_outputs,
|
||||
model_config, name):
|
||||
super().__init__(observation_space, action_space, num_outputs,
|
||||
model_config, name)
|
||||
|
||||
inputs = tf.keras.layers.Input(shape=observation_space.shape)
|
||||
self.fcnet = FullyConnectedNetwork(
|
||||
obs_space=self.obs_space,
|
||||
action_space=self.action_space,
|
||||
num_outputs=self.num_outputs,
|
||||
model_config=self.model_config,
|
||||
name="fc1")
|
||||
out, value_out = self.fcnet.base_model(inputs)
|
||||
|
||||
def lambda_(x):
|
||||
eager_out = tf.py_function(self.forward_eager, [x], tf.float32)
|
||||
with tf.control_dependencies([eager_out]):
|
||||
eager_out.set_shape(x.shape)
|
||||
return eager_out
|
||||
|
||||
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):
|
||||
out, self._value_out = self.base_model(input_dict["obs"], state,
|
||||
seq_lens)
|
||||
return out, []
|
||||
|
||||
@override(ModelV2)
|
||||
def value_function(self):
|
||||
return tf.reshape(self._value_out, [-1])
|
||||
|
||||
def forward_eager(self, feature_layer):
|
||||
assert tf.executing_eagerly()
|
||||
if random.random() > 0.99:
|
||||
print("Eagerly printing the feature layer mean value",
|
||||
tf.reduce_mean(feature_layer))
|
||||
return feature_layer
|
||||
parser.add_argument("--stop-iters", type=int, default=200)
|
||||
parser.add_argument("--stop-timesteps", type=int, default=100000)
|
||||
parser.add_argument("--stop-reward", type=float, default=150)
|
||||
parser.add_argument("--as-test", action="store_true")
|
||||
|
||||
|
||||
def policy_gradient_loss(policy, model, dist_class, train_batch):
|
||||
@@ -119,5 +69,14 @@ if __name__ == "__main__":
|
||||
"custom_model": "eager_model"
|
||||
},
|
||||
}
|
||||
stop = {
|
||||
"timesteps_total": args.stop_timesteps,
|
||||
"training_iteration": args.stop_iters,
|
||||
"episode_reward_mean": args.stop_reward,
|
||||
}
|
||||
|
||||
tune.run(MyTrainer, stop={"training_iteration": args.iters}, config=config)
|
||||
results = tune.run(MyTrainer, stop=stop, config=config)
|
||||
|
||||
if args.as_test:
|
||||
check_learning_achieved(results, args.stop_reward)
|
||||
ray.shutdown()
|
||||
|
||||
Reference in New Issue
Block a user