[rllib] Support agent.get_action in multiagent (#2543)

* support get action on policy id

* comment

* grammar fixes

* Update rllib-algorithms.rst
This commit is contained in:
Eric Liang
2018-08-02 13:35:53 -07:00
committed by GitHub
parent d2ebe4d9a3
commit f7ec292360
3 changed files with 54 additions and 7 deletions
+12 -5
View File
@@ -212,16 +212,23 @@ class Agent(Trainable):
raise NotImplementedError
def compute_action(self, observation, state=None):
"""Computes an action using the current trained policy."""
def compute_action(self, observation, state=None, policy_id="default"):
"""Computes an action for the specified policy.
Arguments:
observation (obj): observation from the environment.
state (list): RNN hidden state, if any.
policy_id (str): policy to query (only applies to multi-agent).
"""
if state is None:
state = []
obs = self.local_evaluator.filters["default"](
filtered_obs = self.local_evaluator.filters[policy_id](
observation, update=False)
return self.local_evaluator.for_policy(
lambda p: p.compute_single_action(obs, state, is_training=False)[0]
)
lambda p: p.compute_single_action(
filtered_obs, state, is_training=False)[0],
policy_id=policy_id)
def get_weights(self, policies=None):
"""Return a dictionary of policy ids to weights.