diff --git a/encodings/alp/src/alp/ops.rs b/encodings/alp/src/alp/ops.rs index 58b5ee99360..11d6edee5eb 100644 --- a/encodings/alp/src/alp/ops.rs +++ b/encodings/alp/src/alp/ops.rs @@ -23,7 +23,7 @@ impl OperationsVTable for ALP { ctx: &mut ExecutionCtx, ) -> VortexResult { if let Some(patches) = array.patches() - && let Some(patch) = patches.get_patched(index)? + && let Some(patch) = patches.get_patched(index, ctx)? { return patch.cast(array.dtype()); } diff --git a/encodings/alp/src/alp_rd/ops.rs b/encodings/alp/src/alp_rd/ops.rs index 1433966a01e..550d5022d80 100644 --- a/encodings/alp/src/alp_rd/ops.rs +++ b/encodings/alp/src/alp_rd/ops.rs @@ -24,7 +24,7 @@ impl OperationsVTable for ALPRD { // The left value can either be a direct value, or an exception. // The exceptions array represents exception positions with non-null values. let maybe_patched_value = match array.left_parts_patches() { - Some(patches) => patches.get_patched(index)?, + Some(patches) => patches.get_patched(index, ctx)?, None => None, }; let left = match maybe_patched_value { diff --git a/encodings/fastlanes/src/bitpacking/vtable/operations.rs b/encodings/fastlanes/src/bitpacking/vtable/operations.rs index 2816407ac03..95d772f73f1 100644 --- a/encodings/fastlanes/src/bitpacking/vtable/operations.rs +++ b/encodings/fastlanes/src/bitpacking/vtable/operations.rs @@ -16,11 +16,11 @@ impl OperationsVTable for BitPacked { fn scalar_at( array: ArrayView<'_, BitPacked>, index: usize, - _ctx: &mut ExecutionCtx, + ctx: &mut ExecutionCtx, ) -> VortexResult { Ok( if let Some(patches) = array.patches() - && let Some(patch) = patches.get_patched(index)? + && let Some(patch) = patches.get_patched(index, ctx)? { patch } else { diff --git a/encodings/sparse/src/ops.rs b/encodings/sparse/src/ops.rs index acabbb103d5..41a64379791 100644 --- a/encodings/sparse/src/ops.rs +++ b/encodings/sparse/src/ops.rs @@ -16,11 +16,11 @@ impl OperationsVTable for Sparse { fn scalar_at( array: ArrayView<'_, Sparse>, index: usize, - _ctx: &mut ExecutionCtx, + ctx: &mut ExecutionCtx, ) -> VortexResult { Ok(array .patches() - .get_patched(index)? + .get_patched(index, ctx)? .unwrap_or_else(|| array.fill_scalar().clone())) } } diff --git a/fuzz/src/array/mod.rs b/fuzz/src/array/mod.rs index e513c5daf81..0fd527863cd 100644 --- a/fuzz/src/array/mod.rs +++ b/fuzz/src/array/mod.rs @@ -63,6 +63,7 @@ use vortex_array::scalar_fn::fns::operators::CompareOperator; use vortex_array::scalar_fn::fns::operators::Operator; use vortex_array::search_sorted::SearchResult; use vortex_array::search_sorted::SearchSorted; +use vortex_array::search_sorted::SearchSortedArray; use vortex_array::search_sorted::SearchSortedSide; use vortex_btrblocks::BtrBlocksCompressor; #[cfg(feature = "zstd")] @@ -636,7 +637,7 @@ pub fn run_fuzz_action(fuzz_action: FuzzArrayAction) -> VortexFuzzResult { if !current_array.is_canonical() { sorted = compress_array(&sorted, CompressorStrategy::Default, &mut ctx); } - assert_search_sorted(sorted, s, side, expected.search(), i)?; + assert_search_sorted(sorted, s, side, expected.search(), i, &mut ctx)?; } Action::Filter(mask_val) => { current_array = current_array @@ -715,8 +716,9 @@ fn assert_search_sorted( side: SearchSortedSide, expected: SearchResult, step: usize, + ctx: &mut ExecutionCtx, ) -> VortexFuzzResult<()> { - let search_result = array + let search_result = SearchSortedArray::new(&array, ctx) .search_sorted(&s, side) .map_err(|e| VortexFuzzError::VortexError(e, Backtrace::capture()))?; if search_result != expected { diff --git a/vortex-array/benches/patches_lookup.rs b/vortex-array/benches/patches_lookup.rs index 262e3144495..cda8ff98238 100644 --- a/vortex-array/benches/patches_lookup.rs +++ b/vortex-array/benches/patches_lookup.rs @@ -9,6 +9,8 @@ use rand::RngExt; use rand::SeedableRng; use rand::rngs::StdRng; use vortex_array::IntoArray; +use vortex_array::VortexSessionExecute; +use vortex_array::array_session; use vortex_array::patches::PATCH_CHUNK_SIZE; use vortex_array::patches::Patches; use vortex_buffer::Buffer; @@ -100,10 +102,10 @@ fn queries_full_range() -> Vec { fn bench_search_index(bencher: Bencher, patches: Patches, queries: Vec) { bencher - .with_inputs(|| (&patches, &queries)) - .bench_refs(|(patches, queries)| { + .with_inputs(|| (&patches, &queries, array_session().create_execution_ctx())) + .bench_refs(|(patches, queries, ctx)| { for &q in queries.iter() { - divan::black_box(patches.search_index(q).unwrap()); + divan::black_box(patches.search_index(q, ctx).unwrap()); } }); } diff --git a/vortex-array/src/patches.rs b/vortex-array/src/patches.rs index 47662e36fe1..f4ac5c7d45f 100644 --- a/vortex-array/src/patches.rs +++ b/vortex-array/src/patches.rs @@ -442,14 +442,14 @@ impl Patches { } /// Get the patched value at a given index if it exists. - #[allow(clippy::disallowed_methods)] - pub fn get_patched(&self, index: usize) -> VortexResult> { - self.search_index(index)? + pub fn get_patched( + &self, + index: usize, + ctx: &mut ExecutionCtx, + ) -> VortexResult> { + self.search_index(index, ctx)? .to_found() - .map(|patch_idx| { - self.values() - .execute_scalar(patch_idx, &mut legacy_session().create_execution_ctx()) - }) + .map(|patch_idx| self.values().execute_scalar(patch_idx, ctx)) .transpose() } @@ -461,6 +461,7 @@ impl Patches { /// /// # Arguments /// * `index` - The index to search for + /// * `ctx` - The execution context used to read `indices` and `chunk_offsets` /// /// # Returns /// * [`SearchResult::Found(patch_idx)`] - If a patch exists at this index, returns the @@ -470,12 +471,12 @@ impl Patches { /// /// [`SearchResult::Found(patch_idx)`]: SearchResult::Found /// [`SearchResult::NotFound(insertion_point)`]: SearchResult::NotFound - pub fn search_index(&self, index: usize) -> VortexResult { + pub fn search_index(&self, index: usize, ctx: &mut ExecutionCtx) -> VortexResult { if self.chunk_offsets.is_some() { - return self.search_index_chunked(index); + return self.search_index_chunked(index, ctx); } - search_index_binary_search(&self.indices, index + self.offset) + search_index_binary_search(&self.indices, index + self.offset, ctx) } /// Constant time searches for `index` in the indices array. @@ -487,7 +488,11 @@ impl Patches { /// or the insertion point if not found. /// /// Returns an error if `chunk_offsets` or `offset_within_chunk` are not set. - fn search_index_chunked(&self, index: usize) -> VortexResult { + fn search_index_chunked( + &self, + index: usize, + ctx: &mut ExecutionCtx, + ) -> VortexResult { let Some(chunk_offsets) = &self.chunk_offsets else { vortex_bail!("chunk_offsets is required to be set") }; @@ -502,10 +507,21 @@ impl Patches { let chunk_idx = (index + self.offset % PATCH_CHUNK_SIZE) / PATCH_CHUNK_SIZE; + // The three reads below are of the same array, so they share one probe rather than + // building one per read as `Self::chunk_offset_at` does. + let mut probe = chunk_offsets.repeated_probe(); + let mut chunk_offset_at = |idx: usize| -> VortexResult { + probe + .execute_scalar(idx, ctx)? + .as_primitive() + .as_::() + .ok_or_else(|| vortex_err!("chunk offset does not fit in usize")) + }; + // Patch index offsets are absolute and need to be offset by the first chunk of the current slice. - let base_offset = self.chunk_offset_at(0)?; + let base_offset = chunk_offset_at(0)?; - let patches_start_idx = (self.chunk_offset_at(chunk_idx)? - base_offset) + let patches_start_idx = (chunk_offset_at(chunk_idx)? - base_offset) // Chunk offsets are only sliced off in case the slice is fully // outside of the chunk range. // @@ -515,7 +531,7 @@ impl Patches { .saturating_sub(offset_within_chunk); let patches_end_idx = if chunk_idx < chunk_offsets.len() - 1 { - (self.chunk_offset_at(chunk_idx + 1)? - base_offset) + (chunk_offset_at(chunk_idx + 1)? - base_offset) .saturating_sub(offset_within_chunk) .min(self.indices.len()) } else { @@ -523,7 +539,7 @@ impl Patches { }; let chunk_indices = self.indices.slice(patches_start_idx..patches_end_idx)?; - let result = search_index_binary_search(&chunk_indices, index + self.offset)?; + let result = search_index_binary_search(&chunk_indices, index + self.offset, ctx)?; Ok(match result { SearchResult::Found(idx) => SearchResult::Found(patches_start_idx + idx), @@ -729,8 +745,9 @@ impl Patches { /// Slice the patches by a range of the patched array. #[allow(clippy::disallowed_methods)] pub fn slice(&self, range: Range) -> VortexResult> { - let slice_start_idx = self.search_index(range.start)?.to_index(); - let slice_end_idx = self.search_index(range.end)?.to_index(); + let mut ctx = legacy_session().create_execution_ctx(); + let slice_start_idx = self.search_index(range.start, &mut ctx)?.to_index(); + let slice_end_idx = self.search_index(range.end, &mut ctx)?.to_index(); if slice_start_idx == slice_end_idx { return Ok(None); @@ -754,7 +771,7 @@ impl Patches { .as_ref() .map(|new_chunk_offsets| -> VortexResult { let new_chunk_base = new_chunk_offsets - .execute_scalar(0, &mut legacy_session().create_execution_ctx())? + .execute_scalar(0, &mut ctx)? .as_primitive() .as_::() .ok_or_else(|| vortex_err!("chunk offset does not fit in usize"))?; @@ -997,7 +1014,11 @@ impl Patches { /// # Returns /// [`SearchResult::Found`] with the position if needle exists, or [`SearchResult::NotFound`] /// with the insertion point if not found. -fn search_index_binary_search(indices: &ArrayRef, needle: usize) -> VortexResult { +fn search_index_binary_search( + indices: &ArrayRef, + needle: usize, + ctx: &mut ExecutionCtx, +) -> VortexResult { if let Some(primitive) = indices.as_opt::() { match_each_unsigned_integer_ptype!(primitive.ptype(), |T| { let Ok(needle) = T::try_from(needle) else { @@ -1013,16 +1034,16 @@ fn search_index_binary_search(indices: &ArrayRef, needle: usize) -> VortexResult }); } - search_index_binary_search_scalar(indices, needle) + search_index_binary_search_scalar(indices, needle, ctx) } -#[allow(clippy::disallowed_methods)] fn search_index_binary_search_scalar( indices: &ArrayRef, needle: usize, + ctx: &mut ExecutionCtx, ) -> VortexResult { match_each_unsigned_integer_ptype!(indices.dtype().as_ptype(), |T| { - SearchSortedPrimitiveArray::::new(indices, &mut legacy_session().create_execution_ctx()) + SearchSortedPrimitiveArray::::new(indices, ctx) .search_sorted(&needle, SearchSortedSide::Left) }) } @@ -1904,6 +1925,7 @@ mod test { #[test] fn test_search_index() { + let mut ctx = array_session().create_execution_ctx(); let patches = Patches::new( 10, 0, @@ -1914,15 +1936,36 @@ mod test { .unwrap(); // Search for exact indices - assert_eq!(patches.search_index(2).unwrap(), SearchResult::Found(0)); - assert_eq!(patches.search_index(5).unwrap(), SearchResult::Found(1)); - assert_eq!(patches.search_index(8).unwrap(), SearchResult::Found(2)); + assert_eq!( + patches.search_index(2, &mut ctx).unwrap(), + SearchResult::Found(0) + ); + assert_eq!( + patches.search_index(5, &mut ctx).unwrap(), + SearchResult::Found(1) + ); + assert_eq!( + patches.search_index(8, &mut ctx).unwrap(), + SearchResult::Found(2) + ); // Search for non-patch indices - assert_eq!(patches.search_index(0).unwrap(), SearchResult::NotFound(0)); - assert_eq!(patches.search_index(3).unwrap(), SearchResult::NotFound(1)); - assert_eq!(patches.search_index(6).unwrap(), SearchResult::NotFound(2)); - assert_eq!(patches.search_index(9).unwrap(), SearchResult::NotFound(3)); + assert_eq!( + patches.search_index(0, &mut ctx).unwrap(), + SearchResult::NotFound(0) + ); + assert_eq!( + patches.search_index(3, &mut ctx).unwrap(), + SearchResult::NotFound(1) + ); + assert_eq!( + patches.search_index(6, &mut ctx).unwrap(), + SearchResult::NotFound(2) + ); + assert_eq!( + patches.search_index(9, &mut ctx).unwrap(), + SearchResult::NotFound(3) + ); } #[test] @@ -2119,6 +2162,7 @@ mod test { #[test] fn test_chunk_offsets_search() { + let mut ctx = array_session().create_execution_ctx(); let indices = buffer![100u64, 200, 3000, 3100].into_array(); let values = buffer![10i32, 20, 30, 40].into_array(); let chunk_offsets = buffer![0u64, 2, 2, 3].into_array(); @@ -2127,31 +2171,44 @@ mod test { assert!(patches.chunk_offsets.is_some()); // chunk 0: patches at 100, 200 - assert_eq!(patches.search_index(100).unwrap(), SearchResult::Found(0)); - assert_eq!(patches.search_index(200).unwrap(), SearchResult::Found(1)); + assert_eq!( + patches.search_index(100, &mut ctx).unwrap(), + SearchResult::Found(0) + ); + assert_eq!( + patches.search_index(200, &mut ctx).unwrap(), + SearchResult::Found(1) + ); // chunks 1, 2: no patches assert_eq!( - patches.search_index(1500).unwrap(), + patches.search_index(1500, &mut ctx).unwrap(), SearchResult::NotFound(2) ); assert_eq!( - patches.search_index(2000).unwrap(), + patches.search_index(2000, &mut ctx).unwrap(), SearchResult::NotFound(2) ); // chunk 3: patches at 3000, 3100 - assert_eq!(patches.search_index(3000).unwrap(), SearchResult::Found(2)); - assert_eq!(patches.search_index(3100).unwrap(), SearchResult::Found(3)); + assert_eq!( + patches.search_index(3000, &mut ctx).unwrap(), + SearchResult::Found(2) + ); + assert_eq!( + patches.search_index(3100, &mut ctx).unwrap(), + SearchResult::Found(3) + ); assert_eq!( - patches.search_index(1024).unwrap(), + patches.search_index(1024, &mut ctx).unwrap(), SearchResult::NotFound(2) ); } #[test] fn test_chunk_offsets_with_slice() { + let mut ctx = array_session().create_execution_ctx(); let indices = buffer![100u64, 500, 1200, 1300, 1500, 1800, 2100, 2500].into_array(); let values = buffer![10i32, 20, 30, 35, 40, 45, 50, 60].into_array(); let chunk_offsets = buffer![0u64, 2, 6].into_array(); @@ -2162,15 +2219,28 @@ mod test { assert!(sliced.chunk_offsets.is_some()); assert_eq!(sliced.offset(), 1000); - assert_eq!(sliced.search_index(200).unwrap(), SearchResult::Found(0)); - assert_eq!(sliced.search_index(500).unwrap(), SearchResult::Found(2)); - assert_eq!(sliced.search_index(1100).unwrap(), SearchResult::Found(4)); + assert_eq!( + sliced.search_index(200, &mut ctx).unwrap(), + SearchResult::Found(0) + ); + assert_eq!( + sliced.search_index(500, &mut ctx).unwrap(), + SearchResult::Found(2) + ); + assert_eq!( + sliced.search_index(1100, &mut ctx).unwrap(), + SearchResult::Found(4) + ); - assert_eq!(sliced.search_index(250).unwrap(), SearchResult::NotFound(1)); + assert_eq!( + sliced.search_index(250, &mut ctx).unwrap(), + SearchResult::NotFound(1) + ); } #[test] fn test_chunk_offsets_with_slice_after_first_chunk() { + let mut ctx = array_session().create_execution_ctx(); let indices = buffer![100u64, 500, 1200, 1300, 1500, 1800, 2100, 2500].into_array(); let values = buffer![10i32, 20, 30, 35, 40, 45, 50, 60].into_array(); let chunk_offsets = buffer![0u64, 2, 6].into_array(); @@ -2181,11 +2251,26 @@ mod test { assert!(sliced.chunk_offsets.is_some()); assert_eq!(sliced.offset(), 1300); - assert_eq!(sliced.search_index(0).unwrap(), SearchResult::Found(0)); - assert_eq!(sliced.search_index(200).unwrap(), SearchResult::Found(1)); - assert_eq!(sliced.search_index(500).unwrap(), SearchResult::Found(2)); - assert_eq!(sliced.search_index(250).unwrap(), SearchResult::NotFound(2)); - assert_eq!(sliced.search_index(900).unwrap(), SearchResult::NotFound(4)); + assert_eq!( + sliced.search_index(0, &mut ctx).unwrap(), + SearchResult::Found(0) + ); + assert_eq!( + sliced.search_index(200, &mut ctx).unwrap(), + SearchResult::Found(1) + ); + assert_eq!( + sliced.search_index(500, &mut ctx).unwrap(), + SearchResult::Found(2) + ); + assert_eq!( + sliced.search_index(250, &mut ctx).unwrap(), + SearchResult::NotFound(2) + ); + assert_eq!( + sliced.search_index(900, &mut ctx).unwrap(), + SearchResult::NotFound(4) + ); } #[test] @@ -2201,6 +2286,7 @@ mod test { #[test] fn test_chunk_offsets_slice_single_patch() { + let mut ctx = array_session().create_execution_ctx(); let indices = buffer![100u64, 1200, 1300, 2500].into_array(); let values = buffer![10i32, 20, 30, 40].into_array(); let chunk_offsets = buffer![0u64, 1, 3].into_array(); @@ -2210,13 +2296,23 @@ mod test { assert_eq!(sliced.num_patches(), 1); assert_eq!(sliced.offset(), 1100); - assert_eq!(sliced.search_index(100).unwrap(), SearchResult::Found(0)); // 1200 - 1100 = 100 - assert_eq!(sliced.search_index(50).unwrap(), SearchResult::NotFound(0)); - assert_eq!(sliced.search_index(150).unwrap(), SearchResult::NotFound(1)); + assert_eq!( + sliced.search_index(100, &mut ctx).unwrap(), + SearchResult::Found(0) + ); // 1200 - 1100 = 100 + assert_eq!( + sliced.search_index(50, &mut ctx).unwrap(), + SearchResult::NotFound(0) + ); + assert_eq!( + sliced.search_index(150, &mut ctx).unwrap(), + SearchResult::NotFound(1) + ); } #[test] fn test_chunk_offsets_slice_across_chunks() { + let mut ctx = array_session().create_execution_ctx(); let indices = buffer![100u64, 200, 1100, 1200, 2100, 2200].into_array(); let values = buffer![10i32, 20, 30, 40, 50, 60].into_array(); let chunk_offsets = buffer![0u64, 2, 4].into_array(); @@ -2227,37 +2323,66 @@ mod test { assert_eq!(sliced.num_patches(), 4); assert_eq!(sliced.offset(), 150); - assert_eq!(sliced.search_index(50).unwrap(), SearchResult::Found(0)); // 200 - 150 = 50 - assert_eq!(sliced.search_index(950).unwrap(), SearchResult::Found(1)); // 1100 - 150 = 950 - assert_eq!(sliced.search_index(1050).unwrap(), SearchResult::Found(2)); // 1200 - 150 = 1050 - assert_eq!(sliced.search_index(1950).unwrap(), SearchResult::Found(3)); // 2100 - 150 = 1950 + assert_eq!( + sliced.search_index(50, &mut ctx).unwrap(), + SearchResult::Found(0) + ); // 200 - 150 = 50 + assert_eq!( + sliced.search_index(950, &mut ctx).unwrap(), + SearchResult::Found(1) + ); // 1100 - 150 = 950 + assert_eq!( + sliced.search_index(1050, &mut ctx).unwrap(), + SearchResult::Found(2) + ); // 1200 - 150 = 1050 + assert_eq!( + sliced.search_index(1950, &mut ctx).unwrap(), + SearchResult::Found(3) + ); // 2100 - 150 = 1950 } #[test] fn test_chunk_offsets_boundary_searches() { + let mut ctx = array_session().create_execution_ctx(); let indices = buffer![1023u64, 1024, 1025, 2047, 2048].into_array(); let values = buffer![10i32, 20, 30, 40, 50].into_array(); let chunk_offsets = buffer![0u64, 1, 4].into_array(); let patches = Patches::new(3000, 0, indices, values, Some(chunk_offsets)).unwrap(); - assert_eq!(patches.search_index(1023).unwrap(), SearchResult::Found(0)); - assert_eq!(patches.search_index(1024).unwrap(), SearchResult::Found(1)); - assert_eq!(patches.search_index(1025).unwrap(), SearchResult::Found(2)); - assert_eq!(patches.search_index(2047).unwrap(), SearchResult::Found(3)); - assert_eq!(patches.search_index(2048).unwrap(), SearchResult::Found(4)); + assert_eq!( + patches.search_index(1023, &mut ctx).unwrap(), + SearchResult::Found(0) + ); + assert_eq!( + patches.search_index(1024, &mut ctx).unwrap(), + SearchResult::Found(1) + ); + assert_eq!( + patches.search_index(1025, &mut ctx).unwrap(), + SearchResult::Found(2) + ); + assert_eq!( + patches.search_index(2047, &mut ctx).unwrap(), + SearchResult::Found(3) + ); + assert_eq!( + patches.search_index(2048, &mut ctx).unwrap(), + SearchResult::Found(4) + ); assert_eq!( - patches.search_index(1022).unwrap(), + patches.search_index(1022, &mut ctx).unwrap(), SearchResult::NotFound(0) ); assert_eq!( - patches.search_index(2046).unwrap(), + patches.search_index(2046, &mut ctx).unwrap(), SearchResult::NotFound(3) ); } #[test] fn test_chunk_offsets_slice_edge_cases() { + let mut ctx = array_session().create_execution_ctx(); let indices = buffer![0u64, 1, 1023, 1024, 2047, 2048].into_array(); let values = buffer![10i32, 20, 30, 40, 50, 60].into_array(); let chunk_offsets = buffer![0u64, 3, 5].into_array(); @@ -2266,18 +2391,31 @@ mod test { // Slice at the very beginning let sliced = patches.slice(0..10).unwrap().unwrap(); assert_eq!(sliced.num_patches(), 2); - assert_eq!(sliced.search_index(0).unwrap(), SearchResult::Found(0)); - assert_eq!(sliced.search_index(1).unwrap(), SearchResult::Found(1)); + assert_eq!( + sliced.search_index(0, &mut ctx).unwrap(), + SearchResult::Found(0) + ); + assert_eq!( + sliced.search_index(1, &mut ctx).unwrap(), + SearchResult::Found(1) + ); // Slice at the very end let sliced = patches.slice(2040..3000).unwrap().unwrap(); assert_eq!(sliced.num_patches(), 2); // patches at 2047 and 2048 - assert_eq!(sliced.search_index(7).unwrap(), SearchResult::Found(0)); // 2047 - 2040 - assert_eq!(sliced.search_index(8).unwrap(), SearchResult::Found(1)); // 2048 - 2040 + assert_eq!( + sliced.search_index(7, &mut ctx).unwrap(), + SearchResult::Found(0) + ); // 2047 - 2040 + assert_eq!( + sliced.search_index(8, &mut ctx).unwrap(), + SearchResult::Found(1) + ); // 2048 - 2040 } #[test] fn test_chunk_offsets_slice_nested() { + let mut ctx = array_session().create_execution_ctx(); let indices = buffer![100u64, 200, 300, 400, 500, 600].into_array(); let values = buffer![10i32, 20, 30, 40, 50, 60].into_array(); let chunk_offsets = buffer![0u64].into_array(); @@ -2290,9 +2428,12 @@ mod test { assert_eq!(sliced2.num_patches(), 1); // 300 assert_eq!(sliced2.offset(), 250); - assert_eq!(sliced2.search_index(50).unwrap(), SearchResult::Found(0)); // 300 - 250 assert_eq!( - sliced2.search_index(150).unwrap(), + sliced2.search_index(50, &mut ctx).unwrap(), + SearchResult::Found(0) + ); // 300 - 250 + assert_eq!( + sliced2.search_index(150, &mut ctx).unwrap(), SearchResult::NotFound(1) ); } @@ -2313,12 +2454,13 @@ mod test { #[test] fn test_index_larger_than_length() { + let mut ctx = array_session().create_execution_ctx(); let chunk_offsets = buffer![0u64].into_array(); let indices = buffer![1023u64].into_array(); let values = buffer![42i32].into_array(); let patches = Patches::new(1024, 0, indices, values, Some(chunk_offsets)).unwrap(); assert_eq!( - patches.search_index(2048).unwrap(), + patches.search_index(2048, &mut ctx).unwrap(), SearchResult::NotFound(1) ); } diff --git a/vortex-array/src/search_sorted/mod.rs b/vortex-array/src/search_sorted/mod.rs index cda27cfc511..e390fb3b7ae 100644 --- a/vortex-array/src/search_sorted/mod.rs +++ b/vortex-array/src/search_sorted/mod.rs @@ -2,6 +2,7 @@ // SPDX-FileCopyrightText: Copyright the Vortex contributors mod primitive; +mod scalar; use std::cmp::Ordering; use std::cmp::Ordering::Equal; @@ -13,13 +14,9 @@ use std::fmt::Formatter; use std::hint; pub use primitive::*; +pub use scalar::*; use vortex_error::VortexResult; -use crate::ArrayRef; -use crate::VortexSessionExecute; -use crate::legacy_session; -use crate::scalar::Scalar; - #[derive(Debug, Copy, Clone, Eq, PartialEq)] pub enum SearchSortedSide { Left, @@ -272,18 +269,6 @@ fn search_sorted_side_idx VortexResult>( } } -impl IndexOrd for ArrayRef { - #[allow(clippy::disallowed_methods)] - fn index_cmp(&self, idx: usize, elem: &Scalar) -> VortexResult> { - let scalar_a = self.execute_scalar(idx, &mut legacy_session().create_execution_ctx())?; - Ok(scalar_a.partial_cmp(elem)) - } - - fn index_len(&self) -> usize { - Self::len(self) - } -} - impl IndexOrd for [T] { #[inline] fn index_cmp(&self, idx: usize, elem: &T) -> VortexResult> { diff --git a/vortex-array/src/search_sorted/primitive.rs b/vortex-array/src/search_sorted/primitive.rs index 3da28c7be2b..dc0d4fc9c91 100644 --- a/vortex-array/src/search_sorted/primitive.rs +++ b/vortex-array/src/search_sorted/primitive.rs @@ -9,20 +9,36 @@ use vortex_error::VortexResult; use crate::ArrayRef; use crate::ExecutionCtx; +use crate::RepeatedArrayProbe; +use crate::arrays::Primitive; +use crate::arrays::primitive::PrimitiveArrayExt; use crate::dtype::NativePType; use crate::search_sorted::IndexOrd; /// A [`SearchSorted`](crate::search_sorted::SearchSorted) adapter over a sorted primitive-typed /// array, comparing elements as the native type `T`. /// +/// Reads go through a [`RepeatedArrayProbe`], so the encoding state and validity resolved on the +/// first comparison are reused by the rest of the search rather than rebuilt per probe. +/// /// Values can be searched as `T`, `Option`, or `usize`. Searching as `T` or `usize` treats /// null elements as `T::zero()`; use `Option` when the array may contain nulls, in which case /// nulls sort before all non-null values. -pub struct SearchSortedPrimitiveArray<'a, T>( - &'a ArrayRef, - RefCell<&'a mut ExecutionCtx>, - PhantomData, -); +pub struct SearchSortedPrimitiveArray<'a, T> { + reader: Reader<'a, T>, + len: usize, + ctx: RefCell<&'a mut ExecutionCtx>, + _ptype: PhantomData, +} + +/// Where a comparison reads its element from. +enum Reader<'a, T> { + /// The array's own buffer. Only for a canonical, non-nullable, host-backed array, where + /// every element is a `T` at a known offset and none of them is null. + Values(&'a [T]), + /// Anything else, through one probe reused by the whole search. + Probe(RefCell), +} impl<'a, T: NativePType> SearchSortedPrimitiveArray<'a, T> { /// Wraps `array` for searching, panicking if the array's [`PType`](crate::dtype::PType) is @@ -33,17 +49,47 @@ impl<'a, T: NativePType> SearchSortedPrimitiveArray<'a, T> { T::PTYPE, "Array PType must match primitive type" ); - Self(array, RefCell::new(ctx), PhantomData) + Self { + reader: Self::reader(array), + len: array.len(), + ctx: RefCell::new(ctx), + _ptype: PhantomData, + } + } + + /// Reads the buffer directly when the array can offer one, and probes otherwise. + /// + /// A nullable array is excluded because the buffer holds an arbitrary value where an element + /// is null, and a device-backed array because its buffer is not addressable here. + fn reader(array: &'a ArrayRef) -> Reader<'a, T> { + if !array.dtype().is_nullable() + && let Some(primitive) = array.as_opt::() + && primitive.buffer_handle().as_host_opt().is_some() + && let Some(values) = primitive.data().as_slice::().get(..array.len()) + { + return Reader::Values(values); + } + Reader::Probe(RefCell::new(array.repeated_probe())) + } + + /// Returns the value at `idx`, or `None` if the element is null. + /// + /// The probe and the context are separate cells, so the two borrows never overlap. + fn typed_value(&self, idx: usize) -> VortexResult> { + let probe = match &self.reader { + Reader::Values(values) => return Ok(Some(values[idx])), + Reader::Probe(probe) => probe, + }; + Ok(probe + .borrow_mut() + .execute_scalar(idx, &mut self.ctx.borrow_mut())? + .as_primitive() + .typed_value::()) } /// Returns the value at `idx`, with nulls mapped to `T::zero()`. fn value(&self, idx: usize) -> VortexResult { - Ok(self - .0 - .execute_scalar(idx, &mut self.1.borrow_mut())? - .as_primitive() - .typed_value::() - .unwrap_or_else(|| T::zero())) + Ok(self.typed_value(idx)?.unwrap_or_else(|| T::zero())) } } @@ -54,15 +100,15 @@ impl IndexOrd for SearchSortedPrimitiveArray<'_, T> { } fn index_len(&self) -> usize { - self.0.len() + self.len } } impl IndexOrd> for SearchSortedPrimitiveArray<'_, T> { fn index_cmp(&self, idx: usize, elem: &Option) -> VortexResult> { - // The borrow must end before `self.value` re-borrows the ctx. - let valid = self.0.is_valid(idx, &mut self.1.borrow_mut())?; - let value = valid.then(|| self.value(idx)).transpose()?; + // A null element reads back as a null scalar, so one read answers both whether the + // element is valid and what its value is. + let value = self.typed_value(idx)?; Ok(match (value, elem.as_ref()) { (Some(l), Some(r)) => Some(l.total_compare(*r)), @@ -73,7 +119,7 @@ impl IndexOrd> for SearchSortedPrimitiveArray<'_, T> { } fn index_len(&self) -> usize { - self.0.len() + self.len } } @@ -89,7 +135,7 @@ impl IndexOrd for SearchSortedPrimitiveArray<'_, T> { } fn index_len(&self) -> usize { - self.0.len() + self.len } } @@ -98,6 +144,7 @@ mod tests { use vortex_buffer::buffer; use vortex_error::VortexResult; + use super::Reader; use crate::IntoArray; use crate::array_session; use crate::arrays::PrimitiveArray; @@ -108,6 +155,34 @@ mod tests { use crate::search_sorted::SearchSortedSide; use crate::validity::Validity; + #[test] + fn search_sorted_reads_canonical_values_directly() -> VortexResult<()> { + let mut ctx = array_session().create_execution_ctx(); + let array = PrimitiveArray::new(buffer![1i32, 2, 3, 4], Validity::NonNullable).into_array(); + let searcher = SearchSortedPrimitiveArray::::new(&array, &mut ctx); + assert!(matches!(searcher.reader, Reader::Values(_))); + assert_eq!( + searcher.search_sorted(&3i32, SearchSortedSide::Left)?, + SearchResult::Found(2) + ); + Ok(()) + } + + #[test] + fn search_sorted_probes_when_values_are_not_addressable() -> VortexResult<()> { + let mut ctx = array_session().create_execution_ctx(); + // Nullable: the buffer holds an arbitrary value wherever an element is null, so the + // search has to consult validity through the probe. + let array = PrimitiveArray::new(buffer![1i32, 2, 3, 4], Validity::AllValid).into_array(); + let searcher = SearchSortedPrimitiveArray::::new(&array, &mut ctx); + assert!(matches!(searcher.reader, Reader::Probe(_))); + assert_eq!( + searcher.search_sorted(&3i32, SearchSortedSide::Left)?, + SearchResult::Found(2) + ); + Ok(()) + } + #[test] fn search_sorted_optional_value() -> VortexResult<()> { let array = PrimitiveArray::new(buffer![1i32, 2, 3, 4], Validity::AllValid).into_array(); diff --git a/vortex-array/src/search_sorted/scalar.rs b/vortex-array/src/search_sorted/scalar.rs new file mode 100644 index 00000000000..907bcba16e5 --- /dev/null +++ b/vortex-array/src/search_sorted/scalar.rs @@ -0,0 +1,105 @@ +// SPDX-License-Identifier: Apache-2.0 +// SPDX-FileCopyrightText: Copyright the Vortex contributors + +use std::cell::RefCell; +use std::cmp::Ordering; + +use vortex_error::VortexResult; + +use crate::ArrayRef; +use crate::ExecutionCtx; +use crate::RepeatedArrayProbe; +use crate::scalar::Scalar; +use crate::search_sorted::IndexOrd; + +/// A [`SearchSorted`](crate::search_sorted::SearchSorted) adapter over a sorted array of any +/// encoding, comparing elements as [`Scalar`]s. +/// +/// Reads go through a [`RepeatedArrayProbe`], so the whole search shares one execution context +/// and whatever the first comparison resolved. +/// +/// Prefer [`SearchSortedPrimitiveArray`](crate::search_sorted::SearchSortedPrimitiveArray) when +/// the element type is known: comparing through `Scalar` cannot read the values directly and +/// builds a scalar per probe. +pub struct SearchSortedArray<'a> { + probe: RefCell, + len: usize, + ctx: RefCell<&'a mut ExecutionCtx>, +} + +impl<'a> SearchSortedArray<'a> { + /// Wraps `array` for searching. + pub fn new(array: &ArrayRef, ctx: &'a mut ExecutionCtx) -> Self { + Self { + probe: RefCell::new(array.repeated_probe()), + len: array.len(), + ctx: RefCell::new(ctx), + } + } +} + +impl IndexOrd for SearchSortedArray<'_> { + /// The probe and the context are separate cells, so the two borrows never overlap. + fn index_cmp(&self, idx: usize, elem: &Scalar) -> VortexResult> { + let scalar = self + .probe + .borrow_mut() + .execute_scalar(idx, &mut self.ctx.borrow_mut())?; + Ok(scalar.partial_cmp(elem)) + } + + fn index_len(&self) -> usize { + self.len + } +} + +#[cfg(test)] +mod tests { + use vortex_buffer::buffer; + use vortex_error::VortexResult; + + use crate::IntoArray; + use crate::array_session; + use crate::arrays::PrimitiveArray; + use crate::executor::VortexSessionExecute; + use crate::scalar::Scalar; + use crate::search_sorted::SearchResult; + use crate::search_sorted::SearchSorted; + use crate::search_sorted::SearchSortedArray; + use crate::search_sorted::SearchSortedSide; + use crate::validity::Validity; + + #[test] + fn search_sorted_scalar() -> VortexResult<()> { + let array = + PrimitiveArray::new(buffer![1i32, 2, 3, 3, 4], Validity::NonNullable).into_array(); + let mut ctx = array_session().create_execution_ctx(); + let searcher = SearchSortedArray::new(&array, &mut ctx); + assert_eq!( + searcher.search_sorted(&Scalar::from(3i32), SearchSortedSide::Left)?, + SearchResult::Found(2) + ); + assert_eq!( + searcher.search_sorted(&Scalar::from(3i32), SearchSortedSide::Right)?, + SearchResult::Found(4) + ); + assert_eq!( + searcher.search_sorted(&Scalar::from(5i32), SearchSortedSide::Left)?, + SearchResult::NotFound(5) + ); + Ok(()) + } + + #[test] + fn search_sorted_scalar_with_nulls() -> VortexResult<()> { + let array = + PrimitiveArray::from_option_iter([None, None, Some(2i32), Some(3)]).into_array(); + let mut ctx = array_session().create_execution_ctx(); + let searcher = SearchSortedArray::new(&array, &mut ctx); + assert_eq!( + searcher.search_sorted(&Scalar::from(Some(2i32)), SearchSortedSide::Left)?, + SearchResult::Found(2) + ); + Ok(()) + } +}