extract.rs

5.4 kB · rust · 167 lines

1use super::models::Network;2use mrlycore::errors::{value_error, Result};3use mrlycore::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.bytes()[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, or an error off 2D and 3D.28pub fn core_graph(grid: &Tensor) -> Result<Network> {29    let dim = grid.shape.len();30    if dim != 2 && dim != 3 {31        return value_error("core_graph expects a 2D or 3D cell.");32    }33    let coords = coords_of(grid);34    let mut index_of = vec![usize::MAX; grid.size()];35    for (i, coord) in coords.iter().enumerate() {36        index_of[grid.index(coord)] = i;37    }38    let mut network = Network::new(dim);39    for coord in &coords {40        network.add_node(center(coord))?;41    }42    for coord in &coords {43        for axis in 0..dim {44            if coord[axis] + 1 >= grid.shape[axis] {45                continue;46            }47            let mut neighbor = coord.clone();48            neighbor[axis] += 1;49            if grid.get(&neighbor) != 0 {50                network.add_branch(51                    index_of[grid.index(coord)],52                    index_of[grid.index(&neighbor)],53                    1.0,54                )?;55            }56        }57    }58    Ok(network)59}6061/// Extracts the core graph of the inverted grid, joining empty sites instead.62pub fn tunnel_graph(grid: &Tensor) -> Result<Network> {63    core_graph(&grid.invert())64}6566/// Extracts the network of corners and edges outlining every filled site, or an error off 2D and 3D.67pub fn edge_graph(grid: &Tensor) -> Result<Network> {68    let dim = grid.shape.len();69    if dim != 2 && dim != 3 {70        return value_error("edge_graph expects a 2D or 3D cell.");71    }72    let corner_shape: Vec<usize> = grid.shape.iter().map(|&s| s + 1).collect();73    let corner_index_dims = corner_shape.clone();74    let flat_corner = |coord: &[usize]| -> usize {75        let mut out = 0;76        for (axis, &c) in coord.iter().enumerate() {77            out = out * corner_index_dims[axis] + c;78        }79        out80    };81    let coords = coords_of(grid);82    let mut used: HashSet<Vec<usize>> = HashSet::new();83    for coord in &coords {84        for offset in 0..(1usize << dim) {85            let corner: Vec<usize> = (0..dim)86                .map(|axis| coord[axis] + ((offset >> (dim - 1 - axis)) & 1))87                .collect();88            used.insert(corner);89        }90    }91    let mut ordered: Vec<Vec<usize>> = used.into_iter().collect();92    ordered.sort();93    let mut corner_node = vec![usize::MAX; corner_shape.iter().product()];94    let mut network = Network::new(dim);95    for corner in &ordered {96        let position: Vec<f64> = corner.iter().rev().map(|&c| c as f64).collect();97        corner_node[flat_corner(corner)] = network.add_node(position)?;98    }99    let mut seen: HashSet<(usize, usize)> = HashSet::new();100    for coord in &coords {101        for (a, b) in cell_edges(coord, dim) {102            let ia = corner_node[flat_corner(&a)];103            let ib = corner_node[flat_corner(&b)];104            let key = (ia.min(ib), ia.max(ib));105            if seen.insert(key) {106                network.add_branch(ia, ib, 1.0)?;107            }108        }109    }110    Ok(network)111}112113fn cell_edges(coord: &[usize], dim: usize) -> Vec<(Vec<usize>, Vec<usize>)> {114    let mut pairs = Vec::new();115    for fixed_axis in 0..dim {116        for combo in 0..(1usize << (dim - 1)) {117            let mut lo = Vec::with_capacity(dim);118            let mut hi = Vec::with_capacity(dim);119            let mut bit = 0;120            for (axis, &c) in coord.iter().enumerate() {121                if axis == fixed_axis {122                    lo.push(c);123                    hi.push(c + 1);124                } else {125                    let v = (combo >> bit) & 1;126                    bit += 1;127                    lo.push(c + v);128                    hi.push(c + v);129                }130            }131            pairs.push((lo, hi));132        }133    }134    pairs135}136137#[cfg(test)]138mod tests {139    use super::*;140    use mrlycore::atoms;141    #[test]142    fn single_square_graphs() {143        let g = atoms::ones_2d(1);144        let core = core_graph(&g).unwrap();145        assert_eq!(core.nodes.len(), 1);146        assert_eq!(core.branches.len(), 0);147        let edges = edge_graph(&g).unwrap();148        assert_eq!(edges.nodes.len(), 4);149        assert_eq!(edges.branches.len(), 4);150    }151    #[test]152    fn carpet_core_graph_counts() {153        let g = atoms::carpet_2d(3);154        let core = core_graph(&g).unwrap();155        assert_eq!(core.nodes.len(), 8);156        assert_eq!(core.branches.len(), 8);157        let tunnels = tunnel_graph(&g).unwrap();158        assert_eq!(tunnels.nodes.len(), 1);159    }160    #[test]161    fn cube_edge_graph() {162        let g = atoms::ones_3d(1);163        let edges = edge_graph(&g).unwrap();164        assert_eq!(edges.nodes.len(), 8);165        assert_eq!(edges.branches.len(), 12);166    }167}