DPPO自定义数据集教程:打造专属机器人控制训练数据的完整流程
【免费下载链接】dppoOfficial implementation of Diffusion Policy Policy Optimization, arxiv 2024项目地址: https://gitcode.com/gh_mirrors/dpp/dppo
DPPO(Diffusion Policy Policy Optimization)作为先进的机器人控制算法,其性能高度依赖高质量的训练数据。本教程将带你完成从数据采集到配置使用的完整流程,轻松创建专属于你的机器人控制数据集,让DPPO模型发挥最佳效果。
一、数据集基础认知:DPPO数据格式解析
DPPO采用结构化的NPZ格式存储训练数据,包含以下核心字段:
- states:环境观测数据,形状为
(总步数, 观测维度) - actions:机器人动作数据,形状为
(总步数, 动作维度) - traj_lengths:轨迹长度数组,标记每个 episode 的步数
- rewards(可选):奖励信号,用于强化学习微调
- terminals(可选): episode 结束标记
核心数据集加载逻辑位于 agent/dataset/sequence.py,其中StitchedSequenceDataset类负责处理轨迹拼接和采样逻辑。代码片段展示了数据加载的关键步骤:
# 从NPZ文件加载数据集 if dataset_path.endswith(".npz"): dataset = np.load(dataset_path, allow_pickle=False) elif dataset_path.endswith(".pkl"): with open(dataset_path, "rb") as f: dataset = pickle.load(f) # 提取核心数据 self.states = torch.from_numpy(dataset["states"][:total_num_steps]).float().to(device) self.actions = torch.from_numpy(dataset["actions"][:total_num_steps]).float().to(device) self.traj_lengths = dataset["traj_lengths"][:max_n_episodes]二、数据采集指南:获取原始机器人交互数据
2.1 传感器数据采集
根据机器人类型选择合适的传感器配置:
- 机械臂系统:需采集末端执行器位姿、关节角度、 gripper 状态
- 移动机器人:需采集里程计数据、IMU读数、激光雷达点云
推荐采样频率:20-100Hz,确保动作序列的连续性。
2.2 数据记录格式
原始数据建议保存为HDF5或ROS bag格式,包含:
- 时间戳(同步多传感器数据)
- 原始观测(未归一化)
- 原始动作(关节空间或任务空间)
- 环境元数据(物体位置、光照条件等)
三、数据预处理:从原始数据到DPPO可用格式
3.1 数据格式转换工具
DPPO提供多种数据集处理脚本,位于 script/dataset/ 目录:
- RoboMimic数据集:process_robomimic_dataset.py
- D3IL数据集:process_d3il_dataset.py
- D4RL数据集:get_d4rl_dataset.py
以RoboMimic处理为例,基本命令:
python script/dataset/process_robomimic_dataset.py \ --load_path=../raw_data/lift_low_dim_v141.hdf5 \ --save_dir=data/robomimic/lift \ --normalize3.2 关键预处理步骤
数据清洗
- 移除异常值(如关节限位外的动作)
- 修复时间戳不连续的轨迹
- 过滤过短轨迹(建议最小长度 > 50步)
特征提取
- 低维观测:关节角度、末端执行器位姿、物体状态
- 图像数据:多视角相机图像(需确保尺寸为8的倍数)
归一化推荐将观测和动作归一化到[-1, 1]范围:
# 归一化公式(来自process_robomimic_dataset.py) obs = 2 * (raw_obs - obs_min) / (obs_max - obs_min + 1e-6) - 1 actions = 2 * (raw_actions - action_min) / (action_max - action_min + 1e-6) - 1数据集划分按轨迹划分训练集和验证集(而非随机打乱):
# 训练集/验证集划分示例 num_train = int(num_traj * (1 - val_split)) train_indices = random.sample(range(num_traj), k=num_train)
四、自定义数据集实现:创建专属数据加载器
4.1 自定义数据集类
创建新的数据集类,继承StitchedSequenceDataset基类:
from agent.dataset.sequence import StitchedSequenceDataset class CustomRobotDataset(StitchedSequenceDataset): def __init__(self, dataset_path, custom_param, **kwargs): super().__init__(dataset_path, **kwargs) self.custom_param = custom_param # 添加自定义参数 def make_indices(self, traj_lengths, horizon_steps): # 重写索引生成逻辑(如特殊轨迹处理) indices = [] # ... 自定义实现 ... return indices4.2 数据加载配置
在配置文件中指定自定义数据集:
# 示例配置:cfg/custom/finetune/custom_env/ft_ppo_diffusion_mlp.yaml train_dataset: _target_: agent.dataset.custom.CustomRobotDataset dataset_path: ${oc.env:DPPO_DATA_DIR}/custom_env/train.npz horizon_steps: 64 cond_steps: 1 max_n_episodes: 500 use_img: false五、数据集使用与调试:确保数据正确加载
5.1 数据集加载验证
使用以下代码验证数据加载是否正确:
# 简单数据加载测试 from agent.dataset.sequence import StitchedSequenceDataset dataset = StitchedSequenceDataset( dataset_path="data/custom/train.npz", horizon_steps=64, device="cpu" ) print(f"数据集大小: {len(dataset)} samples") print(f"状态维度: {dataset.states.shape[1]}") print(f"动作维度: {dataset.actions.shape[1]}")5.2 常见问题排查
数据维度不匹配
- 检查观测/动作维度是否与模型配置一致
- 确保所有轨迹的状态/动作维度相同
内存溢出
- 减少
max_n_episodes参数 - 使用更低精度数据类型(如float32)
- 减少
图像数据问题
- 确保图像尺寸为8的倍数(如96x96, 128x128)
- 检查通道顺序是否为 (C, H, W)
六、高级优化:提升数据集质量的技巧
6.1 数据增强策略
- 状态扰动:添加高斯噪声(如±0.01)增强鲁棒性
- 动作平滑:使用滑动平均减少高频噪声
- 轨迹裁剪:保留任务关键片段,去除冗余部分
6.2 多源数据融合
通过 agent/dataset/sequence.py 中的StitchedSequenceDataset实现多任务数据融合:
# 多数据集拼接配置示例 train_dataset: _target_: agent.dataset.sequence.StitchedSequenceDataset dataset_path: ${oc.env:DPPO_DATA_DIR}/merged/train.npz max_n_episodes: 1000 # 合并多个任务的轨迹6.3 数据集质量评估
关键指标:
- 轨迹多样性:动作空间覆盖率 > 80%
- 数据一致性:状态转移平滑度(速度变化率)
- 任务相关性:与目标任务的动作分布相似度
七、完整工作流示例:从采集到训练
数据采集
# 假设使用ROS采集数据 rosbag record -O raw_data.bag /joint_states /end_effector/pose数据转换
python script/dataset/process_custom_dataset.py \ --load_path=raw_data.bag \ --save_dir=data/custom_robot \ --normalize配置训练
python script/run.py \ agent=pretrain/train_diffusion_agent \ +train_dataset=dataset/custom_robot \ train.max_epochs=100
通过以上步骤,你已成功创建并使用自定义数据集训练DPPO模型。记住,高质量的数据是机器人控制算法成功的关键,花时间优化数据采集和预处理流程将显著提升最终性能。
【免费下载链接】dppoOfficial implementation of Diffusion Policy Policy Optimization, arxiv 2024项目地址: https://gitcode.com/gh_mirrors/dpp/dppo
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考