Download DecMem/utils/action_encoder.py from BonanDing/worldmem-baseline-evals: direct link, hf CLI and curl.
- Browser
- Download file 2.19 kB
-
https://huggingface.co/BonanDing/worldmem-baseline-evals/resolve/main/DecMem/utils/action_encoder.py
- Command line
-
hf download hf://BonanDing/worldmem-baseline-evals/DecMem/utils/action_encoder.py
-
curl -L -o action_encoder.py https://huggingface.co/BonanDing/worldmem-baseline-evals/resolve/main/DecMem/utils/action_encoder.py
2.19 kB
| import torch | |
| import torch.nn as nn | |
| from einops import rearrange | |
| class ActionEncoder(nn.Module): | |
| """Encode raw per-video-frame action vectors and project to latent spatial dim. | |
| Input: (B, T_vid, action_dim) where T_vid = (T_lat - 1) * 4 + 1 | |
| Output: (B, T_lat, C, 1, 1) broadcastable latent-space embeddings. | |
| Packing strategy: | |
| - Frame 0 (first latent frame): the single raw action is repeated 4 times | |
| so it has the same packed dimension as grouped frames. | |
| - Frames 1..(T_vid-1): every 4 consecutive raw actions are flattened into | |
| one latent frame, giving T_lat - 1 grouped frames. | |
| The linear layer therefore receives 4 * action_dim features per latent frame. | |
| """ | |
| def __init__(self, action_dim: int = 25, hidden_dim: int = 64, latent_channels: int = 16): | |
| super().__init__() | |
| self.latent_channels = latent_channels | |
| # self.net = nn.Sequential( | |
| # nn.Linear(action_dim * 4, hidden_dim), | |
| # nn.SiLU(), | |
| # nn.Linear(hidden_dim, latent_channels), | |
| # ) | |
| self.net = nn.Linear(action_dim * 4, latent_channels, bias=False) | |
| def forward(self, actions: torch.Tensor) -> torch.Tensor: | |
| """ | |
| Args: | |
| actions: (B, T_vid, action_dim) T_vid = (T_lat - 1) * 4 + 1 | |
| Returns: | |
| (B, T_lat, C, 1, 1) | |
| """ | |
| B, T_vid, D = actions.shape | |
| # First frame: repeat raw action 4 times to match packed dim of grouped frames | |
| action_first = actions[:, :1].unsqueeze(2).repeat(1, 1, 4, 1) # (B, 1, 4, D) | |
| # Remaining frames: group every 4 raw frames into one latent frame | |
| actions_rest = actions[:, 1:].reshape(B, (T_vid - 1) // 4, 4, D) # (B, T_lat-1, 4, D) | |
| # Concatenate along the latent-frame dimension | |
| actions = torch.cat([action_first, actions_rest], dim=1) # (B, T_lat, 4, D) | |
| # Flatten 4 raw actions into the feature dim | |
| actions = rearrange(actions, "b f p d -> b f (p d)") # (B, T_lat, 4*D) | |
| emb = self.net(actions) # (B, T_lat, C) | |
| return emb.unsqueeze(-1).unsqueeze(-1) # (B, T_lat, C, 1, 1) | |