File size: 2,522 Bytes
952993c | 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 | from typing import Dict
import torch
from torch.nn.functional import pad
class ConcatLeftAlign:
def __init__(
self,
action_target_dim: int | None = None,
state_target_dim: int | None = None
):
self.action_target_dim = action_target_dim
self.state_target_dim = state_target_dim
def set_shape_meta(self, shape_meta):
self.action_meta = shape_meta["action"]
self.state_meta = shape_meta["state"]
def forward(self, batch):
if "action" in batch:
batch["action"] = self._concat(batch["action"], self.action_meta)
batch["action"], batch["action_dim_is_pad"] = self._pad(batch["action"], self.action_target_dim)
batch["state"] = self._concat(batch["state"], self.state_meta)
batch["state"], batch["state_dim_is_pad"] = self._pad(batch["state"], self.state_target_dim)
return batch
def backward(self, batch):
if self.state_target_dim is not None:
assert batch["state"].shape[-1] == self.state_target_dim
batch["state"] = self._crop(batch["state"], self.state_meta)
batch["state"] = self._split(batch["state"], self.state_meta)
if self.action_target_dim is not None:
assert batch["action"].shape[-1] == self.action_target_dim
batch["action"] = self._crop(batch["action"], self.action_meta)
batch["action"] = self._split(batch["action"], self.action_meta)
return batch
@staticmethod
def _pad(x: torch.Tensor, dim: int):
if dim is None:
dim = x.shape[-1]
assert x.ndim == 2 and x.shape[-1] <= dim
pad_dim = dim - x.shape[-1]
x_padded = pad(x, (0, pad_dim))
mask = torch.zeros_like(x[0]).bool()
mask = pad(mask, (0, pad_dim), value=True)
return x_padded, mask
@staticmethod
def _crop(x: torch.Tensor, meta: int):
assert x.ndim == 3
dim = sum([m["shape"] for m in meta])
x = x[:, :, :dim]
return x
@staticmethod
def _concat(x: Dict[str, torch.Tensor], meta: Dict[str, Dict]):
x = torch.cat([x[m["key"]] for m in meta], dim=-1)
assert x.ndim == 2
return x
@staticmethod
def _split(x: torch.Tensor, meta: Dict[str, Dict]):
assert x.ndim == 3
y = {}
idx = 0
for m in meta:
key, dim = m["key"], m["shape"]
y[key] = x[:, :, idx: idx + dim]
idx += dim
return y |