Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
97 changes: 78 additions & 19 deletions vortex-buffer/src/bit/take.rs
Original file line number Diff line number Diff line change
Expand Up @@ -79,9 +79,14 @@ where
*bools.get_unchecked(idx)
})
} else {
#[cfg(target_arch = "x86_64")]
if let Some(taken) = avx2::take(bits, indices, None) {
return taken;
}

let ptr = bits.inner().as_ptr();
let offset = bits.offset();
BitBuffer::collect_bool(indices.len(), |i| unsafe {
BitBuffer::collect_bool_multiversioned(indices.len(), |i| unsafe {
// SAFETY: we're already iterating over every index, a bound
// is excessive
let idx = indices.get_unchecked(i).as_();
Expand All @@ -91,6 +96,11 @@ where
};
};

#[cfg(target_arch = "x86_64")]
if let Some(taken) = avx2::take(bits, indices, Some(validity)) {
return taken;
}

let ptr = validity.inner().as_ptr();
if bits.len() <= COLLECT_TO_VEC {
let offset = validity.offset();
Expand Down Expand Up @@ -132,8 +142,14 @@ fn first_unset(bits: BitBufferView<'_>) -> Option<usize> {
(target < bits.len()).then_some(target)
}

#[cfg(target_arch = "x86_64")]
#[path = "take_avx2.rs"]
mod avx2;

#[cfg(test)]
mod tests {
use num_traits::AsPrimitive;
use num_traits::Zero;
use rand::RngExt;
use rand::SeedableRng;
use rand::rngs::StdRng;
Expand All @@ -152,34 +168,32 @@ mod tests {
(0..len).map(|_| rng.random_range(0..bound)).collect()
}

#[rstest]
#[case(300)]
#[case(5000)]
fn gathers(#[case] bits_len: usize) {
fn check_gathers<I>(bits_len: usize, garbage: [I; 2])
where
I: AsPrimitive<usize> + TryFrom<usize> + Ord + Zero,
<I as TryFrom<usize>>::Error: std::fmt::Debug,
{
let bits = random_bits(bits_len + 11);
let view = bits.as_view().slice(11..);
let len = 400;
let indices: Vec<i64> = random_indices(len, bits_len)
let len = 403;
let indices: Vec<I> = random_indices(len, bits_len)
.into_iter()
.map(|idx| i64::try_from(idx).unwrap())
.map(|idx| I::try_from(idx).unwrap())
.collect();

let taken = take_bits(view, &indices, None);
assert_eq!(taken.len(), len);
for (i, &idx) in indices.iter().enumerate() {
assert_eq!(
taken.value(i),
bits.value(usize::try_from(idx).unwrap() + 11)
);
assert_eq!(taken.value(i), bits.value(idx.as_() + 11));
}

let validity = BitBuffer::collect_bool(len, |i| i % 3 != 0);
let with_garbage: Vec<i64> = indices
let with_garbage: Vec<I> = indices
.iter()
.enumerate()
.map(|(i, &idx)| match i % 6 {
0 => i64::MAX,
3 => -1,
0 => garbage[0],
3 => garbage[1],
_ => idx,
})
.collect();
Expand All @@ -188,14 +202,32 @@ mod tests {
assert_eq!(taken.len(), len);
for (i, &idx) in with_garbage.iter().enumerate() {
if validity.value(i) {
assert_eq!(
taken.value(i),
bits.value(usize::try_from(idx).unwrap() + 11)
);
assert_eq!(taken.value(i), bits.value(idx.as_() + 11));
}
}
}

#[rstest]
#[case(300)]
#[case(5000)]
fn gathers_i64(#[case] bits_len: usize) {
check_gathers::<i64>(bits_len, [i64::MAX, -1]);
}

#[rstest]
#[case(300)]
#[case(5000)]
fn gathers_u16(#[case] bits_len: usize) {
check_gathers::<u16>(bits_len, [u16::MAX, u16::MAX - 1]);
}

#[rstest]
#[case(300)]
#[case(5000)]
fn gathers_u32(#[case] bits_len: usize) {
check_gathers::<u32>(bits_len, [u32::MAX, u32::MAX - 1]);
}

#[rstest]
#[case(BitBuffer::new_set(100), |_: usize| true)]
#[case(BitBuffer::new_unset(100), |_: usize| false)]
Expand Down Expand Up @@ -232,4 +264,31 @@ mod tests {
let taken = take_bits(bits.as_view(), &[7u32, 8, 9], Some(validity.as_view()));
assert_eq!(taken, BitBuffer::new_unset(3));
}

#[test]
#[should_panic(expected = "out of bounds")]
fn valid_index_out_of_bounds() {
let bits = BitBuffer::collect_bool(10_000, |i| i % 2 == 0);
let mut indices: Vec<u32> = random_indices(403, 10_000)
.into_iter()
.map(|idx| u32::try_from(idx).unwrap())
.collect();
indices[80] = 1_000_000;
let validity = BitBuffer::collect_bool(403, |i| i != 200);
let _taken = take_bits(bits.as_view(), &indices, Some(validity.as_view()));
}

#[test]
#[should_panic(expected = "out of bounds")]
fn index_out_of_bounds() {
let bits = BitBuffer::collect_bool(10_000, |i| i % 2 == 0);
let view = bits.as_view().slice(3..);
let mut indices: Vec<u32> = random_indices(403, 9_000)
.into_iter()
.map(|idx| u32::try_from(idx).unwrap())
.collect();
indices[80] = u32::MAX;
let validity = BitBuffer::collect_bool(403, |i| i != 200);
let _taken = take_bits(view, &indices, Some(validity.as_view()));
}
}
240 changes: 240 additions & 0 deletions vortex-buffer/src/bit/take_avx2.rs
Original file line number Diff line number Diff line change
@@ -0,0 +1,240 @@
// SPDX-License-Identifier: Apache-2.0
// SPDX-FileCopyrightText: Copyright the Vortex contributors

use std::arch::x86_64::__m128i;
use std::arch::x86_64::__m256i;
use std::arch::x86_64::_mm_loadu_si128;
use std::arch::x86_64::_mm256_add_epi32;
use std::arch::x86_64::_mm256_and_si256;
use std::arch::x86_64::_mm256_andnot_si256;
use std::arch::x86_64::_mm256_castsi256_ps;
use std::arch::x86_64::_mm256_cmpeq_epi32;
use std::arch::x86_64::_mm256_cvtepu16_epi32;
use std::arch::x86_64::_mm256_i32gather_epi32;
use std::arch::x86_64::_mm256_loadu_si256;
use std::arch::x86_64::_mm256_mask_i32gather_epi32;
use std::arch::x86_64::_mm256_min_epu32;
use std::arch::x86_64::_mm256_movemask_ps;
use std::arch::x86_64::_mm256_or_si256;
use std::arch::x86_64::_mm256_set1_epi32;
use std::arch::x86_64::_mm256_setr_epi32;
use std::arch::x86_64::_mm256_setzero_si256;
use std::arch::x86_64::_mm256_slli_epi32;
use std::arch::x86_64::_mm256_srli_epi32;
use std::arch::x86_64::_mm256_srlv_epi32;
use std::arch::x86_64::_mm256_testz_si256;
use std::slice::from_raw_parts;

use num_traits::AsPrimitive;

use crate::BitBuffer;
use crate::BitBufferView;
use crate::BufferMut;
use crate::bit::get_bit;
use crate::bit::get_bit_unchecked;

/// Caller must verify indices.len() == validity.len()
pub(super) fn take<I: AsPrimitive<usize>>(
bits: BitBufferView<'_>,
indices: &[I],
validity: Option<BitBufferView<'_>>,
) -> Option<BitBuffer> {
if !is_x86_feature_detected!("avx2") {
return None;
}
if bits.len() >= (1 << 31) - 8 {
return None;
}
let last_dword = (bits.offset() + bits.len() - 1) / 32;
if (last_dword + 1) * 4 > bits.inner().len() {
return None;
}

let size = size_of::<I>();
if size == size_of::<u16>() {
// SAFETY: we clamp the index to be valid in take_lanes and check out of
// bound reads there as well, so any invalid index doesn't cause an out
// of bounds read and is reported via panic().
let indices = unsafe { from_raw_parts(indices.as_ptr().cast::<u16>(), indices.len()) };
let indices_ptr: *const u16 = indices.as_ptr();
unsafe {
let group = |group_idx| -> __m256i {
let group_ptr: *const __m128i = indices_ptr.add(group_idx * 8).cast();
// copy 128 bits into __m128i
let vector: __m128i = _mm_loadu_si128(group_ptr);
// zero-extend unsigned 16-bit integers in __m128i to signed
// 32-bit integers in __m256i
_mm256_cvtepu16_epi32(vector)
};
return Some(take_group(bits, indices, validity, group));
}
}
if size == size_of::<u32>() {
// SAFETY: see u16 case above
let indices = unsafe { from_raw_parts(indices.as_ptr().cast::<u32>(), indices.len()) };
let indices_ptr: *const u32 = indices.as_ptr();
unsafe {
let group = |group_idx| -> __m256i {
let group_ptr: *const __m256i = indices_ptr.add(group_idx * 8).cast();
// copy 256 bits into __m256i
_mm256_loadu_si256(group_ptr)
};
return Some(take_group(bits, indices, validity, group));
}
}
None
}

/// By the time we call this function we already handle cases of u16 and u32
/// separately as there's different packing code. Now we also need to branch
/// on validity. take(), the first function, is called from two places, one
/// with validity and one without. Consequently, there are different gather
/// instructions depending on whether validity is present.
#[target_feature(enable = "avx2")]
unsafe fn take_group<I: AsPrimitive<usize>>(
bits: BitBufferView<'_>,
indices: &[I],
validity: Option<BitBufferView<'_>>,
load_group: impl Fn(usize) -> __m256i,
) -> BitBuffer {
let inner = bits.inner();
let base = inner.as_ptr();
let zero: __m256i = _mm256_setzero_si256();

let Some(validity) = validity else {
// If validity isn't present, caller has verified "indices" are valid
// offsets info "bits" (see min/max code in take.rs), so we can use
// unchecked access into "bits".
let gather = |_group: usize, dword: __m256i, _in_bounds: __m256i| -> (__m256i, __m256i) {
let gathered = unsafe { _mm256_i32gather_epi32::<4>(base.cast(), dword) };
(gathered, zero)
};
return unsafe { take_lanes(bits, indices, load_group, gather, None) };
};

// If validity is present, we can't trust "indices" are valid offsets.
// There also may be garbage under a NULL index, and we need to avoid
// reading bits's offset at this garbage.
let valid_bytes = validity.inner();
let validity_offset = validity.offset();
let validity_shift = validity_offset % 8;
let validity_first_byte = validity_offset / 8;
let has_non_byte_shift: usize = (validity_shift != 0).as_();

let lane_bits: __m256i = _mm256_setr_epi32(1, 2, 4, 8, 16, 32, 64, 128);

let gather = |group: usize, dword: __m256i, in_range: __m256i| -> (__m256i, __m256i) {
let byte = validity_first_byte + group;
let next_byte = byte + has_non_byte_shift;

// We've verified in take.rs validity read is in bounds
let lo = unsafe { *valid_bytes.get_unchecked(byte) as u16 };
let hi = unsafe { *valid_bytes.get_unchecked(next_byte) as u16 };

let window = lo | (hi << 8);
let valid = ((window >> validity_shift) & 0xFF) as u8;

// all 1 if element is valid, all 0 otherwise
let valid_vector = _mm256_set1_epi32(i32::from(valid));
// valid_vector & lane_bits
let selected = _mm256_and_si256(valid_vector, lane_bits);
// selected[i] == lane_bits[i]
let lanes = _mm256_cmpeq_epi32(selected, lane_bits);

// !in_range & lanes. We have a valid index which is out of bounds
let violation = _mm256_andnot_si256(in_range, lanes);

let gathered = unsafe { _mm256_mask_i32gather_epi32::<4>(zero, base.cast(), dword, lanes) };
(gathered, violation)
};
unsafe {
take_lanes(
bits,
indices,
load_group,
gather,
Some((valid_bytes, validity_offset)),
)
}
}

#[expect(clippy::cast_possible_truncation)]
#[target_feature(enable = "avx2")]
unsafe fn take_lanes<I>(
bits: BitBufferView<'_>,
indices: &[I],
load_group: impl Fn(usize) -> __m256i,
gather: impl Fn(usize, __m256i, __m256i) -> (__m256i, __m256i),
validity: Option<(&[u8], usize)>,
) -> BitBuffer
where
I: AsPrimitive<usize>,
{
let total = indices.len();
let full_groups = total / 8;

let bit_offset: __m256i = _mm256_set1_epi32(bits.offset() as i32);
let low_bits: __m256i = _mm256_set1_epi32(31);
let max_index: __m256i = _mm256_set1_epi32((bits.len() - 1) as i32);
let mut out_of_bounds: __m256i = _mm256_setzero_si256();

let mut out = BufferMut::<u8>::with_capacity(total.div_ceil(8));
for group in 0..full_groups {
let group_vector: __m256i = load_group(group);

// group_vector[i] = min(group_vector[i], max_index[i])
// We clamp every index so it can never read over bits's buffer, and
// calculate violations separately. If we found a violation after
// looping< we panic.
let clamped = _mm256_min_epu32(group_vector, max_index);

// clamped == group_vector
let in_range = _mm256_cmpeq_epi32(clamped, group_vector);

let bitpos = _mm256_add_epi32(clamped, bit_offset);
let dword = _mm256_srli_epi32::<5>(bitpos);

let (gathered, violation) = gather(group, dword, in_range);
out_of_bounds = _mm256_or_si256(out_of_bounds, violation);

let shift = _mm256_and_si256(bitpos, low_bits);
let shifted = _mm256_srlv_epi32(gathered, shift);
let top = _mm256_slli_epi32::<31>(shifted);
let as_ps = _mm256_castsi256_ps(top);

let group_bits = _mm256_movemask_ps(as_ps);
let group_bits: u8 = (group_bits & 0xFF).as_();

// SAFETY: out has sufficient capacity
unsafe { out.push_unchecked(group_bits) };
}

assert!(
_mm256_testz_si256(out_of_bounds, out_of_bounds) != 0,
"index out of bounds"
);

let tail = total % 8;
if tail == 0 {
return BitBuffer::new(out.freeze(), total);
}

let start = full_groups * 8;
let inner = bits.inner();
let mut byte = 0u8;
for bit in 0..tail {
let pos = start + bit;
let keep = validity.map_or(usize::MAX, |(valid_bytes, validity_offset)| {
// SAFETY: validity has same length as indices
(unsafe { get_bit_unchecked(valid_bytes.as_ptr(), validity_offset + pos) } as usize)
.wrapping_neg()
});
// SAFETY: pos stays within indices
let idx = unsafe { indices.get_unchecked(pos) }.as_() & keep;
let value = get_bit(inner, bits.offset() + idx);
byte |= (value as u8) << bit;
}
// SAFETY: out has sufficient capacity
unsafe { out.push_unchecked(byte) }
BitBuffer::new(out.freeze(), total)
}
Loading