1use 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
37pub struct SimdKernels {
41 pub dot_f32: unsafe extern "C" fn(*const f32, *const f32, u64) -> f32,
43 pub l2sq_f32: unsafe extern "C" fn(*const f32, *const f32, u64) -> f32,
45 pub add_f32: unsafe extern "C" fn(*const f32, *const f32, *mut f32, u64),
47 pub scale_f32: unsafe extern "C" fn(*const f32, f32, *mut f32, u64),
49 _module: JITModule,
50}
51
52unsafe impl Send for SimdKernels {}
56unsafe impl Sync for SimdKernels {}
57
58static KERNELS: OnceLock<Option<SimdKernels>> = OnceLock::new();
59
60pub fn kernels() -> Option<&'static SimdKernels> {
64 KERNELS.get_or_init(|| compile_kernels().ok()).as_ref()
65}
66
67#[derive(Clone, Copy, PartialEq)]
69enum KernelKind {
70 Dot,
72 L2Sq,
74 Add,
76 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 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 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
131fn 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)); sig.params.push(AbiParam::new(types::I64)); sig.params.push(AbiParam::new(types::I64)); sig.returns.push(AbiParam::new(types::F32));
148 }
149 KernelKind::Add => {
150 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)); }
155 KernelKind::Scale => {
156 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)); }
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 let mf = ir::MemFlags::new();
187 let reducing = matches!(kind, KernelKind::Dot | KernelKind::L2Sq);
188
189 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 let vk = if kind == KernelKind::Scale {
196 Some(b.ins().splat(types::F32X4, b_or_k))
197 } else {
198 None
199 };
200
201 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); 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 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 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 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 (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 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}