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>
8.3 KiB
| title | type | date |
|---|---|---|
| feat: Add CorridorKeyMLXEngine integration surface | feat | 2026-03-01 |
feat: Add CorridorKeyMLXEngine integration surface
Overview
Expose a stable CorridorKeyMLXEngine class so the main CorridorKey repo can consume this package as a drop-in MLX backend. Currently only flat functions exist (load_model, infer, infer_and_save) — none accept in-memory arrays, and there is no engine lifecycle or process_frame(...) contract.
Problem Statement
The main CorridorKey repo expects a backend engine with:
- Constructor:
checkpoint_path,device,img_size,use_refiner - Method:
process_frame(image, mask_linear, ...)returning dict withalpha,fg,comp,processed
This repo currently:
- Has no engine class — only loose functions in
pipeline.py infer()accepts file paths, not numpy/PIL arrays- Returns only
alpha+foreground, missingcompandprocessed - Has no
refiner_scale, despill, despeckle, or linear/sRGB handling DEFAULT_CHECKPOINTis a relative path (breaks as installed library)requires-python >= 3.12but main repo targets 3.11
Proposed Solution
Thin adapter class in src/corridorkey_mlx/engine.py wrapping existing model/inference code. No core model changes.
Technical Approach
Phase 1: Engine adapter + packaging fixes
1a. Lower Python version to >=3.11
- All source already uses
from __future__ import annotations— no 3.12-only syntax - Update
pyproject.toml:requires-python = ">=3.11" - Update
tool.ruff.target-versionandtool.mypy.python_version - Check dep version floors are 3.11-compatible (numpy >=2.4.2 may need lowering — numpy 2.0+ supports 3.11)
- Run
uv run ruff check .+uv run mypy src/to verify
1b. Create src/corridorkey_mlx/engine.py
class CorridorKeyMLXEngine:
def __init__(
self,
checkpoint_path: str | Path,
device: str | None = None, # ignored on MLX, accepted for compat
img_size: int = 2048, # production default
use_refiner: bool = True,
compile: bool = True,
) -> None: ...
def process_frame(
self,
image: np.ndarray, # uint8 HWC RGB
mask_linear: np.ndarray, # uint8 HW or HW1 grayscale
refiner_scale: float = 1.0,
input_is_linear: bool = False,
fg_is_straight: bool = True,
despill_strength: float = 1.0,
auto_despeckle: bool = True,
despeckle_size: int = 400,
) -> dict[str, np.ndarray]: ...
Constructor:
- Validates
checkpoint_pathexists (absolute or resolved) - Calls existing
load_model()withimg_sizeandcompile - Stores
use_refiner,img_size - Logs warning if
deviceis not None
process_frame() pipeline:
- Input validation: assert uint8 HWC(3) for image, uint8 HW or HW1 for mask
- Store original resolution for output resize
- Convert to float32 [0,1]:
image / 255.0,mask / 255.0 - Reshape mask: ensure
(H, W, 1) - Resize to
img_size x img_sizevia PIL (bicubic) — reuse existing pattern - Preprocess: call existing
normalize_rgb()+preprocess()fromio/image.py - Forward pass:
self._model(x)— raw output dict - Materialize all outputs
- Select outputs: if
use_refiner=True, usealpha_final/fg_final; else usealpha_coarse/fg_coarse - Apply refiner_scale (output-space lerp):
alpha = lerp(alpha_coarse, alpha_final, refiner_scale) - Postprocess to uint8 via existing
postprocess_alpha()/postprocess_foreground() - Resize outputs back to original input resolution
- Despill/despeckle: no-op stubs for now (documented, warn once)
- Composite:
comp = (fg * alpha_3ch + bg * (1 - alpha_3ch))with black bg, uint8 - Return
{"alpha": ..., "fg": ..., "comp": ..., "processed": fg}—processed= fg until despill/despeckle implemented
1c. Export from __init__.py
from corridorkey_mlx.engine import CorridorKeyMLXEngine
So callers can do from corridorkey_mlx import CorridorKeyMLXEngine.
1d. Fix DEFAULT_CHECKPOINT
Remove relative-path default from pipeline.py. Engine requires explicit checkpoint_path.
Phase 2: Smoke script + tests
2a. scripts/smoke_engine.py
- Takes
--image,--hint,--checkpointargs - Instantiates
CorridorKeyMLXEngine - Runs
process_frame() - Prints output shapes and value ranges
- Optionally saves outputs
2b. Tests in tests/test_engine.py
- test_engine_init_requires_checkpoint: missing path raises error
- test_engine_output_keys: process_frame returns
alpha,fg,comp,processed - test_engine_output_shapes: all outputs match input spatial dims
- test_engine_output_dtypes: all uint8
- test_engine_mask_shape_normalization: HW and HW1 both accepted
Use small synthetic inputs (e.g. 64x64 random) with checkpoint. Mark as @pytest.mark.skipif when checkpoint not available.
Phase 3: Documentation
3a. README section: "Using as a backend"
Cover:
- Editable install:
uv pip install -e ../corridorkey-mlx - Canonical import:
from corridorkey_mlx import CorridorKeyMLXEngine - Constructor params and defaults
- Expected checkpoint format (.safetensors, converted)
- Expected image/hint formats (uint8 HWC RGB, uint8 HW grayscale)
- Output dict keys and shapes
- Smoke command example
- Migration note: standalone script vs backend engine usage
3b. Docstrings
Engine class and process_frame get thorough docstrings covering:
- Input formats (uint8 HWC RGB, uint8 HW mask)
- Output formats and semantics
- Preprocessing chain (resize, ImageNet norm, NHWC)
- Which params are stubs (despill, despeckle)
- Linear vs sRGB assumptions
- img_size: 512 for dev, 2048 for production
Acceptance Criteria
from corridorkey_mlx import CorridorKeyMLXEngineworks- Constructor accepts
checkpoint_path,device,img_size,use_refiner,compile process_frame()accepts numpy uint8 arrays, returns dict withalpha,fg,comp,processed- Outputs resized to original input resolution
use_refiner=Falsereturns coarse predictionsrefiner_scaleblends between coarse and refined- Despill/despeckle are documented stubs (no-op, warn once)
compcomposites fg over black backgroundprocessed= fg (until despill/despeckle implemented)- Python >=3.11 works
- Smoke script runs one frame successfully
- Existing 94 tests still pass
- README documents backend usage + editable install
Design Decisions
| Decision | Choice | Rationale |
|---|---|---|
refiner_scale semantics |
output-space lerp | no model changes needed |
use_refiner=False |
use alpha_coarse/fg_coarse from output dict |
model always runs full forward; adapter selects outputs |
| despill/despeckle | no-op stubs, warn once | algorithms unknown; unblock integration now |
comp background |
black (0,0,0) | common default for matte compositing |
processed meaning |
same as fg for now |
placeholder until despill/despeckle exist |
device param |
accepted, ignored, log warning | compat with Torch engine signature |
default img_size |
2048 (engine), 512 (dev scripts) | model trained at 2048; scripts keep 512 for speed |
input_is_linear |
accepted, no-op for now | ImageNet stats assume sRGB; linearization would break normalization |
| checkpoint default | none — required param | relative paths break as library |
Dependencies & Risks
- numpy version floor:
numpy>=2.4.2may not support 3.11. Need to check and potentially lower to>=1.26or>=2.0. - Despill/despeckle gap: real implementations need the main repo's algorithm. Document as TODO.
- refiner_scale at compile time: output-space lerp works with compiled model since it's post-forward. Logit-space would require model changes and recompilation.
Unresolved Questions
processed— what exactly does main repo return here? Despilled fg? Masked fg?compbg — always black or configurable?use_refiner=False— skip refiner forward pass (perf) or just ignore outputs (simpler)?refiner_scale— logit-space or output-space? (plan assumes output-space)- despill/despeckle algorithms — need from main repo for real impl
input_is_linear— does original model ever receive linear-light inputs?mask_linearnaming — is it actually linear-light or just naming convention?