Skip to content
Closed
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
226 changes: 206 additions & 20 deletions datafusion/common/src/hash_utils.rs
Original file line number Diff line number Diff line change
Expand Up @@ -301,9 +301,10 @@ impl<S: HashState> ChildHashing for HashStateChildHashing<'_, S> {
}
}

/// Builds hash values of PrimitiveArray and writes them into `hashes_buffer`
/// If `rehash==true` this folds the existing hash into the hasher state
/// and hashes only the new value (avoiding a separate combine step).
/// Builds hash values of PrimitiveArray and writes them into `hashes_buffer`.
///
/// If `rehash==true` the previous hash is folded into the hasher state as a
/// per-element seed, avoiding a separate combine step.
#[cfg(not(feature = "force_hash_collisions"))]
fn hash_array_primitive<T>(
array: &PrimitiveArray<T>,
Expand All @@ -319,29 +320,214 @@ fn hash_array_primitive<T>(
"hashes_buffer and array should be of equal length"
);

let values = array.values();

if array.null_count() == 0 {
if rehash {
for (hash, &value) in hashes_buffer.iter_mut().zip(array.values().iter()) {
let mut hasher = random_state.seeded_state(*hash).build_hasher();
value.hash_write(&mut hasher);
*hash = hasher.finish();
}
hash_prim_rehash(values, hashes_buffer, random_state);
} else {
for (hash, &value) in hashes_buffer.iter_mut().zip(array.values().iter()) {
*hash = value.hash_one(random_state);
}
hash_prim_fresh(values, hashes_buffer, random_state);
}
} else if rehash {
for i in array.nulls().unwrap().valid_indices() {
let value = unsafe { array.value_unchecked(i) };
let mut hasher = random_state.seeded_state(hashes_buffer[i]).build_hasher();
value.hash_write(&mut hasher);
hashes_buffer[i] = hasher.finish();
}
hash_prim_with_nulls::<_, _, true>(
values,
hashes_buffer,
array.nulls().unwrap(),
random_state,
);
} else {
for i in array.nulls().unwrap().valid_indices() {
let value = unsafe { array.value_unchecked(i) };
hashes_buffer[i] = value.hash_one(random_state);
hash_prim_with_nulls::<_, _, false>(
values,
hashes_buffer,
array.nulls().unwrap(),
random_state,
);
}
}

#[cfg(not(feature = "force_hash_collisions"))]
const FRESH_LANES: usize = 16;
#[cfg(not(feature = "force_hash_collisions"))]
const REHASH_LANES: usize = 8;

#[cfg(not(feature = "force_hash_collisions"))]
#[inline(always)]
fn hash_prim_fresh<N, S>(values: &[N], hashes: &mut [u64], state: &S)
where
N: HashValue + Copy,
S: HashState,
{
debug_assert_eq!(values.len(), hashes.len());
let len = values.len();
let aligned_end = len & !(FRESH_LANES - 1);
let mut idx = 0;
while idx < aligned_end {
// SAFETY: idx + FRESH_LANES <= aligned_end <= values.len() == hashes.len()
unsafe {
let vals: [N; FRESH_LANES] =
std::array::from_fn(|lane| *values.get_unchecked(idx + lane));
let new_hashes: [u64; FRESH_LANES] =
std::array::from_fn(|lane| vals[lane].hash_one(state));
for lane in 0..FRESH_LANES {
*hashes.get_unchecked_mut(idx + lane) = new_hashes[lane];
}
}
idx += FRESH_LANES;
}
while idx < len {
unsafe {
*hashes.get_unchecked_mut(idx) = values.get_unchecked(idx).hash_one(state);
}
idx += 1;
}
}

// Narrower unroll than `hash_prim_fresh`: each lane keeps both `prev` and
// `value` live, so widening past 8 spills registers in the multi-column case.
#[cfg(not(feature = "force_hash_collisions"))]
#[inline(always)]
fn hash_prim_rehash<N, S>(values: &[N], hashes: &mut [u64], state: &S)
where
N: HashValue + Copy,
S: HashState,
{
debug_assert_eq!(values.len(), hashes.len());
let len = values.len();
let aligned_end = len & !(REHASH_LANES - 1);
let mut idx = 0;
while idx < aligned_end {
// SAFETY: idx + REHASH_LANES <= aligned_end <= values.len() == hashes.len()
unsafe {
let vals: [N; REHASH_LANES] =
std::array::from_fn(|lane| *values.get_unchecked(idx + lane));
let prevs: [u64; REHASH_LANES] =
std::array::from_fn(|lane| *hashes.get_unchecked(idx + lane));
let new_hashes: [u64; REHASH_LANES] =
std::array::from_fn(|lane| rehash_one(state, prevs[lane], vals[lane]));
for lane in 0..REHASH_LANES {
*hashes.get_unchecked_mut(idx + lane) = new_hashes[lane];
}
}
idx += REHASH_LANES;
}
while idx < len {
unsafe {
let value = *values.get_unchecked(idx);
let prev = *hashes.get_unchecked(idx);
*hashes.get_unchecked_mut(idx) = rehash_one(state, prev, value);
}
idx += 1;
}
}

#[cfg(not(feature = "force_hash_collisions"))]
#[inline(always)]
fn rehash_one<N: HashValue, S: HashState>(state: &S, prev: u64, value: N) -> u64 {
let mut hasher = state.seeded_state(prev).build_hasher();
value.hash_write(&mut hasher);
hasher.finish()
}

// Walks the null bitmap in u64 strides so bitmap work is amortized across 64
// rows instead of paid per valid row (as `valid_indices()` does).
#[cfg(not(feature = "force_hash_collisions"))]
#[inline(always)]
fn hash_prim_with_nulls<N, S, const REHASH: bool>(
values: &[N],
hashes: &mut [u64],
nulls: &arrow::buffer::NullBuffer,
state: &S,
) where
N: HashValue + Copy,
S: HashState,
{
debug_assert_eq!(values.len(), hashes.len());
debug_assert_eq!(values.len(), nulls.len());

let bit_chunks = nulls.inner().bit_chunks();
let mut row = 0;

for mask in bit_chunks.iter() {
if mask == u64::MAX {
// SAFETY: row + 64 <= nulls.len() == values.len() == hashes.len()
unsafe {
let vals = values.get_unchecked(row..row + 64);
let out = hashes.get_unchecked_mut(row..row + 64);
if REHASH {
hash_prim_rehash(vals, out, state);
} else {
hash_prim_fresh(vals, out, state);
}
}
} else if mask != 0 {
// SAFETY: same bound as above.
unsafe {
let vals = values.get_unchecked(row..row + 64);
let out = hashes.get_unchecked_mut(row..row + 64);
blend_chunk::<N, S, REHASH>(vals, out, mask, state);
}
}
row += 64;
}

let tail_len = bit_chunks.remainder_len();
if tail_len > 0 {
let tail_mask = (1u64 << tail_len) - 1;
let mut remaining = bit_chunks.remainder_bits() & tail_mask;
while remaining != 0 {
let idx = row + remaining.trailing_zeros() as usize;
unsafe {
let value = *values.get_unchecked(idx);
let prev = *hashes.get_unchecked(idx);
*hashes.get_unchecked_mut(idx) = if REHASH {
rehash_one(state, prev, value)
} else {
value.hash_one(state)
};
}
remaining &= remaining - 1;
}
}
}

#[cfg(not(feature = "force_hash_collisions"))]
const BLEND_LANES: usize = 8;

// Hashes every slot in a 64-row chunk and csel-blends against the mask bit.
// Cheaper than a `trailing_zeros` scan at the null densities we see in practice
// because there's no per-bit bitmap extract — the mask is already in a register.
#[cfg(not(feature = "force_hash_collisions"))]
#[inline(always)]
fn blend_chunk<N, S, const REHASH: bool>(
values: &[N],
hashes: &mut [u64],
mask: u64,
state: &S,
) where
N: HashValue + Copy,
S: HashState,
{
debug_assert_eq!(values.len(), 64);
debug_assert_eq!(hashes.len(), 64);
for offset in (0..64).step_by(BLEND_LANES) {
// SAFETY: offset + BLEND_LANES <= 64 == values.len() == hashes.len()
unsafe {
let vals: [N; BLEND_LANES] =
std::array::from_fn(|lane| *values.get_unchecked(offset + lane));
let prevs: [u64; BLEND_LANES] =
std::array::from_fn(|lane| *hashes.get_unchecked(offset + lane));
let new_hashes: [u64; BLEND_LANES] = std::array::from_fn(|lane| {
if REHASH {
rehash_one(state, prevs[lane], vals[lane])
} else {
vals[lane].hash_one(state)
}
});
for lane in 0..BLEND_LANES {
let valid = (mask >> (offset + lane)) & 1 != 0;
*hashes.get_unchecked_mut(offset + lane) =
if valid { new_hashes[lane] } else { prevs[lane] };
}
}
}
}
Expand Down
Loading