1use std::{error::Error, fmt};
4
5const STOCHASTIC_TOLERANCE: f64 = 1.0e-10;
6
7#[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 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 pub fn states(&self) -> &[S] {
52 &self.states
53 }
54
55 pub fn len(&self) -> usize {
57 self.states.len()
58 }
59
60 pub fn is_empty(&self) -> bool {
62 self.states.is_empty()
63 }
64
65 pub fn rows(&self) -> &[Vec<f64>] {
67 &self.probabilities
68 }
69
70 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 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#[derive(Clone, Debug, PartialEq)]
88pub enum TransitionError {
89 EmptyStates,
91 DuplicateState {
93 index: usize,
95 },
96 RowCount {
98 expected: usize,
100 actual: usize,
102 },
103 ColumnCount {
105 row: usize,
107 expected: usize,
109 actual: usize,
111 },
112 InvalidProbability {
114 distribution: &'static str,
116 row: usize,
118 column: usize,
120 value: f64,
122 },
123 ProbabilityMass {
125 distribution: &'static str,
127 row: usize,
129 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}