diff --git a/CLAUDE.md b/CLAUDE.md index 83d87d7..e7b341c 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -31,6 +31,10 @@ MLX inference port of CorridorKey for Apple Silicon. - Inference only — no training code - Preserve PyTorch behavior before optimizing +## Shell + +- Do not use `cd` — zoxide overrides it and breaks non-interactive shells. Use absolute paths instead. + ## Commands ```bash diff --git a/prompts/phase4-hiera-backbone.md b/prompts/phase4-hiera-backbone.md new file mode 100644 index 0000000..5933c42 --- /dev/null +++ b/prompts/phase4-hiera-backbone.md @@ -0,0 +1,84 @@ +Work only on the Hiera backbone port and its stage-by-stage parity. + +Goal: +Implement an MLX Hiera backbone that reproduces the feature contract expected by CorridorKey’s downstream heads. + +Context: + +- The original PyTorch model uses a timm backbone created with features_only=True. +- The downstream model expects a pyramid of 4 multiscale feature maps. +- The first input projection is patched from 3 input channels to 4 input channels for RGB + alpha hint. +- This phase is about reproducing backbone behavior and feature outputs, not full end-to-end inference. + +Primary deliverables: + +- src/corridorkey_mlx/model/hiera.py +- tests/test_hiera_stage_shapes.py +- tests/test_hiera_stage_parity.py +- updates to converter mapping files as needed +- README notes describing the backbone parity status and any known gaps + +Requirements: + +1. Inspect the original CorridorKey backbone usage and the underlying Hiera/timm feature contract before editing. +2. Treat the PyTorch reference harness as the source of truth for: + - number of feature stages + - output order + - channel counts + - spatial reductions +3. Implement the smallest MLX Hiera subset necessary to produce the correct 4 feature maps. +4. Preserve the patched 4-channel patch-embed behavior exactly: + - do not invent a new initialization rule + - match the PyTorch conversion semantics for the extra alpha-hint channel +5. Keep tensor layout handling explicit and centralized: + - use one canonical boundary between PyTorch-style NCHW fixtures and MLX NHWC internals + - do not scatter transpose logic across the backbone code +6. Add narrow parity checks for: + - patch embed output + - each stage output + - final list of multiscale features +7. If exact numerical parity is not immediately possible, first achieve: + - correct stage count + - correct stage order + - correct shapes + - correct dtype flow + then localize the first divergent block +8. Keep the implementation modular: + - patch embed + - stage/block definition + - downsampling / pooling transitions + - feature collection +9. Update the checkpoint conversion path only as needed to support the backbone weights for completed modules. + +Diagnostics to produce: + +- stage name +- source tensor shape +- destination tensor shape +- layout assumption +- max abs error +- mean abs error +- whether the mismatch begins before or after a stage transition + +Do not: + +- wire the full GreenFormer end-to-end yet +- rewrite decoder/refiner code unless required by an interface mismatch +- optimize with compile() yet +- add training code +- silently accept stage reordering or “close enough” feature contracts + +Working style: + +- Explore first. +- Summarize the exact Hiera feature contract you believe CorridorKey expects before editing. +- After edits, run the narrowest backbone-only tests first. +- When parity fails, stop at the first divergent stage and explain the likely cause before changing more code. + +Definition of done: + +- MLX backbone returns 4 multiscale feature maps in the correct order +- shape tests pass +- stage-level parity tests exist and run +- the 4-channel patch embed path is implemented and documented +- converter mappings exist for completed backbone weights diff --git a/prompts/phase4-inference-pipeline.md b/prompts/phase4-inference-pipeline.md deleted file mode 100644 index e69de29..0000000 diff --git a/prompts/phase5-assembly-parity.md b/prompts/phase5-assembly-parity.md new file mode 100644 index 0000000..afe84e0 --- /dev/null +++ b/prompts/phase5-assembly-parity.md @@ -0,0 +1,96 @@ +Work only on full model assembly and end-to-end inference parity. + +Goal: +Assemble the MLX GreenFormer end-to-end by wiring: + +- Hiera backbone +- alpha decoder head +- foreground decoder head +- upsampling to full resolution +- refiner input construction +- additive delta-logit refinement +- final sigmoid outputs + +Context: + +- Decoder and refiner modules already exist in MLX. +- Backbone work from phase 4 should now provide the expected 4 feature maps. +- This phase is about reproducing the original forward pass structure and validating end-to-end outputs against the PyTorch oracle. + +Primary deliverables: + +- src/corridorkey_mlx/model/greenformer.py +- src/corridorkey_mlx/inference/pipeline.py +- tests/test_greenformer_forward.py +- tests/test_end_to_end_smoke.py +- tests/test_end_to_end_parity.py +- README usage section for single-image inference + +Requirements: + +1. Recreate the PyTorch forward pass structure exactly: + - 4-channel input + - backbone feature extraction + - dual decoder heads + - coarse alpha / foreground logits + - upsample coarse logits to input resolution + - apply sigmoid to obtain coarse probabilities + - concatenate RGB + coarse predictions into the 7-channel refiner input + - predict delta logits + - add delta logits before final sigmoid +2. Preserve the semantic distinction between: + - coarse logits + - coarse probabilities + - delta logits + - final probabilities +3. Do not collapse or reorder operations unless the PyTorch oracle proves equivalence. +4. Add parity checks for: + - coarse alpha logits + - coarse foreground logits + - coarse alpha probabilities + - coarse foreground probabilities + - refiner input tensor + - delta logits + - final alpha + - final foreground +5. Make inference code explicit about layout conversions and output formats. +6. Support reduced-resolution testing first if needed, but document any deviation from the native target resolution. +7. Provide a simple CLI or script entry point for: + - loading MLX weights + - running one image + alpha hint + - saving or printing a concise summary of output tensors +8. Ensure the model can load completed weights with strict checking wherever practical. +9. Keep preprocessing and postprocessing small and auditable. + +Diagnostics to produce: + +- stage name +- tensor shape +- tensor dtype +- layout convention +- max abs error +- mean abs error +- whether mismatch appears first in coarse path or refinement path + +Do not: + +- start performance tuning yet +- introduce tiling unless necessary for a basic smoke path +- expand into video orchestration +- add training/fine-tuning work +- hide parity gaps behind broad tolerances + +Working style: + +- Before editing, summarize the exact full forward-pass contract in plain English. +- Implement the minimal wiring required for correctness. +- Run narrow forward tests first, then end-to-end smoke, then parity tests. +- If parity fails, identify the earliest divergent artifact and focus there. + +Definition of done: + +- MLX GreenFormer forward pass is assembled +- end-to-end smoke test passes +- parity tests exist for coarse path, refiner path, and final outputs +- single-image inference entry point works +- README documents current supported inference workflow diff --git a/prompts/phase5-optimization.md b/prompts/phase5-optimization.md deleted file mode 100644 index e69de29..0000000 diff --git a/prompts/phase6-optimization.md b/prompts/phase6-optimization.md new file mode 100644 index 0000000..14bc467 --- /dev/null +++ b/prompts/phase6-optimization.md @@ -0,0 +1,92 @@ +Work only on optimization, benchmarking, and Apple-silicon inference ergonomics. + +Goal: +Make the MLX CorridorKey port performant and measurable on Apple silicon, with a focus on reliable benchmarking, memory behavior, and practical large-image inference strategy. + +Context: + +- End-to-end correctness should already exist before this phase starts. +- MLX uses lazy evaluation, so timing code must force evaluation. +- MLX compile() behavior depends on shape stability; changing shapes can trigger recompilation unless shapeless=True is used. +- Shape-dependent logic must be treated carefully if testing shapeless compilation. + +Primary deliverables: + +- scripts/bench_mlx.py +- src/corridorkey_mlx/inference/tiling.py +- src/corridorkey_mlx/utils/profiling.py +- tests/test_tiling_consistency.py +- tests/test_compiled_vs_eager_consistency.py +- README performance section +- benchmark output examples or artifacts + +Requirements: + +1. Build a benchmark harness that measures: + - eager/uncompiled inference + - compiled inference with fixed shapes + - optionally shapeless compiled inference if the graph is safe for it +2. Force evaluation in benchmark timing so measurements reflect real execution rather than queued lazy graphs. +3. Separate: + - first-run / compile cost + - warm steady-state latency + - peak practical memory observations +4. Benchmark at several realistic resolutions, for example: + - small sanity size + - medium working size + - high-resolution target or near-target size +5. Add an optional tiled inference path for large images. +6. Validate tiled vs non-tiled consistency where overlap/blending makes sense. +7. Keep the output report concise and decision-oriented: + - resolution + - eager latency + - compiled latency + - compile warmup cost + - batch size + - output parity summary +8. If testing shapeless compile, explicitly verify that no shape-dependent logic breaks correctness. +9. Keep performance changes behind clean flags or config settings so parity remains easy to reproduce. +10. Document recommended settings for an Apple silicon laptop workflow. + +Optimization scope to consider: + +- fixed-shape compile path +- shapeless compile path only if safe +- reduced intermediate materialization +- memory-aware tiling / chunking +- minimizing unnecessary layout conversions +- avoiding debug prints or implicit materialization in hot paths + +Diagnostics to produce: + +- resolution +- batch size +- compile mode +- first-run latency +- steady-state latency +- consistency check against eager mode +- notes on memory behavior and failure thresholds + +Do not: + +- change model behavior to chase speed +- quantize in this phase unless explicitly requested +- add training +- add speculative custom kernels +- benchmark without explicit eval / synchronization semantics +- hide compile warmup inside steady-state numbers + +Working style: + +- Start by auditing the current forward path for obvious avoidable conversions or materializations. +- Implement the benchmark harness before broad optimization changes. +- Compare eager vs compiled numerics after every meaningful optimization. +- Keep high-resolution strategy explicit and documented. + +Definition of done: + +- repeatable benchmark script exists +- eager vs compiled comparisons are implemented +- optional tiling path exists with consistency tests +- README includes practical Apple silicon guidance +- the repo can answer “what settings should I use on a real Mac?” with evidence