Skip to main content

polydat_core/compile/jit/
simd.rs

1// Copyright 2024-2026 Jonathan Shook
2// SPDX-License-Identifier: Apache-2.0
3
4//! Cranelift-SIMD compute kernels for element-wise vector math
5//! (type_system_alignment.md §4 heap-slice plane, §7 execution
6//! eligibility).
7//!
8//! Four f32-lane kernels are compiled once per process through the
9//! same host ISA the graph JIT uses (`host_isa::build_host_isa`);
10//! SIMD types are unconditional in Cranelift IR. Each kernel
11//! processes the slice body in
12//! `F32X4` chunks (unaligned 128-bit loads — `SliceArc<f32>` data
13//! is only 4-aligned) and finishes the remainder in a scalar loop:
14//!
15//! - `dot_f32(a, b, len) -> f32` — F32X4 multiply-accumulate,
16//!   horizontal reduce, scalar tail.
17//! - `l2sq_f32(a, b, len) -> f32` — squared-difference
18//!   accumulate; callers take the square root.
19//! - `add_f32(a, b, out, len)` — element-wise sum.
20//! - `scale_f32(a, k, out, len)` — scalar broadcast multiply
21//!   (`splat`).
22//!
23//! Consumers are the f32 bodies in `crate::numeric::vector`, shared
24//! by the `vec_*` nodes (polydat-nodes `vector_math`) and by the
25//! JIT's vector helpers; each body falls back to its scalar loop
26//! when the `jit` feature is off or the kernels failed to compile.
27//! SIMD accumulation reassociates floating-point addition, so
28//! results may differ from the scalar reference in the final
29//! ulps; the equivalence tests compare with relative tolerance.
30
31use std::sync::OnceLock;
32
33use cranelift_codegen::ir::condcodes::IntCC;
34use cranelift_codegen::ir::{self, AbiParam, InstBuilder, types};
35use cranelift_codegen::settings::{self, Configurable};
36use cranelift_frontend::{FunctionBuilder, FunctionBuilderContext};
37use cranelift_jit::{JITBuilder, JITModule};
38use cranelift_module::{Linkage, Module};
39
40/// Finalized SIMD kernel entry points. The owning [`JITModule`] is
41/// kept alive alongside the pointers (dropping it would unmap the
42/// code pages).
43pub struct SimdKernels {
44    /// Dot product of two `f32` slices of the given length.
45    pub dot_f32: unsafe extern "C" fn(*const f32, *const f32, u64) -> f32,
46    /// Squared L2 distance between two `f32` slices of the given length.
47    pub l2sq_f32: unsafe extern "C" fn(*const f32, *const f32, u64) -> f32,
48    /// Element-wise sum of two `f32` slices into the output slice.
49    pub add_f32: unsafe extern "C" fn(*const f32, *const f32, *mut f32, u64),
50    /// Each element of an `f32` slice times a scalar, into the output slice.
51    pub scale_f32: unsafe extern "C" fn(*const f32, f32, *mut f32, u64),
52    _module: JITModule,
53}
54
55// SAFETY: the function pointers are immutable after construction
56// and the generated code is reentrant (no globals); JITModule is
57// only held to keep the mapping alive.
58unsafe impl Send for SimdKernels {}
59unsafe impl Sync for SimdKernels {}
60
61static KERNELS: OnceLock<Option<SimdKernels>> = OnceLock::new();
62
63/// The process-wide SIMD kernel set, compiled on first use.
64/// `None` when the host ISA cannot be built or the kernels fail
65/// to compile — callers fall back to their scalar loops.
66pub fn kernels() -> Option<&'static SimdKernels> {
67    KERNELS.get_or_init(|| compile_kernels().ok()).as_ref()
68}
69
70/// Which reduction/elementwise body a kernel uses.
71#[derive(Clone, Copy, PartialEq)]
72enum KernelKind {
73    /// acc += a[i] * b[i]; returns acc.
74    Dot,
75    /// d = a[i] - b[i]; acc += d * d; returns acc.
76    L2Sq,
77    /// out[i] = a[i] + b[i].
78    Add,
79    /// out[i] = a[i] * k.
80    Scale,
81}
82
83fn compile_kernels() -> Result<SimdKernels, String> {
84    let mut flag_builder = settings::builder();
85    flag_builder.set("opt_level", "speed").unwrap();
86    // Use the same native feature inference and ISA fingerprint as the graph
87    // JIT. SIMD types are unconditional in Cranelift IR, but their concrete
88    // lowering and instruction selection are target-feature dependent.
89    let isa = super::host_isa::build_host_isa(flag_builder)?;
90
91    let jit_builder = JITBuilder::with_isa(isa, cranelift_module::default_libcall_names());
92    let mut module = JITModule::new(jit_builder);
93    let mut fn_ctx = FunctionBuilderContext::new();
94
95    let dot_id = build_kernel(&mut module, &mut fn_ctx, "simd_dot_f32", KernelKind::Dot)?;
96    let l2_id = build_kernel(&mut module, &mut fn_ctx, "simd_l2sq_f32", KernelKind::L2Sq)?;
97    let add_id = build_kernel(&mut module, &mut fn_ctx, "simd_add_f32", KernelKind::Add)?;
98    let scale_id = build_kernel(
99        &mut module,
100        &mut fn_ctx,
101        "simd_scale_f32",
102        KernelKind::Scale,
103    )?;
104
105    module
106        .finalize_definitions()
107        .map_err(|e| format!("finalize: {e}"))?;
108
109    // SAFETY: signatures below match exactly what build_kernel
110    // declared for each kind.
111    unsafe {
112        Ok(SimdKernels {
113            dot_f32: std::mem::transmute::<
114                *const u8,
115                unsafe extern "C" fn(*const f32, *const f32, u64) -> f32,
116            >(module.get_finalized_function(dot_id)),
117            l2sq_f32: std::mem::transmute::<
118                *const u8,
119                unsafe extern "C" fn(*const f32, *const f32, u64) -> f32,
120            >(module.get_finalized_function(l2_id)),
121            add_f32: std::mem::transmute::<
122                *const u8,
123                unsafe extern "C" fn(*const f32, *const f32, *mut f32, u64),
124            >(module.get_finalized_function(add_id)),
125            scale_f32: std::mem::transmute::<
126                *const u8,
127                unsafe extern "C" fn(*const f32, f32, *mut f32, u64),
128            >(module.get_finalized_function(scale_id)),
129            _module: module,
130        })
131    }
132}
133
134/// Build one kernel function. Reducing kinds (Dot/L2Sq) take
135/// `(a: i64, b: i64, len: i64) -> f32`; element-wise kinds write
136/// through an out pointer and return nothing — Add is
137/// `(a, b, out, len)`, Scale is `(a, k: f32, out, len)`.
138fn build_kernel(
139    module: &mut JITModule,
140    fn_ctx: &mut FunctionBuilderContext,
141    name: &str,
142    kind: KernelKind,
143) -> Result<cranelift_module::FuncId, String> {
144    let mut sig = module.make_signature();
145    match kind {
146        KernelKind::Dot | KernelKind::L2Sq => {
147            sig.params.push(AbiParam::new(types::I64)); // a
148            sig.params.push(AbiParam::new(types::I64)); // b
149            sig.params.push(AbiParam::new(types::I64)); // len
150            sig.returns.push(AbiParam::new(types::F32));
151        }
152        KernelKind::Add => {
153            sig.params.push(AbiParam::new(types::I64)); // a
154            sig.params.push(AbiParam::new(types::I64)); // b
155            sig.params.push(AbiParam::new(types::I64)); // out
156            sig.params.push(AbiParam::new(types::I64)); // len
157        }
158        KernelKind::Scale => {
159            sig.params.push(AbiParam::new(types::I64)); // a
160            sig.params.push(AbiParam::new(types::F32)); // k
161            sig.params.push(AbiParam::new(types::I64)); // out
162            sig.params.push(AbiParam::new(types::I64)); // len
163        }
164    }
165
166    let func_id = module
167        .declare_function(name, Linkage::Local, &sig)
168        .map_err(|e| format!("declare {name}: {e}"))?;
169
170    let mut ctx = module.make_context();
171    ctx.func.signature = sig;
172
173    {
174        let mut b = FunctionBuilder::new(&mut ctx.func, fn_ctx);
175        let entry = b.create_block();
176        b.append_block_params_for_function_params(entry);
177        b.switch_to_block(entry);
178        b.seal_block(entry);
179
180        let params: Vec<ir::Value> = b.block_params(entry).to_vec();
181        let (a_ptr, b_or_k, out_ptr, len) = match kind {
182            KernelKind::Dot | KernelKind::L2Sq => (params[0], params[1], None, params[2]),
183            KernelKind::Add => (params[0], params[1], Some(params[2]), params[3]),
184            KernelKind::Scale => (params[0], params[1], Some(params[2]), params[3]),
185        };
186
187        // Unaligned 128-bit access: SliceArc<f32> data is only
188        // guaranteed 4-aligned.
189        let mf = ir::MemFlags::new();
190        let reducing = matches!(kind, KernelKind::Dot | KernelKind::L2Sq);
191
192        // n4 = len & !3 — the SIMD-chunked prefix.
193        let n4 = b.ins().band_imm(len, !3i64);
194        let zero_i = b.ins().iconst(types::I64, 0);
195        let zero_f = b.ins().f32const(0.0);
196        let vzero = b.ins().splat(types::F32X4, zero_f);
197        // Scale's broadcast operand.
198        let vk = if kind == KernelKind::Scale {
199            Some(b.ins().splat(types::F32X4, b_or_k))
200        } else {
201            None
202        };
203
204        // ── Vector loop ──
205        // head(i: i64, acc: f32x4) — acc unused (carried zero) for
206        // the element-wise kinds; keeping one block shape for all
207        // four kernels keeps this builder small.
208        let vhead = b.create_block();
209        b.append_block_param(vhead, types::I64);
210        b.append_block_param(vhead, types::F32X4);
211        let vbody = b.create_block();
212        let vexit = b.create_block();
213        b.append_block_param(vexit, types::I64);
214        b.append_block_param(vexit, types::F32X4);
215
216        b.ins().jump(vhead, &[zero_i, vzero]);
217
218        b.switch_to_block(vhead);
219        let vi = b.block_params(vhead)[0];
220        let vacc = b.block_params(vhead)[1];
221        let done4 = b.ins().icmp(IntCC::UnsignedGreaterThanOrEqual, vi, n4);
222        b.ins().brif(done4, vexit, &[vi, vacc], vbody, &[]);
223
224        b.switch_to_block(vbody);
225        let byte_off = b.ins().ishl_imm(vi, 2); // i * sizeof(f32)
226        let a_addr = b.ins().iadd(a_ptr, byte_off);
227        let va = b.ins().load(types::F32X4, mf, a_addr, 0);
228        let (next_acc, store_val) = match kind {
229            KernelKind::Dot => {
230                let b_addr = b.ins().iadd(b_or_k, byte_off);
231                let vb = b.ins().load(types::F32X4, mf, b_addr, 0);
232                let prod = b.ins().fmul(va, vb);
233                (b.ins().fadd(vacc, prod), None)
234            }
235            KernelKind::L2Sq => {
236                let b_addr = b.ins().iadd(b_or_k, byte_off);
237                let vb = b.ins().load(types::F32X4, mf, b_addr, 0);
238                let d = b.ins().fsub(va, vb);
239                let sq = b.ins().fmul(d, d);
240                (b.ins().fadd(vacc, sq), None)
241            }
242            KernelKind::Add => {
243                let b_addr = b.ins().iadd(b_or_k, byte_off);
244                let vb = b.ins().load(types::F32X4, mf, b_addr, 0);
245                (vacc, Some(b.ins().fadd(va, vb)))
246            }
247            KernelKind::Scale => (vacc, Some(b.ins().fmul(va, vk.unwrap()))),
248        };
249        if let Some(v) = store_val {
250            let out_addr = b.ins().iadd(out_ptr.unwrap(), byte_off);
251            b.ins().store(mf, v, out_addr, 0);
252        }
253        let vi_next = b.ins().iadd_imm(vi, 4);
254        b.ins().jump(vhead, &[vi_next, next_acc]);
255        b.seal_block(vhead);
256        b.seal_block(vbody);
257
258        // ── Horizontal reduce (reducing kinds) ──
259        b.switch_to_block(vexit);
260        b.seal_block(vexit);
261        let ti = b.block_params(vexit)[0];
262        let facc = b.block_params(vexit)[1];
263        let red = if reducing {
264            let l0 = b.ins().extractlane(facc, 0);
265            let l1 = b.ins().extractlane(facc, 1);
266            let l2 = b.ins().extractlane(facc, 2);
267            let l3 = b.ins().extractlane(facc, 3);
268            let s01 = b.ins().fadd(l0, l1);
269            let s23 = b.ins().fadd(l2, l3);
270            b.ins().fadd(s01, s23)
271        } else {
272            zero_f
273        };
274
275        // ── Scalar tail loop: head(i: i64, s: f32) ──
276        let shead = b.create_block();
277        b.append_block_param(shead, types::I64);
278        b.append_block_param(shead, types::F32);
279        let sbody = b.create_block();
280        let sexit = b.create_block();
281        b.append_block_param(sexit, types::F32);
282
283        b.ins().jump(shead, &[ti, red]);
284
285        b.switch_to_block(shead);
286        let si = b.block_params(shead)[0];
287        let ss = b.block_params(shead)[1];
288        let done = b.ins().icmp(IntCC::UnsignedGreaterThanOrEqual, si, len);
289        b.ins().brif(done, sexit, &[ss], sbody, &[]);
290
291        b.switch_to_block(sbody);
292        let s_off = b.ins().ishl_imm(si, 2);
293        let sa_addr = b.ins().iadd(a_ptr, s_off);
294        let sa = b.ins().load(types::F32, mf, sa_addr, 0);
295        let s_next = match kind {
296            KernelKind::Dot => {
297                let sb_addr = b.ins().iadd(b_or_k, s_off);
298                let sb = b.ins().load(types::F32, mf, sb_addr, 0);
299                let p = b.ins().fmul(sa, sb);
300                b.ins().fadd(ss, p)
301            }
302            KernelKind::L2Sq => {
303                let sb_addr = b.ins().iadd(b_or_k, s_off);
304                let sb = b.ins().load(types::F32, mf, sb_addr, 0);
305                let d = b.ins().fsub(sa, sb);
306                let sq = b.ins().fmul(d, d);
307                b.ins().fadd(ss, sq)
308            }
309            KernelKind::Add => {
310                let sb_addr = b.ins().iadd(b_or_k, s_off);
311                let sb = b.ins().load(types::F32, mf, sb_addr, 0);
312                let v = b.ins().fadd(sa, sb);
313                let out_addr = b.ins().iadd(out_ptr.unwrap(), s_off);
314                b.ins().store(mf, v, out_addr, 0);
315                ss
316            }
317            KernelKind::Scale => {
318                let v = b.ins().fmul(sa, b_or_k);
319                let out_addr = b.ins().iadd(out_ptr.unwrap(), s_off);
320                b.ins().store(mf, v, out_addr, 0);
321                ss
322            }
323        };
324        let si_next = b.ins().iadd_imm(si, 1);
325        b.ins().jump(shead, &[si_next, s_next]);
326        b.seal_block(shead);
327        b.seal_block(sbody);
328
329        b.switch_to_block(sexit);
330        b.seal_block(sexit);
331        let result = b.block_params(sexit)[0];
332        if reducing {
333            b.ins().return_(&[result]);
334        } else {
335            b.ins().return_(&[]);
336        }
337        b.finalize();
338    }
339
340    module
341        .define_function(func_id, &mut ctx)
342        .map_err(|e| format!("define {name}: {e}"))?;
343    module.clear_context(&mut ctx);
344    Ok(func_id)
345}
346
347#[cfg(test)]
348mod tests {
349    /// Deterministic pseudo-random test vectors (no Math::random
350    /// in tests — keep them replayable).
351    fn test_vec(n: usize, seed: u64) -> Vec<f32> {
352        (0..n)
353            .map(|i| {
354                let h = xxhash_rust::xxh3::xxh3_64(&(seed ^ i as u64).to_le_bytes());
355                // map to [-1, 1)
356                (h as f64 / u64::MAX as f64 * 2.0 - 1.0) as f32
357            })
358            .collect()
359    }
360
361    fn assert_close(simd: f32, reference: f32, what: &str) {
362        let denom = reference.abs().max(1e-6);
363        let rel = ((simd - reference).abs()) / denom;
364        assert!(
365            rel < 1e-4,
366            "{what}: simd={simd} reference={reference} rel_err={rel}"
367        );
368    }
369
370    #[test]
371    fn simd_kernels_match_scalar_reference() {
372        let Some(k) = super::kernels() else {
373            panic!(
374                "SIMD kernels failed to compile on this host — \
375                    the host ISA should compile the SIMD kernels"
376            );
377        };
378        // Cover: empty, sub-chunk, exact-chunk, chunk+tail sizes.
379        for &n in &[0usize, 3, 8, 1029] {
380            let a = test_vec(n, 0xA);
381            let b = test_vec(n, 0xB);
382
383            let dot_ref: f32 = a.iter().zip(&b).map(|(x, y)| x * y).sum();
384            let l2_ref: f32 = a.iter().zip(&b).map(|(x, y)| (x - y) * (x - y)).sum();
385            let (dot, l2) = unsafe {
386                (
387                    (k.dot_f32)(a.as_ptr(), b.as_ptr(), n as u64),
388                    (k.l2sq_f32)(a.as_ptr(), b.as_ptr(), n as u64),
389                )
390            };
391            assert_close(dot, dot_ref, &format!("dot n={n}"));
392            assert_close(l2, l2_ref, &format!("l2sq n={n}"));
393
394            let mut out = vec![0.0f32; n];
395            unsafe { (k.add_f32)(a.as_ptr(), b.as_ptr(), out.as_mut_ptr(), n as u64) };
396            for i in 0..n {
397                assert_eq!(out[i], a[i] + b[i], "add lane {i} n={n}");
398            }
399            unsafe { (k.scale_f32)(a.as_ptr(), 2.5, out.as_mut_ptr(), n as u64) };
400            for i in 0..n {
401                assert_eq!(out[i], a[i] * 2.5, "scale lane {i} n={n}");
402            }
403        }
404    }
405}