use smallvec::SmallVec;
use crate::graph::Structure;
use crate::op::Op;
use crate::{Element, Shape, Tensor};
pub(crate) struct View<'plan, Data> {
structure: &'plan Structure<Data>,
wanted: &'plan [bool],
readable: &'plan [bool],
consumers: Vec<usize>,
consumer_of: Vec<SmallVec<[usize; 2]>>,
}
impl<'plan, E: Element> View<'plan, Tensor<E>> {
pub(crate) fn new(
structure: &'plan Structure<Tensor<E>>,
wanted: &'plan [bool],
readable: &'plan [bool],
) -> Self {
let length = structure.len();
let mut consumers = vec![0usize; length];
let mut consumer_of: Vec<SmallVec<[usize; 2]>> = vec![SmallVec::new(); length];
for (index, &wanted_node) in wanted.iter().enumerate() {
if !wanted_node {
continue;
}
let links = structure
.operands
.get(index)
.expect("plan columns are fixed");
for link in links.as_slice() {
consumers[link.index()] += 1;
consumer_of[link.index()].push(index);
}
}
Self {
structure,
wanted,
readable,
consumers,
consumer_of,
}
}
pub(crate) fn len(&self) -> usize {
self.structure.len()
}
pub(crate) fn wanted(&self, index: usize) -> bool {
self.wanted[index]
}
pub(crate) fn interior_ok(&self, index: usize) -> bool {
self.wanted[index] && !self.readable[index] && self.consumers[index] == 1
}
pub(crate) fn closed(&self, root: usize, interiors: &[usize], named: &[usize]) -> bool {
let mut member = vec![false; self.len()];
member[root] = true;
for &node in interiors {
if !self.wanted[node] || self.readable[node] {
return false;
}
member[node] = true;
}
for &node in named {
if !self.wanted[node] || member[node] {
return false;
}
member[node] = true;
}
for &node in interiors.iter().chain(named) {
for &consumer in &self.consumer_of[node] {
if !member[consumer] {
return false;
}
}
}
true
}
pub(crate) fn op(&self, index: usize) -> Option<&'plan Op<Tensor<E>>> {
self.structure.ops.get(index)
}
pub(crate) fn shape(&self, index: usize) -> &'plan Shape {
&self.structure.shapes[index]
}
pub(crate) fn operand(&self, index: usize, position: usize) -> usize {
self.structure
.operands
.get(index)
.expect("plan columns are fixed")
.as_slice()[position]
.index()
}
pub(crate) fn sole_operand(&self, index: usize) -> usize {
self.operand(index, 0)
}
}