#![doc = include_str!(concat!(env!("CARGO_MANIFEST_DIR"), "/src/docs/runtime/trace.md"))]
use crate::core::address::Address;
use crate::error::{FugueError, FugueResult};
use std::collections::BTreeMap;
#[derive(Clone, Debug, PartialEq)]
pub enum ChoiceValue {
F64(f64),
I64(i64),
U64(u64),
Usize(usize),
Bool(bool),
}
impl ChoiceValue {
pub fn as_f64(&self) -> Option<f64> {
match self {
ChoiceValue::F64(v) => Some(*v),
_ => None,
}
}
pub fn as_bool(&self) -> Option<bool> {
match self {
ChoiceValue::Bool(v) => Some(*v),
_ => None,
}
}
pub fn as_u64(&self) -> Option<u64> {
match self {
ChoiceValue::U64(v) => Some(*v),
_ => None,
}
}
pub fn as_usize(&self) -> Option<usize> {
match self {
ChoiceValue::Usize(v) => Some(*v),
_ => None,
}
}
pub fn as_i64(&self) -> Option<i64> {
match self {
ChoiceValue::I64(v) => Some(*v),
_ => None,
}
}
pub fn type_name(&self) -> &'static str {
match self {
ChoiceValue::F64(_) => "f64",
ChoiceValue::Bool(_) => "bool",
ChoiceValue::U64(_) => "u64",
ChoiceValue::Usize(_) => "usize",
ChoiceValue::I64(_) => "i64",
}
}
}
#[derive(Clone, Debug)]
pub struct Choice {
pub addr: Address,
pub value: ChoiceValue,
pub logp: f64,
}
#[derive(Clone, Debug, Default)]
pub struct Trace {
pub choices: BTreeMap<Address, Choice>,
pub log_prior: f64,
pub log_likelihood: f64,
pub log_factors: f64,
}
impl Trace {
pub fn total_log_weight(&self) -> f64 {
self.log_prior + self.log_likelihood + self.log_factors
}
pub fn get_f64(&self, addr: &Address) -> Option<f64> {
self.choices.get(addr)?.value.as_f64()
}
pub fn get_bool(&self, addr: &Address) -> Option<bool> {
self.choices.get(addr)?.value.as_bool()
}
pub fn get_u64(&self, addr: &Address) -> Option<u64> {
self.choices.get(addr)?.value.as_u64()
}
pub fn get_usize(&self, addr: &Address) -> Option<usize> {
self.choices.get(addr)?.value.as_usize()
}
pub fn get_i64(&self, addr: &Address) -> Option<i64> {
self.choices.get(addr)?.value.as_i64()
}
pub fn get_f64_result(&self, addr: &Address) -> FugueResult<f64> {
let choice = self.choices.get(addr).ok_or_else(|| {
FugueError::trace_error(
"get_f64",
Some(addr.clone()),
"Address not found in trace",
crate::error::ErrorCode::TraceAddressNotFound,
)
})?;
choice
.value
.as_f64()
.ok_or_else(|| FugueError::type_mismatch(addr.clone(), "f64", choice.value.type_name()))
}
pub fn get_bool_result(&self, addr: &Address) -> FugueResult<bool> {
let choice = self.choices.get(addr).ok_or_else(|| {
FugueError::trace_error(
"get_bool",
Some(addr.clone()),
"Address not found in trace",
crate::error::ErrorCode::TraceAddressNotFound,
)
})?;
choice.value.as_bool().ok_or_else(|| {
FugueError::type_mismatch(addr.clone(), "bool", choice.value.type_name())
})
}
pub fn get_u64_result(&self, addr: &Address) -> FugueResult<u64> {
let choice = self.choices.get(addr).ok_or_else(|| {
FugueError::trace_error(
"get_u64",
Some(addr.clone()),
"Address not found in trace",
crate::error::ErrorCode::TraceAddressNotFound,
)
})?;
choice
.value
.as_u64()
.ok_or_else(|| FugueError::type_mismatch(addr.clone(), "u64", choice.value.type_name()))
}
pub fn get_usize_result(&self, addr: &Address) -> FugueResult<usize> {
let choice = self.choices.get(addr).ok_or_else(|| {
FugueError::trace_error(
"get_usize",
Some(addr.clone()),
"Address not found in trace",
crate::error::ErrorCode::TraceAddressNotFound,
)
})?;
choice.value.as_usize().ok_or_else(|| {
FugueError::type_mismatch(addr.clone(), "usize", choice.value.type_name())
})
}
pub fn get_i64_result(&self, addr: &Address) -> FugueResult<i64> {
let choice = self.choices.get(addr).ok_or_else(|| {
FugueError::trace_error(
"get_i64",
Some(addr.clone()),
"Address not found in trace",
crate::error::ErrorCode::TraceAddressNotFound,
)
})?;
choice
.value
.as_i64()
.ok_or_else(|| FugueError::type_mismatch(addr.clone(), "i64", choice.value.type_name()))
}
pub fn insert_choice(&mut self, addr: Address, value: ChoiceValue, logp: f64) {
let choice = Choice {
addr: addr.clone(),
value,
logp,
};
self.choices.insert(addr, choice);
}
fn prefix_addresses(&self, prefix: &str) -> Vec<Address> {
self.choices
.range(Address::new(prefix.to_string())..)
.take_while(|(a, _)| a.as_str().starts_with(prefix))
.filter(|(a, _)| a.has_prefix(prefix))
.map(|(a, _)| a.clone())
.collect()
}
pub fn extract_prefix(&self, prefix: &str) -> Trace {
let mut out = Trace::default();
for addr in self.prefix_addresses(prefix) {
out.choices
.insert(addr.clone(), self.choices[&addr].clone());
}
out
}
pub fn truncate_prefix(&mut self, prefix: &str) {
for addr in self.prefix_addresses(prefix) {
self.choices.remove(&addr);
}
}
pub fn graft_prefix(&mut self, prefix: &str, donor: &Trace) {
self.truncate_prefix(prefix);
for addr in donor.prefix_addresses(prefix) {
self.choices
.insert(addr.clone(), donor.choices[&addr].clone());
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::addr;
#[test]
fn insert_and_getters_work() {
let mut t = Trace::default();
t.insert_choice(addr!("a"), ChoiceValue::F64(1.5), -0.5);
t.insert_choice(addr!("b"), ChoiceValue::Bool(true), -0.7);
t.insert_choice(addr!("c"), ChoiceValue::U64(3), -0.2);
t.insert_choice(addr!("d"), ChoiceValue::Usize(4), -0.3);
t.insert_choice(addr!("e"), ChoiceValue::I64(-7), -0.1);
assert_eq!(t.get_f64(&addr!("a")), Some(1.5));
assert_eq!(t.get_bool(&addr!("b")), Some(true));
assert_eq!(t.get_u64(&addr!("c")), Some(3));
assert_eq!(t.get_usize(&addr!("d")), Some(4));
assert_eq!(t.get_i64(&addr!("e")), Some(-7));
assert!(t.get_f64_result(&addr!("a")).is_ok());
assert!(t.get_bool_result(&addr!("b")).is_ok());
assert!(t.get_u64_result(&addr!("c")).is_ok());
assert!(t.get_usize_result(&addr!("d")).is_ok());
assert!(t.get_i64_result(&addr!("e")).is_ok());
let err = t.get_f64_result(&addr!("b")).unwrap_err();
assert!(matches!(err, crate::error::FugueError::TypeMismatch { .. }));
}
#[test]
fn total_log_weight_accumulates() {
let mut t = Trace::default();
t.insert_choice(addr!("x"), ChoiceValue::F64(0.0), -1.0);
t.log_prior = -1.0;
t.log_likelihood = -2.0;
t.log_factors = -3.0;
assert!((t.total_log_weight() - (-6.0)).abs() < 1e-12);
}
#[test]
fn result_accessors_return_errors_for_missing_addresses() {
let t = Trace::default();
let e = t.get_f64_result(&addr!("missing")).unwrap_err();
assert!(matches!(e, crate::error::FugueError::TraceError { .. }));
}
#[test]
fn test_extract_prefix_boundary() {
let mut t = Trace::default();
for (name, v) in [
("a", 0.0),
("a#1", 1.0),
("a/0", 2.0),
("a/1/x", 3.0),
("a::s", 4.0),
("a0", 5.0),
("ab/0", 6.0),
] {
t.insert_choice(Address::new(name), ChoiceValue::F64(v), -0.5);
}
let sub = t.extract_prefix("a");
let got: Vec<&str> = sub.choices.keys().map(|a| a.as_str()).collect();
assert_eq!(got, vec!["a", "a#1", "a/0", "a/1/x", "a::s"]);
let slash = t.extract_prefix("a/");
let got: Vec<&str> = slash.choices.keys().map(|a| a.as_str()).collect();
assert_eq!(got, vec!["a/0", "a/1/x"]);
}
#[test]
fn test_graft_round_trip() {
let mut t = Trace::default();
t.insert_choice(addr!("gene", 0), ChoiceValue::F64(1.0), -0.1);
t.insert_choice(addr!("gene", 1), ChoiceValue::F64(2.0), -0.2);
t.insert_choice(addr!("other"), ChoiceValue::Bool(true), -0.3);
let before: Vec<(String, ChoiceValue)> = t
.choices
.iter()
.map(|(a, c)| (a.as_str().to_string(), c.value.clone()))
.collect();
let block = t.extract_prefix("gene");
t.graft_prefix("gene", &block);
let after: Vec<(String, ChoiceValue)> = t
.choices
.iter()
.map(|(a, c)| (a.as_str().to_string(), c.value.clone()))
.collect();
assert_eq!(before, after);
let mut donor = Trace::default();
donor.insert_choice(addr!("gene", 0), ChoiceValue::F64(9.0), -0.9);
donor.insert_choice(addr!("gene", 1), ChoiceValue::F64(8.0), -0.8);
donor.insert_choice(addr!("other"), ChoiceValue::Bool(false), -0.7);
t.graft_prefix("gene", &donor);
assert_eq!(t.get_f64(&addr!("gene", 0)), Some(9.0));
assert_eq!(t.get_f64(&addr!("gene", 1)), Some(8.0));
assert_eq!(t.get_bool(&addr!("other")), Some(true));
}
#[test]
fn test_extract_zeroes_accumulators() {
let mut t = Trace {
log_prior: -1.0,
log_likelihood: -2.0,
log_factors: -3.0,
..Default::default()
};
t.insert_choice(addr!("x", 0), ChoiceValue::F64(0.5), -0.5);
let sub = t.extract_prefix("x");
assert_eq!(sub.log_prior, 0.0);
assert_eq!(sub.log_likelihood, 0.0);
assert_eq!(sub.log_factors, 0.0);
assert_eq!(sub.choices.len(), 1);
}
}