[rllib] Basic infrastructure for off-policy estimation (IS, WIS) (#3941)

This commit is contained in:
Eric Liang
2019-02-13 16:25:05 -08:00
committed by GitHub
parent 729d0b2825
commit 2dccf383dd
34 changed files with 549 additions and 131 deletions
+32 -6
View File
@@ -24,6 +24,7 @@ class ActionDistribution(object):
@DeveloperAPI
def __init__(self, inputs):
self.inputs = inputs
self.sample_op = self._build_sample_op()
@DeveloperAPI
def logp(self, x):
@@ -37,13 +38,27 @@ class ActionDistribution(object):
@DeveloperAPI
def entropy(self):
"""The entroy of the action distribution."""
"""The entropy of the action distribution."""
raise NotImplementedError
@DeveloperAPI
def _build_sample_op(self):
"""Implement this instead of sample(), to enable op reuse.
This is needed since the sample op is non-deterministic and is shared
between sample() and sampled_action_prob().
"""
raise NotImplementedError
@DeveloperAPI
def sample(self):
"""Draw a sample from the action distribution."""
raise NotImplementedError
return self.sample_op
@DeveloperAPI
def sampled_action_prob(self):
"""Returns the log probability of the sampled action."""
return tf.exp(self.logp(self.sample_op))
class Categorical(ActionDistribution):
@@ -95,7 +110,7 @@ class Categorical(ActionDistribution):
p0 * (a0 - tf.log(z0) - a1 + tf.log(z1)), reduction_indices=[1])
@override(ActionDistribution)
def sample(self):
def _build_sample_op(self):
return tf.squeeze(tf.multinomial(self.inputs, 1), axis=1)
@@ -107,11 +122,11 @@ class DiagGaussian(ActionDistribution):
"""
def __init__(self, inputs):
ActionDistribution.__init__(self, inputs)
mean, log_std = tf.split(inputs, 2, axis=1)
self.mean = mean
self.log_std = log_std
self.std = tf.exp(log_std)
ActionDistribution.__init__(self, inputs)
@override(ActionDistribution)
def logp(self, x):
@@ -136,7 +151,7 @@ class DiagGaussian(ActionDistribution):
reduction_indices=[1])
@override(ActionDistribution)
def sample(self):
def _build_sample_op(self):
return self.mean + self.std * tf.random_normal(tf.shape(self.mean))
@@ -147,7 +162,11 @@ class Deterministic(ActionDistribution):
"""
@override(ActionDistribution)
def sample(self):
def sampled_action_prob(self):
return 1.0
@override(ActionDistribution)
def _build_sample_op(self):
return self.inputs
@@ -205,5 +224,12 @@ class MultiActionDistribution(ActionDistribution):
def sample(self):
return TupleActions([s.sample() for s in self.child_distributions])
@override(ActionDistribution)
def sampled_action_prob(self):
p = self.child_distributions[0].sampled_action_prob()
for c in self.child_distributions[1:]:
p *= c.sampled_action_prob()
return p
TupleActions = namedtuple("TupleActions", ["batches"])