Skip to content

Commit 2ee97ff

Browse files
authored
[diskann-garnet] Optimize reranking (#1308)
This addresses several performance issues with reranking. 1. The existence check turned out to be expensive as it consults the FSM. This reads 1 extra key per Lsearch which is unnecessary as we skip missing vectors anyway. 2. The full vector reads were serial. Using multiread allows Garnet to parallelize these. 3. The reranked candidates buffer was allocated in the post processor; pooling this allocation removes it from the query path. Since we don't have good recall based tests yet, I hand verified the recall was still ok with this change.
1 parent 60c4ac0 commit 2ee97ff

4 files changed

Lines changed: 43 additions & 26 deletions

File tree

‎Cargo.lock‎

Lines changed: 1 addition & 1 deletion
Some generated files are not rendered by default. Learn more about customizing how changed files appear on GitHub.

‎diskann-garnet/Cargo.toml‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,6 @@
11
[package]
22
name = "diskann-garnet"
3-
version = "4.0.3"
3+
version = "4.0.4"
44
edition = "2024"
55
authors.workspace = true
66
license.workspace = true

‎diskann-garnet/diskann-garnet.nuspec‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -2,7 +2,7 @@
22
<package>
33
<metadata>
44
<id>diskann-garnet</id>
5-
<version>4.0.3</version>
5+
<version>4.0.4</version>
66
<readme>docs/README.md</readme>
77
<authors>Microsoft</authors>
88
<projectUrl>https://github.com/microsoft/DiskANN</projectUrl>

‎diskann-garnet/src/provider.rs‎

Lines changed: 40 additions & 23 deletions
Original file line numberDiff line numberDiff line change
@@ -62,6 +62,9 @@ use crate::{
6262
/// bytes are the serialized quant table.
6363
const QUANT_STATE_KEY: u32 = u32::from_be_bytes(*b"_qnt");
6464

65+
/// Starting capacity of the pre-allocated rerank buffers.
66+
const RERANK_BUFFER_LENGTH: usize = 1024;
67+
6568
#[derive(Clone)]
6669
struct AdjList(AdjacencyList<u32>);
6770

@@ -138,6 +141,8 @@ pub(crate) struct GarnetProvider<T: VectorRepr> {
138141
/// Pool of pre-allocated buffers to use for filter decisions during
139142
/// filtered search beam expansion
140143
filtered_decisions_pool: ObjectPool<Vec<bool>>,
144+
/// Pool of pre-allocated buffers to use for reranking
145+
rerank_pool: ObjectPool<Vec<Neighbor<u32>>>,
141146
/// Pool of pre-allocated buffers to use for quantizing vectors
142147
quant_buffer_pool: ObjectPool<Vec<u8>>,
143148
/// Small cache for the start points' neighbors
@@ -173,6 +178,11 @@ impl<T: VectorRepr> GarnetProvider<T> {
173178
parallelism,
174179
Some(parallelism),
175180
);
181+
let rerank_pool = ObjectPool::new(
182+
Undef::new(RERANK_BUFFER_LENGTH),
183+
parallelism,
184+
Some(parallelism),
185+
);
176186

177187
let start_point_cache =
178188
DashMap::with_capacity_and_hasher(1, foldhash::fast::RandomState::default());
@@ -293,6 +303,7 @@ impl<T: VectorRepr> GarnetProvider<T> {
293303
id_buffer_pool,
294304
filtered_ids_pool,
295305
filtered_decisions_pool,
306+
rerank_pool,
296307
quant_buffer_pool,
297308
start_point_cache,
298309
start_point_quant_cache,
@@ -1302,34 +1313,40 @@ impl<'a, 'b, T: VectorRepr> SearchPostProcessStep<DynamicAccessor<'a, T>, &'b [T
13021313
.map_err(|e| GarnetProviderError::PostProcessing(Box::new(e)));
13031314
}
13041315

1305-
let provider = &accessor.provider;
1316+
let provider = accessor.provider;
13061317
let f = T::distance(provider.metric_type, Some(provider.dim));
1307-
let mut v = Poly::broadcast(0u8, provider.dim * mem::size_of::<T>(), AlignToEight)?;
1308-
1309-
// Filter before computing the full precision distances.
1310-
let mut reranked: Vec<_> = candidates
1311-
.filter_map(|n| {
1312-
if !provider.vector_iid_exists(accessor.context, *n.id()) {
1313-
None
1314-
} else if provider.callbacks.read_single_iid(
1315-
&accessor.context.term(Term::Vector),
1316-
*n.id(),
1317-
&mut v,
1318-
) {
1319-
Some(Neighbor::new(
1320-
*n.id(),
1321-
f.evaluate_similarity(query, bytemuck::cast_slice::<u8, T>(&v)),
1322-
))
1323-
} else {
1324-
None
1325-
}
1326-
})
1327-
.collect();
1318+
1319+
let mut reranked = provider
1320+
.rerank_pool
1321+
.get_ref(Undef::new(RERANK_BUFFER_LENGTH));
1322+
reranked.clear();
1323+
1324+
// Use the accessor.filtered_ids pre-allocated buffer to do a multi read from Garnet, placing the results in
1325+
// the rerank buffer.
1326+
accessor.filtered_ids.clear();
1327+
for nbor in candidates {
1328+
accessor.filtered_ids.push(4);
1329+
accessor.filtered_ids.push(*nbor.id());
1330+
}
1331+
1332+
if !accessor.filtered_ids.is_empty() {
1333+
provider.callbacks.read_multi_lpiid(
1334+
&accessor.context.term(Term::Vector),
1335+
&accessor.filtered_ids,
1336+
|i, v| {
1337+
let dist = f.evaluate_similarity(query, bytemuck::cast_slice::<u8, T>(v));
1338+
reranked.push(Neighbor::new(
1339+
accessor.filtered_ids[i as usize * 2 + 1],
1340+
dist,
1341+
));
1342+
},
1343+
);
1344+
}
13281345

13291346
// Sort the full precision distances.
13301347
reranked.sort_unstable_by(diskann::neighbor::ord::fast_distance);
13311348

1332-
next.post_process(accessor, query, reranked.into_iter(), output)
1349+
next.post_process(accessor, query, reranked.iter().copied(), output)
13331350
.await
13341351
.map_err(|e| GarnetProviderError::PostProcessing(Box::new(e)))
13351352
}

0 commit comments

Comments
 (0)