258 lines
10 KiB
Python
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
|