Files
VLA/wall_x/model/action_head.py
T

426 lines
17 KiB
Python
Raw Normal View History

2025-09-07 14:59:17 +08:00
import math
import torch
import torch.nn as nn
from torch.distributions import Beta
from wall_x.utils.constant import action_statistic_dof
2025-09-11 13:18:33 +08:00
2025-09-07 14:59:17 +08:00
class Normalizer(nn.Module):
"""
Action data normalizer for multi-robot systems.
2025-09-11 13:18:33 +08:00
2025-09-07 14:59:17 +08:00
This module handles normalization and denormalization of action data for different robot
configurations. It maintains per-robot statistics (min values and deltas) and applies
normalization to map actions to the [-1, 1] range.
"""
2025-09-11 13:18:33 +08:00
2025-09-07 14:59:17 +08:00
def __init__(self, action_statistic_dof, dof_config):
"""
Initialize the normalizer with robot-specific action statistics.
2025-09-11 13:18:33 +08:00
2025-09-07 14:59:17 +08:00
Args:
action_statistic_dof (dict): Statistical data for each robot's degrees of freedom
dof_config (dict): Configuration mapping for degrees of freedom per robot
"""
super(Normalizer, self).__init__()
action_statistic = {}
2025-09-11 13:18:33 +08:00
2025-09-07 14:59:17 +08:00
# Process statistics for each robot
for robot_name in action_statistic_dof.keys():
action_statistic[robot_name] = {}
all_dof_min = []
all_dof_delta = []
2025-09-11 13:18:33 +08:00
2025-09-07 14:59:17 +08:00
# Collect min and delta values for all DOFs
for k in dof_config:
if k in action_statistic_dof[robot_name]:
all_dof_min.extend(action_statistic_dof[robot_name][k]["min"])
all_dof_delta.extend(action_statistic_dof[robot_name][k]["delta"])
else:
# Use default values if statistics not available
all_dof_min.extend([0.0] * dof_config[k])
all_dof_delta.extend([1.0] * dof_config[k])
2025-09-11 13:18:33 +08:00
2025-09-07 14:59:17 +08:00
all_dof_min = torch.tensor(all_dof_min)
all_dof_delta = torch.tensor(all_dof_delta)
action_statistic[robot_name]["min"] = all_dof_min
action_statistic[robot_name]["delta"] = all_dof_delta
# Register statistics as non-trainable parameters
2025-09-11 13:18:33 +08:00
self.min = nn.ParameterDict(
{
k: nn.Parameter(action_statistic[k]["min"], requires_grad=False)
for k in action_statistic.keys()
}
)
self.delta = nn.ParameterDict(
{
k: nn.Parameter(action_statistic[k]["delta"], requires_grad=False)
for k in action_statistic.keys()
}
)
2025-09-07 14:59:17 +08:00
def normalize_data(self, xs, dataset_names):
"""
Normalize action data to [-1, 1] range using robot-specific statistics.
2025-09-11 13:18:33 +08:00
2025-09-07 14:59:17 +08:00
Args:
xs: Input action data tensors
dataset_names: List of dataset/robot names corresponding to each tensor
2025-09-11 13:18:33 +08:00
2025-09-07 14:59:17 +08:00
Returns:
torch.Tensor: Normalized action data in [-1, 1] range
"""
new_xs = []
# Filter out multimodal dataset entries
dataset_names = [name for name in dataset_names if name != "x2_multimodal"]
2025-09-11 13:18:33 +08:00
2025-09-07 14:59:17 +08:00
for x, dataset_name in zip(xs, dataset_names):
# Apply min-max normalization
x = (x - self.min[dataset_name]) / (self.delta[dataset_name])
# Scale to [-1, 1] range
x = x * 2 - 1
# Clamp to ensure bounds
x = torch.clamp(x, -1, 1)
new_xs.append(x)
2025-09-11 13:18:33 +08:00
2025-09-07 14:59:17 +08:00
new_xs = torch.stack(new_xs)
return new_xs
def unnormalize_data(self, xs, dataset_names, dof_mask=None):
"""
Convert normalized data back to original action space.
2025-09-11 13:18:33 +08:00
2025-09-07 14:59:17 +08:00
Args:
xs: Normalized action data in [-1, 1] range
dataset_names: List of dataset/robot names
dof_mask: Optional mask to select specific degrees of freedom
2025-09-11 13:18:33 +08:00
2025-09-07 14:59:17 +08:00
Returns:
torch.Tensor: Denormalized action data in original scale
"""
new_xs = []
# Filter out multimodal dataset entries
dataset_names = [name for name in dataset_names if name != "x2_multimodal"]
dof_mask = dof_mask if dof_mask is not None else [None] * len(xs)
2025-09-11 13:18:33 +08:00
2025-09-07 14:59:17 +08:00
for x, dataset_name, mask in zip(xs, dataset_names, dof_mask):
# Convert from [-1, 1] to [0, 1] range
x = (x + 1) / 2
2025-09-11 13:18:33 +08:00
2025-09-07 14:59:17 +08:00
# Apply DOF mask if provided
if mask is not None:
mask = mask[0].bool()
action_space_delta = self.delta[dataset_name][mask]
action_space_min = self.min[dataset_name][mask]
else:
action_space_delta = self.delta[dataset_name]
action_space_min = self.min[dataset_name]
2025-09-11 13:18:33 +08:00
2025-09-07 14:59:17 +08:00
# Scale back to original range
x = x * action_space_delta + action_space_min
new_xs.append(x)
2025-09-11 13:18:33 +08:00
2025-09-07 14:59:17 +08:00
new_xs = torch.stack(new_xs)
return new_xs
class SinusoidalPosEmb(nn.Module):
"""
Sinusoidal positional embedding for diffusion timesteps.
2025-09-11 13:18:33 +08:00
2025-09-07 14:59:17 +08:00
Generates sinusoidal embeddings commonly used in diffusion models to encode
timestep information with different frequencies.
"""
2025-09-11 13:18:33 +08:00
2025-09-07 14:59:17 +08:00
def __init__(self, dim):
"""
Initialize sinusoidal positional embedding.
2025-09-11 13:18:33 +08:00
2025-09-07 14:59:17 +08:00
Args:
dim (int): Embedding dimension (must be even)
"""
super().__init__()
self.dim = dim
def forward(self, x):
"""
Generate sinusoidal embeddings for input timesteps.
2025-09-11 13:18:33 +08:00
2025-09-07 14:59:17 +08:00
Args:
x (torch.Tensor): Input timesteps
2025-09-11 13:18:33 +08:00
2025-09-07 14:59:17 +08:00
Returns:
torch.Tensor: Sinusoidal embeddings of shape (..., dim)
"""
device = x.device
half_dim = self.dim // 2
emb = math.log(10000) / (half_dim - 1)
emb = torch.exp(torch.arange(half_dim, device=device) * -emb)
emb = x[:, None] * emb[None, :]
emb = torch.cat((emb.sin(), emb.cos()), dim=-1)
return emb
class ActionProcessor(nn.Module):
"""
Action sequence processor for robotic control with flow matching.
2025-09-11 13:18:33 +08:00
2025-09-07 14:59:17 +08:00
This module handles action sequence processing for robotic systems with the following capabilities:
1. Adds controlled noise to action sequences using Beta distribution scheduling
2. Generates temporal embeddings for timestep conditioning
3. Projects actions to model hidden space for transformer processing
4. Supports proprioceptive data integration and multi-robot configurations
2025-09-11 13:18:33 +08:00
2025-09-07 14:59:17 +08:00
The Beta distribution provides more flexible noise injection strategies compared to
traditional linear schedules, allowing better control over the noise scheduling process.
"""
2025-09-11 13:18:33 +08:00
2025-09-07 14:59:17 +08:00
def __init__(self, config):
"""
Initialize the action processor with multi-robot support.
2025-09-11 13:18:33 +08:00
2025-09-07 14:59:17 +08:00
Args:
config: Configuration object containing:
- dof_config (dict): Degrees of freedom configuration per robot type
- agent_pos_config (dict): Agent position/proprioception configuration
- hidden_size (int): Model hidden layer dimension
- noise_scheduler (dict): Noise scheduler configuration with Beta parameters
"""
super().__init__()
2025-09-11 13:18:33 +08:00
2025-09-07 14:59:17 +08:00
# Calculate action and proprioception dimensions from configuration
self.dof_config = config.dof_config
self.agent_pos_config = config.agent_pos_config
self.action_dim = sum([v for k, v in self.dof_config.items()])
self.propri_dim = sum([v for k, v in self.agent_pos_config.items()])
# Log configuration details for debugging
print("ActionProcessor Configuration:", flush=True)
print(f" Action dimension: {self.action_dim}", flush=True)
print(f" Proprioception dimension: {self.propri_dim}", flush=True)
print(" DOF configuration:", flush=True)
for key, value in self.dof_config.items():
print(f" {key}: {value}", flush=True)
print(" Agent position configuration:", flush=True)
for key, value in self.agent_pos_config.items():
print(f" {key}: {value}", flush=True)
2025-09-11 13:18:33 +08:00
2025-09-07 14:59:17 +08:00
self.hidden_size = config.hidden_size
# Initialize data normalizers for actions and proprioception
self.normalizer_action = Normalizer(action_statistic_dof, config.dof_config)
2025-09-11 13:18:33 +08:00
self.normalizer_propri = Normalizer(
action_statistic_dof, config.agent_pos_config
)
2025-09-07 14:59:17 +08:00
# Proprioception projection layer (includes history/current state)
self.propri_proj = nn.Linear(self.propri_dim * 2, self.hidden_size, bias=False)
# Beta distribution noise scheduler configuration
noise_scheduler_config = config.noise_scheduler
2025-09-11 13:18:33 +08:00
self.beta_alpha = noise_scheduler_config.get(
"beta_alpha", 1.5
) # Beta distribution α parameter
self.beta_beta = noise_scheduler_config.get(
"beta_beta", 1.0
) # Beta distribution β parameter
self.s = noise_scheduler_config.get("s", 0.999) # Scaling factor
2025-09-07 14:59:17 +08:00
# Initialize Beta distribution for noise scheduling
alpha_tensor = torch.tensor(self.beta_alpha, dtype=torch.float32).to("cuda")
beta_tensor = torch.tensor(self.beta_beta, dtype=torch.float32).to("cuda")
self.beta_dist = Beta(alpha_tensor, beta_tensor)
# Sinusoidal positional embedding for timesteps
self.time_embed = SinusoidalPosEmb(config.hidden_size)
# Action embedding network: project to hidden space
2025-09-11 13:18:33 +08:00
self.w1 = nn.Linear(
self.action_dim * 2, self.hidden_size, bias=False
) # *2 for action + DOF mask
self.w2 = nn.Linear(
self.hidden_size * 2, self.hidden_size, bias=False
) # *2 for action + time embeddings
2025-09-07 14:59:17 +08:00
self.w3 = nn.Linear(self.hidden_size, self.hidden_size, bias=False)
self.act_fn = nn.SiLU()
2025-09-11 13:18:33 +08:00
2025-09-07 14:59:17 +08:00
# Project back to action space for flow matching loss
self.action_proj_back = nn.Linear(self.hidden_size, self.action_dim, bias=False)
2025-09-11 13:18:33 +08:00
self.mse_loss = nn.MSELoss(reduction="none")
2025-09-07 14:59:17 +08:00
def sample_time(self, batch_size, device, dtype):
"""
Sample timesteps using Beta distribution for noise scheduling.
2025-09-11 13:18:33 +08:00
2025-09-07 14:59:17 +08:00
Generates random timesteps in [0,1] range using Beta distribution, then scales them.
This provides more flexible control over the noise injection schedule compared to
uniform sampling.
2025-09-11 13:18:33 +08:00
2025-09-07 14:59:17 +08:00
Args:
batch_size (int): Number of timesteps to sample
device: Target device for tensors
dtype: Target data type for tensors
2025-09-11 13:18:33 +08:00
2025-09-07 14:59:17 +08:00
Returns:
torch.Tensor: Sampled timesteps of shape [batch_size]
"""
sample = self.beta_dist.sample([batch_size]).to(device=device, dtype=dtype)
time = (self.s - sample) / self.s
return time
2025-09-11 13:18:33 +08:00
def proprioception_proj(
self, proprioception, dataset_names=None, dof_mask=None, use_history=False
):
2025-09-07 14:59:17 +08:00
"""
Project proprioceptive data (joint positions, orientations) to hidden space.
2025-09-11 13:18:33 +08:00
2025-09-07 14:59:17 +08:00
Args:
proprioception (torch.Tensor): Proprioceptive data of shape [batch_size, seq_len, propri_dim]
dataset_names (list, optional): Dataset names for normalization. Defaults to None.
dof_mask (torch.Tensor, optional): DOF mask of shape [batch_size, propri_dim]. Defaults to None.
use_history (bool, optional): Whether to use historical proprioceptive data. Defaults to False.
2025-09-11 13:18:33 +08:00
2025-09-07 14:59:17 +08:00
Returns:
torch.Tensor: Projected proprioceptive features of shape [batch_size, seq_len, hidden_size]
"""
# Ensure proper device and dtype alignment
2025-09-11 13:18:33 +08:00
proprioception = proprioception.to(device=self.propri_proj.weight.device).to(
dtype=self.propri_proj.weight.dtype
)
2025-09-07 14:59:17 +08:00
if dof_mask is not None:
# Concatenate proprioception with DOF mask
# TODO: Use variable-based dimension checking for better flexibility
if use_history:
proprioception = torch.cat([proprioception, dof_mask], dim=-1)
else:
proprioception = torch.cat([proprioception, dof_mask], dim=-1)
2025-09-11 13:18:33 +08:00
proprioception = proprioception.to(device=self.propri_proj.weight.device).to(
dtype=self.propri_proj.weight.dtype
)
2025-09-07 14:59:17 +08:00
return self.propri_proj(proprioception)
def forward(self, action_chunk, dataset_names, dof_mask=None):
"""
Process action sequences with noise injection and temporal embedding.
2025-09-11 13:18:33 +08:00
2025-09-07 14:59:17 +08:00
This method implements the forward pass for flow matching training:
1. Adds Beta-distributed noise to action sequences
2. Generates sinusoidal timestep embeddings
3. Projects noisy actions to hidden space
4. Combines action and temporal features
2025-09-11 13:18:33 +08:00
2025-09-07 14:59:17 +08:00
Args:
action_chunk (torch.Tensor): Action sequences of shape [batch_size, seq_len, action_dim]
dataset_names (list): Dataset names for normalization
2025-09-11 13:18:33 +08:00
dof_mask (torch.Tensor, optional): DOF mask of shape [batch_size, seq_len, action_dim].
2025-09-07 14:59:17 +08:00
Defaults to None.
2025-09-11 13:18:33 +08:00
2025-09-07 14:59:17 +08:00
Returns:
tuple: (action_embeddings, flow_target) where:
- action_embeddings: Processed action features of shape [batch_size, seq_len, hidden_size]
- flow_target: Flow matching target (action_chunk - noise) for loss computation
"""
batch_size = action_chunk.shape[0]
device = action_chunk.device
dtype = action_chunk.dtype
# 1. Add noise to action sequences using flow matching
noise = torch.randn_like(action_chunk)
time = self.sample_time(batch_size, device, dtype)
t = time.unsqueeze(-1).unsqueeze(-1) # Broadcast to match action dimensions
# Linear interpolation between noise and action (flow matching)
noisy_action = (1 - t) * noise + t * action_chunk
flow = action_chunk - noise # Flow target for loss computation
# 2. Generate sinusoidal positional encoding for timesteps
time_embed = self.time_embed(time)
# 3. Project noisy actions with DOF mask to hidden space
if dof_mask is not None:
noisy_action = torch.cat([noisy_action, dof_mask], dim=-1)
2025-09-11 13:18:33 +08:00
2025-09-07 14:59:17 +08:00
noisy_action = noisy_action.to(dtype=self.w1.weight.dtype)
action_embed = self.w1(noisy_action)
2025-09-11 13:18:33 +08:00
2025-09-07 14:59:17 +08:00
# Repeat time embedding for each sequence position
2025-09-11 13:18:33 +08:00
time_embed = (
time_embed.unsqueeze(1)
.repeat(1, action_embed.shape[1], 1)
.to(dtype=self.w2.weight.dtype)
)
2025-09-07 14:59:17 +08:00
# Combine action and temporal embeddings
concat_embed = torch.cat([action_embed, time_embed], dim=-1)
concat_embed = self.w2(concat_embed)
embed = self.w3(self.act_fn(concat_embed))
return embed, flow
2025-09-11 13:18:33 +08:00
2025-09-07 14:59:17 +08:00
def step(self, timestep, noisy_action, dof_mask=None):
"""
Single denoising step for diffusion inference.
2025-09-11 13:18:33 +08:00
2025-09-07 14:59:17 +08:00
Processes noisy actions at a specific timestep for iterative denoising during inference.
2025-09-11 13:18:33 +08:00
2025-09-07 14:59:17 +08:00
Args:
timestep (torch.Tensor): Current timesteps of shape [batch_size]
noisy_action (torch.Tensor): Noisy actions of shape [batch_size, seq_len, action_dim]
dof_mask (torch.Tensor, optional): DOF mask for action space. Defaults to None.
2025-09-11 13:18:33 +08:00
2025-09-07 14:59:17 +08:00
Returns:
torch.Tensor: Processed action embeddings of shape [batch_size, seq_len, hidden_size]
"""
# Concatenate noisy action with DOF mask if provided
if dof_mask is not None:
noisy_action = torch.cat([noisy_action, dof_mask], dim=-1)
2025-09-11 13:18:33 +08:00
2025-09-07 14:59:17 +08:00
# Generate timestep embeddings
time_embed = self.time_embed(timestep) # [batch_size, hidden_size]
2025-09-11 13:18:33 +08:00
2025-09-07 14:59:17 +08:00
# Project noisy actions
action_embed = self.w1(noisy_action)
2025-09-11 13:18:33 +08:00
2025-09-07 14:59:17 +08:00
# Broadcast time embeddings to sequence length
time_embed = time_embed.unsqueeze(1).repeat(1, action_embed.shape[1], 1)
2025-09-11 13:18:33 +08:00
time_embed = time_embed.to(device=noisy_action.device).to(
dtype=noisy_action.dtype
)
2025-09-07 14:59:17 +08:00
# Combine embeddings and process through MLP
concat_embed = torch.cat([action_embed, time_embed], dim=-1)
concat_embed = self.w2(concat_embed)
embed = self.w3(self.act_fn(concat_embed))
return embed
def flow_loss(self, action_hidden_states, flow, dof_mask=None):
"""
Compute flow matching loss between predicted and target actions.
2025-09-11 13:18:33 +08:00
2025-09-07 14:59:17 +08:00
Args:
action_hidden_states (torch.Tensor): Hidden states from transformer
flow (torch.Tensor): Target flow (action - noise) for matching
dof_mask (torch.Tensor, optional): DOF mask to weight loss per dimension. Defaults to None.
2025-09-11 13:18:33 +08:00
2025-09-07 14:59:17 +08:00
Returns:
torch.Tensor: Flow matching loss (no reduction for channel loss computation)
"""
# Project hidden states back to action space
action_pred = self.action_proj_back(action_hidden_states)
2025-09-11 13:18:33 +08:00
2025-09-07 14:59:17 +08:00
# Compute MSE loss between predicted and target flow
loss = self.mse_loss(action_pred, flow)
2025-09-11 13:18:33 +08:00
2025-09-07 14:59:17 +08:00
# Apply DOF mask if provided
if dof_mask is not None:
dof_mask = dof_mask.reshape(-1, dof_mask.shape[-1])
loss = loss * dof_mask
2025-09-11 13:18:33 +08:00
2025-09-07 14:59:17 +08:00
# Return loss without reduction for channel-wise loss computation
2025-09-11 13:18:33 +08:00
return loss