From 22a5fb3d9ff9dfbe63049b6bfa36ce16755d66a8 Mon Sep 17 00:00:00 2001 From: Wolfgang Schoenberger <221313372+wolfiesch@users.noreply.github.com> Date: Wed, 22 Jul 2026 01:24:37 -0700 Subject: [PATCH] perf(mnemopi): use native batch vector kernels in recall hot loops Swap the batch-shaped scoring loops to the pi-natives kernels behind the existing TS function signatures, so every caller keeps its guards and observable behavior: mmrRerank (default Jaccard path), searchExactVectorIndex, clusterBySimilarity, BinaryVectorStore.search, and FastBinarySearch.search. cosine_similarity_pairs now takes Float64Array so f64 vectors round-trip without f32 narrowing. Scalar one-off call sites stay in TS (query-cache cosine probe, shmr centroid/confidence loops, recall per-row cosine map, custom-similarity mmrRerank). Adds a seeded parity suite: 1e-9 relative tolerance for floats, exact equality for Hamming distances and pair lists, identical MMR index sequences across lambda/topK grids. --- crates/pi-natives/src/vectors.rs | 17 +- .../mnemopi/bench/native-vectors.bench.ts | 122 ++++++++++++ packages/mnemopi/package.json | 1 + packages/mnemopi/src/core/binary-vectors.ts | 44 ++++- packages/mnemopi/src/core/mmr.ts | 16 ++ packages/mnemopi/src/core/shmr.ts | 29 +-- packages/mnemopi/src/core/vector-index.ts | 23 +-- .../mnemopi/test/native-vector-parity.test.ts | 184 ++++++++++++++++++ packages/natives/native/index.d.ts | 13 +- 9 files changed, 402 insertions(+), 47 deletions(-) create mode 100644 packages/mnemopi/bench/native-vectors.bench.ts create mode 100644 packages/mnemopi/test/native-vector-parity.test.ts diff --git a/crates/pi-natives/src/vectors.rs b/crates/pi-natives/src/vectors.rs index 0610e023b..0dc6b66c7 100644 --- a/crates/pi-natives/src/vectors.rs +++ b/crates/pi-natives/src/vectors.rs @@ -92,26 +92,25 @@ pub fn cosine_similarity_batch( /// All pairs `(i, j)` with `i < j` whose cosine similarity meets `threshold`. /// -/// `vectors` is `count` vectors flattened row-major at `dim` `f32` elements -/// per row (zero-padded). Element values are widened to `f64` before any -/// arithmetic, exactly like JS reads from a `Float32Array`, so the similarity -/// is bit-identical to the TS pairwise loop in `clusterBySimilarity`. -/// Returns pairs flattened as `[i0, j0, i1, j1, ...]` in the same `(i, j)` -/// visit order as the TS nested loop. +/// `vectors` is `count` vectors flattened row-major at `dim` `f64` elements +/// per row (zero-padded, which matches the TS `?? 0` missing-element +/// semantics), so the similarity is bit-identical to the TS pairwise loop in +/// `clusterBySimilarity`. Returns pairs flattened as `[i0, j0, i1, j1, ...]` +/// in the same `(i, j)` visit order as the TS nested loop. #[napi] pub fn cosine_similarity_pairs( - vectors: Float32Array, + vectors: Float64Array, count: u32, dim: u32, threshold: f64, ) -> Result { let count = count as usize; let dim = dim as usize; - let data: &[f32] = &vectors; + let data: &[f64] = &vectors; if data.len() != count * dim { return invalid("vectors length must equal count * dim"); } - let widened: Vec = data.iter().map(|&v| f64::from(v)).collect(); + let widened: &[f64] = data; let mut pairs: Vec = Vec::new(); for i in 0..count { let left = &widened[i * dim..(i + 1) * dim]; diff --git a/packages/mnemopi/bench/native-vectors.bench.ts b/packages/mnemopi/bench/native-vectors.bench.ts new file mode 100644 index 000000000..b52de396a --- /dev/null +++ b/packages/mnemopi/bench/native-vectors.bench.ts @@ -0,0 +1,122 @@ +/** + * Crossing-inclusive benchmark of native batch vector kernels vs the TS + * reference loops, at realistic mnemopi recall shapes (fastembed + * bge-small-en-v1.5, dim=384; binarized stride=48 bytes). + * + * Run from the repo root: `bun packages/mnemopi/bench/native-vectors.bench.ts` + */ +import { execSync } from "node:child_process"; +import { cosineSimilarityBatch, hammingDistanceBatch, vectorIndexTopK } from "@oh-my-pi/pi-natives"; +import { hammingDistance } from "../src/core/binary-vectors"; +import { cosineSimilarity } from "../src/core/vector-math"; + +const DIM = 384; +const STRIDE = DIM / 8; +const COUNTS = [10, 100, 1000, 10000]; +const WARMUP = 20; +const ITERATIONS = 200; + +function makeRng(seed: number): () => number { + let state = seed >>> 0; + return () => { + state = (state * 1664525 + 1013904223) >>> 0; + return state / 4294967296; + }; +} + +let sink = 0; + +function timeNs(fn: () => void): number { + for (let i = 0; i < WARMUP; i += 1) fn(); + const start = Bun.nanoseconds(); + for (let i = 0; i < ITERATIONS; i += 1) fn(); + return (Bun.nanoseconds() - start) / ITERATIONS; +} + +interface Row { + kernel: string; + count: number; + tsNs: number; + nativeNs: number; + speedup: number; +} + +const rows: Row[] = []; +const rng = makeRng(0xbe4c4); + +for (const count of COUNTS) { + const query = Float64Array.from({ length: DIM }, () => rng() * 2 - 1); + const flat = new Float64Array(count * DIM); + for (let i = 0; i < flat.length; i += 1) flat[i] = rng() * 2 - 1; + + const tsNs = timeNs(() => { + for (let row = 0; row < count; row += 1) { + sink += cosineSimilarity(query, flat.subarray(row * DIM, (row + 1) * DIM)); + } + }); + const nativeNs = timeNs(() => { + sink += cosineSimilarityBatch(query, flat, DIM)[0] ?? 0; + }); + rows.push({ kernel: "cosineSimilarityBatch", count, tsNs, nativeNs, speedup: tsNs / nativeNs }); +} + +for (const count of COUNTS) { + const matrix = new Float32Array(count * DIM); + for (let i = 0; i < matrix.length; i += 1) matrix[i] = rng() * 2 - 1; + const query = Float64Array.from({ length: DIM }, () => rng() * 2 - 1); + let normSq = 0; + for (const v of query) normSq += v * v; + const norm = Math.sqrt(normSq); + const limit = 10; + + const tsNs = timeNs(() => { + const hits: Array<{ row: number; score: number }> = []; + for (let row = 0; row < count; row += 1) { + let score = 0; + for (let col = 0; col < DIM; col += 1) { + score += (matrix[row * DIM + col] ?? 0) * ((query[col] ?? 0) / norm); + } + hits.push({ row, score }); + } + hits.sort((a, b) => b.score - a.score); + sink += hits[0]?.score ?? 0; + }); + const nativeNs = timeNs(() => { + sink += vectorIndexTopK(matrix, DIM, query, limit).scores[0] ?? 0; + }); + rows.push({ kernel: "vectorIndexTopK", count, tsNs, nativeNs, speedup: tsNs / nativeNs }); +} + +for (const count of COUNTS) { + const query = Uint8Array.from({ length: STRIDE }, () => Math.floor(rng() * 256)); + const packed = new Uint8Array(count * STRIDE); + for (let i = 0; i < packed.length; i += 1) packed[i] = Math.floor(rng() * 256); + const vectors: Uint8Array[] = []; + for (let i = 0; i < count; i += 1) vectors.push(packed.subarray(i * STRIDE, (i + 1) * STRIDE)); + + const tsNs = timeNs(() => { + for (let i = 0; i < count; i += 1) sink += hammingDistance(query, vectors[i] ?? new Uint8Array()); + }); + const nativeNs = timeNs(() => { + sink += hammingDistanceBatch(query, packed, STRIDE)[0] ?? 0; + }); + rows.push({ kernel: "hammingDistanceBatch", count, tsNs, nativeNs, speedup: tsNs / nativeNs }); +} + +const sha = execSync("git rev-parse HEAD").toString().trim(); +const report = { + sha, + date: new Date().toISOString(), + scenario: `dim=${DIM}, stride=${STRIDE}B, warmup=${WARMUP}, iterations=${ITERATIONS}, crossing-inclusive`, + runtime: `bun ${Bun.version}`, + rows: rows.map(r => ({ + kernel: r.kernel, + count: r.count, + ts_us: +(r.tsNs / 1000).toFixed(2), + native_us: +(r.nativeNs / 1000).toFixed(2), + speedup: +r.speedup.toFixed(2), + })), + sink, +}; + +console.log(JSON.stringify(report, null, 2)); diff --git a/packages/mnemopi/package.json b/packages/mnemopi/package.json index 4d5b505af..c38357420 100644 --- a/packages/mnemopi/package.json +++ b/packages/mnemopi/package.json @@ -41,6 +41,7 @@ "dependencies": { "@oh-my-pi/pi-ai": "catalog:", "@oh-my-pi/pi-catalog": "catalog:", + "@oh-my-pi/pi-natives": "catalog:", "@oh-my-pi/pi-utils": "catalog:", "lru-cache": "catalog:" }, diff --git a/packages/mnemopi/src/core/binary-vectors.ts b/packages/mnemopi/src/core/binary-vectors.ts index 6964a5e70..21d26718c 100644 --- a/packages/mnemopi/src/core/binary-vectors.ts +++ b/packages/mnemopi/src/core/binary-vectors.ts @@ -1,4 +1,5 @@ import type { Database } from "bun:sqlite"; +import { hammingDistanceBatch, hammingDistanceForDimBatch } from "@oh-my-pi/pi-natives"; import { embeddingDim, type VecType } from "../config"; import { closeQuietly, type DatabasePath, openDatabase } from "../db"; @@ -147,7 +148,7 @@ export function hammingDistance(binaryA: Uint8Array | ArrayBuffer, binaryB: Uint return distance; } -function hammingDistanceForDimension( +export function hammingDistanceForDimension( binaryA: Uint8Array | ArrayBuffer, binaryB: Uint8Array | ArrayBuffer, dim: number, @@ -226,15 +227,31 @@ export class BinaryVectorStore { const rows = this.conn .query(`SELECT memory_id, binary_vector, original_dim, magnitude FROM ${this.tableName}`) .all() as VectorRow[]; - const results: BinaryVectorSearchResult[] = []; - for (const row of rows) { + // Native batch kernel: pack all rows at a fixed stride (zero-padded, + // matching the TS `?? 0` out-of-range reads) and compute every + // dimension-masked distance in one N-API crossing. + const stride = BYTES_PER_VECTOR; + const packed = new Uint8Array(rows.length * stride); + const dims = new Uint32Array(rows.length); + const comparedDims: number[] = new Array(rows.length); + for (let i = 0; i < rows.length; i += 1) { + const row = rows[i]; + if (row === undefined) continue; const storedDim = Math.max(0, Math.min(EMBEDDING_DIM, Math.trunc(toFiniteNumber(row.original_dim)))); const comparedDim = Math.min(queryDim, storedDim); - const distance = hammingDistanceForDimension(queryBinary, bytesFromBlob(row.binary_vector), comparedDim); + comparedDims[i] = comparedDim; + dims[i] = comparedDim; + const bytes = bytesFromBlob(row.binary_vector); + packed.set(bytes.length > stride ? bytes.subarray(0, stride) : bytes, i * stride); + } + const distances = hammingDistanceForDimBatch(queryBinary, packed, stride, dims); + const results: BinaryVectorSearchResult[] = []; + for (let i = 0; i < rows.length; i += 1) { + const distance = distances[i] ?? 0; results.push({ - memory_id: row.memory_id, + memory_id: rows[i]?.memory_id ?? "", distance, - score: informationTheoreticScore(distance, comparedDim), + score: informationTheoreticScore(distance, comparedDims[i] ?? 0), }); } results.sort((a, b) => b.score - a.score || a.memory_id.localeCompare(b.memory_id)); @@ -302,9 +319,22 @@ export class FastBinarySearch { search(queryBinary: Uint8Array | ArrayBuffer, topK = 10): BinaryVectorSearchResult[] { const query = queryBinary instanceof Uint8Array ? queryBinary : new Uint8Array(queryBinary); + // Native batch kernel: pack all vectors at the max byte length and + // compute every Hamming distance in one N-API crossing. Per-row + // lengths preserve the TS unmatched-tail popcount semantics. + let stride = 0; + for (const vector of this.vectors) if (vector.length > stride) stride = vector.length; + const packed = new Uint8Array(this.vectors.length * stride); + const lengths = new Uint32Array(this.vectors.length); + for (let i = 0; i < this.vectors.length; i += 1) { + const vector = this.vectors[i] ?? new Uint8Array(); + lengths[i] = vector.length; + packed.set(vector, i * stride); + } + const distances = hammingDistanceBatch(query, packed, stride, lengths); const results: BinaryVectorSearchResult[] = []; for (let i = 0; i < this.vectors.length; i += 1) { - const distance = hammingDistance(query, this.vectors[i] ?? new Uint8Array()); + const distance = distances[i] ?? 0; results.push({ memory_id: this.memoryIds[i] ?? "", distance, diff --git a/packages/mnemopi/src/core/mmr.ts b/packages/mnemopi/src/core/mmr.ts index 6cc87322b..9f706943c 100644 --- a/packages/mnemopi/src/core/mmr.ts +++ b/packages/mnemopi/src/core/mmr.ts @@ -1,3 +1,5 @@ +import { mmrRerankIndices } from "@oh-my-pi/pi-natives"; + export interface MmrResult { readonly content?: string; readonly score?: number; @@ -33,6 +35,20 @@ export function mmrRerank( const first = sortedResults[0]; if (first === undefined) return []; + // Native batch kernel: one N-API crossing selects all indices. Only valid + // for the default Jaccard similarity; custom similarity functions stay in TS. + if (similarityFn === jaccardSimilarity) { + const contents = sortedResults.map(result => result.content ?? ""); + const scores = Float64Array.from(sortedResults, result => result.score ?? 0); + const picked = mmrRerankIndices(contents, scores, lambdaParam, limit); + const out: T[] = []; + for (const index of picked) { + const item = sortedResults[index]; + if (item !== undefined) out.push(item); + } + return out; + } + const selected: T[] = [first]; const remaining = sortedResults.slice(1); diff --git a/packages/mnemopi/src/core/shmr.ts b/packages/mnemopi/src/core/shmr.ts index dca63ba24..e6497ed9b 100644 --- a/packages/mnemopi/src/core/shmr.ts +++ b/packages/mnemopi/src/core/shmr.ts @@ -1,5 +1,6 @@ import type { Database } from "bun:sqlite"; import { createHash } from "node:crypto"; +import { cosineSimilarityPairs } from "@oh-my-pi/pi-natives"; import { logger } from "@oh-my-pi/pi-utils"; import * as embeddings from "./embeddings"; import { cosineSimilarity } from "./vector-math"; @@ -166,17 +167,23 @@ export async function clusterBySimilarity(items: readonly ShmrItem[], threshold: if (items.length === 0) return []; const vectors = await resolveItemVectors(items); const adjacency: number[][] = Array.from({ length: items.length }, () => []); - for (let i = 0; i < items.length; i++) { - const leftEmbedding = vectors[i]; - if (leftEmbedding === undefined) continue; - for (let j = i + 1; j < items.length; j++) { - const rightEmbedding = vectors[j]; - if (rightEmbedding === undefined) continue; - if (cosineSimilarity(leftEmbedding, rightEmbedding) >= threshold) { - adjacency[i]?.push(j); - adjacency[j]?.push(i); - } - } + // Native batch kernel: one N-API crossing evaluates all O(n^2) pairs. + // Vectors are zero-padded to a shared dim, matching the TS `?? 0` + // missing-element semantics of cosineSimilarity. + let dim = 0; + for (const vector of vectors) if (vector.length > dim) dim = vector.length; + const flat = new Float64Array(vectors.length * dim); + for (let i = 0; i < vectors.length; i++) { + const vector = vectors[i]; + if (vector === undefined) continue; + for (let col = 0; col < vector.length; col++) flat[i * dim + col] = vector[col] ?? 0; + } + const pairs = cosineSimilarityPairs(flat, vectors.length, dim, threshold); + for (let p = 0; p < pairs.length; p += 2) { + const i = pairs[p] ?? 0; + const j = pairs[p + 1] ?? 0; + adjacency[i]?.push(j); + adjacency[j]?.push(i); } const visited = new Set(); const clusters: ShmrItem[][] = []; diff --git a/packages/mnemopi/src/core/vector-index.ts b/packages/mnemopi/src/core/vector-index.ts index 3092cba52..f3952635c 100644 --- a/packages/mnemopi/src/core/vector-index.ts +++ b/packages/mnemopi/src/core/vector-index.ts @@ -1,3 +1,5 @@ +import { vectorIndexTopK } from "@oh-my-pi/pi-natives"; + export interface ExactVectorSearchHit { id: TId; score: number; @@ -66,19 +68,14 @@ export function searchExactVectorIndex( queryNormSq += value * value; } if (queryNormSq <= 0) return []; - const queryNorm = Math.sqrt(queryNormSq); - const queryDimensions = Math.min(query.length, index.dimensions); + // Native batch kernel: one N-API crossing scores every row and ranks the + // top k with the same stable ordering as the TS sort. TS guards above + // (finite query, positive norm, non-empty index) are preserved. + const topK = vectorIndexTopK(index.matrix, index.dimensions, Float64Array.from(query), k); const hits: ExactVectorSearchHit[] = []; - - for (let row = 0; row < index.count; row += 1) { - const offset = row * index.dimensions; - let score = 0; - for (let col = 0; col < queryDimensions; col += 1) { - score += index.matrix[offset + col] * ((query[col] ?? 0) / queryNorm); - } - hits.push({ id: index.ids[row] as TId, score }); + for (let i = 0; i < topK.indices.length; i += 1) { + const row = topK.indices[i] ?? 0; + hits.push({ id: index.ids[row] as TId, score: topK.scores[i] ?? 0 }); } - - hits.sort((a, b) => b.score - a.score); - return hits.slice(0, Math.min(k, hits.length)); + return hits; } diff --git a/packages/mnemopi/test/native-vector-parity.test.ts b/packages/mnemopi/test/native-vector-parity.test.ts new file mode 100644 index 000000000..490d7ab09 --- /dev/null +++ b/packages/mnemopi/test/native-vector-parity.test.ts @@ -0,0 +1,184 @@ +import { describe, expect, test } from "bun:test"; +import { + cosineSimilarityBatch, + cosineSimilarityPairs, + hammingDistanceBatch, + hammingDistanceForDimBatch, + mmrRerankIndices, + vectorIndexTopK, +} from "@oh-my-pi/pi-natives"; +import { hammingDistance, hammingDistanceForDimension } from "../src/core/binary-vectors"; +import { jaccardSimilarity } from "../src/core/mmr"; +import { cosineSimilarity } from "../src/core/vector-math"; + +/** Deterministic LCG so parity failures reproduce exactly. */ +function makeRng(seed: number): () => number { + let state = seed >>> 0; + return () => { + state = (state * 1664525 + 1013904223) >>> 0; + return state / 4294967296; + }; +} + +const REL_TOL = 1e-9; + +function expectClose(actual: number, expected: number): void { + if (expected === 0) { + expect(actual).toBe(0); + return; + } + expect(Math.abs(actual - expected) / Math.abs(expected)).toBeLessThanOrEqual(REL_TOL); +} + +describe("native vector kernel parity", () => { + test("cosineSimilarityBatch matches TS cosineSimilarity per row", () => { + const rng = makeRng(0xc051e); + const dim = 384; + const count = 200; + const query = Float64Array.from({ length: dim }, () => rng() * 2 - 1); + const candidates = new Float64Array(count * dim); + for (let i = 0; i < candidates.length; i += 1) candidates[i] = rng() * 2 - 1; + // Sprinkle non-finite values to exercise the finite_or_zero path. + candidates[3] = Number.NaN; + candidates[dim + 7] = Number.POSITIVE_INFINITY; + const scores = cosineSimilarityBatch(query, candidates, dim); + expect(scores.length).toBe(count); + for (let row = 0; row < count; row += 1) { + const expected = cosineSimilarity(query, candidates.subarray(row * dim, (row + 1) * dim)); + expectClose(scores[row] ?? Number.NaN, expected); + } + }); + + test("cosineSimilarityPairs matches the TS pairwise threshold loop", () => { + const rng = makeRng(0x9a175); + const dim = 64; + const count = 40; + const flat = new Float64Array(count * dim); + for (let i = 0; i < flat.length; i += 1) flat[i] = rng() * 2 - 1; + const threshold = 0.02; + const expected: number[] = []; + for (let i = 0; i < count; i += 1) { + for (let j = i + 1; j < count; j += 1) { + if (cosineSimilarity(flat.subarray(i * dim, (i + 1) * dim), flat.subarray(j * dim, (j + 1) * dim)) >= threshold) { + expected.push(i, j); + } + } + } + expect(Array.from(cosineSimilarityPairs(flat, count, dim, threshold))).toEqual(expected); + }); + + test("vectorIndexTopK matches the TS scoring loop and stable sort", () => { + const rng = makeRng(0x70b1); + const dims = 384; + const count = 300; + const matrix = new Float32Array(count * dims); + for (let i = 0; i < matrix.length; i += 1) matrix[i] = rng() * 2 - 1; + const query = Float64Array.from({ length: dims }, () => rng() * 2 - 1); + let normSq = 0; + for (const v of query) normSq += v * v; + const norm = Math.sqrt(normSq); + const hits: Array<{ row: number; score: number }> = []; + for (let row = 0; row < count; row += 1) { + let score = 0; + for (let col = 0; col < dims; col += 1) { + score += (matrix[row * dims + col] ?? 0) * ((query[col] ?? 0) / norm); + } + hits.push({ row, score }); + } + hits.sort((a, b) => b.score - a.score); + const limit = 25; + const result = vectorIndexTopK(matrix, dims, query, limit); + expect(Array.from(result.indices)).toEqual(hits.slice(0, limit).map(h => h.row)); + for (let i = 0; i < limit; i += 1) { + expectClose(result.scores[i] ?? Number.NaN, hits[i]?.score ?? Number.NaN); + } + }); + + test("hammingDistanceBatch is exactly equal to TS hammingDistance", () => { + const rng = makeRng(0xba7c4); + const stride = 48; // 384-dim binarized + const count = 128; + const query = Uint8Array.from({ length: stride }, () => Math.floor(rng() * 256)); + const packed = new Uint8Array(count * stride); + const lengths = new Uint32Array(count); + const vectors: Uint8Array[] = []; + for (let i = 0; i < count; i += 1) { + const len = i % 7 === 0 ? Math.floor(rng() * stride) : stride; // ragged rows + const vector = Uint8Array.from({ length: len }, () => Math.floor(rng() * 256)); + vectors.push(vector); + lengths[i] = len; + packed.set(vector, i * stride); + } + const distances = hammingDistanceBatch(query, packed, stride, lengths); + for (let i = 0; i < count; i += 1) { + expect(distances[i]).toBe(hammingDistance(query, vectors[i] ?? new Uint8Array())); + } + }); + + test("hammingDistanceForDimBatch is exactly equal to TS hammingDistanceForDimension", () => { + const rng = makeRng(0xd1235); + const stride = 48; + const count = 96; + const query = Uint8Array.from({ length: stride }, () => Math.floor(rng() * 256)); + const packed = new Uint8Array(count * stride); + const dims = new Uint32Array(count); + for (let i = 0; i < count; i += 1) { + dims[i] = Math.floor(rng() * (stride * 8 + 1)); // includes partial-byte tails and 0 + for (let b = 0; b < stride; b += 1) packed[i * stride + b] = Math.floor(rng() * 256); + } + const distances = hammingDistanceForDimBatch(query, packed, stride, dims); + for (let i = 0; i < count; i += 1) { + const row = packed.subarray(i * stride, (i + 1) * stride); + expect(distances[i]).toBe(hammingDistanceForDimension(query, row, dims[i] ?? 0)); + } + }); + + test("mmrRerankIndices selects identical index sequences to the TS loop", () => { + const rng = makeRng(0x33a11); + const words = ["alpha", "beta", "gamma", "delta", "epsilon", "zeta", "eta", "theta", "iota", "kappa"]; + const count = 60; + const contents: string[] = []; + const scores = new Float64Array(count); + for (let i = 0; i < count; i += 1) { + const n = 1 + Math.floor(rng() * 8); + contents.push( + Array.from({ length: n }, () => words[Math.floor(rng() * words.length)]) + .join(" "), + ); + scores[i] = rng(); + } + scores[5] = scores[6] = 0.5; // exercise strict-> tie keeping the earlier candidate + for (const lambda of [0.0, 0.3, 0.7, 1.0]) { + for (const topK of [1, 10, count, count + 5]) { + // TS reference selection over pre-sorted candidates (mirrors mmrRerank + // after its sort step). + const order = contents.map((_, i) => i).sort((a, b) => (scores[b] ?? 0) - (scores[a] ?? 0)); + const sortedContents = order.map(i => contents[i] ?? ""); + const sortedScores = order.map(i => scores[i] ?? 0); + const selected: number[] = [0]; + const remaining = sortedContents.map((_, i) => i).slice(1); + while (remaining.length > 0 && selected.length < topK) { + let bestIdx = 0; + let bestScore = Number.NEGATIVE_INFINITY; + for (let idx = 0; idx < remaining.length; idx += 1) { + const candidate = remaining[idx] ?? 0; + let maxSimilarity = 0; + for (const picked of selected) { + const sim = jaccardSimilarity(sortedContents[candidate] ?? "", sortedContents[picked] ?? ""); + if (sim > maxSimilarity) maxSimilarity = sim; + } + const mmrScore = lambda * (sortedScores[candidate] ?? 0) - (1 - lambda) * maxSimilarity; + if (mmrScore > bestScore) { + bestScore = mmrScore; + bestIdx = idx; + } + } + selected.push(remaining.splice(bestIdx, 1)[0] ?? 0); + } + if (selected.length < topK) selected.push(...remaining.slice(0, topK - selected.length)); + const native = mmrRerankIndices(sortedContents, Float64Array.from(sortedScores), lambda, topK); + expect(Array.from(native)).toEqual(selected.slice(0, topK)); + } + } + }); +}); diff --git a/packages/natives/native/index.d.ts b/packages/natives/native/index.d.ts index f38776c88..ee5f3df66 100644 --- a/packages/natives/native/index.d.ts +++ b/packages/natives/native/index.d.ts @@ -483,14 +483,13 @@ export declare function cosineSimilarityBatch(query: Float64Array, candidates: F /** * All pairs `(i, j)` with `i < j` whose cosine similarity meets `threshold`. * - * `vectors` is `count` vectors flattened row-major at `dim` `f32` elements - * per row (zero-padded). Element values are widened to `f64` before any - * arithmetic, exactly like JS reads from a `Float32Array`, so the similarity - * is bit-identical to the TS pairwise loop in `clusterBySimilarity`. - * Returns pairs flattened as `[i0, j0, i1, j1, ...]` in the same `(i, j)` - * visit order as the TS nested loop. + * `vectors` is `count` vectors flattened row-major at `dim` `f64` elements + * per row (zero-padded, which matches the TS `?? 0` missing-element + * semantics), so the similarity is bit-identical to the TS pairwise loop in + * `clusterBySimilarity`. Returns pairs flattened as `[i0, j0, i1, j1, ...]` + * in the same `(i, j)` visit order as the TS nested loop. */ -export declare function cosineSimilarityPairs(vectors: Float32Array, count: number, dim: number, threshold: number): Uint32Array +export declare function cosineSimilarityPairs(vectors: Float64Array, count: number, dim: number, threshold: number): Uint32Array /** * Count tokens in `input`.