diff --git a/docs/plans/2026-03-08-feat-mlx-memory-optimizations-plan.md b/docs/plans/2026-03-08-feat-mlx-memory-optimizations-plan.md index 46ebcf3..469aa56 100644 --- a/docs/plans/2026-03-08-feat-mlx-memory-optimizations-plan.md +++ b/docs/plans/2026-03-08-feat-mlx-memory-optimizations-plan.md @@ -310,24 +310,24 @@ uv run pytest tests/test_parity.py tests/test_tiling.py -v ### Functional -- [ ] All existing fp32 parity tests pass unchanged -- [ ] bf16 forward produces outputs within measurable tolerance of fp32 golden references -- [ ] Tiled inference completes without OOM on representative input -- [ ] `GreenFormer(dtype=mx.float32)` is exact same behavior as current code (zero regression) -- [ ] Checkpoint loading unchanged — same safetensors keys +- [x] All existing fp32 parity tests pass unchanged +- [x] bf16 forward produces outputs within measurable tolerance of fp32 golden references +- [x] Tiled inference completes without OOM on representative input +- [x] `GreenFormer(dtype=mx.float32)` is exact same behavior as current code (zero regression) +- [x] Checkpoint loading unchanged — same safetensors keys ### Non-Functional -- [ ] Peak Metal memory measurably reduced (benchmark with `scripts/bench_mlx.py`) -- [ ] No new dependencies -- [ ] GroupNorm `pytorch_compatible=True` preserved in all instances +- [x] Peak Metal memory measurably reduced (benchmark with `scripts/bench_mlx.py`) +- [x] No new dependencies +- [x] GroupNorm `pytorch_compatible=True` preserved in all instances ### Quality Gates -- [ ] `uv run pytest` — all tests pass -- [ ] `uv run ruff check .` — no lint errors -- [ ] `uv run ruff format .` — formatted -- [ ] `uv run ty check` — no type errors +- [x] `uv run pytest` — all tests pass (80 passed, 4 skipped) +- [x] `uv run ruff check .` — no lint errors +- [x] `uv run ruff format .` — formatted +- [x] `uv run ty check` — pre-existing issues only (torch import, tree_flatten typing) --- @@ -343,6 +343,27 @@ uv run pytest tests/test_parity.py tests/test_tiling.py -v --- +## Benchmark Results (2048x2048, M-series Mac) + +| Mode | Time | Peak Memory | Active After | +|------|------|-------------|-------------| +| fp32 full-frame | 4706ms | 27,587MB | 739MB | +| bf16 full-frame | 4707ms | 27,587MB | 739MB | +| fused full-frame | 4794ms | 28,199MB | 739MB | +| **Tiled 512+64 w/ GC** | **3591ms** | **2,310MB** | **423MB** | + +**Key finding:** Tiled + GC pipeline = 12x peak memory reduction (27.6GB → 2.3GB) and 24% faster than full-frame. bf16/fused have negligible impact at 2048 because the backbone (fp32, 24 blocks) dominates. + +### Spike Results + +| Spike | Result | +|-------|--------| +| 0a: mx.compile + bf16 | PASS — works without issues | +| 0b: mx.metal.clear_cache() | EXISTS but deprecated — use `mx.clear_cache()` | +| 0c: fp16 revert history | No revert — fp16 on separate branch, decoder-only stayed within tolerance | + +--- + ## Open Questions - bf16 parity tolerance — exact numbers depend on Spike measurement diff --git a/src/corridorkey_mlx/engine.py b/src/corridorkey_mlx/engine.py index 2fbd4c5..3768b99 100644 --- a/src/corridorkey_mlx/engine.py +++ b/src/corridorkey_mlx/engine.py @@ -79,9 +79,7 @@ class CorridorKeyMLXEngine: self._img_size = img_size self._use_refiner = use_refiner - self._model: GreenFormer = load_model( - checkpoint, img_size=img_size, compile=compile - ) + self._model: GreenFormer = load_model(checkpoint, img_size=img_size, compile=compile) def process_frame( self, @@ -203,11 +201,7 @@ class CorridorKeyMLXEngine: # -- composite over black -- alpha_3ch = alpha_u8[:, :, np.newaxis].astype(np.float32) / 255.0 fg_float = fg_u8.astype(np.float32) - comp = ( - (fg_float * alpha_3ch).astype(np.uint8) - if fg_is_straight - else fg_u8.copy() - ) + comp = (fg_float * alpha_3ch).astype(np.uint8) if fg_is_straight else fg_u8.copy() return { "alpha": alpha_u8, diff --git a/src/corridorkey_mlx/inference/pipeline.py b/src/corridorkey_mlx/inference/pipeline.py index a5e1c10..540a4c0 100644 --- a/src/corridorkey_mlx/inference/pipeline.py +++ b/src/corridorkey_mlx/inference/pipeline.py @@ -33,6 +33,8 @@ def load_model( img_size: int = DEFAULT_IMG_SIZE, compile: bool = False, shapeless: bool = False, + dtype: mx.Dtype = mx.bfloat16, + fused_decode: bool = True, ) -> GreenFormer: """Build GreenFormer and load weights from safetensors checkpoint. @@ -44,8 +46,12 @@ def load_model( 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. + dtype: Compute dtype for decoder activations. bfloat16 reduces memory; + backbone and sigmoid always stay fp32. All outputs are fp32. + fused_decode: If True, batch alpha+fg decoder upsamples to reduce + Metal dispatch calls. Bit-exact with unfused path. """ - model = GreenFormer(img_size=img_size) + model = GreenFormer(img_size=img_size, dtype=dtype, fused_decode=fused_decode) model.load_checkpoint(checkpoint) if compile: model = compile_model(model, shapeless=shapeless) diff --git a/src/corridorkey_mlx/model/decoder.py b/src/corridorkey_mlx/model/decoder.py index 1d70f78..d262593 100644 --- a/src/corridorkey_mlx/model/decoder.py +++ b/src/corridorkey_mlx/model/decoder.py @@ -144,9 +144,7 @@ class FusedDecoderPair(nn.Module): alpha_head._upsampler_8x, ] - def __call__( - self, features: list[mx.array] - ) -> tuple[mx.array, mx.array]: + def __call__(self, features: list[mx.array]) -> tuple[mx.array, mx.array]: """Forward pass with batched upsampling. Returns: @@ -157,9 +155,7 @@ class FusedDecoderPair(nn.Module): alpha_up = [] fg_up = [] - for a_proj, f_proj, up in zip( - alpha_projs, fg_projs, self._upsamplers, strict=True - ): + for a_proj, f_proj, up in zip(alpha_projs, fg_projs, self._upsamplers, strict=True): if up is not None: # Batch upsample: concat along channel axis (NHWC) fused = mx.concatenate([a_proj, f_proj], axis=-1)