[RL] Allow passing tensors of different dtypes for FlattenedTensorBucket (#13413)
This commit is contained in:
@@ -22,6 +22,9 @@ class FlattenedTensorBucket:
|
|||||||
while preserving all metadata needed for reconstruction.
|
while preserving all metadata needed for reconstruction.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
# This field is solely for users of to check whether the class supports this feature
|
||||||
|
supports_multi_dtypes = True
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
named_tensors: List[Tuple[str, torch.Tensor]] = None,
|
named_tensors: List[Tuple[str, torch.Tensor]] = None,
|
||||||
@@ -48,7 +51,7 @@ class FlattenedTensorBucket:
|
|||||||
flattened_tensors: List[torch.Tensor] = [None] * len(named_tensors)
|
flattened_tensors: List[torch.Tensor] = [None] * len(named_tensors)
|
||||||
|
|
||||||
for i, (name, tensor) in enumerate(named_tensors):
|
for i, (name, tensor) in enumerate(named_tensors):
|
||||||
flattened = tensor.flatten()
|
flattened = tensor.flatten().view(torch.uint8)
|
||||||
flattened_tensors[i] = flattened
|
flattened_tensors[i] = flattened
|
||||||
|
|
||||||
# Store metadata
|
# Store metadata
|
||||||
@@ -93,14 +96,12 @@ class FlattenedTensorBucket:
|
|||||||
reconstructed = [None] * len(self.metadata)
|
reconstructed = [None] * len(self.metadata)
|
||||||
|
|
||||||
for i, meta in enumerate(self.metadata):
|
for i, meta in enumerate(self.metadata):
|
||||||
tensor = self.flattened_tensor[meta.start_idx : meta.end_idx].reshape(
|
tensor = (
|
||||||
meta.shape
|
self.flattened_tensor[meta.start_idx : meta.end_idx]
|
||||||
|
.view(meta.dtype)
|
||||||
|
.reshape(meta.shape)
|
||||||
)
|
)
|
||||||
|
|
||||||
# batch dtype conversion (if needed)
|
|
||||||
if tensor.dtype != meta.dtype:
|
|
||||||
tensor = tensor.to(meta.dtype)
|
|
||||||
|
|
||||||
reconstructed[i] = (meta.name, tensor)
|
reconstructed[i] = (meta.name, tensor)
|
||||||
|
|
||||||
return reconstructed
|
return reconstructed
|
||||||
|
|||||||
Reference in New Issue
Block a user