Skip to main content

sim_lib_numbers_stats/
transition.rs

1//! Shared finite row-stochastic transition representation.
2
3use std::{error::Error, fmt};
4
5const STOCHASTIC_TOLERANCE: f64 = 1.0e-10;
6
7/// A finite state vocabulary and row-stochastic transition matrix.
8///
9/// This is the shared transition representation used by observable Markov
10/// models and hidden-state models. State order is caller-owned and retained.
11#[derive(Clone, Debug, PartialEq)]
12pub struct FiniteTransitionMatrix<S> {
13    states: Vec<S>,
14    probabilities: Vec<Vec<f64>>,
15}
16
17impl<S: Eq + Clone> FiniteTransitionMatrix<S> {
18    /// Builds a checked matrix from ordered states and probability rows.
19    pub fn new(states: Vec<S>, probabilities: Vec<Vec<f64>>) -> Result<Self, TransitionError> {
20        if states.is_empty() {
21            return Err(TransitionError::EmptyStates);
22        }
23        for (index, state) in states.iter().enumerate() {
24            if states[..index].contains(state) {
25                return Err(TransitionError::DuplicateState { index });
26            }
27        }
28        if probabilities.len() != states.len() {
29            return Err(TransitionError::RowCount {
30                expected: states.len(),
31                actual: probabilities.len(),
32            });
33        }
34        for (row, probabilities) in probabilities.iter().enumerate() {
35            validate_distribution("transition", row, probabilities, states.len())?;
36        }
37        Ok(Self {
38            states,
39            probabilities,
40        })
41    }
42
43    pub(crate) fn from_normalized(states: Vec<S>, probabilities: Vec<Vec<f64>>) -> Self {
44        Self {
45            states,
46            probabilities,
47        }
48    }
49
50    /// Returns the ordered finite state vocabulary.
51    pub fn states(&self) -> &[S] {
52        &self.states
53    }
54
55    /// Returns the state count and matrix dimension.
56    pub fn len(&self) -> usize {
57        self.states.len()
58    }
59
60    /// Returns whether the state vocabulary is empty.
61    pub fn is_empty(&self) -> bool {
62        self.states.is_empty()
63    }
64
65    /// Returns all row-stochastic probability rows.
66    pub fn rows(&self) -> &[Vec<f64>] {
67        &self.probabilities
68    }
69
70    /// Returns a probability by state indices.
71    pub fn probability_by_index(&self, from: usize, to: usize) -> Option<f64> {
72        self.probabilities
73            .get(from)
74            .and_then(|row| row.get(to))
75            .copied()
76    }
77
78    /// Returns a probability by state values.
79    pub fn probability(&self, from: &S, to: &S) -> Option<f64> {
80        let from = self.states.iter().position(|state| state == from)?;
81        let to = self.states.iter().position(|state| state == to)?;
82        self.probability_by_index(from, to)
83    }
84}
85
86/// Failure while constructing a finite row-stochastic transition matrix.
87#[derive(Clone, Debug, PartialEq)]
88pub enum TransitionError {
89    /// No states were supplied.
90    EmptyStates,
91    /// A state duplicated an earlier state.
92    DuplicateState {
93        /// Position of the duplicate.
94        index: usize,
95    },
96    /// The matrix row count did not match the state count.
97    RowCount {
98        /// Required row count.
99        expected: usize,
100        /// Supplied row count.
101        actual: usize,
102    },
103    /// A row length did not match the state count.
104    ColumnCount {
105        /// Zero-based row index.
106        row: usize,
107        /// Required column count.
108        expected: usize,
109        /// Supplied column count.
110        actual: usize,
111    },
112    /// A probability was non-finite or negative.
113    InvalidProbability {
114        /// Stable matrix or distribution name.
115        distribution: &'static str,
116        /// Zero-based row index.
117        row: usize,
118        /// Zero-based column index.
119        column: usize,
120        /// Rejected probability.
121        value: f64,
122    },
123    /// A probability row did not sum to one.
124    ProbabilityMass {
125        /// Stable matrix or distribution name.
126        distribution: &'static str,
127        /// Zero-based row index.
128        row: usize,
129        /// Observed probability mass.
130        sum: f64,
131    },
132}
133
134impl fmt::Display for TransitionError {
135    fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
136        match self {
137            Self::EmptyStates => write!(formatter, "finite transitions require at least one state"),
138            Self::DuplicateState { index } => {
139                write!(formatter, "finite transition state {index} is duplicated")
140            }
141            Self::RowCount { expected, actual } => write!(
142                formatter,
143                "finite transitions require {expected} rows, got {actual}"
144            ),
145            Self::ColumnCount {
146                row,
147                expected,
148                actual,
149            } => write!(
150                formatter,
151                "finite transition row {row} requires {expected} columns, got {actual}"
152            ),
153            Self::InvalidProbability {
154                distribution,
155                row,
156                column,
157                value,
158            } => write!(
159                formatter,
160                "{distribution} probability {row}:{column} must be finite and nonnegative, got {value}"
161            ),
162            Self::ProbabilityMass {
163                distribution,
164                row,
165                sum,
166            } => write!(
167                formatter,
168                "{distribution} probability row {row} must sum to one, got {sum}"
169            ),
170        }
171    }
172}
173
174impl Error for TransitionError {}
175
176pub(crate) fn validate_distribution(
177    distribution: &'static str,
178    row: usize,
179    probabilities: &[f64],
180    expected: usize,
181) -> Result<(), TransitionError> {
182    if probabilities.len() != expected {
183        return Err(TransitionError::ColumnCount {
184            row,
185            expected,
186            actual: probabilities.len(),
187        });
188    }
189    let mut sum = 0.0;
190    for (column, probability) in probabilities.iter().copied().enumerate() {
191        if !probability.is_finite() || probability < 0.0 {
192            return Err(TransitionError::InvalidProbability {
193                distribution,
194                row,
195                column,
196                value: probability,
197            });
198        }
199        sum += probability;
200    }
201    if (sum - 1.0).abs() > STOCHASTIC_TOLERANCE {
202        return Err(TransitionError::ProbabilityMass {
203            distribution,
204            row,
205            sum,
206        });
207    }
208    Ok(())
209}