mirror of
https://github.com/huggingface/lerobot.git
synced 2026-07-30 13:09:40 +00:00
feat(dataset filtering): adding support for VLM based dataset filtering following lerobot annotation pipeline style
This commit is contained in:
@@ -0,0 +1,52 @@
|
||||
#!/usr/bin/env python
|
||||
|
||||
# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
# The script imports ``lerobot.datasets`` (via the annotation frame provider),
|
||||
# which only ships under the ``dataset`` extra.
|
||||
pytest.importorskip("datasets", reason="datasets is required (install lerobot[dataset])")
|
||||
|
||||
from lerobot.scripts.lerobot_curate_cameras import _to_uint8_frame, _uniform_indices # noqa: E402
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"n,k,expected",
|
||||
[
|
||||
(0, 4, []),
|
||||
(5, 0, []),
|
||||
(3, 5, [0, 1, 2]), # k >= n -> all frames
|
||||
(1, 4, [0]),
|
||||
(10, 1, [0]),
|
||||
(10, 4, [0, 3, 6, 9]), # evenly spaced, endpoints included
|
||||
],
|
||||
)
|
||||
def test_uniform_indices(n, k, expected):
|
||||
assert _uniform_indices(n, k) == expected
|
||||
|
||||
|
||||
def test_to_uint8_frame_scales_floats():
|
||||
frame = torch.ones(3, 4, 4, dtype=torch.float32) # [0,1] float
|
||||
out = _to_uint8_frame(frame)
|
||||
assert out.dtype == torch.uint8
|
||||
assert int(out.max()) == 255
|
||||
|
||||
|
||||
def test_to_uint8_frame_passthrough_uint8():
|
||||
frame = torch.zeros(3, 4, 4, dtype=torch.uint8)
|
||||
out = _to_uint8_frame(frame)
|
||||
assert out is frame # uint8 passes through untouched
|
||||
Reference in New Issue
Block a user