Merge pull request #4 from cmoyates/feat/phase5-assembly-parity
feat(phase5): full model assembly + e2e parity
This commit is contained in:
commit
a14f8a74e8
69
README.md
69
README.md
@ -37,23 +37,55 @@ RGB image + coarse alpha hint (4ch)
|
|||||||
|
|
||||||
| Phase | Scope | Status |
|
| Phase | Scope | Status |
|
||||||
|-------|-------|--------|
|
|-------|-------|--------|
|
||||||
| 1 | PyTorch reference harness + fixture dump | **In progress** |
|
| 1 | PyTorch reference harness + fixture dump | Done |
|
||||||
| 2 | MLX decoder/refiner blocks + parity tests | Not started |
|
| 2 | MLX decoder/refiner blocks + parity tests | Done |
|
||||||
| 3 | Checkpoint conversion (PyTorch → MLX) | Not started |
|
| 3 | Checkpoint conversion (PyTorch → MLX) | Done |
|
||||||
| 4 | Full inference pipeline | Not started |
|
| 4 | Hiera backbone port | Done |
|
||||||
| 5 | Optimization + benchmarking | Not started |
|
| 5 | Full model assembly + e2e parity | Done |
|
||||||
|
|
||||||
See `prompts/` for detailed phase instructions.
|
See `prompts/` for detailed phase instructions.
|
||||||
|
|
||||||
## Setup
|
## Usage
|
||||||
|
|
||||||
|
### Setup
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
uv sync --group dev
|
uv sync --group dev
|
||||||
```
|
```
|
||||||
|
|
||||||
For PyTorch reference work:
|
### Convert weights
|
||||||
|
|
||||||
|
Convert the PyTorch checkpoint to MLX safetensors (one-time):
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
uv sync --group reference
|
uv run python scripts/convert_weights.py \
|
||||||
|
--checkpoint checkpoints/CorridorKey_v1.0.pth \
|
||||||
|
--output checkpoints/corridorkey_mlx.safetensors
|
||||||
|
```
|
||||||
|
|
||||||
|
### Single-image inference
|
||||||
|
|
||||||
|
```bash
|
||||||
|
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
|
||||||
|
|
||||||
|
```python
|
||||||
|
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
|
## Development
|
||||||
@ -65,6 +97,11 @@ uv run ruff format . # format
|
|||||||
uv run mypy src/ # type check
|
uv run mypy src/ # type check
|
||||||
```
|
```
|
||||||
|
|
||||||
|
For PyTorch reference work:
|
||||||
|
```bash
|
||||||
|
uv sync --group reference
|
||||||
|
```
|
||||||
|
|
||||||
## Reference Fixtures
|
## Reference Fixtures
|
||||||
|
|
||||||
Phase 1 generates golden reference tensors from PyTorch for MLX parity testing.
|
Phase 1 generates golden reference tensors from PyTorch for MLX parity testing.
|
||||||
@ -96,6 +133,20 @@ uv run --group reference python scripts/dump_pytorch_reference.py \
|
|||||||
| `alpha_final` | (1, 1, 512, 512) | Final alpha prediction |
|
| `alpha_final` | (1, 1, 512, 512) | Final alpha prediction |
|
||||||
| `fg_final` | (1, 3, 512, 512) | Final FG 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
|
## Current Status
|
||||||
|
|
||||||
Phase 1 in progress — reference harness and fixture dump implemented.
|
Phases 1–5 complete. Full model assembly with end-to-end parity verified.
|
||||||
|
|||||||
@ -185,11 +185,11 @@ source_key -> dest_key | src_shape -> dst_shape | transform
|
|||||||
|
|
||||||
**Deliverables:**
|
**Deliverables:**
|
||||||
|
|
||||||
- [ ] `src/corridorkey_mlx/model/backbone.py` -- Hiera MLX port
|
- [x] `src/corridorkey_mlx/model/backbone.py` -- Hiera MLX port
|
||||||
- [ ] `src/corridorkey_mlx/model/corridorkey.py` -- full model composition
|
- [x] `src/corridorkey_mlx/model/corridorkey.py` -- full model composition
|
||||||
- [ ] `src/corridorkey_mlx/inference/pipeline.py` -- load, preprocess, forward, postprocess, save
|
- [x] `src/corridorkey_mlx/inference/pipeline.py` -- load, preprocess, forward, postprocess, save
|
||||||
- [ ] `src/corridorkey_mlx/io/image.py` -- PIL-based image I/O + preprocessing
|
- [x] `src/corridorkey_mlx/io/image.py` -- PIL-based image I/O + preprocessing
|
||||||
- [ ] End-to-end parity test against PyTorch golden output
|
- [x] End-to-end parity test against PyTorch golden output
|
||||||
|
|
||||||
**Files touched:**
|
**Files touched:**
|
||||||
|
|
||||||
@ -217,9 +217,9 @@ source_key -> dest_key | src_shape -> dst_shape | transform
|
|||||||
|
|
||||||
**Deliverables:**
|
**Deliverables:**
|
||||||
|
|
||||||
- [ ] `scripts/bench_mlx.py` -- latency, throughput, memory reporting
|
- [ ] `scripts/bench_mlx.py` -- latency, throughput, memory reporting (future)
|
||||||
- [ ] `scripts/compare_reference.py` -- side-by-side output comparison
|
- [ ] `scripts/compare_reference.py` -- side-by-side output comparison (future)
|
||||||
- [ ] Performance optimizations (compile, memory layout, batching)
|
- [ ] Performance optimizations (compile, memory layout, batching) (future)
|
||||||
|
|
||||||
**Potential optimizations:**
|
**Potential optimizations:**
|
||||||
- `mx.compile()` on hot paths
|
- `mx.compile()` on hot paths
|
||||||
|
|||||||
64
scripts/infer.py
Normal file
64
scripts/infer.py
Normal file
@ -0,0 +1,64 @@
|
|||||||
|
#!/usr/bin/env python3
|
||||||
|
"""Single-image inference with CorridorKey MLX.
|
||||||
|
|
||||||
|
Usage:
|
||||||
|
uv run python scripts/infer.py --image input.png --hint alpha_hint.png
|
||||||
|
uv run python scripts/infer.py --image input.png --hint alpha_hint.png --output-dir results/
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import argparse
|
||||||
|
import time
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
from corridorkey_mlx.inference.pipeline import (
|
||||||
|
DEFAULT_CHECKPOINT,
|
||||||
|
DEFAULT_IMG_SIZE,
|
||||||
|
infer_and_save,
|
||||||
|
load_model,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def main() -> None:
|
||||||
|
parser = argparse.ArgumentParser(description="CorridorKey MLX inference")
|
||||||
|
parser.add_argument("--image", type=Path, required=True, help="RGB input image")
|
||||||
|
parser.add_argument("--hint", type=Path, required=True, help="Alpha hint (grayscale)")
|
||||||
|
parser.add_argument(
|
||||||
|
"--checkpoint",
|
||||||
|
type=Path,
|
||||||
|
default=DEFAULT_CHECKPOINT,
|
||||||
|
help="MLX safetensors checkpoint",
|
||||||
|
)
|
||||||
|
parser.add_argument("--img-size", type=int, default=DEFAULT_IMG_SIZE, help="Model input size")
|
||||||
|
parser.add_argument("--output-dir", type=Path, default=Path("output"), help="Output directory")
|
||||||
|
args = parser.parse_args()
|
||||||
|
|
||||||
|
if not args.image.exists():
|
||||||
|
print(f"Image not found: {args.image}")
|
||||||
|
raise SystemExit(1)
|
||||||
|
if not args.hint.exists():
|
||||||
|
print(f"Alpha hint not found: {args.hint}")
|
||||||
|
raise SystemExit(1)
|
||||||
|
if not args.checkpoint.exists():
|
||||||
|
print(f"Checkpoint not found: {args.checkpoint}")
|
||||||
|
print("Run: uv run python scripts/convert_weights.py")
|
||||||
|
raise SystemExit(1)
|
||||||
|
|
||||||
|
print(f"Loading model (img_size={args.img_size})...")
|
||||||
|
t0 = time.perf_counter()
|
||||||
|
model = load_model(args.checkpoint, args.img_size)
|
||||||
|
print(f" Model loaded in {time.perf_counter() - t0:.2f}s")
|
||||||
|
|
||||||
|
print(f"Running inference on {args.image}...")
|
||||||
|
t0 = time.perf_counter()
|
||||||
|
results = infer_and_save(model, args.image, args.hint, args.output_dir)
|
||||||
|
print(f" Inference + save in {time.perf_counter() - t0:.2f}s")
|
||||||
|
|
||||||
|
print(f"\nSaved: {args.output_dir / 'alpha.png'}, {args.output_dir / 'foreground.png'}")
|
||||||
|
print(f" Alpha shape: {results['alpha'].shape}")
|
||||||
|
print(f" Foreground shape: {results['foreground'].shape}")
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
main()
|
||||||
@ -1,4 +1,88 @@
|
|||||||
"""Inference pipeline (not yet implemented).
|
"""Inference pipeline.
|
||||||
|
|
||||||
Orchestrates: load image → preprocess → model forward → postprocess → save.
|
Orchestrates: load weights -> load image -> preprocess -> model forward -> postprocess -> save.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import TYPE_CHECKING
|
||||||
|
|
||||||
|
import mlx.core as mx
|
||||||
|
|
||||||
|
from corridorkey_mlx.io.image import (
|
||||||
|
load_alpha_hint,
|
||||||
|
load_image,
|
||||||
|
postprocess_alpha,
|
||||||
|
postprocess_foreground,
|
||||||
|
preprocess,
|
||||||
|
save_alpha,
|
||||||
|
save_foreground,
|
||||||
|
)
|
||||||
|
from corridorkey_mlx.model.corridorkey import GreenFormer
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
import numpy as np
|
||||||
|
|
||||||
|
DEFAULT_CHECKPOINT = Path("checkpoints/corridorkey_mlx.safetensors")
|
||||||
|
DEFAULT_IMG_SIZE = 512
|
||||||
|
|
||||||
|
|
||||||
|
def load_model(
|
||||||
|
checkpoint: str | Path = DEFAULT_CHECKPOINT,
|
||||||
|
img_size: int = DEFAULT_IMG_SIZE,
|
||||||
|
) -> GreenFormer:
|
||||||
|
"""Build GreenFormer and load weights from safetensors checkpoint."""
|
||||||
|
model = GreenFormer(img_size=img_size)
|
||||||
|
model.load_checkpoint(checkpoint)
|
||||||
|
return model
|
||||||
|
|
||||||
|
|
||||||
|
def infer(
|
||||||
|
model: GreenFormer,
|
||||||
|
image_path: str | Path,
|
||||||
|
alpha_hint_path: str | Path,
|
||||||
|
) -> dict[str, mx.array]:
|
||||||
|
"""Run single-image inference.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
model: Loaded GreenFormer model.
|
||||||
|
image_path: Path to RGB input image.
|
||||||
|
alpha_hint_path: Path to coarse alpha hint (grayscale).
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Model output dict with all intermediate and final tensors.
|
||||||
|
"""
|
||||||
|
img_size = model.backbone.img_size
|
||||||
|
rgb = load_image(image_path, img_size=img_size)
|
||||||
|
alpha_hint = load_alpha_hint(alpha_hint_path, img_size=img_size)
|
||||||
|
x = preprocess(rgb, alpha_hint)
|
||||||
|
outputs = model(x)
|
||||||
|
# materialize all outputs
|
||||||
|
mx.eval(outputs) # noqa: S307
|
||||||
|
return outputs
|
||||||
|
|
||||||
|
|
||||||
|
def infer_and_save(
|
||||||
|
model: GreenFormer,
|
||||||
|
image_path: str | Path,
|
||||||
|
alpha_hint_path: str | Path,
|
||||||
|
output_dir: str | Path,
|
||||||
|
) -> dict[str, np.ndarray]:
|
||||||
|
"""Run inference and save alpha + foreground PNGs.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Dict with 'alpha' and 'foreground' as uint8 numpy arrays.
|
||||||
|
"""
|
||||||
|
outputs = infer(model, image_path, alpha_hint_path)
|
||||||
|
|
||||||
|
alpha_arr = postprocess_alpha(outputs["alpha_final"])
|
||||||
|
fg_arr = postprocess_foreground(outputs["fg_final"])
|
||||||
|
|
||||||
|
out = Path(output_dir)
|
||||||
|
out.mkdir(parents=True, exist_ok=True)
|
||||||
|
|
||||||
|
save_alpha(alpha_arr, out / "alpha.png")
|
||||||
|
save_foreground(fg_arr, out / "foreground.png")
|
||||||
|
|
||||||
|
return {"alpha": alpha_arr, "foreground": fg_arr}
|
||||||
|
|||||||
@ -1 +1,104 @@
|
|||||||
"""Image loading, saving, and preprocessing (not yet implemented)."""
|
"""Image loading, saving, and preprocessing for CorridorKey inference.
|
||||||
|
|
||||||
|
All preprocessing produces NHWC tensors suitable for the MLX model.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from typing import TYPE_CHECKING
|
||||||
|
|
||||||
|
import mlx.core as mx
|
||||||
|
import numpy as np
|
||||||
|
from PIL import Image
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
# ImageNet normalization constants
|
||||||
|
IMAGENET_MEAN = np.array([0.485, 0.456, 0.406], dtype=np.float32)
|
||||||
|
IMAGENET_STD = np.array([0.229, 0.224, 0.225], dtype=np.float32)
|
||||||
|
|
||||||
|
|
||||||
|
def load_image(path: str | Path, img_size: int | None = None) -> np.ndarray:
|
||||||
|
"""Load image as RGB float32 array in [0, 1] range, shape (H, W, 3).
|
||||||
|
|
||||||
|
Args:
|
||||||
|
path: Path to image file.
|
||||||
|
img_size: If provided, resize to (img_size, img_size) using bicubic interpolation.
|
||||||
|
"""
|
||||||
|
img = Image.open(path).convert("RGB")
|
||||||
|
if img_size is not None:
|
||||||
|
img = img.resize((img_size, img_size), Image.BICUBIC)
|
||||||
|
return np.asarray(img, dtype=np.float32) / 255.0
|
||||||
|
|
||||||
|
|
||||||
|
def load_alpha_hint(path: str | Path, img_size: int | None = None) -> np.ndarray:
|
||||||
|
"""Load alpha hint as grayscale float32 array in [0, 1], shape (H, W, 1).
|
||||||
|
|
||||||
|
Args:
|
||||||
|
path: Path to alpha hint image file.
|
||||||
|
img_size: If provided, resize to (img_size, img_size) using bicubic interpolation.
|
||||||
|
"""
|
||||||
|
img = Image.open(path).convert("L")
|
||||||
|
if img_size is not None:
|
||||||
|
img = img.resize((img_size, img_size), Image.BICUBIC)
|
||||||
|
return np.asarray(img, dtype=np.float32)[:, :, np.newaxis] / 255.0
|
||||||
|
|
||||||
|
|
||||||
|
def normalize_rgb(rgb: np.ndarray) -> np.ndarray:
|
||||||
|
"""Apply ImageNet normalization to (H, W, 3) float32 RGB in [0, 1]."""
|
||||||
|
return (rgb - IMAGENET_MEAN) / IMAGENET_STD
|
||||||
|
|
||||||
|
|
||||||
|
def preprocess(
|
||||||
|
rgb: np.ndarray,
|
||||||
|
alpha_hint: np.ndarray,
|
||||||
|
) -> mx.array:
|
||||||
|
"""Build 4-channel NHWC input tensor from RGB and alpha hint.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
rgb: (H, W, 3) float32 in [0, 1]
|
||||||
|
alpha_hint: (H, W, 1) float32 in [0, 1]
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
(1, H, W, 4) mx.array — ImageNet-normalized RGB + raw alpha hint.
|
||||||
|
"""
|
||||||
|
normalized_rgb = normalize_rgb(rgb)
|
||||||
|
combined = np.concatenate([normalized_rgb, alpha_hint], axis=-1) # (H, W, 4)
|
||||||
|
return mx.array(combined[np.newaxis]) # (1, H, W, 4)
|
||||||
|
|
||||||
|
|
||||||
|
def postprocess_alpha(alpha: mx.array) -> np.ndarray:
|
||||||
|
"""Convert model alpha output to uint8 numpy array.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
alpha: (1, H, W, 1) probabilities in [0, 1]
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
(H, W) uint8 array.
|
||||||
|
"""
|
||||||
|
arr = np.array(alpha[0, :, :, 0])
|
||||||
|
return (np.clip(arr, 0.0, 1.0) * 255.0).astype(np.uint8)
|
||||||
|
|
||||||
|
|
||||||
|
def postprocess_foreground(fg: mx.array) -> np.ndarray:
|
||||||
|
"""Convert model foreground output to uint8 numpy array.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
fg: (1, H, W, 3) probabilities in [0, 1]
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
(H, W, 3) uint8 array.
|
||||||
|
"""
|
||||||
|
arr = np.array(fg[0])
|
||||||
|
return (np.clip(arr, 0.0, 1.0) * 255.0).astype(np.uint8)
|
||||||
|
|
||||||
|
|
||||||
|
def save_alpha(alpha: np.ndarray, path: str | Path) -> None:
|
||||||
|
"""Save alpha matte as grayscale PNG."""
|
||||||
|
Image.fromarray(alpha, mode="L").save(path)
|
||||||
|
|
||||||
|
|
||||||
|
def save_foreground(fg: np.ndarray, path: str | Path) -> None:
|
||||||
|
"""Save foreground as RGB PNG."""
|
||||||
|
Image.fromarray(fg, mode="RGB").save(path)
|
||||||
|
|||||||
@ -1,4 +1,128 @@
|
|||||||
"""Top-level CorridorKey model — MLX port (not yet implemented).
|
"""Top-level CorridorKey model (GreenFormer) — MLX port.
|
||||||
|
|
||||||
Composes backbone + decoder heads + refiner into full pipeline.
|
Composes Hiera backbone + dual decoder heads + CNN refiner.
|
||||||
|
All internal operations use NHWC layout.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from typing import TYPE_CHECKING
|
||||||
|
|
||||||
|
import mlx.core as mx
|
||||||
|
import mlx.nn as nn
|
||||||
|
from safetensors import safe_open
|
||||||
|
|
||||||
|
from corridorkey_mlx.model.backbone import HieraBackbone
|
||||||
|
from corridorkey_mlx.model.decoder import DecoderHead
|
||||||
|
from corridorkey_mlx.model.hiera import ENCODER_KEY_PREFIX, _interpolate_pos_embed, _prod
|
||||||
|
from corridorkey_mlx.model.refiner import CNNRefinerModule
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
BACKBONE_CHANNELS = [112, 224, 448, 896]
|
||||||
|
EMBED_DIM = 256
|
||||||
|
|
||||||
|
|
||||||
|
class GreenFormer(nn.Module):
|
||||||
|
"""CorridorKey: Hiera encoder + dual decoder heads + CNN refiner.
|
||||||
|
|
||||||
|
Input: (B, H, W, 4) NHWC — ImageNet-normalized RGB + alpha hint [0,1]
|
||||||
|
Output: dict with coarse/final alpha and foreground maps in NHWC.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self, img_size: int = 512) -> None:
|
||||||
|
super().__init__()
|
||||||
|
self.backbone = HieraBackbone(img_size=img_size)
|
||||||
|
self.alpha_decoder = DecoderHead(BACKBONE_CHANNELS, EMBED_DIM, output_dim=1)
|
||||||
|
self.fg_decoder = DecoderHead(BACKBONE_CHANNELS, EMBED_DIM, output_dim=3)
|
||||||
|
self.refiner = CNNRefinerModule()
|
||||||
|
|
||||||
|
def __call__(self, x: mx.array) -> dict[str, mx.array]:
|
||||||
|
"""Forward pass.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
x: (B, H, W, 4) NHWC — ImageNet-normalized RGB + alpha hint.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Dict with keys: alpha_logits, fg_logits, alpha_logits_up, fg_logits_up,
|
||||||
|
alpha_coarse, fg_coarse, delta_logits, alpha_final, fg_final.
|
||||||
|
All tensors in NHWC format.
|
||||||
|
"""
|
||||||
|
input_h, input_w = x.shape[1], x.shape[2]
|
||||||
|
|
||||||
|
# Backbone -> 4 multiscale feature maps in NHWC
|
||||||
|
features = self.backbone(x)
|
||||||
|
|
||||||
|
# Decoder heads -> logits at H/4 resolution
|
||||||
|
alpha_logits = self.alpha_decoder(features) # (B, H/4, W/4, 1)
|
||||||
|
fg_logits = self.fg_decoder(features) # (B, H/4, W/4, 3)
|
||||||
|
|
||||||
|
# Upsample logits to full input resolution
|
||||||
|
scale_h = input_h / alpha_logits.shape[1]
|
||||||
|
scale_w = input_w / alpha_logits.shape[2]
|
||||||
|
upsampler = nn.Upsample(
|
||||||
|
scale_factor=(scale_h, scale_w),
|
||||||
|
mode="linear",
|
||||||
|
align_corners=False,
|
||||||
|
)
|
||||||
|
alpha_logits_up = upsampler(alpha_logits) # (B, H, W, 1)
|
||||||
|
fg_logits_up = upsampler(fg_logits) # (B, H, W, 3)
|
||||||
|
|
||||||
|
# Coarse predictions via sigmoid
|
||||||
|
alpha_coarse = mx.sigmoid(alpha_logits_up)
|
||||||
|
fg_coarse = mx.sigmoid(fg_logits_up)
|
||||||
|
|
||||||
|
# Refiner: RGB + coarse predictions -> delta logits
|
||||||
|
rgb = x[:, :, :, :3] # (B, H, W, 3)
|
||||||
|
coarse_pred = mx.concatenate([alpha_coarse, fg_coarse], axis=-1) # (B, H, W, 4)
|
||||||
|
delta_logits = self.refiner(rgb, coarse_pred) # (B, H, W, 4)
|
||||||
|
|
||||||
|
# Final predictions: additive residual in logit space, then sigmoid
|
||||||
|
alpha_final = mx.sigmoid(alpha_logits_up + delta_logits[:, :, :, 0:1])
|
||||||
|
fg_final = mx.sigmoid(fg_logits_up + delta_logits[:, :, :, 1:4])
|
||||||
|
|
||||||
|
return {
|
||||||
|
"alpha_logits": alpha_logits,
|
||||||
|
"fg_logits": fg_logits,
|
||||||
|
"alpha_logits_up": alpha_logits_up,
|
||||||
|
"fg_logits_up": fg_logits_up,
|
||||||
|
"alpha_coarse": alpha_coarse,
|
||||||
|
"fg_coarse": fg_coarse,
|
||||||
|
"delta_logits": delta_logits,
|
||||||
|
"alpha_final": alpha_final,
|
||||||
|
"fg_final": fg_final,
|
||||||
|
}
|
||||||
|
|
||||||
|
def load_checkpoint(self, path: str | Path) -> None:
|
||||||
|
"""Load all weights from converted safetensors checkpoint.
|
||||||
|
|
||||||
|
Handles:
|
||||||
|
- Backbone keys (encoder.model.* prefix) with pos_embed interpolation
|
||||||
|
- Decoder keys (alpha_decoder.*, fg_decoder.*)
|
||||||
|
- Refiner keys (refiner.*)
|
||||||
|
"""
|
||||||
|
target_tokens = _prod(self.backbone.tokens_spatial_shape)
|
||||||
|
weight_pairs: list[tuple[str, mx.array]] = []
|
||||||
|
|
||||||
|
with safe_open(str(path), framework="numpy") as f:
|
||||||
|
for full_key in f.keys(): # noqa: SIM118
|
||||||
|
tensor = mx.array(f.get_tensor(full_key))
|
||||||
|
|
||||||
|
if full_key.startswith(ENCODER_KEY_PREFIX):
|
||||||
|
# Backbone: strip encoder.model. prefix
|
||||||
|
mlx_key = "backbone." + full_key[len(ENCODER_KEY_PREFIX) :]
|
||||||
|
if mlx_key == "backbone.pos_embed":
|
||||||
|
tensor = _interpolate_pos_embed(tensor, target_tokens)
|
||||||
|
# materialize interpolated embedding
|
||||||
|
mx.eval(tensor)
|
||||||
|
else:
|
||||||
|
# Decoder/refiner: use key as-is
|
||||||
|
mlx_key = full_key
|
||||||
|
|
||||||
|
weight_pairs.append((mlx_key, tensor))
|
||||||
|
|
||||||
|
self.load_weights(weight_pairs)
|
||||||
|
self.eval()
|
||||||
|
# materialize all parameters
|
||||||
|
mx.eval(self.parameters())
|
||||||
|
|||||||
@ -26,9 +26,7 @@ SAFETENSORS_PATH = Path("checkpoints/corridorkey_mlx.safetensors")
|
|||||||
|
|
||||||
|
|
||||||
@pytest.fixture(scope="module")
|
@pytest.fixture(scope="module")
|
||||||
def checkpoint_data() -> (
|
def checkpoint_data() -> tuple[dict[str, np.ndarray], OrderedDict[str, np.ndarray], list]:
|
||||||
tuple[dict[str, np.ndarray], OrderedDict[str, np.ndarray], list]
|
|
||||||
):
|
|
||||||
"""Load and convert checkpoint once for all tests in this module."""
|
"""Load and convert checkpoint once for all tests in this module."""
|
||||||
if not CHECKPOINT_PATH.exists():
|
if not CHECKPOINT_PATH.exists():
|
||||||
pytest.skip("Checkpoint not found — need CorridorKey_v1.0.pth")
|
pytest.skip("Checkpoint not found — need CorridorKey_v1.0.pth")
|
||||||
|
|||||||
94
tests/test_end_to_end_parity.py
Normal file
94
tests/test_end_to_end_parity.py
Normal file
@ -0,0 +1,94 @@
|
|||||||
|
"""End-to-end parity tests: MLX GreenFormer vs PyTorch golden reference (Phase 5).
|
||||||
|
|
||||||
|
Loads golden.npz (PyTorch NCHW), runs same input through MLX model,
|
||||||
|
compares all intermediate and final outputs.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
import mlx.core as mx
|
||||||
|
import numpy as np
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from corridorkey_mlx.model.corridorkey import GreenFormer
|
||||||
|
from corridorkey_mlx.utils.layout import nchw_to_nhwc_np, nhwc_to_nchw_np
|
||||||
|
|
||||||
|
FIXTURE_PATH = Path("reference/fixtures/golden.npz")
|
||||||
|
CHECKPOINT_PATH = Path("checkpoints/corridorkey_mlx.safetensors")
|
||||||
|
IMG_SIZE = 512
|
||||||
|
|
||||||
|
# Tolerance tiers — coarse path inherits backbone drift, refiner adds more.
|
||||||
|
# Backbone stages have up to ~0.01 max abs err (16 sequential blocks).
|
||||||
|
# Final outputs pass through sigmoid which compresses errors.
|
||||||
|
TOLERANCES: dict[str, float] = {
|
||||||
|
"alpha_logits": 1e-3,
|
||||||
|
"fg_logits": 1e-3,
|
||||||
|
"alpha_logits_up": 1e-3,
|
||||||
|
"fg_logits_up": 1e-3,
|
||||||
|
"alpha_coarse": 1e-4,
|
||||||
|
"fg_coarse": 1e-4,
|
||||||
|
"delta_logits": 1e-3,
|
||||||
|
"alpha_final": 1e-4,
|
||||||
|
"fg_final": 1e-4,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def _skip_if_missing() -> None:
|
||||||
|
if not FIXTURE_PATH.exists():
|
||||||
|
pytest.skip("golden.npz not found — run dump_pytorch_reference.py first")
|
||||||
|
if not CHECKPOINT_PATH.exists():
|
||||||
|
pytest.skip("Checkpoint not found — run scripts/convert_weights.py first")
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture(scope="module")
|
||||||
|
def model_outputs_and_fixtures() -> tuple[dict[str, mx.array], dict[str, np.ndarray]]:
|
||||||
|
"""Load model, run forward pass with golden input, return (mlx_outputs, fixtures)."""
|
||||||
|
_skip_if_missing()
|
||||||
|
|
||||||
|
fixtures = dict(np.load(FIXTURE_PATH))
|
||||||
|
|
||||||
|
# Load input: NCHW -> NHWC
|
||||||
|
input_nhwc = mx.array(nchw_to_nhwc_np(fixtures["input"]))
|
||||||
|
|
||||||
|
model = GreenFormer(img_size=IMG_SIZE)
|
||||||
|
model.load_checkpoint(CHECKPOINT_PATH)
|
||||||
|
|
||||||
|
outputs = model(input_nhwc)
|
||||||
|
# materialize — mx.eval is MLX lazy graph evaluation, not Python eval
|
||||||
|
mx.eval(outputs) # noqa: S307
|
||||||
|
|
||||||
|
return outputs, fixtures
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize("key", list(TOLERANCES.keys()))
|
||||||
|
def test_parity(
|
||||||
|
key: str,
|
||||||
|
model_outputs_and_fixtures: tuple[dict[str, mx.array], dict[str, np.ndarray]],
|
||||||
|
) -> None:
|
||||||
|
"""MLX output matches PyTorch golden reference within tolerance."""
|
||||||
|
outputs, fixtures = model_outputs_and_fixtures
|
||||||
|
|
||||||
|
assert key in outputs, f"Missing output key: {key}"
|
||||||
|
assert key in fixtures, f"Missing fixture key: {key}"
|
||||||
|
|
||||||
|
# MLX output is NHWC, fixture is NCHW — convert MLX to NCHW for comparison
|
||||||
|
mlx_nchw = nhwc_to_nchw_np(np.array(outputs[key]))
|
||||||
|
expected_nchw = fixtures[key]
|
||||||
|
|
||||||
|
assert mlx_nchw.shape == expected_nchw.shape, (
|
||||||
|
f"{key} shape mismatch: MLX {mlx_nchw.shape} vs PyTorch {expected_nchw.shape}"
|
||||||
|
)
|
||||||
|
|
||||||
|
abs_err = np.abs(mlx_nchw - expected_nchw)
|
||||||
|
max_abs = float(np.max(abs_err))
|
||||||
|
mean_abs = float(np.mean(abs_err))
|
||||||
|
tol = TOLERANCES[key]
|
||||||
|
|
||||||
|
print(
|
||||||
|
f"\n{key:<20s} | shape={mlx_nchw.shape} | "
|
||||||
|
f"max_abs={max_abs:.6e} | mean_abs={mean_abs:.6e} | tol={tol:.1e}"
|
||||||
|
)
|
||||||
|
|
||||||
|
assert max_abs < tol, f"{key}: max abs error {max_abs:.6e} exceeds tolerance {tol:.1e}"
|
||||||
66
tests/test_end_to_end_smoke.py
Normal file
66
tests/test_end_to_end_smoke.py
Normal file
@ -0,0 +1,66 @@
|
|||||||
|
"""End-to-end smoke test for GreenFormer (Phase 5).
|
||||||
|
|
||||||
|
Verifies the full pipeline works with random weights — no checkpoint needed.
|
||||||
|
Tests model construction, forward pass, and basic output sanity.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import mlx.core as mx
|
||||||
|
import numpy as np
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from corridorkey_mlx.io.image import (
|
||||||
|
postprocess_alpha,
|
||||||
|
postprocess_foreground,
|
||||||
|
preprocess,
|
||||||
|
)
|
||||||
|
from corridorkey_mlx.model.corridorkey import GreenFormer
|
||||||
|
|
||||||
|
IMG_SIZE = 256 # smaller for speed
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture(scope="module")
|
||||||
|
def model() -> GreenFormer:
|
||||||
|
return GreenFormer(img_size=IMG_SIZE)
|
||||||
|
|
||||||
|
|
||||||
|
def test_forward_with_preprocessed_input(model: GreenFormer) -> None:
|
||||||
|
"""Full pipeline: numpy image -> preprocess -> model -> postprocess."""
|
||||||
|
rgb = np.random.rand(IMG_SIZE, IMG_SIZE, 3).astype(np.float32)
|
||||||
|
alpha_hint = np.random.rand(IMG_SIZE, IMG_SIZE, 1).astype(np.float32)
|
||||||
|
|
||||||
|
x = preprocess(rgb, alpha_hint)
|
||||||
|
assert x.shape == (1, IMG_SIZE, IMG_SIZE, 4)
|
||||||
|
|
||||||
|
out = model(x)
|
||||||
|
# materialize — mx.eval is MLX's lazy graph evaluation, not Python eval
|
||||||
|
mx.eval(out)
|
||||||
|
|
||||||
|
alpha = postprocess_alpha(out["alpha_final"])
|
||||||
|
fg = postprocess_foreground(out["fg_final"])
|
||||||
|
|
||||||
|
assert alpha.shape == (IMG_SIZE, IMG_SIZE)
|
||||||
|
assert alpha.dtype == np.uint8
|
||||||
|
assert fg.shape == (IMG_SIZE, IMG_SIZE, 3)
|
||||||
|
assert fg.dtype == np.uint8
|
||||||
|
|
||||||
|
|
||||||
|
def test_deterministic_output(model: GreenFormer) -> None:
|
||||||
|
"""Same input produces same output."""
|
||||||
|
mx.random.seed(123)
|
||||||
|
x = mx.random.normal((1, IMG_SIZE, IMG_SIZE, 4))
|
||||||
|
# materialize — mx.eval is MLX's lazy graph evaluation, not Python eval
|
||||||
|
mx.eval(x)
|
||||||
|
|
||||||
|
out1 = model(x)
|
||||||
|
mx.eval(out1)
|
||||||
|
out2 = model(x)
|
||||||
|
mx.eval(out2)
|
||||||
|
|
||||||
|
for key in out1:
|
||||||
|
np.testing.assert_array_equal(
|
||||||
|
np.array(out1[key]),
|
||||||
|
np.array(out2[key]),
|
||||||
|
err_msg=f"{key} not deterministic",
|
||||||
|
)
|
||||||
77
tests/test_greenformer_forward.py
Normal file
77
tests/test_greenformer_forward.py
Normal file
@ -0,0 +1,77 @@
|
|||||||
|
"""Forward pass shape and smoke tests for GreenFormer (Phase 5).
|
||||||
|
|
||||||
|
Uses random weights — no checkpoint needed.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import mlx.core as mx
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from corridorkey_mlx.model.corridorkey import GreenFormer
|
||||||
|
|
||||||
|
IMG_SIZE = 512
|
||||||
|
BATCH = 1
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture(scope="module")
|
||||||
|
def model_and_output() -> tuple[GreenFormer, dict[str, mx.array]]:
|
||||||
|
"""Build model with random weights and run forward pass once."""
|
||||||
|
model = GreenFormer(img_size=IMG_SIZE)
|
||||||
|
x = mx.random.normal((BATCH, IMG_SIZE, IMG_SIZE, 4))
|
||||||
|
out = model(x)
|
||||||
|
# materialize all outputs — mx.eval is MLX lazy graph evaluation, not Python eval
|
||||||
|
mx.eval(out)
|
||||||
|
return model, out
|
||||||
|
|
||||||
|
|
||||||
|
EXPECTED_SHAPES: dict[str, tuple[int, ...]] = {
|
||||||
|
"alpha_logits": (BATCH, IMG_SIZE // 4, IMG_SIZE // 4, 1),
|
||||||
|
"fg_logits": (BATCH, IMG_SIZE // 4, IMG_SIZE // 4, 3),
|
||||||
|
"alpha_logits_up": (BATCH, IMG_SIZE, IMG_SIZE, 1),
|
||||||
|
"fg_logits_up": (BATCH, IMG_SIZE, IMG_SIZE, 3),
|
||||||
|
"alpha_coarse": (BATCH, IMG_SIZE, IMG_SIZE, 1),
|
||||||
|
"fg_coarse": (BATCH, IMG_SIZE, IMG_SIZE, 3),
|
||||||
|
"delta_logits": (BATCH, IMG_SIZE, IMG_SIZE, 4),
|
||||||
|
"alpha_final": (BATCH, IMG_SIZE, IMG_SIZE, 1),
|
||||||
|
"fg_final": (BATCH, IMG_SIZE, IMG_SIZE, 3),
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize("key,expected_shape", EXPECTED_SHAPES.items())
|
||||||
|
def test_output_shapes(
|
||||||
|
key: str,
|
||||||
|
expected_shape: tuple[int, ...],
|
||||||
|
model_and_output: tuple[GreenFormer, dict[str, mx.array]],
|
||||||
|
) -> None:
|
||||||
|
_, out = model_and_output
|
||||||
|
assert key in out, f"Missing output key: {key}"
|
||||||
|
assert out[key].shape == expected_shape, (
|
||||||
|
f"{key}: expected {expected_shape}, got {out[key].shape}"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_all_keys_present(
|
||||||
|
model_and_output: tuple[GreenFormer, dict[str, mx.array]],
|
||||||
|
) -> None:
|
||||||
|
_, out = model_and_output
|
||||||
|
assert set(out.keys()) == set(EXPECTED_SHAPES.keys())
|
||||||
|
|
||||||
|
|
||||||
|
def test_coarse_probs_in_range(
|
||||||
|
model_and_output: tuple[GreenFormer, dict[str, mx.array]],
|
||||||
|
) -> None:
|
||||||
|
"""Sigmoid outputs must be in [0, 1]."""
|
||||||
|
_, out = model_and_output
|
||||||
|
for key in ("alpha_coarse", "fg_coarse", "alpha_final", "fg_final"):
|
||||||
|
arr = out[key]
|
||||||
|
assert float(mx.min(arr)) >= 0.0, f"{key} has values < 0"
|
||||||
|
assert float(mx.max(arr)) <= 1.0, f"{key} has values > 1"
|
||||||
|
|
||||||
|
|
||||||
|
def test_output_dtype(
|
||||||
|
model_and_output: tuple[GreenFormer, dict[str, mx.array]],
|
||||||
|
) -> None:
|
||||||
|
_, out = model_and_output
|
||||||
|
for key, arr in out.items():
|
||||||
|
assert arr.dtype == mx.float32, f"{key}: expected float32, got {arr.dtype}"
|
||||||
@ -33,9 +33,7 @@ def _skip_if_missing() -> None:
|
|||||||
|
|
||||||
|
|
||||||
@pytest.fixture(scope="module")
|
@pytest.fixture(scope="module")
|
||||||
def backbone_and_fixtures() -> (
|
def backbone_and_fixtures() -> tuple[list[mx.array], dict[str, np.ndarray]]:
|
||||||
tuple[list[mx.array], dict[str, np.ndarray]]
|
|
||||||
):
|
|
||||||
"""Load backbone once, return (mlx_features, fixtures)."""
|
"""Load backbone once, return (mlx_features, fixtures)."""
|
||||||
_skip_if_missing()
|
_skip_if_missing()
|
||||||
|
|
||||||
@ -74,10 +72,7 @@ def test_stage_parity(
|
|||||||
abs_err = np.abs(result_nchw - expected_nchw)
|
abs_err = np.abs(result_nchw - expected_nchw)
|
||||||
max_abs_err = float(np.max(abs_err))
|
max_abs_err = float(np.max(abs_err))
|
||||||
mean_abs_err = float(np.mean(abs_err))
|
mean_abs_err = float(np.mean(abs_err))
|
||||||
print(
|
print(f"\nStage {stage_idx} parity — max_abs: {max_abs_err:.6e}, mean_abs: {mean_abs_err:.6e}")
|
||||||
f"\nStage {stage_idx} parity — "
|
|
||||||
f"max_abs: {max_abs_err:.6e}, mean_abs: {mean_abs_err:.6e}"
|
|
||||||
)
|
|
||||||
|
|
||||||
assert max_abs_err < MAX_ABS_TOL, (
|
assert max_abs_err < MAX_ABS_TOL, (
|
||||||
f"Stage {stage_idx} max abs error {max_abs_err:.6e} exceeds tolerance {MAX_ABS_TOL}"
|
f"Stage {stage_idx} max abs error {max_abs_err:.6e} exceeds tolerance {MAX_ABS_TOL}"
|
||||||
|
|||||||
Loading…
Reference in New Issue
Block a user