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-10-24 17:29:12 +08:00
|
|
|
|
import logging
|
2025-09-07 14:59:17 +08:00
|
|
|
|
|
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-10-24 17:29:12 +08:00
|
|
|
|
def _pad_to_action_dim(self, xs, action_dim):
|
|
|
|
|
|
"""
|
|
|
|
|
|
Pad the action data to the action dimension.
|
|
|
|
|
|
"""
|
|
|
|
|
|
if xs.shape[-1] < action_dim:
|
|
|
|
|
|
padding_shape = list(xs.shape)
|
|
|
|
|
|
padding_shape[-1] = action_dim - padding_shape[-1]
|
|
|
|
|
|
xs = torch.cat([xs, torch.zeros(padding_shape).to(xs.device)], dim=-1)
|
|
|
|
|
|
return xs
|
|
|
|
|
|
|
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-10-24 17:29:12 +08:00
|
|
|
|
# hard code the action dimension to 20
|
|
|
|
|
|
action_dim = 20
|
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
|
2025-10-24 17:29:12 +08:00
|
|
|
|
# raise ValueError(f"Statistics not available for {k} of {robot_name}")
|
|
|
|
|
|
logging.warning(
|
|
|
|
|
|
f"Statistics not available for {k} of {robot_name}, using default values"
|
|
|
|
|
|
)
|
2025-09-07 14:59:17 +08:00
|
|
|
|
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-10-24 17:29:12 +08:00
|
|
|
|
all_dof_min = self._pad_to_action_dim(torch.tensor(all_dof_min), action_dim)
|
|
|
|
|
|
all_dof_delta = self._pad_to_action_dim(
|
|
|
|
|
|
torch.tensor(all_dof_delta), action_dim
|
|
|
|
|
|
)
|
2025-09-07 14:59:17 +08:00
|
|
|
|
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
|
|
|
|
|
2025-10-24 17:29:12 +08:00
|
|
|
|
def normalize_data(self, xs, dataset_names, dof_mask=None):
|
2025-09-07 14:59:17 +08:00
|
|
|
|
"""
|
|
|
|
|
|
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-10-24 17:29:12 +08:00
|
|
|
|
dof_mask = dof_mask if dof_mask is not None else [None] * len(xs)
|
2025-09-11 13:18:33 +08:00
|
|
|
|
|
2025-10-24 17:29:12 +08:00
|
|
|
|
for x, dataset_name, mask in zip(xs, dataset_names, dof_mask):
|
|
|
|
|
|
# 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-07 14:59:17 +08:00
|
|
|
|
# Apply min-max normalization
|
2025-10-24 17:29:12 +08:00
|
|
|
|
x = (x - action_space_min) / (action_space_delta)
|
2025-09-07 14:59:17 +08:00
|
|
|
|
# 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
|
2025-10-24 17:29:12 +08:00
|
|
|
|
self.normalizer_action = Normalizer(
|
|
|
|
|
|
action_statistic_dof,
|
|
|
|
|
|
(
|
|
|
|
|
|
config.customized_dof_config
|
|
|
|
|
|
if hasattr(config, "customized_dof_config")
|
|
|
|
|
|
else config.dof_config
|
|
|
|
|
|
),
|
|
|
|
|
|
)
|
2025-09-11 13:18:33 +08:00
|
|
|
|
self.normalizer_propri = Normalizer(
|
2025-10-24 17:29:12 +08:00
|
|
|
|
action_statistic_dof,
|
|
|
|
|
|
(
|
|
|
|
|
|
config.customized_agent_pos_config
|
|
|
|
|
|
if hasattr(config, "customized_agent_pos_config")
|
|
|
|
|
|
else config.agent_pos_config
|
|
|
|
|
|
),
|
2025-09-11 13:18:33 +08:00
|
|
|
|
)
|
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)
|
2025-10-26 16:02:34 +08:00
|
|
|
|
time = (1 - sample) / self.s
|
2025-09-07 14:59:17 +08:00
|
|
|
|
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
|