corridorkey-mrp-mlx/tests/test_engine.py
cmoyates 4bad586112
feat: add CorridorKeyMLXEngine integration surface
Drop-in MLX backend for main CorridorKey repo. Engine class wraps
existing model/inference with process_frame() API returning alpha,
fg, comp, processed as uint8 numpy arrays. Lowers Python to >=3.11,
removes unused deps, adds smoke script and 16 contract tests.

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
2026-03-01 06:43:07 -03:30

138 lines
5.3 KiB
Python

"""Tests for CorridorKeyMLXEngine contract."""
from __future__ import annotations
from pathlib import Path
import numpy as np
import pytest
from corridorkey_mlx.engine import (
CorridorKeyMLXEngine,
_validate_image,
_validate_mask,
)
CHECKPOINT = Path("checkpoints/corridorkey_mlx.safetensors")
HAS_CHECKPOINT = CHECKPOINT.exists()
# ---------------------------------------------------------------------------
# Input validation (no checkpoint needed)
# ---------------------------------------------------------------------------
class TestValidateImage:
def test_rejects_non_ndarray(self) -> None:
with pytest.raises(TypeError, match="numpy ndarray"):
_validate_image("not_an_array") # type: ignore[arg-type]
def test_rejects_wrong_dtype(self) -> None:
img = np.zeros((64, 64, 3), dtype=np.float32)
with pytest.raises(ValueError, match="uint8"):
_validate_image(img)
def test_rejects_wrong_shape(self) -> None:
img = np.zeros((64, 64), dtype=np.uint8)
with pytest.raises(ValueError, match="\\(H, W, 3\\)"):
_validate_image(img)
def test_rejects_wrong_channels(self) -> None:
img = np.zeros((64, 64, 4), dtype=np.uint8)
with pytest.raises(ValueError, match="\\(H, W, 3\\)"):
_validate_image(img)
def test_accepts_valid_image(self) -> None:
img = np.zeros((64, 64, 3), dtype=np.uint8)
_validate_image(img) # should not raise
class TestValidateMask:
def test_rejects_non_ndarray(self) -> None:
with pytest.raises(TypeError, match="numpy ndarray"):
_validate_mask("not_an_array") # type: ignore[arg-type]
def test_rejects_wrong_dtype(self) -> None:
mask = np.zeros((64, 64), dtype=np.float32)
with pytest.raises(ValueError, match="uint8"):
_validate_mask(mask)
def test_accepts_hw(self) -> None:
mask = np.zeros((64, 64), dtype=np.uint8)
_validate_mask(mask) # should not raise
def test_accepts_hw1(self) -> None:
mask = np.zeros((64, 64, 1), dtype=np.uint8)
_validate_mask(mask) # should not raise
def test_rejects_hw3(self) -> None:
mask = np.zeros((64, 64, 3), dtype=np.uint8)
with pytest.raises(ValueError, match="\\(H, W\\) or \\(H, W, 1\\)"):
_validate_mask(mask)
# ---------------------------------------------------------------------------
# Engine init
# ---------------------------------------------------------------------------
class TestEngineInit:
def test_missing_checkpoint_raises(self) -> None:
with pytest.raises(FileNotFoundError, match="not found"):
CorridorKeyMLXEngine(checkpoint_path="/nonexistent/weights.safetensors")
# ---------------------------------------------------------------------------
# Engine integration (requires checkpoint)
# ---------------------------------------------------------------------------
@pytest.mark.skipif(not HAS_CHECKPOINT, reason="Checkpoint not available")
class TestEngineIntegration:
"""Integration tests that load the real model."""
@pytest.fixture(scope="class")
def engine(self) -> CorridorKeyMLXEngine:
return CorridorKeyMLXEngine(
checkpoint_path=CHECKPOINT,
img_size=512,
compile=False,
)
def test_output_keys(self, engine: CorridorKeyMLXEngine) -> None:
image = np.random.default_rng(42).integers(0, 256, (64, 64, 3), dtype=np.uint8)
mask = np.random.default_rng(42).integers(0, 256, (64, 64), dtype=np.uint8)
result = engine.process_frame(image, mask)
assert set(result.keys()) == {"alpha", "fg", "comp", "processed"}
def test_output_shapes(self, engine: CorridorKeyMLXEngine) -> None:
h, w = 100, 150
image = np.random.default_rng(42).integers(0, 256, (h, w, 3), dtype=np.uint8)
mask = np.random.default_rng(42).integers(0, 256, (h, w), dtype=np.uint8)
result = engine.process_frame(image, mask)
assert result["alpha"].shape == (h, w)
assert result["fg"].shape == (h, w, 3)
assert result["comp"].shape == (h, w, 3)
assert result["processed"].shape == (h, w, 3)
def test_output_dtypes(self, engine: CorridorKeyMLXEngine) -> None:
image = np.random.default_rng(42).integers(0, 256, (64, 64, 3), dtype=np.uint8)
mask = np.random.default_rng(42).integers(0, 256, (64, 64), dtype=np.uint8)
result = engine.process_frame(image, mask)
for key, arr in result.items():
assert arr.dtype == np.uint8, f"{key} dtype is {arr.dtype}"
def test_mask_hw1_accepted(self, engine: CorridorKeyMLXEngine) -> None:
image = np.random.default_rng(42).integers(0, 256, (64, 64, 3), dtype=np.uint8)
mask = np.random.default_rng(42).integers(0, 256, (64, 64, 1), dtype=np.uint8)
result = engine.process_frame(image, mask)
assert "alpha" in result
def test_refiner_scale_zero_returns_coarse(
self, engine: CorridorKeyMLXEngine
) -> None:
image = np.random.default_rng(42).integers(0, 256, (64, 64, 3), dtype=np.uint8)
mask = np.random.default_rng(42).integers(0, 256, (64, 64), dtype=np.uint8)
result = engine.process_frame(image, mask, refiner_scale=0.0)
assert result["alpha"].shape == (64, 64)