Skip to content
Open
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
1 change: 1 addition & 0 deletions Cargo.toml
Original file line number Diff line number Diff line change
Expand Up @@ -129,6 +129,7 @@ name = "graft_masked_branches"
harness = false

[[bench]]
name = "byte_mask"
name = "zipper_head_owned"
harness = false

Expand Down
81 changes: 81 additions & 0 deletions benches/byte_mask.rs
Original file line number Diff line number Diff line change
@@ -0,0 +1,81 @@
use divan::{black_box, Bencher, Divan};

use pathmap::utils::ByteMask;

fn main() {
let divan = Divan::from_args().sample_count(4000);

divan.main();
}

fn spread_mask(on_bits: usize) -> ByteMask {
debug_assert!(on_bits <= 256);

(0..on_bits)
.map(|idx| ((idx * 73 + 19) & 0xFF) as u8)
.collect()
}

fn iter_mask(mask: ByteMask) -> u64 {
let mut acc = 0u64;
let mut count = 0u64;
for byte in black_box(mask).iter() {
let byte = black_box(byte) as u64;
acc = acc.wrapping_mul(257).wrapping_add(byte + count);
count += 1;
}
black_box(acc ^ count)
}

#[divan::bench(args = [0, 1, 2, 4, 8, 16, 150, 200, 256])]
fn bytemask_iter(bencher: Bencher, on_bits: usize) {
let mask = spread_mask(on_bits);
let mut sink = 0u64;

bencher.bench_local(|| {
sink = sink.wrapping_add(iter_mask(mask));
black_box(sink);
});
}

fn recursive_masks(depth: usize, on_bits: usize) -> Vec<ByteMask> {
debug_assert!(on_bits <= 2);

(0..depth)
.map(|level| {
(0..on_bits)
.map(|idx| ((level * 37 + idx * 131 + 11) & 0xFF) as u8)
.collect()
})
.collect()
}

fn recursive_iter(masks: &[ByteMask], level: usize, acc: u64) -> u64 {
if level == masks.len() {
return black_box(acc);
}

let mut out = acc;
let mut iter = black_box(masks[level]).iter();
if let Some(byte) = iter.next() {
out = out.wrapping_add(recursive_iter(masks, level + 1, acc.wrapping_mul(257).wrapping_add(black_box(byte) as u64)));
}
if let Some(byte) = iter.next() {
out = out.wrapping_add(black_box(byte) as u64);
}
black_box(&mut iter);
black_box(out)
}

