vla / tests /test_cil_images.py
anhtld's picture
Initial commit: DoVLA-CIL codebase (h=16 breakthrough) (part 2)
20c251e verified
Raw
History Blame Contribute Delete
3.5 kB
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"