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)