use std::collections::HashMap;
use Code;
use daggy::{Dag, EdgeIndex, NodeIndex, Walker};
use daggy::petgraph::{EdgeDirection};
use rand::random;
use {b1, b3};
pub trait MLDecoder {
fn estimate_states(&self, signal: &Vec<b1>) -> (Dag<DecoderState, usize>,
Vec<NodeIndex>);
fn read_tree(&self, pruned_tree: (Dag<DecoderState, usize>,
Vec<NodeIndex>))
-> Vec<b1>;
fn decode(&self, signal: &Vec<b1>) -> Vec<b1>;
}
pub trait MAPDecoder {
fn estimate_states(signal: Vec<b1>, code: Code) -> (Dag<DecoderState, usize>,
Vec<NodeIndex>);
fn decode(signal: Vec<b1>, code: Code) -> Vec<b1>;
}
#[derive(Clone, Copy, Debug)]
pub struct DecoderState {
pub state_index: usize,
pub dist: usize
}
impl MLDecoder for Code {
fn estimate_states(&self, signal: &Vec<b1>) -> (Dag<DecoderState, usize>,
Vec<NodeIndex>) {
let chunk_width = self.polys.len() + 1;
let start_index: usize = self.start_state[0].into();
let mut decoder_tree = Dag::new();
let root = decoder_tree.add_node(
DecoderState { state_index: start_index, dist: 0 });
let mut previous_nodes = vec![root];
let mut pruned_states: HashMap<usize, (NodeIndex, DecoderState)> =
HashMap::new();
for next_bit_chunk in signal.chunks(chunk_width) {
for next_bit in next_bit_chunk {
let mut descendant_nodes = vec![];
for previous_node in previous_nodes {
let next_state_inds =
self.next_states(decoder_tree[previous_node].state_index);
for (iter_ind, state_ind) in next_state_inds.iter().enumerate() {
let next_bit_guess: usize = iter_ind % (next_bit.max() + 1);
let dist = next_bit_guess ^ usize::from(*next_bit);
let prev_dist = decoder_tree[previous_node].dist;
descendant_nodes.push(
decoder_tree.add_child(
previous_node,
dist,
DecoderState {
state_index: *state_ind,
dist: prev_dist + dist
}
).1);
}
}
previous_nodes = descendant_nodes;
}
pruned_states.clear();
for (node_ind, survivor_node) in previous_nodes
.iter()
.map(|&node_ind| (node_ind, decoder_tree[node_ind])) {
if pruned_states.contains_key(&survivor_node.state_index) {
let cur_node = pruned_states.get_mut(&survivor_node.state_index)
.unwrap();
if cur_node.1.dist > survivor_node.dist {
*cur_node = (node_ind, survivor_node);
}
} else {
pruned_states.insert(survivor_node.state_index,
(node_ind, survivor_node));
}
}
previous_nodes = pruned_states.keys().filter_map(|key| {
if let Some(state_tuple) = pruned_states.get(key) {
Some(state_tuple.0)
} else {
None
}
}).collect();
prune_state_tree(&mut decoder_tree, &previous_nodes);
}
(decoder_tree, previous_nodes)
}
fn read_tree(&self, pruned_tree: (Dag<DecoderState, usize>,
Vec<NodeIndex>))
-> Vec<b1> {
let survivor_node = pruned_tree.1[random::<usize>() % pruned_tree.1.len()];
let tree = pruned_tree.0;
tree.recursive_walk(survivor_node, |tree, node_ind| {
tree.parents(node_ind).iter(tree).nth(1)
}).iter(&tree)
.map(|(_, node_ind)| {
let code_word: b3 = tree[node_ind].state_index.into();
code_word.bits().to_vec()
})
.flat_map(|x| x)
.collect()
}
fn decode(&self, signal: &Vec<b1>) -> Vec<b1> {
self.read_tree(self.estimate_states(signal))
.iter().rev().cloned().collect()
}
}
fn prune_state_tree(tree: &mut Dag<DecoderState, usize>,
survivor_nodes: &Vec<NodeIndex>) {
let end_index = EdgeIndex::<u32>::end();
let dead_leaf_inds: Vec<NodeIndex> = tree.raw_nodes().iter().filter(|&node| {
node.next_edge(EdgeDirection::Outgoing) == end_index
}).filter_map(|leaf| {
let twig = leaf.next_edge(EdgeDirection::Incoming);
let node_id = tree.edge_endpoints(twig).unwrap().1;
if !survivor_nodes.contains(&node_id) {
Some(node_id)
} else {
None
}
}).collect();
let mut good_parents: Vec<NodeIndex> = vec![];
let mut nodes_to_remove: Vec<NodeIndex> = vec![];
let mut edges_to_remove: Vec<EdgeIndex> = vec![];
for leaf_ind in dead_leaf_inds {
tree.recursive_walk(leaf_ind, |tree, node_ind| {
let mut walk_res = None;
for (parent_edge, parent_ind) in tree.parents(node_ind).iter(tree) {
if good_parents.contains(&parent_ind) { continue; }
else if tree.children(parent_ind).any(tree, |_, _, child_ind| {
survivor_nodes.contains(&child_ind)
}) {
good_parents.push(parent_ind);
} else {
nodes_to_remove.push(parent_ind);
edges_to_remove.push(parent_edge);
walk_res = Some((parent_edge, parent_ind));
}
}
walk_res
}).last(&tree);
}
for node_ind in nodes_to_remove.iter() {
tree.remove_node(*node_ind);
}
for edge_ind in edges_to_remove.iter() {
tree.remove_edge(*edge_ind);
}
}