// Identity rules for EGraph optimization
// ============================================================================
// Standard rewrite rules
// ============================================================================
///|
/// Helper to check if an e-class contains a specific integer constant
/// Uses cached value for O(1) lookup
fn EGraph::find_const(self : EGraph, id : EClassId) -> Int64? {
self.get_const(id)
}
///|
/// Helper to check if an e-class contains a specific float constant (as bits)
/// Uses cached value for O(1) lookup
fn EGraph::find_fconst(self : EGraph, id : EClassId) -> UInt64? {
self.get_fconst(id)
}
///|
/// x + 0 = x
fn rule_add_zero() -> RewriteRule {
{
apply: fn(eg, class_id) {
let mut changed = false
for node in eg.get_nodes(class_id) {
if node.op is Add &&
node.children.length() == 2 &&
eg.find_const(node.children[1]) is Some(0L) {
// x + 0 = x
changed = eg.merge_changed(class_id, node.children[0]) || changed
} else if node.op is Add &&
node.children.length() == 2 &&
eg.find_const(node.children[0]) is Some(0L) {
// 0 + x = x
changed = eg.merge_changed(class_id, node.children[1]) || changed
}
}
changed
},
}
}
///|
/// x - 0 = x
fn rule_sub_zero() -> RewriteRule {
{
apply: fn(eg, class_id) {
let mut changed = false
for node in eg.get_nodes(class_id) {
if node.op is Sub &&
node.children.length() == 2 &&
eg.find_const(node.children[1]) is Some(0L) {
// x - 0 = x
changed = eg.merge_changed(class_id, node.children[0]) || changed
}
}
changed
},
}
}
///|
/// x * 1 = x
fn rule_mul_one() -> RewriteRule {
{
apply: fn(eg, class_id) {
let mut changed = false
for node in eg.get_nodes(class_id) {
if node.op is Mul &&
node.children.length() == 2 &&
eg.find_const(node.children[1]) is Some(1L) {
changed = eg.merge_changed(class_id, node.children[0]) || changed
} else if node.op is Mul &&
node.children.length() == 2 &&
eg.find_const(node.children[0]) is Some(1L) {
changed = eg.merge_changed(class_id, node.children[1]) || changed
}
}
changed
},
}
}
///|
/// x * 0 = 0
fn rule_mul_zero() -> RewriteRule {
{
apply: fn(eg, class_id) {
let mut changed = false
for node in eg.get_nodes(class_id) {
if node.op is Mul &&
node.children.length() == 2 &&
(
eg.find_const(node.children[1]) is Some(0L) ||
eg.find_const(node.children[0]) is Some(0L)
) {
let zero = eg.add_const(0L)
changed = eg.merge_changed(class_id, zero) || changed
}
}
changed
},
}
}
///|
/// x & x = x, x | x = x
fn rule_idempotent() -> RewriteRule {
{
apply: fn(eg, class_id) {
let mut changed = false
for node in eg.get_nodes(class_id) {
if (node.op is And || node.op is Or) &&
node.children.length() == 2 &&
eg.equiv(node.children[0], node.children[1]) {
changed = eg.merge_changed(class_id, node.children[0]) || changed
}
}
changed
},
}
}
///|
/// x ^ x = 0
fn rule_xor_self() -> RewriteRule {
{
apply: fn(eg, class_id) {
let mut changed = false
for node in eg.get_nodes(class_id) {
if node.op is Xor &&
node.children.length() == 2 &&
eg.equiv(node.children[0], node.children[1]) {
let zero = eg.add_const(0L)
changed = eg.merge_changed(class_id, zero) || changed
}
}
changed
},
}
}
///|
/// x - x = 0
fn rule_sub_self() -> RewriteRule {
{
apply: fn(eg, class_id) {
let mut changed = false
for node in eg.get_nodes(class_id) {
if node.op is Sub &&
node.children.length() == 2 &&
eg.equiv(node.children[0], node.children[1]) {
let zero = eg.add_const(0L)
changed = eg.merge_changed(class_id, zero) || changed
}
}
changed
},
}
}
///|
/// Helper: check if n is a power of 2 and return log2(n)
fn log2_if_pow2(n : Int64) -> Int? {
if n <= 0L {
return None
}
// Check if n is power of 2: n & (n-1) == 0
if (n & (n - 1L)) != 0L {
return None
}
// Count trailing zeros to get log2
let mut count = 0
let mut val = n
while (val & 1L) == 0L {
count = count + 1
val = val >> 1
}
Some(count)
}
///|
/// x * 2^n = x << n (strength reduction)