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/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 665356d6025..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,14 +471,12 @@ impl Patches { /// /// [`SearchResult::Found(patch_idx)`]: SearchResult::Found /// [`SearchResult::NotFound(insertion_point)`]: SearchResult::NotFound - #[allow(clippy::disallowed_methods)] - 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() { - let mut ctx = legacy_session().create_execution_ctx(); - return self.search_index_chunked(index, &mut ctx); + 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. @@ -540,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), @@ -746,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); @@ -771,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"))?; @@ -1014,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 { @@ -1030,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) }) } @@ -1921,6 +1925,7 @@ mod test { #[test] fn test_search_index() { + let mut ctx = array_session().create_execution_ctx(); let patches = Patches::new( 10, 0, @@ -1931,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] @@ -2136,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(); @@ -2144,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(); @@ -2179,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(); @@ -2198,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] @@ -2218,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(); @@ -2227,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(); @@ -2244,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(); @@ -2283,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(); @@ -2307,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) ); } @@ -2330,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) ); }