worldmem-baseline-evals / DecMem /utils /action_encoder.py
BonanDing's picture
Add isolated Minecraft and RE10K baseline evaluation suite
59630ba verified
Raw History Blame Contribute Delete
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)