feat: 校验ACT episode数据质量
This commit is contained in:
@@ -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
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user