tensor.rs

23.9 kB · rust · 700 lines

1/// The element widths a tensor can hold.2#[derive(Clone, Copy, Debug, PartialEq, Eq)]3pub enum Dtype {4    /// Unsigned 8-bit elements.5    U8,6    /// Unsigned 16-bit elements.7    U16,8    /// Unsigned 32-bit elements.9    U32,10    /// Signed 32-bit elements.11    I32,12}1314impl Dtype {15    /// Returns the largest value the width can hold.16    pub fn max(self) -> i64 {17        match self {18            Dtype::U8 => u8::MAX as i64,19            Dtype::U16 => u16::MAX as i64,20            Dtype::U32 => u32::MAX as i64,21            Dtype::I32 => i32::MAX as i64,22        }23    }24}2526/// The typed storage behind a tensor.27#[derive(Clone, Debug, PartialEq, Eq)]28pub enum Buf {29    /// Unsigned 8-bit storage.30    U8(Vec<u8>),31    /// Unsigned 16-bit storage.32    U16(Vec<u16>),33    /// Unsigned 32-bit storage.34    U32(Vec<u32>),35    /// Signed 32-bit storage.36    I32(Vec<i32>),37}3839impl Buf {40    fn zeros(dtype: Dtype, n: usize) -> Buf {41        Buf::filled(dtype, n, 0)42    }43    fn filled(dtype: Dtype, n: usize, value: i64) -> Buf {44        match dtype {45            Dtype::U8 => Buf::U8(vec![value as u8; n]),46            Dtype::U16 => Buf::U16(vec![value as u16; n]),47            Dtype::U32 => Buf::U32(vec![value as u32; n]),48            Dtype::I32 => Buf::I32(vec![value as i32; n]),49        }50    }51    fn dtype(&self) -> Dtype {52        match self {53            Buf::U8(_) => Dtype::U8,54            Buf::U16(_) => Dtype::U16,55            Buf::U32(_) => Dtype::U32,56            Buf::I32(_) => Dtype::I32,57        }58    }59    fn len(&self) -> usize {60        match self {61            Buf::U8(v) => v.len(),62            Buf::U16(v) => v.len(),63            Buf::U32(v) => v.len(),64            Buf::I32(v) => v.len(),65        }66    }67    fn at(&self, i: usize) -> i64 {68        match self {69            Buf::U8(v) => v[i] as i64,70            Buf::U16(v) => v[i] as i64,71            Buf::U32(v) => v[i] as i64,72            Buf::I32(v) => v[i] as i64,73        }74    }75    fn put(&mut self, i: usize, value: i64) {76        match self {77            Buf::U8(v) => v[i] = value as u8,78            Buf::U16(v) => v[i] = value as u16,79            Buf::U32(v) => v[i] = value as u32,80            Buf::I32(v) => v[i] = value as i32,81        }82    }83    fn gather(&self, idx: &[usize]) -> Buf {84        match self {85            Buf::U8(v) => Buf::U8(idx.iter().map(|&i| v[i]).collect()),86            Buf::U16(v) => Buf::U16(idx.iter().map(|&i| v[i]).collect()),87            Buf::U32(v) => Buf::U32(idx.iter().map(|&i| v[i]).collect()),88            Buf::I32(v) => Buf::I32(idx.iter().map(|&i| v[i]).collect()),89        }90    }91}9293/// An n-dimensional grid of small integers.94#[derive(Clone, Debug, PartialEq, Eq)]95pub struct Tensor {96    data: Buf,97    /// The extent of each axis.98    pub shape: Vec<usize>,99}100101fn strides(shape: &[usize]) -> Vec<usize> {102    let mut out = vec![1; shape.len()];103    for axis in (0..shape.len().saturating_sub(1)).rev() {104        out[axis] = out[axis + 1] * shape[axis + 1];105    }106    out107}108109fn unravel(flat: usize, shape: &[usize]) -> Vec<usize> {110    let s = strides(shape);111    let mut rem = flat;112    s.iter()113        .map(|&st| {114            let i = rem / st;115            rem %= st;116            i117        })118        .collect()119}120121impl Tensor {122    /// Builds a zeroed u8 tensor of the shape.123    pub fn new(shape: Vec<usize>) -> Self {124        Tensor::typed(shape, Dtype::U8)125    }126    /// Builds a u8 tensor filled with one value.127    pub fn full(shape: Vec<usize>, value: u8) -> Self {128        Tensor::filled(shape, value as i64, Dtype::U8)129    }130    /// Builds a zeroed tensor of the shape and width.131    pub fn typed(shape: Vec<usize>, dtype: Dtype) -> Self {132        let size = shape.iter().product();133        Tensor {134            data: Buf::zeros(dtype, size),135            shape,136        }137    }138    /// Builds a tensor of the shape and width filled with one value.139    pub fn filled(shape: Vec<usize>, value: i64, dtype: Dtype) -> Self {140        let size = shape.iter().product();141        Tensor {142            data: Buf::filled(dtype, size, value),143            shape,144        }145    }146    /// Wraps a byte vector as a tensor of the shape.147    pub fn of(data: Vec<u8>, shape: Vec<usize>) -> Self {148        Tensor {149            data: Buf::U8(data),150            shape,151        }152    }153    /// Returns the element width.154    pub fn dtype(&self) -> Dtype {155        self.data.dtype()156    }157    /// Returns the number of elements.158    pub fn size(&self) -> usize {159        self.data.len()160    }161    /// Returns the elements as bytes, panicking for wider tensors.162    pub fn bytes(&self) -> &[u8] {163        match &self.data {164            Buf::U8(v) => v,165            other => panic!("tensor is {:?}, not u8", other.dtype()),166        }167    }168    /// Returns the elements as mutable bytes, panicking for wider tensors.169    pub fn bytes_mut(&mut self) -> &mut [u8] {170        match &mut self.data {171            Buf::U8(v) => v,172            other => panic!("tensor is {:?}, not u8", other.dtype()),173        }174    }175    /// Returns the element at a flat index.176    pub fn at(&self, flat: usize) -> i64 {177        self.data.at(flat)178    }179    /// Writes the element at a flat index.180    pub fn put(&mut self, flat: usize, value: i64) {181        self.data.put(flat, value);182    }183    /// Returns the sum of all elements.184    pub fn sum(&self) -> u64 {185        (0..self.size()).map(|i| self.data.at(i) as u64).sum()186    }187    /// Folds a multi-index into its flat index.188    pub fn index(&self, multi: &[usize]) -> usize {189        let s = strides(&self.shape);190        multi.iter().zip(&s).map(|(m, st)| m * st).sum()191    }192    /// Returns the byte at a multi-index.193    pub fn get(&self, multi: &[usize]) -> u8 {194        self.bytes()[self.index(multi)]195    }196    /// Writes the byte at a multi-index.197    pub fn set(&mut self, multi: &[usize], value: u8) {198        let i = self.index(multi);199        self.bytes_mut()[i] = value;200    }201    /// Builds the Kronecker product of the two tensors.202    pub fn kron(&self, other: &Tensor) -> Tensor {203        let shape: Vec<usize> = self204            .shape205            .iter()206            .zip(&other.shape)207            .map(|(a, b)| a * b)208            .collect();209        let mut out = Tensor::typed(shape.clone(), self.dtype());210        let sa = strides(&self.shape);211        let sb = strides(&other.shape);212        let so = strides(&shape);213        for flat in 0..out.size() {214            let mut ai = 0;215            let mut bi = 0;216            let mut rem = flat;217            for axis in 0..shape.len() {218                let idx = rem / so[axis];219                rem %= so[axis];220                ai += (idx / other.shape[axis]) * sa[axis];221                bi += (idx % other.shape[axis]) * sb[axis];222            }223            out.data.put(flat, self.data.at(ai) * other.data.at(bi));224        }225        out226    }227    /// Folds the tensor into its level-fold Kronecker power.228    ///229    /// ```230    /// use mrlycore::tensor::Tensor;231    /// let f = Tensor::of(vec![1, 1, 0, 1], vec![2, 2]).fractal(2);232    /// assert_eq!(f.shape, vec![4, 4]);233    /// assert_eq!(f.sum(), 9);234    /// ```235    pub fn fractal(&self, level: usize) -> Tensor {236        let mut out = self.clone();237        for _ in 1..level {238            out = out.kron(self);239        }240        out241    }242    fn multi(&self, flat: usize) -> Vec<usize> {243        unravel(flat, &self.shape)244    }245    fn remap(&self, shape: Vec<usize>, map: impl Fn(&[usize]) -> Vec<usize>) -> Tensor {246        let size: usize = shape.iter().product();247        let indices: Vec<usize> = (0..size)248            .map(|flat| self.index(&map(&unravel(flat, &shape))))249            .collect();250        Tensor {251            data: self.data.gather(&indices),252            shape,253        }254    }255    /// Flips every element between zero and one.256    pub fn invert(&self) -> Tensor {257        let mut out = Tensor::typed(self.shape.clone(), self.dtype());258        for i in 0..self.size() {259            out.data.put(i, 1 - self.data.at(i));260        }261        out262    }263    /// Reverses the tensor along one axis.264    pub fn flip(&self, axis: usize) -> Tensor {265        let n = self.shape[axis];266        self.remap(self.shape.clone(), |idx| {267            let mut src = idx.to_vec();268            src[axis] = n - 1 - src[axis];269            src270        })271    }272    /// Swaps two axes.273    pub fn transpose(&self, a: usize, b: usize) -> Tensor {274        let mut shape = self.shape.clone();275        shape.swap(a, b);276        self.remap(shape, |idx| {277            let mut src = idx.to_vec();278            src.swap(a, b);279            src280        })281    }282    /// Drops one axis by fixing it at an index.283    ///284    /// ```285    /// let sponge = mrlycore::atoms::carpet_3d(3);286    /// assert_eq!(sponge.slice(2, 0).unwrap(), mrlycore::atoms::carpet_2d(3));287    /// ```288    pub fn slice(&self, axis: usize, index: usize) -> crate::Result<Tensor> {289        use crate::errors::value_error;290        if axis >= self.shape.len() {291            return value_error("slice axis is past the tensor's rank.");292        }293        if index >= self.shape[axis] {294            return value_error("slice index is past the axis.");295        }296        let mut shape = self.shape.clone();297        shape.remove(axis);298        Ok(self.remap(shape, |idx| {299            let mut src = idx.to_vec();300            src.insert(axis, index);301            src302        }))303    }304    /// Rotates the tensor k quarter turns in the plane of two axes.305    pub fn rot90(&self, k: usize, axes: (usize, usize)) -> Tensor {306        let mut out = self.clone();307        for _ in 0..k % 4 {308            out = out.transpose(axes.0, axes.1).flip(axes.0);309        }310        out311    }312    /// Wraps the tensor in a count-thick border of one value.313    pub fn pad(&self, count: usize, value: u8) -> Tensor {314        let shape: Vec<usize> = self.shape.iter().map(|&n| n + 2 * count).collect();315        let mut out = Tensor::filled(shape, value as i64, self.dtype());316        let so = strides(&out.shape);317        for flat in 0..self.size() {318            let idx = self.multi(flat);319            let target: usize = idx.iter().zip(&so).map(|(i, st)| (i + count) * st).sum();320            out.data.put(target, self.data.at(flat));321        }322        out323    }324    /// Repeats the tensor the given number of times along each axis.325    pub fn tile(&self, reps: &[usize]) -> Tensor {326        let shape: Vec<usize> = self.shape.iter().zip(reps).map(|(n, r)| n * r).collect();327        let inner = self.shape.clone();328        self.remap(shape, |idx| {329            idx.iter().zip(&inner).map(|(i, n)| i % n).collect()330        })331    }332    /// Numbers every position by its concentric ring out from the center.333    pub fn layers(&self, dtype: Dtype) -> Tensor {334        let shape = self.shape.clone();335        let mut out = Tensor::typed(shape.clone(), dtype);336        for flat in 0..out.size() {337            let idx = unravel(flat, &shape);338            let ring = idx339                .iter()340                .zip(&shape)341                .map(|(&i, &n)| {342                    let center = (n as f64 - 1.0) / 2.0;343                    (i as f64 - center).abs().floor() as i64344                })345                .max()346                .unwrap_or(0);347            out.data.put(flat, ring);348        }349        out350    }351    /// Counts each position's masked neighbors holding the target bit, or an error when the mask or count does not fit.352    pub fn neighbors(353        &self,354        mask: &Tensor,355        target: u8,356        wrap: bool,357        dtype: Dtype,358    ) -> crate::Result<Tensor> {359        use crate::errors::value_error;360        if mask.shape.len() != self.shape.len() {361            return value_error("mask must have the same number of dimensions.");362        }363        if mask.shape.iter().any(|n| n.is_multiple_of(2)) {364            return value_error("Neighborhood (mask) dimensions must be odd.");365        }366        if target > 1 {367            return value_error("Bit to count (target) must be 0 or 1.");368        }369        let center: Vec<usize> = mask.shape.iter().map(|&n| n / 2).collect();370        let mut offsets = Vec::new();371        for flat in 0..mask.size() {372            if mask.data.at(flat) == 1 {373                let idx = mask.multi(flat);374                offsets.push(375                    idx.iter()376                        .zip(&center)377                        .map(|(&i, &c)| i as isize - c as isize)378                        .collect::<Vec<isize>>(),379                );380            }381        }382        let mut out = Tensor::typed(self.shape.clone(), dtype);383        for flat in 0..self.size() {384            let idx = self.multi(flat);385            let mut count: u32 = 0;386            for offset in &offsets {387                let mut source = Vec::with_capacity(idx.len());388                let mut inside = true;389                for axis in 0..idx.len() {390                    let n = self.shape[axis] as isize;391                    let mut p = idx[axis] as isize + offset[axis];392                    if wrap {393                        p = p.rem_euclid(n);394                    } else if p < 0 || p >= n {395                        inside = false;396                        break;397                    }398                    source.push(p as usize);399                }400                if inside && self.get(&source) == target {401                    count += 1;402                }403            }404            if count as i64 > dtype.max() {405                return value_error(406                    "neighbor count exceeds dtype range; widen the neighbors dtype.",407                );408            }409            out.data.put(flat, count as i64);410        }411        Ok(out)412    }413    /// Maps every element to one at or above the threshold, zero below.414    pub fn binarize(&self, threshold: u8) -> Tensor {415        let mut out = Tensor::new(self.shape.clone());416        for i in 0..self.size() {417            out.data418                .put(i, i64::from(self.data.at(i) >= threshold as i64));419        }420        out421    }422    /// Returns the Otsu threshold splitting the histogram at greatest variance.423    pub fn otsu_threshold(&self) -> u8 {424        let mut hist = [0u64; 256];425        for i in 0..self.size() {426            hist[self.data.at(i).clamp(0, 255) as usize] += 1;427        }428        let total: u64 = hist.iter().sum();429        if total == 0 {430            return 0;431        }432        let sum_all: f64 = hist433            .iter()434            .enumerate()435            .map(|(v, &c)| v as f64 * c as f64)436            .sum();437        let mut sum_below = 0.0;438        let mut weight_below = 0u64;439        let mut best_variance = -1.0;440        let mut threshold = 0u8;441        for (level, &count) in hist.iter().enumerate() {442            weight_below += count;443            if weight_below == 0 {444                continue;445            }446            let weight_above = total - weight_below;447            if weight_above == 0 {448                break;449            }450            sum_below += level as f64 * count as f64;451            let mean_below = sum_below / weight_below as f64;452            let mean_above = (sum_all - sum_below) / weight_above as f64;453            let variance =454                weight_below as f64 * weight_above as f64 * (mean_below - mean_above).powi(2);455            if variance > best_variance {456                best_variance = variance;457                threshold = level as u8;458            }459        }460        threshold461    }462    /// Binarizes at one above the Otsu threshold.463    pub fn binarize_otsu(&self) -> Tensor {464        self.binarize(self.otsu_threshold().saturating_add(1))465    }466    /// Averages every position over its masked neighborhood, rounded.467    pub fn blur(&self, mask: &Tensor, wrap: bool) -> crate::Result<Tensor> {468        use crate::errors::value_error;469        if mask.shape.len() != self.shape.len() {470            return value_error("mask must have the same number of dimensions.");471        }472        if mask.shape.iter().any(|n| n.is_multiple_of(2)) {473            return value_error("Neighborhood (mask) dimensions must be odd.");474        }475        let center: Vec<usize> = mask.shape.iter().map(|&n| n / 2).collect();476        let mut offsets = Vec::new();477        for flat in 0..mask.size() {478            if mask.data.at(flat) != 0 {479                let idx = mask.multi(flat);480                offsets.push(481                    idx.iter()482                        .zip(&center)483                        .map(|(&i, &c)| i as isize - c as isize)484                        .collect::<Vec<isize>>(),485                );486            }487        }488        if offsets.is_empty() {489            return value_error("blur mask must have at least one nonzero cell.");490        }491        let mut out = Tensor::typed(self.shape.clone(), self.dtype());492        for flat in 0..self.size() {493            let idx = self.multi(flat);494            let mut sum: i64 = 0;495            let mut count: i64 = 0;496            for offset in &offsets {497                let mut source = Vec::with_capacity(idx.len());498                let mut inside = true;499                for axis in 0..idx.len() {500                    let n = self.shape[axis] as isize;501                    let mut p = idx[axis] as isize + offset[axis];502                    if wrap {503                        p = p.rem_euclid(n);504                    } else if p < 0 || p >= n {505                        inside = false;506                        break;507                    }508                    source.push(p as usize);509                }510                if inside {511                    sum += self.data.at(self.index(&source));512                    count += 1;513                }514            }515            let value = if count > 0 {516                (sum as f64 / count as f64).round() as i64517            } else {518                self.data.at(flat)519            };520            out.data.put(flat, value);521        }522        Ok(out)523    }524    /// Stamps the value wherever the tiled mask is nonzero.525    pub fn perforate(&self, mask: &Tensor, value: u8) -> crate::Result<Tensor> {526        use crate::errors::value_error;527        if mask.shape.len() != self.shape.len() {528            return value_error("mask must have the same number of dimensions.");529        }530        for (axis, (&n, &m)) in self.shape.iter().zip(&mask.shape).enumerate() {531            if m == 0 || !n.is_multiple_of(m) {532                return value_error(format!(533                    "mask dimension {axis} must evenly tile the tensor dimension."534                ));535            }536        }537        let mut out = Tensor::typed(self.shape.clone(), self.dtype());538        for flat in 0..self.size() {539            let idx = self.multi(flat);540            let mask_idx: Vec<usize> = idx.iter().zip(&mask.shape).map(|(&i, &m)| i % m).collect();541            let hit = mask.data.at(mask.index(&mask_idx)) != 0;542            let value = if hit {543                value as i64544            } else {545                self.data.at(flat)546            };547            out.data.put(flat, value);548        }549        Ok(out)550    }551}552553#[cfg(test)]554mod tests {555    use super::*;556    #[test]557    fn kron_matches_numpy_semantics() {558        let a = Tensor::of(vec![1, 0, 0, 1], vec![2, 2]);559        let b = Tensor::of(vec![1, 1, 1, 1], vec![2, 2]);560        let k = a.kron(&b);561        assert_eq!(k.shape, vec![4, 4]);562        assert_eq!(k.sum(), 8);563        assert_eq!(k.get(&[0, 0]), 1);564        assert_eq!(k.get(&[0, 2]), 0);565        assert_eq!(k.get(&[3, 3]), 1);566    }567    #[test]568    fn fractal_sum_is_power() {569        let a = Tensor::of(vec![1, 1, 0, 1, 1, 0, 0, 0, 1], vec![3, 3]);570        let f = a.fractal(3);571        assert_eq!(f.shape, vec![27, 27]);572        assert_eq!(f.sum(), a.sum().pow(3));573    }574    #[test]575    fn rot90_matches_numpy() {576        let a = Tensor::of(vec![1, 2, 3, 4], vec![2, 2]);577        assert_eq!(a.rot90(1, (0, 1)).bytes(), &[2, 4, 1, 3]);578        assert_eq!(a.rot90(2, (0, 1)).bytes(), &[4, 3, 2, 1]);579        assert_eq!(a.rot90(4, (0, 1)), a);580    }581    #[test]582    fn pad_and_tile() {583        let a = Tensor::full(vec![2, 2], 1);584        let p = a.pad(1, 0);585        assert_eq!(p.shape, vec![4, 4]);586        assert_eq!(p.sum(), 4);587        assert_eq!(p.get(&[0, 0]), 0);588        assert_eq!(p.get(&[1, 1]), 1);589        let t = a.tile(&[2, 3]);590        assert_eq!(t.shape, vec![4, 6]);591        assert_eq!(t.sum(), 24);592    }593    #[test]594    fn layers_rings() {595        let l = Tensor::new(vec![5, 5]).layers(Dtype::U8);596        assert_eq!(l.get(&[2, 2]), 0);597        assert_eq!(l.get(&[1, 2]), 1);598        assert_eq!(l.get(&[0, 0]), 2);599        assert_eq!(l.get(&[4, 0]), 2);600    }601    #[test]602    fn neighbors_moore() {603        let mut mask = Tensor::full(vec![3, 3], 1);604        mask.set(&[1, 1], 0);605        let ones = Tensor::full(vec![3, 3], 1);606        let n = ones.neighbors(&mask, 1, false, Dtype::U8).unwrap();607        assert_eq!(n.get(&[1, 1]), 8);608        assert_eq!(n.get(&[0, 0]), 3);609        let w = ones.neighbors(&mask, 1, true, Dtype::U8).unwrap();610        assert_eq!(w.get(&[0, 0]), 8);611        assert!(ones612            .neighbors(&Tensor::full(vec![2, 2], 1), 1, false, Dtype::U8)613            .is_err());614    }615    #[test]616    fn neighbors_wide_dtype_survives_large_mask() {617        let mask = Tensor::full(vec![33, 33], 1);618        let grid = Tensor::full(vec![40, 40], 1);619        assert!(grid.neighbors(&mask, 1, true, Dtype::U8).is_err());620        let wide = grid.neighbors(&mask, 1, true, Dtype::U16).unwrap();621        assert_eq!(wide.dtype(), Dtype::U16);622        assert_eq!(wide.at(wide.index(&[20, 20])), 33 * 33);623    }624    #[test]625    fn invert_round_trip() {626        let a = Tensor::of(vec![1, 0, 0, 1], vec![2, 2]);627        assert_eq!(a.invert().invert(), a);628        assert_eq!(a.invert().sum(), 2);629    }630    #[test]631    fn kron_3d() {632        let a = Tensor::full(vec![2, 2, 2], 1);633        let b = Tensor::full(vec![3, 3, 3], 1);634        let k = a.kron(&b);635        assert_eq!(k.shape, vec![6, 6, 6]);636        assert_eq!(k.sum(), 216);637    }638    #[test]639    fn binarize_thresholds_pointwise() {640        let a = Tensor::of(vec![0, 50, 128, 255], vec![2, 2]);641        let b = a.binarize(128);642        assert_eq!(b.bytes(), &[0, 0, 1, 1]);643    }644    #[test]645    fn otsu_splits_bimodal_histogram() {646        let mut data = vec![10u8; 20];647        data.extend(vec![200u8; 20]);648        let a = Tensor::of(data, vec![40, 1]);649        let t = a.otsu_threshold();650        assert!((10..200).contains(&t));651        let b = a.binarize_otsu();652        assert_eq!(b.bytes()[0..20].iter().sum::<u8>(), 0);653        assert_eq!(b.bytes()[20..40].iter().sum::<u8>(), 20);654    }655    #[test]656    fn blur_preserves_mean_under_wrap() {657        let a = Tensor::of((0..25).map(|v| (v * 7) % 251).collect(), vec![5, 5]);658        let mask = Tensor::full(vec![3, 3], 1);659        let b = a.blur(&mask, true).unwrap();660        let mean_a: f64 = a.sum() as f64 / a.size() as f64;661        let mean_b: f64 = b.sum() as f64 / b.size() as f64;662        assert!((mean_a - mean_b).abs() < 1.0);663    }664    #[test]665    fn perforate_zero_mask_is_identity() {666        let a = Tensor::of(vec![1, 2, 3, 4], vec![2, 2]);667        let mask = Tensor::new(vec![2, 2]);668        let p = a.perforate(&mask, 9).unwrap();669        assert_eq!(p, a);670    }671    #[test]672    fn perforate_writes_masked_positions() {673        let a = Tensor::new(vec![4, 4]);674        let mask = Tensor::of(vec![1, 0, 0, 1], vec![2, 2]);675        let p = a.perforate(&mask, 7).unwrap();676        assert_eq!(p.get(&[0, 0]), 7);677        assert_eq!(p.get(&[0, 1]), 0);678        assert_eq!(p.get(&[1, 1]), 7);679        assert_eq!(p.sum(), 8 * 7);680    }681    #[test]682    fn slice_commutes_with_the_fractal() {683        for (n, level) in [(3, 2), (5, 2), (3, 3)] {684            let seed = crate::atoms::xtree_3d(n);685            for axis in 0..3 {686                let deep = seed.fractal(level).slice(axis, 0).unwrap();687                let flat = seed.slice(axis, 0).unwrap().fractal(level);688                assert_eq!(deep, flat);689            }690        }691    }692    #[test]693    fn slice_takes_an_index_and_rejects_a_bad_one() {694        let a = Tensor::of(vec![0, 1, 2, 3, 4, 5], vec![2, 3]);695        assert_eq!(a.slice(0, 1).unwrap(), Tensor::of(vec![3, 4, 5], vec![3]));696        assert_eq!(a.slice(1, 2).unwrap(), Tensor::of(vec![2, 5], vec![2]));697        assert!(a.slice(2, 0).is_err());698        assert!(a.slice(1, 3).is_err());699    }700}