@@ -62,34 +62,36 @@ static inline __m512d mul_j_inv512(__m512d x) {
6262 _mm512_set_pd (-0.0 , 0.0 , -0.0 , 0.0 , -0.0 , 0.0 , -0.0 , 0.0 ));
6363}
6464
65- // Radix-4 butterfly with twiddle (4 complex per ZMM)
65+ // Radix-4 butterfly with twiddle (4 complex per ZMM) — templated on direction
66+ template <bool Fwd>
6667static inline void radix4_dit_butterfly512 (
6768 __m512d a, __m512d b, __m512d c, __m512d d,
68- __m512d w1, __m512d w2, __m512d w3, bool fwd,
69+ __m512d w1, __m512d w2, __m512d w3,
6970 __m512d &out0, __m512d &out1, __m512d &out2, __m512d &out3)
7071{
7172 __m512d apc = _mm512_add_pd (a, c);
7273 __m512d amc = _mm512_sub_pd (a, c);
7374 __m512d bpd = _mm512_add_pd (b, d);
7475 __m512d bmd = _mm512_sub_pd (b, d);
75- __m512d jbmd = fwd ? mul_j_fwd512 (bmd) : mul_j_inv512 (bmd);
76+ __m512d jbmd = Fwd ? mul_j_fwd512 (bmd) : mul_j_inv512 (bmd);
7677
7778 out0 = _mm512_add_pd (apc, bpd);
7879 out1 = cmul512 (_mm512_sub_pd (amc, jbmd), w1);
7980 out2 = cmul512 (_mm512_sub_pd (apc, bpd), w2);
8081 out3 = cmul512 (_mm512_add_pd (amc, jbmd), w3);
8182}
8283
83- // Last radix-4 butterfly (no twiddle)
84+ // Last radix-4 butterfly (no twiddle) — templated on direction
85+ template <bool Fwd>
8486static inline void radix4_last_butterfly512 (
85- __m512d a, __m512d b, __m512d c, __m512d d, bool fwd,
87+ __m512d a, __m512d b, __m512d c, __m512d d,
8688 __m512d &out0, __m512d &out1, __m512d &out2, __m512d &out3)
8789{
8890 __m512d apc = _mm512_add_pd (a, c);
8991 __m512d amc = _mm512_sub_pd (a, c);
9092 __m512d bpd = _mm512_add_pd (b, d);
9193 __m512d bmd = _mm512_sub_pd (b, d);
92- __m512d jbmd = fwd ? mul_j_fwd512 (bmd) : mul_j_inv512 (bmd);
94+ __m512d jbmd = Fwd ? mul_j_fwd512 (bmd) : mul_j_inv512 (bmd);
9395
9496 out0 = _mm512_add_pd (apc, bpd);
9597 out1 = _mm512_sub_pd (amc, jbmd);
@@ -229,59 +231,135 @@ static const __m512i idx_deinl_re = _mm512_set_epi64(14, 12, 10, 8, 6, 4, 2, 0);
229231static const __m512i idx_deinl_im = _mm512_set_epi64(15 , 13 , 11 , 9 , 7 , 5 , 3 , 1 );
230232
231233// ── Stockham radix-4 FFT (AVX512) ──────────────────────────────────────────
232- // Each ZMM processes 4 complex values at a time .
233- // src_ro: read-only initial source (may differ from x for zero-copy FFT) .
234- // x, y: writable work/scratch buffers .
234+ // Templated on direction to eliminate the runtime fwd/inv branch .
235+ // Hand-tuned asm inner loop: w1/w2/w3 preloaded in zmm0-2, jmask in zmm3 .
236+ // zmm4-zmm11 used as temporaries; zmm12-zmm31 remain free for OOO scheduling .
235237
236- static void stockham_r4_from (int32_t ns2, bool fwd, const double *trig,
238+ template <bool Fwd>
239+ static void stockham_r4_from (int32_t ns2, const double *trig,
237240 const double *src_ro, double *x, double *y) {
238241 const double *src = src_ro;
239242 double *dst = y;
240243 const double *tw = trig;
241244 int32_t q = ns2 / 4 ;
242245 int32_t s = 1 ;
243246
247+ // j-multiply sign mask, compile-time selected:
248+ // Fwd: swap re/im then negate even (re) positions → multiply by +j
249+ // Inv: swap re/im then negate odd (im) positions → multiply by -j
250+ const __m512d jmask = Fwd ?
251+ _mm512_set_pd (0.0 , -0.0 , 0.0 , -0.0 , 0.0 , -0.0 , 0.0 , -0.0 ) :
252+ _mm512_set_pd (-0.0 , 0.0 , -0.0 , 0.0 , -0.0 , 0.0 , -0.0 , 0.0 );
253+
244254 while (q >= 4 ) {
245255 int32_t stride = q * s;
246256 if (s == 1 ) {
257+ // First pass: no twiddle multiply; process 4 consecutive p values per ZMM.
247258 for (int32_t p = 0 ; p < q; p += 4 ) {
248259 __m512d a = _mm512_loadu_pd (src + p * 2 );
249260 __m512d b = _mm512_loadu_pd (src + (p + stride) * 2 );
250261 __m512d c = _mm512_loadu_pd (src + (p + 2 *stride) * 2 );
251262 __m512d d = _mm512_loadu_pd (src + (p + 3 *stride) * 2 );
252263 __m512d r0, r1, r2, r3;
253- radix4_last_butterfly512 (a, b, c, d, fwd , r0, r1, r2, r3);
264+ radix4_last_butterfly512<Fwd> (a, b, c, d, r0, r1, r2, r3);
254265 _mm512_storeu_pd (dst + p * 2 , r0);
255266 _mm512_storeu_pd (dst + (p + q) * 2 , r1);
256267 _mm512_storeu_pd (dst + (p + 2 *q) * 2 , r2);
257268 _mm512_storeu_pd (dst + (p + 3 *q) * 2 , r3);
258269 }
259270 } else {
271+ // Subsequent passes: preload w1/w2/w3 into zmm0-2 and jmask into zmm3,
272+ // then run hand-tuned asm p-loop to guarantee no ZMM register spills.
273+ int64_t stride_bytes = (int64_t )stride * 16 ; // stride × 2 doubles × 8 bytes
260274 for (int32_t j = 0 ; j < s; j += 4 ) {
261- __m512d w1 = _mm512_loadu_pd (tw); tw += 8 ;
262- __m512d w2 = _mm512_loadu_pd (tw); tw += 8 ;
263- __m512d w3 = _mm512_loadu_pd (tw); tw += 8 ;
275+ int64_t q_bytes = (int64_t )q * 16 ; // q × 2 doubles × 8 bytes
276+
277+ // Preload twiddles and jmask before the p-loop.
278+ // zmm0-3 clobbered; GCC will not assign C variables to them
279+ // between this block and the inner asm, leaving our values intact.
280+ __asm__ __volatile__ (
281+ " vmovupd (%[tw]), %%zmm0\n\t " // w1: 4 complex twiddles
282+ " vmovupd 64(%[tw]), %%zmm1\n\t " // w2
283+ " vmovupd 128(%[tw]), %%zmm2\n\t " // w3
284+ " vmovapd %[jm], %%zmm3\n\t " // jmask (compile-time constant)
285+ : : [tw] " r" (tw), [jm] " x" (jmask)
286+ : " zmm0" ," zmm1" ," zmm2" ," zmm3"
287+ );
288+
264289 for (int32_t p = 0 ; p < q; p++) {
265- int32_t idx = j + s * p;
266- __m512d a = _mm512_loadu_pd (src + idx * 2 );
267- __m512d b = _mm512_loadu_pd (src + (idx + stride) * 2 );
268- __m512d c = _mm512_loadu_pd (src + (idx + 2 *stride) * 2 );
269- __m512d d = _mm512_loadu_pd (src + (idx + 3 *stride) * 2 );
270- int32_t o = p + q * (4 * j);
271- __m512d r0, r1, r2, r3;
272- radix4_dit_butterfly512 (a, b, c, d, w1, w2, w3, fwd, r0, r1, r2, r3);
273- _mm512_storeu_pd (dst + o * 2 , r0);
274- _mm512_storeu_pd (dst + (o + q) * 2 , r1);
275- _mm512_storeu_pd (dst + (o + 2 *q) * 2 , r2);
276- _mm512_storeu_pd (dst + (o + 3 *q) * 2 , r3);
290+ const double *sptr = src + (int64_t )(j + s * p) * 2 ;
291+ double *dptr = dst + (int64_t )(p + q * (4 * j)) * 2 ;
292+
293+ // zmm0=w1, zmm1=w2, zmm2=w3, zmm3=jmask (preserved across p-loop).
294+ // zmm4-zmm11: temporaries; zmm12-zmm31: untouched.
295+ __asm__ __volatile__ (
296+ // ── Load a, b, c, d ──────────────────────────────────
297+ " vmovupd (%[src]), %%zmm4\n\t " // a
298+ " vmovupd (%[src],%[st]), %%zmm5\n\t " // b
299+ " vmovupd (%[src],%[st],2), %%zmm6\n\t " // c
300+ " vmovupd (%[src],%[st3]), %%zmm7\n\t " // d
301+
302+ // ── Radix-4 butterfly sums/diffs ─────────────────────
303+ " vaddpd %%zmm6, %%zmm4, %%zmm8\n\t " // apc = a+c
304+ " vsubpd %%zmm6, %%zmm4, %%zmm9\n\t " // amc = a-c
305+ " vaddpd %%zmm7, %%zmm5, %%zmm10\n\t " // bpd = b+d
306+ " vsubpd %%zmm7, %%zmm5, %%zmm4\n\t " // bmd = b-d (reuse zmm4)
307+ // j-multiply: swap re/im pairs, then apply sign mask
308+ " vpermilpd $0x55, %%zmm4, %%zmm4\n\t "
309+ " vxorpd %%zmm3, %%zmm4, %%zmm11\n\t " // jbmd (zmm3=jmask)
310+ // live: apc(8), amc(9), bpd(10), jbmd(11)
311+
312+ // ── out0 = apc + bpd (store immediately) ────────────
313+ " vaddpd %%zmm10, %%zmm8, %%zmm4\n\t "
314+ " vmovupd %%zmm4, (%[dst])\n\t "
315+
316+ // ── out2 = (apc − bpd) × w2 ──────────────────────────
317+ " vsubpd %%zmm10, %%zmm8, %%zmm4\n\t " // t = apc-bpd
318+ " vpermilpd $0x55, %%zmm1, %%zmm5\n\t " // w2_swap
319+ " vunpckhpd %%zmm4, %%zmm4, %%zmm6\n\t " // t_im broadcast
320+ " vunpcklpd %%zmm4, %%zmm4, %%zmm4\n\t " // t_re broadcast
321+ " vmulpd %%zmm5, %%zmm6, %%zmm6\n\t " // t_im × w2_swap
322+ " vfmaddsub231pd %%zmm1, %%zmm4, %%zmm6\n\t " // t_re×w2 ± t_im×w2_swap
323+ " vmovupd %%zmm6, (%[dst],%[q2])\n\t "
324+
325+ // ── out1 = (amc − jbmd) × w1 ─────────────────────────
326+ " vsubpd %%zmm11, %%zmm9, %%zmm4\n\t " // t = amc-jbmd
327+ " vpermilpd $0x55, %%zmm0, %%zmm5\n\t " // w1_swap
328+ " vunpckhpd %%zmm4, %%zmm4, %%zmm6\n\t "
329+ " vunpcklpd %%zmm4, %%zmm4, %%zmm4\n\t "
330+ " vmulpd %%zmm5, %%zmm6, %%zmm6\n\t "
331+ " vfmaddsub231pd %%zmm0, %%zmm4, %%zmm6\n\t "
332+ " vmovupd %%zmm6, (%[dst],%[q1])\n\t "
333+
334+ // ── out3 = (amc + jbmd) × w3 ─────────────────────────
335+ " vaddpd %%zmm11, %%zmm9, %%zmm4\n\t "
336+ " vpermilpd $0x55, %%zmm2, %%zmm5\n\t " // w3_swap
337+ " vunpckhpd %%zmm4, %%zmm4, %%zmm6\n\t "
338+ " vunpcklpd %%zmm4, %%zmm4, %%zmm4\n\t "
339+ " vmulpd %%zmm5, %%zmm6, %%zmm6\n\t "
340+ " vfmaddsub231pd %%zmm2, %%zmm4, %%zmm6\n\t "
341+ " vmovupd %%zmm6, (%[dst],%[q3])\n\t "
342+
343+ : : [src] " r" (sptr), [dst] " r" (dptr),
344+ [st] " r" (stride_bytes),
345+ [st3] " r" (stride_bytes * 3 ),
346+ [q1] " r" (q_bytes),
347+ [q2] " r" (2 * q_bytes),
348+ [q3] " r" (3 * q_bytes)
349+ : " zmm4" ," zmm5" ," zmm6" ," zmm7" ,
350+ " zmm8" ," zmm9" ," zmm10" ," zmm11" ," memory"
351+ );
277352 }
353+ tw += 24 ; // advance past this j-group (3 ZMMs × 8 doubles)
278354 }
279355 }
280- // After first pass, switch to mutable buffers
356+ // Switch ping-pong buffers
281357 if (s == 1 ) { src = dst; dst = x; }
282358 else { const double *tmp = src; src = dst; dst = const_cast <double *>(tmp); }
283359 s *= 4 ; q /= 4 ;
284360 }
361+
362+ // Final q=1 or q=2 pass (no twiddle multiply)
285363 if (q == 1 ) {
286364 int32_t stride = s;
287365 for (int32_t j = 0 ; j < s; j += 4 ) {
@@ -290,7 +368,7 @@ static void stockham_r4_from(int32_t ns2, bool fwd, const double *trig,
290368 __m512d c = _mm512_loadu_pd (src + (j + 2 *stride) * 2 );
291369 __m512d d = _mm512_loadu_pd (src + (j + 3 *stride) * 2 );
292370 __m512d r0, r1, r2, r3;
293- radix4_last_butterfly512 (a, b, c, d, fwd , r0, r1, r2, r3);
371+ radix4_last_butterfly512<Fwd> (a, b, c, d, r0, r1, r2, r3);
294372 _mm512_storeu_pd (dst + (4 *j) * 2 , r0);
295373 _mm512_storeu_pd (dst + (4 *j + 4 ) * 2 , r1);
296374 _mm512_storeu_pd (dst + (4 *j + 8 ) * 2 , r2);
@@ -313,9 +391,10 @@ static void stockham_r4_from(int32_t ns2, bool fwd, const double *trig,
313391 if (dst != x) memcpy (x, dst, ns2 * 2 * sizeof (double ));
314392}
315393
316- // Convenience wrapper: src and work buffer are the same (in-place)
317- static void stockham_r4 (int32_t ns2, bool fwd, const double *trig, double *x, double *y) {
318- stockham_r4_from (ns2, fwd, trig, x, x, y);
394+ // Convenience wrapper: in-place (src == x)
395+ template <bool Fwd>
396+ static void stockham_r4 (int32_t ns2, const double *trig, double *x, double *y) {
397+ stockham_r4_from<Fwd>(ns2, trig, x, x, y);
319398}
320399
321400#else // AVX2
@@ -563,7 +642,7 @@ static void intl_fft_from(const INTL_FFT_PRECOMP *tables, const double *src_in,
563642 const double *trig = tables->trig_fwd ;
564643
565644#ifdef USE_AVX512
566- stockham_r4_from (ns2, true , trig, src_in, c, tables->scratch );
645+ stockham_r4_from< true > (ns2, trig, src_in, c, tables->scratch );
567646
568647 const double *tw_twist = trig;
569648 for (int32_t s = 4 ; 4 *s < ns2; s *= 4 )
@@ -618,7 +697,7 @@ static void intl_ifft(const INTL_FFT_PRECOMP *tables, double *c) {
618697 _mm512_store_pd (p, cmul512 (a, w));
619698 }
620699 const double *tw_bf = trig + ns2 * 2 ;
621- stockham_r4 (ns2, false , tw_bf, c, tables->scratch );
700+ stockham_r4< false > (ns2, tw_bf, c, tables->scratch );
622701#else
623702 // AVX2: Separate twist + Stockham inverse with hand-tuned butterfly
624703 for (int32_t j = 0 ; j < ns2; j += 2 ) {
0 commit comments