mirror of
https://github.com/wassname/ray.git
synced 2026-09-09 11:32:43 +08:00
[RLlib] Trajectory view API example script (enhancements and tf2 support). (#13786)
This commit is contained in:
@@ -3,6 +3,8 @@ from ray.rllib.models.torch.misc import SlimFC
|
||||
from ray.rllib.models.torch.torch_modelv2 import TorchModelV2
|
||||
from ray.rllib.policy.view_requirement import ViewRequirement
|
||||
from ray.rllib.utils.framework import try_import_tf, try_import_torch
|
||||
from ray.rllib.utils.tf_ops import one_hot
|
||||
from ray.rllib.utils.torch_ops import one_hot as torch_one_hot
|
||||
|
||||
tf1, tf, tfv = try_import_tf()
|
||||
torch, nn = try_import_torch()
|
||||
@@ -28,27 +30,42 @@ class FrameStackingCartPoleModel(TFModelV2):
|
||||
|
||||
# Construct actual (very simple) FC model.
|
||||
assert len(obs_space.shape) == 1
|
||||
input_ = tf.keras.layers.Input(
|
||||
obs = tf.keras.layers.Input(
|
||||
shape=(self.num_frames, obs_space.shape[0]))
|
||||
reshaped = tf.keras.layers.Reshape(
|
||||
[obs_space.shape[0] * self.num_frames])(input_)
|
||||
layer1 = tf.keras.layers.Dense(64, activation=tf.nn.relu)(reshaped)
|
||||
out = tf.keras.layers.Dense(self.num_outputs)(layer1)
|
||||
obs_reshaped = tf.keras.layers.Reshape(
|
||||
[obs_space.shape[0] * self.num_frames])(obs)
|
||||
rewards = tf.keras.layers.Input(shape=(self.num_frames))
|
||||
rewards_reshaped = tf.keras.layers.Reshape([self.num_frames])(rewards)
|
||||
actions = tf.keras.layers.Input(
|
||||
shape=(self.num_frames, self.action_space.n))
|
||||
actions_reshaped = tf.keras.layers.Reshape(
|
||||
[action_space.n * self.num_frames])(actions)
|
||||
input_ = tf.keras.layers.Concatenate(axis=-1)(
|
||||
[obs_reshaped, actions_reshaped, rewards_reshaped])
|
||||
layer1 = tf.keras.layers.Dense(256, activation=tf.nn.relu)(input_)
|
||||
layer2 = tf.keras.layers.Dense(256, activation=tf.nn.relu)(layer1)
|
||||
out = tf.keras.layers.Dense(self.num_outputs)(layer2)
|
||||
values = tf.keras.layers.Dense(1)(layer1)
|
||||
self.base_model = tf.keras.models.Model([input_], [out, values])
|
||||
|
||||
self.base_model = tf.keras.models.Model([obs, actions, rewards],
|
||||
[out, values])
|
||||
self._last_value = None
|
||||
|
||||
self.view_requirements["prev_n_obs"] = ViewRequirement(
|
||||
data_col="obs",
|
||||
shift="-{}:0".format(num_frames - 1),
|
||||
space=obs_space)
|
||||
self.view_requirements["prev_rewards"] = ViewRequirement(
|
||||
data_col="rewards", shift=-1)
|
||||
self.view_requirements["prev_n_rewards"] = ViewRequirement(
|
||||
data_col="rewards", shift="-{}:-1".format(self.num_frames))
|
||||
self.view_requirements["prev_n_actions"] = ViewRequirement(
|
||||
data_col="actions",
|
||||
shift="-{}:-1".format(self.num_frames),
|
||||
space=self.action_space)
|
||||
|
||||
def forward(self, input_dict, states, seq_lens):
|
||||
obs = input_dict["prev_n_obs"]
|
||||
out, self._last_value = self.base_model(obs)
|
||||
obs = tf.cast(input_dict["prev_n_obs"], tf.float32)
|
||||
rewards = tf.cast(input_dict["prev_n_rewards"], tf.float32)
|
||||
actions = one_hot(input_dict["prev_n_actions"], self.action_space)
|
||||
out, self._last_value = self.base_model([obs, actions, rewards])
|
||||
return out, []
|
||||
|
||||
def value_function(self):
|
||||
@@ -77,13 +94,13 @@ class TorchFrameStackingCartPoleModel(TorchModelV2, nn.Module):
|
||||
|
||||
# Construct actual (very simple) FC model.
|
||||
assert len(obs_space.shape) == 1
|
||||
in_size = self.num_frames * (obs_space.shape[0] + action_space.n + 1)
|
||||
self.layer1 = SlimFC(
|
||||
in_size=obs_space.shape[0] * self.num_frames,
|
||||
out_size=64,
|
||||
activation_fn="relu")
|
||||
in_size=in_size, out_size=256, activation_fn="relu")
|
||||
self.layer2 = SlimFC(in_size=256, out_size=256, activation_fn="relu")
|
||||
self.out = SlimFC(
|
||||
in_size=64, out_size=self.num_outputs, activation_fn="linear")
|
||||
self.values = SlimFC(in_size=64, out_size=1, activation_fn="linear")
|
||||
in_size=256, out_size=self.num_outputs, activation_fn="linear")
|
||||
self.values = SlimFC(in_size=256, out_size=1, activation_fn="linear")
|
||||
|
||||
self._last_value = None
|
||||
|
||||
@@ -91,14 +108,26 @@ class TorchFrameStackingCartPoleModel(TorchModelV2, nn.Module):
|
||||
data_col="obs",
|
||||
shift="-{}:0".format(num_frames - 1),
|
||||
space=obs_space)
|
||||
self.view_requirements["prev_rewards"] = ViewRequirement(
|
||||
data_col="rewards", shift=-1)
|
||||
self.view_requirements["prev_n_rewards"] = ViewRequirement(
|
||||
data_col="rewards", shift="-{}:-1".format(self.num_frames))
|
||||
self.view_requirements["prev_n_actions"] = ViewRequirement(
|
||||
data_col="actions",
|
||||
shift="-{}:-1".format(self.num_frames),
|
||||
space=self.action_space)
|
||||
|
||||
def forward(self, input_dict, states, seq_lens):
|
||||
obs = input_dict["prev_n_obs"]
|
||||
obs = torch.reshape(obs,
|
||||
[-1, self.obs_space.shape[0] * self.num_frames])
|
||||
features = self.layer1(obs)
|
||||
rewards = torch.reshape(input_dict["prev_n_rewards"],
|
||||
[-1, self.num_frames])
|
||||
actions = torch_one_hot(input_dict["prev_n_actions"],
|
||||
self.action_space)
|
||||
actions = torch.reshape(actions,
|
||||
[-1, self.num_frames * actions.shape[-1]])
|
||||
input_ = torch.cat([obs, actions, rewards], dim=-1)
|
||||
features = self.layer1(input_)
|
||||
features = self.layer2(features)
|
||||
out = self.out(features)
|
||||
self._last_value = self.values(features)
|
||||
return out, []
|
||||
|
||||
@@ -2,6 +2,7 @@ import argparse
|
||||
|
||||
import ray
|
||||
from ray import tune
|
||||
from ray.rllib.examples.env.stateless_cartpole import StatelessCartPole
|
||||
from ray.rllib.examples.models.trajectory_view_utilizing_models import \
|
||||
FrameStackingCartPoleModel, TorchFrameStackingCartPoleModel
|
||||
from ray.rllib.models.catalog import ModelCatalog
|
||||
@@ -16,7 +17,7 @@ parser.add_argument(
|
||||
"--framework", choices=["tf2", "tf", "tfe", "torch"], default="tf")
|
||||
parser.add_argument("--as-test", action="store_true")
|
||||
parser.add_argument("--stop-iters", type=int, default=50)
|
||||
parser.add_argument("--stop-timesteps", type=int, default=100000)
|
||||
parser.add_argument("--stop-timesteps", type=int, default=200000)
|
||||
parser.add_argument("--stop-reward", type=float, default=150.0)
|
||||
|
||||
if __name__ == "__main__":
|
||||
@@ -26,13 +27,14 @@ if __name__ == "__main__":
|
||||
ModelCatalog.register_custom_model(
|
||||
"frame_stack_model", FrameStackingCartPoleModel
|
||||
if args.framework != "torch" else TorchFrameStackingCartPoleModel)
|
||||
tune.register_env("stateless_cartpole", lambda c: StatelessCartPole())
|
||||
|
||||
config = {
|
||||
"env": "CartPole-v0",
|
||||
"env": "stateless_cartpole",
|
||||
"model": {
|
||||
"custom_model": "frame_stack_model",
|
||||
"custom_model_config": {
|
||||
"num_frames": 4,
|
||||
"num_frames": 16,
|
||||
}
|
||||
},
|
||||
"framework": args.framework,
|
||||
|
||||
Reference in New Issue
Block a user