Files
roboimi/tests/test_simple_robot_dataset_image_loading.py
T

104 lines
3.7 KiB
Python

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)