[lint] Update lint (#16)
* update lint * update readme * update ruff lint
This commit is contained in:
+105
-80
@@ -1,23 +1,23 @@
|
||||
|
||||
import math
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from torch.distributions import Beta
|
||||
from wall_x.utils.constant import action_statistic_dof
|
||||
|
||||
|
||||
class Normalizer(nn.Module):
|
||||
"""
|
||||
Action data normalizer for multi-robot systems.
|
||||
|
||||
|
||||
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.
|
||||
"""
|
||||
|
||||
|
||||
def __init__(self, action_statistic_dof, dof_config):
|
||||
"""
|
||||
Initialize the normalizer with robot-specific action statistics.
|
||||
|
||||
|
||||
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
|
||||
@@ -25,13 +25,13 @@ class Normalizer(nn.Module):
|
||||
super(Normalizer, self).__init__()
|
||||
|
||||
action_statistic = {}
|
||||
|
||||
|
||||
# Process statistics for each robot
|
||||
for robot_name in action_statistic_dof.keys():
|
||||
action_statistic[robot_name] = {}
|
||||
all_dof_min = []
|
||||
all_dof_delta = []
|
||||
|
||||
|
||||
# Collect min and delta values for all DOFs
|
||||
for k in dof_config:
|
||||
if k in action_statistic_dof[robot_name]:
|
||||
@@ -41,37 +41,41 @@ class Normalizer(nn.Module):
|
||||
# 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])
|
||||
|
||||
|
||||
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
|
||||
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()
|
||||
})
|
||||
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()
|
||||
}
|
||||
)
|
||||
|
||||
def normalize_data(self, xs, dataset_names):
|
||||
"""
|
||||
Normalize action data to [-1, 1] range using robot-specific statistics.
|
||||
|
||||
|
||||
Args:
|
||||
xs: Input action data tensors
|
||||
dataset_names: List of dataset/robot names corresponding to each tensor
|
||||
|
||||
|
||||
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"]
|
||||
|
||||
|
||||
for x, dataset_name in zip(xs, dataset_names):
|
||||
# Apply min-max normalization
|
||||
x = (x - self.min[dataset_name]) / (self.delta[dataset_name])
|
||||
@@ -80,19 +84,19 @@ class Normalizer(nn.Module):
|
||||
# Clamp to ensure bounds
|
||||
x = torch.clamp(x, -1, 1)
|
||||
new_xs.append(x)
|
||||
|
||||
|
||||
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.
|
||||
|
||||
|
||||
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
|
||||
|
||||
|
||||
Returns:
|
||||
torch.Tensor: Denormalized action data in original scale
|
||||
"""
|
||||
@@ -100,11 +104,11 @@ class Normalizer(nn.Module):
|
||||
# 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)
|
||||
|
||||
|
||||
for x, dataset_name, mask in zip(xs, dataset_names, dof_mask):
|
||||
# Convert from [-1, 1] to [0, 1] range
|
||||
x = (x + 1) / 2
|
||||
|
||||
|
||||
# Apply DOF mask if provided
|
||||
if mask is not None:
|
||||
mask = mask[0].bool()
|
||||
@@ -113,11 +117,11 @@ class Normalizer(nn.Module):
|
||||
else:
|
||||
action_space_delta = self.delta[dataset_name]
|
||||
action_space_min = self.min[dataset_name]
|
||||
|
||||
|
||||
# Scale back to original range
|
||||
x = x * action_space_delta + action_space_min
|
||||
new_xs.append(x)
|
||||
|
||||
|
||||
new_xs = torch.stack(new_xs)
|
||||
return new_xs
|
||||
|
||||
@@ -125,15 +129,15 @@ class Normalizer(nn.Module):
|
||||
class SinusoidalPosEmb(nn.Module):
|
||||
"""
|
||||
Sinusoidal positional embedding for diffusion timesteps.
|
||||
|
||||
|
||||
Generates sinusoidal embeddings commonly used in diffusion models to encode
|
||||
timestep information with different frequencies.
|
||||
"""
|
||||
|
||||
|
||||
def __init__(self, dim):
|
||||
"""
|
||||
Initialize sinusoidal positional embedding.
|
||||
|
||||
|
||||
Args:
|
||||
dim (int): Embedding dimension (must be even)
|
||||
"""
|
||||
@@ -143,10 +147,10 @@ class SinusoidalPosEmb(nn.Module):
|
||||
def forward(self, x):
|
||||
"""
|
||||
Generate sinusoidal embeddings for input timesteps.
|
||||
|
||||
|
||||
Args:
|
||||
x (torch.Tensor): Input timesteps
|
||||
|
||||
|
||||
Returns:
|
||||
torch.Tensor: Sinusoidal embeddings of shape (..., dim)
|
||||
"""
|
||||
@@ -159,25 +163,24 @@ class SinusoidalPosEmb(nn.Module):
|
||||
return emb
|
||||
|
||||
|
||||
|
||||
class ActionProcessor(nn.Module):
|
||||
"""
|
||||
Action sequence processor for robotic control with flow matching.
|
||||
|
||||
|
||||
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
|
||||
|
||||
|
||||
The Beta distribution provides more flexible noise injection strategies compared to
|
||||
traditional linear schedules, allowing better control over the noise scheduling process.
|
||||
"""
|
||||
|
||||
|
||||
def __init__(self, config):
|
||||
"""
|
||||
Initialize the action processor with multi-robot support.
|
||||
|
||||
|
||||
Args:
|
||||
config: Configuration object containing:
|
||||
- dof_config (dict): Degrees of freedom configuration per robot type
|
||||
@@ -186,7 +189,7 @@ class ActionProcessor(nn.Module):
|
||||
- noise_scheduler (dict): Noise scheduler configuration with Beta parameters
|
||||
"""
|
||||
super().__init__()
|
||||
|
||||
|
||||
# Calculate action and proprioception dimensions from configuration
|
||||
self.dof_config = config.dof_config
|
||||
self.agent_pos_config = config.agent_pos_config
|
||||
@@ -203,22 +206,28 @@ class ActionProcessor(nn.Module):
|
||||
print(" Agent position configuration:", flush=True)
|
||||
for key, value in self.agent_pos_config.items():
|
||||
print(f" {key}: {value}", flush=True)
|
||||
|
||||
|
||||
self.hidden_size = config.hidden_size
|
||||
|
||||
# Initialize data normalizers for actions and proprioception
|
||||
self.normalizer_action = Normalizer(action_statistic_dof, config.dof_config)
|
||||
self.normalizer_propri = Normalizer(action_statistic_dof, config.agent_pos_config)
|
||||
self.normalizer_propri = Normalizer(
|
||||
action_statistic_dof, config.agent_pos_config
|
||||
)
|
||||
|
||||
# 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
|
||||
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
|
||||
|
||||
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
|
||||
|
||||
# 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")
|
||||
@@ -228,51 +237,59 @@ class ActionProcessor(nn.Module):
|
||||
self.time_embed = SinusoidalPosEmb(config.hidden_size)
|
||||
|
||||
# Action embedding network: project to hidden space
|
||||
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
|
||||
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
|
||||
self.w3 = nn.Linear(self.hidden_size, self.hidden_size, bias=False)
|
||||
self.act_fn = nn.SiLU()
|
||||
|
||||
|
||||
# Project back to action space for flow matching loss
|
||||
self.action_proj_back = nn.Linear(self.hidden_size, self.action_dim, bias=False)
|
||||
self.mse_loss = nn.MSELoss(reduction='none')
|
||||
self.mse_loss = nn.MSELoss(reduction="none")
|
||||
|
||||
def sample_time(self, batch_size, device, dtype):
|
||||
"""
|
||||
Sample timesteps using Beta distribution for noise scheduling.
|
||||
|
||||
|
||||
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.
|
||||
|
||||
|
||||
Args:
|
||||
batch_size (int): Number of timesteps to sample
|
||||
device: Target device for tensors
|
||||
dtype: Target data type for tensors
|
||||
|
||||
|
||||
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
|
||||
|
||||
def proprioception_proj(self, proprioception, dataset_names=None, dof_mask=None, use_history=False):
|
||||
|
||||
def proprioception_proj(
|
||||
self, proprioception, dataset_names=None, dof_mask=None, use_history=False
|
||||
):
|
||||
"""
|
||||
Project proprioceptive data (joint positions, orientations) to hidden space.
|
||||
|
||||
|
||||
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.
|
||||
|
||||
|
||||
Returns:
|
||||
torch.Tensor: Projected proprioceptive features of shape [batch_size, seq_len, hidden_size]
|
||||
"""
|
||||
# Ensure proper device and dtype alignment
|
||||
proprioception = proprioception.to(device=self.propri_proj.weight.device).to(dtype=self.propri_proj.weight.dtype)
|
||||
|
||||
proprioception = proprioception.to(device=self.propri_proj.weight.device).to(
|
||||
dtype=self.propri_proj.weight.dtype
|
||||
)
|
||||
|
||||
if dof_mask is not None:
|
||||
# Concatenate proprioception with DOF mask
|
||||
# TODO: Use variable-based dimension checking for better flexibility
|
||||
@@ -280,26 +297,28 @@ class ActionProcessor(nn.Module):
|
||||
proprioception = torch.cat([proprioception, dof_mask], dim=-1)
|
||||
else:
|
||||
proprioception = torch.cat([proprioception, dof_mask], dim=-1)
|
||||
|
||||
proprioception = proprioception.to(device=self.propri_proj.weight.device).to(dtype=self.propri_proj.weight.dtype)
|
||||
|
||||
proprioception = proprioception.to(device=self.propri_proj.weight.device).to(
|
||||
dtype=self.propri_proj.weight.dtype
|
||||
)
|
||||
return self.propri_proj(proprioception)
|
||||
|
||||
def forward(self, action_chunk, dataset_names, dof_mask=None):
|
||||
"""
|
||||
Process action sequences with noise injection and temporal embedding.
|
||||
|
||||
|
||||
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
|
||||
|
||||
|
||||
Args:
|
||||
action_chunk (torch.Tensor): Action sequences of shape [batch_size, seq_len, action_dim]
|
||||
dataset_names (list): Dataset names for normalization
|
||||
dof_mask (torch.Tensor, optional): DOF mask of shape [batch_size, seq_len, action_dim].
|
||||
dof_mask (torch.Tensor, optional): DOF mask of shape [batch_size, seq_len, action_dim].
|
||||
Defaults to None.
|
||||
|
||||
|
||||
Returns:
|
||||
tuple: (action_embeddings, flow_target) where:
|
||||
- action_embeddings: Processed action features of shape [batch_size, seq_len, hidden_size]
|
||||
@@ -324,48 +343,54 @@ class ActionProcessor(nn.Module):
|
||||
# 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)
|
||||
|
||||
|
||||
noisy_action = noisy_action.to(dtype=self.w1.weight.dtype)
|
||||
action_embed = self.w1(noisy_action)
|
||||
|
||||
|
||||
# Repeat time embedding for each sequence position
|
||||
time_embed = time_embed.unsqueeze(1).repeat(1, action_embed.shape[1], 1).to(dtype=self.w2.weight.dtype)
|
||||
|
||||
time_embed = (
|
||||
time_embed.unsqueeze(1)
|
||||
.repeat(1, action_embed.shape[1], 1)
|
||||
.to(dtype=self.w2.weight.dtype)
|
||||
)
|
||||
|
||||
# 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
|
||||
|
||||
|
||||
def step(self, timestep, noisy_action, dof_mask=None):
|
||||
"""
|
||||
Single denoising step for diffusion inference.
|
||||
|
||||
|
||||
Processes noisy actions at a specific timestep for iterative denoising during inference.
|
||||
|
||||
|
||||
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.
|
||||
|
||||
|
||||
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)
|
||||
|
||||
|
||||
# Generate timestep embeddings
|
||||
time_embed = self.time_embed(timestep) # [batch_size, hidden_size]
|
||||
|
||||
|
||||
# Project noisy actions
|
||||
action_embed = self.w1(noisy_action)
|
||||
|
||||
|
||||
# Broadcast time embeddings to sequence length
|
||||
time_embed = time_embed.unsqueeze(1).repeat(1, action_embed.shape[1], 1)
|
||||
time_embed = time_embed.to(device=noisy_action.device).to(dtype=noisy_action.dtype)
|
||||
|
||||
time_embed = time_embed.to(device=noisy_action.device).to(
|
||||
dtype=noisy_action.dtype
|
||||
)
|
||||
|
||||
# Combine embeddings and process through MLP
|
||||
concat_embed = torch.cat([action_embed, time_embed], dim=-1)
|
||||
concat_embed = self.w2(concat_embed)
|
||||
@@ -376,25 +401,25 @@ class ActionProcessor(nn.Module):
|
||||
def flow_loss(self, action_hidden_states, flow, dof_mask=None):
|
||||
"""
|
||||
Compute flow matching loss between predicted and target actions.
|
||||
|
||||
|
||||
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.
|
||||
|
||||
|
||||
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)
|
||||
|
||||
|
||||
# Compute MSE loss between predicted and target flow
|
||||
loss = self.mse_loss(action_pred, flow)
|
||||
|
||||
|
||||
# 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
|
||||
|
||||
|
||||
# Return loss without reduction for channel-wise loss computation
|
||||
return loss
|
||||
return loss
|
||||
|
||||
@@ -1,2 +1,8 @@
|
||||
from .modeling_qwen2_5_vl_act import Qwen2_5_VLMoEModel,Qwen2_5_VLMoEForAction
|
||||
from .configuration_qwen2_5_vl import Qwen2_5_VLConfig
|
||||
from .modeling_qwen2_5_vl_act import Qwen2_5_VLMoEModel, Qwen2_5_VLMoEForAction
|
||||
from .configuration_qwen2_5_vl import Qwen2_5_VLConfig
|
||||
|
||||
__all__ = [
|
||||
"Qwen2_5_VLMoEModel",
|
||||
"Qwen2_5_VLMoEForAction",
|
||||
"Qwen2_5_VLConfig",
|
||||
]
|
||||
|
||||
@@ -190,7 +190,7 @@ class Qwen2_5_VLConfig(PretrainedConfig):
|
||||
experts=None,
|
||||
dof_config=None,
|
||||
noise_scheduler=None,
|
||||
dim_inputs=(1536,1536),
|
||||
dim_inputs=(1536, 1536),
|
||||
attention_moe=False,
|
||||
mlp_moe=False,
|
||||
**kwargs,
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
Reference in New Issue
Block a user