mirror of
https://github.com/wassname/ray.git
synced 2026-08-03 13:10:57 +08:00
[rllib] tuple space shouldn't assume elements are all the same size (#2637)
* fix * lint
This commit is contained in:
@@ -6,8 +6,6 @@ import tensorflow as tf
|
||||
import numpy as np
|
||||
import distutils.version
|
||||
|
||||
from ray.rllib.utils.reshaper import Reshaper
|
||||
|
||||
use_tf150_api = (distutils.version.LooseVersion(tf.VERSION) >=
|
||||
distutils.version.LooseVersion("1.5.0"))
|
||||
|
||||
@@ -182,10 +180,10 @@ class MultiActionDistribution(ActionDistribution):
|
||||
inputs (Tensor list): A list of tensors from which to compute samples.
|
||||
"""
|
||||
|
||||
def __init__(self, inputs, action_space, child_distributions):
|
||||
# you actually have to instantiate the child distributions
|
||||
self.reshaper = Reshaper(action_space.spaces)
|
||||
split_inputs = self.reshaper.split_tensor(inputs)
|
||||
def __init__(self, inputs, action_space, child_distributions, input_lens):
|
||||
self.input_lens = input_lens
|
||||
inputs = tf.reshape(inputs, [-1, sum(input_lens)])
|
||||
split_inputs = tf.split(inputs, self.input_lens, axis=1)
|
||||
child_list = []
|
||||
for i, distribution in enumerate(child_distributions):
|
||||
child_list.append(distribution(split_inputs[i]))
|
||||
@@ -193,7 +191,7 @@ class MultiActionDistribution(ActionDistribution):
|
||||
|
||||
def logp(self, x):
|
||||
"""The log-likelihood of the action distribution."""
|
||||
split_list = self.reshaper.split_tensor(x)
|
||||
split_list = tf.split(x, len(self.input_lens), axis=1)
|
||||
for i, distribution in enumerate(self.child_distributions):
|
||||
# Remove extra categorical dimension
|
||||
if isinstance(distribution, Categorical):
|
||||
|
||||
Reference in New Issue
Block a user