///|
pub enum NodeOp {
Name(String)
Symbol(String)
Number(Float)
} derive(Eq, Hash, Show)
///|
pub fn NodeOp::label(self : NodeOp) -> String {
match self {
Name(s) => s
Symbol(s) => s
Number(n) => n.to_string()
}
}
///|
struct NodeKey {
op : NodeOp
children : Array[Id]
} derive(Eq, Hash)
///|
pub struct ENode {
op : NodeOp
children : Array[Id]
} derive(Eq, Hash, Show)
///|
pub fn make_enode(op : String, children : Array[Id]) -> ENode {
ENode::{ op: NodeOp::Name(op), children }
}
///|
pub fn make_symbol(sym : String) -> ENode {
ENode::{ op: NodeOp::Symbol(sym), children: [] }
}
///|
pub fn make_number(n : Float) -> ENode {
ENode::{ op: NodeOp::Number(n), children: [] }
}
///|
fn parse_number_literal(token : String) -> Float? {
if token.length() == 0 {
return None
}
let mut idx = 0
let mut neg = false
if token.get_char(0) == Some('-') {
neg = true
idx = 1
}
if idx >= token.length() {
return None
}
let mut acc : Int = 0
while idx < token.length() {
match token.get_char(idx) {
Some(c) if c >= '0' && c <= '9' => {
let digit : Int = c.to_int() - '0'.to_int()
acc = acc * 10 + digit
}
_ => return None
}
idx = idx + 1
}
let val = if neg { -acc } else { acc }
Some(Float::from_int(val))
}
///|
pub struct EClass {
id : Id
nodes : Array[Int]
data : Data
}
///|
pub struct EGraph {
uf : UnionFind
mut classes : Map[Id, EClass]
mut memo : Map[NodeKey, Id]
nodes : Array[ENode]
node_classes : Array[Id]
mut dirty : Bool
analysis : Analysis
mut allow_cycles : Bool
}
///|
pub type SimpleGraph = EGraph
///|
pub fn EGraph::new_with(analysis : Analysis) -> EGraph {
EGraph::{
uf: UnionFind::new(),
classes: Map::new(),
memo: Map::new(),
nodes: Array::new(),
node_classes: Array::new(),
dirty: false,
analysis,
allow_cycles: true,
}
}
///|
pub fn EGraph::new() -> SimpleGraph {
EGraph::new_with(default_analysis())
}
///|
pub fn EGraph::set_allow_cycles(self : EGraph, allow : Bool) -> Unit {
self.allow_cycles = allow
}
///|
pub fn EGraph::allow_cycles(self : EGraph) -> Bool {
self.allow_cycles
}
///|
fn EGraph::canonicalize_children(
self : EGraph,
children : Array[Id],
) -> Array[Id] {
children.map(child => self.uf.find(child))
}
///|
pub fn EGraph::data(self : EGraph, id : Id) -> Data? {
match self.classes.get(self.uf.find_read(id)) {
Some(class) => Some(class.data)
None => None
}
}
///|
pub fn EGraph::find(self : EGraph, id : Id) -> Id {
self.uf.find(id)
}
///|
/// Lookup an existing enode without mutating the e-graph.
pub fn EGraph::lookup(self : EGraph, enode : ENode) -> Id? {
let canon_children = self.canonicalize_children(enode.children)
let key = NodeKey::{ op: enode.op, children: canon_children }
self.memo.get(key)
}
///|
pub fn EGraph::find_read(self : EGraph, id : Id) -> Id {
self.uf.find_read(id)
}
///|
pub fn EGraph::add(self : EGraph, enode : ENode) -> Id {
let canon_children = self.canonicalize_children(enode.children)
let key = NodeKey::{ op: enode.op, children: canon_children }
match self.memo.get(key) {
Some(existing) => self.find(existing)
None => {
let id = self.uf.make_set()
let stored = ENode::{ op: key.op, children: key.children }
let data = (self.analysis.make)(stored, child => self.data(child))
self.nodes.push(stored)
self.node_classes.push(id)
self.classes.set(id, EClass::{
id,
nodes: [self.nodes.length() - 1],
data,
})
self.memo.set(key, id)
self.dirty = true
id
}
}
}
///|
pub fn EGraph::add_expr(self : EGraph, expr : Expr) -> Id {
match expr {
Expr::Leaf(name) =>
match parse_number_literal(name) {
Some(n) => self.add(make_number(n))
None => self.add(make_symbol(name))
}
Expr::Node(op, children) => {
let child_ids = children.map(child => self.add_expr(child))
self.add(make_enode(op, child_ids))
}
}
}
///|
pub fn EGraph::union(self : EGraph, a : Id, b : Id) -> Id {
let root_a = self.find(a)
let root_b = self.find(b)
if root_a == root_b {
return root_a
}
let merged_nodes = {
let left_nodes = match self.classes.get(root_a) {
Some(class) => class.nodes
None => []
}
let right_nodes = match self.classes.get(root_b) {
Some(class) => class.nodes
None => []
}
let combined = left_nodes
combined.append(right_nodes.op_as_view())
combined
}
let merged_data = {
let data_a = self.classes.get(root_a).unwrap().data
let data_b = self.classes.get(root_b).unwrap().data
(self.analysis.merge)(data_a, data_b)
}
let leader = self.uf.union(root_a, root_b)
self.classes.remove(root_b)
self.classes.set(leader, EClass::{
id: leader,
nodes: merged_nodes,
data: merged_data,
})
self.dirty = true
leader
}
///|
pub fn EGraph::rebuild(self : EGraph) -> Unit {
if !self.dirty {
return
}
loop () {
_ => {
let mut merged = false
let seen : Map[NodeKey, Id] = Map::new()
for idx in 0.. {
let before = self.find(owner)
let after = self.union(existing, before)
merged = merged || after != before
}
None => seen.set(key, owner)
}
}
self.memo = seen
if merged {
continue ()
}
break ()
}
}
let next_classes : Map[Id, EClass] = Map::new()
let next_memo : Map[NodeKey, Id] = Map::new()
for idx in 0.. {
let child_root = self.uf.find_read(child)
match next_classes.get(child_root) {
Some(class) => Some(class.data)
None =>
match self.classes.get(child_root) {
Some(class) => Some(class.data)
None => None
}
}
})
next_classes.update(root, existing => match existing {
Some(class) => {
let merged_nodes = class.nodes
merged_nodes.push(idx)
let merged_data = (self.analysis.merge)(class.data, node_data)
Some(EClass::{ id: root, nodes: merged_nodes, data: merged_data })
}
None => Some(EClass::{ id: root, nodes: [idx], data: node_data })
})
next_memo.set(
NodeKey::{ op: canon_node.op, children: canon_node.children },
root,
)
self.node_classes[idx] = root
self.nodes[idx] = canon_node
}
self.classes = next_classes
self.memo = next_memo
self.dirty = false
for class_id in self.class_ids() {
let root = self.find_read(class_id)
if root != class_id {
continue
}
(self.analysis.modify)(self, root)
}
if self.dirty {
self.rebuild()
}
}
///|
pub fn EGraph::are_equivalent(self : EGraph, a : Id, b : Id) -> Bool {
self.find(a) == self.find(b)
}
///|
pub fn EGraph::class_ids(self : EGraph) -> Array[Id] {
Array::from_iter(self.classes.keys())
}
///|
pub fn EGraph::class_for(self : EGraph, id : Id) -> EClass? {
self.classes.get(self.find_read(id))
}