import sys import tempfile import types import unittest from pathlib import Path from unittest import mock import h5py import numpy as np from roboimi.vla.data.simpe_robot_dataset import SimpleRobotDataset class SimpleRobotDatasetImageLoadingTest(unittest.TestCase): def _write_episode(self, dataset_dir: Path, episode_idx: int = 0, action_offset: float = 0.0) -> None: episode_path = dataset_dir / f"episode_{episode_idx}.hdf5" with h5py.File(episode_path, "w") as root: root.create_dataset( "action", data=(np.arange(8, dtype=np.float32).reshape(4, 2) + action_offset), ) root.create_dataset( "observations/qpos", data=np.arange(16, dtype=np.float32).reshape(4, 4), ) root.create_dataset("task", data=np.array([b"sim_transfer"])) root.create_dataset( "observations/images/front", data=np.arange(4 * 8 * 8 * 3, dtype=np.uint8).reshape(4, 8, 8, 3), ) def test_getitem_only_resizes_observation_horizon_images(self): with tempfile.TemporaryDirectory() as tmpdir: dataset_dir = Path(tmpdir) self._write_episode(dataset_dir) dataset = SimpleRobotDataset( dataset_dir, obs_horizon=2, pred_horizon=3, camera_names=["front"], ) resize_calls = [] def fake_resize(image, size, interpolation=None): resize_calls.append( { "shape": tuple(image.shape), "size": size, "interpolation": interpolation, } ) return image fake_cv2 = types.SimpleNamespace(INTER_LINEAR=1, resize=fake_resize) with mock.patch.dict(sys.modules, {"cv2": fake_cv2}): sample = dataset[1] self.assertEqual(len(resize_calls), 2) self.assertEqual(tuple(sample["observation.front"].shape), (2, 3, 8, 8)) def test_getitem_skips_resize_when_image_resize_shape_is_none(self): with tempfile.TemporaryDirectory() as tmpdir: dataset_dir = Path(tmpdir) self._write_episode(dataset_dir) dataset = SimpleRobotDataset( dataset_dir, obs_horizon=2, pred_horizon=3, camera_names=["front"], image_resize_shape=None, ) fake_cv2 = types.SimpleNamespace( INTER_LINEAR=1, resize=mock.Mock(side_effect=AssertionError("resize should be skipped when image_resize_shape=None")), ) with mock.patch.dict(sys.modules, {"cv2": fake_cv2}): sample = dataset[1] fake_cv2.resize.assert_not_called() self.assertEqual(tuple(sample["observation.front"].shape), (2, 3, 8, 8)) def test_dataset_can_filter_by_episode_indices_and_expose_available_episode_indices(self): with tempfile.TemporaryDirectory() as tmpdir: dataset_dir = Path(tmpdir) self._write_episode(dataset_dir, episode_idx=3, action_offset=0.0) self._write_episode(dataset_dir, episode_idx=7, action_offset=100.0) dataset = SimpleRobotDataset( dataset_dir, obs_horizon=2, pred_horizon=3, camera_names=["front"], episode_indices=[7], ) sample = dataset[0] self.assertEqual(dataset.available_episode_indices, [7]) self.assertEqual(len(dataset), 4) self.assertEqual(sample["action"][0, 0].item(), 100.0)