Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions node-graph/gcore/src/lib.rs
Original file line number Diff line number Diff line change
Expand Up @@ -34,6 +34,7 @@ pub use raster::Color;
pub trait Node<'i, Input: 'i>: 'i {
type Output: 'i;
fn eval<'s: 'i>(&'s self, input: Input) -> Self::Output;
fn reset(self: Pin<&mut Self>) {}
}

#[cfg(feature = "alloc")]
Expand Down
5 changes: 5 additions & 0 deletions node-graph/gstd/src/any.rs
Original file line number Diff line number Diff line change
Expand Up @@ -34,7 +34,12 @@ where
Box::new(self.node.eval(*input))
}
}
fn reset(self: std::pin::Pin<&mut Self>) {
let wrapped_node = unsafe { self.map_unchecked_mut(|e| &mut e.node) };
Node::reset(wrapped_node);
}
}

impl<_I, _O, S0> DynAnyRefNode<_I, _O, S0> {
pub const fn new(node: S0) -> Self {
Self { node, _i: core::marker::PhantomData }
Expand Down
24 changes: 21 additions & 3 deletions node-graph/gstd/src/memo.rs
Original file line number Diff line number Diff line change
Expand Up @@ -2,14 +2,16 @@ use graphene_core::Node;

use std::hash::{Hash, Hasher};
use std::marker::PhantomData;
use std::pin::Pin;
use std::sync::atomic::AtomicBool;
use xxhash_rust::xxh3::Xxh3;

/// Caches the output of a given Node and acts as a proxy
#[derive(Default)]
pub struct CacheNode<T, CachedNode> {
// We have to use an append only data structure to make sure the references
// to the cache entries are always valid
cache: boxcar::Vec<(u64, T)>,
cache: boxcar::Vec<(u64, T, AtomicBool)>,
node: CachedNode,
}
impl<'i, T: 'i, I: 'i + Hash, CachedNode: 'i> Node<'i, I> for CacheNode<T, CachedNode>
Expand All @@ -22,17 +24,25 @@ where
input.hash(&mut hasher);
let hash = hasher.finish();

if let Some((_, cached_value)) = self.cache.iter().find(|(h, _)| *h == hash) {
if let Some((_, cached_value, keep)) = self.cache.iter().find(|(h, _, _)| *h == hash) {
keep.store(true, std::sync::atomic::Ordering::Relaxed);
return cached_value;
} else {
trace!("Cache miss");
let output = self.node.eval(input);
let index = self.cache.push((hash, output));
let index = self.cache.push((hash, output, AtomicBool::new(true)));
return &self.cache[index].1;
}
}

fn reset(mut self: Pin<&mut Self>) {
let old_cache = std::mem::take(&mut self.cache);
self.cache = old_cache.into_iter().filter(|(_, _, keep)| keep.swap(false, std::sync::atomic::Ordering::Relaxed)).collect();
}
}

impl<T, CachedNode> std::marker::Unpin for CacheNode<T, CachedNode> {}

impl<T, CachedNode> CacheNode<T, CachedNode> {
pub fn new(node: CachedNode) -> CacheNode<T, CachedNode> {
CacheNode { cache: boxcar::Vec::new(), node }
Expand Down Expand Up @@ -72,8 +82,16 @@ impl<'i, T: 'i + Hash> Node<'i, Option<T>> for LetNode<T> {
None => &self.cache.iter().last().expect("Let node was not initialized").1,
}
}

fn reset(mut self: Pin<&mut Self>) {
if let Some(last) = std::mem::take(&mut self.cache).into_iter().last() {
self.cache = boxcar::vec![last];
}
}
}

impl<T> std::marker::Unpin for LetNode<T> {}

