[lint] Update lint (#16)
* update lint * update readme * update ruff lint
This commit is contained in:
@@ -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
|
||||||
@@ -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
|
||||||
|
```
|
||||||
@@ -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");
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -549,4 +549,3 @@ void launch_multimodal_rope_backward(
|
|||||||
break;
|
break;
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,2 @@
|
|||||||
|
[tool.ruff]
|
||||||
|
per-file-ignores = { "__init__.py" = ["F401", "E402"] }
|
||||||
@@ -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()
|
||||||
|
|||||||
@@ -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()
|
||||||
@@ -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
@@ -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
@@ -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:
|
||||||
|
|||||||
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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)
|
||||||
|
|||||||
@@ -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
@@ -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
@@ -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
@@ -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)
|
||||||
|
|||||||
Reference in New Issue
Block a user