diff --git a/xr_rm_teleop/test/test_act_episode_recorder.py b/xr_rm_teleop/test/test_act_episode_recorder.py index 4913871..18b135f 100644 --- a/xr_rm_teleop/test/test_act_episode_recorder.py +++ b/xr_rm_teleop/test/test_act_episode_recorder.py @@ -14,6 +14,7 @@ from xr_rm_teleop.act_episode_recorder import ( EpisodeMetadata, EpisodeStore, QualityError, + QualityLimits, RecordingSession, RecordingState, TaskDirectoryLock, @@ -23,6 +24,7 @@ from xr_rm_teleop.act_episode_recorder import ( recover_partial_files, select_camera_pair, select_frame, + validate_episode, ) @@ -477,3 +479,122 @@ def test_task_directory_lock_rejects_second_recorder(tmp_path): second.acquire() finally: 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 diff --git a/xr_rm_teleop/xr_rm_teleop/act_episode_recorder.py b/xr_rm_teleop/xr_rm_teleop/act_episode_recorder.py index b34eb4c..fc353c9 100644 --- a/xr_rm_teleop/xr_rm_teleop/act_episode_recorder.py +++ b/xr_rm_teleop/xr_rm_teleop/act_episode_recorder.py @@ -102,6 +102,26 @@ class EpisodeMetadata: 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: def __init__(self, path: Path, root: Any) -> None: self.path = path @@ -305,6 +325,196 @@ def recover_partial_files( 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): pass