(a: &[f32], b: &[f32])
| 12 | |
| 13 | #[target_feature(enable = "avx2,fma")] |
| 14 | unsafe fn l2_squared_impl(a: &[f32], b: &[f32]) -> f32 { |
| 15 | assert_eq!(a.len(), b.len(), "avx2 l2_impl: length mismatch"); |
| 16 | unsafe { |
| 17 | use std::arch::x86_64::*; |
| 18 | let n = a.len(); |
| 19 | let mut sum = _mm256_setzero_ps(); |
| 20 | let chunks = n / 8; |
| 21 | for i in 0..chunks { |
| 22 | let off = i * 8; |
| 23 | let va = _mm256_loadu_ps(a.as_ptr().add(off)); |
| 24 | let vb = _mm256_loadu_ps(b.as_ptr().add(off)); |
| 25 | let diff = _mm256_sub_ps(va, vb); |
| 26 | sum = _mm256_fmadd_ps(diff, diff, sum); |
| 27 | } |
| 28 | let mut result = hsum256(sum); |
| 29 | for i in (chunks * 8)..n { |
| 30 | let d = a[i] - b[i]; |
| 31 | result += d * d; |
| 32 | } |
| 33 | result |
| 34 | } |
| 35 | } |
| 36 | |
| 37 | pub fn cosine_distance(a: &[f32], b: &[f32]) -> f32 { |
| 38 | assert_eq!(a.len(), b.len(), "avx2 cosine: length mismatch"); |
no test coverage detected