2dedc35turing-sphere: reaction-diffusion on the sphere, spectral solver on WebGPUJeremy Magland 1/**
2 * WGSL Fourier-stage kernels (the role cuFFT/VkFFT plays in SHTNS).
3 *
4 * Real fields, band-limited to |m| <= mmax < nphi/2:
5 * - synthesis: assemble a Hermitian spectrum from F_m (m >= 0) and do an
6 * inverse complex FFT along phi; take the real part.
7 * - analysis: forward complex FFT of the (real) row; keep m = 0..mmax.
8 *
9 * Two implementations, selected at plan creation:
10 * - 'fft': radix-2 Stockham in workgroup memory, one workgroup per
11 * latitude row. Requires nphi a power of two and
12 * 2 * 8 * nphi bytes <= maxComputeWorkgroupStorageSize.
13 * - 'dft': direct band-limited trigonometric summation, O(nphi * mmax)
14 * per row. Works for any nphi; also useful as a cross-check.
15 *
16 * All trigonometric factors come from a host-precomputed (f64 -> f32)
17 * table trig[k] = (cos, sin)(2*pi*k/nphi): device sin/cos is only
18 * guaranteed to ~2^-11 absolute error under Vulkan, which would dominate
19 * the fp32 transform error.
20 */
22export interface FourierParams {
23 mmax: number;
24 nlat: number;
25 nphi: number;
26}
28const TRIG_BINDING = /* wgsl */ `
29@group(0) @binding(2) var<storage, read> trig: array<vec2f>; // (cos,sin)(2*pi*k/NPHI), k < NPHI
30`;
32function stockham(nphi: number, threads: number, sign: number): string {
33 const log2n = Math.log2(nphi);
34 if (!Number.isInteger(log2n)) throw new Error('fft requires power-of-two nphi');
35 // twiddle for pass with half-block ns: w = e^{sign*i*pi*j/ns} = T[j * (N/(2*ns))]^sign
36 return /* wgsl */ `
37var<workgroup> bufA: array<vec2f, ${nphi}>;
38var<workgroup> bufB: array<vec2f, ${nphi}>;
40fn cmul(a: vec2f, b: vec2f) -> vec2f {
41 return vec2f(a.x * b.x - a.y * b.y, a.x * b.y + a.y * b.x);
42}
44fn ld(sel: u32, i: u32) -> vec2f {
45 if (sel == 0u) { return bufA[i]; }
46 return bufB[i];
47}
48fn st_(sel: u32, i: u32, v: vec2f) {
49 if (sel == 0u) { bufA[i] = v; } else { bufB[i] = v; }
50}
52// radix-2 Stockham, natural order in and out; data starts in bufA (sel 0)
53// and ends in sel = LOG2N % 2. Unnormalized: X_k = sum_j x_j e^{s*2*pi*i*jk/N}.
54fn fft_inplace(lid: u32) {
55 for (var p = 0u; p < ${log2n}u; p++) {
56 workgroupBarrier();
57 let ns = 1u << p;
58 let sel = p & 1u;
59 let stride = ${nphi / 2}u >> p; // N/(2*ns)
60 for (var t = lid; t < ${nphi / 2}u; t += ${threads}u) {
61 let j = t & (ns - 1u);
62 let tw = trig[j * stride];
63 let w = vec2f(tw.x, ${sign > 0 ? '' : '-'}tw.y);
64 let u = ld(sel, t);
65 let v = cmul(ld(sel, t + ${nphi / 2}u), w);
66 let idst = 2u * (t - j) + j;
67 st_(1u - sel, idst, u + v);
68 st_(1u - sel, idst + ns, u - v);
69 }
70 }
71 workgroupBarrier();
72}
73const FFT_OUT_SEL: u32 = ${log2n % 2}u;
74`;
75}
77/** Choose FFT workgroup size: enough threads for the butterflies, capped at 256. */
78export function fftThreads(nphi: number): number {
79 return Math.max(32, Math.min(256, nphi / 2));
80}
82export function fftSynthWGSL(p: FourierParams): string {
83 const T = fftThreads(p.nphi);
84 return /* wgsl */ `
85const MMAX: u32 = ${p.mmax}u;
86const NLAT: u32 = ${p.nlat}u;
87const NPHI: u32 = ${p.nphi}u;
88@group(0) @binding(0) var<storage, read> fm: array<vec2f>;
89@group(0) @binding(1) var<storage, read_write> spat: array<f32>;
90${TRIG_BINDING}
91${stockham(p.nphi, T, +1)}
93@compute @workgroup_size(${T})
94fn fft_synth(@builtin(local_invocation_id) lid3: vec3u,
95 @builtin(workgroup_id) wid: vec3u) {
96 let lid = lid3.x;
97 let ilat = wid.x;
98 // assemble Hermitian spectrum: X[0] = Re F_0, X[m] = F_m, X[N-m] = conj(F_m)
99 for (var k = lid; k < NPHI; k += ${T}u) {
100 var v = vec2f(0.0);
101 if (k == 0u) {
102 v = vec2f(fm[ilat].x, 0.0);
103 } else if (k <= MMAX) {
104 v = fm[k * NLAT + ilat];
105 } else if (k >= NPHI - MMAX) {
106 let c = fm[(NPHI - k) * NLAT + ilat];
107 v = vec2f(c.x, -c.y);
108 }
109 bufA[k] = v;
110 }
111 fft_inplace(lid);
112 for (var k = lid; k < NPHI; k += ${T}u) {
113 spat[ilat * NPHI + k] = ld(FFT_OUT_SEL, k).x;
114 }
115}
116`;
117}
119export function fftAnalysWGSL(p: FourierParams): string {
120 const T = fftThreads(p.nphi);
121 return /* wgsl */ `
122const MMAX: u32 = ${p.mmax}u;
123const NLAT: u32 = ${p.nlat}u;
124const NPHI: u32 = ${p.nphi}u;
125@group(0) @binding(0) var<storage, read> spat: array<f32>;
126@group(0) @binding(1) var<storage, read_write> fm: array<vec2f>;
127${TRIG_BINDING}
128${stockham(p.nphi, T, -1)}
130@compute @workgroup_size(${T})
131fn fft_analys(@builtin(local_invocation_id) lid3: vec3u,
132 @builtin(workgroup_id) wid: vec3u) {
133 let lid = lid3.x;
134 let ilat = wid.x;
135 for (var k = lid; k < NPHI; k += ${T}u) {
136 bufA[k] = vec2f(spat[ilat * NPHI + k], 0.0);
137 }
138 fft_inplace(lid);
139 for (var m = lid; m <= MMAX; m += ${T}u) {
140 fm[m * NLAT + ilat] = ld(FFT_OUT_SEL, m);
141 }
142}
143`;
144}
146export function dftSynthWGSL(p: FourierParams): string {
147 return /* wgsl */ `
148const MMAX: u32 = ${p.mmax}u;
149const NLAT: u32 = ${p.nlat}u;
150const NPHI: u32 = ${p.nphi}u;
151@group(0) @binding(0) var<storage, read> fm: array<vec2f>;
152@group(0) @binding(1) var<storage, read_write> spat: array<f32>;
153${TRIG_BINDING}
155@compute @workgroup_size(64)
156fn dft_synth(@builtin(global_invocation_id) gid: vec3u) {
157 let iphi = gid.x;
158 let ilat = gid.y;
159 if (iphi >= NPHI) { return; }
160 var v: f32 = fm[ilat].x; // m = 0: real part
161 for (var m = 1u; m <= MMAX; m++) {
162 let w = trig[(m * iphi) % NPHI]; // e^{+i m phi}
163 let c = fm[m * NLAT + ilat];
164 v += 2.0 * (c.x * w.x - c.y * w.y);
165 }
166 spat[ilat * NPHI + iphi] = v;
167}
168`;
169}
171export function dftAnalysWGSL(p: FourierParams): string {
172 return /* wgsl */ `
173const MMAX: u32 = ${p.mmax}u;
174const NLAT: u32 = ${p.nlat}u;
175const NPHI: u32 = ${p.nphi}u;
176@group(0) @binding(0) var<storage, read> spat: array<f32>;
177@group(0) @binding(1) var<storage, read_write> fm: array<vec2f>;
178${TRIG_BINDING}
180@compute @workgroup_size(64)
181fn dft_analys(@builtin(global_invocation_id) gid: vec3u) {
182 let m = gid.x;
183 let ilat = gid.y;
184 if (m > MMAX) { return; }
185 var acc = vec2f(0.0);
186 for (var j = 0u; j < NPHI; j++) {
187 let w = trig[(m * j) % NPHI]; // conj => e^{-i m phi}
188 let f = spat[ilat * NPHI + j];
189 acc += f * vec2f(w.x, -w.y);
190 }
191 fm[m * NLAT + ilat] = acc;
192}
193`;
194}