| """apairo write-back: label preparation, Session.save_apairo, and (if apairo is |
| installed) a real round-trip. The unit tests need no apairo — the writer is only |
| exercised for real in the guarded integration test at the bottom. |
| """ |
|
|
| from __future__ import annotations |
|
|
| from dataclasses import dataclass |
|
|
| import numpy as np |
| import pytest |
|
|
| from splasher.adapters.apairo_source import ApairoSource |
| from splasher.adapters.apairo_writer import project_grid_labels, reference_slice |
| from splasher.core.grid import Grid |
| from splasher.engine import Session |
|
|
|
|
| |
| def test_project_grid_labels_takes_cell_class(): |
| grid = Grid(0.0, 4.0, 0.0, 4.0, 1.0) |
| raster = grid.empty_raster(0) |
| raster[1, 2] = 7 |
| pts = np.array([[2.5, 1.5], |
| [0.5, 0.5], |
| [99.0, 99.0]]) |
| out = project_grid_labels(pts, raster, grid, ignore_id=0) |
| assert out.tolist() == [7, 0, 0] |
|
|
|
|
| def test_reference_slice_extracts_the_right_block(): |
| full = np.array([1, 1, 1, 2, 2, 9], dtype=np.int32) |
| sizes = [("lidarA", 3), ("lidarB", 2), ("extra", 1)] |
| assert reference_slice(full, sizes, "lidarB", 0).tolist() == [2, 2] |
| assert reference_slice(full, sizes, "lidarA", 0).tolist() == [1, 1, 1] |
| assert reference_slice(full, sizes, "absent", 0) is None |
|
|
|
|
| |
| @dataclass |
| class _FakeSample: |
| data: dict |
| timestamp: float | None = None |
|
|
|
|
| class _FakeDataset: |
| """Two lidar frames (points in [0,4)²), synchronous, each with a timestamp.""" |
|
|
| is_synchronous = True |
|
|
| def __init__(self): |
| self.keys = ["lidar"] |
| rng = np.random.RandomState(0) |
| self._frames = [ |
| {"lidar": (rng.rand(20, 4) * np.array([4, 4, 1, 1])).astype(np.float32)} |
| for _ in range(2) |
| ] |
| self._ts = [1000.0, 1000.1] |
|
|
| def __len__(self): |
| return len(self._frames) |
|
|
| def __getitem__(self, i): |
| return _FakeSample(self._frames[i], timestamp=self._ts[i]) |
|
|
|
|
| def _fake_apairo_session(): |
| src = ApairoSource(_FakeDataset(), dataset_root="/nonexistent/dset", reference="lidar") |
| return Session(src), src |
|
|
|
|
| def test_apairo_meta_surfaces_reference_and_write_root(): |
| _, src = _fake_apairo_session() |
| meta = src.apairo_meta() |
| assert meta["is_apairo"] is True |
| assert meta["reference"] == "lidar" |
| assert meta["point_channels"] == ["lidar"] |
| assert meta["write_root"] == "/nonexistent/dset" |
|
|
|
|
| def test_save_apairo_grid_projects_painted_cells(monkeypatch): |
| import splasher.adapters.apairo_writer as W |
|
|
| captured = {} |
|
|
| def _cap(root, ref, ch, by_ts): |
| captured.update(root=root, ref=ref, ch=ch, by_ts=by_ts) |
| return len(by_ts) |
|
|
| monkeypatch.setattr(W, "write_channel", _cap) |
|
|
| session, _ = _fake_apairo_session() |
| cls = session.labelset.paintable[0].id |
| session.set_active_class(cls) |
| session.set_frame(0) |
| session.paint_rect((-1e3, -1e3, 1e3, 1e3)) |
|
|
| out = session.save_apairo(channel="ground_truth", reference="lidar", mode="grid") |
| assert out["channel"] == "ground_truth" and out["mode"] == "grid" |
| assert captured["ref"] == "lidar" and captured["root"] == "/nonexistent/dset" |
| assert set(captured["by_ts"]) == {1000.0} |
| lab = captured["by_ts"][1000.0] |
| assert lab.shape == (20,) and np.all(lab == cls) |
|
|
|
|
| def test_save_apairo_points_writes_reference_slice(monkeypatch): |
| import splasher.adapters.apairo_writer as W |
|
|
| captured = {} |
| monkeypatch.setattr(W, "write_channel", |
| lambda root, ref, ch, by_ts: (captured.update(by_ts=by_ts) or len(by_ts))) |
|
|
| session, _ = _fake_apairo_session() |
| cls = session.labelset.paintable[0].id |
| session.set_active_class(cls) |
| session.set_active_targets({"points"}) |
| session.set_frame(1) |
| session.paint_rect((-1e3, -1e3, 1e3, 1e3)) |
|
|
| session.save_apairo(reference="lidar", mode="points") |
| assert set(captured["by_ts"]) == {1000.1} |
| assert np.all(captured["by_ts"][1000.1] == cls) |
|
|
|
|
| def test_save_apairo_rejects_non_apairo_source(): |
| from splasher.demo import make_demo_source |
|
|
| session = Session(make_demo_source(n_frames=2)) |
| assert session.apairo_meta() == {"is_apairo": False} |
| with pytest.raises(ValueError, match="not an apairo dataset"): |
| session.save_apairo() |
|
|
|
|
| |
| def test_apairo_real_roundtrip(tmp_path): |
| apairo = pytest.importorskip("apairo", reason="`apairo` extra not installed") |
|
|
| |
| lid = tmp_path / "lidar" |
| lid.mkdir() |
| ts = [] |
| for i in range(3): |
| pts = np.zeros((10, 4), np.float32) |
| pts[:, 0] = np.linspace(0.1, 3.9, 10) |
| pts[:, 1] = 1.5 |
| np.save(lid / f"{i:06d}.npy", pts) |
| ts.append(2000.0 + i * 0.1) |
| np.savetxt(lid / "timestamps.txt", ts) |
|
|
| src = ApairoSource.from_path(str(tmp_path), reference="lidar", tolerance=0.05) |
| session = Session(src) |
| cls = session.labelset.paintable[0].id |
| session.set_active_class(cls) |
| session.set_frame(1) |
| session.paint_rect((-1e3, -1e3, 1e3, 1e3)) |
|
|
| report = session.save_apairo(channel="ground_truth", reference="lidar", mode="grid") |
| assert report["frames"] == 1 |
|
|
| |
| |
| gt = tmp_path / "ground_truth" |
| assert sorted(p.name for p in gt.glob("*.npy")) == ["000001.npy"] |
| assert np.unique(np.load(gt / "000001.npy")).tolist() == [cls] |
| import yaml |
|
|
| reg = yaml.safe_load((tmp_path / ".apairo" / "channels.yaml").read_text())["channels"] |
| assert reg["ground_truth"]["loader"] == "npys" |
| assert reg["ground_truth"]["timestamps_from"] == "lidar" |
|
|
| |
| ds = apairo.RawDataset(str(tmp_path), keys=["lidar", "ground_truth"]) |
| sv = ds.synchronize(reference="ground_truth", tolerance=0.05) |
| assert len(sv) == 1 |
| assert round(sv[0].timestamp, 4) == 2000.1 |
| assert np.unique(sv[0].data["ground_truth"]).tolist() == [cls] |
|
|
|
|
| |
| def test_apairo_routes(monkeypatch): |
| pytest.importorskip("fastapi", reason="`api` extra not installed") |
| pytest.importorskip("httpx", reason="httpx required by fastapi.testclient") |
| from fastapi.testclient import TestClient |
|
|
| import splasher.adapters.apairo_writer as W |
| from splasher.server import create_app |
|
|
| captured = {} |
|
|
| def _capture(root, ref, ch, by_ts): |
| captured.update(ch=ch, by_ts=by_ts) |
| return len(by_ts) |
|
|
| monkeypatch.setattr(W, "write_channel", _capture) |
|
|
| session, _ = _fake_apairo_session() |
| client = TestClient(create_app(session)) |
|
|
| info = client.get("/api/apairo/info").json() |
| assert info["is_apairo"] is True and info["reference"] == "lidar" |
|
|
| cls = session.labelset.paintable[0].id |
| client.post("/api/class", json={"id": cls}) |
| client.post("/api/paint", json={"rect": [-1e3, -1e3, 1e3, 1e3]}) |
| r = client.post("/api/apairo/save", json={"channel": "gt", "reference": "lidar", "mode": "grid"}) |
| assert r.status_code == 200 and r.json()["channel"] == "gt" |
| assert captured["ch"] == "gt" |
|
|
| |
| assert client.post("/api/apairo/save", json={"reference": "nope"}).status_code == 422 |
|
|
|
|
| def test_apairo_info_false_for_plain_source(): |
| pytest.importorskip("fastapi") |
| pytest.importorskip("httpx") |
| from fastapi.testclient import TestClient |
|
|
| from splasher.demo import make_demo_source |
| from splasher.server import create_app |
|
|
| client = TestClient(create_app(Session(make_demo_source(n_frames=2)))) |
| assert client.get("/api/apairo/info").json() == {"is_apairo": False} |
|
|