From 9cd800d37b85cb4095307eb56344893ef724c8af Mon Sep 17 00:00:00 2001 From: cmoyates Date: Sun, 1 Mar 2026 09:44:41 -0330 Subject: [PATCH] fix: reverse decoder concat order to match Torch (c4,c3,c2,c1) Torch decoder concatenates feature projections as [c4,c3,c2,c1] but MLX was using [c1,c2,c3,c4]. The linear_fuse conv was trained with the Torch order, producing scrambled features. Verified fix gives correlation=1.0 against Torch on identical inputs. Co-Authored-By: Claude Opus 4.6 --- src/corridorkey_mlx/model/decoder.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/src/corridorkey_mlx/model/decoder.py b/src/corridorkey_mlx/model/decoder.py index f92aff9..fdcdb13 100644 --- a/src/corridorkey_mlx/model/decoder.py +++ b/src/corridorkey_mlx/model/decoder.py @@ -86,8 +86,8 @@ class DecoderHead(nn.Module): x = up(x) projected.append(x) - # Concatenate along channel dim (last dim in NHWC) - fused = mx.concatenate(projected, axis=-1) # (B, H/4, W/4, embed_dim*4) + # Concatenate in c4, c3, c2, c1 order to match Torch decoder + fused = mx.concatenate(projected[::-1], axis=-1) # (B, H/4, W/4, embed_dim*4) fused = self.linear_fuse(fused) fused = self.bn(fused) fused = nn.relu(fused)