Skip to content

Commit bfbe6d1

Browse files
nindanaotoclaude
andcommitted
AVX512 interleaved: template fwd/inv branch + hand-tuned asm: 671ns → 592ns; make default
- Template radix4 butterfly and stockham_r4_from on bool Fwd, eliminating the runtime fwd/inv conditional from the hot loop - Add hand-tuned AVX512 asm inner p-loop (modeled after existing AVX2 path): preloads w1/w2/w3 into zmm0-2 and jmask into zmm3 before the p-loop, uses zmm4-11 as temporaries, leaving zmm12-31 free for OOO scheduling - RawIFFTMulFFT<1024>: 671ns → 592ns (-12%); now 14% faster than split AVX512 - Make USE_SPQLIOS_INTL=ON the default (was OFF) Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
1 parent cf596f0 commit bfbe6d1

2 files changed

Lines changed: 113 additions & 34 deletions

File tree

CMakeLists.txt

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -83,7 +83,7 @@ option(USE_MKL "Use Intel MKL" OFF)
8383
option(USE_FFTW3 "Use FFTW3" OFF)
8484
option(USE_SPQLIOX_AARCH64 "Use spqliox_aarch64" OFF)
8585
option(USE_SPQLIOS_ARITHMETIC "Use spqlios-arithmetic backend" OFF)
86-
option(USE_SPQLIOS_INTL "Use interleaved-format SPQLIOS" OFF)
86+
option(USE_SPQLIOS_INTL "Use interleaved-format SPQLIOS" ON)
8787
option(USE_SPQLIOS_STOCKHAM "Use Stockham radix-4 SPQLIOS (split layout)" OFF)
8888
option(USE_HEXL "Use Intel HEXL" OFF)
8989

thirdparties/spqlios-intl/fft_processor_spqlios_intl.cpp

Lines changed: 112 additions & 33 deletions
Original file line numberDiff line numberDiff line change
@@ -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>
6667
static 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>
8486
static 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);
229231
static 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

Comments
 (0)