feat: 校验ACT episode数据质量

This commit is contained in:
2026-08-10 18:14:10 +08:00
parent 40be5560ee
commit c1ea1a2816
2 changed files with 331 additions and 0 deletions
@@ -14,6 +14,7 @@ from xr_rm_teleop.act_episode_recorder import (
EpisodeMetadata, EpisodeMetadata,
EpisodeStore, EpisodeStore,
QualityError, QualityError,
QualityLimits,
RecordingSession, RecordingSession,
RecordingState, RecordingState,
TaskDirectoryLock, TaskDirectoryLock,
@@ -23,6 +24,7 @@ from xr_rm_teleop.act_episode_recorder import (
recover_partial_files, recover_partial_files,
select_camera_pair, select_camera_pair,
select_frame, select_frame,
validate_episode,
) )
@@ -477,3 +479,122 @@ def test_task_directory_lock_rejects_second_recorder(tmp_path):
second.acquire() second.acquire()
finally: finally:
first.release() first.release()
def _quality_limits():
return QualityLimits(min_samples=3, max_samples=100)
def _valid_episode(tmp_path):
path = tmp_path / "episode_0.partial.hdf5"
store = EpisodeStore.create(path, _metadata())
for index, seq in enumerate((100, 103, 106)):
sample = _episode_sample(seq)
sample["debug/cameras/cam_high_frame_number"] = 200 + index
sample["debug/cameras/cam_wrist_frame_number"] = 300 + index
store.append(sample)
store.close()
with h5py.File(path, "r+") as root:
root.attrs["camera_high_fps"] = 30.0
root.attrs["camera_right_wrist_fps"] = 30.0
root.attrs["camera_high_drop_ratio"] = 0.0
root.attrs["camera_right_wrist_drop_ratio"] = 0.0
return path
@requires_h5py
def test_validate_episode_accepts_valid_file(tmp_path):
report = validate_episode(_valid_episode(tmp_path), _quality_limits())
assert report.accepted
assert report.reason is None
assert report.metrics["control_hz"] == pytest.approx(30.0, rel=1e-5)
def _mutate_episode(path, mutation):
with h5py.File(path, "r+") as root:
if mutation == "short_episode":
root.visititems(
lambda _name, item: item.resize(2, axis=0)
if isinstance(item, h5py.Dataset)
else None
)
elif mutation == "control_seq_gap":
root["debug/control/control_seq"][1] = 104
elif mutation == "nonfinite_qpos":
root["observations/qpos"][1, 0] = np.nan
elif mutation == "joint_limit":
root["action"][1, 0] = 4.0
elif mutation == "invalid_gripper":
root["action"][1, 7] = 0.5
elif mutation == "feedback_age":
control_ns = root["debug/timestamps/control_monotonic_ns"][1]
root["debug/timestamps/feedback_monotonic_ns"][1] = (
control_ns - 60_000_000
)
elif mutation == "action_invalid":
root["debug/control/action_valid"][1] = 0
elif mutation == "control_fault":
root["debug/control/control_fault"][1] = 1
elif mutation == "camera_fps":
root.attrs["camera_high_fps"] = 20.0
elif mutation == "camera_drop":
root.attrs["camera_right_wrist_drop_ratio"] = 0.02
elif mutation == "camera_age":
root["debug/timestamps/cam_high_age_ms"][1] = 60.0
elif mutation == "camera_skew":
root["debug/timestamps/inter_camera_skew_ms"][1] = 60.0
elif mutation == "final_gripper_closed":
root["observations/qpos"][-1, 7] = 0.0
else:
raise AssertionError(f"unknown mutation: {mutation}")
@requires_h5py
@pytest.mark.parametrize(
("mutation", "reason"),
[
("short_episode", "too_few_samples"),
("control_seq_gap", "control_sequence_gap"),
("nonfinite_qpos", "nonfinite_qpos"),
("joint_limit", "joint_limit_violation"),
("invalid_gripper", "invalid_gripper_state"),
("feedback_age", "feedback_too_old"),
("action_invalid", "invalid_action"),
("control_fault", "control_fault"),
("camera_fps", "camera_fps"),
("camera_drop", "camera_drop_ratio"),
("camera_age", "camera_frame_too_old"),
("camera_skew", "camera_skew"),
("final_gripper_closed", "final_gripper_not_open"),
],
)
def test_validate_episode_reports_stable_reason(
tmp_path,
mutation,
reason,
):
path = _valid_episode(tmp_path)
_mutate_episode(path, mutation)
report = validate_episode(path, _quality_limits())
assert not report.accepted
assert report.reason == reason
@requires_h5py
def test_validate_episode_reports_qp_failures_without_rejecting(tmp_path):
path = _valid_episode(tmp_path)
with h5py.File(path, "r+") as root:
root["debug/qp/attempted"][:] = (1, 1, 1)
root["debug/qp/success"][:] = (0, 0, 1)
root["debug/control/target_clamped"][:] = (1, 0, 1)
report = validate_episode(path, _quality_limits())
assert report.accepted
assert report.metrics["qp_failure_count"] == 2
assert report.metrics["qp_failure_ratio"] == pytest.approx(2.0 / 3.0)
assert report.metrics["qp_longest_failure_streak"] == 2
assert report.metrics["target_clamped_count"] == 2
@@ -102,6 +102,26 @@ class EpisodeMetadata:
raise ValueError("joint lower limits must be below upper limits") raise ValueError("joint lower limits must be below upper limits")
@dataclass(frozen=True)
class QualityLimits:
min_samples: int = 60
max_samples: int = 1800
min_control_hz: float = 27.0
max_control_gap_ms: float = 100.0
min_camera_fps: float = 27.0
max_drop_ratio: float = 0.01
max_feedback_age_ms: float = 50.0
max_camera_age_ms: float = 50.0
max_camera_skew_ms: float = 50.0
@dataclass(frozen=True)
class QualityReport:
accepted: bool
reason: str | None
metrics: dict[str, int | float]
class EpisodeStore: class EpisodeStore:
def __init__(self, path: Path, root: Any) -> None: def __init__(self, path: Path, root: Any) -> None:
self.path = path self.path = path
@@ -305,6 +325,196 @@ def recover_partial_files(
return recovered return recovered
def longest_true_run(values: np.ndarray) -> int:
longest = 0
current = 0
for value in values:
current = current + 1 if value else 0
longest = max(longest, current)
return longest
def _quality_failure(
reason: str,
metrics: dict[str, int | float],
) -> QualityReport:
return QualityReport(False, reason, metrics)
def validate_episode(path: Path, limits: QualityLimits) -> QualityReport:
metrics: dict[str, int | float] = {}
try:
root = _h5py().File(path, "r")
except OSError:
return _quality_failure("invalid_hdf5", metrics)
with root:
try:
datasets = {name: root[name] for name in DATA_LAYOUT}
except KeyError:
return _quality_failure("schema_error", metrics)
lengths = {dataset.shape[0] for dataset in datasets.values()}
if len(lengths) != 1:
return _quality_failure("time_axis_length_mismatch", metrics)
sample_count = lengths.pop()
metrics["sample_count"] = sample_count
if sample_count < limits.min_samples:
return _quality_failure("too_few_samples", metrics)
if sample_count >= limits.max_samples:
return _quality_failure("max_duration", metrics)
for name, (dtype, sample_shape, _chunks) in DATA_LAYOUT.items():
dataset = datasets[name]
if dataset.shape != (sample_count, *sample_shape):
return _quality_failure("schema_error", metrics)
if dataset.dtype != np.dtype(dtype):
return _quality_failure("schema_error", metrics)
qpos = datasets["observations/qpos"][:]
action = datasets["action"][:]
if not np.all(np.isfinite(qpos)):
return _quality_failure("nonfinite_qpos", metrics)
if not np.all(np.isfinite(action)):
return _quality_failure("nonfinite_action", metrics)
lower = np.asarray(root.attrs.get("joint_lower_limits", ()))
upper = np.asarray(root.attrs.get("joint_upper_limits", ()))
if lower.shape != (7,) or upper.shape != (7,):
return _quality_failure("joint_limits_missing", metrics)
if (
np.any(qpos[:, :7] < lower)
or np.any(qpos[:, :7] > upper)
or np.any(action[:, :7] < lower)
or np.any(action[:, :7] > upper)
):
return _quality_failure("joint_limit_violation", metrics)
if not (
np.all(np.isin(qpos[:, 7], (0.0, 1.0)))
and np.all(np.isin(action[:, 7], (0.0, 1.0)))
):
return _quality_failure("invalid_gripper_state", metrics)
attempted = datasets["debug/qp/attempted"][:].astype(bool)
success = datasets["debug/qp/success"][:].astype(bool)
failed = attempted & ~success
failure_count = int(failed.sum())
metrics["qp_failure_count"] = failure_count
metrics["qp_failure_ratio"] = float(
failure_count / max(1, int(attempted.sum()))
)
metrics["qp_longest_failure_streak"] = longest_true_run(failed)
metrics["target_clamped_count"] = int(
datasets["debug/control/target_clamped"][:].sum()
)
control_seq = datasets["debug/control/control_seq"][:].astype(
np.int64
)
if not np.all(np.diff(control_seq) == 3):
return _quality_failure("control_sequence_gap", metrics)
control_ns = datasets[
"debug/timestamps/control_monotonic_ns"
][:].astype(np.int64)
control_gap_ns = np.diff(control_ns)
if np.any(control_gap_ns <= 0):
return _quality_failure("control_timestamp_invalid", metrics)
elapsed_ns = int(control_ns[-1] - control_ns[0])
control_hz = (sample_count - 1) * 1e9 / elapsed_ns
max_control_gap_ms = float(control_gap_ns.max() * 1e-6)
metrics["control_hz"] = float(control_hz)
metrics["max_control_gap_ms"] = max_control_gap_ms
if control_hz < limits.min_control_hz:
return _quality_failure("control_rate", metrics)
if max_control_gap_ms > limits.max_control_gap_ms:
return _quality_failure("control_gap", metrics)
feedback_ns = datasets[
"debug/timestamps/feedback_monotonic_ns"
][:].astype(np.int64)
feedback_age_ms = (control_ns - feedback_ns) * 1e-6
metrics["max_feedback_age_ms"] = float(feedback_age_ms.max())
if np.any(feedback_age_ms < 0) or np.any(
feedback_age_ms > limits.max_feedback_age_ms
):
return _quality_failure("feedback_too_old", metrics)
if not np.all(datasets["debug/control/action_valid"][:] == 1):
return _quality_failure("invalid_action", metrics)
if np.any(datasets["debug/control/control_fault"][:] != 0):
return _quality_failure("control_fault", metrics)
if np.any(datasets["debug/gripper/command_failed"][:] != 0):
return _quality_failure("gripper_command_failed", metrics)
if datasets["debug/gripper/command_pending"][-1] != 0:
return _quality_failure("gripper_command_pending", metrics)
camera_stats = (
("camera_high_fps", limits.min_camera_fps, "camera_fps"),
(
"camera_right_wrist_fps",
limits.min_camera_fps,
"camera_fps",
),
(
"camera_high_drop_ratio",
limits.max_drop_ratio,
"camera_drop_ratio",
),
(
"camera_right_wrist_drop_ratio",
limits.max_drop_ratio,
"camera_drop_ratio",
),
)
for name, threshold, reason in camera_stats:
if name not in root.attrs:
return _quality_failure("camera_stats_missing", metrics)
value = float(root.attrs[name])
metrics[name] = value
if not np.isfinite(value):
return _quality_failure("camera_stats_missing", metrics)
if reason == "camera_fps" and value < threshold:
return _quality_failure(reason, metrics)
if reason == "camera_drop_ratio" and value > threshold:
return _quality_failure(reason, metrics)
high_age_ms = datasets[
"debug/timestamps/cam_high_age_ms"
][:]
wrist_age_ms = datasets[
"debug/timestamps/cam_wrist_age_ms"
][:]
max_camera_age_ms = float(max(high_age_ms.max(), wrist_age_ms.max()))
metrics["max_camera_age_ms"] = max_camera_age_ms
if (
np.any(high_age_ms < 0)
or np.any(wrist_age_ms < 0)
or max_camera_age_ms > limits.max_camera_age_ms
):
return _quality_failure("camera_frame_too_old", metrics)
camera_skew_ms = datasets[
"debug/timestamps/inter_camera_skew_ms"
][:]
metrics["max_camera_skew_ms"] = float(camera_skew_ms.max())
if np.any(camera_skew_ms < 0) or np.any(
camera_skew_ms > limits.max_camera_skew_ms
):
return _quality_failure("camera_skew", metrics)
for camera in ("cam_high", "cam_wrist"):
frame_numbers = datasets[
f"debug/cameras/{camera}_frame_number"
][:].astype(np.int64)
discontinuities = int((np.diff(frame_numbers) != 1).sum())
ratio = discontinuities / max(1, sample_count - 1)
metrics[f"{camera}_sample_discontinuity_ratio"] = float(ratio)
if ratio > limits.max_drop_ratio:
return _quality_failure("camera_sample_drop_ratio", metrics)
if qpos[-1, 7] != 1.0:
return _quality_failure("final_gripper_not_open", metrics)
return QualityReport(True, None, metrics)
class QualityError(RuntimeError): class QualityError(RuntimeError):
pass pass