diff --git a/crates/clawhdf5-accel/src/lib.rs b/crates/clawhdf5-accel/src/lib.rs index 6fcd9b3..81a2f03 100644 --- a/crates/clawhdf5-accel/src/lib.rs +++ b/crates/clawhdf5-accel/src/lib.rs @@ -124,11 +124,24 @@ pub fn dot_product(a: &[f32], b: &[f32]) -> f32 { /// Dot product of two `i8` slices, widened to `i32`. /// -/// The kernel behind int8-quantised vector search. Uses the AVX2 path -/// whenever AVX2 is present — including on AVX-512 machines, where it is -/// what the f32 kernels use too on a default build. +/// The kernel behind int8-quantised vector search. On x86-64 it uses the AVX2 +/// path whenever AVX2 is present (including on AVX-512 machines, where it is +/// what the f32 kernels use too on a default build). On aarch64 it uses the +/// ARMv8.2 `SDOT` instruction when the CPU has the dot-product extension, and +/// plain NEON otherwise. pub fn dot_i8(a: &[i8], b: &[i8]) -> i32 { match detect_backend() { + #[cfg(target_arch = "aarch64")] + Backend::Neon => { + if std::arch::is_aarch64_feature_detected!("dotprod") { + // SAFETY: the dotprod extension was just detected at runtime. + unsafe { neon::dot_i8_dotprod(a, b) } + } else { + // SAFETY: NEON is always available on aarch64. + unsafe { neon::dot_i8(a, b) } + } + } + #[cfg(target_arch = "x86_64")] // SAFETY: both variants imply AVX2 was detected at runtime (the // AVX-512 backend is only selected on CPUs that also have AVX2). @@ -760,6 +773,42 @@ mod dot_i8_tests { } } + /// Dispatch only ever takes one path on a given CPU, so on a machine with + /// the dot-product extension the plain-NEON kernel would otherwise go + /// untested. Check each aarch64 kernel against scalar directly. + #[cfg(target_arch = "aarch64")] + #[test] + fn every_aarch64_kernel_matches_scalar_exactly() { + for len in [0, 1, 7, 15, 16, 17, 31, 32, 33, 63, 64, 100, 384, 385, 1536] { + let a = codes(len, 7 + len as u64); + let b = codes(len, 7000 + len as u64); + let want = scalar::dot_i8(&a, &b); + // SAFETY: NEON is always available on aarch64. + assert_eq!(unsafe { neon::dot_i8(&a, &b) }, want, "neon, len {len}"); + if std::arch::is_aarch64_feature_detected!("dotprod") { + // SAFETY: the dotprod extension was just detected. + assert_eq!( + unsafe { neon::dot_i8_dotprod(&a, &b) }, + want, + "dotprod, len {len}" + ); + } + } + // The extremes, through both kernels. + let lo = vec![-128i8; 4096]; + let hi = vec![127i8; 4096]; + // SAFETY: NEON is always available on aarch64. + assert_eq!(unsafe { neon::dot_i8(&lo, &lo) }, 4096 * 128 * 128); + // SAFETY: NEON is always available on aarch64. + assert_eq!(unsafe { neon::dot_i8(&lo, &hi) }, -4096 * 128 * 127); + if std::arch::is_aarch64_feature_detected!("dotprod") { + // SAFETY: the dotprod extension was just detected. + assert_eq!(unsafe { neon::dot_i8_dotprod(&lo, &lo) }, 4096 * 128 * 128); + // SAFETY: the dotprod extension was just detected. + assert_eq!(unsafe { neon::dot_i8_dotprod(&lo, &hi) }, -4096 * 128 * 127); + } + } + #[test] fn extremes_do_not_overflow() { // -128 * -128 is the largest product; a long run of it must still fit. diff --git a/crates/clawhdf5-accel/src/neon.rs b/crates/clawhdf5-accel/src/neon.rs index 495955f..1e195aa 100644 --- a/crates/clawhdf5-accel/src/neon.rs +++ b/crates/clawhdf5-accel/src/neon.rs @@ -180,3 +180,130 @@ pub fn checksum_fletcher32(data: &[u8]) -> u32 { (sum2 << 16) | sum1 } + +/// NEON dot product of two `i8` slices, widened to `i32`, for any aarch64 CPU. +/// +/// `vmull_s8` multiplies eight lanes into `i16` — even `-128 * -128` is 16 384, +/// inside `i16` — and `vpadalq_s16` adds adjacent pairs of those into `i32` +/// accumulators, so nothing can overflow before the final horizontal sum. +/// +/// CPUs with the ARMv8.2 dot-product extension should use +/// [`dot_i8_dotprod`], which does the multiply and the accumulate in one +/// instruction. +/// +/// # Safety +/// Caller must ensure aarch64 target (NEON always available). +// SAFETY: NEON is always available on aarch64 targets; caller guarantees aarch64. +#[target_feature(enable = "neon")] +pub unsafe fn dot_i8(a: &[i8], b: &[i8]) -> i32 { + assert_eq!(a.len(), b.len()); + let len = a.len(); + let mut i = 0; + let mut acc0 = vdupq_n_s32(0); + let mut acc1 = vdupq_n_s32(0); + + while i + 16 <= len { + // SAFETY: NEON is available per the # Safety contract, and both + // 16-byte loads start at an index checked against `len` above. + unsafe { + let va = vld1q_s8(a.as_ptr().add(i)); + let vb = vld1q_s8(b.as_ptr().add(i)); + acc0 = vpadalq_s16(acc0, vmull_s8(vget_low_s8(va), vget_low_s8(vb))); + acc1 = vpadalq_s16(acc1, vmull_high_s8(va, vb)); + } + i += 16; + } + + let mut sum = vaddvq_s32(vaddq_s32(acc0, acc1)); + while i < len { + sum += i32::from(a[i]) * i32::from(b[i]); + i += 1; + } + sum +} + +/// One `SDOT`: for each of the four `i32` lanes of `acc`, add the dot +/// product of the corresponding four `i8` pairs from `a` and `b`. +/// +/// Written as inline assembly because the `vdotq_s32` intrinsic is still +/// behind the unstable `stdarch_neon_dotprod` feature; inline assembly is +/// stable on aarch64. +/// +/// # Safety +/// Caller must ensure the CPU supports the `dotprod` extension. +#[inline] +#[target_feature(enable = "neon,dotprod")] +unsafe fn sdot(acc: int32x4_t, a: int8x16_t, b: int8x16_t) -> int32x4_t { + let mut acc = acc; + // SAFETY: `dotprod` is enabled for this function and the caller + // guarantees the CPU supports it. The instruction reads only its three + // vector registers and touches no memory. + unsafe { + std::arch::asm!( + "sdot {acc:v}.4s, {a:v}.16b, {b:v}.16b", + acc = inout(vreg) acc, + a = in(vreg) a, + b = in(vreg) b, + options(pure, nomem, nostack), + ); + } + acc +} + +/// NEON dot product of two `i8` slices using the ARMv8.2 dot-product +/// extension (`SDOT`): sixteen multiply-accumulates per instruction, straight +/// into `i32` lanes. +/// +/// Present on the cores this crate actually runs on — Cortex-A76 and later +/// (Raspberry Pi 5, current Android phones), Neoverse-N1 (Graviton2, Ampere +/// Altra), and every Apple Silicon generation. +/// +/// # Safety +/// Caller must verify `is_aarch64_feature_detected!("dotprod")`. +// SAFETY: caller has verified the dotprod extension at runtime. +#[target_feature(enable = "neon,dotprod")] +pub unsafe fn dot_i8_dotprod(a: &[i8], b: &[i8]) -> i32 { + assert_eq!(a.len(), b.len()); + let len = a.len(); + let mut i = 0; + let mut acc0 = vdupq_n_s32(0); + let mut acc1 = vdupq_n_s32(0); + + // Two independent accumulators so consecutive SDOTs are not serialised on + // one register. + while i + 32 <= len { + // SAFETY: dotprod is available per the # Safety contract, and every + // 16-byte load starts at an index checked against `len` above. + unsafe { + acc0 = sdot( + acc0, + vld1q_s8(a.as_ptr().add(i)), + vld1q_s8(b.as_ptr().add(i)), + ); + acc1 = sdot( + acc1, + vld1q_s8(a.as_ptr().add(i + 16)), + vld1q_s8(b.as_ptr().add(i + 16)), + ); + } + i += 32; + } + if i + 16 <= len { + // SAFETY: as above; the load is bounds-checked by this condition. + unsafe { + acc0 = sdot( + acc0, + vld1q_s8(a.as_ptr().add(i)), + vld1q_s8(b.as_ptr().add(i)), + ); + } + i += 16; + } + + let mut sum = vaddvq_s32(vaddq_s32(acc0, acc1)); + while i < len { + sum += i32::from(a[i]) * i32::from(b[i]); + i += 1; + } + sum +}