Repository navigation
Expand file tree
/
Copy pathdistribution.py
More file actions
163 lines (127 loc) · 5.63 KB
/
Copy pathdistribution.py
File metadata and controls
163 lines (127 loc) · 5.63 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
# Copyright 2024 The Brax Authors.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
# From https://github.com/google/brax/blob/main/brax/training/distribution.py
"""Probability distributions in JAX."""
import abc
import jax
import jax.numpy as jnp
class ParametricDistribution(abc.ABC):
"""Abstract class for parametric (action) distribution."""
def __init__(self, param_size, postprocessor, event_ndims, reparametrizable):
"""Abstract class for parametric (action) distribution.
Specifies how to transform distribution parameters (i.e. actor output)
into a distribution over actions.
Args:
param_size: size of the parameters for the distribution
postprocessor: bijector which is applied after sampling (in practice, it's
tanh or identity)
event_ndims: rank of the distribution sample (i.e. action)
reparametrizable: is the distribution reparametrizable
"""
self._param_size = param_size
self._postprocessor = postprocessor
self._event_ndims = event_ndims # rank of events
self._reparametrizable = reparametrizable
assert event_ndims in [0, 1]
@abc.abstractmethod
def create_dist(self, parameters):
"""Creates distribution from parameters."""
pass
@property
def param_size(self):
return self._param_size
@property
def reparametrizable(self):
return self._reparametrizable
def postprocess(self, event):
return self._postprocessor.forward(event)
def inverse_postprocess(self, event):
return self._postprocessor.inverse(event)
def sample_no_postprocessing(self, parameters, seed):
return self.create_dist(parameters).sample(seed=seed)
def sample(self, parameters, seed):
"""Returns a sample from the postprocessed distribution."""
return self.postprocess(self.sample_no_postprocessing(parameters, seed))
def mode(self, parameters):
"""Returns the mode of the postprocessed distribution."""
return self.postprocess(self.create_dist(parameters).mode())
def log_prob(self, parameters, actions):
"""Compute the log probability of actions."""
dist = self.create_dist(parameters)
log_probs = dist.log_prob(actions)
log_probs -= self._postprocessor.forward_log_det_jacobian(actions)
if self._event_ndims == 1:
log_probs = jnp.sum(log_probs, axis=-1) # sum over action dimension
return log_probs
def entropy(self, parameters, seed):
"""Return the entropy of the given distribution."""
dist = self.create_dist(parameters)
entropy = dist.entropy()
entropy += self._postprocessor.forward_log_det_jacobian(dist.sample(seed=seed))
if self._event_ndims == 1:
entropy = jnp.sum(entropy, axis=-1)
return entropy
class NormalDistribution:
"""Normal distribution."""
def __init__(self, loc, scale):
self.loc = loc
self.scale = scale
def sample(self, seed):
return jax.random.normal(seed, shape=self.loc.shape) * self.scale + self.loc
def mode(self):
return self.loc
def log_prob(self, x):
log_unnormalized = -0.5 * jnp.square(x / self.scale - self.loc / self.scale)
log_normalization = 0.5 * jnp.log(2.0 * jnp.pi) + jnp.log(self.scale)
return log_unnormalized - log_normalization
def entropy(self):
log_normalization = 0.5 * jnp.log(2.0 * jnp.pi) + jnp.log(self.scale)
entropy = 0.5 + log_normalization
return entropy * jnp.ones_like(self.loc)
class TanhBijector:
"""Tanh Bijector."""
def forward(self, x):
return jnp.tanh(x)
def inverse(self, y):
return jnp.arctanh(y)
def forward_log_det_jacobian(self, x):
return 2.0 * (jnp.log(2.0) - x - jax.nn.softplus(-2.0 * x))
class NormalTanhDistribution(ParametricDistribution):
"""Normal distribution followed by tanh."""
def __init__(self, event_size, min_std=0.001, var_scale=1):
"""Initialize the distribution.
Args:
event_size: the size of events (i.e. actions).
min_std: minimum std for the gaussian.
var_scale: adjust the gaussian's scale parameter.
"""
# We apply tanh to gaussian actions to bound them.
# Normally we would use TransformedDistribution to automatically
# apply tanh to the distribution.
# We can't do it here because of tanh saturation
# which would make log_prob computations impossible. Instead, most
# of the code operate on pre-tanh actions and we take the postprocessor
# jacobian into account in log_prob computations.
super().__init__(
param_size=2 * event_size,
postprocessor=TanhBijector(),
event_ndims=1,
reparametrizable=True,
)
self._min_std = min_std
self._var_scale = var_scale
def create_dist(self, parameters) -> NormalDistribution:
loc, scale = jnp.split(parameters, 2, axis=-1)
scale = (jax.nn.softplus(scale) + self._min_std) * self._var_scale
return NormalDistribution(loc=loc, scale=scale)