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