1use std::cell::{Cell, RefCell};
10use std::f64;
11use std::mem::size_of;
12use std::sync::OnceLock;
13
14type Symbol = bool;
15const HISTORY_WORD_BITS: usize = u64::BITS as usize;
16
17#[inline(always)]
18fn history_word_len(bits: usize) -> usize {
19 bits.div_ceil(HISTORY_WORD_BITS)
20}
21
22trait HistoryAccess {
23 fn len(&self) -> usize;
24 fn bit(&self, index: usize) -> Symbol;
25
26 #[inline(always)]
27 fn recent_bit(&self, _depth: usize) -> Option<Symbol> {
28 None
29 }
30
31 #[inline(always)]
32 fn recent_path_bits(&self, _depth: usize, _len: usize) -> Option<u64> {
33 None
34 }
35}
36
37impl HistoryAccess for [Symbol] {
38 #[inline(always)]
39 fn len(&self) -> usize {
40 <[Symbol]>::len(self)
41 }
42
43 #[inline(always)]
44 fn bit(&self, index: usize) -> Symbol {
45 debug_assert!(index < self.len());
46 unsafe { *self.get_unchecked(index) }
47 }
48}
49
50impl HistoryAccess for Vec<Symbol> {
51 #[inline(always)]
52 fn len(&self) -> usize {
53 self.as_slice().len()
54 }
55
56 #[inline(always)]
57 fn bit(&self, index: usize) -> Symbol {
58 self.as_slice().bit(index)
59 }
60}
61
62#[derive(Clone, Debug, Default, PartialEq, Eq)]
63struct BitHistory {
64 words: Vec<u64>,
65 len: usize,
66 recent: u64,
67}
68
69impl BitHistory {
70 #[inline(always)]
71 fn len(&self) -> usize {
72 self.len
73 }
74
75 #[inline(always)]
76 fn is_empty(&self) -> bool {
77 self.len == 0
78 }
79
80 #[inline(always)]
81 fn memory_usage(&self) -> usize {
82 self.words.capacity() * size_of::<u64>()
83 }
84
85 #[inline]
86 fn rebuild_recent(&mut self) {
87 self.recent = 0;
88 let tail_len = self.len.min(HISTORY_WORD_BITS);
89 for depth in 0..tail_len {
90 let idx = self.len - depth - 1;
91 if self.bit(idx) {
92 self.recent |= 1u64 << depth;
93 }
94 }
95 }
96
97 #[inline]
98 fn reserve_exact(&mut self, additional_bits: usize) {
99 let required_bits = self.len.saturating_add(additional_bits);
100 let required_words = history_word_len(required_bits);
101 if required_words > self.words.capacity() {
102 self.words
103 .reserve_exact(required_words.saturating_sub(self.words.len()));
104 }
105 }
106
107 #[inline]
108 fn push(&mut self, bit: Symbol) {
109 let word_idx = self.len / HISTORY_WORD_BITS;
110 let bit_idx = self.len % HISTORY_WORD_BITS;
111 if word_idx == self.words.len() {
112 self.words.push(0);
113 }
114 let mask = 1u64 << bit_idx;
115 if bit {
116 self.words[word_idx] |= mask;
117 } else {
118 self.words[word_idx] &= !mask;
119 }
120 self.recent = (self.recent << 1) | (bit as u64);
121 self.len += 1;
122 }
123
124 #[inline]
125 fn pop(&mut self) -> Option<Symbol> {
126 if self.len == 0 {
127 return None;
128 }
129 let next_len = self.len - 1;
130 let word_idx = next_len / HISTORY_WORD_BITS;
131 let bit_idx = next_len % HISTORY_WORD_BITS;
132 let mask = 1u64 << bit_idx;
133 let bit = (self.words[word_idx] & mask) != 0;
134 if bit {
135 self.words[word_idx] &= !mask;
136 }
137 self.len = next_len;
138 if bit_idx == 0 {
139 self.words.truncate(word_idx);
140 }
141 self.recent >>= 1;
142 if self.len >= HISTORY_WORD_BITS {
143 let exposed_idx = self.len - HISTORY_WORD_BITS;
144 if self.bit(exposed_idx) {
145 self.recent |= 1u64 << (HISTORY_WORD_BITS - 1);
146 }
147 }
148 Some(bit)
149 }
150
151 #[inline]
152 fn extend_from_slice(&mut self, symbols: &[Symbol]) {
153 self.reserve_exact(symbols.len());
154 for &symbol in symbols {
155 self.push(symbol);
156 }
157 }
158
159 #[inline]
160 fn truncate(&mut self, new_len: usize) {
161 if new_len >= self.len {
162 return;
163 }
164 self.len = new_len;
165 self.words.truncate(history_word_len(new_len));
166 let rem = new_len % HISTORY_WORD_BITS;
167 if rem != 0 {
168 let mask = (1u64 << rem) - 1;
169 if let Some(last) = self.words.last_mut() {
170 *last &= mask;
171 }
172 }
173 self.rebuild_recent();
174 }
175
176 #[inline]
177 fn clear(&mut self) {
178 self.words.clear();
179 self.len = 0;
180 self.recent = 0;
181 }
182
183 #[cfg(test)]
184 fn to_vec(&self) -> Vec<Symbol> {
185 (0..self.len).map(|idx| self.bit(idx)).collect()
186 }
187}
188
189impl HistoryAccess for BitHistory {
190 #[inline(always)]
191 fn len(&self) -> usize {
192 self.len
193 }
194
195 #[inline(always)]
196 fn bit(&self, index: usize) -> Symbol {
197 debug_assert!(index < self.len);
198 let word = unsafe { *self.words.get_unchecked(index / HISTORY_WORD_BITS) };
199 ((word >> (index % HISTORY_WORD_BITS)) & 1) != 0
200 }
201
202 #[inline(always)]
203 fn recent_bit(&self, depth: usize) -> Option<Symbol> {
204 if depth < self.len.min(HISTORY_WORD_BITS) {
205 Some(((self.recent >> depth) & 1) != 0)
206 } else {
207 None
208 }
209 }
210
211 #[inline(always)]
212 fn recent_path_bits(&self, depth: usize, len: usize) -> Option<u64> {
213 let available = self.len.saturating_sub(depth).min(len);
214 if available == 0 {
215 return Some(0);
216 }
217 if depth + available <= self.len.min(HISTORY_WORD_BITS) {
218 let mask = if available >= HISTORY_WORD_BITS {
219 u64::MAX
220 } else {
221 (1u64 << available) - 1
222 };
223 Some((self.recent >> depth) & mask)
224 } else {
225 None
226 }
227 }
228}
229
230const CTW_LOG_CACHE_LIMIT: usize = 1 << 24;
231const CTW_HOT_PREFIX_DEPTH_DEFAULT: usize = 12;
232const CTW_LOG_OVERFLOW_CACHE_SLOTS: usize = 1 << 14;
233
234#[cfg(not(test))]
235#[inline(always)]
236fn ctw_log_cache_limit() -> usize {
237 CTW_LOG_CACHE_LIMIT
238}
239
240#[cfg(test)]
241#[inline(always)]
242fn ctw_log_cache_limit() -> usize {
243 CTW_TEST_LOG_CACHE_LIMIT.with(|limit| limit.borrow().unwrap_or(CTW_LOG_CACHE_LIMIT))
244}
245
246fn ctw_hot_prefix_depth_limit() -> usize {
247 static HOT_PREFIX_DEPTH: OnceLock<usize> = OnceLock::new();
248 *HOT_PREFIX_DEPTH.get_or_init(|| {
249 std::env::var("INFOTHEORY_CTW_HOT_PREFIX_DEPTH")
250 .ok()
251 .and_then(|raw| raw.parse::<usize>().ok())
252 .unwrap_or(CTW_HOT_PREFIX_DEPTH_DEFAULT)
253 })
254}
255
256#[cfg(not(test))]
257#[inline(always)]
258fn ctw_log_overflow_cache_slots() -> usize {
259 CTW_LOG_OVERFLOW_CACHE_SLOTS
260}
261
262#[cfg(test)]
263#[inline(always)]
264fn ctw_log_overflow_cache_slots() -> usize {
265 CTW_TEST_LOG_OVERFLOW_CACHE_SLOTS
266 .with(|slots| slots.borrow().unwrap_or(CTW_LOG_OVERFLOW_CACHE_SLOTS))
267}
268
269#[cfg(test)]
270#[inline(always)]
271fn ensure_log_caches(log_int: &mut Vec<f64>, log_half: &mut Vec<f64>, upto: usize) {
272 if upto < log_int.len() {
273 return;
274 }
275 let required_len = upto + 1;
276 log_int.reserve(required_len - log_int.len());
277 log_half.reserve(required_len - log_half.len());
278 push_log_cache_entries(log_int, log_half, upto);
279}
280
281#[inline(always)]
282fn reserve_bounded_log_cache(cache: &mut Vec<f64>, required_len: usize, max_len: usize) {
283 if cache.capacity() >= required_len {
284 return;
285 }
286 let target_capacity = cache
287 .capacity()
288 .saturating_mul(2)
289 .max(required_len)
290 .min(max_len);
291 cache.reserve_exact(target_capacity - cache.len());
292}
293
294#[inline(always)]
295fn ensure_bounded_log_caches(
296 log_int: &mut Vec<f64>,
297 log_half: &mut Vec<f64>,
298 upto: usize,
299 limit: usize,
300) {
301 let target = upto.min(limit);
302 if target < log_int.len() {
303 return;
304 }
305 let required_len = target + 1;
306 let max_len = limit + 1;
307 reserve_bounded_log_cache(log_int, required_len, max_len);
308 reserve_bounded_log_cache(log_half, required_len, max_len);
309 push_log_cache_entries(log_int, log_half, target);
310}
311
312#[inline(always)]
313fn push_log_cache_entries(log_int: &mut Vec<f64>, log_half: &mut Vec<f64>, upto: usize) {
314 let start = log_int.len();
315 for n in start..=upto {
316 if n == 0 {
317 log_int.push(f64::NEG_INFINITY);
318 } else {
319 log_int.push((n as f64).ln());
320 }
321 log_half.push((n as f64 + 0.5).ln());
322 }
323}
324
325struct LogCacheSlot {
326 key: Cell<usize>,
327 value: Cell<f64>,
328}
329
330impl LogCacheSlot {
331 #[inline(always)]
332 fn empty() -> Self {
333 Self {
334 key: Cell::new(usize::MAX),
335 value: Cell::new(0.0),
336 }
337 }
338}
339
340#[inline(always)]
341fn ensure_overflow_log_cache(cache: &mut Vec<LogCacheSlot>, slots: usize) {
342 if slots == 0 {
343 cache.clear();
344 cache.shrink_to(0);
345 return;
346 }
347 if cache.len() == slots {
348 return;
349 }
350 cache.clear();
351 if cache.capacity() < slots {
352 cache.reserve_exact(slots);
353 }
354 cache.resize_with(slots, LogCacheSlot::empty);
355 if cache.capacity() > slots {
356 cache.shrink_to(slots);
357 }
358}
359
360trait CtLogAccess: Copy {
361 fn log_int(self, n: usize) -> f64;
362 fn log_half(self, n: usize) -> f64;
363}
364
365#[derive(Clone, Copy)]
366struct CachedLogs<'a> {
367 log_int: &'a [f64],
368 log_half: &'a [f64],
369}
370
371impl<'a> CachedLogs<'a> {
372 #[cfg(test)]
373 #[inline(always)]
374 fn new(log_int: &'a [f64], log_half: &'a [f64]) -> Self {
375 Self { log_int, log_half }
376 }
377}
378
379impl CtLogAccess for CachedLogs<'_> {
380 #[inline(always)]
381 fn log_int(self, n: usize) -> f64 {
382 debug_assert!(n < self.log_int.len());
383 unsafe { *self.log_int.get_unchecked(n) }
386 }
387
388 #[inline(always)]
389 fn log_half(self, n: usize) -> f64 {
390 debug_assert!(n < self.log_half.len());
391 unsafe { *self.log_half.get_unchecked(n) }
394 }
395}
396
397#[derive(Clone, Copy)]
398struct BoundedLogs<'a> {
399 log_int: &'a [f64],
400 log_half: &'a [f64],
401 overflow_log_int: &'a [LogCacheSlot],
402 overflow_log_half: &'a [LogCacheSlot],
403}
404
405impl<'a> BoundedLogs<'a> {
406 #[cfg(test)]
407 #[inline(always)]
408 fn new(log_int: &'a [f64], log_half: &'a [f64]) -> Self {
409 Self {
410 log_int,
411 log_half,
412 overflow_log_int: &[],
413 overflow_log_half: &[],
414 }
415 }
416
417 #[inline(always)]
418 fn with_overflow(
419 log_int: &'a [f64],
420 log_half: &'a [f64],
421 overflow_log_int: &'a [LogCacheSlot],
422 overflow_log_half: &'a [LogCacheSlot],
423 ) -> Self {
424 Self {
425 log_int,
426 log_half,
427 overflow_log_int,
428 overflow_log_half,
429 }
430 }
431
432 #[inline(always)]
433 fn lookup_overflow(slots: &'a [LogCacheSlot], n: usize, compute: impl FnOnce() -> f64) -> f64 {
434 if slots.is_empty() {
435 return compute();
436 }
437 let len = slots.len();
438 let slot_idx = if len.is_power_of_two() {
439 n & (len - 1)
440 } else {
441 n % len
442 };
443 let slot = &slots[slot_idx];
444 if slot.key.get() == n {
445 slot.value.get()
446 } else {
447 let value = compute();
448 slot.key.set(n);
449 slot.value.set(value);
450 value
451 }
452 }
453}
454
455impl CtLogAccess for BoundedLogs<'_> {
456 #[inline(always)]
457 fn log_int(self, n: usize) -> f64 {
458 if n < self.log_int.len() {
459 self.log_int[n]
460 } else if n == 0 {
461 f64::NEG_INFINITY
462 } else {
463 Self::lookup_overflow(self.overflow_log_int, n, || (n as f64).ln())
464 }
465 }
466
467 #[inline(always)]
468 fn log_half(self, n: usize) -> f64 {
469 if n < self.log_half.len() {
470 self.log_half[n]
471 } else {
472 Self::lookup_overflow(self.overflow_log_half, n, || (n as f64 + 0.5).ln())
473 }
474 }
475}
476
477#[derive(Default)]
478struct SharedLogCache {
479 log_int: Vec<f64>,
480 log_half: Vec<f64>,
481 overflow_log_int: Vec<LogCacheSlot>,
482 overflow_log_half: Vec<LogCacheSlot>,
483}
484
485impl SharedLogCache {
486 fn new() -> Self {
487 Self {
488 log_int: vec![f64::NEG_INFINITY],
489 log_half: vec![(0.5f64).ln()],
490 overflow_log_int: Vec::new(),
491 overflow_log_half: Vec::new(),
492 }
493 }
494
495 #[inline(always)]
496 fn ensure(&mut self, upto: usize) {
497 let limit = ctw_log_cache_limit();
498 ensure_bounded_log_caches(&mut self.log_int, &mut self.log_half, upto, limit);
499 if upto > limit {
500 let slots = ctw_log_overflow_cache_slots();
501 ensure_overflow_log_cache(&mut self.overflow_log_int, slots);
502 ensure_overflow_log_cache(&mut self.overflow_log_half, slots);
503 }
504 }
505
506 #[inline(always)]
507 fn memory_usage(&self) -> usize {
508 self.log_int.capacity() * size_of::<f64>()
509 + self.log_half.capacity() * size_of::<f64>()
510 + self.overflow_log_int.capacity() * size_of::<LogCacheSlot>()
511 + self.overflow_log_half.capacity() * size_of::<LogCacheSlot>()
512 }
513}
514
515thread_local! {
516 static CTW_SHARED_LOG_CACHE: RefCell<SharedLogCache> =
517 RefCell::new(SharedLogCache::new());
518}
519
520#[cfg(test)]
521thread_local! {
522 static CTW_TEST_LOG_CACHE_LIMIT: RefCell<Option<usize>> = const { RefCell::new(None) };
523}
524
525#[cfg(test)]
526thread_local! {
527 static CTW_TEST_LOG_OVERFLOW_CACHE_SLOTS: RefCell<Option<usize>> = const { RefCell::new(None) };
528}
529
530#[inline]
531fn with_shared_cached_logs<R>(upto: usize, f: impl FnOnce(CachedLogs<'_>) -> R) -> R {
532 debug_assert!(upto <= ctw_log_cache_limit());
533 CTW_SHARED_LOG_CACHE.with(|cache_cell| {
534 let mut cache = cache_cell.borrow_mut();
535 cache.ensure(upto);
536 f(CachedLogs {
537 log_int: &cache.log_int,
538 log_half: &cache.log_half,
539 })
540 })
541}
542
543#[inline]
544fn with_shared_bounded_logs<R>(upto: usize, f: impl FnOnce(BoundedLogs<'_>) -> R) -> R {
545 CTW_SHARED_LOG_CACHE.with(|cache_cell| {
546 let mut cache = cache_cell.borrow_mut();
547 cache.ensure(upto);
548 f(BoundedLogs::with_overflow(
549 &cache.log_int,
550 &cache.log_half,
551 &cache.overflow_log_int,
552 &cache.overflow_log_half,
553 ))
554 })
555}
556
557#[inline]
558fn shared_log_cache_memory_usage() -> usize {
559 CTW_SHARED_LOG_CACHE.with(|cache_cell| cache_cell.borrow().memory_usage())
560}
561
562#[cfg(test)]
563#[inline]
564fn shared_log_cache_lens() -> (usize, usize) {
565 CTW_SHARED_LOG_CACHE.with(|cache_cell| {
566 let cache = cache_cell.borrow();
567 (cache.log_int.len(), cache.log_half.len())
568 })
569}
570
571#[cfg(test)]
572#[inline]
573fn shared_log_overflow_cache_lens() -> (usize, usize) {
574 CTW_SHARED_LOG_CACHE.with(|cache_cell| {
575 let cache = cache_cell.borrow();
576 (cache.overflow_log_int.len(), cache.overflow_log_half.len())
577 })
578}
579
580#[cfg(test)]
581pub(crate) fn reset_shared_log_cache_for_test() {
582 CTW_SHARED_LOG_CACHE.with(|cache_cell| {
583 *cache_cell.borrow_mut() = SharedLogCache::new();
584 });
585}
586
587#[inline(always)]
588fn history_symbol<H: HistoryAccess + ?Sized>(history: &H, depth: usize) -> Symbol {
589 if let Some(bit) = history.recent_bit(depth) {
590 return bit;
591 }
592 let idx = history.len().wrapping_sub(depth + 1);
593 if depth < history.len() {
594 history.bit(idx)
595 } else {
596 false
597 }
598}
599
600#[inline(always)]
601fn history_at_or_zero<H: HistoryAccess + ?Sized>(
602 history: &H,
603 history_len: isize,
604 idx: isize,
605) -> Symbol {
606 if idx >= 0 && idx < history_len {
607 history.bit(idx as usize)
608 } else {
609 false
610 }
611}
612
613const INDEX_BITS: u32 = 31;
614const INDEX_LIMIT: usize = 1usize << INDEX_BITS;
615const CHILD_SEGMENT_TAG: u32 = 1u32 << INDEX_BITS;
616const CHILD_INDEX_MASK: u32 = CHILD_SEGMENT_TAG - 1;
617const SEG_META_MODE_SHIFT: u32 = 30;
618const SEG_META_MODE_MASK: u32 = 0b11 << SEG_META_MODE_SHIFT;
619const SEG_LEN_MASK: u32 = !SEG_META_MODE_MASK;
620const SEG_MODE_EXACT: u32 = 0 << SEG_META_MODE_SHIFT;
621const SEG_MODE_HISTORY: u32 = 1 << SEG_META_MODE_SHIFT;
622const SEG_MODE_HISTORY_INVERT: u32 = 2 << SEG_META_MODE_SHIFT;
623const SEG_MODE_CONST: u32 = 3 << SEG_META_MODE_SHIFT;
624const SEG_EXACT_MAX_LEN: u32 = 64;
625
626#[derive(Clone, Copy, Debug, PartialEq, Eq)]
628pub struct NodeIndex(u32);
629
630impl NodeIndex {
631 #[cold]
632 #[inline(never)]
633 fn overflow() -> ! {
634 panic!("ctw node index overflow");
635 }
636
637 #[inline(always)]
638 fn from_usize(idx: usize) -> Self {
639 if idx >= INDEX_LIMIT {
640 Self::overflow();
641 }
642 Self(idx as u32)
643 }
644
645 #[inline(always)]
646 fn get(self) -> usize {
647 self.0 as usize
648 }
649}
650
651#[derive(Clone, Copy, Debug, PartialEq, Eq)]
653struct SegmentIndex(u32);
654
655impl SegmentIndex {
656 #[cold]
657 #[inline(never)]
658 fn overflow() -> ! {
659 panic!("ctw segment index overflow");
660 }
661
662 #[inline(always)]
663 fn from_usize(idx: usize) -> Self {
664 if idx >= INDEX_LIMIT {
665 Self::overflow();
666 }
667 Self(idx as u32)
668 }
669
670 #[inline(always)]
671 fn get(self) -> usize {
672 self.0 as usize
673 }
674}
675
676#[derive(Clone, Copy, Debug, PartialEq, Eq)]
678struct ChildRef(u32);
679
680impl ChildRef {
681 const NONE: ChildRef = ChildRef(u32::MAX);
682
683 #[inline(always)]
684 fn from_node(idx: NodeIndex) -> Self {
685 debug_assert!(idx.0 < CHILD_SEGMENT_TAG);
686 Self(idx.0)
687 }
688
689 #[inline(always)]
690 fn from_segment(idx: SegmentIndex) -> Self {
691 debug_assert!(idx.0 < CHILD_SEGMENT_TAG);
692 Self(CHILD_SEGMENT_TAG | idx.0)
693 }
694
695 #[inline(always)]
696 fn is_none(self) -> bool {
697 self.0 == u32::MAX
698 }
699
700 #[inline(always)]
701 fn is_some(self) -> bool {
702 self.0 != u32::MAX
703 }
704
705 #[inline(always)]
706 fn as_node(self) -> Option<NodeIndex> {
707 if self.is_none() || (self.0 & CHILD_SEGMENT_TAG) != 0 {
708 None
709 } else {
710 Some(NodeIndex(self.0))
711 }
712 }
713
714 #[inline(always)]
715 fn as_segment(self) -> Option<SegmentIndex> {
716 if self.is_none() || (self.0 & CHILD_SEGMENT_TAG) == 0 {
717 None
718 } else {
719 Some(SegmentIndex(self.0 & CHILD_INDEX_MASK))
720 }
721 }
722}
723
724impl Default for ChildRef {
725 fn default() -> Self {
726 Self::NONE
727 }
728}
729
730#[derive(Clone, Copy, Debug, Default)]
731struct SegmentPayload {
732 repr_lo: u32,
733 repr_hi: u32,
734 meta: u32,
735}
736
737impl SegmentPayload {
738 #[inline(always)]
739 fn exact(bits: u64, len: u32) -> Self {
740 debug_assert!(len <= SEG_EXACT_MAX_LEN);
741 debug_assert!(len <= SEG_LEN_MASK);
742 Self {
743 repr_lo: bits as u32,
744 repr_hi: (bits >> 32) as u32,
745 meta: SEG_MODE_EXACT | len,
746 }
747 }
748
749 #[inline(always)]
750 fn history(anchor: u32, len: u32, invert: bool) -> Self {
751 debug_assert!(len <= SEG_LEN_MASK);
752 Self {
753 repr_lo: anchor,
754 repr_hi: 0,
755 meta: if invert {
756 SEG_MODE_HISTORY_INVERT | len
757 } else {
758 SEG_MODE_HISTORY | len
759 },
760 }
761 }
762
763 #[inline(always)]
764 fn constant(bit: bool, len: u32) -> Self {
765 debug_assert!(len <= SEG_LEN_MASK);
766 Self {
767 repr_lo: bit as u32,
768 repr_hi: 0,
769 meta: SEG_MODE_CONST | len,
770 }
771 }
772
773 #[inline(always)]
774 fn len(self) -> u32 {
775 self.meta & SEG_LEN_MASK
776 }
777
778 #[inline(always)]
779 fn set_len(&mut self, len: u32) {
780 debug_assert!(len <= SEG_LEN_MASK);
781 self.meta = (self.meta & SEG_META_MODE_MASK) | len;
782 }
783
784 #[inline(always)]
785 fn mode(self) -> u32 {
786 self.meta & SEG_META_MODE_MASK
787 }
788
789 #[inline(always)]
790 fn is_exact(self) -> bool {
791 self.mode() == SEG_MODE_EXACT
792 }
793
794 #[inline(always)]
795 fn exact_bits(self) -> u64 {
796 (self.repr_lo as u64) | ((self.repr_hi as u64) << 32)
797 }
798
799 #[inline(always)]
800 fn anchor_or_const(self) -> u32 {
801 self.repr_lo
802 }
803
804 #[inline(always)]
805 fn const_bit(self) -> bool {
806 (self.repr_lo & 1) != 0
807 }
808
809 #[inline(always)]
810 fn prepend_exact(self, edge: usize) -> Option<Self> {
811 if !self.is_exact() || self.len() >= SEG_EXACT_MAX_LEN {
812 return None;
813 }
814 let len = self.len() + 1;
815 let bits = ((edge as u64) & 1) | (self.exact_bits() << 1);
816 Some(Self::exact(bits, len))
817 }
818
819 #[inline(always)]
820 fn prefix(self, len: u32) -> Self {
821 debug_assert!(len <= self.len());
822 match self.mode() {
823 SEG_MODE_EXACT => Self::exact(self.exact_bits() & low_bits_mask_u64(len), len),
824 SEG_MODE_HISTORY | SEG_MODE_HISTORY_INVERT => Self {
825 meta: (self.meta & SEG_META_MODE_MASK) | len,
826 ..self
827 },
828 SEG_MODE_CONST => Self::constant(self.const_bit(), len),
829 _ => unreachable!("invalid ctw segment payload mode"),
830 }
831 }
832
833 #[inline(always)]
834 fn suffix_after(self, skip: u32) -> Self {
835 debug_assert!(skip <= self.len());
836 let new_len = self.len() - skip;
837 match self.mode() {
838 SEG_MODE_EXACT => Self::exact(self.exact_bits() >> skip, new_len),
839 SEG_MODE_HISTORY | SEG_MODE_HISTORY_INVERT => Self {
840 repr_lo: self
841 .anchor_or_const()
842 .checked_sub(skip)
843 .expect("ctw history segment anchor underflow"),
844 meta: (self.meta & SEG_META_MODE_MASK) | new_len,
845 ..self
846 },
847 SEG_MODE_CONST => Self::constant(self.const_bit(), new_len),
848 _ => unreachable!("invalid ctw segment payload mode"),
849 }
850 }
851
852 #[inline(always)]
853 fn from_path<H: HistoryAccess + ?Sized>(history: &H, depth: usize, len: u32) -> Option<Self> {
854 if len > SEG_EXACT_MAX_LEN {
855 return None;
856 }
857 Some(Self::exact(
858 path_bits_from_history(history, depth, len as usize),
859 len,
860 ))
861 }
862}
863
864#[derive(Clone, Copy, Debug, Default)]
865struct LevelState {
866 symbol_count: [u32; 2],
867 log_prob_kt: f64,
868 sibling: ChildRef,
869}
870
871#[cfg(test)]
872#[allow(dead_code)]
873#[derive(Clone, Copy, Debug)]
874struct PredictEntry {
875 symbol_count: [u32; 2],
876 log_prob_kt: f64,
877 log_prob_weighted: f64,
878 sibling_weight: f64,
879 has_sibling: bool,
880}
881
882#[derive(Clone, Copy, Debug)]
883enum Detach {
884 NodeChild { node: NodeIndex, edge: usize },
885 SegmentNext { segment: SegmentIndex, new_len: u32 },
886}
887
888#[derive(Clone, Copy, Debug, PartialEq, Eq)]
889enum ExistingSource {
890 None,
891 Node(NodeIndex),
892 Segment(SegmentIndex, u32),
893}
894
895#[derive(Clone, Copy, Debug, PartialEq, Eq)]
896enum PreparedEnd {
897 MaxDepth,
898 MissingAtRoot,
899 MissingAfterCurrent,
900 MismatchAtCurrentSegment,
901}
902
903#[derive(Clone, Copy, Debug, PartialEq)]
904struct PreparedStep {
905 source: ExistingSource,
906 span: u32,
907 sibling_weight: f64,
908 has_sibling: u8,
909}
910
911#[derive(Clone, Copy, Debug)]
912pub struct CtNode {
914 children: [ChildRef; 2],
915 log_prob_kt: f64,
916 log_prob_weighted: f64,
917 symbol_count: [u32; 2],
918}
919
920#[derive(Clone, Copy, Debug)]
921struct CtSegment {
922 tail: ChildRef,
923 log_prob_kt: f64,
924 head_log_prob_weighted: f64,
925 symbol_count: [u32; 2],
926 payload: SegmentPayload,
927}
928
929impl Default for CtSegment {
930 fn default() -> Self {
931 Self {
932 tail: ChildRef::NONE,
933 log_prob_kt: 0.0,
934 head_log_prob_weighted: 0.0,
935 symbol_count: [0, 0],
936 payload: SegmentPayload::default(),
937 }
938 }
939}
940
941impl CtSegment {
942 #[inline(always)]
943 fn len(self) -> u32 {
944 self.payload.len()
945 }
946
947 #[inline(always)]
948 fn set_len(&mut self, len: u32) {
949 self.payload.set_len(len);
950 }
951}
952
953#[inline(always)]
954fn low_bits_mask_u64(len: u32) -> u64 {
955 if len >= 64 {
956 u64::MAX
957 } else {
958 (1u64 << len) - 1
959 }
960}
961
962#[inline(always)]
963fn path_bits_from_history<H: HistoryAccess + ?Sized>(history: &H, depth: usize, len: usize) -> u64 {
964 if let Some(bits) = history.recent_path_bits(depth, len) {
965 return bits;
966 }
967 let history_len = history.len();
968 let available = history_len.saturating_sub(depth).min(len);
969 if available == 0 {
970 return 0;
971 }
972
973 let mut bits = 0u64;
974 let mut hist_idx = history_len - depth - 1;
975 for offset in 0..available {
976 bits |= (history.bit(hist_idx) as u64) << offset;
977 if hist_idx == 0 {
978 break;
979 }
980 hist_idx -= 1;
981 }
982 bits
983}
984
985#[inline(always)]
986fn shift_path_bits(path_bits: u64, consumed: usize) -> u64 {
987 if consumed >= 64 {
988 0
989 } else {
990 path_bits >> consumed
991 }
992}
993
994#[inline(always)]
995fn first_exact_segment_mismatch(
996 exact_bits: u64,
997 path_bits: u64,
998 comparable_len: usize,
999) -> Option<(usize, bool, bool)> {
1000 if comparable_len == 0 {
1001 return None;
1002 }
1003
1004 let diff = (exact_bits ^ path_bits) & low_bits_mask_u64(comparable_len as u32);
1005 if diff == 0 {
1006 None
1007 } else {
1008 let offset = diff.trailing_zeros() as usize;
1009 Some((
1010 offset,
1011 ((path_bits >> offset) & 1) != 0,
1012 ((exact_bits >> offset) & 1) != 0,
1013 ))
1014 }
1015}
1016
1017#[inline(always)]
1018fn predict_ratio_kt(counts: [u32; 2], sym_idx: usize) -> f64 {
1019 let total = (counts[0] + counts[1]) as f64;
1020 let sym_count = counts[sym_idx] as f64;
1021 (sym_count + 0.5) / (total + 1.0)
1022}
1023
1024#[inline(always)]
1025fn predict_ratio_kt_one(counts: [u32; 2]) -> f64 {
1026 let total = (counts[0] + counts[1]) as f64;
1027 let sym_count = counts[1] as f64;
1028 (sym_count + 0.5) / (total + 1.0)
1029}
1030
1031#[inline(always)]
1032fn update_weighted_log_prob_non_leaf(kt_log_prob: f64, log_prob_w0: f64, log_prob_w1: f64) -> f64 {
1033 let child_log_prob = log_prob_w0 + log_prob_w1;
1034 if child_log_prob.to_bits() == kt_log_prob.to_bits() {
1035 return clamp_log_prob(kt_log_prob);
1036 }
1037 let delta = child_log_prob - kt_log_prob;
1038 let log_prob_weighted = if delta >= 0.0 {
1039 child_log_prob + (-delta).exp().ln_1p() - std::f64::consts::LN_2
1040 } else {
1041 kt_log_prob + delta.exp().ln_1p() - std::f64::consts::LN_2
1042 };
1043 clamp_log_prob(log_prob_weighted)
1044}
1045
1046#[inline(always)]
1047fn update_weighted_log_prob(
1048 kt_log_prob: f64,
1049 log_prob_w0: f64,
1050 log_prob_w1: f64,
1051 is_leaf: bool,
1052) -> f64 {
1053 if is_leaf {
1054 clamp_log_prob(kt_log_prob)
1055 } else {
1056 update_weighted_log_prob_non_leaf(kt_log_prob, log_prob_w0, log_prob_w1)
1057 }
1058}
1059
1060#[inline(always)]
1061fn clamp_log_prob(log_prob: f64) -> f64 {
1062 if log_prob > 1.0e-10 { 0.0 } else { log_prob }
1063}
1064
1065#[inline(always)]
1066fn logsumexp_pair(lhs: f64, rhs: f64) -> f64 {
1067 if lhs == f64::NEG_INFINITY {
1068 return rhs;
1069 }
1070 if rhs == f64::NEG_INFINITY {
1071 return lhs;
1072 }
1073 let pivot = lhs.max(rhs);
1074 pivot + ((lhs - pivot).exp() + (rhs - pivot).exp()).ln()
1075}
1076
1077#[inline(always)]
1078fn unary_chain_log_weight(kt_log_prob: f64, continuation_log_prob: f64, len: u32) -> f64 {
1079 debug_assert!(len > 0);
1080 if kt_log_prob.to_bits() == continuation_log_prob.to_bits() {
1081 return kt_log_prob;
1082 }
1083 let log_alpha = -(len as f64) * std::f64::consts::LN_2;
1084 let alpha = log_alpha.exp();
1085 let log_kt_mass = kt_log_prob + (-alpha).ln_1p();
1086 let log_cont_mass = continuation_log_prob + log_alpha;
1087 clamp_log_prob(logsumexp_pair(log_kt_mass, log_cont_mass))
1088}
1089
1090#[inline(always)]
1091fn combined_weight_ratio_internal(
1092 kt_log_prob: f64,
1093 counts: [u32; 2],
1094 path_child_log_prob: f64,
1095 sibling_log_prob: f64,
1096 child_ratio: f64,
1097 sym_idx: usize,
1098) -> (f64, f64) {
1099 let kt_ratio = predict_ratio_kt(counts, sym_idx);
1100 let child_log_prob = path_child_log_prob + sibling_log_prob;
1101 if child_log_prob.to_bits() == kt_log_prob.to_bits() {
1102 return (clamp_log_prob(kt_log_prob), 0.5 * (kt_ratio + child_ratio));
1103 }
1104 let delta = child_log_prob - kt_log_prob;
1105 if delta >= 0.0 {
1106 let x = (-delta).exp();
1107 (
1108 clamp_log_prob(child_log_prob + x.ln_1p() - std::f64::consts::LN_2),
1109 (kt_ratio * x + child_ratio) / (1.0 + x),
1110 )
1111 } else {
1112 let x = delta.exp();
1113 (
1114 clamp_log_prob(kt_log_prob + x.ln_1p() - std::f64::consts::LN_2),
1115 (kt_ratio + x * child_ratio) / (1.0 + x),
1116 )
1117 }
1118}
1119
1120#[inline(always)]
1121fn combined_weight_ratio_internal_one(
1122 kt_log_prob: f64,
1123 counts: [u32; 2],
1124 path_child_log_prob: f64,
1125 sibling_log_prob: f64,
1126 child_ratio: f64,
1127) -> (f64, f64) {
1128 let kt_ratio = predict_ratio_kt_one(counts);
1129 let child_log_prob = path_child_log_prob + sibling_log_prob;
1130 if child_log_prob.to_bits() == kt_log_prob.to_bits() {
1131 return (clamp_log_prob(kt_log_prob), 0.5 * (kt_ratio + child_ratio));
1132 }
1133 let delta = child_log_prob - kt_log_prob;
1134 if delta >= 0.0 {
1135 let x = (-delta).exp();
1136 (
1137 clamp_log_prob(child_log_prob + x.ln_1p() - std::f64::consts::LN_2),
1138 (kt_ratio * x + child_ratio) / (1.0 + x),
1139 )
1140 } else {
1141 let x = delta.exp();
1142 (
1143 clamp_log_prob(kt_log_prob + x.ln_1p() - std::f64::consts::LN_2),
1144 (kt_ratio + x * child_ratio) / (1.0 + x),
1145 )
1146 }
1147}
1148
1149#[inline(always)]
1150fn unary_chain_log_weight_precomputed(
1151 kt_log_prob: f64,
1152 continuation_log_prob: f64,
1153 alpha: f64,
1154 log_alpha: f64,
1155 log_one_minus_alpha: f64,
1156) -> f64 {
1157 if kt_log_prob.to_bits() == continuation_log_prob.to_bits() {
1158 return clamp_log_prob(kt_log_prob);
1159 }
1160
1161 let delta = continuation_log_prob - kt_log_prob;
1162 let log_prob_weighted = if delta >= 0.0 {
1163 let x = ((1.0 - alpha) / alpha) * (-delta).exp();
1164 continuation_log_prob + log_alpha + x.ln_1p()
1165 } else {
1166 let x = (alpha / (1.0 - alpha)) * delta.exp();
1167 kt_log_prob + log_one_minus_alpha + x.ln_1p()
1168 };
1169 clamp_log_prob(log_prob_weighted)
1170}
1171
1172#[inline(always)]
1173#[allow(clippy::too_many_arguments)]
1177fn unary_chain_ratio_transform_precomputed(
1178 kt_log_prob: f64,
1179 counts: [u32; 2],
1180 continuation_log_prob: f64,
1181 continuation_ratio: f64,
1182 alpha: f64,
1183 log_alpha: f64,
1184 log_one_minus_alpha: f64,
1185 sym_idx: usize,
1186) -> (f64, f64) {
1187 let kt_ratio = predict_ratio_kt(counts, sym_idx);
1188 if kt_log_prob.to_bits() == continuation_log_prob.to_bits()
1189 && kt_ratio.to_bits() == continuation_ratio.to_bits()
1190 {
1191 return (clamp_log_prob(kt_log_prob), kt_ratio);
1192 }
1193 if kt_log_prob.to_bits() == continuation_log_prob.to_bits() {
1194 return (
1195 clamp_log_prob(kt_log_prob),
1196 (1.0 - alpha) * kt_ratio + alpha * continuation_ratio,
1197 );
1198 }
1199
1200 let delta = continuation_log_prob - kt_log_prob;
1201 if delta >= 0.0 {
1202 let x = ((1.0 - alpha) / alpha) * (-delta).exp();
1203 (
1204 clamp_log_prob(continuation_log_prob + log_alpha + x.ln_1p()),
1205 (kt_ratio * x + continuation_ratio) / (1.0 + x),
1206 )
1207 } else {
1208 let x = (alpha / (1.0 - alpha)) * delta.exp();
1209 (
1210 clamp_log_prob(kt_log_prob + log_one_minus_alpha + x.ln_1p()),
1211 (kt_ratio + x * continuation_ratio) / (1.0 + x),
1212 )
1213 }
1214}
1215
1216#[inline(always)]
1217fn unary_chain_ratio_transform_precomputed_one(
1218 kt_log_prob: f64,
1219 counts: [u32; 2],
1220 continuation_log_prob: f64,
1221 continuation_ratio: f64,
1222 alpha: f64,
1223 log_alpha: f64,
1224 log_one_minus_alpha: f64,
1225) -> (f64, f64) {
1226 let kt_ratio = predict_ratio_kt_one(counts);
1227 if kt_log_prob.to_bits() == continuation_log_prob.to_bits()
1228 && kt_ratio.to_bits() == continuation_ratio.to_bits()
1229 {
1230 return (clamp_log_prob(kt_log_prob), kt_ratio);
1231 }
1232 if kt_log_prob.to_bits() == continuation_log_prob.to_bits() {
1233 return (
1234 clamp_log_prob(kt_log_prob),
1235 (1.0 - alpha) * kt_ratio + alpha * continuation_ratio,
1236 );
1237 }
1238
1239 let delta = continuation_log_prob - kt_log_prob;
1240 if delta >= 0.0 {
1241 let x = ((1.0 - alpha) / alpha) * (-delta).exp();
1242 (
1243 clamp_log_prob(continuation_log_prob + log_alpha + x.ln_1p()),
1244 (kt_ratio * x + continuation_ratio) / (1.0 + x),
1245 )
1246 } else {
1247 let x = (alpha / (1.0 - alpha)) * delta.exp();
1248 (
1249 clamp_log_prob(kt_log_prob + log_one_minus_alpha + x.ln_1p()),
1250 (kt_ratio + x * continuation_ratio) / (1.0 + x),
1251 )
1252 }
1253}
1254
1255#[inline(always)]
1256fn predict_ratio_internal(
1257 kt_log_prob: f64,
1258 counts: [u32; 2],
1259 path_child_log_prob: f64,
1260 sibling_log_prob: f64,
1261 child_ratio: f64,
1262 sym_idx: usize,
1263) -> f64 {
1264 let kt_ratio = predict_ratio_kt(counts, sym_idx);
1265 let child_log_prob = path_child_log_prob + sibling_log_prob;
1266 if child_log_prob.to_bits() == kt_log_prob.to_bits() {
1267 return 0.5 * (kt_ratio + child_ratio);
1268 }
1269 let delta = child_log_prob - kt_log_prob;
1270 if delta >= 0.0 {
1271 let inv_rho = (-delta).exp();
1272 (kt_ratio * inv_rho + child_ratio) / (1.0 + inv_rho)
1273 } else {
1274 let rho = delta.exp();
1275 (kt_ratio + rho * child_ratio) / (1.0 + rho)
1276 }
1277}
1278
1279#[inline(always)]
1280fn predict_ratio_internal_one(
1281 kt_log_prob: f64,
1282 counts: [u32; 2],
1283 path_child_log_prob: f64,
1284 sibling_log_prob: f64,
1285 child_ratio: f64,
1286) -> f64 {
1287 let kt_ratio = predict_ratio_kt_one(counts);
1288 let child_log_prob = path_child_log_prob + sibling_log_prob;
1289 if child_log_prob.to_bits() == kt_log_prob.to_bits() {
1290 return 0.5 * (kt_ratio + child_ratio);
1291 }
1292 let delta = child_log_prob - kt_log_prob;
1293 if delta >= 0.0 {
1294 let inv_rho = (-delta).exp();
1295 (kt_ratio * inv_rho + child_ratio) / (1.0 + inv_rho)
1296 } else {
1297 let rho = delta.exp();
1298 (kt_ratio + rho * child_ratio) / (1.0 + rho)
1299 }
1300}
1301
1302#[inline(always)]
1303fn path_edge_at_depth<H: HistoryAccess + ?Sized>(
1304 history: &H,
1305 history_len: usize,
1306 depth: usize,
1307) -> bool {
1308 if depth < history_len {
1309 history.bit(history_len - depth - 1)
1310 } else {
1311 false
1312 }
1313}
1314
1315#[inline(always)]
1316fn segment_edge_from_parts<H: HistoryAccess + ?Sized>(
1317 segment: CtSegment,
1318 offset: usize,
1319 history: &H,
1320 history_len: usize,
1321) -> bool {
1322 match segment.payload.mode() {
1323 SEG_MODE_EXACT => ((segment.payload.exact_bits() >> offset) & 1) != 0,
1324 SEG_MODE_HISTORY | SEG_MODE_HISTORY_INVERT => {
1325 if segment.payload.anchor_or_const() as usize >= offset {
1326 let hist_idx = segment.payload.anchor_or_const() as usize - offset;
1327 if hist_idx < history_len {
1328 let raw = history.bit(hist_idx);
1329 if segment.payload.mode() == SEG_MODE_HISTORY_INVERT {
1330 !raw
1331 } else {
1332 raw
1333 }
1334 } else {
1335 segment.payload.mode() == SEG_MODE_HISTORY_INVERT
1336 }
1337 } else {
1338 segment.payload.mode() == SEG_MODE_HISTORY_INVERT
1339 }
1340 }
1341 SEG_MODE_CONST => segment.payload.const_bit(),
1342 _ => unreachable!("invalid ctw segment payload mode"),
1343 }
1344}
1345
1346#[inline(always)]
1347fn first_segment_mismatch<H: HistoryAccess + ?Sized>(
1348 segment: CtSegment,
1349 depth: usize,
1350 history: &H,
1351 comparable_len: usize,
1352) -> Option<(usize, bool, bool)> {
1353 if comparable_len == 0 {
1354 return None;
1355 }
1356
1357 match segment.payload.mode() {
1358 SEG_MODE_EXACT => first_exact_segment_mismatch(
1359 segment.payload.exact_bits(),
1360 path_bits_from_history(history, depth, comparable_len),
1361 comparable_len,
1362 ),
1363 SEG_MODE_HISTORY | SEG_MODE_HISTORY_INVERT => {
1364 let history_len = history.len() as isize;
1365 let mut path_hist_idx = history_len - depth as isize - 1;
1366 let mut seg_hist_idx = segment.payload.anchor_or_const() as isize;
1367 let invert = segment.payload.mode() == SEG_MODE_HISTORY_INVERT;
1368 for offset in 0..comparable_len {
1369 let path_edge = history_at_or_zero(history, history_len, path_hist_idx);
1370 let existing_raw = history_at_or_zero(history, history_len, seg_hist_idx);
1371 let existing_edge = if invert { !existing_raw } else { existing_raw };
1372 if existing_edge != path_edge {
1373 return Some((offset, path_edge, existing_edge));
1374 }
1375 path_hist_idx -= 1;
1376 seg_hist_idx -= 1;
1377 }
1378 None
1379 }
1380 SEG_MODE_CONST => {
1381 let history_len = history.len() as isize;
1382 let mut path_hist_idx = history_len - depth as isize - 1;
1383 let existing_edge = segment.payload.const_bit();
1384 for offset in 0..comparable_len {
1385 let path_edge = history_at_or_zero(history, history_len, path_hist_idx);
1386 if existing_edge != path_edge {
1387 return Some((offset, path_edge, existing_edge));
1388 }
1389 path_hist_idx -= 1;
1390 }
1391 None
1392 }
1393 _ => unreachable!("invalid ctw segment payload mode"),
1394 }
1395}
1396
1397#[inline]
1398fn apply_update_to_state_raw<L: CtLogAccess>(
1399 logs: L,
1400 symbol_count: &mut [u32; 2],
1401 log_prob_kt: &mut f64,
1402 sym_idx: usize,
1403) {
1404 let total_before = (symbol_count[0] + symbol_count[1]) as usize;
1405 let sym_before = symbol_count[sym_idx] as usize;
1406 debug_assert!(sym_before <= total_before);
1407 let log_half_before = logs.log_half(sym_before);
1408 let log_total_after = logs.log_int(total_before + 1);
1409 *log_prob_kt += log_half_before - log_total_after;
1410 if *log_prob_kt > 1.0e-10 {
1411 *log_prob_kt = 0.0;
1412 }
1413 symbol_count[sym_idx] = symbol_count[sym_idx]
1414 .checked_add(1)
1415 .expect("ctw symbol count overflow");
1416}
1417
1418#[inline]
1419fn apply_revert_to_state_raw<L: CtLogAccess>(
1420 logs: L,
1421 symbol_count: &mut [u32; 2],
1422 log_prob_kt: &mut f64,
1423 sym_idx: usize,
1424) {
1425 let total = (symbol_count[0] + symbol_count[1]) as usize;
1426 let sym_count = symbol_count[sym_idx] as usize;
1427 if sym_count > 0 && total > 0 {
1428 let log_half_before = logs.log_half(sym_count - 1);
1429 let log_total = logs.log_int(total);
1430 *log_prob_kt -= log_half_before - log_total;
1431 symbol_count[sym_idx] -= 1;
1432 }
1433 if *log_prob_kt > 1.0e-10 {
1434 *log_prob_kt = 0.0;
1435 }
1436}
1437
1438#[derive(Clone, Debug)]
1439pub struct CtArena {
1444 nodes: Vec<CtNode>,
1445 segments: Vec<CtSegment>,
1446 free_nodes: Vec<NodeIndex>,
1447 free_segments: Vec<SegmentIndex>,
1448}
1449
1450impl CtArena {
1451 pub fn new() -> Self {
1453 Self {
1454 nodes: Vec::with_capacity(1024),
1455 segments: Vec::with_capacity(1024),
1456 free_nodes: Vec::new(),
1457 free_segments: Vec::new(),
1458 }
1459 }
1460
1461 pub fn with_capacity(cap: usize) -> Self {
1463 Self {
1464 nodes: Vec::with_capacity(cap),
1465 segments: Vec::with_capacity(cap / 4 + 1),
1466 free_nodes: Vec::new(),
1467 free_segments: Vec::new(),
1468 }
1469 }
1470
1471 #[inline]
1472 pub fn reserve_exact(&mut self, additional: usize) {
1474 self.nodes.reserve_exact(additional);
1475 self.segments.reserve_exact(additional / 4 + 1);
1476 }
1477
1478 #[inline(always)]
1479 fn reset_node_slot(&mut self, idx: NodeIndex) {
1480 self.nodes[idx.get()] = CtNode {
1481 children: [ChildRef::NONE, ChildRef::NONE],
1482 log_prob_kt: 0.0,
1483 log_prob_weighted: 0.0,
1484 symbol_count: [0, 0],
1485 };
1486 }
1487
1488 #[inline(always)]
1489 fn reset_segment_slot(&mut self, idx: SegmentIndex) {
1490 self.segments[idx.get()] = CtSegment::default();
1491 }
1492
1493 #[inline(always)]
1494 fn alloc_node(&mut self) -> NodeIndex {
1495 if let Some(idx) = self.free_nodes.pop() {
1496 self.reset_node_slot(idx);
1497 idx
1498 } else {
1499 let idx = NodeIndex::from_usize(self.nodes.len());
1500 self.nodes.push(CtNode {
1501 children: [ChildRef::NONE, ChildRef::NONE],
1502 log_prob_kt: 0.0,
1503 log_prob_weighted: 0.0,
1504 symbol_count: [0, 0],
1505 });
1506 idx
1507 }
1508 }
1509
1510 #[inline(always)]
1511 fn alloc_node_with_state(&mut self, symbol_count: [u32; 2], log_prob_kt: f64) -> NodeIndex {
1512 let idx = self.alloc_node();
1513 self.nodes[idx.get()].symbol_count = symbol_count;
1514 self.nodes[idx.get()].log_prob_kt = log_prob_kt;
1515 idx
1516 }
1517
1518 #[inline(always)]
1519 fn free_node(&mut self, idx: NodeIndex) {
1520 self.free_nodes.push(idx);
1521 }
1522
1523 #[inline(always)]
1524 fn alloc_segment(&mut self) -> SegmentIndex {
1525 if let Some(idx) = self.free_segments.pop() {
1526 self.reset_segment_slot(idx);
1527 idx
1528 } else {
1529 let idx = SegmentIndex::from_usize(self.segments.len());
1530 self.segments.push(CtSegment::default());
1531 idx
1532 }
1533 }
1534
1535 #[inline(always)]
1536 fn free_segment(&mut self, idx: SegmentIndex) {
1537 self.reset_segment_slot(idx);
1538 self.free_segments.push(idx);
1539 }
1540
1541 pub fn clear(&mut self) {
1543 self.nodes.clear();
1544 self.segments.clear();
1545 self.free_nodes.clear();
1546 self.free_segments.clear();
1547 }
1548
1549 #[inline(always)]
1550 fn child(&self, parent_idx: NodeIndex, child_idx: usize) -> ChildRef {
1551 debug_assert!(parent_idx.get() < self.nodes.len());
1552 debug_assert!(child_idx < 2);
1553 unsafe {
1554 *self
1555 .nodes
1556 .get_unchecked(parent_idx.get())
1557 .children
1558 .get_unchecked(child_idx)
1559 }
1560 }
1561
1562 #[inline(always)]
1563 fn set_child(&mut self, parent_idx: NodeIndex, child_idx: usize, child: ChildRef) {
1564 debug_assert!(parent_idx.get() < self.nodes.len());
1565 debug_assert!(child_idx < 2);
1566 unsafe {
1567 *self
1568 .nodes
1569 .get_unchecked_mut(parent_idx.get())
1570 .children
1571 .get_unchecked_mut(child_idx) = child;
1572 }
1573 }
1574
1575 #[inline(always)]
1576 fn set_segment_tail(&mut self, segment_idx: SegmentIndex, child: ChildRef) {
1577 self.segments[segment_idx.get()].tail = child;
1578 }
1579
1580 #[inline(always)]
1581 fn counts(&self, idx: NodeIndex) -> [u32; 2] {
1582 self.nodes[idx.get()].symbol_count
1583 }
1584
1585 #[inline(always)]
1586 fn visits(&self, idx: NodeIndex) -> u32 {
1587 let counts = self.nodes[idx.get()].symbol_count;
1588 counts[0] + counts[1]
1589 }
1590
1591 #[inline(always)]
1592 fn segment_symbol_count(&self, segment_idx: SegmentIndex) -> [u32; 2] {
1593 self.segments[segment_idx.get()].symbol_count
1594 }
1595
1596 #[inline(always)]
1597 fn segment_log_prob_kt(&self, segment_idx: SegmentIndex) -> f64 {
1598 self.segments[segment_idx.get()].log_prob_kt
1599 }
1600
1601 #[inline(always)]
1602 fn segment_len(&self, segment_idx: SegmentIndex) -> u32 {
1603 self.segments[segment_idx.get()].len()
1604 }
1605
1606 #[inline(always)]
1607 fn segment_has_child(&self, segment_idx: SegmentIndex, offset: u32) -> bool {
1608 let segment = self.segments[segment_idx.get()];
1609 offset + 1 < segment.len() || segment.tail.is_some()
1610 }
1611
1612 #[inline(always)]
1613 fn log_prob_weighted(&self, idx: NodeIndex) -> f64 {
1614 self.nodes[idx.get()].log_prob_weighted
1615 }
1616
1617 #[inline(always)]
1618 fn log_prob_kt(&self, idx: NodeIndex) -> f64 {
1619 self.nodes[idx.get()].log_prob_kt
1620 }
1621
1622 #[inline(always)]
1623 unsafe fn child_ref_weighted_unchecked(&self, child: ChildRef) -> f64 {
1624 if child.is_none() {
1625 return 0.0;
1626 }
1627
1628 let raw = child.0;
1629 if (raw & CHILD_SEGMENT_TAG) == 0 {
1630 debug_assert!((raw as usize) < self.nodes.len());
1631 self.nodes.get_unchecked(raw as usize).log_prob_weighted
1632 } else {
1633 let idx = (raw & CHILD_INDEX_MASK) as usize;
1634 debug_assert!(idx < self.segments.len());
1635 self.segments.get_unchecked(idx).head_log_prob_weighted
1636 }
1637 }
1638
1639 #[inline(always)]
1640 fn child_ref_weighted(&self, child: ChildRef) -> f64 {
1641 unsafe { self.child_ref_weighted_unchecked(child) }
1642 }
1643
1644 #[inline(always)]
1645 fn singleton_segment_payload(&self, edge: usize) -> SegmentPayload {
1646 SegmentPayload::exact((edge & 1) as u64, 1)
1647 }
1648
1649 #[inline(always)]
1650 fn segment_edge(
1651 &self,
1652 segment_idx: SegmentIndex,
1653 offset: u32,
1654 history: &(impl HistoryAccess + ?Sized),
1655 ) -> usize {
1656 let segment = self.segments[segment_idx.get()];
1657 segment_edge_from_parts(segment, offset as usize, history, history.len()) as usize
1658 }
1659
1660 fn segment_suffix_weight(&self, segment_idx: SegmentIndex, offset: u32) -> f64 {
1661 let segment = self.segments[segment_idx.get()];
1662 if offset >= segment.len() {
1663 return self.child_ref_weighted(segment.tail);
1664 }
1665 if segment.tail.is_none() {
1666 return segment.log_prob_kt;
1667 }
1668 let remaining = segment.len() - offset;
1669 unary_chain_log_weight(
1670 segment.log_prob_kt,
1671 self.child_ref_weighted(segment.tail),
1672 remaining,
1673 )
1674 }
1675
1676 #[inline(always)]
1677 fn segment_continuation_weight(&self, segment_idx: SegmentIndex, offset: u32) -> f64 {
1678 let segment = self.segments[segment_idx.get()];
1679 if offset + 1 < segment.len() {
1680 self.segment_suffix_weight(segment_idx, offset + 1)
1681 } else {
1682 self.child_ref_weighted(segment.tail)
1683 }
1684 }
1685
1686 fn recompute_segment_head(&mut self, segment_idx: SegmentIndex) {
1687 let segment = self.segments[segment_idx.get()];
1688 let head = if segment.tail.is_some() {
1689 unary_chain_log_weight(
1690 segment.log_prob_kt,
1691 self.child_ref_weighted(segment.tail),
1692 segment.len(),
1693 )
1694 } else {
1695 segment.log_prob_kt
1696 };
1697 self.segments[segment_idx.get()].head_log_prob_weighted = head;
1698 }
1699
1700 fn recompute_node_weight(&mut self, idx: NodeIndex) {
1701 let slot = idx.get();
1702 debug_assert!(slot < self.nodes.len());
1703 let node = unsafe { *self.nodes.get_unchecked(slot) };
1704 let [left, right] = node.children;
1705 let weighted = if left.is_none() && right.is_none() {
1706 clamp_log_prob(node.log_prob_kt)
1707 } else {
1708 let w0 = unsafe { self.child_ref_weighted_unchecked(left) };
1709 let w1 = unsafe { self.child_ref_weighted_unchecked(right) };
1710 update_weighted_log_prob_non_leaf(node.log_prob_kt, w0, w1)
1711 };
1712 unsafe {
1713 self.nodes.get_unchecked_mut(slot).log_prob_weighted = weighted;
1714 }
1715 }
1716
1717 fn alloc_segment_with_parts(
1718 &mut self,
1719 symbol_count: [u32; 2],
1720 log_prob_kt: f64,
1721 tail: ChildRef,
1722 payload: SegmentPayload,
1723 ) -> SegmentIndex {
1724 let segment_idx = self.alloc_segment();
1725 self.segments[segment_idx.get()] = CtSegment {
1726 tail,
1727 log_prob_kt,
1728 head_log_prob_weighted: 0.0,
1729 symbol_count,
1730 payload,
1731 };
1732 if payload.len() == 1 && tail.is_none() {
1733 self.segments[segment_idx.get()].head_log_prob_weighted = log_prob_kt;
1734 } else {
1735 self.recompute_segment_head(segment_idx);
1736 }
1737 segment_idx
1738 }
1739
1740 fn detach_segment_continuation(
1741 &mut self,
1742 segment_idx: SegmentIndex,
1743 offset: u32,
1744 detaches: &mut Vec<Detach>,
1745 ) -> ChildRef {
1746 let segment = self.segments[segment_idx.get()];
1747 if offset + 1 < segment.len() {
1748 let suffix = self.alloc_segment_with_parts(
1749 segment.symbol_count,
1750 segment.log_prob_kt,
1751 segment.tail,
1752 segment.payload.suffix_after(offset + 1),
1753 );
1754 detaches.push(Detach::SegmentNext {
1755 segment: segment_idx,
1756 new_len: offset + 1,
1757 });
1758 ChildRef::from_segment(suffix)
1759 } else {
1760 let tail = segment.tail;
1761 if tail.is_some() {
1762 detaches.push(Detach::SegmentNext {
1763 segment: segment_idx,
1764 new_len: segment.len(),
1765 });
1766 }
1767 tail
1768 }
1769 }
1770
1771 #[allow(clippy::too_many_arguments)]
1775 fn prepend_or_alloc_segment(
1776 &mut self,
1777 history: &(impl HistoryAccess + ?Sized),
1778 depth: usize,
1779 symbol_count: [u32; 2],
1780 log_prob_kt: f64,
1781 child: ChildRef,
1782 edge: usize,
1783 allow_history_pattern: bool,
1784 ) -> ChildRef {
1785 let singleton_payload = self.singleton_segment_payload(edge);
1786
1787 if let Some(segment_idx) = child.as_segment() {
1788 let segment = self.segments[segment_idx.get()];
1789 let same_state = segment.symbol_count == symbol_count
1790 && segment.log_prob_kt.to_bits() == log_prob_kt.to_bits();
1791 if same_state && segment.tail == child {
1792 let segment = &mut self.segments[segment_idx.get()];
1793 let extended_payload = if segment.payload.is_exact() {
1794 segment.payload.prepend_exact(edge)
1795 } else if allow_history_pattern {
1796 let path_payload =
1797 SegmentPayload::from_path(history, depth, segment.len().saturating_add(1));
1798 path_payload.filter(|payload| {
1799 let mut matches = true;
1800 for offset in 0..segment.len() as usize {
1801 let seg_edge =
1802 segment_edge_from_parts(*segment, offset, history, history.len());
1803 let payload_edge = ((payload.exact_bits() >> (offset + 1)) & 1) != 0;
1804 if seg_edge != payload_edge {
1805 matches = false;
1806 break;
1807 }
1808 }
1809 matches
1810 })
1811 } else {
1812 None
1813 };
1814 if let Some(payload) = extended_payload {
1815 let old_head = segment.head_log_prob_weighted;
1816 segment.payload = payload;
1817 segment.head_log_prob_weighted =
1818 update_weighted_log_prob(log_prob_kt, old_head, 0.0, false);
1819 return ChildRef::from_segment(segment_idx);
1820 }
1821 }
1822 }
1823
1824 let segment_idx =
1825 self.alloc_segment_with_parts(symbol_count, log_prob_kt, child, singleton_payload);
1826 ChildRef::from_segment(segment_idx)
1827 }
1828
1829 fn free_child_ref(&mut self, child: ChildRef) {
1830 let mut stack = Vec::with_capacity(16);
1831 if child.is_some() {
1832 stack.push(child);
1833 }
1834 while let Some(next) = stack.pop() {
1835 if let Some(node_idx) = next.as_node() {
1836 let children = self.nodes[node_idx.get()].children;
1837 if children[0].is_some() {
1838 stack.push(children[0]);
1839 }
1840 if children[1].is_some() {
1841 stack.push(children[1]);
1842 }
1843 self.free_node(node_idx);
1844 } else if let Some(segment_idx) = next.as_segment() {
1845 let tail = self.segments[segment_idx.get()].tail;
1846 if tail.is_some() {
1847 stack.push(tail);
1848 }
1849 self.free_segment(segment_idx);
1850 }
1851 }
1852 }
1853
1854 pub fn memory_usage(&self) -> usize {
1856 self.nodes.capacity() * size_of::<CtNode>()
1857 + self.segments.capacity() * size_of::<CtSegment>()
1858 + self.free_nodes.capacity() * size_of::<NodeIndex>()
1859 + self.free_segments.capacity() * size_of::<SegmentIndex>()
1860 }
1861}
1862
1863impl Default for CtArena {
1864 fn default() -> Self {
1865 Self::new()
1866 }
1867}
1868
1869#[derive(Clone)]
1870struct CtEngine {
1871 arena: CtArena,
1872 root: NodeIndex,
1873 max_depth: usize,
1874 segment_alpha: Vec<f64>,
1875 segment_log_alpha: Vec<f64>,
1876 segment_log_one_minus_alpha: Vec<f64>,
1877 levels: Vec<LevelState>,
1878 detaches: Vec<Detach>,
1879 prepared_steps: Vec<PreparedStep>,
1880 prepared_levels: usize,
1881 prepared_end: PreparedEnd,
1882}
1883
1884impl CtEngine {
1885 const RESERVE_MIN_NODES: usize = 4 * 1024;
1886 const RESERVE_MAX_NODES: usize = 1 << 18;
1887
1888 fn new(depth: usize) -> Self {
1889 let mut arena = CtArena::with_capacity(1024.min(1 << depth.min(16)));
1890 let root = arena.alloc_node();
1891 let mut segment_alpha = Vec::with_capacity(depth + 1);
1892 let mut segment_log_alpha = Vec::with_capacity(depth + 1);
1893 let mut segment_log_one_minus_alpha = Vec::with_capacity(depth + 1);
1894 segment_alpha.push(1.0);
1895 segment_log_alpha.push(0.0);
1896 segment_log_one_minus_alpha.push(f64::NEG_INFINITY);
1897 let mut alpha = 1.0f64;
1898 for len in 1..=depth {
1899 alpha *= 0.5;
1900 segment_alpha.push(alpha);
1901 segment_log_alpha.push(-(len as f64) * std::f64::consts::LN_2);
1902 segment_log_one_minus_alpha.push((-alpha).ln_1p());
1903 }
1904 Self {
1905 arena,
1906 root,
1907 max_depth: depth,
1908 segment_alpha,
1909 segment_log_alpha,
1910 segment_log_one_minus_alpha,
1911 levels: vec![LevelState::default(); depth],
1912 detaches: Vec::with_capacity(depth),
1913 prepared_steps: Vec::with_capacity(depth),
1914 prepared_levels: 0,
1915 prepared_end: PreparedEnd::MaxDepth,
1916 }
1917 }
1918
1919 #[inline(always)]
1920 fn root_visits(&self) -> usize {
1921 self.arena.visits(self.root) as usize
1922 }
1923
1924 #[inline(always)]
1925 fn hot_prefix_depth(&self) -> usize {
1926 self.max_depth.min(ctw_hot_prefix_depth_limit())
1927 }
1928
1929 #[inline(always)]
1930 fn push_prepared_segment_step(
1931 &mut self,
1932 segment_idx: SegmentIndex,
1933 offset: usize,
1934 sibling_weight: f64,
1935 has_sibling: u8,
1936 ) {
1937 let span = (offset + 1) as u32;
1938 self.prepared_steps.push(PreparedStep {
1939 source: ExistingSource::Segment(segment_idx, offset as u32),
1940 span,
1941 sibling_weight,
1942 has_sibling,
1943 });
1944 self.prepared_levels += span as usize;
1945 }
1946
1947 #[inline(always)]
1948 fn walk_prepared_exact_segment(
1949 &mut self,
1950 segment_idx: SegmentIndex,
1951 segment: CtSegment,
1952 depth: usize,
1953 path_bits: u64,
1954 ) -> Option<(usize, ExistingSource)> {
1955 let seg_len = segment.len() as usize;
1956 let terminal_offset = self.max_depth.saturating_sub(depth);
1957 let comparable_len = seg_len.min(terminal_offset);
1958 if let Some((offset, _, _)) =
1959 first_exact_segment_mismatch(segment.payload.exact_bits(), path_bits, comparable_len)
1960 {
1961 self.push_prepared_segment_step(
1962 segment_idx,
1963 offset,
1964 self.arena
1965 .segment_continuation_weight(segment_idx, offset as u32),
1966 1,
1967 );
1968 self.prepared_end = PreparedEnd::MismatchAtCurrentSegment;
1969 return None;
1970 }
1971
1972 if terminal_offset < seg_len {
1973 self.push_prepared_segment_step(segment_idx, terminal_offset, 0.0, 0);
1974 return None;
1975 }
1976
1977 if segment.tail.is_none() {
1978 self.push_prepared_segment_step(segment_idx, seg_len - 1, 0.0, 0);
1979 self.prepared_end = PreparedEnd::MissingAfterCurrent;
1980 return None;
1981 }
1982
1983 self.push_prepared_segment_step(segment_idx, seg_len - 1, 0.0, 0);
1984 let tail = segment.tail;
1985 Some((
1986 depth + seg_len,
1987 Self::child_to_existing_source(tail).unwrap_or(ExistingSource::None),
1988 ))
1989 }
1990
1991 fn clear(&mut self) {
1992 self.arena.clear();
1993 self.root = self.arena.alloc_node();
1994 self.levels.fill(LevelState::default());
1995 self.detaches.clear();
1996 self.prepared_steps.clear();
1997 self.prepared_levels = 0;
1998 self.prepared_end = PreparedEnd::MaxDepth;
1999 }
2000
2001 #[inline]
2002 fn reserve_for_symbols(&mut self, total_symbols: usize) {
2003 if total_symbols == 0 {
2004 return;
2005 }
2006
2007 let depth_scale = self.max_depth.saturating_add(1);
2008 let reserve_nodes = total_symbols
2009 .saturating_div(depth_scale)
2010 .clamp(Self::RESERVE_MIN_NODES, Self::RESERVE_MAX_NODES);
2011 let free_nodes = self
2012 .arena
2013 .nodes
2014 .capacity()
2015 .saturating_sub(self.arena.nodes.len());
2016 if reserve_nodes > free_nodes {
2017 self.arena.reserve_exact(reserve_nodes - free_nodes);
2018 }
2019 }
2020
2021 #[inline]
2022 fn get_log_block_probability(&self) -> f64 {
2023 self.arena.log_prob_weighted(self.root)
2024 }
2025
2026 #[inline]
2027 fn with_cached_logs<R>(
2028 &mut self,
2029 upto: usize,
2030 f: impl FnOnce(&mut Self, CachedLogs<'_>) -> R,
2031 ) -> R {
2032 with_shared_cached_logs(upto, |logs| f(self, logs))
2033 }
2034
2035 #[inline]
2036 fn with_bounded_logs<R>(
2037 &mut self,
2038 upto: usize,
2039 f: impl FnOnce(&mut Self, BoundedLogs<'_>) -> R,
2040 ) -> R {
2041 with_shared_bounded_logs(upto, |logs| f(self, logs))
2042 }
2043
2044 #[inline]
2045 fn log_cache_memory_usage(&self) -> usize {
2046 shared_log_cache_memory_usage()
2047 }
2048
2049 #[inline(always)]
2050 fn segment_constants(&self, len: u32) -> (f64, f64, f64) {
2051 let idx = len as usize;
2052 debug_assert!(idx < self.segment_alpha.len());
2053 debug_assert!(idx < self.segment_log_alpha.len());
2054 debug_assert!(idx < self.segment_log_one_minus_alpha.len());
2055 (
2056 unsafe { *self.segment_alpha.get_unchecked(idx) },
2057 unsafe { *self.segment_log_alpha.get_unchecked(idx) },
2058 unsafe { *self.segment_log_one_minus_alpha.get_unchecked(idx) },
2059 )
2060 }
2061
2062 #[inline(always)]
2063 fn source_counts_and_kt_log_prob(&self, source: ExistingSource) -> ([u32; 2], f64) {
2064 match source {
2065 ExistingSource::Node(node_idx) => {
2066 let slot = node_idx.get();
2067 let node = unsafe { *self.arena.nodes.get_unchecked(slot) };
2068 (node.symbol_count, node.log_prob_kt)
2069 }
2070 ExistingSource::Segment(segment_idx, _) => {
2071 let slot = segment_idx.get();
2072 let segment = unsafe { *self.arena.segments.get_unchecked(slot) };
2073 (segment.symbol_count, segment.log_prob_kt)
2074 }
2075 ExistingSource::None => unreachable!("prepared step should never store None"),
2076 }
2077 }
2078
2079 fn build_missing_segment_path(
2080 &mut self,
2081 depth: usize,
2082 history: &(impl HistoryAccess + ?Sized),
2083 sym_idx: usize,
2084 singleton_log_prob_kt: f64,
2085 ) -> ChildRef {
2086 if depth > self.max_depth {
2087 return ChildRef::NONE;
2088 }
2089
2090 let mut counts = [0u32; 2];
2091 counts[sym_idx] = 1;
2092 let log_prob_kt = singleton_log_prob_kt;
2093 let total_len = self.max_depth - depth + 1;
2094
2095 if let Some(payload) = SegmentPayload::from_path(history, depth, total_len as u32) {
2096 let segment =
2097 self.arena
2098 .alloc_segment_with_parts(counts, log_prob_kt, ChildRef::NONE, payload);
2099 return ChildRef::from_segment(segment);
2100 }
2101
2102 let history_nodes = if depth < history.len() {
2103 (self.max_depth.min(history.len() - 1) - depth) + 1
2104 } else {
2105 0
2106 };
2107 let const_nodes = total_len - history_nodes;
2108
2109 let mut built = ChildRef::NONE;
2110 if const_nodes > 0 {
2111 let const_segment = self.arena.alloc_segment_with_parts(
2112 counts,
2113 log_prob_kt,
2114 ChildRef::NONE,
2115 SegmentPayload::constant(false, const_nodes as u32),
2116 );
2117 built = ChildRef::from_segment(const_segment);
2118 }
2119 if history_nodes > 0 {
2120 let history_segment = self.arena.alloc_segment_with_parts(
2121 counts,
2122 log_prob_kt,
2123 built,
2124 SegmentPayload::history(
2125 (history.len() - depth - 1) as u32,
2126 history_nodes as u32,
2127 false,
2128 ),
2129 );
2130 built = ChildRef::from_segment(history_segment);
2131 }
2132 built
2133 }
2134
2135 fn build_missing_path(
2136 &mut self,
2137 depth: usize,
2138 history: &(impl HistoryAccess + ?Sized),
2139 sym_idx: usize,
2140 singleton_log_prob_kt: f64,
2141 ) -> ChildRef {
2142 if depth > self.max_depth {
2143 return ChildRef::NONE;
2144 }
2145
2146 let hot_prefix_depth = self.hot_prefix_depth();
2147 if depth > hot_prefix_depth {
2148 return self.build_missing_segment_path(depth, history, sym_idx, singleton_log_prob_kt);
2149 }
2150
2151 let mut counts = [0u32; 2];
2152 counts[sym_idx] = 1;
2153 let mut built = if hot_prefix_depth < self.max_depth {
2154 self.build_missing_segment_path(
2155 hot_prefix_depth + 1,
2156 history,
2157 sym_idx,
2158 singleton_log_prob_kt,
2159 )
2160 } else {
2161 ChildRef::NONE
2162 };
2163
2164 for node_depth in (depth..=hot_prefix_depth).rev() {
2165 let node = self
2166 .arena
2167 .alloc_node_with_state(counts, singleton_log_prob_kt);
2168 if node_depth < self.max_depth {
2169 let edge = history_symbol(history, node_depth) as usize;
2170 self.arena.set_child(node, edge, built);
2171 }
2172 self.arena.recompute_node_weight(node);
2173 built = ChildRef::from_node(node);
2174 }
2175 built
2176 }
2177
2178 #[inline(always)]
2179 fn build_missing_segment_path_exact_bits(
2180 &mut self,
2181 depth: usize,
2182 path_bits: u64,
2183 sym_idx: usize,
2184 singleton_log_prob_kt: f64,
2185 ) -> ChildRef {
2186 debug_assert!(self.max_depth <= SEG_EXACT_MAX_LEN as usize);
2187 if depth > self.max_depth {
2188 return ChildRef::NONE;
2189 }
2190
2191 let mut counts = [0u32; 2];
2192 counts[sym_idx] = 1;
2193 let total_len = self.max_depth - depth + 1;
2194 let payload = SegmentPayload::exact(
2195 path_bits & low_bits_mask_u64(total_len as u32),
2196 total_len as u32,
2197 );
2198 let segment = self.arena.alloc_segment_with_parts(
2199 counts,
2200 singleton_log_prob_kt,
2201 ChildRef::NONE,
2202 payload,
2203 );
2204 ChildRef::from_segment(segment)
2205 }
2206
2207 #[inline(always)]
2208 fn build_missing_path_exact_bits(
2209 &mut self,
2210 depth: usize,
2211 path_bits: u64,
2212 sym_idx: usize,
2213 singleton_log_prob_kt: f64,
2214 ) -> ChildRef {
2215 debug_assert!(self.max_depth <= SEG_EXACT_MAX_LEN as usize);
2216 if depth > self.max_depth {
2217 return ChildRef::NONE;
2218 }
2219
2220 let hot_prefix_depth = self.hot_prefix_depth();
2221 if depth > hot_prefix_depth {
2222 return self.build_missing_segment_path_exact_bits(
2223 depth,
2224 path_bits,
2225 sym_idx,
2226 singleton_log_prob_kt,
2227 );
2228 }
2229
2230 let mut counts = [0u32; 2];
2231 counts[sym_idx] = 1;
2232 let mut built = if hot_prefix_depth < self.max_depth {
2233 self.build_missing_segment_path_exact_bits(
2234 hot_prefix_depth + 1,
2235 shift_path_bits(path_bits, hot_prefix_depth + 1 - depth),
2236 sym_idx,
2237 singleton_log_prob_kt,
2238 )
2239 } else {
2240 ChildRef::NONE
2241 };
2242
2243 for node_depth in (depth..=hot_prefix_depth).rev() {
2244 let node = self
2245 .arena
2246 .alloc_node_with_state(counts, singleton_log_prob_kt);
2247 if node_depth < self.max_depth {
2248 let edge = ((path_bits >> (node_depth - depth)) & 1) as usize;
2249 self.arena.set_child(node, edge, built);
2250 }
2251 self.arena.recompute_node_weight(node);
2252 built = ChildRef::from_node(node);
2253 }
2254 built
2255 }
2256
2257 #[inline(always)]
2258 fn child_to_existing_source(child: ChildRef) -> Option<ExistingSource> {
2259 if let Some(node) = child.as_node() {
2260 Some(ExistingSource::Node(node))
2261 } else {
2262 child
2263 .as_segment()
2264 .map(|segment| ExistingSource::Segment(segment, 0))
2265 }
2266 }
2267
2268 #[inline(always)]
2269 fn update_source_state<L: CtLogAccess>(
2270 &mut self,
2271 logs: L,
2272 source: ExistingSource,
2273 sym_idx: usize,
2274 ) {
2275 match source {
2276 ExistingSource::Node(node_idx) => {
2277 let slot = node_idx.get();
2278 let mut counts = self.arena.nodes[slot].symbol_count;
2279 let mut log_prob_kt = self.arena.nodes[slot].log_prob_kt;
2280 apply_update_to_state_raw(logs, &mut counts, &mut log_prob_kt, sym_idx);
2281 self.arena.nodes[slot].symbol_count = counts;
2282 self.arena.nodes[slot].log_prob_kt = log_prob_kt;
2283 }
2284 ExistingSource::Segment(segment_idx, _) => {
2285 let slot = segment_idx.get();
2286 let mut counts = self.arena.segments[slot].symbol_count;
2287 let mut log_prob_kt = self.arena.segments[slot].log_prob_kt;
2288 apply_update_to_state_raw(logs, &mut counts, &mut log_prob_kt, sym_idx);
2289 self.arena.segments[slot].symbol_count = counts;
2290 self.arena.segments[slot].log_prob_kt = log_prob_kt;
2291 }
2292 ExistingSource::None => unreachable!("prepared update should never visit None"),
2293 }
2294 }
2295
2296 #[inline(always)]
2297 fn recompute_source_weight(&mut self, source: ExistingSource) {
2298 match source {
2299 ExistingSource::Node(node_idx) => self.arena.recompute_node_weight(node_idx),
2300 ExistingSource::Segment(segment_idx, _) => self.recompute_segment_head(segment_idx),
2301 ExistingSource::None => unreachable!("prepared update should never visit None"),
2302 }
2303 }
2304
2305 #[inline(always)]
2306 fn recompute_segment_head(&mut self, segment_idx: SegmentIndex) {
2307 let segment = self.arena.segments[segment_idx.get()];
2308 let head = if segment.tail.is_some() {
2309 let (alpha, log_alpha, log_one_minus_alpha) = self.segment_constants(segment.len());
2310 unary_chain_log_weight_precomputed(
2311 segment.log_prob_kt,
2312 self.arena.child_ref_weighted(segment.tail),
2313 alpha,
2314 log_alpha,
2315 log_one_minus_alpha,
2316 )
2317 } else {
2318 segment.log_prob_kt
2319 };
2320 self.arena.segments[segment_idx.get()].head_log_prob_weighted = head;
2321 }
2322
2323 fn attach_missing_after_prepared_path(
2324 &mut self,
2325 history: &(impl HistoryAccess + ?Sized),
2326 sym_idx: usize,
2327 singleton_log_prob_kt: f64,
2328 ) {
2329 let Some(last_step) = self.prepared_steps.last().copied() else {
2330 return;
2331 };
2332 let depth = self.prepared_levels;
2333 match last_step.source {
2334 ExistingSource::Node(node_idx) => {
2335 debug_assert!(depth < self.max_depth);
2336 let path_edge = history_symbol(history, depth) as usize;
2337 debug_assert!(self.arena.child(node_idx, path_edge).is_none());
2338 let new_child =
2339 self.build_missing_path(depth + 1, history, sym_idx, singleton_log_prob_kt);
2340 self.arena.set_child(node_idx, path_edge, new_child);
2341 }
2342 ExistingSource::Segment(segment_idx, offset) => {
2343 debug_assert!(depth < self.max_depth);
2344 debug_assert_eq!(offset + 1, self.arena.segment_len(segment_idx));
2345 debug_assert!(self.arena.segments[segment_idx.get()].tail.is_none());
2346 let new_tail =
2347 self.build_missing_path(depth + 1, history, sym_idx, singleton_log_prob_kt);
2348 self.arena.set_segment_tail(segment_idx, new_tail);
2349 }
2350 ExistingSource::None => unreachable!("prepared path should never end in None source"),
2351 }
2352 }
2353
2354 fn replace_prepared_child(
2355 &mut self,
2356 history: &(impl HistoryAccess + ?Sized),
2357 step_index: usize,
2358 current_start_depth: usize,
2359 new_child: ChildRef,
2360 ) {
2361 if step_index == 0 {
2362 let root_edge = history_symbol(history, 0) as usize;
2363 self.arena.set_child(self.root, root_edge, new_child);
2364 return;
2365 }
2366
2367 match self.prepared_steps[step_index - 1].source {
2368 ExistingSource::Node(node_idx) => {
2369 let edge = history_symbol(history, current_start_depth - 1) as usize;
2370 self.arena.set_child(node_idx, edge, new_child);
2371 }
2372 ExistingSource::Segment(segment_idx, offset) => {
2373 debug_assert_eq!(offset + 1, self.arena.segment_len(segment_idx));
2374 self.arena.set_segment_tail(segment_idx, new_child);
2375 }
2376 ExistingSource::None => unreachable!("prepared path should never parent from None"),
2377 }
2378 }
2379
2380 fn update_prepared_mismatch<L: CtLogAccess>(
2381 &mut self,
2382 logs: L,
2383 history: &(impl HistoryAccess + ?Sized),
2384 sym_idx: usize,
2385 singleton_log_prob_kt: f64,
2386 ) -> ChildRef {
2387 let last_index = self.prepared_steps.len() - 1;
2388 for idx in 0..last_index {
2389 self.update_source_state(logs, self.prepared_steps[idx].source, sym_idx);
2390 }
2391
2392 let last_step = self.prepared_steps[last_index];
2393 let ExistingSource::Segment(segment_idx, offset_u32) = last_step.source else {
2394 unreachable!("prepared segment mismatch must end at a segment");
2395 };
2396
2397 let original = self.arena.segments[segment_idx.get()];
2398 let offset = offset_u32 as usize;
2399 let seg_len = original.len() as usize;
2400 let history_len = history.len();
2401 let current_start_depth = self.prepared_levels - last_step.span as usize + 1;
2402 let node_depth = current_start_depth + offset;
2403 let path_edge = path_edge_at_depth(history, history_len, node_depth);
2404 let existing_edge = segment_edge_from_parts(original, offset, history, history_len);
2405 debug_assert_ne!(path_edge, existing_edge);
2406
2407 let old_continuation = if offset + 1 < seg_len {
2408 if offset == 0 {
2409 let segment = &mut self.arena.segments[segment_idx.get()];
2410 segment.payload = original.payload.suffix_after(1);
2411 segment.tail = original.tail;
2412 segment.symbol_count = original.symbol_count;
2413 segment.log_prob_kt = original.log_prob_kt;
2414 self.recompute_segment_head(segment_idx);
2415 ChildRef::from_segment(segment_idx)
2416 } else {
2417 ChildRef::from_segment(self.arena.alloc_segment_with_parts(
2418 original.symbol_count,
2419 original.log_prob_kt,
2420 original.tail,
2421 original.payload.suffix_after(offset as u32 + 1),
2422 ))
2423 }
2424 } else {
2425 original.tail
2426 };
2427
2428 let new_tail =
2429 self.build_missing_path(node_depth + 1, history, sym_idx, singleton_log_prob_kt);
2430 let mut updated_counts = original.symbol_count;
2431 let mut updated_log_prob_kt = original.log_prob_kt;
2432 apply_update_to_state_raw(logs, &mut updated_counts, &mut updated_log_prob_kt, sym_idx);
2433
2434 let branch = self
2435 .arena
2436 .alloc_node_with_state(updated_counts, updated_log_prob_kt);
2437 self.arena
2438 .set_child(branch, existing_edge as usize, old_continuation);
2439 self.arena.set_child(branch, path_edge as usize, new_tail);
2440 self.arena.recompute_node_weight(branch);
2441
2442 if offset == 0 {
2443 if seg_len == 1 {
2444 self.arena.free_segment(segment_idx);
2445 }
2446 self.replace_prepared_child(
2447 history,
2448 last_index,
2449 current_start_depth,
2450 ChildRef::from_node(branch),
2451 );
2452 } else {
2453 let segment = &mut self.arena.segments[segment_idx.get()];
2454 segment.payload = original.payload.prefix(offset as u32);
2455 segment.tail = ChildRef::from_node(branch);
2456 segment.symbol_count = updated_counts;
2457 segment.log_prob_kt = updated_log_prob_kt;
2458 self.recompute_segment_head(segment_idx);
2459 }
2460
2461 for idx in (0..last_index).rev() {
2462 self.recompute_source_weight(self.prepared_steps[idx].source);
2463 }
2464
2465 let root_edge = history_symbol(history, 0) as usize;
2466 self.arena.child(self.root, root_edge)
2467 }
2468
2469 fn update_prepared_cached_path<L: CtLogAccess>(
2470 &mut self,
2471 logs: L,
2472 history: &(impl HistoryAccess + ?Sized),
2473 sym_idx: usize,
2474 singleton_log_prob_kt: f64,
2475 ) {
2476 debug_assert!(!self.prepared_steps.is_empty());
2477 debug_assert!(matches!(
2478 self.prepared_end,
2479 PreparedEnd::MaxDepth | PreparedEnd::MissingAfterCurrent
2480 ));
2481
2482 if self.prepared_end == PreparedEnd::MissingAfterCurrent {
2483 self.attach_missing_after_prepared_path(history, sym_idx, singleton_log_prob_kt);
2484 }
2485
2486 let last_index = self.prepared_steps.len() - 1;
2487 let mut child_weight = if self.prepared_end == PreparedEnd::MissingAfterCurrent {
2488 let last_step = self.prepared_steps[last_index];
2489 match last_step.source {
2490 ExistingSource::Node(node_idx) => {
2491 let depth = self.prepared_levels;
2492 let edge = history_symbol(history, depth) as usize;
2493 self.arena
2494 .child_ref_weighted(self.arena.child(node_idx, edge))
2495 }
2496 ExistingSource::Segment(segment_idx, offset) => {
2497 debug_assert_eq!(offset + 1, self.arena.segment_len(segment_idx));
2498 self.arena
2499 .child_ref_weighted(self.arena.segments[segment_idx.get()].tail)
2500 }
2501 ExistingSource::None => unreachable!("prepared path should never end in None"),
2502 }
2503 } else {
2504 0.0
2505 };
2506
2507 for idx in (0..=last_index).rev() {
2508 let step = self.prepared_steps[idx];
2509 match step.source {
2510 ExistingSource::Node(node_idx) => {
2511 let (mut counts, mut log_prob_kt) =
2512 self.source_counts_and_kt_log_prob(step.source);
2513 apply_update_to_state_raw(logs, &mut counts, &mut log_prob_kt, sym_idx);
2514 let weighted =
2515 if idx == last_index && self.prepared_end == PreparedEnd::MaxDepth {
2516 debug_assert_eq!(step.has_sibling, 0);
2517 clamp_log_prob(log_prob_kt)
2518 } else {
2519 update_weighted_log_prob(
2520 log_prob_kt,
2521 child_weight,
2522 step.sibling_weight,
2523 false,
2524 )
2525 };
2526 let slot = node_idx.get();
2527 self.arena.nodes[slot].symbol_count = counts;
2528 self.arena.nodes[slot].log_prob_kt = log_prob_kt;
2529 self.arena.nodes[slot].log_prob_weighted = weighted;
2530 child_weight = weighted;
2531 }
2532 ExistingSource::Segment(segment_idx, offset) => {
2533 let (mut counts, mut log_prob_kt) =
2534 self.source_counts_and_kt_log_prob(step.source);
2535 apply_update_to_state_raw(logs, &mut counts, &mut log_prob_kt, sym_idx);
2536 let slot = segment_idx.get();
2537 let weighted =
2538 if idx == last_index && self.prepared_end == PreparedEnd::MaxDepth {
2539 debug_assert_eq!(offset + 1, self.arena.segment_len(segment_idx));
2540 debug_assert!(self.arena.segments[slot].tail.is_none());
2541 clamp_log_prob(log_prob_kt)
2542 } else {
2543 let (alpha, log_alpha, log_one_minus_alpha) =
2544 self.segment_constants(step.span);
2545 unary_chain_log_weight_precomputed(
2546 log_prob_kt,
2547 child_weight,
2548 alpha,
2549 log_alpha,
2550 log_one_minus_alpha,
2551 )
2552 };
2553 self.arena.segments[slot].symbol_count = counts;
2554 self.arena.segments[slot].log_prob_kt = log_prob_kt;
2555 self.arena.segments[slot].head_log_prob_weighted = weighted;
2556 child_weight = weighted;
2557 }
2558 ExistingSource::None => unreachable!("prepared update should never visit None"),
2559 }
2560 }
2561 }
2562
2563 fn update_child_fast<L: CtLogAccess>(
2564 &mut self,
2565 logs: L,
2566 child: ChildRef,
2567 depth: usize,
2568 history: &(impl HistoryAccess + ?Sized),
2569 sym_idx: usize,
2570 singleton_log_prob_kt: f64,
2571 ) -> ChildRef {
2572 if depth > self.max_depth {
2573 return child;
2574 }
2575 if child.is_none() {
2576 return self.build_missing_path(depth, history, sym_idx, singleton_log_prob_kt);
2577 }
2578
2579 if let Some(node_idx) = child.as_node() {
2580 if depth < self.max_depth {
2581 let path_edge = history_symbol(history, depth) as usize;
2582 let next = self.arena.child(node_idx, path_edge);
2583 let updated = self.update_child_fast(
2584 logs,
2585 next,
2586 depth + 1,
2587 history,
2588 sym_idx,
2589 singleton_log_prob_kt,
2590 );
2591 if updated != next {
2592 self.arena.set_child(node_idx, path_edge, updated);
2593 }
2594 }
2595 let mut counts = self.arena.nodes[node_idx.get()].symbol_count;
2596 let mut log_prob_kt = self.arena.nodes[node_idx.get()].log_prob_kt;
2597 apply_update_to_state_raw(logs, &mut counts, &mut log_prob_kt, sym_idx);
2598 self.arena.nodes[node_idx.get()].symbol_count = counts;
2599 self.arena.nodes[node_idx.get()].log_prob_kt = log_prob_kt;
2600 self.arena.recompute_node_weight(node_idx);
2601 return ChildRef::from_node(node_idx);
2602 }
2603
2604 let segment_idx = child.as_segment().unwrap();
2605 let original = self.arena.segments[segment_idx.get()];
2606 let seg_len = original.len() as usize;
2607 let mut updated_counts = original.symbol_count;
2608 let mut updated_log_prob_kt = original.log_prob_kt;
2609 apply_update_to_state_raw(logs, &mut updated_counts, &mut updated_log_prob_kt, sym_idx);
2610
2611 let depth_budget = self.max_depth.saturating_sub(depth);
2612 let comparable_len = if original.tail.is_none() {
2613 seg_len.saturating_sub(1)
2614 } else {
2615 seg_len
2616 }
2617 .min(depth_budget);
2618 let mismatch = first_segment_mismatch(original, depth, history, comparable_len).map(
2619 |(offset, path_edge, existing_edge)| (offset, depth + offset, path_edge, existing_edge),
2620 );
2621
2622 if let Some((offset, node_depth, path_edge, existing_edge)) = mismatch {
2623 let old_continuation = if offset + 1 < seg_len {
2624 if offset == 0 {
2625 let segment = &mut self.arena.segments[segment_idx.get()];
2626 segment.payload = original.payload.suffix_after(1);
2627 segment.tail = original.tail;
2628 segment.symbol_count = original.symbol_count;
2629 segment.log_prob_kt = original.log_prob_kt;
2630 self.recompute_segment_head(segment_idx);
2631 ChildRef::from_segment(segment_idx)
2632 } else {
2633 ChildRef::from_segment(self.arena.alloc_segment_with_parts(
2634 original.symbol_count,
2635 original.log_prob_kt,
2636 original.tail,
2637 original.payload.suffix_after(offset as u32 + 1),
2638 ))
2639 }
2640 } else {
2641 original.tail
2642 };
2643
2644 let new_tail =
2645 self.build_missing_path(node_depth + 1, history, sym_idx, singleton_log_prob_kt);
2646 let branch = self
2647 .arena
2648 .alloc_node_with_state(updated_counts, updated_log_prob_kt);
2649 self.arena
2650 .set_child(branch, existing_edge as usize, old_continuation);
2651 self.arena.set_child(branch, path_edge as usize, new_tail);
2652 self.arena.recompute_node_weight(branch);
2653
2654 if offset == 0 {
2655 if offset + 1 >= seg_len {
2656 self.arena.free_segment(segment_idx);
2657 }
2658 return ChildRef::from_node(branch);
2659 }
2660
2661 let segment = &mut self.arena.segments[segment_idx.get()];
2662 segment.payload = original.payload.prefix(offset as u32);
2663 segment.tail = ChildRef::from_node(branch);
2664 segment.symbol_count = updated_counts;
2665 segment.log_prob_kt = updated_log_prob_kt;
2666 self.recompute_segment_head(segment_idx);
2667 return ChildRef::from_segment(segment_idx);
2668 }
2669
2670 if depth_budget < seg_len {
2671 self.arena.segments[segment_idx.get()].symbol_count = updated_counts;
2672 self.arena.segments[segment_idx.get()].log_prob_kt = updated_log_prob_kt;
2673 self.recompute_segment_head(segment_idx);
2674 return ChildRef::from_segment(segment_idx);
2675 }
2676
2677 if original.tail.is_none() {
2678 let new_tail =
2679 self.build_missing_path(depth + seg_len, history, sym_idx, singleton_log_prob_kt);
2680 self.arena.segments[segment_idx.get()].tail = new_tail;
2681 self.arena.segments[segment_idx.get()].symbol_count = updated_counts;
2682 self.arena.segments[segment_idx.get()].log_prob_kt = updated_log_prob_kt;
2683 self.recompute_segment_head(segment_idx);
2684 return ChildRef::from_segment(segment_idx);
2685 }
2686
2687 let tail = original.tail;
2688 let updated_tail = self.update_child_fast(
2689 logs,
2690 tail,
2691 depth + seg_len,
2692 history,
2693 sym_idx,
2694 singleton_log_prob_kt,
2695 );
2696 if updated_tail != tail {
2697 self.arena.set_segment_tail(segment_idx, updated_tail);
2698 }
2699 self.arena.segments[segment_idx.get()].symbol_count = updated_counts;
2700 self.arena.segments[segment_idx.get()].log_prob_kt = updated_log_prob_kt;
2701 self.recompute_segment_head(segment_idx);
2702 ChildRef::from_segment(segment_idx)
2703 }
2704
2705 #[allow(clippy::too_many_arguments)]
2709 fn update_child_fast_exact<L: CtLogAccess>(
2710 &mut self,
2711 logs: L,
2712 child: ChildRef,
2713 depth: usize,
2714 history: &(impl HistoryAccess + ?Sized),
2715 path_bits: u64,
2716 sym_idx: usize,
2717 singleton_log_prob_kt: f64,
2718 ) -> ChildRef {
2719 debug_assert!(self.max_depth <= SEG_EXACT_MAX_LEN as usize);
2720 if depth > self.max_depth {
2721 return child;
2722 }
2723 if child.is_none() {
2724 return self.build_missing_path_exact_bits(
2725 depth,
2726 path_bits,
2727 sym_idx,
2728 singleton_log_prob_kt,
2729 );
2730 }
2731
2732 if let Some(node_idx) = child.as_node() {
2733 if depth < self.max_depth {
2734 let path_edge = (path_bits & 1) as usize;
2735 let next = self.arena.child(node_idx, path_edge);
2736 let updated = self.update_child_fast_exact(
2737 logs,
2738 next,
2739 depth + 1,
2740 history,
2741 shift_path_bits(path_bits, 1),
2742 sym_idx,
2743 singleton_log_prob_kt,
2744 );
2745 if updated != next {
2746 self.arena.set_child(node_idx, path_edge, updated);
2747 }
2748 }
2749 let slot = node_idx.get();
2750 let mut counts = self.arena.nodes[slot].symbol_count;
2751 let mut log_prob_kt = self.arena.nodes[slot].log_prob_kt;
2752 apply_update_to_state_raw(logs, &mut counts, &mut log_prob_kt, sym_idx);
2753 let [left, right] = self.arena.nodes[slot].children;
2754 let weighted = if left.is_none() && right.is_none() {
2755 clamp_log_prob(log_prob_kt)
2756 } else {
2757 let w0 = unsafe { self.arena.child_ref_weighted_unchecked(left) };
2760 let w1 = unsafe { self.arena.child_ref_weighted_unchecked(right) };
2761 update_weighted_log_prob_non_leaf(log_prob_kt, w0, w1)
2762 };
2763 self.arena.nodes[slot].symbol_count = counts;
2764 self.arena.nodes[slot].log_prob_kt = log_prob_kt;
2765 self.arena.nodes[slot].log_prob_weighted = weighted;
2766 return ChildRef::from_node(node_idx);
2767 }
2768
2769 let segment_idx = child.as_segment().unwrap();
2770 let original = self.arena.segments[segment_idx.get()];
2771 if !original.payload.is_exact() {
2772 return self.update_child_fast(
2773 logs,
2774 child,
2775 depth,
2776 history,
2777 sym_idx,
2778 singleton_log_prob_kt,
2779 );
2780 }
2781
2782 let seg_len = original.len() as usize;
2783 let mut updated_counts = original.symbol_count;
2784 let mut updated_log_prob_kt = original.log_prob_kt;
2785 apply_update_to_state_raw(logs, &mut updated_counts, &mut updated_log_prob_kt, sym_idx);
2786
2787 let depth_budget = self.max_depth.saturating_sub(depth);
2788 let comparable_len = if original.tail.is_none() {
2789 seg_len.saturating_sub(1)
2790 } else {
2791 seg_len
2792 }
2793 .min(depth_budget);
2794 let mismatch =
2795 first_exact_segment_mismatch(original.payload.exact_bits(), path_bits, comparable_len)
2796 .map(|(offset, path_edge, existing_edge)| {
2797 (offset, depth + offset, path_edge, existing_edge)
2798 });
2799
2800 if let Some((offset, node_depth, path_edge, existing_edge)) = mismatch {
2801 let old_continuation = if offset + 1 < seg_len {
2802 if offset == 0 {
2803 let segment = &mut self.arena.segments[segment_idx.get()];
2804 segment.payload = original.payload.suffix_after(1);
2805 segment.tail = original.tail;
2806 segment.symbol_count = original.symbol_count;
2807 segment.log_prob_kt = original.log_prob_kt;
2808 self.recompute_segment_head(segment_idx);
2809 ChildRef::from_segment(segment_idx)
2810 } else {
2811 ChildRef::from_segment(self.arena.alloc_segment_with_parts(
2812 original.symbol_count,
2813 original.log_prob_kt,
2814 original.tail,
2815 original.payload.suffix_after(offset as u32 + 1),
2816 ))
2817 }
2818 } else {
2819 original.tail
2820 };
2821
2822 let new_tail = self.build_missing_path_exact_bits(
2823 node_depth + 1,
2824 shift_path_bits(path_bits, offset + 1),
2825 sym_idx,
2826 singleton_log_prob_kt,
2827 );
2828 let branch = self
2829 .arena
2830 .alloc_node_with_state(updated_counts, updated_log_prob_kt);
2831 self.arena
2832 .set_child(branch, existing_edge as usize, old_continuation);
2833 self.arena.set_child(branch, path_edge as usize, new_tail);
2834 self.arena.recompute_node_weight(branch);
2835
2836 if offset == 0 {
2837 if offset + 1 >= seg_len {
2838 self.arena.free_segment(segment_idx);
2839 }
2840 return ChildRef::from_node(branch);
2841 }
2842
2843 let segment = &mut self.arena.segments[segment_idx.get()];
2844 segment.payload = original.payload.prefix(offset as u32);
2845 segment.tail = ChildRef::from_node(branch);
2846 segment.symbol_count = updated_counts;
2847 segment.log_prob_kt = updated_log_prob_kt;
2848 self.recompute_segment_head(segment_idx);
2849 return ChildRef::from_segment(segment_idx);
2850 }
2851
2852 if depth_budget < seg_len {
2853 self.arena.segments[segment_idx.get()].symbol_count = updated_counts;
2854 self.arena.segments[segment_idx.get()].log_prob_kt = updated_log_prob_kt;
2855 self.recompute_segment_head(segment_idx);
2856 return ChildRef::from_segment(segment_idx);
2857 }
2858
2859 if original.tail.is_none() {
2860 let new_tail = self.build_missing_path_exact_bits(
2861 depth + seg_len,
2862 shift_path_bits(path_bits, seg_len),
2863 sym_idx,
2864 singleton_log_prob_kt,
2865 );
2866 self.arena.segments[segment_idx.get()].tail = new_tail;
2867 self.arena.segments[segment_idx.get()].symbol_count = updated_counts;
2868 self.arena.segments[segment_idx.get()].log_prob_kt = updated_log_prob_kt;
2869 self.recompute_segment_head(segment_idx);
2870 return ChildRef::from_segment(segment_idx);
2871 }
2872
2873 let tail = original.tail;
2874 let updated_tail = self.update_child_fast_exact(
2875 logs,
2876 tail,
2877 depth + seg_len,
2878 history,
2879 shift_path_bits(path_bits, seg_len),
2880 sym_idx,
2881 singleton_log_prob_kt,
2882 );
2883 if updated_tail != tail {
2884 self.arena.set_segment_tail(segment_idx, updated_tail);
2885 }
2886 self.arena.segments[segment_idx.get()].symbol_count = updated_counts;
2887 self.arena.segments[segment_idx.get()].log_prob_kt = updated_log_prob_kt;
2888 self.recompute_segment_head(segment_idx);
2889 ChildRef::from_segment(segment_idx)
2890 }
2891
2892 #[inline(always)]
2893 fn update_root_child<L: CtLogAccess>(
2894 &mut self,
2895 logs: L,
2896 child: ChildRef,
2897 history: &(impl HistoryAccess + ?Sized),
2898 sym_idx: usize,
2899 singleton_log_prob_kt: f64,
2900 ) -> ChildRef {
2901 if self.max_depth <= SEG_EXACT_MAX_LEN as usize {
2902 let path_bits = path_bits_from_history(history, 1, self.max_depth);
2903 self.update_child_fast_exact(
2904 logs,
2905 child,
2906 1,
2907 history,
2908 path_bits,
2909 sym_idx,
2910 singleton_log_prob_kt,
2911 )
2912 } else {
2913 self.update_child_fast(logs, child, 1, history, sym_idx, singleton_log_prob_kt)
2914 }
2915 }
2916
2917 fn collect_existing_levels(&mut self, history: &(impl HistoryAccess + ?Sized)) -> ChildRef {
2918 if self.max_depth == 0 {
2919 self.detaches.clear();
2920 return ChildRef::NONE;
2921 }
2922
2923 self.detaches.clear();
2924 self.levels.fill(LevelState::default());
2925
2926 let root_edge = history_symbol(history, 0) as usize;
2927 let old_child = self.arena.child(self.root, root_edge);
2928 let mut source = if let Some(node) = old_child.as_node() {
2929 ExistingSource::Node(node)
2930 } else if let Some(segment) = old_child.as_segment() {
2931 ExistingSource::Segment(segment, 0)
2932 } else {
2933 ExistingSource::None
2934 };
2935
2936 for depth in 1..=self.max_depth {
2937 let slot = depth - 1;
2938 self.levels[slot] = LevelState::default();
2939
2940 match source {
2941 ExistingSource::None => {}
2942 ExistingSource::Node(node_idx) => {
2943 self.levels[slot].symbol_count = self.arena.counts(node_idx);
2944 self.levels[slot].log_prob_kt = self.arena.log_prob_kt(node_idx);
2945 if depth < self.max_depth {
2946 let path_edge = history_symbol(history, depth) as usize;
2947 let sibling_edge = path_edge ^ 1;
2948 let sibling = self.arena.child(node_idx, sibling_edge);
2949 self.levels[slot].sibling = sibling;
2950 if sibling.is_some() {
2951 self.detaches.push(Detach::NodeChild {
2952 node: node_idx,
2953 edge: sibling_edge,
2954 });
2955 }
2956 let next = self.arena.child(node_idx, path_edge);
2957 source = if let Some(next_node) = next.as_node() {
2958 ExistingSource::Node(next_node)
2959 } else if let Some(next_segment) = next.as_segment() {
2960 ExistingSource::Segment(next_segment, 0)
2961 } else {
2962 ExistingSource::None
2963 };
2964 }
2965 }
2966 ExistingSource::Segment(segment_idx, offset) => {
2967 self.levels[slot].symbol_count = self.arena.segment_symbol_count(segment_idx);
2968 self.levels[slot].log_prob_kt = self.arena.segment_log_prob_kt(segment_idx);
2969 if depth < self.max_depth {
2970 let path_edge = history_symbol(history, depth) as usize;
2971 if self.arena.segment_has_child(segment_idx, offset) {
2972 let existing_edge =
2973 self.arena.segment_edge(segment_idx, offset, history);
2974 if path_edge == existing_edge {
2975 let seg_len = self.arena.segment_len(segment_idx);
2976 if offset + 1 < seg_len {
2977 source = ExistingSource::Segment(segment_idx, offset + 1);
2978 } else {
2979 let tail = self.arena.segments[segment_idx.get()].tail;
2980 source = if let Some(next_node) = tail.as_node() {
2981 ExistingSource::Node(next_node)
2982 } else if let Some(next_segment) = tail.as_segment() {
2983 ExistingSource::Segment(next_segment, 0)
2984 } else {
2985 ExistingSource::None
2986 };
2987 }
2988 } else {
2989 let continuation = self.arena.detach_segment_continuation(
2990 segment_idx,
2991 offset,
2992 &mut self.detaches,
2993 );
2994 self.levels[slot].sibling = continuation;
2995 source = ExistingSource::None;
2996 }
2997 } else {
2998 source = ExistingSource::None;
2999 }
3000 }
3001 }
3002 }
3003 }
3004
3005 old_child
3006 }
3007
3008 fn rebuild_path_subtree(&mut self, history: &(impl HistoryAccess + ?Sized)) -> ChildRef {
3009 let mut built = ChildRef::NONE;
3010
3011 for depth in (1..=self.max_depth).rev() {
3012 let level = self.levels[depth - 1];
3013 let visits = level.symbol_count[0] + level.symbol_count[1];
3014 if visits == 0 {
3015 built = ChildRef::NONE;
3016 continue;
3017 }
3018
3019 let path_edge = if depth < self.max_depth {
3020 history_symbol(history, depth) as usize
3021 } else {
3022 0
3023 };
3024 let has_path_child = built.is_some();
3025 let has_sibling = level.sibling.is_some();
3026 let force_node = depth <= self.hot_prefix_depth();
3027
3028 if force_node || (has_path_child && has_sibling) {
3029 let node = self
3030 .arena
3031 .alloc_node_with_state(level.symbol_count, level.log_prob_kt);
3032 if has_path_child {
3033 self.arena.set_child(node, path_edge, built);
3034 }
3035 if has_sibling {
3036 self.arena.set_child(node, path_edge ^ 1, level.sibling);
3037 }
3038 self.arena.recompute_node_weight(node);
3039 built = ChildRef::from_node(node);
3040 } else {
3041 let (edge, child) = if has_path_child {
3042 (path_edge, built)
3043 } else if has_sibling {
3044 (path_edge ^ 1, level.sibling)
3045 } else {
3046 (path_edge, ChildRef::NONE)
3047 };
3048 built = self.arena.prepend_or_alloc_segment(
3049 history,
3050 depth,
3051 level.symbol_count,
3052 level.log_prob_kt,
3053 child,
3054 edge,
3055 false,
3056 );
3057 }
3058 }
3059
3060 built
3061 }
3062
3063 fn apply_detaches(&mut self) {
3064 for detach in self.detaches.drain(..) {
3065 match detach {
3066 Detach::NodeChild { node, edge } => {
3067 self.arena.set_child(node, edge, ChildRef::NONE);
3068 }
3069 Detach::SegmentNext { segment, new_len } => {
3070 self.arena.segments[segment.get()].set_len(new_len);
3071 self.arena.set_segment_tail(segment, ChildRef::NONE);
3072 }
3073 }
3074 }
3075 }
3076
3077 fn update_with_logs<L: CtLogAccess>(
3078 &mut self,
3079 logs: L,
3080 sym: Symbol,
3081 history: &(impl HistoryAccess + ?Sized),
3082 ) {
3083 let sym_idx = sym as usize;
3084 let singleton_log_prob_kt = logs.log_half(0) - logs.log_int(1);
3085 {
3086 let slot = self.root.get();
3087 let mut counts = self.arena.nodes[slot].symbol_count;
3088 let mut log_prob_kt = self.arena.nodes[slot].log_prob_kt;
3089 apply_update_to_state_raw(logs, &mut counts, &mut log_prob_kt, sym_idx);
3090 self.arena.nodes[slot].symbol_count = counts;
3091 self.arena.nodes[slot].log_prob_kt = log_prob_kt;
3092 }
3093
3094 if self.max_depth > 0 {
3095 let root_edge = history_symbol(history, 0) as usize;
3096 let old_child = self.arena.child(self.root, root_edge);
3097 let new_child =
3098 self.update_root_child(logs, old_child, history, sym_idx, singleton_log_prob_kt);
3099 self.arena.set_child(self.root, root_edge, new_child);
3100 }
3101
3102 self.arena.recompute_node_weight(self.root);
3103 }
3104
3105 fn update(&mut self, sym: Symbol, history: &(impl HistoryAccess + ?Sized)) {
3106 let upto = self.root_visits() + 1;
3107 if upto <= ctw_log_cache_limit() {
3108 self.with_cached_logs(upto, |this, logs| {
3109 this.update_with_logs(logs, sym, history);
3110 });
3111 } else {
3112 self.with_bounded_logs(upto, |this, logs| {
3113 this.update_with_logs(logs, sym, history);
3114 });
3115 }
3116 }
3117
3118 fn update_prepared_with_logs<L: CtLogAccess>(
3119 &mut self,
3120 logs: L,
3121 history: &(impl HistoryAccess + ?Sized),
3122 sym_idx: usize,
3123 use_prepared: bool,
3124 ) {
3125 let singleton_log_prob_kt = logs.log_half(0) - logs.log_int(1);
3126 {
3127 let slot = self.root.get();
3128 let mut counts = self.arena.nodes[slot].symbol_count;
3129 let mut log_prob_kt = self.arena.nodes[slot].log_prob_kt;
3130 apply_update_to_state_raw(logs, &mut counts, &mut log_prob_kt, sym_idx);
3131 self.arena.nodes[slot].symbol_count = counts;
3132 self.arena.nodes[slot].log_prob_kt = log_prob_kt;
3133 }
3134
3135 if self.max_depth > 0 {
3136 let root_edge = history_symbol(history, 0) as usize;
3137 let old_child = self.arena.child(self.root, root_edge);
3138 let new_child = if use_prepared {
3139 match self.prepared_end {
3140 PreparedEnd::MissingAtRoot => {
3141 self.build_missing_path(1, history, sym_idx, singleton_log_prob_kt)
3142 }
3143 PreparedEnd::MaxDepth | PreparedEnd::MissingAfterCurrent => {
3144 if !self.prepared_steps.is_empty() {
3145 self.update_prepared_cached_path(
3146 logs,
3147 history,
3148 sym_idx,
3149 singleton_log_prob_kt,
3150 );
3151 }
3152 old_child
3153 }
3154 PreparedEnd::MismatchAtCurrentSegment => {
3155 self.update_prepared_mismatch(logs, history, sym_idx, singleton_log_prob_kt)
3156 }
3157 }
3158 } else {
3159 self.update_root_child(logs, old_child, history, sym_idx, singleton_log_prob_kt)
3160 };
3161 self.arena.set_child(self.root, root_edge, new_child);
3162 }
3163
3164 self.arena.recompute_node_weight(self.root);
3165 }
3166
3167 fn update_prepared(
3168 &mut self,
3169 sym: Symbol,
3170 history: &(impl HistoryAccess + ?Sized),
3171 use_prepared: bool,
3172 ) {
3173 let upto = self.root_visits() + 1;
3174 let sym_idx = sym as usize;
3175 if upto <= ctw_log_cache_limit() {
3176 self.with_cached_logs(upto, |this, logs| {
3177 this.update_prepared_with_logs(logs, history, sym_idx, use_prepared);
3178 });
3179 } else {
3180 self.with_bounded_logs(upto, |this, logs| {
3181 this.update_prepared_with_logs(logs, history, sym_idx, use_prepared);
3182 });
3183 }
3184 }
3185
3186 fn revert_with_logs<L: CtLogAccess>(
3187 &mut self,
3188 logs: L,
3189 history: &(impl HistoryAccess + ?Sized),
3190 sym_idx: usize,
3191 ) {
3192 let old_child = self.collect_existing_levels(history);
3193
3194 {
3195 let slot = self.root.get();
3196 let mut counts = self.arena.nodes[slot].symbol_count;
3197 let mut log_prob_kt = self.arena.nodes[slot].log_prob_kt;
3198 apply_revert_to_state_raw(logs, &mut counts, &mut log_prob_kt, sym_idx);
3199 self.arena.nodes[slot].symbol_count = counts;
3200 self.arena.nodes[slot].log_prob_kt = log_prob_kt;
3201 }
3202
3203 for level in &mut self.levels {
3204 let mut counts = level.symbol_count;
3205 let mut log_prob_kt = level.log_prob_kt;
3206 apply_revert_to_state_raw(logs, &mut counts, &mut log_prob_kt, sym_idx);
3207 level.symbol_count = counts;
3208 level.log_prob_kt = log_prob_kt;
3209 }
3210
3211 if self.max_depth > 0 {
3212 let new_child = self.rebuild_path_subtree(history);
3213 let root_edge = history_symbol(history, 0) as usize;
3214 self.apply_detaches();
3215 self.arena.free_child_ref(old_child);
3216 self.arena.set_child(self.root, root_edge, new_child);
3217 }
3218
3219 self.arena.recompute_node_weight(self.root);
3220 }
3221
3222 fn revert(&mut self, sym: Symbol, history: &(impl HistoryAccess + ?Sized)) {
3223 let upto = self.root_visits();
3224 let sym_idx = sym as usize;
3225 if upto <= ctw_log_cache_limit() {
3226 self.with_cached_logs(upto, |this, logs| {
3227 this.revert_with_logs(logs, history, sym_idx);
3228 });
3229 } else {
3230 self.with_bounded_logs(upto, |this, logs| {
3231 this.revert_with_logs(logs, history, sym_idx);
3232 });
3233 }
3234 }
3235
3236 fn predict(&mut self, sym: Symbol, history: &(impl HistoryAccess + ?Sized)) -> f64 {
3237 self.prepared_steps.clear();
3238 self.prepared_levels = 0;
3239 self.prepared_end = PreparedEnd::MaxDepth;
3240
3241 let (root_sibling, root_has_sibling, mut source) = if self.max_depth > 0 {
3242 let root_edge = history_symbol(history, 0) as usize;
3243 let path_child = self.arena.child(self.root, root_edge);
3244 let sibling = self.arena.child(self.root, root_edge ^ 1);
3245 (
3246 self.arena.child_ref_weighted(sibling),
3247 sibling.is_some() as u8,
3248 Self::child_to_existing_source(path_child).unwrap_or(ExistingSource::None),
3249 )
3250 } else {
3251 (0.0, 0, ExistingSource::None)
3252 };
3253 if self.max_depth > 0 && matches!(source, ExistingSource::None) {
3254 self.prepared_end = PreparedEnd::MissingAtRoot;
3255 }
3256
3257 let history_len = history.len();
3258 let mut depth = 1usize;
3259 'walk: while depth <= self.max_depth {
3260 match source {
3261 ExistingSource::None => break,
3262 ExistingSource::Node(node_idx) => {
3263 if depth == self.max_depth {
3264 self.prepared_steps.push(PreparedStep {
3265 source: ExistingSource::Node(node_idx),
3266 span: 1,
3267 sibling_weight: 0.0,
3268 has_sibling: 0,
3269 });
3270 self.prepared_levels += 1;
3271 break;
3272 }
3273 let path_edge = history_symbol(history, depth) as usize;
3274 let sibling = self.arena.child(node_idx, path_edge ^ 1);
3275 self.prepared_steps.push(PreparedStep {
3276 source: ExistingSource::Node(node_idx),
3277 span: 1,
3278 sibling_weight: self.arena.child_ref_weighted(sibling),
3279 has_sibling: sibling.is_some() as u8,
3280 });
3281 self.prepared_levels += 1;
3282 let next = self.arena.child(node_idx, path_edge);
3283 source = Self::child_to_existing_source(next).unwrap_or(ExistingSource::None);
3284 if matches!(source, ExistingSource::None) {
3285 self.prepared_end = PreparedEnd::MissingAfterCurrent;
3286 break;
3287 }
3288 depth += 1;
3289 }
3290 ExistingSource::Segment(segment_idx, _) => {
3291 let segment = self.arena.segments[segment_idx.get()];
3292 if segment.payload.is_exact() {
3293 let path_bits = path_bits_from_history(
3294 history,
3295 depth,
3296 self.max_depth.saturating_sub(depth),
3297 );
3298 if let Some((next_depth, next_source)) =
3299 self.walk_prepared_exact_segment(segment_idx, segment, depth, path_bits)
3300 {
3301 source = next_source;
3302 if matches!(source, ExistingSource::None) {
3303 self.prepared_end = PreparedEnd::MissingAfterCurrent;
3304 break 'walk;
3305 }
3306 depth = next_depth;
3307 continue 'walk;
3308 }
3309 break 'walk;
3310 }
3311 let seg_len = segment.len() as usize;
3312 for offset in 0..seg_len {
3313 let node_depth = depth + offset;
3314 let span = (offset + 1) as u32;
3315
3316 if node_depth == self.max_depth {
3317 self.prepared_steps.push(PreparedStep {
3318 source: ExistingSource::Segment(segment_idx, offset as u32),
3319 span,
3320 sibling_weight: 0.0,
3321 has_sibling: 0,
3322 });
3323 self.prepared_levels += span as usize;
3324 break 'walk;
3325 }
3326
3327 if offset + 1 >= seg_len && segment.tail.is_none() {
3328 self.prepared_steps.push(PreparedStep {
3329 source: ExistingSource::Segment(segment_idx, offset as u32),
3330 span,
3331 sibling_weight: 0.0,
3332 has_sibling: 0,
3333 });
3334 self.prepared_levels += span as usize;
3335 self.prepared_end = PreparedEnd::MissingAfterCurrent;
3336 break 'walk;
3337 }
3338
3339 let path_edge = path_edge_at_depth(history, history_len, node_depth);
3340 let existing_edge =
3341 segment_edge_from_parts(segment, offset, history, history_len);
3342 if path_edge != existing_edge {
3343 self.prepared_steps.push(PreparedStep {
3344 source: ExistingSource::Segment(segment_idx, offset as u32),
3345 span,
3346 sibling_weight: self
3347 .arena
3348 .segment_continuation_weight(segment_idx, offset as u32),
3349 has_sibling: 1,
3350 });
3351 self.prepared_levels += span as usize;
3352 self.prepared_end = PreparedEnd::MismatchAtCurrentSegment;
3353 break 'walk;
3354 }
3355
3356 if offset + 1 < seg_len {
3357 continue;
3358 }
3359
3360 self.prepared_steps.push(PreparedStep {
3361 source: ExistingSource::Segment(segment_idx, offset as u32),
3362 span,
3363 sibling_weight: 0.0,
3364 has_sibling: 0,
3365 });
3366 self.prepared_levels += span as usize;
3367 let tail = segment.tail;
3368 source =
3369 Self::child_to_existing_source(tail).unwrap_or(ExistingSource::None);
3370 if matches!(source, ExistingSource::None) {
3371 self.prepared_end = PreparedEnd::MissingAfterCurrent;
3372 break 'walk;
3373 }
3374 depth = node_depth + 1;
3375 continue 'walk;
3376 }
3377 }
3378 }
3379 }
3380
3381 let sym_idx = sym as usize;
3382 if self.prepared_levels == 0 {
3383 let counts = self.arena.counts(self.root);
3384 let kt_log_prob = self.arena.log_prob_kt(self.root);
3385 return if self.prepared_end == PreparedEnd::MaxDepth || root_has_sibling == 0 {
3386 predict_ratio_kt(counts, sym_idx)
3387 } else {
3388 predict_ratio_internal(kt_log_prob, counts, 0.0, root_sibling, 0.5, sym_idx)
3389 };
3390 }
3391
3392 let last_step = *self.prepared_steps.last().unwrap();
3393 let (last_counts, last_kt_log_prob) = self.source_counts_and_kt_log_prob(last_step.source);
3394 let (mut child_weight, mut ratio) = if (self.prepared_end == PreparedEnd::MaxDepth
3395 && self.prepared_levels == self.max_depth)
3396 || last_step.has_sibling == 0
3397 {
3398 (last_kt_log_prob, predict_ratio_kt(last_counts, sym_idx))
3399 } else {
3400 combined_weight_ratio_internal(
3401 last_kt_log_prob,
3402 last_counts,
3403 0.0,
3404 last_step.sibling_weight,
3405 0.5,
3406 sym_idx,
3407 )
3408 };
3409
3410 if matches!(last_step.source, ExistingSource::Segment(_, _)) && last_step.span > 1 {
3411 let (alpha, log_alpha, log_one_minus_alpha) =
3412 self.segment_constants(last_step.span - 1);
3413 (child_weight, ratio) = unary_chain_ratio_transform_precomputed(
3414 last_kt_log_prob,
3415 last_counts,
3416 child_weight,
3417 ratio,
3418 alpha,
3419 log_alpha,
3420 log_one_minus_alpha,
3421 sym_idx,
3422 );
3423 }
3424
3425 for idx in (0..self.prepared_steps.len() - 1).rev() {
3426 let step = self.prepared_steps[idx];
3427 let (step_counts, step_kt_log_prob) = self.source_counts_and_kt_log_prob(step.source);
3428 match step.source {
3429 ExistingSource::Node(_) => {
3430 (child_weight, ratio) = combined_weight_ratio_internal(
3431 step_kt_log_prob,
3432 step_counts,
3433 child_weight,
3434 step.sibling_weight,
3435 ratio,
3436 sym_idx,
3437 );
3438 }
3439 ExistingSource::Segment(_, _) => {
3440 let (alpha, log_alpha, log_one_minus_alpha) = self.segment_constants(step.span);
3441 (child_weight, ratio) = unary_chain_ratio_transform_precomputed(
3442 step_kt_log_prob,
3443 step_counts,
3444 child_weight,
3445 ratio,
3446 alpha,
3447 log_alpha,
3448 log_one_minus_alpha,
3449 sym_idx,
3450 );
3451 }
3452 ExistingSource::None => unreachable!("prepared step should never store None"),
3453 }
3454 }
3455
3456 let root_counts = self.arena.counts(self.root);
3457 let root_kt_log_prob = self.arena.log_prob_kt(self.root);
3458 predict_ratio_internal(
3459 root_kt_log_prob,
3460 root_counts,
3461 child_weight,
3462 root_sibling,
3463 ratio,
3464 sym_idx,
3465 )
3466 }
3467
3468 fn predict_one(&mut self, history: &(impl HistoryAccess + ?Sized)) -> f64 {
3469 self.prepared_steps.clear();
3470 self.prepared_levels = 0;
3471 self.prepared_end = PreparedEnd::MaxDepth;
3472
3473 let (root_sibling, root_has_sibling, mut source) = if self.max_depth > 0 {
3474 let root_edge = history_symbol(history, 0) as usize;
3475 let path_child = self.arena.child(self.root, root_edge);
3476 let sibling = self.arena.child(self.root, root_edge ^ 1);
3477 (
3478 self.arena.child_ref_weighted(sibling),
3479 sibling.is_some() as u8,
3480 Self::child_to_existing_source(path_child).unwrap_or(ExistingSource::None),
3481 )
3482 } else {
3483 (0.0, 0, ExistingSource::None)
3484 };
3485 if self.max_depth > 0 && matches!(source, ExistingSource::None) {
3486 self.prepared_end = PreparedEnd::MissingAtRoot;
3487 }
3488
3489 let history_len = history.len();
3490 let mut depth = 1usize;
3491 'walk: while depth <= self.max_depth {
3492 match source {
3493 ExistingSource::None => break,
3494 ExistingSource::Node(node_idx) => {
3495 if depth == self.max_depth {
3496 self.prepared_steps.push(PreparedStep {
3497 source: ExistingSource::Node(node_idx),
3498 span: 1,
3499 sibling_weight: 0.0,
3500 has_sibling: 0,
3501 });
3502 self.prepared_levels += 1;
3503 break;
3504 }
3505 let path_edge = history_symbol(history, depth) as usize;
3506 let sibling = self.arena.child(node_idx, path_edge ^ 1);
3507 self.prepared_steps.push(PreparedStep {
3508 source: ExistingSource::Node(node_idx),
3509 span: 1,
3510 sibling_weight: self.arena.child_ref_weighted(sibling),
3511 has_sibling: sibling.is_some() as u8,
3512 });
3513 self.prepared_levels += 1;
3514 let next = self.arena.child(node_idx, path_edge);
3515 source = Self::child_to_existing_source(next).unwrap_or(ExistingSource::None);
3516 if matches!(source, ExistingSource::None) {
3517 self.prepared_end = PreparedEnd::MissingAfterCurrent;
3518 break;
3519 }
3520 depth += 1;
3521 }
3522 ExistingSource::Segment(segment_idx, _) => {
3523 let segment = self.arena.segments[segment_idx.get()];
3524 if segment.payload.is_exact() {
3525 let path_bits = path_bits_from_history(
3526 history,
3527 depth,
3528 self.max_depth.saturating_sub(depth),
3529 );
3530 if let Some((next_depth, next_source)) =
3531 self.walk_prepared_exact_segment(segment_idx, segment, depth, path_bits)
3532 {
3533 source = next_source;
3534 if matches!(source, ExistingSource::None) {
3535 self.prepared_end = PreparedEnd::MissingAfterCurrent;
3536 break 'walk;
3537 }
3538 depth = next_depth;
3539 continue 'walk;
3540 }
3541 break 'walk;
3542 }
3543 let seg_len = segment.len() as usize;
3544 for offset in 0..seg_len {
3545 let node_depth = depth + offset;
3546 let span = (offset + 1) as u32;
3547
3548 if node_depth == self.max_depth {
3549 self.prepared_steps.push(PreparedStep {
3550 source: ExistingSource::Segment(segment_idx, offset as u32),
3551 span,
3552 sibling_weight: 0.0,
3553 has_sibling: 0,
3554 });
3555 self.prepared_levels += span as usize;
3556 break 'walk;
3557 }
3558
3559 if offset + 1 >= seg_len && segment.tail.is_none() {
3560 self.prepared_steps.push(PreparedStep {
3561 source: ExistingSource::Segment(segment_idx, offset as u32),
3562 span,
3563 sibling_weight: 0.0,
3564 has_sibling: 0,
3565 });
3566 self.prepared_levels += span as usize;
3567 self.prepared_end = PreparedEnd::MissingAfterCurrent;
3568 break 'walk;
3569 }
3570
3571 let path_edge = path_edge_at_depth(history, history_len, node_depth);
3572 let existing_edge =
3573 segment_edge_from_parts(segment, offset, history, history_len);
3574 if path_edge != existing_edge {
3575 self.prepared_steps.push(PreparedStep {
3576 source: ExistingSource::Segment(segment_idx, offset as u32),
3577 span,
3578 sibling_weight: self
3579 .arena
3580 .segment_continuation_weight(segment_idx, offset as u32),
3581 has_sibling: 1,
3582 });
3583 self.prepared_levels += span as usize;
3584 self.prepared_end = PreparedEnd::MismatchAtCurrentSegment;
3585 break 'walk;
3586 }
3587
3588 if offset + 1 < seg_len {
3589 continue;
3590 }
3591
3592 self.prepared_steps.push(PreparedStep {
3593 source: ExistingSource::Segment(segment_idx, offset as u32),
3594 span,
3595 sibling_weight: 0.0,
3596 has_sibling: 0,
3597 });
3598 self.prepared_levels += span as usize;
3599 let tail = segment.tail;
3600 source =
3601 Self::child_to_existing_source(tail).unwrap_or(ExistingSource::None);
3602 if matches!(source, ExistingSource::None) {
3603 self.prepared_end = PreparedEnd::MissingAfterCurrent;
3604 break 'walk;
3605 }
3606 depth = node_depth + 1;
3607 continue 'walk;
3608 }
3609 }
3610 }
3611 }
3612
3613 if self.prepared_levels == 0 {
3614 let counts = self.arena.counts(self.root);
3615 let kt_log_prob = self.arena.log_prob_kt(self.root);
3616 return if self.prepared_end == PreparedEnd::MaxDepth || root_has_sibling == 0 {
3617 predict_ratio_kt_one(counts)
3618 } else {
3619 predict_ratio_internal_one(kt_log_prob, counts, 0.0, root_sibling, 0.5)
3620 };
3621 }
3622
3623 let last_step = *self.prepared_steps.last().unwrap();
3624 let (last_counts, last_kt_log_prob) = self.source_counts_and_kt_log_prob(last_step.source);
3625 let (mut child_weight, mut ratio) = if (self.prepared_end == PreparedEnd::MaxDepth
3626 && self.prepared_levels == self.max_depth)
3627 || last_step.has_sibling == 0
3628 {
3629 (last_kt_log_prob, predict_ratio_kt_one(last_counts))
3630 } else {
3631 combined_weight_ratio_internal_one(
3632 last_kt_log_prob,
3633 last_counts,
3634 0.0,
3635 last_step.sibling_weight,
3636 0.5,
3637 )
3638 };
3639
3640 if let ExistingSource::Segment(_, _) = last_step.source
3641 && last_step.span > 1
3642 {
3643 let (alpha, log_alpha, log_one_minus_alpha) =
3644 self.segment_constants(last_step.span - 1);
3645 (child_weight, ratio) = unary_chain_ratio_transform_precomputed_one(
3646 last_kt_log_prob,
3647 last_counts,
3648 child_weight,
3649 ratio,
3650 alpha,
3651 log_alpha,
3652 log_one_minus_alpha,
3653 );
3654 }
3655
3656 for idx in (0..self.prepared_steps.len() - 1).rev() {
3657 let step = self.prepared_steps[idx];
3658 let (step_counts, step_kt_log_prob) = self.source_counts_and_kt_log_prob(step.source);
3659 match step.source {
3660 ExistingSource::Node(_) => {
3661 (child_weight, ratio) = combined_weight_ratio_internal_one(
3662 step_kt_log_prob,
3663 step_counts,
3664 child_weight,
3665 step.sibling_weight,
3666 ratio,
3667 );
3668 }
3669 ExistingSource::Segment(_, _) => {
3670 let (alpha, log_alpha, log_one_minus_alpha) = self.segment_constants(step.span);
3671 (child_weight, ratio) = unary_chain_ratio_transform_precomputed_one(
3672 step_kt_log_prob,
3673 step_counts,
3674 child_weight,
3675 ratio,
3676 alpha,
3677 log_alpha,
3678 log_one_minus_alpha,
3679 );
3680 }
3681 ExistingSource::None => unreachable!("prepared step should never store None"),
3682 }
3683 }
3684
3685 let root_counts = self.arena.counts(self.root);
3686 let root_kt_log_prob = self.arena.log_prob_kt(self.root);
3687 predict_ratio_internal_one(
3688 root_kt_log_prob,
3689 root_counts,
3690 child_weight,
3691 root_sibling,
3692 ratio,
3693 )
3694 }
3695
3696 fn memory_usage(&self) -> usize {
3697 self.arena.memory_usage()
3698 + self.segment_alpha.capacity() * size_of::<f64>()
3699 + self.segment_log_alpha.capacity() * size_of::<f64>()
3700 + self.segment_log_one_minus_alpha.capacity() * size_of::<f64>()
3701 + self.levels.capacity() * size_of::<LevelState>()
3702 + self.detaches.capacity() * size_of::<Detach>()
3703 + self.prepared_steps.capacity() * size_of::<PreparedStep>()
3704 }
3705
3706 #[cfg(any(test, feature = "research-tooling"))]
3707 fn scratch_memory_usage(&self) -> usize {
3708 self.segment_alpha.capacity() * size_of::<f64>()
3709 + self.segment_log_alpha.capacity() * size_of::<f64>()
3710 + self.segment_log_one_minus_alpha.capacity() * size_of::<f64>()
3711 + self.levels.capacity() * size_of::<LevelState>()
3712 + self.detaches.capacity() * size_of::<Detach>()
3713 + self.prepared_steps.capacity() * size_of::<PreparedStep>()
3714 }
3715
3716 #[cfg(any(test, feature = "research-tooling"))]
3717 fn telemetry(&self, bit_index: usize) -> FacContextTreeTreeTelemetry {
3718 let mut exact_segments: usize = 0;
3719 let mut history_segments: usize = 0;
3720 let mut history_invert_segments: usize = 0;
3721 let mut const_segments: usize = 0;
3722 let mut segment_bits: u64 = 0;
3723 let mut max_segment_len: u32 = 0;
3724 for segment in &self.arena.segments {
3725 let len = segment.len();
3726 segment_bits = segment_bits.saturating_add(len as u64);
3727 max_segment_len = max_segment_len.max(len);
3728 match segment.payload.mode() {
3729 SEG_MODE_EXACT => exact_segments = exact_segments.saturating_add(1),
3730 SEG_MODE_HISTORY => history_segments = history_segments.saturating_add(1),
3731 SEG_MODE_HISTORY_INVERT => {
3732 history_invert_segments = history_invert_segments.saturating_add(1);
3733 }
3734 SEG_MODE_CONST => const_segments = const_segments.saturating_add(1),
3735 _ => unreachable!("invalid ctw segment payload mode"),
3736 }
3737 }
3738
3739 let node_payload_bytes = self.arena.nodes.len() * size_of::<CtNode>();
3740 let node_bytes = self.arena.nodes.capacity() * size_of::<CtNode>();
3741 let segment_payload_bytes = self.arena.segments.len() * size_of::<CtSegment>();
3742 let segment_bytes = self.arena.segments.capacity() * size_of::<CtSegment>();
3743 let free_list_bytes = self.arena.free_nodes.capacity() * size_of::<NodeIndex>()
3744 + self.arena.free_segments.capacity() * size_of::<SegmentIndex>();
3745 let scratch_bytes = self.scratch_memory_usage();
3746 let arena_slack_bytes = node_bytes
3747 .saturating_sub(node_payload_bytes)
3748 .saturating_add(segment_bytes.saturating_sub(segment_payload_bytes));
3749
3750 FacContextTreeTreeTelemetry {
3751 bit_index,
3752 max_depth: self.max_depth,
3753 root_visits: self.root_visits(),
3754 nodes_len: self.arena.nodes.len(),
3755 nodes_capacity: self.arena.nodes.capacity(),
3756 segments_len: self.arena.segments.len(),
3757 segments_capacity: self.arena.segments.capacity(),
3758 free_nodes_len: self.arena.free_nodes.len(),
3759 free_nodes_capacity: self.arena.free_nodes.capacity(),
3760 free_segments_len: self.arena.free_segments.len(),
3761 free_segments_capacity: self.arena.free_segments.capacity(),
3762 node_bytes,
3763 node_payload_bytes,
3764 segment_bytes,
3765 segment_payload_bytes,
3766 free_list_bytes,
3767 scratch_bytes,
3768 total_bytes: node_bytes
3769 .saturating_add(segment_bytes)
3770 .saturating_add(free_list_bytes)
3771 .saturating_add(scratch_bytes),
3772 arena_slack_bytes,
3773 exact_segments,
3774 history_segments,
3775 history_invert_segments,
3776 const_segments,
3777 segment_bits,
3778 max_segment_len,
3779 }
3780 }
3781}
3782
3783#[derive(Clone)]
3785pub struct ContextTree {
3786 engine: CtEngine,
3787 history: BitHistory,
3788 history_version: u64,
3789 prepared_valid: bool,
3790 prepared_history_len: usize,
3791 prepared_history_version: u64,
3792}
3793
3794#[derive(Clone)]
3795pub(crate) struct ContextTreeLifecycleSnapshot {
3796 history: BitHistory,
3797 history_version: u64,
3798}
3799
3800impl ContextTree {
3801 pub fn new(depth: usize) -> Self {
3803 Self {
3804 engine: CtEngine::new(depth),
3805 history: BitHistory::default(),
3806 history_version: 0,
3807 prepared_valid: false,
3808 prepared_history_len: 0,
3809 prepared_history_version: 0,
3810 }
3811 }
3812
3813 #[inline]
3814 fn bump_history_version(&mut self) {
3815 self.history_version = self.history_version.wrapping_add(1);
3816 }
3817
3818 #[inline]
3819 fn clear_prepared_prediction(&mut self) {
3820 self.prepared_valid = false;
3821 }
3822
3823 #[inline]
3824 fn prepared_prediction_matches_history(&self) -> bool {
3825 self.prepared_valid
3826 && self.prepared_history_len == self.history.len()
3827 && self.prepared_history_version == self.history_version
3828 }
3829
3830 pub fn clear(&mut self) {
3832 self.history.clear();
3833 self.engine.clear();
3834 self.history_version = 0;
3835 self.clear_prepared_prediction();
3836 self.prepared_history_len = 0;
3837 self.prepared_history_version = 0;
3838 }
3839
3840 #[inline]
3841 pub(crate) fn reserve_for_symbols(&mut self, total_symbols: usize) {
3842 if total_symbols == 0 {
3843 return;
3844 }
3845 self.engine.reserve_for_symbols(total_symbols);
3846 self.history.reserve_exact(total_symbols);
3847 }
3848
3849 #[inline]
3850 pub fn update(&mut self, sym: Symbol) {
3852 let use_prepared = self.prepared_prediction_matches_history();
3853 self.clear_prepared_prediction();
3854 if use_prepared {
3855 self.engine.update_prepared(sym, &self.history, true);
3856 } else {
3857 self.engine.update(sym, &self.history);
3858 }
3859 self.history.push(sym);
3860 self.bump_history_version();
3861 }
3862
3863 #[inline]
3864 pub fn revert(&mut self) {
3866 let Some(last_sym) = self.history.pop() else {
3867 return;
3868 };
3869 self.clear_prepared_prediction();
3870 self.engine.revert(last_sym, &self.history);
3871 self.bump_history_version();
3872 }
3873
3874 #[inline]
3875 pub fn update_history(&mut self, symbols: &[Symbol]) {
3877 if symbols.is_empty() {
3878 return;
3879 }
3880 self.clear_prepared_prediction();
3881 self.history.extend_from_slice(symbols);
3882 self.bump_history_version();
3883 }
3884
3885 #[inline]
3886 pub fn revert_history(&mut self) {
3888 if self.history.pop().is_some() {
3889 self.clear_prepared_prediction();
3890 self.bump_history_version();
3891 }
3892 }
3893
3894 pub fn truncate_history(&mut self, new_size: usize) {
3896 if new_size < self.history.len() {
3897 self.clear_prepared_prediction();
3898 self.history.truncate(new_size);
3899 self.bump_history_version();
3900 }
3901 }
3902
3903 pub(crate) fn lifecycle_snapshot(&self) -> ContextTreeLifecycleSnapshot {
3909 ContextTreeLifecycleSnapshot {
3910 history: self.history.clone(),
3911 history_version: self.history_version,
3912 }
3913 }
3914
3915 pub(crate) fn restore_lifecycle_snapshot(&mut self, snapshot: ContextTreeLifecycleSnapshot) {
3916 self.history = snapshot.history;
3917 self.history_version = snapshot.history_version;
3918 self.clear_prepared_prediction();
3919 }
3920
3921 #[inline]
3922 pub fn predict(&mut self, sym: Symbol) -> f64 {
3924 let prob = self.engine.predict(sym, &self.history);
3925 self.prepared_valid = true;
3926 self.prepared_history_len = self.history.len();
3927 self.prepared_history_version = self.history_version;
3928 prob
3929 }
3930
3931 #[inline]
3932 pub(crate) fn predict_one(&mut self) -> f64 {
3933 let prob = self.engine.predict_one(&self.history);
3934 self.prepared_valid = true;
3935 self.prepared_history_len = self.history.len();
3936 self.prepared_history_version = self.history_version;
3937 prob
3938 }
3939
3940 #[inline]
3941 pub fn predict_sym_prob(&mut self) -> f64 {
3943 self.predict_one()
3944 }
3945
3946 #[inline]
3947 pub fn get_log_block_probability(&self) -> f64 {
3949 self.engine.get_log_block_probability()
3950 }
3951
3952 #[inline]
3953 pub fn depth(&self) -> usize {
3955 self.engine.max_depth
3956 }
3957
3958 #[inline]
3959 pub fn history_size(&self) -> usize {
3961 self.history.len()
3962 }
3963}
3964
3965#[derive(Clone)]
3966struct ContextTreeCore {
3967 engine: CtEngine,
3968 prepared_valid: bool,
3969 prepared_history_len: usize,
3970 prepared_history_version: u64,
3971}
3972
3973#[derive(Clone, Copy)]
3974struct ContextTreeCorePreparedSnapshot {
3975 prepared_valid: bool,
3976 prepared_history_len: usize,
3977 prepared_history_version: u64,
3978}
3979
3980impl ContextTreeCore {
3981 fn new(depth: usize) -> Self {
3982 Self {
3983 engine: CtEngine::new(depth),
3984 prepared_valid: false,
3985 prepared_history_len: 0,
3986 prepared_history_version: 0,
3987 }
3988 }
3989
3990 fn clear(&mut self) {
3991 self.engine.clear();
3992 self.prepared_valid = false;
3993 self.prepared_history_len = 0;
3994 self.prepared_history_version = 0;
3995 }
3996
3997 #[inline]
3998 fn reserve_for_symbols(&mut self, total_symbols: usize) {
3999 self.engine.reserve_for_symbols(total_symbols);
4000 }
4001
4002 #[inline]
4003 fn update_predicted(
4004 &mut self,
4005 sym: Symbol,
4006 shared_history: &(impl HistoryAccess + ?Sized),
4007 history_version: u64,
4008 ) {
4009 let use_prepared = self.prepared_valid
4010 && self.prepared_history_len == shared_history.len()
4011 && self.prepared_history_version == history_version;
4012 self.prepared_valid = false;
4013 self.engine
4014 .update_prepared(sym, shared_history, use_prepared);
4015 }
4016
4017 #[inline]
4018 fn update_predicted_with_logs<L: CtLogAccess>(
4019 &mut self,
4020 logs: L,
4021 sym: Symbol,
4022 shared_history: &(impl HistoryAccess + ?Sized),
4023 history_version: u64,
4024 ) {
4025 let use_prepared = self.prepared_valid
4026 && self.prepared_history_len == shared_history.len()
4027 && self.prepared_history_version == history_version;
4028 self.prepared_valid = false;
4029 self.engine
4030 .update_prepared_with_logs(logs, shared_history, sym as usize, use_prepared);
4031 }
4032
4033 #[inline]
4034 fn revert(&mut self, last_sym: Symbol, shared_history: &(impl HistoryAccess + ?Sized)) {
4035 self.prepared_valid = false;
4036 self.engine.revert(last_sym, shared_history);
4037 }
4038
4039 #[inline]
4040 fn predict(
4041 &mut self,
4042 sym: Symbol,
4043 shared_history: &(impl HistoryAccess + ?Sized),
4044 history_version: u64,
4045 ) -> f64 {
4046 let prob = self.engine.predict(sym, shared_history);
4047 self.prepared_valid = true;
4048 self.prepared_history_len = shared_history.len();
4049 self.prepared_history_version = history_version;
4050 prob
4051 }
4052
4053 #[inline]
4054 fn predict_one(
4055 &mut self,
4056 shared_history: &(impl HistoryAccess + ?Sized),
4057 history_version: u64,
4058 ) -> f64 {
4059 let prob = self.engine.predict_one(shared_history);
4060 self.prepared_valid = true;
4061 self.prepared_history_len = shared_history.len();
4062 self.prepared_history_version = history_version;
4063 prob
4064 }
4065
4066 #[inline]
4067 fn get_log_block_probability(&self) -> f64 {
4068 self.engine.get_log_block_probability()
4069 }
4070
4071 #[inline]
4072 fn prepared_snapshot(&self) -> ContextTreeCorePreparedSnapshot {
4073 ContextTreeCorePreparedSnapshot {
4074 prepared_valid: self.prepared_valid,
4075 prepared_history_len: self.prepared_history_len,
4076 prepared_history_version: self.prepared_history_version,
4077 }
4078 }
4079
4080 #[inline]
4081 fn restore_prepared_snapshot(&mut self, snapshot: ContextTreeCorePreparedSnapshot) {
4082 self.prepared_valid = snapshot.prepared_valid;
4083 self.prepared_history_len = snapshot.prepared_history_len;
4084 self.prepared_history_version = snapshot.prepared_history_version;
4085 }
4086}
4087
4088#[derive(Clone)]
4090pub struct FacContextTree {
4091 trees: Vec<ContextTreeCore>,
4092 shared_history: BitHistory,
4093 base_depth: usize,
4094 num_bits: usize,
4095 shared_history_version: u64,
4096}
4097
4098#[derive(Clone)]
4099pub(crate) struct FacContextTreeLifecycleSnapshot {
4100 shared_history: BitHistory,
4101 shared_history_version: u64,
4102 prepared: Vec<ContextTreeCorePreparedSnapshot>,
4103}
4104
4105#[cfg(any(test, feature = "research-tooling"))]
4107#[derive(Clone, Copy, Debug, PartialEq, Eq)]
4108pub struct FacContextTreeMemoryUsage {
4109 pub tree_bytes: usize,
4111 pub shared_log_cache_bytes: usize,
4113 pub shared_history_bytes: usize,
4115}
4116
4117#[cfg(any(test, feature = "research-tooling"))]
4119#[derive(Clone, Debug, PartialEq, Eq)]
4120#[non_exhaustive]
4121pub struct FacContextTreeTreeTelemetry {
4122 pub bit_index: usize,
4124 pub max_depth: usize,
4126 pub root_visits: usize,
4128 pub nodes_len: usize,
4130 pub nodes_capacity: usize,
4132 pub segments_len: usize,
4134 pub segments_capacity: usize,
4136 pub free_nodes_len: usize,
4138 pub free_nodes_capacity: usize,
4140 pub free_segments_len: usize,
4142 pub free_segments_capacity: usize,
4144 pub node_bytes: usize,
4146 pub node_payload_bytes: usize,
4148 pub segment_bytes: usize,
4150 pub segment_payload_bytes: usize,
4152 pub free_list_bytes: usize,
4154 pub scratch_bytes: usize,
4156 pub total_bytes: usize,
4158 pub arena_slack_bytes: usize,
4160 pub exact_segments: usize,
4162 pub history_segments: usize,
4164 pub history_invert_segments: usize,
4166 pub const_segments: usize,
4168 pub segment_bits: u64,
4170 pub max_segment_len: u32,
4172}
4173
4174#[cfg(any(test, feature = "research-tooling"))]
4176#[derive(Clone, Debug, PartialEq, Eq)]
4177#[non_exhaustive]
4178pub struct FacContextTreeTelemetry {
4179 pub base_depth: usize,
4181 pub num_bits: usize,
4183 pub shared_history_len_bits: usize,
4185 pub shared_history_capacity_bits: usize,
4187 pub shared_history_bytes: usize,
4189 pub shared_history_payload_bytes: usize,
4191 pub shared_history_slack_bytes: usize,
4193 pub shared_log_cache_bytes: usize,
4195 pub tree_bytes: usize,
4197 pub tree_payload_bytes: usize,
4199 pub tree_arena_slack_bytes: usize,
4201 pub total_bytes: usize,
4203 pub total_slack_bytes: usize,
4205 pub nodes_len: usize,
4207 pub nodes_capacity: usize,
4209 pub segments_len: usize,
4211 pub segments_capacity: usize,
4213 pub free_nodes_len: usize,
4215 pub free_segments_len: usize,
4217 pub exact_segments: usize,
4219 pub history_segments: usize,
4221 pub history_invert_segments: usize,
4223 pub const_segments: usize,
4225 pub segment_bits: u64,
4227 pub max_segment_len: u32,
4229 pub trees: Vec<FacContextTreeTreeTelemetry>,
4231}
4232
4233#[cfg(any(test, feature = "research-tooling"))]
4234impl FacContextTreeMemoryUsage {
4235 #[inline]
4237 pub fn total_bytes(self) -> usize {
4238 self.tree_bytes
4239 .saturating_add(self.shared_log_cache_bytes)
4240 .saturating_add(self.shared_history_bytes)
4241 }
4242}
4243
4244impl FacContextTree {
4245 pub fn new(base_depth: usize, num_percept_bits: usize) -> Self {
4249 let trees = (0..num_percept_bits)
4250 .map(|i| ContextTreeCore::new(base_depth + i))
4251 .collect();
4252 Self {
4253 trees,
4254 shared_history: BitHistory::default(),
4255 base_depth,
4256 num_bits: num_percept_bits,
4257 shared_history_version: 0,
4258 }
4259 }
4260
4261 #[inline(always)]
4262 fn bump_shared_history_version(&mut self) {
4263 self.shared_history_version = self.shared_history_version.wrapping_add(1);
4264 }
4265
4266 #[inline]
4267 pub fn reserve_for_symbols(&mut self, total_symbols: usize) {
4269 if total_symbols == 0 {
4270 return;
4271 }
4272 self.shared_history
4273 .reserve_exact(total_symbols.saturating_mul(self.num_bits));
4274 for tree in &mut self.trees {
4275 tree.reserve_for_symbols(total_symbols);
4276 }
4277 }
4278
4279 #[inline]
4280 pub fn num_bits(&self) -> usize {
4282 self.num_bits
4283 }
4284
4285 #[inline]
4286 pub fn base_depth(&self) -> usize {
4288 self.base_depth
4289 }
4290
4291 #[inline]
4292 pub fn update(&mut self, sym: Symbol, bit_index: usize) {
4294 debug_assert!(bit_index < self.num_bits);
4295 self.trees[bit_index].update_predicted(
4296 sym,
4297 &self.shared_history,
4298 self.shared_history_version,
4299 );
4300 self.shared_history.push(sym);
4301 self.bump_shared_history_version();
4302 }
4303
4304 #[inline]
4305 pub fn update_byte_msb(&mut self, byte: u8) {
4307 if self.num_bits != 8 {
4308 for bit_idx in 0..self.num_bits {
4309 let bit = ((byte >> (7 - bit_idx)) & 1) == 1;
4310 self.update(bit, bit_idx);
4311 }
4312 return;
4313 }
4314
4315 let upto = self.trees[0].engine.root_visits() + 1;
4316 debug_assert!(
4317 self.trees
4318 .iter()
4319 .all(|tree| tree.engine.root_visits() + 1 == upto)
4320 );
4321 if upto <= ctw_log_cache_limit() {
4322 with_shared_cached_logs(upto, |logs| {
4323 for bit_idx in 0..8usize {
4324 let bit = ((byte >> (7 - bit_idx)) & 1) == 1;
4325 let tree = &mut self.trees[bit_idx];
4326 tree.prepared_valid = false;
4327 tree.engine
4328 .update_with_logs(logs, bit, &self.shared_history);
4329 self.shared_history.push(bit);
4330 }
4331 });
4332 } else {
4333 with_shared_bounded_logs(upto, |logs| {
4334 for bit_idx in 0..8usize {
4335 let bit = ((byte >> (7 - bit_idx)) & 1) == 1;
4336 let tree = &mut self.trees[bit_idx];
4337 tree.prepared_valid = false;
4338 tree.engine
4339 .update_with_logs(logs, bit, &self.shared_history);
4340 self.shared_history.push(bit);
4341 }
4342 });
4343 }
4344 self.bump_shared_history_version();
4345 }
4346
4347 #[inline]
4348 fn log_prob_update_byte_msb_with_logs<L: CtLogAccess>(&mut self, logs: L, byte: u8) -> f64 {
4349 let mut logp = 0.0;
4350 for bit_idx in 0..8usize {
4351 let bit = ((byte >> (7 - bit_idx)) & 1) == 1;
4352 let p =
4353 self.trees[bit_idx].predict(bit, &self.shared_history, self.shared_history_version);
4354 if p.is_finite() && p > 0.0 {
4355 logp += p.ln();
4356 } else {
4357 logp = f64::NEG_INFINITY;
4358 }
4359 self.trees[bit_idx].update_predicted_with_logs(
4360 logs,
4361 bit,
4362 &self.shared_history,
4363 self.shared_history_version,
4364 );
4365 self.shared_history.push(bit);
4366 self.bump_shared_history_version();
4367 }
4368 logp
4369 }
4370
4371 #[inline]
4372 pub fn log_prob_update_byte_msb(&mut self, byte: u8) -> f64 {
4374 debug_assert_eq!(self.num_bits, 8);
4375 let upto = self.trees[0].engine.root_visits() + 1;
4376 debug_assert!(
4377 self.trees
4378 .iter()
4379 .all(|tree| tree.engine.root_visits() + 1 == upto)
4380 );
4381 if upto <= ctw_log_cache_limit() {
4382 with_shared_cached_logs(upto, |logs| {
4383 self.log_prob_update_byte_msb_with_logs(logs, byte)
4384 })
4385 } else {
4386 with_shared_bounded_logs(upto, |logs| {
4387 self.log_prob_update_byte_msb_with_logs(logs, byte)
4388 })
4389 }
4390 }
4391
4392 #[inline]
4393 pub fn update_byte_lsb(&mut self, byte: u8) {
4395 let bits = self.num_bits.clamp(1, 8);
4396 let upto = self.trees[0].engine.root_visits() + 1;
4397 debug_assert!(
4398 self.trees
4399 .iter()
4400 .take(bits)
4401 .all(|tree| tree.engine.root_visits() + 1 == upto)
4402 );
4403 if upto <= ctw_log_cache_limit() {
4404 with_shared_cached_logs(upto, |logs| {
4405 for bit_idx in 0..bits {
4406 let bit = ((byte >> bit_idx) & 1) == 1;
4407 let tree = &mut self.trees[bit_idx];
4408 tree.prepared_valid = false;
4409 tree.engine
4410 .update_with_logs(logs, bit, &self.shared_history);
4411 self.shared_history.push(bit);
4412 }
4413 });
4414 } else {
4415 with_shared_bounded_logs(upto, |logs| {
4416 for bit_idx in 0..bits {
4417 let bit = ((byte >> bit_idx) & 1) == 1;
4418 let tree = &mut self.trees[bit_idx];
4419 tree.prepared_valid = false;
4420 tree.engine
4421 .update_with_logs(logs, bit, &self.shared_history);
4422 self.shared_history.push(bit);
4423 }
4424 });
4425 }
4426 self.bump_shared_history_version();
4427 }
4428
4429 #[inline]
4430 pub fn update_predicted(&mut self, sym: Symbol, bit_index: usize) {
4432 debug_assert!(bit_index < self.num_bits);
4433 self.trees[bit_index].update_predicted(
4434 sym,
4435 &self.shared_history,
4436 self.shared_history_version,
4437 );
4438 self.shared_history.push(sym);
4439 self.bump_shared_history_version();
4440 }
4441
4442 #[inline]
4443 pub fn predict(&mut self, sym: Symbol, bit_index: usize) -> f64 {
4445 debug_assert!(bit_index < self.num_bits);
4446 self.trees[bit_index].predict(sym, &self.shared_history, self.shared_history_version)
4447 }
4448
4449 #[inline]
4450 pub(crate) fn predict_one(&mut self, bit_index: usize) -> f64 {
4451 debug_assert!(bit_index < self.num_bits);
4452 self.trees[bit_index].predict_one(&self.shared_history, self.shared_history_version)
4453 }
4454
4455 #[inline]
4456 pub fn revert(&mut self, bit_index: usize) {
4458 debug_assert!(bit_index < self.num_bits);
4459 let Some(last_sym) = self.shared_history.pop() else {
4460 return;
4461 };
4462 self.trees[bit_index].revert(last_sym, &self.shared_history);
4463 self.bump_shared_history_version();
4464 }
4465
4466 #[inline]
4467 pub fn update_history(&mut self, symbols: &[Symbol]) {
4469 if symbols.is_empty() {
4470 return;
4471 }
4472 self.shared_history.extend_from_slice(symbols);
4473 self.bump_shared_history_version();
4474 }
4475
4476 #[inline]
4477 pub fn revert_history(&mut self, count: usize) {
4479 let old_len = self.shared_history.len();
4480 let new_len = self.shared_history.len().saturating_sub(count);
4481 if new_len == old_len {
4482 return;
4483 }
4484 self.shared_history.truncate(new_len);
4485 self.bump_shared_history_version();
4486 }
4487
4488 #[inline]
4489 pub fn reset_history_only(&mut self) {
4491 if self.shared_history.is_empty() {
4492 return;
4493 }
4494 self.shared_history.clear();
4495 self.bump_shared_history_version();
4496 }
4497
4498 pub(crate) fn lifecycle_snapshot(&self) -> FacContextTreeLifecycleSnapshot {
4505 FacContextTreeLifecycleSnapshot {
4506 shared_history: self.shared_history.clone(),
4507 shared_history_version: self.shared_history_version,
4508 prepared: self
4509 .trees
4510 .iter()
4511 .map(ContextTreeCore::prepared_snapshot)
4512 .collect(),
4513 }
4514 }
4515
4516 pub(crate) fn restore_lifecycle_snapshot(&mut self, snapshot: FacContextTreeLifecycleSnapshot) {
4517 debug_assert_eq!(self.trees.len(), snapshot.prepared.len());
4518 self.shared_history = snapshot.shared_history;
4519 self.shared_history_version = snapshot.shared_history_version;
4520 for (tree, prepared) in self.trees.iter_mut().zip(snapshot.prepared) {
4521 tree.restore_prepared_snapshot(prepared);
4522 }
4523 }
4524
4525 #[inline]
4526 pub fn get_log_block_probability(&self) -> f64 {
4528 self.trees
4529 .iter()
4530 .map(|t| t.get_log_block_probability())
4531 .sum()
4532 }
4533
4534 pub fn clear(&mut self) {
4536 for tree in &mut self.trees {
4537 tree.clear();
4538 }
4539 self.shared_history.clear();
4540 self.shared_history_version = 0;
4541 }
4542
4543 #[cfg(any(test, feature = "research-tooling"))]
4545 pub fn memory_usage_breakdown(&self) -> FacContextTreeMemoryUsage {
4546 let tree_mem: usize = self.trees.iter().map(|t| t.engine.memory_usage()).sum();
4547 let log_cache_mem = self
4548 .trees
4549 .first()
4550 .map(|t| t.engine.log_cache_memory_usage())
4551 .unwrap_or(0);
4552 let history_mem = self.shared_history.memory_usage();
4553 FacContextTreeMemoryUsage {
4554 tree_bytes: tree_mem,
4555 shared_log_cache_bytes: log_cache_mem,
4556 shared_history_bytes: history_mem,
4557 }
4558 }
4559
4560 #[cfg(any(test, feature = "research-tooling"))]
4562 pub fn telemetry(&self) -> FacContextTreeTelemetry {
4563 let usage = self.memory_usage_breakdown();
4564 let trees: Vec<FacContextTreeTreeTelemetry> = self
4565 .trees
4566 .iter()
4567 .enumerate()
4568 .map(|(bit_index, tree)| tree.engine.telemetry(bit_index))
4569 .collect();
4570 let shared_history_payload_bytes =
4571 self.shared_history.len().div_ceil(HISTORY_WORD_BITS) * size_of::<u64>();
4572 let shared_history_slack_bytes = usage
4573 .shared_history_bytes
4574 .saturating_sub(shared_history_payload_bytes);
4575 let nodes_len = trees.iter().map(|tree| tree.nodes_len).sum();
4576 let nodes_capacity = trees.iter().map(|tree| tree.nodes_capacity).sum();
4577 let segments_len = trees.iter().map(|tree| tree.segments_len).sum();
4578 let segments_capacity = trees.iter().map(|tree| tree.segments_capacity).sum();
4579 let free_nodes_len = trees.iter().map(|tree| tree.free_nodes_len).sum();
4580 let free_segments_len = trees.iter().map(|tree| tree.free_segments_len).sum();
4581 let exact_segments = trees.iter().map(|tree| tree.exact_segments).sum();
4582 let history_segments = trees.iter().map(|tree| tree.history_segments).sum();
4583 let history_invert_segments = trees.iter().map(|tree| tree.history_invert_segments).sum();
4584 let const_segments = trees.iter().map(|tree| tree.const_segments).sum();
4585 let segment_bits = trees.iter().map(|tree| tree.segment_bits).sum();
4586 let max_segment_len = trees
4587 .iter()
4588 .map(|tree| tree.max_segment_len)
4589 .max()
4590 .unwrap_or(0);
4591 let tree_payload_bytes = trees
4592 .iter()
4593 .map(|tree| {
4594 tree.node_payload_bytes
4595 .saturating_add(tree.segment_payload_bytes)
4596 })
4597 .sum();
4598 let tree_arena_slack_bytes: usize = trees.iter().map(|tree| tree.arena_slack_bytes).sum();
4599 let total_slack_bytes = tree_arena_slack_bytes.saturating_add(shared_history_slack_bytes);
4600
4601 FacContextTreeTelemetry {
4602 base_depth: self.base_depth,
4603 num_bits: self.num_bits,
4604 shared_history_len_bits: self.shared_history.len(),
4605 shared_history_capacity_bits: self.shared_history.words.capacity() * HISTORY_WORD_BITS,
4606 shared_history_bytes: usage.shared_history_bytes,
4607 shared_history_payload_bytes,
4608 shared_history_slack_bytes,
4609 shared_log_cache_bytes: usage.shared_log_cache_bytes,
4610 tree_bytes: usage.tree_bytes,
4611 tree_payload_bytes,
4612 tree_arena_slack_bytes,
4613 total_bytes: usage.total_bytes(),
4614 total_slack_bytes,
4615 nodes_len,
4616 nodes_capacity,
4617 segments_len,
4618 segments_capacity,
4619 free_nodes_len,
4620 free_segments_len,
4621 exact_segments,
4622 history_segments,
4623 history_invert_segments,
4624 const_segments,
4625 segment_bits,
4626 max_segment_len,
4627 trees,
4628 }
4629 }
4630
4631 pub fn memory_usage(&self) -> usize {
4633 let tree_mem: usize = self.trees.iter().map(|t| t.engine.memory_usage()).sum();
4634 let log_cache_mem = self
4635 .trees
4636 .first()
4637 .map(|t| t.engine.log_cache_memory_usage())
4638 .unwrap_or(0);
4639 let history_mem = self.shared_history.memory_usage();
4640 tree_mem
4641 .saturating_add(log_cache_mem)
4642 .saturating_add(history_mem)
4643 }
4644}
4645
4646#[inline]
4647fn compact_symbol_msb_shift(bits_per_symbol: usize, bit_idx: usize) -> usize {
4648 bits_per_symbol.saturating_sub(1).saturating_sub(bit_idx)
4649}
4650
4651#[inline]
4652pub(crate) fn ctw_symbol_bit_msb(symbol: u8, bits_per_symbol: usize, bit_idx: usize) -> bool {
4653 let bits = bits_per_symbol.clamp(1, 8);
4654 let shift = compact_symbol_msb_shift(bits, bit_idx);
4655 ((symbol >> shift) & 1) == 1
4656}
4657
4658#[inline]
4659pub(crate) fn ctw_log_prob_msb(
4660 tree: &mut ContextTree,
4661 symbol: u8,
4662 bits_per_symbol: usize,
4663 min_prob: f64,
4664) -> f64 {
4665 let bits = bits_per_symbol.clamp(1, 8);
4666 let mut logp = 0.0;
4667 for bit_idx in 0..bits {
4668 let bit = ctw_symbol_bit_msb(symbol, bits, bit_idx);
4669 let p = tree.predict(bit);
4670 if p.is_finite() && p > 0.0 {
4671 logp += p.ln();
4672 } else {
4673 logp = f64::NEG_INFINITY;
4674 }
4675 tree.update(bit);
4676 }
4677 for _ in 0..bits {
4678 tree.revert();
4679 }
4680 if logp.is_finite() {
4681 logp.max(min_prob.ln())
4682 } else {
4683 min_prob.ln()
4684 }
4685}
4686
4687#[inline]
4688pub(crate) fn ctw_log_prob_update_msb(
4689 tree: &mut ContextTree,
4690 symbol: u8,
4691 bits_per_symbol: usize,
4692 min_prob: f64,
4693) -> f64 {
4694 let bits = bits_per_symbol.clamp(1, 8);
4695 let mut logp = 0.0;
4696 for bit_idx in 0..bits {
4697 let bit = ctw_symbol_bit_msb(symbol, bits, bit_idx);
4698 let p = tree.predict(bit);
4699 if p.is_finite() && p > 0.0 {
4700 logp += p.ln();
4701 } else {
4702 logp = f64::NEG_INFINITY;
4703 }
4704 tree.update(bit);
4705 }
4706 if logp.is_finite() {
4707 logp.max(min_prob.ln())
4708 } else {
4709 min_prob.ln()
4710 }
4711}
4712
4713#[inline]
4714pub(crate) fn ctw_log_prob_update_lsb(
4715 tree: &mut FacContextTree,
4716 symbol: u8,
4717 bits_per_symbol: usize,
4718 min_prob: f64,
4719) -> f64 {
4720 let mut logp = 0.0;
4721 for bit_idx in 0..bits_per_symbol {
4722 let bit = ((symbol >> bit_idx) & 1) == 1;
4723 let p = tree.predict(bit, bit_idx);
4724 if p.is_finite() && p > 0.0 {
4725 logp += p.ln();
4726 } else {
4727 logp = f64::NEG_INFINITY;
4728 }
4729 tree.update_predicted(bit, bit_idx);
4730 }
4731 if logp.is_finite() {
4732 logp.max(min_prob.ln())
4733 } else {
4734 min_prob.ln()
4735 }
4736}
4737
4738pub(crate) fn fill_ctw_tree_log_probs(
4739 tree: &mut ContextTree,
4740 bits_per_symbol: usize,
4741 min_logp: f64,
4742 out: &mut [f64; 256],
4743) {
4744 let bits = bits_per_symbol.clamp(1, 8);
4745 let patterns = 1usize << bits;
4746 let mut pattern_logps = [f64::NEG_INFINITY; 256];
4747 let log_before = tree.get_log_block_probability();
4748
4749 fn rec(
4750 tree: &mut ContextTree,
4751 depth: usize,
4752 bits: usize,
4753 log_before: f64,
4754 min_logp: f64,
4755 symbol_acc: u8,
4756 pattern_logps: &mut [f64; 256],
4757 ) {
4758 if depth == bits {
4759 let pat = symbol_acc as usize;
4760 let logp = (tree.get_log_block_probability() - log_before).max(min_logp);
4761 pattern_logps[pat] = logp;
4762 return;
4763 }
4764
4765 for bit in [false, true] {
4766 tree.update(bit);
4767 let shift = compact_symbol_msb_shift(bits, depth);
4768 let next_symbol = if bit {
4769 symbol_acc | (1u8 << shift)
4770 } else {
4771 symbol_acc
4772 };
4773 rec(
4774 tree,
4775 depth + 1,
4776 bits,
4777 log_before,
4778 min_logp,
4779 next_symbol,
4780 pattern_logps,
4781 );
4782 tree.revert();
4783 }
4784 }
4785
4786 rec(tree, 0, bits, log_before, min_logp, 0, &mut pattern_logps);
4787
4788 if bits == 8 {
4789 out.copy_from_slice(&pattern_logps);
4790 } else {
4791 let aliases = 1usize << (8 - bits);
4792 let alias_ln = (aliases as f64).ln();
4793 let mask = patterns - 1;
4794 for byte in 0..256usize {
4795 out[byte] = pattern_logps[byte & mask] - alias_ln;
4796 }
4797 }
4798}
4799
4800pub(crate) fn fill_fac_tree_log_probs(
4801 tree: &mut FacContextTree,
4802 bits_per_symbol: usize,
4803 msb_first: bool,
4804 min_logp: f64,
4805 out: &mut [f64; 256],
4806) {
4807 struct RecParams {
4808 bits: usize,
4809 msb_first: bool,
4810 log_before: f64,
4811 min_logp: f64,
4812 }
4813
4814 let bits = bits_per_symbol.clamp(1, 8);
4815 let patterns = 1usize << bits;
4816 let mut pattern_logps = [f64::NEG_INFINITY; 256];
4817 let params = RecParams {
4818 bits,
4819 msb_first,
4820 log_before: tree.get_log_block_probability(),
4821 min_logp,
4822 };
4823
4824 fn rec(
4825 tree: &mut FacContextTree,
4826 depth: usize,
4827 params: &RecParams,
4828 symbol_acc: u8,
4829 pattern_logps: &mut [f64; 256],
4830 ) {
4831 if depth == params.bits {
4832 let pat = symbol_acc as usize;
4833 let logp = (tree.get_log_block_probability() - params.log_before).max(params.min_logp);
4834 pattern_logps[pat] = logp;
4835 return;
4836 }
4837
4838 for bit in [false, true] {
4839 tree.update(bit, depth);
4840 let mut next_symbol = symbol_acc;
4841 if params.msb_first {
4842 let shift = compact_symbol_msb_shift(params.bits, depth);
4845 if bit {
4846 next_symbol |= 1u8 << shift;
4847 }
4848 } else if bit {
4849 next_symbol |= 1u8 << depth;
4850 }
4851 rec(tree, depth + 1, params, next_symbol, pattern_logps);
4852 tree.revert(depth);
4853 }
4854 }
4855
4856 rec(tree, 0, ¶ms, 0, &mut pattern_logps);
4857
4858 if bits == 8 {
4859 out.copy_from_slice(&pattern_logps);
4860 } else {
4861 let aliases = 1usize << (8 - bits);
4862 let alias_ln = (aliases as f64).ln();
4863 let mask = patterns - 1;
4864 for byte in 0..256usize {
4865 out[byte] = pattern_logps[byte & mask] - alias_ln;
4866 }
4867 }
4868}
4869
4870#[cfg(test)]
4871mod tests {
4872 use super::*;
4873
4874 struct ScopedLogCacheLimit {
4875 previous: Option<usize>,
4876 }
4877
4878 struct ScopedLogOverflowCacheSlots {
4879 previous: Option<usize>,
4880 }
4881
4882 impl ScopedLogCacheLimit {
4883 fn set(limit: usize) -> Self {
4884 CTW_TEST_LOG_CACHE_LIMIT.with(|cell| {
4885 let previous = cell.replace(Some(limit));
4886 Self { previous }
4887 })
4888 }
4889 }
4890
4891 impl ScopedLogOverflowCacheSlots {
4892 fn set(slots: usize) -> Self {
4893 CTW_TEST_LOG_OVERFLOW_CACHE_SLOTS.with(|cell| {
4894 let previous = cell.replace(Some(slots));
4895 Self { previous }
4896 })
4897 }
4898 }
4899
4900 impl Drop for ScopedLogCacheLimit {
4901 fn drop(&mut self) {
4902 CTW_TEST_LOG_CACHE_LIMIT.with(|cell| {
4903 cell.replace(self.previous);
4904 });
4905 }
4906 }
4907
4908 impl Drop for ScopedLogOverflowCacheSlots {
4909 fn drop(&mut self) {
4910 CTW_TEST_LOG_OVERFLOW_CACHE_SLOTS.with(|cell| {
4911 cell.replace(self.previous);
4912 });
4913 }
4914 }
4915
4916 #[derive(Clone)]
4917 struct RefNode {
4918 children: [Option<Box<RefNode>>; 2],
4919 log_prob_kt: f64,
4920 log_prob_weighted: f64,
4921 symbol_count: [u32; 2],
4922 }
4923
4924 impl Default for RefNode {
4925 fn default() -> Self {
4926 Self {
4927 children: [None, None],
4928 log_prob_kt: 0.0,
4929 log_prob_weighted: 0.0,
4930 symbol_count: [0, 0],
4931 }
4932 }
4933 }
4934
4935 #[derive(Clone)]
4936 struct RefContextTree {
4937 root: RefNode,
4938 history: Vec<Symbol>,
4939 max_depth: usize,
4940 log_int: Vec<f64>,
4941 log_half: Vec<f64>,
4942 }
4943
4944 impl RefContextTree {
4945 fn new(depth: usize) -> Self {
4946 Self {
4947 root: RefNode::default(),
4948 history: Vec::new(),
4949 max_depth: depth,
4950 log_int: vec![f64::NEG_INFINITY],
4951 log_half: vec![(0.5f64).ln()],
4952 }
4953 }
4954
4955 fn root_visits(&self) -> usize {
4956 (self.root.symbol_count[0] + self.root.symbol_count[1]) as usize
4957 }
4958
4959 fn recompute(node: &mut RefNode) {
4960 let w0 = node.children[0]
4961 .as_ref()
4962 .map(|c| c.log_prob_weighted)
4963 .unwrap_or(0.0);
4964 let w1 = node.children[1]
4965 .as_ref()
4966 .map(|c| c.log_prob_weighted)
4967 .unwrap_or(0.0);
4968 let is_leaf = node.children[0].is_none() && node.children[1].is_none();
4969 node.log_prob_weighted = update_weighted_log_prob(node.log_prob_kt, w0, w1, is_leaf);
4970 }
4971
4972 fn update(&mut self, sym: Symbol) {
4973 let upto = self.root_visits() + 1;
4974 ensure_log_caches(&mut self.log_int, &mut self.log_half, upto);
4975 let sym_idx = sym as usize;
4976 Self::update_node(
4977 &mut self.root,
4978 0,
4979 self.max_depth,
4980 &self.history,
4981 sym_idx,
4982 CachedLogs::new(&self.log_int, &self.log_half),
4983 );
4984 self.history.push(sym);
4985 }
4986
4987 fn revert(&mut self) {
4988 let Some(last_sym) = self.history.pop() else {
4989 return;
4990 };
4991 let upto = self.root_visits();
4992 ensure_log_caches(&mut self.log_int, &mut self.log_half, upto);
4993 let sym_idx = last_sym as usize;
4994 let _ = Self::revert_node(
4995 &mut self.root,
4996 0,
4997 self.max_depth,
4998 &self.history,
4999 sym_idx,
5000 CachedLogs::new(&self.log_int, &self.log_half),
5001 );
5002 }
5003
5004 fn predict(&mut self, sym: Symbol) -> f64 {
5005 let sym_idx = sym as usize;
5006 let mut entries = Vec::with_capacity(self.max_depth + 1);
5007 let reached_max_depth = Self::collect_predict_entries(
5008 &self.root,
5009 0,
5010 self.max_depth,
5011 &self.history,
5012 &mut entries,
5013 );
5014
5015 let deepest = entries.len() - 1;
5016 let mut ratio = if reached_max_depth && deepest == self.max_depth {
5017 predict_ratio_kt(entries[deepest].symbol_count, sym_idx)
5018 } else {
5019 0.5
5020 };
5021 for idx in (0..=deepest).rev() {
5022 if reached_max_depth && idx == deepest {
5023 continue;
5024 }
5025 let child_weight = if idx < deepest {
5026 entries[idx + 1].log_prob_weighted
5027 } else {
5028 0.0
5029 };
5030 ratio = predict_ratio_internal(
5031 entries[idx].log_prob_kt,
5032 entries[idx].symbol_count,
5033 child_weight,
5034 entries[idx].sibling_weight,
5035 ratio,
5036 sym_idx,
5037 );
5038 }
5039 ratio
5040 }
5041
5042 fn get_log_block_probability(&self) -> f64 {
5043 self.root.log_prob_weighted
5044 }
5045
5046 fn update_node<L: CtLogAccess>(
5047 node: &mut RefNode,
5048 depth: usize,
5049 max_depth: usize,
5050 history: &(impl HistoryAccess + ?Sized),
5051 sym_idx: usize,
5052 logs: L,
5053 ) {
5054 if depth < max_depth {
5055 let edge = history_symbol(history, depth) as usize;
5056 if node.children[edge].is_none() {
5057 node.children[edge] = Some(Box::new(RefNode::default()));
5058 }
5059 Self::update_node(
5060 node.children[edge].as_deref_mut().unwrap(),
5061 depth + 1,
5062 max_depth,
5063 history,
5064 sym_idx,
5065 logs,
5066 );
5067 }
5068 apply_update_to_state_raw(logs, &mut node.symbol_count, &mut node.log_prob_kt, sym_idx);
5069 Self::recompute(node);
5070 }
5071
5072 fn revert_node<L: CtLogAccess>(
5073 node: &mut RefNode,
5074 depth: usize,
5075 max_depth: usize,
5076 history: &(impl HistoryAccess + ?Sized),
5077 sym_idx: usize,
5078 logs: L,
5079 ) -> bool {
5080 if depth < max_depth {
5081 let edge = history_symbol(history, depth) as usize;
5082 let remove_child = if let Some(child) = node.children[edge].as_deref_mut() {
5083 Self::revert_node(child, depth + 1, max_depth, history, sym_idx, logs)
5084 } else {
5085 false
5086 };
5087 if remove_child {
5088 node.children[edge] = None;
5089 }
5090 }
5091 apply_revert_to_state_raw(logs, &mut node.symbol_count, &mut node.log_prob_kt, sym_idx);
5092 Self::recompute(node);
5093 node.symbol_count[0] + node.symbol_count[1] == 0
5094 }
5095
5096 fn collect_predict_entries(
5097 node: &RefNode,
5098 depth: usize,
5099 max_depth: usize,
5100 history: &(impl HistoryAccess + ?Sized),
5101 entries: &mut Vec<PredictEntry>,
5102 ) -> bool {
5103 let sibling_weight = if depth < max_depth {
5104 let path_edge = history_symbol(history, depth) as usize;
5105 node.children[path_edge ^ 1]
5106 .as_ref()
5107 .map(|c| c.log_prob_weighted)
5108 .unwrap_or(0.0)
5109 } else {
5110 0.0
5111 };
5112 entries.push(PredictEntry {
5113 symbol_count: node.symbol_count,
5114 log_prob_kt: node.log_prob_kt,
5115 log_prob_weighted: node.log_prob_weighted,
5116 sibling_weight,
5117 has_sibling: depth < max_depth
5118 && node.children[(history_symbol(history, depth) as usize) ^ 1].is_some(),
5119 });
5120 if depth == max_depth {
5121 return true;
5122 }
5123 let edge = history_symbol(history, depth) as usize;
5124 let Some(child) = node.children[edge].as_ref() else {
5125 return false;
5126 };
5127 Self::collect_predict_entries(child, depth + 1, max_depth, history, entries)
5128 }
5129 }
5130
5131 #[derive(Clone)]
5132 struct RefFacContextTree {
5133 trees: Vec<RefContextTree>,
5134 history: Vec<Symbol>,
5135 }
5136
5137 impl RefFacContextTree {
5138 fn new(base_depth: usize, num_bits: usize) -> Self {
5139 Self {
5140 trees: (0..num_bits)
5141 .map(|i| RefContextTree::new(base_depth + i))
5142 .collect(),
5143 history: Vec::new(),
5144 }
5145 }
5146
5147 fn update(&mut self, sym: Symbol, bit_index: usize) {
5148 let tree = &mut self.trees[bit_index];
5149 tree.history = self.history.clone();
5150 tree.update(sym);
5151 self.history.push(sym);
5152 }
5153
5154 fn predict(&mut self, sym: Symbol, bit_index: usize) -> f64 {
5155 let tree = &mut self.trees[bit_index];
5156 tree.history = self.history.clone();
5157 tree.predict(sym)
5158 }
5159
5160 fn revert(&mut self, bit_index: usize) {
5161 let Some(last_sym) = self.history.pop() else {
5162 return;
5163 };
5164 let tree = &mut self.trees[bit_index];
5165 tree.history = self.history.clone();
5166 tree.history.push(last_sym);
5167 tree.revert();
5168 }
5169
5170 fn get_log_block_probability(&self) -> f64 {
5171 self.trees
5172 .iter()
5173 .map(RefContextTree::get_log_block_probability)
5174 .sum()
5175 }
5176 }
5177
5178 fn assert_close(a: f64, b: f64) {
5179 let diff = (a - b).abs();
5180 let scale = a.abs().max(b.abs()).max(1.0);
5181 assert!(diff <= 1e-12 * scale, "a={a} b={b} diff={diff}");
5182 }
5183
5184 fn child_after_hot_prefix(
5185 tree: &ContextTree,
5186 history_before_update: &(impl HistoryAccess + ?Sized),
5187 ) -> ChildRef {
5188 let hot_prefix_depth = tree.engine.hot_prefix_depth();
5189 if hot_prefix_depth == 0 {
5190 return ChildRef::NONE;
5191 }
5192
5193 let root_edge = history_symbol(history_before_update, 0) as usize;
5194 let mut current = tree
5195 .engine
5196 .arena
5197 .child(tree.engine.root, root_edge)
5198 .as_node()
5199 .expect("hot-prefix node");
5200 for node_depth in 1..hot_prefix_depth {
5201 let edge = history_symbol(history_before_update, node_depth) as usize;
5202 current = tree
5203 .engine
5204 .arena
5205 .child(current, edge)
5206 .as_node()
5207 .expect("next hot-prefix node");
5208 }
5209 let tail_edge = history_symbol(history_before_update, hot_prefix_depth) as usize;
5210 tree.engine.arena.child(current, tail_edge)
5211 }
5212
5213 #[test]
5214 #[should_panic(expected = "ctw node index overflow")]
5215 fn node_index_from_usize_rejects_overflow() {
5216 let _ = NodeIndex::from_usize(INDEX_LIMIT);
5217 }
5218
5219 #[test]
5220 #[should_panic(expected = "ctw node index overflow")]
5221 fn node_index_from_usize_rejects_large_values() {
5222 let _ = NodeIndex::from_usize(u32::MAX as usize);
5223 }
5224
5225 #[test]
5226 fn ctw_count_lane_stays_packed() {
5227 assert_eq!(std::mem::size_of::<CtNode>(), 32);
5228 }
5229
5230 #[test]
5231 fn ctw_segment_payload_stays_packed() {
5232 assert_eq!(std::mem::size_of::<CtSegment>(), 40);
5233 }
5234
5235 #[test]
5236 fn log_lookup_matches_direct_log_formulas_past_cache_limit() {
5237 let log_int = vec![f64::NEG_INFINITY, 0.0, (2.0f64).ln()];
5238 let log_half = vec![(0.5f64).ln(), (1.5f64).ln(), (2.5f64).ln()];
5239 let lookup = BoundedLogs::new(&log_int, &log_half);
5240
5241 for n in 0..16usize {
5242 let expected_int = if n == 0 {
5243 f64::NEG_INFINITY
5244 } else {
5245 (n as f64).ln()
5246 };
5247 assert_eq!(lookup.log_int(n).to_bits(), expected_int.to_bits());
5248 assert_eq!(
5249 lookup.log_half(n).to_bits(),
5250 (n as f64 + 0.5).ln().to_bits()
5251 );
5252 }
5253 }
5254
5255 #[test]
5256 fn log_lookup_overflow_cache_reuses_hot_exact_values() {
5257 let slots = [LogCacheSlot::empty(), LogCacheSlot::empty()];
5258 let misses = Cell::new(0usize);
5259 let first = BoundedLogs::lookup_overflow(&slots, 17, || {
5260 misses.set(misses.get() + 1);
5261 (17.0f64).ln()
5262 });
5263 let second = BoundedLogs::lookup_overflow(&slots, 17, || {
5264 misses.set(misses.get() + 1);
5265 f64::NAN
5266 });
5267
5268 assert_eq!(first.to_bits(), (17.0f64).ln().to_bits());
5269 assert_eq!(second.to_bits(), first.to_bits());
5270 assert_eq!(misses.get(), 1);
5271 }
5272
5273 #[test]
5274 fn shared_log_cache_respects_test_limit() {
5275 reset_shared_log_cache_for_test();
5276 let _limit = ScopedLogCacheLimit::set(3);
5277 let _overflow_slots = ScopedLogOverflowCacheSlots::set(4);
5278 with_shared_bounded_logs(32, |lookup| {
5279 assert_eq!(lookup.log_int(32).to_bits(), (32.0f64).ln().to_bits());
5280 assert_eq!(lookup.log_half(32).to_bits(), (32.5f64).ln().to_bits());
5281 });
5282
5283 let (log_int_len, log_half_len) = shared_log_cache_lens();
5284 let (overflow_log_int_len, overflow_log_half_len) = shared_log_overflow_cache_lens();
5285 assert!(log_int_len <= 4, "log_int_len={log_int_len}");
5286 assert!(log_half_len <= 4, "log_half_len={log_half_len}");
5287 assert_eq!(overflow_log_int_len, 4);
5288 assert_eq!(overflow_log_half_len, 4);
5289 }
5290
5291 #[test]
5292 fn fac_ctw_memory_usage_breakdown_sums_to_existing_total() {
5293 let mut fac = FacContextTree::new(7, 8);
5294 let payload = b"ctw memory usage breakdown payload";
5295 for &byte in payload {
5296 fac.update_byte_msb(byte);
5297 }
5298
5299 let usage = fac.memory_usage_breakdown();
5300 assert_eq!(fac.memory_usage(), usage.total_bytes());
5301 assert!(usage.tree_bytes > 0);
5302 assert!(usage.shared_log_cache_bytes > 0);
5303 assert!(usage.shared_history_bytes > 0);
5304
5305 let telemetry = fac.telemetry();
5306 assert_eq!(telemetry.total_bytes, usage.total_bytes());
5307 assert_eq!(telemetry.tree_bytes, usage.tree_bytes);
5308 assert_eq!(
5309 telemetry.shared_log_cache_bytes,
5310 usage.shared_log_cache_bytes
5311 );
5312 assert_eq!(telemetry.shared_history_bytes, usage.shared_history_bytes);
5313 assert_eq!(telemetry.shared_history_len_bits, payload.len() * 8);
5314 assert_eq!(telemetry.trees.len(), 8);
5315 assert_eq!(
5316 telemetry.nodes_len,
5317 telemetry
5318 .trees
5319 .iter()
5320 .map(|tree| tree.nodes_len)
5321 .sum::<usize>()
5322 );
5323 assert_eq!(
5324 telemetry.segments_len,
5325 telemetry
5326 .trees
5327 .iter()
5328 .map(|tree| tree.segments_len)
5329 .sum::<usize>()
5330 );
5331 assert_eq!(
5332 telemetry.segments_len,
5333 telemetry.exact_segments
5334 + telemetry.history_segments
5335 + telemetry.history_invert_segments
5336 + telemetry.const_segments
5337 );
5338 }
5339
5340 #[test]
5341 fn bit_history_preserves_logical_bits_and_packs_memory() {
5342 let mut history = BitHistory::default();
5343 let bits: Vec<Symbol> = (0..130usize).map(|idx| (idx * 17 + 5) % 7 < 3).collect();
5344 history.extend_from_slice(&bits);
5345 assert_eq!(history.len(), bits.len());
5346 assert_eq!(history.to_vec(), bits);
5347 assert_eq!(history.memory_usage(), 3 * size_of::<u64>());
5348 for depth in 0..160usize {
5349 assert_eq!(
5350 history_symbol(&history, depth),
5351 history_symbol(&bits, depth)
5352 );
5353 for len in [0usize, 1, 2, 7, 31, 32, 33, 63, 64] {
5354 assert_eq!(
5355 path_bits_from_history(&history, depth, len),
5356 path_bits_from_history(&bits, depth, len),
5357 "depth={depth} len={len}",
5358 );
5359 }
5360 }
5361
5362 let last = history.pop();
5363 assert_eq!(last, bits.last().copied());
5364 assert_eq!(history.len(), bits.len() - 1);
5365 let popped_bits = &bits[..bits.len() - 1];
5366 for depth in 0..160usize {
5367 assert_eq!(
5368 history_symbol(&history, depth),
5369 history_symbol(popped_bits, depth)
5370 );
5371 for len in [0usize, 1, 2, 7, 31, 32, 33, 63, 64] {
5372 assert_eq!(
5373 path_bits_from_history(&history, depth, len),
5374 path_bits_from_history(popped_bits, depth, len),
5375 "after pop depth={depth} len={len}",
5376 );
5377 }
5378 }
5379
5380 history.truncate(65);
5381 assert_eq!(history.to_vec(), bits[..65].to_vec());
5382 assert_eq!(history.memory_usage(), 3 * size_of::<u64>());
5383
5384 history.clear();
5385 assert!(history.is_empty());
5386 assert_eq!(history.memory_usage(), 3 * size_of::<u64>());
5387 }
5388
5389 #[test]
5390 fn fac_ctw_history_memory_is_bit_packed() {
5391 let mut fac = FacContextTree::new(4, 8);
5392 fac.reserve_for_symbols(1_000);
5393 let usage = fac.memory_usage_breakdown();
5394 assert_eq!(
5395 usage.shared_history_bytes,
5396 history_word_len(8_000) * size_of::<u64>()
5397 );
5398 }
5399
5400 #[test]
5401 fn fac_ctw_memory_usage_breakdown_reports_bounded_log_cache_component() {
5402 reset_shared_log_cache_for_test();
5403 let _limit = ScopedLogCacheLimit::set(2);
5404 let _overflow_slots = ScopedLogOverflowCacheSlots::set(4);
5405 let mut fac = FacContextTree::new(7, 8);
5406 for &byte in b"bounded log cache memory component payload" {
5407 fac.update_byte_msb(byte);
5408 }
5409
5410 let usage = fac.memory_usage_breakdown();
5411 let (log_int_len, log_half_len) = shared_log_cache_lens();
5412 let (overflow_log_int_len, overflow_log_half_len) = shared_log_overflow_cache_lens();
5413 assert!(log_int_len <= 3, "log_int_len={log_int_len}");
5414 assert!(log_half_len <= 3, "log_half_len={log_half_len}");
5415 assert_eq!(overflow_log_int_len, 4);
5416 assert_eq!(overflow_log_half_len, 4);
5417 assert!(
5418 usage.shared_log_cache_bytes <= 16 * size_of::<f64>() + 32 * size_of::<LogCacheSlot>()
5419 );
5420 assert_eq!(fac.memory_usage(), usage.total_bytes());
5421 }
5422
5423 #[test]
5424 fn bounded_log_lookup_preserves_context_tree_updates_and_reverts() {
5425 let mut bounded = ContextTree::new(9);
5426 let mut unbounded = bounded.clone();
5427 let stream = b"bounded exact log lookup context-tree parity payload";
5428
5429 {
5430 let _limit = ScopedLogCacheLimit::set(3);
5431 for &byte in stream {
5432 for bit_idx in 0..8usize {
5433 let bit = ((byte >> (7 - bit_idx)) & 1) == 1;
5434 bounded.update(bit);
5435 }
5436 }
5437 }
5438
5439 for &byte in stream {
5440 for bit_idx in 0..8usize {
5441 let bit = ((byte >> (7 - bit_idx)) & 1) == 1;
5442 unbounded.update(bit);
5443 }
5444 }
5445
5446 assert_eq!(
5447 bounded.get_log_block_probability().to_bits(),
5448 unbounded.get_log_block_probability().to_bits()
5449 );
5450 for &sym in &[false, true] {
5451 assert_eq!(
5452 bounded.predict(sym).to_bits(),
5453 unbounded.predict(sym).to_bits()
5454 );
5455 }
5456
5457 {
5458 let _limit = ScopedLogCacheLimit::set(3);
5459 for _ in 0..16usize {
5460 bounded.revert();
5461 }
5462 }
5463 for _ in 0..16usize {
5464 unbounded.revert();
5465 }
5466
5467 assert_eq!(
5468 bounded.get_log_block_probability().to_bits(),
5469 unbounded.get_log_block_probability().to_bits()
5470 );
5471 for &sym in &[false, true] {
5472 assert_eq!(
5473 bounded.predict(sym).to_bits(),
5474 unbounded.predict(sym).to_bits()
5475 );
5476 }
5477 }
5478
5479 #[test]
5480 fn bounded_log_lookup_preserves_fac_byte_fast_paths() {
5481 let mut bounded_msb = FacContextTree::new(7, 8);
5482 let mut unbounded_msb = bounded_msb.clone();
5483 let mut bounded_lsb = FacContextTree::new(7, 5);
5484 let mut unbounded_lsb = bounded_lsb.clone();
5485 let stream = b"bounded exact log lookup fac byte parity payload";
5486
5487 {
5488 let _limit = ScopedLogCacheLimit::set(2);
5489 for &byte in stream {
5490 bounded_msb.update_byte_msb(byte);
5491 bounded_lsb.update_byte_lsb(byte);
5492 }
5493 }
5494 for &byte in stream {
5495 unbounded_msb.update_byte_msb(byte);
5496 unbounded_lsb.update_byte_lsb(byte);
5497 }
5498
5499 assert_eq!(
5500 bounded_msb.get_log_block_probability().to_bits(),
5501 unbounded_msb.get_log_block_probability().to_bits()
5502 );
5503 assert_eq!(
5504 bounded_lsb.get_log_block_probability().to_bits(),
5505 unbounded_lsb.get_log_block_probability().to_bits()
5506 );
5507 for bit_idx in 0..bounded_msb.num_bits() {
5508 assert_eq!(
5509 bounded_msb.predict(false, bit_idx).to_bits(),
5510 unbounded_msb.predict(false, bit_idx).to_bits()
5511 );
5512 assert_eq!(
5513 bounded_msb.predict_one(bit_idx).to_bits(),
5514 unbounded_msb.predict_one(bit_idx).to_bits()
5515 );
5516 }
5517 for bit_idx in 0..bounded_lsb.num_bits() {
5518 assert_eq!(
5519 bounded_lsb.predict(false, bit_idx).to_bits(),
5520 unbounded_lsb.predict(false, bit_idx).to_bits()
5521 );
5522 assert_eq!(
5523 bounded_lsb.predict_one(bit_idx).to_bits(),
5524 unbounded_lsb.predict_one(bit_idx).to_bits()
5525 );
5526 }
5527 }
5528
5529 #[test]
5530 fn context_tree_singleton_paths_use_hot_prefix_nodes() {
5531 let mut tree = ContextTree::new(13);
5532 tree.update(false);
5533
5534 let hot_prefix_depth = tree.engine.hot_prefix_depth();
5535 let child = tree.engine.arena.child(tree.engine.root, 0);
5536 let mut current = child.as_node().expect("hot-prefix node");
5537 let mut visited_hot_prefix_nodes = 1usize;
5538 for depth in 1..hot_prefix_depth {
5539 let next = tree.engine.arena.child(current, 0);
5540 current = next.as_node().expect("next hot-prefix node");
5541 visited_hot_prefix_nodes += 1;
5542 assert!(depth < hot_prefix_depth);
5543 }
5544 assert_eq!(visited_hot_prefix_nodes, hot_prefix_depth);
5545 let segment = tree
5546 .engine
5547 .arena
5548 .child(current, 0)
5549 .as_segment()
5550 .expect("segment tail");
5551 assert!(tree.engine.arena.child(current, 1).is_none());
5552 assert!(tree.engine.arena.segments[segment.get()].tail.is_none());
5553 assert_close(
5554 tree.engine.arena.segments[segment.get()].head_log_prob_weighted,
5555 -std::f64::consts::LN_2,
5556 );
5557 assert_close(tree.get_log_block_probability(), -std::f64::consts::LN_2);
5558 }
5559
5560 #[test]
5561 fn context_tree_missing_path_tail_uses_exact_segment_payloads() {
5562 let mut tree = ContextTree::new(13);
5563 tree.update(true);
5564 let child = tree.engine.arena.child(tree.engine.root, 0);
5565 let mut current = child.as_node().expect("hot-prefix node");
5566 for _ in 1..tree.engine.hot_prefix_depth() {
5567 current = tree
5568 .engine
5569 .arena
5570 .child(current, 0)
5571 .as_node()
5572 .expect("next hot-prefix node");
5573 }
5574 let segment = tree
5575 .engine
5576 .arena
5577 .child(current, 0)
5578 .as_segment()
5579 .expect("segment tail");
5580 let payload = tree.engine.arena.segments[segment.get()].payload;
5581 assert!(payload.is_exact());
5582 assert_eq!(
5583 payload.len() as usize,
5584 tree.engine.max_depth - tree.engine.hot_prefix_depth()
5585 );
5586 assert_eq!(payload.exact_bits() & low_bits_mask_u64(payload.len()), 0);
5587 }
5588
5589 #[test]
5590 fn context_tree_missing_path_tail_uses_const_payload_beyond_exact_limit() {
5591 let mut tree = ContextTree::new(80);
5592 let history_before = tree.history.clone();
5593 tree.update(false);
5594
5595 let segment = child_after_hot_prefix(&tree, &history_before)
5596 .as_segment()
5597 .expect("segment tail");
5598 let segment = tree.engine.arena.segments[segment.get()];
5599 assert_eq!(segment.payload.mode(), SEG_MODE_CONST);
5600 assert_eq!(
5601 segment.payload.len() as usize,
5602 tree.engine.max_depth - tree.engine.hot_prefix_depth()
5603 );
5604 assert!(!segment.payload.const_bit());
5605 assert!(segment.tail.is_none());
5606 }
5607
5608 #[test]
5609 fn context_tree_missing_path_tail_uses_history_and_const_payloads_beyond_exact_limit() {
5610 let mut tree = ContextTree::new(80);
5611 let seeded_history: Vec<Symbol> = (0..80).map(|i| (i & 1) == 1).collect();
5612 tree.update_history(&seeded_history);
5613 let history_before = tree.history.clone();
5614 tree.update(false);
5615
5616 let first_segment = child_after_hot_prefix(&tree, &history_before)
5617 .as_segment()
5618 .expect("history-backed segment tail");
5619 let first_segment = tree.engine.arena.segments[first_segment.get()];
5620 assert_eq!(first_segment.payload.mode(), SEG_MODE_HISTORY);
5621 assert_eq!(
5622 first_segment.payload.len() as usize,
5623 tree.engine.max_depth - tree.engine.hot_prefix_depth() - 1
5624 );
5625 for offset in [0usize, 1, 7, 31, 66] {
5626 assert_eq!(
5627 segment_edge_from_parts(
5628 first_segment,
5629 offset,
5630 &history_before,
5631 history_before.len()
5632 ),
5633 history_symbol(&history_before, tree.engine.hot_prefix_depth() + 1 + offset)
5634 );
5635 }
5636
5637 let tail_segment = first_segment
5638 .tail
5639 .as_segment()
5640 .expect("constant fallback tail");
5641 let tail_segment = tree.engine.arena.segments[tail_segment.get()];
5642 assert_eq!(tail_segment.payload.mode(), SEG_MODE_CONST);
5643 assert_eq!(tail_segment.payload.len(), 1);
5644 assert!(!tail_segment.payload.const_bit());
5645 assert!(tail_segment.tail.is_none());
5646 }
5647
5648 #[test]
5649 fn context_tree_matches_reference_on_short_sequences() {
5650 for depth in 0..=6usize {
5651 for len in 0..=6usize {
5652 for mask in 0..(1usize << len) {
5653 let mut prod = ContextTree::new(depth);
5654 let mut reference = RefContextTree::new(depth);
5655 for step in 0..len {
5656 let p_prod_0 = prod.predict(false);
5657 let p_ref_0 = reference.predict(false);
5658 assert!(
5659 (p_prod_0 - p_ref_0).abs()
5660 <= 1e-12 * p_prod_0.abs().max(p_ref_0.abs()).max(1.0),
5661 "predict0 mismatch depth={depth} len={len} mask={mask} step={step} prod={p_prod_0} ref={p_ref_0} history={:?}",
5662 prod.history
5663 );
5664 let p_prod_1 = prod.predict(true);
5665 let p_ref_1 = reference.predict(true);
5666 assert!(
5667 (p_prod_1 - p_ref_1).abs()
5668 <= 1e-12 * p_prod_1.abs().max(p_ref_1.abs()).max(1.0),
5669 "predict1 mismatch depth={depth} len={len} mask={mask} step={step} prod={p_prod_1} ref={p_ref_1} history={:?}",
5670 prod.history
5671 );
5672 let log_prod = prod.get_log_block_probability();
5673 let log_ref = reference.get_log_block_probability();
5674 assert!(
5675 (log_prod - log_ref).abs()
5676 <= 1e-12 * log_prod.abs().max(log_ref.abs()).max(1.0),
5677 "log mismatch before update depth={depth} len={len} mask={mask} step={step} prod={log_prod} ref={log_ref} history={:?}",
5678 prod.history
5679 );
5680 let bit = ((mask >> step) & 1) == 1;
5681 prod.update(bit);
5682 reference.update(bit);
5683 let log_prod = prod.get_log_block_probability();
5684 let log_ref = reference.get_log_block_probability();
5685 assert!(
5686 (log_prod - log_ref).abs()
5687 <= 1e-12 * log_prod.abs().max(log_ref.abs()).max(1.0),
5688 "log mismatch after update depth={depth} len={len} mask={mask} step={step} bit={bit} prod={log_prod} ref={log_ref} history={:?}",
5689 prod.history
5690 );
5691 }
5692 while prod.history_size() > 0 {
5693 let p_prod_0 = prod.predict(false);
5694 let p_ref_0 = reference.predict(false);
5695 assert!(
5696 (p_prod_0 - p_ref_0).abs()
5697 <= 1e-12 * p_prod_0.abs().max(p_ref_0.abs()).max(1.0),
5698 "revert predict0 mismatch depth={depth} len={len} mask={mask} prod={p_prod_0} ref={p_ref_0} history={:?}",
5699 prod.history
5700 );
5701 let p_prod_1 = prod.predict(true);
5702 let p_ref_1 = reference.predict(true);
5703 assert!(
5704 (p_prod_1 - p_ref_1).abs()
5705 <= 1e-12 * p_prod_1.abs().max(p_ref_1.abs()).max(1.0),
5706 "revert predict1 mismatch depth={depth} len={len} mask={mask} prod={p_prod_1} ref={p_ref_1} history={:?}",
5707 prod.history
5708 );
5709 prod.revert();
5710 reference.revert();
5711 let log_prod = prod.get_log_block_probability();
5712 let log_ref = reference.get_log_block_probability();
5713 assert!(
5714 (log_prod - log_ref).abs()
5715 <= 1e-12 * log_prod.abs().max(log_ref.abs()).max(1.0),
5716 "revert log mismatch depth={depth} len={len} mask={mask} prod={log_prod} ref={log_ref} history={:?}",
5717 prod.history
5718 );
5719 }
5720 }
5721 }
5722 }
5723 }
5724
5725 #[test]
5726 fn context_tree_long_depth_matches_reference_on_short_sequences() {
5727 for &depth in &[65usize, 80usize] {
5728 for len in 0..=6usize {
5729 for mask in 0..(1usize << len) {
5730 let mut prod = ContextTree::new(depth);
5731 let mut reference = RefContextTree::new(depth);
5732 for step in 0..len {
5733 let p_prod_0 = prod.predict(false);
5734 let p_ref_0 = reference.predict(false);
5735 assert!(
5736 (p_prod_0 - p_ref_0).abs()
5737 <= 1e-12 * p_prod_0.abs().max(p_ref_0.abs()).max(1.0),
5738 "long-depth predict0 mismatch depth={depth} len={len} mask={mask} step={step} prod={p_prod_0} ref={p_ref_0} history={:?}",
5739 prod.history
5740 );
5741 let p_prod_1 = prod.predict(true);
5742 let p_ref_1 = reference.predict(true);
5743 assert!(
5744 (p_prod_1 - p_ref_1).abs()
5745 <= 1e-12 * p_prod_1.abs().max(p_ref_1.abs()).max(1.0),
5746 "long-depth predict1 mismatch depth={depth} len={len} mask={mask} step={step} prod={p_prod_1} ref={p_ref_1} history={:?}",
5747 prod.history
5748 );
5749 assert_close(
5750 prod.get_log_block_probability(),
5751 reference.get_log_block_probability(),
5752 );
5753 let bit = ((mask >> step) & 1) == 1;
5754 prod.update(bit);
5755 reference.update(bit);
5756 assert_close(
5757 prod.get_log_block_probability(),
5758 reference.get_log_block_probability(),
5759 );
5760 }
5761
5762 while prod.history_size() > 0 {
5763 assert_close(prod.predict(false), reference.predict(false));
5764 assert_close(prod.predict(true), reference.predict(true));
5765 prod.revert();
5766 reference.revert();
5767 assert_close(
5768 prod.get_log_block_probability(),
5769 reference.get_log_block_probability(),
5770 );
5771 }
5772 }
5773 }
5774 }
5775 }
5776
5777 #[test]
5778 fn fac_ctw_matches_reference_on_short_sequences() {
5779 let mut fac = FacContextTree::new(4, 4);
5780 let mut reference = RefFacContextTree::new(4, 4);
5781 let stream = [
5782 (true, 0usize),
5783 (false, 1usize),
5784 (true, 2usize),
5785 (true, 3usize),
5786 (false, 0usize),
5787 (false, 1usize),
5788 (true, 2usize),
5789 (false, 3usize),
5790 ];
5791
5792 for &(bit, idx) in &stream {
5793 assert_close(fac.predict(false, idx), reference.predict(false, idx));
5794 assert_close(fac.predict(true, idx), reference.predict(true, idx));
5795 fac.update(bit, idx);
5796 reference.update(bit, idx);
5797 assert_close(
5798 fac.get_log_block_probability(),
5799 reference.get_log_block_probability(),
5800 );
5801 }
5802
5803 for &(_, idx) in stream.iter().rev() {
5804 fac.revert(idx);
5805 reference.revert(idx);
5806 assert_close(
5807 fac.get_log_block_probability(),
5808 reference.get_log_block_probability(),
5809 );
5810 }
5811 }
5812
5813 #[test]
5814 fn fac_ctw_history_consistency() {
5815 let mut fac = FacContextTree::new(4, 4);
5816
5817 fac.update_history(&[true, false, true]);
5818 assert_eq!(fac.shared_history.len(), 3);
5819
5820 fac.update(true, 0);
5821 fac.update(false, 1);
5822 assert_eq!(fac.shared_history.len(), 5);
5823
5824 fac.revert(1);
5825 assert_eq!(fac.shared_history.len(), 4);
5826
5827 fac.revert(0);
5828 assert_eq!(fac.shared_history.len(), 3);
5829 }
5830
5831 #[test]
5832 fn fac_ctw_predict_one_matches_predict_true() {
5833 let mut fac = FacContextTree::new(6, 8);
5834 for &byte in b"predict-one exactness regression payload" {
5835 for bit_idx in 0..8usize {
5836 let p_generic = fac.predict(true, bit_idx);
5837 let p_one = fac.predict_one(bit_idx);
5838 assert_close(p_generic, p_one);
5839 let bit = ((byte >> (7 - bit_idx)) & 1) == 1;
5840 fac.update_predicted(bit, bit_idx);
5841 }
5842 }
5843 }
5844
5845 #[test]
5846 fn fac_ctw_long_depth_predict_one_matches_predict_true() {
5847 let mut fac = FacContextTree::new(78, 4);
5848 for step in 0..24usize {
5849 for bit_idx in 0..fac.num_bits() {
5850 let p_generic = fac.predict(true, bit_idx);
5851 let p_one = fac.predict_one(bit_idx);
5852 assert_close(p_generic, p_one);
5853 let bit = ((step * 5 + bit_idx * 3) & 1) == 1;
5854 fac.update_predicted(bit, bit_idx);
5855 }
5856 }
5857 }
5858
5859 #[test]
5860 fn fac_ctw_update_byte_msb_matches_bit_updates() {
5861 let mut by_byte = FacContextTree::new(6, 8);
5862 let mut by_bits = FacContextTree::new(6, 8);
5863 for &byte in b"byte update msb regression payload" {
5864 by_byte.update_byte_msb(byte);
5865 for bit_idx in 0..8usize {
5866 let bit = ((byte >> (7 - bit_idx)) & 1) == 1;
5867 by_bits.update(bit, bit_idx);
5868 }
5869 assert_close(
5870 by_byte.get_log_block_probability(),
5871 by_bits.get_log_block_probability(),
5872 );
5873 assert_eq!(by_byte.shared_history, by_bits.shared_history);
5874 }
5875 }
5876
5877 #[test]
5878 fn fac_ctw_log_prob_update_byte_msb_matches_manual_fast_path() {
5879 let mut batched = FacContextTree::new(6, 8);
5880 for &byte in b"log prob update byte msb regression payload" {
5881 let mut manual = batched.clone();
5882 let observed = batched.log_prob_update_byte_msb(byte);
5883 let mut expected = 0.0;
5884 for bit_idx in 0..8usize {
5885 let bit = ((byte >> (7 - bit_idx)) & 1) == 1;
5886 let p = manual.predict(bit, bit_idx);
5887 if p.is_finite() && p > 0.0 {
5888 expected += p.ln();
5889 } else {
5890 expected = f64::NEG_INFINITY;
5891 }
5892 manual.update_predicted(bit, bit_idx);
5893 }
5894 assert_eq!(observed.to_bits(), expected.to_bits());
5895 assert_eq!(batched.shared_history, manual.shared_history);
5896 assert_eq!(
5897 batched.get_log_block_probability().to_bits(),
5898 manual.get_log_block_probability().to_bits(),
5899 );
5900 }
5901 }
5902
5903 #[test]
5904 fn fac_ctw_update_byte_lsb_matches_bit_updates() {
5905 let mut by_byte = FacContextTree::new(6, 5);
5906 let mut by_bits = FacContextTree::new(6, 5);
5907 for &byte in b"byte update lsb regression payload" {
5908 by_byte.update_byte_lsb(byte);
5909 for bit_idx in 0..5usize {
5910 let bit = ((byte >> bit_idx) & 1) == 1;
5911 by_bits.update(bit, bit_idx);
5912 }
5913 assert_close(
5914 by_byte.get_log_block_probability(),
5915 by_bits.get_log_block_probability(),
5916 );
5917 assert_eq!(by_byte.shared_history, by_bits.shared_history);
5918 }
5919 }
5920
5921 #[test]
5922 fn fac_ctw_log_cache_tracks_tree_visits_not_shared_history() {
5923 let mut fac = FacContextTree::new(8, 8);
5924 let updates_per_tree = 512usize;
5925 let (log_int_before, log_half_before) = shared_log_cache_lens();
5926
5927 for step in 0..updates_per_tree {
5928 let bit = (step & 1) == 1;
5929 for bit_idx in 0..8usize {
5930 fac.update(bit, bit_idx);
5931 }
5932 }
5933
5934 assert_eq!(fac.shared_history.len(), updates_per_tree * 8);
5935 for tree in &fac.trees {
5936 let visits = tree.engine.arena.visits(tree.engine.root) as usize;
5937 assert_eq!(visits, updates_per_tree);
5938 }
5939
5940 let (log_int_after, log_half_after) = shared_log_cache_lens();
5941 let expected_len = updates_per_tree + 1;
5942 assert!(
5943 log_int_after <= log_int_before.max(expected_len),
5944 "log_int grew to {log_int_after} (before={log_int_before}, expected_len={expected_len})"
5945 );
5946 assert!(
5947 log_half_after <= log_half_before.max(expected_len),
5948 "log_half grew to {log_half_after} (before={log_half_before}, expected_len={expected_len})"
5949 );
5950 }
5951
5952 fn seed_fac_cache_regression_state(fac: &mut FacContextTree) {
5953 for step in 0..24usize {
5954 for bit_idx in 0..fac.num_bits() {
5955 let bit = ((step * 3 + bit_idx) & 1) == 1;
5956 fac.update(bit, bit_idx);
5957 }
5958 }
5959 }
5960
5961 fn assert_update_predicted_matches_fresh_after_history_rewrite<F>(mut rewrite: F)
5962 where
5963 F: FnMut(&mut FacContextTree),
5964 {
5965 let mut predicted = FacContextTree::new(6, 4);
5966 seed_fac_cache_regression_state(&mut predicted);
5967 let mut fresh = predicted.clone();
5968 let original_history = predicted.shared_history.clone();
5969 let target_bit = 2usize;
5970
5971 let _ = predicted.predict(true, target_bit);
5972 rewrite(&mut predicted);
5973 rewrite(&mut fresh);
5974
5975 assert_eq!(predicted.shared_history.len(), original_history.len());
5976 assert_ne!(predicted.shared_history, original_history);
5977 assert_eq!(predicted.shared_history, fresh.shared_history);
5978
5979 predicted.update_predicted(false, target_bit);
5980 fresh.update(false, target_bit);
5981
5982 assert_eq!(predicted.shared_history, fresh.shared_history);
5983 assert_close(
5984 predicted.get_log_block_probability(),
5985 fresh.get_log_block_probability(),
5986 );
5987 for bit_idx in 0..predicted.num_bits() {
5988 assert_close(
5989 predicted.predict(false, bit_idx),
5990 fresh.predict(false, bit_idx),
5991 );
5992 assert_close(
5993 predicted.predict(true, bit_idx),
5994 fresh.predict(true, bit_idx),
5995 );
5996 }
5997 }
5998
5999 #[test]
6000 fn fac_ctw_update_predicted_ignores_stale_cache_after_reset_and_rewrite() {
6001 assert_update_predicted_matches_fresh_after_history_rewrite(|fac| {
6002 let mut rewritten = fac.shared_history.to_vec();
6003 for bit in &mut rewritten {
6004 *bit = !*bit;
6005 }
6006 fac.reset_history_only();
6007 fac.update_history(&rewritten);
6008 });
6009 }
6010
6011 #[test]
6012 fn fac_ctw_update_predicted_ignores_stale_cache_after_revert_and_rewrite() {
6013 assert_update_predicted_matches_fresh_after_history_rewrite(|fac| {
6014 let original = fac.shared_history.to_vec();
6015 let keep = original.len() / 3;
6016 let remove = original.len() - keep;
6017 let mut rewritten_suffix = original[keep..].to_vec();
6018 for bit in &mut rewritten_suffix {
6019 *bit = !*bit;
6020 }
6021 fac.revert_history(remove);
6022 fac.update_history(&rewritten_suffix);
6023 });
6024 }
6025
6026 #[test]
6027 fn fac_ctw_shared_history_version_tracks_mutations() {
6028 let mut fac = FacContextTree::new(4, 2);
6029 let mut version = fac.shared_history_version;
6030
6031 fac.update_history(&[]);
6032 assert_eq!(fac.shared_history_version, version);
6033
6034 fac.update_history(&[true, false]);
6035 assert_ne!(fac.shared_history_version, version);
6036 version = fac.shared_history_version;
6037
6038 fac.revert_history(0);
6039 assert_eq!(fac.shared_history_version, version);
6040
6041 fac.revert_history(1);
6042 assert_ne!(fac.shared_history_version, version);
6043 version = fac.shared_history_version;
6044
6045 let _ = fac.predict(true, 0);
6046 assert_eq!(fac.shared_history_version, version);
6047
6048 fac.update_predicted(true, 0);
6049 assert_ne!(fac.shared_history_version, version);
6050 version = fac.shared_history_version;
6051
6052 fac.reset_history_only();
6053 assert_ne!(fac.shared_history_version, version);
6054 }
6055
6056 #[test]
6057 fn context_tree_predict_preserves_state() {
6058 let mut tree = ContextTree::new(6);
6059 for &bit in &[true, false, true, true, false, false, true, false] {
6060 tree.update(bit);
6061 }
6062 let p0_before = tree.predict(false);
6063 let p1_before = tree.predict(true);
6064 let log_before = tree.get_log_block_probability();
6065 let history_before = tree.history.clone();
6066 let _ = tree.predict(true);
6067
6068 assert_eq!(tree.history, history_before);
6069 assert_close(tree.get_log_block_probability(), log_before);
6070 assert_close(tree.predict(false), p0_before);
6071 assert_close(tree.predict(true), p1_before);
6072 }
6073
6074 #[test]
6075 fn context_tree_predict_matches_update_ratio() {
6076 let mut tree = ContextTree::new(7);
6077 for &bit in &[true, false, true, false, true, true, false, true, false] {
6078 tree.update(bit);
6079 }
6080 for &sym in &[false, true] {
6081 let predicted = tree.predict(sym);
6082 let mut reference = tree.clone();
6083 let before = reference.get_log_block_probability();
6084 reference.update(sym);
6085 let after = reference.get_log_block_probability();
6086 assert_close(predicted, (after - before).exp());
6087 }
6088 }
6089
6090 #[test]
6091 fn fac_ctw_predict_preserves_state() {
6092 let mut fac = FacContextTree::new(5, 8);
6093 for &byte in b"fac ctw state preservation" {
6094 for bit_idx in 0..8usize {
6095 let bit = ((byte >> (7 - bit_idx)) & 1) == 1;
6096 fac.update(bit, bit_idx);
6097 }
6098 }
6099 let p0_before = fac.predict(false, 3);
6100 let p1_before = fac.predict(true, 3);
6101 let log_before = fac.get_log_block_probability();
6102 let history_before = fac.shared_history.clone();
6103 let _ = fac.predict(true, 3);
6104
6105 assert_eq!(fac.shared_history, history_before);
6106 assert_close(fac.get_log_block_probability(), log_before);
6107 assert_close(fac.predict(false, 3), p0_before);
6108 assert_close(fac.predict(true, 3), p1_before);
6109 }
6110
6111 #[test]
6112 fn ctw_inline_node_weight_recompute_matches_recompute_node_weight() {
6113 fn inline_weight(arena: &CtArena, idx: NodeIndex) -> f64 {
6114 let slot = idx.get();
6115 let node = arena.nodes[slot];
6116 let [left, right] = node.children;
6117 if left.is_none() && right.is_none() {
6118 clamp_log_prob(node.log_prob_kt)
6119 } else {
6120 let w0 = arena.child_ref_weighted(left);
6121 let w1 = arena.child_ref_weighted(right);
6122 update_weighted_log_prob_non_leaf(node.log_prob_kt, w0, w1)
6123 }
6124 }
6125
6126 let mut arena = CtArena::new();
6127 let parent = arena.alloc_node_with_state([3, 5], -1.75);
6128 let left_node = arena.alloc_node_with_state([2, 1], -0.25);
6129 let right_node = arena.alloc_node_with_state([1, 2], -0.50);
6130 arena.nodes[left_node.get()].log_prob_weighted = -0.333_333_333_f64;
6131 arena.nodes[right_node.get()].log_prob_weighted = -0.777_777_777_f64;
6132
6133 let seg = arena.alloc_segment();
6134 arena.segments[seg.get()].head_log_prob_weighted = -0.125_f64;
6135
6136 let cases = [
6137 (ChildRef::NONE, ChildRef::NONE, -2.0_f64),
6138 (
6139 ChildRef::from_node(left_node),
6140 ChildRef::from_node(right_node),
6141 -1.0_f64,
6142 ),
6143 (
6144 ChildRef::from_node(left_node),
6145 ChildRef::from_segment(seg),
6146 -0.625_f64,
6147 ),
6148 (
6149 ChildRef::from_segment(seg),
6150 ChildRef::from_node(right_node),
6151 -0.3125_f64,
6152 ),
6153 ];
6154
6155 for (left, right, kt) in cases {
6156 let slot = parent.get();
6157 arena.nodes[slot].children = [left, right];
6158 arena.nodes[slot].log_prob_kt = kt;
6159 arena.nodes[slot].log_prob_weighted = f64::NAN;
6160
6161 let expected = inline_weight(&arena, parent);
6162 arena.recompute_node_weight(parent);
6163 let actual = arena.nodes[slot].log_prob_weighted;
6164 assert_eq!(
6165 actual.to_bits(),
6166 expected.to_bits(),
6167 "node-weight recompute mismatch for children=({left:?},{right:?}) kt={kt}"
6168 );
6169 }
6170 }
6171
6172 #[test]
6173 fn fac_ctw_predict_matches_update_ratio() {
6174 let mut fac = FacContextTree::new(6, 8);
6175 for &byte in b"fac ctw exact predictive ratio" {
6176 for bit_idx in 0..8usize {
6177 let bit = ((byte >> (7 - bit_idx)) & 1) == 1;
6178 fac.update(bit, bit_idx);
6179 }
6180 }
6181 for &sym in &[false, true] {
6182 let predicted = fac.predict(sym, 4);
6183 let mut reference = fac.clone();
6184 let before = reference.get_log_block_probability();
6185 reference.update(sym, 4);
6186 let after = reference.get_log_block_probability();
6187 assert_close(predicted, (after - before).exp());
6188 }
6189 }
6190
6191 #[test]
6192 fn fac_ctw_update_predicted_matches_fresh_update_on_byte_stream() {
6193 let mut predicted = FacContextTree::new(6, 8);
6194 let mut fresh = predicted.clone();
6195 let stream = b"fac-ctw prepared update exactness regression";
6196
6197 for (byte_pos, &byte) in stream.iter().enumerate() {
6198 for bit_idx in 0..8usize {
6199 let bit = ((byte >> (7 - bit_idx)) & 1) == 1;
6200 let _ = predicted.predict(true, bit_idx);
6201 predicted.update_predicted(bit, bit_idx);
6202 fresh.update(bit, bit_idx);
6203
6204 let predicted_log = predicted.get_log_block_probability();
6205 let fresh_log = fresh.get_log_block_probability();
6206 assert!(
6207 (predicted_log - fresh_log).abs()
6208 <= 1e-12 * predicted_log.abs().max(fresh_log.abs()).max(1.0),
6209 "log mismatch byte_pos={byte_pos} bit_idx={bit_idx} bit={bit} predicted={predicted_log} fresh={fresh_log}\nshared_history={:?}\npredicted_arena={:#?}\nfresh_arena={:#?}\npredicted_steps={:?}\nfresh_steps={:?}",
6210 predicted.shared_history,
6211 predicted.trees[bit_idx].engine.arena,
6212 fresh.trees[bit_idx].engine.arena,
6213 predicted.trees[bit_idx].engine.prepared_steps,
6214 fresh.trees[bit_idx].engine.prepared_steps,
6215 );
6216 for probe_idx in 0..8usize {
6217 let p_pred_0 = predicted.predict(false, probe_idx);
6218 let p_fresh_0 = fresh.predict(false, probe_idx);
6219 assert!(
6220 (p_pred_0 - p_fresh_0).abs()
6221 <= 1e-12 * p_pred_0.abs().max(p_fresh_0.abs()).max(1.0),
6222 "predict0 mismatch byte_pos={byte_pos} bit_idx={bit_idx} probe_idx={probe_idx} predicted={p_pred_0} fresh={p_fresh_0}",
6223 );
6224 let p_pred_1 = predicted.predict(true, probe_idx);
6225 let p_fresh_1 = fresh.predict(true, probe_idx);
6226 assert!(
6227 (p_pred_1 - p_fresh_1).abs()
6228 <= 1e-12 * p_pred_1.abs().max(p_fresh_1.abs()).max(1.0),
6229 "predict1 mismatch byte_pos={byte_pos} bit_idx={bit_idx} probe_idx={probe_idx} predicted={p_pred_1} fresh={p_fresh_1}",
6230 );
6231 }
6232 }
6233 }
6234 }
6235
6236 #[test]
6237 fn fac_ctw_long_depth_update_predicted_matches_fresh_update_on_bit_stream() {
6238 let mut predicted = FacContextTree::new(78, 4);
6239 let mut fresh = predicted.clone();
6240
6241 for step in 0..20usize {
6242 for bit_idx in 0..predicted.num_bits() {
6243 let bit = ((step * 7 + bit_idx * 11) & 1) == 1;
6244 let _ = predicted.predict(true, bit_idx);
6245 predicted.update_predicted(bit, bit_idx);
6246 fresh.update(bit, bit_idx);
6247 assert_eq!(predicted.shared_history, fresh.shared_history);
6248 assert_close(
6249 predicted.get_log_block_probability(),
6250 fresh.get_log_block_probability(),
6251 );
6252 }
6253 }
6254
6255 for bit_idx in 0..predicted.num_bits() {
6256 assert_close(
6257 predicted.predict(false, bit_idx),
6258 fresh.predict(false, bit_idx),
6259 );
6260 assert_close(
6261 predicted.predict(true, bit_idx),
6262 fresh.predict(true, bit_idx),
6263 );
6264 }
6265 }
6266
6267 fn scan_symbol_space(tree: &mut FacContextTree, bits: usize) {
6268 fn rec(tree: &mut FacContextTree, bits: usize, depth: usize) {
6269 if depth == bits {
6270 return;
6271 }
6272 for bit in [false, true] {
6273 let bit_idx = depth;
6274 tree.update(bit, bit_idx);
6275 rec(tree, bits, depth + 1);
6276 tree.revert(bit_idx);
6277 }
6278 }
6279 rec(tree, bits, 0);
6280 }
6281
6282 fn fac_symbol_bit(symbol: u8, msb_first: bool, bits: usize, bit_idx: usize) -> bool {
6283 if msb_first {
6284 ctw_symbol_bit_msb(symbol, bits, bit_idx)
6285 } else {
6286 ((symbol >> bit_idx) & 1) == 1
6287 }
6288 }
6289
6290 fn symbol_log_prob(tree: &mut FacContextTree, symbol: u8, msb_first: bool, bits: usize) -> f64 {
6291 let before = tree.get_log_block_probability();
6292 for bit_idx in 0..bits {
6293 let bit = fac_symbol_bit(symbol, msb_first, bits, bit_idx);
6294 tree.update(bit, bit_idx);
6295 }
6296 let after = tree.get_log_block_probability();
6297 for bit_idx in (0..bits).rev() {
6298 tree.revert(bit_idx);
6299 }
6300 after - before
6301 }
6302
6303 fn assert_symbol_scan_then_update_matches_plain(msb_first: bool) {
6304 let bits = 8usize;
6305 let mut with_scan = FacContextTree::new(7, bits);
6306 let mut plain = with_scan.clone();
6307 for &byte in b"pdf then update parity payload" {
6308 for bit_idx in 0..bits {
6309 let bit = fac_symbol_bit(byte, msb_first, bits, bit_idx);
6310 with_scan.update(bit, bit_idx);
6311 plain.update(bit, bit_idx);
6312 }
6313 }
6314
6315 scan_symbol_space(&mut with_scan, bits);
6316
6317 let observed = b'n';
6318 for bit_idx in 0..bits {
6319 let bit = fac_symbol_bit(observed, msb_first, bits, bit_idx);
6320 with_scan.update(bit, bit_idx);
6321 plain.update(bit, bit_idx);
6322 }
6323
6324 for sym in 0u8..=255u8 {
6325 let lp_scan = symbol_log_prob(&mut with_scan, sym, msb_first, bits);
6326 let lp_plain = symbol_log_prob(&mut plain, sym, msb_first, bits);
6327 let diff = (lp_scan - lp_plain).abs();
6328 assert!(
6329 diff < 1e-12,
6330 "symbol={sym} lp_scan={lp_scan} lp_plain={lp_plain} diff={diff}",
6331 );
6332 }
6333 }
6334
6335 fn assert_fill_fac_tree_log_probs_matches_direct_symbol_probs(msb_first: bool, bits: usize) {
6336 let bits = bits.clamp(1, 8);
6337 let min_logp = 1e-12f64.ln();
6338 let patterns = 1usize << bits;
6339 let alias_ln = if bits == 8 {
6340 0.0
6341 } else {
6342 ((1usize << (8 - bits)) as f64).ln()
6343 };
6344 let training = [0x0u8, 0x3, 0x5, 0x6, 0x9, 0xA, 0xC, 0xF, 0x7, 0x1];
6345 let mut tree = FacContextTree::new(7, bits);
6346 for &symbol in &training {
6347 for bit_idx in 0..bits {
6348 tree.update(fac_symbol_bit(symbol, msb_first, bits, bit_idx), bit_idx);
6349 }
6350 }
6351
6352 let log_before = tree.get_log_block_probability();
6353 let mut predict_zero_before = vec![0.0; bits];
6354 let mut predict_one_before = vec![0.0; bits];
6355 for bit_idx in 0..bits {
6356 predict_zero_before[bit_idx] = tree.predict(false, bit_idx);
6357 predict_one_before[bit_idx] = tree.predict(true, bit_idx);
6358 }
6359
6360 let mut out = [0.0; 256];
6361 fill_fac_tree_log_probs(&mut tree, bits, msb_first, min_logp, &mut out);
6362
6363 assert_close(tree.get_log_block_probability(), log_before);
6364 for bit_idx in 0..bits {
6365 assert_close(tree.predict(false, bit_idx), predict_zero_before[bit_idx]);
6366 assert_close(tree.predict(true, bit_idx), predict_one_before[bit_idx]);
6367 }
6368
6369 let mask = patterns - 1;
6370 for (byte, &actual) in out.iter().enumerate() {
6371 let symbol = if bits == 8 {
6372 byte as u8
6373 } else {
6374 (byte & mask) as u8
6375 };
6376 let expected =
6377 symbol_log_prob(&mut tree, symbol, msb_first, bits).max(min_logp) - alias_ln;
6378 assert_close(actual, expected);
6379 }
6380 }
6381
6382 #[test]
6383 fn fac_ctw_symbol_scan_then_update_matches_plain_msb() {
6384 assert_symbol_scan_then_update_matches_plain(true);
6385 }
6386
6387 #[test]
6388 fn fac_ctw_symbol_scan_then_update_matches_plain_lsb() {
6389 assert_symbol_scan_then_update_matches_plain(false);
6390 }
6391
6392 #[test]
6393 fn fill_fac_tree_log_probs_matches_direct_symbol_probs_for_subbyte_msb() {
6394 assert_fill_fac_tree_log_probs_matches_direct_symbol_probs(true, 4);
6395 }
6396
6397 #[test]
6398 fn fill_fac_tree_log_probs_matches_direct_symbol_probs_for_subbyte_lsb() {
6399 assert_fill_fac_tree_log_probs_matches_direct_symbol_probs(false, 4);
6400 }
6401}