corridorkey-mrp-mlx/scripts/smoke_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

60 lines
1.9 KiB
Python

"""Smoke test for the CorridorKeyMLXEngine integration surface.
Instantiates the engine, runs one frame, prints output shapes and value ranges.
"""
from __future__ import annotations
import argparse
from pathlib import Path
import numpy as np
from PIL import Image
from corridorkey_mlx import CorridorKeyMLXEngine
def main() -> None:
parser = argparse.ArgumentParser(description="Smoke test: CorridorKeyMLXEngine")
parser.add_argument("--image", type=Path, required=True, help="RGB input image")
parser.add_argument("--hint", type=Path, required=True, help="Grayscale alpha hint")
parser.add_argument(
"--checkpoint",
type=Path,
default=Path("checkpoints/corridorkey_mlx.safetensors"),
)
parser.add_argument("--img-size", type=int, default=512)
parser.add_argument("--output-dir", type=Path, default=None)
args = parser.parse_args()
print(f"Loading engine (img_size={args.img_size})...")
engine = CorridorKeyMLXEngine(
checkpoint_path=args.checkpoint,
img_size=args.img_size,
compile=False,
)
rgb = np.asarray(Image.open(args.image).convert("RGB"))
mask = np.asarray(Image.open(args.hint).convert("L"))
print(f"Input image: {rgb.shape} {rgb.dtype}")
print(f"Input mask: {mask.shape} {mask.dtype}")
result = engine.process_frame(rgb, mask)
for key, arr in result.items():
print(f" {key}: shape={arr.shape} dtype={arr.dtype} range=[{arr.min()}, {arr.max()}]")
if args.output_dir is not None:
out = Path(args.output_dir)
out.mkdir(parents=True, exist_ok=True)
Image.fromarray(result["alpha"], mode="L").save(out / "alpha.png")
Image.fromarray(result["fg"], mode="RGB").save(out / "fg.png")
Image.fromarray(result["comp"], mode="RGB").save(out / "comp.png")
print(f"Saved outputs to {out}")
print("Smoke test passed.")
if __name__ == "__main__":
main()