Skip to main content

zerotrie/byte_phf/
builder.rs

1// This file is part of ICU4X. For terms of use, please see the file
2// called LICENSE at the top level of the ICU4X source tree
3// (online at: https://github.com/unicode-org/icu4x/blob/main/LICENSE ).
4
5use super::*;
6use crate::error::ZeroTrieBuildError;
7use alloc::vec;
8use alloc::vec::Vec;
9
10/// To speed up the search algorithm, we limit the number of times the level-2 parameter (q)
11/// can hit its max value (initially `Q_FAST_MAX`) before we try the next level-1 parameter (p).
12/// In practice, this has a small impact on the resulting perfect hash, resulting in about
13/// 1 in 10000 hash maps that fall back to the slow path.
14const MAX_L2_SEARCH_MISSES: usize = 24;
15
16/// Directly compute the perfect hash function.
17///
18/// Returns `(p, [q_0, q_1, ..., q_(N-1)])`, or an error if the PHF could not be computed.
19#[allow(unused_labels)] // for readability
20#[allow(clippy::indexing_slicing)] // carefully reviewed to not panic
21pub fn find(bytes: &[u8]) -> Result<(u8, Vec<u8>), ZeroTrieBuildError> {
22    let n_usize = bytes.len();
23
24    let mut p = 0u8;
25    let mut qq = vec![0u8; n_usize];
26
27    let mut bqs = vec![0u8; n_usize];
28    let mut seen = vec![false; n_usize];
29    let max_allowable_p = P_FAST_MAX;
30    let mut max_allowable_q = Q_FAST_MAX;
31
32    #[allow(non_snake_case)]
33    let N = if n_usize > 0 && n_usize < 256 {
34        n_usize as u8
35    } else {
36        debug_assert!(n_usize == 0 || n_usize == 256);
37        return Ok((p, qq));
38    };
39
40    'p_loop: loop {
41        // Vec of tuples: (index, bucket count)
42        let mut buckets: Vec<(usize, Vec<u8>)> = (0..n_usize).map(|i| (i, vec![])).collect();
43        for byte in bytes {
44            let l1 = f1(*byte, p, N) as usize;
45            buckets[l1].1.push(*byte);
46        }
47        buckets.sort_by_key(|(_, v)| -(v.len() as isize));
48        // println!("New P: p={p:?}, buckets={buckets:?}");
49        let mut i = 0;
50        let mut num_max_q = 0;
51        bqs.fill(0);
52        seen.fill(false);
53        'q_loop: loop {
54            // Loop condition: exit when i is beyond the buckets length
55            if i == buckets.len() {
56                for (local_j, real_j) in buckets.iter().map(|(j, _)| *j).enumerate() {
57                    debug_assert!(local_j < n_usize); // comes from .enumerate()
58                    debug_assert!(real_j < n_usize); // first item of bucket tuple is an index
59                    qq[real_j] = bqs[local_j];
60                }
61                // println!("Success: p={p:?}, num_max_q={num_max_q:?}, bqs={bqs:?}, qq={qq:?}");
62                // if num_max_q > 0 {
63                //     println!("num_max_q={num_max_q:?}");
64                // }
65                return Ok((p, qq));
66            }
67            let mut bucket = buckets[i].1.as_slice();
68            'byte_loop: for (j, byte) in bucket.iter().enumerate() {
69                let l2 = f2(*byte, bqs[i], N) as usize;
70                if seen[l2] {
71                    // println!("Skipping Q: p={p:?}, i={i:?}, byte={byte:}, q={i:?}, l2={:?}", f2(*byte, bqs[i], N));
72                    for k_byte in &bucket[0..j] {
73                        let l2 = f2(*k_byte, bqs[i], N) as usize;
74                        assert!(seen[l2]);
75                        seen[l2] = false;
76                    }
77                    'reset_loop: loop {
78                        if bqs[i] < max_allowable_q {
79                            bqs[i] += 1;
80                            continue 'q_loop;
81                        }
82                        num_max_q += 1;
83                        bqs[i] = 0;
84                        if i == 0 || num_max_q > MAX_L2_SEARCH_MISSES {
85                            if p == max_allowable_p && max_allowable_q != Q_REAL_MAX {
86                                // println!("Could not solve fast function: trying again: {bytes:?}");
87                                max_allowable_q = Q_REAL_MAX;
88                                p = 0;
89                                continue 'p_loop;
90                            } else if p == max_allowable_p {
91                                // If a fallback algorithm for `p` is added, relax this assertion
92                                // and re-run the loop with a higher `max_allowable_p`.
93                                debug_assert_eq!(max_allowable_p, P_REAL_MAX);
94                                // println!("Could not solve PHF function");
95                                return Err(ZeroTrieBuildError::CouldNotSolvePerfectHash);
96                            } else {
97                                p += 1;
98                                continue 'p_loop;
99                            }
100                        }
101                        i -= 1;
102                        bucket = buckets[i].1.as_slice();
103                        for byte in bucket {
104                            let l2 = f2(*byte, bqs[i], N) as usize;
105                            assert!(seen[l2]);
106                            seen[l2] = false;
107                        }
108                    }
109                } else {
110                    // println!("Marking as seen: i={i:?}, byte={byte:}, l2={:?}", f2(*byte, bqs[i], N));
111                    let l2 = f2(*byte, bqs[i], N) as usize;
112                    seen[l2] = true;
113                }
114            }
115            // println!("Found Q: i={i:?}, q={:?}", bqs[i]);
116            i += 1;
117        }
118    }
119}
120
121impl PerfectByteHashMap<Vec<u8>> {
122    /// Computes a new [`PerfectByteHashMap`].
123    ///
124    /// (this is a doc-hidden API)
125    #[allow(clippy::indexing_slicing)] // carefully reviewed to not panic
126    pub fn try_new(keys: &[u8]) -> Result<Self, ZeroTrieBuildError> {
127        let n_usize = keys.len();
128        let n = n_usize as u8;
129        let (p, mut qq) = find(keys)?;
130        let mut keys_permuted = vec![0; n_usize];
131        for key in keys {
132            let l1 = f1(*key, p, n) as usize;
133            let q = qq[l1];
134            let l2 = f2(*key, q, n) as usize;
135            keys_permuted[l2] = *key;
136        }
137        let mut result = Vec::with_capacity(n_usize * 2 + 1);
138        result.push(p);
139        result.append(&mut qq);
140        result.append(&mut keys_permuted);
141        Ok(Self(result))
142    }
143}
144
145#[cfg(test)]
146mod tests {
147    use super::*;
148
149    extern crate std;
150    use std::print;
151    use std::println;
152
153    fn print_byte_to_stdout(byte: u8) {
154        let c = char::from(byte);
155        if c.is_ascii_alphanumeric() {
156            print!("'{c}'");
157        } else {
158            print!("0x{byte:X}");
159        }
160    }
161
162    fn random_alphanums(seed: u64, len: usize) -> Vec<u8> {
163        use rand::seq::SliceRandom;
164        use rand::SeedableRng;
165        let mut bytes: Vec<u8> =
166            b"ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789".into();
167        let mut rng = rand_pcg::Lcg64Xsh32::seed_from_u64(seed);
168        bytes.partial_shuffle(&mut rng, len).0.into()
169    }
170
171    #[test]
172    fn test_random_distributions() {
173        let mut p_distr = vec![0; 256];
174        let mut q_distr = vec![0; 256];
175        for len in 0..50 {
176            for seed in 0..50 {
177                let bytes = random_alphanums(seed, len);
178                let (p, qq) = find(bytes.as_slice()).unwrap();
179                p_distr[p as usize] += 1;
180                for q in qq {
181                    q_distr[q as usize] += 1;
182                }
183            }
184        }
185        println!("p_distr: {p_distr:?}");
186        println!("q_distr: {q_distr:?}");
187
188        let fast_p = p_distr[0..=P_FAST_MAX as usize].iter().sum::<usize>();
189        let slow_p = p_distr[(P_FAST_MAX + 1) as usize..].iter().sum::<usize>();
190        let fast_q = q_distr[0..=Q_FAST_MAX as usize].iter().sum::<usize>();
191        let slow_q = q_distr[(Q_FAST_MAX + 1) as usize..].iter().sum::<usize>();
192
193        assert_eq!(2500, fast_p);
194        assert_eq!(0, slow_p);
195        assert_eq!(61243, fast_q);
196        assert_eq!(7, slow_q);
197
198        let bytes = random_alphanums(0, 16);
199
200        #[allow(non_snake_case)]
201        let N = u8::try_from(bytes.len()).unwrap();
202
203        let (p, qq) = find(bytes.as_slice()).unwrap();
204
205        println!("Results:");
206        for byte in bytes.iter() {
207            print_byte_to_stdout(*byte);
208            let l1 = f1(*byte, p, N) as usize;
209            let q = qq[l1];
210            let l2 = f2(*byte, q, N) as usize;
211            println!(" => l1 {l1} => q {q} => l2 {l2}");
212        }
213    }
214}