mirror of
https://github.com/huggingface/lerobot.git
synced 2026-07-24 02:06:15 +00:00
feat(datasets): warn when skipping stats for zero-width features
Per review, log a warning when compute_episode_stats skips a feature with a zero-width shape, so users know stats were intentionally not computed for it.
This commit is contained in:
@@ -520,6 +520,9 @@ def compute_episode_stats(
|
|||||||
continue
|
continue
|
||||||
|
|
||||||
if any(d == 0 for d in features[key].get("shape", ())):
|
if any(d == 0 for d in features[key].get("shape", ())):
|
||||||
|
logging.warning(
|
||||||
|
f"Skipping stats for feature '{key}' with a zero-width shape {features[key]['shape']}."
|
||||||
|
)
|
||||||
continue
|
continue
|
||||||
|
|
||||||
if features[key]["dtype"] in ["image", "video"]:
|
if features[key]["dtype"] in ["image", "video"]:
|
||||||
|
|||||||
@@ -13,6 +13,7 @@
|
|||||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||||
# See the License for the specific language governing permissions and
|
# See the License for the specific language governing permissions and
|
||||||
# limitations under the License.
|
# limitations under the License.
|
||||||
|
import logging
|
||||||
from unittest.mock import patch
|
from unittest.mock import patch
|
||||||
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
@@ -687,8 +688,8 @@ def test_compute_episode_stats_string_features_skipped():
|
|||||||
assert "q01" in stats["action"]
|
assert "q01" in stats["action"]
|
||||||
|
|
||||||
|
|
||||||
def test_compute_episode_stats_zero_width_features_skipped():
|
def test_compute_episode_stats_zero_width_features_skipped(caplog):
|
||||||
"""Test that features with a zero-width dim (e.g. shape=(0,)) are skipped."""
|
"""Test that features with a zero-width dim (e.g. shape=(0,)) are skipped with a warning."""
|
||||||
episode_data = {
|
episode_data = {
|
||||||
"empty": np.zeros((100, 0), dtype=np.float32), # Zero-width feature
|
"empty": np.zeros((100, 0), dtype=np.float32), # Zero-width feature
|
||||||
"action": np.random.normal(0, 1, (100, 5)),
|
"action": np.random.normal(0, 1, (100, 5)),
|
||||||
@@ -698,13 +699,12 @@ def test_compute_episode_stats_zero_width_features_skipped():
|
|||||||
"action": {"dtype": "float32", "shape": (5,)},
|
"action": {"dtype": "float32", "shape": (5,)},
|
||||||
}
|
}
|
||||||
|
|
||||||
stats = compute_episode_stats(
|
with caplog.at_level(logging.WARNING):
|
||||||
episode_data,
|
stats = compute_episode_stats(episode_data, features)
|
||||||
features,
|
|
||||||
)
|
|
||||||
|
|
||||||
# Zero-width features should be skipped
|
# Zero-width features should be skipped with a warning, others computed as usual
|
||||||
assert "empty" not in stats
|
assert "empty" not in stats
|
||||||
|
assert "empty" in caplog.text
|
||||||
assert "action" in stats
|
assert "action" in stats
|
||||||
assert "q01" in stats["action"]
|
assert "q01" in stats["action"]
|
||||||
assert stats["action"]["mean"].shape == (5,)
|
assert stats["action"]["mean"].shape == (5,)
|
||||||
|
|||||||
Reference in New Issue
Block a user