Go to file
cmoyates 0481955cfd
feat(phase5): full model assembly + e2e parity
Wire GreenFormer (backbone + decoders + refiner), add image I/O,
inference pipeline, CLI entry point, and 3 test files (23 new tests).

E2e parity vs PyTorch: max abs err < 1.5e-4 across all outputs.

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
2026-03-01 06:07:43 -03:30
docs/plans feat(phase5): full model assembly + e2e parity 2026-03-01 06:07:43 -03:30
prompts chore: add phase 4-6 prompts, fix naming, note zoxide cd issue 2026-03-01 05:20:40 -03:30
scripts feat(phase5): full model assembly + e2e parity 2026-03-01 06:07:43 -03:30
src/corridorkey_mlx feat(phase5): full model assembly + e2e parity 2026-03-01 06:07:43 -03:30
tests feat(phase5): full model assembly + e2e parity 2026-03-01 06:07:43 -03:30
.gitignore init: repo scaffolding, deps, prompts, plan 2026-03-01 04:48:51 -03:30
.python-version init: repo scaffolding, deps, prompts, plan 2026-03-01 04:48:51 -03:30
CLAUDE.md chore: add phase 4-6 prompts, fix naming, note zoxide cd issue 2026-03-01 05:20:40 -03:30
main.py init: repo scaffolding, deps, prompts, plan 2026-03-01 04:48:51 -03:30
pyproject.toml init: repo scaffolding, deps, prompts, plan 2026-03-01 04:48:51 -03:30
README.md feat(phase5): full model assembly + e2e parity 2026-03-01 06:07:43 -03:30
uv.lock init: repo scaffolding, deps, prompts, plan 2026-03-01 04:48:51 -03:30

corridorkey-mlx

MLX inference port of CorridorKey for Apple Silicon.

Architecture

RGB image + coarse alpha hint (4ch)
        │
        ▼
┌──────────────────┐
│  Hiera backbone   │  (timm, features_only)
│  → 4 multiscale   │
│    feature maps    │
└──────────────────┘
        │
   ┌────┴────┐
   ▼         ▼
┌───────┐ ┌───────┐
│ Alpha │ │  FG   │
│ head  │ │ head  │
│ (1ch) │ │ (3ch) │
└───────┘ └───────┘
   │         │
   └────┬────┘
        ▼
┌──────────────────┐
│   CNN Refiner     │  RGB + coarse preds (7ch)
│   → delta logits  │  → sigmoid
└──────────────────┘
        │
        ▼
  final alpha + fg

Phased Roadmap

Phase Scope Status
1 PyTorch reference harness + fixture dump Done
2 MLX decoder/refiner blocks + parity tests Done
3 Checkpoint conversion (PyTorch → MLX) Done
4 Hiera backbone port Done
5 Full model assembly + e2e parity Done

See prompts/ for detailed phase instructions.

Usage

Setup

uv sync --group dev

Convert weights

Convert the PyTorch checkpoint to MLX safetensors (one-time):

uv run python scripts/convert_weights.py \
    --checkpoint checkpoints/CorridorKey_v1.0.pth \
    --output checkpoints/corridorkey_mlx.safetensors

Single-image inference

uv run python scripts/infer.py \
    --image input.png \
    --hint alpha_hint.png \
    --output-dir output/

Outputs output/alpha.png (alpha matte) and output/foreground.png (foreground).

Options:

  • --checkpoint PATH — MLX safetensors file (default: checkpoints/corridorkey_mlx.safetensors)
  • --img-size N — model input resolution (default: 512)
  • --output-dir DIR — output directory (default: output/)

Python API

from corridorkey_mlx.inference.pipeline import load_model, infer_and_save

model = load_model("checkpoints/corridorkey_mlx.safetensors", img_size=512)
results = infer_and_save(model, "input.png", "alpha_hint.png", "output/")

Development

uv run pytest              # tests
uv run ruff check .        # lint
uv run ruff format .       # format
uv run mypy src/           # type check

For PyTorch reference work:

uv sync --group reference

Reference Fixtures

Phase 1 generates golden reference tensors from PyTorch for MLX parity testing.

Format: single reference/fixtures/golden.npz (numpy compressed archive)

Generate:

uv run --group reference python scripts/dump_pytorch_reference.py \
    --checkpoint checkpoints/CorridorKey_v1.0.pth

Contents (all float32, NCHW, batch=1, img_size=512):

Key Shape Description
input (1, 4, 512, 512) Random input (seed=42)
encoder_feature_0 (1, 112, 128, 128) Backbone stride-4
encoder_feature_1 (1, 224, 64, 64) Backbone stride-8
encoder_feature_2 (1, 448, 32, 32) Backbone stride-16
encoder_feature_3 (1, 896, 16, 16) Backbone stride-32
alpha_logits (1, 1, 128, 128) Alpha decoder output (H/4)
fg_logits (1, 3, 128, 128) FG decoder output (H/4)
alpha_logits_up (1, 1, 512, 512) Alpha logits upsampled
fg_logits_up (1, 3, 512, 512) FG logits upsampled
alpha_coarse (1, 1, 512, 512) sigmoid(alpha_logits_up)
fg_coarse (1, 3, 512, 512) sigmoid(fg_logits_up)
delta_logits (1, 4, 512, 512) Refiner output (10x scaled)
alpha_final (1, 1, 512, 512) Final alpha prediction
fg_final (1, 3, 512, 512) Final FG prediction

Parity Results

End-to-end parity vs PyTorch reference (512×512, float32):

Tensor Max Abs Error Mean Abs Error
alpha_logits 8.8e-05 1.6e-05
fg_logits 1.5e-04 7.2e-06
alpha_coarse 9.7e-06 1.1e-06
fg_coarse 6.7e-06 1.1e-06
delta_logits 1.1e-04 4.3e-06
alpha_final 2.6e-05 8.7e-08
fg_final 9.5e-06 1.1e-06

Current Status

Phases 15 complete. Full model assembly with end-to-end parity verified.