mirror of
https://github.com/wassname/ray.git
synced 2026-08-15 12:45:23 +08:00
[RLlib] Policy.compute_single_action() broken for nested actions (Issue 8411). (#8514)
This commit is contained in:
@@ -92,3 +92,39 @@ def flatten_to_single_ndarray(input_):
|
||||
expanded.append(np.reshape(in_, [-1]))
|
||||
input_ = np.concatenate(expanded, axis=0).flatten()
|
||||
return input_
|
||||
|
||||
|
||||
def unbatch(batches_struct):
|
||||
"""Converts input from (nested) struct of batches to batch of structs.
|
||||
|
||||
Input: Struct of different batches (each batch has size=3):
|
||||
{"a": [1, 2, 3], "b": ([4, 5, 6], [7.0, 8.0, 9.0])}
|
||||
Output: Batch (list) of structs (each of these structs representing a
|
||||
single action):
|
||||
[
|
||||
{"a": 1, "b": (4, 7.0)}, <- action 1
|
||||
{"a": 2, "b": (5, 8.0)}, <- action 2
|
||||
{"a": 3, "b": (6, 9.0)}, <- action 3
|
||||
]
|
||||
|
||||
Args:
|
||||
batches_struct (any): The struct of component batches. Each leaf item
|
||||
in this struct represents the batch for a single component
|
||||
(in case struct is tuple/dict).
|
||||
Alternatively, `batches_struct` may also simply be a batch of
|
||||
primitives (non tuple/dict).
|
||||
|
||||
Returns:
|
||||
List[struct[components]]: The list of rows. Each item
|
||||
in the returned list represents a single (maybe complex) struct.
|
||||
"""
|
||||
flat_batches = tree.flatten(batches_struct)
|
||||
|
||||
out = []
|
||||
for batch_pos in range(len(flat_batches[0])):
|
||||
out.append(
|
||||
tree.unflatten_as(
|
||||
batches_struct,
|
||||
[flat_batches[i][batch_pos]
|
||||
for i in range(len(flat_batches))]))
|
||||
return out
|
||||
|
||||
Reference in New Issue
Block a user