use rand::prelude::*;
use serde::{Deserialize, Serialize};
use weighted_rand::builder::*;
use weighted_rand::table::WalkerTable;
#[derive(Serialize, Deserialize, Debug, PartialEq)]
pub struct MarkovChain<T> {
state_space: Vec<T>,
wa_table: Vec<WalkerTable>,
prev_index: usize,
}
impl<T> MarkovChain<T>
where
T: Clone,
T: Eq,
T: Ord,
T: PartialOrd,
T: PartialEq,
{
fn new(state_space: Vec<T>, wa_table: Vec<WalkerTable>, prev_index: usize) -> MarkovChain<T> {
MarkovChain {
state_space: state_space,
wa_table: wa_table,
prev_index: prev_index,
}
}
pub fn from(elements: &[T]) -> MarkovChain<T> {
let mut state_space = elements.to_vec();
state_space.sort();
state_space.dedup();
let space_len = state_space.len();
let mut freq_table = vec![vec![0; space_len]; space_len];
let mut prev_index: Option<usize> = None;
for element in elements {
let cur_index = state_space
.iter()
.position(|state| *element == *state)
.expect("There is no state that should exist.");
if let Some(i) = prev_index {
freq_table[i][cur_index] += 1;
}
prev_index = Some(cur_index);
}
let mut wa_table = Vec::with_capacity(space_len);
for row in freq_table {
let builder = WalkerTableBuilder::new(&row);
wa_table.push(builder.build());
}
MarkovChain::new(state_space, wa_table, space_len)
}
pub fn next(&mut self) -> &T {
let mut rng = rand::thread_rng();
self.next_rng(&mut rng)
}
pub fn next_rng(&mut self, rng: &mut ThreadRng) -> &T {
let row = {
if self.prev_index == self.state_space.len() {
self.prev_index = rng.gen_range(0..self.state_space.len());
}
self.prev_index
};
let elem_index = self.wa_table[row].next_rng(rng);
self.prev_index = elem_index;
&self.state_space[elem_index]
}
pub fn initialize(&mut self) {
self.prev_index = self.state_space.len();
}
}
#[cfg(test)]
mod markov_test {
use crate::MarkovChain;
use weighted_rand::table::WalkerTable;
const TEXT: [&str; 11] = [
"I", "think", "that", "that", "that", "that", "that", "boy", "wrote", "is", "wrong",
];
#[test]
fn make_markov_model() {
let actual = MarkovChain::from(&TEXT);
let expected = MarkovChain {
state_space: vec!["I", "boy", "is", "that", "think", "wrong", "wrote"],
wa_table: vec![
WalkerTable::new(
vec![4, 4, 4, 4, 4, 4, 4],
vec![1.0, 1.0, 1.0, 1.0, 1.0, 1.0, 1.0],
),
WalkerTable::new(
vec![6, 6, 6, 6, 6, 6, 6],
vec![1.0, 1.0, 1.0, 1.0, 1.0, 1.0, 1.0],
),
WalkerTable::new(
vec![5, 5, 5, 5, 5, 5, 5],
vec![1.0, 1.0, 1.0, 1.0, 1.0, 1.0, 1.0],
),
WalkerTable::new(
vec![3, 1, 3, 1, 3, 3, 3],
vec![1.0, 1.0, 1.0, 0.4, 1.0, 1.0, 1.0],
),
WalkerTable::new(
vec![3, 3, 3, 3, 3, 3, 3],
vec![1.0, 1.0, 1.0, 1.0, 1.0, 1.0, 1.0],
),
WalkerTable::new(
vec![0, 0, 0, 0, 0, 0, 0],
vec![0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0],
),
WalkerTable::new(
vec![2, 2, 2, 2, 2, 2, 2],
vec![1.0, 1.0, 1.0, 1.0, 1.0, 1.0, 1.0],
),
],
prev_index: 7,
};
assert_eq!(actual, expected)
}
#[test]
fn generate_element() {
let mut model = MarkovChain::from(&TEXT);
let element = model.next();
let include = TEXT
.iter()
.fold(false, |acc, cur| if acc { acc } else { element == cur });
assert!(include)
}
#[test]
fn initialize() {
let mut model = MarkovChain::from(&TEXT);
model.next();
let before = model.prev_index;
model.initialize();
let after = model.prev_index;
assert!(before != after);
assert_eq!(after, 7);
}
}