1use 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
40pub struct SimdKernels {
44 pub dot_f32: unsafe extern "C" fn(*const f32, *const f32, u64) -> f32,
46 pub l2sq_f32: unsafe extern "C" fn(*const f32, *const f32, u64) -> f32,
48 pub add_f32: unsafe extern "C" fn(*const f32, *const f32, *mut f32, u64),
50 pub scale_f32: unsafe extern "C" fn(*const f32, f32, *mut f32, u64),
52 _module: JITModule,
53}
54
55unsafe impl Send for SimdKernels {}
59unsafe impl Sync for SimdKernels {}
60
61static KERNELS: OnceLock<Option<SimdKernels>> = OnceLock::new();
62
63pub fn kernels() -> Option<&'static SimdKernels> {
67 KERNELS.get_or_init(|| compile_kernels().ok()).as_ref()
68}
69
70#[derive(Clone, Copy, PartialEq)]
72enum KernelKind {
73 Dot,
75 L2Sq,
77 Add,
79 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 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 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
134fn 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)); sig.params.push(AbiParam::new(types::I64)); sig.params.push(AbiParam::new(types::I64)); sig.returns.push(AbiParam::new(types::F32));
151 }
152 KernelKind::Add => {
153 sig.params.push(AbiParam::new(types::I64)); sig.params.push(AbiParam::new(types::I64)); sig.params.push(AbiParam::new(types::I64)); sig.params.push(AbiParam::new(types::I64)); }
158 KernelKind::Scale => {
159 sig.params.push(AbiParam::new(types::I64)); sig.params.push(AbiParam::new(types::F32)); sig.params.push(AbiParam::new(types::I64)); sig.params.push(AbiParam::new(types::I64)); }
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 let mf = ir::MemFlags::new();
190 let reducing = matches!(kind, KernelKind::Dot | KernelKind::L2Sq);
191
192 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 let vk = if kind == KernelKind::Scale {
199 Some(b.ins().splat(types::F32X4, b_or_k))
200 } else {
201 None
202 };
203
204 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); 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 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 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 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 (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 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}