#[divan::bench(args = [1, 2])]
fn bytemask_iter_recursive_stack(bencher: Bencher, on_bits: usize) {
const DEPTH: usize = 50;

let masks = recursive_masks(DEPTH, on_bits);
let mut sink = 0u64;

bencher.bench_local(|| {
sink = sink.wrapping_add(recursive_iter(&masks, 0, 0));
black_box(sink);
});
}
6 changes: 3 additions & 3 deletions src/bridge_node.rs
Original file line number Diff line number Diff line change
Expand Up @@ -398,11 +398,11 @@ impl<V: Clone + Send + Sync> TrieNode<V> for BridgeNode<V> {
self.is_empty()
}
#[inline(always)]
fn new_iter_token(&self) -> u128 {
fn new_iter_token(&self) -> IterToken {
0
}
#[inline(always)]
fn iter_token_for_path(&self, key: &[u8]) -> (u128, &[u8]) {
fn iter_token_for_path(&self, key: &[u8]) -> (IterToken, &[u8]) {
let node_key = self.key();
if key.len() <= node_key.len() {
let short_key = &node_key[..key.len()];
Expand All @@ -416,7 +416,7 @@ impl<V: Clone + Send + Sync> TrieNode<V> for BridgeNode<V> {
(NODE_ITER_FINISHED, &[])
}
#[inline(always)]
fn next_items(&self, token: u128) -> (u128, &[u8], Option<&TrieNodeODRc<V>>, Option<&V>) {
fn next_items(&self, token: IterToken) -> (IterToken, &[u8], Option<&TrieNodeODRc<V>>, Option<&V>) {
if token == 0 {
let node_key = self.key();
if self.is_used_child() {
Expand Down
142 changes: 108 additions & 34 deletions src/dense_byte_node.rs
Original file line number Diff line number Diff line change
Expand Up @@ -303,6 +303,52 @@ impl<V: Clone + Send + Sync, A: Allocator, Cf: CoFree<V=V, A=A>> ByteNode<Cf, A>
unsafe{ self.values.get_unchecked_mut(ix) }
}

// ------ IterToken format for ByteNode ------
// Iter tokens pack the next byte position to inspect with the corresponding index into `values`.
// The low 9 bits store one more than the last returned byte, or 0 before iteration starts. The
// upper bits store the index of the next CoFree in `values`, avoiding a `mask.index_of` recompute
// on every item. `iter_token_for_path` computes the values index once for the requested path.
const ITER_TOKEN_BYTE_BITS: u64 = 9;
const ITER_TOKEN_BYTE_MASK: IterToken = (1 << Self::ITER_TOKEN_BYTE_BITS) - 1;

#[inline(always)]
fn iter_token(next_byte: u16, values_idx: u16) -> IterToken {
(values_idx as IterToken) << Self::ITER_TOKEN_BYTE_BITS | next_byte as IterToken
}

#[inline(always)]
fn iter_token_next_byte(token: IterToken) -> u16 {
(token & Self::ITER_TOKEN_BYTE_MASK) as u16
}

#[inline(always)]
fn iter_token_values_idx(token: IterToken) -> usize {
(token >> Self::ITER_TOKEN_BYTE_BITS) as usize
}

#[inline(always)]
fn next_iter_item_from(&self, token: IterToken) -> Option<(u8, usize)> {
let start = Self::iter_token_next_byte(token);
if start >= 256 {
return None;
}

let mut word_idx = (start >> 6) as usize;
let mut bit_idx = (start & 0x3F) as u32;
loop {
let word = unsafe { *self.mask.0.get_unchecked(word_idx) } & (!0u64 << bit_idx);
if word != 0 {
return Some(((word_idx as u8) * 64 + word.trailing_zeros() as u8, Self::iter_token_values_idx(token)));
}

word_idx += 1;
if word_idx == 4 {
return None;
}
bit_idx = 0;
}
}

#[inline]
fn is_empty(&self) -> bool {
self.mask.is_empty_mask()
Expand Down Expand Up @@ -923,49 +969,32 @@ impl<V: Clone + Send + Sync, A: Allocator, Cf: CoFree<V=V, A=A>> TrieNode<V, A>
self.values.len() == 0
}
#[inline(always)]
fn new_iter_token(&self) -> u128 {
self.mask.0[0] as u128
fn new_iter_token(&self) -> IterToken {
Self::iter_token(0, 0)
}
#[inline(always)]
fn iter_token_for_path(&self, key: &[u8]) -> u128 {
fn iter_token_for_path(&self, key: &[u8]) -> IterToken {
if key.len() != 1 {
self.new_iter_token()
} else {
let k = *unsafe{ key.get_unchecked(0) } as usize;
let idx = (k & 0b11000000) >> 6;
let bit_i = k & 0b00111111;
debug_assert!(idx < 4);
let mask: u64 = if bit_i+1 < 64 {
(0xFFFFFFFFFFFFFFFF << bit_i+1) & unsafe{ self.mask.0.get_unchecked(idx) }
} else {
0
};
((idx as u128) << 64) | (mask as u128)
let key_byte = unsafe{ *key.get_unchecked(0) };
let mut values_idx = self.mask.index_of(key_byte);
if self.mask.test_bit(key_byte) {
values_idx += 1;
}
Self::iter_token(key_byte as u16 + 1, values_idx as u16)
}
}
#[inline(always)]
fn next_items(&self, token: u128) -> (u128, &[u8], Option<&TrieNodeODRc<V, A>>, Option<&V>) {
let mut i = (token >> 64) as u8;
let mut w = token as u64;
loop {
if w != 0 {
let wi = w.trailing_zeros() as u8;
w ^= 1u64 << wi;
let k = i*64 + wi;

let new_token = ((i as u128) << 64) | (w as u128);
let cf = unsafe{ self.get_unchecked(k) };
let k = k as usize;
return (new_token, &ALL_BYTES[k..=k], cf.rec(), cf.val())

} else if i < 3 {
i += 1;
fn next_items(&self, token: IterToken) -> (IterToken, &[u8], Option<&TrieNodeODRc<V, A>>, Option<&V>) {
let Some((k, values_idx)) = self.next_iter_item_from(token) else {
return (NODE_ITER_FINISHED, &[], None, None);
};

w = unsafe { *self.mask.0.get_unchecked(i as usize) };
} else {
return (NODE_ITER_FINISHED, &[], None, None)
}
}
let next_token = Self::iter_token(k as u16 + 1, values_idx as u16 + 1);
let cf = unsafe{ self.values.get_unchecked(values_idx) };
let k = k as usize;
(next_token, &ALL_BYTES[k..=k], cf.rec(), cf.val())
}
fn node_val_count(&self, cache: &mut HashMap<u64, usize>) -> usize {
//Discussion: These two implementations do the same thing but with a slightly different ordering of
Expand Down Expand Up @@ -2378,3 +2407,48 @@ fn bit_siblings() {
assert_eq!(63, bit_sibling(63, 1u64 << 63, false));
assert_eq!(63, bit_sibling(63, 1u64 << 63, true));
}

#[test]
fn byte_node_iter_token_crosses_mask_word_boundaries() {
let mut node = DenseByteNode::new_in(crate::alloc::global_alloc());
for byte in [0, 63, 64, 127, 128, 191, 192, 255] {
node.set_val(byte, byte);
}

let mut token = node.new_iter_token();
let mut visited = Vec::new();
while token != NODE_ITER_FINISHED {
let (next_token, path, _child, value) = node.next_items(token);
token = next_token;
if token != NODE_ITER_FINISHED {
assert_eq!(path.len(), 1);
assert_eq!(Some(&path[0]), value);
visited.push(path[0]);
}
}
assert_eq!(visited, [0, 63, 64, 127, 128, 191, 192, 255]);

let mut token = node.iter_token_for_path(&[127]);
let mut visited_after_127 = Vec::new();
while token != NODE_ITER_FINISHED {
let (next_token, path, _child, value) = node.next_items(token);
token = next_token;
if token != NODE_ITER_FINISHED {
assert_eq!(Some(&path[0]), value);
visited_after_127.push(path[0]);
}
}
assert_eq!(visited_after_127, [128, 191, 192, 255]);

let mut token = node.iter_token_for_path(&[65]);
let mut visited_after_missing_65 = Vec::new();
while token != NODE_ITER_FINISHED {
let (next_token, path, _child, value) = node.next_items(token);
token = next_token;
if token != NODE_ITER_FINISHED {
assert_eq!(Some(&path[0]), value);
visited_after_missing_65.push(path[0]);
}
}
assert_eq!(visited_after_missing_65, [127, 128, 191, 192, 255]);
}
6 changes: 3 additions & 3 deletions src/empty_node.rs
Original file line number Diff line number Diff line change
Expand Up @@ -58,13 +58,13 @@ impl<V: Clone + Send + Sync, A: Allocator> TrieNode<V, A> for EmptyNode {
}
fn node_remove_unmasked_branches(&mut self, _key: &[u8], _mask: ByteMask, _prune: bool) {}
fn node_is_empty(&self) -> bool { true }
fn new_iter_token(&self) -> u128 {
fn new_iter_token(&self) -> IterToken {
0
}
fn iter_token_for_path(&self, _key: &[u8]) -> u128 {
fn iter_token_for_path(&self, _key: &[u8]) -> IterToken {
0
}
fn next_items(&self, _token: u128) -> (u128, &[u8], Option<&TrieNodeODRc<V, A>>, Option<&V>) {
fn next_items(&self, _token: IterToken) -> (IterToken, &[u8], Option<&TrieNodeODRc<V, A>>, Option<&V>) {
(NODE_ITER_FINISHED, &[], None, None)
}
fn node_val_count(&self, _cache: &mut HashMap<u64, usize>) -> usize {
Expand Down
6 changes: 3 additions & 3 deletions src/line_list_node.rs
Original file line number Diff line number Diff line change
Expand Up @@ -1893,7 +1893,7 @@ impl<V: Clone + Send + Sync, A: Allocator> TrieNode<V, A> for LineListNode<V, A>
// * NODE_ITER_FINISHED
// *==--==**==--==**==--==**==--==**==--==**==--==**==--==**==--==**==--==**==--==**==--==**==--==*
#[inline(always)]
fn new_iter_token(&self) -> u128 {
fn new_iter_token(&self) -> IterToken {
0
}
/// Explanation of logic: The ListNode contains a sorted list of keys (up to 2 of them), and the
Expand All @@ -1904,7 +1904,7 @@ impl<V: Clone + Send + Sync, A: Allocator> TrieNode<V, A> for LineListNode<V, A>
/// - == key1, we should return (2, key1)
/// - > key1, (NODE_ITER_FINISHED, &[])
#[inline(always)]
fn iter_token_for_path(&self, key: &[u8]) -> u128 {
fn iter_token_for_path(&self, key: &[u8]) -> IterToken {
if key.len() == 0 {
return 0
}
Expand All @@ -1921,7 +1921,7 @@ impl<V: Clone + Send + Sync, A: Allocator> TrieNode<V, A> for LineListNode<V, A>
NODE_ITER_FINISHED
}
#[inline(always)]
fn next_items(&self, token: u128) -> (u128, &[u8], Option<&TrieNodeODRc<V, A>>, Option<&V>) {
fn next_items(&self, token: IterToken) -> (IterToken, &[u8], Option<&TrieNodeODRc<V, A>>, Option<&V>) {
match token {
0 => {
if !self.is_used::<0>() {
Expand Down
2 changes: 1 addition & 1 deletion src/old_cursor.rs
Original file line number Diff line number Diff line change
Expand Up @@ -249,7 +249,7 @@ impl <'a, V : Clone + Send + Sync> Iterator for ByteTrieNodeIter<'a, V> {

pub struct PathMapCursor<'a, V: Clone + Send + Sync> {
prefix_buf: Vec<u8>,
btnis: Vec<(TaggedNodeRef<'a, V, GlobalAlloc>, u128, usize)>,
btnis: Vec<(TaggedNodeRef<'a, V, GlobalAlloc>, IterToken, usize)>,
}

impl <'a, V : Clone + Send + Sync + Unpin> PathMapCursor<'a, V> {
Expand Down
6 changes: 3 additions & 3 deletions src/tiny_node.rs
Original file line number Diff line number Diff line change
Expand Up @@ -214,9 +214,9 @@ impl<'a, V: Clone + Send + Sync, A: Allocator> TrieNode<V, A> for TinyRefNode<'a
fn node_is_empty(&self) -> bool {
self.header & (1 << 7) == 0
}
fn new_iter_token(&self) -> u128 { unreachable!() }
fn iter_token_for_path(&self, _key: &[u8]) -> u128 { unreachable!() }
fn next_items(&self, _token: u128) -> (u128, &'a[u8], Option<&TrieNodeODRc<V, A>>, Option<&V>) { unreachable!() }
fn new_iter_token(&self) -> IterToken { unreachable!() }
fn iter_token_for_path(&self, _key: &[u8]) -> IterToken { unreachable!() }
fn next_items(&self, _token: IterToken) -> (IterToken, &'a[u8], Option<&TrieNodeODRc<V, A>>, Option<&V>) { unreachable!() }
fn node_val_count(&self, cache: &mut HashMap<u64, usize>) -> usize {
let temp_node = self.into_full().unwrap();
temp_node.node_val_count(cache)
Expand Down
Loading