///|
/// Analysis payload supporting numeric and boolean constants plus free-variable sets.
pub(all) enum Value {
Num(Float)
Bool(Bool)
} derive(Eq, Show)
///|
/// Per-eclass analysis data.
pub(all) struct Data {
free : Map[Id, Bool]
constant : Value?
} derive(Show)
///|
pub fn data_num(d : Data?) -> Float? {
match d {
Some(payload) =>
match payload.constant {
Some(Value::Num(n)) => Some(n)
_ => None
}
None => None
}
}
///|
pub fn data_bool(d : Data?) -> Bool? {
match d {
Some(payload) =>
match payload.constant {
Some(Value::Bool(b)) => Some(b)
_ => None
}
None => None
}
}
///|
fn empty_data() -> Data {
let free : Map[Id, Bool] = Map::new()
Data::{ free, constant: None }
}
///|
pub fn intersect_free(a : Map[Id, Bool], b : Map[Id, Bool]) -> Map[Id, Bool] {
let result = Map::new()
for pair in a.iter() {
let (k, v) = pair
if v && b.contains(k) {
result.set(k, true)
}
}
result
}
///|
fn union_free(children : Array[Data?]) -> Map[Id, Bool] {
let free : Map[Id, Bool] = Map::new()
for child_data in children {
match child_data {
Some(payload) =>
for pair in payload.free.iter() {
let (k, v) = pair
if v {
free.set(k, true)
}
}
None => ignore(())
}
}
free
}
///|
pub fn merge_data(left : Data, right : Data) -> Data {
let constant = match (left.constant, right.constant) {
(Some(a), Some(b)) => if a == b { Some(a) } else { None }
(Some(a), None) => Some(a)
(None, Some(b)) => Some(b)
(None, None) => None
}
Data::{ free: intersect_free(left.free, right.free), constant }
}
///|
/// Keep only leaf nodes inside a constant class to curb blowup.
fn prune_to_leaves(egraph : EGraph, id : Id) -> Unit {
let root = egraph.find(id)
match egraph.class_for(root) {
Some(class) => {
let leaf_nodes = class.nodes.filter(idx => egraph.nodes[idx].children.is_empty())
egraph.classes.set(root, EClass::{
id: root,
nodes: leaf_nodes,
data: class.data,
})
}
None => ()
}
}
///|
/// Per-eclass analysis callbacks.
pub(all) struct Analysis {
make : (ENode, (Id) -> Data?) -> Data
merge : (Data, Data) -> Data
modify : (EGraph, Id) -> Unit
}
///|
pub fn default_analysis() -> Analysis {
Analysis::{
make: (_, _) => empty_data(),
merge: (a, b) => merge_data(a, b),
modify: (_, _) => (),
}
}
///|
pub fn constant_analysis() -> Analysis {
Analysis::{
make: (node, child_data) => {
let free = union_free(node.children.map(child => child_data(child)))
let constant = match node.op {
NodeOp::Number(n) => Some(Value::Num(n))
NodeOp::Symbol(name) if name == "true" => Some(Value::Bool(true))
NodeOp::Symbol(name) if name == "false" => Some(Value::Bool(false))
NodeOp::Name(op) =>
if op == "add" || op == "+" {
match (child_data(node.children[0]), child_data(node.children[1])) {
(Some(payload_a), Some(payload_b)) =>
match (payload_a.constant, payload_b.constant) {
(Some(Value::Num(a)), Some(Value::Num(b))) =>
Some(Value::Num(a + b))
_ => None
}
_ => None
}
} else if op == "sub" || op == "-" {
match (child_data(node.children[0]), child_data(node.children[1])) {
(Some(payload_a), Some(payload_b)) =>
match (payload_a.constant, payload_b.constant) {
(Some(Value::Num(a)), Some(Value::Num(b))) =>
Some(Value::Num(a - b))
_ => None
}
_ => None
}
} else if op == "mul" || op == "*" {
match (child_data(node.children[0]), child_data(node.children[1])) {
(Some(payload_a), Some(payload_b)) =>
match (payload_a.constant, payload_b.constant) {
(Some(Value::Num(a)), Some(Value::Num(b))) =>
Some(Value::Num(a * b))
_ => None
}
_ => None
}
} else if op == "div" || op == "/" {
match (child_data(node.children[0]), child_data(node.children[1])) {
(Some(payload_a), Some(payload_b)) =>
match (payload_a.constant, payload_b.constant) {
(Some(Value::Num(a)), Some(Value::Num(b))) =>
if b != 0.0 {
Some(Value::Num(a / b))
} else {
None
}
_ => None
}
_ => None
}
} else if op == "=" {
match (child_data(node.children[0]), child_data(node.children[1])) {
(Some(payload_a), Some(payload_b)) =>
match (payload_a.constant, payload_b.constant) {
(Some(Value::Num(a)), Some(Value::Num(b))) =>
Some(Value::Bool(a == b))
(Some(Value::Bool(a)), Some(Value::Bool(b))) =>
Some(Value::Bool(a == b))
_ => None
}
_ => None
}
} else {
None
}
_ => None
}
Data::{ free, constant }
},
merge: (a, b) => merge_data(a, b),
modify: (egraph, id) => match egraph.data(id) {
Some(payload) =>
match payload.constant {
Some(Value::Num(n)) => {
let num = egraph.add(make_number(n))
ignore(egraph.union(id, num))
prune_to_leaves(egraph, id)
}
Some(Value::Bool(true)) => {
let t = egraph.add(make_symbol("true"))
ignore(egraph.union(id, t))
prune_to_leaves(egraph, id)
}
Some(Value::Bool(false)) => {
let f = egraph.add(make_symbol("false"))
ignore(egraph.union(id, f))
prune_to_leaves(egraph, id)
}
None => ()
}
None => ()
},
}
}
///|
/// Boolean constant folding for propositional logic tests.
pub fn bool_analysis() -> Analysis {
Analysis::{
make: (node, child_data) => {
let free = union_free(node.children.map(child => child_data(child)))
let constant = match node.op {
NodeOp::Symbol(sym) if sym == "true" => Some(Value::Bool(true))
NodeOp::Symbol(sym) if sym == "false" => Some(Value::Bool(false))
NodeOp::Name(op) =>
if op == "&" {
match (child_data(node.children[0]), child_data(node.children[1])) {
(Some(payload_a), Some(payload_b)) =>
match (payload_a.constant, payload_b.constant) {
(Some(Value::Bool(a)), Some(Value::Bool(b))) =>
Some(Value::Bool(a && b))
_ => None
}
_ => None
}
} else if op == "|" {
match (child_data(node.children[0]), child_data(node.children[1])) {
(Some(payload_a), Some(payload_b)) =>
match (payload_a.constant, payload_b.constant) {
(Some(Value::Bool(a)), Some(Value::Bool(b))) =>
Some(Value::Bool(a || b))
_ => None
}
_ => None
}
} else if op == "~" {
match child_data(node.children[0]) {
Some(payload) =>
match payload.constant {
Some(Value::Bool(v)) => Some(Value::Bool(!v))
_ => None
}
None => None
}
} else if op == "->" {
match (child_data(node.children[0]), child_data(node.children[1])) {
(Some(payload_a), Some(payload_b)) =>
match (payload_a.constant, payload_b.constant) {
(Some(Value::Bool(a)), Some(Value::Bool(b))) =>
Some(Value::Bool(!a || b))
_ => None
}
_ => None
}
} else {
None
}
_ => None
}
Data::{ free, constant }
},
merge: (a, b) => merge_data(a, b),
modify: (egraph, id) => match egraph.data(id) {
Some(payload) =>
match payload.constant {
Some(Value::Bool(true)) => {
let t = egraph.add(make_symbol("true"))
ignore(egraph.union(id, t))
}
Some(Value::Bool(false)) => {
let f = egraph.add(make_symbol("false"))
ignore(egraph.union(id, f))
}
_ => ()
}
None => ()
},
}
}
///|
/// Lambda analysis tracking free variables and simple constant folding.
pub fn lambda_analysis() -> Analysis {
Analysis::{
make: (node, child_data) => {
let free = union_free(node.children.map(child => child_data(child)))
let constant = match node.op {
NodeOp::Number(n) => Some(Value::Num(n))
NodeOp::Symbol(sym) if sym == "true" => Some(Value::Bool(true))
NodeOp::Symbol(sym) if sym == "false" => Some(Value::Bool(false))
NodeOp::Name(op) =>
if op == "+" {
match (child_data(node.children[0]), child_data(node.children[1])) {
(Some(a), Some(b)) =>
match (a.constant, b.constant) {
(Some(Value::Num(x)), Some(Value::Num(y))) =>
Some(Value::Num(x + y))
_ => None
}
_ => None
}
} else if op == "=" {
match (child_data(node.children[0]), child_data(node.children[1])) {
(Some(a), Some(b)) =>
match (a.constant, b.constant) {
(Some(Value::Num(x)), Some(Value::Num(y))) =>
Some(Value::Bool(x == y))
(Some(Value::Bool(x)), Some(Value::Bool(y))) =>
Some(Value::Bool(x == y))
_ => None
}
_ => None
}
} else {
None
}
_ => None
}
// adjust free variables for binders
match node.op {
NodeOp::Name(op) if op == "var" =>
if node.children.length() > 0 {
free.set(node.children[0], true)
}
NodeOp::Name(op) if op == "lam" =>
if node.children.length() > 0 {
free.remove(node.children[0])
}
NodeOp::Name(op) if op == "fix" =>
if node.children.length() > 0 {
free.remove(node.children[0])
}
NodeOp::Name(op) if op == "let" =>
if node.children.length() > 0 {
free.remove(node.children[0])
}
_ => ignore(())
}
Data::{ free, constant }
},
merge: (a, b) => merge_data(a, b),
modify: (egraph, id) => {
match egraph.data(id) {
Some(payload) =>
match payload.constant {
Some(Value::Num(n)) => {
let num = egraph.add(make_number(n))
ignore(egraph.union(id, num))
prune_to_leaves(egraph, id)
}
Some(Value::Bool(true)) => {
let t = egraph.add(make_symbol("true"))
ignore(egraph.union(id, t))
prune_to_leaves(egraph, id)
}
Some(Value::Bool(false)) => {
let f = egraph.add(make_symbol("false"))
ignore(egraph.union(id, f))
prune_to_leaves(egraph, id)
}
None => ()
}
None => ()
}
match egraph.class_for(id) {
Some(class) =>
for node_idx in class.nodes {
let node = egraph.nodes[node_idx]
match node.op {
NodeOp::Name(op) if op == "if" && node.children.length() == 3 =>
match egraph.data(node.children[0]) {
Some(payload) =>
match payload.constant {
Some(Value::Bool(true)) =>
ignore(egraph.union(id, node.children[1]))
Some(Value::Bool(false)) =>
ignore(egraph.union(id, node.children[2]))
_ => ()
}
None => ()
}
_ => ()
}
}
None => ()
}
},
}
}