#![allow(clippy::module_name_repetitions)]
use std::fmt;
pub const DEFAULT_ALPHA: f64 = 0.01;
pub const DEFAULT_P0: f64 = 0.97;
pub const DEFAULT_LAMBDA: f64 = 0.5;
pub const DEFAULT_AUDIT_CADENCE: u32 = 8;
pub const DEFAULT_MARGIN_THRESHOLD: f64 = 4.0;
pub const MIN_CALIBRATION_SAMPLES: usize = 200;
const E_FLOOR: f64 = f64::MIN_POSITIVE;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Action {
AcceptDraft,
FallBack,
}
impl Action {
#[must_use]
pub fn is_verified(self) -> bool {
matches!(self, Action::FallBack)
}
}
impl fmt::Display for Action {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Action::AcceptDraft => f.write_str("accept_draft"),
Action::FallBack => f.write_str("fall_back"),
}
}
}
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct LossMatrix {
pub c_save: f64,
pub l_flip: f64,
pub c_full: f64,
}
impl Default for LossMatrix {
fn default() -> Self {
Self {
c_save: 1.0,
l_flip: 1.0e6,
c_full: 2.0,
}
}
}
impl LossMatrix {
#[must_use]
pub fn expected_accept_loss(&self, p_flip: f64) -> f64 {
(1.0 - p_flip) * (-self.c_save) + p_flip * self.l_flip
}
#[must_use]
pub fn fallback_loss(&self) -> f64 {
self.c_full
}
#[must_use]
pub fn break_even_flip_rate(&self) -> f64 {
let denom = self.c_save + self.l_flip;
if denom <= 0.0 {
return 0.0;
}
((self.c_full + self.c_save) / denom).clamp(0.0, 1.0)
}
}
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct Calibration {
alpha: f64,
p0: f64,
lambda: f64,
audit_cadence: u32,
margin_threshold: f64,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum CalibrationError {
AlphaOutOfRange,
P0OutOfRange,
LambdaOutOfRange,
AuditCadenceZero,
}
impl fmt::Display for CalibrationError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
CalibrationError::AlphaOutOfRange => {
f.write_str("alpha out of (0,1] (alpha-unset => deterministic fallback)")
}
CalibrationError::P0OutOfRange => f.write_str("p0 out of (0,1)"),
CalibrationError::LambdaOutOfRange => {
f.write_str("lambda out of [0, 1/q0) (would break e_t >= 0)")
}
CalibrationError::AuditCadenceZero => f.write_str("audit_cadence must be >= 1"),
}
}
}
impl std::error::Error for CalibrationError {}
impl Calibration {
pub fn new(
alpha: f64,
p0: f64,
lambda: f64,
audit_cadence: u32,
margin_threshold: f64,
) -> Result<Self, CalibrationError> {
if !(alpha.is_finite() && alpha > 0.0 && alpha <= 1.0) {
return Err(CalibrationError::AlphaOutOfRange);
}
if !(p0.is_finite() && p0 > 0.0 && p0 < 1.0) {
return Err(CalibrationError::P0OutOfRange);
}
let q0 = 1.0 - p0;
let lambda_cap = 1.0 / q0;
if !(lambda.is_finite() && lambda >= 0.0 && lambda < lambda_cap) {
return Err(CalibrationError::LambdaOutOfRange);
}
if audit_cadence == 0 {
return Err(CalibrationError::AuditCadenceZero);
}
if !margin_threshold.is_finite() || margin_threshold < 0.0 {
return Ok(Self {
alpha,
p0,
lambda,
audit_cadence,
margin_threshold: f64::INFINITY,
});
}
Ok(Self {
alpha,
p0,
lambda,
audit_cadence,
margin_threshold,
})
}
#[must_use]
pub fn worked_example() -> Self {
Self::new(
DEFAULT_ALPHA,
DEFAULT_P0,
DEFAULT_LAMBDA,
DEFAULT_AUDIT_CADENCE,
DEFAULT_MARGIN_THRESHOLD,
)
.expect("worked-example calibration constants are in-range")
}
#[must_use]
pub fn alpha(&self) -> f64 {
self.alpha
}
#[must_use]
pub fn q0(&self) -> f64 {
1.0 - self.p0
}
#[must_use]
pub fn threshold(&self) -> f64 {
1.0 / self.alpha
}
}
#[derive(Debug, Clone, PartialEq)]
pub struct EProcess {
q0: f64,
lambda: f64,
alpha: f64,
e_value: f64,
obs_count: u64,
rejected_at: Option<u64>,
}
impl EProcess {
#[must_use]
pub fn new(cal: &Calibration) -> Self {
Self {
q0: cal.q0(),
lambda: cal.lambda,
alpha: cal.alpha,
e_value: 1.0,
obs_count: 0,
rejected_at: None,
}
}
pub fn observe(&mut self, disagree: bool) -> bool {
self.obs_count += 1;
let d_t = if disagree { 1.0 } else { 0.0 };
let e_t = 1.0 + self.lambda * (d_t - self.q0);
self.e_value *= e_t;
self.e_value = self.e_value.clamp(E_FLOOR, f64::MAX / 2.0);
if self.e_value >= self.threshold() && self.rejected_at.is_none() {
self.rejected_at = Some(self.obs_count);
return true;
}
false
}
#[must_use]
pub fn e_value(&self) -> f64 {
self.e_value
}
#[must_use]
pub fn threshold(&self) -> f64 {
1.0 / self.alpha
}
#[must_use]
pub fn obs_count(&self) -> u64 {
self.obs_count
}
#[must_use]
pub fn rejected(&self) -> bool {
self.rejected_at.is_some()
}
#[must_use]
pub fn rejected_at(&self) -> Option<u64> {
self.rejected_at
}
}
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct DecisionRecord {
pub step: u64,
pub action: Action,
pub margin: f64,
pub audited: bool,
pub agree: Option<bool>,
pub e_value: f64,
pub emitted_token: u32,
}
#[derive(Debug, Clone, Default, PartialEq)]
pub struct SpeculativeLedger {
records: Vec<DecisionRecord>,
accepted: u64,
audited: u64,
disagreements: u64,
tripped: bool,
fallback_active: bool,
}
impl SpeculativeLedger {
#[must_use]
pub fn fallback() -> Self {
Self {
fallback_active: true,
..Self::default()
}
}
fn push(&mut self, rec: DecisionRecord) {
match rec.action {
Action::AcceptDraft => self.accepted += 1,
Action::FallBack => {}
}
if rec.audited {
self.audited += 1;
if rec.agree == Some(false) {
self.disagreements += 1;
}
}
self.records.push(rec);
}
#[must_use]
pub fn records(&self) -> &[DecisionRecord] {
&self.records
}
#[must_use]
pub fn accepted(&self) -> u64 {
self.accepted
}
#[must_use]
pub fn audited(&self) -> u64 {
self.audited
}
#[must_use]
pub fn disagreements(&self) -> u64 {
self.disagreements
}
#[must_use]
pub fn steps(&self) -> u64 {
self.records.len() as u64
}
#[must_use]
pub fn acceptance_rate(&self) -> f64 {
if self.records.is_empty() {
return 0.0;
}
self.accepted as f64 / self.records.len() as f64
}
#[must_use]
pub fn measured_disagreement_rate(&self) -> f64 {
if self.audited == 0 {
return 0.0;
}
self.disagreements as f64 / self.audited as f64
}
#[must_use]
pub fn tripped(&self) -> bool {
self.tripped
}
#[must_use]
pub fn fallback_active(&self) -> bool {
self.fallback_active
}
}
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct GuardLevel {
pub streams: u64,
pub breaches: u64,
pub alpha: f64,
}
impl GuardLevel {
#[must_use]
pub fn breach_fraction(&self) -> f64 {
if self.streams == 0 {
return 0.0;
}
self.breaches as f64 / self.streams as f64
}
#[must_use]
pub fn coverage(&self) -> f64 {
1.0 - self.breach_fraction()
}
#[must_use]
pub fn holds_level(&self) -> bool {
self.breach_fraction() <= self.alpha + 1e-12
}
}
pub struct SpeculativeGuard {
eprocess: Option<EProcess>,
cal: Option<Calibration>,
since_audit: u32,
tripped: bool,
ledger: SpeculativeLedger,
}
impl SpeculativeGuard {
#[must_use]
pub fn enabled(cal: Calibration) -> Self {
Self {
eprocess: Some(EProcess::new(&cal)),
cal: Some(cal),
since_audit: 0,
tripped: false,
ledger: SpeculativeLedger::default(),
}
}
#[must_use]
pub fn disabled() -> Self {
Self {
eprocess: None,
cal: None,
since_audit: 0,
tripped: false,
ledger: SpeculativeLedger::fallback(),
}
}
#[must_use]
pub fn from_calibration(
cal: Option<Result<Calibration, CalibrationError>>,
calibration_samples: usize,
) -> Self {
match cal {
Some(Ok(c)) if calibration_samples >= MIN_CALIBRATION_SAMPLES => Self::enabled(c),
_ => Self::disabled(),
}
}
#[must_use]
pub fn speculation_active(&self) -> bool {
self.eprocess.is_some() && !self.tripped
}
#[must_use]
pub fn tripped(&self) -> bool {
self.tripped
}
#[must_use]
pub fn should_verify(&self, margin: f64) -> bool {
let Some(cal) = self.cal.as_ref() else {
return true; };
if self.tripped {
return true; }
if !margin.is_finite() || margin < cal.margin_threshold {
return true;
}
self.since_audit + 1 >= cal.audit_cadence
}
pub fn decide(&mut self, margin: f64, full_token: u32, draft_token: u32) -> Decision {
let step = self.ledger.steps() + 1;
if self.should_verify(margin) {
self.since_audit = 0;
let disagree = draft_token != full_token;
let (audited, e_value, agree) = if let Some(ep) = self.eprocess.as_mut() {
if !self.tripped {
let crossed = ep.observe(disagree);
if crossed {
self.tripped = true;
self.ledger.tripped = true;
}
(true, ep.e_value(), Some(!disagree))
} else {
(false, ep.e_value(), None)
}
} else {
(false, 1.0, None)
};
self.ledger.push(DecisionRecord {
step,
action: Action::FallBack,
margin,
audited,
agree,
e_value,
emitted_token: full_token,
});
Decision {
action: Action::FallBack,
emitted_token: full_token,
}
} else {
self.since_audit += 1;
let e_value = self.eprocess.as_ref().map_or(1.0, EProcess::e_value);
let draft_agrees_with_full = draft_token.eq(&full_token);
if draft_agrees_with_full {
self.ledger.push(DecisionRecord {
step,
action: Action::AcceptDraft,
margin,
audited: false,
agree: None,
e_value,
emitted_token: draft_token,
});
Decision {
action: Action::AcceptDraft,
emitted_token: draft_token,
}
} else {
self.ledger.push(DecisionRecord {
step,
action: Action::FallBack,
margin,
audited: false,
agree: None,
e_value,
emitted_token: full_token,
});
Decision {
action: Action::FallBack,
emitted_token: full_token,
}
}
}
}
#[must_use]
pub fn ledger(&self) -> &SpeculativeLedger {
&self.ledger
}
#[must_use]
pub fn e_value(&self) -> Option<f64> {
self.eprocess.as_ref().map(EProcess::e_value)
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct Decision {
pub action: Action,
pub emitted_token: u32,
}
#[must_use]
pub fn argmax(logits: &[f32]) -> u32 {
let mut best = 0usize;
let mut best_v = f32::NEG_INFINITY;
for (i, &v) in logits.iter().enumerate() {
if v > best_v {
best_v = v;
best = i;
}
}
best as u32
}
#[must_use]
pub fn top1_top2_margin(logits: &[f32]) -> f64 {
let mut first = f32::NEG_INFINITY;
let mut second = f32::NEG_INFINITY;
for &v in logits {
if v > first {
second = first;
first = v;
} else if v > second {
second = v;
}
}
f64::from(first - second)
}
#[cfg(test)]
mod tests {
use super::*;
struct SplitMix64(u64);
impl SplitMix64 {
fn next_u64(&mut self) -> u64 {
self.0 = self.0.wrapping_add(0x9E37_79B9_7F4A_7C15);
let mut z = self.0;
z = (z ^ (z >> 30)).wrapping_mul(0xBF58_476D_1CE4_E5B9);
z = (z ^ (z >> 27)).wrapping_mul(0x94D0_49BB_1331_11EB);
z ^ (z >> 31)
}
fn unit(&mut self) -> f64 {
(self.next_u64() >> 11) as f64 / (1u64 << 53) as f64
}
}
fn approx(a: f64, b: f64, tol: f64) -> bool {
(a - b).abs() <= tol
}
#[test]
fn calibration_rejects_out_of_range_params() {
assert_eq!(
Calibration::new(0.0, 0.97, 0.5, 8, 4.0),
Err(CalibrationError::AlphaOutOfRange)
);
assert_eq!(
Calibration::new(1.5, 0.97, 0.5, 8, 4.0),
Err(CalibrationError::AlphaOutOfRange)
);
assert_eq!(
Calibration::new(0.01, 1.0, 0.5, 8, 4.0),
Err(CalibrationError::P0OutOfRange)
);
let q0_hi = 1.0 - 0.97_f64;
let cap_hi = 1.0 / q0_hi;
assert!(Calibration::new(0.01, 0.97, cap_hi, 8, 4.0).is_err());
assert!(Calibration::new(0.01, 0.97, cap_hi - 1e-6, 8, 4.0).is_ok());
assert_eq!(
Calibration::new(0.01, 0.4, 2.0, 8, 4.0),
Err(CalibrationError::LambdaOutOfRange)
);
let q0_lo = 1.0 - 0.4_f64;
assert!(Calibration::new(0.01, 0.4, 1.0 / q0_lo - 1e-6, 8, 4.0).is_ok());
assert_eq!(
Calibration::new(0.01, 0.97, 0.5, 0, 4.0),
Err(CalibrationError::AuditCadenceZero)
);
}
#[test]
fn worked_example_matches_doc_2_3() {
let cal = Calibration::worked_example();
assert!(approx(cal.threshold(), 100.0, 1e-9));
assert!(approx(cal.q0(), 0.03, 1e-12));
let mut ep = EProcess::new(&cal);
ep.observe(false);
assert!(approx(ep.e_value(), 0.985, 1e-12), "agree e_t");
let mut ep2 = EProcess::new(&cal);
ep2.observe(true);
assert!(approx(ep2.e_value(), 1.485, 1e-12), "disagree e_t");
}
#[test]
fn e_value_property_fair_at_boundary() {
let cal = Calibration::worked_example();
let q0 = cal.q0();
let lambda = DEFAULT_LAMBDA;
let e_agree = 1.0 + lambda * (0.0 - q0);
let e_dis = 1.0 + lambda * (1.0 - q0);
let expected = (1.0 - q0) * e_agree + q0 * e_dis;
assert!(approx(expected, 1.0, 1e-12), "E[e_t] = {expected}");
}
#[test]
fn agreements_shrink_disagreements_grow() {
let cal = Calibration::worked_example();
let mut ep = EProcess::new(&cal);
let start = ep.e_value();
ep.observe(false);
assert!(ep.e_value() < start, "agreement shrinks wealth");
let mut ep2 = EProcess::new(&cal);
ep2.observe(true);
assert!(ep2.e_value() > 1.0, "disagreement grows wealth");
}
#[test]
fn healthy_stream_never_rejects() {
let cal = Calibration::worked_example();
let mut ep = EProcess::new(&cal);
for _ in 0..10_000 {
ep.observe(false);
}
assert!(!ep.rejected());
assert!(ep.e_value() < ep.threshold());
}
#[test]
fn disagreement_burst_trips_guard() {
let cal = Calibration::worked_example();
let mut ep = EProcess::new(&cal);
let mut fired_at = None;
for i in 1..=20u64 {
if ep.observe(true) {
fired_at = Some(i);
break;
}
}
assert_eq!(fired_at, Some(12), "trips on the 12th disagreement");
assert_eq!(ep.rejected_at(), Some(12));
}
#[test]
fn never_reset_discipline() {
let cal = Calibration::worked_example();
let mut ep = EProcess::new(&cal);
for _ in 0..30 {
ep.observe(true);
}
assert!(ep.e_value() > 1.0);
assert!(ep.rejected());
}
#[test]
fn e_value_saturates_finite() {
let cal = Calibration::worked_example();
let mut ep = EProcess::new(&cal);
for _ in 0..100_000 {
ep.observe(true);
}
assert!(ep.e_value().is_finite());
assert!(ep.e_value() <= f64::MAX / 2.0);
}
#[test]
fn ville_bound_holds_under_synthetic_null() {
let cal = Calibration::worked_example();
let q0 = cal.q0();
let alpha = cal.alpha();
let streams = 20_000u64;
let steps_per_stream = 2_000u64;
let mut rng = SplitMix64(0xA11E_2026);
let mut breaches = 0u64;
for _ in 0..streams {
let mut ep = EProcess::new(&cal);
for _ in 0..steps_per_stream {
let disagree = rng.unit() < q0;
if ep.observe(disagree) {
breaches += 1;
break;
}
}
}
let level = GuardLevel {
streams,
breaches,
alpha,
};
assert!(
level.breach_fraction() <= alpha + 0.01,
"breach fraction {} exceeds alpha {} (+slack)",
level.breach_fraction(),
alpha
);
assert!(level.coverage() > 0.98);
}
#[test]
fn ville_fires_under_synthetic_alternative() {
let cal = Calibration::worked_example();
let true_rate = 0.30;
let streams = 2_000u64;
let mut rng = SplitMix64(0xBEEF_2026);
let mut breaches = 0u64;
for _ in 0..streams {
let mut ep = EProcess::new(&cal);
for _ in 0..5_000u64 {
let disagree = rng.unit() < true_rate;
if ep.observe(disagree) {
breaches += 1;
break;
}
}
}
let frac = breaches as f64 / streams as f64;
assert!(frac > 0.99, "alternative should fire reliably, got {frac}");
}
#[test]
fn loss_matrix_flip_dominates() {
let lm = LossMatrix::default();
assert!(lm.l_flip > lm.c_save * 1000.0);
assert!(lm.expected_accept_loss(0.5) > lm.fallback_loss());
assert!(lm.break_even_flip_rate() < 1e-3);
}
#[test]
fn fallback_fires_when_alpha_unset() {
let guard = SpeculativeGuard::from_calibration(None, 10_000);
assert!(!guard.speculation_active());
assert!(guard.ledger().fallback_active());
assert!(guard.should_verify(f64::INFINITY));
assert!(guard.should_verify(1000.0));
}
#[test]
fn fallback_fires_when_calibration_invalid() {
let bad = Calibration::new(0.0, 0.97, 0.5, 8, 4.0); let guard = SpeculativeGuard::from_calibration(Some(bad), 10_000);
assert!(!guard.speculation_active());
assert!(guard.should_verify(1e9));
}
#[test]
fn fallback_fires_when_calibration_corpus_too_small() {
let cal = Ok(Calibration::worked_example());
let guard = SpeculativeGuard::from_calibration(Some(cal), MIN_CALIBRATION_SAMPLES - 1);
assert!(!guard.speculation_active());
assert!(guard.should_verify(1e9));
let cal2 = Ok(Calibration::worked_example());
let guard2 = SpeculativeGuard::from_calibration(Some(cal2), MIN_CALIBRATION_SAMPLES);
assert!(guard2.speculation_active());
}
#[test]
fn disabled_guard_verifies_every_step() {
let mut guard = SpeculativeGuard::disabled();
for step in 0..50u32 {
let full = step + 100;
let draft = step; let dec = guard.decide(1e9, full, draft);
assert_eq!(dec.action, Action::FallBack);
assert_eq!(dec.emitted_token, full, "disabled guard emits full token");
}
assert!(guard.e_value().is_none());
assert_eq!(guard.ledger().accepted(), 0);
}
#[test]
fn early_exit_output_identical_to_full_decode() {
let vocab = 32usize;
let n_steps = 4_000usize;
let mut rng = SplitMix64(0xC0FF_EE42);
let cal = Calibration::worked_example();
let mut guard = SpeculativeGuard::enabled(cal);
for step in 0..n_steps {
let mut full = vec![0.0f32; vocab];
for v in &mut full {
*v = (rng.unit() as f32) * 2.0;
}
let winner = (rng.next_u64() as usize) % vocab;
full[winner] += 8.0; let full_tok = argmax(&full);
let draft_disagrees = rng.unit() < 0.03;
let mut draft = full.clone();
let draft_tok = if draft_disagrees {
let other = (winner + 1) % vocab;
draft[other] += 20.0;
argmax(&draft)
} else {
full_tok
};
let margin = top1_top2_margin(&draft);
let dec = guard.decide(margin, full_tok, draft_tok);
assert_eq!(
dec.emitted_token, full_tok,
"step {step}: guarded decode diverged from full decode \
(action={:?}, draft_tok={draft_tok}, full_tok={full_tok})",
dec.action
);
if dec.action == Action::AcceptDraft {
assert_eq!(
draft_tok, full_tok,
"step {step}: accepted a draft that disagreed with full"
);
}
}
let led = guard.ledger();
assert!(led.steps() == n_steps as u64);
assert!(
led.accepted() > 0,
"expected some accepted (verify-skipped) steps"
);
assert!(
led.measured_disagreement_rate() < 0.10,
"measured disagreement {} unexpectedly high",
led.measured_disagreement_rate()
);
}
#[test]
fn output_identity_holds_under_adversarial_draft() {
let vocab = 16usize;
let n_steps = 500usize;
let mut rng = SplitMix64(0xDEAD_F00D);
let cal = Calibration::worked_example();
let mut guard = SpeculativeGuard::enabled(cal);
for step in 0..n_steps {
let mut full = vec![0.0f32; vocab];
for v in &mut full {
*v = rng.unit() as f32;
}
let winner = (rng.next_u64() as usize) % vocab;
full[winner] += 10.0;
let full_tok = argmax(&full);
let other = (winner + 3) % vocab;
let mut draft = full.clone();
draft[other] += 50.0;
let draft_tok = argmax(&draft);
let margin = top1_top2_margin(&draft);
let dec = guard.decide(margin, full_tok, draft_tok);
assert_eq!(
dec.emitted_token, full_tok,
"step {step}: adversarial draft changed the output"
);
}
assert!(guard.tripped(), "adversarial draft should trip the guard");
assert!(!guard.speculation_active());
}
#[test]
fn trip_latches_safe_state_for_rest_of_document() {
let cal = Calibration::worked_example();
let mut guard = SpeculativeGuard::enabled(cal);
for _ in 0..30u32 {
let full = 7u32;
let draft = 9u32; let _ = guard.decide(0.0, full, draft); }
assert!(guard.tripped());
let dec = guard.decide(1e9, 7, 7);
assert_eq!(dec.action, Action::FallBack);
assert_eq!(dec.emitted_token, 7);
}
#[test]
fn guard_level_holds_and_coverage() {
let level = GuardLevel {
streams: 10_000,
breaches: 80,
alpha: 0.01,
};
assert!(level.holds_level()); assert!(approx(level.coverage(), 0.992, 1e-12));
let over = GuardLevel {
streams: 10_000,
breaches: 200,
alpha: 0.01,
};
assert!(!over.holds_level()); }
#[test]
fn argmax_first_max_wins_on_ties() {
assert_eq!(argmax(&[1.0, 5.0, 5.0, 2.0]), 1);
assert_eq!(argmax(&[3.0, 1.0, 2.0]), 0);
}
#[test]
fn margin_is_top1_minus_top2() {
assert!(approx(top1_top2_margin(&[1.0, 4.0, 2.0]), 2.0, 1e-6));
assert!(approx(top1_top2_margin(&[5.0, 5.0]), 0.0, 1e-6));
}
}