/ concept-collection / turing-sphere-2
Sign in
concept-collection / turing-sphere-2
turing-sphere-2 / src / sht / wgsl / leg.ts
182 lines · 5.2 KBBlameHistoryRaw
1/**
2 * WGSL Legendre-transform kernels, modeled on leg_m_kernel / ileg_m_kernel
3 * in SHT/cuda_legendre.gen.cu (non-Ishioka fp32 path: SHTNS disables the
4 * Ishioka recurrence for fp32 because it loses too much accuracy).
5 *
6 * Synthesis: F_m(theta_i) = sum_{l=m..lmax} Q_lm * ytilde_l^m(theta_i)
7 * - one thread per latitude, one workgroup row per m (workgroup_id.y).
8 * Analysis: Q_lm = sum_i w_i * G_m(theta_i) * ytilde_l^m(theta_i)
9 * - one workgroup per m; threads own latitudes (strided); per-l pair
10 * workgroup tree reduction (portable stand-in for the CUDA warp
11 * shuffles).
12 *
13 * The associated Legendre functions are generated on the fly by the
14 * standard 3-term recurrence over l (coefficients a,b precomputed on the
15 * host in f64), with the SHTNS fp32 rescaling scheme for sin(theta)^m
16 * underflow (see common.ts).
17 */
18import { RESCALE_WGSL } from './common.ts';
20export interface LegParams {
21 lmax: number;
22 mmax: number;
23 nlat: number;
24 wgSynth: number; // workgroup size for synthesis (threads over latitude)
25 wgAnalys: number; // workgroup size for analysis (power of two)
28const BINDINGS = /* wgsl */ `
29@group(0) @binding(0) var<storage, read> ab: array<vec2f>; // (a_l^m, b_l^m) per lm
30@group(0) @binding(1) var<storage, read> amm: array<f32>; // seed per m
31@group(0) @binding(2) var<storage, read> ctstw: array<f32>; // [ct | st | w], each NLAT
32`;
34export function legSynthWGSL(p: LegParams): string {
35 return /* wgsl */ `
36${RESCALE_WGSL}
37const LMAX: u32 = ${p.lmax}u;
38const NLAT: u32 = ${p.nlat}u;
39${BINDINGS}
40@group(0) @binding(3) var<storage, read> qlm: array<vec2f>;
41@group(0) @binding(4) var<storage, read_write> fm: array<vec2f>; // [(m)*NLAT + ilat]
43@compute @workgroup_size(${p.wgSynth})
44fn leg_synth(@builtin(global_invocation_id) gid: vec3u,
45 @builtin(workgroup_id) wid: vec3u) {
46 let ilat = gid.x;
47 let m = wid.y;
48 if (ilat >= NLAT) { return; }
50 let ct = ctstw[ilat];
51 let st = ctstw[NLAT + ilat];
52 let base = m * (LMAX + 1u) - (m * (m - 1u)) / 2u; // lm index of (l=m, m)
54 var seed = sinpow_rescaled(st, m);
55 var y0 = seed.y0 * amm[m];
56 var ny = seed.ny;
57 var y1: f32 = 0.0;
58 if (m < LMAX) {
59 y1 = ab[base + 1u].x * ct * y0;
60 }
62 var acc = vec2f(0.0);
63 var l = m;
64 loop {
65 if (ny == 0) {
66 acc += y0 * qlm[base + (l - m)];
67 if (l + 1u <= LMAX) {
68 acc += y1 * qlm[base + (l + 1u - m)];
69 }
70 } else if (abs(y0) > RESCALE_THR) {
71 ny += 1;
72 y0 *= INV_SCALE;
73 y1 *= INV_SCALE;
74 }
75 if (l + 2u > LMAX) { break; }
76 let c0 = ab[base + (l + 2u - m)];
77 y0 = c0.x * ct * y1 + c0.y * y0;
78 if (l + 3u <= LMAX) {
79 let c1 = ab[base + (l + 3u - m)];
80 y1 = c1.x * ct * y0 + c1.y * y1;
81 }
82 l += 2u;
83 }
84 fm[m * NLAT + ilat] = acc;
86`;
89export function legAnalysWGSL(p: LegParams): string {
90 const K = Math.ceil(p.nlat / p.wgAnalys); // latitudes per thread
91 return /* wgsl */ `
92${RESCALE_WGSL}
93const LMAX: u32 = ${p.lmax}u;
94const NLAT: u32 = ${p.nlat}u;
95const WG: u32 = ${p.wgAnalys}u;
96const K: u32 = ${K}u;
97${BINDINGS}
98@group(0) @binding(3) var<storage, read> fm: array<vec2f>; // [(m)*NLAT + ilat]
99@group(0) @binding(4) var<storage, read_write> qout: array<vec2f>;
101var<workgroup> red: array<vec4f, ${p.wgAnalys}>;
103@compute @workgroup_size(${p.wgAnalys})
104fn leg_analys(@builtin(local_invocation_id) lid3: vec3u,
105 @builtin(workgroup_id) wid: vec3u) {
106 let lid = lid3.x;
107 let m = wid.x;
108 let base = m * (LMAX + 1u) - (m * (m - 1u)) / 2u;
110 // per-thread recurrence state for K latitudes
111 var y0v: array<f32, ${K}>;
112 var y1v: array<f32, ${K}>;
113 var nyv: array<i32, ${K}>;
114 var ctv: array<f32, ${K}>;
115 var wfv: array<vec2f, ${K}>;
117 for (var k = 0u; k < K; k++) {
118 let lat = lid + k * WG;
119 var ct: f32 = 0.0;
120 var st: f32 = 0.0;
121 var wf = vec2f(0.0);
122 if (lat < NLAT) {
123 ct = ctstw[lat];
124 st = ctstw[NLAT + lat];
125 wf = fm[m * NLAT + lat] * ctstw[2u * NLAT + lat]; // Gauss weight (incl. 2*pi/nphi)
126 }
127 ctv[k] = ct;
128 let seed = sinpow_rescaled(st, m);
129 y0v[k] = seed.y0 * amm[m];
130 nyv[k] = seed.ny;
131 y1v[k] = 0.0;
132 if (m < LMAX) {
133 y1v[k] = ab[base + 1u].x * ct * y0v[k];
134 }
135 wfv[k] = wf;
136 }
138 var l = m;
139 loop {
140 var c0 = vec2f(0.0);
141 var c1 = vec2f(0.0);
142 for (var k = 0u; k < K; k++) {
143 if (nyv[k] == 0) {
144 c0 += wfv[k] * y0v[k];
145 c1 += wfv[k] * y1v[k];
146 } else if (abs(y0v[k]) > RESCALE_THR) {
147 nyv[k] += 1;
148 y0v[k] *= INV_SCALE;
149 y1v[k] *= INV_SCALE;
150 }
151 }
152 // workgroup tree reduction of (c0, c1)
153 red[lid] = vec4f(c0, c1);
154 workgroupBarrier();
155 var s = WG / 2u;
156 while (s > 0u) {
157 if (lid < s) { red[lid] += red[lid + s]; }
158 workgroupBarrier();
159 s = s >> 1u;
160 }
161 if (lid == 0u) {
162 qout[base + (l - m)] = red[0].xy;
163 if (l + 1u <= LMAX) {
164 qout[base + (l + 1u - m)] = red[0].zw;
165 }
166 }
167 if (l + 2u > LMAX) { break; }
168 let a0 = ab[base + (l + 2u - m)];
169 var a1 = vec2f(0.0);
170 if (l + 3u <= LMAX) {
171 a1 = ab[base + (l + 3u - m)];
172 }
173 for (var k = 0u; k < K; k++) {
174 let t0 = a0.x * ctv[k] * y1v[k] + a0.y * y0v[k];
175 y0v[k] = t0;
176 y1v[k] = a1.x * ctv[k] * t0 + a1.y * y1v[k];
177 }
178 l += 2u;
179 }
181`;
moveopenescclose