用 HDF5 数据集训练
平台导出的 HDF5 数据集由平台之外的训练框架读取。接入训练前需完成数据准备、字段读取与结构校验。
适用角色与前提
角色与关注点
| 角色 | 关注点 |
|---|---|
| 算法工程师 | 字段形状、帧对齐、训练框架接入 |
| 训练运维 | 数据集挂载、显存与磁盘占用 |
使用前提
| 项 | 要求 |
|---|---|
| 数据集 | 已完成 HDF5 导出并下载归档,结构见 HDF5 数据集 |
| 文件 | 解压后包含 chunk_*.hdf5 |
| 字段 | 每个 episode 至少含 action 与 observation.state |
| 运行环境 | Python 环境含 h5py、numpy,以及目标框架(PyTorch、TensorFlow、JAX 之一) |
| 图像解码 | 可解码 JPEG(Pillow、OpenCV 或 torchvision) |
| 资源 | 训练机磁盘可容纳解压后的分块文件与 checkpoint |
操作步骤
- 解压归档,确认全部分块文件位于同一目录。
- 打开一个分块文件,读取
data下的 episode 列表与各数据集形状,核对与 HDF5 数据集 的字段表和形状一致。
import h5py
with h5py.File("chunk_001.hdf5", "r") as f:
for episode_name in f["data"]:
episode = f[f"data/{episode_name}"]
print(episode_name, episode.attrs["task"].decode())
for key in episode:
print(" ", key, episode[key].shape, episode[key].dtype)
- 建立
(分块文件路径, episode 名)索引,避免每个 epoch 重新扫描文件。 - 按 episode 划分训练集与验证集,同一分块文件不跨集合,避免同源样本泄漏。
- 读取时对
observation.images.*逐帧 JPEG 解码,并与action、observation.state、observation.gripper按帧下标对齐。 - 将 episode 封装为目标框架的数据集对象,按批输出图像张量与状态、动作向量。
- 读取
task、task_zh、score属性,用于按任务或质量筛选样本。
最小加载器示例:
import io
import h5py
import numpy as np
import torch
from PIL import Image
from torch.utils.data import Dataset
class Hdf5EpisodeDataset(Dataset):
def __init__(self, files, transform=None):
self.index = []
self.transform = transform
for path in files:
with h5py.File(path, "r") as f:
self.index += [(path, name) for name in f["data"]]
def __len__(self):
return len(self.index)
def __getitem__(self, i):
path, name = self.index[i]
with h5py.File(path, "r") as f:
ep = f[f"data/{name}"]
images = [Image.open(io.BytesIO(frame.tobytes()))
for frame in ep["observation.images.camera_01"][:]]
if self.transform:
images = [self.transform(img) for img in images]
return {
"images": torch.stack(images),
"state": torch.as_tensor(np.asarray(ep["observation.state"][:]), dtype=torch.float32),
"action": torch.as_tensor(np.asarray(ep["action"][:]), dtype=torch.float32),
"task": ep.attrs["task"].decode(),
}
结果校验
| 校验项 | 方法 | 通过标准 |
|---|---|---|
| 文件完整 | 列出解压目录内的分块文件 | 命名连续,无缺号 |
| 字段齐全 | 遍历 data 下各 episode | 含 action、observation.state,且含至少一个 observation.images.* |
| 形状一致 | 读取各数据集 shape | 各数据集首维等于该 episode 的帧数 |
| 图像可解码 | 抽样 observation.images.* 元素 | JPEG 解码成功 |
| 帧对齐 | 比较 action 与 observation.state 的行数 | 行数相等 |
异常处置
| 现象 | 可能原因 | 处置 | 责任方 |
|---|---|---|---|
| 文件无法打开 | 下载中断或解压不完整 | 重新下载并核对文件大小 | 使用者 |
缺少 observation.gripper | 导出源无夹爪话题 | 训练时忽略该字段,或在动作维度配置中去除 | 算法工程师 |
| 图像解码失败 | 该帧为原始数组而非 JPEG | 按 uint8 数组自行解码,或剔除该帧 | 算法工程师 |
| 各数据集帧数不一致 | sidecar JSON 子任务区间与消息时间不对齐 | 以 action 的时间戳为准重采样 | 算法工程师 |
| 显存不足 | 单批加载图像过多 | 降低批大小或降低图像分辨率 | 算法工程师 |
| 训练指标不收敛 | 帧率抽样过低或状态与动作错位 | 提高导出 hz,检查话题映射与帧对齐 | 算法工程师 |