* feat(phase4a): MLX Hiera backbone with unroll/reroll parity Port timm hiera_base_plus_224 to MLX: PatchEmbed, MaskUnitAttention, HieraBlock, unroll/reroll, HieraBackbone. 4 NHWC feature maps at strides 4/8/16/32. backbone.py now re-exports from hiera.py. Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com> * feat(phase4c): weight loading + pos_embed bicubic interpolation HieraBackbone.load_checkpoint() loads safetensors, strips encoder.model. prefix, bicubic-interpolates pos_embed from 512x512 to 128x128 tokens. Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com> * feat(phase4d): shape + parity tests, fix attention transpose Fix MaskUnitAttention output transpose (0,2,3,1,4 → 0,3,2,1,4) to match PyTorch transpose(1,3) token ordering for windowed attention. Parity results (4 stages): Stage 0: max_abs 2.9e-4, mean 1.6e-5 Stage 1: max_abs 1.3e-4, mean 1.3e-5 Stage 2: max_abs 1.1e-2, mean 5.2e-5 (16 blocks, expected drift) Stage 3: max_abs 5.8e-4, mean 2.6e-5 Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com> * docs: clarify pos_embed parameter intent in HieraBackbone Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com> --------- Co-authored-by: Claude Opus 4.6 <noreply@anthropic.com>
5.0 KiB
5.0 KiB
Phase 4: Hiera Backbone MLX Port
Context
Phase 4 of the CorridorKey MLX port. Phases 1-3 complete (reference harness, decoder/refiner, converter). Now need the Hiera backbone — the most complex component. Once done, all model pieces exist for end-to-end inference.
Architecture Summary
Hiera is a hierarchical vision transformer from timm. Key components:
- PatchEmbed: Conv2d(4->112, 7x7, stride=4) + reshape to [B, N, C]
- pos_embed: Learned (1, N, 112), bicubic interpolated from training res
- Unroll: Permutes tokens for mask-unit windowed attention
- 24 HieraBlocks in 4 stages:
- Stage 0 (blocks 0-1): dim=112, heads=2, window=64, mask_unit_attn=True
- Stage 1 (blocks 2-4): dim=224, heads=4, window=16, mask_unit_attn=True
- Stage 2 (blocks 5-20): dim=448, heads=8, window=4, mask_unit_attn=False (global)
- Stage 3 (blocks 21-23): dim=896, heads=16, window=1, mask_unit_attn=False (global)
- Reroll: Undoes permutation -> spatial [B, H, W, C]
- Stage transitions (blocks 2, 5, 21): proj Linear + max-pool reduces tokens 4x
- Features: Collected at stage_ends [1, 4, 20, 23], rerolled to NHWC
MaskUnitAttention
- qkv = Linear(dim_in, 3*dim_out), reshaped per window
- q_stride > 1 at transitions: max-pool over q_stride dim in queries
- Scaled dot-product attention
- proj = Linear(dim_out, dim_out)
MLP
- fc1(dim->4dim) -> GELU -> fc2(4dim->dim)
Checkpoint Details (from safetensors)
- 297 encoder keys total
encoder.model.patch_embed.proj.weight: (112, 7, 7, 4) — already transposed by converterencoder.model.pos_embed: (1, 262144, 112) — raw, needs bicubic interpolation- Transition blocks (2, 5, 21) have extra
proj.weight/bias - All block weights are Linear (no transpose needed)
- No LayerScale weights (init_values=None for base_plus)
Implementation Plan
Sub-phase 4a: Core backbone module
File: src/corridorkey_mlx/model/hiera.py
Implement (all in [B, N, C] sequence format, NHWC only at boundaries):
HieraPatchEmbed— Conv2d(4->112, 7x7, stride=4) + flatten to [B, N, C]unroll(x, spatial_size, schedule)— faithful port of timm's Unroll.forwardreroll(x, block_idx, schedule_map)— faithful port of timm's Reroll.forward (no-mask path)undo_windowing(x, shape, mu_shape)— helper for rerollMaskUnitAttention— windowed multi-head attention with optional q max-poolHieraMLP— fc1 -> GELU -> fc2HieraBlock— norm1 -> [proj+maxpool if transition] -> attn + residual -> norm2 -> mlp + residualHieraBackbone— full assembly: patch_embed + pos_embed + unroll + 24 blocks + reroll at stage_ends
Sub-phase 4b: Converter updates
File: src/corridorkey_mlx/convert/converter.py
- Encoder keys are already handled (patch_embed.proj in CONV_WEIGHT_KEYS, all else passthrough)
- No changes needed — verified all 297 keys pass through correctly
Sub-phase 4c: Weight loading + pos_embed interpolation
In: src/corridorkey_mlx/model/hiera.py (method on HieraBackbone)
- Load safetensors weights into MLX model
- Bicubic interpolation for pos_embed: (1, 262144, 112) -> (1, N', 112)
- Reshape to (1, H_train, W_train, 112) -> bilinear/bicubic resize -> flatten
Sub-phase 4d: Tests
Files:
tests/test_hiera_stage_shapes.py— shape contract tests (no checkpoint needed)tests/test_hiera_stage_parity.py— numerical parity vs golden fixtures
Shape tests:
- PatchEmbed output shape
- Each stage feature map shape (128x128x112, 64x64x224, 32x32x448, 16x16x896)
- Correct number of features (4)
Parity tests (require checkpoint + fixtures):
- Load golden
encoder_feature_{0-3}fromreference/fixtures/golden.npz - Load checkpoint weights into MLX backbone
- Compare each stage output (NCHW fixtures -> NHWC comparison)
- Report max_abs and mean_abs error per stage
Key Files
| File | Action |
|---|---|
src/corridorkey_mlx/model/hiera.py |
create (replace placeholder backbone.py) |
src/corridorkey_mlx/model/backbone.py |
keep as thin re-export |
tests/test_hiera_stage_shapes.py |
create |
tests/test_hiera_stage_parity.py |
create |
Key Decisions
- Faithful Unroll/Reroll port — required for exact parity
- File naming:
hiera.py— more descriptive;backbone.pystays as re-export - No masking support — inference only, skip masked token paths
- pos_embed interpolation at load time — converter outputs raw checkpoint pos_embed
- All backbone ops in [B, N, C] — NHWC only at boundaries (input image, output features)
- MLX key prefix:
encoder.model.stripped when loading into HieraBackbone
Verification
uv run pytest tests/test_hiera_stage_shapes.py -v
uv run pytest tests/test_hiera_stage_parity.py -v -s
uv run ruff check src/corridorkey_mlx/model/hiera.py
uv run mypy src/corridorkey_mlx/model/hiera.py
Unresolved Questions
- E2e parity tolerance? Expecting ~1e-4 max_abs, Metal float32 may drift more through 24 blocks
mx.fast.scaled_dot_product_attentionavailability/API — fallback to manual if needed