Skip to main content

infotheory/compression/
mod.rs

1//! Rate-coded compression helpers (AC/rANS) with optional framing.
2//!
3//! The functions in this module implement lossless byte compression by combining:
4//! - a predictive rate model (`RateBackend`) that emits per-symbol PDFs,
5//! - an entropy coder (`AC` or `rANS`),
6//! - optional framing metadata for robust decompression.
7#![cfg_attr(
8    not(feature = "all-backends"),
9    allow(dead_code, unused_imports, unused_variables, unused_mut)
10)]
11
12use anyhow::{Result, bail};
13
14use crate::api::{MixtureKind, MixtureScheduleMode};
15#[cfg(test)]
16use crate::api::{MixtureSpec, RateBackend};
17#[cfg(feature = "backend-calibrated")]
18use crate::backends::calibration::CalibratorCore;
19#[cfg(feature = "backend-ctw")]
20use crate::backends::ctw::{ContextTree, FacContextTree, ctw_symbol_bit_msb};
21#[cfg(feature = "backend-match")]
22use crate::backends::match_model::MatchModel;
23#[cfg(feature = "backend-ppmd")]
24use crate::backends::ppmd::PpmdModel;
25#[cfg(feature = "backend-rosa")]
26use crate::backends::rosaplus::RosaPlus;
27#[cfg(feature = "backend-sequitur")]
28use crate::backends::sequitur::SequiturModel;
29#[cfg(feature = "backend-match")]
30use crate::backends::sparse_match::SparseMatchModel;
31use crate::backends::text_context::TextContextAnalyzer;
32#[cfg(feature = "backend-zpaq")]
33use crate::backends::zpaq_rate::ZpaqRateModel;
34#[cfg(all(test, feature = "all-backends"))]
35use crate::byte_prefix::zeroed_prefix_cdf;
36use crate::byte_prefix::{
37    BytePrefixCdf, MsbPrefixRange, fill_prefix_cdf_from_pdf, normalize_pdf, zeroed_prefix_cdf_box,
38};
39use crate::coders::{
40    ANS_TOTAL, ArithmeticDecoder, ArithmeticEncoder, BlockedRansDecoder, BlockedRansEncoder,
41    CDF_TOTAL, Cdf, CoderType, crc32, quantize_pdf_to_rans_cdf_with_buffer,
42};
43#[cfg(feature = "backend-mamba")]
44use crate::mambazip;
45use crate::mixture::{
46    DEFAULT_MIN_PROB, convex_step_size_for_update, project_simplex_with_scratch,
47    switching_alpha_for_update,
48};
49use crate::neural_mix::NeuralMixCore;
50#[cfg(feature = "backend-rwkv")]
51use crate::rwkvzip;
52use crate::spec::CompiledRateBackend;
53use rayon::{ThreadPool, prelude::*};
54
55const FRAMED_MAGIC: u32 = 0x4354_4946; // "FITC"
56const FRAMED_VERSION: u8 = 1;
57const PDF_MIN: f64 = DEFAULT_MIN_PROB;
58const DIAGNOSTIC_PARALLEL_THRESHOLD: usize = 4;
59
60#[derive(Clone, Copy, Debug, Default, Eq, PartialEq)]
61/// Wire format mode for rate-coded payloads.
62pub enum FramingMode {
63    /// Emit only coder payload bytes (no integrity/length header).
64    Raw,
65    /// Emit framed payload with magic/version/length/checksum header.
66    #[default]
67    Framed,
68}
69
70#[derive(Clone, Copy, Debug)]
71struct FramedHeader {
72    magic: u32,
73    version: u8,
74    coder: u8,
75    original_len: u64,
76    crc32: u32,
77}
78
79impl FramedHeader {
80    const SIZE: usize = 4 + 1 + 1 + 8 + 4;
81
82    fn new(coder: CoderType, original_len: u64, crc32: u32) -> Self {
83        let coder = match coder {
84            CoderType::AC => 0,
85            CoderType::RANS => 1,
86        };
87        Self {
88            magic: FRAMED_MAGIC,
89            version: FRAMED_VERSION,
90            coder,
91            original_len,
92            crc32,
93        }
94    }
95
96    fn write(&self, out: &mut Vec<u8>) {
97        out.extend_from_slice(&self.magic.to_le_bytes());
98        out.push(self.version);
99        out.push(self.coder);
100        out.extend_from_slice(&self.original_len.to_le_bytes());
101        out.extend_from_slice(&self.crc32.to_le_bytes());
102    }
103
104    fn read(input: &[u8]) -> Result<Self> {
105        if input.len() < Self::SIZE {
106            bail!("framed payload too short");
107        }
108        let magic = u32::from_le_bytes([input[0], input[1], input[2], input[3]]);
109        if magic != FRAMED_MAGIC {
110            bail!("invalid framed magic: expected 0x{FRAMED_MAGIC:08X}, got 0x{magic:08X}");
111        }
112        let version = input[4];
113        if version != FRAMED_VERSION {
114            bail!("unsupported framed version: {version}");
115        }
116        let coder = input[5];
117        let original_len = u64::from_le_bytes([
118            input[6], input[7], input[8], input[9], input[10], input[11], input[12], input[13],
119        ]);
120        let crc32 = u32::from_le_bytes([input[14], input[15], input[16], input[17]]);
121        Ok(Self {
122            magic,
123            version,
124            coder,
125            original_len,
126            crc32,
127        })
128    }
129
130    fn coder_type(&self) -> CoderType {
131        match self.coder {
132            0 => CoderType::AC,
133            _ => CoderType::RANS,
134        }
135    }
136}
137
138#[cfg(feature = "backend-ctw")]
139#[derive(Clone)]
140enum CtwCompressionTree {
141    Ac(Box<ContextTree>),
142    Fac(FacContextTree),
143}
144
145#[cfg(feature = "backend-ctw")]
146#[derive(Clone)]
147pub(crate) struct CtwPredictor {
148    tree: CtwCompressionTree,
149    bits_per_symbol: usize,
150    msb_first: bool,
151    pdf: Vec<f64>,
152    pattern_logps: Vec<f64>,
153    valid: bool,
154}
155
156#[cfg(feature = "backend-ctw")]
157impl CtwPredictor {
158    pub(crate) fn new_ctw(depth: usize) -> Self {
159        Self {
160            tree: CtwCompressionTree::Ac(Box::new(ContextTree::new(depth))),
161            bits_per_symbol: 8,
162            msb_first: true,
163            pdf: vec![0.0; 256],
164            pattern_logps: vec![f64::NEG_INFINITY; 256],
165            valid: false,
166        }
167    }
168
169    pub(crate) fn new_fac(
170        base_depth: usize,
171        bits_per_symbol: usize,
172        msb_first: Option<bool>,
173    ) -> Self {
174        let effective_msb_first: bool = msb_first.unwrap_or(bits_per_symbol == 8);
175        Self {
176            tree: CtwCompressionTree::Fac(FacContextTree::new(base_depth, bits_per_symbol)),
177            bits_per_symbol,
178            msb_first: effective_msb_first,
179            pdf: vec![0.0; 256],
180            pattern_logps: vec![f64::NEG_INFINITY; 256],
181            valid: false,
182        }
183    }
184
185    fn fill_pattern_log_probs(&mut self) -> usize {
186        let bits = self.bits_per_symbol.clamp(1, 8);
187        let patterns = 1usize << bits;
188        self.pattern_logps[..patterns].fill(f64::NEG_INFINITY);
189        match &mut self.tree {
190            CtwCompressionTree::Ac(tree) => {
191                fn rec(
192                    tree: &mut ContextTree,
193                    depth: usize,
194                    bits: usize,
195                    pattern: usize,
196                    log_before: f64,
197                    out: &mut [f64],
198                ) {
199                    if depth == bits {
200                        out[pattern] = tree.get_log_block_probability() - log_before;
201                        return;
202                    }
203                    for bit in [false, true] {
204                        tree.update(bit);
205                        rec(
206                            tree,
207                            depth + 1,
208                            bits,
209                            (pattern << 1) | (bit as usize),
210                            log_before,
211                            out,
212                        );
213                        tree.revert();
214                    }
215                }
216
217                let log_before = tree.get_log_block_probability();
218                rec(
219                    tree,
220                    0,
221                    bits,
222                    0,
223                    log_before,
224                    &mut self.pattern_logps[..patterns],
225                );
226            }
227            CtwCompressionTree::Fac(tree) => {
228                fn rec(
229                    tree: &mut FacContextTree,
230                    bits: usize,
231                    msb_first: bool,
232                    depth: usize,
233                    pattern: usize,
234                    log_before: f64,
235                    out: &mut [f64],
236                ) {
237                    if depth == bits {
238                        out[pattern] = tree.get_log_block_probability() - log_before;
239                        return;
240                    }
241                    for bit in [false, true] {
242                        tree.update(bit, depth);
243                        let next_pattern = if msb_first {
244                            (pattern << 1) | (bit as usize)
245                        } else {
246                            pattern | ((bit as usize) << depth)
247                        };
248                        rec(
249                            tree,
250                            bits,
251                            msb_first,
252                            depth + 1,
253                            next_pattern,
254                            log_before,
255                            out,
256                        );
257                        tree.revert(depth);
258                    }
259                }
260
261                let log_before = tree.get_log_block_probability();
262                rec(
263                    tree,
264                    bits,
265                    self.msb_first,
266                    0,
267                    0,
268                    log_before,
269                    &mut self.pattern_logps[..patterns],
270                );
271            }
272        }
273        patterns
274    }
275
276    #[cfg(test)]
277    fn log_prob_symbol_bruteforce(&mut self, symbol: u8) -> f64 {
278        let bits = self.bits_per_symbol.clamp(1, 8);
279        match &mut self.tree {
280            CtwCompressionTree::Ac(tree) => {
281                debug_assert!(self.msb_first);
282                let before = tree.get_log_block_probability();
283                for bit_idx in 0..bits {
284                    tree.update(ctw_symbol_bit_msb(symbol, bits, bit_idx));
285                }
286                let after = tree.get_log_block_probability();
287                for _ in 0..bits {
288                    tree.revert();
289                }
290                after - before
291            }
292            CtwCompressionTree::Fac(tree) => {
293                let before = tree.get_log_block_probability();
294                if self.msb_first {
295                    for bit_idx in 0..bits {
296                        let bit = ((symbol >> (7 - bit_idx)) & 1) == 1;
297                        tree.update(bit, bit_idx);
298                    }
299                    let after = tree.get_log_block_probability();
300                    for bit_idx in (0..bits).rev() {
301                        tree.revert(bit_idx);
302                    }
303                    after - before
304                } else {
305                    for bit_idx in 0..bits {
306                        let bit = ((symbol >> bit_idx) & 1) == 1;
307                        tree.update(bit, bit_idx);
308                    }
309                    let after = tree.get_log_block_probability();
310                    for bit_idx in (0..bits).rev() {
311                        tree.revert(bit_idx);
312                    }
313                    after - before
314                }
315            }
316        }
317    }
318
319    fn pdf_next(&mut self) -> &[f64] {
320        if !self.valid {
321            let bits = self.bits_per_symbol.clamp(1, 8);
322            let patterns = self.fill_pattern_log_probs();
323            if bits == 8 {
324                for sym in 0..256usize {
325                    self.pdf[sym] = self.pattern_logps[sym].exp();
326                }
327            } else {
328                let aliases = 1usize << (8 - bits);
329                for byte in 0..256usize {
330                    let pat = if self.msb_first {
331                        byte >> (8 - bits)
332                    } else {
333                        byte & (patterns - 1)
334                    };
335                    self.pdf[byte] = self.pattern_logps[pat].exp() / (aliases as f64);
336                }
337            }
338            normalize_pdf(&mut self.pdf, PDF_MIN);
339            self.valid = true;
340        }
341        &self.pdf
342    }
343
344    fn update(&mut self, symbol: u8) {
345        match &mut self.tree {
346            CtwCompressionTree::Ac(tree) => {
347                debug_assert!(self.msb_first);
348                for bit_idx in 0..self.bits_per_symbol {
349                    tree.update(ctw_symbol_bit_msb(symbol, self.bits_per_symbol, bit_idx));
350                }
351            }
352            CtwCompressionTree::Fac(tree) => {
353                if self.msb_first {
354                    for bit_idx in 0..self.bits_per_symbol {
355                        let bit = ((symbol >> (7 - bit_idx)) & 1) == 1;
356                        tree.update(bit, bit_idx);
357                    }
358                } else {
359                    for bit_idx in 0..self.bits_per_symbol {
360                        let bit = ((symbol >> bit_idx) & 1) == 1;
361                        tree.update(bit, bit_idx);
362                    }
363                }
364            }
365        }
366        self.valid = false;
367    }
368
369    fn reserve_for_symbols(&mut self, total_symbols: usize) {
370        match &mut self.tree {
371            CtwCompressionTree::Ac(tree) => tree.reserve_for_symbols(
372                total_symbols.saturating_mul(self.bits_per_symbol.clamp(1, 8)),
373            ),
374            CtwCompressionTree::Fac(tree) => tree.reserve_for_symbols(total_symbols),
375        }
376    }
377
378    #[inline]
379    fn can_fast_ac_bitwise(&self) -> bool {
380        self.bits_per_symbol == 8 && self.msb_first
381    }
382
383    #[inline]
384    fn bit_prob_one_msb(&mut self, bit_idx: usize) -> f64 {
385        debug_assert!(self.can_fast_ac_bitwise());
386        match &mut self.tree {
387            CtwCompressionTree::Ac(tree) => tree.predict(true).clamp(PDF_MIN, 1.0 - PDF_MIN),
388            CtwCompressionTree::Fac(tree) => {
389                tree.predict_one(bit_idx).clamp(PDF_MIN, 1.0 - PDF_MIN)
390            }
391        }
392    }
393
394    #[inline]
395    fn update_bit_msb(&mut self, bit_idx: usize, bit: bool) {
396        debug_assert!(self.can_fast_ac_bitwise());
397        match &mut self.tree {
398            CtwCompressionTree::Ac(tree) => tree.update(bit),
399            CtwCompressionTree::Fac(tree) => tree.update_predicted(bit, bit_idx),
400        }
401        self.valid = false;
402    }
403}
404
405#[cfg(feature = "backend-rosa")]
406#[derive(Clone)]
407pub(crate) struct RosaPredictor {
408    model: RosaPlus,
409    pdf: Vec<f64>,
410    cdf: [f64; 257],
411    valid: bool,
412    cdf_valid: bool,
413}
414
415#[cfg(feature = "backend-rosa")]
416impl RosaPredictor {
417    pub(crate) fn new(max_order: i64) -> Self {
418        let mut model = RosaPlus::new(max_order, false, 0, 42);
419        model.build_lm_full_bytes_no_finalize_endpos();
420        Self {
421            model,
422            pdf: vec![0.0; 256],
423            cdf: uniform_cdf_row(),
424            valid: false,
425            cdf_valid: false,
426        }
427    }
428
429    fn pdf_next(&mut self) -> &[f64] {
430        self.ensure_pdf(false);
431        &self.pdf
432    }
433
434    fn cdf_next(&mut self) -> &[f64; 257] {
435        self.ensure_pdf(true);
436        &self.cdf
437    }
438
439    fn ensure_pdf(&mut self, want_cdf: bool) {
440        if self.valid {
441            if want_cdf && !self.cdf_valid {
442                build_cdf_row_from_pdf_slice(&self.pdf, &mut self.cdf);
443                self.cdf_valid = true;
444            }
445            return;
446        }
447        self.model.fill_probs_for_last_bytes(&mut self.pdf);
448        normalize_pdf_vec_and_maybe_build_cdf(
449            &mut self.pdf,
450            if want_cdf { Some(&mut self.cdf) } else { None },
451        );
452        self.valid = true;
453        self.cdf_valid = want_cdf;
454    }
455
456    fn update(&mut self, symbol: u8) {
457        self.model.train_byte(symbol);
458        self.valid = false;
459        self.cdf_valid = false;
460    }
461
462    fn begin_stream(&mut self, total_len: usize) {
463        self.model.reserve_for_stream(total_len);
464    }
465}
466
467#[derive(Clone)]
468#[cfg(feature = "backend-mamba")]
469pub(crate) struct MambaPredictor {
470    compressor: mambazip::Compressor,
471    primed: bool,
472    pdf: Vec<f64>,
473    cdf: [f64; 257],
474    valid: bool,
475    cdf_valid: bool,
476}
477
478#[derive(Clone)]
479#[cfg(feature = "backend-rwkv")]
480pub(crate) struct RwkvPredictor {
481    compressor: rwkvzip::Compressor,
482    primed: bool,
483    cdf: [f64; 257],
484    cdf_valid: bool,
485}
486
487#[cfg(feature = "backend-zpaq")]
488#[derive(Clone)]
489pub(crate) struct ZpaqPredictor {
490    method: String,
491    history: Vec<u8>,
492    pdf: Vec<f64>,
493    valid: bool,
494}
495
496#[cfg(feature = "backend-zpaq")]
497impl ZpaqPredictor {
498    pub(crate) fn new(method: String) -> Self {
499        Self {
500            method,
501            history: Vec::new(),
502            pdf: vec![0.0; 256],
503            valid: false,
504        }
505    }
506
507    fn pdf_next(&mut self) -> &[f64] {
508        if !self.valid {
509            for sym in 0..256usize {
510                let mut model = ZpaqRateModel::new(self.method.clone(), PDF_MIN);
511                if !self.history.is_empty() {
512                    let _ = model.update_and_score(&self.history);
513                }
514                let logp = model.log_prob(sym as u8);
515                self.pdf[sym] = logp.exp().max(PDF_MIN);
516            }
517            normalize_pdf(&mut self.pdf, PDF_MIN);
518            self.valid = true;
519        }
520        &self.pdf
521    }
522
523    fn update(&mut self, symbol: u8) {
524        self.history.push(symbol);
525        self.valid = false;
526    }
527}
528
529#[cfg(feature = "backend-mamba")]
530impl MambaPredictor {
531    #[cfg(test)]
532    fn from_method(method: &str) -> Result<Self> {
533        let spec = mambazip::parse_method_spec(method)?;
534        Self::from_method_spec(&spec)
535    }
536
537    pub(crate) fn from_method_spec(method: &mambazip::MethodSpec) -> Result<Self> {
538        let compressor = mambazip::Compressor::new_from_method_spec(method)?;
539        let vocab = compressor.vocab_size();
540        Ok(Self {
541            compressor,
542            primed: false,
543            pdf: vec![0.0; vocab],
544            cdf: uniform_cdf_row(),
545            valid: false,
546            cdf_valid: false,
547        })
548    }
549
550    fn ensure_predicted(&mut self, want_cdf: bool) {
551        if self.valid {
552            if want_cdf && !self.cdf_valid {
553                debug_assert!(self.pdf.len() >= 256);
554                build_cdf_row_from_pdf_slice(&self.pdf[..256], &mut self.cdf);
555                self.cdf_valid = true;
556            }
557            return;
558        }
559        if !self.primed {
560            self.compressor.forward_to_pdf(0, &mut self.pdf);
561            self.primed = true;
562            self.valid = true;
563            self.cdf_valid = false;
564            if want_cdf {
565                debug_assert!(self.pdf.len() >= 256);
566                build_cdf_row_from_pdf_slice(&self.pdf[..256], &mut self.cdf);
567                self.cdf_valid = true;
568            }
569            return;
570        }
571        self.valid = true;
572        self.cdf_valid = false;
573        if want_cdf {
574            debug_assert!(self.pdf.len() >= 256);
575            build_cdf_row_from_pdf_slice(&self.pdf[..256], &mut self.cdf);
576            self.cdf_valid = true;
577        }
578    }
579
580    fn pdf_next(&mut self) -> &[f64] {
581        self.ensure_predicted(false);
582        &self.pdf
583    }
584
585    fn cdf_next(&mut self) -> &[f64; 257] {
586        self.ensure_predicted(true);
587        &self.cdf
588    }
589
590    fn update(&mut self, symbol: u8) -> Result<()> {
591        self.ensure_predicted(false);
592        self.compressor.online_update_from_pdf(symbol, &self.pdf)?;
593        self.compressor.forward_to_pdf(symbol as u32, &mut self.pdf);
594        self.valid = true;
595        self.cdf_valid = false;
596        Ok(())
597    }
598
599    fn begin_stream(&mut self, total_len: usize) -> Result<()> {
600        self.compressor
601            .begin_online_policy_stream(Some(total_len as u64))
602    }
603}
604
605#[cfg(feature = "backend-rwkv")]
606impl RwkvPredictor {
607    #[cfg(test)]
608    fn from_method(method: &str) -> Result<Self> {
609        let spec = rwkvzip::parse_method_spec(method)?;
610        Self::from_method_spec(&spec)
611    }
612
613    pub(crate) fn from_method_spec(method: &rwkvzip::MethodSpec) -> Result<Self> {
614        let compressor = rwkvzip::Compressor::new_from_method_spec(method)?;
615        Ok(Self {
616            compressor,
617            primed: false,
618            cdf: uniform_cdf_row(),
619            cdf_valid: false,
620        })
621    }
622
623    fn ensure_predicted(&mut self, want_cdf: bool) {
624        if !self.primed {
625            self.compressor.reset_and_prime();
626            self.primed = true;
627            self.cdf_valid = false;
628        }
629        if want_cdf && !self.cdf_valid {
630            debug_assert!(self.compressor.pdf_buffer.len() >= 256);
631            build_cdf_row_from_pdf_slice(&self.compressor.pdf_buffer[..256], &mut self.cdf);
632            self.cdf_valid = true;
633        }
634    }
635
636    fn pdf_next(&mut self) -> &[f64] {
637        self.ensure_predicted(false);
638        &self.compressor.pdf_buffer
639    }
640
641    fn cdf_next(&mut self) -> &[f64; 257] {
642        self.ensure_predicted(true);
643        &self.cdf
644    }
645
646    fn update(&mut self, symbol: u8) -> Result<()> {
647        self.ensure_predicted(false);
648        self.compressor.observe_symbol_from_current_pdf(symbol)?;
649        self.cdf_valid = false;
650        Ok(())
651    }
652
653    fn begin_stream(&mut self, total_len: usize) -> Result<()> {
654        self.compressor
655            .begin_online_policy_stream(Some(total_len as u64))
656    }
657
658    fn finish_stream(&mut self) -> Result<()> {
659        self.compressor.finish_online_policy_stream()
660    }
661}
662
663#[derive(Clone)]
664struct MixExpert {
665    predictor: Box<RatePdfPredictor>,
666    log_weight: f64,
667    log_prior: f64,
668    cum_log_loss: f64,
669}
670
671#[derive(Clone, Debug)]
672pub(crate) enum PredictorBitwiseStepState {
673    NativeRecursive,
674    CachedCdf {
675        range: MsbPrefixRange,
676    },
677    PdfPrefix {
678        cdf: Box<BytePrefixCdf>,
679        range: MsbPrefixRange,
680    },
681}
682
683impl Default for PredictorBitwiseStepState {
684    fn default() -> Self {
685        Self::PdfPrefix {
686            cdf: zeroed_prefix_cdf_box(),
687            range: MsbPrefixRange::FULL,
688        }
689    }
690}
691
692impl PredictorBitwiseStepState {
693    fn prepare(&mut self, predictor: &mut RatePdfPredictor) -> Result<()> {
694        if predictor.begin_native_recursive_bitwise_byte_step()? {
695            *self = Self::NativeRecursive;
696            return Ok(());
697        }
698        if predictor.prepare_cached_cdf_fast_bitwise()? {
699            *self = Self::CachedCdf {
700                range: MsbPrefixRange::FULL,
701            };
702            return Ok(());
703        }
704
705        let mut cdf = match std::mem::take(self) {
706            Self::PdfPrefix { cdf, .. } => cdf,
707            _ => zeroed_prefix_cdf_box(),
708        };
709        fill_prefix_cdf_from_pdf(&mut cdf, predictor.pdf_next()?, PDF_MIN);
710        *self = Self::PdfPrefix {
711            cdf,
712            range: MsbPrefixRange::FULL,
713        };
714        Ok(())
715    }
716
717    fn bit_prob_one_msb(
718        &mut self,
719        predictor: &mut RatePdfPredictor,
720        bit_idx: usize,
721    ) -> Result<f64> {
722        match self {
723            Self::NativeRecursive => predictor.native_recursive_bit_prob_one_msb(bit_idx),
724            Self::CachedCdf { range } => Ok(predictor
725                .cached_cdf_bit_prob_one_msb(*range)
726                .expect("CachedCdf state invariant violated: missing cached CDF entry")),
727            Self::PdfPrefix { cdf, range } => Ok(range.prob_one(cdf.as_ref(), PDF_MIN)),
728        }
729    }
730
731    fn observe_bit_msb(
732        &mut self,
733        predictor: &mut RatePdfPredictor,
734        bit_idx: usize,
735        bit: bool,
736    ) -> Result<()> {
737        match self {
738            Self::NativeRecursive => predictor.native_recursive_observe_bit_msb(bit_idx, bit),
739            Self::CachedCdf { range } | Self::PdfPrefix { range, .. } => {
740                range.observe(bit);
741                Ok(())
742            }
743        }
744    }
745
746    fn finish_symbol(&mut self, predictor: &mut RatePdfPredictor, symbol: u8) -> Result<()> {
747        match self {
748            Self::NativeRecursive => predictor.finish_native_recursive_bitwise_byte_step(symbol),
749            Self::CachedCdf { .. } | Self::PdfPrefix { .. } => predictor.update(symbol),
750        }
751    }
752}
753
754#[derive(Clone, Copy, Debug, Default)]
755pub(crate) struct AcLogLossNodeValue {
756    pub(crate) prob: f64,
757    pub(crate) local_weight: f64,
758    pub(crate) effective_weight: f64,
759}
760
761#[derive(Clone, Debug, Default)]
762pub(crate) struct AcLogLossSubtreeSnapshot {
763    pub(crate) prob: f64,
764    pub(crate) rows: Vec<AcLogLossNodeValue>,
765}
766
767#[derive(Clone, Copy, Debug, Default)]
768pub(crate) struct AcLogLossRootSnapshot {
769    pub(crate) mix_prob: f64,
770    pub(crate) root_weight_entropy_bits: f64,
771    pub(crate) root_top1_child_index: Option<usize>,
772    pub(crate) root_top1_weight: f64,
773    pub(crate) root_top2_child_index: Option<usize>,
774    pub(crate) root_top2_weight: f64,
775}
776
777#[derive(Clone)]
778pub(crate) struct MixturePredictor {
779    kind: MixtureKind,
780    schedule: MixtureScheduleMode,
781    alpha: f64,
782    decay: f64,
783    experts: Vec<MixExpert>,
784    prior_weights: Vec<f64>,
785    neural: NeuralMixCore,
786    analyzer: TextContextAnalyzer,
787    bitwise_expert_states: Vec<PredictorBitwiseStepState>,
788    // Reused per-expert observation scratch: symbol paths store log p(symbol),
789    // while bitwise AC temporarily stages p(bit = 1) before collapsing back to
790    // the symbol log-probability at byte completion.
791    expert_observation_scratch: Vec<f64>,
792    scratch: Vec<f64>,
793    scratch2: Vec<f64>,
794    projection_scratch: Vec<f64>,
795    pdf: Vec<f64>,
796    valid: bool,
797    switch_updates: u64,
798    convex_updates: u64,
799}
800
801impl MixturePredictor {
802    pub(crate) fn new_from_compiled(backend: &CompiledRateBackend) -> Result<Self> {
803        let crate::spec::core::RateBackendPlan::Mixture {
804            kind,
805            schedule,
806            alpha,
807            decay,
808            experts: plan_experts,
809            ..
810        } = backend.plan()
811        else {
812            bail!("compiled backend is not a mixture backend");
813        };
814        let mut experts = Vec::with_capacity(plan_experts.len());
815        for expert_plan in plan_experts.iter() {
816            let compiled =
817                crate::spec::core::compiled_rate_backend_from_plan(expert_plan.backend.clone())
818                    .map_err(anyhow::Error::msg)?;
819            experts.push(MixExpert {
820                predictor: Box::new(crate::runtime::build_rate_pdf_predictor(&compiled)?),
821                log_weight: expert_plan.log_prior,
822                log_prior: expert_plan.log_prior,
823                cum_log_loss: 0.0,
824            });
825        }
826        let m = logsumexp_expert_weights(&experts);
827        for e in &mut experts {
828            e.log_weight -= m;
829        }
830
831        let mut prior_weights = vec![0.0; experts.len()];
832        normalized_mix_expert_prior_weights(&experts, &mut prior_weights);
833        let mut neural_prior_weights = prior_weights.clone();
834        for weight in &mut neural_prior_weights {
835            *weight = weight.clamp(PDF_MIN, 1.0 - PDF_MIN);
836        }
837
838        let base_lr = alpha.abs().clamp(1e-6, 1.0);
839        let effective_lr = (base_lr * 25.0).clamp(1e-6, 1.0);
840        let analyzer = TextContextAnalyzer::new();
841        let mut neural = NeuralMixCore::new(
842            experts.len(),
843            &neural_prior_weights,
844            effective_lr * 0.5,
845            effective_lr,
846            1e-5,
847        );
848        neural.set_context_state(analyzer.state());
849        Ok(Self {
850            kind: *kind,
851            schedule: *schedule,
852            alpha: *alpha,
853            decay: decay.unwrap_or(1.0).clamp(0.0, 1.0),
854            experts,
855            prior_weights,
856            neural,
857            analyzer,
858            bitwise_expert_states: Vec::new(),
859            expert_observation_scratch: vec![0.0; plan_experts.len()],
860            scratch: Vec::new(),
861            scratch2: Vec::new(),
862            projection_scratch: Vec::new(),
863            pdf: vec![0.0; 256],
864            valid: false,
865            switch_updates: 0,
866            convex_updates: 0,
867        })
868    }
869
870    fn best_expert_index(&self) -> Option<usize> {
871        let mut best_idx = None;
872        let mut best_loss = f64::INFINITY;
873        for (index, expert) in self.experts.iter().enumerate() {
874            if expert.cum_log_loss < best_loss {
875                best_loss = expert.cum_log_loss;
876                best_idx = Some(index);
877            }
878        }
879        best_idx
880    }
881
882    fn predictive_weights(&mut self) -> Vec<f64> {
883        if self.experts.is_empty() {
884            return Vec::new();
885        }
886
887        match self.kind {
888            MixtureKind::Neural => {
889                if self.experts.len() == 1 {
890                    return vec![1.0];
891                }
892                self.neural.set_context_state(self.analyzer.state());
893                self.neural.evaluate_expert_weights();
894                let mut weights = self.neural.expert_weights().to_vec();
895                normalize_simplex_weights(&mut weights);
896                weights
897            }
898            MixtureKind::Mdl => {
899                let mut weights = vec![0.0; self.experts.len()];
900                if let Some(best_idx) = self.best_expert_index() {
901                    weights[best_idx] = 1.0;
902                }
903                weights
904            }
905            MixtureKind::FadingBayes => {
906                let max_log = self
907                    .experts
908                    .iter()
909                    .map(|expert| self.decay * expert.log_weight)
910                    .fold(f64::NEG_INFINITY, f64::max);
911                let mut weights = self
912                    .experts
913                    .iter()
914                    .map(|expert| {
915                        if max_log.is_finite() {
916                            (self.decay * expert.log_weight - max_log).exp()
917                        } else {
918                            0.0
919                        }
920                    })
921                    .collect::<Vec<_>>();
922                normalize_simplex_weights(&mut weights);
923                weights
924            }
925            MixtureKind::Convex => {
926                let mut weights = self
927                    .experts
928                    .iter()
929                    .map(|expert| expert.log_weight.exp())
930                    .collect::<Vec<_>>();
931                normalize_simplex_weights(&mut weights);
932                weights
933            }
934            MixtureKind::Bayes | MixtureKind::Switching => {
935                let max_log = self
936                    .experts
937                    .iter()
938                    .map(|expert| expert.log_weight)
939                    .fold(f64::NEG_INFINITY, f64::max);
940                let mut weights = self
941                    .experts
942                    .iter()
943                    .map(|expert| {
944                        if max_log.is_finite() {
945                            (expert.log_weight - max_log).exp()
946                        } else {
947                            0.0
948                        }
949                    })
950                    .collect::<Vec<_>>();
951                normalize_simplex_weights(&mut weights);
952                weights
953            }
954        }
955    }
956
957    fn ensure_pdf(&mut self) -> Result<&[f64]> {
958        if self.valid {
959            return Ok(&self.pdf);
960        }
961        let weights = self.predictive_weights();
962        self.pdf.fill(0.0);
963        for (index, expert) in self.experts.iter_mut().enumerate() {
964            let weight = weights.get(index).copied().unwrap_or(0.0);
965            if weight <= 0.0 {
966                continue;
967            }
968            let epdf = expert.predictor.pdf_next()?;
969            for (slot, &p) in self.pdf.iter_mut().zip(epdf.iter()) {
970                *slot += weight * p;
971            }
972        }
973
974        normalize_pdf(&mut self.pdf, PDF_MIN);
975        self.valid = true;
976        Ok(&self.pdf)
977    }
978
979    fn begin_stream(&mut self, total_len: usize) -> Result<()> {
980        for expert in &mut self.experts {
981            match &mut *expert.predictor {
982                // Direct CTW benefits from pre-reserving, but inside mixtures that extra
983                // headroom can dominate peak RSS without a proportional runtime gain.
984                #[cfg(feature = "backend-ctw")]
985                RatePdfPredictor::Ctw(_) | RatePdfPredictor::FacCtw(_) => {}
986                _ => expert.predictor.begin_stream(total_len)?,
987            }
988        }
989        Ok(())
990    }
991
992    fn diagnostic_collect_children(
993        &mut self,
994        symbol: u8,
995        weights: &[f64],
996        effective_prefix: f64,
997        pool: Option<&ThreadPool>,
998    ) -> Result<Vec<AcLogLossSubtreeSnapshot>> {
999        let use_parallel = pool.is_some() && self.experts.len() >= DIAGNOSTIC_PARALLEL_THRESHOLD;
1000        if use_parallel {
1001            let pool = pool.expect("checked is_some");
1002            pool.install(|| {
1003                self.experts
1004                    .par_iter_mut()
1005                    .enumerate()
1006                    .map(|(index, expert)| {
1007                        let local_weight = weights.get(index).copied().unwrap_or(0.0);
1008                        let effective_weight = effective_prefix * local_weight;
1009                        expert.predictor.diagnostic_snapshot_subtree(
1010                            symbol,
1011                            local_weight,
1012                            effective_weight,
1013                            None,
1014                        )
1015                    })
1016                    .collect()
1017            })
1018        } else {
1019            let mut children = Vec::with_capacity(self.experts.len());
1020            for (index, expert) in self.experts.iter_mut().enumerate() {
1021                let local_weight = weights.get(index).copied().unwrap_or(0.0);
1022                let effective_weight = effective_prefix * local_weight;
1023                children.push(expert.predictor.diagnostic_snapshot_subtree(
1024                    symbol,
1025                    local_weight,
1026                    effective_weight,
1027                    pool,
1028                )?);
1029            }
1030            Ok(children)
1031        }
1032    }
1033
1034    fn diagnostic_subtree_snapshot(
1035        &mut self,
1036        symbol: u8,
1037        local_weight: f64,
1038        effective_weight: f64,
1039        pool: Option<&ThreadPool>,
1040    ) -> Result<AcLogLossSubtreeSnapshot> {
1041        let weights = self.predictive_weights();
1042        let children =
1043            self.diagnostic_collect_children(symbol, &weights, effective_weight, pool)?;
1044        let mix_prob = children
1045            .iter()
1046            .enumerate()
1047            .map(|(index, child)| weights.get(index).copied().unwrap_or(0.0) * child.prob)
1048            .sum::<f64>()
1049            .max(PDF_MIN);
1050        let total_rows = 1 + children.iter().map(|child| child.rows.len()).sum::<usize>();
1051        let mut rows = Vec::with_capacity(total_rows);
1052        rows.push(AcLogLossNodeValue {
1053            prob: mix_prob,
1054            local_weight,
1055            effective_weight,
1056        });
1057        for child in children {
1058            rows.extend(child.rows);
1059        }
1060        Ok(AcLogLossSubtreeSnapshot {
1061            prob: mix_prob,
1062            rows,
1063        })
1064    }
1065
1066    fn diagnostic_root_snapshot(
1067        &mut self,
1068        symbol: u8,
1069        pool: Option<&ThreadPool>,
1070        out: &mut Vec<AcLogLossNodeValue>,
1071    ) -> Result<AcLogLossRootSnapshot> {
1072        let weights = self.predictive_weights();
1073        let children = self.diagnostic_collect_children(symbol, &weights, 1.0, pool)?;
1074        out.clear();
1075        out.reserve(children.iter().map(|child| child.rows.len()).sum::<usize>());
1076        for child in &children {
1077            out.extend_from_slice(&child.rows);
1078        }
1079
1080        let mix_prob = children
1081            .iter()
1082            .enumerate()
1083            .map(|(index, child)| weights.get(index).copied().unwrap_or(0.0) * child.prob)
1084            .sum::<f64>()
1085            .max(PDF_MIN);
1086
1087        let mut top1 = None;
1088        let mut top2 = None;
1089        for (index, &weight) in weights.iter().enumerate() {
1090            match top1 {
1091                None => top1 = Some((index, weight)),
1092                Some((best_idx, best_weight)) if weight > best_weight => {
1093                    top2 = Some((best_idx, best_weight));
1094                    top1 = Some((index, weight));
1095                }
1096                _ => match top2 {
1097                    None => top2 = Some((index, weight)),
1098                    Some((_, second_weight)) if weight > second_weight => {
1099                        top2 = Some((index, weight));
1100                    }
1101                    _ => {}
1102                },
1103            }
1104        }
1105
1106        let root_weight_entropy_bits = weights
1107            .iter()
1108            .copied()
1109            .filter(|weight| *weight > 0.0)
1110            .map(|weight| -weight * weight.log2())
1111            .sum::<f64>();
1112
1113        Ok(AcLogLossRootSnapshot {
1114            mix_prob,
1115            root_weight_entropy_bits,
1116            root_top1_child_index: top1.map(|(index, _)| index),
1117            root_top1_weight: top1.map(|(_, weight)| weight).unwrap_or(0.0),
1118            root_top2_child_index: top2.map(|(index, _)| index),
1119            root_top2_weight: top2.map(|(_, weight)| weight).unwrap_or(0.0),
1120        })
1121    }
1122
1123    fn update(&mut self, symbol: u8) -> Result<()> {
1124        let _ = self.ensure_pdf()?;
1125
1126        match self.kind {
1127            MixtureKind::Bayes => {
1128                let n = self.experts.len();
1129                self.scratch.resize(n, 0.0);
1130                self.scratch2.resize(n, 0.0);
1131                for (i, e) in self.experts.iter_mut().enumerate() {
1132                    let p = e.predictor.pdf_next()?[symbol as usize].max(PDF_MIN);
1133                    let lp = p.ln();
1134                    self.scratch[i] = lp;
1135                    self.scratch2[i] = e.log_weight + lp;
1136                }
1137                let log_mix = logsumexp_slice(&self.scratch2[..n]);
1138                for (i, e) in self.experts.iter_mut().enumerate() {
1139                    e.log_weight = e.log_weight + self.scratch[i] - log_mix;
1140                    e.cum_log_loss -= self.scratch[i];
1141                    e.predictor.update(symbol)?;
1142                }
1143            }
1144            MixtureKind::FadingBayes => {
1145                let n = self.experts.len();
1146                self.scratch.resize(n, 0.0);
1147                self.scratch2.resize(n, 0.0);
1148                for (i, e) in self.experts.iter_mut().enumerate() {
1149                    let p = e.predictor.pdf_next()?[symbol as usize].max(PDF_MIN);
1150                    let lp = p.ln();
1151                    self.scratch[i] = lp;
1152                    self.scratch2[i] = e.log_weight + lp;
1153                }
1154                for (i, e) in self.experts.iter_mut().enumerate() {
1155                    self.scratch2[i] = self.decay * e.log_weight + self.scratch[i];
1156                }
1157                let log_mix = logsumexp_slice(&self.scratch2[..n]);
1158                for (i, e) in self.experts.iter_mut().enumerate() {
1159                    e.log_weight = self.decay * e.log_weight + self.scratch[i] - log_mix;
1160                    e.cum_log_loss -= self.scratch[i];
1161                    e.predictor.update(symbol)?;
1162                }
1163            }
1164            MixtureKind::Switching => {
1165                let n = self.experts.len();
1166                self.scratch.resize(n, 0.0);
1167                self.scratch2.resize(n, 0.0);
1168                for (i, e) in self.experts.iter_mut().enumerate() {
1169                    let p = e.predictor.pdf_next()?[symbol as usize].max(PDF_MIN);
1170                    let lp = p.ln();
1171                    self.scratch[i] = lp;
1172                    self.scratch2[i] = e.log_weight + lp;
1173                }
1174                let log_mix = logsumexp_slice(&self.scratch2[..n]);
1175                for (i, e) in self.experts.iter_mut().enumerate() {
1176                    self.scratch2[i] = (self.scratch2[i] - log_mix).exp();
1177                    e.cum_log_loss -= self.scratch[i];
1178                    e.predictor.update(symbol)?;
1179                }
1180                let alpha =
1181                    switching_alpha_for_update(self.schedule, self.alpha, self.switch_updates);
1182                self.switch_updates = self.switch_updates.saturating_add(1);
1183                apply_switching_weights(
1184                    &mut self.experts,
1185                    &self.prior_weights[..n],
1186                    alpha,
1187                    &mut self.scratch2[..n],
1188                    &mut self.scratch[..n],
1189                );
1190            }
1191            MixtureKind::Convex => {
1192                let n = self.experts.len();
1193                self.scratch.resize(n, 0.0);
1194                self.scratch2.resize(n, 0.0);
1195                for (i, e) in self.experts.iter_mut().enumerate() {
1196                    let p = e.predictor.pdf_next()?[symbol as usize].max(PDF_MIN);
1197                    let lp = p.ln();
1198                    self.scratch[i] = lp;
1199                    self.scratch2[i] = e.log_weight.exp();
1200                    e.cum_log_loss -= lp;
1201                    e.predictor.update(symbol)?;
1202                }
1203                let mix_prob = self
1204                    .scratch
1205                    .iter()
1206                    .zip(self.scratch2.iter())
1207                    .map(|(&lp, &w)| w * lp.exp())
1208                    .sum::<f64>()
1209                    .max(PDF_MIN);
1210                let log_mix = mix_prob.ln();
1211                self.convex_updates = self.convex_updates.saturating_add(1);
1212                let eta =
1213                    convex_step_size_for_update(self.schedule, self.alpha, self.convex_updates);
1214                for i in 0..n {
1215                    let grad = -(self.scratch[i] - log_mix).exp();
1216                    self.scratch2[i] -= eta * grad;
1217                }
1218                project_simplex_with_scratch(&mut self.scratch2[..n], &mut self.projection_scratch);
1219                for i in 0..n {
1220                    self.experts[i].log_weight = self.scratch2[i].max(PDF_MIN).ln();
1221                }
1222            }
1223            MixtureKind::Mdl => {
1224                let n = self.experts.len();
1225                self.scratch.resize(n, 0.0);
1226                for (i, e) in self.experts.iter_mut().enumerate() {
1227                    let p = e.predictor.pdf_next()?[symbol as usize].max(PDF_MIN);
1228                    let lp = p.ln();
1229                    self.scratch[i] = lp;
1230                }
1231                for (i, e) in self.experts.iter_mut().enumerate() {
1232                    e.cum_log_loss -= self.scratch[i];
1233                    e.predictor.update(symbol)?;
1234                }
1235            }
1236            MixtureKind::Neural => {
1237                let y = symbol as usize;
1238                if self.experts.len() == 1 {
1239                    let lp = self.experts[0].predictor.pdf_next()?[y].max(PDF_MIN).ln();
1240                    self.experts[0].cum_log_loss -= lp;
1241                    self.experts[0].predictor.update(symbol)?;
1242                    self.analyzer.update(symbol);
1243                    self.neural.set_context_state(self.analyzer.state());
1244                    self.valid = false;
1245                    return Ok(());
1246                }
1247                let n = self.experts.len();
1248                self.neural.set_context_state(self.analyzer.state());
1249                self.expert_observation_scratch.resize(n, 0.0);
1250                for i in 0..n {
1251                    let p = self.experts[i].predictor.pdf_next()?[y].max(PDF_MIN);
1252                    let lp = p.ln();
1253                    self.expert_observation_scratch[i] = lp;
1254                    self.experts[i].cum_log_loss -= lp;
1255                }
1256                self.neural
1257                    .evaluate_symbol(&self.expert_observation_scratch, PDF_MIN);
1258                self.neural
1259                    .update_weights_symbol(&self.expert_observation_scratch, PDF_MIN);
1260                for e in &mut self.experts {
1261                    e.predictor.update(symbol)?;
1262                }
1263                self.analyzer.update(symbol);
1264                self.neural.set_context_state(self.analyzer.state());
1265            }
1266        }
1267
1268        self.valid = false;
1269        Ok(())
1270    }
1271
1272    fn finish_stream(&mut self) -> Result<()> {
1273        for expert in &mut self.experts {
1274            expert.predictor.finish_stream()?;
1275        }
1276        Ok(())
1277    }
1278
1279    #[inline]
1280    fn has_recursive_native_bitwise_expert(&self) -> bool {
1281        self.experts
1282            .iter()
1283            .any(|expert| expert.predictor.has_recursive_native_bitwise_path())
1284    }
1285
1286    fn begin_bitwise_byte_step(&mut self) -> Result<bool> {
1287        if !self.has_recursive_native_bitwise_expert() {
1288            return Ok(false);
1289        }
1290
1291        let n = self.experts.len();
1292        self.scratch.resize(n, 0.0);
1293        match self.kind {
1294            MixtureKind::Neural if n > 1 => {
1295                self.neural.set_context_state(self.analyzer.state());
1296                self.neural.evaluate_expert_weights();
1297                self.scratch.copy_from_slice(self.neural.expert_weights());
1298            }
1299            _ => {
1300                let weights = self.predictive_weights();
1301                self.scratch.copy_from_slice(&weights);
1302            }
1303        }
1304        self.scratch2.resize(n, 1.0);
1305        self.scratch2.fill(1.0);
1306        self.expert_observation_scratch.resize(n, 0.0);
1307        self.bitwise_expert_states
1308            .resize_with(n, PredictorBitwiseStepState::default);
1309        for i in 0..n {
1310            self.bitwise_expert_states[i].prepare(&mut self.experts[i].predictor)?;
1311        }
1312        Ok(true)
1313    }
1314
1315    fn bit_prob_one_msb(&mut self, bit_idx: usize) -> Result<f64> {
1316        let mut denom = 0.0;
1317        let mut numer1 = 0.0;
1318        for i in 0..self.experts.len() {
1319            let p1 = self.bitwise_expert_states[i]
1320                .bit_prob_one_msb(&mut self.experts[i].predictor, bit_idx)?;
1321            self.expert_observation_scratch[i] = p1;
1322            let wp = self.scratch[i] * self.scratch2[i];
1323            denom += wp;
1324            numer1 += wp * p1;
1325        }
1326        Ok(if denom.is_finite() && denom > 0.0 {
1327            (numer1 / denom).clamp(PDF_MIN, 1.0 - PDF_MIN)
1328        } else {
1329            panic!(
1330                "MixturePredictor bit_prob_one_msb: invalid denom (finite>0 violated); \
1331                 this is an internal invariant failure in the bitwise mixture state machine"
1332            )
1333        })
1334    }
1335
1336    fn observe_bit_msb(&mut self, bit_idx: usize, bit: bool) -> Result<()> {
1337        for i in 0..self.experts.len() {
1338            let p1 = self.expert_observation_scratch[i];
1339            let pb = if bit { p1 } else { 1.0 - p1 };
1340            self.scratch2[i] = (self.scratch2[i] * pb).max(PDF_MIN);
1341            self.bitwise_expert_states[i].observe_bit_msb(
1342                &mut self.experts[i].predictor,
1343                bit_idx,
1344                bit,
1345            )?;
1346        }
1347        Ok(())
1348    }
1349
1350    fn finish_bitwise_symbol(&mut self, symbol: u8) -> Result<()> {
1351        let n = self.experts.len();
1352        for i in 0..n {
1353            let lp = self.scratch2[i].max(PDF_MIN).ln();
1354            self.expert_observation_scratch[i] = lp;
1355            self.experts[i].cum_log_loss -= lp;
1356            self.bitwise_expert_states[i].finish_symbol(&mut self.experts[i].predictor, symbol)?;
1357        }
1358
1359        match self.kind {
1360            MixtureKind::Bayes => {
1361                for i in 0..n {
1362                    self.scratch[i] =
1363                        self.experts[i].log_weight + self.expert_observation_scratch[i];
1364                }
1365                let log_mix = logsumexp_slice(&self.scratch[..n]);
1366                for i in 0..n {
1367                    self.experts[i].log_weight += self.expert_observation_scratch[i] - log_mix;
1368                }
1369            }
1370            MixtureKind::FadingBayes => {
1371                for i in 0..n {
1372                    self.scratch[i] = self.decay * self.experts[i].log_weight
1373                        + self.expert_observation_scratch[i];
1374                }
1375                let log_mix = logsumexp_slice(&self.scratch[..n]);
1376                for i in 0..n {
1377                    self.experts[i].log_weight = self.scratch[i] - log_mix;
1378                }
1379            }
1380            MixtureKind::Switching => {
1381                for i in 0..n {
1382                    self.scratch[i] =
1383                        self.experts[i].log_weight + self.expert_observation_scratch[i];
1384                }
1385                let log_mix = logsumexp_slice(&self.scratch[..n]);
1386                for weight in &mut self.scratch[..n] {
1387                    *weight = (*weight - log_mix).exp();
1388                }
1389                let alpha =
1390                    switching_alpha_for_update(self.schedule, self.alpha, self.switch_updates);
1391                self.switch_updates = self.switch_updates.saturating_add(1);
1392                apply_switching_weights(
1393                    &mut self.experts,
1394                    &self.prior_weights[..n],
1395                    alpha,
1396                    &mut self.scratch[..n],
1397                    &mut self.scratch2[..n],
1398                );
1399            }
1400            MixtureKind::Convex => {
1401                self.scratch.resize(n, 0.0);
1402                self.scratch2.resize(n, 0.0);
1403                for i in 0..n {
1404                    self.scratch2[i] = self.experts[i].log_weight.exp();
1405                }
1406                let mix_prob = self
1407                    .expert_observation_scratch
1408                    .iter()
1409                    .zip(self.scratch2.iter())
1410                    .map(|(&lp, &w)| w * lp.exp())
1411                    .sum::<f64>()
1412                    .max(PDF_MIN);
1413                let log_mix = mix_prob.ln();
1414                self.convex_updates = self.convex_updates.saturating_add(1);
1415                let eta =
1416                    convex_step_size_for_update(self.schedule, self.alpha, self.convex_updates);
1417                for i in 0..n {
1418                    let grad = -(self.expert_observation_scratch[i] - log_mix).exp();
1419                    self.scratch2[i] -= eta * grad;
1420                }
1421                project_simplex_with_scratch(&mut self.scratch2[..n], &mut self.projection_scratch);
1422                for i in 0..n {
1423                    self.experts[i].log_weight = self.scratch2[i].max(PDF_MIN).ln();
1424                }
1425            }
1426            MixtureKind::Mdl => {}
1427            MixtureKind::Neural => {
1428                if n > 1 {
1429                    self.neural.set_context_state(self.analyzer.state());
1430                    self.neural
1431                        .evaluate_symbol(&self.expert_observation_scratch, PDF_MIN);
1432                    self.neural
1433                        .update_weights_symbol(&self.expert_observation_scratch, PDF_MIN);
1434                }
1435                self.analyzer.update(symbol);
1436                self.neural.set_context_state(self.analyzer.state());
1437            }
1438        }
1439        self.valid = false;
1440        Ok(())
1441    }
1442}
1443
1444pub(crate) struct DiagnosticRatePredictor {
1445    inner: RatePdfPredictor,
1446}
1447
1448impl DiagnosticRatePredictor {
1449    #[cfg(test)]
1450    pub(crate) fn from_rate_backend(backend: RateBackend) -> Result<Self> {
1451        let compiled = backend.compile().map_err(anyhow::Error::msg)?;
1452        Self::from_compiled(&compiled)
1453    }
1454
1455    pub(crate) fn from_compiled(backend: &CompiledRateBackend) -> Result<Self> {
1456        Ok(Self {
1457            inner: crate::runtime::build_rate_pdf_predictor(backend)?,
1458        })
1459    }
1460
1461    pub(crate) fn begin_stream(&mut self, total_len: usize) -> Result<()> {
1462        self.inner.begin_stream(total_len)
1463    }
1464
1465    pub(crate) fn finish_stream(&mut self) -> Result<()> {
1466        self.inner.finish_stream()
1467    }
1468
1469    #[cfg(test)]
1470    pub(crate) fn pdf_next(&mut self) -> Result<&[f64]> {
1471        self.inner.pdf_next()
1472    }
1473
1474    #[cfg(test)]
1475    pub(crate) fn update(&mut self, symbol: u8) -> Result<()> {
1476        self.inner.update(symbol)
1477    }
1478
1479    pub(crate) fn diagnostic_root_snapshot(
1480        &mut self,
1481        symbol: u8,
1482        pool: Option<&ThreadPool>,
1483        out: &mut Vec<AcLogLossNodeValue>,
1484    ) -> Result<AcLogLossRootSnapshot> {
1485        self.inner.diagnostic_root_snapshot(symbol, pool, out)
1486    }
1487
1488    pub(crate) fn encode_symbol_ac_step<W: std::io::Write>(
1489        &mut self,
1490        symbol: u8,
1491        encoder: &mut ArithmeticEncoder<W>,
1492        cdf: &mut [u32; 257],
1493    ) -> Result<()> {
1494        self.inner.encode_symbol_ac_step(symbol, encoder, cdf)
1495    }
1496}
1497
1498#[derive(Clone)]
1499#[allow(clippy::large_enum_variant)]
1500pub(crate) enum RatePdfPredictor {
1501    #[cfg(feature = "backend-rosa")]
1502    Rosa(RosaPredictor),
1503    #[cfg(feature = "backend-match")]
1504    Match { model: MatchModel },
1505    #[cfg(feature = "backend-match")]
1506    SparseMatch { model: SparseMatchModel },
1507    #[cfg(feature = "backend-ppmd")]
1508    Ppmd { model: PpmdModel },
1509    #[cfg(feature = "backend-sequitur")]
1510    Sequitur { model: SequiturModel },
1511    #[cfg(feature = "backend-ctw")]
1512    Ctw(CtwPredictor),
1513    #[cfg(feature = "backend-ctw")]
1514    FacCtw(CtwPredictor),
1515    #[cfg(feature = "backend-mamba")]
1516    Mamba(MambaPredictor),
1517    #[cfg(feature = "backend-rwkv")]
1518    Rwkv(RwkvPredictor),
1519    #[cfg(feature = "backend-zpaq")]
1520    Zpaq(ZpaqPredictor),
1521    #[cfg(feature = "backend-mixture")]
1522    Mixture(MixturePredictor),
1523    #[cfg(feature = "backend-particle")]
1524    Particle(crate::backends::particle::ParticleRuntime),
1525    #[cfg(feature = "backend-calibrated")]
1526    Calibrated {
1527        base: Box<RatePdfPredictor>,
1528        core: CalibratorCore,
1529        bitwise: PredictorBitwiseStepState,
1530        pdf: Vec<f64>,
1531        valid: bool,
1532    },
1533    #[allow(dead_code)]
1534    Disabled { reason: String },
1535}
1536
1537impl RatePdfPredictor {
1538    #[cfg(test)]
1539    fn from_compiled(backend: &CompiledRateBackend) -> Result<Self> {
1540        crate::runtime::build_rate_pdf_predictor(backend)
1541    }
1542
1543    #[cfg(test)]
1544    pub(crate) fn from_rate_backend(backend: RateBackend) -> Result<Self> {
1545        let compiled = backend.compile().map_err(anyhow::Error::msg)?;
1546        Self::from_compiled(&compiled)
1547    }
1548
1549    fn begin_stream(&mut self, total_len: usize) -> Result<()> {
1550        self.finish_stream()?;
1551        match self {
1552            #[cfg(feature = "backend-rosa")]
1553            Self::Rosa(m) => {
1554                m.begin_stream(total_len);
1555                Ok(())
1556            }
1557            #[cfg(feature = "backend-match")]
1558            Self::Match { .. } => Ok(()),
1559            #[cfg(feature = "backend-match")]
1560            Self::SparseMatch { .. } => Ok(()),
1561            #[cfg(feature = "backend-ppmd")]
1562            Self::Ppmd { .. } => Ok(()),
1563            #[cfg(feature = "backend-zpaq")]
1564            Self::Zpaq(_) => Ok(()),
1565            #[cfg(feature = "backend-particle")]
1566            Self::Particle(_) => Ok(()),
1567            #[cfg(feature = "backend-sequitur")]
1568            Self::Sequitur { model } => {
1569                model.begin_stream(Some(total_len as u64));
1570                Ok(())
1571            }
1572            #[cfg(feature = "backend-ctw")]
1573            Self::Ctw(m) | Self::FacCtw(m) => {
1574                m.reserve_for_symbols(total_len);
1575                Ok(())
1576            }
1577            #[cfg(feature = "backend-mamba")]
1578            Self::Mamba(m) => m.begin_stream(total_len),
1579            #[cfg(feature = "backend-rwkv")]
1580            Self::Rwkv(m) => m.begin_stream(total_len),
1581            #[cfg(feature = "backend-mixture")]
1582            Self::Mixture(m) => m.begin_stream(total_len),
1583            #[cfg(feature = "backend-calibrated")]
1584            Self::Calibrated {
1585                base,
1586                bitwise,
1587                valid,
1588                ..
1589            } => {
1590                *bitwise = PredictorBitwiseStepState::default();
1591                *valid = false;
1592                base.begin_stream(total_len)
1593            }
1594            Self::Disabled { reason } => bail!("{reason}"),
1595        }
1596    }
1597
1598    fn finish_stream(&mut self) -> Result<()> {
1599        match self {
1600            #[cfg(feature = "backend-rosa")]
1601            Self::Rosa(_) => Ok(()),
1602            #[cfg(feature = "backend-match")]
1603            Self::Match { .. } => Ok(()),
1604            #[cfg(feature = "backend-match")]
1605            Self::SparseMatch { .. } => Ok(()),
1606            #[cfg(feature = "backend-ppmd")]
1607            Self::Ppmd { .. } => Ok(()),
1608            #[cfg(feature = "backend-ctw")]
1609            Self::Ctw(_) => Ok(()),
1610            #[cfg(feature = "backend-ctw")]
1611            Self::FacCtw(_) => Ok(()),
1612            #[cfg(feature = "backend-zpaq")]
1613            Self::Zpaq(_) => Ok(()),
1614            #[cfg(feature = "backend-particle")]
1615            Self::Particle(_) => Ok(()),
1616            #[cfg(feature = "backend-sequitur")]
1617            Self::Sequitur { .. } => Ok(()),
1618            #[cfg(feature = "backend-mamba")]
1619            Self::Mamba(m) => m.compressor.finish_online_policy_stream(),
1620            #[cfg(feature = "backend-rwkv")]
1621            Self::Rwkv(m) => m.finish_stream(),
1622            #[cfg(feature = "backend-mixture")]
1623            Self::Mixture(m) => m.finish_stream(),
1624            #[cfg(feature = "backend-calibrated")]
1625            Self::Calibrated {
1626                base,
1627                bitwise,
1628                valid,
1629                ..
1630            } => {
1631                *bitwise = PredictorBitwiseStepState::default();
1632                *valid = false;
1633                base.finish_stream()
1634            }
1635            Self::Disabled { .. } => Ok(()),
1636        }
1637    }
1638
1639    fn pdf_next(&mut self) -> Result<&[f64]> {
1640        match self {
1641            #[cfg(feature = "backend-rosa")]
1642            Self::Rosa(m) => Ok(m.pdf_next()),
1643            #[cfg(feature = "backend-match")]
1644            Self::Match { model } => Ok(model.pdf()),
1645            #[cfg(feature = "backend-ctw")]
1646            Self::Ctw(m) => Ok(m.pdf_next()),
1647            #[cfg(feature = "backend-ctw")]
1648            Self::FacCtw(m) => Ok(m.pdf_next()),
1649            #[cfg(feature = "backend-mamba")]
1650            Self::Mamba(m) => Ok(m.pdf_next()),
1651            #[cfg(feature = "backend-rwkv")]
1652            Self::Rwkv(m) => Ok(m.pdf_next()),
1653            #[cfg(feature = "backend-zpaq")]
1654            Self::Zpaq(m) => Ok(m.pdf_next()),
1655            #[cfg(feature = "backend-mixture")]
1656            Self::Mixture(m) => m.ensure_pdf(),
1657            #[cfg(feature = "backend-particle")]
1658            Self::Particle(m) => Ok(m.pdf_next()),
1659            #[cfg(feature = "backend-match")]
1660            Self::SparseMatch { model } => Ok(model.pdf()),
1661            #[cfg(feature = "backend-ppmd")]
1662            Self::Ppmd { model } => Ok(model.pdf()),
1663            #[cfg(feature = "backend-sequitur")]
1664            Self::Sequitur { model } => Ok(model.pdf()),
1665            #[cfg(feature = "backend-calibrated")]
1666            Self::Calibrated {
1667                base,
1668                core,
1669                bitwise: _,
1670                pdf,
1671                valid,
1672            } => {
1673                if !*valid {
1674                    let base_pdf = base.pdf_next()?;
1675                    core.apply_pdf(base_pdf, pdf);
1676                    normalize_pdf(pdf, PDF_MIN);
1677                    *valid = true;
1678                }
1679                Ok(pdf)
1680            }
1681            Self::Disabled { reason } => bail!("{reason}"),
1682        }
1683    }
1684
1685    fn update(&mut self, symbol: u8) -> Result<()> {
1686        match self {
1687            #[cfg(feature = "backend-rosa")]
1688            Self::Rosa(m) => {
1689                m.update(symbol);
1690                Ok(())
1691            }
1692            #[cfg(feature = "backend-match")]
1693            Self::Match { model } => {
1694                model.update(symbol);
1695                Ok(())
1696            }
1697            #[cfg(feature = "backend-match")]
1698            Self::SparseMatch { model } => {
1699                model.update(symbol);
1700                Ok(())
1701            }
1702            #[cfg(feature = "backend-ppmd")]
1703            Self::Ppmd { model } => {
1704                model.update(symbol);
1705                Ok(())
1706            }
1707            #[cfg(feature = "backend-sequitur")]
1708            Self::Sequitur { model } => {
1709                model.update(symbol);
1710                Ok(())
1711            }
1712            #[cfg(feature = "backend-ctw")]
1713            Self::Ctw(m) => {
1714                m.update(symbol);
1715                Ok(())
1716            }
1717            #[cfg(feature = "backend-ctw")]
1718            Self::FacCtw(m) => {
1719                m.update(symbol);
1720                Ok(())
1721            }
1722            #[cfg(feature = "backend-mamba")]
1723            Self::Mamba(m) => m.update(symbol),
1724            #[cfg(feature = "backend-rwkv")]
1725            Self::Rwkv(m) => m.update(symbol),
1726            #[cfg(feature = "backend-zpaq")]
1727            Self::Zpaq(m) => {
1728                m.update(symbol);
1729                Ok(())
1730            }
1731            #[cfg(feature = "backend-mixture")]
1732            Self::Mixture(m) => m.update(symbol),
1733            #[cfg(feature = "backend-particle")]
1734            Self::Particle(m) => {
1735                m.step(symbol);
1736                Ok(())
1737            }
1738            #[cfg(feature = "backend-calibrated")]
1739            Self::Calibrated {
1740                base,
1741                core,
1742                bitwise: _,
1743                valid,
1744                ..
1745            } => {
1746                let base_pdf = base.pdf_next()?;
1747                core.observe_symbol_from_base_pdf(symbol, base_pdf)
1748                    .map_err(anyhow::Error::msg)?;
1749                base.update(symbol)?;
1750                *valid = false;
1751                Ok(())
1752            }
1753            Self::Disabled { reason } => bail!("{reason}"),
1754        }
1755    }
1756
1757    fn prepare_cached_cdf_fast_bitwise(&mut self) -> Result<bool> {
1758        match self {
1759            #[cfg(feature = "backend-rosa")]
1760            Self::Rosa(m) => {
1761                let _ = m.cdf_next();
1762                Ok(true)
1763            }
1764            #[cfg(feature = "backend-match")]
1765            Self::Match { model } => {
1766                let _ = model.cdf();
1767                Ok(true)
1768            }
1769            #[cfg(feature = "backend-match")]
1770            Self::SparseMatch { model } => {
1771                let _ = model.cdf();
1772                Ok(true)
1773            }
1774            #[cfg(feature = "backend-ppmd")]
1775            Self::Ppmd { model } => {
1776                let _ = model.cdf();
1777                Ok(true)
1778            }
1779            #[cfg(feature = "backend-mamba")]
1780            Self::Mamba(m) => {
1781                let _ = m.cdf_next();
1782                Ok(true)
1783            }
1784            #[cfg(feature = "backend-rwkv")]
1785            Self::Rwkv(m) => {
1786                let _ = m.cdf_next();
1787                Ok(true)
1788            }
1789            _ => Ok(false),
1790        }
1791    }
1792
1793    fn cached_cdf_bit_prob_one_msb(&mut self, range: MsbPrefixRange) -> Option<f64> {
1794        match self {
1795            #[cfg(feature = "backend-rosa")]
1796            Self::Rosa(m) => Some(range.prob_one(&m.cdf, PDF_MIN)),
1797            #[cfg(feature = "backend-match")]
1798            Self::Match { model } => Some(range.prob_one(model.cdf(), PDF_MIN)),
1799            #[cfg(feature = "backend-match")]
1800            Self::SparseMatch { model } => Some(range.prob_one(model.cdf(), PDF_MIN)),
1801            #[cfg(feature = "backend-ppmd")]
1802            Self::Ppmd { model } => Some(range.prob_one(model.cdf(), PDF_MIN)),
1803            #[cfg(feature = "backend-mamba")]
1804            Self::Mamba(m) => Some(range.prob_one(m.cdf_next(), PDF_MIN)),
1805            #[cfg(feature = "backend-rwkv")]
1806            Self::Rwkv(m) => Some(range.prob_one(m.cdf_next(), PDF_MIN)),
1807            _ => None,
1808        }
1809    }
1810
1811    #[inline]
1812    fn has_recursive_native_bitwise_path(&self) -> bool {
1813        match self {
1814            #[cfg(feature = "backend-ctw")]
1815            Self::Ctw(m) | Self::FacCtw(m) => m.can_fast_ac_bitwise(),
1816            #[cfg(feature = "backend-mixture")]
1817            Self::Mixture(m) => m.has_recursive_native_bitwise_expert(),
1818            #[cfg(feature = "backend-calibrated")]
1819            Self::Calibrated { .. } => true,
1820            _ => false,
1821        }
1822    }
1823
1824    fn begin_native_recursive_bitwise_byte_step(&mut self) -> Result<bool> {
1825        match self {
1826            #[cfg(feature = "backend-ctw")]
1827            Self::Ctw(m) | Self::FacCtw(m) => Ok(m.can_fast_ac_bitwise()),
1828            #[cfg(feature = "backend-mixture")]
1829            Self::Mixture(m) => m.begin_bitwise_byte_step(),
1830            #[cfg(feature = "backend-calibrated")]
1831            Self::Calibrated {
1832                base,
1833                core,
1834                bitwise,
1835                valid,
1836                ..
1837            } => {
1838                core.begin_byte().map_err(anyhow::Error::msg)?;
1839                if let Err(err) = bitwise.prepare(base) {
1840                    let _ = core.abort_empty_byte();
1841                    return Err(err);
1842                }
1843                *valid = false;
1844                Ok(true)
1845            }
1846            _ => Ok(false),
1847        }
1848    }
1849
1850    fn native_recursive_bit_prob_one_msb(&mut self, bit_idx: usize) -> Result<f64> {
1851        match self {
1852            #[cfg(feature = "backend-ctw")]
1853            Self::Ctw(m) | Self::FacCtw(m) => Ok(m.bit_prob_one_msb(bit_idx)),
1854            #[cfg(feature = "backend-mixture")]
1855            Self::Mixture(m) => m.bit_prob_one_msb(bit_idx),
1856            #[cfg(feature = "backend-calibrated")]
1857            Self::Calibrated {
1858                base,
1859                core,
1860                bitwise,
1861                ..
1862            } => {
1863                let base_p1: f64 = bitwise.bit_prob_one_msb(base, bit_idx)?;
1864                debug_assert!(core.byte_is_active());
1865                Ok(core.predict_bit_unchecked(base_p1))
1866            }
1867            _ => bail!("native recursive bitwise stepping is unavailable for this predictor"),
1868        }
1869    }
1870
1871    fn native_recursive_observe_bit_msb(&mut self, bit_idx: usize, bit: bool) -> Result<()> {
1872        match self {
1873            #[cfg(feature = "backend-ctw")]
1874            Self::Ctw(m) | Self::FacCtw(m) => {
1875                m.update_bit_msb(bit_idx, bit);
1876                Ok(())
1877            }
1878            #[cfg(feature = "backend-mixture")]
1879            Self::Mixture(m) => m.observe_bit_msb(bit_idx, bit),
1880            #[cfg(feature = "backend-calibrated")]
1881            Self::Calibrated {
1882                base,
1883                core,
1884                bitwise,
1885                valid,
1886                ..
1887            } => {
1888                debug_assert!(core.byte_is_active());
1889                core.observe_bit_unchecked(bit);
1890                bitwise.observe_bit_msb(base, bit_idx, bit)?;
1891                *valid = false;
1892                Ok(())
1893            }
1894            _ => bail!("native recursive bitwise stepping is unavailable for this predictor"),
1895        }
1896    }
1897
1898    fn finish_native_recursive_bitwise_byte_step(&mut self, symbol: u8) -> Result<()> {
1899        match self {
1900            #[cfg(feature = "backend-ctw")]
1901            Self::Ctw(_) | Self::FacCtw(_) => Ok(()),
1902            #[cfg(feature = "backend-mixture")]
1903            Self::Mixture(m) => m.finish_bitwise_symbol(symbol),
1904            #[cfg(feature = "backend-calibrated")]
1905            Self::Calibrated {
1906                base,
1907                core,
1908                bitwise,
1909                valid,
1910                ..
1911            } => {
1912                core.validate_complete_byte().map_err(anyhow::Error::msg)?;
1913                bitwise.finish_symbol(base, symbol)?;
1914                core.finish_byte().map_err(anyhow::Error::msg)?;
1915                *valid = false;
1916                Ok(())
1917            }
1918            _ => bail!("native recursive bitwise stepping is unavailable for this predictor"),
1919        }
1920    }
1921
1922    #[inline]
1923    fn can_fast_ac_bitwise(&self) -> bool {
1924        self.has_recursive_native_bitwise_path()
1925    }
1926
1927    // Keep this separate from the live AC payload path so framed AC preserves
1928    // the v1 wire contract while still allowing backend-agnostic bitwise
1929    // stepping as an internal utility.
1930    fn ac_step_bitwise<F>(&mut self, mut choose_bit: F) -> Result<u8>
1931    where
1932        F: FnMut(usize, f64) -> Result<u8>,
1933    {
1934        let mut state = PredictorBitwiseStepState::default();
1935        state.prepare(self)?;
1936        let mut symbol = 0u8;
1937        for bit_idx in 0..8usize {
1938            let p1 = state.bit_prob_one_msb(self, bit_idx)?;
1939            let bit = choose_bit(bit_idx, p1)? & 1;
1940            if bit == 1 {
1941                symbol |= 1u8 << (7 - bit_idx);
1942            }
1943            state.observe_bit_msb(self, bit_idx, bit == 1)?;
1944        }
1945        state.finish_symbol(self, symbol)?;
1946        Ok(symbol)
1947    }
1948
1949    fn ac_step_fast_bitwise<F>(&mut self, choose_bit: F) -> Result<u8>
1950    where
1951        F: FnMut(usize, f64) -> Result<u8>,
1952    {
1953        debug_assert!(self.can_fast_ac_bitwise());
1954        self.ac_step_bitwise(choose_bit)
1955    }
1956
1957    fn diagnostic_snapshot_subtree(
1958        &mut self,
1959        symbol: u8,
1960        local_weight: f64,
1961        effective_weight: f64,
1962        pool: Option<&ThreadPool>,
1963    ) -> Result<AcLogLossSubtreeSnapshot> {
1964        match self {
1965            #[cfg(feature = "backend-mixture")]
1966            Self::Mixture(m) => {
1967                m.diagnostic_subtree_snapshot(symbol, local_weight, effective_weight, pool)
1968            }
1969            _ => {
1970                let prob = self.pdf_next()?[symbol as usize].max(PDF_MIN);
1971                Ok(AcLogLossSubtreeSnapshot {
1972                    prob,
1973                    rows: vec![AcLogLossNodeValue {
1974                        prob,
1975                        local_weight,
1976                        effective_weight,
1977                    }],
1978                })
1979            }
1980        }
1981    }
1982
1983    // The mixture implementation mutates the Vec allocation; no-mixture builds
1984    // only see this forwarding signature and would otherwise flag it as `ptr_arg`.
1985    #[allow(clippy::ptr_arg)]
1986    fn diagnostic_root_snapshot(
1987        &mut self,
1988        symbol: u8,
1989        pool: Option<&ThreadPool>,
1990        out: &mut Vec<AcLogLossNodeValue>,
1991    ) -> Result<AcLogLossRootSnapshot> {
1992        match self {
1993            #[cfg(feature = "backend-mixture")]
1994            Self::Mixture(m) => m.diagnostic_root_snapshot(symbol, pool, out),
1995            _ => anyhow::bail!("AC log-loss diagnostics require a top-level mixture backend"),
1996        }
1997    }
1998
1999    fn encode_symbol_ac_step<W: std::io::Write>(
2000        &mut self,
2001        symbol: u8,
2002        encoder: &mut ArithmeticEncoder<W>,
2003        cdf: &mut [u32; 257],
2004    ) -> Result<()> {
2005        if self.can_fast_ac_bitwise() {
2006            self.ac_step_fast_bitwise(|bit_idx, p1_mix| {
2007                let bit = (symbol >> (7 - bit_idx)) & 1;
2008                let split = binary_split_from_prob_one(p1_mix);
2009                if bit == 0 {
2010                    encoder.encode_counts(0, split as u64, CDF_TOTAL as u64)?;
2011                } else {
2012                    encoder.encode_counts(split as u64, CDF_TOTAL as u64, CDF_TOTAL as u64)?;
2013                }
2014                Ok(bit)
2015            })?;
2016            return Ok(());
2017        }
2018
2019        let pdf = self.pdf_next()?;
2020        crate::coders::quantize_pdf_to_integer_cdf_dense_positive_with_buffer(
2021            pdf,
2022            CDF_TOTAL,
2023            cdf.as_mut_slice(),
2024        );
2025        let sym = symbol as usize;
2026        encoder.encode_counts(cdf[sym] as u64, cdf[sym + 1] as u64, CDF_TOTAL as u64)?;
2027        self.update(symbol)
2028    }
2029}
2030
2031#[inline]
2032fn binary_split_from_prob_one(p1: f64) -> u32 {
2033    let p1 = p1.clamp(PDF_MIN, 1.0 - PDF_MIN);
2034    let p0 = 1.0 - p1;
2035    let mut split = (p0 * (CDF_TOTAL as f64)) as u32;
2036    if split == 0 {
2037        split = 1;
2038    } else if split >= CDF_TOTAL {
2039        split = CDF_TOTAL - 1;
2040    }
2041    split
2042}
2043
2044/// This function should be considered when fine-tuning Compression/decompression for a particular runtime case. In particular, my benchmarking has shown that inlining is non-obvious in how it affects performance
2045/// Inlining both encode and decode seems to cause performance issues with Match+AC decompression specifically, hence the odd configuration here for balance.
2046/// Encode default: inline
2047/// Technical note: this fast-path preserves the same bit ordering and CDF split mapping as the generic AC path (MSB-first with `binary_split_from_prob_one`).
2048#[cfg_attr(not(infotheory_ac_encode_deinline), inline(always))]
2049#[cfg_attr(infotheory_ac_encode_deinline, inline(never))]
2050fn encode_payload_ac_fast_bitwise(
2051    data: &[u8],
2052    predictor: &mut RatePdfPredictor,
2053) -> Result<Vec<u8>> {
2054    let mut out = Vec::new();
2055    {
2056        let mut enc = ArithmeticEncoder::new(&mut out);
2057        for &symbol in data {
2058            predictor.ac_step_fast_bitwise(|bit_idx, p1_mix| {
2059                let bit = (symbol >> (7 - bit_idx)) & 1;
2060                let split = binary_split_from_prob_one(p1_mix);
2061                if bit == 0 {
2062                    enc.encode_counts(0, split as u64, CDF_TOTAL as u64)?;
2063                } else {
2064                    enc.encode_counts(split as u64, CDF_TOTAL as u64, CDF_TOTAL as u64)?;
2065                }
2066                Ok(bit)
2067            })?;
2068        }
2069        let _ = enc.finish()?;
2070    }
2071    Ok(out)
2072}
2073
2074/// This function should be considered when fine-tuning Compression/decompression for a particular runtime case. In particular, my benchmarking has shown that inlining is non-obvious in how it affects performance
2075/// Inlining both encode and decode seems to cause performance issues with Match+AC decompression specifically, hence the odd configuration here for balance.
2076/// Decode default: deinline
2077/// Technical note: this decodes exactly `out_len` symbols from the same binary CDF domain (`CDF_TOTAL`) used by the paired encode fast-path.
2078#[cfg_attr(infotheory_ac_decode_inline, inline(always))]
2079#[cfg_attr(not(infotheory_ac_decode_inline), inline(never))]
2080fn decode_payload_ac_fast_bitwise(
2081    payload: &[u8],
2082    out_len: usize,
2083    predictor: &mut RatePdfPredictor,
2084) -> Result<Vec<u8>> {
2085    let mut dec = ArithmeticDecoder::new(payload)?;
2086    let mut out = Vec::with_capacity(out_len);
2087    for _ in 0..out_len {
2088        let symbol = predictor.ac_step_fast_bitwise(|_, p1_mix| {
2089            let split = binary_split_from_prob_one(p1_mix);
2090            dec.decode_binary_counts(split, CDF_TOTAL)
2091        })?;
2092        out.push(symbol);
2093    }
2094    Ok(out)
2095}
2096
2097fn encode_payload_ac(data: &[u8], predictor: &mut RatePdfPredictor) -> Result<Vec<u8>> {
2098    predictor.begin_stream(data.len())?;
2099
2100    if predictor.can_fast_ac_bitwise() {
2101        let out = encode_payload_ac_fast_bitwise(data, predictor)?;
2102        predictor.finish_stream()?;
2103        return Ok(out);
2104    }
2105
2106    let mut out = Vec::new();
2107    {
2108        let mut enc = ArithmeticEncoder::new(&mut out);
2109        // Reuse one CDF scratch buffer for the full stream to avoid per-symbol allocation.
2110        let mut cdf = [0u32; 257];
2111        for &symbol in data {
2112            let pdf = predictor.pdf_next()?;
2113            crate::coders::quantize_pdf_to_integer_cdf_dense_positive_with_buffer(
2114                pdf,
2115                CDF_TOTAL,
2116                cdf.as_mut_slice(),
2117            );
2118            let sym = symbol as usize;
2119            enc.encode_counts(cdf[sym] as u64, cdf[sym + 1] as u64, CDF_TOTAL as u64)?;
2120            predictor.update(symbol)?;
2121        }
2122        let _ = enc.finish()?;
2123    }
2124    predictor.finish_stream()?;
2125    Ok(out)
2126}
2127
2128fn decode_payload_ac(
2129    payload: &[u8],
2130    out_len: usize,
2131    predictor: &mut RatePdfPredictor,
2132) -> Result<Vec<u8>> {
2133    predictor.begin_stream(out_len)?;
2134
2135    if predictor.can_fast_ac_bitwise() {
2136        let out = decode_payload_ac_fast_bitwise(payload, out_len, predictor)?;
2137        predictor.finish_stream()?;
2138        return Ok(out);
2139    }
2140
2141    let mut dec = ArithmeticDecoder::new(payload)?;
2142    let mut out = Vec::with_capacity(out_len);
2143    let mut cdf = vec![0u32; 257];
2144    for _ in 0..out_len {
2145        let pdf = predictor.pdf_next()?;
2146        crate::coders::quantize_pdf_to_integer_cdf_dense_positive_with_buffer(
2147            pdf, CDF_TOTAL, &mut cdf,
2148        );
2149        let sym = dec.decode_symbol_counts(&cdf, CDF_TOTAL)? as u8;
2150        out.push(sym);
2151        predictor.update(sym)?;
2152    }
2153    predictor.finish_stream()?;
2154    Ok(out)
2155}
2156
2157fn encode_payload_rans(data: &[u8], predictor: &mut RatePdfPredictor) -> Result<Vec<u8>> {
2158    predictor.begin_stream(data.len())?;
2159    let mut encoder = BlockedRansEncoder::new();
2160    let mut cdf = vec![0u32; 257];
2161    let mut freq = vec![0i64; 256];
2162
2163    for &b in data {
2164        let pdf = predictor.pdf_next()?;
2165        quantize_pdf_to_rans_cdf_with_buffer(pdf, &mut cdf, &mut freq);
2166        let s = b as usize;
2167        encoder.encode(Cdf::new(cdf[s], cdf[s + 1], ANS_TOTAL));
2168        predictor.update(b)?;
2169    }
2170
2171    let blocks = encoder.finish();
2172    let mut out = Vec::new();
2173    out.extend_from_slice(&(blocks.len() as u32).to_le_bytes());
2174    for block in blocks {
2175        out.extend_from_slice(&(block.len() as u32).to_le_bytes());
2176        out.extend_from_slice(&block);
2177    }
2178    predictor.finish_stream()?;
2179    Ok(out)
2180}
2181
2182fn decode_payload_rans(
2183    payload: &[u8],
2184    out_len: usize,
2185    predictor: &mut RatePdfPredictor,
2186) -> Result<Vec<u8>> {
2187    predictor.begin_stream(out_len)?;
2188    if payload.len() < 4 {
2189        bail!("rANS payload too short");
2190    }
2191    let block_count = u32::from_le_bytes([payload[0], payload[1], payload[2], payload[3]]) as usize;
2192    let mut pos = 4usize;
2193    let mut blocks = Vec::with_capacity(block_count);
2194    for _ in 0..block_count {
2195        if pos + 4 > payload.len() {
2196            bail!("truncated rANS block header");
2197        }
2198        let len = u32::from_le_bytes([
2199            payload[pos],
2200            payload[pos + 1],
2201            payload[pos + 2],
2202            payload[pos + 3],
2203        ]) as usize;
2204        pos += 4;
2205        if pos + len > payload.len() {
2206            bail!("truncated rANS block data");
2207        }
2208        blocks.push(&payload[pos..pos + len]);
2209        pos += len;
2210    }
2211
2212    let mut dec = BlockedRansDecoder::new(blocks, out_len)?;
2213    let mut out = Vec::with_capacity(out_len);
2214    let mut cdf = vec![0u32; 257];
2215    let mut freq = vec![0i64; 256];
2216
2217    for _ in 0..out_len {
2218        let pdf = predictor.pdf_next()?;
2219        quantize_pdf_to_rans_cdf_with_buffer(pdf, &mut cdf, &mut freq);
2220        let sym = dec.decode(&cdf)? as u8;
2221        out.push(sym);
2222        predictor.update(sym)?;
2223    }
2224    predictor.finish_stream()?;
2225    Ok(out)
2226}
2227
2228/// Compress bytes using a predictive rate backend and entropy coder.
2229///
2230/// When `framing` is [`FramingMode::Framed`], output includes a compact header
2231/// with payload metadata and CRC for safer transport/storage.
2232pub fn compress_rate_bytes(
2233    data: &[u8],
2234    rate_backend: &CompiledRateBackend,
2235    coder: CoderType,
2236    framing: FramingMode,
2237) -> Result<Vec<u8>> {
2238    let mut predictor = crate::runtime::build_rate_pdf_predictor(rate_backend)?;
2239    let payload = match coder {
2240        CoderType::AC => encode_payload_ac(data, &mut predictor)?,
2241        CoderType::RANS => encode_payload_rans(data, &mut predictor)?,
2242    };
2243
2244    if framing == FramingMode::Raw {
2245        return Ok(payload);
2246    }
2247
2248    let mut out = Vec::with_capacity(FramedHeader::SIZE + payload.len());
2249    let hdr = FramedHeader::new(coder, data.len() as u64, crc32(data));
2250    hdr.write(&mut out);
2251    out.extend_from_slice(&payload);
2252    Ok(out)
2253}
2254
2255/// Return compressed size (in bytes) for `data` using rate coding.
2256pub fn compress_rate_size(
2257    data: &[u8],
2258    rate_backend: &CompiledRateBackend,
2259    coder: CoderType,
2260    framing: FramingMode,
2261) -> Result<u64> {
2262    let encoded = compress_rate_bytes(data, rate_backend, coder, framing)?;
2263    Ok(encoded.len() as u64)
2264}
2265
2266/// Return compressed size (in bytes) for concatenated slices under one stream.
2267pub fn compress_rate_size_chain(
2268    parts: &[&[u8]],
2269    rate_backend: &CompiledRateBackend,
2270    coder: CoderType,
2271    framing: FramingMode,
2272) -> Result<u64> {
2273    let total = parts.iter().map(|p| p.len()).sum();
2274    let mut data = Vec::with_capacity(total);
2275    for p in parts {
2276        data.extend_from_slice(p);
2277    }
2278    compress_rate_size(&data, rate_backend, coder, framing)
2279}
2280
2281/// Decompress bytes produced by [`compress_rate_bytes`].
2282pub fn decompress_rate_bytes(
2283    input: &[u8],
2284    rate_backend: &CompiledRateBackend,
2285    _coder: CoderType,
2286    framing: FramingMode,
2287) -> Result<Vec<u8>> {
2288    let (payload, coder, out_len, expected_crc) = if framing == FramingMode::Framed {
2289        let hdr = FramedHeader::read(input)?;
2290        (
2291            &input[FramedHeader::SIZE..],
2292            hdr.coder_type(),
2293            hdr.original_len as usize,
2294            Some(hdr.crc32),
2295        )
2296    } else {
2297        bail!("raw payload decompression requires explicit output length and is not supported");
2298    };
2299
2300    let _ = coder;
2301    let mut predictor = crate::runtime::build_rate_pdf_predictor(rate_backend)?;
2302    let decoded = match coder {
2303        CoderType::AC => decode_payload_ac(payload, out_len, &mut predictor)?,
2304        CoderType::RANS => decode_payload_rans(payload, out_len, &mut predictor)?,
2305    };
2306
2307    if let Some(crc) = expected_crc {
2308        let got = crc32(&decoded);
2309        if got != crc {
2310            bail!("CRC32 mismatch: expected 0x{crc:08X}, got 0x{got:08X}");
2311        }
2312    }
2313
2314    Ok(decoded)
2315}
2316
2317#[inline]
2318fn uniform_cdf_row() -> [f64; 257] {
2319    let mut cdf = [0.0; 257];
2320    let inv = 1.0 / 256.0;
2321    for (i, slot) in cdf.iter_mut().enumerate() {
2322        *slot = (i as f64) * inv;
2323    }
2324    cdf
2325}
2326
2327#[inline]
2328fn build_cdf_row_from_pdf_slice(pdf: &[f64], cdf: &mut [f64; 257]) {
2329    cdf[0] = 0.0;
2330    let mut acc = 0.0;
2331    for i in 0..256 {
2332        acc += pdf[i];
2333        cdf[i + 1] = acc;
2334    }
2335}
2336
2337fn normalize_pdf_vec_and_maybe_build_cdf(pdf: &mut [f64], cdf: Option<&mut [f64; 257]>) {
2338    let mut sum = 0.0;
2339    for p in pdf.iter_mut() {
2340        *p = if p.is_finite() {
2341            (*p).max(PDF_MIN)
2342        } else {
2343            PDF_MIN
2344        };
2345        sum += *p;
2346    }
2347    if !(sum.is_finite()) || sum <= 0.0 {
2348        let u = 1.0 / (pdf.len() as f64);
2349        pdf.fill(u);
2350        if let Some(cdf) = cdf {
2351            *cdf = uniform_cdf_row();
2352        }
2353        return;
2354    }
2355    let inv = 1.0 / sum;
2356    if let Some(cdf) = cdf {
2357        cdf[0] = 0.0;
2358        let mut acc = 0.0;
2359        for i in 0..256 {
2360            pdf[i] *= inv;
2361            acc += pdf[i];
2362            cdf[i + 1] = acc;
2363        }
2364    } else {
2365        for p in pdf.iter_mut() {
2366            *p *= inv;
2367        }
2368    }
2369}
2370
2371#[inline]
2372fn logsumexp_slice(vals: &[f64]) -> f64 {
2373    let mut m = f64::NEG_INFINITY;
2374    for &v in vals {
2375        if v > m {
2376            m = v;
2377        }
2378    }
2379    if !m.is_finite() {
2380        return m;
2381    }
2382    let mut s = 0.0;
2383    for &v in vals {
2384        s += (v - m).exp();
2385    }
2386    m + s.ln()
2387}
2388
2389#[inline]
2390fn logsumexp_expert_weights(experts: &[MixExpert]) -> f64 {
2391    let mut m = f64::NEG_INFINITY;
2392    for e in experts {
2393        if e.log_weight > m {
2394            m = e.log_weight;
2395        }
2396    }
2397    if !m.is_finite() {
2398        return m;
2399    }
2400    let mut s = 0.0;
2401    for e in experts {
2402        s += (e.log_weight - m).exp();
2403    }
2404    m + s.ln()
2405}
2406
2407fn normalize_simplex_weights(weights: &mut [f64]) {
2408    if weights.is_empty() {
2409        return;
2410    }
2411    let mut sum = 0.0;
2412    for weight in weights.iter_mut() {
2413        if !weight.is_finite() || *weight < 0.0 {
2414            *weight = 0.0;
2415        }
2416        sum += *weight;
2417    }
2418    if !sum.is_finite() || sum <= 0.0 {
2419        let uniform = 1.0 / (weights.len() as f64);
2420        weights.fill(uniform);
2421        return;
2422    }
2423    for weight in weights.iter_mut() {
2424        *weight /= sum;
2425    }
2426}
2427
2428fn normalized_mix_expert_prior_weights(experts: &[MixExpert], out: &mut [f64]) {
2429    debug_assert_eq!(experts.len(), out.len());
2430    let max_log = experts
2431        .iter()
2432        .map(|expert| expert.log_prior)
2433        .fold(f64::NEG_INFINITY, f64::max);
2434    for (slot, expert) in out.iter_mut().zip(experts.iter()) {
2435        *slot = if max_log.is_finite() {
2436            (expert.log_prior - max_log).exp()
2437        } else {
2438            0.0
2439        };
2440    }
2441    normalize_simplex_weights(out);
2442}
2443
2444fn set_mix_expert_log_weights_from_linear(experts: &mut [MixExpert], weights: &[f64]) {
2445    for (expert, &weight) in experts.iter_mut().zip(weights.iter()) {
2446        expert.log_weight = if weight > 0.0 {
2447            weight.ln()
2448        } else {
2449            f64::NEG_INFINITY
2450        };
2451    }
2452}
2453
2454fn apply_switching_weights(
2455    experts: &mut [MixExpert],
2456    prior_weights: &[f64],
2457    alpha: f64,
2458    posterior: &mut [f64],
2459    scratch: &mut [f64],
2460) {
2461    if experts.is_empty() {
2462        return;
2463    }
2464    debug_assert_eq!(experts.len(), prior_weights.len());
2465
2466    normalize_simplex_weights(posterior);
2467    if experts.len() == 1 || alpha <= 0.0 {
2468        set_mix_expert_log_weights_from_linear(experts, posterior);
2469        return;
2470    }
2471
2472    let num_switch_targets = prior_weights.iter().filter(|&&prior| prior < 1.0).count();
2473    if num_switch_targets <= 1 {
2474        set_mix_expert_log_weights_from_linear(experts, posterior);
2475        return;
2476    }
2477
2478    let mut switch_out_sum = 0.0;
2479    for i in 0..experts.len() {
2480        let denom = 1.0 - prior_weights[i];
2481        if denom > 0.0 {
2482            switch_out_sum += posterior[i] / denom;
2483        }
2484    }
2485
2486    for i in 0..experts.len() {
2487        let prior = prior_weights[i];
2488        let stay = (1.0 - alpha) * posterior[i];
2489        let switch_in = if prior > 0.0 {
2490            let denom = 1.0 - prior;
2491            let switchable_mass = if denom > 0.0 {
2492                switch_out_sum - posterior[i] / denom
2493            } else {
2494                0.0
2495            };
2496            alpha * prior * switchable_mass
2497        } else {
2498            0.0
2499        };
2500        scratch[i] = stay + switch_in;
2501    }
2502
2503    normalize_simplex_weights(scratch);
2504    set_mix_expert_log_weights_from_linear(experts, scratch);
2505}
2506
2507#[allow(dead_code)]
2508#[cfg(feature = "backend-zpaq")]
2509fn _zpaq_marker(_: &ZpaqRateModel) {}
2510
2511#[cfg(all(test, feature = "all-backends"))]
2512mod tests {
2513    use super::*;
2514    use std::sync::Arc;
2515
2516    fn compiled_rate_backend(backend: &RateBackend) -> CompiledRateBackend {
2517        backend
2518            .compile()
2519            .unwrap_or_else(|err| panic!("failed to compile rate backend for test: {err}"))
2520    }
2521
2522    fn compress_rate_bytes(
2523        data: &[u8],
2524        rate_backend: &RateBackend,
2525        coder: CoderType,
2526        framing: FramingMode,
2527    ) -> Result<Vec<u8>> {
2528        super::compress_rate_bytes(data, &compiled_rate_backend(rate_backend), coder, framing)
2529    }
2530
2531    fn compress_rate_size(
2532        data: &[u8],
2533        rate_backend: &RateBackend,
2534        coder: CoderType,
2535        framing: FramingMode,
2536    ) -> Result<u64> {
2537        super::compress_rate_size(data, &compiled_rate_backend(rate_backend), coder, framing)
2538    }
2539
2540    fn decompress_rate_bytes(
2541        input: &[u8],
2542        rate_backend: &RateBackend,
2543        coder: CoderType,
2544        framing: FramingMode,
2545    ) -> Result<Vec<u8>> {
2546        super::decompress_rate_bytes(input, &compiled_rate_backend(rate_backend), coder, framing)
2547    }
2548
2549    fn assert_pdf_close(lhs: &[f64], rhs: &[f64], tol: f64) {
2550        assert_eq!(lhs.len(), rhs.len());
2551        for (idx, (&a, &b)) in lhs.iter().zip(rhs.iter()).enumerate() {
2552            let delta = (a - b).abs();
2553            assert!(
2554                delta <= tol,
2555                "pdf mismatch at symbol {idx}: lhs={a} rhs={b} delta={delta}"
2556            );
2557        }
2558    }
2559
2560    fn brute_force_pdf(predictor: &mut CtwPredictor) -> Vec<f64> {
2561        let bits = predictor.bits_per_symbol.clamp(1, 8);
2562        let mut out = vec![0.0; 256];
2563
2564        if bits == 8 {
2565            for (sym, slot) in out.iter_mut().enumerate().take(256usize) {
2566                *slot = predictor.log_prob_symbol_bruteforce(sym as u8).exp();
2567            }
2568        } else {
2569            let patterns = 1usize << bits;
2570            let aliases = 1usize << (8 - bits);
2571            let mut pat_prob = vec![0.0; patterns];
2572            for (pat, value) in pat_prob.iter_mut().enumerate() {
2573                let symbol = if predictor.msb_first {
2574                    (pat as u8) << (8 - bits)
2575                } else {
2576                    pat as u8
2577                };
2578                *value = predictor.log_prob_symbol_bruteforce(symbol).exp();
2579            }
2580            for (byte, slot) in out.iter_mut().enumerate().take(256usize) {
2581                let pat = if predictor.msb_first {
2582                    byte >> (8 - bits)
2583                } else {
2584                    byte & (patterns - 1)
2585                };
2586                *slot = pat_prob[pat] / (aliases as f64);
2587            }
2588        }
2589
2590        normalize_pdf(&mut out, PDF_MIN);
2591        out
2592    }
2593
2594    #[test]
2595    fn ctw_pdf_fast_matches_bruteforce() {
2596        let mut predictor = CtwPredictor::new_ctw(6);
2597        for &b in b"ctw fast-path regression corpus 1234567890" {
2598            predictor.update(b);
2599        }
2600
2601        let fast = predictor.pdf_next().to_vec();
2602        predictor.valid = false;
2603        let brute = brute_force_pdf(&mut predictor);
2604
2605        for i in 0..256usize {
2606            let delta = (fast[i] - brute[i]).abs();
2607            assert!(
2608                delta < 1e-12,
2609                "symbol={i} fast={} brute={} delta={delta}",
2610                fast[i],
2611                brute[i]
2612            );
2613        }
2614    }
2615
2616    #[test]
2617    fn fac_pdf_fast_matches_bruteforce_subbyte() {
2618        let mut predictor = CtwPredictor::new_fac(5, 5, None);
2619        for &b in b"fac ctw subbyte regression corpus abcdefghijklmnopqrstuvwxyz" {
2620            predictor.update(b);
2621        }
2622
2623        let fast = predictor.pdf_next().to_vec();
2624        predictor.valid = false;
2625        let brute = brute_force_pdf(&mut predictor);
2626
2627        for i in 0..256usize {
2628            let delta = (fast[i] - brute[i]).abs();
2629            assert!(
2630                delta < 1e-12,
2631                "symbol={i} fast={} brute={} delta={delta}",
2632                fast[i],
2633                brute[i]
2634            );
2635        }
2636    }
2637
2638    #[test]
2639    fn fac_ctw_default_bit_order_is_byte_msb_and_subbyte_lsb() {
2640        let byte_default = CtwPredictor::new_fac(5, 8, None);
2641        assert!(
2642            byte_default.can_fast_ac_bitwise(),
2643            "8-bit FacCtw without explicit order should use MSB-first native bitwise path"
2644        );
2645
2646        let subbyte_default = CtwPredictor::new_fac(5, 5, None);
2647        assert!(
2648            !subbyte_default.can_fast_ac_bitwise(),
2649            "non-byte FacCtw without explicit order keeps legacy LSB-first behavior"
2650        );
2651
2652        let explicit_lsb = CtwPredictor::new_fac(5, 8, Some(false));
2653        assert!(
2654            !explicit_lsb.can_fast_ac_bitwise(),
2655            "explicit msb_first=false must preserve legacy LSB-first behavior"
2656        );
2657
2658        let explicit_subbyte_msb = CtwPredictor::new_fac(5, 5, Some(true));
2659        assert!(
2660            !explicit_subbyte_msb.can_fast_ac_bitwise(),
2661            "subbyte widths do not use byte-packed native fast path even when MSB-first"
2662        );
2663    }
2664
2665    fn assert_ctw_pdf_next_preserves_state(mut predictor: CtwPredictor) {
2666        for &b in b"ctw predictor state preservation payload" {
2667            predictor.update(b);
2668        }
2669        let mut baseline = [0.0f64; 256];
2670        for (sym, slot) in baseline.iter_mut().enumerate() {
2671            *slot = predictor.log_prob_symbol_bruteforce(sym as u8);
2672        }
2673        let _ = predictor.pdf_next();
2674        for (sym, &expected) in baseline.iter().enumerate() {
2675            let after = predictor.log_prob_symbol_bruteforce(sym as u8);
2676            assert!(
2677                (expected - after).abs() < 1e-12,
2678                "symbol {sym} drift: {expected} vs {after}"
2679            );
2680        }
2681    }
2682
2683    #[test]
2684    fn ctw_pdf_next_preserves_state() {
2685        assert_ctw_pdf_next_preserves_state(CtwPredictor::new_ctw(7));
2686    }
2687
2688    #[test]
2689    fn fac_pdf_next_preserves_state() {
2690        assert_ctw_pdf_next_preserves_state(CtwPredictor::new_fac(7, 8, None));
2691    }
2692
2693    fn assert_fill_pattern_preserves_symbol_log_probs(mut predictor: CtwPredictor) {
2694        for &b in b"fill-pattern preservation regression payload" {
2695            predictor.update(b);
2696        }
2697        let mut baseline = [0.0f64; 256];
2698        for (sym, slot) in baseline.iter_mut().enumerate() {
2699            *slot = predictor.log_prob_symbol_bruteforce(sym as u8);
2700        }
2701        let _ = predictor.fill_pattern_log_probs();
2702        for (sym, &expected) in baseline.iter().enumerate() {
2703            let got = predictor.log_prob_symbol_bruteforce(sym as u8);
2704            let diff = (expected - got).abs();
2705            assert!(
2706                diff < 1e-12,
2707                "symbol={sym} expected={expected} got={got} diff={diff}"
2708            );
2709        }
2710    }
2711
2712    #[test]
2713    fn ctw_fill_pattern_preserves_symbol_log_probs() {
2714        assert_fill_pattern_preserves_symbol_log_probs(CtwPredictor::new_ctw(7));
2715    }
2716
2717    #[test]
2718    fn fac_fill_pattern_preserves_symbol_log_probs() {
2719        assert_fill_pattern_preserves_symbol_log_probs(CtwPredictor::new_fac(7, 8, None));
2720    }
2721
2722    fn assert_pdf_then_update_matches_plain_update(mut base: CtwPredictor) {
2723        for &b in b"pdf then update parity payload" {
2724            base.update(b);
2725        }
2726        let observed = b'n';
2727        let mut with_pdf = base.clone();
2728        let mut plain = base;
2729
2730        let _ = with_pdf.pdf_next();
2731        with_pdf.update(observed);
2732        plain.update(observed);
2733
2734        for sym in 0u8..=255u8 {
2735            let lp_with_pdf = with_pdf.log_prob_symbol_bruteforce(sym);
2736            let lp_plain = plain.log_prob_symbol_bruteforce(sym);
2737            let diff = (lp_with_pdf - lp_plain).abs();
2738            assert!(
2739                diff < 1e-12,
2740                "symbol={sym} with_pdf={lp_with_pdf} plain={lp_plain} diff={diff}"
2741            );
2742        }
2743    }
2744
2745    #[test]
2746    fn ctw_pdf_then_update_matches_plain_update() {
2747        assert_pdf_then_update_matches_plain_update(CtwPredictor::new_ctw(7));
2748    }
2749
2750    #[test]
2751    fn fac_pdf_then_update_matches_plain_update() {
2752        assert_pdf_then_update_matches_plain_update(CtwPredictor::new_fac(7, 8, None));
2753    }
2754
2755    #[test]
2756    fn roundtrip_rate_ac_ctw() {
2757        let data = b"ctw backend roundtrip payload";
2758        let backend = RateBackend::Ctw { depth: 8 };
2759        let enc = compress_rate_bytes(data, &backend, CoderType::AC, FramingMode::Framed).unwrap();
2760        let dec =
2761            decompress_rate_bytes(&enc, &backend, CoderType::AC, FramingMode::Framed).unwrap();
2762        assert_eq!(dec, data);
2763    }
2764
2765    #[test]
2766    fn roundtrip_rate_ac_match_family_and_ppmd() {
2767        let data = b"repeat repeat repeat sparse sparse repeat payload";
2768        for backend in [
2769            RateBackend::Match {
2770                hash_bits: 20,
2771                min_len: 4,
2772                max_len: 255,
2773                base_mix: 0.02,
2774                confidence_scale: 1.0,
2775            },
2776            RateBackend::SparseMatch {
2777                hash_bits: 19,
2778                min_len: 3,
2779                max_len: 64,
2780                gap_min: 1,
2781                gap_max: 2,
2782                base_mix: 0.05,
2783                confidence_scale: 1.0,
2784            },
2785            RateBackend::Ppmd {
2786                order: 8,
2787                memory_mb: 8,
2788            },
2789        ] {
2790            let enc =
2791                compress_rate_bytes(data, &backend, CoderType::AC, FramingMode::Framed).unwrap();
2792            let dec =
2793                decompress_rate_bytes(&enc, &backend, CoderType::AC, FramingMode::Framed).unwrap();
2794            assert_eq!(dec, data);
2795        }
2796    }
2797
2798    #[test]
2799    fn framed_rate_ac_keeps_v1_coder_byte_for_ctw() {
2800        let data = b"legacy framed ac header payload";
2801        let backend = RateBackend::Ctw { depth: 8 };
2802        let enc = compress_rate_bytes(data, &backend, CoderType::AC, FramingMode::Framed).unwrap();
2803        let hdr = FramedHeader::read(&enc).expect("framed header");
2804        assert_eq!(hdr.coder_type(), CoderType::AC);
2805        assert_eq!(hdr.coder, 0);
2806        let dec =
2807            decompress_rate_bytes(&enc, &backend, CoderType::AC, FramingMode::Framed).unwrap();
2808        assert_eq!(dec, data);
2809    }
2810
2811    #[test]
2812    fn framed_rate_rans_keeps_v1_coder_byte() {
2813        let data = b"legacy framed rans header payload";
2814        let backend = RateBackend::Ctw { depth: 8 };
2815        let enc =
2816            compress_rate_bytes(data, &backend, CoderType::RANS, FramingMode::Framed).unwrap();
2817        let hdr = FramedHeader::read(&enc).expect("framed header");
2818        assert_eq!(hdr.coder_type(), CoderType::RANS);
2819        assert_eq!(hdr.coder, 1);
2820        let dec =
2821            decompress_rate_bytes(&enc, &backend, CoderType::RANS, FramingMode::Framed).unwrap();
2822        assert_eq!(dec, data);
2823    }
2824
2825    #[test]
2826    fn framed_rate_ac_keeps_byte_prefix_models_on_legacy_path() {
2827        let data = b"byte prefix adapter exists but byte ac remains default";
2828        let backend = RateBackend::Match {
2829            hash_bits: 20,
2830            min_len: 4,
2831            max_len: 255,
2832            base_mix: 0.02,
2833            confidence_scale: 1.0,
2834        };
2835        let enc = compress_rate_bytes(data, &backend, CoderType::AC, FramingMode::Framed).unwrap();
2836        let hdr = FramedHeader::read(&enc).expect("framed header");
2837        assert_eq!(hdr.coder_type(), CoderType::AC);
2838        assert_eq!(hdr.coder, 0);
2839        let dec =
2840            decompress_rate_bytes(&enc, &backend, CoderType::AC, FramingMode::Framed).unwrap();
2841        assert_eq!(dec, data);
2842    }
2843
2844    #[test]
2845    fn roundtrip_rate_ac_ppmd_high_order_text_payload() {
2846        let seed = include_bytes!("../../../../README.md");
2847        let mut data = Vec::with_capacity(4096);
2848        while data.len() < 4096 {
2849            data.extend_from_slice(seed);
2850        }
2851        data.truncate(4096);
2852
2853        let backend = RateBackend::Ppmd {
2854            order: 12,
2855            memory_mb: 256,
2856        };
2857        let enc = compress_rate_bytes(&data, &backend, CoderType::AC, FramingMode::Framed)
2858            .expect("ppmd high-order compression");
2859        let dec = decompress_rate_bytes(&enc, &backend, CoderType::AC, FramingMode::Framed)
2860            .expect("ppmd high-order decompression");
2861        assert_eq!(dec, data);
2862    }
2863
2864    #[test]
2865    fn roundtrip_rate_ac_calibrated_backend() {
2866        let data = b"calibration wrapper payload calibration wrapper payload";
2867        let backend = RateBackend::Calibrated {
2868            spec: Arc::new(crate::CalibratedSpec::new(
2869                RateBackend::Ctw { depth: 8 },
2870                crate::CalibrationContextKind::Text,
2871            )),
2872        };
2873        let predictor = RatePdfPredictor::from_rate_backend(backend.clone()).unwrap();
2874        assert!(
2875            predictor.can_fast_ac_bitwise(),
2876            "calibrated CTW should expose the SSE bitwise AC path"
2877        );
2878        let enc = compress_rate_bytes(data, &backend, CoderType::AC, FramingMode::Framed).unwrap();
2879        let dec =
2880            decompress_rate_bytes(&enc, &backend, CoderType::AC, FramingMode::Framed).unwrap();
2881        assert_eq!(dec, data);
2882    }
2883
2884    #[test]
2885    fn roundtrip_rate_ac_calibrated_byte_pdf_base_backend() {
2886        let data = b"calibrated byte pdf base payload calibrated byte pdf base payload";
2887        let backend = RateBackend::Calibrated {
2888            spec: Arc::new(crate::CalibratedSpec::new(
2889                RateBackend::Match {
2890                    hash_bits: 18,
2891                    min_len: 3,
2892                    max_len: 64,
2893                    base_mix: 0.08,
2894                    confidence_scale: 1.0,
2895                },
2896                crate::CalibrationContextKind::ByteClass,
2897            )),
2898        };
2899        let predictor = RatePdfPredictor::from_rate_backend(backend.clone()).unwrap();
2900        assert!(
2901            predictor.can_fast_ac_bitwise(),
2902            "calibrated byte-PDF bases should use the generic SSE bitwise adapter"
2903        );
2904        let enc = compress_rate_bytes(data, &backend, CoderType::AC, FramingMode::Framed).unwrap();
2905        let dec =
2906            decompress_rate_bytes(&enc, &backend, CoderType::AC, FramingMode::Framed).unwrap();
2907        assert_eq!(dec, data);
2908    }
2909
2910    #[test]
2911    fn roundtrip_rate_ac_single_expert_ctw_neural_mixture() {
2912        let data = b"single expert neural ctw fast path payload";
2913        let spec = MixtureSpec::new(
2914            MixtureKind::Neural,
2915            vec![crate::MixtureExpertSpec {
2916                name: Some("ctw".to_string()),
2917                log_prior: 0.0,
2918                backend: RateBackend::Ctw { depth: 8 },
2919            }],
2920        )
2921        .with_alpha(0.03);
2922        let backend = RateBackend::Mixture {
2923            spec: Arc::new(spec),
2924        };
2925        let enc = compress_rate_bytes(data, &backend, CoderType::AC, FramingMode::Framed).unwrap();
2926        let dec =
2927            decompress_rate_bytes(&enc, &backend, CoderType::AC, FramingMode::Framed).unwrap();
2928        assert_eq!(dec, data);
2929    }
2930
2931    #[test]
2932    fn roundtrip_rate_ac_single_expert_ctw_bayes_mixture() {
2933        let data = b"single expert bayes ctw fast path payload";
2934        let spec = MixtureSpec::new(
2935            MixtureKind::Bayes,
2936            vec![crate::MixtureExpertSpec {
2937                name: Some("ctw".to_string()),
2938                log_prior: 0.0,
2939                backend: RateBackend::Ctw { depth: 8 },
2940            }],
2941        )
2942        .with_alpha(0.03);
2943        let backend = RateBackend::Mixture {
2944            spec: Arc::new(spec),
2945        };
2946        let enc = compress_rate_bytes(data, &backend, CoderType::AC, FramingMode::Framed).unwrap();
2947        let dec =
2948            decompress_rate_bytes(&enc, &backend, CoderType::AC, FramingMode::Framed).unwrap();
2949        assert_eq!(dec, data);
2950    }
2951
2952    #[test]
2953    fn roundtrip_rate_rans_recursive_mixture() {
2954        let data = b"recursive mixture payload";
2955        let nested = MixtureSpec::new(
2956            MixtureKind::Bayes,
2957            vec![
2958                crate::MixtureExpertSpec {
2959                    name: Some("ctw".to_string()),
2960                    log_prior: 0.0,
2961                    backend: RateBackend::Ctw { depth: 6 },
2962                },
2963                crate::MixtureExpertSpec {
2964                    name: Some("fac".to_string()),
2965                    log_prior: 0.0,
2966                    backend: RateBackend::FacCtw {
2967                        base_depth: 6,
2968                        num_percept_bits: 8,
2969                        encoding_bits: 8,
2970                        msb_first: None,
2971                    },
2972                },
2973            ],
2974        );
2975        let root = MixtureSpec::new(
2976            MixtureKind::Switching,
2977            vec![
2978                crate::MixtureExpertSpec {
2979                    name: Some("nested".to_string()),
2980                    log_prior: 0.0,
2981                    backend: RateBackend::Mixture {
2982                        spec: Arc::new(nested),
2983                    },
2984                },
2985                crate::MixtureExpertSpec {
2986                    name: Some("zpaq".to_string()),
2987                    log_prior: 0.0,
2988                    backend: RateBackend::Zpaq {
2989                        method: crate::api::ZpaqMethodSpec::literal("1"),
2990                    },
2991                },
2992            ],
2993        )
2994        .with_alpha(0.05);
2995
2996        let backend = RateBackend::Mixture {
2997            spec: Arc::new(root),
2998        };
2999        let enc =
3000            compress_rate_bytes(data, &backend, CoderType::RANS, FramingMode::Framed).unwrap();
3001        let dec =
3002            decompress_rate_bytes(&enc, &backend, CoderType::RANS, FramingMode::Framed).unwrap();
3003        assert_eq!(dec, data);
3004    }
3005
3006    #[test]
3007    fn roundtrip_rate_ac_recursive_neural_mixture() {
3008        let data = b"neural recursive mixture payload for ac coder";
3009        let inner = MixtureSpec::new(
3010            MixtureKind::Bayes,
3011            vec![
3012                crate::MixtureExpertSpec {
3013                    name: Some("ctw".to_string()),
3014                    log_prior: 0.0,
3015                    backend: RateBackend::Ctw { depth: 6 },
3016                },
3017                crate::MixtureExpertSpec {
3018                    name: Some("fac".to_string()),
3019                    log_prior: 0.0,
3020                    backend: RateBackend::FacCtw {
3021                        base_depth: 6,
3022                        num_percept_bits: 8,
3023                        encoding_bits: 8,
3024                        msb_first: None,
3025                    },
3026                },
3027            ],
3028        );
3029        let root = MixtureSpec::new(
3030            MixtureKind::Neural,
3031            vec![
3032                crate::MixtureExpertSpec {
3033                    name: Some("nested".to_string()),
3034                    log_prior: 0.0,
3035                    backend: RateBackend::Mixture {
3036                        spec: Arc::new(inner),
3037                    },
3038                },
3039                crate::MixtureExpertSpec {
3040                    name: Some("zpaq".to_string()),
3041                    log_prior: 0.0,
3042                    backend: RateBackend::Zpaq {
3043                        method: crate::api::ZpaqMethodSpec::literal("1"),
3044                    },
3045                },
3046            ],
3047        )
3048        .with_alpha(0.03);
3049
3050        let backend = RateBackend::Mixture {
3051            spec: Arc::new(root),
3052        };
3053        let enc = compress_rate_bytes(data, &backend, CoderType::AC, FramingMode::Framed).unwrap();
3054        let dec =
3055            decompress_rate_bytes(&enc, &backend, CoderType::AC, FramingMode::Framed).unwrap();
3056        assert_eq!(dec, data);
3057    }
3058
3059    #[test]
3060    fn roundtrip_rate_ac_recursive_native_bitwise_mixture() {
3061        let data = b"recursive native bitwise mixture payload";
3062        let backend = recursive_native_bitwise_backend();
3063        let predictor = RatePdfPredictor::from_rate_backend(backend.clone()).unwrap();
3064        assert!(predictor.can_fast_ac_bitwise());
3065
3066        let enc = compress_rate_bytes(data, &backend, CoderType::AC, FramingMode::Framed).unwrap();
3067        let dec =
3068            decompress_rate_bytes(&enc, &backend, CoderType::AC, FramingMode::Framed).unwrap();
3069        assert_eq!(dec, data);
3070    }
3071
3072    fn assert_runtime_and_compression_predictor_align(spec: MixtureSpec, data: &[u8], tol: f64) {
3073        let backend = RateBackend::Mixture {
3074            spec: Arc::new(spec.clone()),
3075        };
3076        let mut predictor = RatePdfPredictor::from_rate_backend(backend).unwrap();
3077        let experts = spec.build_experts();
3078        let mut runtime = crate::mixture::build_mixture_runtime(&spec, &experts).unwrap();
3079
3080        for (t, &symbol) in data.iter().enumerate() {
3081            let pdf = predictor.pdf_next().unwrap();
3082            let p_comp = pdf[symbol as usize];
3083            let p_runtime = runtime.peek_log_prob(symbol).exp();
3084            assert!(
3085                (p_comp - p_runtime).abs() < tol,
3086                "t={t} p_comp={p_comp} p_runtime={p_runtime} symbol={symbol}"
3087            );
3088            predictor.update(symbol).unwrap();
3089            runtime.step(symbol);
3090        }
3091    }
3092
3093    fn alignment_experts() -> Vec<crate::MixtureExpertSpec> {
3094        vec![
3095            crate::MixtureExpertSpec {
3096                name: Some("ctw".to_string()),
3097                log_prior: 0.0,
3098                backend: RateBackend::Ctw { depth: 7 },
3099            },
3100            crate::MixtureExpertSpec {
3101                name: Some("fac".to_string()),
3102                log_prior: -0.7,
3103                backend: RateBackend::FacCtw {
3104                    base_depth: 7,
3105                    num_percept_bits: 8,
3106                    encoding_bits: 8,
3107                    msb_first: None,
3108                },
3109            },
3110        ]
3111    }
3112
3113    #[test]
3114    fn bayes_runtime_and_compression_predictor_align() {
3115        let spec = MixtureSpec::new(MixtureKind::Bayes, alignment_experts());
3116        assert_runtime_and_compression_predictor_align(
3117            spec,
3118            b"bayes predictor alignment check sequence",
3119            1e-8,
3120        );
3121    }
3122
3123    #[test]
3124    fn fading_runtime_and_compression_predictor_align() {
3125        let spec = MixtureSpec::new(MixtureKind::FadingBayes, alignment_experts()).with_decay(0.97);
3126        assert_runtime_and_compression_predictor_align(
3127            spec,
3128            b"fading predictor alignment check sequence",
3129            1e-8,
3130        );
3131    }
3132
3133    #[test]
3134    fn switching_runtime_and_compression_predictor_align() {
3135        let spec = MixtureSpec::new(MixtureKind::Switching, alignment_experts()).with_alpha(0.17);
3136        assert_runtime_and_compression_predictor_align(
3137            spec,
3138            b"switching predictor alignment check sequence",
3139            1e-8,
3140        );
3141    }
3142
3143    #[test]
3144    fn switching_theorem_runtime_and_compression_predictor_align() {
3145        let spec = MixtureSpec::new(MixtureKind::Switching, alignment_experts())
3146            .with_schedule(MixtureScheduleMode::Theorem)
3147            .with_alpha(0.91);
3148        assert_runtime_and_compression_predictor_align(
3149            spec,
3150            b"switching theorem predictor alignment check sequence",
3151            1e-8,
3152        );
3153    }
3154
3155    #[test]
3156    fn convex_runtime_and_compression_predictor_align_for_alpha_above_one() {
3157        let spec = MixtureSpec::new(MixtureKind::Convex, alignment_experts()).with_alpha(1.25);
3158        assert_runtime_and_compression_predictor_align(
3159            spec,
3160            b"convex predictor alignment check sequence",
3161            1e-8,
3162        );
3163    }
3164
3165    #[test]
3166    fn convex_theorem_runtime_and_compression_predictor_align() {
3167        let spec = MixtureSpec::new(MixtureKind::Convex, alignment_experts())
3168            .with_schedule(MixtureScheduleMode::Theorem)
3169            .with_alpha(7.5);
3170        assert_runtime_and_compression_predictor_align(
3171            spec,
3172            b"convex theorem predictor alignment check sequence",
3173            1e-8,
3174        );
3175    }
3176
3177    #[test]
3178    fn neural_runtime_and_compression_predictor_align() {
3179        let spec = MixtureSpec::new(MixtureKind::Neural, alignment_experts()).with_alpha(0.03);
3180        assert_runtime_and_compression_predictor_align(
3181            spec,
3182            b"neural alignment check sequence",
3183            1e-8,
3184        );
3185    }
3186
3187    #[test]
3188    fn mdl_runtime_and_compression_predictor_align() {
3189        let spec = MixtureSpec::new(MixtureKind::Mdl, alignment_experts());
3190        assert_runtime_and_compression_predictor_align(spec, b"mdl alignment check sequence", 1e-8);
3191    }
3192
3193    #[test]
3194    fn nested_runtime_and_compression_predictor_align() {
3195        let nested = MixtureSpec::new(MixtureKind::Bayes, alignment_experts());
3196        let spec = MixtureSpec::new(
3197            MixtureKind::Switching,
3198            vec![
3199                crate::MixtureExpertSpec {
3200                    name: Some("nested".to_string()),
3201                    log_prior: 0.0,
3202                    backend: RateBackend::Mixture {
3203                        spec: Arc::new(nested),
3204                    },
3205                },
3206                crate::MixtureExpertSpec {
3207                    name: Some("ppmd".to_string()),
3208                    log_prior: -0.2,
3209                    backend: RateBackend::Ppmd {
3210                        order: 5,
3211                        memory_mb: 8,
3212                    },
3213                },
3214            ],
3215        )
3216        .with_alpha(0.13);
3217        assert_runtime_and_compression_predictor_align(
3218            spec,
3219            b"nested mixture predictor alignment check sequence",
3220            1e-8,
3221        );
3222    }
3223
3224    fn recursive_native_bitwise_backend() -> RateBackend {
3225        let nested = MixtureSpec::new(MixtureKind::Bayes, alignment_experts());
3226        let root = MixtureSpec::new(
3227            MixtureKind::Switching,
3228            vec![
3229                crate::MixtureExpertSpec {
3230                    name: Some("nested".to_string()),
3231                    log_prior: 0.0,
3232                    backend: RateBackend::Mixture {
3233                        spec: Arc::new(nested),
3234                    },
3235                },
3236                crate::MixtureExpertSpec {
3237                    name: Some("ppmd".to_string()),
3238                    log_prior: -0.2,
3239                    backend: RateBackend::Ppmd {
3240                        order: 5,
3241                        memory_mb: 8,
3242                    },
3243                },
3244            ],
3245        )
3246        .with_alpha(0.13);
3247        RateBackend::Mixture {
3248            spec: Arc::new(root),
3249        }
3250    }
3251
3252    fn assert_bitwise_byte_step_matches_pdf_and_plain_update(
3253        mut predictor: RatePdfPredictor,
3254        data: &[u8],
3255        tol_prob: f64,
3256        tol_pdf: f64,
3257    ) {
3258        for &symbol in data {
3259            let expected_pdf = predictor.pdf_next().unwrap().to_vec();
3260            let expected_prob = expected_pdf[symbol as usize];
3261
3262            let mut stepped = predictor.clone();
3263            let mut product = 1.0f64;
3264            let produced = stepped
3265                .ac_step_bitwise(|bit_idx, p1| {
3266                    let bit = (symbol >> (7 - bit_idx)) & 1;
3267                    let pb = if bit == 1 { p1 } else { 1.0 - p1 };
3268                    product *= pb;
3269                    Ok(bit)
3270                })
3271                .unwrap();
3272            assert_eq!(produced, symbol);
3273            assert!(
3274                (product - expected_prob).abs() <= tol_prob,
3275                "symbol={symbol} product={product} expected_prob={expected_prob}"
3276            );
3277
3278            let mut plain = predictor.clone();
3279            plain.update(symbol).unwrap();
3280            let stepped_pdf = stepped.pdf_next().unwrap().to_vec();
3281            let plain_pdf = plain.pdf_next().unwrap().to_vec();
3282            assert_pdf_close(&stepped_pdf, &plain_pdf, tol_pdf);
3283
3284            predictor.update(symbol).unwrap();
3285        }
3286    }
3287
3288    #[test]
3289    fn bitwise_byte_step_matches_pdf_and_plain_update_for_native_and_recursive_mixtures() {
3290        assert_bitwise_byte_step_matches_pdf_and_plain_update(
3291            RatePdfPredictor::from_rate_backend(RateBackend::Ctw { depth: 7 }).unwrap(),
3292            b"direct ctw bitwise byte step parity",
3293            1e-12,
3294            1e-12,
3295        );
3296
3297        let direct_fac = RateBackend::FacCtw {
3298            base_depth: 7,
3299            num_percept_bits: 8,
3300            encoding_bits: 8,
3301            msb_first: None,
3302        };
3303        let direct_fac_predictor = RatePdfPredictor::from_rate_backend(direct_fac).unwrap();
3304        assert!(
3305            direct_fac_predictor.can_fast_ac_bitwise(),
3306            "byte-wide MSB fac-ctw must expose its native recursive AC path"
3307        );
3308        assert_bitwise_byte_step_matches_pdf_and_plain_update(
3309            direct_fac_predictor,
3310            b"direct fac ctw bitwise byte step parity",
3311            1e-12,
3312            1e-12,
3313        );
3314
3315        let single_expert = RateBackend::Mixture {
3316            spec: Arc::new(MixtureSpec::new(
3317                MixtureKind::Bayes,
3318                vec![crate::MixtureExpertSpec {
3319                    name: Some("ctw".to_string()),
3320                    log_prior: 0.0,
3321                    backend: RateBackend::Ctw { depth: 7 },
3322                }],
3323            )),
3324        };
3325        assert_bitwise_byte_step_matches_pdf_and_plain_update(
3326            RatePdfPredictor::from_rate_backend(single_expert).unwrap(),
3327            b"single expert ctw mixture bitwise byte step parity",
3328            1e-12,
3329            1e-12,
3330        );
3331
3332        let single_fac_neural = RateBackend::Mixture {
3333            spec: Arc::new(MixtureSpec::new(
3334                MixtureKind::Neural,
3335                vec![crate::MixtureExpertSpec {
3336                    name: Some("fac-ctw".to_string()),
3337                    log_prior: 0.0,
3338                    backend: RateBackend::FacCtw {
3339                        base_depth: 7,
3340                        num_percept_bits: 8,
3341                        encoding_bits: 8,
3342                        msb_first: None,
3343                    },
3344                }],
3345            )),
3346        };
3347        let single_fac_neural_predictor =
3348            RatePdfPredictor::from_rate_backend(single_fac_neural).unwrap();
3349        assert!(
3350            single_fac_neural_predictor.can_fast_ac_bitwise(),
3351            "neural mixtures containing only byte-wide MSB fac-ctw must not fall back to PDF-prefix AC"
3352        );
3353        assert_bitwise_byte_step_matches_pdf_and_plain_update(
3354            single_fac_neural_predictor,
3355            b"single expert fac neural mixture bitwise byte step parity",
3356            1e-12,
3357            1e-12,
3358        );
3359
3360        let mixed_direct = RateBackend::Mixture {
3361            spec: Arc::new(MixtureSpec::new(
3362                MixtureKind::Bayes,
3363                vec![
3364                    crate::MixtureExpertSpec {
3365                        name: Some("ctw".to_string()),
3366                        log_prior: 0.0,
3367                        backend: RateBackend::Ctw { depth: 7 },
3368                    },
3369                    crate::MixtureExpertSpec {
3370                        name: Some("match".to_string()),
3371                        log_prior: -0.3,
3372                        backend: RateBackend::Match {
3373                            hash_bits: 20,
3374                            min_len: 4,
3375                            max_len: 255,
3376                            base_mix: 0.02,
3377                            confidence_scale: 1.0,
3378                        },
3379                    },
3380                ],
3381            )),
3382        };
3383        assert_bitwise_byte_step_matches_pdf_and_plain_update(
3384            RatePdfPredictor::from_rate_backend(mixed_direct).unwrap(),
3385            b"mixed direct mixture bitwise byte step parity",
3386            1e-11,
3387            1e-11,
3388        );
3389
3390        let recursive = recursive_native_bitwise_backend();
3391        let predictor = RatePdfPredictor::from_rate_backend(recursive).unwrap();
3392        assert!(predictor.can_fast_ac_bitwise());
3393        assert_bitwise_byte_step_matches_pdf_and_plain_update(
3394            predictor,
3395            b"recursive nested mixture bitwise byte step parity",
3396            1e-10,
3397            1e-10,
3398        );
3399    }
3400
3401    fn assert_cached_cdf_fast_bitwise_matches_pdf_rows(mut predictor: RatePdfPredictor) {
3402        let data = b"cached cdf parity check payload";
3403        for &symbol in data {
3404            let pdf = predictor.pdf_next().unwrap().to_vec();
3405            assert!(predictor.prepare_cached_cdf_fast_bitwise().unwrap());
3406
3407            let mut row = zeroed_prefix_cdf();
3408            fill_prefix_cdf_from_pdf(&mut row, &pdf, PDF_MIN);
3409
3410            let mut stack = vec![MsbPrefixRange::FULL];
3411            while let Some(range) = stack.pop() {
3412                if range.hi() - range.lo() <= 1 {
3413                    continue;
3414                }
3415                let expected = range.prob_one(&row, PDF_MIN);
3416                let got = predictor
3417                    .cached_cdf_bit_prob_one_msb(range)
3418                    .expect("cached cdf branch probability");
3419                let diff = (expected - got).abs();
3420                assert!(
3421                    diff <= 1e-12,
3422                    "lo={} hi={} expected={expected} got={got} diff={diff}",
3423                    range.lo(),
3424                    range.hi()
3425                );
3426                stack.push(range.observed(false));
3427                stack.push(range.observed(true));
3428            }
3429
3430            predictor.update(symbol).unwrap();
3431        }
3432    }
3433
3434    #[test]
3435    fn cached_cdf_fast_bitwise_matches_pdf_rows_for_specialized_predictors() {
3436        assert_cached_cdf_fast_bitwise_matches_pdf_rows(
3437            RatePdfPredictor::from_rate_backend(RateBackend::RosaPlus { max_order: -1 }).unwrap(),
3438        );
3439        assert_cached_cdf_fast_bitwise_matches_pdf_rows(
3440            RatePdfPredictor::from_rate_backend(RateBackend::Ppmd {
3441                order: 6,
3442                memory_mb: 8,
3443            })
3444            .unwrap(),
3445        );
3446        assert_cached_cdf_fast_bitwise_matches_pdf_rows(
3447            RatePdfPredictor::from_rate_backend(RateBackend::Match {
3448                hash_bits: 20,
3449                min_len: 4,
3450                max_len: 255,
3451                base_mix: 0.02,
3452                confidence_scale: 1.0,
3453            })
3454            .unwrap(),
3455        );
3456        #[cfg(feature = "backend-rwkv")]
3457        assert_cached_cdf_fast_bitwise_matches_pdf_rows(
3458            RatePdfPredictor::from_rate_backend(RateBackend::Rwkv7Method {
3459                method: crate::rwkvzip::parse_method_spec("cfg:hidden=64,layers=1,intermediate=64,decay_rank=8,a_rank=8,v_rank=8,g_rank=8,seed=11,train=none,lr=0.0,stride=1;policy:schedule=0..100:infer").expect("rwkv method spec"),
3460            })
3461            .unwrap(),
3462        );
3463        #[cfg(feature = "backend-mamba")]
3464        assert_cached_cdf_fast_bitwise_matches_pdf_rows(
3465            RatePdfPredictor::from_rate_backend(RateBackend::MambaMethod {
3466                method: crate::mambazip::parse_method_spec("cfg:hidden=64,layers=1,intermediate=64,state=8,conv=3,dt_rank=4,seed=7,train=none,lr=0.0,stride=1;policy:schedule=0..100:infer").expect("mamba method spec"),
3467            })
3468            .unwrap(),
3469        );
3470    }
3471
3472    #[test]
3473    fn raw_size_not_larger_than_framed_size() {
3474        let data = b"raw/framed size check payload";
3475        let backend = RateBackend::RosaPlus { max_order: 8 };
3476        let raw = compress_rate_size(data, &backend, CoderType::AC, FramingMode::Raw).unwrap();
3477        let framed =
3478            compress_rate_size(data, &backend, CoderType::AC, FramingMode::Framed).unwrap();
3479        assert!(framed >= raw);
3480    }
3481
3482    #[cfg(feature = "backend-rwkv")]
3483    #[test]
3484    fn roundtrip_rate_rwkv_method_cfg() {
3485        let data = b"rwkv cfg method backend";
3486        let backend = RateBackend::Rwkv7Method {
3487            method: crate::rwkvzip::parse_method_spec("cfg:hidden=64,layers=1,intermediate=64,decay_rank=8,a_rank=8,v_rank=8,g_rank=8,seed=11,train=none,lr=0.0,stride=1;policy:schedule=0..100:infer").expect("rwkv method spec"),
3488        };
3489        let enc = compress_rate_bytes(data, &backend, CoderType::AC, FramingMode::Framed).unwrap();
3490        let dec =
3491            decompress_rate_bytes(&enc, &backend, CoderType::AC, FramingMode::Framed).unwrap();
3492        assert_eq!(dec, data);
3493    }
3494
3495    #[cfg(feature = "backend-rwkv")]
3496    #[test]
3497    fn rwkv_rate_predictor_preserves_backend_pdf_exactly() {
3498        let method = "cfg:hidden=64,layers=1,intermediate=64,decay_rank=8,a_rank=8,v_rank=8,g_rank=8,seed=11,train=none,lr=0.0,stride=1;policy:schedule=0..100:infer";
3499        let mut predictor = RwkvPredictor::from_method(method).expect("rwkv predictor");
3500        let mut backend = rwkvzip::Compressor::new_from_method(method).expect("rwkv backend");
3501        let mut direct = vec![0.0; backend.vocab_size()];
3502
3503        let predicted = predictor.pdf_next().to_vec();
3504        backend.forward_to_pdf(0, &mut direct);
3505        assert_pdf_close(&predicted, &direct, 1e-18);
3506
3507        predictor.update(b'x').expect("predictor update");
3508        backend
3509            .online_update_from_pdf(b'x', &direct)
3510            .expect("backend update");
3511        backend.forward_to_pdf(u32::from(b'x'), &mut direct);
3512        assert_pdf_close(predictor.pdf_next(), &direct, 1e-18);
3513    }
3514
3515    #[cfg(feature = "backend-rwkv")]
3516    #[test]
3517    fn compiled_rwkv_rate_pdf_predictor_preserves_backend_pdf_exactly() {
3518        let method = "cfg:hidden=64,layers=1,intermediate=64,decay_rank=8,a_rank=8,v_rank=8,g_rank=8,seed=11,train=none,lr=0.0,stride=1;policy:schedule=0..100:infer";
3519        let backend = RateBackend::Rwkv7Method {
3520            method: crate::rwkvzip::parse_method_spec(method).expect("rwkv method spec"),
3521        }
3522        .compile()
3523        .expect("compiled rwkv backend");
3524        let spec = rwkvzip::parse_method_spec(method).expect("parsed rwkv spec");
3525        let mut predictor =
3526            RatePdfPredictor::from_compiled(&backend).expect("compiled rwkv predictor");
3527        let mut direct =
3528            rwkvzip::Compressor::new_from_method_spec(&spec).expect("rwkv backend from spec");
3529        let mut pdf = vec![0.0; direct.vocab_size()];
3530
3531        let predicted = predictor.pdf_next().expect("predictor pdf").to_vec();
3532        direct.forward_to_pdf(0, &mut pdf);
3533        assert_pdf_close(&predicted, &pdf, 1e-18);
3534
3535        predictor.update(b'x').expect("predictor update");
3536        direct
3537            .online_update_from_pdf(b'x', &pdf)
3538            .expect("backend update");
3539        direct.forward_to_pdf(u32::from(b'x'), &mut pdf);
3540        assert_pdf_close(predictor.pdf_next().expect("predictor pdf"), &pdf, 1e-18);
3541    }
3542
3543    #[cfg(feature = "backend-rwkv")]
3544    #[test]
3545    fn rwkv_rate_predictor_matches_backend_after_partial_tbptt_stream() {
3546        let method = "cfg:hidden=64,layers=1,intermediate=64,decay_rank=8,a_rank=8,v_rank=8,g_rank=8,seed=29,train=adam,lr=0.0008,stride=1;policy:schedule=0..100:train(scope=all,opt=adam,lr=0.0008,stride=1,bptt=8,clip=0,momentum=0.9)";
3547        let data = b"abcdefghij";
3548        let mut predictor = RwkvPredictor::from_method(method).expect("rwkv predictor");
3549        let mut backend = rwkvzip::Compressor::new_from_method(method).expect("rwkv backend");
3550        let mut direct = vec![0.0; backend.vocab_size()];
3551
3552        predictor
3553            .begin_stream(data.len())
3554            .expect("begin predictor stream");
3555        backend
3556            .begin_online_policy_stream(Some(data.len() as u64))
3557            .expect("begin backend stream");
3558        backend.reset_and_prime();
3559
3560        for &byte in data {
3561            let predicted = predictor.pdf_next().to_vec();
3562            backend.copy_current_pdf_to(&mut direct);
3563            assert_pdf_close(&predicted, &direct, 1e-18);
3564
3565            predictor.update(byte).expect("predictor update");
3566            backend
3567                .observe_symbol_from_current_pdf(byte)
3568                .expect("backend update");
3569        }
3570
3571        predictor.finish_stream().expect("finish predictor stream");
3572        backend
3573            .finish_online_policy_stream()
3574            .expect("finish backend stream");
3575        backend.copy_current_pdf_to(&mut direct);
3576        assert_pdf_close(predictor.pdf_next(), &direct, 1e-18);
3577    }
3578
3579    #[cfg(feature = "backend-rwkv")]
3580    #[test]
3581    fn roundtrip_rate_rwkv_two_json_method_2m() {
3582        let two_json: serde_json::Value =
3583            serde_json::from_str(include_str!("../../../../configs/bench/two.json")).unwrap();
3584        let experts = two_json["experts"]
3585            .as_array()
3586            .expect("two.json must define experts array");
3587        let method = experts
3588            .iter()
3589            .find(|expert| expert["kind"].as_str() == Some("rwkv7"))
3590            .and_then(|expert| expert["method"].as_str())
3591            .expect("two.json must include rwkv7 expert with string method")
3592            .to_string();
3593
3594        let backend = RateBackend::Rwkv7Method {
3595            method: crate::rwkvzip::parse_method_spec(&method).expect("rwkv method spec"),
3596        };
3597        let seed = include_bytes!("../../../../README.md");
3598        let target_len = 2_097_152usize;
3599        let mut data = Vec::with_capacity(target_len);
3600        while data.len() < target_len {
3601            let remaining = target_len - data.len();
3602            data.extend_from_slice(&seed[..seed.len().min(remaining)]);
3603        }
3604
3605        let enc = compress_rate_bytes(&data, &backend, CoderType::AC, FramingMode::Framed).unwrap();
3606        let dec =
3607            decompress_rate_bytes(&enc, &backend, CoderType::AC, FramingMode::Framed).unwrap();
3608        assert_eq!(dec, data);
3609    }
3610
3611    #[test]
3612    fn benchmark_two_json_matches_examples_and_historical_alpha() {
3613        let canonical: serde_json::Value =
3614            serde_json::from_str(include_str!("../../../../configs/bench/two.json")).unwrap();
3615        let example: serde_json::Value =
3616            serde_json::from_str(include_str!("../../../../examples/two.json")).unwrap();
3617
3618        assert_eq!(canonical, example, "benchmark specs drifted");
3619        assert_eq!(canonical["kind"].as_str(), Some("neural"));
3620        let alpha = canonical["alpha"].as_f64().expect("neural alpha");
3621        assert!(
3622            (alpha - 0.03).abs() <= 1e-12,
3623            "expected historical neural alpha 0.03, got {alpha}"
3624        );
3625    }
3626
3627    #[cfg(feature = "backend-mamba")]
3628    #[test]
3629    fn mamba_rate_predictor_preserves_backend_pdf_exactly() {
3630        let method = "cfg:hidden=64,layers=1,intermediate=64,state=8,conv=3,dt_rank=4,seed=7,train=none,lr=0.0,stride=1;policy:schedule=0..100:infer";
3631        let mut predictor = MambaPredictor::from_method(method).expect("mamba predictor");
3632        let mut backend = mambazip::Compressor::new_from_method(method).expect("mamba backend");
3633        let mut direct = vec![0.0; backend.vocab_size()];
3634
3635        let predicted = predictor.pdf_next().to_vec();
3636        backend.forward_to_pdf(0, &mut direct);
3637        assert_pdf_close(&predicted, &direct, 1e-18);
3638
3639        predictor.update(b'x').expect("predictor update");
3640        backend
3641            .online_update_from_pdf(b'x', &direct)
3642            .expect("backend update");
3643        backend.forward_to_pdf(u32::from(b'x'), &mut direct);
3644        assert_pdf_close(predictor.pdf_next(), &direct, 1e-18);
3645    }
3646
3647    #[cfg(feature = "backend-mamba")]
3648    #[test]
3649    fn compiled_mamba_rate_pdf_predictor_preserves_backend_pdf_exactly() {
3650        let method = "cfg:hidden=64,layers=1,intermediate=64,state=8,conv=3,dt_rank=4,seed=7,train=none,lr=0.0,stride=1;policy:schedule=0..100:infer";
3651        let backend = RateBackend::MambaMethod {
3652            method: crate::mambazip::parse_method_spec(method).expect("mamba method spec"),
3653        }
3654        .compile()
3655        .expect("compiled mamba backend");
3656        let spec = mambazip::parse_method_spec(method).expect("parsed mamba spec");
3657        let mut predictor =
3658            RatePdfPredictor::from_compiled(&backend).expect("compiled mamba predictor");
3659        let mut direct =
3660            mambazip::Compressor::new_from_method_spec(&spec).expect("mamba backend from spec");
3661        let mut pdf = vec![0.0; direct.vocab_size()];
3662
3663        let predicted = predictor.pdf_next().expect("predictor pdf").to_vec();
3664        direct.forward_to_pdf(0, &mut pdf);
3665        assert_pdf_close(&predicted, &pdf, 1e-18);
3666
3667        predictor.update(b'x').expect("predictor update");
3668        direct
3669            .online_update_from_pdf(b'x', &pdf)
3670            .expect("backend update");
3671        direct.forward_to_pdf(u32::from(b'x'), &mut pdf);
3672        assert_pdf_close(predictor.pdf_next().expect("predictor pdf"), &pdf, 1e-18);
3673    }
3674
3675    #[test]
3676    fn roundtrip_rate_ac_particle() {
3677        let spec = crate::ParticleSpec {
3678            num_particles: 4,
3679            num_cells: 4,
3680            cell_dim: 8,
3681            num_rules: 2,
3682            selector_hidden: 16,
3683            rule_hidden: 16,
3684            context_window: 8,
3685            unroll_steps: 1,
3686            ..crate::ParticleSpec::default()
3687        };
3688        let data = b"particle ac roundtrip payload";
3689        let backend = RateBackend::Particle {
3690            spec: Arc::new(spec),
3691        };
3692        let enc = compress_rate_bytes(data, &backend, CoderType::AC, FramingMode::Framed).unwrap();
3693        let dec =
3694            decompress_rate_bytes(&enc, &backend, CoderType::AC, FramingMode::Framed).unwrap();
3695        assert_eq!(dec, data);
3696    }
3697
3698    #[test]
3699    fn roundtrip_rate_rans_particle() {
3700        let spec = crate::ParticleSpec {
3701            num_particles: 4,
3702            num_cells: 4,
3703            cell_dim: 8,
3704            num_rules: 2,
3705            selector_hidden: 16,
3706            rule_hidden: 16,
3707            context_window: 8,
3708            unroll_steps: 1,
3709            ..crate::ParticleSpec::default()
3710        };
3711        let data = b"particle rans roundtrip payload";
3712        let backend = RateBackend::Particle {
3713            spec: Arc::new(spec),
3714        };
3715        let enc =
3716            compress_rate_bytes(data, &backend, CoderType::RANS, FramingMode::Framed).unwrap();
3717        let dec =
3718            decompress_rate_bytes(&enc, &backend, CoderType::RANS, FramingMode::Framed).unwrap();
3719        assert_eq!(dec, data);
3720    }
3721
3722    #[test]
3723    fn mixture_with_particle_expert_roundtrip() {
3724        let particle_spec = crate::ParticleSpec {
3725            num_particles: 4,
3726            num_cells: 4,
3727            cell_dim: 8,
3728            num_rules: 2,
3729            selector_hidden: 16,
3730            rule_hidden: 16,
3731            context_window: 8,
3732            unroll_steps: 1,
3733            ..crate::ParticleSpec::default()
3734        };
3735        let spec = MixtureSpec::new(
3736            MixtureKind::Bayes,
3737            vec![
3738                crate::MixtureExpertSpec {
3739                    name: Some("particle".to_string()),
3740                    log_prior: 0.0,
3741                    backend: RateBackend::Particle {
3742                        spec: Arc::new(particle_spec),
3743                    },
3744                },
3745                crate::MixtureExpertSpec {
3746                    name: Some("ctw".to_string()),
3747                    log_prior: 0.0,
3748                    backend: RateBackend::Ctw { depth: 6 },
3749                },
3750            ],
3751        );
3752        let backend = RateBackend::Mixture {
3753            spec: Arc::new(spec),
3754        };
3755        let data = b"mixture with particle expert roundtrip";
3756        let enc = compress_rate_bytes(data, &backend, CoderType::AC, FramingMode::Framed).unwrap();
3757        let dec =
3758            decompress_rate_bytes(&enc, &backend, CoderType::AC, FramingMode::Framed).unwrap();
3759        assert_eq!(dec, data);
3760    }
3761}