use super::*;
#[inline(always)]
pub fn morton_cascade<R: UnsignedIntegerRegister, const N: usize>(values: [Storage<R>; N]) -> Storage<R> {
if const { N == 1 } {
return values[0];
}
let width = (size_of::<R::Element>() * 8) as u32;
let bits = morton_bits_per_lane(width, N); let passes = morton_passes(bits); let low = R::shr(R::not(R::EMPTY), width - bits);
let mut code = R::EMPTY;
let mut d = 0;
while d < N {
let mut x = R::bitand(values[d], low);
let mut k = passes;
while k > 0 {
k -= 1;
let shift = (N as u32 - 1) << k; let mask = morton_block_mask::<R>(k, N);
x = R::bitand(R::bitor(x, R::shl(x, shift)), mask);
}
code = R::bitor(code, R::shl(x, d as u32));
d += 1;
}
code
}
#[inline(always)]
pub fn morton_pack2<R: UnsignedIntegerRegister, const N: usize>(x: Storage<R>, y: Storage<R>) -> [Storage<R>; N] {
let cols = [x, y];
let mut out = [R::EMPTY; N];
let mut d = 0;
while d < 2 && d < N {
out[d] = cols[d];
d += 1;
}
out
}
#[inline(always)]
pub fn reverse_morton_cascade<R: UnsignedIntegerRegister, const N: usize>(code: Storage<R>) -> [Storage<R>; N] {
if const { N == 1 } {
return [code; N];
}
let width = (size_of::<R::Element>() * 8) as u32;
let bits = morton_bits_per_lane(width, N);
let passes = morton_passes(bits);
let low = R::shr(R::not(R::EMPTY), width - bits);
let stride = morton_block_mask::<R>(0, N);
let mut out = [R::EMPTY; N];
let mut d = 0;
while d < N {
let mut x = R::bitand(R::shr(code, d as u32), stride);
let mut k = 0;
while k < passes {
let shift = (N as u32 - 1) << k; let mask = if k + 1 >= passes {
low
} else {
morton_block_mask::<R>(k + 1, N)
};
x = R::bitand(R::bitor(x, R::shr(x, shift)), mask);
k += 1;
}
out[d] = x;
d += 1;
}
out
}
pub const fn morton_bits_per_lane(width: u32, dims: usize) -> u32 {
let bits = width / dims as u32;
if bits == 0 { 1 } else { bits }
}
pub const fn morton_passes(bits: u32) -> u32 {
if bits <= 1 {
0
} else {
u32::BITS - (bits - 1).leading_zeros()
}
}
#[inline(always)]
pub fn morton_block_mask<R: UnsignedIntegerRegister>(block_log2: u32, dims: usize) -> Storage<R> {
let width = (size_of::<R::Element>() * 8) as u32;
let block = 1u32 << block_log2; let period = dims as u32 * block;
let mut mask = R::shr(R::not(R::EMPTY), width - block);
let mut s = period;
while s < width {
mask = R::bitor(mask, R::shl(mask, s));
s <<= 1;
}
mask
}