Files
ray/python/ray/rllib/keras_policy.py
T
Eric LiangandRichard Liaw 02583a8598 [rllib] Rename PolicyGraph => Policy, move from evaluation/ to policy/ (#4819)
This implements some of the renames proposed in #4813
We leave behind backwards-compatibility aliases for *PolicyGraph and SampleBatch.
2019-05-20 16:46:05 -07:00

66 lines
2.0 KiB
Python

from __future__ import absolute_import
from __future__ import division
from __future__ import print_function
import numpy as np
from ray.rllib.policy.policy import Policy
def _sample(probs):
return [np.random.choice(len(pr), p=pr) for pr in probs]
class KerasPolicy(Policy):
"""Initialize the Keras Policy.
This is a Policy used for models with actor and critics.
Note: This class is built for specific usage of Actor-Critic models,
and is less general compared to TFPolicy and TorchPolicies.
Args:
observation_space (gym.Space): Observation space of the policy.
action_space (gym.Space): Action space of the policy.
config (dict): Policy-specific configuration data.
actor (Model): A model that holds the policy.
critic (Model): A model that holds the value function.
"""
def __init__(self,
observation_space,
action_space,
config,
actor=None,
critic=None):
Policy.__init__(self, observation_space, action_space, config)
self.actor = actor
self.critic = critic
self.models = [self.actor, self.critic]
def compute_actions(self, obs, *args, **kwargs):
state = np.array(obs)
policy = self.actor.predict(state)
value = self.critic.predict(state)
return _sample(policy), [], {"vf_preds": value.flatten()}
def learn_on_batch(self, batch, *args):
self.actor.fit(
batch["obs"],
batch["adv_targets"],
epochs=1,
verbose=0,
steps_per_epoch=20)
self.critic.fit(
batch["obs"],
batch["value_targets"],
epochs=1,
verbose=0,
steps_per_epoch=20)
return {}
def get_weights(self):
return [model.get_weights() for model in self.models]
def set_weights(self, weights):
return [model.set_weights(w) for model, w in zip(self.models, weights)]