mirror of
https://github.com/huggingface/lerobot.git
synced 2026-07-27 19:56:09 +00:00
chore(dataset): add check dataset shape
This commit is contained in:
@@ -687,21 +687,16 @@ def test_compute_episode_stats_string_features_skipped():
|
||||
assert "q01" in stats["action"]
|
||||
|
||||
|
||||
def test_compute_episode_stats_zero_width_feature_skipped():
|
||||
"""Features with a zero-width dimension carry no values and must be skipped, not crash.
|
||||
|
||||
Regression test for https://github.com/huggingface/lerobot/issues/3654:
|
||||
a feature declared with shape=(0,) previously reached RunningQuantileStats.update,
|
||||
where ``batch.reshape(-1, batch.shape[-1])`` raised
|
||||
"ValueError: cannot reshape array of size 0 into shape (0)".
|
||||
"""
|
||||
@pytest.mark.parametrize("shape", [(0,), (0, 2), (2, 0), (1, 0, 2)])
|
||||
def test_compute_episode_stats_zero_width_feature_skipped(shape):
|
||||
"""Features with any zero-width dimension carry no values and are skipped."""
|
||||
episode_data = {
|
||||
"action": np.random.normal(0, 1, (100, 5)).astype(np.float32),
|
||||
"target": np.zeros((100, 0), dtype=np.float32), # zero-width feature
|
||||
"target": np.zeros((100, *shape), dtype=np.float32),
|
||||
}
|
||||
features = {
|
||||
"action": {"dtype": "float32", "shape": (5,)},
|
||||
"target": {"dtype": "float32", "shape": (0,)},
|
||||
"target": {"dtype": "float32", "shape": shape},
|
||||
}
|
||||
|
||||
stats = compute_episode_stats(episode_data, features)
|
||||
|
||||
Reference in New Issue
Block a user