use std::collections::BTreeSet;
use rustc_hash::FxHashMap;
use thiserror::Error;
use super::program::Program;
use super::spec_types::{BufferAccess, DataType};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, PartialOrd, Ord)]
pub struct GraphValueId(pub u32);
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, PartialOrd, Ord)]
pub struct GraphNodeId(pub u32);
#[derive(Debug, Clone, PartialEq, Eq, Hash, PartialOrd, Ord)]
pub enum ShapeDim {
Known(u64),
Symbol(String),
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, PartialOrd, Ord)]
pub enum ValueLifetime {
Constant,
Invocation,
Retained,
Output,
}
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
pub struct ValueContract {
pub dtype: DataType,
pub shape: Vec<ShapeDim>,
pub access: BufferAccess,
pub lifetime: ValueLifetime,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct GraphInput {
pub buffer: String,
pub value: GraphValueId,
pub contract: ValueContract,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct GraphOutput {
pub buffer: String,
pub name: String,
pub contract: ValueContract,
pub retained_successor_of: Option<GraphValueId>,
}
#[derive(Debug, Clone)]
pub struct ProgramGraphNode {
pub id: GraphNodeId,
pub name: String,
pub program: Program,
pub inputs: Vec<GraphInput>,
pub outputs: Vec<GraphValueId>,
pub output_ports: Vec<GraphOutput>,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct ProgramGraphValue {
pub id: GraphValueId,
pub name: String,
pub contract: ValueContract,
pub producer: Option<GraphNodeId>,
pub consumers: Vec<GraphNodeId>,
pub retained_successor_of: Option<GraphValueId>,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct LivenessInterval {
pub value: GraphValueId,
pub start: usize,
pub end: usize,
}
#[derive(Debug, Clone, PartialEq, Eq, Error)]
pub enum ProgramGraphError {
#[error("duplicate graph name `{0}`; use one stable identity per node or value")]
DuplicateName(String),
#[error("graph value {0:?} does not exist")]
MissingValue(GraphValueId),
#[error("program node `{node}` has no buffer `{buffer}`")]
MissingBuffer {
node: String,
buffer: String,
},
#[error("program node `{node}` buffer `{buffer}` disagrees with its value contract: {reason}")]
BufferContract {
node: String,
buffer: String,
reason: String,
},
#[error(
"program node `{node}` buffer `{buffer}` expects {expected:?}, but graph value {value:?} provides {actual:?}"
)]
InputContract {
node: String,
buffer: String,
value: GraphValueId,
actual: ValueContract,
expected: ValueContract,
},
#[error("retained output `{output}` is not a type-preserving successor of {prior:?}")]
InvalidRetainedTransition {
output: String,
prior: GraphValueId,
},
#[error("program node `{node}` binds buffer `{buffer}` more than once")]
DuplicatePort {
node: String,
buffer: String,
},
#[error("program node `{node}` binds graph value {value:?} more than once")]
DuplicateValueInput {
node: String,
value: GraphValueId,
},
#[error("retained output `{output}` names {prior:?} without consuming that prior value")]
MissingRetainedInput {
output: String,
prior: GraphValueId,
},
#[error("ProgramGraph has more than {0} addressable values or nodes")]
IdentityOverflow(u32),
#[error("invalid ProgramGraph wire data: {0}")]
Wire(String),
}
#[derive(Debug, Default)]
pub struct ProgramGraph {
nodes: Vec<ProgramGraphNode>,
values: Vec<ProgramGraphValue>,
names: FxHashMap<String, ()>,
}
impl ProgramGraph {
#[must_use]
pub fn new() -> Self {
Self::default()
}
pub fn from_program(
node_name: impl Into<String>,
program: Program,
) -> Result<Self, ProgramGraphError> {
let mut graph = Self::new();
let mut inputs = Vec::new();
let mut outputs = Vec::new();
for buffer in program.buffers() {
if buffer.access() == BufferAccess::Workgroup {
continue;
}
let contract = ValueContract {
dtype: buffer.element(),
shape: vec![ShapeDim::Known(u64::from(buffer.count()))],
access: buffer.access(),
lifetime: if buffer.is_output() || buffer.access() == BufferAccess::WriteOnly {
ValueLifetime::Output
} else if buffer.access() == BufferAccess::ReadWrite {
ValueLifetime::Retained
} else {
ValueLifetime::Invocation
},
};
if contract.lifetime == ValueLifetime::Output {
outputs.push(GraphOutput {
buffer: buffer.name().to_string(),
name: buffer.name().to_string(),
contract,
retained_successor_of: None,
});
} else {
let value = graph.add_external_value(buffer.name(), contract.clone())?;
inputs.push(GraphInput {
buffer: buffer.name().to_string(),
value,
contract,
});
}
}
graph.add_node(node_name, program, inputs, outputs)?;
Ok(graph)
}
pub fn add_external_value(
&mut self,
name: impl Into<String>,
contract: ValueContract,
) -> Result<GraphValueId, ProgramGraphError> {
self.push_value(name.into(), contract, None, None)
}
pub fn add_external_values(
&mut self,
values: Vec<(String, ValueContract)>,
) -> Result<Vec<GraphValueId>, ProgramGraphError> {
let mut batch_names = BTreeSet::new();
let mut ids = Vec::with_capacity(values.len());
for (offset, (name, _)) in values.iter().enumerate() {
self.ensure_name_available(name)?;
if !batch_names.insert(name.as_str()) {
return Err(ProgramGraphError::DuplicateName(name.clone()));
}
let index = self
.values
.len()
.checked_add(offset)
.ok_or(ProgramGraphError::IdentityOverflow(u32::MAX))?;
ids.push(GraphValueId(
u32::try_from(index).map_err(|_| ProgramGraphError::IdentityOverflow(u32::MAX))?,
));
}
for ((name, contract), id) in values.into_iter().zip(ids.iter().copied()) {
self.names.insert(name.clone(), ());
self.values.push(ProgramGraphValue {
id,
name,
contract,
producer: None,
consumers: Vec::new(),
retained_successor_of: None,
});
}
Ok(ids)
}
pub fn add_node(
&mut self,
name: impl Into<String>,
program: Program,
inputs: Vec<GraphInput>,
outputs: Vec<GraphOutput>,
) -> Result<(GraphNodeId, Vec<GraphValueId>), ProgramGraphError> {
let name = name.into();
self.ensure_name_available(&name)?;
let node_id = GraphNodeId(
u32::try_from(self.nodes.len())
.map_err(|_| ProgramGraphError::IdentityOverflow(u32::MAX))?,
);
let mut new_names = BTreeSet::new();
new_names.insert(name.as_str());
let mut output_ids = Vec::with_capacity(outputs.len());
for (offset, output) in outputs.iter().enumerate() {
self.ensure_name_available(&output.name)?;
if !new_names.insert(output.name.as_str()) {
return Err(ProgramGraphError::DuplicateName(output.name.clone()));
}
let index = self
.values
.len()
.checked_add(offset)
.ok_or(ProgramGraphError::IdentityOverflow(u32::MAX))?;
output_ids.push(GraphValueId(
u32::try_from(index).map_err(|_| ProgramGraphError::IdentityOverflow(u32::MAX))?,
));
}
let mut bound = BTreeSet::new();
let mut bound_values = BTreeSet::new();
for input in &inputs {
if !bound.insert(input.buffer.as_str()) {
return Err(ProgramGraphError::DuplicatePort {
node: name,
buffer: input.buffer.clone(),
});
}
if !bound_values.insert(input.value) {
return Err(ProgramGraphError::DuplicateValueInput {
node: name,
value: input.value,
});
}
let value = self
.values
.get(input.value.0 as usize)
.ok_or(ProgramGraphError::MissingValue(input.value))?;
if value.contract.dtype != input.contract.dtype
|| value.contract.shape != input.contract.shape
|| value.contract.lifetime != input.contract.lifetime
{
return Err(ProgramGraphError::InputContract {
node: name,
buffer: input.buffer.clone(),
value: input.value,
actual: value.contract.clone(),
expected: input.contract.clone(),
});
}
validate_buffer(
&name,
&program,
&input.buffer,
&input.contract,
PortRole::Input,
)?;
}
for output in &outputs {
if let Some(prior_id) = output.retained_successor_of {
let prior = self
.values
.get(prior_id.0 as usize)
.ok_or(ProgramGraphError::MissingValue(prior_id))?;
if !inputs.iter().any(|input| input.value == prior_id) {
return Err(ProgramGraphError::MissingRetainedInput {
output: output.name.clone(),
prior: prior_id,
});
}
if prior.contract.lifetime != ValueLifetime::Retained
|| output.contract.lifetime != ValueLifetime::Retained
|| prior.contract != output.contract
{
return Err(ProgramGraphError::InvalidRetainedTransition {
output: output.name.clone(),
prior: prior_id,
});
}
}
let retained_rebind = output.retained_successor_of.is_some_and(|prior| {
inputs
.iter()
.any(|input| input.buffer == output.buffer && input.value == prior)
});
if !bound.insert(output.buffer.as_str()) && !retained_rebind {
return Err(ProgramGraphError::DuplicatePort {
node: name,
buffer: output.buffer.clone(),
});
}
validate_buffer(
&name,
&program,
&output.buffer,
&output.contract,
PortRole::Output,
)?;
}
self.names.insert(name.clone(), ());
for output in &outputs {
self.names.insert(output.name.clone(), ());
}
let mut consumed = BTreeSet::new();
for input in &inputs {
if consumed.insert(input.value) {
self.values[input.value.0 as usize].consumers.push(node_id);
}
}
let output_ports = outputs.clone();
for (output, id) in outputs.into_iter().zip(output_ids.iter().copied()) {
self.values.push(ProgramGraphValue {
id,
name: output.name,
contract: output.contract,
producer: Some(node_id),
consumers: Vec::new(),
retained_successor_of: output.retained_successor_of,
});
}
self.nodes.push(ProgramGraphNode {
id: node_id,
name,
program,
inputs,
outputs: output_ids.clone(),
output_ports,
});
Ok((node_id, output_ids))
}
#[must_use]
pub fn nodes(&self) -> &[ProgramGraphNode] {
&self.nodes
}
#[must_use]
pub fn values(&self) -> &[ProgramGraphValue] {
&self.values
}
#[must_use]
pub fn schedule(&self) -> Vec<GraphNodeId> {
self.nodes.iter().map(|node| node.id).collect()
}
#[must_use]
pub fn liveness_intervals(&self) -> Vec<LivenessInterval> {
self.values
.iter()
.map(|value| {
let start = value.producer.map_or(0, |producer| producer.0 as usize);
let end = value
.consumers
.iter()
.map(|consumer| consumer.0 as usize)
.max()
.unwrap_or(start);
LivenessInterval {
value: value.id,
start,
end,
}
})
.collect()
}
fn push_value(
&mut self,
name: String,
contract: ValueContract,
producer: Option<GraphNodeId>,
retained_successor_of: Option<GraphValueId>,
) -> Result<GraphValueId, ProgramGraphError> {
self.ensure_name_available(&name)?;
let id = GraphValueId(
u32::try_from(self.values.len())
.map_err(|_| ProgramGraphError::IdentityOverflow(u32::MAX))?,
);
self.names.insert(name.clone(), ());
self.values.push(ProgramGraphValue {
id,
name,
contract,
producer,
consumers: Vec::new(),
retained_successor_of,
});
Ok(id)
}
fn ensure_name_available(&self, name: &str) -> Result<(), ProgramGraphError> {
if self.names.contains_key(name) {
return Err(ProgramGraphError::DuplicateName(name.to_string()));
}
Ok(())
}
}
#[derive(Debug, Clone, Copy)]
enum PortRole {
Input,
Output,
}
fn validate_buffer(
node: &str,
program: &Program,
buffer_name: &str,
contract: &ValueContract,
role: PortRole,
) -> Result<(), ProgramGraphError> {
let buffer = program
.buffers()
.iter()
.find(|buffer| buffer.name() == buffer_name)
.ok_or_else(|| ProgramGraphError::MissingBuffer {
node: node.to_string(),
buffer: buffer_name.to_string(),
})?;
if buffer.element() != contract.dtype {
return Err(ProgramGraphError::BufferContract {
node: node.to_string(),
buffer: buffer_name.to_string(),
reason: format!(
"Program uses {:?}, graph uses {:?}",
buffer.element(),
contract.dtype
),
});
}
if let Some(elements) = static_element_count(&contract.shape).map_err(|reason| {
ProgramGraphError::BufferContract {
node: node.to_string(),
buffer: buffer_name.to_string(),
reason,
}
})? {
if buffer.count() != 0 && elements != u64::from(buffer.count()) {
return Err(ProgramGraphError::BufferContract {
node: node.to_string(),
buffer: buffer_name.to_string(),
reason: format!(
"Program declares {} elements, graph shape requires {elements}",
buffer.count()
),
});
}
}
let access_satisfies_contract = match contract.access {
BufferAccess::ReadOnly => matches!(
buffer.access(),
BufferAccess::ReadOnly | BufferAccess::ReadWrite | BufferAccess::Uniform
),
BufferAccess::ReadWrite => buffer.access() == BufferAccess::ReadWrite,
BufferAccess::WriteOnly => {
matches!(
buffer.access(),
BufferAccess::WriteOnly | BufferAccess::ReadWrite
)
}
BufferAccess::Uniform => buffer.access() == BufferAccess::Uniform,
_ => false,
};
if !access_satisfies_contract {
return Err(ProgramGraphError::BufferContract {
node: node.to_string(),
buffer: buffer_name.to_string(),
reason: format!(
"Program access {:?} does not satisfy graph access {:?}",
buffer.access(),
contract.access
),
});
}
let access = buffer.access();
let compatible = match role {
PortRole::Input => matches!(
access,
BufferAccess::ReadOnly | BufferAccess::ReadWrite | BufferAccess::Uniform
),
PortRole::Output => matches!(access, BufferAccess::ReadWrite | BufferAccess::WriteOnly),
};
if !compatible {
return Err(ProgramGraphError::BufferContract {
node: node.to_string(),
buffer: buffer_name.to_string(),
reason: format!("{role:?} port cannot use {access:?} access"),
});
}
Ok(())
}
fn static_element_count(shape: &[ShapeDim]) -> Result<Option<u64>, String> {
let mut elements = 1_u64;
for dimension in shape {
match dimension {
ShapeDim::Known(extent) => {
elements = elements.checked_mul(*extent).ok_or_else(|| {
"graph shape element count overflows u64; reduce or shard dimensions"
.to_string()
})?;
}
ShapeDim::Symbol(_) => return Ok(None),
}
}
Ok(Some(elements))
}