use std::collections::HashSet;
use std::fmt;
use crate::arena::ArenaKey;
use crate::device::DeviceId;
use crate::dtype::DataType;
use crate::layout::TensorLayout;
use crate::node::NodeId;
use crate::shape::Shape;
#[derive(Clone, Copy, PartialEq, Eq, Hash, Debug)]
pub struct ValueId(pub u32);
impl ArenaKey for ValueId {
fn from_raw(raw: u32) -> Self {
ValueId(raw)
}
fn to_raw(self) -> u32 {
self.0
}
}
pub type Usage = (NodeId, u32);
#[derive(Clone, Default, PartialEq, Eq)]
pub struct Consumers {
uses: HashSet<Usage>,
}
impl Consumers {
pub(crate) fn insert(&mut self, node: NodeId, input_index: u32) {
self.uses.insert((node, input_index));
}
pub(crate) fn remove(&mut self, node: NodeId, input_index: u32) -> bool {
self.uses.remove(&(node, input_index))
}
pub(crate) fn contains(&self, node: NodeId, input_index: u32) -> bool {
self.uses.contains(&(node, input_index))
}
pub fn len(&self) -> usize {
self.uses.len()
}
pub fn is_empty(&self) -> bool {
self.uses.is_empty()
}
pub fn uses(&self) -> Vec<Usage> {
let mut uses: Vec<_> = self.uses.iter().copied().collect();
uses.sort_unstable_by_key(|&(node, input_index)| (node.0, input_index));
uses
}
pub fn nodes(&self) -> Vec<NodeId> {
let mut nodes: Vec<_> = self.uses.iter().map(|&(node, _)| node).collect();
nodes.sort_unstable_by_key(|node| node.0);
nodes.dedup();
nodes
}
}
impl fmt::Debug for Consumers {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_list().entries(self.uses()).finish()
}
}
#[derive(Clone, Debug, PartialEq)]
pub struct Value {
pub id: ValueId,
pub name: Option<String>,
pub dtype: DataType,
pub shape: Shape,
pub layout: TensorLayout,
pub device: Option<DeviceId>,
pub producer: Option<NodeId>,
pub consumers: Consumers,
pub is_graph_input: bool,
pub is_graph_output: bool,
}
impl Value {
pub fn new(id: ValueId, dtype: DataType, shape: Shape) -> Self {
Self {
id,
name: None,
dtype,
shape,
layout: TensorLayout::contiguous(),
device: None,
producer: None,
consumers: Consumers::default(),
is_graph_input: false,
is_graph_output: false,
}
}
pub fn is_source(&self) -> bool {
self.producer.is_none()
}
pub fn rank(&self) -> usize {
self.shape.len()
}
}