-
Notifications
You must be signed in to change notification settings - Fork 7.8k
[rllib] example and docs on how to use parametric actions with DQN / PG algorithms #3384
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
Merged
Merged
Changes from 3 commits
Commits
Show all changes
22 commits
Select commit
Hold shift + click to select a range
95b2aa0
pa model
ericl 2c3ca54
docs
ericl 9633107
lint
ericl a95f0db
update
ericl caca39a
doc
ericl de5cd1b
large
ericl 26dcc8f
Update rllib-training.rst
ericl fd2058d
Update parametric_action_cartpole.py
ericl 0a50f72
doc dqn
ericl 3877d6a
doc
ericl b1dd888
Merge branch 'pa-model' of github.com:ericl/ray into pa-model
ericl 85d9fb8
stop at 50
ericl 6187de0
0.2 gpu
ericl b707bd7
masked
ericl ee79f73
improve docs
ericl 90b52e1
gym env
ericl 084e8c7
custom lstm doc
ericl 115977c
token
ericl 1ca2992
Merge remote-tracking branch 'upstream/master' into pa-model
ericl 78c765c
support dueling
ericl 06148b4
Update rllib-models.rst
ericl 02b886d
128
ericl File filter
Filter by extension
Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
There are no files selected for viewing
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
183 changes: 183 additions & 0 deletions
183
python/ray/rllib/examples/parametric_action_cartpole.py
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
| Original file line number | Diff line number | Diff line change |
|---|---|---|
| @@ -0,0 +1,183 @@ | ||
| """Example of handling variable length and/or parametric action spaces. | ||
|
|
||
| This is a toy example of the action-embedding based approach for handling large | ||
| discrete action spaces (potentially infinite in size), similar to how | ||
| OpenAI Five works: | ||
|
|
||
| https://neuro.cs.ut.ee/the-use-of-embeddings-in-openai-five/ | ||
|
|
||
| Note: this currently only works with RLlib's policy gradient style algorithms | ||
| i.e., PG / PPO and not DQN / DDPG. | ||
| """ | ||
|
|
||
| from __future__ import absolute_import | ||
| from __future__ import division | ||
| from __future__ import print_function | ||
|
|
||
| import argparse | ||
| import random | ||
| import numpy as np | ||
| import gym | ||
| from gym.spaces import Box, Discrete, Dict | ||
| import tensorflow as tf | ||
| import tensorflow.contrib.slim as slim | ||
|
|
||
| import ray | ||
| from ray.rllib.models import Model, ModelCatalog | ||
| from ray.rllib.models.misc import normc_initializer | ||
| from ray.tune import run_experiments | ||
| from ray.tune.registry import register_env | ||
|
|
||
| parser = argparse.ArgumentParser() | ||
| parser.add_argument("--num-iters", type=int, default=200) | ||
| parser.add_argument("--run", type=str, default="PPO") | ||
|
|
||
|
|
||
| class ParametricActionCartpole(gym.Env): | ||
| """Parametric action version of CartPole. | ||
|
|
||
| In this env there are only ever two valid actions, but we pretend there are | ||
| actually up to `max_avail_actions` actions that can be taken, and the two | ||
| valid actions are randomly hidden among this set. | ||
|
|
||
| At each step, we emit a dict of: | ||
| - the actual cart observation | ||
| - a mask of valid actions (e.g., [0, 1, 1, 0, 0, 1] for 3 / 6 active) | ||
| - the list of action embeddings (w/ zeroes for invalid actions) (e.g., | ||
| [[0, 0], | ||
| [-1.4414, 0.9071], | ||
| [-0.2322, -0.2569], | ||
| [0, 0], | ||
| [0, 0], | ||
| [0.7878, 1.2297]] for 3 active of 6 max) | ||
|
|
||
| In a real environment, the actions embeddings would be larger than two | ||
| units of course, and also there would be a variable number of valid actions | ||
| per step instead of always [LEFT, RIGHT]. | ||
| """ | ||
|
|
||
| def __init__(self, max_avail_actions): | ||
| # Use simple random 2-unit action embeddings for [LEFT, RIGHT] | ||
| self.left_action_embed = np.random.randn(2) | ||
| self.right_action_embed = np.random.randn(2) | ||
| self.action_space = Discrete(max_avail_actions) | ||
| self.wrapped = gym.make("CartPole-v0") | ||
| self.observation_space = Dict({ | ||
| "action_mask": Box(0, 1, shape=(max_avail_actions, )), | ||
| "avail_actions": Box(-1, 1, shape=(max_avail_actions, 2)), | ||
| "cart": self.wrapped.observation_space, | ||
| }) | ||
|
|
||
| def update_avail_actions(self): | ||
| self.action_assignments = [[0, 0]] * self.action_space.n | ||
| self.action_mask = [0] * self.action_space.n | ||
| self.left_idx, self.right_idx = random.sample( | ||
| range(self.action_space.n), 2) | ||
| self.action_assignments[self.left_idx] = self.left_action_embed | ||
| self.action_assignments[self.right_idx] = self.right_action_embed | ||
| self.action_mask[self.left_idx] = 1 | ||
| self.action_mask[self.right_idx] = 1 | ||
|
|
||
| def reset(self): | ||
| self.update_avail_actions() | ||
| return { | ||
| "action_mask": self.action_mask, | ||
| "avail_actions": self.action_assignments, | ||
| "cart": self.wrapped.reset(), | ||
| } | ||
|
|
||
| def step(self, action): | ||
| if action == self.left_idx: | ||
| actual_action = 0 | ||
| elif action == self.right_idx: | ||
| actual_action = 1 | ||
| else: | ||
| raise ValueError( | ||
| "Chosen action was not one of the non-zero action embeddings", | ||
| action, self.action_assignments, self.action_mask, | ||
| self.left_idx, self.right_idx) | ||
| orig_obs, rew, done, info = self.wrapped.step(actual_action) | ||
| self.update_avail_actions() | ||
| obs = { | ||
| "action_mask": self.action_mask, | ||
| "avail_actions": self.action_assignments, | ||
| "cart": orig_obs, | ||
| } | ||
| return obs, rew, done, info | ||
|
|
||
|
|
||
| class ParametricActionsModel(Model): | ||
| """Parametric action model that handles the dot product and masking. | ||
|
|
||
| This assumes the outputs are logits for a Categorical action dist.""" | ||
|
|
||
| def _build_layers_v2(self, input_dict, num_outputs, options): | ||
| # Extract the available actions tensor from the observation. | ||
| avail_actions = input_dict["obs"]["avail_actions"] | ||
| action_mask = input_dict["obs"]["action_mask"] | ||
| action_embed_size = avail_actions.shape[2].value | ||
| if num_outputs != avail_actions.shape[1].value: | ||
| raise ValueError( | ||
| "This model assumes num outputs is equal to max avail actions", | ||
| num_outputs, avail_actions) | ||
|
|
||
| # Standard FC net component. | ||
| last_layer = input_dict["obs"]["cart"] | ||
| hiddens = [256, 256] | ||
| for i, size in enumerate(hiddens): | ||
| label = "fc{}".format(i) | ||
| last_layer = slim.fully_connected( | ||
| last_layer, | ||
| size, | ||
| weights_initializer=normc_initializer(1.0), | ||
| activation_fn=tf.nn.tanh, | ||
| scope=label) | ||
| output = slim.fully_connected( | ||
| last_layer, | ||
| action_embed_size, | ||
| weights_initializer=normc_initializer(0.01), | ||
| activation_fn=None, | ||
| scope="fc_out") | ||
|
|
||
| # Expand the model output to [BATCH, 1, EMBED_SIZE]. Note that the | ||
| # avail actions tensor is of shape [BATCH, MAX_ACTIONS, EMBED_SIZE]. | ||
| intent_vector = tf.expand_dims(output, 1) | ||
|
|
||
| # Shape of logits is [BATCH, MAX_ACTIONS]. | ||
| action_logits = tf.reduce_sum(avail_actions * intent_vector, axis=2) | ||
|
|
||
| # Mask out invalid actions (use tf.float32.min for stability) | ||
| inf_mask = tf.maximum(tf.log(action_mask), tf.float32.min) | ||
| masked_logits = inf_mask + action_logits | ||
|
|
||
| return masked_logits, last_layer | ||
|
|
||
|
|
||
| if __name__ == "__main__": | ||
| args = parser.parse_args() | ||
| ray.init() | ||
|
|
||
| ModelCatalog.register_custom_model("pa_model", ParametricActionsModel) | ||
| register_env("pa_cartpole", lambda _: ParametricActionCartpole(10)) | ||
| if args.run == "PPO": | ||
| cfg = { | ||
| "observation_filter": "NoFilter", # don't filter the action list | ||
| "vf_share_layers": True, # don't create duplicate value model | ||
| } | ||
| else: | ||
| cfg = {} | ||
| run_experiments({ | ||
| "parametric_cartpole": { | ||
| "run": args.run, | ||
| "env": "pa_cartpole", | ||
| "stop": { | ||
| "training_iteration": args.num_iters | ||
| }, | ||
| "config": dict({ | ||
| "model": { | ||
| "custom_model": "pa_model", | ||
| }, | ||
| "num_workers": 0, | ||
| }, **cfg), | ||
| }, | ||
| }) |
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Add this suggestion to a batch that can be applied as a single commit.
This suggestion is invalid because no changes were made to the code.
Suggestions cannot be applied while the pull request is closed.
Suggestions cannot be applied while viewing a subset of changes.
Only one suggestion per line can be applied in a batch.
Add this suggestion to a batch that can be applied as a single commit.
Applying suggestions on deleted lines is not supported.
You must change the existing code in this line in order to create a valid suggestion.
Outdated suggestions cannot be applied.
This suggestion has been applied or marked resolved.
Suggestions cannot be applied from pending reviews.
Suggestions cannot be applied on multi-line comments.
Suggestions cannot be applied while the pull request is queued to merge.
Suggestion cannot be applied right now. Please check back later.
There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
oops this is kind of an important bug fix
There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
do you need to match the rest of the checks? https://github.com/openai/gym/blob/master/gym/spaces/dict_space.py#L34
There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
No, note that that is a check against the space spec, this is sorting the observation dict.