use eggplant::{prelude::*, tx_rx_vt_pr};
use log::info;
#[eggplant::dsl]
enum Expression {
MNum {
num: i64,
},
MVar {
name: String,
},
MAdd {
l: Expression,
r: Expression,
},
MSub {
l: Expression,
r: Expression,
},
MMul {
l: Expression,
r: Expression,
},
MDiv {
l: Expression,
r: Expression,
},
MMod {
l: Expression,
r: Expression,
},
MMin {
l: Expression,
r: Expression,
},
MMax {
l: Expression,
r: Expression,
},
MAnd {
l: Expression,
r: Expression,
},
MOr {
l: Expression,
r: Expression,
},
MGte {
l: Expression,
r: Expression,
},
MLt {
l: Expression,
r: Expression,
},
MFloorTo {
l: Expression,
r: Expression,
},
MReplace {
l: Expression,
r: Expression,
rpl: Expression,
},
MAccum {
name: String,
},
}
#[eggplant::base_ty]
#[derive(Hash)]
enum UnOp {
Exp2,
Log2,
Sqrt,
Sin,
Recip,
Neg,
}
#[eggplant::base_ty]
enum BinOp {
Add,
Mul,
Max,
}
#[eggplant::dsl]
enum IR {
GMEM {
name: String,
},
LoopIn {
ir: IR,
l: Expression,
r: Expression,
},
LoopOut {
ir: IR,
l: Expression,
r: Expression,
},
Unary {
op: UnOp,
ir: IR,
},
Binary {
op: BinOp,
l: IR,
r: IR,
},
SwapLoops {
ir: IR,
level: i64,
},
TileLoop {
ir: IR,
level: i64,
},
MergeLoops {
ir: IR,
level: i64,
},
TCMatmul {
inp_a: IR,
inp_b: IR,
a_k_stride: Expression,
b_k_stride: Expression,
a_inner_stride: Expression,
b_inner_stride: Expression,
c_inner_stride: Expression,
num_k_loops: Expression,
},
TiledMatmulInputA {
ir: IR,
num: i64,
expr: Expression,
},
TiledMatmulInputB {
ir: IR,
num: i64,
expr: Expression,
},
}
fn main() {
env_logger::init();
let default_ruleset = MyTx::new_ruleset("default_ruleset");
tx_rx_vt_pr!(MyTx, MyPatRec);
#[eggplant::pat_vars]
struct AddCommuPat {
a: Expression,
b: Expression,
m_add_node1: MAdd,
}
MyTx::add_rule(
"AddCommu",
default_ruleset,
|| {
let a = Expression::query_leaf();
let b = Expression::query_leaf();
let m_add_node1 = MAdd::query(&a, &b);
AddCommuPat::new(a, b, m_add_node1)
},
|ctx, pat| {
let result = ctx.insert_m_add(pat.b, pat.a);
ctx.union(pat.m_add_node1, result);
},
);
#[eggplant::pat_vars]
struct MulCommuPat {
a: Expression,
b: Expression,
m_mul_node1: MMul,
}
MyTx::add_rule(
"MulCommu",
default_ruleset,
|| {
let a = Expression::query_leaf();
let b = Expression::query_leaf();
let m_mul_node1 = MMul::query(&a, &b);
MulCommuPat::new(a, b, m_mul_node1)
},
|ctx, pat| {
let result = ctx.insert_m_mul(pat.b, pat.a);
ctx.union(pat.m_mul_node1, result);
},
);
#[eggplant::pat_vars]
struct AddAssocPat {
a: Expression,
b: Expression,
c: Expression,
m_add_node2: MAdd,
m_add_node1: MAdd,
}
MyTx::add_rule(
"AddAssoc",
default_ruleset,
|| {
let a = Expression::query_leaf();
let b = Expression::query_leaf();
let m_add_node2 = MAdd::query(&a, &b);
let c = Expression::query_leaf();
let m_add_node1 = MAdd::query(&m_add_node2, &c);
AddAssocPat::new(a, b, c, m_add_node2, m_add_node1)
},
|ctx, pat| {
let result = ctx.insert_m_add(pat.a, ctx.insert_m_add(pat.b, pat.c));
ctx.union(pat.m_add_node1, result);
},
);
#[eggplant::pat_vars]
struct MulAssocPat {
a: Expression,
b: Expression,
c: Expression,
m_mul_node2: MMul,
m_mul_node1: MMul,
}
MyTx::add_rule(
"MulAssoc",
default_ruleset,
|| {
let a = Expression::query_leaf();
let b = Expression::query_leaf();
let m_mul_node2 = MMul::query(&a, &b);
let c = Expression::query_leaf();
let m_mul_node1 = MMul::query(&m_mul_node2, &c);
MulAssocPat::new(a, b, c, m_mul_node2, m_mul_node1)
},
|ctx, pat| {
let result = ctx.insert_m_mul(pat.a, ctx.insert_m_mul(pat.b, pat.c));
ctx.union(pat.m_mul_node1, result);
},
);
#[eggplant::pat_vars]
struct AddFoldPat {
m_num_node2: MNum,
m_num_node3: MNum,
m_add_node1: MAdd,
}
MyTx::add_rule(
"AddFold",
default_ruleset,
|| {
let m_num_node2 = MNum::query();
let m_num_node3 = MNum::query();
let m_add_node1 = MAdd::query(&m_num_node2, &m_num_node3);
AddFoldPat::new(m_num_node2, m_num_node3, m_add_node1)
},
|ctx, pat| {
let result = ctx.insert_m_num(
(ctx.devalue(pat.m_num_node2.num) + ctx.devalue(pat.m_num_node3.num)),
);
ctx.union(pat.m_add_node1, result);
},
);
#[eggplant::pat_vars]
struct SubFoldPat {
m_num_node2: MNum,
m_num_node3: MNum,
m_sub_node1: MSub,
}
MyTx::add_rule(
"SubFold",
default_ruleset,
|| {
let m_num_node2 = MNum::query();
let m_num_node3 = MNum::query();
let m_sub_node1 = MSub::query(&m_num_node2, &m_num_node3);
SubFoldPat::new(m_num_node2, m_num_node3, m_sub_node1)
},
|ctx, pat| {
let result = ctx.insert_m_num(
(ctx.devalue(pat.m_num_node2.num) - ctx.devalue(pat.m_num_node3.num)),
);
ctx.union(pat.m_sub_node1, result);
},
);
#[eggplant::pat_vars]
struct MulFoldPat {
m_num_node2: MNum,
m_num_node3: MNum,
m_mul_node1: MMul,
}
MyTx::add_rule(
"MulFold",
default_ruleset,
|| {
let m_num_node2 = MNum::query();
let m_num_node3 = MNum::query();
let m_mul_node1 = MMul::query(&m_num_node2, &m_num_node3);
MulFoldPat::new(m_num_node2, m_num_node3, m_mul_node1)
},
|ctx, pat| {
let result = ctx.insert_m_num(
(ctx.devalue(pat.m_num_node2.num) * ctx.devalue(pat.m_num_node3.num)),
);
ctx.union(pat.m_mul_node1, result);
},
);
#[eggplant::pat_vars]
struct MaxFoldPat {
m_num_node2: MNum,
m_num_node3: MNum,
m_max_node1: MMax,
}
MyTx::add_rule(
"MaxFold",
default_ruleset,
|| {
let m_num_node2 = MNum::query();
let m_num_node3 = MNum::query();
let m_max_node1 = MMax::query(&m_num_node2, &m_num_node3);
MaxFoldPat::new(m_num_node2, m_num_node3, m_max_node1)
},
|ctx, pat| {
let result = ctx.insert_m_num(std::cmp::max(
ctx.devalue(pat.m_num_node2.num),
ctx.devalue(pat.m_num_node3.num),
));
ctx.union(pat.m_max_node1, result);
},
);
#[eggplant::pat_vars]
struct MinFoldPat {
m_num_node2: MNum,
m_num_node3: MNum,
m_min_node1: MMin,
}
MyTx::add_rule(
"MinFold",
default_ruleset,
|| {
let m_num_node2 = MNum::query();
let m_num_node3 = MNum::query();
let m_min_node1 = MMin::query(&m_num_node2, &m_num_node3);
MinFoldPat::new(m_num_node2, m_num_node3, m_min_node1)
},
|ctx, pat| {
let result = ctx.insert_m_num(std::cmp::min(
ctx.devalue(pat.m_num_node2.num),
ctx.devalue(pat.m_num_node3.num),
));
ctx.union(pat.m_min_node1, result);
},
);
#[eggplant::pat_vars]
struct AndFoldPat {
m_num_node3: MNum,
m_num_node2: MNum,
m_and_node1: MAnd,
}
MyTx::add_rule(
"AndFold",
default_ruleset,
|| {
let m_num_node2 = MNum::query();
let m_num_node3 = MNum::query();
let m_and_node1 = MAnd::query(&m_num_node2, &m_num_node3);
AndFoldPat::new(m_num_node3, m_num_node2, m_and_node1)
},
|ctx, pat| {
let result = ctx.insert_m_num(std::ops::BitAnd::bitand(
ctx.devalue(pat.m_num_node2.num),
ctx.devalue(pat.m_num_node3.num),
));
ctx.union(pat.m_and_node1, result);
},
);
#[eggplant::pat_vars]
struct AddZeroPat {
a: Expression,
m_add_node1: MAdd,
}
MyTx::add_rule(
"AddZero",
default_ruleset,
|| {
let a = Expression::query_leaf();
let m_num_node2 = MNum::query();
let m_add_node1 = MAdd::query(&a, &m_num_node2);
AddZeroPat::new(a, m_add_node1)
},
|ctx, pat| {
let result = pat.a;
ctx.union(pat.m_add_node1, result);
},
);
#[eggplant::pat_vars]
struct MulOnePat {
a: Expression,
m_mul_node1: MMul,
}
MyTx::add_rule(
"MulOne",
default_ruleset,
|| {
let a = Expression::query_leaf();
let m_num_node2 = MNum::query();
let m_mul_node1 = MMul::query(&a, &m_num_node2);
MulOnePat::new(a, m_mul_node1)
},
|ctx, pat| {
let result = pat.a;
ctx.union(pat.m_mul_node1, result);
},
);
#[eggplant::pat_vars]
struct MulZeroPat {
a: Expression,
m_mul_node1: MMul,
}
MyTx::add_rule(
"MulZero",
default_ruleset,
|| {
let a = Expression::query_leaf();
let m_num_node2 = MNum::query();
let m_mul_node1 = MMul::query(&a, &m_num_node2);
MulZeroPat::new(a, m_mul_node1)
},
|ctx, pat| {
let result = ctx.insert_m_num(0);
ctx.union(pat.m_mul_node1, result);
},
);
#[eggplant::pat_vars]
struct DivOnePat {
a: Expression,
m_div_node1: MDiv,
}
MyTx::add_rule(
"DivOne",
default_ruleset,
|| {
let a = Expression::query_leaf();
let m_div_node1 = MDiv::query(&a, &MNum::query());
DivOnePat::new(a, m_div_node1)
},
|ctx, pat| {
let result = pat.a;
ctx.union(pat.m_div_node1, result);
},
);
#[eggplant::pat_vars]
struct ModMulToZeroPat {
_x: Expression,
_y: Expression,
m_mod_node1: MMod,
m_mul_node2: MMul,
}
MyTx::add_rule(
"ModMulToZero",
default_ruleset,
|| {
let _x = Expression::query_leaf();
let _y = Expression::query_leaf();
let m_mul_node2 = MMul::query(&_x, &_y);
let m_mod_node1 = MMod::query(&m_mul_node2, &_y);
ModMulToZeroPat::new(_x, _y, m_mod_node1, m_mul_node2)
},
|ctx, pat| {
let result = ctx.insert_m_num(0);
ctx.union(pat.m_mod_node1, result);
},
);
#[eggplant::pat_vars]
struct ModMod1Pat {
_x: Expression,
m_num_node4: MNum,
m_num_node3: MNum,
m_mod_node2: MMod,
m_mod_node1: MMod,
}
MyTx::add_rule(
"ModMod1",
default_ruleset,
|| {
let _x = Expression::query_leaf();
let m_num_node3 = MNum::query();
let m_mod_node2 = MMod::query(&_x, &m_num_node3);
let m_num_node4 = MNum::query();
let m_mod_node1 = MMod::query(&m_mod_node2, &m_num_node4);
let cond_z_y = { m_num_node4.handle_num().ge(&m_num_node3.handle_num()) };
ModMod1Pat::new(_x, m_num_node4, m_num_node3, m_mod_node2, m_mod_node1).assert(cond_z_y)
},
|ctx, pat| {
if ctx.devalue(pat.m_num_node3.num) % ctx.devalue(pat.m_num_node4.num) != 0 {
return;
}
let result =
ctx.insert_m_mod(pat._x, ctx.insert_m_num(ctx.devalue(pat.m_num_node3.num)));
ctx.union(pat.m_mod_node1, result);
},
);
#[eggplant::pat_vars]
struct ModMod2Pat {
_x: Expression,
m_mod_node2: MMod,
m_num_node3: MNum,
m_num_node4: MNum,
m_mod_node1: MMod,
}
MyTx::add_rule(
"ModMod2",
default_ruleset,
|| {
let _x = Expression::query_leaf();
let m_num_node3 = MNum::query();
let m_mod_node2 = MMod::query(&_x, &m_num_node3);
let m_num_node4 = MNum::query();
let m_mod_node1 = MMod::query(&m_mod_node2, &m_num_node4);
ModMod2Pat::new(_x, m_mod_node2, m_num_node3, m_num_node4, m_mod_node1)
},
|ctx, pat| {
let y = ctx.devalue(pat.m_num_node3.num);
let z = ctx.devalue(pat.m_num_node4.num);
if (y >= z) && z % y == 0 {
let result =
ctx.insert_m_mod(pat._x, ctx.insert_m_num(ctx.devalue(pat.m_num_node4.num)));
ctx.union(pat.m_mod_node1, result);
}
},
);
#[eggplant::pat_vars]
struct ReplacePat {
_x: Expression,
_y: Expression,
_z: Expression,
m_replace_node1: MReplace,
}
MyTx::add_rule(
"Replace",
default_ruleset,
|| {
let _x = Expression::query_leaf();
let _y = Expression::query_leaf();
let _z = Expression::query_leaf();
let m_replace_node1 = MReplace::query(&_x, &_y, &_z);
let cond__x__y = { _x.handle().eq(&_y.handle()) };
ReplacePat::new(_x, _y, _z, m_replace_node1).assert(cond__x__y)
},
|ctx, pat| {
let result = pat._z;
ctx.union(pat.m_replace_node1, result);
},
);
#[eggplant::pat_vars]
struct ReplaceRecurAddPat {
_a: Expression,
_b: Expression,
_x: Expression,
_y: Expression,
m_add_node2: MAdd,
m_replace_node1: MReplace,
}
MyTx::add_rule(
"ReplaceRecurAdd",
default_ruleset,
|| {
let _a = Expression::query_leaf();
let _b = Expression::query_leaf();
let m_add_node2 = MAdd::query(&_a, &_b);
let _x = Expression::query_leaf();
let _y = Expression::query_leaf();
let m_replace_node1 = MReplace::query(&m_add_node2, &_x, &_y);
ReplaceRecurAddPat::new(_a, _b, _x, _y, m_add_node2, m_replace_node1)
},
|ctx, pat| {
let result = ctx.insert_m_add(
ctx.insert_m_replace(pat._a, pat._x, pat._y),
ctx.insert_m_replace(pat._b, pat._x, pat._y),
);
ctx.union(pat.m_replace_node1, result);
},
);
#[eggplant::pat_vars]
struct ReplaceRecurSubPat {
_a: Expression,
_b: Expression,
_x: Expression,
_y: Expression,
m_replace_node1: MReplace,
m_sub_node2: MSub,
}
MyTx::add_rule(
"ReplaceRecurSub",
default_ruleset,
|| {
let _a = Expression::query_leaf();
let _b = Expression::query_leaf();
let m_sub_node2 = MSub::query(&_a, &_b);
let _x = Expression::query_leaf();
let _y = Expression::query_leaf();
let m_replace_node1 = MReplace::query(&m_sub_node2, &_x, &_y);
ReplaceRecurSubPat::new(_a, _b, _x, _y, m_replace_node1, m_sub_node2)
},
|ctx, pat| {
let result = ctx.insert_m_sub(
ctx.insert_m_replace(pat._a, pat._x, pat._y),
ctx.insert_m_replace(pat._b, pat._x, pat._y),
);
ctx.union(pat.m_replace_node1, result);
},
);
#[eggplant::pat_vars]
struct ReplaceRecurMulPat {
_a: Expression,
_b: Expression,
_x: Expression,
_y: Expression,
m_mul_node2: MMul,
m_replace_node1: MReplace,
}
MyTx::add_rule(
"ReplaceRecurMul",
default_ruleset,
|| {
let _a = Expression::query_leaf();
let _b = Expression::query_leaf();
let m_mul_node2 = MMul::query(&_a, &_b);
let _x = Expression::query_leaf();
let _y = Expression::query_leaf();
let m_replace_node1 = MReplace::query(&m_mul_node2, &_x, &_y);
ReplaceRecurMulPat::new(_a, _b, _x, _y, m_mul_node2, m_replace_node1)
},
|ctx, pat| {
let result = ctx.insert_m_mul(
ctx.insert_m_replace(pat._a, pat._x, pat._y),
ctx.insert_m_replace(pat._b, pat._x, pat._y),
);
ctx.union(pat.m_replace_node1, result);
},
);
#[eggplant::pat_vars]
struct ReplaceRecurDivPat {
_a: Expression,
_b: Expression,
_x: Expression,
_y: Expression,
m_replace_node1: MReplace,
m_div_node2: MDiv,
}
MyTx::add_rule(
"ReplaceRecurDiv",
default_ruleset,
|| {
let _a = Expression::query_leaf();
let _b = Expression::query_leaf();
let m_div_node2 = MDiv::query(&_a, &_b);
let _x = Expression::query_leaf();
let _y = Expression::query_leaf();
let m_replace_node1 = MReplace::query(&m_div_node2, &_x, &_y);
ReplaceRecurDivPat::new(_a, _b, _x, _y, m_replace_node1, m_div_node2)
},
|ctx, pat| {
let result = ctx.insert_m_div(
ctx.insert_m_replace(pat._a, pat._x, pat._y),
ctx.insert_m_replace(pat._b, pat._x, pat._y),
);
ctx.union(pat.m_replace_node1, result);
},
);
#[eggplant::pat_vars]
struct ReplaceRecurModPat {
_a: Expression,
_b: Expression,
_x: Expression,
_y: Expression,
m_replace_node1: MReplace,
m_mod_node2: MMod,
}
MyTx::add_rule(
"ReplaceRecurMod",
default_ruleset,
|| {
let _a = Expression::query_leaf();
let _b = Expression::query_leaf();
let m_mod_node2 = MMod::query(&_a, &_b);
let _x = Expression::query_leaf();
let _y = Expression::query_leaf();
let m_replace_node1 = MReplace::query(&m_mod_node2, &_x, &_y);
ReplaceRecurModPat::new(_a, _b, _x, _y, m_replace_node1, m_mod_node2)
},
|ctx, pat| {
let result = ctx.insert_m_mod(
ctx.insert_m_replace(pat._a, pat._x, pat._y),
ctx.insert_m_replace(pat._b, pat._x, pat._y),
);
ctx.union(pat.m_replace_node1, result);
},
);
#[eggplant::pat_vars]
struct ReplaceRecurMinPat {
_a: Expression,
_b: Expression,
_x: Expression,
_y: Expression,
m_replace_node1: MReplace,
m_min_node2: MMin,
}
MyTx::add_rule(
"ReplaceRecurMin",
default_ruleset,
|| {
let _a = Expression::query_leaf();
let _b = Expression::query_leaf();
let m_min_node2 = MMin::query(&_a, &_b);
let _x = Expression::query_leaf();
let _y = Expression::query_leaf();
let m_replace_node1 = MReplace::query(&m_min_node2, &_x, &_y);
ReplaceRecurMinPat::new(_a, _b, _x, _y, m_replace_node1, m_min_node2)
},
|ctx, pat| {
let result = ctx.insert_m_min(
ctx.insert_m_replace(pat._a, pat._x, pat._y),
ctx.insert_m_replace(pat._b, pat._x, pat._y),
);
ctx.union(pat.m_replace_node1, result);
},
);
#[eggplant::pat_vars]
struct ReplaceRecurMaxPat {
_a: Expression,
_b: Expression,
_x: Expression,
_y: Expression,
m_max_node2: MMax,
m_replace_node1: MReplace,
}
MyTx::add_rule(
"ReplaceRecurMax",
default_ruleset,
|| {
let _a = Expression::query_leaf();
let _b = Expression::query_leaf();
let m_max_node2 = MMax::query(&_a, &_b);
let _x = Expression::query_leaf();
let _y = Expression::query_leaf();
let m_replace_node1 = MReplace::query(&m_max_node2, &_x, &_y);
ReplaceRecurMaxPat::new(_a, _b, _x, _y, m_max_node2, m_replace_node1)
},
|ctx, pat| {
let result = ctx.insert_m_max(
ctx.insert_m_replace(pat._a, pat._x, pat._y),
ctx.insert_m_replace(pat._b, pat._x, pat._y),
);
ctx.union(pat.m_replace_node1, result);
},
);
#[eggplant::pat_vars]
struct ReplaceRecurFloorPat {
_a: Expression,
_b: Expression,
_x: Expression,
_y: Expression,
m_replace_node1: MReplace,
m_floor_to_node2: MFloorTo,
}
MyTx::add_rule(
"ReplaceRecurFloor",
default_ruleset,
|| {
let _a = Expression::query_leaf();
let _b = Expression::query_leaf();
let m_floor_to_node2 = MFloorTo::query(&_a, &_b);
let _x = Expression::query_leaf();
let _y = Expression::query_leaf();
let m_replace_node1 = MReplace::query(&m_floor_to_node2, &_x, &_y);
ReplaceRecurFloorPat::new(_a, _b, _x, _y, m_replace_node1, m_floor_to_node2)
},
|ctx, pat| {
let result = ctx.insert_m_floor_to(
ctx.insert_m_replace(pat._a, pat._x, pat._y),
ctx.insert_m_replace(pat._b, pat._x, pat._y),
);
ctx.union(pat.m_replace_node1, result);
},
);
#[eggplant::pat_vars]
struct ReplaceRecurNumPat {
_x: Expression,
_y: Expression,
m_replace_node1: MReplace,
m_num_node2: MNum,
}
MyTx::add_rule(
"ReplaceRecurNum",
default_ruleset,
|| {
let m_num_node2 = MNum::query();
let _x = Expression::query_leaf();
let _y = Expression::query_leaf();
let m_replace_node1 = MReplace::query(&m_num_node2, &_x, &_y);
ReplaceRecurNumPat::new(_x, _y, m_replace_node1, m_num_node2)
},
|ctx, pat| {
let result = ctx.insert_m_num(ctx.devalue(pat.m_num_node2.num));
ctx.union(pat.m_replace_node1, result);
},
);
#[eggplant::pat_vars]
struct ReplaceRecurAccumPat {
_x: Expression,
_y: Expression,
m_replace_node1: MReplace,
m_accum_node2: MAccum,
}
MyTx::add_rule(
"ReplaceRecurAccum",
default_ruleset,
|| {
let m_accum_node2 = MAccum::query();
let _x = Expression::query_leaf();
let _y = Expression::query_leaf();
let m_replace_node1 = MReplace::query(&m_accum_node2, &_x, &_y);
ReplaceRecurAccumPat::new(_x, _y, m_replace_node1, m_accum_node2)
},
|ctx, pat| {
let result = ctx.insert_m_accum(ctx.devalue(pat.m_accum_node2.name));
ctx.union(pat.m_replace_node1, result);
},
);
#[eggplant::pat_vars]
struct ReplaceRecurVarPat {
_y: Expression,
m_replace_node1: MReplace,
m_var_node3: MVar,
m_var_node2: MVar,
}
MyTx::add_rule(
"ReplaceRecurVar",
default_ruleset,
|| {
let m_var_node2 = MVar::query();
let m_var_node3 = MVar::query();
let _y = Expression::query_leaf();
let m_replace_node1 = MReplace::query(&m_var_node2, &m_var_node3, &_y);
let cond_v_x = { m_var_node2.handle_name().ne(&m_var_node3.handle_name()) };
ReplaceRecurVarPat::new(_y, m_replace_node1, m_var_node3, m_var_node2).assert(cond_v_x)
},
|ctx, pat| {
let result = ctx.insert_m_var(ctx.devalue(pat.m_var_node2.name));
ctx.union(pat.m_replace_node1, result);
},
);
MyTx::run_ruleset(default_ruleset, RunConfig::Sat);
info!("Eggplant program executed successfully!");
}