@@ -780,24 +780,55 @@ void dequantize_turbo3_0_t4(device const block_turbo3_0 * xb, short il, thread t
780780 // TURBO_USE_4MAG=1 (pre-M5): 4-entry magnitude LUT + XOR sign (+38-45% on M2)
781781 // TURBO_USE_4MAG=0 (M5+): 8-entry full LUT (best on M5, 0.905x q8_0)
782782#if TURBO_USE_4MAG
783- // 4-mag LUT + per-element norm multiply: PROVEN BEST on M2 Pro (+38-45%).
784- // 4 divergent constant reads + 0 branches = optimal for M2 hardware.
785- // Approaches with fewer constant reads all add branches/ALU that cost more.
786- const uint8_t mi0 = q0 ^ (s0 ? 0u : 0x3u );
787- const uint8_t mi1 = q1 ^ (s1 ? 0u : 0x3u );
788- const uint8_t mi2 = q2 ^ (s2 ? 0u : 0x3u );
789- const uint8_t mi3 = q3 ^ (s3 ? 0u : 0x3u );
790-
791- const float v0 = float (turbo_mag_3bit_h[mi0]) * norm;
792- const float v1 = float (turbo_mag_3bit_h[mi1]) * norm;
793- const float v2 = float (turbo_mag_3bit_h[mi2]) * norm;
794- const float v3 = float (turbo_mag_3bit_h[mi3]) * norm;
783+ // FMA ARITHMETIC DECODE: compute centroid from bits using fused multiply-add.
784+ // ZERO memory access (no constant, no stack). All compile-time constants.
785+ // Uses fma() which is a single hardware instruction on Apple GPUs.
786+ // Sign computed branchlessly: s = 1.0 - 2.0 * float(sign_bit)
787+ //
788+ // Previous bit-arithmetic used separate multiply+add (11.6 tok/s on M2).
789+ // FMA version chains 3 fma ops which may pipeline better on Apple8.
790+ // 4-mag LUT was 15.1 — need to beat that.
791+ //
792+ // Magnitude from 2-bit index via bilinear interpolation:
793+ // mag = M0 + b0*D1 + b1*D2 + b0*b1*D3
794+ // Implemented as: fma(b0*b1, D3, fma(b1, D2, fma(b0, D1, M0)))
795+
796+ // FULLY BRANCHLESS: zero ternaries, zero selects, zero branches.
797+ // XOR mask: sign_bit=1 → mask=0, sign_bit=0 → mask=3
798+ // Computed as: mask = 3 * (1 - sign_bit) = 3 - 3*sign_bit
799+ const uint xm0 = 3u - 3u * uint (s0);
800+ const uint xm1 = 3u - 3u * uint (s1);
801+ const uint xm2 = 3u - 3u * uint (s2);
802+ const uint xm3 = 3u - 3u * uint (s3);
803+
804+ const uint mi0 = uint (q0) ^ xm0;
805+ const uint mi1 = uint (q1) ^ xm1;
806+ const uint mi2 = uint (q2) ^ xm2;
807+ const uint mi3 = uint (q3) ^ xm3;
808+
809+ // Extract bits
810+ const float b00 = float (mi0 & 1u ), b01 = float ((mi0 >> 1u ) & 1u );
811+ const float b10 = float (mi1 & 1u ), b11 = float ((mi1 >> 1u ) & 1u );
812+ const float b20 = float (mi2 & 1u ), b21 = float ((mi2 >> 1u ) & 1u );
813+ const float b30 = float (mi3 & 1u ), b31 = float ((mi3 >> 1u ) & 1u );
814+
815+ // FMA chain for magnitude (3 fma + 1 multiply per element)
816+ const float mag0 = fma (b00*b01, 0 .028596f , fma (b01, 0 .096372f , fma (b00, 0 .044257f , 0 .021460f )));
817+ const float mag1 = fma (b10*b11, 0 .028596f , fma (b11, 0 .096372f , fma (b10, 0 .044257f , 0 .021460f )));
818+ const float mag2 = fma (b20*b21, 0 .028596f , fma (b21, 0 .096372f , fma (b20, 0 .044257f , 0 .021460f )));
819+ const float mag3 = fma (b30*b31, 0 .028596f , fma (b31, 0 .096372f , fma (b30, 0 .044257f , 0 .021460f )));
820+
821+ // Branchless sign: 2*sign_bit - 1 → +1 or -1 (no ternary, no branch)
822+ const float sg0 = 2 .0f * float (s0) - 1 .0f ;
823+ const float sg1 = 2 .0f * float (s1) - 1 .0f ;
824+ const float sg2 = 2 .0f * float (s2) - 1 .0f ;
825+ const float sg3 = 2 .0f * float (s3) - 1 .0f ;
795826
796827 reg = type4 (float4 (
797- s0 ? v0 : -v0 ,
798- s1 ? v1 : -v1 ,
799- s2 ? v2 : -v2 ,
800- s3 ? v3 : -v3
828+ sg0 * mag0 * norm ,
829+ sg1 * mag1 * norm ,
830+ sg2 * mag2 * norm ,
831+ sg3 * mag3 * norm
801832 ));
802833#else
803834 // 8-entry full LUT: best on M5 Max (0.905x q8_0, 77.4 tok/s)
0 commit comments