InfiniSplat / tests /test_batch_inputs.py
PLUS-WAVE's picture
Add meerkat demo and Node 22 runtime
a6825eb verified
Raw
History Blame Contribute Delete
4.91 kB
from argparse import Namespace
from pathlib import Path
import pytest
import src.demo.infer_batch_images as batch_module
from src.demo.infer_batch_images import (
_collect_images,
_resolve_output_dir,
_resolve_prompt_depth_path,
run_batch,
)
def _args(input_path, **overrides) -> Namespace:
values = {
"input_path": input_path,
"mode": "rgb",
"limit": 0,
"prompt_depth": None,
"prompt_depth_dir": None,
}
values.update(overrides)
return Namespace(**values)
def test_collect_images_accepts_single_image(tmp_path) -> None:
image = tmp_path / "input.jpg"
image.touch()
assert _collect_images(_args(image)) == [image]
def test_collect_images_accepts_directory(tmp_path) -> None:
first = tmp_path / "a.png"
second = tmp_path / "b.jpg"
ignored = tmp_path / "notes.txt"
first.touch()
second.touch()
ignored.touch()
assert _collect_images(_args(tmp_path)) == [first, second]
def test_collect_images_does_not_scan_subdirectories(tmp_path) -> None:
image = tmp_path / "input.jpg"
nested = tmp_path / "nested"
nested.mkdir()
nested_image = nested / "nested.jpg"
image.touch()
nested_image.touch()
assert _collect_images(_args(tmp_path)) == [image]
def test_collect_images_rejects_duplicate_stems(tmp_path) -> None:
(tmp_path / "frame.jpg").touch()
(tmp_path / "frame.png").touch()
with pytest.raises(ValueError, match="unique filename stems"):
_collect_images(_args(tmp_path))
def test_lidar_directory_keeps_png_rgb_with_npz_depth(tmp_path) -> None:
image = tmp_path / "frame.png"
depth = tmp_path / "frame.npz"
image.touch()
depth.touch()
assert _collect_images(_args(tmp_path, mode="lidar")) == [image]
def test_default_output_layout() -> None:
rgb_output = _resolve_output_dir(Namespace(output_dir=None, mode="rgb"))
lidar_output = _resolve_output_dir(Namespace(output_dir=None, mode="lidar"))
assert rgb_output == Path("outputs/demo/rgb")
assert lidar_output == Path("outputs/demo/lidar")
def test_lidar_mode_matches_adjacent_depth_by_stem(tmp_path) -> None:
image = tmp_path / "frame.jpg"
depth = tmp_path / "frame.npy"
image.touch()
depth.touch()
args = _args(
image,
mode="lidar",
prompt_depth=None,
prompt_depth_dir=None,
)
assert _resolve_prompt_depth_path(args, image) == depth
def test_lidar_mode_does_not_use_input_png_as_depth(tmp_path) -> None:
image = tmp_path / "frame.png"
image.touch()
args = _args(
image,
mode="lidar",
prompt_depth=None,
prompt_depth_dir=None,
)
with pytest.raises(FileNotFoundError):
_resolve_prompt_depth_path(args, image)
def _run_args(input_path: Path, output_dir: Path, **overrides) -> Namespace:
values = {
"input_path": input_path,
"mode": "rgb",
"checkpoint": None,
"output_dir": output_dir,
"limit": 0,
"overwrite": False,
"device": "auto",
"intrinsics_file": None,
"focal_px": None,
"focal_mm": None,
"prompt_depth": None,
"prompt_depth_dir": None,
"disable_floater_filter": False,
"no_video": True,
"export_html": False,
}
values.update(overrides)
return Namespace(**values)
def test_completed_batch_does_not_load_model(tmp_path, monkeypatch) -> None:
image = tmp_path / "frame.jpg"
output_dir = tmp_path / "outputs"
scene_ply = output_dir / "frame" / "frame.ply"
image.touch()
scene_ply.parent.mkdir(parents=True)
scene_ply.touch()
monkeypatch.setattr(
batch_module,
"load_demo_config",
lambda *_args, **_kwargs: pytest.fail("completed batch loaded the model config"),
)
result = run_batch(_run_args(image, output_dir))
assert result["succeeded"] == 0
assert result["skipped"] == 1
def test_html_only_batch_does_not_load_model(tmp_path, monkeypatch) -> None:
image = tmp_path / "frame.jpg"
output_dir = tmp_path / "outputs"
scene_ply = output_dir / "frame" / "frame.ply"
html = output_dir / "frame" / "frame.html"
image.touch()
scene_ply.parent.mkdir(parents=True)
scene_ply.touch()
monkeypatch.setattr(batch_module.shutil, "which", lambda _command: "/usr/bin/splat-transform")
monkeypatch.setattr(
batch_module,
"load_demo_config",
lambda *_args, **_kwargs: pytest.fail("HTML-only batch loaded the model config"),
)
monkeypatch.setattr(
batch_module,
"_convert_scene_if_requested",
lambda _export_html, _paths, _viewer_settings: html.touch(),
)
result = run_batch(_run_args(image, output_dir, export_html=True))
assert result["succeeded"] == 1
assert result["skipped"] == 0
assert html.exists()