|
| 1 | +//! AVX2+FMA kernels for x86_64. |
| 2 | +
|
| 3 | +#![cfg(target_arch = "x86_64")] |
| 4 | + |
| 5 | +pub fn l2_squared(a: &[f32], b: &[f32]) -> f32 { |
| 6 | + assert_eq!(a.len(), b.len(), "avx2 l2: length mismatch"); |
| 7 | + // SAFETY: caller verified avx2+fma via is_x86_feature_detected. |
| 8 | + unsafe { l2_squared_impl(a, b) } |
| 9 | +} |
| 10 | + |
| 11 | +#[target_feature(enable = "avx2,fma")] |
| 12 | +unsafe fn l2_squared_impl(a: &[f32], b: &[f32]) -> f32 { |
| 13 | + assert_eq!(a.len(), b.len(), "avx2 l2_impl: length mismatch"); |
| 14 | + unsafe { |
| 15 | + use std::arch::x86_64::*; |
| 16 | + let n = a.len(); |
| 17 | + let mut sum = _mm256_setzero_ps(); |
| 18 | + let chunks = n / 8; |
| 19 | + for i in 0..chunks { |
| 20 | + let off = i * 8; |
| 21 | + let va = _mm256_loadu_ps(a.as_ptr().add(off)); |
| 22 | + let vb = _mm256_loadu_ps(b.as_ptr().add(off)); |
| 23 | + let diff = _mm256_sub_ps(va, vb); |
| 24 | + sum = _mm256_fmadd_ps(diff, diff, sum); |
| 25 | + } |
| 26 | + let mut result = hsum256(sum); |
| 27 | + for i in (chunks * 8)..n { |
| 28 | + let d = a[i] - b[i]; |
| 29 | + result += d * d; |
| 30 | + } |
| 31 | + result |
| 32 | + } |
| 33 | +} |
| 34 | + |
| 35 | +pub fn cosine_distance(a: &[f32], b: &[f32]) -> f32 { |
| 36 | + assert_eq!(a.len(), b.len(), "avx2 cosine: length mismatch"); |
| 37 | + unsafe { cosine_impl(a, b) } |
| 38 | +} |
| 39 | + |
| 40 | +#[target_feature(enable = "avx2,fma")] |
| 41 | +unsafe fn cosine_impl(a: &[f32], b: &[f32]) -> f32 { |
| 42 | + assert_eq!(a.len(), b.len(), "avx2 cosine_impl: length mismatch"); |
| 43 | + unsafe { |
| 44 | + use std::arch::x86_64::*; |
| 45 | + let n = a.len(); |
| 46 | + let mut vdot = _mm256_setzero_ps(); |
| 47 | + let mut vna = _mm256_setzero_ps(); |
| 48 | + let mut vnb = _mm256_setzero_ps(); |
| 49 | + let chunks = n / 8; |
| 50 | + for i in 0..chunks { |
| 51 | + let off = i * 8; |
| 52 | + let va = _mm256_loadu_ps(a.as_ptr().add(off)); |
| 53 | + let vb = _mm256_loadu_ps(b.as_ptr().add(off)); |
| 54 | + vdot = _mm256_fmadd_ps(va, vb, vdot); |
| 55 | + vna = _mm256_fmadd_ps(va, va, vna); |
| 56 | + vnb = _mm256_fmadd_ps(vb, vb, vnb); |
| 57 | + } |
| 58 | + let mut dot = hsum256(vdot); |
| 59 | + let mut na = hsum256(vna); |
| 60 | + let mut nb = hsum256(vnb); |
| 61 | + for i in (chunks * 8)..n { |
| 62 | + dot += a[i] * b[i]; |
| 63 | + na += a[i] * a[i]; |
| 64 | + nb += b[i] * b[i]; |
| 65 | + } |
| 66 | + let denom = (na * nb).sqrt(); |
| 67 | + if denom < f32::EPSILON { |
| 68 | + 1.0 |
| 69 | + } else { |
| 70 | + (1.0 - dot / denom).max(0.0) |
| 71 | + } |
| 72 | + } |
| 73 | +} |
| 74 | + |
| 75 | +pub fn neg_inner_product(a: &[f32], b: &[f32]) -> f32 { |
| 76 | + assert_eq!(a.len(), b.len(), "avx2 ip: length mismatch"); |
| 77 | + unsafe { ip_impl(a, b) } |
| 78 | +} |
| 79 | + |
| 80 | +#[target_feature(enable = "avx2,fma")] |
| 81 | +unsafe fn ip_impl(a: &[f32], b: &[f32]) -> f32 { |
| 82 | + assert_eq!(a.len(), b.len(), "avx2 ip_impl: length mismatch"); |
| 83 | + unsafe { |
| 84 | + use std::arch::x86_64::*; |
| 85 | + let n = a.len(); |
| 86 | + let mut vdot = _mm256_setzero_ps(); |
| 87 | + let chunks = n / 8; |
| 88 | + for i in 0..chunks { |
| 89 | + let off = i * 8; |
| 90 | + let va = _mm256_loadu_ps(a.as_ptr().add(off)); |
| 91 | + let vb = _mm256_loadu_ps(b.as_ptr().add(off)); |
| 92 | + vdot = _mm256_fmadd_ps(va, vb, vdot); |
| 93 | + } |
| 94 | + let mut dot = hsum256(vdot); |
| 95 | + for i in (chunks * 8)..n { |
| 96 | + dot += a[i] * b[i]; |
| 97 | + } |
| 98 | + -dot |
| 99 | + } |
| 100 | +} |
| 101 | + |
| 102 | +/// Horizontal sum of 8 × f32 in a __m256. |
| 103 | +#[target_feature(enable = "avx2")] |
| 104 | +unsafe fn hsum256(v: std::arch::x86_64::__m256) -> f32 { |
| 105 | + use std::arch::x86_64::*; |
| 106 | + let hi = _mm256_extractf128_ps(v, 1); |
| 107 | + let lo = _mm256_castps256_ps128(v); |
| 108 | + let sum128 = _mm_add_ps(lo, hi); |
| 109 | + let shuf = _mm_movehdup_ps(sum128); |
| 110 | + let sums = _mm_add_ps(sum128, shuf); |
| 111 | + let shuf2 = _mm_movehl_ps(sums, sums); |
| 112 | + let sums2 = _mm_add_ss(sums, shuf2); |
| 113 | + _mm_cvtss_f32(sums2) |
| 114 | +} |
0 commit comments