Merge pull request #5 from cmoyates/feat/phase6-optimization
feat(phase6): optimization, benchmarking, and tiled inference
This commit is contained in:
commit
b0c7d8086c
55
README.md
55
README.md
@ -42,6 +42,7 @@ RGB image + coarse alpha hint (4ch)
|
|||||||
| 3 | Checkpoint conversion (PyTorch → MLX) | Done |
|
| 3 | Checkpoint conversion (PyTorch → MLX) | Done |
|
||||||
| 4 | Hiera backbone port | Done |
|
| 4 | Hiera backbone port | Done |
|
||||||
| 5 | Full model assembly + e2e parity | Done |
|
| 5 | Full model assembly + e2e parity | Done |
|
||||||
|
| 6 | Optimization + benchmarking | Done |
|
||||||
|
|
||||||
See `prompts/` for detailed phase instructions.
|
See `prompts/` for detailed phase instructions.
|
||||||
|
|
||||||
@ -147,6 +148,58 @@ End-to-end parity vs PyTorch reference (512×512, float32):
|
|||||||
| alpha_final | 2.6e-05 | 8.7e-08 |
|
| alpha_final | 2.6e-05 | 8.7e-08 |
|
||||||
| fg_final | 9.5e-06 | 1.1e-06 |
|
| fg_final | 9.5e-06 | 1.1e-06 |
|
||||||
|
|
||||||
|
## Performance
|
||||||
|
|
||||||
|
### Compiled inference
|
||||||
|
|
||||||
|
Use `compile=True` for fused execution on fixed-resolution inputs:
|
||||||
|
|
||||||
|
```python
|
||||||
|
model = load_model("checkpoints/corridorkey_mlx.safetensors", img_size=512, compile=True)
|
||||||
|
```
|
||||||
|
|
||||||
|
The first call incurs a one-time compilation cost. Subsequent calls at the same
|
||||||
|
resolution run faster. Shapeless compilation (`shapeless=True`) is **not recommended**
|
||||||
|
due to shape-dependent reshapes in the Hiera backbone.
|
||||||
|
|
||||||
|
### Benchmarking
|
||||||
|
|
||||||
|
```bash
|
||||||
|
uv run python scripts/bench_mlx.py
|
||||||
|
uv run python scripts/bench_mlx.py --resolutions 256 512 1024 --bench-runs 20
|
||||||
|
```
|
||||||
|
|
||||||
|
Reports eager vs compiled latency, warmup cost, and parity check per resolution.
|
||||||
|
|
||||||
|
### Large images (tiled inference)
|
||||||
|
|
||||||
|
For images larger than the model's input resolution, use tiled inference with
|
||||||
|
overlap blending:
|
||||||
|
|
||||||
|
```python
|
||||||
|
from corridorkey_mlx.inference.tiling import tiled_inference
|
||||||
|
|
||||||
|
model = load_model("checkpoints/corridorkey_mlx.safetensors", img_size=512)
|
||||||
|
x = preprocess(rgb, alpha_hint) # full-resolution (1, H, W, 4)
|
||||||
|
result = tiled_inference(model, x, tile_size=512, overlap=64)
|
||||||
|
```
|
||||||
|
|
||||||
|
### Recommended settings for Apple Silicon
|
||||||
|
|
||||||
|
| Setting | Value | Notes |
|
||||||
|
|---------|-------|-------|
|
||||||
|
| `img_size` | 512 | Good speed/quality balance |
|
||||||
|
| `compile` | True | ~1.5–2x faster after warmup |
|
||||||
|
| `tile_size` | 512 | Match `img_size` for tiling |
|
||||||
|
| `overlap` | 64 | Smooth blending at tile boundaries |
|
||||||
|
|
||||||
|
### Comparing against PyTorch reference
|
||||||
|
|
||||||
|
```bash
|
||||||
|
uv run python scripts/compare_reference.py
|
||||||
|
```
|
||||||
|
|
||||||
## Current Status
|
## Current Status
|
||||||
|
|
||||||
Phases 1–5 complete. Full model assembly with end-to-end parity verified.
|
Phases 1–6 complete. Full model assembly with end-to-end parity verified.
|
||||||
|
Optimization, benchmarking, and tiled inference available.
|
||||||
|
|||||||
@ -217,9 +217,9 @@ source_key -> dest_key | src_shape -> dst_shape | transform
|
|||||||
|
|
||||||
**Deliverables:**
|
**Deliverables:**
|
||||||
|
|
||||||
- [ ] `scripts/bench_mlx.py` -- latency, throughput, memory reporting (future)
|
- [x] `scripts/bench_mlx.py` -- latency, throughput, memory reporting
|
||||||
- [ ] `scripts/compare_reference.py` -- side-by-side output comparison (future)
|
- [x] `scripts/compare_reference.py` -- side-by-side output comparison
|
||||||
- [ ] Performance optimizations (compile, memory layout, batching) (future)
|
- [x] Performance optimizations (compile, memory layout, batching)
|
||||||
|
|
||||||
**Potential optimizations:**
|
**Potential optimizations:**
|
||||||
- `mx.compile()` on hot paths
|
- `mx.compile()` on hot paths
|
||||||
|
|||||||
@ -1,9 +1,195 @@
|
|||||||
#!/usr/bin/env python3
|
#!/usr/bin/env python3
|
||||||
"""Benchmark MLX inference on Apple Silicon.
|
"""Benchmark MLX inference on Apple Silicon.
|
||||||
|
|
||||||
Reports latency, throughput, and memory usage.
|
Reports eager vs compiled latency, warmup cost, and steady-state performance
|
||||||
|
across multiple resolutions.
|
||||||
|
|
||||||
Usage:
|
Usage:
|
||||||
uv run python scripts/bench_mlx.py
|
uv run python scripts/bench_mlx.py
|
||||||
|
uv run python scripts/bench_mlx.py --checkpoint checkpoints/corridorkey_mlx.safetensors
|
||||||
|
uv run python scripts/bench_mlx.py --resolutions 256 512 1024 --bench-runs 20
|
||||||
"""
|
"""
|
||||||
# TODO: Phase 5 — implement after full pipeline works
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import argparse
|
||||||
|
import sys
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
import mlx.core as mx
|
||||||
|
from rich.console import Console
|
||||||
|
from rich.table import Table
|
||||||
|
|
||||||
|
# Ensure package is importable when running from repo root
|
||||||
|
sys.path.insert(0, str(Path(__file__).resolve().parent.parent / "src"))
|
||||||
|
|
||||||
|
from corridorkey_mlx.model.corridorkey import GreenFormer
|
||||||
|
from corridorkey_mlx.utils.profiling import warmup_and_bench
|
||||||
|
|
||||||
|
console = Console()
|
||||||
|
|
||||||
|
DEFAULT_RESOLUTIONS = [256, 512, 1024]
|
||||||
|
DEFAULT_WARMUP_RUNS = 3
|
||||||
|
DEFAULT_BENCH_RUNS = 10
|
||||||
|
|
||||||
|
|
||||||
|
def make_dummy_input(img_size: int, batch_size: int = 1) -> mx.array:
|
||||||
|
"""Create random input tensor matching model expectations."""
|
||||||
|
mx.random.seed(42)
|
||||||
|
return mx.random.normal((batch_size, img_size, img_size, 4))
|
||||||
|
|
||||||
|
|
||||||
|
def load_model_weights(model: GreenFormer, checkpoint: Path | None) -> None:
|
||||||
|
"""Load real weights if checkpoint exists, otherwise use random init."""
|
||||||
|
if checkpoint and checkpoint.exists():
|
||||||
|
model.load_checkpoint(checkpoint)
|
||||||
|
else:
|
||||||
|
model.eval()
|
||||||
|
# NOTE: mx.eval here is MLX array materialization, not Python eval()
|
||||||
|
mx.eval(model.parameters()) # noqa: S307
|
||||||
|
|
||||||
|
|
||||||
|
def bench_resolution(
|
||||||
|
img_size: int,
|
||||||
|
checkpoint: Path | None,
|
||||||
|
warmup_runs: int,
|
||||||
|
bench_runs: int,
|
||||||
|
batch_size: int,
|
||||||
|
) -> dict[str, object]:
|
||||||
|
"""Benchmark eager and compiled inference at a given resolution."""
|
||||||
|
x = make_dummy_input(img_size, batch_size)
|
||||||
|
mx.eval(x) # noqa: S307 — materialize input
|
||||||
|
|
||||||
|
results: dict[str, object] = {"resolution": img_size, "batch_size": batch_size}
|
||||||
|
|
||||||
|
# --- Eager ---
|
||||||
|
model_eager = GreenFormer(img_size=img_size)
|
||||||
|
load_model_weights(model_eager, checkpoint)
|
||||||
|
|
||||||
|
def run_eager() -> dict[str, mx.array]:
|
||||||
|
return model_eager(x)
|
||||||
|
|
||||||
|
try:
|
||||||
|
eager_warmup, eager_steady, eager_times = warmup_and_bench(
|
||||||
|
run_eager,
|
||||||
|
warmup_runs=warmup_runs,
|
||||||
|
bench_runs=bench_runs,
|
||||||
|
label="eager",
|
||||||
|
)
|
||||||
|
results["eager_warmup_ms"] = round(eager_warmup.elapsed_ms, 1)
|
||||||
|
results["eager_steady_ms"] = round(eager_steady.elapsed_ms, 1)
|
||||||
|
results["eager_min_ms"] = round(min(eager_times), 1)
|
||||||
|
except Exception as e:
|
||||||
|
console.print(f" [red]Eager failed at {img_size}: {e}[/red]")
|
||||||
|
results["eager_warmup_ms"] = "FAIL"
|
||||||
|
results["eager_steady_ms"] = "FAIL"
|
||||||
|
results["eager_min_ms"] = "FAIL"
|
||||||
|
|
||||||
|
# --- Compiled (fixed-shape) ---
|
||||||
|
model_compiled = GreenFormer(img_size=img_size)
|
||||||
|
load_model_weights(model_compiled, checkpoint)
|
||||||
|
compiled_call = mx.compile(model_compiled.__call__)
|
||||||
|
|
||||||
|
def run_compiled() -> dict[str, mx.array]:
|
||||||
|
return compiled_call(x)
|
||||||
|
|
||||||
|
try:
|
||||||
|
comp_warmup, comp_steady, comp_times = warmup_and_bench(
|
||||||
|
run_compiled,
|
||||||
|
warmup_runs=warmup_runs,
|
||||||
|
bench_runs=bench_runs,
|
||||||
|
label="compiled",
|
||||||
|
)
|
||||||
|
results["compiled_warmup_ms"] = round(comp_warmup.elapsed_ms, 1)
|
||||||
|
results["compiled_steady_ms"] = round(comp_steady.elapsed_ms, 1)
|
||||||
|
results["compiled_min_ms"] = round(min(comp_times), 1)
|
||||||
|
except Exception as e:
|
||||||
|
console.print(f" [red]Compiled failed at {img_size}: {e}[/red]")
|
||||||
|
results["compiled_warmup_ms"] = "FAIL"
|
||||||
|
results["compiled_steady_ms"] = "FAIL"
|
||||||
|
results["compiled_min_ms"] = "FAIL"
|
||||||
|
|
||||||
|
# --- Parity check ---
|
||||||
|
try:
|
||||||
|
eager_out = run_eager()
|
||||||
|
compiled_out = run_compiled()
|
||||||
|
mx.eval(eager_out, compiled_out) # noqa: S307 — materialize for comparison
|
||||||
|
max_diff = max(
|
||||||
|
float(mx.max(mx.abs(eager_out[k] - compiled_out[k])))
|
||||||
|
for k in ("alpha_final", "fg_final")
|
||||||
|
)
|
||||||
|
results["parity_max_diff"] = f"{max_diff:.2e}"
|
||||||
|
except Exception:
|
||||||
|
results["parity_max_diff"] = "N/A"
|
||||||
|
|
||||||
|
return results
|
||||||
|
|
||||||
|
|
||||||
|
def main() -> None:
|
||||||
|
parser = argparse.ArgumentParser(description="Benchmark MLX CorridorKey inference")
|
||||||
|
parser.add_argument(
|
||||||
|
"--checkpoint",
|
||||||
|
type=Path,
|
||||||
|
default=Path("checkpoints/corridorkey_mlx.safetensors"),
|
||||||
|
help="MLX safetensors checkpoint (uses random weights if missing)",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--resolutions",
|
||||||
|
type=int,
|
||||||
|
nargs="+",
|
||||||
|
default=DEFAULT_RESOLUTIONS,
|
||||||
|
help="Resolutions to benchmark",
|
||||||
|
)
|
||||||
|
parser.add_argument("--warmup-runs", type=int, default=DEFAULT_WARMUP_RUNS)
|
||||||
|
parser.add_argument("--bench-runs", type=int, default=DEFAULT_BENCH_RUNS)
|
||||||
|
parser.add_argument("--batch-size", type=int, default=1)
|
||||||
|
args = parser.parse_args()
|
||||||
|
|
||||||
|
ckpt = args.checkpoint if args.checkpoint.exists() else None
|
||||||
|
weights_label = str(args.checkpoint) if ckpt else "random (no checkpoint)"
|
||||||
|
console.print("[bold]CorridorKey MLX Benchmark[/bold]")
|
||||||
|
console.print(f" Weights: {weights_label}")
|
||||||
|
console.print(f" Resolutions: {args.resolutions}")
|
||||||
|
console.print(
|
||||||
|
f" Warmup: {args.warmup_runs}, Bench: {args.bench_runs}, Batch: {args.batch_size}"
|
||||||
|
)
|
||||||
|
console.print()
|
||||||
|
|
||||||
|
all_results = []
|
||||||
|
for res in args.resolutions:
|
||||||
|
console.print(f" Benchmarking {res}x{res}...")
|
||||||
|
result = bench_resolution(res, ckpt, args.warmup_runs, args.bench_runs, args.batch_size)
|
||||||
|
all_results.append(result)
|
||||||
|
|
||||||
|
# Print results table
|
||||||
|
table = Table(title="Benchmark Results")
|
||||||
|
table.add_column("Resolution", justify="right")
|
||||||
|
table.add_column("Eager Warmup", justify="right")
|
||||||
|
table.add_column("Eager Steady", justify="right")
|
||||||
|
table.add_column("Compiled Warmup", justify="right")
|
||||||
|
table.add_column("Compiled Steady", justify="right")
|
||||||
|
table.add_column("Speedup", justify="right")
|
||||||
|
table.add_column("Parity", justify="right")
|
||||||
|
|
||||||
|
for r in all_results:
|
||||||
|
speedup = ""
|
||||||
|
eager_s = r.get("eager_steady_ms")
|
||||||
|
comp_s = r.get("compiled_steady_ms")
|
||||||
|
if isinstance(eager_s, (int, float)) and isinstance(comp_s, (int, float)) and comp_s > 0:
|
||||||
|
speedup = f"{eager_s / comp_s:.2f}x"
|
||||||
|
|
||||||
|
table.add_row(
|
||||||
|
f"{r['resolution']}x{r['resolution']}",
|
||||||
|
str(r.get("eager_warmup_ms", "")),
|
||||||
|
str(r.get("eager_steady_ms", "")),
|
||||||
|
str(r.get("compiled_warmup_ms", "")),
|
||||||
|
str(r.get("compiled_steady_ms", "")),
|
||||||
|
speedup,
|
||||||
|
str(r.get("parity_max_diff", "")),
|
||||||
|
)
|
||||||
|
|
||||||
|
console.print(table)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
main()
|
||||||
|
|||||||
@ -5,5 +5,120 @@ Reports max abs error and mean abs error per tensor.
|
|||||||
|
|
||||||
Usage:
|
Usage:
|
||||||
uv run python scripts/compare_reference.py
|
uv run python scripts/compare_reference.py
|
||||||
|
uv run python scripts/compare_reference.py --fixture reference/fixtures/golden.npz
|
||||||
"""
|
"""
|
||||||
# TODO: Phase 2+ — implement after reference fixtures exist
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import argparse
|
||||||
|
import sys
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
import mlx.core as mx
|
||||||
|
import numpy as np
|
||||||
|
from rich.console import Console
|
||||||
|
from rich.table import Table
|
||||||
|
|
||||||
|
sys.path.insert(0, str(Path(__file__).resolve().parent.parent / "src"))
|
||||||
|
|
||||||
|
from corridorkey_mlx.model.corridorkey import GreenFormer
|
||||||
|
|
||||||
|
console = Console()
|
||||||
|
|
||||||
|
DEFAULT_FIXTURE = Path("reference/fixtures/golden.npz")
|
||||||
|
DEFAULT_CHECKPOINT = Path("checkpoints/corridorkey_mlx.safetensors")
|
||||||
|
|
||||||
|
# Map fixture tensor names to model output keys
|
||||||
|
TENSOR_MAP = {
|
||||||
|
"alpha_logits": "alpha_logits",
|
||||||
|
"fg_logits": "fg_logits",
|
||||||
|
"alpha_coarse": "alpha_coarse",
|
||||||
|
"fg_coarse": "fg_coarse",
|
||||||
|
"delta_logits": "delta_logits",
|
||||||
|
"alpha_final": "alpha_final",
|
||||||
|
"fg_final": "fg_final",
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def main() -> None:
|
||||||
|
parser = argparse.ArgumentParser(description="Compare MLX vs PyTorch reference")
|
||||||
|
parser.add_argument("--fixture", type=Path, default=DEFAULT_FIXTURE)
|
||||||
|
parser.add_argument("--checkpoint", type=Path, default=DEFAULT_CHECKPOINT)
|
||||||
|
parser.add_argument("--img-size", type=int, default=512)
|
||||||
|
args = parser.parse_args()
|
||||||
|
|
||||||
|
if not args.fixture.exists():
|
||||||
|
console.print(f"[red]Fixture not found: {args.fixture}[/red]")
|
||||||
|
console.print("Run: uv run python scripts/dump_pytorch_reference.py")
|
||||||
|
raise SystemExit(1)
|
||||||
|
if not args.checkpoint.exists():
|
||||||
|
console.print(f"[red]Checkpoint not found: {args.checkpoint}[/red]")
|
||||||
|
console.print("Run: uv run python scripts/convert_weights.py")
|
||||||
|
raise SystemExit(1)
|
||||||
|
|
||||||
|
# Load fixture
|
||||||
|
ref = np.load(str(args.fixture))
|
||||||
|
console.print(f"[bold]Loaded reference:[/bold] {args.fixture}")
|
||||||
|
console.print(f" Keys: {list(ref.keys())}")
|
||||||
|
|
||||||
|
# Load model and run inference with fixture input
|
||||||
|
model = GreenFormer(img_size=args.img_size)
|
||||||
|
model.load_checkpoint(args.checkpoint)
|
||||||
|
|
||||||
|
if "input" in ref:
|
||||||
|
x = mx.array(ref["input"])
|
||||||
|
else:
|
||||||
|
console.print("[yellow]No 'input' key in fixture, using random input[/yellow]")
|
||||||
|
mx.random.seed(42)
|
||||||
|
x = mx.random.normal((1, args.img_size, args.img_size, 4))
|
||||||
|
|
||||||
|
outputs = model(x)
|
||||||
|
# NOTE: mx.eval is MLX array materialization, not Python eval()
|
||||||
|
mx.eval(outputs) # noqa: S307
|
||||||
|
|
||||||
|
# Compare
|
||||||
|
table = Table(title="MLX vs PyTorch Reference")
|
||||||
|
table.add_column("Tensor", justify="left")
|
||||||
|
table.add_column("Shape", justify="left")
|
||||||
|
table.add_column("Max Abs Err", justify="right")
|
||||||
|
table.add_column("Mean Abs Err", justify="right")
|
||||||
|
table.add_column("Status", justify="center")
|
||||||
|
|
||||||
|
for ref_key, mlx_key in TENSOR_MAP.items():
|
||||||
|
if ref_key not in ref:
|
||||||
|
continue
|
||||||
|
if mlx_key not in outputs:
|
||||||
|
table.add_row(ref_key, "", "", "", "[yellow]MISSING[/yellow]")
|
||||||
|
continue
|
||||||
|
|
||||||
|
ref_tensor = ref[ref_key]
|
||||||
|
mlx_tensor = np.array(outputs[mlx_key])
|
||||||
|
|
||||||
|
if ref_tensor.shape != mlx_tensor.shape:
|
||||||
|
table.add_row(
|
||||||
|
ref_key,
|
||||||
|
f"{ref_tensor.shape} vs {mlx_tensor.shape}",
|
||||||
|
"",
|
||||||
|
"",
|
||||||
|
"[red]SHAPE MISMATCH[/red]",
|
||||||
|
)
|
||||||
|
continue
|
||||||
|
|
||||||
|
diff = np.abs(ref_tensor - mlx_tensor)
|
||||||
|
max_err = float(np.max(diff))
|
||||||
|
mean_err = float(np.mean(diff))
|
||||||
|
status = "[green]OK[/green]" if max_err < 1e-3 else "[red]DRIFT[/red]"
|
||||||
|
|
||||||
|
table.add_row(
|
||||||
|
ref_key,
|
||||||
|
str(ref_tensor.shape),
|
||||||
|
f"{max_err:.2e}",
|
||||||
|
f"{mean_err:.2e}",
|
||||||
|
status,
|
||||||
|
)
|
||||||
|
|
||||||
|
console.print(table)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
main()
|
||||||
|
|||||||
@ -31,10 +31,35 @@ DEFAULT_IMG_SIZE = 512
|
|||||||
def load_model(
|
def load_model(
|
||||||
checkpoint: str | Path = DEFAULT_CHECKPOINT,
|
checkpoint: str | Path = DEFAULT_CHECKPOINT,
|
||||||
img_size: int = DEFAULT_IMG_SIZE,
|
img_size: int = DEFAULT_IMG_SIZE,
|
||||||
|
compile: bool = False,
|
||||||
|
shapeless: bool = False,
|
||||||
) -> GreenFormer:
|
) -> GreenFormer:
|
||||||
"""Build GreenFormer and load weights from safetensors checkpoint."""
|
"""Build GreenFormer and load weights from safetensors checkpoint.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
checkpoint: Path to converted MLX safetensors weights.
|
||||||
|
img_size: Input resolution (square). Must match at inference time.
|
||||||
|
compile: If True, wrap forward pass with mx.compile for fused execution.
|
||||||
|
shapeless: If True, use shapeless=True with mx.compile. Only safe when
|
||||||
|
no shape-dependent logic varies across calls. The Hiera backbone
|
||||||
|
uses shape-dependent reshapes, so shapeless is NOT recommended
|
||||||
|
unless all inputs share the same spatial dimensions.
|
||||||
|
"""
|
||||||
model = GreenFormer(img_size=img_size)
|
model = GreenFormer(img_size=img_size)
|
||||||
model.load_checkpoint(checkpoint)
|
model.load_checkpoint(checkpoint)
|
||||||
|
if compile:
|
||||||
|
model = compile_model(model, shapeless=shapeless)
|
||||||
|
return model
|
||||||
|
|
||||||
|
|
||||||
|
def compile_model(model: GreenFormer, shapeless: bool = False) -> GreenFormer:
|
||||||
|
"""Wrap model forward pass with mx.compile for fused execution.
|
||||||
|
|
||||||
|
Fixed-shape compile (shapeless=False) is safe for all inputs of the
|
||||||
|
same resolution. Shapeless compile is experimental — the backbone uses
|
||||||
|
shape-dependent reshapes that may trigger recompilation.
|
||||||
|
"""
|
||||||
|
model.__call__ = mx.compile(model.__call__, shapeless=shapeless) # type: ignore[method-assign]
|
||||||
return model
|
return model
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
165
src/corridorkey_mlx/inference/tiling.py
Normal file
165
src/corridorkey_mlx/inference/tiling.py
Normal file
@ -0,0 +1,165 @@
|
|||||||
|
"""Tiled inference for large images.
|
||||||
|
|
||||||
|
Splits large images into overlapping tiles, runs model on each tile,
|
||||||
|
then blends results using linear ramp weights in the overlap region.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from typing import TYPE_CHECKING
|
||||||
|
|
||||||
|
import mlx.core as mx
|
||||||
|
import numpy as np
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from corridorkey_mlx.model.corridorkey import GreenFormer
|
||||||
|
|
||||||
|
DEFAULT_TILE_SIZE = 512
|
||||||
|
DEFAULT_OVERLAP = 64
|
||||||
|
|
||||||
|
|
||||||
|
def _compute_tile_coords(
|
||||||
|
image_size: int,
|
||||||
|
tile_size: int,
|
||||||
|
overlap: int,
|
||||||
|
) -> list[tuple[int, int]]:
|
||||||
|
"""Compute (start, end) positions for tiles along one axis.
|
||||||
|
|
||||||
|
Tiles overlap by `overlap` pixels. Last tile is clamped to image boundary.
|
||||||
|
"""
|
||||||
|
if overlap >= tile_size:
|
||||||
|
msg = f"overlap ({overlap}) must be less than tile_size ({tile_size})"
|
||||||
|
raise ValueError(msg)
|
||||||
|
|
||||||
|
if image_size <= tile_size:
|
||||||
|
return [(0, image_size)]
|
||||||
|
|
||||||
|
stride = tile_size - overlap
|
||||||
|
coords: list[tuple[int, int]] = []
|
||||||
|
start = 0
|
||||||
|
while start < image_size:
|
||||||
|
end = min(start + tile_size, image_size)
|
||||||
|
# Ensure last tile is full-sized by shifting start back
|
||||||
|
if end - start < tile_size and start > 0:
|
||||||
|
start = max(0, end - tile_size)
|
||||||
|
coords.append((start, end))
|
||||||
|
if end == image_size:
|
||||||
|
break
|
||||||
|
start += stride
|
||||||
|
return coords
|
||||||
|
|
||||||
|
|
||||||
|
def _make_blend_weights_2d(
|
||||||
|
tile_h: int,
|
||||||
|
tile_w: int,
|
||||||
|
overlap: int,
|
||||||
|
position: tuple[bool, bool, bool, bool],
|
||||||
|
) -> np.ndarray:
|
||||||
|
"""Create 2D linear blend weights for a tile.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
tile_h, tile_w: Tile spatial dimensions.
|
||||||
|
overlap: Overlap size in pixels.
|
||||||
|
position: (has_top_neighbor, has_bottom_neighbor, has_left_neighbor, has_right_neighbor)
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
(tile_h, tile_w) float32 weight array with linear ramps in overlap regions.
|
||||||
|
"""
|
||||||
|
weights = np.ones((tile_h, tile_w), dtype=np.float32)
|
||||||
|
has_top, has_bottom, has_left, has_right = position
|
||||||
|
|
||||||
|
if overlap <= 0:
|
||||||
|
return weights
|
||||||
|
|
||||||
|
ramp = np.linspace(0.0, 1.0, overlap, dtype=np.float32)
|
||||||
|
|
||||||
|
if has_top and overlap <= tile_h:
|
||||||
|
weights[:overlap, :] *= ramp[:, None]
|
||||||
|
if has_bottom and overlap <= tile_h:
|
||||||
|
weights[-overlap:, :] *= ramp[::-1, None]
|
||||||
|
if has_left and overlap <= tile_w:
|
||||||
|
weights[:, :overlap] *= ramp[None, :]
|
||||||
|
if has_right and overlap <= tile_w:
|
||||||
|
weights[:, -overlap:] *= ramp[None, ::-1]
|
||||||
|
|
||||||
|
return weights
|
||||||
|
|
||||||
|
|
||||||
|
def tiled_inference(
|
||||||
|
model: GreenFormer,
|
||||||
|
x: mx.array,
|
||||||
|
tile_size: int = DEFAULT_TILE_SIZE,
|
||||||
|
overlap: int = DEFAULT_OVERLAP,
|
||||||
|
) -> dict[str, mx.array]:
|
||||||
|
"""Run model on overlapping tiles and blend the results.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
model: Loaded GreenFormer (must accept tile_size x tile_size input).
|
||||||
|
x: Full-resolution input (1, H, W, 4) NHWC.
|
||||||
|
tile_size: Size of each square tile. Must match model.backbone.img_size.
|
||||||
|
overlap: Overlap in pixels between adjacent tiles.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Dict with 'alpha_final' and 'fg_final' blended at full resolution.
|
||||||
|
"""
|
||||||
|
if x.shape[0] != 1:
|
||||||
|
msg = f"Tiled inference only supports batch_size=1, got {x.shape[0]}"
|
||||||
|
raise ValueError(msg)
|
||||||
|
|
||||||
|
_, full_h, full_w, _ = x.shape
|
||||||
|
|
||||||
|
# If image fits in one tile, just run normally
|
||||||
|
if full_h <= tile_size and full_w <= tile_size:
|
||||||
|
return model(x)
|
||||||
|
|
||||||
|
y_coords = _compute_tile_coords(full_h, tile_size, overlap)
|
||||||
|
x_coords = _compute_tile_coords(full_w, tile_size, overlap)
|
||||||
|
|
||||||
|
# Accumulators for weighted blending (numpy for simplicity)
|
||||||
|
alpha_accum = np.zeros((full_h, full_w, 1), dtype=np.float32)
|
||||||
|
fg_accum = np.zeros((full_h, full_w, 3), dtype=np.float32)
|
||||||
|
weight_accum = np.zeros((full_h, full_w, 1), dtype=np.float32)
|
||||||
|
|
||||||
|
for yi, (y_start, y_end) in enumerate(y_coords):
|
||||||
|
for xi, (x_start, x_end) in enumerate(x_coords):
|
||||||
|
tile = x[:, y_start:y_end, x_start:x_end, :]
|
||||||
|
|
||||||
|
# Pad to tile_size if needed (edge tiles may be smaller)
|
||||||
|
pad_h = tile_size - (y_end - y_start)
|
||||||
|
pad_w = tile_size - (x_end - x_start)
|
||||||
|
if pad_h > 0 or pad_w > 0:
|
||||||
|
tile = mx.pad(tile, [(0, 0), (0, pad_h), (0, pad_w), (0, 0)])
|
||||||
|
|
||||||
|
out = model(tile)
|
||||||
|
# NOTE: mx.eval is MLX array materialization, not Python eval()
|
||||||
|
mx.eval(out) # noqa: S307
|
||||||
|
|
||||||
|
alpha_tile = np.array(out["alpha_final"][0]) # (tile_h, tile_w, 1)
|
||||||
|
fg_tile = np.array(out["fg_final"][0]) # (tile_h, tile_w, 3)
|
||||||
|
|
||||||
|
# Crop padding
|
||||||
|
actual_h = y_end - y_start
|
||||||
|
actual_w = x_end - x_start
|
||||||
|
alpha_tile = alpha_tile[:actual_h, :actual_w, :]
|
||||||
|
fg_tile = fg_tile[:actual_h, :actual_w, :]
|
||||||
|
|
||||||
|
# Blend weights
|
||||||
|
position = (
|
||||||
|
yi > 0, # has top neighbor
|
||||||
|
yi < len(y_coords) - 1, # has bottom neighbor
|
||||||
|
xi > 0, # has left neighbor
|
||||||
|
xi < len(x_coords) - 1, # has right neighbor
|
||||||
|
)
|
||||||
|
w = _make_blend_weights_2d(actual_h, actual_w, overlap, position)
|
||||||
|
w3d = w[:, :, None] # (H, W, 1)
|
||||||
|
|
||||||
|
alpha_accum[y_start:y_end, x_start:x_end, :] += alpha_tile * w3d
|
||||||
|
fg_accum[y_start:y_end, x_start:x_end, :] += fg_tile * w3d
|
||||||
|
weight_accum[y_start:y_end, x_start:x_end, :] += w3d
|
||||||
|
|
||||||
|
# Normalize by accumulated weights
|
||||||
|
weight_accum = np.maximum(weight_accum, 1e-8)
|
||||||
|
alpha_final = mx.array(alpha_accum / weight_accum)[None] # (1, H, W, 1)
|
||||||
|
fg_final = mx.array(fg_accum / weight_accum)[None] # (1, H, W, 3)
|
||||||
|
|
||||||
|
return {"alpha_final": alpha_final, "fg_final": fg_final}
|
||||||
@ -38,6 +38,11 @@ class GreenFormer(nn.Module):
|
|||||||
self.fg_decoder = DecoderHead(BACKBONE_CHANNELS, EMBED_DIM, output_dim=3)
|
self.fg_decoder = DecoderHead(BACKBONE_CHANNELS, EMBED_DIM, output_dim=3)
|
||||||
self.refiner = CNNRefinerModule()
|
self.refiner = CNNRefinerModule()
|
||||||
|
|
||||||
|
# Decoder outputs at stride-4 (H/4, W/4); upsampler is always 4x
|
||||||
|
self._logit_upsampler = nn.Upsample(
|
||||||
|
scale_factor=(4.0, 4.0), mode="linear", align_corners=False
|
||||||
|
)
|
||||||
|
|
||||||
def __call__(self, x: mx.array) -> dict[str, mx.array]:
|
def __call__(self, x: mx.array) -> dict[str, mx.array]:
|
||||||
"""Forward pass.
|
"""Forward pass.
|
||||||
|
|
||||||
@ -49,8 +54,6 @@ class GreenFormer(nn.Module):
|
|||||||
alpha_coarse, fg_coarse, delta_logits, alpha_final, fg_final.
|
alpha_coarse, fg_coarse, delta_logits, alpha_final, fg_final.
|
||||||
All tensors in NHWC format.
|
All tensors in NHWC format.
|
||||||
"""
|
"""
|
||||||
input_h, input_w = x.shape[1], x.shape[2]
|
|
||||||
|
|
||||||
# Backbone -> 4 multiscale feature maps in NHWC
|
# Backbone -> 4 multiscale feature maps in NHWC
|
||||||
features = self.backbone(x)
|
features = self.backbone(x)
|
||||||
|
|
||||||
@ -58,16 +61,9 @@ class GreenFormer(nn.Module):
|
|||||||
alpha_logits = self.alpha_decoder(features) # (B, H/4, W/4, 1)
|
alpha_logits = self.alpha_decoder(features) # (B, H/4, W/4, 1)
|
||||||
fg_logits = self.fg_decoder(features) # (B, H/4, W/4, 3)
|
fg_logits = self.fg_decoder(features) # (B, H/4, W/4, 3)
|
||||||
|
|
||||||
# Upsample logits to full input resolution
|
# Upsample logits to full input resolution (4x from stride-4 decoder)
|
||||||
scale_h = input_h / alpha_logits.shape[1]
|
alpha_logits_up = self._logit_upsampler(alpha_logits) # (B, H, W, 1)
|
||||||
scale_w = input_w / alpha_logits.shape[2]
|
fg_logits_up = self._logit_upsampler(fg_logits) # (B, H, W, 3)
|
||||||
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
|
# Coarse predictions via sigmoid
|
||||||
alpha_coarse = mx.sigmoid(alpha_logits_up)
|
alpha_coarse = mx.sigmoid(alpha_logits_up)
|
||||||
|
|||||||
@ -42,6 +42,13 @@ class DecoderHead(nn.Module):
|
|||||||
self.linear_c3 = MLP(in_channels[2], embed_dim)
|
self.linear_c3 = MLP(in_channels[2], embed_dim)
|
||||||
self.linear_c4 = MLP(in_channels[3], embed_dim)
|
self.linear_c4 = MLP(in_channels[3], embed_dim)
|
||||||
|
|
||||||
|
# Pre-build upsamplers for feature maps at strides 2x, 4x, 8x
|
||||||
|
# relative to the first (stride-4) feature map.
|
||||||
|
upsample_kwargs = {"mode": "linear", "align_corners": False}
|
||||||
|
self._upsampler_2x = nn.Upsample(scale_factor=(2.0, 2.0), **upsample_kwargs)
|
||||||
|
self._upsampler_4x = nn.Upsample(scale_factor=(4.0, 4.0), **upsample_kwargs)
|
||||||
|
self._upsampler_8x = nn.Upsample(scale_factor=(8.0, 8.0), **upsample_kwargs)
|
||||||
|
|
||||||
fused_channels = embed_dim * len(in_channels)
|
fused_channels = embed_dim * len(in_channels)
|
||||||
self.linear_fuse = nn.Conv2d(fused_channels, embed_dim, kernel_size=1, bias=False)
|
self.linear_fuse = nn.Conv2d(fused_channels, embed_dim, kernel_size=1, bias=False)
|
||||||
self.bn = nn.BatchNorm(embed_dim)
|
self.bn = nn.BatchNorm(embed_dim)
|
||||||
@ -61,12 +68,13 @@ class DecoderHead(nn.Module):
|
|||||||
Logits in NHWC: (B, H/4, W/4, output_dim)
|
Logits in NHWC: (B, H/4, W/4, output_dim)
|
||||||
"""
|
"""
|
||||||
c1, c2, c3, c4 = features
|
c1, c2, c3, c4 = features
|
||||||
target_h, target_w = c1.shape[1], c1.shape[2] # H/4, W/4
|
upsamplers = [None, self._upsampler_2x, self._upsampler_4x, self._upsampler_8x]
|
||||||
|
|
||||||
projected = []
|
projected = []
|
||||||
for feat, linear in zip(
|
for feat, linear, up in zip(
|
||||||
[c1, c2, c3, c4],
|
[c1, c2, c3, c4],
|
||||||
[self.linear_c1, self.linear_c2, self.linear_c3, self.linear_c4],
|
[self.linear_c1, self.linear_c2, self.linear_c3, self.linear_c4],
|
||||||
|
upsamplers,
|
||||||
strict=True,
|
strict=True,
|
||||||
):
|
):
|
||||||
b, h, w, _c = feat.shape
|
b, h, w, _c = feat.shape
|
||||||
@ -74,16 +82,8 @@ class DecoderHead(nn.Module):
|
|||||||
x = feat.reshape(b, h * w, _c)
|
x = feat.reshape(b, h * w, _c)
|
||||||
x = linear(x) # (B, H*W, embed_dim)
|
x = linear(x) # (B, H*W, embed_dim)
|
||||||
x = x.reshape(b, h, w, -1) # (B, H, W, embed_dim)
|
x = x.reshape(b, h, w, -1) # (B, H, W, embed_dim)
|
||||||
# Upsample to target spatial size
|
if up is not None:
|
||||||
if h != target_h or w != target_w:
|
x = up(x)
|
||||||
scale_h = target_h / h
|
|
||||||
scale_w = target_w / w
|
|
||||||
upsampler = nn.Upsample(
|
|
||||||
scale_factor=(scale_h, scale_w),
|
|
||||||
mode="linear",
|
|
||||||
align_corners=False,
|
|
||||||
)
|
|
||||||
x = upsampler(x)
|
|
||||||
projected.append(x)
|
projected.append(x)
|
||||||
|
|
||||||
# Concatenate along channel dim (last dim in NHWC)
|
# Concatenate along channel dim (last dim in NHWC)
|
||||||
|
|||||||
79
src/corridorkey_mlx/utils/profiling.py
Normal file
79
src/corridorkey_mlx/utils/profiling.py
Normal file
@ -0,0 +1,79 @@
|
|||||||
|
"""Profiling utilities for MLX inference benchmarking.
|
||||||
|
|
||||||
|
Provides timing helpers that force mx.eval() for accurate measurement,
|
||||||
|
since MLX uses lazy evaluation.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import time
|
||||||
|
from dataclasses import dataclass
|
||||||
|
|
||||||
|
import mlx.core as mx
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class TimingResult:
|
||||||
|
"""Result from a timed inference run."""
|
||||||
|
|
||||||
|
elapsed_ms: float
|
||||||
|
label: str
|
||||||
|
|
||||||
|
def __repr__(self) -> str:
|
||||||
|
return f"{self.label}: {self.elapsed_ms:.1f} ms"
|
||||||
|
|
||||||
|
|
||||||
|
def time_fn(fn: object, *args: object, label: str = "run", **kwargs: object) -> TimingResult:
|
||||||
|
"""Time a function that returns MLX arrays, forcing evaluation.
|
||||||
|
|
||||||
|
Forces mx.eval on the result to ensure the computation graph
|
||||||
|
is fully executed before stopping the timer.
|
||||||
|
"""
|
||||||
|
start = time.perf_counter()
|
||||||
|
result = fn(*args, **kwargs) # type: ignore[operator]
|
||||||
|
|
||||||
|
# Force evaluation of all MLX arrays in the result
|
||||||
|
# NOTE: mx.eval is MLX's array materialization, not Python eval()
|
||||||
|
if isinstance(result, mx.array):
|
||||||
|
mx.eval(result) # noqa: S307
|
||||||
|
elif isinstance(result, dict):
|
||||||
|
arrays = [v for v in result.values() if isinstance(v, mx.array)]
|
||||||
|
if arrays:
|
||||||
|
mx.eval(*arrays) # noqa: S307
|
||||||
|
elif isinstance(result, (list, tuple)):
|
||||||
|
arrays = [v for v in result if isinstance(v, mx.array)]
|
||||||
|
if arrays:
|
||||||
|
mx.eval(*arrays) # noqa: S307
|
||||||
|
|
||||||
|
elapsed_ms = (time.perf_counter() - start) * 1000.0
|
||||||
|
return TimingResult(elapsed_ms=elapsed_ms, label=label)
|
||||||
|
|
||||||
|
|
||||||
|
def warmup_and_bench(
|
||||||
|
fn: object,
|
||||||
|
*args: object,
|
||||||
|
warmup_runs: int = 3,
|
||||||
|
bench_runs: int = 10,
|
||||||
|
label: str = "bench",
|
||||||
|
**kwargs: object,
|
||||||
|
) -> tuple[TimingResult, TimingResult, list[float]]:
|
||||||
|
"""Run warmup iterations then benchmark, reporting both.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
(warmup_first_run, steady_state_median, all_bench_times_ms)
|
||||||
|
"""
|
||||||
|
# Warmup — first run captures compile cost if using mx.compile
|
||||||
|
warmup_first = time_fn(fn, *args, label=f"{label}/warmup_first", **kwargs)
|
||||||
|
for _ in range(warmup_runs - 1):
|
||||||
|
time_fn(fn, *args, label="warmup", **kwargs)
|
||||||
|
|
||||||
|
# Benchmark runs
|
||||||
|
times: list[float] = []
|
||||||
|
for _ in range(bench_runs):
|
||||||
|
result = time_fn(fn, *args, label=label, **kwargs)
|
||||||
|
times.append(result.elapsed_ms)
|
||||||
|
|
||||||
|
times.sort()
|
||||||
|
median_ms = times[len(times) // 2]
|
||||||
|
steady = TimingResult(elapsed_ms=median_ms, label=f"{label}/steady_median")
|
||||||
|
return warmup_first, steady, times
|
||||||
57
tests/test_compiled_vs_eager_consistency.py
Normal file
57
tests/test_compiled_vs_eager_consistency.py
Normal file
@ -0,0 +1,57 @@
|
|||||||
|
"""Test that mx.compile() produces numerically consistent results vs eager mode."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import mlx.core as mx
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from corridorkey_mlx.model.corridorkey import GreenFormer
|
||||||
|
|
||||||
|
IMG_SIZE = 256
|
||||||
|
TOLERANCE = 1e-4
|
||||||
|
OUTPUT_KEYS = ("alpha_final", "fg_final", "alpha_coarse", "fg_coarse", "delta_logits")
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture()
|
||||||
|
def model() -> GreenFormer:
|
||||||
|
model = GreenFormer(img_size=IMG_SIZE)
|
||||||
|
model.eval()
|
||||||
|
# NOTE: mx.eval is MLX array materialization, not Python eval()
|
||||||
|
mx.eval(model.parameters()) # noqa: S307
|
||||||
|
return model
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture()
|
||||||
|
def dummy_input() -> mx.array:
|
||||||
|
mx.random.seed(42)
|
||||||
|
x = mx.random.normal((1, IMG_SIZE, IMG_SIZE, 4))
|
||||||
|
mx.eval(x) # noqa: S307
|
||||||
|
return x
|
||||||
|
|
||||||
|
|
||||||
|
def test_compiled_matches_eager(model: GreenFormer, dummy_input: mx.array) -> None:
|
||||||
|
"""Fixed-shape compiled output matches eager output within tolerance."""
|
||||||
|
eager_out = model(dummy_input)
|
||||||
|
mx.eval(eager_out) # noqa: S307
|
||||||
|
|
||||||
|
compiled_fn = mx.compile(model.__call__)
|
||||||
|
compiled_out = compiled_fn(dummy_input)
|
||||||
|
mx.eval(compiled_out) # noqa: S307
|
||||||
|
|
||||||
|
for key in OUTPUT_KEYS:
|
||||||
|
diff = float(mx.max(mx.abs(eager_out[key] - compiled_out[key])))
|
||||||
|
assert diff < TOLERANCE, f"{key}: max_abs_diff={diff:.2e} > {TOLERANCE}"
|
||||||
|
|
||||||
|
|
||||||
|
def test_compiled_deterministic(model: GreenFormer, dummy_input: mx.array) -> None:
|
||||||
|
"""Compiled model produces identical results across consecutive calls."""
|
||||||
|
compiled_fn = mx.compile(model.__call__)
|
||||||
|
|
||||||
|
out1 = compiled_fn(dummy_input)
|
||||||
|
mx.eval(out1) # noqa: S307
|
||||||
|
out2 = compiled_fn(dummy_input)
|
||||||
|
mx.eval(out2) # noqa: S307
|
||||||
|
|
||||||
|
for key in OUTPUT_KEYS:
|
||||||
|
diff = float(mx.max(mx.abs(out1[key] - out2[key])))
|
||||||
|
assert diff == 0.0, f"{key}: non-deterministic, max_diff={diff:.2e}"
|
||||||
124
tests/test_tiling_consistency.py
Normal file
124
tests/test_tiling_consistency.py
Normal file
@ -0,0 +1,124 @@
|
|||||||
|
"""Test tiled inference consistency.
|
||||||
|
|
||||||
|
Verifies that:
|
||||||
|
1. Single-tile images produce same results as non-tiled inference.
|
||||||
|
2. Tiled results on larger images have reasonable blending behavior.
|
||||||
|
3. Tile coordinate computation is correct.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import mlx.core as mx
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from corridorkey_mlx.inference.tiling import (
|
||||||
|
_compute_tile_coords,
|
||||||
|
_make_blend_weights_2d,
|
||||||
|
tiled_inference,
|
||||||
|
)
|
||||||
|
from corridorkey_mlx.model.corridorkey import GreenFormer
|
||||||
|
|
||||||
|
TILE_SIZE = 256
|
||||||
|
TOLERANCE = 1e-5
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture()
|
||||||
|
def model() -> GreenFormer:
|
||||||
|
model = GreenFormer(img_size=TILE_SIZE)
|
||||||
|
model.eval()
|
||||||
|
# NOTE: mx.eval is MLX array materialization, not Python eval()
|
||||||
|
mx.eval(model.parameters()) # noqa: S307
|
||||||
|
return model
|
||||||
|
|
||||||
|
|
||||||
|
class TestTileCoords:
|
||||||
|
def test_overlap_gte_tile_size_raises(self) -> None:
|
||||||
|
with pytest.raises(ValueError, match="overlap.*must be less than tile_size"):
|
||||||
|
_compute_tile_coords(512, 256, 256)
|
||||||
|
with pytest.raises(ValueError, match="overlap.*must be less than tile_size"):
|
||||||
|
_compute_tile_coords(512, 256, 300)
|
||||||
|
|
||||||
|
def test_single_tile(self) -> None:
|
||||||
|
coords = _compute_tile_coords(200, 256, 32)
|
||||||
|
assert coords == [(0, 200)]
|
||||||
|
|
||||||
|
def test_exact_fit(self) -> None:
|
||||||
|
coords = _compute_tile_coords(256, 256, 32)
|
||||||
|
assert coords == [(0, 256)]
|
||||||
|
|
||||||
|
def test_two_tiles_with_overlap(self) -> None:
|
||||||
|
coords = _compute_tile_coords(400, 256, 32)
|
||||||
|
assert len(coords) == 2
|
||||||
|
# All tiles cover full range
|
||||||
|
assert coords[0][0] == 0
|
||||||
|
assert coords[-1][1] == 400
|
||||||
|
# Overlap exists between tiles
|
||||||
|
assert coords[0][1] > coords[1][0]
|
||||||
|
|
||||||
|
def test_full_coverage(self) -> None:
|
||||||
|
"""Every pixel is covered by at least one tile."""
|
||||||
|
for image_size in [300, 512, 700, 1024]:
|
||||||
|
coords = _compute_tile_coords(image_size, 256, 32)
|
||||||
|
covered = set()
|
||||||
|
for start, end in coords:
|
||||||
|
covered.update(range(start, end))
|
||||||
|
assert covered == set(range(image_size)), f"Gap in coverage at size={image_size}"
|
||||||
|
|
||||||
|
|
||||||
|
class TestBlendWeights:
|
||||||
|
def test_interior_tile(self) -> None:
|
||||||
|
"""Interior tile (all neighbors) has ramps on all edges."""
|
||||||
|
w = _make_blend_weights_2d(256, 256, 32, (True, True, True, True))
|
||||||
|
assert w.shape == (256, 256)
|
||||||
|
# Corners should be near zero (double ramp)
|
||||||
|
assert w[0, 0] < 0.01
|
||||||
|
# Center should be 1.0
|
||||||
|
assert w[128, 128] == 1.0
|
||||||
|
|
||||||
|
def test_corner_tile(self) -> None:
|
||||||
|
"""Top-left corner tile (no top/left neighbors) has full weight at origin."""
|
||||||
|
w = _make_blend_weights_2d(256, 256, 32, (False, True, False, True))
|
||||||
|
assert w[0, 0] == 1.0
|
||||||
|
# Bottom-right edge should ramp
|
||||||
|
assert w[-1, -1] < 1.0
|
||||||
|
|
||||||
|
def test_no_overlap(self) -> None:
|
||||||
|
w = _make_blend_weights_2d(256, 256, 0, (True, True, True, True))
|
||||||
|
assert (w == 1.0).all()
|
||||||
|
|
||||||
|
|
||||||
|
class TestTiledInference:
|
||||||
|
def test_single_tile_matches_direct(self, model: GreenFormer) -> None:
|
||||||
|
"""Image that fits in one tile produces identical results to direct inference."""
|
||||||
|
mx.random.seed(42)
|
||||||
|
x = mx.random.normal((1, TILE_SIZE, TILE_SIZE, 4))
|
||||||
|
mx.eval(x) # noqa: S307
|
||||||
|
|
||||||
|
direct_out = model(x)
|
||||||
|
mx.eval(direct_out) # noqa: S307
|
||||||
|
|
||||||
|
tiled_out = tiled_inference(model, x, tile_size=TILE_SIZE, overlap=32)
|
||||||
|
mx.eval(tiled_out) # noqa: S307
|
||||||
|
|
||||||
|
for key in ("alpha_final", "fg_final"):
|
||||||
|
diff = float(mx.max(mx.abs(direct_out[key] - tiled_out[key])))
|
||||||
|
assert diff < TOLERANCE, f"{key}: max_diff={diff:.2e}"
|
||||||
|
|
||||||
|
def test_larger_image_runs(self, model: GreenFormer) -> None:
|
||||||
|
"""Tiled inference runs without error on larger-than-tile images."""
|
||||||
|
mx.random.seed(42)
|
||||||
|
x = mx.random.normal((1, 400, 400, 4))
|
||||||
|
mx.eval(x) # noqa: S307
|
||||||
|
|
||||||
|
result = tiled_inference(model, x, tile_size=TILE_SIZE, overlap=32)
|
||||||
|
mx.eval(result) # noqa: S307
|
||||||
|
|
||||||
|
assert result["alpha_final"].shape == (1, 400, 400, 1)
|
||||||
|
assert result["fg_final"].shape == (1, 400, 400, 3)
|
||||||
|
|
||||||
|
def test_batch_size_validation(self, model: GreenFormer) -> None:
|
||||||
|
"""Batch size > 1 raises ValueError."""
|
||||||
|
mx.random.seed(42)
|
||||||
|
x = mx.random.normal((2, TILE_SIZE, TILE_SIZE, 4))
|
||||||
|
with pytest.raises(ValueError, match="batch_size=1"):
|
||||||
|
tiled_inference(model, x, tile_size=TILE_SIZE)
|
||||||
Loading…
Reference in New Issue
Block a user