[lint] Update lint (#16)

* update lint

* update readme

* update ruff lint
This commit is contained in:
Lufang Chen
2025-09-11 13:18:33 +08:00
committed by GitHub
parent a89dce95aa
commit e9332a283d
28 changed files with 2406 additions and 1074 deletions
+27
View File
@@ -0,0 +1,27 @@
# See https://pre-commit.com for more information
# See https://pre-commit.com/hooks.html for more hooks
exclude: ".git"
repos:
- repo: https://github.com/astral-sh/ruff-pre-commit
rev: v0.2.2
hooks:
- id: ruff
args: [ --fix, --exit-non-zero-on-fix ]
- repo: https://github.com/psf/black
rev: 24.2.0
hooks:
- id: black
- repo: https://github.com/pre-commit/pre-commit-hooks
rev: v4.5.0
hooks:
- id: check-added-large-files
- id: check-ast
- id: check-case-conflict
- id: check-merge-conflict
- id: check-toml
- id: check-yaml
- id: end-of-file-fixer
- id: trailing-whitespace
+17
View File
@@ -0,0 +1,17 @@
# Contributing to Wall-x
## Submit a Pull Request
Before opening a pull request, please make sure your code passes the lint checks.
```bash
# Install pre-commit hooks (run once)
pre-commit install
```
Or
```bash
# Manually run all checks
pre-commit run --all-files
```
-1
View File
@@ -13,4 +13,3 @@ PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.def("rope", &launch_multimodal_rope_forward, "Multimodal RoPE forward kernel"); m.def("rope", &launch_multimodal_rope_forward, "Multimodal RoPE forward kernel");
m.def("rope_bwd", &launch_multimodal_rope_backward, "Multimodal RoPE backward kernel"); m.def("rope_bwd", &launch_multimodal_rope_backward, "Multimodal RoPE backward kernel");
} }
-1
View File
@@ -549,4 +549,3 @@ void launch_multimodal_rope_backward(
break; break;
} }
} }
+2
View File
@@ -0,0 +1,2 @@
[tool.ruff]
per-file-ignores = { "__init__.py" = ["F401", "E402"] }
+18 -13
View File
@@ -1,6 +1,7 @@
import os import os
import yaml import yaml
import torch import torch
import matplotlib.pyplot as plt
from wall_x.model.qwen2_5_based.modeling_qwen2_5_vl_act import Qwen2_5_VLMoEForAction from wall_x.model.qwen2_5_based.modeling_qwen2_5_vl_act import Qwen2_5_VLMoEForAction
from wall_x.data.load_lerobot_dataset import load_test_dataset, get_data_configs from wall_x.data.load_lerobot_dataset import load_test_dataset, get_data_configs
@@ -8,11 +9,14 @@ from wall_x.data.load_lerobot_dataset import load_test_dataset, get_data_configs
model_path = "path/to/model" model_path = "path/to/model"
action_tokenizer_path = "path/to/action_tokenizer" action_tokenizer_path = "path/to/action_tokenizer"
save_dir = "path/to/plot" save_dir = "path/to/plot"
model = Qwen2_5_VLMoEForAction.from_pretrained(model_path, action_tokenizer_path=action_tokenizer_path) model = Qwen2_5_VLMoEForAction.from_pretrained(
model_path, action_tokenizer_path=action_tokenizer_path
)
model.eval() model.eval()
model = model.to("cuda") model = model.to("cuda")
model = model.bfloat16() model = model.bfloat16()
def load_config(config_path): def load_config(config_path):
"""Load configuration from YAML file.""" """Load configuration from YAML file."""
with open(config_path, "r") as f: with open(config_path, "r") as f:
@@ -22,6 +26,7 @@ def load_config(config_path):
return config return config
# get test dataloader # get test dataloader
path = "path/to/config" path = "path/to/config"
config = load_config(path) config = load_config(path)
@@ -46,14 +51,16 @@ for idx, batch in enumerate(dataloader):
action_dim=action_dim, action_dim=action_dim,
pred_horizon=pred_horizon, pred_horizon=pred_horizon,
mode="predict", mode="predict",
predict_mode="fast" predict_mode="fast",
) )
pred_traj[idx : idx + pred_horizon] = outputs['predict_action'].detach().cpu() pred_traj[idx : idx + pred_horizon] = outputs["predict_action"].detach().cpu()
# Denormalize ground truth actions # Denormalize ground truth actions
gt_action_chunk = batch['action_chunk'][:, :, :action_dim] gt_action_chunk = batch["action_chunk"][:, :, :action_dim]
dof_mask = batch["dof_mask"].to(gt_action_chunk.dtype) dof_mask = batch["dof_mask"].to(gt_action_chunk.dtype)
denormalized_gt = model.action_preprocessor.normalizer_action.unnormalize_data(gt_action_chunk, ["x2_normal"], dof_mask) denormalized_gt = model.action_preprocessor.normalizer_action.unnormalize_data(
gt_action_chunk, ["x2_normal"], dof_mask
)
gt_traj[idx : idx + pred_horizon] = denormalized_gt.detach().cpu() gt_traj[idx : idx + pred_horizon] = denormalized_gt.detach().cpu()
@@ -62,20 +69,18 @@ pred_traj_np = pred_traj.numpy()
timesteps = gt_traj.shape[0] timesteps = gt_traj.shape[0]
import matplotlib.pyplot as plt
fig, axs = plt.subplots(action_dim, 1, figsize=(15, 5 * action_dim), sharex=True) fig, axs = plt.subplots(action_dim, 1, figsize=(15, 5 * action_dim), sharex=True)
fig.suptitle(f'Action Comparison for lerobot', fontsize=16) fig.suptitle("Action Comparison for lerobot", fontsize=16)
for i in range(action_dim): for i in range(action_dim):
axs[i].plot(range(timesteps), gt_traj_np[:, i], label='Ground Truth') axs[i].plot(range(timesteps), gt_traj_np[:, i], label="Ground Truth")
axs[i].plot(range(timesteps), pred_traj_np[:, i], label='Prediction') axs[i].plot(range(timesteps), pred_traj_np[:, i], label="Prediction")
axs[i].set_ylabel(f'Action Dim {i+1}') axs[i].set_ylabel(f"Action Dim {i+1}")
axs[i].legend() axs[i].legend()
axs[i].grid(True) axs[i].grid(True)
axs[-1].set_xlabel('Timestep') axs[-1].set_xlabel("Timestep")
plt.tight_layout(rect=[0, 0.03, 1, 0.95]) plt.tight_layout(rect=[0, 0.03, 1, 0.95])
os.makedirs(save_dir, exist_ok=True) os.makedirs(save_dir, exist_ok=True)
plt.savefig(os.path.join(save_dir, f"lerobot_comparison.png")) plt.savefig(os.path.join(save_dir, "lerobot_comparison.png"))
plt.close() plt.close()
+9 -4
View File
@@ -10,10 +10,14 @@ batch_size = 1
seq_length = 50 seq_length = 50
torch.manual_seed(0) torch.manual_seed(0)
fake_input_ids = torch.randint(0, len(model.processor.tokenizer), (batch_size, seq_length), dtype=torch.long) fake_input_ids = torch.randint(
0, len(model.processor.tokenizer), (batch_size, seq_length), dtype=torch.long
)
fake_attention_mask = torch.ones((batch_size, seq_length), dtype=torch.long) fake_attention_mask = torch.ones((batch_size, seq_length), dtype=torch.long)
fake_moe_token_types = torch.zeros((batch_size, seq_length), dtype=torch.long) fake_moe_token_types = torch.zeros((batch_size, seq_length), dtype=torch.long)
fake_position_ids = torch.arange(seq_length, dtype=torch.long).unsqueeze(0).expand(batch_size, -1) fake_position_ids = (
torch.arange(seq_length, dtype=torch.long).unsqueeze(0).expand(batch_size, -1)
)
fake_proprioception = torch.randn((batch_size, 1, 20), dtype=torch.float32) fake_proprioception = torch.randn((batch_size, 1, 20), dtype=torch.float32)
fake_agent_pos_mask = torch.ones((batch_size, 1, 20), dtype=torch.float32) fake_agent_pos_mask = torch.ones((batch_size, 1, 20), dtype=torch.float32)
fake_dof_mask = torch.ones((batch_size, 32, 20), dtype=torch.float32) fake_dof_mask = torch.ones((batch_size, 32, 20), dtype=torch.float32)
@@ -44,7 +48,7 @@ try:
agent_pos_mask=fake_agent_pos_mask, agent_pos_mask=fake_agent_pos_mask,
dof_mask=fake_dof_mask, dof_mask=fake_dof_mask,
dataset_names=fake_dataset_names, dataset_names=fake_dataset_names,
mode="validate" mode="validate",
) )
print("✅ Fake inference test successful!") print("✅ Fake inference test successful!")
@@ -68,7 +72,7 @@ try:
else: else:
print("❌ Output contains infinity values") print("❌ Output contains infinity values")
print(f"Output logits statistics:") print("Output logits statistics:")
print(f" Min value: {outputs.logits.min().item():.4f}") print(f" Min value: {outputs.logits.min().item():.4f}")
print(f" Max value: {outputs.logits.max().item():.4f}") print(f" Max value: {outputs.logits.max().item():.4f}")
print(f" Mean: {outputs.logits.mean().item():.4f}") print(f" Mean: {outputs.logits.mean().item():.4f}")
@@ -77,4 +81,5 @@ try:
except Exception as e: except Exception as e:
print(f"❌ Fake inference test failed: {e}") print(f"❌ Fake inference test failed: {e}")
import traceback import traceback
traceback.print_exc() traceback.print_exc()
+4 -3
View File
@@ -8,13 +8,15 @@ use_fast_tokenizer = True
processor = AutoProcessor.from_pretrained(processor_path, use_fast=True) processor = AutoProcessor.from_pretrained(processor_path, use_fast=True)
processor.tokenizer.padding_side = "left" processor.tokenizer.padding_side = "left"
action_tokenizer = AutoProcessor.from_pretrained(action_tokenizer_path, trust_remote_code=True) action_tokenizer = AutoProcessor.from_pretrained(
action_tokenizer_path, trust_remote_code=True
)
new_tokens = ["<|propri|>", "<|action|>"] new_tokens = ["<|propri|>", "<|action|>"]
new_tokens += [f"<|action_token_{i}|>" for i in range(action_tokenizer.vocab_size)] new_tokens += [f"<|action_token_{i}|>" for i in range(action_tokenizer.vocab_size)]
num_added_tokens = processor.tokenizer.add_tokens(new_tokens) num_added_tokens = processor.tokenizer.add_tokens(new_tokens)
begin_idx_token = f"<|action_token_0|>" begin_idx_token = "<|action_token_0|>"
token_id = processor.tokenizer.convert_tokens_to_ids(begin_idx_token) token_id = processor.tokenizer.convert_tokens_to_ids(begin_idx_token)
processor.tokenizer.init_kwargs["action_token_start_index"] = token_id processor.tokenizer.init_kwargs["action_token_start_index"] = token_id
processor.tokenizer.init_kwargs["action_token_vocab_size"] = action_tokenizer.vocab_size processor.tokenizer.init_kwargs["action_token_vocab_size"] = action_tokenizer.vocab_size
@@ -22,4 +24,3 @@ processor.tokenizer.init_kwargs["action_token_vocab_size"] = action_tokenizer.vo
new_tokenizer_dir = "/path/to/new_tokenizer" new_tokenizer_dir = "/path/to/new_tokenizer"
os.makedirs(new_tokenizer_dir, exist_ok=True) os.makedirs(new_tokenizer_dir, exist_ok=True)
processor.save_pretrained(new_tokenizer_dir) processor.save_pretrained(new_tokenizer_dir)
+19 -7
View File
@@ -4,7 +4,11 @@ import time
import yaml import yaml
import wandb import wandb
from argparse import ArgumentParser from argparse import ArgumentParser
from accelerate import Accelerator, DistributedDataParallelKwargs, DataLoaderConfiguration from accelerate import (
Accelerator,
DistributedDataParallelKwargs,
DataLoaderConfiguration,
)
from wall_x.trainer.qwen_vl_act_trainer import QwenVlAct_Trainer from wall_x.trainer.qwen_vl_act_trainer import QwenVlAct_Trainer
@@ -27,7 +31,9 @@ def load_config(config_path):
def setup_accelerator(config): def setup_accelerator(config):
"""Initialize and configure the accelerator for distributed training.""" """Initialize and configure the accelerator for distributed training."""
print(f"[{time.strftime('%Y-%m-%d %H:%M:%S', time.localtime())}] Preparing accelerator") print(
f"[{time.strftime('%Y-%m-%d %H:%M:%S', time.localtime())}] Preparing accelerator"
)
ddp_kwargs = DistributedDataParallelKwargs(find_unused_parameters=True) ddp_kwargs = DistributedDataParallelKwargs(find_unused_parameters=True)
accelerator_dataloader_config = DataLoaderConfiguration(dispatch_batches=False) accelerator_dataloader_config = DataLoaderConfiguration(dispatch_batches=False)
@@ -36,10 +42,12 @@ def setup_accelerator(config):
kwargs_handlers=[ddp_kwargs], kwargs_handlers=[ddp_kwargs],
mixed_precision="bf16", mixed_precision="bf16",
dataloader_config=accelerator_dataloader_config, dataloader_config=accelerator_dataloader_config,
gradient_accumulation_steps=config.get("gradient_accumulation_steps", 1) gradient_accumulation_steps=config.get("gradient_accumulation_steps", 1),
) )
print(f"[{time.strftime('%Y-%m-%d %H:%M:%S', time.localtime())}] Accelerator initialization complete") print(
f"[{time.strftime('%Y-%m-%d %H:%M:%S', time.localtime())}] Accelerator initialization complete"
)
return accelerator return accelerator
@@ -97,10 +105,14 @@ def main(args):
trainer.fit() trainer.fit()
if __name__ == '__main__': if __name__ == "__main__":
parser = ArgumentParser(description="Training script for Wall-X model") parser = ArgumentParser(description="Training script for Wall-X model")
parser.add_argument("--config", type=str, required=True, help="Path to configuration YAML file") parser.add_argument(
parser.add_argument("--seed", type=int, default=42, help="Random seed for reproducibility") "--config", type=str, required=True, help="Path to configuration YAML file"
)
parser.add_argument(
"--seed", type=int, default=42, help="Random seed for reproducibility"
)
args = parser.parse_args() args = parser.parse_args()
main(args) main(args)
+42 -13
View File
@@ -6,24 +6,51 @@ from qwen_vl_utils.vision_process import MIN_PIXELS, MAX_PIXELS, IMAGE_FACTOR
# Tactile sensor file mapping for data processing # Tactile sensor file mapping for data processing
TACTILE_FILE_MAPPING = { TACTILE_FILE_MAPPING = {
"tactile_data_left": "left_tactile", "tactile_data_left": "left_tactile",
"tactile_data_right": "right_tactile" "tactile_data_right": "right_tactile",
} }
# Supported action datasets # Supported action datasets
ACTION_DATASET_NAMES = [ ACTION_DATASET_NAMES = [
"x2_normal", "agibotworld_alpha", "droid", "fractal", "bridge_data_v2", "x2_normal",
"DobbE", "RH20T", "UMI-biarm", "austin_buds", "austin_sailor", "austin_sirius", "agibotworld_alpha",
"bc_z", "berkeley_autolab_ur5", "berkeley_cable_routing", "berkeley_fanuc_manipulation", "droid",
"dlr_edan_shared_control", "fmb", "furniture_bench", "jaco_play", "nyu_rot", "fractal",
"stanford_hydra", "stanford_kuka_multimodal", "taco_play", "utaustin_mutex", "viola" "bridge_data_v2",
"DobbE",
"RH20T",
"UMI-biarm",
"austin_buds",
"austin_sailor",
"austin_sirius",
"bc_z",
"berkeley_autolab_ur5",
"berkeley_cable_routing",
"berkeley_fanuc_manipulation",
"dlr_edan_shared_control",
"fmb",
"furniture_bench",
"jaco_play",
"nyu_rot",
"stanford_hydra",
"stanford_kuka_multimodal",
"taco_play",
"utaustin_mutex",
"viola",
] ]
# Supported multimodal datasets # Supported multimodal datasets
MULTIMODAL_DATASET_NAMES = [ MULTIMODAL_DATASET_NAMES = [
"x2_multimodal_from_action", "x2_multimodal", "x2_subtask_generation", "x2_multimodal_from_action",
"multimodal_CapsFusion", "multimodal_Robo2VLM", "multimodal_RoboPoint", "x2_multimodal",
"multimodal_EQA", "multimodal_Cambrian", "multimodal_pixmo", "x2_subtask_generation",
"multimodal_VQAv2", "multimodal_COCO" "multimodal_CapsFusion",
"multimodal_Robo2VLM",
"multimodal_RoboPoint",
"multimodal_EQA",
"multimodal_Cambrian",
"multimodal_pixmo",
"multimodal_VQAv2",
"multimodal_COCO",
] ]
@@ -45,7 +72,7 @@ class X2RDataProcessingConfig:
default_factory=lambda: { default_factory=lambda: {
"face_view": -1, "face_view": -1,
"left_wrist_view": 128, "left_wrist_view": 128,
"right_wrist_view": 128 "right_wrist_view": 128,
} }
) )
@@ -68,7 +95,9 @@ class X2RDataProcessingConfig:
"""Post-initialization validation and setup.""" """Post-initialization validation and setup."""
# Validate train/test split # Validate train/test split
if not 0 < self.train_test_split < 1: if not 0 < self.train_test_split < 1:
raise ValueError(f"train_test_split must be between 0 and 1, got {self.train_test_split}") raise ValueError(
f"train_test_split must be between 0 and 1, got {self.train_test_split}"
)
def as_dict(self) -> Dict: def as_dict(self) -> Dict:
"""Convert configuration to dictionary format. """Convert configuration to dictionary format.
@@ -78,7 +107,7 @@ class X2RDataProcessingConfig:
""" """
return self.__dict__ return self.__dict__
def update(self, **kwargs) -> 'X2RDataProcessingConfig': def update(self, **kwargs) -> "X2RDataProcessingConfig":
"""Update configuration parameters. """Update configuration parameters.
Args: Args:
+116 -25
View File
@@ -8,7 +8,12 @@ from lerobot.datasets.lerobot_dataset import LeRobotDataset
from typing import Protocol, SupportsIndex, TypeVar from typing import Protocol, SupportsIndex, TypeVar
from qwen_vl_utils.vision_process import smart_resize from qwen_vl_utils.vision_process import smart_resize
from wall_x.data.config import X2RDataProcessingConfig from wall_x.data.config import X2RDataProcessingConfig
from wall_x.data.utils import process_grounding_points, get_wallx_normal_text, replace_action_token, preprocesser_call from wall_x.data.utils import (
process_grounding_points,
get_wallx_normal_text,
replace_action_token,
preprocesser_call,
)
from transformers import AutoProcessor from transformers import AutoProcessor
@@ -67,7 +72,9 @@ class PreprocessedDataset(Dataset[T_co]):
img_pil = Image.fromarray((current_obs * 255).to(torch.uint8).cpu().numpy()) img_pil = Image.fromarray((current_obs * 255).to(torch.uint8).cpu().numpy())
orig_width, orig_height = img_pil.size orig_width, orig_height = img_pil.size
# 2. Apply resolution constraints (if config is not -1) # 2. Apply resolution constraints (if config is not -1)
target_size = self.data_config.resolution.get(self._cam_key_mapping[key], -1) target_size = self.data_config.resolution.get(
self._cam_key_mapping[key], -1
)
if target_size != -1: if target_size != -1:
# Maintain aspect ratio logic # Maintain aspect ratio logic
if orig_width > orig_height: # Landscape image if orig_width > orig_height: # Landscape image
@@ -108,7 +115,9 @@ class PreprocessedDataset(Dataset[T_co]):
self._cam_key_mapping, self._cam_key_mapping,
generate_subtask_ratio=generate_subtask_ratio, generate_subtask_ratio=generate_subtask_ratio,
) )
text = process_grounding_points(complete_text, h, w, resize_h, resize_w, self.data_config.model_type) text = process_grounding_points(
complete_text, h, w, resize_h, resize_w, self.data_config.model_type
)
result = { result = {
"image_inputs": image_inputs, "image_inputs": image_inputs,
"text": text, "text": text,
@@ -150,7 +159,9 @@ class PreprocessedDataset(Dataset[T_co]):
batch_size=batch_size, batch_size=batch_size,
sampler=sampler, # Use distributed sampler instead of shuffle=True sampler=sampler, # Use distributed sampler instead of shuffle=True
num_workers=num_workers, num_workers=num_workers,
collate_fn=DataCollator(self.config, self.dataload_config, self._dataset.meta.stats), collate_fn=DataCollator(
self.config, self.dataload_config, self._dataset.meta.stats
),
pin_memory=True, # Enable for GPU training pin_memory=True, # Enable for GPU training
persistent_workers=num_workers > 0, # Only if num_workers > 0 persistent_workers=num_workers > 0, # Only if num_workers > 0
prefetch_factor=2, # Reduce memory usage prefetch_factor=2, # Reduce memory usage
@@ -164,7 +175,9 @@ class PreprocessedDataset(Dataset[T_co]):
Get distributed evaluation dataloader (no shuffling for consistent evaluation) Get distributed evaluation dataloader (no shuffling for consistent evaluation)
""" """
batch_size = self.config.get("eval_batch_size_per_gpu", self.config.get("batch_size_per_gpu", 8)) batch_size = self.config.get(
"eval_batch_size_per_gpu", self.config.get("batch_size_per_gpu", 8)
)
num_workers = self.config.get("num_workers", 4) num_workers = self.config.get("num_workers", 4)
# Create distributed sampler for evaluation (no shuffle) # Create distributed sampler for evaluation (no shuffle)
@@ -181,7 +194,9 @@ class PreprocessedDataset(Dataset[T_co]):
batch_size=batch_size, batch_size=batch_size,
sampler=sampler, sampler=sampler,
num_workers=num_workers, num_workers=num_workers,
collate_fn=DataCollator(self.config, self.dataload_config, self._dataset.meta.stats), collate_fn=DataCollator(
self.config, self.dataload_config, self._dataset.meta.stats
),
pin_memory=True, pin_memory=True,
persistent_workers=num_workers > 0, persistent_workers=num_workers > 0,
prefetch_factor=2, prefetch_factor=2,
@@ -212,19 +227,30 @@ class DataCollator:
# Use cached processors if available # Use cached processors if available
if processor_path not in self._processor_cache: if processor_path not in self._processor_cache:
self._processor_cache[processor_path] = AutoProcessor.from_pretrained(processor_path, use_fast=True) self._processor_cache[processor_path] = AutoProcessor.from_pretrained(
processor_path, use_fast=True
)
if self.config.get("padding_side", "left") == "left": if self.config.get("padding_side", "left") == "left":
self._processor_cache[processor_path].tokenizer.padding_side = "left" self._processor_cache[processor_path].tokenizer.padding_side = "left"
if self.use_fast_tokenizer and action_tokenizer_path not in self._action_tokenizer_cache: if (
self._action_tokenizer_cache[action_tokenizer_path] = AutoProcessor.from_pretrained(action_tokenizer_path, trust_remote_code=True) self.use_fast_tokenizer
and action_tokenizer_path not in self._action_tokenizer_cache
):
self._action_tokenizer_cache[action_tokenizer_path] = (
AutoProcessor.from_pretrained(
action_tokenizer_path, trust_remote_code=True
)
)
self.processor = self._processor_cache[processor_path] self.processor = self._processor_cache[processor_path]
if not self.use_fast_tokenizer: if not self.use_fast_tokenizer:
self.train_action_tokenizer = None self.train_action_tokenizer = None
else: else:
self.train_action_tokenizer = self._action_tokenizer_cache[action_tokenizer_path] self.train_action_tokenizer = self._action_tokenizer_cache[
action_tokenizer_path
]
if self.use_fast_tokenizer: if self.use_fast_tokenizer:
self.action_mapper = {} self.action_mapper = {}
@@ -254,9 +280,27 @@ class DataCollator:
agent_pos.nan_to_num_(nan=0.0) agent_pos.nan_to_num_(nan=0.0)
agent_pos = self._normalize(agent_pos, self.min_stat, self.delta) agent_pos = self._normalize(agent_pos, self.min_stat, self.delta)
if agent_pos.shape[-1] != 20: if agent_pos.shape[-1] != 20:
agent_pos = torch.cat([agent_pos, torch.zeros(agent_pos.shape[0], agent_pos.shape[1], 20 - agent_pos.shape[-1])], dim=-1) agent_pos = torch.cat(
[
agent_pos,
torch.zeros(
agent_pos.shape[0],
agent_pos.shape[1],
20 - agent_pos.shape[-1],
),
],
dim=-1,
)
agent_pos_mask = torch.cat( agent_pos_mask = torch.cat(
[agent_pos_mask, torch.zeros(agent_pos_mask.shape[0], agent_pos_mask.shape[1], 20 - agent_pos_mask.shape[-1])], dim=-1 [
agent_pos_mask,
torch.zeros(
agent_pos_mask.shape[0],
agent_pos_mask.shape[1],
20 - agent_pos_mask.shape[-1],
),
],
dim=-1,
) )
additional_inputs["proprioception"] = agent_pos additional_inputs["proprioception"] = agent_pos
additional_inputs["agent_pos_mask"] = agent_pos_mask additional_inputs["agent_pos_mask"] = agent_pos_mask
@@ -268,18 +312,42 @@ class DataCollator:
action.nan_to_num_(nan=0.0) action.nan_to_num_(nan=0.0)
action = self._normalize(action, self.min_stat, self.delta) action = self._normalize(action, self.min_stat, self.delta)
if action.shape[-1] != 20: if action.shape[-1] != 20:
action = torch.cat([action, torch.zeros(action.shape[0], action.shape[1], 20 - action.shape[-1])], dim=-1) action = torch.cat(
dof_mask = torch.cat([dof_mask, torch.zeros(dof_mask.shape[0], dof_mask.shape[1], 20 - dof_mask.shape[-1])], dim=-1) [
action,
torch.zeros(
action.shape[0], action.shape[1], 20 - action.shape[-1]
),
],
dim=-1,
)
dof_mask = torch.cat(
[
dof_mask,
torch.zeros(
dof_mask.shape[0],
dof_mask.shape[1],
20 - dof_mask.shape[-1],
),
],
dim=-1,
)
additional_inputs["action_chunk"] = action additional_inputs["action_chunk"] = action
additional_inputs["dof_mask"] = dof_mask additional_inputs["dof_mask"] = dof_mask
elif key == "image_inputs": elif key == "image_inputs":
additional_inputs["image_inputs"] = [item["image_inputs"] for item in batch] additional_inputs["image_inputs"] = [
item["image_inputs"] for item in batch
]
elif key == "text": elif key == "text":
additional_inputs["text"] = [item["text"] for item in batch] additional_inputs["text"] = [item["text"] for item in batch]
elif key == "frame_index": elif key == "frame_index":
additional_inputs["frame_index"] = torch.stack([item["frame_index"] for item in batch]) additional_inputs["frame_index"] = torch.stack(
[item["frame_index"] for item in batch]
)
else: else:
raise NotImplementedError(f"{key} input not implemented in preprocesser") raise NotImplementedError(
f"{key} input not implemented in preprocesser"
)
additional_inputs["text"] = replace_action_token( additional_inputs["text"] = replace_action_token(
additional_inputs["text"], additional_inputs["text"],
@@ -342,20 +410,27 @@ def load_lerobot_data(
delta_timestamps = { delta_timestamps = {
# action chunk # action chunk
"action": [t / dataset_fps for t in range(dataload_config.get("action_horizon", 32) - 1)], "action": [
t / dataset_fps
for t in range(dataload_config.get("action_horizon", 32) - 1)
],
} }
batch_size = config.get("batch_size_per_gpu", 8) batch_size = config.get("batch_size_per_gpu", 8)
# repo_id = "lerobot/aloha_mobile_cabinet" # repo_id = "lerobot/aloha_mobile_cabinet"
repo_id = lerobot_config.get("repo_id", "lerobot/aloha_mobile_cabinet") repo_id = lerobot_config.get("repo_id", "lerobot/aloha_mobile_cabinet")
dataset = LeRobotDataset(repo_id, delta_timestamps=delta_timestamps, video_backend="pyav") dataset = LeRobotDataset(
repo_id, delta_timestamps=delta_timestamps, video_backend="pyav"
)
if rank == 0: if rank == 0:
print(f"Selected episodes: {dataset.episodes}") print(f"Selected episodes: {dataset.episodes}")
print(f"Number of episodes selected: {dataset.num_episodes}") print(f"Number of episodes selected: {dataset.num_episodes}")
print(f"Number of frames selected: {dataset.num_frames}") print(f"Number of frames selected: {dataset.num_frames}")
dataset = PreprocessedDataset(dataset, config, dataload_config, seed=seed, rank=rank, world_size=world_size) dataset = PreprocessedDataset(
dataset, config, dataload_config, seed=seed, rank=rank, world_size=world_size
)
# Calculate samples per process # Calculate samples per process
if world_size > 1: if world_size > 1:
@@ -385,7 +460,9 @@ def load_lerobot_data(
return dataset, train_num return dataset, train_num
def get_distributed_dataloader(dataset, config, rank=0, world_size=1, seed=42, is_train=True): def get_distributed_dataloader(
dataset, config, rank=0, world_size=1, seed=42, is_train=True
):
""" """
Helper function to get distributed dataloader Helper function to get distributed dataloader
@@ -429,9 +506,12 @@ def get_data_configs(config):
return data_config return data_config
class TestDataset(PreprocessedDataset): class TestDataset(PreprocessedDataset):
def __init__(self, dataset, config, dataload_config, seed=42): def __init__(self, dataset, config, dataload_config, seed=42):
super().__init__(dataset, config, dataload_config, seed=seed, rank=0, world_size=1) super().__init__(
dataset, config, dataload_config, seed=seed, rank=0, world_size=1
)
def get_dataloader(self): def get_dataloader(self):
""" """
@@ -441,11 +521,14 @@ class TestDataset(PreprocessedDataset):
dataloader = torch.utils.data.DataLoader( dataloader = torch.utils.data.DataLoader(
self, self,
batch_size=1, batch_size=1,
collate_fn=DataCollator(self.config, self.dataload_config, self._dataset.meta.stats), collate_fn=DataCollator(
self.config, self.dataload_config, self._dataset.meta.stats
),
) )
return dataloader return dataloader
def load_test_dataset( def load_test_dataset(
config, config,
lerobot_config, lerobot_config,
@@ -471,11 +554,19 @@ def load_test_dataset(
delta_timestamps = { delta_timestamps = {
# action chunk # action chunk
"action": [t / dataset_fps for t in range(dataload_config.get("action_horizon", 32) - 1)], "action": [
t / dataset_fps
for t in range(dataload_config.get("action_horizon", 32) - 1)
],
} }
repo_id = lerobot_config.get("repo_id", "lerobot/aloha_mobile_cabinet") repo_id = lerobot_config.get("repo_id", "lerobot/aloha_mobile_cabinet")
dataset = LeRobotDataset(repo_id, episodes=[episode], delta_timestamps=delta_timestamps, video_backend="pyav") dataset = LeRobotDataset(
repo_id,
episodes=[episode],
delta_timestamps=delta_timestamps,
video_backend="pyav",
)
print(f"Selected episodes: {dataset.episodes}") print(f"Selected episodes: {dataset.episodes}")
print(f"Number of episodes selected: {dataset.num_episodes}") print(f"Number of episodes selected: {dataset.num_episodes}")
+74 -23
View File
@@ -137,9 +137,11 @@ def preprocesser_call(
while "<|image_pad|>" in text[i]: while "<|image_pad|>" in text[i]:
# Add bounds checking to avoid index overflow # Add bounds checking to avoid index overflow
if index >= len(image_grid_thw): if index >= len(image_grid_thw):
print(f"Warning: Number of image placeholders ({index + 1}) " print(
f"Warning: Number of image placeholders ({index + 1}) "
f"exceeds actual images ({len(image_grid_thw)}), " f"exceeds actual images ({len(image_grid_thw)}), "
f"skipping remaining placeholder processing") f"skipping remaining placeholder processing"
)
break break
# Replace image placeholder with actual token count # Replace image placeholder with actual token count
token_count = image_grid_thw[index].prod() // merge_length token_count = image_grid_thw[index].prod() // merge_length
@@ -169,7 +171,7 @@ def preprocesser_call(
return_tensors=return_tensors, return_tensors=return_tensors,
padding=padding, padding=padding,
truncation=truncation, truncation=truncation,
max_length=max_length max_length=max_length,
) )
# Get pad token ID for label generation # Get pad token ID for label generation
@@ -209,9 +211,9 @@ def preprocesser_call(
# From second part onwards, each part starts with assistant response # From second part onwards, each part starts with assistant response
for k in range(current_pos + 1, len(text_inputs.input_ids[i])): for k in range(current_pos + 1, len(text_inputs.input_ids[i])):
if text_inputs.input_ids[i][k] == im_end_token_id: if text_inputs.input_ids[i][k] == im_end_token_id:
assistant_regions.append(( assistant_regions.append(
current_pos + len(assistant_tokens), k + 2 (current_pos + len(assistant_tokens), k + 2)
)) )
break break
current_pos += len(part_tokens) + 3 current_pos += len(part_tokens) + 3
@@ -235,7 +237,14 @@ def preprocesser_call(
return BatchFeature(data={**text_inputs, **image_inputs, **videos_inputs}) return BatchFeature(data={**text_inputs, **image_inputs, **videos_inputs})
def process_grounding_points(text: str, orig_height: int, orig_width: int, resized_height: int, resized_width: int, model_type: str) -> str: def process_grounding_points(
text: str,
orig_height: int,
orig_width: int,
resized_height: int,
resized_width: int,
model_type: str,
) -> str:
"""Process grounding point coordinates in text based on image resizing. """Process grounding point coordinates in text based on image resizing.
Adjusts coordinate values in <point> tags to match resized image dimensions Adjusts coordinate values in <point> tags to match resized image dimensions
@@ -309,7 +318,9 @@ def process_grounding_points(text: str, orig_height: int, orig_width: int, resiz
def get_frame_instruction( def get_frame_instruction(
instruction_info: Dict[str, Any], frame_idx: Optional[int] = None, truncate_keys: Optional[List[str]] = None instruction_info: Dict[str, Any],
frame_idx: Optional[int] = None,
truncate_keys: Optional[List[str]] = None,
) -> Tuple[Dict[str, Any], Optional[int]]: ) -> Tuple[Dict[str, Any], Optional[int]]:
"""Extract frame-specific instruction from instruction dictionary. """Extract frame-specific instruction from instruction dictionary.
@@ -322,7 +333,12 @@ def get_frame_instruction(
Tuple of (frame_instruction_dict, split_end_frame) Tuple of (frame_instruction_dict, split_end_frame)
""" """
if truncate_keys is None: if truncate_keys is None:
truncate_keys = ["subtask_generation", "distribute", "subtask_generation_zh", "distribute_zh"] truncate_keys = [
"subtask_generation",
"distribute",
"subtask_generation_zh",
"distribute_zh",
]
instruction_for_frame = {} instruction_for_frame = {}
split_end = None split_end = None
@@ -334,7 +350,11 @@ def get_frame_instruction(
start_frame, end_frame = map(int, frame_range.split(" ")) start_frame, end_frame = map(int, frame_range.split(" "))
if start_frame <= frame_idx < end_frame or (start_frame == frame_idx): if start_frame <= frame_idx < end_frame or (start_frame == frame_idx):
instruction_for_frame[key] = frame_instruction instruction_for_frame[key] = frame_instruction
if truncate_keys is not None and split_end is None and key in truncate_keys: if (
truncate_keys is not None
and split_end is None
and key in truncate_keys
):
split_end = end_frame + 1 split_end = end_frame + 1
break break
else: else:
@@ -343,7 +363,9 @@ def get_frame_instruction(
return instruction_for_frame, split_end return instruction_for_frame, split_end
def get_task_instruction(frame_instruction_info: Dict[str, Any], priority_order: Optional[OrderedDict] = None) -> str: def get_task_instruction(
frame_instruction_info: Dict[str, Any], priority_order: Optional[OrderedDict] = None
) -> str:
"""Construct task instruction from available instruction fields using priority sampling. """Construct task instruction from available instruction fields using priority sampling.
Args: Args:
@@ -428,7 +450,9 @@ def get_wallx_normal_text(
action_fast_symbol = "<|action_fast|>" action_fast_symbol = "<|action_fast|>"
# System prologue # System prologue
prologue = f"{role_start_symbol}system\nYou are a helpful assistant.{role_end_symbol}\n" prologue = (
f"{role_start_symbol}system\nYou are a helpful assistant.{role_end_symbol}\n"
)
# User request with observation # User request with observation
user_request = f"{role_start_symbol}user\nObservation:" user_request = f"{role_start_symbol}user\nObservation:"
@@ -439,13 +463,18 @@ def get_wallx_normal_text(
user_request += "\nInstruction:" user_request += "\nInstruction:"
# Get frame-specific instruction # Get frame-specific instruction
frame_instruction_info, _ = get_frame_instruction(instruction_info, frame_idx=frame_idx) frame_instruction_info, _ = get_frame_instruction(
instruction_info, frame_idx=frame_idx
)
generate_subtask = False generate_subtask = False
priority_keys = ["subtask_generation", "distribute"] priority_keys = ["subtask_generation", "distribute"]
# Decide whether to generate subtask or actions # Decide whether to generate subtask or actions
if bool(set(frame_instruction_info.keys()) & set(priority_keys)) and random.random() < generate_subtask_ratio: if (
bool(set(frame_instruction_info.keys()) & set(priority_keys))
and random.random() < generate_subtask_ratio
):
# Generate subtask (equivalent to VQA task) # Generate subtask (equivalent to VQA task)
instruction = frame_instruction_info.get("instruction", "") instruction = frame_instruction_info.get("instruction", "")
text_prompt = "\nPredict the next action in language.\n" text_prompt = "\nPredict the next action in language.\n"
@@ -457,11 +486,15 @@ def get_wallx_normal_text(
output_instruction = frame_instruction_info[key] output_instruction = frame_instruction_info[key]
break break
assistant_output = f"{role_start_symbol}assistant\n{output_instruction}\n{role_end_symbol}" assistant_output = (
f"{role_start_symbol}assistant\n{output_instruction}\n{role_end_symbol}"
)
generate_subtask = True generate_subtask = True
else: else:
# Generate actions # Generate actions
instruction = get_task_instruction(frame_instruction_info, priority_order=priority_order) instruction = get_task_instruction(
frame_instruction_info, priority_order=priority_order
)
text_prompt = f"\nPredict the next action in robot action.\nProprioception: {propri_symbol}\n" text_prompt = f"\nPredict the next action in robot action.\nProprioception: {propri_symbol}\n"
user_message = f"{user_request} {instruction}{text_prompt}{role_end_symbol}\n" user_message = f"{user_request} {instruction}{text_prompt}{role_end_symbol}\n"
assistant_output = f"{role_start_symbol}assistant\n{action_fast_symbol}{role_end_symbol}\n{action_symbol * action_chunk_size}" assistant_output = f"{role_start_symbol}assistant\n{action_fast_symbol}{role_end_symbol}\n{action_symbol * action_chunk_size}"
@@ -470,7 +503,9 @@ def get_wallx_normal_text(
return complete_text, generate_subtask return complete_text, generate_subtask
def get_action_tokens(normalized_actions: Union[torch.Tensor, List], action_tokenizer) -> List[List[str]]: def get_action_tokens(
normalized_actions: Union[torch.Tensor, List], action_tokenizer
) -> List[List[str]]:
"""Convert normalized actions to action token strings. """Convert normalized actions to action token strings.
Args: Args:
@@ -495,7 +530,9 @@ def get_action_tokens(normalized_actions: Union[torch.Tensor, List], action_toke
return all_action_tokens return all_action_tokens
def pad_action_token_strs(actions_token_lists: List[List[str]], pad_token: str = "<|endoftext|>") -> List[str]: def pad_action_token_strs(
actions_token_lists: List[List[str]], pad_token: str = "<|endoftext|>"
) -> List[str]:
"""Pad action token lists to same length and join as strings. """Pad action token lists to same length and join as strings.
Args: Args:
@@ -509,14 +546,20 @@ def pad_action_token_strs(actions_token_lists: List[List[str]], pad_token: str =
padded_action_strs = [] padded_action_strs = []
for tokens in actions_token_lists: for tokens in actions_token_lists:
padded_tokens = tokens + ["<|im_end|>\n"] + [pad_token] * (max_len - len(tokens)) padded_tokens = (
tokens + ["<|im_end|>\n"] + [pad_token] * (max_len - len(tokens))
)
padded_action_strs.append("".join(padded_tokens)) padded_action_strs.append("".join(padded_tokens))
return padded_action_strs return padded_action_strs
def replace_action_token( def replace_action_token(
text: List[str], norm_action: Optional[torch.Tensor], action_tokenizer, dataset_names: List[str], dof_masks: Optional[torch.Tensor] = None text: List[str],
norm_action: Optional[torch.Tensor],
action_tokenizer,
dataset_names: List[str],
dof_masks: Optional[torch.Tensor] = None,
) -> List[str]: ) -> List[str]:
"""Replace action placeholders in text with actual action tokens. """Replace action placeholders in text with actual action tokens.
@@ -531,14 +574,19 @@ def replace_action_token(
List of text strings with action tokens replaced List of text strings with action tokens replaced
""" """
# Filter out multimodal dataset names # Filter out multimodal dataset names
dataset_names = [name for name in dataset_names if name not in MULTIMODAL_DATASET_NAMES] dataset_names = [
name for name in dataset_names if name not in MULTIMODAL_DATASET_NAMES
]
# Get required action chunk sizes # Get required action chunk sizes
required_chunk_sizes = [FREQUENCY_MAPPING.get(name, 32) for name in dataset_names] required_chunk_sizes = [FREQUENCY_MAPPING.get(name, 32) for name in dataset_names]
if action_tokenizer is not None and norm_action is not None: if action_tokenizer is not None and norm_action is not None:
# Extract actions based on chunk sizes and DOF masks # Extract actions based on chunk sizes and DOF masks
norm_action = [action[: required_chunk_sizes[i], dof_masks[i, 0].bool()] for i, action in enumerate(norm_action)] norm_action = [
action[: required_chunk_sizes[i], dof_masks[i, 0].bool()]
for i, action in enumerate(norm_action)
]
# Convert to action tokens and pad # Convert to action tokens and pad
actions_fast_tokens = get_action_tokens(norm_action, action_tokenizer) actions_fast_tokens = get_action_tokens(norm_action, action_tokenizer)
@@ -548,7 +596,10 @@ def replace_action_token(
actions_fast_token_idx = 0 actions_fast_token_idx = 0
for i in range(len(text)): for i in range(len(text)):
if "<|action_fast|>" in text[i]: if "<|action_fast|>" in text[i]:
text[i] = text[i].replace("<|action_fast|><|im_end|>\n", actions_fast_token_strs[actions_fast_token_idx]) text[i] = text[i].replace(
"<|action_fast|><|im_end|>\n",
actions_fast_token_strs[actions_fast_token_idx],
)
actions_fast_token_idx += 1 actions_fast_token_idx += 1
# Remove remaining action placeholders # Remove remaining action placeholders
+57 -28
View File
@@ -13,11 +13,12 @@ from typing import Tuple, Optional
import wallx_csrc as backend import wallx_csrc as backend
def _allocate_asymmetric_dual_outputs(
def _allocate_asymmetric_dual_outputs(input_expert0: torch.Tensor, input_expert0: torch.Tensor,
input_expert1: torch.Tensor, input_expert1: torch.Tensor,
weight_expert0: torch.Tensor, weight_expert0: torch.Tensor,
weight_expert1: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]: weight_expert1: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor]:
""" """
Allocate output tensors for asymmetric dual expert GEMM operations. Allocate output tensors for asymmetric dual expert GEMM operations.
@@ -45,30 +46,38 @@ def _allocate_asymmetric_dual_outputs(input_expert0: torch.Tensor,
assert weight_expert1.ndim == 2, "Expected 2D tensor for weight_expert1" assert weight_expert1.ndim == 2, "Expected 2D tensor for weight_expert1"
# Verify dimension compatibility for matrix multiplication # Verify dimension compatibility for matrix multiplication
assert input_expert0.size(1) == weight_expert0.size(0), \ assert input_expert0.size(1) == weight_expert0.size(
f"Input expert0 K dimension {input_expert0.size(1)} != weight expert0 K dimension {weight_expert0.size(0)}" 0
assert input_expert1.size(1) == weight_expert1.size(0), \ ), f"Input expert0 K dimension {input_expert0.size(1)} != weight expert0 K dimension {weight_expert0.size(0)}"
f"Input expert1 K dimension {input_expert1.size(1)} != weight expert1 K dimension {weight_expert1.size(0)}" assert input_expert1.size(1) == weight_expert1.size(
0
), f"Input expert1 K dimension {input_expert1.size(1)} != weight expert1 K dimension {weight_expert1.size(0)}"
# Calculate output shapes: [m, k] × [k, n] = [m, n] # Calculate output shapes: [m, k] × [k, n] = [m, n]
m0, n0 = input_expert0.size(0), weight_expert0.size(1) m0, n0 = input_expert0.size(0), weight_expert0.size(1)
m1, n1 = input_expert1.size(0), weight_expert1.size(1) m1, n1 = input_expert1.size(0), weight_expert1.size(1)
# Allocate output tensors with matching device and dtype # Allocate output tensors with matching device and dtype
output_expert0 = torch.empty(m0, n0, device=input_expert0.device, dtype=input_expert0.dtype) output_expert0 = torch.empty(
output_expert1 = torch.empty(m1, n1, device=input_expert1.device, dtype=input_expert1.dtype) m0, n0, device=input_expert0.device, dtype=input_expert0.dtype
)
output_expert1 = torch.empty(
m1, n1, device=input_expert1.device, dtype=input_expert1.dtype
)
return output_expert0, output_expert1 return output_expert0, output_expert1
def asym_dual_gmm_separated(input_expert0: torch.Tensor, def asym_dual_gmm_separated(
input_expert0: torch.Tensor,
input_expert1: torch.Tensor, input_expert1: torch.Tensor,
weight_expert0: torch.Tensor, weight_expert0: torch.Tensor,
weight_expert1: torch.Tensor, weight_expert1: torch.Tensor,
output_expert0: Optional[torch.Tensor] = None, output_expert0: Optional[torch.Tensor] = None,
output_expert1: Optional[torch.Tensor] = None, output_expert1: Optional[torch.Tensor] = None,
trans_a: bool = False, trans_a: bool = False,
trans_b: bool = False) -> Tuple[torch.Tensor, torch.Tensor]: trans_b: bool = False,
) -> Tuple[torch.Tensor, torch.Tensor]:
""" """
Asymmetric dual expert grouped GEMM with separated inputs and outputs. Asymmetric dual expert grouped GEMM with separated inputs and outputs.
@@ -113,20 +122,26 @@ def asym_dual_gmm_separated(input_expert0: torch.Tensor,
# Call optimized C++ backend kernel # Call optimized C++ backend kernel
backend.asym_dual_gmm( backend.asym_dual_gmm(
input_expert0, input_expert1, input_expert0,
weight_expert0, weight_expert1, input_expert1,
output_expert0, output_expert1, weight_expert0,
trans_a, trans_b weight_expert1,
output_expert0,
output_expert1,
trans_a,
trans_b,
) )
return output_expert0, output_expert1 return output_expert0, output_expert1
def permute(input: torch.Tensor, def permute(
input: torch.Tensor,
indices: torch.Tensor, indices: torch.Tensor,
num_out_tokens: int, num_out_tokens: int,
workspace: torch.Tensor, workspace: torch.Tensor,
max_expanded_token_num: int) -> torch.Tensor: max_expanded_token_num: int,
) -> torch.Tensor:
""" """
Permute input tokens according to expert assignment indices for MoE routing. Permute input tokens according to expert assignment indices for MoE routing.
@@ -147,14 +162,18 @@ def permute(input: torch.Tensor,
This is typically used with top-k expert selection where each token This is typically used with top-k expert selection where each token
can be routed to multiple experts. can be routed to multiple experts.
""" """
return backend.permute(input, indices, num_out_tokens, workspace, max_expanded_token_num) return backend.permute(
input, indices, num_out_tokens, workspace, max_expanded_token_num
)
def unpermute(input: torch.Tensor, def unpermute(
input: torch.Tensor,
row_id_map: torch.Tensor, row_id_map: torch.Tensor,
prob: torch.Tensor, prob: torch.Tensor,
max_tokens: int, max_tokens: int,
num_topK: int) -> torch.Tensor: num_topK: int,
) -> torch.Tensor:
""" """
Unpermute expert outputs back to original token order with probability weighting. Unpermute expert outputs back to original token order with probability weighting.
@@ -178,10 +197,12 @@ def unpermute(input: torch.Tensor,
return backend.unpermute(input, row_id_map, prob, max_tokens, num_topK) return backend.unpermute(input, row_id_map, prob, max_tokens, num_topK)
def unpermute_bwd(input_bwd: torch.Tensor, def unpermute_bwd(
input_bwd: torch.Tensor,
input_fwd: torch.Tensor, input_fwd: torch.Tensor,
row_id_map: torch.Tensor, row_id_map: torch.Tensor,
prob: Optional[torch.Tensor]) -> torch.Tensor: prob: Optional[torch.Tensor],
) -> torch.Tensor:
""" """
Backward pass for unpermute operation with gradient flow. Backward pass for unpermute operation with gradient flow.
@@ -202,18 +223,22 @@ def unpermute_bwd(input_bwd: torch.Tensor,
""" """
# Handle case where probabilities are not provided # Handle case where probabilities are not provided
if prob is None: if prob is None:
prob = torch.ones([input_bwd.size(0), 1], dtype=torch.float32, device=input_bwd.device) prob = torch.ones(
[input_bwd.size(0), 1], dtype=torch.float32, device=input_bwd.device
)
return backend.unpermute_bwd(input_bwd, input_fwd, row_id_map, prob) return backend.unpermute_bwd(input_bwd, input_fwd, row_id_map, prob)
def rope(q: torch.Tensor, def rope(
q: torch.Tensor,
k: torch.Tensor, k: torch.Tensor,
cos: torch.Tensor, cos: torch.Tensor,
sin: torch.Tensor, sin: torch.Tensor,
q_out: torch.Tensor, q_out: torch.Tensor,
k_out: torch.Tensor, k_out: torch.Tensor,
mrope_section_doubled: bool) -> None: mrope_section_doubled: bool,
) -> None:
""" """
Apply RoPE (Rotary Position Embedding) to query and key tensors. Apply RoPE (Rotary Position Embedding) to query and key tensors.
@@ -237,7 +262,8 @@ def rope(q: torch.Tensor,
return backend.rope(q, k, cos, sin, q_out, k_out, mrope_section_doubled) return backend.rope(q, k, cos, sin, q_out, k_out, mrope_section_doubled)
def rope_bwd(grad_q_out: torch.Tensor, def rope_bwd(
grad_q_out: torch.Tensor,
grad_k_out: torch.Tensor, grad_k_out: torch.Tensor,
q: torch.Tensor, q: torch.Tensor,
k: torch.Tensor, k: torch.Tensor,
@@ -245,7 +271,8 @@ def rope_bwd(grad_q_out: torch.Tensor,
sin: torch.Tensor, sin: torch.Tensor,
grad_q: torch.Tensor, grad_q: torch.Tensor,
grad_k: torch.Tensor, grad_k: torch.Tensor,
mrope_section_doubled: bool) -> None: mrope_section_doubled: bool,
) -> None:
""" """
Backward pass for RoPE operation with gradient computation. Backward pass for RoPE operation with gradient computation.
@@ -267,4 +294,6 @@ def rope_bwd(grad_q_out: torch.Tensor,
This function computes the analytical gradient of the RoPE operation, This function computes the analytical gradient of the RoPE operation,
which involves the inverse rotation compared to the forward pass. which involves the inverse rotation compared to the forward pass.
""" """
return backend.rope_bwd(grad_q_out, grad_k_out, q, k, cos, sin, grad_q, grad_k, mrope_section_doubled) return backend.rope_bwd(
grad_q_out, grad_k_out, q, k, cos, sin, grad_q, grad_k, mrope_section_doubled
)
+137 -37
View File
@@ -5,7 +5,9 @@ from wall_x.fusions import backend
class AsymmetricDualExpertGemm(torch.autograd.Function): class AsymmetricDualExpertGemm(torch.autograd.Function):
@staticmethod @staticmethod
def forward(ctx, input_expert0, input_expert1, weight_expert0, weight_expert1, trans_b=False): def forward(
ctx, input_expert0, input_expert1, weight_expert0, weight_expert1, trans_b=False
):
""" """
Forward pass for asymmetric dual expert GEMM. Forward pass for asymmetric dual expert GEMM.
@@ -27,14 +29,24 @@ class AsymmetricDualExpertGemm(torch.autograd.Function):
# Dimension validation depends on trans_b # Dimension validation depends on trans_b
if trans_b: if trans_b:
assert input_expert0.size(1) == weight_expert0.size(1), "Expert 0 dimension mismatch (trans_b=True)" assert input_expert0.size(1) == weight_expert0.size(
assert input_expert1.size(1) == weight_expert1.size(1), "Expert 1 dimension mismatch (trans_b=True)" 1
), "Expert 0 dimension mismatch (trans_b=True)"
assert input_expert1.size(1) == weight_expert1.size(
1
), "Expert 1 dimension mismatch (trans_b=True)"
else: else:
assert input_expert0.size(1) == weight_expert0.size(0), "Expert 0 dimension mismatch (trans_b=False)" assert input_expert0.size(1) == weight_expert0.size(
assert input_expert1.size(1) == weight_expert1.size(0), "Expert 1 dimension mismatch (trans_b=False)" 0
), "Expert 0 dimension mismatch (trans_b=False)"
assert input_expert1.size(1) == weight_expert1.size(
0
), "Expert 1 dimension mismatch (trans_b=False)"
# Save tensors and trans_b for backward pass # Save tensors and trans_b for backward pass
ctx.save_for_backward(input_expert0, input_expert1, weight_expert0, weight_expert1) ctx.save_for_backward(
input_expert0, input_expert1, weight_expert0, weight_expert1
)
ctx.trans_b = trans_b ctx.trans_b = trans_b
# Allocate output tensors # Allocate output tensors
@@ -43,11 +55,23 @@ class AsymmetricDualExpertGemm(torch.autograd.Function):
n0 = weight_expert0.size(0) if trans_b else weight_expert0.size(1) n0 = weight_expert0.size(0) if trans_b else weight_expert0.size(1)
n1 = weight_expert1.size(0) if trans_b else weight_expert1.size(1) n1 = weight_expert1.size(0) if trans_b else weight_expert1.size(1)
output_expert0 = torch.empty(m0, n0, device=input_expert0.device, dtype=input_expert0.dtype) output_expert0 = torch.empty(
output_expert1 = torch.empty(m1, n1, device=input_expert1.device, dtype=input_expert1.dtype) m0, n0, device=input_expert0.device, dtype=input_expert0.dtype
)
output_expert1 = torch.empty(
m1, n1, device=input_expert1.device, dtype=input_expert1.dtype
)
# Call the backend C++ function # Call the backend C++ function
backend.asym_dual_gmm_separated(input_expert0, input_expert1, weight_expert0, weight_expert1, output_expert0, output_expert1, trans_b=trans_b) backend.asym_dual_gmm_separated(
input_expert0,
input_expert1,
weight_expert0,
weight_expert1,
output_expert0,
output_expert1,
trans_b=trans_b,
)
return output_expert0, output_expert1 return output_expert0, output_expert1
@@ -111,10 +135,18 @@ class AsymmetricDualExpertGemm(torch.autograd.Function):
trans_b=False, trans_b=False,
) )
return grad_input_expert0, grad_input_expert1, grad_weight_expert0, grad_weight_expert1, None return (
grad_input_expert0,
grad_input_expert1,
grad_weight_expert0,
grad_weight_expert1,
None,
)
def asym_dual_gmm(input_expert0, input_expert1, weight_expert0, weight_expert1, trans_b=False): def asym_dual_gmm(
input_expert0, input_expert1, weight_expert0, weight_expert1, trans_b=False
):
""" """
Convenience function for asymmetric dual expert GEMM. Convenience function for asymmetric dual expert GEMM.
@@ -128,7 +160,9 @@ def asym_dual_gmm(input_expert0, input_expert1, weight_expert0, weight_expert1,
Returns: Returns:
Tuple of (output_expert0, output_expert1) Tuple of (output_expert0, output_expert1)
""" """
return AsymmetricDualExpertGemm.apply(input_expert0, input_expert1, weight_expert0, weight_expert1, trans_b) return AsymmetricDualExpertGemm.apply(
input_expert0, input_expert1, weight_expert0, weight_expert1, trans_b
)
################################################################################################ ################################################################################################
@@ -145,7 +179,13 @@ class PermuteMoE_topK(torch.autograd.Function):
max_expanded_token_num = 0 max_expanded_token_num = 0
@staticmethod @staticmethod
def forward(ctx, input_act: torch.Tensor, indices: torch.Tensor, num_out_tokens: int, max_token_num: int): def forward(
ctx,
input_act: torch.Tensor,
indices: torch.Tensor,
num_out_tokens: int,
max_token_num: int,
):
""" """
indices: for topK=1, indices in a 1-d tensor of shape [num_tokens], indices: for topK=1, indices in a 1-d tensor of shape [num_tokens],
otherwise, it's a 2-d tensor of shape [num_tokens, topK] otherwise, it's a 2-d tensor of shape [num_tokens, topK]
@@ -160,18 +200,27 @@ class PermuteMoE_topK(torch.autograd.Function):
# Device check # Device check
if input_act.is_cpu: if input_act.is_cpu:
raise RuntimeError("[Error] The input `input_act` of permute_topK op is on the device: CPU!") raise RuntimeError(
"[Error] The input `input_act` of permute_topK op is on the device: CPU!"
)
if indices.is_cpu: if indices.is_cpu:
warnings.warn("The input `indices` of permute_topK op is on the device: CPU!") warnings.warn(
expert_for_rows = expert_for_rows.cuda() "The input `indices` of permute_topK op is on the device: CPU!"
)
# Shape check # Shape check
if input_act.size(0) != indices.size(0): if input_act.size(0) != indices.size(0):
raise RuntimeError(f"[Error] permute_topK op input `indices` shape mismatch! " f"Expect {input_act.size(0)}, but got {indices.size(0)}.") raise RuntimeError(
f"[Error] permute_topK op input `indices` shape mismatch! "
f"Expect {input_act.size(0)}, but got {indices.size(0)}."
)
# Data type check # Data type check
if indices.dtype != torch.int32: if indices.dtype != torch.int32:
warnings.warn(f"The data type of the input `indices` of permute_topK op is {indices.dtype}! " "The recommended type is torch.int32.") warnings.warn(
f"The data type of the input `indices` of permute_topK op is {indices.dtype}! "
"The recommended type is torch.int32."
)
indices = indices.to(torch.int32) indices = indices.to(torch.int32)
# Contiguous check # Contiguous check
@@ -194,7 +243,11 @@ class PermuteMoE_topK(torch.autograd.Function):
PermuteMoE_topK.workspace_fw = [] PermuteMoE_topK.workspace_fw = []
permuted_act, row_id_map, PermuteMoE_topK.workspace_fw = backend.permute( permuted_act, row_id_map, PermuteMoE_topK.workspace_fw = backend.permute(
input_act, indices, num_out_tokens, PermuteMoE_topK.workspace_fw, PermuteMoE_topK.max_expanded_token_num input_act,
indices,
num_out_tokens,
PermuteMoE_topK.workspace_fw,
PermuteMoE_topK.max_expanded_token_num,
) )
ctx.row_id_map = row_id_map ctx.row_id_map = row_id_map
@@ -215,7 +268,9 @@ class PermuteMoE_topK(torch.autograd.Function):
num_tokens = ctx.num_tokens num_tokens = ctx.num_tokens
num_topK = ctx.num_topK num_topK = ctx.num_topK
unpermuted_act_grad = backend.unpermute(permuted_act_grad, row_id_map, torch.tensor([]), num_tokens, num_topK) unpermuted_act_grad = backend.unpermute(
permuted_act_grad, row_id_map, torch.tensor([]), num_tokens, num_topK
)
return unpermuted_act_grad, None, None, None return unpermuted_act_grad, None, None, None
@@ -229,7 +284,12 @@ class PermuteMoE_topK(torch.autograd.Function):
class UnpermuteMoE_topK(torch.autograd.Function): class UnpermuteMoE_topK(torch.autograd.Function):
@staticmethod @staticmethod
def forward(ctx, input_act: torch.Tensor, row_id_map: torch.Tensor, probs: torch.Tensor = None): def forward(
ctx,
input_act: torch.Tensor,
row_id_map: torch.Tensor,
probs: torch.Tensor = None,
):
# Empty input check # Empty input check
if not input_act.numel(): if not input_act.numel():
ctx.probs = probs ctx.probs = probs
@@ -237,36 +297,51 @@ class UnpermuteMoE_topK(torch.autograd.Function):
# Device check # Device check
if input_act.is_cpu: if input_act.is_cpu:
raise RuntimeError("[Error] The input `input_act` of unpermute_topK op is on the device: CPU!") raise RuntimeError(
"[Error] The input `input_act` of unpermute_topK op is on the device: CPU!"
)
if row_id_map.is_cpu: if row_id_map.is_cpu:
warnings.warn("The input `row_id_map` of unpermute_topK op is on the device: CPU!") warnings.warn(
"The input `row_id_map` of unpermute_topK op is on the device: CPU!"
)
row_id_map = row_id_map.cuda() row_id_map = row_id_map.cuda()
if probs is not None and probs.is_cpu: if probs is not None and probs.is_cpu:
warnings.warn("The input `probs` of unpermute_topK op is on the device: CPU!") warnings.warn(
"The input `probs` of unpermute_topK op is on the device: CPU!"
)
probs = probs.cuda() probs = probs.cuda()
# Shape check # Shape check
if probs is not None and row_id_map.size(0) != probs.size(0) * probs.size(1): if probs is not None and row_id_map.size(0) != probs.size(0) * probs.size(1):
raise RuntimeError( raise RuntimeError(
f"[Error] unpermute_topK op input `probs` shape mismatch! " f"Expect {row_id_map.size(0)}, but got {probs.size(0) * probs.size(1)}." f"[Error] unpermute_topK op input `probs` shape mismatch! "
f"Expect {row_id_map.size(0)}, but got {probs.size(0) * probs.size(1)}."
) )
# Data type check # Data type check
if row_id_map.dtype != torch.int32: if row_id_map.dtype != torch.int32:
warnings.warn( warnings.warn(
f"The data type of the input `row_id_map` of unpermute_topK op is {row_id_map.dtype}! " "The recommended type is torch.int32." f"The data type of the input `row_id_map` of unpermute_topK op is {row_id_map.dtype}! "
"The recommended type is torch.int32."
) )
row_id_map = row_id_map.to(torch.int32) row_id_map = row_id_map.to(torch.int32)
if probs is not None and probs.dtype != torch.float32: if probs is not None and probs.dtype != torch.float32:
warnings.warn(f"The data type of the input `probs` of unpermute_topK op is {probs.dtype}! " "The recommended type is torch.float32.") warnings.warn(
f"The data type of the input `probs` of unpermute_topK op is {probs.dtype}! "
"The recommended type is torch.float32."
)
probs = probs.to(torch.float32) probs = probs.to(torch.float32)
# Contiguous check # Contiguous check
if not input_act.is_contiguous(): if not input_act.is_contiguous():
warnings.warn("The input `input_act` of unpermute_topK op is discontiguous!") warnings.warn(
"The input `input_act` of unpermute_topK op is discontiguous!"
)
input_act = input_act.contiguous() input_act = input_act.contiguous()
if not row_id_map.is_contiguous(): if not row_id_map.is_contiguous():
warnings.warn("The input `row_id_map` of unpermute_topK op is discontiguous!") warnings.warn(
"The input `row_id_map` of unpermute_topK op is discontiguous!"
)
row_id_map = row_id_map.contiguous() row_id_map = row_id_map.contiguous()
if probs is not None and not probs.is_contiguous(): if probs is not None and not probs.is_contiguous():
warnings.warn("The input `probs` of unpermute_topK op is discontiguous!") warnings.warn("The input `probs` of unpermute_topK op is discontiguous!")
@@ -275,7 +350,13 @@ class UnpermuteMoE_topK(torch.autograd.Function):
num_tokens = probs.size(0) if probs is not None else input_act.size(0) num_tokens = probs.size(0) if probs is not None else input_act.size(0)
num_topK = probs.size(1) if probs is not None else 1 num_topK = probs.size(1) if probs is not None else 1
unpermuted_output = backend.unpermute(input_act, row_id_map, probs if probs is not None else torch.tensor([]), num_tokens, num_topK) unpermuted_output = backend.unpermute(
input_act,
row_id_map,
probs if probs is not None else torch.tensor([]),
num_tokens,
num_topK,
)
ctx.save_for_backward(input_act, row_id_map, probs) ctx.save_for_backward(input_act, row_id_map, probs)
return unpermuted_output return unpermuted_output
@@ -293,7 +374,9 @@ class UnpermuteMoE_topK(torch.autograd.Function):
act_grad = None act_grad = None
if ctx.needs_input_grad[0]: if ctx.needs_input_grad[0]:
act_grad, prob_grad = backend.unpermute_bwd(unpermuted_act_grad, input_act, row_id_map, probs) act_grad, prob_grad = backend.unpermute_bwd(
unpermuted_act_grad, input_act, row_id_map, probs
)
if not ctx.needs_input_grad[2]: if not ctx.needs_input_grad[2]:
prob_grad = None prob_grad = None
@@ -319,19 +402,36 @@ def unpermute(input_act, row_id_map, probs=None):
class MultimodalRoPE(torch.autograd.Function): class MultimodalRoPE(torch.autograd.Function):
@staticmethod @staticmethod
def forward(ctx, q: torch.Tensor, k: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor, mrope_section: list): def forward(
ctx,
q: torch.Tensor,
k: torch.Tensor,
cos: torch.Tensor,
sin: torch.Tensor,
mrope_section: list,
):
# Device check # Device check
if q.is_cpu: if q.is_cpu:
raise RuntimeError("[Error] The input `q` of multimodal_rope op is on the device: CPU!") raise RuntimeError(
"[Error] The input `q` of multimodal_rope op is on the device: CPU!"
)
if k.is_cpu: if k.is_cpu:
raise RuntimeError("[Error] The input `k` of multimodal_rope op is on the device: CPU!") raise RuntimeError(
"[Error] The input `k` of multimodal_rope op is on the device: CPU!"
)
if cos.is_cpu: if cos.is_cpu:
raise RuntimeError("[Error] The input `cos` of multimodal_rope op is on the device: CPU!") raise RuntimeError(
"[Error] The input `cos` of multimodal_rope op is on the device: CPU!"
)
if sin.is_cpu: if sin.is_cpu:
raise RuntimeError("[Error] The input `sin` of multimodal_rope op is on the device: CPU!") raise RuntimeError(
"[Error] The input `sin` of multimodal_rope op is on the device: CPU!"
)
if len(mrope_section) != 3: if len(mrope_section) != 3:
raise RuntimeError("[Error] The input `mrope_section` of multimodal_rope op must be a list of 3 integers!") raise RuntimeError(
"[Error] The input `mrope_section` of multimodal_rope op must be a list of 3 integers!"
)
# Contiguous check # Contiguous check
if not q.is_contiguous(): if not q.is_contiguous():
+43 -18
View File
@@ -1,10 +1,10 @@
import math import math
import torch import torch
import torch.nn as nn import torch.nn as nn
from torch.distributions import Beta from torch.distributions import Beta
from wall_x.utils.constant import action_statistic_dof from wall_x.utils.constant import action_statistic_dof
class Normalizer(nn.Module): class Normalizer(nn.Module):
""" """
Action data normalizer for multi-robot systems. Action data normalizer for multi-robot systems.
@@ -48,14 +48,18 @@ class Normalizer(nn.Module):
action_statistic[robot_name]["delta"] = all_dof_delta action_statistic[robot_name]["delta"] = all_dof_delta
# Register statistics as non-trainable parameters # Register statistics as non-trainable parameters
self.min = nn.ParameterDict({ self.min = nn.ParameterDict(
{
k: nn.Parameter(action_statistic[k]["min"], requires_grad=False) k: nn.Parameter(action_statistic[k]["min"], requires_grad=False)
for k in action_statistic.keys() for k in action_statistic.keys()
}) }
self.delta = nn.ParameterDict({ )
self.delta = nn.ParameterDict(
{
k: nn.Parameter(action_statistic[k]["delta"], requires_grad=False) k: nn.Parameter(action_statistic[k]["delta"], requires_grad=False)
for k in action_statistic.keys() for k in action_statistic.keys()
}) }
)
def normalize_data(self, xs, dataset_names): def normalize_data(self, xs, dataset_names):
""" """
@@ -159,7 +163,6 @@ class SinusoidalPosEmb(nn.Module):
return emb return emb
class ActionProcessor(nn.Module): class ActionProcessor(nn.Module):
""" """
Action sequence processor for robotic control with flow matching. Action sequence processor for robotic control with flow matching.
@@ -208,16 +211,22 @@ class ActionProcessor(nn.Module):
# Initialize data normalizers for actions and proprioception # Initialize data normalizers for actions and proprioception
self.normalizer_action = Normalizer(action_statistic_dof, config.dof_config) 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) # Proprioception projection layer (includes history/current state)
self.propri_proj = nn.Linear(self.propri_dim * 2, self.hidden_size, bias=False) self.propri_proj = nn.Linear(self.propri_dim * 2, self.hidden_size, bias=False)
# Beta distribution noise scheduler configuration # Beta distribution noise scheduler configuration
noise_scheduler_config = config.noise_scheduler noise_scheduler_config = config.noise_scheduler
self.beta_alpha = noise_scheduler_config.get('beta_alpha', 1.5) # Beta distribution α parameter self.beta_alpha = noise_scheduler_config.get(
self.beta_beta = noise_scheduler_config.get('beta_beta', 1.0) # Beta distribution β parameter "beta_alpha", 1.5
self.s = noise_scheduler_config.get('s', 0.999) # Scaling factor ) # 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 # Initialize Beta distribution for noise scheduling
alpha_tensor = torch.tensor(self.beta_alpha, dtype=torch.float32).to("cuda") alpha_tensor = torch.tensor(self.beta_alpha, dtype=torch.float32).to("cuda")
@@ -228,14 +237,18 @@ class ActionProcessor(nn.Module):
self.time_embed = SinusoidalPosEmb(config.hidden_size) self.time_embed = SinusoidalPosEmb(config.hidden_size)
# Action embedding network: project to hidden space # 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.w1 = nn.Linear(
self.w2 = nn.Linear(self.hidden_size * 2, self.hidden_size, bias=False) # *2 for action + time embeddings 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.w3 = nn.Linear(self.hidden_size, self.hidden_size, bias=False)
self.act_fn = nn.SiLU() self.act_fn = nn.SiLU()
# Project back to action space for flow matching loss # Project back to action space for flow matching loss
self.action_proj_back = nn.Linear(self.hidden_size, self.action_dim, bias=False) 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): def sample_time(self, batch_size, device, dtype):
""" """
@@ -257,7 +270,9 @@ class ActionProcessor(nn.Module):
time = (self.s - sample) / self.s time = (self.s - sample) / self.s
return time 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. Project proprioceptive data (joint positions, orientations) to hidden space.
@@ -271,7 +286,9 @@ class ActionProcessor(nn.Module):
torch.Tensor: Projected proprioceptive features of shape [batch_size, seq_len, hidden_size] torch.Tensor: Projected proprioceptive features of shape [batch_size, seq_len, hidden_size]
""" """
# Ensure proper device and dtype alignment # 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: if dof_mask is not None:
# Concatenate proprioception with DOF mask # Concatenate proprioception with DOF mask
@@ -281,7 +298,9 @@ class ActionProcessor(nn.Module):
else: else:
proprioception = torch.cat([proprioception, dof_mask], dim=-1) 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) return self.propri_proj(proprioception)
def forward(self, action_chunk, dataset_names, dof_mask=None): def forward(self, action_chunk, dataset_names, dof_mask=None):
@@ -329,7 +348,11 @@ class ActionProcessor(nn.Module):
action_embed = self.w1(noisy_action) action_embed = self.w1(noisy_action)
# Repeat time embedding for each sequence position # 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 # Combine action and temporal embeddings
concat_embed = torch.cat([action_embed, time_embed], dim=-1) concat_embed = torch.cat([action_embed, time_embed], dim=-1)
@@ -364,7 +387,9 @@ class ActionProcessor(nn.Module):
# Broadcast time embeddings to sequence length # Broadcast time embeddings to sequence length
time_embed = time_embed.unsqueeze(1).repeat(1, action_embed.shape[1], 1) 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 # Combine embeddings and process through MLP
concat_embed = torch.cat([action_embed, time_embed], dim=-1) concat_embed = torch.cat([action_embed, time_embed], dim=-1)
+6
View File
@@ -1,2 +1,8 @@
from .modeling_qwen2_5_vl_act import Qwen2_5_VLMoEModel, Qwen2_5_VLMoEForAction from .modeling_qwen2_5_vl_act import Qwen2_5_VLMoEModel, Qwen2_5_VLMoEForAction
from .configuration_qwen2_5_vl import Qwen2_5_VLConfig from .configuration_qwen2_5_vl import Qwen2_5_VLConfig
__all__ = [
"Qwen2_5_VLMoEModel",
"Qwen2_5_VLMoEForAction",
"Qwen2_5_VLConfig",
]
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
+205 -52
View File
@@ -12,13 +12,16 @@ from datetime import datetime
from torch.optim import AdamW from torch.optim import AdamW
from accelerate import Accelerator from accelerate import Accelerator
from safetensors.torch import load_file from safetensors.torch import load_file
from accelerate.utils import DistributedType
from transformers.optimization import get_cosine_with_min_lr_schedule_with_warmup from transformers.optimization import get_cosine_with_min_lr_schedule_with_warmup
from wall_x.utils.timers import Timers from wall_x.utils.timers import Timers
from wall_x.model.qwen2_5_based import Qwen2_5_VLMoEForAction from wall_x.model.qwen2_5_based import Qwen2_5_VLMoEForAction
from wall_x.data.config import ACTION_DATASET_NAMES, MULTIMODAL_DATASET_NAMES from wall_x.data.config import ACTION_DATASET_NAMES, MULTIMODAL_DATASET_NAMES
from wall_x.data.load_lerobot_dataset import PreprocessedDataset, get_data_configs, load_lerobot_data from wall_x.data.load_lerobot_dataset import (
PreprocessedDataset,
get_data_configs,
load_lerobot_data,
)
def timer(func): def timer(func):
@@ -31,6 +34,7 @@ def timer(func):
Returns: Returns:
Wrapped function with timing functionality Wrapped function with timing functionality
""" """
@wraps(func) @wraps(func)
def wrapper(*args, **kwargs): def wrapper(*args, **kwargs):
start_time = time.time() start_time = time.time()
@@ -40,6 +44,7 @@ def timer(func):
f"\033[92m[current time: {time.strftime('%Y-%m-%d %H:%M:%S', time.localtime())}] Function {func.__name__} took {end_time - start_time:.2f} seconds to execute\033[0m" f"\033[92m[current time: {time.strftime('%Y-%m-%d %H:%M:%S', time.localtime())}] Function {func.__name__} took {end_time - start_time:.2f} seconds to execute\033[0m"
) )
return result return result
return wrapper return wrapper
@@ -87,7 +92,14 @@ class QwenVlAct_Trainer:
""" """
@timer @timer
def __init__(self, config, logger, accelerator: Accelerator = None, seed=42, data_config_path=None): def __init__(
self,
config,
logger,
accelerator: Accelerator = None,
seed=42,
data_config_path=None,
):
""" """
Initialize the Vision-Language-Action trainer. Initialize the Vision-Language-Action trainer.
@@ -139,19 +151,29 @@ class QwenVlAct_Trainer:
# Distributed training setup # Distributed training setup
self.rank = self.accelerator.process_index self.rank = self.accelerator.process_index
self.world_size = self.accelerator.num_processes self.world_size = self.accelerator.num_processes
print(f"rank {self.accelerator.process_index} after load model memory usage: {torch.cuda.memory_allocated() / 1024 ** 3:.2f} GB", flush=True) print(
f"rank {self.accelerator.process_index} after load model memory usage: {torch.cuda.memory_allocated() / 1024 ** 3:.2f} GB",
flush=True,
)
# Load training data # Load training data
self.load_qact_data() self.load_qact_data()
print(f"rank {self.accelerator.process_index} after load qact data usage: {torch.cuda.memory_allocated() / 1024 ** 3:.2f} GB", flush=True) print(
f"rank {self.accelerator.process_index} after load qact data usage: {torch.cuda.memory_allocated() / 1024 ** 3:.2f} GB",
flush=True,
)
# Resume from checkpoint if specified # Resume from checkpoint if specified
if "resume" in self.config: if "resume" in self.config:
self.resume_from_checkpoint() self.resume_from_checkpoint()
# Initialize special token IDs # Initialize special token IDs
self.propri_token_id = self.processor.tokenizer.convert_tokens_to_ids("<|propri|>") self.propri_token_id = self.processor.tokenizer.convert_tokens_to_ids(
self.action_token_id = self.processor.tokenizer.convert_tokens_to_ids("<|action|>") "<|propri|>"
)
self.action_token_id = self.processor.tokenizer.convert_tokens_to_ids(
"<|action|>"
)
# Initialize evaluation metrics # Initialize evaluation metrics
self.base_l1_loss = None self.base_l1_loss = None
@@ -162,7 +184,9 @@ class QwenVlAct_Trainer:
# Adjust global step if resuming from checkpoint # Adjust global step if resuming from checkpoint
if self.initial_step != 0: if self.initial_step != 0:
self.global_step = self.initial_step // self.config.get("gradient_accumulation_steps", 1) self.global_step = self.initial_step // self.config.get(
"gradient_accumulation_steps", 1
)
def print_rank0(self, msg, flush=True): def print_rank0(self, msg, flush=True):
""" """
@@ -188,7 +212,9 @@ class QwenVlAct_Trainer:
self.accelerator.wait_for_everyone() self.accelerator.wait_for_everyone()
# Optional validation before training starts # Optional validation before training starts
if self.config.get("resume", None) is not None and self.config["resume"].get("validate_first", False): if self.config.get("resume", None) is not None and self.config["resume"].get(
"validate_first", False
):
self.val_loop() self.val_loop()
self.accelerator.wait_for_everyone() self.accelerator.wait_for_everyone()
@@ -227,7 +253,9 @@ class QwenVlAct_Trainer:
if getattr(self, "train_dataloader", None) is not None: if getattr(self, "train_dataloader", None) is not None:
self.train_sampler.set_epoch(epoch) self.train_sampler.set_epoch(epoch)
else: else:
self.train_dataloader, self.train_sampler = self.dataset.get_train_dataloader() self.train_dataloader, self.train_sampler = (
self.dataset.get_train_dataloader()
)
self.train_sampler.set_epoch(epoch) self.train_sampler.set_epoch(epoch)
else: else:
self.train_dataloader = self.dataset.get_train_dataloader() self.train_dataloader = self.dataset.get_train_dataloader()
@@ -236,16 +264,23 @@ class QwenVlAct_Trainer:
grad_accum_steps = self.config.get("gradient_accumulation_steps", 1) grad_accum_steps = self.config.get("gradient_accumulation_steps", 1)
total = len(self.train_dataloader) total = len(self.train_dataloader)
t0 = time.time() t0 = time.time()
enable_profiling = self.config['profile'] enable_profiling = self.config["profile"]
# Optional PyTorch profiler for performance analysis # Optional PyTorch profiler for performance analysis
if enable_profiling: if enable_profiling:
profiler = torch.profiler.profile( profiler = torch.profiler.profile(
activities=[torch.profiler.ProfilerActivity.CPU, torch.profiler.ProfilerActivity.CUDA], activities=[
schedule=torch.profiler.schedule(wait=self.config['profile_wait_iters'], torch.profiler.ProfilerActivity.CPU,
warmup=self.config['profile_warmup_iters'], torch.profiler.ProfilerActivity.CUDA,
active=self.config['profile_active_iters']), ],
on_trace_ready=torch.profiler.tensorboard_trace_handler(self.config['profile_save_path'], worker_name="worker0"), schedule=torch.profiler.schedule(
wait=self.config["profile_wait_iters"],
warmup=self.config["profile_warmup_iters"],
active=self.config["profile_active_iters"],
),
on_trace_ready=torch.profiler.tensorboard_trace_handler(
self.config["profile_save_path"], worker_name="worker0"
),
record_shapes=True, record_shapes=True,
profile_memory=True, profile_memory=True,
with_stack=True, with_stack=True,
@@ -261,7 +296,14 @@ class QwenVlAct_Trainer:
for i, batch in enumerate(self.train_dataloader, self.initial_step): for i, batch in enumerate(self.train_dataloader, self.initial_step):
# Move batch to device # Move batch to device
if isinstance(self.dataset, PreprocessedDataset): if isinstance(self.dataset, PreprocessedDataset):
batch = {k: v.to(self.accelerator.device, non_blocking=True) if isinstance(v, torch.Tensor) else v for k, v in batch.items()} batch = {
k: (
v.to(self.accelerator.device, non_blocking=True)
if isinstance(v, torch.Tensor)
else v
)
for k, v in batch.items()
}
self.timers("data-load").stop() self.timers("data-load").stop()
@@ -275,7 +317,10 @@ class QwenVlAct_Trainer:
# Check for NaN loss # Check for NaN loss
if torch.isnan(loss): if torch.isnan(loss):
print(f"Warning: NaN loss detected in epoch: {epoch}, step: {i}", flush=True) print(
f"Warning: NaN loss detected in epoch: {epoch}, step: {i}",
flush=True,
)
continue continue
# Backward pass # Backward pass
@@ -301,31 +346,77 @@ class QwenVlAct_Trainer:
lr = self.lr_scheduler.get_last_lr()[0] lr = self.lr_scheduler.get_last_lr()[0]
# Gather loss across all processes for logging # Gather loss across all processes for logging
train_loss = self.accelerator.gather(loss.detach()).mean().item() train_loss = (
self.accelerator.gather(loss.detach()).mean().item()
)
_log_dict = { _log_dict = {
"lr": lr, "lr": lr,
"train_loss": train_loss, "train_loss": train_loss,
} }
# Log component losses # Log component losses
if "cross_entropy_loss" in outputs and outputs.cross_entropy_loss is not None: if (
_log_dict["cross_entropy_loss"] = self.accelerator.gather(outputs.cross_entropy_loss.detach()).mean().item() "cross_entropy_loss" in outputs
and outputs.cross_entropy_loss is not None
):
_log_dict["cross_entropy_loss"] = (
self.accelerator.gather(
outputs.cross_entropy_loss.detach()
)
.mean()
.item()
)
if "flow_loss" in outputs and outputs.flow_loss is not None: if "flow_loss" in outputs and outputs.flow_loss is not None:
_log_dict["flow_loss"] = self.accelerator.gather(outputs.flow_loss.detach()).mean().item() _log_dict["flow_loss"] = (
self.accelerator.gather(outputs.flow_loss.detach())
.mean()
.item()
)
# Log per-dataset channel losses # Log per-dataset channel losses
if "channel_loss_dict" in outputs and outputs.channel_loss_dict is not None: if (
for dataset_name_i in ACTION_DATASET_NAMES + MULTIMODAL_DATASET_NAMES: "channel_loss_dict" in outputs
count_sum = self.accelerator.gather(outputs.channel_loss_count_dict[dataset_name_i]).sum().item() and outputs.channel_loss_dict is not None
):
for dataset_name_i in (
ACTION_DATASET_NAMES + MULTIMODAL_DATASET_NAMES
):
count_sum = (
self.accelerator.gather(
outputs.channel_loss_count_dict[dataset_name_i]
)
.sum()
.item()
)
if count_sum > 0: if count_sum > 0:
channel_loss = self.accelerator.gather(outputs.channel_loss_dict[dataset_name_i].detach()).sum().item() / count_sum channel_loss = (
_log_dict[f"channel_loss_{dataset_name_i}"] = channel_loss self.accelerator.gather(
outputs.channel_loss_dict[
dataset_name_i
].detach()
)
.sum()
.item()
/ count_sum
)
_log_dict[f"channel_loss_{dataset_name_i}"] = (
channel_loss
)
# Log action accuracy for fast tokenizer # Log action accuracy for fast tokenizer
if "action_accuracy" in outputs.channel_loss_dict and self.use_fast_tokenizer: if (
"action_accuracy" in outputs.channel_loss_dict
and self.use_fast_tokenizer
):
_log_dict["action_accuracy"] = ( _log_dict["action_accuracy"] = (
self.accelerator.gather(outputs.channel_loss_dict["action_accuracy"].detach()).mean().item() self.accelerator.gather(
outputs.channel_loss_dict[
"action_accuracy"
].detach()
)
.mean()
.item()
) )
# Log metrics # Log metrics
@@ -334,7 +425,9 @@ class QwenVlAct_Trainer:
# Log gradient norm # Log gradient norm
if self.logger is not None and self.accelerator.sync_gradients: if self.logger is not None and self.accelerator.sync_gradients:
self.logger.log({"total_norm": total_norm}, step=self.global_step) self.logger.log(
{"total_norm": total_norm}, step=self.global_step
)
self.timers("interval-time").stop() self.timers("interval-time").stop()
@@ -343,12 +436,13 @@ class QwenVlAct_Trainer:
self.timers("interval-time", log_level=0).start(barrier=False) self.timers("interval-time", log_level=0).start(barrier=False)
self.timers("data-load", log_level=0).start(barrier=False) self.timers("data-load", log_level=0).start(barrier=False)
# Periodic logging # Periodic logging
t1 = time.time() t1 = time.time()
if i % 1 == 0: if i % 1 == 0:
lr = self.lr_scheduler.get_last_lr()[0] lr = self.lr_scheduler.get_last_lr()[0]
self.training_log(epoch, self.num_epoch, i, total, loss, lr, t1 - t0) self.training_log(
epoch, self.num_epoch, i, total, loss, lr, t1 - t0
)
t0 = time.time() t0 = time.time()
if enable_profiling: if enable_profiling:
@@ -377,11 +471,22 @@ class QwenVlAct_Trainer:
# Validation loop # Validation loop
for i, batch in enumerate( for i, batch in enumerate(
tqdm(self.val_dataloader, desc="Validating", total=len(self.val_dataloader), tqdm(
disable=not self.accelerator.is_main_process) self.val_dataloader,
desc="Validating",
total=len(self.val_dataloader),
disable=not self.accelerator.is_main_process,
)
): ):
if isinstance(self.dataset, PreprocessedDataset): if isinstance(self.dataset, PreprocessedDataset):
batch = {k: v.to(self.accelerator.device, non_blocking=True) if isinstance(v, torch.Tensor) else v for k, v in batch.items()} batch = {
k: (
v.to(self.accelerator.device, non_blocking=True)
if isinstance(v, torch.Tensor)
else v
)
for k, v in batch.items()
}
with torch.no_grad(): with torch.no_grad():
outputs = self.model(**batch, mode="train") outputs = self.model(**batch, mode="train")
@@ -412,7 +517,7 @@ class QwenVlAct_Trainer:
# Load pretrained model # Load pretrained model
model = Qwen2_5_VLMoEForAction.from_pretrained( model = Qwen2_5_VLMoEForAction.from_pretrained(
self.config["pretrained_wallx_path"], self.config["pretrained_wallx_path"],
**{"use_fast_tokenizer": self.use_fast_tokenizer} **{"use_fast_tokenizer": self.use_fast_tokenizer},
) )
self.processor = model.processor self.processor = model.processor
model = model.to(torch.bfloat16) model = model.to(torch.bfloat16)
@@ -442,17 +547,26 @@ class QwenVlAct_Trainer:
# Configure parameter groups # Configure parameter groups
if self.config.get("train_action_expert_only", False): if self.config.get("train_action_expert_only", False):
self.print_rank0("Training action expert only", flush=True) self.print_rank0("Training action expert only", flush=True)
param_groups = [{"params": moe_params, "lr": self.config["action_expert_learning_rate"]}] param_groups = [
{
"params": moe_params,
"lr": self.config["action_expert_learning_rate"],
}
]
else: else:
param_groups = [ param_groups = [
{"params": vlm_params, "lr": self.config["learning_rate"]}, {"params": vlm_params, "lr": self.config["learning_rate"]},
{"params": moe_params, "lr": self.config["action_expert_learning_rate"]}, {
"params": moe_params,
"lr": self.config["action_expert_learning_rate"],
},
] ]
self.optimizer = AdamW(param_groups, weight_decay=0.1) self.optimizer = AdamW(param_groups, weight_decay=0.1)
self.print_rank0( self.print_rank0(
f"Setting MoE learning rate to {self.config['action_expert_learning_rate']}, " f"Setting MoE learning rate to {self.config['action_expert_learning_rate']}, "
f"VLM learning rate to {self.config['learning_rate']}", flush=True f"VLM learning rate to {self.config['learning_rate']}",
flush=True,
) )
else: else:
# Standard optimizer configuration # Standard optimizer configuration
@@ -479,9 +593,13 @@ class QwenVlAct_Trainer:
if hasattr(model, "enable_input_require_grads"): if hasattr(model, "enable_input_require_grads"):
self.model.enable_input_require_grads() self.model.enable_input_require_grads()
else: else:
def make_inputs_require_grad(module, input, output): def make_inputs_require_grad(module, input, output):
output.requires_grad_(True) output.requires_grad_(True)
self.model.get_input_embeddings().register_forward_hook(make_inputs_require_grad)
self.model.get_input_embeddings().register_forward_hook(
make_inputs_require_grad
)
# Prepare model, optimizer, and scheduler for distributed training # Prepare model, optimizer, and scheduler for distributed training
self.model, self.optimizer, self.lr_scheduler = self.accelerator.prepare( self.model, self.optimizer, self.lr_scheduler = self.accelerator.prepare(
@@ -522,7 +640,9 @@ class QwenVlAct_Trainer:
Handles weight key renaming for MoE architecture compatibility. Handles weight key renaming for MoE architecture compatibility.
""" """
# Load all safetensors files # Load all safetensors files
weight_files = sorted([f for f in os.listdir(pretrain_weight_path) if f.endswith(".safetensors")]) weight_files = sorted(
[f for f in os.listdir(pretrain_weight_path) if f.endswith(".safetensors")]
)
merged_weights = {} merged_weights = {}
# Merge weights from all files # Merge weights from all files
@@ -534,19 +654,31 @@ class QwenVlAct_Trainer:
# Rename weights for MoE compatibility # Rename weights for MoE compatibility
renamed_weights = {} renamed_weights = {}
for key, value in merged_weights.items(): for key, value in merged_weights.items():
if key.startswith("model.layers") and "mlp." in key and model.config.mlp_moe: if (
key.startswith("model.layers")
and "mlp." in key
and model.config.mlp_moe
):
# Rename MLP weights for MoE structure # Rename MLP weights for MoE structure
layer_num = key.split(".layers.")[1].split(".mlp")[0] layer_num = key.split(".layers.")[1].split(".mlp")[0]
new_key = key.replace(f"layers.{layer_num}.mlp.", f"layers.{layer_num}.moe.experts.0.") new_key = key.replace(
f"layers.{layer_num}.mlp.", f"layers.{layer_num}.moe.experts.0."
)
renamed_weights[new_key] = value renamed_weights[new_key] = value
elif key.startswith("model.layers") and "self_attn." in key and model.config.attention_moe: elif (
key.startswith("model.layers")
and "self_attn." in key
and model.config.attention_moe
):
# Rename attention weights for MoE structure # Rename attention weights for MoE structure
layer_num = key.split(".layers.")[1].split(".self_attn")[0] layer_num = key.split(".layers.")[1].split(".self_attn")[0]
proj_types = ["q_proj", "k_proj", "v_proj", "o_proj"] proj_types = ["q_proj", "k_proj", "v_proj", "o_proj"]
for proj in proj_types: for proj in proj_types:
if proj in key: if proj in key:
new_key = key.replace(f"layers.{layer_num}.self_attn.{proj}", new_key = key.replace(
f"layers.{layer_num}.self_attn.{proj}_experts.0") f"layers.{layer_num}.self_attn.{proj}",
f"layers.{layer_num}.self_attn.{proj}_experts.0",
)
renamed_weights[new_key] = value renamed_weights[new_key] = value
break break
else: else:
@@ -560,7 +692,16 @@ class QwenVlAct_Trainer:
return model return model
def training_log(self, current_epoch, total_epoch, current_train_iter, total_train_iter, loss, lr, time_per_step): def training_log(
self,
current_epoch,
total_epoch,
current_train_iter,
total_train_iter,
loss,
lr,
time_per_step,
):
""" """
Log training progress and performance metrics. Log training progress and performance metrics.
@@ -573,7 +714,13 @@ class QwenVlAct_Trainer:
lr (float): Current learning rate lr (float): Current learning rate
time_per_step (float): Time taken for current step time_per_step (float): Time taken for current step
""" """
timers_to_log = ["interval-time", "data-load", "forward-compute", "backward-compute", "optimizer"] timers_to_log = [
"interval-time",
"data-load",
"forward-compute",
"backward-compute",
"optimizer",
]
log_string = f" [{datetime.now().strftime('%Y-%m-%d %H:%M:%S')}]" log_string = f" [{datetime.now().strftime('%Y-%m-%d %H:%M:%S')}]"
log_string += " epoch {:3d}/{:3d} |".format(current_epoch, total_epoch) log_string += " epoch {:3d}/{:3d} |".format(current_epoch, total_epoch)
@@ -609,7 +756,9 @@ class QwenVlAct_Trainer:
if isinstance(self.dataset, PreprocessedDataset): if isinstance(self.dataset, PreprocessedDataset):
torch.save( torch.save(
{"epoch": epoch, "step": step}, {"epoch": epoch, "step": step},
os.path.join(ckpt_path, f"epoch_{epoch}_step_{step}_rank_{_rank}.pth") os.path.join(
ckpt_path, f"epoch_{epoch}_step_{step}_rank_{_rank}.pth"
),
) )
def resume_from_checkpoint(self): def resume_from_checkpoint(self):
@@ -632,7 +781,7 @@ class QwenVlAct_Trainer:
new_key = "module." + key new_key = "module." + key
new_state_dict[new_key] = state_dict[key] new_state_dict[new_key] = state_dict[key]
err = self.model.load_state_dict(new_state_dict, strict=False) self.model.load_state_dict(new_state_dict, strict=False)
else: else:
# Load full checkpoint including optimizer and scheduler states # Load full checkpoint including optimizer and scheduler states
self.accelerator.load_state(checkpoint_path) self.accelerator.load_state(checkpoint_path)
@@ -662,7 +811,9 @@ class QwenVlAct_Trainer:
mean_action = all_label.mean(dim=0) mean_action = all_label.mean(dim=0)
self.base_l1_loss = nn.functional.l1_loss(all_label, mean_action) self.base_l1_loss = nn.functional.l1_loss(all_label, mean_action)
self.logger.log({"base_l1_loss": self.base_l1_loss.item()}, step=self.global_step) self.logger.log(
{"base_l1_loss": self.base_l1_loss.item()}, step=self.global_step
)
# Log L1 loss for each DOF component # Log L1 loss for each DOF component
start_idx = 0 start_idx = 0
@@ -674,6 +825,8 @@ class QwenVlAct_Trainer:
dof_l1 = nn.functional.l1_loss(dof_pred, dof_label) dof_l1 = nn.functional.l1_loss(dof_pred, dof_label)
self.print_rank0(f"DOF {dof}, L1 loss: {dof_l1.item()}", flush=True) self.print_rank0(f"DOF {dof}, L1 loss: {dof_l1.item()}", flush=True)
self.logger.log({f"detail/l1_loss_{dof}": dof_l1.item()}, step=self.global_step) self.logger.log(
{f"detail/l1_loss_{dof}": dof_l1.item()}, step=self.global_step
)
start_idx = end_idx start_idx = end_idx
+228 -57
View File
@@ -9,133 +9,292 @@ action_statistic_dof = {
"min": [-3.6176], "min": [-3.6176],
"delta": [8.5015], "delta": [8.5015],
}, },
"follow_left_ee_cartesian_pos": {"min": [-0.036, -0.3241, -0.1245], "delta": [0.4389, 0.557, 0.479]}, "follow_left_ee_cartesian_pos": {
"follow_left_ee_rotation": {"min": [-1.2373, -0.1929, -1.5182], "delta": [2.2009, 1.5669, 2.0936]}, "min": [-0.036, -0.3241, -0.1245],
"delta": [0.4389, 0.557, 0.479],
},
"follow_left_ee_rotation": {
"min": [-1.2373, -0.1929, -1.5182],
"delta": [2.2009, 1.5669, 2.0936],
},
"follow_left_gripper": {"min": [-0.1196], "delta": [4.5226]}, "follow_left_gripper": {"min": [-0.1196], "delta": [4.5226]},
"follow_right_ee_cartesian_pos": {"min": [-0.0326, -0.2273, -0.1377], "delta": [0.4574, 0.5704, 0.4743]}, "follow_right_ee_cartesian_pos": {
"follow_right_ee_rotation": {"min": [-1.2201, -0.2611, -0.7427], "delta": [2.6623, 1.6622, 2.4186]}, "min": [-0.0326, -0.2273, -0.1377],
"delta": [0.4574, 0.5704, 0.4743],
},
"follow_right_ee_rotation": {
"min": [-1.2201, -0.2611, -0.7427],
"delta": [2.6623, 1.6622, 2.4186],
},
"follow_right_gripper": {"min": [-0.1208], "delta": [4.5261]}, "follow_right_gripper": {"min": [-0.1208], "delta": [4.5261]},
"height": {"min": [-0.0001], "delta": [0.5051]}, "height": {"min": [-0.0001], "delta": [0.5051]},
"head_actions": {"min": [-1.5000, -1.4167], "delta": [2.5000, 1.8879]}, "head_actions": {"min": [-1.5000, -1.4167], "delta": [2.5000, 1.8879]},
"base_velocity": {"min": [-0.0359, -0.084, -0.0162], "delta": [0.1539, 0.1848, 0.0322]}, "base_velocity": {
"min": [-0.0359, -0.084, -0.0162],
"delta": [0.1539, 0.1848, 0.0322],
},
}, },
"DobbE": { "DobbE": {
"follow_right_ee_cartesian_pos": {"min": [-0.6107, -0.3272, -0.4282], "delta": [1.2629, 1.5297, 0.8349]}, "follow_right_ee_cartesian_pos": {
"follow_right_ee_rotation": {"min": [-1.7378, -1.4597, -1.8712], "delta": [2.7031, 2.8182, 3.5921]}, "min": [-0.6107, -0.3272, -0.4282],
"delta": [1.2629, 1.5297, 0.8349],
},
"follow_right_ee_rotation": {
"min": [-1.7378, -1.4597, -1.8712],
"delta": [2.7031, 2.8182, 3.5921],
},
"follow_right_gripper": {"min": [0.0], "delta": [0.9983]}, "follow_right_gripper": {"min": [0.0], "delta": [0.9983]},
}, },
"RH20T": { "RH20T": {
"follow_right_ee_cartesian_pos": {"min": [0.3646, -0.2722, 0.0066], "delta": [0.3813, 0.5973, 0.3277]}, "follow_right_ee_cartesian_pos": {
"follow_right_ee_rotation": {"min": [-1.8716, -0.4398, -3.1414], "delta": [3.4145, 1.0225, 6.2828]}, "min": [0.3646, -0.2722, 0.0066],
"delta": [0.3813, 0.5973, 0.3277],
},
"follow_right_ee_rotation": {
"min": [-1.8716, -0.4398, -3.1414],
"delta": [3.4145, 1.0225, 6.2828],
},
"follow_right_gripper": {"min": [0.0], "delta": [95.0]}, "follow_right_gripper": {"min": [0.0], "delta": [95.0]},
}, },
"agibotworld_alpha": { "agibotworld_alpha": {
"follow_left_ee_cartesian_pos": {"min": [0.4954, 0.0166, 0.1729], "delta": [0.3336, 0.5123, 0.9189]}, "follow_left_ee_cartesian_pos": {
"follow_left_ee_rotation": {"min": [-3.1064, -1.2629, -3.1238], "delta": [6.2127, 2.5923, 6.2496]}, "min": [0.4954, 0.0166, 0.1729],
"delta": [0.3336, 0.5123, 0.9189],
},
"follow_left_ee_rotation": {
"min": [-3.1064, -1.2629, -3.1238],
"delta": [6.2127, 2.5923, 6.2496],
},
"follow_left_gripper": {"min": [34.6222], "delta": [86.1921]}, "follow_left_gripper": {"min": [34.6222], "delta": [86.1921]},
"follow_right_ee_cartesian_pos": {"min": [0.4615, -0.5975, 0.1638], "delta": [0.3823, 0.5577, 0.8873]}, "follow_right_ee_cartesian_pos": {
"follow_right_ee_rotation": {"min": [-3.0891, -1.0739, -2.5091], "delta": [6.1707, 2.3074, 3.8533]}, "min": [0.4615, -0.5975, 0.1638],
"delta": [0.3823, 0.5577, 0.8873],
},
"follow_right_ee_rotation": {
"min": [-3.0891, -1.0739, -2.5091],
"delta": [6.1707, 2.3074, 3.8533],
},
"follow_right_gripper": {"min": [34.6222], "delta": [85.7635]}, "follow_right_gripper": {"min": [34.6222], "delta": [85.7635]},
"height": {"min": [0.0], "delta": [0.4535]}, "height": {"min": [0.0], "delta": [0.4535]},
"head_actions": {"min": [-0.1746, 0.0523], "delta": [0.2444, 0.4713]}, "head_actions": {"min": [-0.1746, 0.0523], "delta": [0.2444, 0.4713]},
}, },
"austin_buds": { "austin_buds": {
"follow_right_ee_cartesian_pos": {"min": [0.3496, -0.2855, 0.0105], "delta": [0.3748, 0.492, 0.3116]}, "follow_right_ee_cartesian_pos": {
"follow_right_ee_rotation": {"min": [-3.1405, -0.151, -0.0737], "delta": [6.2813, 0.3218, 0.1536]}, "min": [0.3496, -0.2855, 0.0105],
"delta": [0.3748, 0.492, 0.3116],
},
"follow_right_ee_rotation": {
"min": [-3.1405, -0.151, -0.0737],
"delta": [6.2813, 0.3218, 0.1536],
},
"follow_right_gripper": {"min": [0.0076], "delta": [0.0724]}, "follow_right_gripper": {"min": [0.0076], "delta": [0.0724]},
}, },
"austin_sailor": { "austin_sailor": {
"follow_right_ee_cartesian_pos": {"min": [0.387, -0.3165, 0.0244], "delta": [0.2999, 0.5252, 0.2308]}, "follow_right_ee_cartesian_pos": {
"follow_right_ee_rotation": {"min": [-3.1402, -0.1618, -1.5918], "delta": [6.2804, 0.337, 2.9478]}, "min": [0.387, -0.3165, 0.0244],
"delta": [0.2999, 0.5252, 0.2308],
},
"follow_right_ee_rotation": {
"min": [-3.1402, -0.1618, -1.5918],
"delta": [6.2804, 0.337, 2.9478],
},
"follow_right_gripper": {"min": [0.0005], "delta": [0.0773]}, "follow_right_gripper": {"min": [0.0005], "delta": [0.0773]},
}, },
"austin_sirius": { "austin_sirius": {
"follow_right_ee_cartesian_pos": {"min": [0.0, -0.1182, 0.0], "delta": [0.5329, 0.3812, 0.2723]}, "follow_right_ee_cartesian_pos": {
"follow_right_ee_rotation": {"min": [-3.1407, -0.1243, -1.7434], "delta": [6.2823, 0.1975, 1.8073]}, "min": [0.0, -0.1182, 0.0],
"delta": [0.5329, 0.3812, 0.2723],
},
"follow_right_ee_rotation": {
"min": [-3.1407, -0.1243, -1.7434],
"delta": [6.2823, 0.1975, 1.8073],
},
"follow_right_gripper": {"min": [0.0334], "delta": [0.046]}, "follow_right_gripper": {"min": [0.0334], "delta": [0.046]},
}, },
"bc_z": { "bc_z": {
"follow_right_ee_cartesian_pos": {"min": [-0.3883, -0.1116, 0.6113], "delta": [0.7199, 0.4288, 0.3709]}, "follow_right_ee_cartesian_pos": {
"follow_right_ee_rotation": {"min": [-1.056, -1.0587, -2.6295], "delta": [1.9142, 1.9455, 4.8064]}, "min": [-0.3883, -0.1116, 0.6113],
"delta": [0.7199, 0.4288, 0.3709],
},
"follow_right_ee_rotation": {
"min": [-1.056, -1.0587, -2.6295],
"delta": [1.9142, 1.9455, 4.8064],
},
"follow_right_gripper": {"min": [0.2], "delta": [0.8]}, "follow_right_gripper": {"min": [0.2], "delta": [0.8]},
}, },
"berkeley_autolab_ur5": { "berkeley_autolab_ur5": {
"follow_right_ee_cartesian_pos": {"min": [0.3018, -0.2129, -0.1888], "delta": [0.3121, 0.52, 0.3107]}, "follow_right_ee_cartesian_pos": {
"follow_right_ee_rotation": {"min": [-3.1396, -0.2278, 1.1413], "delta": [6.279, 0.454, 0.9841]}, "min": [0.3018, -0.2129, -0.1888],
"delta": [0.3121, 0.52, 0.3107],
},
"follow_right_ee_rotation": {
"min": [-3.1396, -0.2278, 1.1413],
"delta": [6.279, 0.454, 0.9841],
},
"follow_right_gripper": {"min": [0.0], "delta": [1.0]}, "follow_right_gripper": {"min": [0.0], "delta": [1.0]},
}, },
"berkeley_cable_routing": { "berkeley_cable_routing": {
"follow_right_ee_cartesian_pos": {"min": [0.4617, -0.28, 0.03], "delta": [0.1838, 0.5665, 0.1272]}, "follow_right_ee_cartesian_pos": {
"follow_right_ee_rotation": {"min": [-3.1413, -0.0299, -0.7665], "delta": [6.2826, 0.0692, 3.322]}, "min": [0.4617, -0.28, 0.03],
"delta": [0.1838, 0.5665, 0.1272],
},
"follow_right_ee_rotation": {
"min": [-3.1413, -0.0299, -0.7665],
"delta": [6.2826, 0.0692, 3.322],
},
}, },
"berkeley_fanuc_manipulation": { "berkeley_fanuc_manipulation": {
"follow_right_ee_cartesian_pos": {"min": [0.3718, -0.4072, 0.0184], "delta": [0.3483, 0.7201, 0.5229]}, "follow_right_ee_cartesian_pos": {
"follow_right_ee_rotation": {"min": [-3.1399, -1.0166, -1.6988], "delta": [6.2802, 1.4498, 3.2074]}, "min": [0.3718, -0.4072, 0.0184],
"delta": [0.3483, 0.7201, 0.5229],
},
"follow_right_ee_rotation": {
"min": [-3.1399, -1.0166, -1.6988],
"delta": [6.2802, 1.4498, 3.2074],
},
"follow_right_gripper": {"min": [0.0], "delta": [1.0]}, "follow_right_gripper": {"min": [0.0], "delta": [1.0]},
}, },
"bridge_data_v2": { "bridge_data_v2": {
"follow_right_ee_cartesian_pos": {"min": [0.1498, -0.2178, -0.0901], "delta": [0.3012, 0.469, 0.298]}, "follow_right_ee_cartesian_pos": {
"follow_right_ee_rotation": {"min": [-0.3279, -0.6105, -1.0578], "delta": [0.7378, 1.0353, 2.2552]}, "min": [0.1498, -0.2178, -0.0901],
"delta": [0.3012, 0.469, 0.298],
},
"follow_right_ee_rotation": {
"min": [-0.3279, -0.6105, -1.0578],
"delta": [0.7378, 1.0353, 2.2552],
},
"follow_right_gripper": {"min": [0.0692], "delta": [0.9426]}, "follow_right_gripper": {"min": [0.0692], "delta": [0.9426]},
}, },
"dlr_edan_shared_control": { "dlr_edan_shared_control": {
"follow_right_ee_cartesian_pos": {"min": [-0.8387, 0.1473, -0.3934], "delta": [0.6579, 0.6025, 1.1566]}, "follow_right_ee_cartesian_pos": {
"follow_right_ee_rotation": {"min": [-3.1217, -1.5197, -2.2516], "delta": [6.2505, 1.5594, 4.2831]}, "min": [-0.8387, 0.1473, -0.3934],
"delta": [0.6579, 0.6025, 1.1566],
},
"follow_right_ee_rotation": {
"min": [-3.1217, -1.5197, -2.2516],
"delta": [6.2505, 1.5594, 4.2831],
},
"follow_right_gripper": {"min": [0.0], "delta": [1.0]}, "follow_right_gripper": {"min": [0.0], "delta": [1.0]},
}, },
"droid": { "droid": {
"follow_right_ee_cartesian_pos": {"min": [0.2667, -0.4396, -0.0472], "delta": [0.5159, 0.8806, 0.8331]}, "follow_right_ee_cartesian_pos": {
"follow_right_ee_rotation": {"min": [-3.1374, -1.216, -2.1741], "delta": [6.2749, 2.1075, 4.2259]}, "min": [0.2667, -0.4396, -0.0472],
"delta": [0.5159, 0.8806, 0.8331],
},
"follow_right_ee_rotation": {
"min": [-3.1374, -1.216, -2.1741],
"delta": [6.2749, 2.1075, 4.2259],
},
"follow_right_gripper": {"min": [0.0], "delta": [0.9912]}, "follow_right_gripper": {"min": [0.0], "delta": [0.9912]},
}, },
"fmb": { "fmb": {
"follow_right_ee_cartesian_pos": {"min": [0.3554, -0.2844, 0.0354], "delta": [0.336, 0.4961, 0.2943]}, "follow_right_ee_cartesian_pos": {
"follow_right_ee_rotation": {"min": [-3.1404, -0.9302, -0.0599], "delta": [6.2807, 1.724, 1.8284]}, "min": [0.3554, -0.2844, 0.0354],
"delta": [0.336, 0.4961, 0.2943],
},
"follow_right_ee_rotation": {
"min": [-3.1404, -0.9302, -0.0599],
"delta": [6.2807, 1.724, 1.8284],
},
"follow_right_gripper": {"min": [0.0], "delta": [1.0]}, "follow_right_gripper": {"min": [0.0], "delta": [1.0]},
}, },
"fractal": { "fractal": {
"follow_right_ee_cartesian_pos": {"min": [0.3242, -0.2836, 0.1405], "delta": [0.5518, 0.4963, 0.9328]}, "follow_right_ee_cartesian_pos": {
"follow_right_ee_rotation": {"min": [-3.1308, -0.2421, -2.9685], "delta": [6.2609, 1.7343, 5.819]}, "min": [0.3242, -0.2836, 0.1405],
"delta": [0.5518, 0.4963, 0.9328],
},
"follow_right_ee_rotation": {
"min": [-3.1308, -0.2421, -2.9685],
"delta": [6.2609, 1.7343, 5.819],
},
"follow_right_gripper": {"min": [0.0], "delta": [1.0]}, "follow_right_gripper": {"min": [0.0], "delta": [1.0]},
}, },
"furniture_bench": { "furniture_bench": {
"follow_right_ee_cartesian_pos": {"min": [0.3691, -0.181, 0.0058], "delta": [0.2962, 0.3582, 0.1775]}, "follow_right_ee_cartesian_pos": {
"follow_right_ee_rotation": {"min": [-3.1394, -0.6121, -1.9958], "delta": [6.2786, 1.6114, 3.7748]}, "min": [0.3691, -0.181, 0.0058],
"delta": [0.2962, 0.3582, 0.1775],
},
"follow_right_ee_rotation": {
"min": [-3.1394, -0.6121, -1.9958],
"delta": [6.2786, 1.6114, 3.7748],
},
"follow_right_gripper": {"min": [0.0035], "delta": [0.0762]}, "follow_right_gripper": {"min": [0.0035], "delta": [0.0762]},
}, },
"jaco_play": { "jaco_play": {
"follow_right_ee_cartesian_pos": {"min": [-0.3787, -0.6294, 0.1682], "delta": [0.5898, 0.3587, 0.2183]}, "follow_right_ee_cartesian_pos": {
"follow_right_ee_rotation": {"min": [0.9792, -0.0668, -0.0498], "delta": [0.0175, 0.1277, 0.0686]}, "min": [-0.3787, -0.6294, 0.1682],
"delta": [0.5898, 0.3587, 0.2183],
},
"follow_right_ee_rotation": {
"min": [0.9792, -0.0668, -0.0498],
"delta": [0.0175, 0.1277, 0.0686],
},
"follow_right_gripper": {"min": [0.0791], "delta": [0.1033]}, "follow_right_gripper": {"min": [0.0791], "delta": [0.1033]},
}, },
"nyu_rot": { "nyu_rot": {
"follow_right_ee_cartesian_pos": {"min": [0.25, -1.0, -0.2], "delta": [0.75, 2.0, 1.2]}, "follow_right_ee_cartesian_pos": {
"follow_right_ee_rotation": {"min": [-3.1416, -3.1416, 6.2831], "delta": [9.4248, 4.1416, 0.0]}, "min": [0.25, -1.0, -0.2],
"delta": [0.75, 2.0, 1.2],
},
"follow_right_ee_rotation": {
"min": [-3.1416, -3.1416, 6.2831],
"delta": [9.4248, 4.1416, 0.0],
},
"follow_right_gripper": {"min": [0.0], "delta": [1.0]}, "follow_right_gripper": {"min": [0.0], "delta": [1.0]},
}, },
"stanford_hydra": { "stanford_hydra": {
"follow_right_ee_cartesian_pos": {"min": [0.2068, -0.274, 0.1317], "delta": [0.4929, 0.4981, 0.4588]}, "follow_right_ee_cartesian_pos": {
"follow_right_ee_rotation": {"min": [-3.1321, -0.7496, -3.0269], "delta": [6.2658, 1.5176, 5.8261]}, "min": [0.2068, -0.274, 0.1317],
"delta": [0.4929, 0.4981, 0.4588],
},
"follow_right_ee_rotation": {
"min": [-3.1321, -0.7496, -3.0269],
"delta": [6.2658, 1.5176, 5.8261],
},
"follow_right_gripper": {"min": [0.0], "delta": [0.0811]}, "follow_right_gripper": {"min": [0.0], "delta": [0.0811]},
}, },
"stanford_kuka_multimodal": { "stanford_kuka_multimodal": {
"follow_right_ee_cartesian_pos": {"min": [0.4781, -0.0659, 0.3424], "delta": [0.0868, 0.0864, 0.1863]}, "follow_right_ee_cartesian_pos": {
"follow_right_ee_rotation": {"min": [-3.136, -0.0521, -3.1413], "delta": [6.2727, 0.1109, 6.2825]}, "min": [0.4781, -0.0659, 0.3424],
"delta": [0.0868, 0.0864, 0.1863],
},
"follow_right_ee_rotation": {
"min": [-3.136, -0.0521, -3.1413],
"delta": [6.2727, 0.1109, 6.2825],
},
"follow_right_gripper": {"min": [-0.4713], "delta": [0.9485]}, "follow_right_gripper": {"min": [-0.4713], "delta": [0.9485]},
}, },
"taco_play": { "taco_play": {
"follow_right_ee_cartesian_pos": {"min": [0.1375, -0.4291, 0.2052], "delta": [0.5327, 1.0237, 0.3913]}, "follow_right_ee_cartesian_pos": {
"follow_right_ee_rotation": {"min": [-3.1391, -0.6946, -1.2808], "delta": [6.2784, 0.8196, 3.0856]}, "min": [0.1375, -0.4291, 0.2052],
"delta": [0.5327, 1.0237, 0.3913],
},
"follow_right_ee_rotation": {
"min": [-3.1391, -0.6946, -1.2808],
"delta": [6.2784, 0.8196, 3.0856],
},
"follow_right_gripper": {"min": [0.0001], "delta": [0.0806]}, "follow_right_gripper": {"min": [0.0001], "delta": [0.0806]},
}, },
"utaustin_mutex": { "utaustin_mutex": {
"follow_right_ee_cartesian_pos": {"min": [0.3213, -0.4734, 0.0141], "delta": [0.2108, 0.8471, 0.5644]}, "follow_right_ee_cartesian_pos": {
"follow_right_ee_rotation": {"min": [-3.1404, -0.2202, -1.5489], "delta": [6.2805, 0.582, 1.9282]}, "min": [0.3213, -0.4734, 0.0141],
"delta": [0.2108, 0.8471, 0.5644],
},
"follow_right_ee_rotation": {
"min": [-3.1404, -0.2202, -1.5489],
"delta": [6.2805, 0.582, 1.9282],
},
"follow_right_gripper": {"min": [0.0019], "delta": [0.0738]}, "follow_right_gripper": {"min": [0.0019], "delta": [0.0738]},
}, },
"viola": { "viola": {
"follow_right_ee_cartesian_pos": {"min": [0.4011, -0.2521, 0.0103], "delta": [0.2444, 0.4305, 0.4355]}, "follow_right_ee_cartesian_pos": {
"follow_right_ee_rotation": {"min": [-3.1403, -0.2737, -1.8626], "delta": [6.2804, 0.4901, 2.0618]}, "min": [0.4011, -0.2521, 0.0103],
"delta": [0.2444, 0.4305, 0.4355],
},
"follow_right_ee_rotation": {
"min": [-3.1403, -0.2737, -1.8626],
"delta": [6.2804, 0.4901, 2.0618],
},
"follow_right_gripper": {"min": [0.0002], "delta": [0.0773]}, "follow_right_gripper": {"min": [0.0002], "delta": [0.0773]},
}, },
"kuka": { "kuka": {
@@ -179,11 +338,23 @@ action_statistic_dof = {
}, },
}, },
"agibotworld_beta": { "agibotworld_beta": {
"follow_left_ee_cartesian_pos": {"min": [0.4954, 0.0166, 0.1729], "delta": [0.3336, 0.5123, 0.9189]}, "follow_left_ee_cartesian_pos": {
"follow_left_ee_rotation": {"min": [-3.1064, -1.2629, -3.1238], "delta": [6.2127, 2.5923, 6.2496]}, "min": [0.4954, 0.0166, 0.1729],
"delta": [0.3336, 0.5123, 0.9189],
},
"follow_left_ee_rotation": {
"min": [-3.1064, -1.2629, -3.1238],
"delta": [6.2127, 2.5923, 6.2496],
},
"follow_left_gripper": {"min": [34.6222], "delta": [86.1921]}, "follow_left_gripper": {"min": [34.6222], "delta": [86.1921]},
"follow_right_ee_cartesian_pos": {"min": [0.4615, -0.5975, 0.1638], "delta": [0.3823, 0.5577, 0.8873]}, "follow_right_ee_cartesian_pos": {
"follow_right_ee_rotation": {"min": [-3.0891, -1.0739, -2.5091], "delta": [6.1707, 2.3074, 3.8533]}, "min": [0.4615, -0.5975, 0.1638],
"delta": [0.3823, 0.5577, 0.8873],
},
"follow_right_ee_rotation": {
"min": [-3.0891, -1.0739, -2.5091],
"delta": [6.1707, 2.3074, 3.8533],
},
"follow_right_gripper": {"min": [34.6222], "delta": [85.7635]}, "follow_right_gripper": {"min": [34.6222], "delta": [85.7635]},
"height": {"min": [0.0], "delta": [0.4535]}, "height": {"min": [0.0], "delta": [0.4535]},
"head_actions": {"min": [-0.1746, 0.0523], "delta": [0.2444, 0.4713]}, "head_actions": {"min": [-0.1746, 0.0523], "delta": [0.2444, 0.4713]},
+54 -35
View File
@@ -4,23 +4,28 @@ from torch.cuda import nvtx
from abc import ABC, abstractmethod from abc import ABC, abstractmethod
from typing import List from typing import List
def _is_distributed(): def _is_distributed():
return torch.distributed.is_available() and torch.distributed.is_initialized() return torch.distributed.is_available() and torch.distributed.is_initialized()
def _get_world_size(): def _get_world_size():
if _is_distributed(): if _is_distributed():
return torch.distributed.get_world_size() return torch.distributed.get_world_size()
return 1 return 1
def _get_rank(): def _get_rank():
if _is_distributed(): if _is_distributed():
return torch.distributed.get_rank() return torch.distributed.get_rank()
return 0 return 0
def _barrier(group=None): def _barrier(group=None):
if _is_distributed(): if _is_distributed():
torch.distributed.barrier(group=group) torch.distributed.barrier(group=group)
if torch.distributed.is_available(): if torch.distributed.is_available():
try: try:
dist_all_gather_func = torch.distributed.all_gather_into_tensor dist_all_gather_func = torch.distributed.all_gather_into_tensor
@@ -29,6 +34,7 @@ if torch.distributed.is_available():
else: else:
dist_all_gather_func = None dist_all_gather_func = None
class TimerBase(ABC): class TimerBase(ABC):
"""Timer base class.""" """Timer base class."""
@@ -76,7 +82,7 @@ class DummyTimer(TimerBase):
"""Dummy Timer.""" """Dummy Timer."""
def __init__(self): def __init__(self):
super().__init__('dummy timer') super().__init__("dummy timer")
def start(self, barrier=False, nvtx_push=False): def start(self, barrier=False, nvtx_push=False):
return return
@@ -89,8 +95,8 @@ class DummyTimer(TimerBase):
def elapsed(self, reset=True, barrier=False): def elapsed(self, reset=True, barrier=False):
raise Exception( raise Exception(
'dummy timer should not be used to calculate elapsed time, ' "dummy timer should not be used to calculate elapsed time, "
'check if timer\'s log_level <= self._log_level.' "check if timer's log_level <= self._log_level."
) )
def active_time(self): def active_time(self):
@@ -98,8 +104,8 @@ class DummyTimer(TimerBase):
Note: Not supported for DummyTimer. Note: Not supported for DummyTimer.
""" """
raise Exception( raise Exception(
'active timer should not be used to calculate elapsed time, ' "active timer should not be used to calculate elapsed time, "
'check if timer\'s log_level <= self._log_level.' "check if timer's log_level <= self._log_level."
) )
@@ -144,7 +150,7 @@ class Timer(TimerBase):
Args: Args:
barrier (bool, optional): Synchronizes ranks before starting. Defaults to False. barrier (bool, optional): Synchronizes ranks before starting. Defaults to False.
""" """
assert not self._started, 'timer has already been started' assert not self._started, "timer has already been started"
if barrier: if barrier:
_barrier(group=self._barrier_group) _barrier(group=self._barrier_group)
if torch.cuda.is_available(): if torch.cuda.is_available():
@@ -155,7 +161,6 @@ class Timer(TimerBase):
nvtx.range_push("{}".format(self.name)) nvtx.range_push("{}".format(self.name))
self.nvtx = True self.nvtx = True
def stop(self, barrier=False, sync=False): def stop(self, barrier=False, sync=False):
"""Stop the timer. """Stop the timer.
@@ -164,7 +169,7 @@ class Timer(TimerBase):
""" """
if self.nvtx: if self.nvtx:
nvtx.range_pop() nvtx.range_pop()
assert self._started, 'timer is not started' assert self._started, "timer is not started"
if barrier: if barrier:
_barrier(group=self._barrier_group) _barrier(group=self._barrier_group)
if torch.cuda.is_available() and sync: if torch.cuda.is_available() and sync:
@@ -221,10 +226,10 @@ class Timers:
Allowed: ['max', 'minmax', 'all']. Allowed: ['max', 'minmax', 'all'].
""" """
self._log_level = log_level self._log_level = log_level
allowed_log_options = set(['max', 'minmax', 'all']) allowed_log_options = set(["max", "minmax", "all"])
assert ( assert (
log_option in allowed_log_options log_option in allowed_log_options
), 'input log option {} is invalid. It must be one of {}'.format( ), "input log option {} is invalid. It must be one of {}".format(
log_option, allowed_log_options log_option, allowed_log_options
) )
self._log_option = log_option self._log_option = log_option
@@ -240,8 +245,10 @@ class Timers:
if name in self._timers: if name in self._timers:
if log_level is not None: if log_level is not None:
assert log_level == self._log_levels[name], ( assert log_level == self._log_levels[name], (
'input log level {} does not match already existing ' "input log level {} does not match already existing "
'log level {} for {} timer'.format(log_level, self._log_levels[name], name) "log level {} for {} timer".format(
log_level, self._log_levels[name], name
)
) )
return self._timers[name] return self._timers[name]
# If timer does not exist and no log level is provided, # If timer does not exist and no log level is provided,
@@ -250,7 +257,7 @@ class Timers:
log_level = self._max_log_level log_level = self._max_log_level
assert ( assert (
log_level <= self._max_log_level log_level <= self._max_log_level
), 'log level {} is larger than max supported log level {}'.format( ), "log level {} is larger than max supported log level {}".format(
log_level, self._max_log_level log_level, self._max_log_level
) )
# Now if the input log level is larger than the one set for # Now if the input log level is larger than the one set for
@@ -284,7 +291,7 @@ class Timers:
if torch.cuda.is_available(): if torch.cuda.is_available():
device = torch.cuda.current_device() device = torch.cuda.current_device()
else: else:
device = torch.device('cpu') device = torch.device("cpu")
rank_name_to_time = torch.zeros( rank_name_to_time = torch.zeros(
(world_size, len(names)), dtype=torch.float, device=device (world_size, len(names)), dtype=torch.float, device=device
@@ -296,7 +303,9 @@ class Timers:
if world_size > 1 and _is_distributed() and dist_all_gather_func is not None: if world_size > 1 and _is_distributed() and dist_all_gather_func is not None:
try: try:
dist_all_gather_func(rank_name_to_time.view(-1), rank_name_to_time[rank, :].view(-1)) dist_all_gather_func(
rank_name_to_time.view(-1), rank_name_to_time[rank, :].view(-1)
)
except Exception as e: except Exception as e:
print(f"Warning: all_gather failed: {e}. Using single rank timing.") print(f"Warning: all_gather failed: {e}. Using single rank timing.")
@@ -319,30 +328,38 @@ class Timers:
) )
return name_to_min_max_time return name_to_min_max_time
def _get_global_min_max_time_string(self, names, reset, barrier, normalizer, max_only): def _get_global_min_max_time_string(
self, names, reset, barrier, normalizer, max_only
):
"""Report strings for max/minmax times across all ranks.""" """Report strings for max/minmax times across all ranks."""
name_to_min_max_time = self._get_global_min_max_time(names, reset, barrier, normalizer) name_to_min_max_time = self._get_global_min_max_time(
names, reset, barrier, normalizer
)
if not name_to_min_max_time: if not name_to_min_max_time:
return None return None
world_size = _get_world_size() world_size = _get_world_size()
if world_size == 1: if world_size == 1:
output_string = 'time (ms):' output_string = "time (ms):"
for name in name_to_min_max_time: for name in name_to_min_max_time:
_, max_time = name_to_min_max_time[name] _, max_time = name_to_min_max_time[name]
output_string += '\n {}: {:.2f}'.format((name + ' ').ljust(48, '.'), max_time) output_string += "\n {}: {:.2f}".format(
(name + " ").ljust(48, "."), max_time
)
else: else:
if max_only: if max_only:
output_string = 'max time across ranks (ms):' output_string = "max time across ranks (ms):"
else: else:
output_string = '(min, max) time across ranks (ms):' output_string = "(min, max) time across ranks (ms):"
for name in name_to_min_max_time: for name in name_to_min_max_time:
min_time, max_time = name_to_min_max_time[name] min_time, max_time = name_to_min_max_time[name]
if max_only: if max_only:
output_string += '\n {}: {:.2f}'.format((name + ' ').ljust(48, '.'), max_time) output_string += "\n {}: {:.2f}".format(
(name + " ").ljust(48, "."), max_time
)
else: else:
output_string += '\n {}: ({:.2f}, {:.2f})'.format( output_string += "\n {}: ({:.2f}, {:.2f})".format(
(name + ' ').ljust(48, '.'), min_time, max_time (name + " ").ljust(48, "."), min_time, max_time
) )
return output_string return output_string
@@ -351,7 +368,7 @@ class Timers:
rank_name_to_time = self._get_elapsed_time_all_ranks(names, reset, barrier) rank_name_to_time = self._get_elapsed_time_all_ranks(names, reset, barrier)
world_size = _get_world_size() world_size = _get_world_size()
output_string = 'times across ranks (ms):' output_string = "times across ranks (ms):"
no_reported_timing = True no_reported_timing = True
for i, name in enumerate(names): for i, name in enumerate(names):
not_yet_found = True not_yet_found = True
@@ -360,13 +377,13 @@ class Timers:
no_reported_timing = False no_reported_timing = False
if not_yet_found: if not_yet_found:
not_yet_found = False not_yet_found = False
output_string += '\n {}:'.format(name) output_string += "\n {}:".format(name)
if world_size == 1: if world_size == 1:
output_string += '\n {:.2f}'.format( output_string += "\n {:.2f}".format(
rank_name_to_time[rank, i] / normalizer rank_name_to_time[rank, i] / normalizer
) )
else: else:
output_string += '\n rank {:2d}: {:.2f}'.format( output_string += "\n rank {:2d}: {:.2f}".format(
rank, rank_name_to_time[rank, i] / normalizer rank, rank_name_to_time[rank, i] / normalizer
) )
if no_reported_timing: if no_reported_timing:
@@ -398,23 +415,23 @@ class Timers:
str: Formatted string with the timer values. str: Formatted string with the timer values.
""" """
if names == None: # get all registered timers if names is None: # get all registered timers
names = list(self._timers.keys()) names = list(self._timers.keys())
assert normalizer > 0.0 assert normalizer > 0.0
if self._log_option in ['max', 'minmax']: if self._log_option in ["max", "minmax"]:
max_only = False max_only = False
if self._log_option == 'max': if self._log_option == "max":
max_only = True max_only = True
output_string = self._get_global_min_max_time_string( output_string = self._get_global_min_max_time_string(
names, reset, barrier, normalizer / 1000.0, max_only names, reset, barrier, normalizer / 1000.0, max_only
) )
elif self._log_option == 'all': elif self._log_option == "all":
output_string = self._get_all_ranks_time_string( output_string = self._get_all_ranks_time_string(
names, reset, barrier, normalizer / 1000.0 names, reset, barrier, normalizer / 1000.0
) )
else: else:
raise Exception('unknown timing log option {}'.format(self._log_option)) raise Exception("unknown timing log option {}".format(self._log_option))
return output_string return output_string
def log( def log(
@@ -476,8 +493,10 @@ class Timers:
# torch.utils.add_scalars makes each timer its own run, which # torch.utils.add_scalars makes each timer its own run, which
# polutes the runs list, so we just add each as a scalar # polutes the runs list, so we just add each as a scalar
assert normalizer > 0.0 assert normalizer > 0.0
name_to_min_max_time = self._get_global_min_max_time(names, reset, barrier, normalizer) name_to_min_max_time = self._get_global_min_max_time(
names, reset, barrier, normalizer
)
if writer is not None: if writer is not None:
for name in name_to_min_max_time: for name in name_to_min_max_time:
_, max_time = name_to_min_max_time[name] _, max_time = name_to_min_max_time[name]
writer.add_scalar(name + '-time', max_time, iteration) writer.add_scalar(name + "-time", max_time, iteration)