impl<T> LetNode<T> {
pub fn new() -> LetNode<T> {
LetNode { cache: boxcar::Vec::new() }
Comment on lines 95 to 97

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

This is actually the most relevant part, we can basically replace the box car with a single unsafe cell because this should during the runtime of the graph only ever have one value which can be thrown away at the end

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Nevermind didn't see this was already implemented for the cache node

Expand Down
39 changes: 24 additions & 15 deletions node-graph/interpreted-executor/src/executor.rs
Original file line number Diff line number Diff line change
@@ -1,6 +1,6 @@
use std::collections::HashSet;
use std::collections::{HashMap, HashSet};
use std::error::Error;
use std::{collections::HashMap, sync::Arc};
use std::sync::{Arc, RwLock};

use dyn_any::StaticType;
use graph_craft::document::value::UpcastNode;
Expand Down Expand Up @@ -59,7 +59,7 @@ impl Executor for DynamicExecutor {
pub struct NodeContainer<'n> {
pub node: TypeErasedPinned<'n>,
// the dependencies are only kept to ensure that the nodes are not dropped while still in use
_dependencies: Vec<Arc<NodeContainer<'static>>>,
_dependencies: Vec<Arc<RwLock<NodeContainer<'static>>>>,
}

impl<'a> core::fmt::Debug for NodeContainer<'a> {
Expand All @@ -69,7 +69,7 @@ impl<'a> core::fmt::Debug for NodeContainer<'a> {
}

impl<'a> NodeContainer<'a> {
pub fn new(node: TypeErasedPinned<'a>, _dependencies: Vec<Arc<NodeContainer<'static>>>) -> Self {
pub fn new(node: TypeErasedPinned<'a>, _dependencies: Vec<Arc<RwLock<NodeContainer<'static>>>>) -> Self {
Self { node, _dependencies }
}

Expand All @@ -89,7 +89,7 @@ impl NodeContainer<'static> {

#[derive(Default, Debug, Clone)]
pub struct BorrowTree {
nodes: HashMap<NodeId, Arc<NodeContainer<'static>>>,
nodes: HashMap<NodeId, Arc<RwLock<NodeContainer<'static>>>>,
}

impl BorrowTree {
Expand All @@ -107,36 +107,45 @@ impl BorrowTree {
for (id, node) in proto_network.nodes {
if !self.nodes.contains_key(&id) {
self.push_node(id, node, typing_context)?;
} else {
let Some(node_container) = self.nodes.get_mut(&id) else { continue };
let mut node_container_writer = node_container.write().unwrap();
let node = node_container_writer.node.as_mut();
node.reset();
}
old_nodes.remove(&id);
}
Ok(old_nodes.into_iter().collect())
}

fn node_refs(&self, nodes: &[NodeId]) -> Vec<TypeErasedPinnedRef<'static>> {
self.node_deps(nodes).into_iter().map(|node| unsafe { node.as_ref().static_ref() }).collect()
self.node_deps(nodes).into_iter().map(|node| unsafe { node.read().unwrap().static_ref() }).collect()
}
fn node_deps(&self, nodes: &[NodeId]) -> Vec<Arc<NodeContainer<'static>>> {
fn node_deps(&self, nodes: &[NodeId]) -> Vec<Arc<RwLock<NodeContainer<'static>>>> {
nodes.iter().map(|node| self.nodes.get(node).unwrap().clone()).collect()
}

fn store_node(&mut self, node: Arc<NodeContainer<'static>>, id: NodeId) -> Arc<NodeContainer<'static>> {
fn store_node(&mut self, node: Arc<RwLock<NodeContainer<'static>>>, id: NodeId) -> Arc<RwLock<NodeContainer<'static>>> {
self.nodes.insert(id, node.clone());
node
}

pub fn get(&self, id: NodeId) -> Option<Arc<NodeContainer<'static>>> {
pub fn get(&self, id: NodeId) -> Option<Arc<RwLock<NodeContainer<'static>>>> {
self.nodes.get(&id).cloned()
}

pub fn eval<'i, I: StaticType + 'i, O: StaticType + 'i>(&self, id: NodeId, input: I) -> Option<O> {
pub fn eval<'i, I: StaticType + 'i, O: StaticType + 'i>(&'i self, id: NodeId, input: I) -> Option<O> {
let node = self.nodes.get(&id).cloned()?;
let output = node.node.eval(Box::new(input));
let reader = node.read().unwrap();
let output = reader.node.eval(Box::new(input));
dyn_any::downcast::<O>(output).ok().map(|o| *o)
}
pub fn eval_any<'i, 's: 'i>(&'s self, id: NodeId, input: Any<'i>) -> Option<Any<'i>> {
pub fn eval_any<'i>(&'i self, id: NodeId, input: Any<'i>) -> Option<Any<'i>> {
let node = self.nodes.get(&id)?;
Some(node.node.eval(input))
// TODO: Comments by @TrueDoctor before this was merged:
// TODO: Oof I dislike the evaluation being an unsafe operation but I guess its fine because it only is a lifetime extension
// TODO: We should ideally let miri run on a test that evaluates the nodegraph multiple times to check if this contains any subtle UB but this looks fine for now
Some(unsafe { (*((&*node.read().unwrap()) as *const NodeContainer)).node.eval(input) })

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Oof I dislike the evaluation being an unsafe operation but I guess its fine because it only is a lifetime extension

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

That is a big ugly - I wasn't quite sure how to resolve this though.

}

pub fn free_node(&mut self, id: NodeId) {
Expand All @@ -152,7 +161,7 @@ impl BorrowTree {
let node = Box::pin(upcasted) as TypeErasedPinned<'_>;
let node = NodeContainer { node, _dependencies: vec![] };
let node = unsafe { node.erase_lifetime() };
self.store_node(Arc::new(node), id);
self.store_node(Arc::new(node.into()), id);
}
ConstructionArgs::Nodes(ids) => {
let ids: Vec<_> = ids.iter().map(|(id, _)| *id).collect();
Expand All @@ -164,7 +173,7 @@ impl BorrowTree {
_dependencies: self.node_deps(&ids),
};
let node = unsafe { node.erase_lifetime() };
self.store_node(Arc::new(node), id);
self.store_node(Arc::new(node.into()), id);
}
};
Ok(())
Expand Down