extract.rs

5.5 kB · rust · 179 lines

1use super::models::Network;2use crate::core::error::{value_error, Result};3use crate::core::tensor::Tensor;4use std::collections::HashSet;56fn center(coord: &[usize]) -> Vec<f64> {7    coord.iter().rev().map(|&c| c as f64 + 0.5).collect()8}910fn coords_of(grid: &Tensor) -> Vec<Vec<usize>> {11    let mut out = Vec::new();12    for flat in 0..grid.size() {13        if grid.at(flat) != 0 {14            let mut rem = flat;15            let mut multi = Vec::with_capacity(grid.shape.len());16            for axis in 0..grid.shape.len() {17                let stride: usize = grid.shape[(axis + 1)..].iter().product();18                multi.push(rem / stride);19                rem %= stride;20            }21            out.push(multi);22        }23    }24    out25}2627/// Extracts the network of filled sites joined to their axis neighbors.28///29/// # Errors30///31/// Errors off two and three dimensions.32pub fn core_graph(grid: &Tensor) -> Result<Network> {33    let dim = grid.shape.len();34    if dim != 2 && dim != 3 {35        return value_error("core_graph expects a 2D or 3D cell.");36    }37    let coords = coords_of(grid);38    let mut index_of = vec![usize::MAX; grid.size()];39    for (i, coord) in coords.iter().enumerate() {40        index_of[grid.index(coord)] = i;41    }42    let mut network = Network::new(dim);43    for coord in &coords {44        network.add_node(center(coord))?;45    }46    for coord in &coords {47        for axis in 0..dim {48            if coord[axis] + 1 >= grid.shape[axis] {49                continue;50            }51            let mut neighbor = coord.clone();52            neighbor[axis] += 1;53            if grid.get(&neighbor)? != 0 {54                network.add_branch(55                    index_of[grid.index(coord)],56                    index_of[grid.index(&neighbor)],57                    1.0,58                )?;59            }60        }61    }62    Ok(network)63}6465/// Extracts the core graph of the inverted grid, joining empty sites instead.66///67/// # Errors68///69/// Errors off two and three dimensions.70pub fn tunnel_graph(grid: &Tensor) -> Result<Network> {71    core_graph(&grid.invert())72}7374/// Extracts the network of corners and edges outlining every filled site.75///76/// # Errors77///78/// Errors off two and three dimensions.79pub fn edge_graph(grid: &Tensor) -> Result<Network> {80    let dim = grid.shape.len();81    if dim != 2 && dim != 3 {82        return value_error("edge_graph expects a 2D or 3D cell.");83    }84    let corner_shape: Vec<usize> = grid.shape.iter().map(|&s| s + 1).collect();85    let corner_index_dims = corner_shape.clone();86    let flat_corner = |coord: &[usize]| -> usize {87        let mut out = 0;88        for (axis, &c) in coord.iter().enumerate() {89            out = out * corner_index_dims[axis] + c;90        }91        out92    };93    let coords = coords_of(grid);94    let mut used: HashSet<Vec<usize>> = HashSet::new();95    for coord in &coords {96        for offset in 0..(1usize << dim) {97            let corner: Vec<usize> = (0..dim)98                .map(|axis| coord[axis] + ((offset >> (dim - 1 - axis)) & 1))99                .collect();100            used.insert(corner);101        }102    }103    let mut ordered: Vec<Vec<usize>> = used.into_iter().collect();104    ordered.sort();105    let mut corner_node = vec![usize::MAX; corner_shape.iter().product()];106    let mut network = Network::new(dim);107    for corner in &ordered {108        let position: Vec<f64> = corner.iter().rev().map(|&c| c as f64).collect();109        corner_node[flat_corner(corner)] = network.add_node(position)?;110    }111    let mut seen: HashSet<(usize, usize)> = HashSet::new();112    for coord in &coords {113        for (a, b) in cell_edges(coord, dim) {114            let ia = corner_node[flat_corner(&a)];115            let ib = corner_node[flat_corner(&b)];116            let key = (ia.min(ib), ia.max(ib));117            if seen.insert(key) {118                network.add_branch(ia, ib, 1.0)?;119            }120        }121    }122    Ok(network)123}124125fn cell_edges(coord: &[usize], dim: usize) -> Vec<(Vec<usize>, Vec<usize>)> {126    let mut pairs = Vec::new();127    for fixed_axis in 0..dim {128        for combo in 0..(1usize << (dim - 1)) {129            let mut lo = Vec::with_capacity(dim);130            let mut hi = Vec::with_capacity(dim);131            let mut bit = 0;132            for (axis, &c) in coord.iter().enumerate() {133                if axis == fixed_axis {134                    lo.push(c);135                    hi.push(c + 1);136                } else {137                    let v = (combo >> bit) & 1;138                    bit += 1;139                    lo.push(c + v);140                    hi.push(c + v);141                }142            }143            pairs.push((lo, hi));144        }145    }146    pairs147}148149#[cfg(test)]150mod tests {151    use super::*;152    use crate::math::atoms;153    #[test]154    fn single_square_graphs() {155        let g = atoms::ones_2d(1);156        let core = core_graph(&g).unwrap();157        assert_eq!(core.nodes.len(), 1);158        assert_eq!(core.branches.len(), 0);159        let edges = edge_graph(&g).unwrap();160        assert_eq!(edges.nodes.len(), 4);161        assert_eq!(edges.branches.len(), 4);162    }163    #[test]164    fn carpet_core_graph_counts() {165        let g = atoms::carpet_2d(3);166        let core = core_graph(&g).unwrap();167        assert_eq!(core.nodes.len(), 8);168        assert_eq!(core.branches.len(), 8);169        let tunnels = tunnel_graph(&g).unwrap();170        assert_eq!(tunnels.nodes.len(), 1);171    }172    #[test]173    fn cube_edge_graph() {174        let g = atoms::ones_3d(1);175        let edges = edge_graph(&g).unwrap();176        assert_eq!(edges.nodes.len(), 8);177        assert_eq!(edges.branches.len(), 12);178    }179}