Skip to main content

vole_document/entropy/
model.rs

1//! Canonical entropy models: a deterministic, integer-only (no floating point)
2//! frequency-table normalizer and its wire serialization.
3//!
4//! A model is a frequency table over a fixed byte alphabet ([`ALPHABET`] = 256)
5//! whose entries sum to exactly `1 << scale_bits`. Normalization from raw symbol
6//! counts is a *pure function of the counts and `scale_bits`*: identical inputs
7//! always produce identical output bytes. This is the substrate rANS will later
8//! consume; no entropy coder lives here yet.
9
10use crate::error::{Error, Result};
11
12/// Supported alphabet size (bytes).
13pub const ALPHABET: usize = 256;
14/// Model wire version.
15pub const MODEL_VERSION_1: u8 = 1;
16/// Minimum / maximum scale bits. Frequencies must fit in u16, so <= 15.
17pub const MIN_SCALE_BITS: u8 = 1;
18pub const MAX_SCALE_BITS: u8 = 15;
19
20/// Bytes preceding the frequency table: `version`, `scale_bits`, `count`.
21const MODEL_HEADER_LEN: usize = 4;
22/// Total canonical encoded length: header plus `ALPHABET` little-endian u16s.
23const MODEL_ENCODED_LEN: usize = MODEL_HEADER_LEN + ALPHABET * 2;
24
25/// A normalized frequency table over the byte alphabet.
26#[derive(Debug, Clone, PartialEq, Eq)]
27pub struct EntropyModel {
28    /// Number of scale bits; the table sums to `1 << scale_bits`.
29    pub scale_bits: u8,
30    /// Exactly `ALPHABET` entries; sum == `1 << scale_bits`.
31    pub frequencies: Vec<u32>,
32}
33
34impl EntropyModel {
35    /// Total frequency (== `1 << scale_bits`).
36    pub fn total(&self) -> u32 {
37        1u32 << self.scale_bits
38    }
39
40    /// Sum of all frequencies as a widened integer (avoids debug overflow on
41    /// hand-built, out-of-contract tables).
42    fn frequency_sum(&self) -> u64 {
43        self.frequencies.iter().map(|&f| u64::from(f)).sum()
44    }
45
46    /// Canonical normalization of observed symbol counts into a frequency table.
47    ///
48    /// Requirements (all hold):
49    /// - integer-only; no floating point anywhere;
50    /// - deterministic: same input yields identical output bytes;
51    /// - every symbol with `count > 0` gets frequency `>= 1`;
52    /// - symbols with `count == 0` get frequency `0`;
53    /// - `sum(frequencies) == 1 << scale_bits` exactly;
54    /// - if all counts are `0`, return the canonical uniform model.
55    ///
56    /// Algorithm: after granting each *present* symbol a guaranteed minimum of
57    /// 1, the remaining budget is apportioned by the largest-remainder (Hare)
58    /// method using `u128` products, with ties broken by lower symbol index.
59    pub fn from_counts(counts: &[u64; ALPHABET], scale_bits: u8) -> Result<EntropyModel> {
60        if !(MIN_SCALE_BITS..=MAX_SCALE_BITS).contains(&scale_bits) {
61            let (min, max) = (MIN_SCALE_BITS, MAX_SCALE_BITS);
62            return Err(Error::invalid_model(format!(
63                "scale_bits {scale_bits} outside {min}..={max}"
64            )));
65        }
66        let target: u32 = 1u32 << scale_bits;
67
68        // 2. Present symbols.
69        let present: Vec<usize> = (0..ALPHABET).filter(|&i| counts[i] > 0).collect();
70        if present.is_empty() {
71            return Self::uniform(scale_bits);
72        }
73
74        // 3. Alphabet (as a set of present symbols) must fit in the target.
75        if present.len() as u64 > u64::from(target) {
76            return Err(Error::invalid_model(format!(
77                "{} present symbols exceed target {target}",
78                present.len()
79            )));
80        }
81
82        // 4. Guaranteed minimum of 1 for each present symbol.
83        let mut frequencies = vec![0u32; ALPHABET];
84        for &i in &present {
85            frequencies[i] = 1;
86        }
87        let remaining: u32 = target - present.len() as u32;
88
89        if remaining > 0 {
90            let total: u128 = present.iter().map(|&i| u128::from(counts[i])).sum();
91
92            // 5. Integer quota + remainder per present symbol.
93            let mut quota_sum: u64 = 0;
94            let mut rems: Vec<(usize, u128)> = Vec::with_capacity(present.len());
95            for &i in &present {
96                let product = u128::from(counts[i]) * u128::from(remaining);
97                let quota = product / total;
98                let rem = product % total;
99                frequencies[i] += u32::try_from(quota)
100                    .map_err(|_| Error::internal_invariant("quota exceeded 32-bit range"))?;
101                quota_sum += u64::try_from(quota)
102                    .map_err(|_| Error::internal_invariant("quota exceeded 64-bit range"))?;
103                rems.push((i, rem));
104            }
105            let leftover: u64 = u64::from(remaining) - quota_sum;
106
107            // 6. Largest remainder, ties by lower index. There are always at
108            // least `leftover` symbols with a nonzero remainder, so this pass
109            // consumes the whole leftover in practice.
110            rems.sort_by(|a, b| b.1.cmp(&a.1).then_with(|| a.0.cmp(&b.0)));
111            let mut added: u64 = 0;
112            for &(i, _) in &rems {
113                if added >= leftover {
114                    break;
115                }
116                frequencies[i] += 1;
117                added += 1;
118            }
119
120            // Deterministic fallback, kept for total-sum safety. It is
121            // unreachable for well-formed inputs (the remainder pass above
122            // always suffices), and adds to the present symbol with the largest
123            // count, breaking ties by lower index.
124            while added < leftover {
125                let i = present
126                    .iter()
127                    .copied()
128                    .max_by_key(|&i| (counts[i], core::cmp::Reverse(i)))
129                    .expect("present is non-empty");
130                frequencies[i] += 1;
131                added += 1;
132            }
133        }
134
135        let model = EntropyModel {
136            scale_bits,
137            frequencies,
138        };
139        if model.frequency_sum() != u64::from(target) {
140            return Err(Error::internal_invariant(
141                "normalized frequencies do not sum to target",
142            ));
143        }
144        Ok(model)
145    }
146
147    /// Uniform model with every symbol getting an equal frequency.
148    ///
149    /// Requires `scale_bits >= 8`, i.e. `1 << scale_bits >= ALPHABET`; below
150    /// that no integer table of 256 equal entries can sum to the target, so the
151    /// request is rejected as [`ErrorClass::InvalidModel`](crate::error::ErrorClass::InvalidModel).
152    pub fn uniform(scale_bits: u8) -> Result<EntropyModel> {
153        if !(MIN_SCALE_BITS..=MAX_SCALE_BITS).contains(&scale_bits) {
154            let (min, max) = (MIN_SCALE_BITS, MAX_SCALE_BITS);
155            return Err(Error::invalid_model(format!(
156                "scale_bits {scale_bits} outside {min}..={max}"
157            )));
158        }
159        let target: u32 = 1u32 << scale_bits;
160        if target < ALPHABET as u32 {
161            return Err(Error::invalid_model(format!(
162                "scale_bits {scale_bits} cannot hold {ALPHABET} equal frequencies"
163            )));
164        }
165        // `target` is a power of two >= 256, hence exactly divisible by 256.
166        let per = target / ALPHABET as u32;
167        Ok(EntropyModel {
168            scale_bits,
169            frequencies: vec![per; ALPHABET],
170        })
171    }
172
173    /// Canonical wire encoding (little-endian):
174    /// `[version u8 = 1][scale_bits u8][count u16 = 256][freq u16 x 256]`.
175    ///
176    /// The model is validated before serialization so that `encode` cannot emit
177    /// a non-canonical table.
178    pub fn encode(&self) -> Result<Vec<u8>> {
179        if !(MIN_SCALE_BITS..=MAX_SCALE_BITS).contains(&self.scale_bits) {
180            let (min, max) = (MIN_SCALE_BITS, MAX_SCALE_BITS);
181            return Err(Error::invalid_model(format!(
182                "scale_bits {} outside {min}..={max}",
183                self.scale_bits
184            )));
185        }
186        if self.frequencies.len() != ALPHABET {
187            return Err(Error::invalid_model(format!(
188                "expected {ALPHABET} frequencies, got {}",
189                self.frequencies.len()
190            )));
191        }
192        let target = u64::from(self.total());
193        if self.frequency_sum() != target {
194            return Err(Error::invalid_model(
195                "frequencies do not sum to 1 << scale_bits",
196            ));
197        }
198
199        let mut out = Vec::with_capacity(MODEL_ENCODED_LEN);
200        out.push(MODEL_VERSION_1);
201        out.push(self.scale_bits);
202        out.extend_from_slice(&(ALPHABET as u16).to_le_bytes());
203        for &f in &self.frequencies {
204            let f =
205                u16::try_from(f).map_err(|_| Error::invalid_model("frequency does not fit u16"))?;
206            out.extend_from_slice(&f.to_le_bytes());
207        }
208        Ok(out)
209    }
210
211    /// Parse and validate a canonical model; `bytes` must be exactly consumed.
212    pub fn decode(bytes: &[u8]) -> Result<EntropyModel> {
213        if bytes.len() < MODEL_HEADER_LEN {
214            return Err(Error::invalid_model("model shorter than header"));
215        }
216        let version = bytes[0];
217        if version != MODEL_VERSION_1 {
218            return Err(Error::unsupported_version(format!(
219                "model version {version}, expected {MODEL_VERSION_1}"
220            )));
221        }
222        let scale_bits = bytes[1];
223        if !(MIN_SCALE_BITS..=MAX_SCALE_BITS).contains(&scale_bits) {
224            let (min, max) = (MIN_SCALE_BITS, MAX_SCALE_BITS);
225            return Err(Error::invalid_model(format!(
226                "scale_bits {scale_bits} outside {min}..={max}"
227            )));
228        }
229        let count = u16::from_le_bytes([bytes[2], bytes[3]]);
230        if count as usize != ALPHABET {
231            return Err(Error::invalid_model(format!(
232                "declared count {count}, expected {ALPHABET}"
233            )));
234        }
235        if bytes.len() != MODEL_ENCODED_LEN {
236            return Err(Error::invalid_model(
237                "model length does not match declared count (trailing or truncated)",
238            ));
239        }
240
241        let target = u64::from(1u32 << scale_bits);
242        let mut frequencies = Vec::with_capacity(ALPHABET);
243        let mut sum: u64 = 0;
244        let (chunks, _rest) = bytes[MODEL_HEADER_LEN..].as_chunks::<2>();
245        for chunk in chunks {
246            let f = u32::from(u16::from_le_bytes(*chunk));
247            sum += u64::from(f);
248            frequencies.push(f);
249        }
250        if sum != target {
251            return Err(Error::invalid_model(
252                "frequencies do not sum to 1 << scale_bits",
253            ));
254        }
255        Ok(EntropyModel {
256            scale_bits,
257            frequencies,
258        })
259    }
260}
261
262#[cfg(test)]
263mod tests {
264    use super::*;
265
266    fn counts_with(pairs: &[(usize, u64)]) -> [u64; ALPHABET] {
267        let mut counts = [0u64; ALPHABET];
268        for &(i, v) in pairs {
269            counts[i] = v;
270        }
271        counts
272    }
273
274    #[test]
275    fn uniform_sums_to_target() {
276        for bits in [8u8, 12, 15] {
277            let model = EntropyModel::uniform(bits).expect("uniform must succeed");
278            assert_eq!(model.frequencies.len(), ALPHABET);
279            assert_eq!(model.frequency_sum(), u64::from(1u32 << bits));
280            let per = model.frequencies[0];
281            assert!(model.frequencies.iter().all(|&f| f == per));
282        }
283    }
284
285    #[test]
286    fn from_counts_sums_exactly() {
287        let inputs: &[&[(usize, u64)]] = &[
288            &[(0, 1)],
289            &[(0, 1), (1, 1)],
290            &[(0, 300), (1, 1)],
291            &[(0, 1), (1, 2), (2, 3), (3, 4)],
292            &[(5, 10), (200, 1), (255, 7)],
293            &[(0, u64::MAX), (255, 1)],
294        ];
295        for bits in [8u8, 12] {
296            for input in inputs {
297                let counts = counts_with(input);
298                let model =
299                    EntropyModel::from_counts(&counts, bits).expect("normalization must succeed");
300                assert_eq!(model.scale_bits, bits);
301                assert_eq!(model.frequency_sum(), u64::from(1u32 << bits));
302                for &(i, _) in *input {
303                    assert!(model.frequencies[i] >= 1);
304                }
305            }
306        }
307    }
308
309    #[test]
310    fn determinism() {
311        let counts = counts_with(&[(0, 17), (3, 5), (7, 250), (255, 1)]);
312        let a = EntropyModel::from_counts(&counts, 12)
313            .expect("ok")
314            .encode()
315            .expect("ok");
316        let b = EntropyModel::from_counts(&counts, 12)
317            .expect("ok")
318            .encode()
319            .expect("ok");
320        assert_eq!(a, b);
321    }
322
323    #[test]
324    fn present_symbols_get_at_least_one() {
325        // Every symbol present; at scale_bits 12 the target (4096) comfortably
326        // exceeds the alphabet size (256).
327        let counts = [1u64; ALPHABET];
328        let model = EntropyModel::from_counts(&counts, 12).expect("ok");
329        assert!(model.frequencies.iter().all(|&f| f >= 1));
330        assert_eq!(model.frequency_sum(), 4096);
331    }
332
333    #[test]
334    fn zero_symbols_get_zero() {
335        let counts = counts_with(&[(0, 5), (10, 9), (255, 3)]);
336        let model = EntropyModel::from_counts(&counts, 12).expect("ok");
337        for (i, &f) in model.frequencies.iter().enumerate() {
338            if counts[i] == 0 {
339                assert_eq!(f, 0, "symbol {i} had zero count but nonzero frequency");
340            }
341        }
342    }
343
344    #[test]
345    fn alphabet_too_big_errors() {
346        let counts = [1u64; ALPHABET];
347        for bits in 1u8..=7 {
348            assert!(EntropyModel::from_counts(&counts, bits).is_err());
349        }
350    }
351
352    #[test]
353    fn roundtrip_encode_decode() {
354        let counts = counts_with(&[(0, 1), (1, 2), (2, 3), (100, 400), (255, 7)]);
355        for bits in [8u8, 12, 15] {
356            let model = EntropyModel::from_counts(&counts, bits).expect("ok");
357            let bytes = model.encode().expect("ok");
358            assert_eq!(bytes.len(), MODEL_ENCODED_LEN);
359            let decoded = EntropyModel::decode(&bytes).expect("ok");
360            assert_eq!(decoded, model);
361        }
362    }
363
364    #[test]
365    fn decode_rejects_bad_sum() {
366        let mut bytes = EntropyModel::uniform(8).expect("ok").encode().expect("ok");
367        // Bump the last frequency so the total no longer equals 1 << 8.
368        let last = bytes.len();
369        let value = u16::from_le_bytes([bytes[last - 2], bytes[last - 1]]);
370        let bumped = value.checked_add(1).expect("fits");
371        bytes[last - 2..].copy_from_slice(&bumped.to_le_bytes());
372        assert!(EntropyModel::decode(&bytes).is_err());
373    }
374
375    #[test]
376    fn decode_rejects_bad_version() {
377        let mut bytes = EntropyModel::uniform(8).expect("ok").encode().expect("ok");
378        bytes[0] = 2;
379        let err = EntropyModel::decode(&bytes).expect_err("must reject");
380        assert_eq!(err.class(), crate::error::ErrorClass::UnsupportedVersion);
381    }
382
383    #[test]
384    fn decode_rejects_trailing() {
385        let mut bytes = EntropyModel::uniform(8).expect("ok").encode().expect("ok");
386        bytes.push(0);
387        assert!(EntropyModel::decode(&bytes).is_err());
388    }
389
390    #[test]
391    fn all_zero_is_uniform() {
392        let zero = [0u64; ALPHABET];
393        for bits in [8u8, 12, 15] {
394            let model = EntropyModel::from_counts(&zero, bits).expect("ok");
395            assert_eq!(model, EntropyModel::uniform(bits).expect("ok"));
396        }
397    }
398}