a2c2 / src /model.py
dennis96's picture
Upload folder using huggingface_hub
0adab2f verified
Raw
History Blame Contribute Delete
8.72 kB
"""A2C2 correction head architecture for cached BEHAVIOR/OpenPI features."""
from __future__ import annotations
from dataclasses import dataclass
import math
import torch
from torch import Tensor, nn
@dataclass(frozen=True)
class A2C2CorrectionHeadConfig:
state_dim: int = 256
action_dim: int = 23
action_horizon: int = 32
base_policy_z_dim: int = 2048
use_base_policy_z: bool = True
time_dim: int = 2
dim_model: int = 512
n_heads: int = 8
n_encoder_layers: int = 6
dim_feedforward: int = 2048
dropout: float = 0.1
mlp_hidden_dim: int = 1024
def _sinusoidal_positions(length: int, dim: int) -> Tensor:
if dim % 2 != 0:
raise ValueError("dim must be even for sinusoidal positional encoding.")
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) / dim))
pe = torch.zeros(length, dim, dtype=torch.float32)
pe[:, 0::2] = torch.sin(position * div_term)
pe[:, 1::2] = torch.cos(position * div_term)
return pe
class A2C2CorrectionHead(nn.Module):
"""Transformer + MLP correction head following the A2C2 residual design."""
def __init__(self, config: A2C2CorrectionHeadConfig | None = None) -> None:
super().__init__()
self.config = config or A2C2CorrectionHeadConfig()
cfg = self.config
self.cls_token = nn.Parameter(torch.zeros(1, 1, cfg.dim_model))
self.type_embedding = nn.Parameter(torch.zeros(6, cfg.dim_model))
self.state_proj = nn.Linear(cfg.state_dim, cfg.dim_model)
if cfg.use_base_policy_z:
self.z_proj = nn.Linear(cfg.base_policy_z_dim, cfg.dim_model)
self.time_proj = nn.Linear(cfg.time_dim, cfg.dim_model)
self.action_proj = nn.Linear(cfg.action_dim, cfg.dim_model)
chunk_pos = _sinusoidal_positions(cfg.action_horizon, cfg.dim_model)
self.register_buffer("chunk_pos_embedding", chunk_pos, persistent=False)
encoder_layer = nn.TransformerEncoderLayer(
d_model=cfg.dim_model,
nhead=cfg.n_heads,
dim_feedforward=cfg.dim_feedforward,
dropout=cfg.dropout,
activation="gelu",
batch_first=True,
norm_first=True,
)
self.encoder = nn.TransformerEncoder(encoder_layer, num_layers=cfg.n_encoder_layers)
self.encoder_norm = nn.LayerNorm(cfg.dim_model)
head_input_token_count = 5 if cfg.use_base_policy_z else 4
head_input_dim = cfg.dim_model * head_input_token_count + cfg.action_dim
self.residual_head = nn.Sequential(
nn.Linear(head_input_dim, cfg.mlp_hidden_dim),
nn.GELU(),
nn.Dropout(cfg.dropout),
nn.Linear(cfg.mlp_hidden_dim, cfg.mlp_hidden_dim),
nn.GELU(),
nn.Dropout(cfg.dropout),
nn.Linear(cfg.mlp_hidden_dim, cfg.action_dim),
)
self._reset_parameters()
@staticmethod
def make_time_feature(chunk_index: Tensor, horizon: int) -> Tensor:
"""Create [sin, cos] phase features from chunk indices."""
idx = chunk_index.to(dtype=torch.float32)
denom = max(horizon - 1, 1)
phase = 2.0 * math.pi * idx / denom
return torch.stack([torch.sin(phase), torch.cos(phase)], dim=-1)
def forward(
self,
observation_state: Tensor,
selected_base_action: Tensor,
base_action_chunk: Tensor,
base_policy_z: Tensor,
time_feature: Tensor,
valid_action_mask: Tensor | None = None,
) -> Tensor:
"""Predict residual action delta.
Args:
observation_state: [B, state_dim]
selected_base_action: [B, action_dim], the current action being corrected.
base_action_chunk: [B, H, action_dim]
base_policy_z: [B, z_dim]
time_feature: [B, 2]
valid_action_mask: optional bool tensor [B, H], True for valid chunk
entries. Invalid chunk entries are ignored by transformer attention.
Returns:
Tensor [B, action_dim], the predicted residual delta.
"""
cfg = self.config
batch_size = observation_state.shape[0]
device = observation_state.device
dtype = observation_state.dtype
self._validate_inputs(observation_state, selected_base_action, base_action_chunk, base_policy_z, time_feature)
cls = self.cls_token.to(device=device, dtype=dtype).expand(batch_size, -1, -1)
cls = cls + self.type_embedding[0].to(device=device, dtype=dtype)
state_token = self.state_proj(observation_state).unsqueeze(1)
state_token = state_token + self.type_embedding[1].to(device=device, dtype=dtype)
time_token = self.time_proj(time_feature).unsqueeze(1)
time_token = time_token + self.type_embedding[3].to(device=device, dtype=dtype)
selected_action_token = self.action_proj(selected_base_action).unsqueeze(1)
selected_action_token = selected_action_token + self.type_embedding[4].to(device=device, dtype=dtype)
chunk_tokens = self.action_proj(base_action_chunk)
chunk_pos = self.chunk_pos_embedding[: base_action_chunk.shape[1]].to(device=device, dtype=dtype)
chunk_tokens = chunk_tokens + chunk_pos.unsqueeze(0)
chunk_tokens = chunk_tokens + self.type_embedding[5].to(device=device, dtype=dtype)
prefix_tokens = [cls, state_token]
if cfg.use_base_policy_z:
z_token = self.z_proj(base_policy_z).unsqueeze(1)
z_token = z_token + self.type_embedding[2].to(device=device, dtype=dtype)
prefix_tokens.append(z_token)
prefix_tokens.extend([time_token, selected_action_token])
tokens = torch.cat([*prefix_tokens, chunk_tokens], dim=1)
padding_mask = None
if valid_action_mask is not None:
valid_action_mask = valid_action_mask.to(device=device, dtype=torch.bool)
prefix_mask = torch.zeros(batch_size, len(prefix_tokens), device=device, dtype=torch.bool)
padding_mask = torch.cat([prefix_mask, ~valid_action_mask], dim=1)
encoded = self.encoder(tokens, src_key_padding_mask=padding_mask)
encoded = self.encoder_norm(encoded)
cls_state = encoded[:, 0]
state_state = encoded[:, 1]
if cfg.use_base_policy_z:
z_state = encoded[:, 2]
time_state = encoded[:, 3]
selected_action_state = encoded[:, 4]
head_states = [cls_state, state_state, z_state, time_state, selected_action_state]
else:
time_state = encoded[:, 2]
selected_action_state = encoded[:, 3]
head_states = [cls_state, state_state, time_state, selected_action_state]
head_input = torch.cat(
[*head_states, selected_base_action],
dim=-1,
)
return self.residual_head(head_input)
def _validate_inputs(
self,
observation_state: Tensor,
selected_base_action: Tensor,
base_action_chunk: Tensor,
base_policy_z: Tensor,
time_feature: Tensor,
) -> None:
cfg = self.config
if observation_state.ndim != 2 or observation_state.shape[-1] != cfg.state_dim:
raise ValueError(f"observation_state must have shape [B, {cfg.state_dim}].")
if selected_base_action.ndim != 2 or selected_base_action.shape[-1] != cfg.action_dim:
raise ValueError(f"selected_base_action must have shape [B, {cfg.action_dim}].")
if base_action_chunk.ndim != 3 or base_action_chunk.shape[-1] != cfg.action_dim:
raise ValueError(f"base_action_chunk must have shape [B, H, {cfg.action_dim}].")
if base_action_chunk.shape[1] > cfg.action_horizon:
raise ValueError(f"base_action_chunk horizon cannot exceed {cfg.action_horizon}.")
if cfg.use_base_policy_z and (base_policy_z.ndim != 2 or base_policy_z.shape[-1] != cfg.base_policy_z_dim):
raise ValueError(f"base_policy_z must have shape [B, {cfg.base_policy_z_dim}].")
if time_feature.ndim != 2 or time_feature.shape[-1] != cfg.time_dim:
raise ValueError(f"time_feature must have shape [B, {cfg.time_dim}].")
def _reset_parameters(self) -> None:
nn.init.trunc_normal_(self.cls_token, std=0.02)
nn.init.trunc_normal_(self.type_embedding, 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)