@harperfast/hnsw 0.1.0
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- package/Cargo.lock +229 -0
- package/Cargo.toml +30 -0
- package/DESIGN.md +319 -0
- package/LICENSE +202 -0
- package/README.md +88 -0
- package/build.mjs +46 -0
- package/build.rs +6 -0
- package/index.d.ts +104 -0
- package/index.js +57 -0
- package/package.json +44 -0
- package/prebuilds/darwin-arm64/hnsw-plane.node +0 -0
- package/prebuilds/linux-arm64/hnsw-plane.node +0 -0
- package/prebuilds/linux-x64/hnsw-plane.node +0 -0
- package/prebuilds/win32-x64/hnsw-plane.node +0 -0
- package/src/bin/bench.rs +246 -0
- package/src/distance.rs +138 -0
- package/src/format.rs +556 -0
- package/src/graph.rs +650 -0
- package/src/insert.rs +292 -0
- package/src/lib.rs +15 -0
- package/src/napi.rs +486 -0
- package/src/search.rs +477 -0
- package/src/seqlock.rs +217 -0
package/src/bin/bench.rs
ADDED
|
@@ -0,0 +1,246 @@
|
|
|
1
|
+
//! Standalone cost benchmark: build an N-node graph in the plane file, run queries, report
|
|
2
|
+
//! per-visit cost — the number that decides whether the native plane hits its 0.25–0.4 µs
|
|
3
|
+
//! budget (JS baseline: 4.34 µs/visit at 5M/ef 512).
|
|
4
|
+
//!
|
|
5
|
+
//! Usage: bench [n=100000] [dims=768] [queries=200] [ef=512] [path=/tmp/bench.hnsw] [cap=128] [threads=0]
|
|
6
|
+
//! threads > 0 adds a concurrent-throughput pass: T searcher threads (queries each) + one
|
|
7
|
+
//! background writer inserting throughout, reporting aggregate QPS and per-thread p50/p99.
|
|
8
|
+
|
|
9
|
+
use hnsw_plane::distance::Query;
|
|
10
|
+
use hnsw_plane::insert::{insert, InsertParams};
|
|
11
|
+
use hnsw_plane::search::{search, SearchScratch};
|
|
12
|
+
use hnsw_plane::{Graph, PlaneFile};
|
|
13
|
+
use std::path::PathBuf;
|
|
14
|
+
use std::time::Instant;
|
|
15
|
+
|
|
16
|
+
// xorshift for reproducible synthetic vectors without a rand dependency
|
|
17
|
+
struct Rng(u64);
|
|
18
|
+
impl Rng {
|
|
19
|
+
fn next_unit(&mut self) -> f32 {
|
|
20
|
+
self.0 ^= self.0 << 13;
|
|
21
|
+
self.0 ^= self.0 >> 7;
|
|
22
|
+
self.0 ^= self.0 << 17;
|
|
23
|
+
(self.0 >> 40) as f32 / (1u64 << 24) as f32
|
|
24
|
+
}
|
|
25
|
+
// Box-Muller
|
|
26
|
+
fn next_gauss(&mut self) -> f32 {
|
|
27
|
+
let u1 = self.next_unit().max(f32::MIN_POSITIVE);
|
|
28
|
+
let u2 = self.next_unit();
|
|
29
|
+
(-2.0 * u1.ln()).sqrt() * (2.0 * std::f32::consts::PI * u2).cos()
|
|
30
|
+
}
|
|
31
|
+
}
|
|
32
|
+
|
|
33
|
+
/// Gaussian-mixture corpus matching benchmarks/hnsw-scale.js: unit centroids, per-dim noise
|
|
34
|
+
/// derived from an intra-cluster cosine target of 0.75 (uniform-random 768-d is a corpus
|
|
35
|
+
/// "no ANN can index" per that benchmark's own calibration notes).
|
|
36
|
+
struct Corpus {
|
|
37
|
+
centroids: Vec<f32>,
|
|
38
|
+
n_clusters: usize,
|
|
39
|
+
dims: usize,
|
|
40
|
+
noise: f32,
|
|
41
|
+
}
|
|
42
|
+
|
|
43
|
+
impl Corpus {
|
|
44
|
+
fn new(n: u64, dims: usize, rng: &mut Rng) -> Self {
|
|
45
|
+
let intra_cos = 0.75f32;
|
|
46
|
+
let noise = ((1.0 / (intra_cos * intra_cos) - 1.0) / dims as f32).sqrt();
|
|
47
|
+
let n_clusters = 8.max((n as f64 / 500.0).round() as usize);
|
|
48
|
+
let mut centroids = vec![0.0f32; n_clusters * dims];
|
|
49
|
+
for c in 0..n_clusters {
|
|
50
|
+
let mut mag = 0.0f32;
|
|
51
|
+
for d in 0..dims {
|
|
52
|
+
let x = rng.next_gauss();
|
|
53
|
+
centroids[c * dims + d] = x;
|
|
54
|
+
mag += x * x;
|
|
55
|
+
}
|
|
56
|
+
let mag = mag.sqrt().max(f32::MIN_POSITIVE);
|
|
57
|
+
for d in 0..dims {
|
|
58
|
+
centroids[c * dims + d] /= mag;
|
|
59
|
+
}
|
|
60
|
+
}
|
|
61
|
+
Corpus { centroids, n_clusters, dims, noise }
|
|
62
|
+
}
|
|
63
|
+
|
|
64
|
+
fn row(&self, rng: &mut Rng) -> Vec<f32> {
|
|
65
|
+
let c = (rng.next_unit() * self.n_clusters as f32) as usize % self.n_clusters;
|
|
66
|
+
let mut v = vec![0.0f32; self.dims];
|
|
67
|
+
let mut mag = 0.0f32;
|
|
68
|
+
for d in 0..self.dims {
|
|
69
|
+
let x = self.centroids[c * self.dims + d] + rng.next_gauss() * self.noise;
|
|
70
|
+
v[d] = x;
|
|
71
|
+
mag += x * x;
|
|
72
|
+
}
|
|
73
|
+
let mag = mag.sqrt().max(f32::MIN_POSITIVE);
|
|
74
|
+
for d in 0..self.dims {
|
|
75
|
+
v[d] /= mag;
|
|
76
|
+
}
|
|
77
|
+
v
|
|
78
|
+
}
|
|
79
|
+
}
|
|
80
|
+
|
|
81
|
+
fn main() {
|
|
82
|
+
let args: Vec<String> = std::env::args().collect();
|
|
83
|
+
let n: u64 = args.get(1).and_then(|a| a.parse().ok()).unwrap_or(100_000);
|
|
84
|
+
let dims: usize = args.get(2).and_then(|a| a.parse().ok()).unwrap_or(768);
|
|
85
|
+
let queries: usize = args.get(3).and_then(|a| a.parse().ok()).unwrap_or(200);
|
|
86
|
+
let ef: usize = args.get(4).and_then(|a| a.parse().ok()).unwrap_or(512);
|
|
87
|
+
let path: PathBuf = args.get(5).map(Into::into).unwrap_or_else(|| "/tmp/bench.hnsw".into());
|
|
88
|
+
let layer0_cap: usize = args.get(6).and_then(|a| a.parse().ok()).unwrap_or(128);
|
|
89
|
+
|
|
90
|
+
// Reuse an existing plane file when it already holds exactly n nodes at the same cap
|
|
91
|
+
// (ef sweeps without rebuilding). The corpus RNG below replays identically.
|
|
92
|
+
let reuse = PlaneFile::open(&path)
|
|
93
|
+
.ok()
|
|
94
|
+
.filter(|f| f.id_high_water() == n && f.layer0_cap == layer0_cap)
|
|
95
|
+
.is_some();
|
|
96
|
+
let file = if reuse {
|
|
97
|
+
println!("reusing existing plane at {}", path.display());
|
|
98
|
+
PlaneFile::open(&path).expect("open")
|
|
99
|
+
} else {
|
|
100
|
+
PlaneFile::create(&path, dims, layer0_cap, n + 1024).expect("create")
|
|
101
|
+
};
|
|
102
|
+
println!(
|
|
103
|
+
"plane: {} nodes x {} dims, slot {} B, file {:.1} GB (sparse)",
|
|
104
|
+
n,
|
|
105
|
+
dims,
|
|
106
|
+
file.slot_size,
|
|
107
|
+
(n * file.slot_size as u64) as f64 / 1e9
|
|
108
|
+
);
|
|
109
|
+
let graph = Graph::new(file);
|
|
110
|
+
let params = InsertParams::default();
|
|
111
|
+
let mut scratch = SearchScratch::new();
|
|
112
|
+
let mut rng = Rng(0x1234_5678_9abc_def0);
|
|
113
|
+
let corpus = Corpus::new(n, dims, &mut rng);
|
|
114
|
+
|
|
115
|
+
if reuse {
|
|
116
|
+
// replay the build's RNG draws so query rows match a fresh run; the upper region
|
|
117
|
+
// persists inside the plane file
|
|
118
|
+
for _ in 0..n {
|
|
119
|
+
let _ = corpus.row(&mut rng);
|
|
120
|
+
}
|
|
121
|
+
} else {
|
|
122
|
+
let build_start = Instant::now();
|
|
123
|
+
for i in 0..n {
|
|
124
|
+
let v = corpus.row(&mut rng);
|
|
125
|
+
insert(&graph, &v, ¶ms, &mut scratch);
|
|
126
|
+
if (i + 1) % 50_000 == 0 {
|
|
127
|
+
let rate = (i + 1) as f64 / build_start.elapsed().as_secs_f64();
|
|
128
|
+
println!(" built {} ({:.0} inserts/s)", i + 1, rate);
|
|
129
|
+
}
|
|
130
|
+
}
|
|
131
|
+
let build = build_start.elapsed();
|
|
132
|
+
println!("build: {:.1}s ({:.0} inserts/s)", build.as_secs_f64(), n as f64 / build.as_secs_f64());
|
|
133
|
+
graph.file.msync().expect("msync");
|
|
134
|
+
}
|
|
135
|
+
|
|
136
|
+
// Query with held-out vectors; measure latency and set-recall@10 vs brute-force truth
|
|
137
|
+
// (same asymmetric metric, so recall isolates graph quality, not quantization).
|
|
138
|
+
let mut latencies = Vec::with_capacity(queries);
|
|
139
|
+
let mut total_visits = 0u64;
|
|
140
|
+
let mut recall_hits = 0usize;
|
|
141
|
+
let mut recall_total = 0usize;
|
|
142
|
+
for _ in 0..queries {
|
|
143
|
+
let q = Query::new(corpus.row(&mut rng));
|
|
144
|
+
let start = Instant::now();
|
|
145
|
+
let (results, stats) = search(&graph, &q, 10, ef, &mut scratch);
|
|
146
|
+
latencies.push(start.elapsed());
|
|
147
|
+
total_visits += stats.visits;
|
|
148
|
+
assert!(!results.is_empty());
|
|
149
|
+
|
|
150
|
+
let mut truth: Vec<(u32, f32)> = (0..n as u32)
|
|
151
|
+
.filter_map(|id| graph.distance_to(id, &q).map(|d| (id, d)))
|
|
152
|
+
.collect();
|
|
153
|
+
truth.sort_by(|a, b| a.1.partial_cmp(&b.1).unwrap());
|
|
154
|
+
truth.truncate(10);
|
|
155
|
+
recall_total += truth.len();
|
|
156
|
+
recall_hits += truth.iter().filter(|(tid, _)| results.iter().any(|(rid, _)| rid == tid)).count();
|
|
157
|
+
}
|
|
158
|
+
latencies.sort();
|
|
159
|
+
let p50 = latencies[queries / 2];
|
|
160
|
+
let p95 = latencies[queries * 95 / 100];
|
|
161
|
+
let p99 = latencies[(queries * 99 / 100).min(queries - 1)];
|
|
162
|
+
let mean_visits = total_visits as f64 / queries as f64;
|
|
163
|
+
let us_per_visit = p50.as_micros() as f64 / mean_visits;
|
|
164
|
+
println!(
|
|
165
|
+
"search (ef {}): p50 {:.2} ms p95 {:.2} ms p99 {:.2} ms visits/query {:.0} -> {:.3} us/visit (JS baseline 4.34)",
|
|
166
|
+
ef,
|
|
167
|
+
p50.as_secs_f64() * 1e3,
|
|
168
|
+
p95.as_secs_f64() * 1e3,
|
|
169
|
+
p99.as_secs_f64() * 1e3,
|
|
170
|
+
mean_visits,
|
|
171
|
+
us_per_visit
|
|
172
|
+
);
|
|
173
|
+
println!("recall@10 (set): {:.3}", recall_hits as f64 / recall_total as f64);
|
|
174
|
+
|
|
175
|
+
let threads: usize = args.get(7).and_then(|a| a.parse().ok()).unwrap_or(0);
|
|
176
|
+
if threads > 0 {
|
|
177
|
+
use std::sync::atomic::{AtomicBool, Ordering};
|
|
178
|
+
use std::sync::Arc;
|
|
179
|
+
let graph = Arc::new(graph);
|
|
180
|
+
let corpus = Arc::new(corpus);
|
|
181
|
+
let stop = Arc::new(AtomicBool::new(false));
|
|
182
|
+
let per_thread = queries.max(100);
|
|
183
|
+
let start = Instant::now();
|
|
184
|
+
let mut handles = Vec::new();
|
|
185
|
+
for t in 0..threads {
|
|
186
|
+
let graph = graph.clone();
|
|
187
|
+
let corpus = corpus.clone();
|
|
188
|
+
handles.push(std::thread::spawn(move || {
|
|
189
|
+
let mut scratch = SearchScratch::new();
|
|
190
|
+
let mut rng = Rng(0x9e37_79b9 ^ (t as u64 + 1) * 0x1234_5677);
|
|
191
|
+
let mut lat: Vec<std::time::Duration> = Vec::with_capacity(per_thread);
|
|
192
|
+
for _ in 0..per_thread {
|
|
193
|
+
let q = Query::new(corpus.row(&mut rng));
|
|
194
|
+
let s = Instant::now();
|
|
195
|
+
let (r, _) = search(&graph, &q, 10, ef, &mut scratch);
|
|
196
|
+
lat.push(s.elapsed());
|
|
197
|
+
assert!(!r.is_empty());
|
|
198
|
+
}
|
|
199
|
+
lat.sort();
|
|
200
|
+
(lat[per_thread / 2], lat[(per_thread * 99 / 100).min(per_thread - 1)])
|
|
201
|
+
}));
|
|
202
|
+
}
|
|
203
|
+
// background writer: sustained inserts while searchers run
|
|
204
|
+
let writer = {
|
|
205
|
+
let graph = graph.clone();
|
|
206
|
+
let corpus = corpus.clone();
|
|
207
|
+
let stop = stop.clone();
|
|
208
|
+
std::thread::spawn(move || {
|
|
209
|
+
let params = InsertParams::default();
|
|
210
|
+
let mut scratch = SearchScratch::new();
|
|
211
|
+
let mut rng = Rng(0xdead_beef_cafe_f00d);
|
|
212
|
+
let mut count = 0u64;
|
|
213
|
+
while !stop.load(Ordering::Relaxed) {
|
|
214
|
+
let v = corpus.row(&mut rng);
|
|
215
|
+
if insert(&graph, &v, ¶ms, &mut scratch).is_err() {
|
|
216
|
+
break; // plane full
|
|
217
|
+
}
|
|
218
|
+
count += 1;
|
|
219
|
+
}
|
|
220
|
+
count
|
|
221
|
+
})
|
|
222
|
+
};
|
|
223
|
+
let mut p50s = Vec::new();
|
|
224
|
+
let mut p99s = Vec::new();
|
|
225
|
+
for h in handles {
|
|
226
|
+
let (p50, p99) = h.join().unwrap();
|
|
227
|
+
p50s.push(p50);
|
|
228
|
+
p99s.push(p99);
|
|
229
|
+
}
|
|
230
|
+
let wall = start.elapsed();
|
|
231
|
+
stop.store(true, Ordering::Relaxed);
|
|
232
|
+
let inserted = writer.join().unwrap();
|
|
233
|
+
let total_q = (threads * per_thread) as f64;
|
|
234
|
+
p50s.sort();
|
|
235
|
+
p99s.sort();
|
|
236
|
+
println!(
|
|
237
|
+
"concurrent: {} threads x {} queries + writer -> {:.0} QPS aggregate p50(med) {:.2} ms p99(worst) {:.2} ms writer {:.0} inserts/s",
|
|
238
|
+
threads,
|
|
239
|
+
per_thread,
|
|
240
|
+
total_q / wall.as_secs_f64(),
|
|
241
|
+
p50s[threads / 2].as_secs_f64() * 1e3,
|
|
242
|
+
p99s[threads - 1].as_secs_f64() * 1e3,
|
|
243
|
+
inserted as f64 / wall.as_secs_f64()
|
|
244
|
+
);
|
|
245
|
+
}
|
|
246
|
+
}
|
package/src/distance.rs
ADDED
|
@@ -0,0 +1,138 @@
|
|
|
1
|
+
//! Distance kernels. Asymmetric: full-precision f32 query × int8-stored vector (matches the JS
|
|
2
|
+
//! quantizeInt8 scale + cached 1/|v| model). Symmetric int8×int8 for construction-time
|
|
3
|
+
//! neighbor↔neighbor checks (stored per-edge distances were dropped from the format; recompute).
|
|
4
|
+
//! AVX2 with scalar fallback; Linux x86_64 is the performance target, other platforms take the
|
|
5
|
+
//! scalar path (fine for dev).
|
|
6
|
+
|
|
7
|
+
/// Precomputed query state, built once per search.
|
|
8
|
+
pub struct Query {
|
|
9
|
+
pub vector: Vec<f32>,
|
|
10
|
+
pub inv_mag: f32,
|
|
11
|
+
}
|
|
12
|
+
|
|
13
|
+
impl Query {
|
|
14
|
+
pub fn new(vector: Vec<f32>) -> Self {
|
|
15
|
+
let mag_sq: f32 = vector.iter().map(|v| v * v).sum();
|
|
16
|
+
let inv_mag = 1.0 / mag_sq.sqrt().max(f32::MIN_POSITIVE);
|
|
17
|
+
Query { vector, inv_mag }
|
|
18
|
+
}
|
|
19
|
+
}
|
|
20
|
+
|
|
21
|
+
#[inline]
|
|
22
|
+
fn dot_f32_i8_scalar(q: &[f32], v: *const i8) -> f32 {
|
|
23
|
+
let mut acc = [0.0f32; 8];
|
|
24
|
+
let chunks = q.len() / 8;
|
|
25
|
+
for c in 0..chunks {
|
|
26
|
+
let base = c * 8;
|
|
27
|
+
for lane in 0..8 {
|
|
28
|
+
acc[lane] += q[base + lane] * unsafe { *v.add(base + lane) } as f32;
|
|
29
|
+
}
|
|
30
|
+
}
|
|
31
|
+
let mut dot: f32 = acc.iter().sum();
|
|
32
|
+
for i in chunks * 8..q.len() {
|
|
33
|
+
dot += q[i] * unsafe { *v.add(i) } as f32;
|
|
34
|
+
}
|
|
35
|
+
dot
|
|
36
|
+
}
|
|
37
|
+
|
|
38
|
+
#[cfg(target_arch = "x86_64")]
|
|
39
|
+
#[target_feature(enable = "avx2", enable = "fma")]
|
|
40
|
+
unsafe fn dot_f32_i8_avx2(q: &[f32], v: *const i8) -> f32 {
|
|
41
|
+
use std::arch::x86_64::*;
|
|
42
|
+
let mut acc0 = _mm256_setzero_ps();
|
|
43
|
+
let mut acc1 = _mm256_setzero_ps();
|
|
44
|
+
let chunks = q.len() / 16;
|
|
45
|
+
for c in 0..chunks {
|
|
46
|
+
let base = c * 16;
|
|
47
|
+
let v16 = _mm_loadu_si128(v.add(base) as *const __m128i);
|
|
48
|
+
let lo = _mm256_cvtepi32_ps(_mm256_cvtepi8_epi32(v16));
|
|
49
|
+
let hi = _mm256_cvtepi32_ps(_mm256_cvtepi8_epi32(_mm_srli_si128(v16, 8)));
|
|
50
|
+
acc0 = _mm256_fmadd_ps(_mm256_loadu_ps(q.as_ptr().add(base)), lo, acc0);
|
|
51
|
+
acc1 = _mm256_fmadd_ps(_mm256_loadu_ps(q.as_ptr().add(base + 8)), hi, acc1);
|
|
52
|
+
}
|
|
53
|
+
let acc = _mm256_add_ps(acc0, acc1);
|
|
54
|
+
let s = _mm_add_ps(_mm256_extractf128_ps(acc, 1), _mm256_castps256_ps128(acc));
|
|
55
|
+
let s = _mm_hadd_ps(s, s);
|
|
56
|
+
let s = _mm_hadd_ps(s, s);
|
|
57
|
+
let mut dot = _mm_cvtss_f32(s);
|
|
58
|
+
for i in chunks * 16..q.len() {
|
|
59
|
+
dot += q[i] * *v.add(i) as f32;
|
|
60
|
+
}
|
|
61
|
+
dot
|
|
62
|
+
}
|
|
63
|
+
|
|
64
|
+
#[inline]
|
|
65
|
+
fn dot_f32_i8(q: &[f32], v: *const i8) -> f32 {
|
|
66
|
+
#[cfg(target_arch = "x86_64")]
|
|
67
|
+
{
|
|
68
|
+
if std::arch::is_x86_feature_detected!("avx2") && std::arch::is_x86_feature_detected!("fma") {
|
|
69
|
+
return unsafe { dot_f32_i8_avx2(q, v) };
|
|
70
|
+
}
|
|
71
|
+
}
|
|
72
|
+
dot_f32_i8_scalar(q, v)
|
|
73
|
+
}
|
|
74
|
+
|
|
75
|
+
/// Cosine distance: f32 query × raw int8 vector at `stored` (dims = query.vector.len()).
|
|
76
|
+
/// Zero-copy: `stored` points into the mmap; the caller's seqlock read discards torn results.
|
|
77
|
+
#[inline]
|
|
78
|
+
pub fn cosine_int8_raw(query: &Query, stored: *const i8, scale: f32, stored_inv_mag: f32) -> f32 {
|
|
79
|
+
let dot = dot_f32_i8(&query.vector, stored);
|
|
80
|
+
1.0 - dot * scale * stored_inv_mag * query.inv_mag
|
|
81
|
+
}
|
|
82
|
+
|
|
83
|
+
#[inline]
|
|
84
|
+
fn dot_i8_i8_scalar(a: *const i8, b: *const i8, len: usize) -> i32 {
|
|
85
|
+
let mut dot = 0i32;
|
|
86
|
+
for i in 0..len {
|
|
87
|
+
dot += unsafe { *a.add(i) as i32 * *b.add(i) as i32 };
|
|
88
|
+
}
|
|
89
|
+
dot
|
|
90
|
+
}
|
|
91
|
+
|
|
92
|
+
#[cfg(target_arch = "x86_64")]
|
|
93
|
+
#[target_feature(enable = "avx2")]
|
|
94
|
+
unsafe fn dot_i8_i8_avx2(a: *const i8, b: *const i8, len: usize) -> i32 {
|
|
95
|
+
use std::arch::x86_64::*;
|
|
96
|
+
let mut acc = _mm256_setzero_si256();
|
|
97
|
+
let chunks = len / 16;
|
|
98
|
+
for c in 0..chunks {
|
|
99
|
+
let av = _mm256_cvtepi8_epi16(_mm_loadu_si128(a.add(c * 16) as *const __m128i));
|
|
100
|
+
let bv = _mm256_cvtepi8_epi16(_mm_loadu_si128(b.add(c * 16) as *const __m128i));
|
|
101
|
+
acc = _mm256_add_epi32(acc, _mm256_madd_epi16(av, bv));
|
|
102
|
+
}
|
|
103
|
+
let lo = _mm256_castsi256_si128(acc);
|
|
104
|
+
let hi = _mm256_extracti128_si256(acc, 1);
|
|
105
|
+
let s = _mm_add_epi32(lo, hi);
|
|
106
|
+
let s = _mm_add_epi32(s, _mm_srli_si128(s, 8));
|
|
107
|
+
let s = _mm_add_epi32(s, _mm_srli_si128(s, 4));
|
|
108
|
+
let mut dot = _mm_cvtsi128_si32(s);
|
|
109
|
+
for i in chunks * 16..len {
|
|
110
|
+
dot += *a.add(i) as i32 * *b.add(i) as i32;
|
|
111
|
+
}
|
|
112
|
+
dot
|
|
113
|
+
}
|
|
114
|
+
|
|
115
|
+
/// Cosine distance between two int8-stored vectors (construction-time neighbor checks).
|
|
116
|
+
#[inline]
|
|
117
|
+
pub fn cosine_i8_i8_raw(a: *const i8, scale_a: f32, inv_mag_a: f32, b: *const i8, scale_b: f32, inv_mag_b: f32, len: usize) -> f32 {
|
|
118
|
+
#[cfg(target_arch = "x86_64")]
|
|
119
|
+
let dot = if std::arch::is_x86_feature_detected!("avx2") {
|
|
120
|
+
unsafe { dot_i8_i8_avx2(a, b, len) }
|
|
121
|
+
} else {
|
|
122
|
+
dot_i8_i8_scalar(a, b, len)
|
|
123
|
+
};
|
|
124
|
+
#[cfg(not(target_arch = "x86_64"))]
|
|
125
|
+
let dot = dot_i8_i8_scalar(a, b, len);
|
|
126
|
+
1.0 - dot as f32 * scale_a * scale_b * inv_mag_a * inv_mag_b
|
|
127
|
+
}
|
|
128
|
+
|
|
129
|
+
/// Symmetric int8 quantization matching the JS quantizeInt8: scale maps max |component| to 127.
|
|
130
|
+
pub fn quantize_int8(vector: &[f32]) -> (Vec<i8>, f32, f32) {
|
|
131
|
+
let max_abs = vector.iter().fold(0.0f32, |m, v| m.max(v.abs()));
|
|
132
|
+
let scale = if max_abs == 0.0 { 1.0 } else { max_abs / 127.0 };
|
|
133
|
+
let inv_scale = 1.0 / scale;
|
|
134
|
+
let bytes: Vec<i8> = vector.iter().map(|v| (v * inv_scale).round().clamp(-127.0, 127.0) as i8).collect();
|
|
135
|
+
let mag_sq: f32 = vector.iter().map(|v| v * v).sum();
|
|
136
|
+
let inv_mag = 1.0 / mag_sq.sqrt().max(f32::MIN_POSITIVE);
|
|
137
|
+
(bytes, scale, inv_mag)
|
|
138
|
+
}
|