use std::fmt::Write;
use approx::{AbsDiffEq, RelativeEq};
use itertools::Itertools;
use ndarray::prelude::*;
use serde::{
Deserialize, Deserializer, Serialize, Serializer,
de::{MapAccess, Visitor},
ser::SerializeMap,
};
use crate::{
datasets::{CatEv, CatIncTable, CatSample, CatTable, CatWtdTable},
impl_json_io,
inference::TopologicalOrder,
io::{BifIO, BifParser},
models::{BN, CPD, CatCPD, DiGraph, Graph, Labelled},
set,
types::{Error, Labels, Map, Result, Set, States},
};
#[derive(Clone, Debug)]
pub struct CatBN {
name: Option<String>,
description: Option<String>,
labels: Labels,
states: States,
shape: Array1<usize>,
graph: DiGraph,
cpds: Map<String, CatCPD>,
topological_order: Vec<usize>,
}
impl CatBN {
#[inline]
pub const fn states(&self) -> &States {
&self.states
}
#[inline]
pub fn shape(&self) -> &Array1<usize> {
&self.shape
}
}
impl PartialEq for CatBN {
fn eq(&self, other: &Self) -> bool {
self.labels.eq(&other.labels)
&& self.states.eq(&other.states)
&& self.shape.eq(&other.shape)
&& self.graph.eq(&other.graph)
&& self.topological_order.eq(&other.topological_order)
&& self.cpds.eq(&other.cpds)
}
}
impl AbsDiffEq for CatBN {
type Epsilon = f64;
fn default_epsilon() -> Self::Epsilon {
Self::Epsilon::default_epsilon()
}
fn abs_diff_eq(&self, other: &Self, epsilon: Self::Epsilon) -> bool {
self.labels.eq(&other.labels)
&& self.states.eq(&other.states)
&& self.shape.eq(&other.shape)
&& self.graph.eq(&other.graph)
&& self.topological_order.eq(&other.topological_order)
&& self
.cpds
.iter()
.zip(&other.cpds)
.all(|((label, cpd), (other_label, other_cpd))| {
label.eq(other_label) && cpd.abs_diff_eq(other_cpd, epsilon)
})
}
}
impl RelativeEq for CatBN {
fn default_max_relative() -> Self::Epsilon {
Self::Epsilon::default_max_relative()
}
fn relative_eq(
&self,
other: &Self,
epsilon: Self::Epsilon,
max_relative: Self::Epsilon,
) -> bool {
self.labels.eq(&other.labels)
&& self.states.eq(&other.states)
&& self.shape.eq(&other.shape)
&& self.graph.eq(&other.graph)
&& self.topological_order.eq(&other.topological_order)
&& self
.cpds
.iter()
.zip(&other.cpds)
.all(|((label, cpd), (other_label, other_cpd))| {
label.eq(other_label) && cpd.relative_eq(other_cpd, epsilon, max_relative)
})
}
}
impl Labelled for CatBN {
#[inline]
fn labels(&self) -> &Labels {
&self.labels
}
}
impl BN for CatBN {
type CPD = CatCPD;
type Evidence = CatEv;
type Sample = CatSample;
type Samples = CatTable;
type IncSamples = CatIncTable;
type WtdSamples = CatWtdTable;
fn new<I>(graph: DiGraph, cpds: I) -> Result<Self>
where
I: IntoIterator<Item = Self::CPD>,
{
let mut cpds: Map<_, _> = cpds
.into_iter()
.map(|x| {
if x.labels().len() != 1 {
return Err(Error::InvalidParameter(
"cpd",
"CPD must contain exactly one label.",
));
}
Ok((x.labels()[0].to_owned(), x))
})
.collect::<Result<_>>()?;
cpds.sort_keys();
if !graph.labels().iter().eq(cpds.keys()) {
return Err(Error::LabelMismatch("graph labels", "distributions labels"));
}
let mut states: States = Default::default();
cpds.values().try_for_each(|cpd| {
cpd.states()
.iter()
.chain(cpd.conditioning_states())
.try_for_each(|(l, s)| {
if let Some(existing_states) = states.get(l) {
if existing_states != s {
return Err(Error::InvalidParameter(
"cpds",
&format!("States of `{l}` must be the same across CPDs."),
));
}
} else {
states.insert(l.to_owned(), s.clone());
}
Ok(())
})
})?;
states.sort_keys();
let labels: Labels = states.keys().cloned().collect();
let shape: Array1<usize> = states.values().map(|s| s.len()).collect();
graph.vertices().into_iter().try_for_each(|i| {
let pa_i = graph.parents(&set![i])?.into_iter();
let pa_i: &Labels = &pa_i.map(|j| labels[j].to_owned()).collect(); let pa_j = cpds[&labels[i]].conditioning_labels();
if pa_i != pa_j {
return Err(Error::LabelMismatch(
&format!("{pa_i:?}"),
&format!("{pa_j:?}"),
));
}
Ok(())
})?;
let topological_order = graph.topological_order().ok_or_else(|| Error::NotADag())?;
Ok(Self {
name: None,
description: None,
labels,
states,
shape,
graph,
cpds,
topological_order,
})
}
#[inline]
fn name(&self) -> Option<&str> {
self.name.as_deref()
}
#[inline]
fn description(&self) -> Option<&str> {
self.description.as_deref()
}
#[inline]
fn graph(&self) -> &DiGraph {
&self.graph
}
#[inline]
fn cpds(&self) -> &Map<String, Self::CPD> {
&self.cpds
}
#[inline]
fn parameters_size(&self) -> usize {
self.cpds.iter().map(|(_, x)| x.parameters_size()).sum()
}
fn select(&self, x: &Set<usize>) -> Result<Self>
where
Self: Sized,
{
x.iter().try_for_each(|&i| {
if i >= self.labels.len() {
return Err(Error::IndexOutOfBounds(i));
}
Ok(())
})?;
let mut x = x.clone();
x.sort();
let graph = self.graph.select(&x)?;
let cpds = x.iter().map(|&i| self.cpds[i].clone());
Self::with_optionals(
self.name.clone(),
self.description.clone(),
graph,
cpds,
)
}
#[inline]
fn topological_order(&self) -> &[usize] {
&self.topological_order
}
fn with_optionals<I>(
name: Option<String>,
description: Option<String>,
graph: DiGraph,
cpds: I,
) -> Result<Self>
where
I: IntoIterator<Item = Self::CPD>,
{
if let Some(name) = &name
&& name.is_empty()
{
return Err(Error::InvalidParameter("name", "cannot be empty"));
}
if let Some(description) = &description
&& description.is_empty()
{
return Err(Error::InvalidParameter("description", "cannot be empty"));
}
let mut bn = Self::new(graph, cpds)?;
bn.name = name;
bn.description = description;
Ok(bn)
}
}
impl Serialize for CatBN {
fn serialize<S>(&self, serializer: S) -> std::result::Result<S::Ok, S::Error>
where
S: Serializer,
{
let mut size = 3;
size += self.name.is_some() as usize;
size += self.description.is_some() as usize;
let mut map = serializer.serialize_map(Some(size))?;
if let Some(name) = &self.name {
map.serialize_entry("name", name)?;
}
if let Some(description) = &self.description {
map.serialize_entry("description", description)?;
}
map.serialize_entry("graph", &self.graph)?;
let cpds: Vec<_> = self.cpds.values().cloned().collect();
map.serialize_entry("cpds", &cpds)?;
map.serialize_entry("type", "catbn")?;
map.end()
}
}
impl<'de> Deserialize<'de> for CatBN {
fn deserialize<D>(deserializer: D) -> std::result::Result<Self, D::Error>
where
D: Deserializer<'de>,
{
#[derive(Deserialize)]
#[serde(field_identifier, rename_all = "snake_case")]
enum Field {
Name,
Description,
Graph,
Cpds,
Type,
}
struct CatBNVisitor;
impl<'de> Visitor<'de> for CatBNVisitor {
type Value = CatBN;
fn expecting(&self, formatter: &mut std::fmt::Formatter) -> std::fmt::Result {
formatter.write_str("struct CatBN")
}
fn visit_map<V>(self, mut map: V) -> std::result::Result<CatBN, V::Error>
where
V: MapAccess<'de>,
{
use serde::de::Error as E;
let mut name = None;
let mut description = None;
let mut graph = None;
let mut cpds = None;
let mut type_ = None;
while let Some(key) = map.next_key()? {
match key {
Field::Name => {
if name.is_some() {
return Err(E::duplicate_field("name"));
}
name = Some(map.next_value()?);
}
Field::Description => {
if description.is_some() {
return Err(E::duplicate_field("description"));
}
description = Some(map.next_value()?);
}
Field::Graph => {
if graph.is_some() {
return Err(E::duplicate_field("graph"));
}
graph = Some(map.next_value()?);
}
Field::Cpds => {
if cpds.is_some() {
return Err(E::duplicate_field("cpds"));
}
cpds = Some(map.next_value()?);
}
Field::Type => {
if type_.is_some() {
return Err(E::duplicate_field("type"));
}
type_ = Some(map.next_value()?);
}
}
}
let graph = graph.ok_or_else(|| E::missing_field("graph"))?;
let cpds = cpds.ok_or_else(|| E::missing_field("cpds"))?;
let type_: String = type_.ok_or_else(|| E::missing_field("type"))?;
if type_ != "catbn" {
return Err(E::custom(format!(
"Invalid type for CatBN: expected 'catbn', found '{type_}'"
)));
}
let cpds: Vec<_> = cpds;
CatBN::with_optionals(name, description, graph, cpds)
.map_err(serde::de::Error::custom)
}
}
const FIELDS: &[&str] = &["name", "description", "graph", "cpds", "type"];
deserializer.deserialize_struct("CatBN", FIELDS, CatBNVisitor)
}
}
impl_json_io!(CatBN);
impl BifIO for CatBN {
fn from_bif_string(bif: &str) -> Result<Self> {
BifParser::parse_str(bif)
}
fn to_bif_string(&self) -> Result<String> {
let mut f = String::new();
writeln!(
f,
"network {} {{",
self.name.as_deref().unwrap_or("Network")
)
.map_err(|e| Error::Parsing(&e.to_string()))?;
if let Some(description) = &self.description {
writeln!(f, " property description \"{}\";", description)
.map_err(|e| Error::Parsing(&e.to_string()))?;
}
writeln!(f, "}}").map_err(|e| Error::Parsing(&e.to_string()))?;
for label in self.labels() {
writeln!(f, "variable {} {{", label).map_err(|e| Error::Parsing(&e.to_string()))?;
let states = &self.states()[label];
let states_str = states.iter().map(|x| x.to_string()).join(", ");
writeln!(
f,
" type discrete [ {} ] {{ {} }};",
states.len(),
states_str
)
.map_err(|e| Error::Parsing(&e.to_string()))?;
writeln!(f, "}}").map_err(|e| Error::Parsing(&e.to_string()))?;
}
for (label, cpd) in &self.cpds {
let parents = cpd.conditioning_labels();
if parents.is_empty() {
writeln!(f, "probability ( {} ) {{", label)
.map_err(|e| Error::Parsing(&e.to_string()))?;
} else {
let parents_str = parents.iter().map(|x| x.to_string()).join(", ");
writeln!(f, "probability ( {} | {} ) {{", label, parents_str)
.map_err(|e| Error::Parsing(&e.to_string()))?;
}
if parents.is_empty() {
let values = cpd.parameters().iter().map(|x| x.to_string()).join(", ");
writeln!(f, " table {};", values).map_err(|e| Error::Parsing(&e.to_string()))?;
} else {
let conditioning_states = cpd.conditioning_states();
let combinations = parents
.iter()
.map(|p| &conditioning_states[p])
.multi_cartesian_product();
for (i, states) in combinations.enumerate() {
let states_str = states.iter().map(|x| x.to_string()).join(", ");
let probs_str = cpd
.parameters()
.row(i)
.iter()
.map(|x| x.to_string())
.join(", ");
writeln!(f, " ({}) {};", states_str, probs_str)
.map_err(|e| Error::Parsing(&e.to_string()))?;
}
}
writeln!(f, "}}").map_err(|e| Error::Parsing(&e.to_string()))?;
}
Ok(f)
}
fn from_bif_file(path: &str) -> Result<Self> {
Self::from_bif_string(&std::fs::read_to_string(path).map_err(Error::from)?)
}
fn to_bif_file(&self, path: &str) -> Result<()> {
std::fs::write(path, self.to_bif_string()?).map_err(Error::from)
}
}