1#![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; const 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)]
61pub enum FramingMode {
63 Raw,
65 #[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 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 #[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 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 #[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#[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#[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 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
2228pub 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
2255pub 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
2266pub 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
2281pub 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}