Files
roboimi/roboimi/vla/models/heads/act.py
T
2026-07-31 10:11:04 +08:00

258 lines
10 KiB
Python

"""Local ACT-style CVAE policy head.
This module intentionally reimplements only the ACT model logic needed by the
RoboIMI VLA training stack. It does not import or vendor the external ACT repo.
"""
from __future__ import annotations
import math
from typing import Optional, Tuple
import torch
import torch.nn as nn
def kl_divergence(mu: torch.Tensor, logvar: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
"""KL divergence between diagonal posterior N(mu, exp(logvar)) and N(0, I)."""
if mu.ndim > 2:
mu = mu.view(mu.size(0), -1)
if logvar.ndim > 2:
logvar = logvar.view(logvar.size(0), -1)
klds = -0.5 * (1 + logvar - mu.pow(2) - logvar.exp())
total_kld = klds.sum(1).mean(0, keepdim=True)
dimension_wise_kld = klds.mean(0)
mean_kld = klds.mean(1).mean(0, keepdim=True)
return total_kld, dimension_wise_kld, mean_kld
def _build_sinusoidal_table(length: int, dim: int) -> torch.Tensor:
if length <= 0:
raise ValueError(f"length must be positive, got {length}")
if dim <= 0:
raise ValueError(f"dim must be positive, got {dim}")
position = torch.arange(length, dtype=torch.float32).unsqueeze(1)
div_term = torch.exp(
torch.arange(0, dim, 2, dtype=torch.float32) * (-math.log(10000.0) / max(dim, 1))
)
table = torch.zeros(length, dim, dtype=torch.float32)
table[:, 0::2] = torch.sin(position * div_term)
if dim > 1:
table[:, 1::2] = torch.cos(position * div_term[: table[:, 1::2].shape[1]])
return table.unsqueeze(0)
class ACTPolicyHead(nn.Module):
"""ACT CVAE head using native PyTorch transformer blocks.
Args follow the existing Hydra style and are intentionally configurable so
this implementation is not tied to ACT's original 14-DoF ALOHA setup.
"""
def __init__(
self,
action_dim: int,
obs_dim: int,
vision_dim: int,
num_cams: int,
pred_horizon: int,
obs_horizon: int = 1,
hidden_dim: int = 256,
nheads: int = 8,
enc_layers: int = 4,
dec_layers: int = 6,
dim_feedforward: int = 2048,
latent_dim: int = 32,
dropout: float = 0.1,
kl_weight: float = 10.0,
activation: str = "gelu",
**_: object,
) -> None:
super().__init__()
self.action_dim = int(action_dim)
self.obs_dim = int(obs_dim)
self.vision_dim = int(vision_dim)
self.num_cams = int(num_cams)
self.pred_horizon = int(pred_horizon)
self.obs_horizon = int(obs_horizon)
self.hidden_dim = int(hidden_dim)
self.nheads = int(nheads)
self.latent_dim = int(latent_dim)
self.kl_weight = float(kl_weight)
if self.pred_horizon <= 0:
raise ValueError(f"pred_horizon must be positive, got {self.pred_horizon}")
if self.hidden_dim % self.nheads != 0:
raise ValueError(
f"hidden_dim({self.hidden_dim}) must be divisible by nheads({self.nheads})"
)
self.cls_embed = nn.Parameter(torch.zeros(1, 1, self.hidden_dim))
self.encoder_qpos_proj = nn.Linear(self.obs_dim, self.hidden_dim)
self.encoder_action_proj = nn.Linear(self.action_dim, self.hidden_dim)
encoder_layer = nn.TransformerEncoderLayer(
d_model=self.hidden_dim,
nhead=self.nheads,
dim_feedforward=int(dim_feedforward),
dropout=float(dropout),
activation=activation,
batch_first=True,
norm_first=False,
)
self.posterior_encoder = nn.TransformerEncoder(encoder_layer, num_layers=int(enc_layers))
self.latent_proj = nn.Linear(self.hidden_dim, self.latent_dim * 2)
self.latent_out_proj = nn.Linear(self.latent_dim, self.hidden_dim)
self.decoder_qpos_proj = nn.Linear(self.obs_dim, self.hidden_dim)
self.visual_proj = nn.Linear(self.vision_dim, self.hidden_dim)
self.memory_type_embed = nn.Embedding(3, self.hidden_dim) # latent, qpos, visual
self.query_embed = nn.Embedding(self.pred_horizon, self.hidden_dim)
decoder_layer = nn.TransformerDecoderLayer(
d_model=self.hidden_dim,
nhead=self.nheads,
dim_feedforward=int(dim_feedforward),
dropout=float(dropout),
activation=activation,
batch_first=True,
norm_first=False,
)
self.decoder = nn.TransformerDecoder(decoder_layer, num_layers=int(dec_layers))
self.action_head = nn.Linear(self.hidden_dim, self.action_dim)
self.register_buffer(
"posterior_pos_table",
_build_sinusoidal_table(self.pred_horizon + 2, self.hidden_dim),
persistent=False,
)
self.register_buffer(
"memory_pos_table",
_build_sinusoidal_table(1024, self.hidden_dim),
persistent=False,
)
self._reset_parameters()
def _reset_parameters(self) -> None:
nn.init.normal_(self.cls_embed, mean=0.0, std=0.02)
for module in self.modules():
if isinstance(module, nn.Linear):
nn.init.xavier_uniform_(module.weight)
if module.bias is not None:
nn.init.zeros_(module.bias)
elif isinstance(module, nn.Embedding):
nn.init.normal_(module.weight, mean=0.0, std=0.02)
def _validate_inputs(
self,
qpos: torch.Tensor,
visual_tokens: torch.Tensor,
actions: Optional[torch.Tensor],
action_is_pad: Optional[torch.Tensor],
) -> None:
if qpos.ndim != 2 or qpos.shape[-1] != self.obs_dim:
raise ValueError(f"qpos must have shape (B,{self.obs_dim}), got {tuple(qpos.shape)}")
if visual_tokens.ndim != 3 or visual_tokens.shape[-1] != self.vision_dim:
raise ValueError(
f"visual_tokens must have shape (B,N,{self.vision_dim}), got {tuple(visual_tokens.shape)}"
)
if visual_tokens.shape[0] != qpos.shape[0]:
raise ValueError("qpos and visual_tokens batch dimensions must match")
if actions is not None:
expected = (qpos.shape[0], self.pred_horizon, self.action_dim)
if tuple(actions.shape) != expected:
raise ValueError(f"actions must have shape {expected}, got {tuple(actions.shape)}")
if action_is_pad is not None and tuple(action_is_pad.shape) != expected[:2]:
raise ValueError(
f"action_is_pad must have shape {expected[:2]}, got {tuple(action_is_pad.shape)}"
)
def _posterior(
self,
qpos: torch.Tensor,
actions: Optional[torch.Tensor],
action_is_pad: Optional[torch.Tensor],
) -> Tuple[torch.Tensor, Optional[torch.Tensor], Optional[torch.Tensor], Optional[torch.Tensor]]:
batch_size = qpos.shape[0]
if actions is None:
latent = torch.zeros(batch_size, self.latent_dim, device=qpos.device, dtype=qpos.dtype)
return latent, None, None, None
cls = self.cls_embed.to(dtype=qpos.dtype).expand(batch_size, -1, -1)
qpos_token = self.encoder_qpos_proj(qpos).unsqueeze(1)
action_tokens = self.encoder_action_proj(actions)
tokens = torch.cat([cls, qpos_token, action_tokens], dim=1)
pos = self.posterior_pos_table[:, : tokens.shape[1]].to(device=tokens.device, dtype=tokens.dtype)
tokens = tokens + pos
padding_mask = None
if action_is_pad is not None:
prefix = torch.zeros(
batch_size,
2,
dtype=torch.bool,
device=action_is_pad.device,
)
padding_mask = torch.cat([prefix, action_is_pad.to(torch.bool)], dim=1)
encoded = self.posterior_encoder(tokens, src_key_padding_mask=padding_mask)
latent_info = self.latent_proj(encoded[:, 0])
mu, logvar = torch.chunk(latent_info, 2, dim=-1)
std = torch.exp(0.5 * logvar)
eps = torch.randn_like(std)
latent = mu + eps * std
kl, _, _ = kl_divergence(mu, logvar)
return latent, mu, logvar, kl
def _memory(self, qpos: torch.Tensor, visual_tokens: torch.Tensor, latent: torch.Tensor) -> torch.Tensor:
batch_size = qpos.shape[0]
latent_token = self.latent_out_proj(latent).unsqueeze(1)
qpos_token = self.decoder_qpos_proj(qpos).unsqueeze(1)
visual = self.visual_proj(visual_tokens)
memory = torch.cat([latent_token, qpos_token, visual], dim=1)
if memory.shape[1] > self.memory_pos_table.shape[1]:
pos = _build_sinusoidal_table(memory.shape[1], self.hidden_dim).to(
device=memory.device,
dtype=memory.dtype,
)
else:
pos = self.memory_pos_table[:, : memory.shape[1]].to(
device=memory.device,
dtype=memory.dtype,
)
memory = memory + pos
type_ids = torch.cat(
[
torch.zeros(1, dtype=torch.long, device=memory.device),
torch.ones(1, dtype=torch.long, device=memory.device),
torch.full((memory.shape[1] - 2,), 2, dtype=torch.long, device=memory.device),
]
)
memory = memory + self.memory_type_embed(type_ids).to(dtype=memory.dtype).unsqueeze(0)
if memory.shape[0] != batch_size:
raise RuntimeError("internal memory batch shape mismatch")
return memory
def forward(
self,
qpos: torch.Tensor,
visual_tokens: torch.Tensor,
actions: Optional[torch.Tensor] = None,
action_is_pad: Optional[torch.Tensor] = None,
):
self._validate_inputs(qpos, visual_tokens, actions, action_is_pad)
latent, mu, logvar, kl = self._posterior(qpos, actions, action_is_pad)
memory = self._memory(qpos, visual_tokens, latent)
query = self.query_embed.weight.to(dtype=memory.dtype).unsqueeze(0).expand(qpos.shape[0], -1, -1)
target = torch.zeros_like(query)
decoded = self.decoder(target + query, memory)
pred_actions = self.action_head(decoded)
if kl is None:
kl = torch.zeros(1, device=qpos.device, dtype=qpos.dtype)
latent_info = {
"mu": mu,
"logvar": logvar,
"kl": kl,
"kl_weight": self.kl_weight,
}
return pred_actions, latent_info