Skip to content

Commit 9683089

Browse files
authored
[PERF]: Parallelize fetching blocks for brute force regex (#5051)
## Description of changes _Summarize the changes made by this PR._ - Improvements & Bug fixes - Brute forcing list of ids or all the ids was sequential. We now parallelize all the block gets and then perform regex - New functionality - ... ## Test plan _How are these changes tested?_ - [x] Tests pass locally with `pytest` for python, `yarn test` for js, `cargo test` for rust ## Documentation Changes None
1 parent 80f3783 commit 9683089

5 files changed

Lines changed: 82 additions & 69 deletions

File tree

rust/blockstore/src/arrow/blockfile.rs

Lines changed: 11 additions & 16 deletions
Original file line numberDiff line numberDiff line change
@@ -15,7 +15,7 @@ use chroma_cache::AysncPartitionedMutex;
1515
use chroma_error::ChromaError;
1616
use chroma_error::ErrorCodes;
1717
use chroma_storage::admissioncontrolleds3::StorageRequestPriority;
18-
use futures::future::join_all;
18+
use futures::future::{join_all, try_join_all};
1919
use futures::{Stream, StreamExt, TryStreamExt};
2020
use parking_lot::{Mutex, RwLock};
2121
use std::collections::HashSet;
@@ -625,24 +625,19 @@ impl<'me, K: ArrowReadableKey<'me> + Into<KeyWrapper>, V: ArrowReadableValue<'me
625625
.sparse_index
626626
.get_block_ids_range(prefix_range.clone());
627627

628-
let mut result: Vec<(&str, K, V)> = vec![];
629-
for block_id in block_ids {
630-
let block_opt = match self.get_block(block_id, StorageRequestPriority::P0).await {
631-
Ok(Some(block)) => Some(block),
628+
let block_futures = block_ids.into_iter().map(|block_id| async move {
629+
match self.get_block(block_id, StorageRequestPriority::P0).await {
630+
Ok(Some(block)) => Ok(block),
632631
Ok(None) => {
633-
return Err(Box::new(ArrowBlockfileError::BlockNotFound));
634-
}
635-
Err(e) => {
636-
return Err(Box::new(e));
632+
Err(Box::new(ArrowBlockfileError::BlockNotFound) as Box<dyn ChromaError>)
637633
}
638-
};
634+
Err(e) => Err(Box::new(e) as Box<dyn ChromaError>),
635+
}
636+
});
639637

640-
let block = match block_opt {
641-
Some(b) => b,
642-
None => {
643-
return Err(Box::new(ArrowBlockfileError::BlockNotFound));
644-
}
645-
};
638+
let blocks = try_join_all(block_futures).await?;
639+
let mut result: Vec<(&str, K, V)> = vec![];
640+
for block in blocks {
646641
result.extend(block.get_range(prefix_range.clone(), key_range.clone()));
647642
}
648643

rust/segment/src/blockfile_metadata.rs

Lines changed: 10 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -1598,16 +1598,16 @@ mod test {
15981598
.await
15991599
.expect("Record segment get all data failed");
16001600
assert_eq!(res.len(), 2);
1601-
res.sort_by(|x, y| x.id.cmp(y.id));
1601+
res.sort_by(|x, y| x.1.id.cmp(y.1.id));
16021602
let mut id1_mt = HashMap::new();
16031603
id1_mt.insert(
16041604
String::from("hello"),
16051605
MetadataValue::Str(String::from("new world")),
16061606
);
1607-
assert_eq!(res.first().as_ref().unwrap().metadata, Some(id1_mt));
1607+
assert_eq!(res.first().as_ref().unwrap().1.metadata, Some(id1_mt));
16081608
let mut id2_mt = HashMap::new();
16091609
id2_mt.insert(String::from("hello"), MetadataValue::Float(1.0));
1610-
assert_eq!(res.get(1).as_ref().unwrap().metadata, Some(id2_mt));
1610+
assert_eq!(res.get(1).as_ref().unwrap().1.metadata, Some(id2_mt));
16111611
}
16121612

16131613
#[tokio::test]
@@ -1840,13 +1840,13 @@ mod test {
18401840
.await
18411841
.expect("Record segment get all data failed");
18421842
assert_eq!(res.len(), 1);
1843-
res.sort_by(|x, y| x.id.cmp(y.id));
1843+
res.sort_by(|x, y| x.1.id.cmp(y.1.id));
18441844
let mut id1_mt = HashMap::new();
18451845
id1_mt.insert(
18461846
String::from("bye"),
18471847
MetadataValue::Str(String::from("world")),
18481848
);
1849-
assert_eq!(res.first().as_ref().unwrap().metadata, Some(id1_mt));
1849+
assert_eq!(res.first().as_ref().unwrap().1.metadata, Some(id1_mt));
18501850
}
18511851

18521852
#[tokio::test]
@@ -2069,9 +2069,9 @@ mod test {
20692069
.await
20702070
.expect("Record segment get all data failed");
20712071
assert_eq!(res.len(), 1);
2072-
res.sort_by(|x, y| x.id.cmp(y.id));
2072+
res.sort_by(|x, y| x.1.id.cmp(y.1.id));
20732073
assert_eq!(
2074-
res.first().as_ref().unwrap().document,
2074+
res.first().as_ref().unwrap().1.document,
20752075
Some(String::from("bye").as_str())
20762076
);
20772077
}
@@ -2269,13 +2269,13 @@ mod test {
22692269
.await
22702270
.expect("Record segment get all data failed");
22712271
assert_eq!(res.len(), 2);
2272-
res.sort_by(|x, y| x.id.cmp(y.id));
2272+
res.sort_by(|x, y| x.1.id.cmp(y.1.id));
22732273
assert_eq!(
2274-
res.first().as_ref().unwrap().document,
2274+
res.first().as_ref().unwrap().1.document,
22752275
Some(String::from("hello").as_str())
22762276
);
22772277
assert_eq!(
2278-
res.get(1).as_ref().unwrap().document,
2278+
res.get(1).as_ref().unwrap().1.document,
22792279
Some(String::from("world").as_str())
22802280
);
22812281
}

rust/segment/src/blockfile_record.rs

Lines changed: 6 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -951,12 +951,12 @@ impl RecordSegmentReader<'_> {
951951
}
952952

953953
/// Returns all data in the record segment, sorted by their offset ids
954-
#[allow(dead_code)]
955-
pub async fn get_all_data(&self) -> Result<Vec<DataRecord>, Box<dyn ChromaError>> {
956-
self.id_to_data
957-
.get_range(""..="", ..)
958-
.await
959-
.map(|vec| vec.into_iter().map(|(_, _, data)| data).collect())
954+
pub async fn get_all_data(&self) -> Result<Vec<(u32, DataRecord)>, Box<dyn ChromaError>> {
955+
self.id_to_data.get_range(""..="", ..).await.map(|vec| {
956+
vec.into_iter()
957+
.map(|(_, offset, data)| (offset, data))
958+
.collect()
959+
})
960960
}
961961

962962
pub async fn get_data_stream<'me>(

rust/segment/src/types.rs

Lines changed: 37 additions & 24 deletions
Original file line numberDiff line numberDiff line change
@@ -1337,10 +1337,10 @@ mod tests {
13371337
.expect("Get all data failed");
13381338
assert_eq!(all_data.len(), 1);
13391339
let record = &all_data[0];
1340-
assert_eq!(record.id, "embedding_id_1");
1341-
assert_eq!(record.document, Some("number"));
1342-
assert_eq!(record.embedding, &[7.0, 8.0, 9.0]);
1343-
assert_eq!(record.metadata, Some(res_metadata));
1340+
assert_eq!(record.1.id, "embedding_id_1");
1341+
assert_eq!(record.1.document, Some("number"));
1342+
assert_eq!(record.1.embedding, &[7.0, 8.0, 9.0]);
1343+
assert_eq!(record.1.metadata, Some(res_metadata));
13441344
// Search by metadata filter.
13451345
let metadata_segment_reader =
13461346
MetadataSegmentReader::from_segment(&metadata_segment, &blockfile_provider)
@@ -1621,10 +1621,10 @@ mod tests {
16211621
.expect("Get all data failed");
16221622
assert_eq!(all_data.len(), 1);
16231623
let record = &all_data[0];
1624-
assert_eq!(record.id, "embedding_id_1");
1625-
assert_eq!(record.document, Some("doc1"));
1626-
assert_eq!(record.embedding, &[7.0, 8.0, 9.0]);
1627-
assert_eq!(record.metadata, Some(res_metadata));
1624+
assert_eq!(record.1.id, "embedding_id_1");
1625+
assert_eq!(record.1.document, Some("doc1"));
1626+
assert_eq!(record.1.embedding, &[7.0, 8.0, 9.0]);
1627+
assert_eq!(record.1.metadata, Some(res_metadata));
16281628
// Search by metadata filter.
16291629
let metadata_segment_reader =
16301630
MetadataSegmentReader::from_segment(&metadata_segment, &blockfile_provider)
@@ -1926,10 +1926,10 @@ mod tests {
19261926
.expect("Get all data failed");
19271927
assert_eq!(all_data.len(), 1);
19281928
let record = &all_data[0];
1929-
assert_eq!(record.id, "embedding_id_1");
1930-
assert_eq!(record.document, Some("number"));
1931-
assert_eq!(record.embedding, &[7.0, 8.0, 9.0]);
1932-
assert_eq!(record.metadata, Some(res_metadata));
1929+
assert_eq!(record.1.id, "embedding_id_1");
1930+
assert_eq!(record.1.document, Some("number"));
1931+
assert_eq!(record.1.embedding, &[7.0, 8.0, 9.0]);
1932+
assert_eq!(record.1.metadata, Some(res_metadata));
19331933
// Search by metadata filter.
19341934
let metadata_segment_reader =
19351935
MetadataSegmentReader::from_segment(&metadata_segment, &blockfile_provider)
@@ -2271,70 +2271,83 @@ mod tests {
22712271
.await
22722272
.expect("Get all data failed");
22732273
for data in all_data {
2274-
assert_ne!(data.id, "embedding_id_2");
2275-
if data.id == "embedding_id_1" {
2274+
assert_ne!(data.1.id, "embedding_id_2");
2275+
if data.1.id == "embedding_id_1" {
22762276
assert!(data
2277+
.1
22772278
.metadata
22782279
.clone()
22792280
.expect("Metadata is empty")
22802281
.contains_key("hello"),);
22812282
assert_eq!(
2282-
data.metadata
2283+
data.1
2284+
.metadata
22832285
.clone()
22842286
.expect("Metadata is empty")
22852287
.get("hello"),
22862288
Some(&MetadataValue::Str(String::from("new_world")))
22872289
);
22882290
assert!(data
2291+
.1
22892292
.metadata
22902293
.clone()
22912294
.expect("Metadata is empty")
22922295
.contains_key("bye"),);
22932296
assert_eq!(
2294-
data.metadata.clone().expect("Metadata is empty").get("bye"),
2297+
data.1
2298+
.metadata
2299+
.clone()
2300+
.expect("Metadata is empty")
2301+
.get("bye"),
22952302
Some(&MetadataValue::Str(String::from("world")))
22962303
);
22972304
assert!(data
2305+
.1
22982306
.metadata
22992307
.clone()
23002308
.expect("Metadata is empty")
23012309
.contains_key("hello_again"),);
23022310
assert_eq!(
2303-
data.metadata
2311+
data.1
2312+
.metadata
23042313
.clone()
23052314
.expect("Metadata is empty")
23062315
.get("hello_again"),
23072316
Some(&MetadataValue::Str(String::from("new_world")))
23082317
);
2309-
assert_eq!(data.document.expect("Non empty document"), "doc1");
2310-
assert_eq!(data.embedding, vec![1.0, 2.0, 3.0]);
2311-
} else if data.id == "embedding_id_3" {
2318+
assert_eq!(data.1.document.expect("Non empty document"), "doc1");
2319+
assert_eq!(data.1.embedding, vec![1.0, 2.0, 3.0]);
2320+
} else if data.1.id == "embedding_id_3" {
23122321
assert!(data
2322+
.1
23132323
.metadata
23142324
.clone()
23152325
.expect("Metadata is empty")
23162326
.contains_key("hello"),);
23172327
assert_eq!(
2318-
data.metadata
2328+
data.1
2329+
.metadata
23192330
.clone()
23202331
.expect("Metadata is empty")
23212332
.get("hello"),
23222333
Some(&MetadataValue::Str(String::from("new_world")))
23232334
);
23242335
assert!(data
2336+
.1
23252337
.metadata
23262338
.clone()
23272339
.expect("Metadata is empty")
23282340
.contains_key("hello_again"),);
23292341
assert_eq!(
2330-
data.metadata
2342+
data.1
2343+
.metadata
23312344
.clone()
23322345
.expect("Metadata is empty")
23332346
.get("hello_again"),
23342347
Some(&MetadataValue::Str(String::from("new_world")))
23352348
);
2336-
assert_eq!(data.document.expect("Non empty document"), "doc3");
2337-
assert_eq!(data.embedding, vec![7.0, 8.0, 9.0]);
2349+
assert_eq!(data.1.document.expect("Non empty document"), "doc3");
2350+
assert_eq!(data.1.embedding, vec![7.0, 8.0, 9.0]);
23382351
}
23392352
}
23402353
}

rust/worker/src/execution/operators/filter.rs

Lines changed: 18 additions & 13 deletions
Original file line numberDiff line numberDiff line change
@@ -19,11 +19,11 @@ use chroma_types::{
1919
literal_expr::{LiteralExpr, NgramLiteralProvider},
2020
ChromaRegex, ChromaRegexError,
2121
},
22-
BooleanOperator, Chunk, CompositeExpression, DocumentExpression, DocumentOperator, LogRecord,
23-
MaterializedLogOperation, MetadataComparison, MetadataExpression, MetadataSetValue,
22+
BooleanOperator, Chunk, CompositeExpression, DataRecord, DocumentExpression, DocumentOperator,
23+
LogRecord, MaterializedLogOperation, MetadataComparison, MetadataExpression, MetadataSetValue,
2424
MetadataValue, PrimitiveOperator, Segment, SetOperator, SignedRoaringBitmap, Where,
2525
};
26-
use futures::TryStreamExt;
26+
use futures::future::try_join_all;
2727
use roaring::RoaringBitmap;
2828
use thiserror::Error;
2929
use tracing::{Instrument, Span};
@@ -245,22 +245,27 @@ impl<'me> MetadataProvider<'me> {
245245
Some(offset_ids)
246246
if offset_ids.len() < rec_reader.count().await? as u64 / 10 =>
247247
{
248-
for id in offset_ids {
249-
if rec_reader.get_data_for_offset_id(id).await?.is_some_and(
250-
|rec| rec.document.is_some_and(|doc| regex.is_match(doc)),
251-
) {
248+
let fetch_futures: Vec<_> = offset_ids
249+
.into_iter()
250+
.map(|id| async move {
251+
let data = rec_reader.get_data_for_offset_id(id).await?;
252+
Ok::<(u32, Option<DataRecord>), Box<dyn ChromaError>>((
253+
id, data,
254+
))
255+
})
256+
.collect();
257+
let data_results = try_join_all(fetch_futures).await?;
258+
for (id, data_opt) in data_results {
259+
if data_opt.is_some_and(|rec| {
260+
rec.document.is_some_and(|doc| regex.is_match(doc))
261+
}) {
252262
exact_matching_offset_ids.insert(id);
253263
}
254264
}
255265
}
256266
// Perform range scan of all documents
257267
candidate_offsets => {
258-
for (offset, record) in rec_reader
259-
.get_data_stream(..)
260-
.await
261-
.try_collect::<Vec<_>>()
262-
.await?
263-
{
268+
for (offset, record) in rec_reader.get_all_data().await? {
264269
if (candidate_offsets.is_none()
265270
|| candidate_offsets
266271
.as_ref()

0 commit comments

Comments
 (0)