Fix normalizer (#57)
* fix normalizer * fix val * update compute stats * delete norm * update readme * minor fix * fix action normalizer * fix * fix * update * update * update * update * update * lint * lint * lint * lint
This commit is contained in:
@@ -53,6 +53,7 @@ MAX_JOBS=4 pip install flash-attn==2.7.4.post1 --no-build-isolation
|
|||||||
Install lerobot:
|
Install lerobot:
|
||||||
```bash
|
```bash
|
||||||
git clone https://github.com/huggingface/lerobot.git
|
git clone https://github.com/huggingface/lerobot.git
|
||||||
|
git checkout c66cd401767e60baece16e1cf68da2824227e076
|
||||||
cd lerobot
|
cd lerobot
|
||||||
pip install -e .
|
pip install -e .
|
||||||
```
|
```
|
||||||
|
|||||||
@@ -0,0 +1,76 @@
|
|||||||
|
import yaml
|
||||||
|
import torch
|
||||||
|
import tqdm
|
||||||
|
from lerobot.datasets.lerobot_dataset import LeRobotDataset, LeRobotDatasetMetadata
|
||||||
|
from wall_x.data.load_lerobot_dataset import KEY_MAPPINGS
|
||||||
|
import normalize
|
||||||
|
import numpy as np
|
||||||
|
import argparse
|
||||||
|
|
||||||
|
|
||||||
|
def load_config(config_path):
|
||||||
|
"""Load configuration from YAML file."""
|
||||||
|
with open(config_path, "r") as f:
|
||||||
|
config = yaml.load(f, Loader=yaml.FullLoader)
|
||||||
|
|
||||||
|
config["data"]["model_type"] = config.get("model_type")
|
||||||
|
|
||||||
|
return config
|
||||||
|
|
||||||
|
|
||||||
|
def load_lerobot_dataset(repo_id, action_horizon, args):
|
||||||
|
dataset_meta = LeRobotDatasetMetadata(repo_id)
|
||||||
|
dataset = LeRobotDataset(
|
||||||
|
repo_id,
|
||||||
|
delta_timestamps={
|
||||||
|
key: [t / dataset_meta.fps for t in range(action_horizon)]
|
||||||
|
for key in [KEY_MAPPINGS[repo_id]["action"]]
|
||||||
|
},
|
||||||
|
)
|
||||||
|
num_batches = len(dataset) // args.batch_size
|
||||||
|
generator = torch.Generator()
|
||||||
|
generator.manual_seed(args.seed)
|
||||||
|
data_loader = torch.utils.data.DataLoader(
|
||||||
|
dataset,
|
||||||
|
batch_size=args.batch_size,
|
||||||
|
shuffle=False,
|
||||||
|
drop_last=True,
|
||||||
|
generator=generator,
|
||||||
|
num_workers=args.num_workers,
|
||||||
|
persistent_workers=True if args.num_workers > 0 else False,
|
||||||
|
)
|
||||||
|
return data_loader, num_batches
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
# set args
|
||||||
|
parser = argparse.ArgumentParser()
|
||||||
|
parser.add_argument("--batch_size", type=int, default=256)
|
||||||
|
parser.add_argument("--num_workers", type=int, default=2)
|
||||||
|
parser.add_argument("--seed", type=int, default=0)
|
||||||
|
args = parser.parse_args()
|
||||||
|
|
||||||
|
# Configs
|
||||||
|
path = "/path/to/config.yml"
|
||||||
|
output_path = "/path/to/output"
|
||||||
|
config = load_config(path)
|
||||||
|
lerobot_config = config["data"]["lerobot_config"]
|
||||||
|
repo_id = lerobot_config.get("repo_id", None)
|
||||||
|
assert repo_id is not None, "repo id is required"
|
||||||
|
action_horizon = config["data"].get("action_horizon", 32)
|
||||||
|
|
||||||
|
data_loader, num_batches = load_lerobot_dataset(repo_id, action_horizon, args)
|
||||||
|
|
||||||
|
keys = ["state", "action"]
|
||||||
|
stats = {key: normalize.RunningStats() for key in keys}
|
||||||
|
for batch in tqdm.tqdm(data_loader, total=num_batches, desc="Computing stats"):
|
||||||
|
for key in keys:
|
||||||
|
stats[key].update(np.asarray(batch[KEY_MAPPINGS[repo_id][key]]))
|
||||||
|
norm_stats = {
|
||||||
|
KEY_MAPPINGS[repo_id][key]: stats.get_statistics()
|
||||||
|
for key, stats in stats.items()
|
||||||
|
}
|
||||||
|
|
||||||
|
output_path = output_path + "/" + repo_id
|
||||||
|
print(f"Writing stats to: {output_path}")
|
||||||
|
normalize.save(output_path, norm_stats)
|
||||||
@@ -1,22 +1,13 @@
|
|||||||
import os
|
import os
|
||||||
import yaml
|
import yaml
|
||||||
import torch
|
import torch
|
||||||
|
import argparse
|
||||||
|
from tqdm import tqdm
|
||||||
import matplotlib.pyplot as plt
|
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
|
||||||
|
|
||||||
|
|
||||||
model_path = "path/to/model"
|
|
||||||
action_tokenizer_path = "path/to/action_tokenizer"
|
|
||||||
save_dir = "path/to/plot"
|
|
||||||
model = Qwen2_5_VLMoEForAction.from_pretrained(
|
|
||||||
model_path, action_tokenizer_path=action_tokenizer_path
|
|
||||||
)
|
|
||||||
model.eval()
|
|
||||||
model = model.to("cuda")
|
|
||||||
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:
|
||||||
@@ -27,60 +18,97 @@ def load_config(config_path):
|
|||||||
return config
|
return config
|
||||||
|
|
||||||
|
|
||||||
# get test dataloader
|
if __name__ == "__main__":
|
||||||
path = "path/to/config"
|
parser = argparse.ArgumentParser()
|
||||||
config = load_config(path)
|
parser.add_argument("--pred_horizon", type=int, default=32)
|
||||||
dataload_config = get_data_configs(config["data"])
|
parser.add_argument("--origin_action_dim", type=int, default=7)
|
||||||
lerobot_config = dataload_config.get("lerobot_config", {})
|
args = parser.parse_args()
|
||||||
dataset = load_test_dataset(config, lerobot_config, seed=42)
|
|
||||||
dataloader = dataset.get_dataloader()
|
|
||||||
|
|
||||||
total_frames = len(dataloader)
|
origin_action_dim = args.origin_action_dim
|
||||||
|
pred_horizon = args.pred_horizon
|
||||||
|
|
||||||
pred_horizon = 32
|
# get train config
|
||||||
action_dim = 14
|
model_path = "/path/to/model"
|
||||||
gt_traj = torch.zeros((total_frames, action_dim))
|
action_tokenizer_path = "/path/to/action/tokenizer"
|
||||||
pred_traj = torch.zeros((total_frames, action_dim))
|
save_dir = "/path/to/save/dir"
|
||||||
|
path = "/path/to/train/config"
|
||||||
|
config = load_config(path)
|
||||||
|
|
||||||
for idx, batch in enumerate(dataloader):
|
# load model with customized robot config
|
||||||
if idx % pred_horizon == 0 and idx + pred_horizon < total_frames:
|
model = Qwen2_5_VLMoEForAction.from_pretrained(
|
||||||
batch = batch.to("cuda")
|
model_path, train_config=config, action_tokenizer_path=action_tokenizer_path
|
||||||
with torch.no_grad():
|
)
|
||||||
outputs = model(
|
model.eval()
|
||||||
**batch,
|
model = model.to("cuda")
|
||||||
action_dim=action_dim,
|
model = model.bfloat16()
|
||||||
pred_horizon=pred_horizon,
|
|
||||||
mode="predict",
|
# get test dataloader
|
||||||
predict_mode="fast",
|
dataload_config = get_data_configs(config["data"])
|
||||||
|
lerobot_config = dataload_config.get("lerobot_config", {})
|
||||||
|
dataset = load_test_dataset(config, lerobot_config, seed=42)
|
||||||
|
dataloader = dataset.get_dataloader()
|
||||||
|
|
||||||
|
total_frames = len(dataloader)
|
||||||
|
|
||||||
|
predict_mode = "fast" if config.get("use_fast_tokenizer", False) else "diffusion"
|
||||||
|
action_dim = 20 if predict_mode == "diffusion" else origin_action_dim
|
||||||
|
gt_traj = torch.zeros((total_frames, origin_action_dim))
|
||||||
|
pred_traj = torch.zeros((total_frames, origin_action_dim))
|
||||||
|
|
||||||
|
# use tqdm to show the progress
|
||||||
|
for idx, batch in tqdm(
|
||||||
|
enumerate(dataloader), total=total_frames, desc="predicting"
|
||||||
|
):
|
||||||
|
if idx % pred_horizon == 0 and idx + pred_horizon < total_frames:
|
||||||
|
batch = batch.to("cuda")
|
||||||
|
with torch.no_grad():
|
||||||
|
outputs = model(
|
||||||
|
**batch,
|
||||||
|
action_dim=action_dim,
|
||||||
|
pred_horizon=pred_horizon,
|
||||||
|
mode="predict",
|
||||||
|
predict_mode=predict_mode,
|
||||||
|
)
|
||||||
|
pred_traj[idx : idx + pred_horizon] = (
|
||||||
|
outputs["predict_action"][:, :, :origin_action_dim]
|
||||||
|
.detach()
|
||||||
|
.cpu()
|
||||||
|
.squeeze(0)
|
||||||
|
)
|
||||||
|
|
||||||
|
# Denormalize ground truth actions
|
||||||
|
gt_action_chunk = batch["action_chunk"][:, :, :origin_action_dim]
|
||||||
|
dof_mask = batch["dof_mask"].to(gt_action_chunk.dtype)
|
||||||
|
denormalized_gt = (
|
||||||
|
model.action_preprocessor.normalizer_action.unnormalize_data(
|
||||||
|
gt_action_chunk,
|
||||||
|
[lerobot_config.get("repo_id", "physical-intelligence/libero")],
|
||||||
|
dof_mask,
|
||||||
|
).squeeze(0)
|
||||||
)
|
)
|
||||||
pred_traj[idx : idx + pred_horizon] = outputs["predict_action"].detach().cpu()
|
gt_traj[idx : idx + pred_horizon] = denormalized_gt.detach().cpu()
|
||||||
|
|
||||||
# Denormalize ground truth actions
|
gt_traj_np = gt_traj.numpy()
|
||||||
gt_action_chunk = batch["action_chunk"][:, :, :action_dim]
|
pred_traj_np = pred_traj.numpy()
|
||||||
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
|
|
||||||
)
|
|
||||||
gt_traj[idx : idx + pred_horizon] = denormalized_gt.detach().cpu()
|
|
||||||
|
|
||||||
|
timesteps = gt_traj.shape[0]
|
||||||
|
|
||||||
gt_traj_np = gt_traj.numpy()
|
fig, axs = plt.subplots(
|
||||||
pred_traj_np = pred_traj.numpy()
|
origin_action_dim, 1, figsize=(15, 5 * origin_action_dim), sharex=True
|
||||||
|
)
|
||||||
|
fig.suptitle("Action Comparison for lerobot", fontsize=16)
|
||||||
|
|
||||||
timesteps = gt_traj.shape[0]
|
for i in range(origin_action_dim):
|
||||||
|
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].set_ylabel(f"Action Dim {i+1}")
|
||||||
|
axs[i].legend()
|
||||||
|
axs[i].grid(True)
|
||||||
|
|
||||||
fig, axs = plt.subplots(action_dim, 1, figsize=(15, 5 * action_dim), sharex=True)
|
axs[-1].set_xlabel("Timestep")
|
||||||
fig.suptitle("Action Comparison for lerobot", fontsize=16)
|
plt.tight_layout(rect=[0, 0.03, 1, 0.95])
|
||||||
|
os.makedirs(save_dir, exist_ok=True)
|
||||||
for i in range(action_dim):
|
save_path = os.path.join(save_dir, "lerobot_comparison.png")
|
||||||
axs[i].plot(range(timesteps), gt_traj_np[:, i], label="Ground Truth")
|
plt.savefig(save_path)
|
||||||
axs[i].plot(range(timesteps), pred_traj_np[:, i], label="Prediction")
|
print(f"Saved plot to {save_path}")
|
||||||
axs[i].set_ylabel(f"Action Dim {i+1}")
|
plt.close()
|
||||||
axs[i].legend()
|
|
||||||
axs[i].grid(True)
|
|
||||||
|
|
||||||
axs[-1].set_xlabel("Timestep")
|
|
||||||
plt.tight_layout(rect=[0, 0.03, 1, 0.95])
|
|
||||||
os.makedirs(save_dir, exist_ok=True)
|
|
||||||
plt.savefig(os.path.join(save_dir, "lerobot_comparison.png"))
|
|
||||||
plt.close()
|
|
||||||
|
|||||||
@@ -0,0 +1,161 @@
|
|||||||
|
# This file is copied from openpi
|
||||||
|
import json
|
||||||
|
import pathlib
|
||||||
|
|
||||||
|
import numpy as np
|
||||||
|
import numpydantic
|
||||||
|
import pydantic
|
||||||
|
|
||||||
|
|
||||||
|
@pydantic.dataclasses.dataclass
|
||||||
|
class NormStats:
|
||||||
|
mean: numpydantic.NDArray
|
||||||
|
std: numpydantic.NDArray
|
||||||
|
q01: numpydantic.NDArray | None = None # 1st quantile
|
||||||
|
q99: numpydantic.NDArray | None = None # 99th quantile
|
||||||
|
|
||||||
|
|
||||||
|
class RunningStats:
|
||||||
|
"""Compute running statistics of a batch of vectors."""
|
||||||
|
|
||||||
|
def __init__(self):
|
||||||
|
self._count = 0
|
||||||
|
self._mean = None
|
||||||
|
self._mean_of_squares = None
|
||||||
|
self._min = None
|
||||||
|
self._max = None
|
||||||
|
self._histograms = None
|
||||||
|
self._bin_edges = None
|
||||||
|
self._num_quantile_bins = 5000 # for computing quantiles on the fly
|
||||||
|
|
||||||
|
def update(self, batch: np.ndarray) -> None:
|
||||||
|
"""
|
||||||
|
Update the running statistics with a batch of vectors.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
vectors (np.ndarray): An array where all dimensions except the last are batch dimensions.
|
||||||
|
"""
|
||||||
|
batch = batch.reshape(-1, batch.shape[-1])
|
||||||
|
num_elements, vector_length = batch.shape
|
||||||
|
if self._count == 0:
|
||||||
|
self._mean = np.mean(batch, axis=0)
|
||||||
|
self._mean_of_squares = np.mean(batch**2, axis=0)
|
||||||
|
self._min = np.min(batch, axis=0)
|
||||||
|
self._max = np.max(batch, axis=0)
|
||||||
|
self._histograms = [
|
||||||
|
np.zeros(self._num_quantile_bins) for _ in range(vector_length)
|
||||||
|
]
|
||||||
|
self._bin_edges = [
|
||||||
|
np.linspace(
|
||||||
|
self._min[i] - 1e-10,
|
||||||
|
self._max[i] + 1e-10,
|
||||||
|
self._num_quantile_bins + 1,
|
||||||
|
)
|
||||||
|
for i in range(vector_length)
|
||||||
|
]
|
||||||
|
else:
|
||||||
|
if vector_length != self._mean.size:
|
||||||
|
raise ValueError(
|
||||||
|
"The length of new vectors does not match the initialized vector length."
|
||||||
|
)
|
||||||
|
new_max = np.max(batch, axis=0)
|
||||||
|
new_min = np.min(batch, axis=0)
|
||||||
|
max_changed = np.any(new_max > self._max)
|
||||||
|
min_changed = np.any(new_min < self._min)
|
||||||
|
self._max = np.maximum(self._max, new_max)
|
||||||
|
self._min = np.minimum(self._min, new_min)
|
||||||
|
|
||||||
|
if max_changed or min_changed:
|
||||||
|
self._adjust_histograms()
|
||||||
|
|
||||||
|
self._count += num_elements
|
||||||
|
|
||||||
|
batch_mean = np.mean(batch, axis=0)
|
||||||
|
batch_mean_of_squares = np.mean(batch**2, axis=0)
|
||||||
|
|
||||||
|
# Update running mean and mean of squares.
|
||||||
|
self._mean += (batch_mean - self._mean) * (num_elements / self._count)
|
||||||
|
self._mean_of_squares += (batch_mean_of_squares - self._mean_of_squares) * (
|
||||||
|
num_elements / self._count
|
||||||
|
)
|
||||||
|
|
||||||
|
self._update_histograms(batch)
|
||||||
|
|
||||||
|
def get_statistics(self) -> NormStats:
|
||||||
|
"""
|
||||||
|
Compute and return the statistics of the vectors processed so far.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
dict: A dictionary containing the computed statistics.
|
||||||
|
"""
|
||||||
|
if self._count < 2:
|
||||||
|
raise ValueError("Cannot compute statistics for less than 2 vectors.")
|
||||||
|
|
||||||
|
variance = self._mean_of_squares - self._mean**2
|
||||||
|
stddev = np.sqrt(np.maximum(0, variance))
|
||||||
|
q01, q99 = self._compute_quantiles([0.01, 0.99])
|
||||||
|
return NormStats(mean=self._mean, std=stddev, q01=q01, q99=q99)
|
||||||
|
|
||||||
|
def _adjust_histograms(self):
|
||||||
|
"""Adjust histograms when min or max changes."""
|
||||||
|
for i in range(len(self._histograms)):
|
||||||
|
old_edges = self._bin_edges[i]
|
||||||
|
new_edges = np.linspace(
|
||||||
|
self._min[i], self._max[i], self._num_quantile_bins + 1
|
||||||
|
)
|
||||||
|
|
||||||
|
# Redistribute the existing histogram counts to the new bins
|
||||||
|
new_hist, _ = np.histogram(
|
||||||
|
old_edges[:-1], bins=new_edges, weights=self._histograms[i]
|
||||||
|
)
|
||||||
|
|
||||||
|
self._histograms[i] = new_hist
|
||||||
|
self._bin_edges[i] = new_edges
|
||||||
|
|
||||||
|
def _update_histograms(self, batch: np.ndarray) -> None:
|
||||||
|
"""Update histograms with new vectors."""
|
||||||
|
for i in range(batch.shape[1]):
|
||||||
|
hist, _ = np.histogram(batch[:, i], bins=self._bin_edges[i])
|
||||||
|
self._histograms[i] += hist
|
||||||
|
|
||||||
|
def _compute_quantiles(self, quantiles):
|
||||||
|
"""Compute quantiles based on histograms."""
|
||||||
|
results = []
|
||||||
|
for q in quantiles:
|
||||||
|
target_count = q * self._count
|
||||||
|
q_values = []
|
||||||
|
for hist, edges in zip(self._histograms, self._bin_edges, strict=True):
|
||||||
|
cumsum = np.cumsum(hist)
|
||||||
|
idx = np.searchsorted(cumsum, target_count)
|
||||||
|
q_values.append(edges[idx])
|
||||||
|
results.append(np.array(q_values))
|
||||||
|
return results
|
||||||
|
|
||||||
|
|
||||||
|
class _NormStatsDict(pydantic.BaseModel):
|
||||||
|
norm_stats: dict[str, NormStats]
|
||||||
|
|
||||||
|
|
||||||
|
def serialize_json(norm_stats: dict[str, NormStats]) -> str:
|
||||||
|
"""Serialize the running statistics to a JSON string."""
|
||||||
|
return _NormStatsDict(norm_stats=norm_stats).model_dump_json(indent=2)
|
||||||
|
|
||||||
|
|
||||||
|
def deserialize_json(data: str) -> dict[str, NormStats]:
|
||||||
|
"""Deserialize the running statistics from a JSON string."""
|
||||||
|
return _NormStatsDict(**json.loads(data)).norm_stats
|
||||||
|
|
||||||
|
|
||||||
|
def save(directory: pathlib.Path | str, norm_stats: dict[str, NormStats]) -> None:
|
||||||
|
"""Save the normalization stats to a directory."""
|
||||||
|
path = pathlib.Path(directory) / "norm_stats.json"
|
||||||
|
path.parent.mkdir(parents=True, exist_ok=True)
|
||||||
|
path.write_text(serialize_json(norm_stats))
|
||||||
|
|
||||||
|
|
||||||
|
def load(directory: pathlib.Path | str) -> dict[str, NormStats]:
|
||||||
|
"""Load the normalization stats from a directory."""
|
||||||
|
path = pathlib.Path(directory) / "norm_stats.json"
|
||||||
|
if not path.exists():
|
||||||
|
raise FileNotFoundError(f"Norm stats file not found at: {path}")
|
||||||
|
return deserialize_json(path.read_text())
|
||||||
@@ -36,6 +36,8 @@ ACTION_DATASET_NAMES = [
|
|||||||
"taco_play",
|
"taco_play",
|
||||||
"utaustin_mutex",
|
"utaustin_mutex",
|
||||||
"viola",
|
"viola",
|
||||||
|
"physical-intelligence/libero",
|
||||||
|
"lerobot/aloha_mobile_cabinet",
|
||||||
]
|
]
|
||||||
|
|
||||||
# Supported multimodal datasets
|
# Supported multimodal datasets
|
||||||
|
|||||||
@@ -17,17 +17,10 @@ from wall_x.data.utils import (
|
|||||||
)
|
)
|
||||||
|
|
||||||
from transformers import AutoProcessor
|
from transformers import AutoProcessor
|
||||||
|
from .utils import load_norm_stats, KEY_MAPPINGS
|
||||||
|
|
||||||
T_co = TypeVar("T_co", covariant=True)
|
T_co = TypeVar("T_co", covariant=True)
|
||||||
|
|
||||||
CAMERA_KEY_MAPPINGS = {
|
|
||||||
"lerobot/aloha_mobile_cabinet": {
|
|
||||||
"observation.images.cam_high": "face_view",
|
|
||||||
"observation.images.cam_left_wrist": "left_wrist_view",
|
|
||||||
"observation.images.cam_right_wrist": "right_wrist_view",
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|
||||||
|
|
||||||
# Abstract class for dataset
|
# Abstract class for dataset
|
||||||
class Dataset(Protocol[T_co]):
|
class Dataset(Protocol[T_co]):
|
||||||
@@ -46,6 +39,8 @@ class PreprocessedDataset(Dataset[T_co]):
|
|||||||
dataset,
|
dataset,
|
||||||
config,
|
config,
|
||||||
dataload_config,
|
dataload_config,
|
||||||
|
norm_stats,
|
||||||
|
lerobot_config,
|
||||||
seed=42,
|
seed=42,
|
||||||
rank=0,
|
rank=0,
|
||||||
world_size=1,
|
world_size=1,
|
||||||
@@ -72,6 +67,8 @@ class PreprocessedDataset(Dataset[T_co]):
|
|||||||
self.config = config
|
self.config = config
|
||||||
self.use_fast_tokenizer = self.config.get("use_fast_tokenizer", False)
|
self.use_fast_tokenizer = self.config.get("use_fast_tokenizer", False)
|
||||||
self.dataload_config = dataload_config
|
self.dataload_config = dataload_config
|
||||||
|
self.norm_stats = norm_stats
|
||||||
|
self.lerobot_config = lerobot_config
|
||||||
|
|
||||||
self.data_config = X2RDataProcessingConfig().update(
|
self.data_config = X2RDataProcessingConfig().update(
|
||||||
train_test_split=self.dataload_config["train_test_split"],
|
train_test_split=self.dataload_config["train_test_split"],
|
||||||
@@ -82,7 +79,9 @@ class PreprocessedDataset(Dataset[T_co]):
|
|||||||
priority_order=self.dataload_config.get("priority_order", None),
|
priority_order=self.dataload_config.get("priority_order", None),
|
||||||
)
|
)
|
||||||
|
|
||||||
self._cam_key_mapping = CAMERA_KEY_MAPPINGS[self.hf_dataset.meta.repo_id]
|
self._cam_key_mapping = KEY_MAPPINGS[self.hf_dataset.meta.repo_id]["camera"]
|
||||||
|
self._state_key_mapping = KEY_MAPPINGS[self.hf_dataset.meta.repo_id]
|
||||||
|
self._action_key_mapping = KEY_MAPPINGS[self.hf_dataset.meta.repo_id]
|
||||||
|
|
||||||
def _vision_preprocess(self, frames):
|
def _vision_preprocess(self, frames):
|
||||||
processed_frames = []
|
processed_frames = []
|
||||||
@@ -124,14 +123,14 @@ class PreprocessedDataset(Dataset[T_co]):
|
|||||||
def __getitem__(self, index):
|
def __getitem__(self, index):
|
||||||
data = self._dataset[index]
|
data = self._dataset[index]
|
||||||
image_inputs, h, w, resize_h, resize_w = self._vision_preprocess(data)
|
image_inputs, h, w, resize_h, resize_w = self._vision_preprocess(data)
|
||||||
agent_pos = data["observation.state"]
|
agent_pos = data[self._state_key_mapping["state"]]
|
||||||
action = data["action"]
|
action = data[self._action_key_mapping["action"]]
|
||||||
frame_index = data["frame_index"]
|
frame_index = data["frame_index"]
|
||||||
instruction_info = {"instruction": data["task"]}
|
instruction_info = {"instruction": data["task"]}
|
||||||
generate_subtask_ratio = self.data_config.generate_subtask_ratio
|
generate_subtask_ratio = self.data_config.generate_subtask_ratio
|
||||||
complete_text, generate_subtask = get_wallx_normal_text(
|
complete_text, generate_subtask = get_wallx_normal_text(
|
||||||
instruction_info,
|
instruction_info,
|
||||||
33 - 1,
|
self.dataload_config.get("action_horizon", 33) - 1,
|
||||||
frame_index,
|
frame_index,
|
||||||
self.data_config.priority_order,
|
self.data_config.priority_order,
|
||||||
self._cam_key_mapping,
|
self._cam_key_mapping,
|
||||||
@@ -189,7 +188,7 @@ class PreprocessedDataset(Dataset[T_co]):
|
|||||||
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(
|
collate_fn=DataCollator(
|
||||||
self.config, self.dataload_config, self.hf_dataset.meta.stats
|
self.config, self.dataload_config, self.norm_stats, self.lerobot_config
|
||||||
),
|
),
|
||||||
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
|
||||||
@@ -225,7 +224,7 @@ class PreprocessedDataset(Dataset[T_co]):
|
|||||||
sampler=sampler,
|
sampler=sampler,
|
||||||
num_workers=num_workers,
|
num_workers=num_workers,
|
||||||
collate_fn=DataCollator(
|
collate_fn=DataCollator(
|
||||||
self.config, self.dataload_config, self.hf_dataset.meta.stats
|
self.config, self.dataload_config, self.norm_stats, self.lerobot_config
|
||||||
),
|
),
|
||||||
pin_memory=True,
|
pin_memory=True,
|
||||||
persistent_workers=num_workers > 0,
|
persistent_workers=num_workers > 0,
|
||||||
@@ -241,13 +240,16 @@ class DataCollator:
|
|||||||
_processor_cache = {}
|
_processor_cache = {}
|
||||||
_action_tokenizer_cache = {}
|
_action_tokenizer_cache = {}
|
||||||
|
|
||||||
def __init__(self, config, dataload_config, stats):
|
def __init__(self, config, dataload_config, stats, lerobot_config):
|
||||||
self.config = config
|
self.config = config
|
||||||
self.dataload_config = dataload_config
|
self.dataload_config = dataload_config
|
||||||
self.stats = stats
|
self.stats = stats
|
||||||
self.min_stat = stats["action"]["min"]
|
self.action_min_stat = stats["action"].min
|
||||||
self.max_stat = stats["action"]["max"]
|
self.action_delta = stats["action"].delta
|
||||||
self.delta = self.max_stat - self.min_stat
|
self.state_min_stat = stats["state"].min
|
||||||
|
self.state_delta = stats["state"].delta
|
||||||
|
self.lerobot_config = lerobot_config
|
||||||
|
|
||||||
self.use_fast_tokenizer = self.config.get("use_fast_tokenizer", False)
|
self.use_fast_tokenizer = self.config.get("use_fast_tokenizer", False)
|
||||||
self.load_processor()
|
self.load_processor()
|
||||||
|
|
||||||
@@ -271,10 +273,11 @@ class DataCollator:
|
|||||||
if self.config.get("padding_side", "left") == "left":
|
if self.config.get("padding_side", "left") == "left":
|
||||||
processor.tokenizer.padding_side = "left"
|
processor.tokenizer.padding_side = "left"
|
||||||
|
|
||||||
|
new_tokens = ["<|propri|>", "<|action|>"]
|
||||||
|
processor.tokenizer.add_tokens(new_tokens)
|
||||||
if self.use_fast_tokenizer and self.config.get("model_type") == "qwen2_5":
|
if self.use_fast_tokenizer and self.config.get("model_type") == "qwen2_5":
|
||||||
action_tokenizer = self._action_tokenizer_cache[action_tokenizer_path]
|
action_tokenizer = self._action_tokenizer_cache[action_tokenizer_path]
|
||||||
new_tokens = ["<|propri|>", "<|action|>"]
|
new_tokens = [
|
||||||
new_tokens += [
|
|
||||||
f"<|action_token_{i}|>" for i in range(action_tokenizer.vocab_size)
|
f"<|action_token_{i}|>" for i in range(action_tokenizer.vocab_size)
|
||||||
]
|
]
|
||||||
processor.tokenizer.add_tokens(new_tokens)
|
processor.tokenizer.add_tokens(new_tokens)
|
||||||
@@ -301,7 +304,6 @@ class DataCollator:
|
|||||||
"""
|
"""
|
||||||
Normalize action data using min-max normalization.
|
Normalize action data using min-max normalization.
|
||||||
"""
|
"""
|
||||||
delta = torch.from_numpy(delta)
|
|
||||||
delta = torch.where(delta == 0, torch.ones_like(delta), delta)
|
delta = torch.where(delta == 0, torch.ones_like(delta), delta)
|
||||||
x = (action - min_stat) / delta
|
x = (action - min_stat) / delta
|
||||||
x = x * 2 - 1
|
x = x * 2 - 1
|
||||||
@@ -318,7 +320,9 @@ class DataCollator:
|
|||||||
agent_pos = agent_pos.unsqueeze(1)
|
agent_pos = agent_pos.unsqueeze(1)
|
||||||
agent_pos_mask = (~torch.isnan(agent_pos)).float()
|
agent_pos_mask = (~torch.isnan(agent_pos)).float()
|
||||||
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.state_min_stat, self.state_delta
|
||||||
|
)
|
||||||
if agent_pos.shape[-1] != 20:
|
if agent_pos.shape[-1] != 20:
|
||||||
agent_pos = torch.cat(
|
agent_pos = torch.cat(
|
||||||
[
|
[
|
||||||
@@ -350,7 +354,9 @@ class DataCollator:
|
|||||||
action = action.unsqueeze(1)
|
action = action.unsqueeze(1)
|
||||||
dof_mask = (~torch.isnan(action)).float()
|
dof_mask = (~torch.isnan(action)).float()
|
||||||
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.action_min_stat, self.action_delta
|
||||||
|
)
|
||||||
if action.shape[-1] != 20:
|
if action.shape[-1] != 20:
|
||||||
action = torch.cat(
|
action = torch.cat(
|
||||||
[
|
[
|
||||||
@@ -393,7 +399,7 @@ class DataCollator:
|
|||||||
additional_inputs["text"],
|
additional_inputs["text"],
|
||||||
additional_inputs["action_chunk"],
|
additional_inputs["action_chunk"],
|
||||||
self.train_action_tokenizer if self.use_fast_tokenizer else None,
|
self.train_action_tokenizer if self.use_fast_tokenizer else None,
|
||||||
["x2_normal"] * additional_inputs["text"].__len__(),
|
[self.lerobot_config["repo_id"]] * additional_inputs["text"].__len__(),
|
||||||
additional_inputs["dof_mask"],
|
additional_inputs["dof_mask"],
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -415,7 +421,9 @@ class DataCollator:
|
|||||||
|
|
||||||
inputs.update(additional_inputs)
|
inputs.update(additional_inputs)
|
||||||
|
|
||||||
inputs["dataset_names"] = ["x2_normal"] * inputs["action_chunk"].shape[0]
|
inputs["dataset_names"] = [self.lerobot_config["repo_id"]] * inputs[
|
||||||
|
"action_chunk"
|
||||||
|
].shape[0]
|
||||||
|
|
||||||
return inputs
|
return inputs
|
||||||
|
|
||||||
@@ -447,18 +455,24 @@ def load_lerobot_data(
|
|||||||
|
|
||||||
dataload_config = get_data_configs(config["data"])
|
dataload_config = get_data_configs(config["data"])
|
||||||
|
|
||||||
# repo_id = "lerobot/aloha_mobile_cabinet"
|
repo_id = lerobot_config.get("repo_id", None)
|
||||||
repo_id = lerobot_config.get("repo_id", "lerobot/aloha_mobile_cabinet")
|
assert repo_id is not None, "repo id is required"
|
||||||
root = lerobot_config.get("root", None)
|
root = lerobot_config.get("root", None)
|
||||||
meta_info = LeRobotDatasetMetadata(repo_id)
|
meta_info = LeRobotDatasetMetadata(repo_id, root=root)
|
||||||
dataset_fps = meta_info.fps
|
dataset_fps = meta_info.fps
|
||||||
episodes_num = meta_info.total_episodes
|
episodes_num = meta_info.total_episodes
|
||||||
|
|
||||||
|
norm_stats_path = config.get("norm_stats_path", None)
|
||||||
|
assert (
|
||||||
|
norm_stats_path is not None
|
||||||
|
), "norm stats is required, please refer to 'wall-x/scripts/compute_norm_stats.py' to compute stats"
|
||||||
|
norm_stats = load_norm_stats(norm_stats_path, repo_id)
|
||||||
|
|
||||||
delta_timestamps = {
|
delta_timestamps = {
|
||||||
# action chunk
|
# action chunk
|
||||||
"action": [
|
KEY_MAPPINGS[repo_id]["action"]: [
|
||||||
t / dataset_fps
|
t / dataset_fps
|
||||||
for t in range(dataload_config.get("action_horizon", 32) - 1)
|
for t in range(dataload_config.get("action_horizon", 33) - 1)
|
||||||
],
|
],
|
||||||
}
|
}
|
||||||
batch_size = config.get("batch_size_per_gpu", 8)
|
batch_size = config.get("batch_size_per_gpu", 8)
|
||||||
@@ -486,6 +500,8 @@ def load_lerobot_data(
|
|||||||
train_dataset,
|
train_dataset,
|
||||||
config,
|
config,
|
||||||
dataload_config,
|
dataload_config,
|
||||||
|
norm_stats,
|
||||||
|
lerobot_config,
|
||||||
seed=seed,
|
seed=seed,
|
||||||
rank=rank,
|
rank=rank,
|
||||||
world_size=world_size,
|
world_size=world_size,
|
||||||
@@ -567,11 +583,15 @@ def get_data_configs(config):
|
|||||||
|
|
||||||
|
|
||||||
class TestDataset(PreprocessedDataset):
|
class TestDataset(PreprocessedDataset):
|
||||||
def __init__(self, dataset, config, dataload_config, seed=42):
|
def __init__(
|
||||||
|
self, dataset, config, dataload_config, norm_stats, lerobot_config, seed=42
|
||||||
|
):
|
||||||
super().__init__(
|
super().__init__(
|
||||||
dataset,
|
dataset,
|
||||||
config,
|
config,
|
||||||
dataload_config,
|
dataload_config,
|
||||||
|
norm_stats,
|
||||||
|
lerobot_config,
|
||||||
seed=seed,
|
seed=seed,
|
||||||
rank=0,
|
rank=0,
|
||||||
world_size=1,
|
world_size=1,
|
||||||
@@ -587,7 +607,7 @@ class TestDataset(PreprocessedDataset):
|
|||||||
self,
|
self,
|
||||||
batch_size=1,
|
batch_size=1,
|
||||||
collate_fn=DataCollator(
|
collate_fn=DataCollator(
|
||||||
self.config, self.dataload_config, self.hf_dataset.meta.stats
|
self.config, self.dataload_config, self.norm_stats, self.lerobot_config
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -614,29 +634,41 @@ def load_test_dataset(
|
|||||||
# Set seed for reproducibility
|
# Set seed for reproducibility
|
||||||
torch.manual_seed(seed)
|
torch.manual_seed(seed)
|
||||||
|
|
||||||
dataset_fps = 50
|
repo_id = lerobot_config.get("repo_id", None)
|
||||||
|
assert repo_id is not None, "repo id is required"
|
||||||
|
root = lerobot_config.get("root", None)
|
||||||
|
meta_info = LeRobotDatasetMetadata(repo_id, root=root)
|
||||||
|
dataset_fps = meta_info.fps
|
||||||
dataload_config = get_data_configs(config["data"])
|
dataload_config = get_data_configs(config["data"])
|
||||||
|
|
||||||
|
norm_stats_path = config.get("norm_stats_path", None)
|
||||||
|
assert (
|
||||||
|
norm_stats_path is not None
|
||||||
|
), "norm stats is required, please refer to 'wall-x/scripts/compute_norm_stats.py' to compute stats"
|
||||||
|
norm_stats = load_norm_stats(norm_stats_path, repo_id)
|
||||||
|
|
||||||
delta_timestamps = {
|
delta_timestamps = {
|
||||||
# action chunk
|
# action chunk
|
||||||
"action": [
|
KEY_MAPPINGS[repo_id]["action"]: [
|
||||||
t / dataset_fps
|
t / dataset_fps
|
||||||
for t in range(dataload_config.get("action_horizon", 32) - 1)
|
for t in range(dataload_config.get("action_horizon", 33) - 1)
|
||||||
],
|
],
|
||||||
}
|
}
|
||||||
|
|
||||||
repo_id = lerobot_config.get("repo_id", "lerobot/aloha_mobile_cabinet")
|
|
||||||
dataset = LeRobotDataset(
|
dataset = LeRobotDataset(
|
||||||
repo_id,
|
repo_id,
|
||||||
episodes=[episode],
|
episodes=[episode],
|
||||||
delta_timestamps=delta_timestamps,
|
delta_timestamps=delta_timestamps,
|
||||||
video_backend="pyav",
|
video_backend="pyav",
|
||||||
|
root=root,
|
||||||
)
|
)
|
||||||
|
|
||||||
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 = TestDataset(dataset, config, dataload_config, seed=seed)
|
dataset = TestDataset(
|
||||||
|
dataset, config, dataload_config, norm_stats, lerobot_config, seed=seed
|
||||||
|
)
|
||||||
|
|
||||||
return dataset
|
return dataset
|
||||||
|
|||||||
@@ -11,7 +11,28 @@ import random
|
|||||||
from collections import OrderedDict
|
from collections import OrderedDict
|
||||||
from typing import List, Dict, Any, Optional, Union, Tuple
|
from typing import List, Dict, Any, Optional, Union, Tuple
|
||||||
from transformers import BatchFeature
|
from transformers import BatchFeature
|
||||||
|
from dataclasses import dataclass
|
||||||
|
import json
|
||||||
|
|
||||||
|
KEY_MAPPINGS = {
|
||||||
|
"lerobot/aloha_mobile_cabinet": {
|
||||||
|
"camera": {
|
||||||
|
"observation.images.cam_high": "face_view",
|
||||||
|
"observation.images.cam_left_wrist": "left_wrist_view",
|
||||||
|
"observation.images.cam_right_wrist": "right_wrist_view",
|
||||||
|
},
|
||||||
|
"state": "observation.state",
|
||||||
|
"action": "action",
|
||||||
|
},
|
||||||
|
"physical-intelligence/libero": {
|
||||||
|
"camera": {
|
||||||
|
"image": "face_view",
|
||||||
|
"wrist_image": "left_wrist_view",
|
||||||
|
},
|
||||||
|
"state": "state",
|
||||||
|
"action": "actions",
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
CAMERA_NAME_MAPPING = {
|
CAMERA_NAME_MAPPING = {
|
||||||
"face_view": "front view",
|
"face_view": "front view",
|
||||||
@@ -609,3 +630,35 @@ def replace_action_token(
|
|||||||
text = [t.replace("<|action_fast|><|im_end|>\n", "") for t in text]
|
text = [t.replace("<|action_fast|><|im_end|>\n", "") for t in text]
|
||||||
|
|
||||||
return text
|
return text
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class NormStats:
|
||||||
|
min: torch.Tensor
|
||||||
|
max: torch.Tensor
|
||||||
|
delta: torch.Tensor
|
||||||
|
|
||||||
|
|
||||||
|
def load_norm_stats(norm_stats_path, dataset_name):
|
||||||
|
with open(norm_stats_path, "r") as f:
|
||||||
|
norm_stats = json.load(f)
|
||||||
|
action_key = KEY_MAPPINGS[dataset_name]["action"]
|
||||||
|
state_key = KEY_MAPPINGS[dataset_name]["state"]
|
||||||
|
q01 = torch.tensor(norm_stats["norm_stats"][action_key]["q01"])
|
||||||
|
q99 = torch.tensor(norm_stats["norm_stats"][action_key]["q99"])
|
||||||
|
delta = q99 - q01
|
||||||
|
action_norm_stats = NormStats(
|
||||||
|
min=q01,
|
||||||
|
max=q99,
|
||||||
|
delta=delta,
|
||||||
|
)
|
||||||
|
q01 = torch.tensor(norm_stats["norm_stats"][state_key]["q01"])
|
||||||
|
q99 = torch.tensor(norm_stats["norm_stats"][state_key]["q99"])
|
||||||
|
delta = q99 - q01
|
||||||
|
state_norm_stats = NormStats(
|
||||||
|
min=q01,
|
||||||
|
max=q99,
|
||||||
|
delta=delta,
|
||||||
|
)
|
||||||
|
|
||||||
|
return {"action": action_norm_stats, "state": state_norm_stats}
|
||||||
|
|||||||
@@ -3,6 +3,7 @@ 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
|
||||||
|
import logging
|
||||||
|
|
||||||
|
|
||||||
class Normalizer(nn.Module):
|
class Normalizer(nn.Module):
|
||||||
@@ -14,6 +15,16 @@ class Normalizer(nn.Module):
|
|||||||
normalization to map actions to the [-1, 1] range.
|
normalization to map actions to the [-1, 1] range.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
def _pad_to_action_dim(self, xs, action_dim):
|
||||||
|
"""
|
||||||
|
Pad the action data to the action dimension.
|
||||||
|
"""
|
||||||
|
if xs.shape[-1] < action_dim:
|
||||||
|
padding_shape = list(xs.shape)
|
||||||
|
padding_shape[-1] = action_dim - padding_shape[-1]
|
||||||
|
xs = torch.cat([xs, torch.zeros(padding_shape).to(xs.device)], dim=-1)
|
||||||
|
return xs
|
||||||
|
|
||||||
def __init__(self, action_statistic_dof, dof_config):
|
def __init__(self, action_statistic_dof, dof_config):
|
||||||
"""
|
"""
|
||||||
Initialize the normalizer with robot-specific action statistics.
|
Initialize the normalizer with robot-specific action statistics.
|
||||||
@@ -25,6 +36,8 @@ class Normalizer(nn.Module):
|
|||||||
super(Normalizer, self).__init__()
|
super(Normalizer, self).__init__()
|
||||||
|
|
||||||
action_statistic = {}
|
action_statistic = {}
|
||||||
|
# hard code the action dimension to 20
|
||||||
|
action_dim = 20
|
||||||
|
|
||||||
# Process statistics for each robot
|
# Process statistics for each robot
|
||||||
for robot_name in action_statistic_dof.keys():
|
for robot_name in action_statistic_dof.keys():
|
||||||
@@ -39,11 +52,17 @@ class Normalizer(nn.Module):
|
|||||||
all_dof_delta.extend(action_statistic_dof[robot_name][k]["delta"])
|
all_dof_delta.extend(action_statistic_dof[robot_name][k]["delta"])
|
||||||
else:
|
else:
|
||||||
# Use default values if statistics not available
|
# Use default values if statistics not available
|
||||||
|
# raise ValueError(f"Statistics not available for {k} of {robot_name}")
|
||||||
|
logging.warning(
|
||||||
|
f"Statistics not available for {k} of {robot_name}, using default values"
|
||||||
|
)
|
||||||
all_dof_min.extend([0.0] * dof_config[k])
|
all_dof_min.extend([0.0] * dof_config[k])
|
||||||
all_dof_delta.extend([1.0] * dof_config[k])
|
all_dof_delta.extend([1.0] * dof_config[k])
|
||||||
|
|
||||||
all_dof_min = torch.tensor(all_dof_min)
|
all_dof_min = self._pad_to_action_dim(torch.tensor(all_dof_min), action_dim)
|
||||||
all_dof_delta = torch.tensor(all_dof_delta)
|
all_dof_delta = self._pad_to_action_dim(
|
||||||
|
torch.tensor(all_dof_delta), action_dim
|
||||||
|
)
|
||||||
action_statistic[robot_name]["min"] = all_dof_min
|
action_statistic[robot_name]["min"] = all_dof_min
|
||||||
action_statistic[robot_name]["delta"] = all_dof_delta
|
action_statistic[robot_name]["delta"] = all_dof_delta
|
||||||
|
|
||||||
@@ -61,7 +80,7 @@ class Normalizer(nn.Module):
|
|||||||
}
|
}
|
||||||
)
|
)
|
||||||
|
|
||||||
def normalize_data(self, xs, dataset_names):
|
def normalize_data(self, xs, dataset_names, dof_mask=None):
|
||||||
"""
|
"""
|
||||||
Normalize action data to [-1, 1] range using robot-specific statistics.
|
Normalize action data to [-1, 1] range using robot-specific statistics.
|
||||||
|
|
||||||
@@ -75,10 +94,19 @@ class Normalizer(nn.Module):
|
|||||||
new_xs = []
|
new_xs = []
|
||||||
# Filter out multimodal dataset entries
|
# Filter out multimodal dataset entries
|
||||||
dataset_names = [name for name in dataset_names if name != "x2_multimodal"]
|
dataset_names = [name for name in dataset_names if name != "x2_multimodal"]
|
||||||
|
dof_mask = dof_mask if dof_mask is not None else [None] * len(xs)
|
||||||
|
|
||||||
for x, dataset_name in zip(xs, dataset_names):
|
for x, dataset_name, mask in zip(xs, dataset_names, dof_mask):
|
||||||
|
# Apply DOF mask if provided
|
||||||
|
if mask is not None:
|
||||||
|
mask = mask[0].bool()
|
||||||
|
action_space_delta = self.delta[dataset_name][mask]
|
||||||
|
action_space_min = self.min[dataset_name][mask]
|
||||||
|
else:
|
||||||
|
action_space_delta = self.delta[dataset_name]
|
||||||
|
action_space_min = self.min[dataset_name]
|
||||||
# Apply min-max normalization
|
# Apply min-max normalization
|
||||||
x = (x - self.min[dataset_name]) / (self.delta[dataset_name])
|
x = (x - action_space_min) / (action_space_delta)
|
||||||
# Scale to [-1, 1] range
|
# Scale to [-1, 1] range
|
||||||
x = x * 2 - 1
|
x = x * 2 - 1
|
||||||
# Clamp to ensure bounds
|
# Clamp to ensure bounds
|
||||||
@@ -210,9 +238,21 @@ class ActionProcessor(nn.Module):
|
|||||||
self.hidden_size = config.hidden_size
|
self.hidden_size = config.hidden_size
|
||||||
|
|
||||||
# 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.customized_dof_config
|
||||||
|
if hasattr(config, "customized_dof_config")
|
||||||
|
else config.dof_config
|
||||||
|
),
|
||||||
|
)
|
||||||
self.normalizer_propri = Normalizer(
|
self.normalizer_propri = Normalizer(
|
||||||
action_statistic_dof, config.agent_pos_config
|
action_statistic_dof,
|
||||||
|
(
|
||||||
|
config.customized_agent_pos_config
|
||||||
|
if hasattr(config, "customized_agent_pos_config")
|
||||||
|
else config.agent_pos_config
|
||||||
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
# Proprioception projection layer (includes history/current state)
|
# Proprioception projection layer (includes history/current state)
|
||||||
|
|||||||
@@ -44,7 +44,9 @@ from wall_x.model.qwen2_5_based.modeling_qwen2_5_vl import (
|
|||||||
Qwen2_5_VLSdpaAttention,
|
Qwen2_5_VLSdpaAttention,
|
||||||
)
|
)
|
||||||
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.utils.constant import action_statistic_dof
|
||||||
|
from wall_x.data.utils import load_norm_stats
|
||||||
|
from pprint import pprint
|
||||||
|
|
||||||
logger = logging.get_logger(__name__)
|
logger = logging.get_logger(__name__)
|
||||||
|
|
||||||
@@ -744,10 +746,77 @@ class Qwen2_5_VLMoEForAction(Qwen2_5_VLForConditionalGeneration):
|
|||||||
config_class = Qwen2_5_VLConfig
|
config_class = Qwen2_5_VLConfig
|
||||||
_no_split_modules = ["Qwen2_5_VLDecoderLayer_with_MoE", "Qwen2_5_VLVisionBlock"]
|
_no_split_modules = ["Qwen2_5_VLDecoderLayer_with_MoE", "Qwen2_5_VLVisionBlock"]
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def _set_customized_config(cls, config):
|
||||||
|
"""
|
||||||
|
Processing norm_stats.json and reconstruct the DoF mapping
|
||||||
|
"""
|
||||||
|
dataload_config = config["data"]
|
||||||
|
if not dataload_config.get("use_lerobot", False):
|
||||||
|
raise NotImplementedError(
|
||||||
|
"Not implemented for non-lerobot dataset currently"
|
||||||
|
)
|
||||||
|
|
||||||
|
enable_customized_robot_config = config.get(
|
||||||
|
"enable_customized_robot_config", False
|
||||||
|
)
|
||||||
|
assert (
|
||||||
|
enable_customized_robot_config
|
||||||
|
), "enable_customized_robot_config must be true when use lerobot dataset"
|
||||||
|
|
||||||
|
customized_dof_config = config["customized_robot_config"][
|
||||||
|
"customized_dof_config"
|
||||||
|
]
|
||||||
|
customized_agent_pos_config = config["customized_robot_config"][
|
||||||
|
"customized_agent_pos_config"
|
||||||
|
]
|
||||||
|
norm_stats_path = config["norm_stats_path"]
|
||||||
|
norm_stats = load_norm_stats(
|
||||||
|
norm_stats_path, config["data"]["lerobot_config"]["repo_id"]
|
||||||
|
)
|
||||||
|
action_min = norm_stats["action"].min.numpy().tolist()
|
||||||
|
action_delta = norm_stats["action"].delta.numpy().tolist()
|
||||||
|
state_min = norm_stats["state"].min.numpy().tolist()
|
||||||
|
state_delta = norm_stats["state"].delta.numpy().tolist()
|
||||||
|
|
||||||
|
name = config["customized_robot_config"]["name"]
|
||||||
|
|
||||||
|
dof_key = []
|
||||||
|
agent_pos_key = []
|
||||||
|
dof_value = []
|
||||||
|
agent_pos_value = []
|
||||||
|
stats_dict = {}
|
||||||
|
for k, v in customized_dof_config.items():
|
||||||
|
dof_key.append(k)
|
||||||
|
dof_value.append(v)
|
||||||
|
for k, v in customized_agent_pos_config.items():
|
||||||
|
agent_pos_key.append(k)
|
||||||
|
agent_pos_value.append(v)
|
||||||
|
|
||||||
|
dof_idx = np.array([0] + dof_value).cumsum()
|
||||||
|
for i in range(len(dof_idx) - 1):
|
||||||
|
stats_dict[dof_key[i]] = {
|
||||||
|
"min": action_min[dof_idx[i] : dof_idx[i + 1]],
|
||||||
|
"delta": action_delta[dof_idx[i] : dof_idx[i + 1]],
|
||||||
|
}
|
||||||
|
|
||||||
|
agent_pos_idx = np.array([0] + agent_pos_value).cumsum()
|
||||||
|
for i in range(len(agent_pos_idx) - 1):
|
||||||
|
stats_dict[agent_pos_key[i]] = {
|
||||||
|
"min": state_min[agent_pos_idx[i] : agent_pos_idx[i + 1]],
|
||||||
|
"delta": state_delta[agent_pos_idx[i] : agent_pos_idx[i + 1]],
|
||||||
|
}
|
||||||
|
|
||||||
|
action_statistic_dof[name] = stats_dict
|
||||||
|
|
||||||
|
print("Customized robot config added")
|
||||||
|
pprint(action_statistic_dof)
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def from_pretrained(
|
def from_pretrained(
|
||||||
cls,
|
cls,
|
||||||
pretrained_model_path,
|
pretrained_model_path,
|
||||||
|
train_config,
|
||||||
config_path=None,
|
config_path=None,
|
||||||
processor_path=None,
|
processor_path=None,
|
||||||
action_tokenizer_path=None,
|
action_tokenizer_path=None,
|
||||||
@@ -766,7 +835,6 @@ class Qwen2_5_VLMoEForAction(Qwen2_5_VLForConditionalGeneration):
|
|||||||
Returns:
|
Returns:
|
||||||
Qwen2_5_VLMoEForAction: Loaded model instance
|
Qwen2_5_VLMoEForAction: Loaded model instance
|
||||||
"""
|
"""
|
||||||
|
|
||||||
# Load model components from pretrained path
|
# Load model components from pretrained path
|
||||||
config_path = os.path.join(pretrained_model_path, "config.json")
|
config_path = os.path.join(pretrained_model_path, "config.json")
|
||||||
config = cls.config_class.from_pretrained(config_path)
|
config = cls.config_class.from_pretrained(config_path)
|
||||||
@@ -776,6 +844,18 @@ class Qwen2_5_VLMoEForAction(Qwen2_5_VLForConditionalGeneration):
|
|||||||
action_tokenizer_path, trust_remote_code=True
|
action_tokenizer_path, trust_remote_code=True
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# Set the customized robot configuration to ensure consistency between cross-embodiment
|
||||||
|
# representations and the Wall-X action dimensionality.
|
||||||
|
cls._set_customized_config(train_config)
|
||||||
|
customized_dof_config = train_config["customized_robot_config"][
|
||||||
|
"customized_dof_config"
|
||||||
|
]
|
||||||
|
customized_agent_pos_config = train_config["customized_robot_config"][
|
||||||
|
"customized_agent_pos_config"
|
||||||
|
]
|
||||||
|
setattr(config, "customized_dof_config", customized_dof_config)
|
||||||
|
setattr(config, "customized_agent_pos_config", customized_agent_pos_config)
|
||||||
|
|
||||||
# Initialize model with configuration and processor
|
# Initialize model with configuration and processor
|
||||||
model = cls(config, processor=processor, **kwargs)
|
model = cls(config, processor=processor, **kwargs)
|
||||||
|
|
||||||
@@ -789,6 +869,14 @@ class Qwen2_5_VLMoEForAction(Qwen2_5_VLForConditionalGeneration):
|
|||||||
state_dict = {}
|
state_dict = {}
|
||||||
for file in safetensor_files:
|
for file in safetensor_files:
|
||||||
sd = load_file(file, device="cpu")
|
sd = load_file(file, device="cpu")
|
||||||
|
# filter normalizer statistic params
|
||||||
|
del_keys = []
|
||||||
|
for key in sd.keys():
|
||||||
|
if "action_preprocessor.normalizer" in key:
|
||||||
|
print(f"filter load model weight {key}")
|
||||||
|
del_keys.append(key)
|
||||||
|
for key in del_keys:
|
||||||
|
del sd[key]
|
||||||
state_dict.update(sd)
|
state_dict.update(sd)
|
||||||
|
|
||||||
model.load_state_dict(state_dict, strict=False)
|
model.load_state_dict(state_dict, strict=False)
|
||||||
@@ -860,11 +948,14 @@ class Qwen2_5_VLMoEForAction(Qwen2_5_VLForConditionalGeneration):
|
|||||||
"""
|
"""
|
||||||
# Create list of fast action token IDs
|
# Create list of fast action token IDs
|
||||||
fast_action_token_list = []
|
fast_action_token_list = []
|
||||||
for i in range(self.processor.tokenizer.init_kwargs["action_token_vocab_size"]):
|
if self.use_fast_tokenizer:
|
||||||
action_token_id = self.processor.tokenizer.convert_tokens_to_ids(
|
for i in range(
|
||||||
f"<|action_token_{i}|>"
|
self.processor.tokenizer.init_kwargs["action_token_vocab_size"]
|
||||||
)
|
):
|
||||||
fast_action_token_list.append(action_token_id)
|
action_token_id = self.processor.tokenizer.convert_tokens_to_ids(
|
||||||
|
f"<|action_token_{i}|>"
|
||||||
|
)
|
||||||
|
fast_action_token_list.append(action_token_id)
|
||||||
|
|
||||||
# Get special action token IDs
|
# Get special action token IDs
|
||||||
action_token_id = self.processor.tokenizer.convert_tokens_to_ids("<|action|>")
|
action_token_id = self.processor.tokenizer.convert_tokens_to_ids("<|action|>")
|
||||||
|
|||||||
@@ -263,7 +263,6 @@ class QwenVlAct_Trainer:
|
|||||||
self.train_dataloader = self.dataset.get_train_dataloader()
|
self.train_dataloader = self.dataset.get_train_dataloader()
|
||||||
|
|
||||||
self.model.train()
|
self.model.train()
|
||||||
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"]
|
||||||
@@ -341,7 +340,7 @@ class QwenVlAct_Trainer:
|
|||||||
self.timers("optimizer").stop()
|
self.timers("optimizer").stop()
|
||||||
|
|
||||||
# Update global step and learning rate after gradient accumulation
|
# Update global step and learning rate after gradient accumulation
|
||||||
if (i + 1) % grad_accum_steps == 0:
|
if self.accelerator.sync_gradients:
|
||||||
self.lr_scheduler.step()
|
self.lr_scheduler.step()
|
||||||
self.global_step += 1
|
self.global_step += 1
|
||||||
lr = self.lr_scheduler.get_last_lr()[0]
|
lr = self.lr_scheduler.get_last_lr()[0]
|
||||||
@@ -522,7 +521,12 @@ class QwenVlAct_Trainer:
|
|||||||
if model_type == "wall-oss":
|
if model_type == "wall-oss":
|
||||||
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},
|
train_config=self.config,
|
||||||
|
action_tokenizer_path=(
|
||||||
|
self.config["action_tokenizer_path"]
|
||||||
|
if self.use_fast_tokenizer
|
||||||
|
else None
|
||||||
|
),
|
||||||
)
|
)
|
||||||
self.processor = model.processor
|
self.processor = model.processor
|
||||||
model = model.to(torch.bfloat16)
|
model = model.to(torch.bfloat16)
|
||||||
@@ -535,14 +539,15 @@ class QwenVlAct_Trainer:
|
|||||||
self.processor = AutoProcessor.from_pretrained(
|
self.processor = AutoProcessor.from_pretrained(
|
||||||
self.config["pretrained_wallx_path"], use_fast=True
|
self.config["pretrained_wallx_path"], use_fast=True
|
||||||
)
|
)
|
||||||
|
new_tokens = ["<|propri|>", "<|action|>"]
|
||||||
|
self.processor.tokenizer.add_tokens(new_tokens)
|
||||||
if self.config.get("use_fast_tokenizer", False):
|
if self.config.get("use_fast_tokenizer", False):
|
||||||
action_tokenizer_path = self.config["action_tokenizer_path"]
|
action_tokenizer_path = self.config["action_tokenizer_path"]
|
||||||
action_tokenizer = AutoProcessor.from_pretrained(
|
action_tokenizer = AutoProcessor.from_pretrained(
|
||||||
action_tokenizer_path, trust_remote_code=True
|
action_tokenizer_path, trust_remote_code=True
|
||||||
)
|
)
|
||||||
# process for use fast
|
# process for use fast
|
||||||
new_tokens = ["<|propri|>", "<|action|>"]
|
new_tokens = [
|
||||||
new_tokens += [
|
|
||||||
f"<|action_token_{i}|>" for i in range(action_tokenizer.vocab_size)
|
f"<|action_token_{i}|>" for i in range(action_tokenizer.vocab_size)
|
||||||
]
|
]
|
||||||
self.processor.tokenizer.add_tokens(new_tokens)
|
self.processor.tokenizer.add_tokens(new_tokens)
|
||||||
@@ -557,6 +562,19 @@ class QwenVlAct_Trainer:
|
|||||||
action_tokenizer.vocab_size
|
action_tokenizer.vocab_size
|
||||||
)
|
)
|
||||||
self.processor.action_processor = action_tokenizer
|
self.processor.action_processor = action_tokenizer
|
||||||
|
|
||||||
|
# Set the customized robot configuration to ensure consistency between cross-embodiment
|
||||||
|
# representations and the Wall-X action dimensionality.
|
||||||
|
Qwen2_5_VLMoEForAction._set_customized_config(self.config)
|
||||||
|
customized_dof_config = self.config["customized_robot_config"][
|
||||||
|
"customized_dof_config"
|
||||||
|
]
|
||||||
|
customized_agent_pos_config = self.config["customized_robot_config"][
|
||||||
|
"customized_agent_pos_config"
|
||||||
|
]
|
||||||
|
setattr(config, "customized_dof_config", customized_dof_config)
|
||||||
|
setattr(config, "customized_agent_pos_config", customized_agent_pos_config)
|
||||||
|
|
||||||
model = Qwen2_5_VLMoEForAction(
|
model = Qwen2_5_VLMoEForAction(
|
||||||
config,
|
config,
|
||||||
self.use_fast_tokenizer,
|
self.use_fast_tokenizer,
|
||||||
@@ -803,7 +821,8 @@ class QwenVlAct_Trainer:
|
|||||||
# merge checkpoint section to merge the weights into a single safetensors if needed.
|
# merge checkpoint section to merge the weights into a single safetensors if needed.
|
||||||
self.accelerator.save_state(ckpt_path)
|
self.accelerator.save_state(ckpt_path)
|
||||||
|
|
||||||
self.processor.save_pretrained(os.path.join(ckpt_path, "processor"))
|
if self.accelerator.is_main_process:
|
||||||
|
self.processor.save_pretrained(os.path.join(ckpt_path, "processor"))
|
||||||
|
|
||||||
# Save current iteration steps for dataset resuming
|
# Save current iteration steps for dataset resuming
|
||||||
if step != 0:
|
if step != 0:
|
||||||
@@ -845,7 +864,7 @@ class QwenVlAct_Trainer:
|
|||||||
# 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)
|
||||||
|
|
||||||
self.print_rank0(f"Resumed from checkpoint: {checkpoint_path}")
|
self.print_rank0(f"\033[32mResumed from checkpoint: {checkpoint_path}\033[0m")
|
||||||
|
|
||||||
def _load_fsdp_state_dict_with_distribute_tensor(self):
|
def _load_fsdp_state_dict_with_distribute_tensor(self):
|
||||||
|
|
||||||
|
|||||||
+22
-1
@@ -27,7 +27,7 @@ bash ./workspace/lerobot_example/run.sh
|
|||||||
```
|
```
|
||||||
|
|
||||||
## Enable FAST tokenizer
|
## Enable FAST tokenizer
|
||||||
To fine-tune using the FAST tokenizer, please download the repository and update the `action_tokenizer_path`. Make sure to set `use_fast_tokenizer` to `true`:
|
To fine-tune using the FAST tokenizer, please download the repository and update the `action_tokenizer_path`. Make sure to set `use_fast_tokenizer` to `true` and q01 and q99 to normalize the dataset, refer to `wall-x/scripts/compute_norm_stats.py`:
|
||||||
```bash
|
```bash
|
||||||
git clone https://huggingface.co/physical-intelligence/fast
|
git clone https://huggingface.co/physical-intelligence/fast
|
||||||
```
|
```
|
||||||
@@ -38,7 +38,26 @@ pretrained_wallx_path: "/path/to/wallx_model/" # Path to pretrained wallx m
|
|||||||
save_path: "/path/to/workspace/" # Path to save training outputs
|
save_path: "/path/to/workspace/" # Path to save training outputs
|
||||||
use_fast_tokenizer: False # True: train FAST, False: train Flow
|
use_fast_tokenizer: False # True: train FAST, False: train Flow
|
||||||
action_tokenizer_path: "/path/to/fast/" # Must set if use_fast_tokenizer is True
|
action_tokenizer_path: "/path/to/fast/" # Must set if use_fast_tokenizer is True
|
||||||
|
norm_stats_path: "/path/to/stats/" # Must set for normalize dataset
|
||||||
```
|
```
|
||||||
|
## Customize your robot configuration
|
||||||
|
Ensure that the sum of the configuration dimensions corresponds to the values specified in norm_stats.json, and that each key is unique. The maximum dimensionality is set to 20, consistent with our robot configuration.
|
||||||
|
```yaml
|
||||||
|
customized_dof_config:
|
||||||
|
"action_eef": 6
|
||||||
|
"action_gripper": 1
|
||||||
|
|
||||||
|
customized_agent_pos_config:
|
||||||
|
"state_eef_with_gripper": 7
|
||||||
|
```
|
||||||
|
|
||||||
|
## Compute stats
|
||||||
|
```bash
|
||||||
|
python wall-x/scripts/compute_norm_stats.py
|
||||||
|
```
|
||||||
|
|
||||||
|
## Configuration Explain
|
||||||
|
- `agent_pos_config` corresponds to `obs_action_keys` and subsequently to state, while `dof_config` corresponds to `predict_action_keys` and subsequently to action. Note that the state and action may not necessarily share the same set of DoF.
|
||||||
|
|
||||||
## Training Parameters (Commonly Modified)
|
## Training Parameters (Commonly Modified)
|
||||||
|
|
||||||
@@ -96,6 +115,8 @@ Keep `agent_pos_config` consistent with `dof_config`.
|
|||||||
```bash
|
```bash
|
||||||
# refer to accelerate/commands/merge.py
|
# refer to accelerate/commands/merge.py
|
||||||
accelerate merge-weights /path/to/sharded_tensors /path/to/model.safetensors
|
accelerate merge-weights /path/to/sharded_tensors /path/to/model.safetensors
|
||||||
|
# copy the saved processor files
|
||||||
|
cp /path/to/saved_processor_dir/* /path/to/model.safetensors
|
||||||
```
|
```
|
||||||
|
|
||||||
## Memory Usage
|
## Memory Usage
|
||||||
|
|||||||
@@ -20,7 +20,7 @@ profile_active_iters: 2
|
|||||||
# Training hyperparameters
|
# Training hyperparameters
|
||||||
num_warmup_steps: 100
|
num_warmup_steps: 100
|
||||||
num_training_steps: 64000000
|
num_training_steps: 64000000
|
||||||
learning_rate: 0.00009
|
learning_rate: 0.00005
|
||||||
min_lr: 0.00005
|
min_lr: 0.00005
|
||||||
num_epoch: 100
|
num_epoch: 100
|
||||||
gradient_accumulation_steps: 32
|
gradient_accumulation_steps: 32
|
||||||
@@ -61,6 +61,41 @@ agent_pos_config:
|
|||||||
# ckpt: "/path/to/resume_model/"
|
# ckpt: "/path/to/resume_model/"
|
||||||
# load_ckpt_only: true
|
# load_ckpt_only: true
|
||||||
|
|
||||||
|
norm_stats_path: "/path/to/norm_stats.json"
|
||||||
|
|
||||||
|
enable_customized_robot_config: true
|
||||||
|
customized_robot_config:
|
||||||
|
name: "lerobot/aloha_mobile_cabinet"
|
||||||
|
customized_dof_config:
|
||||||
|
"action_left_shoulder" : 1
|
||||||
|
"action_left_elbow" : 1
|
||||||
|
"action_left_forearm_roll" : 1
|
||||||
|
"action_left_wrist_angle" : 1
|
||||||
|
"action_left_wrist_rotate" : 1
|
||||||
|
"action_left_gripper" : 1
|
||||||
|
"action_right_waist" : 1
|
||||||
|
"action_right_shoulder" : 1
|
||||||
|
"action_right_elbow" : 1
|
||||||
|
"action_right_forearm_roll" : 1
|
||||||
|
"action_right_wrist_angle" : 1
|
||||||
|
"action_right_wrist_rotate" : 1
|
||||||
|
"action_right_gripper" : 1
|
||||||
|
|
||||||
|
customized_agent_pos_config:
|
||||||
|
"state_left_shoulder" : 1
|
||||||
|
"state_left_elbow" : 1
|
||||||
|
"state_left_forearm_roll" : 1
|
||||||
|
"state_left_wrist_angle" : 1
|
||||||
|
"state_left_wrist_rotate" : 1
|
||||||
|
"state_left_gripper" : 1
|
||||||
|
"state_right_waist" : 1
|
||||||
|
"state_right_shoulder" : 1
|
||||||
|
"state_right_elbow" : 1
|
||||||
|
"state_right_forearm_roll" : 1
|
||||||
|
"state_right_wrist_angle" : 1
|
||||||
|
"state_right_wrist_rotate" : 1
|
||||||
|
"state_right_gripper" : 1
|
||||||
|
|
||||||
# Data configuration
|
# Data configuration
|
||||||
data:
|
data:
|
||||||
use_lerobot: true
|
use_lerobot: true
|
||||||
|
|||||||
@@ -62,6 +62,41 @@ agent_pos_config:
|
|||||||
# ckpt: "/path/to/resume_model/"
|
# ckpt: "/path/to/resume_model/"
|
||||||
# load_ckpt_only: true
|
# load_ckpt_only: true
|
||||||
|
|
||||||
|
norm_stats_path: "/path/to/norm_stats.json"
|
||||||
|
|
||||||
|
enable_customized_robot_config: true
|
||||||
|
customized_robot_config:
|
||||||
|
name: "physical-intelligence/libero"
|
||||||
|
customized_dof_config:
|
||||||
|
"action_left_shoulder" : 1
|
||||||
|
"action_left_elbow" : 1
|
||||||
|
"action_left_forearm_roll" : 1
|
||||||
|
"action_left_wrist_angle" : 1
|
||||||
|
"action_left_wrist_rotate" : 1
|
||||||
|
"action_left_gripper" : 1
|
||||||
|
"action_right_waist" : 1
|
||||||
|
"action_right_shoulder" : 1
|
||||||
|
"action_right_elbow" : 1
|
||||||
|
"action_right_forearm_roll" : 1
|
||||||
|
"action_right_wrist_angle" : 1
|
||||||
|
"action_right_wrist_rotate" : 1
|
||||||
|
"action_right_gripper" : 1
|
||||||
|
|
||||||
|
customized_agent_pos_config:
|
||||||
|
"state_left_shoulder" : 1
|
||||||
|
"state_left_elbow" : 1
|
||||||
|
"state_left_forearm_roll" : 1
|
||||||
|
"state_left_wrist_angle" : 1
|
||||||
|
"state_left_wrist_rotate" : 1
|
||||||
|
"state_left_gripper" : 1
|
||||||
|
"state_right_waist" : 1
|
||||||
|
"state_right_shoulder" : 1
|
||||||
|
"state_right_elbow" : 1
|
||||||
|
"state_right_forearm_roll" : 1
|
||||||
|
"state_right_wrist_angle" : 1
|
||||||
|
"state_right_wrist_rotate" : 1
|
||||||
|
"state_right_gripper" : 1
|
||||||
|
|
||||||
# Data configuration
|
# Data configuration
|
||||||
data:
|
data:
|
||||||
use_lerobot: true
|
use_lerobot: true
|
||||||
|
|||||||
@@ -0,0 +1,124 @@
|
|||||||
|
# Training Configuration for Wall-X Robotic Multi-Modal Learning
|
||||||
|
# This configuration supports multi-modal learning with vision, language, and action data
|
||||||
|
|
||||||
|
# Model and paths configuration
|
||||||
|
log_name: "opensource_training"
|
||||||
|
log_project: "libero"
|
||||||
|
model_type: qwen2_5
|
||||||
|
use_fast_tokenizer: true
|
||||||
|
pretrained_wallx_path: "/path/to/qwen/"
|
||||||
|
action_tokenizer_path: "/path/to/fast/"
|
||||||
|
qwen_vl_act_config_path: "/path/to/qwen25_config.json"
|
||||||
|
|
||||||
|
save_path: "/path/to/save"
|
||||||
|
# Torch Profile
|
||||||
|
profile: False
|
||||||
|
profile_save_path: /path/to/profile/
|
||||||
|
profile_wait_iters: 10
|
||||||
|
profile_warmup_iters: 5
|
||||||
|
profile_active_iters: 2
|
||||||
|
|
||||||
|
# Training hyperparameters
|
||||||
|
num_warmup_steps: 100
|
||||||
|
num_training_steps: 64000000
|
||||||
|
learning_rate: 0.00005
|
||||||
|
min_lr: 0.00005
|
||||||
|
num_epoch: 100
|
||||||
|
gradient_accumulation_steps: 1
|
||||||
|
batch_size_per_gpu: 8
|
||||||
|
padding_side: left
|
||||||
|
epoch_save_interval: 1
|
||||||
|
|
||||||
|
# Robot configuration - Define degrees of freedom for each component
|
||||||
|
dof_config:
|
||||||
|
follow_left_ee_cartesian_pos: 3 # Left end-effector Cartesian position
|
||||||
|
follow_left_ee_rotation: 3 # Left end-effector rotation
|
||||||
|
follow_left_gripper: 1 # Left gripper control
|
||||||
|
follow_right_ee_cartesian_pos: 3 # Right end-effector Cartesian position
|
||||||
|
follow_right_ee_rotation: 3 # Right end-effector rotation
|
||||||
|
follow_right_gripper: 1 # Right gripper control
|
||||||
|
head_actions: 2 # Head/camera movement
|
||||||
|
height: 1 # Mobile base height control
|
||||||
|
car_pose: 3 # Mobile base pose (x, y, theta)
|
||||||
|
|
||||||
|
# Agent proprioception configuration (typically matches DOF config)
|
||||||
|
agent_pos_config:
|
||||||
|
follow_left_ee_cartesian_pos: 3
|
||||||
|
follow_left_ee_rotation: 3
|
||||||
|
follow_left_gripper: 1
|
||||||
|
follow_right_ee_cartesian_pos: 3
|
||||||
|
follow_right_ee_rotation: 3
|
||||||
|
follow_right_gripper: 1
|
||||||
|
head_actions: 2
|
||||||
|
height: 1
|
||||||
|
car_pose: 3
|
||||||
|
|
||||||
|
norm_stats_path: "wall-x/workspace/lerobot_example/libero/libero_norm_stats.json"
|
||||||
|
|
||||||
|
enable_customized_robot_config: true
|
||||||
|
customized_robot_config:
|
||||||
|
name: "physical-intelligence/libero"
|
||||||
|
customized_dof_config:
|
||||||
|
"panda_action_eef_with_gripper": 7
|
||||||
|
|
||||||
|
customized_agent_pos_config:
|
||||||
|
"panda_state_eef_with_gripper": 8
|
||||||
|
|
||||||
|
# Checkpoint resuming configuration
|
||||||
|
# resume:
|
||||||
|
# ckpt: "/path/to/ckpt"
|
||||||
|
# load_ckpt_only: false
|
||||||
|
|
||||||
|
# Data configuration
|
||||||
|
data:
|
||||||
|
use_lerobot: true
|
||||||
|
|
||||||
|
# LeRobot dataset configuration
|
||||||
|
lerobot_config:
|
||||||
|
repo_id: "physical-intelligence/libero"
|
||||||
|
root: null
|
||||||
|
episodes: null
|
||||||
|
image_transforms: null
|
||||||
|
delta_timestamps: null
|
||||||
|
tolerance_s: 1e-4
|
||||||
|
revision: null
|
||||||
|
force_cache_sync: false
|
||||||
|
download_videos: true
|
||||||
|
video_backend: null
|
||||||
|
|
||||||
|
action_horizon: 32
|
||||||
|
train_test_split: 0.95
|
||||||
|
|
||||||
|
# Action keys for observation and prediction
|
||||||
|
obs_action_keys:
|
||||||
|
- follow_left_ee_cartesian_pos
|
||||||
|
- follow_left_ee_rotation
|
||||||
|
- follow_left_gripper
|
||||||
|
- follow_right_ee_cartesian_pos
|
||||||
|
- follow_right_ee_rotation
|
||||||
|
- follow_right_gripper
|
||||||
|
- head_actions
|
||||||
|
- height
|
||||||
|
- car_pose
|
||||||
|
|
||||||
|
predict_action_keys:
|
||||||
|
- follow_left_ee_cartesian_pos
|
||||||
|
- follow_left_ee_rotation
|
||||||
|
- follow_left_gripper
|
||||||
|
- follow_right_ee_cartesian_pos
|
||||||
|
- follow_right_ee_rotation
|
||||||
|
- follow_right_gripper
|
||||||
|
- head_actions
|
||||||
|
- height
|
||||||
|
- car_pose
|
||||||
|
|
||||||
|
# Image resolution configuration for different camera views
|
||||||
|
resolution:
|
||||||
|
face_view: 256
|
||||||
|
left_wrist_view: 256
|
||||||
|
right_wrist_view: 256
|
||||||
|
move1_view: 256
|
||||||
|
move2_view: 256
|
||||||
|
top_view: 256
|
||||||
|
wall_view: 256
|
||||||
|
multi_modal: 256
|
||||||
Reference in New Issue
Block a user