from __future__ import annotations import io import json from types import SimpleNamespace import numpy as np import pytest from dovla_cil.data.images import CILImageReader from dovla_cil.generation.pipeline import generate_cil_dataset from dovla_cil.tasks.library import built_in_toy_tasks from dovla_cil.training.trainer import DoVLATrainer, TrainerConfig def test_image_reader_resolves_collection_source_and_caches_archive(tmp_path) -> None: h5py = pytest.importorskip("h5py") image_module = pytest.importorskip("PIL.Image") source = tmp_path / "source" source.mkdir() expected = np.full((6, 8, 3), 117, dtype=np.uint8) buffer = io.BytesIO() image_module.fromarray(expected).save(buffer, format="JPEG", quality=95) encoded = np.frombuffer(buffer.getvalue(), dtype=np.uint8) with h5py.File(source / "observations.h5", "w") as handle: dataset = handle.create_dataset( "initial_rgb_jpeg", shape=(1,), dtype=h5py.vlen_dtype(np.dtype("uint8")), ) dataset[0] = encoded record = SimpleNamespace( record_id="record-0", observation_ref="observations.h5#initial_rgb_jpeg/0", next_observation_ref=None, metadata={"source_dataset": str(source)}, ) with CILImageReader(tmp_path / "collection") as reader: actual = reader.read(record) cached = reader.read(record) assert len(reader._handles) == 1 assert actual.shape == expected.shape assert actual.dtype == np.uint8 assert np.array_equal(actual, cached) def test_rgb_trainer_smoke_writes_checkpoint(tmp_path) -> None: h5py = pytest.importorskip("h5py") image_module = pytest.importorskip("PIL.Image") torch = pytest.importorskip("torch") dataset_dir = tmp_path / "data" generate_cil_dataset( backend="toy", tasks=built_in_toy_tasks()[:2], out_dir=dataset_dir, num_states_per_task=1, k=2, seed=9, shard_size=8, inline_observations=True, ) image = np.zeros((24, 24, 3), dtype=np.uint8) image[..., 0] = 180 buffer = io.BytesIO() image_module.fromarray(image).save(buffer, format="JPEG", quality=90) encoded = np.frombuffer(buffer.getvalue(), dtype=np.uint8) with h5py.File(dataset_dir / "observations.h5", "w") as handle: images = handle.create_dataset( "initial_rgb_jpeg", shape=(1,), dtype=h5py.vlen_dtype(np.dtype("uint8")), ) images[0] = encoded for shard in (dataset_dir / "shards").glob("*.jsonl"): records = [json.loads(line) for line in shard.read_text().splitlines()] for record in records: record["observation_ref"] = "observations.h5#initial_rgb_jpeg/0" shard.write_text("".join(json.dumps(record) + "\n" for record in records)) run_dir = tmp_path / "run" result = DoVLATrainer( TrainerConfig( dataset_dir=dataset_dir, output_dir=run_dir, epochs=1, batch_groups=1, records_per_group=2, pair_count_per_group=1, hidden_dim=24, action_horizon=2, effect_dim=8, observation_mode="rgb", device="cpu", ) ).train() checkpoint = torch.load(run_dir / "best.pt", map_location="cpu", weights_only=False) assert result["best"] assert checkpoint["model_config"]["observation_mode"] == "rgb"