| 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" |
|
|