Skip to content

Commit

Permalink
Custom fn util
Browse files Browse the repository at this point in the history
  • Loading branch information
jafioti committed Jan 16, 2024
1 parent 54912c4 commit fa04b05
Show file tree
Hide file tree
Showing 2 changed files with 21 additions and 16 deletions.
20 changes: 4 additions & 16 deletions src/compilers/metal/elementwise_fusion.rs
Original file line number Diff line number Diff line change
Expand Up @@ -64,14 +64,8 @@ impl<T: MetalFloat> Compiler for ElementwiseFusionCompiler<T> {
// Fuse into a FusedElementwiseOp
let new_op;
let mut a_equation = graph
.graph
.node_weight_mut(a)
.unwrap()
.custom("elementwise", Box::<()>::default())
.unwrap()
.downcast_ref::<String>()
.unwrap()
.clone();
.node_custom::<String, _>(a, "elementwise", ())
.unwrap();
let mut n_edges = graph
.graph
.edges_directed(a, Direction::Incoming)
Expand Down Expand Up @@ -131,14 +125,8 @@ impl<T: MetalFloat> Compiler for ElementwiseFusionCompiler<T> {
}
} else {
let mut b_equation = graph
.graph
.node_weight_mut(b)
.unwrap()
.custom("elementwise", Box::<()>::default())
.unwrap()
.downcast_ref::<String>()
.unwrap()
.clone();
.node_custom::<String, _>(b, "elementwise", ())
.unwrap();
b_equation = b_equation.replace(&format!("input{to_input}"), &a_equation);
// Since we are removing the input from a, we must decrement all inputs larger than that
for i in to_input..n_edges {
Expand Down
17 changes: 17 additions & 0 deletions src/core/compiler_utils.rs
Original file line number Diff line number Diff line change
Expand Up @@ -242,6 +242,23 @@ impl Graph {
self.graph.add_edge(a, b, Dependency::Schedule);
}

/// Run the custom function on a node and get an output
pub fn node_custom<O: 'static, I: 'static>(
&mut self,
node: NodeIndex,
key: &str,
input: I,
) -> Option<O> {
let Some(node_weight) = self.graph.node_weight_mut(node) else {
return None;
};

node_weight
.custom(key, Box::new(input))
.map(|o| o.downcast::<O>().ok().map(|o| *o))
.flatten()
}

/// Convert to debug-viewable graph
pub fn debug_graph(
&self,
Expand Down

0 comments on commit fa04b05

Please sign in to comment.