mirror of
https://github.com/wassname/ray.git
synced 2026-07-25 13:30:52 +08:00
* Fix typo * Fix A3C PyTorch agent initialization `registry` needs to be passed as an argument or else the `super` init will fail.
79 lines
2.9 KiB
Python
79 lines
2.9 KiB
Python
from __future__ import absolute_import
|
|
from __future__ import division
|
|
from __future__ import print_function
|
|
|
|
import torch
|
|
from torch.autograd import Variable
|
|
import torch.nn.functional as F
|
|
|
|
from ray.rllib.a3c.torchpolicy import TorchPolicy
|
|
from ray.rllib.models.pytorch.misc import var_to_np, convert_batch
|
|
from ray.rllib.models.catalog import ModelCatalog
|
|
|
|
|
|
class SharedTorchPolicy(TorchPolicy):
|
|
"""Assumes nonrecurrent."""
|
|
|
|
other_output = ["vf_preds"]
|
|
is_recurrent = False
|
|
|
|
def __init__(self, registry, ob_space, ac_space, config, **kwargs):
|
|
super(SharedTorchPolicy, self).__init__(
|
|
registry, ob_space, ac_space, config, **kwargs)
|
|
|
|
def _setup_graph(self, ob_space, ac_space):
|
|
_, self.logit_dim = ModelCatalog.get_action_dist(ac_space)
|
|
self._model = ModelCatalog.get_torch_model(
|
|
self.registry, ob_space, self.logit_dim, self.config["model"])
|
|
self.optimizer = torch.optim.Adam(
|
|
self._model.parameters(), lr=self.config["lr"])
|
|
|
|
def compute(self, ob, *args):
|
|
"""Should take in a SINGLE ob"""
|
|
with self.lock:
|
|
ob = Variable(torch.from_numpy(ob).float().unsqueeze(0))
|
|
logits, values = self._model(ob)
|
|
samples = self._model.probs(logits).multinomial().squeeze()
|
|
values = values.squeeze(0)
|
|
return var_to_np(samples), {"vf_preds": var_to_np(values)}
|
|
|
|
def compute_logits(self, ob, *args):
|
|
with self.lock:
|
|
ob = Variable(torch.from_numpy(ob).float().unsqueeze(0))
|
|
res = self._model.hidden_layers(ob)
|
|
return var_to_np(self._model.logits(res))
|
|
|
|
def value(self, ob, *args):
|
|
with self.lock:
|
|
ob = Variable(torch.from_numpy(ob).float().unsqueeze(0))
|
|
res = self._model.hidden_layers(ob)
|
|
res = self._model.value_branch(res)
|
|
res = res.squeeze(0)
|
|
return var_to_np(res)
|
|
|
|
def _evaluate(self, obs, actions):
|
|
"""Passes in multiple obs."""
|
|
logits, values = self._model(obs)
|
|
log_probs = F.log_softmax(logits)
|
|
probs = self._model.probs(logits)
|
|
action_log_probs = log_probs.gather(1, actions.view(-1, 1))
|
|
entropy = -(log_probs * probs).sum(-1).sum()
|
|
return values, action_log_probs, entropy
|
|
|
|
def _backward(self, batch):
|
|
"""Loss is encoded in here. Defining a new loss function
|
|
would start by rewriting this function"""
|
|
|
|
states, acs, advs, rs, _ = convert_batch(batch)
|
|
values, ac_logprobs, entropy = self._evaluate(states, acs)
|
|
pi_err = -(advs * ac_logprobs).sum()
|
|
value_err = 0.5 * (values - rs).pow(2).sum()
|
|
|
|
self.optimizer.zero_grad()
|
|
overall_err = (pi_err +
|
|
value_err * self.config["vf_loss_coeff"] +
|
|
entropy * self.config["entropy_coeff"])
|
|
overall_err.backward()
|
|
torch.nn.utils.clip_grad_norm(
|
|
self._model.parameters(), self.config["grad_clip"])
|