perf(agent): maintain a persistent flat embedding buffer for BLAS/Accelerate search

blas_cosine_batch and accelerate_cosine_batch_vecs re-flattened the
entire Vec<Vec<f32>> corpus into a fresh Vec<f32> on every single
query before running the batch matmul — an O(N·dim) copy paid per
query when fast-math/accelerate/openblas is enabled, even though a
flat fast-path (blas_cosine_batch_flat / accelerate_cosine_batch)
already existed for pre-flattened input.

Add MemoryCache::embeddings_flat, a contiguous [N × embedding_dim]
buffer maintained incrementally in push/update/compact (O(1) amortized
append, O(dim) in-place overwrite, O(n) rebuild only on compact/bulk
load). schema.rs's direct-push load path calls the new rebuild_flat()
explicitly. flat_embeddings() now just clones the already-maintained
buffer instead of rebuilding it.

Thread the flat buffer through strategy::search_with_metrics as a new
vectors_flat parameter, used only by the Blas/Accelerate arms (now
calling the *_flat variants); other strategies are unaffected. No
current caller wires search_with_metrics into the production query
path yet (only its own tests exercise it) — this fixes the identified
per-query re-flatten and makes the flat buffer available for whenever
that wiring lands.

INT-16
This commit is contained in:
ClawHDF5 Coding Agent
2026-08-17 00:41:23 +00:00
parent 1efd82c841
commit 45a38ba260
3 changed files with 185 additions and 8 deletions
+145 -5
View File
@@ -7,6 +7,11 @@ use crate::vector_search;
pub struct MemoryCache {
pub chunks: Vec<String>,
pub embeddings: Vec<Vec<f32>>,
/// `embeddings` flattened into one contiguous `[N × embedding_dim]`
/// buffer, maintained incrementally alongside `embeddings` (push/update/
/// compact) so BLAS/Accelerate batch search can read it directly instead
/// of re-flattening the whole corpus on every query.
pub embeddings_flat: Vec<f32>,
pub source_channels: Vec<String>,
pub timestamps: Vec<f64>,
pub session_ids: Vec<String>,
@@ -24,6 +29,7 @@ impl MemoryCache {
Self {
chunks: Vec::new(),
embeddings: Vec::new(),
embeddings_flat: Vec::new(),
source_channels: Vec::new(),
timestamps: Vec::new(),
session_ids: Vec::new(),
@@ -35,6 +41,17 @@ impl MemoryCache {
}
}
/// Rebuild `embeddings_flat` from `embeddings` from scratch. Callers that
/// populate `embeddings` directly (bulk loads) must call this afterward.
pub fn rebuild_flat(&mut self) {
self.embeddings_flat.clear();
self.embeddings_flat
.reserve(self.embeddings.len() * self.embedding_dim);
for emb in &self.embeddings {
self.embeddings_flat.extend_from_slice(emb);
}
}
/// Total number of entries (including tombstoned).
pub fn len(&self) -> usize {
self.chunks.len()
@@ -62,6 +79,7 @@ impl MemoryCache {
let idx = self.chunks.len();
let norm = vector_search::compute_norm(&embedding);
self.chunks.push(chunk);
self.embeddings_flat.extend_from_slice(&embedding);
self.embeddings.push(embedding);
self.source_channels.push(source_channel);
self.timestamps.push(timestamp);
@@ -100,7 +118,20 @@ impl MemoryCache {
if idx < self.chunks.len() {
let norm = vector_search::compute_norm(&embedding);
self.chunks[idx] = chunk;
let dim = self.embedding_dim;
let flat_start = idx * dim;
let matches_dim =
embedding.len() == dim && flat_start + dim <= self.embeddings_flat.len();
self.embeddings[idx] = embedding;
if matches_dim {
self.embeddings_flat[flat_start..flat_start + dim]
.copy_from_slice(&self.embeddings[idx]);
} else {
// Embedding length doesn't match embedding_dim (shouldn't
// happen in practice) — fall back to a full rebuild rather
// than leave embeddings_flat misaligned with embeddings.
self.rebuild_flat();
}
self.source_channels[idx] = source_channel;
self.timestamps[idx] = timestamp;
self.session_ids[idx] = session_id;
@@ -173,16 +204,125 @@ impl MemoryCache {
self.tombstones = new_tombstones;
self.norms = new_norms;
self.activation_weights = new_activation_weights;
self.rebuild_flat();
(removed, index_map)
}
/// Flatten all embeddings into a single Vec<f32> for HDF5 storage.
/// `embeddings_flat` is already maintained incrementally, so this just
/// clones it — kept as a method for callers that want an owned copy.
pub fn flat_embeddings(&self) -> Vec<f32> {
let mut flat = Vec::with_capacity(self.embeddings.len() * self.embedding_dim);
for emb in &self.embeddings {
flat.extend_from_slice(emb);
}
flat
self.embeddings_flat.clone()
}
}
#[cfg(test)]
mod tests {
use super::*;
/// `embeddings_flat` must always equal a from-scratch flatten of `embeddings`.
fn assert_flat_in_sync(cache: &MemoryCache) {
let expected: Vec<f32> = cache.embeddings.iter().flatten().copied().collect();
assert_eq!(cache.embeddings_flat, expected);
}
#[test]
fn push_keeps_flat_buffer_in_sync() {
let mut cache = MemoryCache::new(3);
cache.push(
"a".into(),
vec![1.0, 2.0, 3.0],
"chan".into(),
0.0,
"s1".into(),
String::new(),
);
cache.push(
"b".into(),
vec![4.0, 5.0, 6.0],
"chan".into(),
1.0,
"s1".into(),
String::new(),
);
assert_flat_in_sync(&cache);
assert_eq!(cache.embeddings_flat, vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0]);
}
#[test]
fn update_keeps_flat_buffer_in_sync() {
let mut cache = MemoryCache::new(3);
cache.push(
"a".into(),
vec![1.0, 2.0, 3.0],
"chan".into(),
0.0,
"s1".into(),
String::new(),
);
cache.push(
"b".into(),
vec![4.0, 5.0, 6.0],
"chan".into(),
1.0,
"s1".into(),
String::new(),
);
cache.update(
0,
"a2".into(),
vec![7.0, 8.0, 9.0],
"chan".into(),
2.0,
"s1".into(),
);
assert_flat_in_sync(&cache);
assert_eq!(
cache.embeddings_flat,
vec![7.0, 8.0, 9.0, 4.0, 5.0, 6.0],
"update must overwrite the correct flat slice, not just append"
);
}
#[test]
fn compact_keeps_flat_buffer_in_sync() {
let mut cache = MemoryCache::new(2);
cache.push(
"a".into(),
vec![1.0, 1.0],
"chan".into(),
0.0,
"s1".into(),
String::new(),
);
cache.push(
"b".into(),
vec![2.0, 2.0],
"chan".into(),
1.0,
"s1".into(),
String::new(),
);
cache.push(
"c".into(),
vec![3.0, 3.0],
"chan".into(),
2.0,
"s1".into(),
String::new(),
);
cache.mark_deleted(1);
cache.compact();
assert_flat_in_sync(&cache);
assert_eq!(cache.embeddings_flat, vec![1.0, 1.0, 3.0, 3.0]);
}
#[test]
fn rebuild_flat_matches_manual_flatten() {
let mut cache = MemoryCache::new(2);
cache.embeddings = vec![vec![1.0, 2.0], vec![3.0, 4.0]];
cache.rebuild_flat();
assert_eq!(cache.embeddings_flat, vec![1.0, 2.0, 3.0, 4.0]);
}
}