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(¢er)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(¢er)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}