(ruleset matmul_flatten)
(ruleset kernel_lower)
(ruleset direct_kernel)
(ruleset kernel_specialize)
(ruleset buffer_reuse)
(ruleset matmul_backend)
(ruleset glumoe)
(ruleset fusion_pair)
(ruleset fusion_grow)
(ruleset fusion_merge)
(ruleset expr)
(ruleset dtype_prop)
(ruleset cleanup)
(ruleset post_cleanup)
(datatype*
(Expression
(MNum i64)
(MFloat f64)
(MIter)
(MVar String)
(MAdd Expression Expression)
(MSub Expression Expression)
(MMul Expression Expression)
(MCeilDiv Expression Expression)
(MDiv Expression Expression)
(MMod Expression Expression)
(MMin Expression Expression)
(MMax Expression Expression)
(MAnd Expression Expression)
(MOr Expression Expression)
(MGte Expression Expression)
(MLt Expression Expression)
(MFloorTo Expression Expression)
(MReplace Expression Expression Expression)
)
(EList
(ECons Expression EList)
(ENil)
(MReplaceList EList Expression Expression)
(ReplaceNthFromEnd EList Expression i64)
(RemoveNthFromEnd EList i64)
(RowMajor EList)
)
(DType
(F32)
(F16)
(Bf16)
(Int)
(Bool)
(F4E2M1)
(F8E4M3)
(F8E5M2)
(F8UE8M0)
(I4)
(TF32)
)
)
(rule ((= ?__rw (MMul ?a ?b))) ((union ?__rw (MMul ?b ?a))) :ruleset expr :name "mul-comm")
(rule ((= ?__rw (MAdd ?a ?b))) ((union ?__rw (MAdd ?b ?a))) :ruleset expr :name "add-comm")
(rule ((= ?e (MAdd (MNum ?a) (MNum ?b))) (= ?ans (+ ?a ?b))) ((union ?e (MNum ?ans)) (subsume (MAdd (MNum ?a) (MNum ?b)))) :ruleset expr)
(rule ((= ?__rw (MSub (MNum ?a) (MNum ?b)))) ((union ?__rw (MNum (- ?a ?b)))) :ruleset expr :name "sub-const")
(rule ((= ?e (MMul (MNum ?a) (MNum ?b))) (= ?prod (* ?a ?b))) ((union ?e (MNum ?prod)) (subsume (MMul (MNum ?a) (MNum ?b)))) :ruleset expr)
(rule ((= ?__rw (MDiv (MNum ?a) (MNum ?b))) (!= 0 ?b) (= 0 (% ?a ?b))) ((union ?__rw (MNum (/ ?a ?b)))) :ruleset expr :name "div-const")
(rule ((= ?__rw (MDiv ?a ?a))) ((union ?__rw (MNum 1))) :ruleset expr :name "div-self")
(rule ((= ?__rw (MCeilDiv (MNum ?a) (MNum ?b))) (!= 0 ?b) (= 0 (% ?a ?b))) ((union ?__rw (MNum (/ ?a ?b)))) :ruleset expr :name "ceildiv-const")
(rule ((= ?__rw (MMax (MNum ?a) (MNum ?b)))) ((union ?__rw (MNum (max ?a ?b)))) :ruleset expr :name "max-const")
(rule ((= ?__rw (MMin (MNum ?a) (MNum ?b)))) ((union ?__rw (MNum (min ?a ?b)))) :ruleset expr :name "min-const")
(rule ((= ?__rw (MAnd (MNum ?a) (MNum ?b)))) ((union ?__rw (MNum (& ?a ?b)))) :ruleset expr :name "and-const")
(rule ((= ?__rw (MFloat -1.0))) ((union ?__rw (MNum -1))) :ruleset expr :name "float-neg1-to-num")
(rule ((= ?__rw (MNum -1))) ((union ?__rw (MFloat -1.0))) :ruleset expr :name "num-neg1-to-float")
(rule ((= ?__rw (MAdd ?a (MNum 0)))) ((union ?__rw ?a)) :ruleset expr :name "add-zero")
(rule ((= ?e (MMul ?a (MNum 1)))) ((union ?e ?a)) :ruleset expr)
(rule ((= ?e (MMul ?a (MNum 0)))) ((union ?e (MNum 0)) (subsume (MMul ?a (MNum 0)))) :ruleset expr)
(rule ((= ?__rw (MDiv ?a (MNum 1)))) ((union ?__rw ?a)) :ruleset expr :name "div-one")
(rule ((= ?__rw (MMod (MMul ?x ?y) ?y))) ((union ?__rw (MNum 0))) :ruleset expr :name "mod-mul-self")
(rule ((= ?__rw (MMod (MMod ?x (MNum ?y)) (MNum ?z))) (>= ?z ?y) (= 0 (% ?y ?z))) ((union ?__rw (MMod ?x (MNum ?y)))) :ruleset expr :name "mod-mod-larger")
(rule ((= ?__rw (MMod (MMod ?x (MNum ?y)) (MNum ?z))) (>= ?y ?z) (= 0 (% ?z ?y))) ((union ?__rw (MMod ?x (MNum ?z)))) :ruleset expr :name "mod-mod-smaller")
(rule ((= ?__rw (MAdd (MMul (MDiv ?z ?x) ?x) (MMod ?z ?x)))) ((union ?__rw ?z)) :ruleset expr :name "merge-dims")
(rule ((= ?__rw (MDiv (MDiv ?a (MNum ?b)) (MNum ?c)))) ((union ?__rw (MDiv ?a (MNum (* ?b ?c))))) :ruleset expr :name "div-div-num")
(rule ((= ?__rw (MAdd (MDiv ?a ?b) ?c))) ((union ?__rw (MDiv (MAdd ?a (MMul ?c ?b)) ?b))) :ruleset expr :name "add-div")
(rule ((= ?__rw (MAdd ?a (MSub ?b ?a)))) ((union ?__rw ?b)) :ruleset expr :name "add-sub-cancel")
(rule ((= ?__rw (MAdd (MSub ?b ?a) ?a))) ((union ?__rw ?b)) :ruleset expr :name "add-sub-cancel2")
(rule ((= ?__rw (MSub ?a ?a))) ((union ?__rw (MNum 0))) :ruleset expr :name "sub-self")
(rule ((= ?__rw (MAdd (MSub ?a (MNum ?b)) (MNum ?c)))) ((union ?__rw (MSub ?a (MNum (- ?b ?c))))) :ruleset expr :name "add-sub-const")
(rule ((= ?__rw (MAdd (MNum ?c) (MSub ?a (MNum ?b))))) ((union ?__rw (MSub ?a (MNum (- ?b ?c))))) :ruleset expr :name "add-sub-const2")
(rule ((= ?__rw (MSub (MAdd ?a (MNum ?b)) (MNum ?c)))) ((union ?__rw (MAdd ?a (MNum (- ?b ?c))))) :ruleset expr :name "sub-add-const")
(rule ((= ?__rw (MSub (MSub ?a (MNum ?b)) (MNum ?c)))) ((union ?__rw (MSub ?a (MNum (+ ?b ?c))))) :ruleset expr :name "sub-sub-const")
(rule ((= ?__rw (MAdd (MMul ?a ?b) (MMul ?a ?c)))) ((union ?__rw (MMul ?a (MAdd ?b ?c)))) :ruleset expr :name "factor")
(rule ((= ?__rw (MAdd ?a ?a))) ((union ?__rw (MMul (MNum 2) ?a))) :ruleset expr :name "double")
(rule ((= ?e (MAdd (MAdd ?a (MNum ?b)) (MNum ?c))) (= ?ans (+ ?b ?c))) ((union ?e (MAdd ?a (MNum ?ans))) (subsume (MAdd (MAdd ?a (MNum ?b)) (MNum ?c)))) :ruleset expr)
(rule ((= ?__rw (MAdd (MAdd (MNum ?b) (MVar ?v)) (MNum ?c)))) ((union ?__rw (MAdd (MVar ?v) (MNum (+ ?b ?c))))) :ruleset expr :name "add-assoc-var")
(rule ((= ?__rw (MAdd (MAdd (MNum ?b) (MMul ?n ?a)) (MNum ?c)))) ((union ?__rw (MAdd (MMul ?n ?a) (MNum (+ ?b ?c))))) :ruleset expr :name "add-assoc-mul")
(rule ((= ?__rw (MAdd (MMul (MNum ?n) ?a) ?a))) ((union ?__rw (MMul (MNum (+ ?n 1)) ?a)) (subsume (MAdd (MMul (MNum ?n) ?a) ?a))) :ruleset expr :name "combine-like-1")
(rule ((= ?__rw (MAdd ?a (MMul (MNum ?n) ?a)))) ((union ?__rw (MMul (MNum (+ ?n 1)) ?a)) (subsume (MAdd ?a (MMul (MNum ?n) ?a)))) :ruleset expr :name "combine-like-2")
(rule ((= ?__rw (MAdd (MMul ?a (MNum ?n)) ?a))) ((union ?__rw (MMul (MNum (+ ?n 1)) ?a)) (subsume (MAdd (MMul ?a (MNum ?n)) ?a))) :ruleset expr :name "combine-like-3")
(rule ((= ?__rw (MAdd ?a (MMul ?a (MNum ?n))))) ((union ?__rw (MMul (MNum (+ ?n 1)) ?a)) (subsume (MAdd ?a (MMul ?a (MNum ?n))))) :ruleset expr :name "combine-like-4")
(rule ((= ?__rw (MAdd (MAdd ?a (MVar ?v)) (MVar ?v)))) ((union ?__rw (MAdd ?a (MMul (MNum 2) (MVar ?v)))) (subsume (MAdd (MAdd ?a (MVar ?v)) (MVar ?v)))) :ruleset expr :name "combine-var-1")
(rule ((= ?__rw (MAdd (MAdd (MVar ?v) ?a) (MVar ?v)))) ((union ?__rw (MAdd ?a (MMul (MNum 2) (MVar ?v)))) (subsume (MAdd (MAdd (MVar ?v) ?a) (MVar ?v)))) :ruleset expr :name "combine-var-2")
(rule ((= ?__rw (MAdd (MAdd (MMul (MNum ?n) ?a) ?b) ?a))) ((union ?__rw (MAdd (MMul (MNum (+ ?n 1)) ?a) ?b)) (subsume (MAdd (MAdd (MMul (MNum ?n) ?a) ?b) ?a))) :ruleset expr :name "accum-1")
(rule ((= ?__rw (MAdd (MAdd ?b (MMul (MNum ?n) ?a)) ?a))) ((union ?__rw (MAdd ?b (MMul (MNum (+ ?n 1)) ?a))) (subsume (MAdd (MAdd ?b (MMul (MNum ?n) ?a)) ?a))) :ruleset expr :name "accum-2")
(rule ((= ?__rw (MReplace ?x ?y ?z)) (= ?x ?y)) ((union ?__rw ?z)) :ruleset expr :name "replace-match")
(rule ((= ?__rw (MReplace (MAdd ?a ?b) ?x ?y))) ((union ?__rw (MAdd (MReplace ?a ?x ?y) (MReplace ?b ?x ?y)))) :ruleset expr :name "replace-MAdd")
(rule ((= ?__rw (MReplace (MSub ?a ?b) ?x ?y))) ((union ?__rw (MSub (MReplace ?a ?x ?y) (MReplace ?b ?x ?y)))) :ruleset expr :name "replace-MSub")
(rule ((= ?__rw (MReplace (MMul ?a ?b) ?x ?y))) ((union ?__rw (MMul (MReplace ?a ?x ?y) (MReplace ?b ?x ?y)))) :ruleset expr :name "replace-MMul")
(rule ((= ?__rw (MReplace (MDiv ?a ?b) ?x ?y))) ((union ?__rw (MDiv (MReplace ?a ?x ?y) (MReplace ?b ?x ?y)))) :ruleset expr :name "replace-MDiv")
(rule ((= ?__rw (MReplace (MCeilDiv ?a ?b) ?x ?y))) ((union ?__rw (MCeilDiv (MReplace ?a ?x ?y) (MReplace ?b ?x ?y)))) :ruleset expr :name "replace-MCeilDiv")
(rule ((= ?__rw (MReplace (MMod ?a ?b) ?x ?y))) ((union ?__rw (MMod (MReplace ?a ?x ?y) (MReplace ?b ?x ?y)))) :ruleset expr :name "replace-MMod")
(rule ((= ?__rw (MReplace (MMin ?a ?b) ?x ?y))) ((union ?__rw (MMin (MReplace ?a ?x ?y) (MReplace ?b ?x ?y)))) :ruleset expr :name "replace-MMin")
(rule ((= ?__rw (MReplace (MMax ?a ?b) ?x ?y))) ((union ?__rw (MMax (MReplace ?a ?x ?y) (MReplace ?b ?x ?y)))) :ruleset expr :name "replace-MMax")
(rule ((= ?__rw (MReplace (MFloorTo ?a ?b) ?x ?y))) ((union ?__rw (MFloorTo (MReplace ?a ?x ?y) (MReplace ?b ?x ?y)))) :ruleset expr :name "replace-MFloorTo")
(rule ((= ?__rw (MReplace (MNum ?n) ?x ?y))) ((union ?__rw (MNum ?n))) :ruleset expr :name "replace-num")
(rule ((= ?__rw (MReplace (MVar ?z) ?find ?replace)) (!= ?find (MVar ?z))) ((union ?__rw (MVar ?z))) :ruleset expr :name "replace-var-miss")
(rule ((= ?__rw (MReplace (MIter) ?find ?replace)) (!= ?find (MIter))) ((union ?__rw (MIter))) :ruleset expr :name "replace-iter-miss")
(function len (EList) i64 :merge new)
(rule ((= ?e (ENil))) ((set (len ?e) 0)) :ruleset expr)
(rule ((= ?e (ECons ?expr ?list)) (= ?prev_len (len ?list))) ((set (len ?e) (+ ?prev_len 1))) :ruleset expr)
(function nth_from_end (EList i64) Expression :merge new)
(rule ((= ?e (ECons ?expr ?list)) (= ?list_len (len ?list))) ((set (nth_from_end ?e ?list_len) ?expr)) :ruleset expr)
(rule ((= ?e (ECons ?expr ?list)) (= ?other_nth (nth_from_end ?list ?n))) ((set (nth_from_end ?e ?n) ?other_nth)) :ruleset expr)
(function n_elements (EList) Expression :merge new)
(rule ((= ?e (ENil))) ((set (n_elements ?e) (MNum 1))) :ruleset expr)
(rule ((= ?e (ECons ?dim ?other)) (= ?other_elems (n_elements ?other))) ((set (n_elements ?e) (MMul ?dim ?other_elems))) :ruleset expr)
(rule ((= ?other (ECons ?other_dim ?other_other)) (= ?list (ECons ?d ?other)) (= ?e (RowMajor ?list)) (= ?n_elems (n_elements ?other))) ((union ?e (ECons (MMul ?n_elems (MIter)) (RowMajor ?other)))) :ruleset expr)
(rule ((= ?__rw (RowMajor (ECons ?dim (ENil))))) ((union ?__rw (ECons (MIter) (ENil)))) :ruleset expr :name "rowmajor-base")
(rule ((= ?__rw (MReplaceList (ECons ?expr ?list) ?from ?to))) ((union ?__rw (ECons (MReplace ?expr ?from ?to) (MReplaceList ?list ?from ?to)))) :ruleset expr :name "replace-list-cons")
(rule ((= ?e (ReplaceNthFromEnd (ECons ?expr ?list) ?to ?ind)) (= ?ind (len ?list))) ((union ?e (ECons ?to ?list))) :ruleset expr)
(rule ((= ?e (ReplaceNthFromEnd (ECons ?expr ?list) ?to ?ind)) (< ?ind (len ?list))) ((union ?e (ECons ?expr (ReplaceNthFromEnd ?list ?to ?ind)))) :ruleset expr)
(rule ((= ?e (RemoveNthFromEnd (ECons ?expr ?list) ?ind)) (= ?ind (len ?list))) ((union ?e ?list)) :ruleset expr)
(rule ((= ?e (RemoveNthFromEnd (ECons ?expr ?list) ?ind)) (< ?ind (len ?list))) ((union ?e (ECons ?expr (RemoveNthFromEnd ?list ?ind)))) :ruleset expr)
(datatype*
(IR
(OutputJoin IR IR)
(Op OpKind IList)
(ConsumedBuffer IR)
(Input i64 String DType)
(Output IR i64)
(LoopStart IR i64 i64 Expression DType)
(LoopEnd IR i64 i64 DType)
)
(OpKind
(KernelAdd EList EList EList EList DType)
(KernelMul EList EList EList EList DType)
(KernelMod EList EList EList EList DType)
(KernelLessThan EList EList EList EList DType)
(KernelIota Expression Expression)
(KernelGather EList EList EList EList EList DType)
(KernelScatter EList EList EList EList EList EList DType)
(KernelSum EList Expression EList Expression EList DType)
(KernelMax EList Expression EList Expression EList DType)
(KernelExp2 EList EList EList DType)
(KernelLog2 EList EList EList DType)
(KernelSin EList EList EList DType)
(KernelRecip EList EList EList DType)
(KernelSqrt EList EList EList DType)
(KernelConstant f64)
(KernelCast Expression DType DType)
(KernelEmbed EList EList EList Expression)
(KernelMean EList Expression EList Expression EList DType)
(KernelBatchMatVec EList Expression EList Expression EList Expression EList DType)
(KernelBatchMatMul EList Expression EList Expression EList Expression EList DType)
(KernelScatterNoCopy EList EList EList EList EList EList DType)
(KernelSoftmax EList EList EList Expression Expression DType)
(KernelExp EList EList EList DType)
(KernelSigmoid EList EList EList DType)
(FusionStart EList EList DType)
(FusionEnd EList EList DType)
(FusedSin EList EList EList DType)
(FusedSqrt EList EList EList DType)
(FusedExp EList EList EList DType)
(FusedExp2 EList EList EList DType)
(FusedLog2 EList EList EList DType)
(FusedRecip EList EList EList DType)
(FusedAdd EList EList EList EList DType)
(FusedMul EList EList EList EList DType)
(cublaslt Expression Expression Expression String String String String String String Expression Expression Expression Expression Expression Expression Expression Expression Expression DType DType DType DType String String f64 f64 String)
(GLUMoE Expression Expression Expression Expression Expression Expression Expression Expression)
(ComputeAttnMask Expression Expression)
(FlashInferAttention Expression Expression Expression Expression Expression)
(CustomOpKind i64 DType)
(LoopInput i64 i64 DType)
(LoopInputStatic i64 i64 DType)
(LoopOutput i64 i64 DType)
(LoopOutputSelect i64 i64 i64 DType)
(Constant f64)
(Cast Expression DType)
(Iota Expression Expression)
(Exp2 EList EList EList)
(Log2 EList EList EList)
(Sin EList EList EList)
(Recip EList EList EList)
(Sqrt EList EList EList)
(Add EList EList EList EList)
(Mul EList EList EList EList)
(Mod EList EList EList EList)
(LessThan EList EList EList EList)
(Gather EList EList EList EList)
(Scatter EList EList EList EList EList)
(Sum EList Expression EList Expression EList)
(Max EList Expression EList Expression EList)
(Softmax EList EList EList Expression Expression)
)
(IList
(ICons IR IList)
(INil)
)
)
(function dtype (IR) DType :merge new)
(ruleset base_cleanup)
(rule ((= ?m (MReplace ?a ?b ?c))) ((delete (MReplace ?a ?b ?c))) :ruleset base_cleanup)
(rule ((= ?m (MReplaceList ?a ?b ?c))) ((delete (MReplaceList ?a ?b ?c))) :ruleset base_cleanup)
(rule ((= ?m (ReplaceNthFromEnd ?a ?b ?c))) ((delete (ReplaceNthFromEnd ?a ?b ?c))) :ruleset base_cleanup)
(rule ((= ?m (RemoveNthFromEnd ?a ?b))) ((delete (RemoveNthFromEnd ?a ?b))) :ruleset base_cleanup)
(rule ((= ?m (RowMajor ?x))) ((delete (RowMajor ?x))) :ruleset base_cleanup)
(rule ((= ?m (len ?x))) ((delete (len ?x))) :ruleset base_cleanup)
(rule ((= ?m (nth_from_end ?x ?y))) ((delete (nth_from_end ?x ?y))) :ruleset base_cleanup)
(rule ((= ?m (n_elements ?x))) ((delete (n_elements ?x))) :ruleset base_cleanup)
(rule ((= ?__rw0 (Op (Add ?v43_shape ?v43_a_strides ?v43_b_strides ?v43_out_strides) ?__inputs)) (= ?__dt (dtype (Op (Add ?v43_shape ?v43_a_strides ?v43_b_strides ?v43_out_strides) ?__inputs)))) ((union ?__rw0 (Op (KernelAdd ?v43_shape ?v43_a_strides ?v43_b_strides ?v43_out_strides ?__dt) ?__inputs))) :ruleset kernel_lower)
(rule ((= ?__rw0 (Op (Mul ?v44_shape ?v44_a_strides ?v44_b_strides ?v44_out_strides) (ICons ?__first_inp ?__tail))) (= ?__dt (dtype ?__first_inp))) ((union ?__rw0 (Op (KernelMul ?v44_shape ?v44_a_strides ?v44_b_strides ?v44_out_strides ?__dt) (ICons ?__first_inp ?__tail)))) :ruleset kernel_lower)
(rule ((= ?__rw0 (Op (Mod ?v45_shape ?v45_a_strides ?v45_b_strides ?v45_out_strides) ?__inputs)) (= ?__dt (dtype (Op (Mod ?v45_shape ?v45_a_strides ?v45_b_strides ?v45_out_strides) ?__inputs)))) ((union ?__rw0 (Op (KernelMod ?v45_shape ?v45_a_strides ?v45_b_strides ?v45_out_strides ?__dt) ?__inputs))) :ruleset kernel_lower)
(rule ((= ?__rw0 (Op (LessThan ?v46_shape ?v46_a_strides ?v46_b_strides ?v46_out_strides) (ICons ?__inp_a (ICons ?__inp_b (INil))))) (= ?__dt (dtype ?__inp_a))) ((union ?__rw0 (Op (KernelLessThan ?v46_shape ?v46_a_strides ?v46_b_strides ?v46_out_strides ?__dt) (ICons ?__inp_a (ICons ?__inp_b (INil)))))) :ruleset kernel_lower)
(rule ((= ?__rw0 (Op (Iota ?v47_expr ?v47_range) ?__inputs))) ((union ?__rw0 (Op (KernelIota ?v47_expr ?v47_range) ?__inputs)) (set (dtype (Op (KernelIota ?v47_expr ?v47_range) ?__inputs)) (Int))) :ruleset kernel_lower)
(rule ((= ?__rw0 (Op (Gather ?v48_index_shape ?v48_index_strides ?v48_data_shape ?v48_data_strides) (ICons ?__indexes (ICons ?__data ?__tail)))) (= ?__dt (dtype ?__data))) ((union ?__rw0 (Op (KernelGather ?v48_index_shape ?v48_index_strides ?v48_data_shape ?v48_data_strides (RowMajor ?v48_index_shape) ?__dt) (ICons ?__indexes (ICons ?__data (INil)))))) :ruleset kernel_lower)
(rule ((= ?__rw0 (Op (Scatter ?v49_dest_shape ?v49_dest_strides ?v49_index_shape ?v49_index_strides ?v49_src_strides) (ICons ?__dest (ICons ?__indexes (ICons ?__src (INil)))))) (= ?__dt (dtype ?__src))) ((union ?__rw0 (Op (KernelScatter ?v49_dest_shape ?v49_dest_strides ?v49_index_shape ?v49_index_strides ?v49_src_strides (RowMajor ?v49_dest_shape) ?__dt) (ICons ?__dest (ICons ?__indexes (ICons ?__src (INil))))))) :ruleset kernel_lower)
(rule ((= ?__rw0 (Op (Sum ?v50_shape ?v50_iters ?v50_strides ?v50_iter_stride ?v50_out_strides) ?__inputs)) (= ?__dt (dtype (Op (Sum ?v50_shape ?v50_iters ?v50_strides ?v50_iter_stride ?v50_out_strides) ?__inputs)))) ((union ?__rw0 (Op (KernelSum ?v50_shape ?v50_iters ?v50_strides ?v50_iter_stride ?v50_out_strides ?__dt) ?__inputs))) :ruleset kernel_lower)
(rule ((= ?__rw0 (Op (Max ?v51_shape ?v51_iters ?v51_strides ?v51_iter_stride ?v51_out_strides) ?__inputs)) (= ?__dt (dtype (Op (Max ?v51_shape ?v51_iters ?v51_strides ?v51_iter_stride ?v51_out_strides) ?__inputs)))) ((union ?__rw0 (Op (KernelMax ?v51_shape ?v51_iters ?v51_strides ?v51_iter_stride ?v51_out_strides ?__dt) ?__inputs))) :ruleset kernel_lower)
(rule ((= ?__rw0 (Op (Exp2 ?v52_shape ?v52_strides ?v52_out_strides) ?__inputs)) (= ?__dt (dtype (Op (Exp2 ?v52_shape ?v52_strides ?v52_out_strides) ?__inputs)))) ((union ?__rw0 (Op (KernelExp2 ?v52_shape ?v52_strides ?v52_out_strides ?__dt) ?__inputs))) :ruleset kernel_lower)
(rule ((= ?__rw0 (Op (Log2 ?v53_shape ?v53_strides ?v53_out_strides) ?__inputs)) (= ?__dt (dtype (Op (Log2 ?v53_shape ?v53_strides ?v53_out_strides) ?__inputs)))) ((union ?__rw0 (Op (KernelLog2 ?v53_shape ?v53_strides ?v53_out_strides ?__dt) ?__inputs))) :ruleset kernel_lower)
(rule ((= ?__rw0 (Op (Sin ?v54_shape ?v54_strides ?v54_out_strides) ?__inputs)) (= ?__dt (dtype (Op (Sin ?v54_shape ?v54_strides ?v54_out_strides) ?__inputs)))) ((union ?__rw0 (Op (KernelSin ?v54_shape ?v54_strides ?v54_out_strides ?__dt) ?__inputs))) :ruleset kernel_lower)
(rule ((= ?__rw0 (Op (Recip ?v55_shape ?v55_strides ?v55_out_strides) ?__inputs)) (= ?__dt (dtype (Op (Recip ?v55_shape ?v55_strides ?v55_out_strides) ?__inputs)))) ((union ?__rw0 (Op (KernelRecip ?v55_shape ?v55_strides ?v55_out_strides ?__dt) ?__inputs))) :ruleset kernel_lower)
(rule ((= ?__rw0 (Op (Sqrt ?v56_shape ?v56_strides ?v56_out_strides) ?__inputs)) (= ?__dt (dtype (Op (Sqrt ?v56_shape ?v56_strides ?v56_out_strides) ?__inputs)))) ((union ?__rw0 (Op (KernelSqrt ?v56_shape ?v56_strides ?v56_out_strides ?__dt) ?__inputs))) :ruleset kernel_lower)
(rule ((= ?__rw0 (Op (Constant ?v57_value) ?__inputs))) ((union ?__rw0 (Op (KernelConstant ?v57_value) ?__inputs)) (set (dtype (Op (KernelConstant ?v57_value) ?__inputs)) (F32))) :ruleset kernel_lower)
(rule ((= ?__rw0 (Op (Cast ?v58_size ?v58_dtype) (ICons ?__inp (INil)))) (= ?__in_dt (dtype ?__inp))) ((union ?__rw0 (Op (KernelCast ?v58_size ?__in_dt ?v58_dtype) (ICons ?__inp (INil))))) :ruleset kernel_lower)
(rule
(
(= ?gather (Op (Gather ?idx_shape ?idx_stride ?embed_shape ?embed_stride) (ICons ?indices (ICons ?embed_table (INil)))))
(= (len ?idx_shape) 2)
(= ?indices (Op (Add ?add_shape ?mul_stride ?iota_stride ?add_out_stride) (ICons ?mul_result (ICons ?iota_result (INil)))))
(= ?mul_result (Op (Mul ?mul_shape ?token_cast_stride ?mul_const_stride ?mul_out_stride) (ICons ?token_ids_cast (ICons ?mul_const (INil)))))
(= ?token_ids_cast (Op (Cast ?cast_size ?cast_dtype) (ICons ?token_ids (INil))))
(= ?embed_dim (nth_from_end ?embed_shape 0))
(= ?batch_shape (RemoveNthFromEnd ?idx_shape 0))
(= ?out_stride_batch (RemoveNthFromEnd ?add_out_stride 0))
)
(
(let ?ke (Op (KernelEmbed ?batch_shape ?token_cast_stride ?out_stride_batch ?embed_dim) (ICons ?token_ids_cast (ICons ?embed_table (INil)))))
(union ?gather ?ke)
(set (dtype ?ke) (F32))
)
:ruleset kernel_specialize
:name "kernel embed with cast mul"
)
(rule
(
(= ?gather (Op (Gather ?idx_shape ?idx_stride ?embed_shape ?embed_stride) (ICons ?indices (ICons ?embed_table (INil)))))
(= (len ?idx_shape) 2)
(= ?indices (Op (Add ?add_shape ?iota_stride ?mul_stride ?add_out_stride) (ICons ?iota_result (ICons ?mul_result (INil)))))
(= ?mul_result (Op (Mul ?mul_shape ?token_cast_stride ?mul_const_stride ?mul_out_stride) (ICons ?token_ids_cast (ICons ?mul_const (INil)))))
(= ?token_ids_cast (Op (Cast ?cast_size ?cast_dtype) (ICons ?token_ids (INil))))
(= ?embed_dim (nth_from_end ?embed_shape 0))
(= ?batch_shape (RemoveNthFromEnd ?idx_shape 0))
(= ?out_stride_batch (RemoveNthFromEnd ?add_out_stride 0))
)
(
(let ?ke (Op (KernelEmbed ?batch_shape ?token_cast_stride ?out_stride_batch ?embed_dim) (ICons ?token_ids_cast (ICons ?embed_table (INil)))))
(union ?gather ?ke)
(set (dtype ?ke) (F32))
)
:ruleset kernel_specialize
:name "kernel embed with cast mul reversed"
)
(rule
(
(= ?gather (Op (Gather ?idx_shape ?idx_stride ?embed_shape ?embed_stride) (ICons ?indices (ICons ?embed_table (INil)))))
(= (len ?idx_shape) 2)
(= ?indices (Op (Add ?add_shape ?mul_stride ?iota_stride ?add_out_stride) (ICons ?mul_result (ICons ?iota_result (INil)))))
(= ?mul_result (Op (Mul ?mul_shape ?token_stride ?mul_const_stride ?mul_out_stride) (ICons ?token_ids (ICons ?mul_const (INil)))))
(= ?embed_dim (nth_from_end ?embed_shape 0))
(= ?batch_shape (RemoveNthFromEnd ?idx_shape 0))
(= ?out_stride_batch (RemoveNthFromEnd ?add_out_stride 0))
)
(
(let ?ke (Op (KernelEmbed ?batch_shape ?token_stride ?out_stride_batch ?embed_dim) (ICons ?token_ids (ICons ?embed_table (INil)))))
(union ?gather ?ke)
(set (dtype ?ke) (F32))
)
:ruleset kernel_specialize
:name "kernel embed with mul"
)
(rule
(
(= ?gather (Op (Gather ?idx_shape ?idx_stride ?embed_shape ?embed_stride) (ICons ?indices (ICons ?embed_table (INil)))))
(= (len ?idx_shape) 2)
(= ?indices (Op (Add ?add_shape ?iota_stride ?mul_stride ?add_out_stride) (ICons ?iota_result (ICons ?mul_result (INil)))))
(= ?mul_result (Op (Mul ?mul_shape ?token_stride ?mul_const_stride ?mul_out_stride) (ICons ?token_ids (ICons ?mul_const (INil)))))
(= ?embed_dim (nth_from_end ?embed_shape 0))
(= ?batch_shape (RemoveNthFromEnd ?idx_shape 0))
(= ?out_stride_batch (RemoveNthFromEnd ?add_out_stride 0))
)
(
(let ?ke (Op (KernelEmbed ?batch_shape ?token_stride ?out_stride_batch ?embed_dim) (ICons ?token_ids (ICons ?embed_table (INil)))))
(union ?gather ?ke)
(set (dtype ?ke) (F32))
)
:ruleset kernel_specialize
:name "kernel embed with mul reversed"
)
(rule
(
; Match Mul node (broadcast multiply)
(= ?mul (Op (Mul ?mul_shape ?a_stride ?b_stride ?mul_out_stride) (ICons ?a (ICons ?b (INil)))))
; Match Sum that reduces the Mul (k dimension)
(= ?sum (Op (Sum ?out_shape ?k ?sum_in_stride ?k_stride ?sum_out_stride) (ICons ?mul (INil))))
; Output shape must have 3+ dimensions (batched)
(= ?out_shape (ECons ?batch_or_d0 (ECons ?d1 (ECons ?d2 ?rest))))
; k_stride must be contiguous
(= ?k_stride (MIter))
; Get A's k-dimension stride (second from end in Mul's a_stride)
(= ?a_k_stride (nth_from_end ?a_stride 1))
; Get B's k-dimension stride (second from end in Mul's b_stride)
(= ?b_k_stride (nth_from_end ?b_stride 1))
; A's k stride must be contiguous (row-major A)
(= ?a_k_stride (MIter))
; B's k stride must be contiguous (col-major B)
(= ?b_k_stride (MIter))
; Must be F32
(= (F32) (dtype ?a))
(= (F32) (dtype ?b))
)
(
; Remove the k-dimension from A strides for the kernel
(let ?a_kern_stride (RemoveNthFromEnd ?a_stride 1))
; Remove the k-dimension from B strides
(let ?b_kern_stride (RemoveNthFromEnd ?b_stride 1))
(let ?bmv (Op (KernelBatchMatVec
?out_shape ?k
?a_kern_stride ?a_k_stride
?b_kern_stride ?b_k_stride
?sum_out_stride (F32)) (ICons ?a (ICons ?b (INil)))))
(union ?sum ?bmv)
(set (dtype ?bmv) (F32))
)
:ruleset matmul_backend
:name "batch mat-vec"
)
(rule
(
; Match Mul node (broadcast multiply)
(= ?mul (Op (Mul ?mul_shape ?a_stride ?b_stride ?mul_out_stride) (ICons ?a (ICons ?b (INil)))))
; Match Sum that reduces the Mul (k dimension)
(= ?sum (Op (Sum ?out_shape ?k ?sum_in_stride ?k_stride ?sum_out_stride) (ICons ?mul (INil))))
; Output shape must have 3+ dimensions (batched)
(= ?out_shape (ECons ?batch_or_d0 (ECons ?d1 (ECons ?d2 ?rest))))
; k_stride must be contiguous in the Sum output
(= ?k_stride (MIter))
; K must be > 1 (K=1 is a degenerate outer product, not a real matmul)
(!= ?k (MNum 1))
; Get A's and B's k-dimension strides (no contiguity requirement)
(= ?a_k_stride (nth_from_end ?a_stride 1))
(= ?b_k_stride (nth_from_end ?b_stride 1))
; One of A's non-k strides must be 0 (broadcast along n)
(= (MNum 0) (nth_from_end ?a_stride 0))
; One of B's non-k strides must be 0 (broadcast along m)
(= (MNum 0) (nth_from_end ?b_stride 2))
; Must be F32
(= (F32) (dtype ?a))
(= (F32) (dtype ?b))
)
(
(let ?a_kern_stride (RemoveNthFromEnd ?a_stride 1))
(let ?b_kern_stride (RemoveNthFromEnd ?b_stride 1))
(let ?bmm (Op (KernelBatchMatMul
?out_shape ?k
?a_kern_stride ?a_k_stride
?b_kern_stride ?b_k_stride
?sum_out_stride (F32)) (ICons ?a (ICons ?b (INil)))))
(union ?sum ?bmm)
(set (dtype ?bmm) (F32))
)
:ruleset matmul_backend
:name "batch matmul"
)
(relation consumed_buffer_ilist_contains (IList IR))
(rule
((= ?list (ICons ?head ?tail)))
((consumed_buffer_ilist_contains ?list ?head))
:ruleset cleanup
:name "consumed-buffer-ilist-contains-head"
)
(rule
((= ?list (ICons ?head ?tail))
(consumed_buffer_ilist_contains ?tail ?item))
((consumed_buffer_ilist_contains ?list ?item))
:ruleset cleanup
:name "consumed-buffer-ilist-contains-tail"
)
(rule
(
(= ?scatter (Op (KernelScatter ?ds ?dst ?is ?istr ?ss ?os ?dt)
(ICons ?dest (ICons ?indexes (ICons ?src (INil))))))
(= ?dst ?os)
(= ?dty (dtype ?src))
)
(
(let ?consumed (ConsumedBuffer ?dest))
(let ?nocopy (Op (KernelScatterNoCopy ?ds ?dst ?is ?istr ?ss ?os ?dt)
(ICons ?consumed (ICons ?indexes (ICons ?src (INil))))))
(union ?scatter ?nocopy)
(set (dtype ?nocopy) ?dty)
)
:ruleset buffer_reuse
:name "scatter to scatter-no-copy"
)
(rule
((= ?cb (ConsumedBuffer ?a))
(= ?dt (dtype ?a)))
((set (dtype ?cb) ?dt))
:ruleset dtype_prop
:name "consumed-buffer-dtype"
)
(rule
((= ?cb (ConsumedBuffer ?a))
(= ?op1 (Op ?k1 ?ilist1))
(consumed_buffer_ilist_contains ?ilist1 ?cb)
(= ?op2 (Op ?k2 ?ilist2))
(!= ?op1 ?op2)
(consumed_buffer_ilist_contains ?ilist2 ?a))
((delete (ConsumedBuffer ?a)))
:ruleset cleanup
:name "consumed-buffer-cleanup-shared-op-use"
)
(rule
((= ?cb (ConsumedBuffer ?dest))
(= ?scatter (Op (KernelScatter ?ds ?dst ?is ?istr ?ss ?os ?dt)
(ICons ?dest (ICons ?indexes (ICons ?src (INil))))))
(= ?nocopy (Op (KernelScatterNoCopy ?ds ?dst ?is ?istr ?ss ?os ?dt)
(ICons ?cb (ICons ?indexes (ICons ?src (INil)))))))
((delete (Op (KernelScatter ?ds ?dst ?is ?istr ?ss ?os ?dt)
(ICons ?dest (ICons ?indexes (ICons ?src (INil)))))))
:ruleset post_cleanup
:name "scatter-no-copy-dominates-valid-consumed-buffer"
)
(rule
((= ?cb (ConsumedBuffer ?a)))
((union ?cb ?a)
(delete (ConsumedBuffer ?a)))
:ruleset base_cleanup
:name "consumed-buffer-resolve"
)
(rule ((= ?__rw0 (Op (Softmax ?v59_shape ?v59_in_strides ?v59_out_strides ?v59_reduce_dim ?v59_reduce_stride) ?__inputs)) (= ?__dt (dtype (Op (Softmax ?v59_shape ?v59_in_strides ?v59_out_strides ?v59_reduce_dim ?v59_reduce_stride) ?__inputs)))) ((union ?__rw0 (Op (KernelSoftmax ?v59_shape ?v59_in_strides ?v59_out_strides ?v59_reduce_dim ?v59_reduce_stride ?__dt) ?__inputs))) :ruleset kernel_lower)
(rule
(
(= ?sm (Op (Softmax ?shape ?in_strides ?out_strides ?reduce_dim ?reduce_stride) ?inputs))
)
(
(let ?ksm (Op (KernelSoftmax ?shape ?in_strides ?out_strides ?reduce_dim ?reduce_stride (F32)) ?inputs))
(union ?sm ?ksm)
(set (dtype ?ksm) (F32))
)
:ruleset kernel_lower
:name "softmax-to-kernel-f32"
)
(rule
(
(= ?mul (Op (Mul ?shape ?x_stride ?const_stride ?inter_stride) (ICons ?x (ICons ?exp_const (INil)))))
(= ?exp2 (Op (Exp2 ?shape ?inter_stride ?out_stride) (ICons ?mul (INil))))
(= ?dt (dtype ?x))
(= ?cv (Op (Constant ?val) (INil)))
(= ?exp_const ?cv)
(> ?val 1.44)
(< ?val 1.45)
)
(
(let ?kexp (Op (KernelExp ?shape ?x_stride ?out_stride ?dt) (ICons ?x (INil))))
(union ?exp2 ?kexp)
(set (dtype ?kexp) ?dt)
)
:ruleset direct_kernel
:name "direct-exp-fusion"
)
(datatype*
(KernelSigmoidScaledState
(MkKernelSigmoidScaledState IR EList EList DType)
)
)
(function kernel_sigmoid_scaled (IR) KernelSigmoidScaledState :merge new)
(rule
(
(= ?neg1 (Op (Constant ?nv) (INil)))
(< ?nv -0.99)
(> ?nv -1.01)
(= ?neg_x (Op (Mul ?shape ?x_stride ?neg_stride ?neg_out_stride) (ICons ?x (ICons ?neg1 (INil)))))
(= ?log2e (Op (Constant ?lv) (INil)))
(> ?lv 1.44)
(< ?lv 1.45)
(= ?scaled (Op (Mul ?shape ?neg_out_stride ?log2e_stride ?scaled_stride) (ICons ?neg_x (ICons ?log2e (INil)))))
(= ?dt (dtype ?x))
)
(
(set (kernel_sigmoid_scaled ?scaled)
(MkKernelSigmoidScaledState ?x ?shape ?x_stride ?dt))
)
:ruleset direct_kernel
:name "direct-sigmoid-scaled-marker"
)
(rule
(
(= ?scaled_state (kernel_sigmoid_scaled ?scaled))
(= ?scaled_state (MkKernelSigmoidScaledState ?x ?shape ?x_stride ?dt))
(= ?exp2 (Op (Exp2 ?shape ?scaled_stride ?exp_stride) (ICons ?scaled (INil))))
(= ?one (Op (Constant ?ov) (INil)))
(> ?ov 0.99)
(< ?ov 1.01)
(= ?plus_one (Op (Add ?shape ?exp_stride ?one_stride ?add_stride) (ICons ?exp2 (ICons ?one (INil)))))
(= ?sig_out (Op (Recip ?shape ?add_stride ?out_stride) (ICons ?plus_one (INil))))
)
(
(let ?ksig (Op (KernelSigmoid ?shape ?x_stride ?out_stride ?dt) (ICons ?x (INil))))
(union ?sig_out ?ksig)
(set (dtype ?ksig) ?dt)
)
:ruleset direct_kernel
:name "direct-sigmoid-fusion"
)
(rule (
(= ?u1 (Op (KernelSin ?shape ?s ?s ?dt) (ICons ?x (INil))))
(= ?u2 (Op (KernelSin ?shape ?s ?s ?dt) (ICons ?u1 (INil))))
) (
(let ?fs (Op (FusionStart ?shape ?s ?dt) (ICons ?x (INil))))
(let ?fu1 (Op (FusedSin ?shape ?s ?s ?dt) (ICons ?fs (INil))))
(let ?fu2 (Op (FusedSin ?shape ?s ?s ?dt) (ICons ?fu1 (INil))))
(let ?fe (Op (FusionEnd ?shape ?s ?dt) (ICons ?fu2 (INil))))
(union ?u2 ?fe)
) :ruleset fusion_pair :name "pair-fuse-U-U-KernelSin-KernelSin")
(rule (
(= ?u1 (Op (KernelSin ?shape ?s ?s ?dt) (ICons ?x (INil))))
(= ?u2 (Op (KernelSqrt ?shape ?s ?s ?dt) (ICons ?u1 (INil))))
) (
(let ?fs (Op (FusionStart ?shape ?s ?dt) (ICons ?x (INil))))
(let ?fu1 (Op (FusedSin ?shape ?s ?s ?dt) (ICons ?fs (INil))))
(let ?fu2 (Op (FusedSqrt ?shape ?s ?s ?dt) (ICons ?fu1 (INil))))
(let ?fe (Op (FusionEnd ?shape ?s ?dt) (ICons ?fu2 (INil))))
(union ?u2 ?fe)
) :ruleset fusion_pair :name "pair-fuse-U-U-KernelSin-KernelSqrt")
(rule (
(= ?u1 (Op (KernelSin ?shape ?s ?s ?dt) (ICons ?x (INil))))
(= ?u2 (Op (KernelExp ?shape ?s ?s ?dt) (ICons ?u1 (INil))))
) (
(let ?fs (Op (FusionStart ?shape ?s ?dt) (ICons ?x (INil))))
(let ?fu1 (Op (FusedSin ?shape ?s ?s ?dt) (ICons ?fs (INil))))
(let ?fu2 (Op (FusedExp ?shape ?s ?s ?dt) (ICons ?fu1 (INil))))
(let ?fe (Op (FusionEnd ?shape ?s ?dt) (ICons ?fu2 (INil))))
(union ?u2 ?fe)
) :ruleset fusion_pair :name "pair-fuse-U-U-KernelSin-KernelExp")
(rule (
(= ?u1 (Op (KernelSin ?shape ?s ?s ?dt) (ICons ?x (INil))))
(= ?u2 (Op (KernelExp2 ?shape ?s ?s ?dt) (ICons ?u1 (INil))))
) (
(let ?fs (Op (FusionStart ?shape ?s ?dt) (ICons ?x (INil))))
(let ?fu1 (Op (FusedSin ?shape ?s ?s ?dt) (ICons ?fs (INil))))
(let ?fu2 (Op (FusedExp2 ?shape ?s ?s ?dt) (ICons ?fu1 (INil))))
(let ?fe (Op (FusionEnd ?shape ?s ?dt) (ICons ?fu2 (INil))))
(union ?u2 ?fe)
) :ruleset fusion_pair :name "pair-fuse-U-U-KernelSin-KernelExp2")
(rule (
(= ?u1 (Op (KernelSin ?shape ?s ?s ?dt) (ICons ?x (INil))))
(= ?u2 (Op (KernelLog2 ?shape ?s ?s ?dt) (ICons ?u1 (INil))))
) (
(let ?fs (Op (FusionStart ?shape ?s ?dt) (ICons ?x (INil))))
(let ?fu1 (Op (FusedSin ?shape ?s ?s ?dt) (ICons ?fs (INil))))
(let ?fu2 (Op (FusedLog2 ?shape ?s ?s ?dt) (ICons ?fu1 (INil))))
(let ?fe (Op (FusionEnd ?shape ?s ?dt) (ICons ?fu2 (INil))))
(union ?u2 ?fe)
) :ruleset fusion_pair :name "pair-fuse-U-U-KernelSin-KernelLog2")
(rule (
(= ?u1 (Op (KernelSin ?shape ?s ?s ?dt) (ICons ?x (INil))))
(= ?u2 (Op (KernelRecip ?shape ?s ?s ?dt) (ICons ?u1 (INil))))
) (
(let ?fs (Op (FusionStart ?shape ?s ?dt) (ICons ?x (INil))))
(let ?fu1 (Op (FusedSin ?shape ?s ?s ?dt) (ICons ?fs (INil))))
(let ?fu2 (Op (FusedRecip ?shape ?s ?s ?dt) (ICons ?fu1 (INil))))
(let ?fe (Op (FusionEnd ?shape ?s ?dt) (ICons ?fu2 (INil))))
(union ?u2 ?fe)
) :ruleset fusion_pair :name "pair-fuse-U-U-KernelSin-KernelRecip")
(rule (
(= ?u1 (Op (KernelSqrt ?shape ?s ?s ?dt) (ICons ?x (INil))))
(= ?u2 (Op (KernelSin ?shape ?s ?s ?dt) (ICons ?u1 (INil))))
) (
(let ?fs (Op (FusionStart ?shape ?s ?dt) (ICons ?x (INil))))
(let ?fu1 (Op (FusedSqrt ?shape ?s ?s ?dt) (ICons ?fs (INil))))
(let ?fu2 (Op (FusedSin ?shape ?s ?s ?dt) (ICons ?fu1 (INil))))
(let ?fe (Op (FusionEnd ?shape ?s ?dt) (ICons ?fu2 (INil))))
(union ?u2 ?fe)
) :ruleset fusion_pair :name "pair-fuse-U-U-KernelSqrt-KernelSin")
(rule (
(= ?u1 (Op (KernelSqrt ?shape ?s ?s ?dt) (ICons ?x (INil))))
(= ?u2 (Op (KernelSqrt ?shape ?s ?s ?dt) (ICons ?u1 (INil))))
) (
(let ?fs (Op (FusionStart ?shape ?s ?dt) (ICons ?x (INil))))
(let ?fu1 (Op (FusedSqrt ?shape ?s ?s ?dt) (ICons ?fs (INil))))
(let ?fu2 (Op (FusedSqrt ?shape ?s ?s ?dt) (ICons ?fu1 (INil))))
(let ?fe (Op (FusionEnd ?shape ?s ?dt) (ICons ?fu2 (INil))))
(union ?u2 ?fe)
) :ruleset fusion_pair :name "pair-fuse-U-U-KernelSqrt-KernelSqrt")
(rule (
(= ?u1 (Op (KernelSqrt ?shape ?s ?s ?dt) (ICons ?x (INil))))
(= ?u2 (Op (KernelExp ?shape ?s ?s ?dt) (ICons ?u1 (INil))))
) (
(let ?fs (Op (FusionStart ?shape ?s ?dt) (ICons ?x (INil))))
(let ?fu1 (Op (FusedSqrt ?shape ?s ?s ?dt) (ICons ?fs (INil))))
(let ?fu2 (Op (FusedExp ?shape ?s ?s ?dt) (ICons ?fu1 (INil))))
(let ?fe (Op (FusionEnd ?shape ?s ?dt) (ICons ?fu2 (INil))))
(union ?u2 ?fe)
) :ruleset fusion_pair :name "pair-fuse-U-U-KernelSqrt-KernelExp")
(rule (
(= ?u1 (Op (KernelSqrt ?shape ?s ?s ?dt) (ICons ?x (INil))))
(= ?u2 (Op (KernelExp2 ?shape ?s ?s ?dt) (ICons ?u1 (INil))))
) (
(let ?fs (Op (FusionStart ?shape ?s ?dt) (ICons ?x (INil))))
(let ?fu1 (Op (FusedSqrt ?shape ?s ?s ?dt) (ICons ?fs (INil))))
(let ?fu2 (Op (FusedExp2 ?shape ?s ?s ?dt) (ICons ?fu1 (INil))))
(let ?fe (Op (FusionEnd ?shape ?s ?dt) (ICons ?fu2 (INil))))
(union ?u2 ?fe)
) :ruleset fusion_pair :name "pair-fuse-U-U-KernelSqrt-KernelExp2")
(rule (
(= ?u1 (Op (KernelSqrt ?shape ?s ?s ?dt) (ICons ?x (INil))))
(= ?u2 (Op (KernelLog2 ?shape ?s ?s ?dt) (ICons ?u1 (INil))))
) (
(let ?fs (Op (FusionStart ?shape ?s ?dt) (ICons ?x (INil))))
(let ?fu1 (Op (FusedSqrt ?shape ?s ?s ?dt) (ICons ?fs (INil))))
(let ?fu2 (Op (FusedLog2 ?shape ?s ?s ?dt) (ICons ?fu1 (INil))))
(let ?fe (Op (FusionEnd ?shape ?s ?dt) (ICons ?fu2 (INil))))
(union ?u2 ?fe)
) :ruleset fusion_pair :name "pair-fuse-U-U-KernelSqrt-KernelLog2")
(rule (
(= ?u1 (Op (KernelSqrt ?shape ?s ?s ?dt) (ICons ?x (INil))))
(= ?u2 (Op (KernelRecip ?shape ?s ?s ?dt) (ICons ?u1 (INil))))
) (
(let ?fs (Op (FusionStart ?shape ?s ?dt) (ICons ?x (INil))))
(let ?fu1 (Op (FusedSqrt ?shape ?s ?s ?dt) (ICons ?fs (INil))))
(let ?fu2 (Op (FusedRecip ?shape ?s ?s ?dt) (ICons ?fu1 (INil))))
(let ?fe (Op (FusionEnd ?shape ?s ?dt) (ICons ?fu2 (INil))))
(union ?u2 ?fe)
) :ruleset fusion_pair :name "pair-fuse-U-U-KernelSqrt-KernelRecip")
(rule (
(= ?u1 (Op (KernelExp ?shape ?s ?s ?dt) (ICons ?x (INil))))
(= ?u2 (Op (KernelSin ?shape ?s ?s ?dt) (ICons ?u1 (INil))))
) (
(let ?fs (Op (FusionStart ?shape ?s ?dt) (ICons ?x (INil))))
(let ?fu1 (Op (FusedExp ?shape ?s ?s ?dt) (ICons ?fs (INil))))
(let ?fu2 (Op (FusedSin ?shape ?s ?s ?dt) (ICons ?fu1 (INil))))
(let ?fe (Op (FusionEnd ?shape ?s ?dt) (ICons ?fu2 (INil))))
(union ?u2 ?fe)
) :ruleset fusion_pair :name "pair-fuse-U-U-KernelExp-KernelSin")
(rule (
(= ?u1 (Op (KernelExp ?shape ?s ?s ?dt) (ICons ?x (INil))))
(= ?u2 (Op (KernelSqrt ?shape ?s ?s ?dt) (ICons ?u1 (INil))))
) (
(let ?fs (Op (FusionStart ?shape ?s ?dt) (ICons ?x (INil))))
(let ?fu1 (Op (FusedExp ?shape ?s ?s ?dt) (ICons ?fs (INil))))
(let ?fu2 (Op (FusedSqrt ?shape ?s ?s ?dt) (ICons ?fu1 (INil))))
(let ?fe (Op (FusionEnd ?shape ?s ?dt) (ICons ?fu2 (INil))))
(union ?u2 ?fe)
) :ruleset fusion_pair :name "pair-fuse-U-U-KernelExp-KernelSqrt")
(rule (
(= ?u1 (Op (KernelExp ?shape ?s ?s ?dt) (ICons ?x (INil))))
(= ?u2 (Op (KernelExp ?shape ?s ?s ?dt) (ICons ?u1 (INil))))
) (
(let ?fs (Op (FusionStart ?shape ?s ?dt) (ICons ?x (INil))))
(let ?fu1 (Op (FusedExp ?shape ?s ?s ?dt) (ICons ?fs (INil))))
(let ?fu2 (Op (FusedExp ?shape ?s ?s ?dt) (ICons ?fu1 (INil))))
(let ?fe (Op (FusionEnd ?shape ?s ?dt) (ICons ?fu2 (INil))))
(union ?u2 ?fe)
) :ruleset fusion_pair :name "pair-fuse-U-U-KernelExp-KernelExp")
(rule (
(= ?u1 (Op (KernelExp ?shape ?s ?s ?dt) (ICons ?x (INil))))
(= ?u2 (Op (KernelExp2 ?shape ?s ?s ?dt) (ICons ?u1 (INil))))
) (
(let ?fs (Op (FusionStart ?shape ?s ?dt) (ICons ?x (INil))))
(let ?fu1 (Op (FusedExp ?shape ?s ?s ?dt) (ICons ?fs (INil))))
(let ?fu2 (Op (FusedExp2 ?shape ?s ?s ?dt) (ICons ?fu1 (INil))))
(let ?fe (Op (FusionEnd ?shape ?s ?dt) (ICons ?fu2 (INil))))
(union ?u2 ?fe)
) :ruleset fusion_pair :name "pair-fuse-U-U-KernelExp-KernelExp2")
(rule (
(= ?u1 (Op (KernelExp ?shape ?s ?s ?dt) (ICons ?x (INil))))
(= ?u2 (Op (KernelLog2 ?shape ?s ?s ?dt) (ICons ?u1 (INil))))
) (
(let ?fs (Op (FusionStart ?shape ?s ?dt) (ICons ?x (INil))))
(let ?fu1 (Op (FusedExp ?shape ?s ?s ?dt) (ICons ?fs (INil))))
(let ?fu2 (Op (FusedLog2 ?shape ?s ?s ?dt) (ICons ?fu1 (INil))))
(let ?fe (Op (FusionEnd ?shape ?s ?dt) (ICons ?fu2 (INil))))
(union ?u2 ?fe)
) :ruleset fusion_pair :name "pair-fuse-U-U-KernelExp-KernelLog2")
(rule (
(= ?u1 (Op (KernelExp ?shape ?s ?s ?dt) (ICons ?x (INil))))
(= ?u2 (Op (KernelRecip ?shape ?s ?s ?dt) (ICons ?u1 (INil))))
) (
(let ?fs (Op (FusionStart ?shape ?s ?dt) (ICons ?x (INil))))
(let ?fu1 (Op (FusedExp ?shape ?s ?s ?dt) (ICons ?fs (INil))))
(let ?fu2 (Op (FusedRecip ?shape ?s ?s ?dt) (ICons ?fu1 (INil))))
(let ?fe (Op (FusionEnd ?shape ?s ?dt) (ICons ?fu2 (INil))))
(union ?u2 ?fe)
) :ruleset fusion_pair :name "pair-fuse-U-U-KernelExp-KernelRecip")
(rule (
(= ?u1 (Op (KernelExp2 ?shape ?s ?s ?dt) (ICons ?x (INil))))
(= ?u2 (Op (KernelSin ?shape ?s ?s ?dt) (ICons ?u1 (INil))))
) (
(let ?fs (Op (FusionStart ?shape ?s ?dt) (ICons ?x (INil))))
(let ?fu1 (Op (FusedExp2 ?shape ?s ?s ?dt) (ICons ?fs (INil))))
(let ?fu2 (Op (FusedSin ?shape ?s ?s ?dt) (ICons ?fu1 (INil))))
(let ?fe (Op (FusionEnd ?shape ?s ?dt) (ICons ?fu2 (INil))))
(union ?u2 ?fe)
) :ruleset fusion_pair :name "pair-fuse-U-U-KernelExp2-KernelSin")
(rule (
(= ?u1 (Op (KernelExp2 ?shape ?s ?s ?dt) (ICons ?x (INil))))
(= ?u2 (Op (KernelSqrt ?shape ?s ?s ?dt) (ICons ?u1 (INil))))
) (
(let ?fs (Op (FusionStart ?shape ?s ?dt) (ICons ?x (INil))))
(let ?fu1 (Op (FusedExp2 ?shape ?s ?s ?dt) (ICons ?fs (INil))))
(let ?fu2 (Op (FusedSqrt ?shape ?s ?s ?dt) (ICons ?fu1 (INil))))
(let ?fe (Op (FusionEnd ?shape ?s ?dt) (ICons ?fu2 (INil))))
(union ?u2 ?fe)
) :ruleset fusion_pair :name "pair-fuse-U-U-KernelExp2-KernelSqrt")
(rule (
(= ?u1 (Op (KernelExp2 ?shape ?s ?s ?dt) (ICons ?x (INil))))
(= ?u2 (Op (KernelExp ?shape ?s ?s ?dt) (ICons ?u1 (INil))))
) (
(let ?fs (Op (FusionStart ?shape ?s ?dt) (ICons ?x (INil))))
(let ?fu1 (Op (FusedExp2 ?shape ?s ?s ?dt) (ICons ?fs (INil))))
(let ?fu2 (Op (FusedExp ?shape ?s ?s ?dt) (ICons ?fu1 (INil))))
(let ?fe (Op (FusionEnd ?shape ?s ?dt) (ICons ?fu2 (INil))))
(union ?u2 ?fe)
) :ruleset fusion_pair :name "pair-fuse-U-U-KernelExp2-KernelExp")
(rule (
(= ?u1 (Op (KernelExp2 ?shape ?s ?s ?dt) (ICons ?x (INil))))
(= ?u2 (Op (KernelExp2 ?shape ?s ?s ?dt) (ICons ?u1 (INil))))
) (
(let ?fs (Op (FusionStart ?shape ?s ?dt) (ICons ?x (INil))))
(let ?fu1 (Op (FusedExp2 ?shape ?s ?s ?dt) (ICons ?fs (INil))))
(let ?fu2 (Op (FusedExp2 ?shape ?s ?s ?dt) (ICons ?fu1 (INil))))
(let ?fe (Op (FusionEnd ?shape ?s ?dt) (ICons ?fu2 (INil))))
(union ?u2 ?fe)
) :ruleset fusion_pair :name "pair-fuse-U-U-KernelExp2-KernelExp2")
(rule (
(= ?u1 (Op (KernelExp2 ?shape ?s ?s ?dt) (ICons ?x (INil))))
(= ?u2 (Op (KernelLog2 ?shape ?s ?s ?dt) (ICons ?u1 (INil))))
) (
(let ?fs (Op (FusionStart ?shape ?s ?dt) (ICons ?x (INil))))
(let ?fu1 (Op (FusedExp2 ?shape ?s ?s ?dt) (ICons ?fs (INil))))
(let ?fu2 (Op (FusedLog2 ?shape ?s ?s ?dt) (ICons ?fu1 (INil))))
(let ?fe (Op (FusionEnd ?shape ?s ?dt) (ICons ?fu2 (INil))))
(union ?u2 ?fe)
) :ruleset fusion_pair :name "pair-fuse-U-U-KernelExp2-KernelLog2")
(rule (
(= ?u1 (Op (KernelExp2 ?shape ?s ?s ?dt) (ICons ?x (INil))))
(= ?u2 (Op (KernelRecip ?shape ?s ?s ?dt) (ICons ?u1 (INil))))
) (
(let ?fs (Op (FusionStart ?shape ?s ?dt) (ICons ?x (INil))))
(let ?fu1 (Op (FusedExp2 ?shape ?s ?s ?dt) (ICons ?fs (INil))))
(let ?fu2 (Op (FusedRecip ?shape ?s ?s ?dt) (ICons ?fu1 (INil))))
(let ?fe (Op (FusionEnd ?shape ?s ?dt) (ICons ?fu2 (INil))))
(union ?u2 ?fe)
) :ruleset fusion_pair :name "pair-fuse-U-U-KernelExp2-KernelRecip")
(rule (
(= ?u1 (Op (KernelLog2 ?shape ?s ?s ?dt) (ICons ?x (INil))))
(= ?u2 (Op (KernelSin ?shape ?s ?s ?dt) (ICons ?u1 (INil))))
) (
(let ?fs (Op (FusionStart ?shape ?s ?dt) (ICons ?x (INil))))
(let ?fu1 (Op (FusedLog2 ?shape ?s ?s ?dt) (ICons ?fs (INil))))
(let ?fu2 (Op (FusedSin ?shape ?s ?s ?dt) (ICons ?fu1 (INil))))
(let ?fe (Op (FusionEnd ?shape ?s ?dt) (ICons ?fu2 (INil))))
(union ?u2 ?fe)
) :ruleset fusion_pair :name "pair-fuse-U-U-KernelLog2-KernelSin")
(rule (
(= ?u1 (Op (KernelLog2 ?shape ?s ?s ?dt) (ICons ?x (INil))))
(= ?u2 (Op (KernelSqrt ?shape ?s ?s ?dt) (ICons ?u1 (INil))))
) (
(let ?fs (Op (FusionStart ?shape ?s ?dt) (ICons ?x (INil))))
(let ?fu1 (Op (FusedLog2 ?shape ?s ?s ?dt) (ICons ?fs (INil))))
(let ?fu2 (Op (FusedSqrt ?shape ?s ?s ?dt) (ICons ?fu1 (INil))))
(let ?fe (Op (FusionEnd ?shape ?s ?dt) (ICons ?fu2 (INil))))
(union ?u2 ?fe)
) :ruleset fusion_pair :name "pair-fuse-U-U-KernelLog2-KernelSqrt")
(rule (
(= ?u1 (Op (KernelLog2 ?shape ?s ?s ?dt) (ICons ?x (INil))))
(= ?u2 (Op (KernelExp ?shape ?s ?s ?dt) (ICons ?u1 (INil))))
) (
(let ?fs (Op (FusionStart ?shape ?s ?dt) (ICons ?x (INil))))
(let ?fu1 (Op (FusedLog2 ?shape ?s ?s ?dt) (ICons ?fs (INil))))
(let ?fu2 (Op (FusedExp ?shape ?s ?s ?dt) (ICons ?fu1 (INil))))
(let ?fe (Op (FusionEnd ?shape ?s ?dt) (ICons ?fu2 (INil))))
(union ?u2 ?fe)
) :ruleset fusion_pair :name "pair-fuse-U-U-KernelLog2-KernelExp")
(rule (
(= ?u1 (Op (KernelLog2 ?shape ?s ?s ?dt) (ICons ?x (INil))))
(= ?u2 (Op (KernelExp2 ?shape ?s ?s ?dt) (ICons ?u1 (INil))))
) (
(let ?fs (Op (FusionStart ?shape ?s ?dt) (ICons ?x (INil))))
(let ?fu1 (Op (FusedLog2 ?shape ?s ?s ?dt) (ICons ?fs (INil))))
(let ?fu2 (Op (FusedExp2 ?shape ?s ?s ?dt) (ICons ?fu1 (INil))))
(let ?fe (Op (FusionEnd ?shape ?s ?dt) (ICons ?fu2 (INil))))
(union ?u2 ?fe)
) :ruleset fusion_pair :name "pair-fuse-U-U-KernelLog2-KernelExp2")
(rule (
(= ?u1 (Op (KernelLog2 ?shape ?s ?s ?dt) (ICons ?x (INil))))
(= ?u2 (Op (KernelLog2 ?shape ?s ?s ?dt) (ICons ?u1 (INil))))
) (
(let ?fs (Op (FusionStart ?shape ?s ?dt) (ICons ?x (INil))))
(let ?fu1 (Op (FusedLog2 ?shape ?s ?s ?dt) (ICons ?fs (INil))))
(let ?fu2 (Op (FusedLog2 ?shape ?s ?s ?dt) (ICons ?fu1 (INil))))
(let ?fe (Op (FusionEnd ?shape ?s ?dt) (ICons ?fu2 (INil))))
(union ?u2 ?fe)
) :ruleset fusion_pair :name "pair-fuse-U-U-KernelLog2-KernelLog2")
(rule (
(= ?u1 (Op (KernelLog2 ?shape ?s ?s ?dt) (ICons ?x (INil))))
(= ?u2 (Op (KernelRecip ?shape ?s ?s ?dt) (ICons ?u1 (INil))))
) (
(let ?fs (Op (FusionStart ?shape ?s ?dt) (ICons ?x (INil))))
(let ?fu1 (Op (FusedLog2 ?shape ?s ?s ?dt) (ICons ?fs (INil))))
(let ?fu2 (Op (FusedRecip ?shape ?s ?s ?dt) (ICons ?fu1 (INil))))
(let ?fe (Op (FusionEnd ?shape ?s ?dt) (ICons ?fu2 (INil))))
(union ?u2 ?fe)
) :ruleset fusion_pair :name "pair-fuse-U-U-KernelLog2-KernelRecip")
(rule (
(= ?u1 (Op (KernelRecip ?shape ?s ?s ?dt) (ICons ?x (INil))))
(= ?u2 (Op (KernelSin ?shape ?s ?s ?dt) (ICons ?u1 (INil))))
) (
(let ?fs (Op (FusionStart ?shape ?s ?dt) (ICons ?x (INil))))
(let ?fu1 (Op (FusedRecip ?shape ?s ?s ?dt) (ICons ?fs (INil))))
(let ?fu2 (Op (FusedSin ?shape ?s ?s ?dt) (ICons ?fu1 (INil))))
(let ?fe (Op (FusionEnd ?shape ?s ?dt) (ICons ?fu2 (INil))))
(union ?u2 ?fe)
) :ruleset fusion_pair :name "pair-fuse-U-U-KernelRecip-KernelSin")
(rule (
(= ?u1 (Op (KernelRecip ?shape ?s ?s ?dt) (ICons ?x (INil))))
(= ?u2 (Op (KernelSqrt ?shape ?s ?s ?dt) (ICons ?u1 (INil))))
) (
(let ?fs (Op (FusionStart ?shape ?s ?dt) (ICons ?x (INil))))
(let ?fu1 (Op (FusedRecip ?shape ?s ?s ?dt) (ICons ?fs (INil))))
(let ?fu2 (Op (FusedSqrt ?shape ?s ?s ?dt) (ICons ?fu1 (INil))))
(let ?fe (Op (FusionEnd ?shape ?s ?dt) (ICons ?fu2 (INil))))
(union ?u2 ?fe)
) :ruleset fusion_pair :name "pair-fuse-U-U-KernelRecip-KernelSqrt")
(rule (
(= ?u1 (Op (KernelRecip ?shape ?s ?s ?dt) (ICons ?x (INil))))
(= ?u2 (Op (KernelExp ?shape ?s ?s ?dt) (ICons ?u1 (INil))))
) (
(let ?fs (Op (FusionStart ?shape ?s ?dt) (ICons ?x (INil))))
(let ?fu1 (Op (FusedRecip ?shape ?s ?s ?dt) (ICons ?fs (INil))))
(let ?fu2 (Op (FusedExp ?shape ?s ?s ?dt) (ICons ?fu1 (INil))))
(let ?fe (Op (FusionEnd ?shape ?s ?dt) (ICons ?fu2 (INil))))
(union ?u2 ?fe)
) :ruleset fusion_pair :name "pair-fuse-U-U-KernelRecip-KernelExp")
(rule (
(= ?u1 (Op (KernelRecip ?shape ?s ?s ?dt) (ICons ?x (INil))))
(= ?u2 (Op (KernelExp2 ?shape ?s ?s ?dt) (ICons ?u1 (INil))))
) (
(let ?fs (Op (FusionStart ?shape ?s ?dt) (ICons ?x (INil))))
(let ?fu1 (Op (FusedRecip ?shape ?s ?s ?dt) (ICons ?fs (INil))))
(let ?fu2 (Op (FusedExp2 ?shape ?s ?s ?dt) (ICons ?fu1 (INil))))
(let ?fe (Op (FusionEnd ?shape ?s ?dt) (ICons ?fu2 (INil))))
(union ?u2 ?fe)
) :ruleset fusion_pair :name "pair-fuse-U-U-KernelRecip-KernelExp2")
(rule (
(= ?u1 (Op (KernelRecip ?shape ?s ?s ?dt) (ICons ?x (INil))))
(= ?u2 (Op (KernelLog2 ?shape ?s ?s ?dt) (ICons ?u1 (INil))))
) (
(let ?fs (Op (FusionStart ?shape ?s ?dt) (ICons ?x (INil))))
(let ?fu1 (Op (FusedRecip ?shape ?s ?s ?dt) (ICons ?fs (INil))))
(let ?fu2 (Op (FusedLog2 ?shape ?s ?s ?dt) (ICons ?fu1 (INil))))
(let ?fe (Op (FusionEnd ?shape ?s ?dt) (ICons ?fu2 (INil))))
(union ?u2 ?fe)
) :ruleset fusion_pair :name "pair-fuse-U-U-KernelRecip-KernelLog2")
(rule (
(= ?u1 (Op (KernelRecip ?shape ?s ?s ?dt) (ICons ?x (INil))))
(= ?u2 (Op (KernelRecip ?shape ?s ?s ?dt) (ICons ?u1 (INil))))
) (
(let ?fs (Op (FusionStart ?shape ?s ?dt) (ICons ?x (INil))))
(let ?fu1 (Op (FusedRecip ?shape ?s ?s ?dt) (ICons ?fs (INil))))
(let ?fu2 (Op (FusedRecip ?shape ?s ?s ?dt) (ICons ?fu1 (INil))))
(let ?fe (Op (FusionEnd ?shape ?s ?dt) (ICons ?fu2 (INil))))
(union ?u2 ?fe)
) :ruleset fusion_pair :name "pair-fuse-U-U-KernelRecip-KernelRecip")
(rule (
(= ?bin (Op (KernelAdd ?shape ?a_s ?b_s ?o_s ?dt)
(ICons ?a (ICons ?b (INil)))))
(= ?u (Op (KernelSin ?shape ?o_s ?o_s ?dt) (ICons ?bin (INil))))
) (
(let ?fs_a (Op (FusionStart ?shape ?a_s ?dt) (ICons ?a (INil))))
(let ?fs_b (Op (FusionStart ?shape ?b_s ?dt) (ICons ?b (INil))))
(let ?fbin (Op (FusedAdd ?shape ?a_s ?b_s ?o_s ?dt)
(ICons ?fs_a (ICons ?fs_b (INil)))))
(let ?fu (Op (FusedSin ?shape ?o_s ?o_s ?dt) (ICons ?fbin (INil))))
(let ?fe (Op (FusionEnd ?shape ?o_s ?dt) (ICons ?fu (INil))))
(union ?u ?fe)
) :ruleset fusion_pair :name "pair-fuse-B-U-Add-KernelSin")
(rule (
(= ?bin (Op (KernelAdd ?shape ?a_s ?b_s ?o_s ?dt)
(ICons ?a (ICons ?b (INil)))))
(= ?u (Op (KernelSqrt ?shape ?o_s ?o_s ?dt) (ICons ?bin (INil))))
) (
(let ?fs_a (Op (FusionStart ?shape ?a_s ?dt) (ICons ?a (INil))))
(let ?fs_b (Op (FusionStart ?shape ?b_s ?dt) (ICons ?b (INil))))
(let ?fbin (Op (FusedAdd ?shape ?a_s ?b_s ?o_s ?dt)
(ICons ?fs_a (ICons ?fs_b (INil)))))
(let ?fu (Op (FusedSqrt ?shape ?o_s ?o_s ?dt) (ICons ?fbin (INil))))
(let ?fe (Op (FusionEnd ?shape ?o_s ?dt) (ICons ?fu (INil))))
(union ?u ?fe)
) :ruleset fusion_pair :name "pair-fuse-B-U-Add-KernelSqrt")
(rule (
(= ?bin (Op (KernelAdd ?shape ?a_s ?b_s ?o_s ?dt)
(ICons ?a (ICons ?b (INil)))))
(= ?u (Op (KernelExp ?shape ?o_s ?o_s ?dt) (ICons ?bin (INil))))
) (
(let ?fs_a (Op (FusionStart ?shape ?a_s ?dt) (ICons ?a (INil))))
(let ?fs_b (Op (FusionStart ?shape ?b_s ?dt) (ICons ?b (INil))))
(let ?fbin (Op (FusedAdd ?shape ?a_s ?b_s ?o_s ?dt)
(ICons ?fs_a (ICons ?fs_b (INil)))))
(let ?fu (Op (FusedExp ?shape ?o_s ?o_s ?dt) (ICons ?fbin (INil))))
(let ?fe (Op (FusionEnd ?shape ?o_s ?dt) (ICons ?fu (INil))))
(union ?u ?fe)
) :ruleset fusion_pair :name "pair-fuse-B-U-Add-KernelExp")
(rule (
(= ?bin (Op (KernelAdd ?shape ?a_s ?b_s ?o_s ?dt)
(ICons ?a (ICons ?b (INil)))))
(= ?u (Op (KernelExp2 ?shape ?o_s ?o_s ?dt) (ICons ?bin (INil))))
) (
(let ?fs_a (Op (FusionStart ?shape ?a_s ?dt) (ICons ?a (INil))))
(let ?fs_b (Op (FusionStart ?shape ?b_s ?dt) (ICons ?b (INil))))
(let ?fbin (Op (FusedAdd ?shape ?a_s ?b_s ?o_s ?dt)
(ICons ?fs_a (ICons ?fs_b (INil)))))
(let ?fu (Op (FusedExp2 ?shape ?o_s ?o_s ?dt) (ICons ?fbin (INil))))
(let ?fe (Op (FusionEnd ?shape ?o_s ?dt) (ICons ?fu (INil))))
(union ?u ?fe)
) :ruleset fusion_pair :name "pair-fuse-B-U-Add-KernelExp2")
(rule (
(= ?bin (Op (KernelAdd ?shape ?a_s ?b_s ?o_s ?dt)
(ICons ?a (ICons ?b (INil)))))
(= ?u (Op (KernelLog2 ?shape ?o_s ?o_s ?dt) (ICons ?bin (INil))))
) (
(let ?fs_a (Op (FusionStart ?shape ?a_s ?dt) (ICons ?a (INil))))
(let ?fs_b (Op (FusionStart ?shape ?b_s ?dt) (ICons ?b (INil))))
(let ?fbin (Op (FusedAdd ?shape ?a_s ?b_s ?o_s ?dt)
(ICons ?fs_a (ICons ?fs_b (INil)))))
(let ?fu (Op (FusedLog2 ?shape ?o_s ?o_s ?dt) (ICons ?fbin (INil))))
(let ?fe (Op (FusionEnd ?shape ?o_s ?dt) (ICons ?fu (INil))))
(union ?u ?fe)
) :ruleset fusion_pair :name "pair-fuse-B-U-Add-KernelLog2")
(rule (
(= ?bin (Op (KernelAdd ?shape ?a_s ?b_s ?o_s ?dt)
(ICons ?a (ICons ?b (INil)))))
(= ?u (Op (KernelRecip ?shape ?o_s ?o_s ?dt) (ICons ?bin (INil))))
) (
(let ?fs_a (Op (FusionStart ?shape ?a_s ?dt) (ICons ?a (INil))))
(let ?fs_b (Op (FusionStart ?shape ?b_s ?dt) (ICons ?b (INil))))
(let ?fbin (Op (FusedAdd ?shape ?a_s ?b_s ?o_s ?dt)
(ICons ?fs_a (ICons ?fs_b (INil)))))
(let ?fu (Op (FusedRecip ?shape ?o_s ?o_s ?dt) (ICons ?fbin (INil))))
(let ?fe (Op (FusionEnd ?shape ?o_s ?dt) (ICons ?fu (INil))))
(union ?u ?fe)
) :ruleset fusion_pair :name "pair-fuse-B-U-Add-KernelRecip")
(rule (
(= ?bin (Op (KernelMul ?shape ?a_s ?b_s ?o_s ?dt)
(ICons ?a (ICons ?b (INil)))))
(= ?u (Op (KernelSin ?shape ?o_s ?o_s ?dt) (ICons ?bin (INil))))
) (
(let ?fs_a (Op (FusionStart ?shape ?a_s ?dt) (ICons ?a (INil))))
(let ?fs_b (Op (FusionStart ?shape ?b_s ?dt) (ICons ?b (INil))))
(let ?fbin (Op (FusedMul ?shape ?a_s ?b_s ?o_s ?dt)
(ICons ?fs_a (ICons ?fs_b (INil)))))
(let ?fu (Op (FusedSin ?shape ?o_s ?o_s ?dt) (ICons ?fbin (INil))))
(let ?fe (Op (FusionEnd ?shape ?o_s ?dt) (ICons ?fu (INil))))
(union ?u ?fe)
) :ruleset fusion_pair :name "pair-fuse-B-U-Mul-KernelSin")
(rule (
(= ?bin (Op (KernelMul ?shape ?a_s ?b_s ?o_s ?dt)
(ICons ?a (ICons ?b (INil)))))
(= ?u (Op (KernelSqrt ?shape ?o_s ?o_s ?dt) (ICons ?bin (INil))))
) (
(let ?fs_a (Op (FusionStart ?shape ?a_s ?dt) (ICons ?a (INil))))
(let ?fs_b (Op (FusionStart ?shape ?b_s ?dt) (ICons ?b (INil))))
(let ?fbin (Op (FusedMul ?shape ?a_s ?b_s ?o_s ?dt)
(ICons ?fs_a (ICons ?fs_b (INil)))))
(let ?fu (Op (FusedSqrt ?shape ?o_s ?o_s ?dt) (ICons ?fbin (INil))))
(let ?fe (Op (FusionEnd ?shape ?o_s ?dt) (ICons ?fu (INil))))
(union ?u ?fe)
) :ruleset fusion_pair :name "pair-fuse-B-U-Mul-KernelSqrt")
(rule (
(= ?bin (Op (KernelMul ?shape ?a_s ?b_s ?o_s ?dt)
(ICons ?a (ICons ?b (INil)))))
(= ?u (Op (KernelExp ?shape ?o_s ?o_s ?dt) (ICons ?bin (INil))))
) (
(let ?fs_a (Op (FusionStart ?shape ?a_s ?dt) (ICons ?a (INil))))
(let ?fs_b (Op (FusionStart ?shape ?b_s ?dt) (ICons ?b (INil))))
(let ?fbin (Op (FusedMul ?shape ?a_s ?b_s ?o_s ?dt)
(ICons ?fs_a (ICons ?fs_b (INil)))))
(let ?fu (Op (FusedExp ?shape ?o_s ?o_s ?dt) (ICons ?fbin (INil))))
(let ?fe (Op (FusionEnd ?shape ?o_s ?dt) (ICons ?fu (INil))))
(union ?u ?fe)
) :ruleset fusion_pair :name "pair-fuse-B-U-Mul-KernelExp")
(rule (
(= ?bin (Op (KernelMul ?shape ?a_s ?b_s ?o_s ?dt)
(ICons ?a (ICons ?b (INil)))))
(= ?u (Op (KernelExp2 ?shape ?o_s ?o_s ?dt) (ICons ?bin (INil))))
) (
(let ?fs_a (Op (FusionStart ?shape ?a_s ?dt) (ICons ?a (INil))))
(let ?fs_b (Op (FusionStart ?shape ?b_s ?dt) (ICons ?b (INil))))
(let ?fbin (Op (FusedMul ?shape ?a_s ?b_s ?o_s ?dt)
(ICons ?fs_a (ICons ?fs_b (INil)))))
(let ?fu (Op (FusedExp2 ?shape ?o_s ?o_s ?dt) (ICons ?fbin (INil))))
(let ?fe (Op (FusionEnd ?shape ?o_s ?dt) (ICons ?fu (INil))))
(union ?u ?fe)
) :ruleset fusion_pair :name "pair-fuse-B-U-Mul-KernelExp2")
(rule (
(= ?bin (Op (KernelMul ?shape ?a_s ?b_s ?o_s ?dt)
(ICons ?a (ICons ?b (INil)))))
(= ?u (Op (KernelLog2 ?shape ?o_s ?o_s ?dt) (ICons ?bin (INil))))
) (
(let ?fs_a (Op (FusionStart ?shape ?a_s ?dt) (ICons ?a (INil))))
(let ?fs_b (Op (FusionStart ?shape ?b_s ?dt) (ICons ?b (INil))))
(let ?fbin (Op (FusedMul ?shape ?a_s ?b_s ?o_s ?dt)
(ICons ?fs_a (ICons ?fs_b (INil)))))
(let ?fu (Op (FusedLog2 ?shape ?o_s ?o_s ?dt) (ICons ?fbin (INil))))
(let ?fe (Op (FusionEnd ?shape ?o_s ?dt) (ICons ?fu (INil))))
(union ?u ?fe)
) :ruleset fusion_pair :name "pair-fuse-B-U-Mul-KernelLog2")
(rule (
(= ?bin (Op (KernelMul ?shape ?a_s ?b_s ?o_s ?dt)
(ICons ?a (ICons ?b (INil)))))
(= ?u (Op (KernelRecip ?shape ?o_s ?o_s ?dt) (ICons ?bin (INil))))
) (
(let ?fs_a (Op (FusionStart ?shape ?a_s ?dt) (ICons ?a (INil))))
(let ?fs_b (Op (FusionStart ?shape ?b_s ?dt) (ICons ?b (INil))))
(let ?fbin (Op (FusedMul ?shape ?a_s ?b_s ?o_s ?dt)
(ICons ?fs_a (ICons ?fs_b (INil)))))
(let ?fu (Op (FusedRecip ?shape ?o_s ?o_s ?dt) (ICons ?fbin (INil))))
(let ?fe (Op (FusionEnd ?shape ?o_s ?dt) (ICons ?fu (INil))))
(union ?u ?fe)
) :ruleset fusion_pair :name "pair-fuse-B-U-Mul-KernelRecip")
(rule (
(= ?u (Op (KernelSin ?shape ?u_s ?u_s ?dt) (ICons ?a (INil))))
(= ?bin (Op (KernelAdd ?shape ?u_s ?b_s ?o_s ?dt)
(ICons ?u (ICons ?b (INil)))))
) (
(let ?fs_a (Op (FusionStart ?shape ?u_s ?dt) (ICons ?a (INil))))
(let ?fs_b (Op (FusionStart ?shape ?b_s ?dt) (ICons ?b (INil))))
(let ?fu (Op (FusedSin ?shape ?u_s ?u_s ?dt) (ICons ?fs_a (INil))))
(let ?fbin (Op (FusedAdd ?shape ?u_s ?b_s ?o_s ?dt)
(ICons ?fu (ICons ?fs_b (INil)))))
(let ?fe (Op (FusionEnd ?shape ?o_s ?dt) (ICons ?fbin (INil))))
(union ?bin ?fe)
) :ruleset fusion_pair :name "pair-fuse-U-B-lhs-KernelSin-Add")
(rule (
(= ?u (Op (KernelSin ?shape ?u_s ?u_s ?dt) (ICons ?b (INil))))
(= ?bin (Op (KernelAdd ?shape ?a_s ?u_s ?o_s ?dt)
(ICons ?a (ICons ?u (INil)))))
) (
(let ?fs_a (Op (FusionStart ?shape ?a_s ?dt) (ICons ?a (INil))))
(let ?fs_b (Op (FusionStart ?shape ?u_s ?dt) (ICons ?b (INil))))
(let ?fu (Op (FusedSin ?shape ?u_s ?u_s ?dt) (ICons ?fs_b (INil))))
(let ?fbin (Op (FusedAdd ?shape ?a_s ?u_s ?o_s ?dt)
(ICons ?fs_a (ICons ?fu (INil)))))
(let ?fe (Op (FusionEnd ?shape ?o_s ?dt) (ICons ?fbin (INil))))
(union ?bin ?fe)
) :ruleset fusion_pair :name "pair-fuse-U-B-rhs-KernelSin-Add")
(rule (
(= ?u (Op (KernelSin ?shape ?u_s ?u_s ?dt) (ICons ?a (INil))))
(= ?bin (Op (KernelMul ?shape ?u_s ?b_s ?o_s ?dt)
(ICons ?u (ICons ?b (INil)))))
) (
(let ?fs_a (Op (FusionStart ?shape ?u_s ?dt) (ICons ?a (INil))))
(let ?fs_b (Op (FusionStart ?shape ?b_s ?dt) (ICons ?b (INil))))
(let ?fu (Op (FusedSin ?shape ?u_s ?u_s ?dt) (ICons ?fs_a (INil))))
(let ?fbin (Op (FusedMul ?shape ?u_s ?b_s ?o_s ?dt)
(ICons ?fu (ICons ?fs_b (INil)))))
(let ?fe (Op (FusionEnd ?shape ?o_s ?dt) (ICons ?fbin (INil))))
(union ?bin ?fe)
) :ruleset fusion_pair :name "pair-fuse-U-B-lhs-KernelSin-Mul")
(rule (
(= ?u (Op (KernelSin ?shape ?u_s ?u_s ?dt) (ICons ?b (INil))))
(= ?bin (Op (KernelMul ?shape ?a_s ?u_s ?o_s ?dt)
(ICons ?a (ICons ?u (INil)))))
) (
(let ?fs_a (Op (FusionStart ?shape ?a_s ?dt) (ICons ?a (INil))))
(let ?fs_b (Op (FusionStart ?shape ?u_s ?dt) (ICons ?b (INil))))
(let ?fu (Op (FusedSin ?shape ?u_s ?u_s ?dt) (ICons ?fs_b (INil))))
(let ?fbin (Op (FusedMul ?shape ?a_s ?u_s ?o_s ?dt)
(ICons ?fs_a (ICons ?fu (INil)))))
(let ?fe (Op (FusionEnd ?shape ?o_s ?dt) (ICons ?fbin (INil))))
(union ?bin ?fe)
) :ruleset fusion_pair :name "pair-fuse-U-B-rhs-KernelSin-Mul")
(rule (
(= ?u (Op (KernelSqrt ?shape ?u_s ?u_s ?dt) (ICons ?a (INil))))
(= ?bin (Op (KernelAdd ?shape ?u_s ?b_s ?o_s ?dt)
(ICons ?u (ICons ?b (INil)))))
) (
(let ?fs_a (Op (FusionStart ?shape ?u_s ?dt) (ICons ?a (INil))))
(let ?fs_b (Op (FusionStart ?shape ?b_s ?dt) (ICons ?b (INil))))
(let ?fu (Op (FusedSqrt ?shape ?u_s ?u_s ?dt) (ICons ?fs_a (INil))))
(let ?fbin (Op (FusedAdd ?shape ?u_s ?b_s ?o_s ?dt)
(ICons ?fu (ICons ?fs_b (INil)))))
(let ?fe (Op (FusionEnd ?shape ?o_s ?dt) (ICons ?fbin (INil))))
(union ?bin ?fe)
) :ruleset fusion_pair :name "pair-fuse-U-B-lhs-KernelSqrt-Add")
(rule (
(= ?u (Op (KernelSqrt ?shape ?u_s ?u_s ?dt) (ICons ?b (INil))))
(= ?bin (Op (KernelAdd ?shape ?a_s ?u_s ?o_s ?dt)
(ICons ?a (ICons ?u (INil)))))
) (
(let ?fs_a (Op (FusionStart ?shape ?a_s ?dt) (ICons ?a (INil))))
(let ?fs_b (Op (FusionStart ?shape ?u_s ?dt) (ICons ?b (INil))))
(let ?fu (Op (FusedSqrt ?shape ?u_s ?u_s ?dt) (ICons ?fs_b (INil))))
(let ?fbin (Op (FusedAdd ?shape ?a_s ?u_s ?o_s ?dt)
(ICons ?fs_a (ICons ?fu (INil)))))
(let ?fe (Op (FusionEnd ?shape ?o_s ?dt) (ICons ?fbin (INil))))
(union ?bin ?fe)
) :ruleset fusion_pair :name "pair-fuse-U-B-rhs-KernelSqrt-Add")
(rule (
(= ?u (Op (KernelSqrt ?shape ?u_s ?u_s ?dt) (ICons ?a (INil))))
(= ?bin (Op (KernelMul ?shape ?u_s ?b_s ?o_s ?dt)
(ICons ?u (ICons ?b (INil)))))
) (
(let ?fs_a (Op (FusionStart ?shape ?u_s ?dt) (ICons ?a (INil))))
(let ?fs_b (Op (FusionStart ?shape ?b_s ?dt) (ICons ?b (INil))))
(let ?fu (Op (FusedSqrt ?shape ?u_s ?u_s ?dt) (ICons ?fs_a (INil))))
(let ?fbin (Op (FusedMul ?shape ?u_s ?b_s ?o_s ?dt)
(ICons ?fu (ICons ?fs_b (INil)))))
(let ?fe (Op (FusionEnd ?shape ?o_s ?dt) (ICons ?fbin (INil))))
(union ?bin ?fe)
) :ruleset fusion_pair :name "pair-fuse-U-B-lhs-KernelSqrt-Mul")
(rule (
(= ?u (Op (KernelSqrt ?shape ?u_s ?u_s ?dt) (ICons ?b (INil))))
(= ?bin (Op (KernelMul ?shape ?a_s ?u_s ?o_s ?dt)
(ICons ?a (ICons ?u (INil)))))
) (
(let ?fs_a (Op (FusionStart ?shape ?a_s ?dt) (ICons ?a (INil))))
(let ?fs_b (Op (FusionStart ?shape ?u_s ?dt) (ICons ?b (INil))))
(let ?fu (Op (FusedSqrt ?shape ?u_s ?u_s ?dt) (ICons ?fs_b (INil))))
(let ?fbin (Op (FusedMul ?shape ?a_s ?u_s ?o_s ?dt)
(ICons ?fs_a (ICons ?fu (INil)))))
(let ?fe (Op (FusionEnd ?shape ?o_s ?dt) (ICons ?fbin (INil))))
(union ?bin ?fe)
) :ruleset fusion_pair :name "pair-fuse-U-B-rhs-KernelSqrt-Mul")
(rule (
(= ?u (Op (KernelExp ?shape ?u_s ?u_s ?dt) (ICons ?a (INil))))
(= ?bin (Op (KernelAdd ?shape ?u_s ?b_s ?o_s ?dt)
(ICons ?u (ICons ?b (INil)))))
) (
(let ?fs_a (Op (FusionStart ?shape ?u_s ?dt) (ICons ?a (INil))))
(let ?fs_b (Op (FusionStart ?shape ?b_s ?dt) (ICons ?b (INil))))
(let ?fu (Op (FusedExp ?shape ?u_s ?u_s ?dt) (ICons ?fs_a (INil))))
(let ?fbin (Op (FusedAdd ?shape ?u_s ?b_s ?o_s ?dt)
(ICons ?fu (ICons ?fs_b (INil)))))
(let ?fe (Op (FusionEnd ?shape ?o_s ?dt) (ICons ?fbin (INil))))
(union ?bin ?fe)
) :ruleset fusion_pair :name "pair-fuse-U-B-lhs-KernelExp-Add")
(rule (
(= ?u (Op (KernelExp ?shape ?u_s ?u_s ?dt) (ICons ?b (INil))))
(= ?bin (Op (KernelAdd ?shape ?a_s ?u_s ?o_s ?dt)
(ICons ?a (ICons ?u (INil)))))
) (
(let ?fs_a (Op (FusionStart ?shape ?a_s ?dt) (ICons ?a (INil))))
(let ?fs_b (Op (FusionStart ?shape ?u_s ?dt) (ICons ?b (INil))))
(let ?fu (Op (FusedExp ?shape ?u_s ?u_s ?dt) (ICons ?fs_b (INil))))
(let ?fbin (Op (FusedAdd ?shape ?a_s ?u_s ?o_s ?dt)
(ICons ?fs_a (ICons ?fu (INil)))))
(let ?fe (Op (FusionEnd ?shape ?o_s ?dt) (ICons ?fbin (INil))))
(union ?bin ?fe)
) :ruleset fusion_pair :name "pair-fuse-U-B-rhs-KernelExp-Add")
(rule (
(= ?u (Op (KernelExp ?shape ?u_s ?u_s ?dt) (ICons ?a (INil))))
(= ?bin (Op (KernelMul ?shape ?u_s ?b_s ?o_s ?dt)
(ICons ?u (ICons ?b (INil)))))
) (
(let ?fs_a (Op (FusionStart ?shape ?u_s ?dt) (ICons ?a (INil))))
(let ?fs_b (Op (FusionStart ?shape ?b_s ?dt) (ICons ?b (INil))))
(let ?fu (Op (FusedExp ?shape ?u_s ?u_s ?dt) (ICons ?fs_a (INil))))
(let ?fbin (Op (FusedMul ?shape ?u_s ?b_s ?o_s ?dt)
(ICons ?fu (ICons ?fs_b (INil)))))
(let ?fe (Op (FusionEnd ?shape ?o_s ?dt) (ICons ?fbin (INil))))
(union ?bin ?fe)
) :ruleset fusion_pair :name "pair-fuse-U-B-lhs-KernelExp-Mul")
(rule (
(= ?u (Op (KernelExp ?shape ?u_s ?u_s ?dt) (ICons ?b (INil))))
(= ?bin (Op (KernelMul ?shape ?a_s ?u_s ?o_s ?dt)
(ICons ?a (ICons ?u (INil)))))
) (
(let ?fs_a (Op (FusionStart ?shape ?a_s ?dt) (ICons ?a (INil))))
(let ?fs_b (Op (FusionStart ?shape ?u_s ?dt) (ICons ?b (INil))))
(let ?fu (Op (FusedExp ?shape ?u_s ?u_s ?dt) (ICons ?fs_b (INil))))
(let ?fbin (Op (FusedMul ?shape ?a_s ?u_s ?o_s ?dt)
(ICons ?fs_a (ICons ?fu (INil)))))
(let ?fe (Op (FusionEnd ?shape ?o_s ?dt) (ICons ?fbin (INil))))
(union ?bin ?fe)
) :ruleset fusion_pair :name "pair-fuse-U-B-rhs-KernelExp-Mul")
(rule (
(= ?u (Op (KernelExp2 ?shape ?u_s ?u_s ?dt) (ICons ?a (INil))))
(= ?bin (Op (KernelAdd ?shape ?u_s ?b_s ?o_s ?dt)
(ICons ?u (ICons ?b (INil)))))
) (
(let ?fs_a (Op (FusionStart ?shape ?u_s ?dt) (ICons ?a (INil))))
(let ?fs_b (Op (FusionStart ?shape ?b_s ?dt) (ICons ?b (INil))))
(let ?fu (Op (FusedExp2 ?shape ?u_s ?u_s ?dt) (ICons ?fs_a (INil))))
(let ?fbin (Op (FusedAdd ?shape ?u_s ?b_s ?o_s ?dt)
(ICons ?fu (ICons ?fs_b (INil)))))
(let ?fe (Op (FusionEnd ?shape ?o_s ?dt) (ICons ?fbin (INil))))
(union ?bin ?fe)
) :ruleset fusion_pair :name "pair-fuse-U-B-lhs-KernelExp2-Add")
(rule (
(= ?u (Op (KernelExp2 ?shape ?u_s ?u_s ?dt) (ICons ?b (INil))))
(= ?bin (Op (KernelAdd ?shape ?a_s ?u_s ?o_s ?dt)
(ICons ?a (ICons ?u (INil)))))
) (
(let ?fs_a (Op (FusionStart ?shape ?a_s ?dt) (ICons ?a (INil))))
(let ?fs_b (Op (FusionStart ?shape ?u_s ?dt) (ICons ?b (INil))))
(let ?fu (Op (FusedExp2 ?shape ?u_s ?u_s ?dt) (ICons ?fs_b (INil))))
(let ?fbin (Op (FusedAdd ?shape ?a_s ?u_s ?o_s ?dt)
(ICons ?fs_a (ICons ?fu (INil)))))
(let ?fe (Op (FusionEnd ?shape ?o_s ?dt) (ICons ?fbin (INil))))
(union ?bin ?fe)
) :ruleset fusion_pair :name "pair-fuse-U-B-rhs-KernelExp2-Add")
(rule (
(= ?u (Op (KernelExp2 ?shape ?u_s ?u_s ?dt) (ICons ?a (INil))))
(= ?bin (Op (KernelMul ?shape ?u_s ?b_s ?o_s ?dt)
(ICons ?u (ICons ?b (INil)))))
) (
(let ?fs_a (Op (FusionStart ?shape ?u_s ?dt) (ICons ?a (INil))))
(let ?fs_b (Op (FusionStart ?shape ?b_s ?dt) (ICons ?b (INil))))
(let ?fu (Op (FusedExp2 ?shape ?u_s ?u_s ?dt) (ICons ?fs_a (INil))))
(let ?fbin (Op (FusedMul ?shape ?u_s ?b_s ?o_s ?dt)
(ICons ?fu (ICons ?fs_b (INil)))))
(let ?fe (Op (FusionEnd ?shape ?o_s ?dt) (ICons ?fbin (INil))))
(union ?bin ?fe)
) :ruleset fusion_pair :name "pair-fuse-U-B-lhs-KernelExp2-Mul")
(rule (
(= ?u (Op (KernelExp2 ?shape ?u_s ?u_s ?dt) (ICons ?b (INil))))
(= ?bin (Op (KernelMul ?shape ?a_s ?u_s ?o_s ?dt)
(ICons ?a (ICons ?u (INil)))))
) (
(let ?fs_a (Op (FusionStart ?shape ?a_s ?dt) (ICons ?a (INil))))
(let ?fs_b (Op (FusionStart ?shape ?u_s ?dt) (ICons ?b (INil))))
(let ?fu (Op (FusedExp2 ?shape ?u_s ?u_s ?dt) (ICons ?fs_b (INil))))
(let ?fbin (Op (FusedMul ?shape ?a_s ?u_s ?o_s ?dt)
(ICons ?fs_a (ICons ?fu (INil)))))
(let ?fe (Op (FusionEnd ?shape ?o_s ?dt) (ICons ?fbin (INil))))
(union ?bin ?fe)
) :ruleset fusion_pair :name "pair-fuse-U-B-rhs-KernelExp2-Mul")
(rule (
(= ?u (Op (KernelLog2 ?shape ?u_s ?u_s ?dt) (ICons ?a (INil))))
(= ?bin (Op (KernelAdd ?shape ?u_s ?b_s ?o_s ?dt)
(ICons ?u (ICons ?b (INil)))))
) (
(let ?fs_a (Op (FusionStart ?shape ?u_s ?dt) (ICons ?a (INil))))
(let ?fs_b (Op (FusionStart ?shape ?b_s ?dt) (ICons ?b (INil))))
(let ?fu (Op (FusedLog2 ?shape ?u_s ?u_s ?dt) (ICons ?fs_a (INil))))
(let ?fbin (Op (FusedAdd ?shape ?u_s ?b_s ?o_s ?dt)
(ICons ?fu (ICons ?fs_b (INil)))))
(let ?fe (Op (FusionEnd ?shape ?o_s ?dt) (ICons ?fbin (INil))))
(union ?bin ?fe)
) :ruleset fusion_pair :name "pair-fuse-U-B-lhs-KernelLog2-Add")
(rule (
(= ?u (Op (KernelLog2 ?shape ?u_s ?u_s ?dt) (ICons ?b (INil))))
(= ?bin (Op (KernelAdd ?shape ?a_s ?u_s ?o_s ?dt)
(ICons ?a (ICons ?u (INil)))))
) (
(let ?fs_a (Op (FusionStart ?shape ?a_s ?dt) (ICons ?a (INil))))
(let ?fs_b (Op (FusionStart ?shape ?u_s ?dt) (ICons ?b (INil))))
(let ?fu (Op (FusedLog2 ?shape ?u_s ?u_s ?dt) (ICons ?fs_b (INil))))
(let ?fbin (Op (FusedAdd ?shape ?a_s ?u_s ?o_s ?dt)
(ICons ?fs_a (ICons ?fu (INil)))))
(let ?fe (Op (FusionEnd ?shape ?o_s ?dt) (ICons ?fbin (INil))))
(union ?bin ?fe)
) :ruleset fusion_pair :name "pair-fuse-U-B-rhs-KernelLog2-Add")
(rule (
(= ?u (Op (KernelLog2 ?shape ?u_s ?u_s ?dt) (ICons ?a (INil))))
(= ?bin (Op (KernelMul ?shape ?u_s ?b_s ?o_s ?dt)
(ICons ?u (ICons ?b (INil)))))
) (
(let ?fs_a (Op (FusionStart ?shape ?u_s ?dt) (ICons ?a (INil))))
(let ?fs_b (Op (FusionStart ?shape ?b_s ?dt) (ICons ?b (INil))))
(let ?fu (Op (FusedLog2 ?shape ?u_s ?u_s ?dt) (ICons ?fs_a (INil))))
(let ?fbin (Op (FusedMul ?shape ?u_s ?b_s ?o_s ?dt)
(ICons ?fu (ICons ?fs_b (INil)))))
(let ?fe (Op (FusionEnd ?shape ?o_s ?dt) (ICons ?fbin (INil))))
(union ?bin ?fe)
) :ruleset fusion_pair :name "pair-fuse-U-B-lhs-KernelLog2-Mul")
(rule (
(= ?u (Op (KernelLog2 ?shape ?u_s ?u_s ?dt) (ICons ?b (INil))))
(= ?bin (Op (KernelMul ?shape ?a_s ?u_s ?o_s ?dt)
(ICons ?a (ICons ?u (INil)))))
) (
(let ?fs_a (Op (FusionStart ?shape ?a_s ?dt) (ICons ?a (INil))))
(let ?fs_b (Op (FusionStart ?shape ?u_s ?dt) (ICons ?b (INil))))
(let ?fu (Op (FusedLog2 ?shape ?u_s ?u_s ?dt) (ICons ?fs_b (INil))))
(let ?fbin (Op (FusedMul ?shape ?a_s ?u_s ?o_s ?dt)
(ICons ?fs_a (ICons ?fu (INil)))))
(let ?fe (Op (FusionEnd ?shape ?o_s ?dt) (ICons ?fbin (INil))))
(union ?bin ?fe)
) :ruleset fusion_pair :name "pair-fuse-U-B-rhs-KernelLog2-Mul")
(rule (
(= ?u (Op (KernelRecip ?shape ?u_s ?u_s ?dt) (ICons ?a (INil))))
(= ?bin (Op (KernelAdd ?shape ?u_s ?b_s ?o_s ?dt)
(ICons ?u (ICons ?b (INil)))))
) (
(let ?fs_a (Op (FusionStart ?shape ?u_s ?dt) (ICons ?a (INil))))
(let ?fs_b (Op (FusionStart ?shape ?b_s ?dt) (ICons ?b (INil))))
(let ?fu (Op (FusedRecip ?shape ?u_s ?u_s ?dt) (ICons ?fs_a (INil))))
(let ?fbin (Op (FusedAdd ?shape ?u_s ?b_s ?o_s ?dt)
(ICons ?fu (ICons ?fs_b (INil)))))
(let ?fe (Op (FusionEnd ?shape ?o_s ?dt) (ICons ?fbin (INil))))
(union ?bin ?fe)
) :ruleset fusion_pair :name "pair-fuse-U-B-lhs-KernelRecip-Add")
(rule (
(= ?u (Op (KernelRecip ?shape ?u_s ?u_s ?dt) (ICons ?b (INil))))
(= ?bin (Op (KernelAdd ?shape ?a_s ?u_s ?o_s ?dt)
(ICons ?a (ICons ?u (INil)))))
) (
(let ?fs_a (Op (FusionStart ?shape ?a_s ?dt) (ICons ?a (INil))))
(let ?fs_b (Op (FusionStart ?shape ?u_s ?dt) (ICons ?b (INil))))
(let ?fu (Op (FusedRecip ?shape ?u_s ?u_s ?dt) (ICons ?fs_b (INil))))
(let ?fbin (Op (FusedAdd ?shape ?a_s ?u_s ?o_s ?dt)
(ICons ?fs_a (ICons ?fu (INil)))))
(let ?fe (Op (FusionEnd ?shape ?o_s ?dt) (ICons ?fbin (INil))))
(union ?bin ?fe)
) :ruleset fusion_pair :name "pair-fuse-U-B-rhs-KernelRecip-Add")
(rule (
(= ?u (Op (KernelRecip ?shape ?u_s ?u_s ?dt) (ICons ?a (INil))))
(= ?bin (Op (KernelMul ?shape ?u_s ?b_s ?o_s ?dt)
(ICons ?u (ICons ?b (INil)))))
) (
(let ?fs_a (Op (FusionStart ?shape ?u_s ?dt) (ICons ?a (INil))))
(let ?fs_b (Op (FusionStart ?shape ?b_s ?dt) (ICons ?b (INil))))
(let ?fu (Op (FusedRecip ?shape ?u_s ?u_s ?dt) (ICons ?fs_a (INil))))
(let ?fbin (Op (FusedMul ?shape ?u_s ?b_s ?o_s ?dt)
(ICons ?fu (ICons ?fs_b (INil)))))
(let ?fe (Op (FusionEnd ?shape ?o_s ?dt) (ICons ?fbin (INil))))
(union ?bin ?fe)
) :ruleset fusion_pair :name "pair-fuse-U-B-lhs-KernelRecip-Mul")
(rule (
(= ?u (Op (KernelRecip ?shape ?u_s ?u_s ?dt) (ICons ?b (INil))))
(= ?bin (Op (KernelMul ?shape ?a_s ?u_s ?o_s ?dt)
(ICons ?a (ICons ?u (INil)))))
) (
(let ?fs_a (Op (FusionStart ?shape ?a_s ?dt) (ICons ?a (INil))))
(let ?fs_b (Op (FusionStart ?shape ?u_s ?dt) (ICons ?b (INil))))
(let ?fu (Op (FusedRecip ?shape ?u_s ?u_s ?dt) (ICons ?fs_b (INil))))
(let ?fbin (Op (FusedMul ?shape ?a_s ?u_s ?o_s ?dt)
(ICons ?fs_a (ICons ?fu (INil)))))
(let ?fe (Op (FusionEnd ?shape ?o_s ?dt) (ICons ?fbin (INil))))
(union ?bin ?fe)
) :ruleset fusion_pair :name "pair-fuse-U-B-rhs-KernelRecip-Mul")
(rule (
(= ?bi (Op (KernelAdd ?shape ?ai_s ?bi_s ?oi_s ?dt)
(ICons ?a (ICons ?b (INil)))))
(= ?bo (Op (KernelAdd ?shape ?oi_s ?co_s ?oo_s ?dt)
(ICons ?bi (ICons ?c (INil)))))
) (
(let ?fs_a (Op (FusionStart ?shape ?ai_s ?dt) (ICons ?a (INil))))
(let ?fs_b (Op (FusionStart ?shape ?bi_s ?dt) (ICons ?b (INil))))
(let ?fs_c (Op (FusionStart ?shape ?co_s ?dt) (ICons ?c (INil))))
(let ?fbi (Op (FusedAdd ?shape ?ai_s ?bi_s ?oi_s ?dt)
(ICons ?fs_a (ICons ?fs_b (INil)))))
(let ?fbo (Op (FusedAdd ?shape ?oi_s ?co_s ?oo_s ?dt)
(ICons ?fbi (ICons ?fs_c (INil)))))
(let ?fe (Op (FusionEnd ?shape ?oo_s ?dt) (ICons ?fbo (INil))))
(union ?bo ?fe)
) :ruleset fusion_pair :name "pair-fuse-B-B-lhs-Add-Add")
(rule (
(= ?bi (Op (KernelAdd ?shape ?ai_s ?bi_s ?oi_s ?dt)
(ICons ?a (ICons ?b (INil)))))
(= ?bo (Op (KernelAdd ?shape ?co_s ?oi_s ?oo_s ?dt)
(ICons ?c (ICons ?bi (INil)))))
) (
(let ?fs_a (Op (FusionStart ?shape ?ai_s ?dt) (ICons ?a (INil))))
(let ?fs_b (Op (FusionStart ?shape ?bi_s ?dt) (ICons ?b (INil))))
(let ?fs_c (Op (FusionStart ?shape ?co_s ?dt) (ICons ?c (INil))))
(let ?fbi (Op (FusedAdd ?shape ?ai_s ?bi_s ?oi_s ?dt)
(ICons ?fs_a (ICons ?fs_b (INil)))))
(let ?fbo (Op (FusedAdd ?shape ?co_s ?oi_s ?oo_s ?dt)
(ICons ?fs_c (ICons ?fbi (INil)))))
(let ?fe (Op (FusionEnd ?shape ?oo_s ?dt) (ICons ?fbo (INil))))
(union ?bo ?fe)
) :ruleset fusion_pair :name "pair-fuse-B-B-rhs-Add-Add")
(rule (
(= ?bi (Op (KernelAdd ?shape ?ai_s ?bi_s ?oi_s ?dt)
(ICons ?a (ICons ?b (INil)))))
(= ?bo (Op (KernelMul ?shape ?oi_s ?co_s ?oo_s ?dt)
(ICons ?bi (ICons ?c (INil)))))
) (
(let ?fs_a (Op (FusionStart ?shape ?ai_s ?dt) (ICons ?a (INil))))
(let ?fs_b (Op (FusionStart ?shape ?bi_s ?dt) (ICons ?b (INil))))
(let ?fs_c (Op (FusionStart ?shape ?co_s ?dt) (ICons ?c (INil))))
(let ?fbi (Op (FusedAdd ?shape ?ai_s ?bi_s ?oi_s ?dt)
(ICons ?fs_a (ICons ?fs_b (INil)))))
(let ?fbo (Op (FusedMul ?shape ?oi_s ?co_s ?oo_s ?dt)
(ICons ?fbi (ICons ?fs_c (INil)))))
(let ?fe (Op (FusionEnd ?shape ?oo_s ?dt) (ICons ?fbo (INil))))
(union ?bo ?fe)
) :ruleset fusion_pair :name "pair-fuse-B-B-lhs-Add-Mul")
(rule (
(= ?bi (Op (KernelAdd ?shape ?ai_s ?bi_s ?oi_s ?dt)
(ICons ?a (ICons ?b (INil)))))
(= ?bo (Op (KernelMul ?shape ?co_s ?oi_s ?oo_s ?dt)
(ICons ?c (ICons ?bi (INil)))))
) (
(let ?fs_a (Op (FusionStart ?shape ?ai_s ?dt) (ICons ?a (INil))))
(let ?fs_b (Op (FusionStart ?shape ?bi_s ?dt) (ICons ?b (INil))))
(let ?fs_c (Op (FusionStart ?shape ?co_s ?dt) (ICons ?c (INil))))
(let ?fbi (Op (FusedAdd ?shape ?ai_s ?bi_s ?oi_s ?dt)
(ICons ?fs_a (ICons ?fs_b (INil)))))
(let ?fbo (Op (FusedMul ?shape ?co_s ?oi_s ?oo_s ?dt)
(ICons ?fs_c (ICons ?fbi (INil)))))
(let ?fe (Op (FusionEnd ?shape ?oo_s ?dt) (ICons ?fbo (INil))))
(union ?bo ?fe)
) :ruleset fusion_pair :name "pair-fuse-B-B-rhs-Add-Mul")
(rule (
(= ?bi (Op (KernelMul ?shape ?ai_s ?bi_s ?oi_s ?dt)
(ICons ?a (ICons ?b (INil)))))
(= ?bo (Op (KernelAdd ?shape ?oi_s ?co_s ?oo_s ?dt)
(ICons ?bi (ICons ?c (INil)))))
) (
(let ?fs_a (Op (FusionStart ?shape ?ai_s ?dt) (ICons ?a (INil))))
(let ?fs_b (Op (FusionStart ?shape ?bi_s ?dt) (ICons ?b (INil))))
(let ?fs_c (Op (FusionStart ?shape ?co_s ?dt) (ICons ?c (INil))))
(let ?fbi (Op (FusedMul ?shape ?ai_s ?bi_s ?oi_s ?dt)
(ICons ?fs_a (ICons ?fs_b (INil)))))
(let ?fbo (Op (FusedAdd ?shape ?oi_s ?co_s ?oo_s ?dt)
(ICons ?fbi (ICons ?fs_c (INil)))))
(let ?fe (Op (FusionEnd ?shape ?oo_s ?dt) (ICons ?fbo (INil))))
(union ?bo ?fe)
) :ruleset fusion_pair :name "pair-fuse-B-B-lhs-Mul-Add")
(rule (
(= ?bi (Op (KernelMul ?shape ?ai_s ?bi_s ?oi_s ?dt)
(ICons ?a (ICons ?b (INil)))))
(= ?bo (Op (KernelAdd ?shape ?co_s ?oi_s ?oo_s ?dt)
(ICons ?c (ICons ?bi (INil)))))
) (
(let ?fs_a (Op (FusionStart ?shape ?ai_s ?dt) (ICons ?a (INil))))
(let ?fs_b (Op (FusionStart ?shape ?bi_s ?dt) (ICons ?b (INil))))
(let ?fs_c (Op (FusionStart ?shape ?co_s ?dt) (ICons ?c (INil))))
(let ?fbi (Op (FusedMul ?shape ?ai_s ?bi_s ?oi_s ?dt)
(ICons ?fs_a (ICons ?fs_b (INil)))))
(let ?fbo (Op (FusedAdd ?shape ?co_s ?oi_s ?oo_s ?dt)
(ICons ?fs_c (ICons ?fbi (INil)))))
(let ?fe (Op (FusionEnd ?shape ?oo_s ?dt) (ICons ?fbo (INil))))
(union ?bo ?fe)
) :ruleset fusion_pair :name "pair-fuse-B-B-rhs-Mul-Add")
(rule (
(= ?bi (Op (KernelMul ?shape ?ai_s ?bi_s ?oi_s ?dt)
(ICons ?a (ICons ?b (INil)))))
(= ?bo (Op (KernelMul ?shape ?oi_s ?co_s ?oo_s ?dt)
(ICons ?bi (ICons ?c (INil)))))
) (
(let ?fs_a (Op (FusionStart ?shape ?ai_s ?dt) (ICons ?a (INil))))
(let ?fs_b (Op (FusionStart ?shape ?bi_s ?dt) (ICons ?b (INil))))
(let ?fs_c (Op (FusionStart ?shape ?co_s ?dt) (ICons ?c (INil))))
(let ?fbi (Op (FusedMul ?shape ?ai_s ?bi_s ?oi_s ?dt)
(ICons ?fs_a (ICons ?fs_b (INil)))))
(let ?fbo (Op (FusedMul ?shape ?oi_s ?co_s ?oo_s ?dt)
(ICons ?fbi (ICons ?fs_c (INil)))))
(let ?fe (Op (FusionEnd ?shape ?oo_s ?dt) (ICons ?fbo (INil))))
(union ?bo ?fe)
) :ruleset fusion_pair :name "pair-fuse-B-B-lhs-Mul-Mul")
(rule (
(= ?bi (Op (KernelMul ?shape ?ai_s ?bi_s ?oi_s ?dt)
(ICons ?a (ICons ?b (INil)))))
(= ?bo (Op (KernelMul ?shape ?co_s ?oi_s ?oo_s ?dt)
(ICons ?c (ICons ?bi (INil)))))
) (
(let ?fs_a (Op (FusionStart ?shape ?ai_s ?dt) (ICons ?a (INil))))
(let ?fs_b (Op (FusionStart ?shape ?bi_s ?dt) (ICons ?b (INil))))
(let ?fs_c (Op (FusionStart ?shape ?co_s ?dt) (ICons ?c (INil))))
(let ?fbi (Op (FusedMul ?shape ?ai_s ?bi_s ?oi_s ?dt)
(ICons ?fs_a (ICons ?fs_b (INil)))))
(let ?fbo (Op (FusedMul ?shape ?co_s ?oi_s ?oo_s ?dt)
(ICons ?fs_c (ICons ?fbi (INil)))))
(let ?fe (Op (FusionEnd ?shape ?oo_s ?dt) (ICons ?fbo (INil))))
(union ?bo ?fe)
) :ruleset fusion_pair :name "pair-fuse-B-B-rhs-Mul-Mul")
(rule (
(= ?fe (Op (FusionEnd ?shape ?s ?dt) (ICons ?inner (INil))))
(= ?u (Op (KernelSin ?shape ?s ?s ?dt) (ICons ?fe (INil))))
) (
(let ?fu (Op (FusedSin ?shape ?s ?s ?dt) (ICons ?inner (INil))))
(let ?new_fe (Op (FusionEnd ?shape ?s ?dt) (ICons ?fu (INil))))
(union ?u ?new_fe)
) :ruleset fusion_grow :name "grow-FE-U-KernelSin")
(rule (
(= ?fe (Op (FusionEnd ?shape ?s ?dt) (ICons ?inner (INil))))
(= ?u (Op (KernelSqrt ?shape ?s ?s ?dt) (ICons ?fe (INil))))
) (
(let ?fu (Op (FusedSqrt ?shape ?s ?s ?dt) (ICons ?inner (INil))))
(let ?new_fe (Op (FusionEnd ?shape ?s ?dt) (ICons ?fu (INil))))
(union ?u ?new_fe)
) :ruleset fusion_grow :name "grow-FE-U-KernelSqrt")
(rule (
(= ?fe (Op (FusionEnd ?shape ?s ?dt) (ICons ?inner (INil))))
(= ?u (Op (KernelExp ?shape ?s ?s ?dt) (ICons ?fe (INil))))
) (
(let ?fu (Op (FusedExp ?shape ?s ?s ?dt) (ICons ?inner (INil))))
(let ?new_fe (Op (FusionEnd ?shape ?s ?dt) (ICons ?fu (INil))))
(union ?u ?new_fe)
) :ruleset fusion_grow :name "grow-FE-U-KernelExp")
(rule (
(= ?fe (Op (FusionEnd ?shape ?s ?dt) (ICons ?inner (INil))))
(= ?u (Op (KernelExp2 ?shape ?s ?s ?dt) (ICons ?fe (INil))))
) (
(let ?fu (Op (FusedExp2 ?shape ?s ?s ?dt) (ICons ?inner (INil))))
(let ?new_fe (Op (FusionEnd ?shape ?s ?dt) (ICons ?fu (INil))))
(union ?u ?new_fe)
) :ruleset fusion_grow :name "grow-FE-U-KernelExp2")
(rule (
(= ?fe (Op (FusionEnd ?shape ?s ?dt) (ICons ?inner (INil))))
(= ?u (Op (KernelLog2 ?shape ?s ?s ?dt) (ICons ?fe (INil))))
) (
(let ?fu (Op (FusedLog2 ?shape ?s ?s ?dt) (ICons ?inner (INil))))
(let ?new_fe (Op (FusionEnd ?shape ?s ?dt) (ICons ?fu (INil))))
(union ?u ?new_fe)
) :ruleset fusion_grow :name "grow-FE-U-KernelLog2")
(rule (
(= ?fe (Op (FusionEnd ?shape ?s ?dt) (ICons ?inner (INil))))
(= ?u (Op (KernelRecip ?shape ?s ?s ?dt) (ICons ?fe (INil))))
) (
(let ?fu (Op (FusedRecip ?shape ?s ?s ?dt) (ICons ?inner (INil))))
(let ?new_fe (Op (FusionEnd ?shape ?s ?dt) (ICons ?fu (INil))))
(union ?u ?new_fe)
) :ruleset fusion_grow :name "grow-FE-U-KernelRecip")
(rule (
(= ?fe (Op (FusionEnd ?shape ?a_s ?dt) (ICons ?inner_a (INil))))
(= ?bin (Op (KernelAdd ?shape ?a_s ?b_s ?o_s ?dt)
(ICons ?fe (ICons ?b (INil)))))
) (
(let ?fs_b (Op (FusionStart ?shape ?b_s ?dt) (ICons ?b (INil))))
(let ?fbin (Op (FusedAdd ?shape ?a_s ?b_s ?o_s ?dt)
(ICons ?inner_a (ICons ?fs_b (INil)))))
(let ?new_fe (Op (FusionEnd ?shape ?o_s ?dt) (ICons ?fbin (INil))))
(union ?bin ?new_fe)
) :ruleset fusion_grow :name "grow-FE-B-lhs-Add")
(rule (
(= ?fe (Op (FusionEnd ?shape ?b_s ?dt) (ICons ?inner_b (INil))))
(= ?bin (Op (KernelAdd ?shape ?a_s ?b_s ?o_s ?dt)
(ICons ?a (ICons ?fe (INil)))))
) (
(let ?fs_a (Op (FusionStart ?shape ?a_s ?dt) (ICons ?a (INil))))
(let ?fbin (Op (FusedAdd ?shape ?a_s ?b_s ?o_s ?dt)
(ICons ?fs_a (ICons ?inner_b (INil)))))
(let ?new_fe (Op (FusionEnd ?shape ?o_s ?dt) (ICons ?fbin (INil))))
(union ?bin ?new_fe)
) :ruleset fusion_grow :name "grow-FE-B-rhs-Add")
(rule (
(= ?fe (Op (FusionEnd ?shape ?a_s ?dt) (ICons ?inner_a (INil))))
(= ?bin (Op (KernelMul ?shape ?a_s ?b_s ?o_s ?dt)
(ICons ?fe (ICons ?b (INil)))))
) (
(let ?fs_b (Op (FusionStart ?shape ?b_s ?dt) (ICons ?b (INil))))
(let ?fbin (Op (FusedMul ?shape ?a_s ?b_s ?o_s ?dt)
(ICons ?inner_a (ICons ?fs_b (INil)))))
(let ?new_fe (Op (FusionEnd ?shape ?o_s ?dt) (ICons ?fbin (INil))))
(union ?bin ?new_fe)
) :ruleset fusion_grow :name "grow-FE-B-lhs-Mul")
(rule (
(= ?fe (Op (FusionEnd ?shape ?b_s ?dt) (ICons ?inner_b (INil))))
(= ?bin (Op (KernelMul ?shape ?a_s ?b_s ?o_s ?dt)
(ICons ?a (ICons ?fe (INil)))))
) (
(let ?fs_a (Op (FusionStart ?shape ?a_s ?dt) (ICons ?a (INil))))
(let ?fbin (Op (FusedMul ?shape ?a_s ?b_s ?o_s ?dt)
(ICons ?fs_a (ICons ?inner_b (INil)))))
(let ?new_fe (Op (FusionEnd ?shape ?o_s ?dt) (ICons ?fbin (INil))))
(union ?bin ?new_fe)
) :ruleset fusion_grow :name "grow-FE-B-rhs-Mul")
(rule (
(= ?fe_a (Op (FusionEnd ?shape ?a_s ?dt) (ICons ?inner_a (INil))))
(= ?fe_b (Op (FusionEnd ?shape ?b_s ?dt) (ICons ?inner_b (INil))))
(= ?bin (Op (KernelAdd ?shape ?a_s ?b_s ?o_s ?dt)
(ICons ?fe_a (ICons ?fe_b (INil)))))
) (
(let ?fbin (Op (FusedAdd ?shape ?a_s ?b_s ?o_s ?dt)
(ICons ?inner_a (ICons ?inner_b (INil)))))
(let ?new_fe (Op (FusionEnd ?shape ?o_s ?dt) (ICons ?fbin (INil))))
(union ?bin ?new_fe)
(subsume (Op (FusionEnd ?shape ?a_s ?dt) (ICons ?inner_a (INil))))
(subsume (Op (FusionEnd ?shape ?b_s ?dt) (ICons ?inner_b (INil))))
) :ruleset fusion_merge :name "merge-FE-FE-Add")
(rule (
(= ?fe_a (Op (FusionEnd ?shape ?a_s ?dt) (ICons ?inner_a (INil))))
(= ?fe_b (Op (FusionEnd ?shape ?b_s ?dt) (ICons ?inner_b (INil))))
(= ?bin (Op (KernelMul ?shape ?a_s ?b_s ?o_s ?dt)
(ICons ?fe_a (ICons ?fe_b (INil)))))
) (
(let ?fbin (Op (FusedMul ?shape ?a_s ?b_s ?o_s ?dt)
(ICons ?inner_a (ICons ?inner_b (INil)))))
(let ?new_fe (Op (FusionEnd ?shape ?o_s ?dt) (ICons ?fbin (INil))))
(union ?bin ?new_fe)
(subsume (Op (FusionEnd ?shape ?a_s ?dt) (ICons ?inner_a (INil))))
(subsume (Op (FusionEnd ?shape ?b_s ?dt) (ICons ?inner_b (INil))))
) :ruleset fusion_merge :name "merge-FE-FE-Mul")
(relation cublaslt_base_dtype (DType))
(cublaslt_base_dtype (F32))
(cublaslt_base_dtype (F16))
(cublaslt_base_dtype (Bf16))
(cublaslt_base_dtype (TF32))
; Row-major matmul: C[m,n] = A[m,k] × B[k,n]
; A[m,k] row-major → expand to [m, n, k] with strides [k, 0, MIter]
; B[k,n] row-major → permute to [n,k] then expand to [m, n, k] with strides [0, MIter, n]
;
; Row-major viewed as column-major (swap trick):
; Row-major A[m,k] ≡ column-major [k,m] with lda=k
; Row-major B[k,n] ≡ column-major [n,k] with ldb=n
; Row-major C[m,n] ≡ column-major [n,m] with ldc=n
;
; cuBLAS computes: C_col[n,m] = B_col[n,k] × A_col[k,m]
; cublasSgemm(OP_N, OP_N, n, m, k, α, B, n, A, k, β, C, n)
(rule
(
; Match Mul node
(= ?mul (Op (Mul ?mul_shape ?a_stride ?b_stride ?mul_out_stride) (ICons ?a (ICons ?b (INil)))))
; Match Sum that reduces the Mul (k dimension)
(= ?sum (Op (Sum ?out_shape ?k ?sum_in_stride ?k_stride ?sum_out_stride) (ICons ?mul (INil))))
; Match exactly 2D output shape
(= ?out_shape (ECons ?m (ECons ?n (ENil))))
(!= ?m (MNum 0))
(!= ?n (MNum 0))
(!= ?k (MNum 1))
; Match exactly 3D strides [m, n, k]
(= ?a_stride (ECons ?a_m_stride (ECons ?a_n_stride (ECons ?a_k_stride (ENil)))))
(= ?b_stride (ECons ?b_m_stride (ECons ?b_n_stride (ECons ?b_k_stride (ENil)))))
; Assert contiguous k stride on output (required for reduction)
(= ?k_stride (MIter))
; Assert A has strides [k*MIter, 0, MIter] (row-major A[m,k] broadcast to [m,n,k])
(= ?a_m_stride (MMul (MIter) ?k))
(= ?a_n_stride (MNum 0))
(= ?a_k_stride (MIter))
; Assert B has strides [0, MIter, n*MIter] (row-major B[k,n] permuted to [n,k] then broadcast to [m,n,k])
(= ?b_m_stride (MNum 0))
(= ?b_n_stride (MIter))
(= ?b_k_stride (MMul (MIter) ?n))
(= ?dt (dtype ?a))
(= ?dt (dtype ?b))
(cublaslt_base_dtype ?dt)
)
(
; For row-major C = A × B with cuBLAS (column-major):
; cublasSgemm(OP_N, OP_N, n, m, k, α, B, n, A, k, β, C, n)
(let ?sgemm (Op (cublaslt
?n ; cuBLAS m = our n (swapped)
?m ; cuBLAS n = our m (swapped)
?k ; k unchanged
"N" ; transa = No transpose
"N" ; transb = No transpose
"COL" "COL" "COL" "COL" ; A/B/C/D matrix orders
?b_k_stride ; lda = B's row stride (resolves to n after z→1)
?a_m_stride ; ldb = A's row stride (resolves to k after z→1)
?n ; ldc = n (row-major C[m,n] viewed as col-major [n,m])
?n ; ldd = ldc for current row-major output rewrites
(MNum 1) ; batch_count = 1
(MNum 0) ; stride_a = 0
(MNum 0) ; stride_b = 0
(MNum 0) ; stride_c = 0
(MNum 0) ; stride_d = 0
?dt ?dt ?dt ?dt "default" "default" 1.0 0.0 "DEFAULT") ; type tuple, alpha, beta
(ICons ?b (ICons ?a (INil)))))
(union ?sum ?sgemm)
(set (dtype ?sgemm) ?dt)
)
:ruleset matmul_backend
:name "cublaslt row-major x row-major"
)
; Batched Row-major × Row-major: C[batch,m,n] = A[batch,m,k] × B[batch,k,n]
; In broadcast [batch, m, n, k] space:
; A row-major per batch: a_k_stride=MIter, a_n_stride=0
; B row-major per batch: b_n_stride=MIter, b_m_stride=0
; Leading dimensions may differ from k/n when batch slices are non-contiguous.
(rule
(
(= ?mul (Op (Mul ?mul_shape ?a_stride ?b_stride ?mul_out_stride) (ICons ?a (ICons ?b (INil)))))
(= ?sum (Op (Sum ?out_shape ?k ?sum_in_stride ?k_stride ?sum_out_stride) (ICons ?mul (INil))))
; Output shape: [batch, m, n]
(= ?batch (nth_from_end ?out_shape 2))
(= ?m (nth_from_end ?out_shape 1))
(= ?n (nth_from_end ?out_shape 0))
(!= ?m (MNum 0))
(!= ?n (MNum 0))
(!= ?k (MNum 1))
(!= ?batch (MNum 0))
; A strides in [batch, m, n, k]
(= ?a_batch_stride (nth_from_end ?a_stride 3))
(= ?a_m_stride (nth_from_end ?a_stride 2))
(= ?a_n_stride (nth_from_end ?a_stride 1))
(= ?a_k_stride (nth_from_end ?a_stride 0))
; B strides in [batch, m, n, k]
(= ?b_batch_stride (nth_from_end ?b_stride 3))
(= ?b_m_stride (nth_from_end ?b_stride 2))
(= ?b_n_stride (nth_from_end ?b_stride 1))
(= ?b_k_stride (nth_from_end ?b_stride 0))
(= ?k_stride (MIter))
; A row-major: k=MIter, n=0, m_stride=k*MIter
(= ?a_k_stride (MIter))
(= ?a_n_stride (MNum 0))
(= ?a_m_stride (MMul (MIter) ?k))
; B row-major: n=MIter, m=0, k_stride=n*MIter
(= ?b_n_stride (MIter))
(= ?b_m_stride (MNum 0))
(= ?b_k_stride (MMul (MIter) ?n))
; Uniform batch strides (contiguous per batch, no GQA-style repetition)
(= ?a_batch_stride (MMul ?m ?a_m_stride))
(= ?b_batch_stride (MMul ?k ?b_k_stride))
(= ?dt (dtype ?a))
(= ?dt (dtype ?b))
(cublaslt_base_dtype ?dt)
)
(
; cuBLAS swap: C^T[n,m] = B^T[n,k] × A^T[k,m] per batch
; cublas(OP_N, OP_N, n, m, k, B, lda=b_k_stride, A, ldb=a_m_stride, C, ldc=n)
(let ?sgemm (Op (cublaslt
?n ?m ?k
"N" "N"
"COL" "COL" "COL" "COL"
?b_k_stride ; lda (cuBLAS A = our B, row stride)
?a_m_stride ; ldb (cuBLAS B = our A, row stride)
?n ; ldc (contiguous output per batch)
?n ; ldd
?batch ; batch_count
?b_batch_stride ; stride_a (cuBLAS A = our B)
?a_batch_stride ; stride_b (cuBLAS B = our A)
(MMul ?m ?n) ; stride_c
(MMul ?m ?n) ; stride_d
?dt ?dt ?dt ?dt "default" "default" 1.0 0.0 "DEFAULT")
(ICons ?b (ICons ?a (INil)))))
(union ?sum ?sgemm)
(set (dtype ?sgemm) ?dt)
)
:ruleset matmul_backend
:name "cublaslt batched row-major × row-major"
)
; Row-major × Column-major matmul: C[m,n] = A[m,k] × B[k,n]
; A[m,k] row-major → expand to [m, n, k] with strides [k, 0, MIter]
; B[k,n] column-major → expand to [m, n, k] with strides [0, k, MIter]
;
; Row-major viewed as column-major (swap trick):
; Row-major A[m,k] ≡ column-major A^T[k,m] with lda=k
; Column-major B[k,n] is already column-major with ldb=k
; Row-major C[m,n] ≡ column-major C^T[n,m] with ldc=n
;
; C^T[n,m] = (A × B)^T = B^T[n,k] × A^T[k,m]
; cuBLAS: cublasSgemm(OP_T, OP_N, n, m, k, α, B, k, A, k, β, C, n)
(rule
(
; Match Mul node
(= ?mul (Op (Mul ?mul_shape ?a_stride ?b_stride ?mul_out_stride) (ICons ?a (ICons ?b (INil)))))
; Match Sum that reduces the Mul (k dimension)
(= ?sum (Op (Sum ?out_shape ?k ?sum_in_stride ?k_stride ?sum_out_stride) (ICons ?mul (INil))))
; Match exactly 2D output shape
(= ?out_shape (ECons ?m (ECons ?n (ENil))))
(!= ?m (MNum 0))
(!= ?n (MNum 0))
(!= ?k (MNum 1))
; Match exactly 3D strides [m, n, k]
(= ?a_stride (ECons ?a_m_stride (ECons ?a_n_stride (ECons ?a_k_stride (ENil)))))
(= ?b_stride (ECons ?b_m_stride (ECons ?b_n_stride (ECons ?b_k_stride (ENil)))))
; Assert contiguous k stride on output (required for reduction)
(= ?k_stride (MIter))
; Assert A has strides [k*MIter, 0, MIter] (row-major A[m,k] broadcast to [m,n,k])
(= ?a_m_stride (MMul (MIter) ?k))
(= ?a_n_stride (MNum 0))
(= ?a_k_stride (MIter))
; Assert B has strides [0, k*MIter, MIter] (column-major B[k,n] broadcast to [m,n,k])
(= ?b_m_stride (MNum 0))
(= ?b_n_stride (MMul (MIter) ?k))
(= ?b_k_stride (MIter))
(= ?dt (dtype ?a))
(= ?dt (dtype ?b))
(cublaslt_base_dtype ?dt)
)
(
; For row-major A × column-major B with cuBLAS:
; C^T = B^T × A^T → cublasSgemm(OP_T, OP_N, n, m, k, α, B, k, A, k, β, C, n)
(let ?sgemm (Op (cublaslt
?n ; cuBLAS m = our n (swapped)
?m ; cuBLAS n = our m (swapped)
?k ; k unchanged
"T" ; transa = Transpose (B is column-major, need B^T)
"N" ; transb = No transpose
"COL" "COL" "COL" "COL" ; A/B/C/D matrix orders
?b_n_stride ; lda = B's column stride (resolves to k after z→1)
?a_m_stride ; ldb = A's row stride (resolves to k after z→1)
?n ; ldc = n (row-major C[m,n] viewed as col-major [n,m])
?n ; ldd = ldc for current row-major output rewrites
(MNum 1) ; batch_count = 1
(MNum 0) ; stride_a = 0
(MNum 0) ; stride_b = 0
(MNum 0) ; stride_c = 0
(MNum 0) ; stride_d = 0
?dt ?dt ?dt ?dt "default" "default" 1.0 0.0 "DEFAULT") ; type tuple, alpha, beta
(ICons ?b (ICons ?a (INil)))))
(union ?sum ?sgemm)
(set (dtype ?sgemm) ?dt)
)
:ruleset matmul_backend
:name "cublaslt row-major × column-major"
)
; Batched Row-major × Column-major: C[batch,m,n] = A[batch,m,k] × B[batch,k,n]
; A row-major per batch: a_k_stride=MIter, a_n_stride=0
; B column-major per batch: b_k_stride=MIter, b_m_stride=0
(rule
(
(= ?mul (Op (Mul ?mul_shape ?a_stride ?b_stride ?mul_out_stride) (ICons ?a (ICons ?b (INil)))))
(= ?sum (Op (Sum ?out_shape ?k ?sum_in_stride ?k_stride ?sum_out_stride) (ICons ?mul (INil))))
(= ?batch (nth_from_end ?out_shape 2))
(= ?m (nth_from_end ?out_shape 1))
(= ?n (nth_from_end ?out_shape 0))
(!= ?m (MNum 0))
(!= ?n (MNum 0))
(!= ?k (MNum 1))
(!= ?batch (MNum 0))
(= ?a_batch_stride (nth_from_end ?a_stride 3))
(= ?a_m_stride (nth_from_end ?a_stride 2))
(= ?a_n_stride (nth_from_end ?a_stride 1))
(= ?a_k_stride (nth_from_end ?a_stride 0))
(= ?b_batch_stride (nth_from_end ?b_stride 3))
(= ?b_m_stride (nth_from_end ?b_stride 2))
(= ?b_n_stride (nth_from_end ?b_stride 1))
(= ?b_k_stride (nth_from_end ?b_stride 0))
(= ?k_stride (MIter))
; A row-major: k=MIter, n=0, m_stride=k*MIter
(= ?a_k_stride (MIter))
(= ?a_n_stride (MNum 0))
(= ?a_m_stride (MMul (MIter) ?k))
; B column-major: k=MIter, m=0, n_stride=k*MIter
(= ?b_k_stride (MIter))
(= ?b_m_stride (MNum 0))
(= ?b_n_stride (MMul (MIter) ?k))
; Uniform batch strides (contiguous per batch)
(= ?a_batch_stride (MMul ?m ?a_m_stride))
(= ?b_batch_stride (MMul ?n ?b_n_stride))
(= ?dt (dtype ?a))
(= ?dt (dtype ?b))
(cublaslt_base_dtype ?dt)
)
(
; cuBLAS: cublas(OP_T, OP_N, n, m, k, B, lda=b_n_stride, A, ldb=a_m_stride, C, ldc=n)
(let ?sgemm (Op (cublaslt
?n ?m ?k
"T" "N"
"COL" "COL" "COL" "COL"
?b_n_stride ; lda (cuBLAS A = our B, column stride)
?a_m_stride ; ldb (cuBLAS B = our A, row stride)
?n ; ldc
?n ; ldd
?batch
?b_batch_stride ; stride_a (cuBLAS A = our B)
?a_batch_stride ; stride_b (cuBLAS B = our A)
(MMul ?m ?n) ; stride_c
(MMul ?m ?n) ; stride_d
?dt ?dt ?dt ?dt "default" "default" 1.0 0.0 "DEFAULT")
(ICons ?b (ICons ?a (INil)))))
(union ?sum ?sgemm)
(set (dtype ?sgemm) ?dt)
)
:ruleset matmul_backend
:name "cublaslt batched row-major × column-major"
)
; Column-major × Row-major matmul: C[m,n] = A[m,k] × B[k,n]
; A[m,k] column-major → expand to [m, n, k] with strides [MIter, 0, m]
; B[k,n] row-major → permute to [n,k] then expand to [m, n, k] with strides [0, MIter, n]
;
; Row-major viewed as column-major (swap trick):
; Column-major A[m,k] is already column-major with lda=m
; Row-major B[k,n] ≡ column-major B^T[n,k] with ldb=n
; Row-major C[m,n] ≡ column-major C^T[n,m] with ldc=n
;
; C^T[n,m] = (A × B)^T = B^T[n,k] × A^T[k,m]
; cuBLAS: cublasSgemm(OP_N, OP_T, n, m, k, α, B, n, A, m, β, C, n)
(rule
(
; Match Mul node
(= ?mul (Op (Mul ?mul_shape ?a_stride ?b_stride ?mul_out_stride) (ICons ?a (ICons ?b (INil)))))
; Match Sum that reduces the Mul (k dimension)
(= ?sum (Op (Sum ?out_shape ?k ?sum_in_stride ?k_stride ?sum_out_stride) (ICons ?mul (INil))))
; Match exactly 2D output shape
(= ?out_shape (ECons ?m (ECons ?n (ENil))))
(!= ?m (MNum 0))
(!= ?n (MNum 0))
(!= ?k (MNum 1))
; Match exactly 3D strides [m, n, k]
(= ?a_stride (ECons ?a_m_stride (ECons ?a_n_stride (ECons ?a_k_stride (ENil)))))
(= ?b_stride (ECons ?b_m_stride (ECons ?b_n_stride (ECons ?b_k_stride (ENil)))))
; Assert contiguous k stride on output (required for reduction)
(= ?k_stride (MIter))
; Assert A has strides [MIter, 0, m*MIter] (column-major A[m,k] broadcast to [m,n,k])
(= ?a_m_stride (MIter))
(= ?a_n_stride (MNum 0))
(= ?a_k_stride (MMul (MIter) ?m))
; Assert B has strides [0, MIter, n*MIter] (row-major B[k,n] permuted to [n,k] then broadcast to [m,n,k])
(= ?b_m_stride (MNum 0))
(= ?b_n_stride (MIter))
(= ?b_k_stride (MMul (MIter) ?n))
(= ?dt (dtype ?a))
(= ?dt (dtype ?b))
(cublaslt_base_dtype ?dt)
)
(
; For column-major A × row-major B with cuBLAS:
; C^T = B^T × A^T → cublasSgemm(OP_N, OP_T, n, m, k, α, B, n, A, m, β, C, n)
(let ?sgemm (Op (cublaslt
?n ; cuBLAS m = our n (swapped)
?m ; cuBLAS n = our m (swapped)
?k ; k unchanged
"N" ; transa = No transpose (B is row-major, viewed as col-major [n,k])
"T" ; transb = Transpose (A is column-major [m,k], need A^T[k,m])
"COL" "COL" "COL" "COL" ; A/B/C/D matrix orders
?b_k_stride ; lda = B's row stride (resolves to n after z→1)
?a_k_stride ; ldb = A's column stride (resolves to m after z→1)
?n ; ldc = n (row-major C[m,n] viewed as col-major [n,m])
?n ; ldd = ldc for current row-major output rewrites
(MNum 1) ; batch_count = 1
(MNum 0) ; stride_a = 0
(MNum 0) ; stride_b = 0
(MNum 0) ; stride_c = 0
(MNum 0) ; stride_d = 0
?dt ?dt ?dt ?dt "default" "default" 1.0 0.0 "DEFAULT") ; type tuple, alpha, beta
(ICons ?b (ICons ?a (INil)))))
(union ?sum ?sgemm)
(set (dtype ?sgemm) ?dt)
)
:ruleset matmul_backend
:name "cublaslt column-major × row-major"
)
; Batched Column-major × Row-major: C[batch,m,n] = A[batch,m,k] × B[batch,k,n]
; A column-major per batch: a_m_stride=MIter, a_n_stride=0
; B row-major per batch: b_n_stride=MIter, b_m_stride=0
(rule
(
(= ?mul (Op (Mul ?mul_shape ?a_stride ?b_stride ?mul_out_stride) (ICons ?a (ICons ?b (INil)))))
(= ?sum (Op (Sum ?out_shape ?k ?sum_in_stride ?k_stride ?sum_out_stride) (ICons ?mul (INil))))
(= ?batch (nth_from_end ?out_shape 2))
(= ?m (nth_from_end ?out_shape 1))
(= ?n (nth_from_end ?out_shape 0))
(!= ?m (MNum 0))
(!= ?n (MNum 0))
(!= ?k (MNum 1))
(!= ?batch (MNum 0))
(= ?a_batch_stride (nth_from_end ?a_stride 3))
(= ?a_m_stride (nth_from_end ?a_stride 2))
(= ?a_n_stride (nth_from_end ?a_stride 1))
(= ?a_k_stride (nth_from_end ?a_stride 0))
(= ?b_batch_stride (nth_from_end ?b_stride 3))
(= ?b_m_stride (nth_from_end ?b_stride 2))
(= ?b_n_stride (nth_from_end ?b_stride 1))
(= ?b_k_stride (nth_from_end ?b_stride 0))
(= ?k_stride (MIter))
; A column-major: m=MIter, n=0, k_stride=m*MIter
(= ?a_m_stride (MIter))
(= ?a_n_stride (MNum 0))
(= ?a_k_stride (MMul (MIter) ?m))
; B row-major: n=MIter, m=0, k_stride=n*MIter
(= ?b_n_stride (MIter))
(= ?b_m_stride (MNum 0))
(= ?b_k_stride (MMul (MIter) ?n))
; Uniform batch strides (contiguous per batch)
(= ?a_batch_stride (MMul ?k ?a_k_stride))
(= ?b_batch_stride (MMul ?k ?b_k_stride))
(= ?dt (dtype ?a))
(= ?dt (dtype ?b))
(cublaslt_base_dtype ?dt)
)
(
; cuBLAS: cublas(OP_N, OP_T, n, m, k, B, lda=b_k_stride, A, ldb=a_k_stride, C, ldc=n)
(let ?sgemm (Op (cublaslt
?n ?m ?k
"N" "T"
"COL" "COL" "COL" "COL"
?b_k_stride ; lda (cuBLAS A = our B, row stride)
?a_k_stride ; ldb (cuBLAS B = our A, column stride)
?n ; ldc
?n ; ldd
?batch
?b_batch_stride ; stride_a (cuBLAS A = our B)
?a_batch_stride ; stride_b (cuBLAS B = our A)
(MMul ?m ?n) ; stride_c
(MMul ?m ?n) ; stride_d
?dt ?dt ?dt ?dt "default" "default" 1.0 0.0 "DEFAULT")
(ICons ?b (ICons ?a (INil)))))
(union ?sum ?sgemm)
(set (dtype ?sgemm) ?dt)
)
:ruleset matmul_backend
:name "cublaslt batched column-major × row-major"
)
; Column-major × Column-major matmul: C[m,n] = A[m,k] × B[k,n]
; A[m,k] column-major → expand to [m, n, k] with strides [MIter, 0, m]
; B[k,n] column-major → expand to [m, n, k] with strides [0, k, MIter]
;
; Row-major viewed as column-major (swap trick):
; Column-major A[m,k] is already column-major with lda=m
; Column-major B[k,n] is already column-major with ldb=k
; Row-major C[m,n] ≡ column-major C^T[n,m] with ldc=n
;
; C^T[n,m] = (A × B)^T = B^T[n,k] × A^T[k,m]
; cuBLAS: cublasSgemm(OP_T, OP_T, n, m, k, α, B, k, A, m, β, C, n)
(rule
(
; Match Mul node
(= ?mul (Op (Mul ?mul_shape ?a_stride ?b_stride ?mul_out_stride) (ICons ?a (ICons ?b (INil)))))
; Match Sum that reduces the Mul (k dimension)
(= ?sum (Op (Sum ?out_shape ?k ?sum_in_stride ?k_stride ?sum_out_stride) (ICons ?mul (INil))))
; Match exactly 2D output shape
(= ?out_shape (ECons ?m (ECons ?n (ENil))))
(!= ?m (MNum 0))
(!= ?n (MNum 0))
(!= ?k (MNum 1))
; Match exactly 3D strides [m, n, k]
(= ?a_stride (ECons ?a_m_stride (ECons ?a_n_stride (ECons ?a_k_stride (ENil)))))
(= ?b_stride (ECons ?b_m_stride (ECons ?b_n_stride (ECons ?b_k_stride (ENil)))))
; Assert contiguous k stride on output (required for reduction)
(= ?k_stride (MIter))
; Assert A has strides [MIter, 0, m*MIter] (column-major A[m,k] broadcast to [m,n,k])
(= ?a_m_stride (MIter))
(= ?a_n_stride (MNum 0))
(= ?a_k_stride (MMul (MIter) ?m))
; Assert B has strides [0, k*MIter, MIter] (column-major B[k,n] broadcast to [m,n,k])
(= ?b_m_stride (MNum 0))
(= ?b_n_stride (MMul (MIter) ?k))
(= ?b_k_stride (MIter))
(= ?dt (dtype ?a))
(= ?dt (dtype ?b))
(cublaslt_base_dtype ?dt)
)
(
; For column-major A × column-major B with cuBLAS:
; C^T = B^T × A^T → cublasSgemm(OP_T, OP_T, n, m, k, α, B, k, A, m, β, C, n)
(let ?sgemm (Op (cublaslt
?n ; cuBLAS m = our n (swapped)
?m ; cuBLAS n = our m (swapped)
?k ; k unchanged
"T" ; transa = Transpose (B is column-major [k,n], need B^T[n,k])
"T" ; transb = Transpose (A is column-major [m,k], need A^T[k,m])
"COL" "COL" "COL" "COL" ; A/B/C/D matrix orders
?b_n_stride ; lda = B's column stride (resolves to k after z→1)
?a_k_stride ; ldb = A's column stride (resolves to m after z→1)
?n ; ldc = n (row-major C[m,n] viewed as col-major [n,m])
?n ; ldd = ldc for current row-major output rewrites
(MNum 1) ; batch_count = 1
(MNum 0) ; stride_a = 0
(MNum 0) ; stride_b = 0
(MNum 0) ; stride_c = 0
(MNum 0) ; stride_d = 0
?dt ?dt ?dt ?dt "default" "default" 1.0 0.0 "DEFAULT") ; type tuple, alpha, beta
(ICons ?b (ICons ?a (INil)))))
(union ?sum ?sgemm)
(set (dtype ?sgemm) ?dt)
)
:ruleset matmul_backend
:name "cublaslt column-major × column-major"
)
; Batched Column-major × Column-major: C[batch,m,n] = A[batch,m,k] × B[batch,k,n]
; A column-major per batch: a_m_stride=MIter, a_n_stride=0
; B column-major per batch: b_k_stride=MIter, b_m_stride=0
(rule
(
(= ?mul (Op (Mul ?mul_shape ?a_stride ?b_stride ?mul_out_stride) (ICons ?a (ICons ?b (INil)))))
(= ?sum (Op (Sum ?out_shape ?k ?sum_in_stride ?k_stride ?sum_out_stride) (ICons ?mul (INil))))
(= ?batch (nth_from_end ?out_shape 2))
(= ?m (nth_from_end ?out_shape 1))
(= ?n (nth_from_end ?out_shape 0))
(!= ?m (MNum 0))
(!= ?n (MNum 0))
(!= ?k (MNum 1))
(!= ?batch (MNum 0))
(= ?a_batch_stride (nth_from_end ?a_stride 3))
(= ?a_m_stride (nth_from_end ?a_stride 2))
(= ?a_n_stride (nth_from_end ?a_stride 1))
(= ?a_k_stride (nth_from_end ?a_stride 0))
(= ?b_batch_stride (nth_from_end ?b_stride 3))
(= ?b_m_stride (nth_from_end ?b_stride 2))
(= ?b_n_stride (nth_from_end ?b_stride 1))
(= ?b_k_stride (nth_from_end ?b_stride 0))
(= ?k_stride (MIter))
; A column-major: m=MIter, n=0, k_stride=m*MIter
(= ?a_m_stride (MIter))
(= ?a_n_stride (MNum 0))
(= ?a_k_stride (MMul (MIter) ?m))
; B column-major: k=MIter, m=0, n_stride=k*MIter
(= ?b_k_stride (MIter))
(= ?b_m_stride (MNum 0))
(= ?b_n_stride (MMul (MIter) ?k))
; Uniform batch strides (contiguous per batch)
(= ?a_batch_stride (MMul ?k ?a_k_stride))
(= ?b_batch_stride (MMul ?n ?b_n_stride))
(= ?dt (dtype ?a))
(= ?dt (dtype ?b))
(cublaslt_base_dtype ?dt)
)
(
; cuBLAS: cublas(OP_T, OP_T, n, m, k, B, lda=b_n_stride, A, ldb=a_k_stride, C, ldc=n)
(let ?sgemm (Op (cublaslt
?n ?m ?k
"T" "T"
"COL" "COL" "COL" "COL"
?b_n_stride ; lda (cuBLAS A = our B, column stride)
?a_k_stride ; ldb (cuBLAS B = our A, column stride)
?n ; ldc
?n ; ldd
?batch
?b_batch_stride ; stride_a (cuBLAS A = our B)
?a_batch_stride ; stride_b (cuBLAS B = our A)
(MMul ?m ?n) ; stride_c
(MMul ?m ?n) ; stride_d
?dt ?dt ?dt ?dt "default" "default" 1.0 0.0 "DEFAULT")
(ICons ?b (ICons ?a (INil)))))
(union ?sum ?sgemm)
(set (dtype ?sgemm) ?dt)
)
:ruleset matmul_backend
:name "cublaslt batched column-major × column-major"
)
; FP8 support is narrower than "any FP8 x any FP8". cuBLASLt's regular FP8
; matmul table supports these A/B descriptor pairs for F32 outputs:
; E4M3 x E4M3
; E4M3 x E5M2
; E5M2 x E4M3
; and requires TN format on Ada/Hopper-class GPUs. These rules therefore match
; row-major x column-major Luminal matmuls, which the existing COL-order lowering
; describes as descriptor A = logical B, descriptor B = logical A, transa=T,
; transb=N.
(rule
(
(= ?mul (Op (Mul ?mul_shape ?a_stride ?b_stride ?mul_out_stride) (ICons ?a (ICons ?b (INil)))))
(= ?sum (Op (Sum ?out_shape ?k ?sum_in_stride ?k_stride ?sum_out_stride) (ICons ?mul (INil))))
(= ?cast (Op (Cast ?size (F32)) (ICons ?sum (INil))))
(= ?out_shape (ECons ?m (ECons ?n (ENil))))
(!= ?m (MNum 0))
(!= ?n (MNum 0))
(!= ?k (MNum 1))
(= ?a_stride (ECons ?a_m_stride (ECons ?a_n_stride (ECons ?a_k_stride (ENil)))))
(= ?b_stride (ECons ?b_m_stride (ECons ?b_n_stride (ECons ?b_k_stride (ENil)))))
(= ?k_stride (MIter))
(= ?a_m_stride (MMul (MIter) ?k))
(= ?a_n_stride (MNum 0))
(= ?a_k_stride (MIter))
(= ?b_m_stride (MNum 0))
(= ?b_n_stride (MMul (MIter) ?k))
(= ?b_k_stride (MIter))
(= (F8E4M3) (dtype ?a))
(= (F8E4M3) (dtype ?b))
)
(
(let ?sgemm (Op (cublaslt
?n ?m ?k
"T" "N"
"COL" "COL" "COL" "COL"
?b_n_stride
?a_m_stride
?n
?n
(MNum 1)
(MNum 0)
(MNum 0)
(MNum 0)
(MNum 0)
(F8E4M3) (F8E4M3) (F32) (F32) "32F" "F32" 1.0 0.0 "DEFAULT")
(ICons ?b (ICons ?a (INil)))))
(union ?cast ?sgemm)
(set (dtype ?sgemm) (F32))
)
:ruleset matmul_backend
:name "cublaslt fp8 e4m3/e4m3 row-major x column-major f32 output"
)
(rule
(
(= ?mul (Op (Mul ?mul_shape ?a_stride ?b_stride ?mul_out_stride) (ICons ?a (ICons ?b (INil)))))
(= ?sum (Op (Sum ?out_shape ?k ?sum_in_stride ?k_stride ?sum_out_stride) (ICons ?mul (INil))))
(= ?cast (Op (Cast ?size (F32)) (ICons ?sum (INil))))
(= ?out_shape (ECons ?m (ECons ?n (ENil))))
(!= ?m (MNum 0))
(!= ?n (MNum 0))
(!= ?k (MNum 1))
(= ?a_stride (ECons ?a_m_stride (ECons ?a_n_stride (ECons ?a_k_stride (ENil)))))
(= ?b_stride (ECons ?b_m_stride (ECons ?b_n_stride (ECons ?b_k_stride (ENil)))))
(= ?k_stride (MIter))
(= ?a_m_stride (MMul (MIter) ?k))
(= ?a_n_stride (MNum 0))
(= ?a_k_stride (MIter))
(= ?b_m_stride (MNum 0))
(= ?b_n_stride (MMul (MIter) ?k))
(= ?b_k_stride (MIter))
(= (F8E4M3) (dtype ?a))
(= (F8E5M2) (dtype ?b))
)
(
(let ?sgemm (Op (cublaslt
?n ?m ?k
"T" "N"
"COL" "COL" "COL" "COL"
?b_n_stride
?a_m_stride
?n
?n
(MNum 1)
(MNum 0)
(MNum 0)
(MNum 0)
(MNum 0)
(F8E5M2) (F8E4M3) (F32) (F32) "32F" "F32" 1.0 0.0 "DEFAULT")
(ICons ?b (ICons ?a (INil)))))
(union ?cast ?sgemm)
(set (dtype ?sgemm) (F32))
)
:ruleset matmul_backend
:name "cublaslt fp8 e5m2/e4m3 row-major x column-major f32 output"
)
(rule
(
(= ?mul (Op (Mul ?mul_shape ?a_stride ?b_stride ?mul_out_stride) (ICons ?a (ICons ?b (INil)))))
(= ?sum (Op (Sum ?out_shape ?k ?sum_in_stride ?k_stride ?sum_out_stride) (ICons ?mul (INil))))
(= ?cast (Op (Cast ?size (F32)) (ICons ?sum (INil))))
(= ?out_shape (ECons ?m (ECons ?n (ENil))))
(!= ?m (MNum 0))
(!= ?n (MNum 0))
(!= ?k (MNum 1))
(= ?a_stride (ECons ?a_m_stride (ECons ?a_n_stride (ECons ?a_k_stride (ENil)))))
(= ?b_stride (ECons ?b_m_stride (ECons ?b_n_stride (ECons ?b_k_stride (ENil)))))
(= ?k_stride (MIter))
(= ?a_m_stride (MMul (MIter) ?k))
(= ?a_n_stride (MNum 0))
(= ?a_k_stride (MIter))
(= ?b_m_stride (MNum 0))
(= ?b_n_stride (MMul (MIter) ?k))
(= ?b_k_stride (MIter))
(= (F8E5M2) (dtype ?a))
(= (F8E4M3) (dtype ?b))
)
(
(let ?sgemm (Op (cublaslt
?n ?m ?k
"T" "N"
"COL" "COL" "COL" "COL"
?b_n_stride
?a_m_stride
?n
?n
(MNum 1)
(MNum 0)
(MNum 0)
(MNum 0)
(MNum 0)
(F8E4M3) (F8E5M2) (F32) (F32) "32F" "F32" 1.0 0.0 "DEFAULT")
(ICons ?b (ICons ?a (INil)))))
(union ?cast ?sgemm)
(set (dtype ?sgemm) (F32))
)
:ruleset matmul_backend
:name "cublaslt fp8 e4m3/e5m2 row-major x column-major f32 output"
)
(rule
(
(= ?mul (Op (Mul ?mul_shape ?a_stride ?b_stride ?mul_out_stride) (ICons ?a (ICons ?b (INil)))))
(= ?sum (Op (Sum ?out_shape ?k ?sum_in_stride ?k_stride ?sum_out_stride) (ICons ?mul (INil))))
(= ?cast (Op (Cast ?size (F32)) (ICons ?sum (INil))))
(= ?batch (nth_from_end ?out_shape 2))
(= ?m (nth_from_end ?out_shape 1))
(= ?n (nth_from_end ?out_shape 0))
(!= ?m (MNum 0))
(!= ?n (MNum 0))
(!= ?k (MNum 1))
(!= ?batch (MNum 0))
(= ?a_batch_stride (nth_from_end ?a_stride 3))
(= ?a_m_stride (nth_from_end ?a_stride 2))
(= ?a_n_stride (nth_from_end ?a_stride 1))
(= ?a_k_stride (nth_from_end ?a_stride 0))
(= ?b_batch_stride (nth_from_end ?b_stride 3))
(= ?b_m_stride (nth_from_end ?b_stride 2))
(= ?b_n_stride (nth_from_end ?b_stride 1))
(= ?b_k_stride (nth_from_end ?b_stride 0))
(= ?k_stride (MIter))
(= ?a_k_stride (MIter))
(= ?a_n_stride (MNum 0))
(= ?a_m_stride (MMul (MIter) ?k))
(= ?b_k_stride (MIter))
(= ?b_m_stride (MNum 0))
(= ?b_n_stride (MMul (MIter) ?k))
(= ?a_batch_stride (MMul ?m ?a_m_stride))
(= ?b_batch_stride (MMul ?n ?b_n_stride))
(= (F8E4M3) (dtype ?a))
(= (F8E4M3) (dtype ?b))
)
(
(let ?sgemm (Op (cublaslt
?n ?m ?k
"T" "N"
"COL" "COL" "COL" "COL"
?b_n_stride
?a_m_stride
?n
?n
?batch
?b_batch_stride
?a_batch_stride
(MMul ?m ?n)
(MMul ?m ?n)
(F8E4M3) (F8E4M3) (F32) (F32) "32F" "F32" 1.0 0.0 "DEFAULT")
(ICons ?b (ICons ?a (INil)))))
(union ?cast ?sgemm)
(set (dtype ?sgemm) (F32))
)
:ruleset matmul_backend
:name "cublaslt fp8 e4m3/e4m3 batched row-major x column-major f32 output"
)
(rule
(
(= ?mul (Op (Mul ?mul_shape ?a_stride ?b_stride ?mul_out_stride) (ICons ?a (ICons ?b (INil)))))
(= ?sum (Op (Sum ?out_shape ?k ?sum_in_stride ?k_stride ?sum_out_stride) (ICons ?mul (INil))))
(= ?cast (Op (Cast ?size (F32)) (ICons ?sum (INil))))
(= ?batch (nth_from_end ?out_shape 2))
(= ?m (nth_from_end ?out_shape 1))
(= ?n (nth_from_end ?out_shape 0))
(!= ?m (MNum 0))
(!= ?n (MNum 0))
(!= ?k (MNum 1))
(!= ?batch (MNum 0))
(= ?a_batch_stride (nth_from_end ?a_stride 3))
(= ?a_m_stride (nth_from_end ?a_stride 2))
(= ?a_n_stride (nth_from_end ?a_stride 1))
(= ?a_k_stride (nth_from_end ?a_stride 0))
(= ?b_batch_stride (nth_from_end ?b_stride 3))
(= ?b_m_stride (nth_from_end ?b_stride 2))
(= ?b_n_stride (nth_from_end ?b_stride 1))
(= ?b_k_stride (nth_from_end ?b_stride 0))
(= ?k_stride (MIter))
(= ?a_k_stride (MIter))
(= ?a_n_stride (MNum 0))
(= ?a_m_stride (MMul (MIter) ?k))
(= ?b_k_stride (MIter))
(= ?b_m_stride (MNum 0))
(= ?b_n_stride (MMul (MIter) ?k))
(= ?a_batch_stride (MMul ?m ?a_m_stride))
(= ?b_batch_stride (MMul ?n ?b_n_stride))
(= (F8E4M3) (dtype ?a))
(= (F8E5M2) (dtype ?b))
)
(
(let ?sgemm (Op (cublaslt
?n ?m ?k
"T" "N"
"COL" "COL" "COL" "COL"
?b_n_stride
?a_m_stride
?n
?n
?batch
?b_batch_stride
?a_batch_stride
(MMul ?m ?n)
(MMul ?m ?n)
(F8E5M2) (F8E4M3) (F32) (F32) "32F" "F32" 1.0 0.0 "DEFAULT")
(ICons ?b (ICons ?a (INil)))))
(union ?cast ?sgemm)
(set (dtype ?sgemm) (F32))
)
:ruleset matmul_backend
:name "cublaslt fp8 e5m2/e4m3 batched row-major x column-major f32 output"
)
(rule
(
(= ?mul (Op (Mul ?mul_shape ?a_stride ?b_stride ?mul_out_stride) (ICons ?a (ICons ?b (INil)))))
(= ?sum (Op (Sum ?out_shape ?k ?sum_in_stride ?k_stride ?sum_out_stride) (ICons ?mul (INil))))
(= ?cast (Op (Cast ?size (F32)) (ICons ?sum (INil))))
(= ?batch (nth_from_end ?out_shape 2))
(= ?m (nth_from_end ?out_shape 1))
(= ?n (nth_from_end ?out_shape 0))
(!= ?m (MNum 0))
(!= ?n (MNum 0))
(!= ?k (MNum 1))
(!= ?batch (MNum 0))
(= ?a_batch_stride (nth_from_end ?a_stride 3))
(= ?a_m_stride (nth_from_end ?a_stride 2))
(= ?a_n_stride (nth_from_end ?a_stride 1))
(= ?a_k_stride (nth_from_end ?a_stride 0))
(= ?b_batch_stride (nth_from_end ?b_stride 3))
(= ?b_m_stride (nth_from_end ?b_stride 2))
(= ?b_n_stride (nth_from_end ?b_stride 1))
(= ?b_k_stride (nth_from_end ?b_stride 0))
(= ?k_stride (MIter))
(= ?a_k_stride (MIter))
(= ?a_n_stride (MNum 0))
(= ?a_m_stride (MMul (MIter) ?k))
(= ?b_k_stride (MIter))
(= ?b_m_stride (MNum 0))
(= ?b_n_stride (MMul (MIter) ?k))
(= ?a_batch_stride (MMul ?m ?a_m_stride))
(= ?b_batch_stride (MMul ?n ?b_n_stride))
(= (F8E5M2) (dtype ?a))
(= (F8E4M3) (dtype ?b))
)
(
(let ?sgemm (Op (cublaslt
?n ?m ?k
"T" "N"
"COL" "COL" "COL" "COL"
?b_n_stride
?a_m_stride
?n
?n
?batch
?b_batch_stride
?a_batch_stride
(MMul ?m ?n)
(MMul ?m ?n)
(F8E4M3) (F8E5M2) (F32) (F32) "32F" "F32" 1.0 0.0 "DEFAULT")
(ICons ?b (ICons ?a (INil)))))
(union ?cast ?sgemm)
(set (dtype ?sgemm) (F32))
)
:ruleset matmul_backend
:name "cublaslt fp8 e4m3/e5m2 batched row-major x column-major f32 output"
)
; Natural cuBLASLt row-order output rewrites. These keep Luminal's logical
; output C[m,n] as a cuBLASLt ROW-ordered D[m,n] instead of using the older
; swapped COL-ordered D[n,m] view. A and B orders mirror their matched logical
; layouts, so this family is the legal base for future ROW-ordered beta fusions.
(rule
(
(= ?mul (Op (Mul ?mul_shape ?a_stride ?b_stride ?mul_out_stride) (ICons ?a (ICons ?b (INil)))))
(= ?sum (Op (Sum ?out_shape ?k ?sum_in_stride ?k_stride ?sum_out_stride) (ICons ?mul (INil))))
(= ?out_shape (ECons ?m (ECons ?n (ENil))))
(!= ?m (MNum 0))
(!= ?n (MNum 0))
(!= ?k (MNum 1))
(= ?a_stride (ECons ?a_m_stride (ECons ?a_n_stride (ECons ?a_k_stride (ENil)))))
(= ?b_stride (ECons ?b_m_stride (ECons ?b_n_stride (ECons ?b_k_stride (ENil)))))
(= ?k_stride (MIter))
(= ?a_m_stride (MMul (MIter) ?k))
(= ?a_n_stride (MNum 0))
(= ?a_k_stride (MIter))
(= ?b_m_stride (MNum 0))
(= ?b_n_stride (MIter))
(= ?b_k_stride (MMul (MIter) ?n))
(= ?dt (dtype ?a))
(= ?dt (dtype ?b))
(cublaslt_base_dtype ?dt)
)
(
(let ?sgemm (Op (cublaslt
?m ?n ?k
"N" "N"
"ROW" "ROW" "ROW" "ROW"
?a_m_stride
?b_k_stride
?n
?n
(MNum 1)
(MNum 0)
(MNum 0)
(MNum 0)
(MNum 0)
?dt ?dt ?dt ?dt "default" "default" 1.0 0.0 "DEFAULT")
(ICons ?a (ICons ?b (INil)))))
(union ?sum ?sgemm)
(set (dtype ?sgemm) ?dt)
)
:ruleset matmul_backend
:name "cublaslt row-order row-major x row-major"
)
(rule
(
(= ?mul (Op (Mul ?mul_shape ?a_stride ?b_stride ?mul_out_stride) (ICons ?a (ICons ?b (INil)))))
(= ?sum (Op (Sum ?out_shape ?k ?sum_in_stride ?k_stride ?sum_out_stride) (ICons ?mul (INil))))
(= ?out_shape (ECons ?m (ECons ?n (ENil))))
(!= ?m (MNum 0))
(!= ?n (MNum 0))
(!= ?k (MNum 1))
(= ?a_stride (ECons ?a_m_stride (ECons ?a_n_stride (ECons ?a_k_stride (ENil)))))
(= ?b_stride (ECons ?b_m_stride (ECons ?b_n_stride (ECons ?b_k_stride (ENil)))))
(= ?k_stride (MIter))
(= ?a_m_stride (MMul (MIter) ?k))
(= ?a_n_stride (MNum 0))
(= ?a_k_stride (MIter))
(= ?b_m_stride (MNum 0))
(= ?b_n_stride (MMul (MIter) ?k))
(= ?b_k_stride (MIter))
(= ?dt (dtype ?a))
(= ?dt (dtype ?b))
(cublaslt_base_dtype ?dt)
)
(
(let ?sgemm (Op (cublaslt
?m ?n ?k
"N" "N"
"ROW" "COL" "ROW" "ROW"
?a_m_stride
?b_n_stride
?n
?n
(MNum 1)
(MNum 0)
(MNum 0)
(MNum 0)
(MNum 0)
?dt ?dt ?dt ?dt "default" "default" 1.0 0.0 "DEFAULT")
(ICons ?a (ICons ?b (INil)))))
(union ?sum ?sgemm)
(set (dtype ?sgemm) ?dt)
)
:ruleset matmul_backend
:name "cublaslt row-order row-major x column-major"
)
(rule
(
(= ?mul (Op (Mul ?mul_shape ?a_stride ?b_stride ?mul_out_stride) (ICons ?a (ICons ?b (INil)))))
(= ?sum (Op (Sum ?out_shape ?k ?sum_in_stride ?k_stride ?sum_out_stride) (ICons ?mul (INil))))
(= ?out_shape (ECons ?m (ECons ?n (ENil))))
(!= ?m (MNum 0))
(!= ?n (MNum 0))
(!= ?k (MNum 1))
(= ?a_stride (ECons ?a_m_stride (ECons ?a_n_stride (ECons ?a_k_stride (ENil)))))
(= ?b_stride (ECons ?b_m_stride (ECons ?b_n_stride (ECons ?b_k_stride (ENil)))))
(= ?k_stride (MIter))
(= ?a_m_stride (MIter))
(= ?a_n_stride (MNum 0))
(= ?a_k_stride (MMul (MIter) ?m))
(= ?b_m_stride (MNum 0))
(= ?b_n_stride (MIter))
(= ?b_k_stride (MMul (MIter) ?n))
(= ?dt (dtype ?a))
(= ?dt (dtype ?b))
(cublaslt_base_dtype ?dt)
)
(
(let ?sgemm (Op (cublaslt
?m ?n ?k
"N" "N"
"COL" "ROW" "ROW" "ROW"
?a_k_stride
?b_k_stride
?n
?n
(MNum 1)
(MNum 0)
(MNum 0)
(MNum 0)
(MNum 0)
?dt ?dt ?dt ?dt "default" "default" 1.0 0.0 "DEFAULT")
(ICons ?a (ICons ?b (INil)))))
(union ?sum ?sgemm)
(set (dtype ?sgemm) ?dt)
)
:ruleset matmul_backend
:name "cublaslt row-order column-major x row-major"
)
(rule
(
(= ?mul (Op (Mul ?mul_shape ?a_stride ?b_stride ?mul_out_stride) (ICons ?a (ICons ?b (INil)))))
(= ?sum (Op (Sum ?out_shape ?k ?sum_in_stride ?k_stride ?sum_out_stride) (ICons ?mul (INil))))
(= ?out_shape (ECons ?m (ECons ?n (ENil))))
(!= ?m (MNum 0))
(!= ?n (MNum 0))
(!= ?k (MNum 1))
(= ?a_stride (ECons ?a_m_stride (ECons ?a_n_stride (ECons ?a_k_stride (ENil)))))
(= ?b_stride (ECons ?b_m_stride (ECons ?b_n_stride (ECons ?b_k_stride (ENil)))))
(= ?k_stride (MIter))
(= ?a_m_stride (MIter))
(= ?a_n_stride (MNum 0))
(= ?a_k_stride (MMul (MIter) ?m))
(= ?b_m_stride (MNum 0))
(= ?b_n_stride (MMul (MIter) ?k))
(= ?b_k_stride (MIter))
(= ?dt (dtype ?a))
(= ?dt (dtype ?b))
(cublaslt_base_dtype ?dt)
)
(
(let ?sgemm (Op (cublaslt
?m ?n ?k
"N" "N"
"COL" "COL" "ROW" "ROW"
?a_k_stride
?b_n_stride
?n
?n
(MNum 1)
(MNum 0)
(MNum 0)
(MNum 0)
(MNum 0)
?dt ?dt ?dt ?dt "default" "default" 1.0 0.0 "DEFAULT")
(ICons ?a (ICons ?b (INil)))))
(union ?sum ?sgemm)
(set (dtype ?sgemm) ?dt)
)
:ruleset matmul_backend
:name "cublaslt row-order column-major x column-major"
)
(rule
(
(= ?mul (Op (Mul ?mul_shape ?a_stride ?b_stride ?mul_out_stride) (ICons ?a (ICons ?b (INil)))))
(= ?sum (Op (Sum ?out_shape ?k ?sum_in_stride ?k_stride ?sum_out_stride) (ICons ?mul (INil))))
(= ?batch (nth_from_end ?out_shape 2))
(= ?m (nth_from_end ?out_shape 1))
(= ?n (nth_from_end ?out_shape 0))
(!= ?m (MNum 0))
(!= ?n (MNum 0))
(!= ?k (MNum 1))
(!= ?batch (MNum 0))
(= ?a_batch_stride (nth_from_end ?a_stride 3))
(= ?a_m_stride (nth_from_end ?a_stride 2))
(= ?a_n_stride (nth_from_end ?a_stride 1))
(= ?a_k_stride (nth_from_end ?a_stride 0))
(= ?b_batch_stride (nth_from_end ?b_stride 3))
(= ?b_m_stride (nth_from_end ?b_stride 2))
(= ?b_n_stride (nth_from_end ?b_stride 1))
(= ?b_k_stride (nth_from_end ?b_stride 0))
(= ?k_stride (MIter))
(= ?a_k_stride (MIter))
(= ?a_n_stride (MNum 0))
(= ?a_m_stride (MMul (MIter) ?k))
(= ?b_n_stride (MIter))
(= ?b_m_stride (MNum 0))
(= ?b_k_stride (MMul (MIter) ?n))
(= ?a_batch_stride (MMul ?m ?a_m_stride))
(= ?b_batch_stride (MMul ?k ?b_k_stride))
(= ?dt (dtype ?a))
(= ?dt (dtype ?b))
(cublaslt_base_dtype ?dt)
)
(
(let ?sgemm (Op (cublaslt
?m ?n ?k
"N" "N"
"ROW" "ROW" "ROW" "ROW"
?a_m_stride
?b_k_stride
?n
?n
?batch
?a_batch_stride
?b_batch_stride
(MMul ?m ?n)
(MMul ?m ?n)
?dt ?dt ?dt ?dt "default" "default" 1.0 0.0 "DEFAULT")
(ICons ?a (ICons ?b (INil)))))
(union ?sum ?sgemm)
(set (dtype ?sgemm) ?dt)
)
:ruleset matmul_backend
:name "cublaslt row-order batched row-major x row-major"
)
(rule
(
(= ?mul (Op (Mul ?mul_shape ?a_stride ?b_stride ?mul_out_stride) (ICons ?a (ICons ?b (INil)))))
(= ?sum (Op (Sum ?out_shape ?k ?sum_in_stride ?k_stride ?sum_out_stride) (ICons ?mul (INil))))
(= ?batch (nth_from_end ?out_shape 2))
(= ?m (nth_from_end ?out_shape 1))
(= ?n (nth_from_end ?out_shape 0))
(!= ?m (MNum 0))
(!= ?n (MNum 0))
(!= ?k (MNum 1))
(!= ?batch (MNum 0))
(= ?a_batch_stride (nth_from_end ?a_stride 3))
(= ?a_m_stride (nth_from_end ?a_stride 2))
(= ?a_n_stride (nth_from_end ?a_stride 1))
(= ?a_k_stride (nth_from_end ?a_stride 0))
(= ?b_batch_stride (nth_from_end ?b_stride 3))
(= ?b_m_stride (nth_from_end ?b_stride 2))
(= ?b_n_stride (nth_from_end ?b_stride 1))
(= ?b_k_stride (nth_from_end ?b_stride 0))
(= ?k_stride (MIter))
(= ?a_k_stride (MIter))
(= ?a_n_stride (MNum 0))
(= ?a_m_stride (MMul (MIter) ?k))
(= ?b_k_stride (MIter))
(= ?b_m_stride (MNum 0))
(= ?b_n_stride (MMul (MIter) ?k))
(= ?a_batch_stride (MMul ?m ?a_m_stride))
(= ?b_batch_stride (MMul ?n ?b_n_stride))
(= ?dt (dtype ?a))
(= ?dt (dtype ?b))
(cublaslt_base_dtype ?dt)
)
(
(let ?sgemm (Op (cublaslt
?m ?n ?k
"N" "N"
"ROW" "COL" "ROW" "ROW"
?a_m_stride
?b_n_stride
?n
?n
?batch
?a_batch_stride
?b_batch_stride
(MMul ?m ?n)
(MMul ?m ?n)
?dt ?dt ?dt ?dt "default" "default" 1.0 0.0 "DEFAULT")
(ICons ?a (ICons ?b (INil)))))
(union ?sum ?sgemm)
(set (dtype ?sgemm) ?dt)
)
:ruleset matmul_backend
:name "cublaslt row-order batched row-major x column-major"
)
(rule
(
(= ?mul (Op (Mul ?mul_shape ?a_stride ?b_stride ?mul_out_stride) (ICons ?a (ICons ?b (INil)))))
(= ?sum (Op (Sum ?out_shape ?k ?sum_in_stride ?k_stride ?sum_out_stride) (ICons ?mul (INil))))
(= ?batch (nth_from_end ?out_shape 2))
(= ?m (nth_from_end ?out_shape 1))
(= ?n (nth_from_end ?out_shape 0))
(!= ?m (MNum 0))
(!= ?n (MNum 0))
(!= ?k (MNum 1))
(!= ?batch (MNum 0))
(= ?a_batch_stride (nth_from_end ?a_stride 3))
(= ?a_m_stride (nth_from_end ?a_stride 2))
(= ?a_n_stride (nth_from_end ?a_stride 1))
(= ?a_k_stride (nth_from_end ?a_stride 0))
(= ?b_batch_stride (nth_from_end ?b_stride 3))
(= ?b_m_stride (nth_from_end ?b_stride 2))
(= ?b_n_stride (nth_from_end ?b_stride 1))
(= ?b_k_stride (nth_from_end ?b_stride 0))
(= ?k_stride (MIter))
(= ?a_m_stride (MIter))
(= ?a_n_stride (MNum 0))
(= ?a_k_stride (MMul (MIter) ?m))
(= ?b_n_stride (MIter))
(= ?b_m_stride (MNum 0))
(= ?b_k_stride (MMul (MIter) ?n))
(= ?a_batch_stride (MMul ?k ?a_k_stride))
(= ?b_batch_stride (MMul ?k ?b_k_stride))
(= ?dt (dtype ?a))
(= ?dt (dtype ?b))
(cublaslt_base_dtype ?dt)
)
(
(let ?sgemm (Op (cublaslt
?m ?n ?k
"N" "N"
"COL" "ROW" "ROW" "ROW"
?a_k_stride
?b_k_stride
?n
?n
?batch
?a_batch_stride
?b_batch_stride
(MMul ?m ?n)
(MMul ?m ?n)
?dt ?dt ?dt ?dt "default" "default" 1.0 0.0 "DEFAULT")
(ICons ?a (ICons ?b (INil)))))
(union ?sum ?sgemm)
(set (dtype ?sgemm) ?dt)
)
:ruleset matmul_backend
:name "cublaslt row-order batched column-major x row-major"
)
(rule
(
(= ?mul (Op (Mul ?mul_shape ?a_stride ?b_stride ?mul_out_stride) (ICons ?a (ICons ?b (INil)))))
(= ?sum (Op (Sum ?out_shape ?k ?sum_in_stride ?k_stride ?sum_out_stride) (ICons ?mul (INil))))
(= ?batch (nth_from_end ?out_shape 2))
(= ?m (nth_from_end ?out_shape 1))
(= ?n (nth_from_end ?out_shape 0))
(!= ?m (MNum 0))
(!= ?n (MNum 0))
(!= ?k (MNum 1))
(!= ?batch (MNum 0))
(= ?a_batch_stride (nth_from_end ?a_stride 3))
(= ?a_m_stride (nth_from_end ?a_stride 2))
(= ?a_n_stride (nth_from_end ?a_stride 1))
(= ?a_k_stride (nth_from_end ?a_stride 0))
(= ?b_batch_stride (nth_from_end ?b_stride 3))
(= ?b_m_stride (nth_from_end ?b_stride 2))
(= ?b_n_stride (nth_from_end ?b_stride 1))
(= ?b_k_stride (nth_from_end ?b_stride 0))
(= ?k_stride (MIter))
(= ?a_m_stride (MIter))
(= ?a_n_stride (MNum 0))
(= ?a_k_stride (MMul (MIter) ?m))
(= ?b_k_stride (MIter))
(= ?b_m_stride (MNum 0))
(= ?b_n_stride (MMul (MIter) ?k))
(= ?a_batch_stride (MMul ?k ?a_k_stride))
(= ?b_batch_stride (MMul ?n ?b_n_stride))
(= ?dt (dtype ?a))
(= ?dt (dtype ?b))
(cublaslt_base_dtype ?dt)
)
(
(let ?sgemm (Op (cublaslt
?m ?n ?k
"N" "N"
"COL" "COL" "ROW" "ROW"
?a_k_stride
?b_n_stride
?n
?n
?batch
?a_batch_stride
?b_batch_stride
(MMul ?m ?n)
(MMul ?m ?n)
?dt ?dt ?dt ?dt "default" "default" 1.0 0.0 "DEFAULT")
(ICons ?a (ICons ?b (INil)))))
(union ?sum ?sgemm)
(set (dtype ?sgemm) ?dt)
)
:ruleset matmul_backend
:name "cublaslt row-order batched column-major x column-major"
)
; Mixed output dtype rewrites for cuBLASLt.
;
; The first mixed mode we need for low-precision matmuls is:
;
; D[f32] = A[fp16/bf16] * B[fp16/bf16]
;
; Luminal graphs express this today as a Cast(F32) around a low-precision
; matmul. cuBLASLt can write the f32 output directly, so expose that candidate
; before beta fusion tries to consume an f32 C input.
(rule
(
(= ?matmul (Op (cublaslt
?m ?n ?k
?a_layout ?b_layout
?a_order ?b_order ?c_order ?d_order
?lda ?ldb ?ldc ?ldd
?batch
?stride_a ?stride_b ?stride_c ?stride_d
(F16) (F16) (F16) (F16)
?compute_type ?scale_dtype
?alpha ?beta ?epilogue)
?inputs))
(= ?cast (Op (Cast ?size (F32)) (ICons ?matmul (INil))))
)
(
(let ?fused (Op (cublaslt
?m ?n ?k
?a_layout ?b_layout ?a_order ?b_order ?c_order ?d_order
?lda ?ldb ?ldc ?ldd
?batch
?stride_a ?stride_b ?stride_c ?stride_d
(F16) (F16) (F32) (F32)
?compute_type ?scale_dtype
?alpha ?beta ?epilogue)
?inputs))
(union ?cast ?fused)
(set (dtype ?fused) (F32))
)
:ruleset matmul_backend
:name "cublaslt f16 matmul cast f32 output"
)
(rule
(
(= ?matmul (Op (cublaslt
?m ?n ?k
?a_layout ?b_layout
?a_order ?b_order ?c_order ?d_order
?lda ?ldb ?ldc ?ldd
?batch
?stride_a ?stride_b ?stride_c ?stride_d
(Bf16) (Bf16) (Bf16) (Bf16)
?compute_type ?scale_dtype
?alpha ?beta ?epilogue)
?inputs))
(= ?cast (Op (Cast ?size (F32)) (ICons ?matmul (INil))))
)
(
(let ?fused (Op (cublaslt
?m ?n ?k
?a_layout ?b_layout ?a_order ?b_order ?c_order ?d_order
?lda ?ldb ?ldc ?ldd
?batch
?stride_a ?stride_b ?stride_c ?stride_d
(Bf16) (Bf16) (F32) (F32)
?compute_type ?scale_dtype
?alpha ?beta ?epilogue)
?inputs))
(union ?cast ?fused)
(set (dtype ?fused) (F32))
)
:ruleset matmul_backend
:name "cublaslt bf16 matmul cast f32 output"
)
; Scalar alpha/beta rewrites for cuBLASLt. These rules target scalar constants
; expanded across the matmul/add shape, i.e. zero strides on every logical axis.
(rule
(
(= ?matmul (Op (cublaslt
?m ?n ?k
?a_layout ?b_layout
?a_order ?b_order ?c_order ?d_order
?lda ?ldb ?ldc ?ldd
?batch
?stride_a ?stride_b ?stride_c ?stride_d
?a_dtype ?b_dtype ?c_dtype ?d_dtype
?compute_type ?scale_dtype
1.0 0.0 "DEFAULT")
(ICons ?a (ICons ?b ?matmul_tail))))
(= ?scale (Op (Constant ?alpha) (INil)))
; alpha=1.0 hash-conses ?fused == ?matmul; the union merges Mul into ?matmul's eclass and saturate diverges.
(!= ?alpha 1.0)
(= ?scaled (Op (Mul ?shape
?matmul_strides
(ECons (MNum 0) (ECons (MNum 0) (ENil)))
?scaled_out_strides)
(ICons ?matmul (ICons ?scale (INil)))))
(= ?matmul_strides ?scaled_out_strides)
)
(
(let ?fused (Op (cublaslt
?m ?n ?k
?a_layout ?b_layout
?a_order ?b_order ?c_order ?d_order
?lda ?ldb ?ldc ?ldd
?batch
?stride_a ?stride_b ?stride_c ?stride_d
?a_dtype ?b_dtype ?c_dtype ?d_dtype
?compute_type ?scale_dtype
?alpha 0.0 "DEFAULT")
(ICons ?a (ICons ?b ?matmul_tail))))
(union ?scaled ?fused)
(set (dtype ?fused) ?d_dtype)
)
:ruleset matmul_backend
:name "cublaslt 2d alpha scale"
)
(rule
(
(= ?matmul (Op (cublaslt
?m ?n ?k
?a_layout ?b_layout
?a_order ?b_order ?c_order ?d_order
?lda ?ldb ?ldc ?ldd
?batch
?stride_a ?stride_b ?stride_c ?stride_d
?a_dtype ?b_dtype ?c_dtype ?d_dtype
?compute_type ?scale_dtype
1.0 0.0 "DEFAULT")
(ICons ?a (ICons ?b ?matmul_tail))))
(= ?scale (Op (Constant ?alpha) (INil)))
; See 2d alpha scale: alpha=1.0 makes (saturate ...) diverge.
(!= ?alpha 1.0)
(= ?scaled (Op (Mul ?shape
?matmul_strides
(ECons (MNum 0) (ECons (MNum 0) (ECons (MNum 0) (ENil))))
?scaled_out_strides)
(ICons ?matmul (ICons ?scale (INil)))))
(= ?matmul_strides ?scaled_out_strides)
)
(
(let ?fused (Op (cublaslt
?m ?n ?k
?a_layout ?b_layout
?a_order ?b_order ?c_order ?d_order
?lda ?ldb ?ldc ?ldd
?batch
?stride_a ?stride_b ?stride_c ?stride_d
?a_dtype ?b_dtype ?c_dtype ?d_dtype
?compute_type ?scale_dtype
?alpha 0.0 "DEFAULT")
(ICons ?a (ICons ?b ?matmul_tail))))
(union ?scaled ?fused)
(set (dtype ?fused) ?d_dtype)
)
:ruleset matmul_backend
:name "cublaslt batched alpha scale"
)
(rule
(
(= ?matmul (Op (cublaslt
?m ?n ?k
?a_layout ?b_layout
?a_order ?b_order ?matmul_c_order "ROW"
?lda ?ldb ?matmul_ldc ?ldd
(MNum 1)
?stride_a ?stride_b ?matmul_stride_c ?stride_d
?a_dtype ?b_dtype ?c_dtype ?d_dtype
?compute_type ?scale_dtype
?alpha 0.0 ?epilogue)
(ICons ?a (ICons ?b ?matmul_tail))))
(= ?beta_node (Op (Constant ?beta) (INil)))
(= ?scaled_c (Op (Mul
(ECons ?m (ECons ?n (ENil)))
?c_strides
(ECons (MNum 0) (ECons (MNum 0) (ENil)))
?scaled_c_out_strides)
(ICons ?c (ICons ?beta_node (INil)))))
(= ?add (Op (Add
(ECons ?m (ECons ?n (ENil)))
?matmul_add_strides
?scaled_c_add_strides
?add_out_strides)
(ICons ?matmul (ICons ?scaled_c (INil)))))
(= ?matmul_add_strides (ECons ?d_row_stride (ECons ?d_col_stride (ENil))))
(= ?c_strides (ECons ?c_row_stride (ECons ?c_col_stride (ENil))))
(= ?add_out_strides (ECons ?d_row_stride (ECons ?d_col_stride (ENil))))
(= ?scaled_c_add_strides ?scaled_c_out_strides)
(= ?c_col_stride (MIter))
(!= ?c_row_stride (MNum 0))
(= ?matmul_add_strides ?add_out_strides)
(= ?c_dtype (dtype ?c))
)
(
(let ?fused (Op (cublaslt
?m ?n ?k
?a_layout ?b_layout
?a_order ?b_order "ROW" "ROW"
?lda ?ldb ?c_row_stride ?ldd
(MNum 1)
?stride_a ?stride_b (MNum 0) ?stride_d
?a_dtype ?b_dtype ?c_dtype ?d_dtype
?compute_type ?scale_dtype
?alpha ?beta ?epilogue)
(ICons ?a (ICons ?b (ICons ?c ?matmul_tail)))))
(union ?add ?fused)
(set (dtype ?fused) ?d_dtype)
)
:ruleset matmul_backend
:name "cublaslt row-order 2d scaled c beta"
)
(rule
(
(= ?matmul (Op (cublaslt
?m ?n ?k
?a_layout ?b_layout
?a_order ?b_order ?matmul_c_order "ROW"
?lda ?ldb ?matmul_ldc ?ldd
(MNum 1)
?stride_a ?stride_b ?matmul_stride_c ?stride_d
?a_dtype ?b_dtype ?c_dtype ?d_dtype
?compute_type ?scale_dtype
?alpha 0.0 ?epilogue)
(ICons ?a (ICons ?b ?matmul_tail))))
(= ?beta_node (Op (Constant ?beta) (INil)))
(= ?scaled_c (Op (Mul
(ECons ?m (ECons ?n (ENil)))
?c_strides
(ECons (MNum 0) (ECons (MNum 0) (ENil)))
?scaled_c_out_strides)
(ICons ?c (ICons ?beta_node (INil)))))
(= ?add (Op (Add
(ECons ?m (ECons ?n (ENil)))
?scaled_c_add_strides
?matmul_add_strides
?add_out_strides)
(ICons ?scaled_c (ICons ?matmul (INil)))))
(= ?matmul_add_strides (ECons ?d_row_stride (ECons ?d_col_stride (ENil))))
(= ?c_strides (ECons ?c_row_stride (ECons ?c_col_stride (ENil))))
(= ?add_out_strides (ECons ?d_row_stride (ECons ?d_col_stride (ENil))))
(= ?scaled_c_add_strides ?scaled_c_out_strides)
(= ?c_col_stride (MIter))
(!= ?c_row_stride (MNum 0))
(= ?matmul_add_strides ?add_out_strides)
(= ?c_dtype (dtype ?c))
)
(
(let ?fused (Op (cublaslt
?m ?n ?k
?a_layout ?b_layout
?a_order ?b_order "ROW" "ROW"
?lda ?ldb ?c_row_stride ?ldd
(MNum 1)
?stride_a ?stride_b (MNum 0) ?stride_d
?a_dtype ?b_dtype ?c_dtype ?d_dtype
?compute_type ?scale_dtype
?alpha ?beta ?epilogue)
(ICons ?a (ICons ?b (ICons ?c ?matmul_tail)))))
(union ?add ?fused)
(set (dtype ?fused) ?d_dtype)
)
:ruleset matmul_backend
:name "cublaslt row-order 2d scaled c plus matmul beta"
)
(rule
(
(= ?matmul (Op (cublaslt
?m ?n ?k
?a_layout ?b_layout
?a_order ?b_order ?matmul_c_order "ROW"
?lda ?ldb ?matmul_ldc ?ldd
?batch
?stride_a ?stride_b ?matmul_stride_c ?stride_d
?a_dtype ?b_dtype ?c_dtype ?d_dtype
?compute_type ?scale_dtype
?alpha 0.0 ?epilogue)
(ICons ?a (ICons ?b ?matmul_tail))))
(= ?beta_node (Op (Constant ?beta) (INil)))
(= ?scaled_c (Op (Mul
(ECons ?batch (ECons ?m (ECons ?n (ENil))))
?c_strides
(ECons (MNum 0) (ECons (MNum 0) (ECons (MNum 0) (ENil))))
?scaled_c_out_strides)
(ICons ?c (ICons ?beta_node (INil)))))
(= ?add (Op (Add
(ECons ?batch (ECons ?m (ECons ?n (ENil))))
?matmul_add_strides
?scaled_c_add_strides
?add_out_strides)
(ICons ?matmul (ICons ?scaled_c (INil)))))
(= ?matmul_add_strides (ECons ?d_batch_stride (ECons ?d_row_stride (ECons ?d_col_stride (ENil)))))
(= ?c_strides (ECons ?c_batch_stride (ECons ?c_row_stride (ECons ?c_col_stride (ENil)))))
(= ?add_out_strides (ECons ?d_batch_stride (ECons ?d_row_stride (ECons ?d_col_stride (ENil)))))
(= ?scaled_c_add_strides ?scaled_c_out_strides)
(= ?c_col_stride (MIter))
(!= ?c_row_stride (MNum 0))
(= ?matmul_add_strides ?add_out_strides)
(= ?c_dtype (dtype ?c))
)
(
(let ?fused (Op (cublaslt
?m ?n ?k
?a_layout ?b_layout
?a_order ?b_order "ROW" "ROW"
?lda ?ldb ?c_row_stride ?ldd
?batch
?stride_a ?stride_b ?c_batch_stride ?stride_d
?a_dtype ?b_dtype ?c_dtype ?d_dtype
?compute_type ?scale_dtype
?alpha ?beta ?epilogue)
(ICons ?a (ICons ?b (ICons ?c ?matmul_tail)))))
(union ?add ?fused)
(set (dtype ?fused) ?d_dtype)
)
:ruleset matmul_backend
:name "cublaslt row-order batched scaled c beta"
)
(rule
(
(= ?matmul (Op (cublaslt
?m ?n ?k
?a_layout ?b_layout
?a_order ?b_order ?matmul_c_order "ROW"
?lda ?ldb ?matmul_ldc ?ldd
?batch
?stride_a ?stride_b ?matmul_stride_c ?stride_d
?a_dtype ?b_dtype ?c_dtype ?d_dtype
?compute_type ?scale_dtype
?alpha 0.0 ?epilogue)
(ICons ?a (ICons ?b ?matmul_tail))))
(= ?beta_node (Op (Constant ?beta) (INil)))
(= ?scaled_c (Op (Mul
(ECons ?batch (ECons ?m (ECons ?n (ENil))))
?c_strides
(ECons (MNum 0) (ECons (MNum 0) (ECons (MNum 0) (ENil))))
?scaled_c_out_strides)
(ICons ?c (ICons ?beta_node (INil)))))
(= ?add (Op (Add
(ECons ?batch (ECons ?m (ECons ?n (ENil))))
?scaled_c_add_strides
?matmul_add_strides
?add_out_strides)
(ICons ?scaled_c (ICons ?matmul (INil)))))
(= ?matmul_add_strides (ECons ?d_batch_stride (ECons ?d_row_stride (ECons ?d_col_stride (ENil)))))
(= ?c_strides (ECons ?c_batch_stride (ECons ?c_row_stride (ECons ?c_col_stride (ENil)))))
(= ?add_out_strides (ECons ?d_batch_stride (ECons ?d_row_stride (ECons ?d_col_stride (ENil)))))
(= ?scaled_c_add_strides ?scaled_c_out_strides)
(= ?c_col_stride (MIter))
(!= ?c_row_stride (MNum 0))
(= ?matmul_add_strides ?add_out_strides)
(= ?c_dtype (dtype ?c))
)
(
(let ?fused (Op (cublaslt
?m ?n ?k
?a_layout ?b_layout
?a_order ?b_order "ROW" "ROW"
?lda ?ldb ?c_row_stride ?ldd
?batch
?stride_a ?stride_b ?c_batch_stride ?stride_d
?a_dtype ?b_dtype ?c_dtype ?d_dtype
?compute_type ?scale_dtype
?alpha ?beta ?epilogue)
(ICons ?a (ICons ?b (ICons ?c ?matmul_tail)))))
(union ?add ?fused)
(set (dtype ?fused) ?d_dtype)
)
:ruleset matmul_backend
:name "cublaslt row-order batched scaled c plus matmul beta"
)
; Fuse a row-major Add on top of an existing cuBLASLt matmul into
; D = alpha * A * B + beta * C.
;
; The existing matmul rewrites view Luminal's row-major output [m,n] as a
; column-major cuBLASLt matrix [n,m]. A row-major C input with logical strides
; [row_stride, 1] therefore maps to ldc=row_stride. This lets a C slice from a
; wider parent tensor use a larger ldc while D keeps the matmul output layout.
; cuBLASLt requires out-of-place C and D to have the same matrix order, so these
; beta rules only fuse C layouts that map to the current COL-ordered D layout.
(rule
(
(= ?matmul (Op (cublaslt
?m ?n ?k
?a_layout ?b_layout
?a_order ?b_order ?matmul_c_order "COL"
?lda ?ldb ?matmul_ldc ?ldd
(MNum 1)
?stride_a ?stride_b ?matmul_stride_c ?stride_d
?a_dtype ?b_dtype ?c_dtype ?d_dtype
?compute_type ?scale_dtype
?alpha 0.0 ?epilogue)
(ICons ?a (ICons ?b ?matmul_tail))))
(!= ?epilogue "RELU")
(!= ?epilogue "RELU_BIAS")
(!= ?epilogue "GELU")
(!= ?epilogue "GELU_BIAS")
(= ?add (Op (Add
(ECons ?n (ECons ?m (ENil)))
?matmul_add_strides
?c_add_strides
?add_out_strides)
(ICons ?matmul (ICons ?c (INil)))))
(= ?matmul_add_strides (ECons ?d_row_stride (ECons ?d_col_stride (ENil))))
(= ?c_add_strides (ECons ?c_row_stride (ECons ?c_col_stride (ENil))))
(= ?add_out_strides (ECons ?d_row_stride (ECons ?d_col_stride (ENil))))
(= ?c_col_stride (MIter))
(!= ?c_row_stride (MNum 0))
(= ?matmul_add_strides ?add_out_strides)
(= ?c_dtype (dtype ?c))
)
(
(let ?fused (Op (cublaslt
?m ?n ?k
?a_layout ?b_layout
?a_order ?b_order "COL" "COL"
?lda ?ldb ?c_row_stride ?ldd
(MNum 1)
?stride_a ?stride_b (MNum 0) ?stride_d
?a_dtype ?b_dtype ?c_dtype ?d_dtype
?compute_type ?scale_dtype
?alpha 1.0 ?epilogue)
(ICons ?a (ICons ?b (ICons ?c ?matmul_tail)))))
(union ?add ?fused)
(set (dtype ?fused) ?d_dtype)
)
:ruleset matmul_backend
:name "cublaslt 2d matmul plus c beta"
)
(rule
(
(= ?matmul (Op (cublaslt
?m ?n ?k
?a_layout ?b_layout
?a_order ?b_order ?matmul_c_order "COL"
?lda ?ldb ?matmul_ldc ?ldd
(MNum 1)
?stride_a ?stride_b ?matmul_stride_c ?stride_d
?a_dtype ?b_dtype ?c_dtype ?d_dtype
?compute_type ?scale_dtype
?alpha 0.0 ?epilogue)
(ICons ?a (ICons ?b ?matmul_tail))))
(!= ?epilogue "RELU")
(!= ?epilogue "RELU_BIAS")
(!= ?epilogue "GELU")
(!= ?epilogue "GELU_BIAS")
(= ?add (Op (Add
(ECons ?n (ECons ?m (ENil)))
?c_add_strides
?matmul_add_strides
?add_out_strides)
(ICons ?c (ICons ?matmul (INil)))))
(= ?matmul_add_strides (ECons ?d_row_stride (ECons ?d_col_stride (ENil))))
(= ?c_add_strides (ECons ?c_row_stride (ECons ?c_col_stride (ENil))))
(= ?add_out_strides (ECons ?d_row_stride (ECons ?d_col_stride (ENil))))
(= ?c_col_stride (MIter))
(!= ?c_row_stride (MNum 0))
(= ?matmul_add_strides ?add_out_strides)
(= ?c_dtype (dtype ?c))
)
(
(let ?fused (Op (cublaslt
?m ?n ?k
?a_layout ?b_layout
?a_order ?b_order "COL" "COL"
?lda ?ldb ?c_row_stride ?ldd
(MNum 1)
?stride_a ?stride_b (MNum 0) ?stride_d
?a_dtype ?b_dtype ?c_dtype ?d_dtype
?compute_type ?scale_dtype
?alpha 1.0 ?epilogue)
(ICons ?a (ICons ?b (ICons ?c ?matmul_tail)))))
(union ?add ?fused)
(set (dtype ?fused) ?d_dtype)
)
:ruleset matmul_backend
:name "cublaslt 2d c plus matmul beta"
)
(rule
(
(= ?matmul (Op (cublaslt
?m ?n ?k
?a_layout ?b_layout
?a_order ?b_order ?matmul_c_order "COL"
?lda ?ldb ?matmul_ldc ?ldd
?batch
?stride_a ?stride_b ?matmul_stride_c ?stride_d
?a_dtype ?b_dtype ?c_dtype ?d_dtype
?compute_type ?scale_dtype
?alpha 0.0 ?epilogue)
(ICons ?a (ICons ?b ?matmul_tail))))
(!= ?epilogue "RELU")
(!= ?epilogue "RELU_BIAS")
(!= ?epilogue "GELU")
(!= ?epilogue "GELU_BIAS")
(= ?add (Op (Add
(ECons ?batch (ECons ?n (ECons ?m (ENil))))
?matmul_add_strides
?c_add_strides
?add_out_strides)
(ICons ?matmul (ICons ?c (INil)))))
(= ?matmul_add_strides (ECons ?d_batch_stride (ECons ?d_row_stride (ECons ?d_col_stride (ENil)))))
(= ?c_add_strides (ECons ?c_batch_stride (ECons ?c_row_stride (ECons ?c_col_stride (ENil)))))
(= ?add_out_strides (ECons ?d_batch_stride (ECons ?d_row_stride (ECons ?d_col_stride (ENil)))))
(= ?c_col_stride (MIter))
(!= ?c_row_stride (MNum 0))
(= ?matmul_add_strides ?add_out_strides)
(= ?c_dtype (dtype ?c))
)
(
(let ?fused (Op (cublaslt
?m ?n ?k
?a_layout ?b_layout
?a_order ?b_order "COL" "COL"
?lda ?ldb ?c_row_stride ?ldd
?batch
?stride_a ?stride_b ?c_batch_stride ?stride_d
?a_dtype ?b_dtype ?c_dtype ?d_dtype
?compute_type ?scale_dtype
?alpha 1.0 ?epilogue)
(ICons ?a (ICons ?b (ICons ?c ?matmul_tail)))))
(union ?add ?fused)
(set (dtype ?fused) ?d_dtype)
)
:ruleset matmul_backend
:name "cublaslt batched matmul plus c beta"
)
(rule
(
(= ?matmul (Op (cublaslt
?m ?n ?k
?a_layout ?b_layout
?a_order ?b_order ?matmul_c_order "COL"
?lda ?ldb ?matmul_ldc ?ldd
?batch
?stride_a ?stride_b ?matmul_stride_c ?stride_d
?a_dtype ?b_dtype ?c_dtype ?d_dtype
?compute_type ?scale_dtype
?alpha 0.0 ?epilogue)
(ICons ?a (ICons ?b ?matmul_tail))))
(!= ?epilogue "RELU")
(!= ?epilogue "RELU_BIAS")
(!= ?epilogue "GELU")
(!= ?epilogue "GELU_BIAS")
(= ?add (Op (Add
(ECons ?batch (ECons ?n (ECons ?m (ENil))))
?c_add_strides
?matmul_add_strides
?add_out_strides)
(ICons ?c (ICons ?matmul (INil)))))
(= ?matmul_add_strides (ECons ?d_batch_stride (ECons ?d_row_stride (ECons ?d_col_stride (ENil)))))
(= ?c_add_strides (ECons ?c_batch_stride (ECons ?c_row_stride (ECons ?c_col_stride (ENil)))))
(= ?add_out_strides (ECons ?d_batch_stride (ECons ?d_row_stride (ECons ?d_col_stride (ENil)))))
(= ?c_col_stride (MIter))
(!= ?c_row_stride (MNum 0))
(= ?matmul_add_strides ?add_out_strides)
(= ?c_dtype (dtype ?c))
)
(
(let ?fused (Op (cublaslt
?m ?n ?k
?a_layout ?b_layout
?a_order ?b_order "COL" "COL"
?lda ?ldb ?c_row_stride ?ldd
?batch
?stride_a ?stride_b ?c_batch_stride ?stride_d
?a_dtype ?b_dtype ?c_dtype ?d_dtype
?compute_type ?scale_dtype
?alpha 1.0 ?epilogue)
(ICons ?a (ICons ?b (ICons ?c ?matmul_tail)))))
(union ?add ?fused)
(set (dtype ?fused) ?d_dtype)
)
:ruleset matmul_backend
:name "cublaslt batched c plus matmul beta"
)
; ROW-ordered D beta fusions. These pair with cublaslt_row_order_rewrite.egg,
; where the cuBLASLt problem dimensions match Luminal's logical output [m,n].
; A row-major C input with logical strides [row_stride, 1] maps directly to a
; ROW-ordered cuBLASLt C[m,n] descriptor with ldc=row_stride.
(rule
(
(= ?matmul (Op (cublaslt
?m ?n ?k
?a_layout ?b_layout
?a_order ?b_order ?matmul_c_order "ROW"
?lda ?ldb ?matmul_ldc ?ldd
(MNum 1)
?stride_a ?stride_b ?matmul_stride_c ?stride_d
?a_dtype ?b_dtype ?c_dtype ?d_dtype
?compute_type ?scale_dtype
?alpha 0.0 ?epilogue)
(ICons ?a (ICons ?b ?matmul_tail))))
(!= ?epilogue "RELU")
(!= ?epilogue "RELU_BIAS")
(!= ?epilogue "GELU")
(!= ?epilogue "GELU_BIAS")
(= ?add (Op (Add
(ECons ?m (ECons ?n (ENil)))
?matmul_add_strides
?c_add_strides
?add_out_strides)
(ICons ?matmul (ICons ?c (INil)))))
(= ?matmul_add_strides (ECons ?d_row_stride (ECons ?d_col_stride (ENil))))
(= ?c_add_strides (ECons ?c_row_stride (ECons ?c_col_stride (ENil))))
(= ?add_out_strides (ECons ?d_row_stride (ECons ?d_col_stride (ENil))))
(= ?c_col_stride (MIter))
(!= ?c_row_stride (MNum 0))
(= ?matmul_add_strides ?add_out_strides)
(= ?c_dtype (dtype ?c))
)
(
(let ?fused (Op (cublaslt
?m ?n ?k
?a_layout ?b_layout
?a_order ?b_order "ROW" "ROW"
?lda ?ldb ?c_row_stride ?ldd
(MNum 1)
?stride_a ?stride_b (MNum 0) ?stride_d
?a_dtype ?b_dtype ?c_dtype ?d_dtype
?compute_type ?scale_dtype
?alpha 1.0 ?epilogue)
(ICons ?a (ICons ?b (ICons ?c ?matmul_tail)))))
(union ?add ?fused)
(set (dtype ?fused) ?d_dtype)
)
:ruleset matmul_backend
:name "cublaslt row-order 2d matmul plus c beta"
)
(rule
(
(= ?matmul (Op (cublaslt
?m ?n ?k
?a_layout ?b_layout
?a_order ?b_order ?matmul_c_order "ROW"
?lda ?ldb ?matmul_ldc ?ldd
(MNum 1)
?stride_a ?stride_b ?matmul_stride_c ?stride_d
?a_dtype ?b_dtype ?c_dtype ?d_dtype
?compute_type ?scale_dtype
?alpha 0.0 ?epilogue)
(ICons ?a (ICons ?b ?matmul_tail))))
(!= ?epilogue "RELU")
(!= ?epilogue "RELU_BIAS")
(!= ?epilogue "GELU")
(!= ?epilogue "GELU_BIAS")
(= ?add (Op (Add
(ECons ?m (ECons ?n (ENil)))
?c_add_strides
?matmul_add_strides
?add_out_strides)
(ICons ?c (ICons ?matmul (INil)))))
(= ?matmul_add_strides (ECons ?d_row_stride (ECons ?d_col_stride (ENil))))
(= ?c_add_strides (ECons ?c_row_stride (ECons ?c_col_stride (ENil))))
(= ?add_out_strides (ECons ?d_row_stride (ECons ?d_col_stride (ENil))))
(= ?c_col_stride (MIter))
(!= ?c_row_stride (MNum 0))
(= ?matmul_add_strides ?add_out_strides)
(= ?c_dtype (dtype ?c))
)
(
(let ?fused (Op (cublaslt
?m ?n ?k
?a_layout ?b_layout
?a_order ?b_order "ROW" "ROW"
?lda ?ldb ?c_row_stride ?ldd
(MNum 1)
?stride_a ?stride_b (MNum 0) ?stride_d
?a_dtype ?b_dtype ?c_dtype ?d_dtype
?compute_type ?scale_dtype
?alpha 1.0 ?epilogue)
(ICons ?a (ICons ?b (ICons ?c ?matmul_tail)))))
(union ?add ?fused)
(set (dtype ?fused) ?d_dtype)
)
:ruleset matmul_backend
:name "cublaslt row-order 2d c plus matmul beta"
)
(rule
(
(= ?matmul (Op (cublaslt
?m ?n ?k
?a_layout ?b_layout
?a_order ?b_order ?matmul_c_order "ROW"
?lda ?ldb ?matmul_ldc ?ldd
?batch
?stride_a ?stride_b ?matmul_stride_c ?stride_d
?a_dtype ?b_dtype ?c_dtype ?d_dtype
?compute_type ?scale_dtype
?alpha 0.0 ?epilogue)
(ICons ?a (ICons ?b ?matmul_tail))))
(!= ?epilogue "RELU")
(!= ?epilogue "RELU_BIAS")
(!= ?epilogue "GELU")
(!= ?epilogue "GELU_BIAS")
(= ?add (Op (Add
(ECons ?batch (ECons ?m (ECons ?n (ENil))))
?matmul_add_strides
?c_add_strides
?add_out_strides)
(ICons ?matmul (ICons ?c (INil)))))
(= ?matmul_add_strides (ECons ?d_batch_stride (ECons ?d_row_stride (ECons ?d_col_stride (ENil)))))
(= ?c_add_strides (ECons ?c_batch_stride (ECons ?c_row_stride (ECons ?c_col_stride (ENil)))))
(= ?add_out_strides (ECons ?d_batch_stride (ECons ?d_row_stride (ECons ?d_col_stride (ENil)))))
(= ?c_col_stride (MIter))
(!= ?c_row_stride (MNum 0))
(= ?matmul_add_strides ?add_out_strides)
(= ?c_dtype (dtype ?c))
)
(
(let ?fused (Op (cublaslt
?m ?n ?k
?a_layout ?b_layout
?a_order ?b_order "ROW" "ROW"
?lda ?ldb ?c_row_stride ?ldd
?batch
?stride_a ?stride_b ?c_batch_stride ?stride_d
?a_dtype ?b_dtype ?c_dtype ?d_dtype
?compute_type ?scale_dtype
?alpha 1.0 ?epilogue)
(ICons ?a (ICons ?b (ICons ?c ?matmul_tail)))))
(union ?add ?fused)
(set (dtype ?fused) ?d_dtype)
)
:ruleset matmul_backend
:name "cublaslt row-order batched matmul plus c beta"
)
(rule
(
(= ?matmul (Op (cublaslt
?m ?n ?k
?a_layout ?b_layout
?a_order ?b_order ?matmul_c_order "ROW"
?lda ?ldb ?matmul_ldc ?ldd
?batch
?stride_a ?stride_b ?matmul_stride_c ?stride_d
?a_dtype ?b_dtype ?c_dtype ?d_dtype
?compute_type ?scale_dtype
?alpha 0.0 ?epilogue)
(ICons ?a (ICons ?b ?matmul_tail))))
(!= ?epilogue "RELU")
(!= ?epilogue "RELU_BIAS")
(!= ?epilogue "GELU")
(!= ?epilogue "GELU_BIAS")
(= ?add (Op (Add
(ECons ?batch (ECons ?m (ECons ?n (ENil))))
?c_add_strides
?matmul_add_strides
?add_out_strides)
(ICons ?c (ICons ?matmul (INil)))))
(= ?matmul_add_strides (ECons ?d_batch_stride (ECons ?d_row_stride (ECons ?d_col_stride (ENil)))))
(= ?c_add_strides (ECons ?c_batch_stride (ECons ?c_row_stride (ECons ?c_col_stride (ENil)))))
(= ?add_out_strides (ECons ?d_batch_stride (ECons ?d_row_stride (ECons ?d_col_stride (ENil)))))
(= ?c_col_stride (MIter))
(!= ?c_row_stride (MNum 0))
(= ?matmul_add_strides ?add_out_strides)
(= ?c_dtype (dtype ?c))
)
(
(let ?fused (Op (cublaslt
?m ?n ?k
?a_layout ?b_layout
?a_order ?b_order "ROW" "ROW"
?lda ?ldb ?c_row_stride ?ldd
?batch
?stride_a ?stride_b ?c_batch_stride ?stride_d
?a_dtype ?b_dtype ?c_dtype ?d_dtype
?compute_type ?scale_dtype
?alpha 1.0 ?epilogue)
(ICons ?a (ICons ?b (ICons ?c ?matmul_tail)))))
(union ?add ?fused)
(set (dtype ?fused) ?d_dtype)
)
:ruleset matmul_backend
:name "cublaslt row-order batched c plus matmul beta"
)
; cuBLASLt epilogue rewrites.
;
; ReLU in the frontend lowers through maximum_f32(0.0):
;
; (matmul < 0) * 0 + cast(cast((-cast(matmul < 0) + 1) as bool) as f32) * matmul
;
; These rules fuse that expression back into CUBLASLT_EPILOGUE_RELU.
(rule
(
(= ?matmul (Op (cublaslt
?m ?n ?k
?a_layout ?b_layout
?a_order ?b_order ?c_order ?d_order
?lda ?ldb ?ldc ?ldd
?batch
?stride_a ?stride_b ?stride_c ?stride_d
?a_dtype ?b_dtype ?c_dtype (F32)
?compute_type ?scale_dtype
?alpha 0.0 "DEFAULT")
(ICons ?a (ICons ?b ?matmul_tail))))
(= ?zero (Op (Constant 0.0) (INil)))
(= ?neg_one (Op (Constant -1.0) (INil)))
(= ?one (Op (Constant 1.0) (INil)))
(= ?lt (Op (LessThan
?shape
?matmul_strides
(ECons (MNum 0) (ECons (MNum 0) (ENil)))
?mask_strides)
(ICons ?matmul (ICons ?zero (INil)))))
(= ?lt_f32 (Op (Cast ?size (F32)) (ICons ?lt (INil))))
(= ?zeroed (Op (Mul
?shape
?mask_strides
(ECons (MNum 0) (ECons (MNum 0) (ENil)))
?zeroed_strides)
(ICons ?lt_f32 (ICons ?zero (INil)))))
(= ?neg_mask (Op (Mul
?shape
?mask_strides
(ECons (MNum 0) (ECons (MNum 0) (ENil)))
?neg_mask_strides)
(ICons ?lt_f32 (ICons ?neg_one (INil)))))
(= ?not_mask_f32 (Op (Add
?shape
?neg_mask_strides
(ECons (MNum 0) (ECons (MNum 0) (ENil)))
?not_mask_f32_strides)
(ICons ?neg_mask (ICons ?one (INil)))))
(= ?not_mask_bool (Op (Cast ?size (Bool)) (ICons ?not_mask_f32 (INil))))
(= ?not_mask (Op (Cast ?size (F32)) (ICons ?not_mask_bool (INil))))
(= ?positive (Op (Mul
?shape
?not_mask_f32_strides
?matmul_strides
?positive_strides)
(ICons ?not_mask (ICons ?matmul (INil)))))
(= ?relu (Op (Add
?shape
?zeroed_strides
?positive_strides
?relu_strides)
(ICons ?zeroed (ICons ?positive (INil)))))
)
(
(let ?fused (Op (cublaslt
?m ?n ?k
?a_layout ?b_layout
?a_order ?b_order ?c_order ?d_order
?lda ?ldb ?ldc ?ldd
?batch
?stride_a ?stride_b ?stride_c ?stride_d
?a_dtype ?b_dtype ?c_dtype (F32)
?compute_type ?scale_dtype
?alpha 0.0 "RELU")
(ICons ?a (ICons ?b ?matmul_tail))))
(union ?relu ?fused)
(set (dtype ?fused) (F32))
)
:ruleset matmul_backend
:name "cublaslt 2d relu epilogue"
)
(rule
(
(= ?matmul (Op (cublaslt
?m ?n ?k
?a_layout ?b_layout
?a_order ?b_order ?c_order ?d_order
?lda ?ldb ?ldc ?ldd
?batch
?stride_a ?stride_b ?stride_c ?stride_d
?a_dtype ?b_dtype ?c_dtype (F32)
?compute_type ?scale_dtype
?alpha 0.0 "DEFAULT")
(ICons ?a (ICons ?b ?matmul_tail))))
(= ?zero (Op (Constant 0.0) (INil)))
(= ?neg_one (Op (Constant -1.0) (INil)))
(= ?one (Op (Constant 1.0) (INil)))
(= ?lt (Op (LessThan
?shape
?matmul_strides
(ECons (MNum 0) (ECons (MNum 0) (ECons (MNum 0) (ENil))))
?mask_strides)
(ICons ?matmul (ICons ?zero (INil)))))
(= ?lt_f32 (Op (Cast ?size (F32)) (ICons ?lt (INil))))
(= ?zeroed (Op (Mul
?shape
?mask_strides
(ECons (MNum 0) (ECons (MNum 0) (ECons (MNum 0) (ENil))))
?zeroed_strides)
(ICons ?lt_f32 (ICons ?zero (INil)))))
(= ?neg_mask (Op (Mul
?shape
?mask_strides
(ECons (MNum 0) (ECons (MNum 0) (ECons (MNum 0) (ENil))))
?neg_mask_strides)
(ICons ?lt_f32 (ICons ?neg_one (INil)))))
(= ?not_mask_f32 (Op (Add
?shape
?neg_mask_strides
(ECons (MNum 0) (ECons (MNum 0) (ECons (MNum 0) (ENil))))
?not_mask_f32_strides)
(ICons ?neg_mask (ICons ?one (INil)))))
(= ?not_mask_bool (Op (Cast ?size (Bool)) (ICons ?not_mask_f32 (INil))))
(= ?not_mask (Op (Cast ?size (F32)) (ICons ?not_mask_bool (INil))))
(= ?positive (Op (Mul
?shape
?not_mask_f32_strides
?matmul_strides
?positive_strides)
(ICons ?not_mask (ICons ?matmul (INil)))))
(= ?relu (Op (Add
?shape
?zeroed_strides
?positive_strides
?relu_strides)
(ICons ?zeroed (ICons ?positive (INil)))))
)
(
(let ?fused (Op (cublaslt
?m ?n ?k
?a_layout ?b_layout
?a_order ?b_order ?c_order ?d_order
?lda ?ldb ?ldc ?ldd
?batch
?stride_a ?stride_b ?stride_c ?stride_d
?a_dtype ?b_dtype ?c_dtype (F32)
?compute_type ?scale_dtype
?alpha 0.0 "RELU")
(ICons ?a (ICons ?b ?matmul_tail))))
(union ?relu ?fused)
(set (dtype ?fused) (F32))
)
:ruleset matmul_backend
:name "cublaslt batched relu epilogue"
)
(rule
(
(= ?matmul (Op (cublaslt
?m ?n ?k
?a_layout ?b_layout
?a_order ?b_order ?c_order ?d_order
?lda ?ldb ?ldc ?ldd
?batch
?stride_a ?stride_b ?stride_c ?stride_d
?a_dtype ?b_dtype ?c_dtype (F32)
?compute_type ?scale_dtype
?alpha 0.0 "BIAS")
(ICons ?a (ICons ?b ?matmul_tail))))
(= ?zero (Op (Constant 0.0) (INil)))
(= ?neg_one (Op (Constant -1.0) (INil)))
(= ?one (Op (Constant 1.0) (INil)))
(= ?lt (Op (LessThan
?shape
?matmul_strides
(ECons (MNum 0) (ECons (MNum 0) (ENil)))
?mask_strides)
(ICons ?matmul (ICons ?zero (INil)))))
(= ?lt_f32 (Op (Cast ?size (F32)) (ICons ?lt (INil))))
(= ?zeroed (Op (Mul
?shape
?mask_strides
(ECons (MNum 0) (ECons (MNum 0) (ENil)))
?zeroed_strides)
(ICons ?lt_f32 (ICons ?zero (INil)))))
(= ?neg_mask (Op (Mul
?shape
?mask_strides
(ECons (MNum 0) (ECons (MNum 0) (ENil)))
?neg_mask_strides)
(ICons ?lt_f32 (ICons ?neg_one (INil)))))
(= ?not_mask_f32 (Op (Add
?shape
?neg_mask_strides
(ECons (MNum 0) (ECons (MNum 0) (ENil)))
?not_mask_f32_strides)
(ICons ?neg_mask (ICons ?one (INil)))))
(= ?not_mask_bool (Op (Cast ?size (Bool)) (ICons ?not_mask_f32 (INil))))
(= ?not_mask (Op (Cast ?size (F32)) (ICons ?not_mask_bool (INil))))
(= ?positive (Op (Mul
?shape
?not_mask_f32_strides
?matmul_strides
?positive_strides)
(ICons ?not_mask (ICons ?matmul (INil)))))
(= ?relu (Op (Add
?shape
?zeroed_strides
?positive_strides
?relu_strides)
(ICons ?zeroed (ICons ?positive (INil)))))
)
(
(let ?fused (Op (cublaslt
?m ?n ?k
?a_layout ?b_layout
?a_order ?b_order ?c_order ?d_order
?lda ?ldb ?ldc ?ldd
?batch
?stride_a ?stride_b ?stride_c ?stride_d
?a_dtype ?b_dtype ?c_dtype (F32)
?compute_type ?scale_dtype
?alpha 0.0 "RELU_BIAS")
(ICons ?a (ICons ?b ?matmul_tail))))
(union ?relu ?fused)
(set (dtype ?fused) (F32))
)
:ruleset matmul_backend
:name "cublaslt 2d relu bias epilogue"
)
(rule
(
(= ?matmul (Op (cublaslt
?m ?n ?k
?a_layout ?b_layout
?a_order ?b_order ?c_order ?d_order
?lda ?ldb ?ldc ?ldd
?batch
?stride_a ?stride_b ?stride_c ?stride_d
?a_dtype ?b_dtype ?c_dtype (F32)
?compute_type ?scale_dtype
?alpha 0.0 "BIAS")
(ICons ?a (ICons ?b ?matmul_tail))))
(= ?zero (Op (Constant 0.0) (INil)))
(= ?neg_one (Op (Constant -1.0) (INil)))
(= ?one (Op (Constant 1.0) (INil)))
(= ?lt (Op (LessThan
?shape
?matmul_strides
(ECons (MNum 0) (ECons (MNum 0) (ECons (MNum 0) (ENil))))
?mask_strides)
(ICons ?matmul (ICons ?zero (INil)))))
(= ?lt_f32 (Op (Cast ?size (F32)) (ICons ?lt (INil))))
(= ?zeroed (Op (Mul
?shape
?mask_strides
(ECons (MNum 0) (ECons (MNum 0) (ECons (MNum 0) (ENil))))
?zeroed_strides)
(ICons ?lt_f32 (ICons ?zero (INil)))))
(= ?neg_mask (Op (Mul
?shape
?mask_strides
(ECons (MNum 0) (ECons (MNum 0) (ECons (MNum 0) (ENil))))
?neg_mask_strides)
(ICons ?lt_f32 (ICons ?neg_one (INil)))))
(= ?not_mask_f32 (Op (Add
?shape
?neg_mask_strides
(ECons (MNum 0) (ECons (MNum 0) (ECons (MNum 0) (ENil))))
?not_mask_f32_strides)
(ICons ?neg_mask (ICons ?one (INil)))))
(= ?not_mask_bool (Op (Cast ?size (Bool)) (ICons ?not_mask_f32 (INil))))
(= ?not_mask (Op (Cast ?size (F32)) (ICons ?not_mask_bool (INil))))
(= ?positive (Op (Mul
?shape
?not_mask_f32_strides
?matmul_strides
?positive_strides)
(ICons ?not_mask (ICons ?matmul (INil)))))
(= ?relu (Op (Add
?shape
?zeroed_strides
?positive_strides
?relu_strides)
(ICons ?zeroed (ICons ?positive (INil)))))
)
(
(let ?fused (Op (cublaslt
?m ?n ?k
?a_layout ?b_layout
?a_order ?b_order ?c_order ?d_order
?lda ?ldb ?ldc ?ldd
?batch
?stride_a ?stride_b ?stride_c ?stride_d
?a_dtype ?b_dtype ?c_dtype (F32)
?compute_type ?scale_dtype
?alpha 0.0 "RELU_BIAS")
(ICons ?a (ICons ?b ?matmul_tail))))
(union ?relu ?fused)
(set (dtype ?fused) (F32))
)
:ruleset matmul_backend
:name "cublaslt batched relu bias epilogue"
)
; Canonical tanh-approx GELU can also appear directly as:
;
; x * sigmoid(1.5957691216 * x * (1 + 0.044715 * x * x))
;
; Match that sigmoid form and fuse it into the cuBLASLt GELU epilogues.
(rule
(
(= ?matmul (Op (cublaslt
?m ?n ?k
?a_layout ?b_layout
?a_order ?b_order ?c_order ?d_order
?lda ?ldb ?ldc ?ldd
?batch
?stride_a ?stride_b ?stride_c ?stride_d
?a_dtype ?b_dtype ?c_dtype (F32)
?compute_type ?scale_dtype
?alpha 0.0 "DEFAULT")
(ICons ?a (ICons ?b ?matmul_tail))))
(= ?gelu_coeff_inner (Op (Constant 0.044715) (INil)))
(= ?gelu_inner_scaled (Op (Mul ?gelu_inner_scaled_shape ?gelu_inner_scaled_a_stride ?gelu_inner_scaled_b_stride ?gelu_inner_scaled_out_stride) (ICons ?matmul (ICons ?gelu_coeff_inner (INil)))))
(= ?gelu_inner_quad (Op (Mul ?gelu_inner_quad_shape ?gelu_inner_quad_a_stride ?gelu_inner_quad_b_stride ?gelu_inner_quad_out_stride) (ICons ?gelu_inner_scaled (ICons ?matmul (INil)))))
(= ?gelu_one (Op (Constant 1.000000) (INil)))
(= ?gelu_poly (Op (Add ?gelu_poly_shape ?gelu_poly_a_stride ?gelu_poly_b_stride ?gelu_poly_out_stride) (ICons ?gelu_inner_quad (ICons ?gelu_one (INil)))))
(= ?gelu_coeff_outer (Op (Constant 1.595769) (INil)))
(= ?gelu_outer_scaled (Op (Mul ?gelu_outer_scaled_shape ?gelu_outer_scaled_a_stride ?gelu_outer_scaled_b_stride ?gelu_outer_scaled_out_stride) (ICons ?matmul (ICons ?gelu_coeff_outer (INil)))))
(= ?gelu_scaled (Op (Mul ?gelu_scaled_shape ?gelu_scaled_a_stride ?gelu_scaled_b_stride ?gelu_scaled_out_stride) (ICons ?gelu_outer_scaled (ICons ?gelu_poly (INil)))))
(= ?neg1 (Op (Constant -1.000000) (INil)))
(= ?gelu_neg (Op (Mul ?gelu_neg_shape ?gelu_neg_a_stride ?gelu_neg_b_stride ?gelu_neg_out_stride) (ICons ?gelu_scaled (ICons ?neg1 (INil)))))
(= ?log2e (Op (Constant 1.442695) (INil)))
(= ?gelu_exp_scaled (Op (Mul ?gelu_exp_scaled_shape ?gelu_exp_scaled_a_stride ?gelu_exp_scaled_b_stride ?gelu_exp_scaled_out_stride) (ICons ?gelu_neg (ICons ?log2e (INil)))))
(= ?gelu_exp2_val (Op (Exp2 ?gelu_exp_shape ?gelu_exp_in_stride ?gelu_exp_out_stride) (ICons ?gelu_exp_scaled (INil))))
(= ?gelu_plus1 (Op (Add ?gelu_plus1_shape ?gelu_plus1_a_stride ?gelu_plus1_b_stride ?gelu_plus1_out_stride) (ICons ?gelu_exp2_val (ICons ?gelu_one (INil)))))
(= ?gelu_sigmoid (Op (Recip ?gelu_sigmoid_shape ?gelu_sigmoid_in_stride ?gelu_sigmoid_out_stride) (ICons ?gelu_plus1 (INil))))
(= ?gelu_out (Op (Mul ?gelu_out_shape ?gelu_out_a_stride ?gelu_out_b_stride ?gelu_out_out_stride) (ICons ?matmul (ICons ?gelu_sigmoid (INil)))))
)
(
(let ?fused (Op (cublaslt
?m ?n ?k
?a_layout ?b_layout
?a_order ?b_order ?c_order ?d_order
?lda ?ldb ?ldc ?ldd
?batch
?stride_a ?stride_b ?stride_c ?stride_d
?a_dtype ?b_dtype ?c_dtype (F32)
?compute_type ?scale_dtype
?alpha 0.0 "GELU")
(ICons ?a (ICons ?b ?matmul_tail))))
(union ?gelu_out ?fused)
(set (dtype ?fused) (F32))
)
:ruleset matmul_backend
:name "cublaslt gelu epilogue"
)
(rule
(
(= ?matmul (Op (cublaslt
?m ?n ?k
?a_layout ?b_layout
?a_order ?b_order ?c_order ?d_order
?lda ?ldb ?ldc ?ldd
?batch
?stride_a ?stride_b ?stride_c ?stride_d
?a_dtype ?b_dtype ?c_dtype (F32)
?compute_type ?scale_dtype
?alpha 0.0 "BIAS")
(ICons ?a (ICons ?b ?matmul_tail))))
(= ?gelu_coeff_inner (Op (Constant 0.044715) (INil)))
(= ?gelu_inner_scaled (Op (Mul ?gelu_inner_scaled_shape ?gelu_inner_scaled_a_stride ?gelu_inner_scaled_b_stride ?gelu_inner_scaled_out_stride) (ICons ?matmul (ICons ?gelu_coeff_inner (INil)))))
(= ?gelu_inner_quad (Op (Mul ?gelu_inner_quad_shape ?gelu_inner_quad_a_stride ?gelu_inner_quad_b_stride ?gelu_inner_quad_out_stride) (ICons ?gelu_inner_scaled (ICons ?matmul (INil)))))
(= ?gelu_one (Op (Constant 1.000000) (INil)))
(= ?gelu_poly (Op (Add ?gelu_poly_shape ?gelu_poly_a_stride ?gelu_poly_b_stride ?gelu_poly_out_stride) (ICons ?gelu_inner_quad (ICons ?gelu_one (INil)))))
(= ?gelu_coeff_outer (Op (Constant 1.595769) (INil)))
(= ?gelu_outer_scaled (Op (Mul ?gelu_outer_scaled_shape ?gelu_outer_scaled_a_stride ?gelu_outer_scaled_b_stride ?gelu_outer_scaled_out_stride) (ICons ?matmul (ICons ?gelu_coeff_outer (INil)))))
(= ?gelu_scaled (Op (Mul ?gelu_scaled_shape ?gelu_scaled_a_stride ?gelu_scaled_b_stride ?gelu_scaled_out_stride) (ICons ?gelu_outer_scaled (ICons ?gelu_poly (INil)))))
(= ?neg1 (Op (Constant -1.000000) (INil)))
(= ?gelu_neg (Op (Mul ?gelu_neg_shape ?gelu_neg_a_stride ?gelu_neg_b_stride ?gelu_neg_out_stride) (ICons ?gelu_scaled (ICons ?neg1 (INil)))))
(= ?log2e (Op (Constant 1.442695) (INil)))
(= ?gelu_exp_scaled (Op (Mul ?gelu_exp_scaled_shape ?gelu_exp_scaled_a_stride ?gelu_exp_scaled_b_stride ?gelu_exp_scaled_out_stride) (ICons ?gelu_neg (ICons ?log2e (INil)))))
(= ?gelu_exp2_val (Op (Exp2 ?gelu_exp_shape ?gelu_exp_in_stride ?gelu_exp_out_stride) (ICons ?gelu_exp_scaled (INil))))
(= ?gelu_plus1 (Op (Add ?gelu_plus1_shape ?gelu_plus1_a_stride ?gelu_plus1_b_stride ?gelu_plus1_out_stride) (ICons ?gelu_exp2_val (ICons ?gelu_one (INil)))))
(= ?gelu_sigmoid (Op (Recip ?gelu_sigmoid_shape ?gelu_sigmoid_in_stride ?gelu_sigmoid_out_stride) (ICons ?gelu_plus1 (INil))))
(= ?gelu_out (Op (Mul ?gelu_out_shape ?gelu_out_a_stride ?gelu_out_b_stride ?gelu_out_out_stride) (ICons ?matmul (ICons ?gelu_sigmoid (INil)))))
)
(
(let ?fused (Op (cublaslt
?m ?n ?k
?a_layout ?b_layout
?a_order ?b_order ?c_order ?d_order
?lda ?ldb ?ldc ?ldd
?batch
?stride_a ?stride_b ?stride_c ?stride_d
?a_dtype ?b_dtype ?c_dtype (F32)
?compute_type ?scale_dtype
?alpha 0.0 "GELU_BIAS")
(ICons ?a (ICons ?b ?matmul_tail))))
(union ?gelu_out ?fused)
(set (dtype ?fused) (F32))
)
:ruleset matmul_backend
:name "cublaslt gelu bias epilogue"
)
; This first slice fuses column-bias adds into CUBLASLT_EPILOGUE_BIAS for the
; older COL-ordered output view. In that view Luminal's logical [m,n] output is
; represented as a cuBLASLt [n,m] matrix, so cuBLASLt's row-broadcast bias maps
; to the common logical column bias of length n.
(rule
(
(= ?matmul (Op (cublaslt
?m ?n ?k
?a_layout ?b_layout
?a_order ?b_order ?c_order "COL"
?lda ?ldb ?ldc ?ldd
(MNum 1)
?stride_a ?stride_b ?stride_c ?stride_d
?a_dtype ?b_dtype ?c_dtype ?d_dtype
?compute_type ?scale_dtype
?alpha 0.0 "DEFAULT")
(ICons ?a (ICons ?b (INil)))))
(= ?add (Op (Add
(ECons ?n (ECons ?m (ENil)))
?matmul_add_strides
?bias_add_strides
?add_out_strides)
(ICons ?matmul (ICons ?bias (INil)))))
(= ?bias_add_strides (ECons (MNum 0) (ECons (MIter) (ENil))))
(= ?matmul_add_strides ?add_out_strides)
(= ?d_dtype (dtype ?bias))
)
(
(let ?fused (Op (cublaslt
?m ?n ?k
?a_layout ?b_layout
?a_order ?b_order ?c_order "COL"
?lda ?ldb ?ldc ?ldd
(MNum 1)
?stride_a ?stride_b ?stride_c ?stride_d
?a_dtype ?b_dtype ?c_dtype ?d_dtype
?compute_type ?scale_dtype
?alpha 0.0 "BIAS")
(ICons ?a (ICons ?b (ICons ?bias (INil))))))
(union ?add ?fused)
(set (dtype ?fused) ?d_dtype)
)
:ruleset matmul_backend
:name "cublaslt 2d matmul plus column bias epilogue"
)
(rule
(
(= ?matmul (Op (cublaslt
?m ?n ?k
?a_layout ?b_layout
?a_order ?b_order ?c_order "COL"
?lda ?ldb ?ldc ?ldd
(MNum 1)
?stride_a ?stride_b ?stride_c ?stride_d
?a_dtype ?b_dtype ?c_dtype ?d_dtype
?compute_type ?scale_dtype
?alpha 0.0 "DEFAULT")
(ICons ?a (ICons ?b (INil)))))
(= ?add (Op (Add
(ECons ?n (ECons ?m (ENil)))
?bias_add_strides
?matmul_add_strides
?add_out_strides)
(ICons ?bias (ICons ?matmul (INil)))))
(= ?bias_add_strides (ECons (MNum 0) (ECons (MIter) (ENil))))
(= ?matmul_add_strides ?add_out_strides)
(= ?d_dtype (dtype ?bias))
)
(
(let ?fused (Op (cublaslt
?m ?n ?k
?a_layout ?b_layout
?a_order ?b_order ?c_order "COL"
?lda ?ldb ?ldc ?ldd
(MNum 1)
?stride_a ?stride_b ?stride_c ?stride_d
?a_dtype ?b_dtype ?c_dtype ?d_dtype
?compute_type ?scale_dtype
?alpha 0.0 "BIAS")
(ICons ?a (ICons ?b (ICons ?bias (INil))))))
(union ?add ?fused)
(set (dtype ?fused) ?d_dtype)
)
:ruleset matmul_backend
:name "cublaslt 2d column bias plus matmul epilogue"
)
(rule
(
(= ?matmul (Op (cublaslt
?m ?n ?k
?a_layout ?b_layout
?a_order ?b_order ?c_order "COL"
?lda ?ldb ?ldc ?ldd
?batch
?stride_a ?stride_b ?stride_c ?stride_d
?a_dtype ?b_dtype ?c_dtype ?d_dtype
?compute_type ?scale_dtype
?alpha 0.0 "DEFAULT")
(ICons ?a (ICons ?b (INil)))))
(= ?add (Op (Add
(ECons ?batch (ECons ?n (ECons ?m (ENil))))
?matmul_add_strides
?bias_add_strides
?add_out_strides)
(ICons ?matmul (ICons ?bias (INil)))))
(= ?bias_add_strides (ECons (MNum 0) (ECons (MNum 0) (ECons (MIter) (ENil)))))
(= ?matmul_add_strides ?add_out_strides)
(= ?d_dtype (dtype ?bias))
)
(
(let ?fused (Op (cublaslt
?m ?n ?k
?a_layout ?b_layout
?a_order ?b_order ?c_order "COL"
?lda ?ldb ?ldc ?ldd
?batch
?stride_a ?stride_b ?stride_c ?stride_d
?a_dtype ?b_dtype ?c_dtype ?d_dtype
?compute_type ?scale_dtype
?alpha 0.0 "BIAS")
(ICons ?a (ICons ?b (ICons ?bias (INil))))))
(union ?add ?fused)
(set (dtype ?fused) ?d_dtype)
)
:ruleset matmul_backend
:name "cublaslt batched matmul plus column bias epilogue"
)
(rule
(
(= ?matmul (Op (cublaslt
?m ?n ?k
?a_layout ?b_layout
?a_order ?b_order ?c_order "COL"
?lda ?ldb ?ldc ?ldd
?batch
?stride_a ?stride_b ?stride_c ?stride_d
?a_dtype ?b_dtype ?c_dtype ?d_dtype
?compute_type ?scale_dtype
?alpha 0.0 "DEFAULT")
(ICons ?a (ICons ?b (INil)))))
(= ?add (Op (Add
(ECons ?batch (ECons ?n (ECons ?m (ENil))))
?bias_add_strides
?matmul_add_strides
?add_out_strides)
(ICons ?bias (ICons ?matmul (INil)))))
(= ?bias_add_strides (ECons (MNum 0) (ECons (MNum 0) (ECons (MIter) (ENil)))))
(= ?matmul_add_strides ?add_out_strides)
(= ?d_dtype (dtype ?bias))
)
(
(let ?fused (Op (cublaslt
?m ?n ?k
?a_layout ?b_layout
?a_order ?b_order ?c_order "COL"
?lda ?ldb ?ldc ?ldd
?batch
?stride_a ?stride_b ?stride_c ?stride_d
?a_dtype ?b_dtype ?c_dtype ?d_dtype
?compute_type ?scale_dtype
?alpha 0.0 "BIAS")
(ICons ?a (ICons ?b (ICons ?bias (INil))))))
(union ?add ?fused)
(set (dtype ?fused) ?d_dtype)
)
:ruleset matmul_backend
:name "cublaslt batched column bias plus matmul epilogue"
)
(rule
((= ?mul (Op (KernelMul ?shape ?as ?bs ?os ?dt) ?inputs))
(= (MNum 0) (nth_from_end ?as 1))
(= (MNum 0) (nth_from_end ?bs 2))
(= ?sum (Op (Sum ?sshape ?sk ?ssi ?sks ?sso) (ICons ?mul (INil))))
(= ?sum (Op (cublaslt ?cm ?cn ?ck ?cta ?ctb ?cao ?cbo ?cco ?cdo ?clda ?cldb ?cldc ?cldd ?cbc ?csa ?csb ?csc ?csd ?cadt ?cbdt ?ccdt ?cddt ?ccompute ?cscale ?calpha ?cbeta ?cepilogue) ?ci)))
((delete (Op (KernelMul ?shape ?as ?bs ?os ?dt) ?inputs)))
:ruleset cleanup
)
(rule
((= ?mul (Op (KernelMul ?shape ?as ?bs ?os ?dt) ?inputs))
(= (MNum 0) (nth_from_end ?as 1))
(= (MNum 0) (nth_from_end ?bs 2))
(= ?sum (Op (Sum ?sshape ?sk ?ssi ?sks ?sso) (ICons ?mul (INil))))
(= ?sum (Op (KernelBatchMatMul ?bos ?bk ?bas ?baks ?bbs ?bbks ?bouts ?bdt) ?bi)))
((delete (Op (KernelMul ?shape ?as ?bs ?os ?dt) ?inputs)))
:ruleset cleanup
)
(rule
(
(= ?e (Op (GLUMoE ?gu_io ?dn_io ?gu_matmul_k ?dn_matmul_k ?output_k ?gu_within_range ?dn_within_range ?mode) ?inputs))
)
(
(set (dtype ?e) (F32))
)
:ruleset dtype_prop
)
; GLUMoE: Match the expert computation subgraph of a gated MoE.
;
; One fused op supports two activation modes:
; mode=0: Qwen-style SwiGLU (silu(gate) * up)
; mode=1: Gemma-style GELU (gate * sigmoid(1.595769 * gate * (1 + 0.044715 * gate^2)))
;
; To keep matching fast, we stage through marker states:
; 1) Shared expert index/gather markers
; 2) Shared gate-up matmul marker
; 3) Activation marker (separate swiglu / gemma_gelu paths)
; 4) Down matmul marker (separate swiglu / gemma_gelu paths)
; 5) Final GLUMoE fusion (separate swiglu / gemma_gelu rules)
(datatype*
(GLUMoEExpertIndexState
(MkGLUMoEExpertIndexState Expression Expression IR)
)
(GLUMoEExpertGatherState
(MkGLUMoEExpertGatherState Expression Expression IR IR)
)
(GLUMoEGateUpState
(MkGLUMoEGateUpState Expression Expression Expression IR IR IR)
)
(GLUMoESwiGLUState
(MkGLUMoESwiGLUState GLUMoEGateUpState)
)
(GLUMoEGemmaGELUState
(MkGLUMoEGemmaGELUState GLUMoEGateUpState)
)
(GLUMoESwiGLUDownState
(MkGLUMoESwiGLUDownState Expression Expression Expression GLUMoESwiGLUState IR IR)
)
(GLUMoEGemmaDownState
(MkGLUMoEGemmaDownState Expression Expression Expression GLUMoEGemmaGELUState IR IR)
)
)
(function glumoe_expert_index (IR) GLUMoEExpertIndexState :merge new)
(function glumoe_expert_gather (IR) GLUMoEExpertGatherState :merge new)
(function glumoe_gate_up (IR) GLUMoEGateUpState :merge new)
(function glumoe_swiglu (IR) GLUMoESwiGLUState :merge new)
(function glumoe_gemma_gelu (IR) GLUMoEGemmaGELUState :merge new)
(function glumoe_swiglu_down (IR) GLUMoESwiGLUDownState :merge new)
(function glumoe_gemma_down (IR) GLUMoEGemmaDownState :merge new)
(rule
(
(= ?iota_base (Op (Iota ?io ?iota_base_range) (INil)))
(= ?mul_base (Op (Mul ?mul_base_shape ?mul_base_a_stride ?mul_base_b_stride ?mul_base_out_stride) (ICons ?topk_idx (ICons ?iota_base (INil)))))
(= ?iota_within (Op (Iota (MIter) ?iota_within_range) (INil)))
(= ?add_idx (Op (Add ?add_shape ?add_a_stride ?add_b_stride ?add_out_stride) (ICons ?mul_base (ICons ?iota_within (INil)))))
)
(
(set (glumoe_expert_index ?add_idx)
(MkGLUMoEExpertIndexState ?io ?iota_within_range ?topk_idx))
)
:ruleset glumoe
:name "GLUMoE expert index marker"
)
(rule
(
(= ?index_state (glumoe_expert_index ?idx))
(= ?index_state (MkGLUMoEExpertIndexState ?io ?within_range ?topk_idx))
(= ?gathered (Op (Gather ?gather_idx_shape ?gather_idx_stride ?gather_data_shape ?gather_data_stride) (ICons ?idx (ICons ?weights (INil)))))
(= ?f32 (Op (Cast ?f32_size (F32)) (ICons ?gathered (INil))))
)
(
(set (glumoe_expert_gather ?f32)
(MkGLUMoEExpertGatherState ?io ?within_range ?topk_idx ?weights))
)
:ruleset glumoe
:name "GLUMoE expert gather marker"
)
(rule
(
(= ?gather_state (glumoe_expert_gather ?gu_f32))
(= ?gather_state (MkGLUMoEExpertGatherState ?gu_io ?gu_iota_within_range ?topk_idx ?gate_up_w))
(= ?gu_matmul_mul (Op (Mul ?gu_matmul_mul_shape ?gu_matmul_a_stride ?gu_matmul_b_stride ?gu_matmul_mul_out_stride) (ICons ?x (ICons ?gu_f32 (INil)))))
(= ?gu_matmul (Op (Sum ?gu_matmul_out_shape ?gu_matmul_k ?gu_matmul_in_stride ?gu_matmul_k_stride ?gu_matmul_out_stride) (ICons ?gu_matmul_mul (INil))))
)
(
(set (glumoe_gate_up ?gu_matmul)
(MkGLUMoEGateUpState ?gu_io ?gu_matmul_k ?gu_iota_within_range ?x ?topk_idx ?gate_up_w))
)
:ruleset glumoe
:name "GLUMoE gate-up matmul marker"
)
; ===== SwiGLU activation marker =====
(rule
(
(= ?gate_up_state (glumoe_gate_up ?gu_matmul))
(= ?gate_up_state (MkGLUMoEGateUpState ?gu_io ?gu_matmul_k ?gu_within_range ?x ?topk_idx ?gate_up_w))
(= ?up_iota (Op (Iota ?up_iota_expr ?up_iota_range) (INil)))
(= ?up_slice (Op (Gather ?up_gather_idx_shape ?up_gather_idx_stride ?up_gather_data_shape ?up_gather_data_stride) (ICons ?up_iota (ICons ?gu_matmul (INil)))))
(= ?neg1 (Op (Constant -1.000000) (INil)))
(= ?neg_gate (Op (Mul ?silu_shape1 ?silu_a_stride1 ?silu_b_stride1 ?silu_out_stride1) (ICons ?gu_matmul (ICons ?neg1 (INil)))))
(= ?log2e (Op (Constant 1.442695) (INil)))
(= ?scaled (Op (Mul ?silu_shape2 ?silu_a_stride2 ?silu_b_stride2 ?silu_out_stride2) (ICons ?neg_gate (ICons ?log2e (INil)))))
(= ?exp2_val (Op (Exp2 ?silu_shape3 ?silu_in_stride3 ?silu_out_stride3) (ICons ?scaled (INil))))
(= ?one (Op (Constant 1.000000) (INil)))
(= ?plus1 (Op (Add ?silu_shape4 ?silu_a_stride4 ?silu_b_stride4 ?silu_out_stride4) (ICons ?exp2_val (ICons ?one (INil)))))
(= ?sigmoid (Op (Recip ?silu_shape5 ?silu_in_stride5 ?silu_out_stride5) (ICons ?plus1 (INil))))
(= ?silu_out (Op (Mul ?silu_shape6 ?silu_a_stride6 ?silu_b_stride6 ?silu_out_stride6) (ICons ?gu_matmul (ICons ?sigmoid (INil)))))
(= ?swiglu_out (Op (Mul ?swiglu_shape ?swiglu_a_stride ?swiglu_b_stride ?swiglu_out_stride) (ICons ?silu_out (ICons ?up_slice (INil)))))
)
(
(set (glumoe_swiglu ?swiglu_out) (MkGLUMoESwiGLUState ?gate_up_state))
)
:ruleset glumoe
:name "GLUMoE swiglu marker"
)
; ===== Gemma GELU activation marker =====
(rule
(
(= ?gate_up_state (glumoe_gate_up ?gu_matmul))
(= ?gate_up_state (MkGLUMoEGateUpState ?gu_io ?gu_matmul_k ?gu_within_range ?x ?topk_idx ?gate_up_w))
(= ?up_iota (Op (Iota ?up_iota_expr ?up_iota_range) (INil)))
(= ?up_slice (Op (Gather ?up_gather_idx_shape ?up_gather_idx_stride ?up_gather_data_shape ?up_gather_data_stride) (ICons ?up_iota (ICons ?gu_matmul (INil)))))
(= ?gelu_coeff_inner (Op (Constant 0.044715) (INil)))
(= ?gelu_inner_scaled (Op (Mul ?gelu_inner_scaled_shape ?gelu_inner_scaled_a_stride ?gelu_inner_scaled_b_stride ?gelu_inner_scaled_out_stride) (ICons ?gu_matmul (ICons ?gelu_coeff_inner (INil)))))
(= ?gelu_inner_quad (Op (Mul ?gelu_inner_quad_shape ?gelu_inner_quad_a_stride ?gelu_inner_quad_b_stride ?gelu_inner_quad_out_stride) (ICons ?gelu_inner_scaled (ICons ?gu_matmul (INil)))))
(= ?gelu_one (Op (Constant 1.000000) (INil)))
(= ?gelu_poly (Op (Add ?gelu_poly_shape ?gelu_poly_a_stride ?gelu_poly_b_stride ?gelu_poly_out_stride) (ICons ?gelu_inner_quad (ICons ?gelu_one (INil)))))
(= ?gelu_coeff_outer (Op (Constant 1.595769) (INil)))
(= ?gelu_outer_scaled (Op (Mul ?gelu_outer_scaled_shape ?gelu_outer_scaled_a_stride ?gelu_outer_scaled_b_stride ?gelu_outer_scaled_out_stride) (ICons ?gu_matmul (ICons ?gelu_coeff_outer (INil)))))
(= ?gelu_scaled (Op (Mul ?gelu_scaled_shape ?gelu_scaled_a_stride ?gelu_scaled_b_stride ?gelu_scaled_out_stride) (ICons ?gelu_outer_scaled (ICons ?gelu_poly (INil)))))
(= ?neg1 (Op (Constant -1.000000) (INil)))
(= ?gelu_neg (Op (Mul ?gelu_neg_shape ?gelu_neg_a_stride ?gelu_neg_b_stride ?gelu_neg_out_stride) (ICons ?gelu_scaled (ICons ?neg1 (INil)))))
(= ?log2e (Op (Constant 1.442695) (INil)))
(= ?gelu_exp_scaled (Op (Mul ?gelu_exp_scaled_shape ?gelu_exp_scaled_a_stride ?gelu_exp_scaled_b_stride ?gelu_exp_scaled_out_stride) (ICons ?gelu_neg (ICons ?log2e (INil)))))
(= ?gelu_exp2_val (Op (Exp2 ?gelu_exp_shape ?gelu_exp_in_stride ?gelu_exp_out_stride) (ICons ?gelu_exp_scaled (INil))))
(= ?gelu_plus1 (Op (Add ?gelu_plus1_shape ?gelu_plus1_a_stride ?gelu_plus1_b_stride ?gelu_plus1_out_stride) (ICons ?gelu_exp2_val (ICons ?gelu_one (INil)))))
(= ?gelu_sigmoid (Op (Recip ?gelu_sigmoid_shape ?gelu_sigmoid_in_stride ?gelu_sigmoid_out_stride) (ICons ?gelu_plus1 (INil))))
(= ?gelu_out (Op (Mul ?gelu_out_shape ?gelu_out_a_stride ?gelu_out_b_stride ?gelu_out_out_stride) (ICons ?gu_matmul (ICons ?gelu_sigmoid (INil)))))
(= ?gemma_out (Op (Mul ?geglu_shape ?geglu_a_stride ?geglu_b_stride ?geglu_out_stride) (ICons ?gelu_out (ICons ?up_slice (INil)))))
)
(
(set (glumoe_gemma_gelu ?gemma_out) (MkGLUMoEGemmaGELUState ?gate_up_state))
)
:ruleset glumoe
:name "GLUMoE gemma gelu marker"
)
; ===== SwiGLU down marker =====
(rule
(
(= ?swiglu_state (glumoe_swiglu ?swiglu_out))
(= ?swiglu_state (MkGLUMoESwiGLUState ?gate_up_state))
(= ?gather_state (glumoe_expert_gather ?dn_f32))
(= ?gather_state (MkGLUMoEExpertGatherState ?dn_io ?dn_iota_within_range ?topk_idx ?down_w))
(= ?dn_matmul_mul (Op (Mul ?dn_matmul_mul_shape ?dn_matmul_a_stride ?dn_matmul_b_stride ?dn_matmul_mul_out_stride) (ICons ?swiglu_out (ICons ?dn_f32 (INil)))))
(= ?dn_matmul (Op (Sum ?dn_matmul_out_shape ?dn_matmul_k ?dn_matmul_in_stride ?dn_matmul_k_stride ?dn_matmul_out_stride) (ICons ?dn_matmul_mul (INil))))
)
(
(set (glumoe_swiglu_down ?dn_matmul)
(MkGLUMoESwiGLUDownState ?dn_io ?dn_matmul_k ?dn_iota_within_range ?swiglu_state ?topk_idx ?down_w))
)
:ruleset glumoe
:name "GLUMoE swiglu down marker"
)
; ===== Gemma GELU down marker =====
(rule
(
(= ?gemma_state (glumoe_gemma_gelu ?gemma_out))
(= ?gemma_state (MkGLUMoEGemmaGELUState ?gate_up_state))
(= ?gather_state (glumoe_expert_gather ?dn_f32))
(= ?gather_state (MkGLUMoEExpertGatherState ?dn_io ?dn_iota_within_range ?topk_idx ?down_w))
(= ?dn_matmul_mul (Op (Mul ?dn_matmul_mul_shape ?dn_matmul_a_stride ?dn_matmul_b_stride ?dn_matmul_mul_out_stride) (ICons ?gemma_out (ICons ?dn_f32 (INil)))))
(= ?dn_matmul (Op (Sum ?dn_matmul_out_shape ?dn_matmul_k ?dn_matmul_in_stride ?dn_matmul_k_stride ?dn_matmul_out_stride) (ICons ?dn_matmul_mul (INil))))
)
(
(set (glumoe_gemma_down ?dn_matmul)
(MkGLUMoEGemmaDownState ?dn_io ?dn_matmul_k ?dn_iota_within_range ?gemma_state ?topk_idx ?down_w))
)
:ruleset glumoe
:name "GLUMoE gemma down marker"
)
; ===== Final fusion: mode 0 (SwiGLU) =====
(rule
(
(= ?down_state (glumoe_swiglu_down ?dn_matmul))
(= ?down_state (MkGLUMoESwiGLUDownState ?dn_io ?dn_matmul_k ?dn_within_range ?swiglu_state ?topk_idx ?down_w))
(= ?swiglu_state (MkGLUMoESwiGLUState ?gate_up_state))
(= ?gate_up_state (MkGLUMoEGateUpState ?gu_io ?gu_matmul_k ?gu_within_range ?x ?topk_idx ?gate_up_w))
(= ?weighted (Op (Mul ?weighted_shape ?weighted_a_stride ?weighted_b_stride ?weighted_out_stride) (ICons ?dn_matmul (ICons ?topk_vals (INil)))))
(= ?output (Op (Sum ?output_shape ?output_k ?output_in_stride ?output_k_stride ?output_out_stride) (ICons ?weighted (INil))))
)
(
(let ?glumoe (Op (GLUMoE
?gu_io ?dn_io ?gu_matmul_k ?dn_matmul_k ?output_k
?gu_within_range ?dn_within_range (MNum 0))
(ICons ?x (ICons ?topk_idx (ICons ?topk_vals (ICons ?gate_up_w (ICons ?down_w (ICons ?topk_vals (INil)))))))))
(union ?output ?glumoe)
(subsume (Op (Sum ?output_shape ?output_k ?output_in_stride ?output_k_stride ?output_out_stride) (ICons ?weighted (INil))))
(subsume (Op (KernelSum ?output_shape ?output_k ?output_in_stride ?output_k_stride ?output_out_stride (F32)) (ICons ?weighted (INil))))
)
:ruleset glumoe
:name "GLUMoE fused expert computation (swiglu)"
)
; ===== Final fusion: mode 1 (Gemma GELU) =====
(rule
(
(= ?down_state (glumoe_gemma_down ?dn_matmul))
(= ?down_state (MkGLUMoEGemmaDownState ?dn_io ?dn_matmul_k ?dn_within_range ?gemma_state ?topk_idx ?down_w))
(= ?gemma_state (MkGLUMoEGemmaGELUState ?gate_up_state))
(= ?gate_up_state (MkGLUMoEGateUpState ?gu_io ?gu_matmul_k ?gu_within_range ?x ?topk_idx ?gate_up_w))
; Gemma expert weights: topk_weights = normed_topk * per_expert_scale.gather(topk_idx)
(= ?per_expert_vals (Op (Gather ?scale_gather_idx_shape ?scale_gather_idx_stride ?scale_gather_data_shape ?scale_gather_data_stride) (ICons ?topk_idx (ICons ?per_expert_scale (INil)))))
(= ?topk_row_offsets (Op (Iota ?topk_row_offsets_expr ?topk_row_offsets_range) (INil)))
(= ?topk_flat_idx (Op (Add ?topk_flat_idx_shape ?topk_flat_idx_a_stride ?topk_flat_idx_b_stride ?topk_flat_idx_out_stride) (ICons ?topk_row_offsets (ICons ?topk_idx (INil)))))
(= ?topk_vals (Op (Gather ?topk_vals_gather_idx_shape ?topk_vals_gather_idx_stride ?topk_vals_gather_data_shape ?topk_vals_gather_data_stride) (ICons ?topk_flat_idx (ICons ?routing_weights (INil)))))
(= ?topk_norm (Op (Sum ?topk_norm_shape ?output_k ?topk_norm_in_stride ?topk_norm_k_stride ?topk_norm_out_stride) (ICons ?topk_vals (INil))))
(= ?topk_norm_factor (Op (Recip ?topk_norm_recip_shape ?topk_norm_recip_in_stride ?topk_norm_recip_out_stride) (ICons ?topk_norm (INil))))
(= ?normed_topk (Op (Mul ?normed_topk_shape ?normed_topk_a_stride ?normed_topk_b_stride ?normed_topk_out_stride) (ICons ?topk_vals (ICons ?topk_norm_factor (INil)))))
(= ?expert_weights (Op (Mul ?expert_weights_shape ?expert_weights_a_stride ?expert_weights_b_stride ?expert_weights_out_stride) (ICons ?normed_topk (ICons ?per_expert_vals (INil)))))
(= ?weighted (Op (Mul ?weighted_shape ?weighted_a_stride ?weighted_b_stride ?weighted_out_stride) (ICons ?dn_matmul (ICons ?expert_weights (INil)))))
(= ?output (Op (Sum ?output_shape ?output_k ?output_in_stride ?output_k_stride ?output_out_stride) (ICons ?weighted (INil))))
)
(
(let ?glumoe (Op (GLUMoE
?gu_io ?dn_io ?gu_matmul_k ?dn_matmul_k ?output_k
?gu_within_range ?dn_within_range (MNum 1))
(ICons ?x (ICons ?topk_idx (ICons ?topk_vals (ICons ?gate_up_w (ICons ?down_w (ICons ?per_expert_scale (INil)))))))))
(union ?output ?glumoe)
(subsume (Op (Sum ?output_shape ?output_k ?output_in_stride ?output_k_stride ?output_out_stride) (ICons ?weighted (INil))))
(subsume (Op (KernelSum ?output_shape ?output_k ?output_in_stride ?output_k_stride ?output_out_stride (F32)) (ICons ?weighted (INil))))
)
:ruleset glumoe
:name "GLUMoE fused expert computation (gemma_gelu)"
)
; FlashInfer batch decode attention rewrite rule.
;
; Matches the paged attention pattern for ANY model with GQA:
; Gather(K_cache) → GQA broadcast → Q*K^T matmul → scale → add mask → softmax → attn*V matmul
; Gather(V_cache) → GQA broadcast ──────────────────────────────────────────→ attn*V matmul
;
; Structural anchors (prevent false matches on MLP/other ops):
; - Gather ops from 2D cache pools (MLP never uses Gather)
; - GQA broadcast via Mul(gathered, Constant(1.0)) with all-zero strides
; - Scale Mul(QK, constant) connecting QK scores to mask Add
; - Mask Add with zero-stride broadcast in first dim (nheads broadcast)
; - Data flow: two sequential matmul+reduce pairs connected through softmax
;
; The egglog rule captures the mask as 5th input. During extract(), a Rust
; function walks the mask's computation chain in the e-graph to locate the
; qo_indptr and kv_indptr Input nodes (validated via the Constant(1e10) anchor
; and structural checks). These are appended as inputs 5 and 6 so FlashInfer
; can build the CSR page table directly — no runtime derivation needed.
;
; Shape dimensions are egglog variables, not pinned constants.
; Dynamic dims "s" (batch/seq) and "c" (context) stay pinned as MVar.
(rule
(
; ── Second matmul: Mul(softmax_out, V_gqa) ──
; Shape: (nheads, s, hdim, c) — 4D
(= ?mul2 (Op (Mul
(ECons ?nheads (ECons (MVar "s") (ECons ?hdim (ECons (MVar "c") (ENil)))))
?mul2_a_strides
?mul2_b_strides
?mul2_out_strides)
(ICons ?soft (ICons ?v_gqa (INil)))))
; ── Second matmul: Sum (reduction over c) → output ──
; Shape: (nheads, s, hdim) — reduces c
(= ?output (Op (Sum
(ECons ?nheads2 (ECons (MVar "s") (ECons ?hdim2 (ENil))))
(MVar "c")
?out_in_strides
(MIter)
?out_out_strides)
(ICons ?mul2 (INil))))
; ── V GQA broadcast: Mul(V_gathered, 1.0) with zero-stride constant ──
; Shape: (nheads, c, hdim) — 3D
(= ?v_gqa_const (Op (Constant 1.000000) (INil)))
(= ?v_gqa (Op (Mul
(ECons ?nheads3 (ECons (MVar "c") (ECons ?hdim3 (ENil))))
?v_gqa_a_strides
(ECons (MNum 0) (ECons (MNum 0) (ECons (MNum 0) (ENil))))
?v_gqa_out_strides)
(ICons ?v_gathered (ICons ?v_gqa_const (INil)))))
; ── V Gather: rows from V_cache (2D) ──
; Shape: (c, kvdim), Source: (num_slots, kvdim)
(= ?v_gathered (Op (Gather
(ECons (MVar "c") (ECons ?kvdim (ENil)))
?v_gather_strides
(ECons ?num_slots_v (ECons ?kvdim2 (ENil)))
?v_src_strides)
(ICons ?v_idx (ICons ?v_cache (INil)))))
; ── First matmul: Mul(Q, K_gqa) ──
; Shape: (nheads, s, c, hdim) — 4D
(= ?mul1 (Op (Mul
(ECons ?nheads4 (ECons (MVar "s") (ECons (MVar "c") (ECons ?hdim4 (ENil)))))
?mul1_a_strides
?mul1_b_strides
?mul1_out_strides)
(ICons ?q (ICons ?k_gqa (INil)))))
; ── First matmul: Sum (reduction over hdim) → QK scores ──
; Shape: (nheads, s, c) — reduces hdim
(= ?qk (Op (Sum
(ECons ?nheads5 (ECons (MVar "s") (ECons (MVar "c") (ENil))))
?hdim5
?qk_in_strides
(MIter)
?qk_out_strides)
(ICons ?mul1 (INil))))
; ── Mask Add: Add(scaled_QK, mask) ──
; Shape: (nheads, s, c) — 3D
; Mask is broadcast from (s, c) via zero-stride in first dim (nheads).
(= ?masked (Op (Add
(ECons ?nheads8 (ECons (MVar "s") (ECons (MVar "c") (ENil))))
?mask_add_a_strides
(ECons (MNum 0) ?mask_rest_strides)
?mask_add_out_strides)
(ICons ?scaled_qk (ICons ?mask (INil)))))
; ── K GQA broadcast: Mul(K_gathered, 1.0) with zero-stride constant ──
; Shape: (nheads, hdim, c) — 3D
(= ?k_gqa_const (Op (Constant 1.000000) (INil)))
(= ?k_gqa (Op (Mul
(ECons ?nheads6 (ECons ?hdim6 (ECons (MVar "c") (ENil))))
?k_gqa_a_strides
(ECons (MNum 0) (ECons (MNum 0) (ECons (MNum 0) (ENil))))
?k_gqa_out_strides)
(ICons ?k_gathered (ICons ?k_gqa_const (INil)))))
; ── K Gather: rows from K_cache (2D) ──
; Shape: (c, kvdim), Source: (num_slots, kvdim)
(= ?k_gathered (Op (Gather
(ECons (MVar "c") (ECons ?kvdim3 (ENil)))
?k_gather_strides
(ECons ?num_slots_k (ECons ?kvdim4 (ENil)))
?k_src_strides)
(ICons ?k_idx (ICons ?k_cache (INil)))))
; ── Dtype consistency ──
(= ?dt (dtype ?q))
(= ?dt (dtype ?k_cache))
(= ?dt (dtype ?v_cache))
)
(
(let ?fi (Op (FlashInferAttention
?nheads (MDiv ?kvdim ?hdim) ?hdim (MNum 1) (MVar "s"))
(ICons ?q (ICons ?k_cache (ICons ?v_cache (ICons ?k_idx (ICons ?mask (INil))))))))
(union ?output ?fi)
(set (dtype ?fi) ?dt)
)
:ruleset matmul_backend
:name "FlashInfer batch decode attention"
)
(rule ((= ?__e (Input ?v60_node ?v60_label ?v60_dtype))) ((set (dtype ?__e) ?v60_dtype)) :ruleset dtype_prop)
(rule ((= ?__e (Output ?v61_inp ?v61_node)) (= ?__dty (dtype ?v61_inp))) ((set (dtype ?__e) ?__dty)) :ruleset dtype_prop)
(rule ((= ?__e (Op (CustomOpKind ?v62_id ?v62_dtype) ?__inputs))) ((set (dtype ?__e) ?v62_dtype)) :ruleset dtype_prop)
(rule ((= ?__e (LoopStart ?v63_inp ?v63_loop_id ?v63_slot_idx ?v63_iters ?v63_dtype))) ((set (dtype ?__e) ?v63_dtype)) :ruleset dtype_prop)
(rule ((= ?__e (LoopEnd ?v64_inp ?v64_loop_id ?v64_slot_idx ?v64_dtype))) ((set (dtype ?__e) ?v64_dtype)) :ruleset dtype_prop)
(rule ((= ?__e (Op (LoopInput ?v65_loop_id ?v65_stream_id ?v65_dtype) ?__inputs))) ((set (dtype ?__e) ?v65_dtype)) :ruleset dtype_prop)
(relation identical_inputs (IList))
; All four rules live in the `expr` ruleset, which the schedule
; saturates each iteration. Default-ruleset scheduling only runs
; each rule once per outer step, which is not enough to propagate
; `identical_inputs` through an N-element IList.
; Base: single-element list is trivially identical.
(rule ((= ?l (ICons ?x (INil))))
((identical_inputs ?l))
:ruleset expr
:name "identical_inputs base")
; Inductive: head equals next-head, and the tail starting at next-head is identical.
(rule ((= ?l (ICons ?x (ICons ?x ?tail)))
(identical_inputs (ICons ?x ?tail)))
((identical_inputs ?l))
:ruleset expr
:name "identical_inputs ind")
; LoopInput with an identical IList is equivalent to LoopInputStatic over a single copy.
(rule ((= ?e (Op (LoopInput ?id ?stream ?dt) (ICons ?x ?cont)))
(identical_inputs (ICons ?x ?cont)))
((let ?static (Op (LoopInputStatic ?id ?stream ?dt) (ICons ?x (INil))))
(union ?e ?static))
:ruleset expr
:name "LoopInput to LoopInputStatic")
; LoopInputStatic is equivalent to its single inner value — collapses the boundary
; wrapper for pattern-matching and extraction purposes.
(rule ((= ?e (Op (LoopInputStatic ?id ?stream ?dt) (ICons ?x (INil)))))
((union ?e ?x))
:ruleset expr
:name "LoopInputStatic inline")
(rule ((= ?__e (Op (LoopInputStatic ?v66_loop_id ?v66_stream_id ?v66_dtype) ?__inputs))) ((set (dtype ?__e) ?v66_dtype)) :ruleset dtype_prop)
(rule ((= ?__e (Op (LoopOutput ?v67_loop_id ?v67_stream_id ?v67_dtype) ?__inputs))) ((set (dtype ?__e) ?v67_dtype)) :ruleset dtype_prop)
(rule ((= ?__e (Op (LoopOutputSelect ?v68_loop_id ?v68_stream_id ?v68_iter ?v68_dtype) ?__inputs))) ((set (dtype ?__e) ?v68_dtype)) :ruleset dtype_prop)
(rule ((= ?__e (Op (Constant ?v69_value) ?__inputs))) ((set (dtype ?__e) (F32))) :ruleset dtype_prop)
(rule ((= ?__e (Op (Cast ?v70_size ?v70_dtype) ?__inputs))) ((set (dtype ?__e) ?v70_dtype)) :ruleset dtype_prop)
(rule ((= ?__e (Op (Iota ?v71_expr ?v71_range) ?__inputs))) ((set (dtype ?__e) (Int))) :ruleset dtype_prop)
(rule ((= ?__e (Op (Exp2 ?v72_shape ?v72_strides ?v72_out_strides) (ICons ?__first_inp ?__tail))) (= ?__dty (dtype ?__first_inp))) ((set (dtype ?__e) ?__dty)) :ruleset dtype_prop)
(rule ((= ?__e (Op (Log2 ?v73_shape ?v73_strides ?v73_out_strides) (ICons ?__first_inp ?__tail))) (= ?__dty (dtype ?__first_inp))) ((set (dtype ?__e) ?__dty)) :ruleset dtype_prop)
(rule ((= ?__e (Op (Sin ?v74_shape ?v74_strides ?v74_out_strides) (ICons ?__first_inp ?__tail))) (= ?__dty (dtype ?__first_inp))) ((set (dtype ?__e) ?__dty)) :ruleset dtype_prop)
(rule ((= ?__e (Op (Recip ?v75_shape ?v75_strides ?v75_out_strides) (ICons ?__first_inp ?__tail))) (= ?__dty (dtype ?__first_inp))) ((set (dtype ?__e) ?__dty)) :ruleset dtype_prop)
(rule ((= ?__e (Op (Sqrt ?v76_shape ?v76_strides ?v76_out_strides) (ICons ?__first_inp ?__tail))) (= ?__dty (dtype ?__first_inp))) ((set (dtype ?__e) ?__dty)) :ruleset dtype_prop)
(rule ((= ?__e (Op (Add ?v77_shape ?v77_a_strides ?v77_b_strides ?v77_out_strides) (ICons ?__first_inp ?__tail))) (= ?__dty (dtype ?__first_inp))) ((set (dtype ?__e) ?__dty)) :ruleset dtype_prop)
(rule
(
(= ?ls (LoopStart ?initial ?loop_id ?slot_idx (MNum 2) ?dt))
(= ?li (Op (LoopInput ?loop_id ?stream ?dt) (ICons ?s0 (ICons ?s1 (INil)))))
(= ?body (Op (Add ?sh ?as ?bs ?os) (ICons ?ls (ICons ?li (INil)))))
(= ?le (LoopEnd ?body ?loop_id ?slot_idx ?dt))
)
(
(let ?u0 (Op (Add ?sh ?as ?bs ?os) (ICons ?initial (ICons ?s0 (INil)))))
(let ?u1 (Op (Add ?sh ?as ?bs ?os) (ICons ?u0 (ICons ?s1 (INil)))))
(union ?le ?u1)
)
:ruleset expr
:name "unroll Add body trips=2 state=0"
)
(rule
(
(= ?ls (LoopStart ?initial ?loop_id ?slot_idx (MNum 2) ?dt))
(= ?li (Op (LoopInput ?loop_id ?stream ?dt) (ICons ?s0 (ICons ?s1 (INil)))))
(= ?body (Op (Add ?sh ?as ?bs ?os) (ICons ?li (ICons ?ls (INil)))))
(= ?le (LoopEnd ?body ?loop_id ?slot_idx ?dt))
)
(
(let ?u0 (Op (Add ?sh ?as ?bs ?os) (ICons ?s0 (ICons ?initial (INil)))))
(let ?u1 (Op (Add ?sh ?as ?bs ?os) (ICons ?s1 (ICons ?u0 (INil)))))
(union ?le ?u1)
)
:ruleset expr
:name "unroll Add body trips=2 state=1"
)
(rule
(
(= ?ls (LoopStart ?initial ?loop_id ?slot_idx (MNum 3) ?dt))
(= ?li (Op (LoopInput ?loop_id ?stream ?dt) (ICons ?s0 (ICons ?s1 (ICons ?s2 (INil))))))
(= ?body (Op (Add ?sh ?as ?bs ?os) (ICons ?ls (ICons ?li (INil)))))
(= ?le (LoopEnd ?body ?loop_id ?slot_idx ?dt))
)
(
(let ?u0 (Op (Add ?sh ?as ?bs ?os) (ICons ?initial (ICons ?s0 (INil)))))
(let ?u1 (Op (Add ?sh ?as ?bs ?os) (ICons ?u0 (ICons ?s1 (INil)))))
(let ?u2 (Op (Add ?sh ?as ?bs ?os) (ICons ?u1 (ICons ?s2 (INil)))))
(union ?le ?u2)
)
:ruleset expr
:name "unroll Add body trips=3 state=0"
)
(rule
(
(= ?ls (LoopStart ?initial ?loop_id ?slot_idx (MNum 3) ?dt))
(= ?li (Op (LoopInput ?loop_id ?stream ?dt) (ICons ?s0 (ICons ?s1 (ICons ?s2 (INil))))))
(= ?body (Op (Add ?sh ?as ?bs ?os) (ICons ?li (ICons ?ls (INil)))))
(= ?le (LoopEnd ?body ?loop_id ?slot_idx ?dt))
)
(
(let ?u0 (Op (Add ?sh ?as ?bs ?os) (ICons ?s0 (ICons ?initial (INil)))))
(let ?u1 (Op (Add ?sh ?as ?bs ?os) (ICons ?s1 (ICons ?u0 (INil)))))
(let ?u2 (Op (Add ?sh ?as ?bs ?os) (ICons ?s2 (ICons ?u1 (INil)))))
(union ?le ?u2)
)
:ruleset expr
:name "unroll Add body trips=3 state=1"
)
(rule
(
(= ?ls (LoopStart ?initial ?loop_id ?slot_idx (MNum 4) ?dt))
(= ?li (Op (LoopInput ?loop_id ?stream ?dt) (ICons ?s0 (ICons ?s1 (ICons ?s2 (ICons ?s3 (INil)))))))
(= ?body (Op (Add ?sh ?as ?bs ?os) (ICons ?ls (ICons ?li (INil)))))
(= ?le (LoopEnd ?body ?loop_id ?slot_idx ?dt))
)
(
(let ?u0 (Op (Add ?sh ?as ?bs ?os) (ICons ?initial (ICons ?s0 (INil)))))
(let ?u1 (Op (Add ?sh ?as ?bs ?os) (ICons ?u0 (ICons ?s1 (INil)))))
(let ?u2 (Op (Add ?sh ?as ?bs ?os) (ICons ?u1 (ICons ?s2 (INil)))))
(let ?u3 (Op (Add ?sh ?as ?bs ?os) (ICons ?u2 (ICons ?s3 (INil)))))
(union ?le ?u3)
)
:ruleset expr
:name "unroll Add body trips=4 state=0"
)
(rule
(
(= ?ls (LoopStart ?initial ?loop_id ?slot_idx (MNum 4) ?dt))
(= ?li (Op (LoopInput ?loop_id ?stream ?dt) (ICons ?s0 (ICons ?s1 (ICons ?s2 (ICons ?s3 (INil)))))))
(= ?body (Op (Add ?sh ?as ?bs ?os) (ICons ?li (ICons ?ls (INil)))))
(= ?le (LoopEnd ?body ?loop_id ?slot_idx ?dt))
)
(
(let ?u0 (Op (Add ?sh ?as ?bs ?os) (ICons ?s0 (ICons ?initial (INil)))))
(let ?u1 (Op (Add ?sh ?as ?bs ?os) (ICons ?s1 (ICons ?u0 (INil)))))
(let ?u2 (Op (Add ?sh ?as ?bs ?os) (ICons ?s2 (ICons ?u1 (INil)))))
(let ?u3 (Op (Add ?sh ?as ?bs ?os) (ICons ?s3 (ICons ?u2 (INil)))))
(union ?le ?u3)
)
:ruleset expr
:name "unroll Add body trips=4 state=1"
)
(rule ((= ?__e (Op (Mul ?v78_shape ?v78_a_strides ?v78_b_strides ?v78_out_strides) (ICons ?__first_inp ?__tail))) (= ?__dty (dtype ?__first_inp))) ((set (dtype ?__e) ?__dty)) :ruleset dtype_prop)
(rule
(
(= ?ls (LoopStart ?initial ?loop_id ?slot_idx (MNum 2) ?dt))
(= ?li (Op (LoopInput ?loop_id ?stream ?dt) (ICons ?s0 (ICons ?s1 (INil)))))
(= ?body (Op (Mul ?sh ?as ?bs ?os) (ICons ?ls (ICons ?li (INil)))))
(= ?le (LoopEnd ?body ?loop_id ?slot_idx ?dt))
)
(
(let ?u0 (Op (Mul ?sh ?as ?bs ?os) (ICons ?initial (ICons ?s0 (INil)))))
(let ?u1 (Op (Mul ?sh ?as ?bs ?os) (ICons ?u0 (ICons ?s1 (INil)))))
(union ?le ?u1)
)
:ruleset expr
:name "unroll Mul body trips=2 state=0"
)
(rule
(
(= ?ls (LoopStart ?initial ?loop_id ?slot_idx (MNum 2) ?dt))
(= ?li (Op (LoopInput ?loop_id ?stream ?dt) (ICons ?s0 (ICons ?s1 (INil)))))
(= ?body (Op (Mul ?sh ?as ?bs ?os) (ICons ?li (ICons ?ls (INil)))))
(= ?le (LoopEnd ?body ?loop_id ?slot_idx ?dt))
)
(
(let ?u0 (Op (Mul ?sh ?as ?bs ?os) (ICons ?s0 (ICons ?initial (INil)))))
(let ?u1 (Op (Mul ?sh ?as ?bs ?os) (ICons ?s1 (ICons ?u0 (INil)))))
(union ?le ?u1)
)
:ruleset expr
:name "unroll Mul body trips=2 state=1"
)
(rule
(
(= ?ls (LoopStart ?initial ?loop_id ?slot_idx (MNum 3) ?dt))
(= ?li (Op (LoopInput ?loop_id ?stream ?dt) (ICons ?s0 (ICons ?s1 (ICons ?s2 (INil))))))
(= ?body (Op (Mul ?sh ?as ?bs ?os) (ICons ?ls (ICons ?li (INil)))))
(= ?le (LoopEnd ?body ?loop_id ?slot_idx ?dt))
)
(
(let ?u0 (Op (Mul ?sh ?as ?bs ?os) (ICons ?initial (ICons ?s0 (INil)))))
(let ?u1 (Op (Mul ?sh ?as ?bs ?os) (ICons ?u0 (ICons ?s1 (INil)))))
(let ?u2 (Op (Mul ?sh ?as ?bs ?os) (ICons ?u1 (ICons ?s2 (INil)))))
(union ?le ?u2)
)
:ruleset expr
:name "unroll Mul body trips=3 state=0"
)
(rule
(
(= ?ls (LoopStart ?initial ?loop_id ?slot_idx (MNum 3) ?dt))
(= ?li (Op (LoopInput ?loop_id ?stream ?dt) (ICons ?s0 (ICons ?s1 (ICons ?s2 (INil))))))
(= ?body (Op (Mul ?sh ?as ?bs ?os) (ICons ?li (ICons ?ls (INil)))))
(= ?le (LoopEnd ?body ?loop_id ?slot_idx ?dt))
)
(
(let ?u0 (Op (Mul ?sh ?as ?bs ?os) (ICons ?s0 (ICons ?initial (INil)))))
(let ?u1 (Op (Mul ?sh ?as ?bs ?os) (ICons ?s1 (ICons ?u0 (INil)))))
(let ?u2 (Op (Mul ?sh ?as ?bs ?os) (ICons ?s2 (ICons ?u1 (INil)))))
(union ?le ?u2)
)
:ruleset expr
:name "unroll Mul body trips=3 state=1"
)
(rule
(
(= ?ls (LoopStart ?initial ?loop_id ?slot_idx (MNum 4) ?dt))
(= ?li (Op (LoopInput ?loop_id ?stream ?dt) (ICons ?s0 (ICons ?s1 (ICons ?s2 (ICons ?s3 (INil)))))))
(= ?body (Op (Mul ?sh ?as ?bs ?os) (ICons ?ls (ICons ?li (INil)))))
(= ?le (LoopEnd ?body ?loop_id ?slot_idx ?dt))
)
(
(let ?u0 (Op (Mul ?sh ?as ?bs ?os) (ICons ?initial (ICons ?s0 (INil)))))
(let ?u1 (Op (Mul ?sh ?as ?bs ?os) (ICons ?u0 (ICons ?s1 (INil)))))
(let ?u2 (Op (Mul ?sh ?as ?bs ?os) (ICons ?u1 (ICons ?s2 (INil)))))
(let ?u3 (Op (Mul ?sh ?as ?bs ?os) (ICons ?u2 (ICons ?s3 (INil)))))
(union ?le ?u3)
)
:ruleset expr
:name "unroll Mul body trips=4 state=0"
)
(rule
(
(= ?ls (LoopStart ?initial ?loop_id ?slot_idx (MNum 4) ?dt))
(= ?li (Op (LoopInput ?loop_id ?stream ?dt) (ICons ?s0 (ICons ?s1 (ICons ?s2 (ICons ?s3 (INil)))))))
(= ?body (Op (Mul ?sh ?as ?bs ?os) (ICons ?li (ICons ?ls (INil)))))
(= ?le (LoopEnd ?body ?loop_id ?slot_idx ?dt))
)
(
(let ?u0 (Op (Mul ?sh ?as ?bs ?os) (ICons ?s0 (ICons ?initial (INil)))))
(let ?u1 (Op (Mul ?sh ?as ?bs ?os) (ICons ?s1 (ICons ?u0 (INil)))))
(let ?u2 (Op (Mul ?sh ?as ?bs ?os) (ICons ?s2 (ICons ?u1 (INil)))))
(let ?u3 (Op (Mul ?sh ?as ?bs ?os) (ICons ?s3 (ICons ?u2 (INil)))))
(union ?le ?u3)
)
:ruleset expr
:name "unroll Mul body trips=4 state=1"
)
(rule ((= ?__e (Op (Mod ?v79_shape ?v79_a_strides ?v79_b_strides ?v79_out_strides) (ICons ?__first_inp ?__tail))) (= ?__dty (dtype ?__first_inp))) ((set (dtype ?__e) ?__dty)) :ruleset dtype_prop)
(rule
(
(= ?ls (LoopStart ?initial ?loop_id ?slot_idx (MNum 2) ?dt))
(= ?li (Op (LoopInput ?loop_id ?stream ?dt) (ICons ?s0 (ICons ?s1 (INil)))))
(= ?body (Op (Mod ?sh ?as ?bs ?os) (ICons ?ls (ICons ?li (INil)))))
(= ?le (LoopEnd ?body ?loop_id ?slot_idx ?dt))
)
(
(let ?u0 (Op (Mod ?sh ?as ?bs ?os) (ICons ?initial (ICons ?s0 (INil)))))
(let ?u1 (Op (Mod ?sh ?as ?bs ?os) (ICons ?u0 (ICons ?s1 (INil)))))
(union ?le ?u1)
)
:ruleset expr
:name "unroll Mod body trips=2 state=0"
)
(rule
(
(= ?ls (LoopStart ?initial ?loop_id ?slot_idx (MNum 2) ?dt))
(= ?li (Op (LoopInput ?loop_id ?stream ?dt) (ICons ?s0 (ICons ?s1 (INil)))))
(= ?body (Op (Mod ?sh ?as ?bs ?os) (ICons ?li (ICons ?ls (INil)))))
(= ?le (LoopEnd ?body ?loop_id ?slot_idx ?dt))
)
(
(let ?u0 (Op (Mod ?sh ?as ?bs ?os) (ICons ?s0 (ICons ?initial (INil)))))
(let ?u1 (Op (Mod ?sh ?as ?bs ?os) (ICons ?s1 (ICons ?u0 (INil)))))
(union ?le ?u1)
)
:ruleset expr
:name "unroll Mod body trips=2 state=1"
)
(rule
(
(= ?ls (LoopStart ?initial ?loop_id ?slot_idx (MNum 3) ?dt))
(= ?li (Op (LoopInput ?loop_id ?stream ?dt) (ICons ?s0 (ICons ?s1 (ICons ?s2 (INil))))))
(= ?body (Op (Mod ?sh ?as ?bs ?os) (ICons ?ls (ICons ?li (INil)))))
(= ?le (LoopEnd ?body ?loop_id ?slot_idx ?dt))
)
(
(let ?u0 (Op (Mod ?sh ?as ?bs ?os) (ICons ?initial (ICons ?s0 (INil)))))
(let ?u1 (Op (Mod ?sh ?as ?bs ?os) (ICons ?u0 (ICons ?s1 (INil)))))
(let ?u2 (Op (Mod ?sh ?as ?bs ?os) (ICons ?u1 (ICons ?s2 (INil)))))
(union ?le ?u2)
)
:ruleset expr
:name "unroll Mod body trips=3 state=0"
)
(rule
(
(= ?ls (LoopStart ?initial ?loop_id ?slot_idx (MNum 3) ?dt))
(= ?li (Op (LoopInput ?loop_id ?stream ?dt) (ICons ?s0 (ICons ?s1 (ICons ?s2 (INil))))))
(= ?body (Op (Mod ?sh ?as ?bs ?os) (ICons ?li (ICons ?ls (INil)))))
(= ?le (LoopEnd ?body ?loop_id ?slot_idx ?dt))
)
(
(let ?u0 (Op (Mod ?sh ?as ?bs ?os) (ICons ?s0 (ICons ?initial (INil)))))
(let ?u1 (Op (Mod ?sh ?as ?bs ?os) (ICons ?s1 (ICons ?u0 (INil)))))
(let ?u2 (Op (Mod ?sh ?as ?bs ?os) (ICons ?s2 (ICons ?u1 (INil)))))
(union ?le ?u2)
)
:ruleset expr
:name "unroll Mod body trips=3 state=1"
)
(rule
(
(= ?ls (LoopStart ?initial ?loop_id ?slot_idx (MNum 4) ?dt))
(= ?li (Op (LoopInput ?loop_id ?stream ?dt) (ICons ?s0 (ICons ?s1 (ICons ?s2 (ICons ?s3 (INil)))))))
(= ?body (Op (Mod ?sh ?as ?bs ?os) (ICons ?ls (ICons ?li (INil)))))
(= ?le (LoopEnd ?body ?loop_id ?slot_idx ?dt))
)
(
(let ?u0 (Op (Mod ?sh ?as ?bs ?os) (ICons ?initial (ICons ?s0 (INil)))))
(let ?u1 (Op (Mod ?sh ?as ?bs ?os) (ICons ?u0 (ICons ?s1 (INil)))))
(let ?u2 (Op (Mod ?sh ?as ?bs ?os) (ICons ?u1 (ICons ?s2 (INil)))))
(let ?u3 (Op (Mod ?sh ?as ?bs ?os) (ICons ?u2 (ICons ?s3 (INil)))))
(union ?le ?u3)
)
:ruleset expr
:name "unroll Mod body trips=4 state=0"
)
(rule
(
(= ?ls (LoopStart ?initial ?loop_id ?slot_idx (MNum 4) ?dt))
(= ?li (Op (LoopInput ?loop_id ?stream ?dt) (ICons ?s0 (ICons ?s1 (ICons ?s2 (ICons ?s3 (INil)))))))
(= ?body (Op (Mod ?sh ?as ?bs ?os) (ICons ?li (ICons ?ls (INil)))))
(= ?le (LoopEnd ?body ?loop_id ?slot_idx ?dt))
)
(
(let ?u0 (Op (Mod ?sh ?as ?bs ?os) (ICons ?s0 (ICons ?initial (INil)))))
(let ?u1 (Op (Mod ?sh ?as ?bs ?os) (ICons ?s1 (ICons ?u0 (INil)))))
(let ?u2 (Op (Mod ?sh ?as ?bs ?os) (ICons ?s2 (ICons ?u1 (INil)))))
(let ?u3 (Op (Mod ?sh ?as ?bs ?os) (ICons ?s3 (ICons ?u2 (INil)))))
(union ?le ?u3)
)
:ruleset expr
:name "unroll Mod body trips=4 state=1"
)
(rule ((= ?__e (Op (LessThan ?v80_shape ?v80_a_strides ?v80_b_strides ?v80_out_strides) ?__inputs))) ((set (dtype ?__e) (Bool))) :ruleset dtype_prop)
(rule
(
(= ?ls (LoopStart ?initial ?loop_id ?slot_idx (MNum 2) ?dt))
(= ?li (Op (LoopInput ?loop_id ?stream ?dt) (ICons ?s0 (ICons ?s1 (INil)))))
(= ?body (Op (LessThan ?sh ?as ?bs ?os) (ICons ?ls (ICons ?li (INil)))))
(= ?le (LoopEnd ?body ?loop_id ?slot_idx ?dt))
)
(
(let ?u0 (Op (LessThan ?sh ?as ?bs ?os) (ICons ?initial (ICons ?s0 (INil)))))
(let ?u1 (Op (LessThan ?sh ?as ?bs ?os) (ICons ?u0 (ICons ?s1 (INil)))))
(union ?le ?u1)
)
:ruleset expr
:name "unroll LessThan body trips=2 state=0"
)
(rule
(
(= ?ls (LoopStart ?initial ?loop_id ?slot_idx (MNum 2) ?dt))
(= ?li (Op (LoopInput ?loop_id ?stream ?dt) (ICons ?s0 (ICons ?s1 (INil)))))
(= ?body (Op (LessThan ?sh ?as ?bs ?os) (ICons ?li (ICons ?ls (INil)))))
(= ?le (LoopEnd ?body ?loop_id ?slot_idx ?dt))
)
(
(let ?u0 (Op (LessThan ?sh ?as ?bs ?os) (ICons ?s0 (ICons ?initial (INil)))))
(let ?u1 (Op (LessThan ?sh ?as ?bs ?os) (ICons ?s1 (ICons ?u0 (INil)))))
(union ?le ?u1)
)
:ruleset expr
:name "unroll LessThan body trips=2 state=1"
)
(rule
(
(= ?ls (LoopStart ?initial ?loop_id ?slot_idx (MNum 3) ?dt))
(= ?li (Op (LoopInput ?loop_id ?stream ?dt) (ICons ?s0 (ICons ?s1 (ICons ?s2 (INil))))))
(= ?body (Op (LessThan ?sh ?as ?bs ?os) (ICons ?ls (ICons ?li (INil)))))
(= ?le (LoopEnd ?body ?loop_id ?slot_idx ?dt))
)
(
(let ?u0 (Op (LessThan ?sh ?as ?bs ?os) (ICons ?initial (ICons ?s0 (INil)))))
(let ?u1 (Op (LessThan ?sh ?as ?bs ?os) (ICons ?u0 (ICons ?s1 (INil)))))
(let ?u2 (Op (LessThan ?sh ?as ?bs ?os) (ICons ?u1 (ICons ?s2 (INil)))))
(union ?le ?u2)
)
:ruleset expr
:name "unroll LessThan body trips=3 state=0"
)
(rule
(
(= ?ls (LoopStart ?initial ?loop_id ?slot_idx (MNum 3) ?dt))
(= ?li (Op (LoopInput ?loop_id ?stream ?dt) (ICons ?s0 (ICons ?s1 (ICons ?s2 (INil))))))
(= ?body (Op (LessThan ?sh ?as ?bs ?os) (ICons ?li (ICons ?ls (INil)))))
(= ?le (LoopEnd ?body ?loop_id ?slot_idx ?dt))
)
(
(let ?u0 (Op (LessThan ?sh ?as ?bs ?os) (ICons ?s0 (ICons ?initial (INil)))))
(let ?u1 (Op (LessThan ?sh ?as ?bs ?os) (ICons ?s1 (ICons ?u0 (INil)))))
(let ?u2 (Op (LessThan ?sh ?as ?bs ?os) (ICons ?s2 (ICons ?u1 (INil)))))
(union ?le ?u2)
)
:ruleset expr
:name "unroll LessThan body trips=3 state=1"
)
(rule
(
(= ?ls (LoopStart ?initial ?loop_id ?slot_idx (MNum 4) ?dt))
(= ?li (Op (LoopInput ?loop_id ?stream ?dt) (ICons ?s0 (ICons ?s1 (ICons ?s2 (ICons ?s3 (INil)))))))
(= ?body (Op (LessThan ?sh ?as ?bs ?os) (ICons ?ls (ICons ?li (INil)))))
(= ?le (LoopEnd ?body ?loop_id ?slot_idx ?dt))
)
(
(let ?u0 (Op (LessThan ?sh ?as ?bs ?os) (ICons ?initial (ICons ?s0 (INil)))))
(let ?u1 (Op (LessThan ?sh ?as ?bs ?os) (ICons ?u0 (ICons ?s1 (INil)))))
(let ?u2 (Op (LessThan ?sh ?as ?bs ?os) (ICons ?u1 (ICons ?s2 (INil)))))
(let ?u3 (Op (LessThan ?sh ?as ?bs ?os) (ICons ?u2 (ICons ?s3 (INil)))))
(union ?le ?u3)
)
:ruleset expr
:name "unroll LessThan body trips=4 state=0"
)
(rule
(
(= ?ls (LoopStart ?initial ?loop_id ?slot_idx (MNum 4) ?dt))
(= ?li (Op (LoopInput ?loop_id ?stream ?dt) (ICons ?s0 (ICons ?s1 (ICons ?s2 (ICons ?s3 (INil)))))))
(= ?body (Op (LessThan ?sh ?as ?bs ?os) (ICons ?li (ICons ?ls (INil)))))
(= ?le (LoopEnd ?body ?loop_id ?slot_idx ?dt))
)
(
(let ?u0 (Op (LessThan ?sh ?as ?bs ?os) (ICons ?s0 (ICons ?initial (INil)))))
(let ?u1 (Op (LessThan ?sh ?as ?bs ?os) (ICons ?s1 (ICons ?u0 (INil)))))
(let ?u2 (Op (LessThan ?sh ?as ?bs ?os) (ICons ?s2 (ICons ?u1 (INil)))))
(let ?u3 (Op (LessThan ?sh ?as ?bs ?os) (ICons ?s3 (ICons ?u2 (INil)))))
(union ?le ?u3)
)
:ruleset expr
:name "unroll LessThan body trips=4 state=1"
)
(rule ((= ?__e (Op (Gather ?v81_index_shape ?v81_index_strides ?v81_data_shape ?v81_data_strides) (ICons ?__indexes (ICons ?__data ?__tail)))) (= ?__dty (dtype ?__data))) ((set (dtype ?__e) ?__dty)) :ruleset dtype_prop)
(rule ((= ?__e (Op (Scatter ?v82_dest_shape ?v82_dest_strides ?v82_index_shape ?v82_index_strides ?v82_src_strides) (ICons ?__dest (ICons ?__indexes (ICons ?__src ?__tail))))) (= ?__dty (dtype ?__src))) ((set (dtype ?__e) ?__dty)) :ruleset dtype_prop)
(rule ((= ?__e (Op (Sum ?v83_shape ?v83_iters ?v83_strides ?v83_iter_stride ?v83_out_strides) (ICons ?__first_inp ?__tail))) (= ?__dty (dtype ?__first_inp))) ((set (dtype ?__e) ?__dty)) :ruleset dtype_prop)
; Squeeze: remove outermost dim=1 from Mul+Sum
; [1, d1, d2, ...] → [d1, d2, ...] via direct union
;
; When the outermost dimension of the Sum output is 1, it can always be
; removed regardless of strides (index 0 contributes 0*stride = 0 to
; every address). This enables downstream 2D matmul rules (cuBLAS) to match.
; The rule fires recursively: 4D → 3D → 2D.
;
; Direct union is safe because the total element count is preserved
; (1*N = N) and downstream ops use their own strides to index into the
; flat buffer, which has the same layout regardless of the leading dim=1.
(rule
(
; Match Mul + Sum pattern (Op wrapper syntax)
(= ?mul (Op (Mul ?mul_shape ?a_stride ?b_stride ?mul_out_stride) (ICons ?a (ICons ?b (INil)))))
(= ?sum (Op (Sum ?sum_shape ?k ?sum_in_stride ?k_stride ?sum_out_stride) (ICons ?mul (INil))))
; Sum output shape starts with 1
(= ?sum_shape (ECons (MNum 1) ?rest_sum_shape))
; Result must have >= 2 dims (don't collapse below 2D)
(!= ?rest_sum_shape (ENil))
(= ?rest_sum_len (len ?rest_sum_shape))
(>= ?rest_sum_len 2)
; Mul shape also starts with 1
(= ?mul_shape (ECons (MNum 1) ?rest_mul_shape))
; Destructure all stride lists to drop first element
(= ?a_stride (ECons ?as0 ?rest_as))
(= ?b_stride (ECons ?bs0 ?rest_bs))
(= ?mul_out_stride (ECons ?mos0 ?rest_mos))
(= ?sum_in_stride (ECons ?sis0 ?rest_sis))
(= ?sum_out_stride (ECons ?sos0 ?rest_sos))
; Get dtype
(= ?dt (dtype ?a))
)
(
; Create collapsed Mul (drop outermost dim from all lists)
(let ?new_mul (Op (Mul ?rest_mul_shape
?rest_as
?rest_bs
?rest_mos)
(ICons ?a (ICons ?b (INil)))))
; Create collapsed Sum
(let ?new_sum (Op (Sum ?rest_sum_shape ?k
?rest_sis ?k_stride ?rest_sos)
(ICons ?new_mul (INil))))
; Direct union — no wrapper needed
(union ?sum ?new_sum)
; Propagate dtype
(set (dtype ?new_mul) ?dt)
(set (dtype ?new_sum) ?dt)
)
:ruleset matmul_flatten
:name "batch-collapse squeeze dim=1"
)
; Batch-merge: collapse outermost two dims when A is contiguous, B is broadcast
; [d0, d1, ...] → [d0*d1, ...] via direct union
;
; Preconditions:
; - A's outermost stride is contiguous: a_stride[0] = a_stride[1] * dim[1]
; - B is broadcast on both leading dims: b_stride[0] = 0, b_stride[1] = 0
; - Output strides are contiguous on the leading dims
;
; Direct union is safe because the total element count is preserved
; (d0*d1 = d0*d1) and the flat buffer layout is identical.
(rule
(
(= ?mul (Op (Mul ?mul_shape ?a_stride ?b_stride ?mul_out_stride) (ICons ?a (ICons ?b (INil)))))
(= ?sum (Op (Sum ?sum_shape ?k ?sum_in_stride ?k_stride ?sum_out_stride) (ICons ?mul (INil))))
; Sum shape has >= 3 dims (result will have >= 2)
(= ?sum_shape (ECons ?sd0 (ECons ?sd1 ?rest_sum_shape)))
(!= ?rest_sum_shape (ENil))
; Mul shape
(= ?mul_shape (ECons ?d0 (ECons ?d1 ?rest_mul_shape)))
; A strides: a_s0 = a_s1 * d1 (contiguous)
(= ?a_stride (ECons ?as0 (ECons ?as1 ?rest_as)))
(= ?as0 (MMul ?as1 ?d1))
; B strides: broadcast on both leading dims
(= ?b_stride (ECons (MNum 0) (ECons (MNum 0) ?rest_bs)))
; Mul output strides: contiguous (os0 = os1 * d1)
(= ?mul_out_stride (ECons ?mos0 (ECons ?mos1 ?rest_mos)))
(= ?mos0 (MMul ?mos1 ?d1))
; Sum input strides: contiguous leading
(= ?sum_in_stride (ECons ?sis0 (ECons ?sis1 ?rest_sis)))
(= ?sis0 (MMul ?sis1 ?sd1))
; Sum output strides: contiguous leading
(= ?sum_out_stride (ECons ?sos0 (ECons ?sos1 ?rest_sos)))
(= ?sos0 (MMul ?sos1 ?sd1))
(= ?dt (dtype ?a))
)
(
; Merged dimensions
(let ?new_d (MMul ?d0 ?d1))
(let ?new_sd (MMul ?sd0 ?sd1))
; Collapsed Mul: merged leading dim, A uses as1, B uses 0
(let ?new_mul (Op (Mul (ECons ?new_d ?rest_mul_shape)
(ECons ?as1 ?rest_as)
(ECons (MNum 0) ?rest_bs)
(ECons ?mos1 ?rest_mos))
(ICons ?a (ICons ?b (INil)))))
; Collapsed Sum
(let ?new_sum (Op (Sum (ECons ?new_sd ?rest_sum_shape)
?k
(ECons ?sis1 ?rest_sis)
?k_stride
(ECons ?sos1 ?rest_sos))
(ICons ?new_mul (INil))))
; Direct union
(union ?sum ?new_sum)
(set (dtype ?new_mul) ?dt)
(set (dtype ?new_sum) ?dt)
)
:ruleset matmul_flatten
:name "batch-collapse merge A-contiguous B-broadcast"
)
; Batch-merge: collapse outermost two dims when B is contiguous, A is broadcast
; [d0, d1, ...] → [d0*d1, ...] via direct union
;
; Symmetric case of batch_merge_a_contig:
; - B's outermost stride is contiguous: b_stride[0] = b_stride[1] * dim[1]
; - A is broadcast on both leading dims: a_stride[0] = 0, a_stride[1] = 0
; - Output strides are contiguous on the leading dims
(rule
(
(= ?mul (Op (Mul ?mul_shape ?a_stride ?b_stride ?mul_out_stride) (ICons ?a (ICons ?b (INil)))))
(= ?sum (Op (Sum ?sum_shape ?k ?sum_in_stride ?k_stride ?sum_out_stride) (ICons ?mul (INil))))
; Sum shape has >= 3 dims
(= ?sum_shape (ECons ?sd0 (ECons ?sd1 ?rest_sum_shape)))
(!= ?rest_sum_shape (ENil))
; Mul shape
(= ?mul_shape (ECons ?d0 (ECons ?d1 ?rest_mul_shape)))
; A strides: broadcast on both leading dims
(= ?a_stride (ECons (MNum 0) (ECons (MNum 0) ?rest_as)))
; B strides: b_s0 = b_s1 * d1 (contiguous)
(= ?b_stride (ECons ?bs0 (ECons ?bs1 ?rest_bs)))
(= ?bs0 (MMul ?bs1 ?d1))
; Mul output strides: contiguous (os0 = os1 * d1)
(= ?mul_out_stride (ECons ?mos0 (ECons ?mos1 ?rest_mos)))
(= ?mos0 (MMul ?mos1 ?d1))
; Sum input strides: contiguous leading
(= ?sum_in_stride (ECons ?sis0 (ECons ?sis1 ?rest_sis)))
(= ?sis0 (MMul ?sis1 ?sd1))
; Sum output strides: contiguous leading
(= ?sum_out_stride (ECons ?sos0 (ECons ?sos1 ?rest_sos)))
(= ?sos0 (MMul ?sos1 ?sd1))
(= ?dt (dtype ?a))
)
(
; Merged dimensions
(let ?new_d (MMul ?d0 ?d1))
(let ?new_sd (MMul ?sd0 ?sd1))
; Collapsed Mul: merged leading dim, A uses 0, B uses bs1
(let ?new_mul (Op (Mul (ECons ?new_d ?rest_mul_shape)
(ECons (MNum 0) ?rest_as)
(ECons ?bs1 ?rest_bs)
(ECons ?mos1 ?rest_mos))
(ICons ?a (ICons ?b (INil)))))
; Collapsed Sum
(let ?new_sum (Op (Sum (ECons ?new_sd ?rest_sum_shape)
?k
(ECons ?sis1 ?rest_sis)
?k_stride
(ECons ?sos1 ?rest_sos))
(ICons ?new_mul (INil))))
; Direct union
(union ?sum ?new_sum)
(set (dtype ?new_mul) ?dt)
(set (dtype ?new_sum) ?dt)
)
:ruleset matmul_flatten
:name "batch-collapse merge B-contiguous A-broadcast"
)
(rule ((= ?__e (Op (Max ?v84_shape ?v84_iters ?v84_strides ?v84_iter_stride ?v84_out_strides) (ICons ?__first_inp ?__tail))) (= ?__dty (dtype ?__first_inp))) ((set (dtype ?__e) ?__dty)) :ruleset dtype_prop)
(rule ((= ?__e (Op (Softmax ?v85_shape ?v85_in_strides ?v85_out_strides ?v85_reduce_dim ?v85_reduce_stride) (ICons ?__first_inp ?__tail))) (= ?__dty (dtype ?__first_inp))) ((set (dtype ?__e) ?__dty)) :ruleset dtype_prop)
(ruleset cuda_memory_analysis)
(relation cuda_output_bytes (OpKind Expression))
(relation cuda_local_memory (IR Expression))
(rule ((= ?node (Input ?id ?label ?dtype)))
((cuda_local_memory ?node (MNum 0)))
:ruleset cuda_memory_analysis
:name "cuda-memory-input")
(rule ((= ?node (Output ?inp ?id)))
((cuda_local_memory ?node (MNum 0)))
:ruleset cuda_memory_analysis
:name "cuda-memory-output")
(rule ((= ?node (OutputJoin ?a ?b)))
((cuda_local_memory ?node (MNum 0)))
:ruleset cuda_memory_analysis
:name "cuda-memory-output-join")
(rule ((= ?node (Op ?kind ?inputs))
(cuda_output_bytes ?kind ?bytes))
((cuda_local_memory ?node ?bytes))
:ruleset cuda_memory_analysis
:name "cuda-memory-op-local")
(rule
((= ?kind (FusedAdd ?shape ?a_strides ?b_strides ?out_strides ?dtype)))
((cuda_output_bytes ?kind (MNum 0)))
:ruleset cuda_memory_analysis
:name "cuda-memory-FusedAdd-zero"
)
(rule
((= ?kind (FusedExp ?shape ?strides ?out_strides ?dtype)))
((cuda_output_bytes ?kind (MNum 0)))
:ruleset cuda_memory_analysis
:name "cuda-memory-FusedExp-zero"
)
(rule
((= ?kind (FusedExp2 ?shape ?strides ?out_strides ?dtype)))
((cuda_output_bytes ?kind (MNum 0)))
:ruleset cuda_memory_analysis
:name "cuda-memory-FusedExp2-zero"
)
(rule
((= ?kind (FusedLog2 ?shape ?strides ?out_strides ?dtype)))
((cuda_output_bytes ?kind (MNum 0)))
:ruleset cuda_memory_analysis
:name "cuda-memory-FusedLog2-zero"
)
(rule
((= ?kind (FusedMul ?shape ?a_strides ?b_strides ?out_strides ?dtype)))
((cuda_output_bytes ?kind (MNum 0)))
:ruleset cuda_memory_analysis
:name "cuda-memory-FusedMul-zero"
)
(rule
((= ?kind (FusedRecip ?shape ?strides ?out_strides ?dtype)))
((cuda_output_bytes ?kind (MNum 0)))
:ruleset cuda_memory_analysis
:name "cuda-memory-FusedRecip-zero"
)
(rule
((= ?kind (FusedSin ?shape ?strides ?out_strides ?dtype)))
((cuda_output_bytes ?kind (MNum 0)))
:ruleset cuda_memory_analysis
:name "cuda-memory-FusedSin-zero"
)
(rule
((= ?kind (FusedSqrt ?shape ?strides ?out_strides ?dtype)))
((cuda_output_bytes ?kind (MNum 0)))
:ruleset cuda_memory_analysis
:name "cuda-memory-FusedSqrt-zero"
)
(rule
((= ?kind (FusionEnd ?shape ?strides (F32)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 32)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-FusionEnd-F32-bytes"
)
(rule
((= ?kind (FusionEnd ?shape ?strides (F16)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 16)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-FusionEnd-F16-bytes"
)
(rule
((= ?kind (FusionEnd ?shape ?strides (Bf16)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 16)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-FusionEnd-Bf16-bytes"
)
(rule
((= ?kind (FusionEnd ?shape ?strides (Int)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 32)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-FusionEnd-Int-bytes"
)
(rule
((= ?kind (FusionEnd ?shape ?strides (Bool)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 8)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-FusionEnd-Bool-bytes"
)
(rule
((= ?kind (FusionEnd ?shape ?strides (I4)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 4)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-FusionEnd-I4-bytes"
)
(rule
((= ?kind (FusionEnd ?shape ?strides (TF32)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 19)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-FusionEnd-TF32-bytes"
)
(rule
((= ?kind (FusionStart ?shape ?strides ?dtype)))
((cuda_output_bytes ?kind (MNum 0)))
:ruleset cuda_memory_analysis
:name "cuda-memory-FusionStart-zero"
)
(rule
((= ?kind (GLUMoE ?gu_io ?dn_io ?gu_matmul_k ?dn_matmul_k ?output_k ?gu_within_range ?dn_within_range ?mode)))
((cuda_output_bytes ?kind (MMul (MMul (MVar "s") ?gu_matmul_k) (MNum 4))))
:ruleset cuda_memory_analysis
:name "cuda-memory-GLUMoE-f32-glumoe"
)
(rule
((= ?kind (KernelAdd ?shape ?a_strides ?b_strides ?out_strides (F32)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 32)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelAdd-F32-bytes"
)
(rule
((= ?kind (KernelAdd ?shape ?a_strides ?b_strides ?out_strides (F16)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 16)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelAdd-F16-bytes"
)
(rule
((= ?kind (KernelAdd ?shape ?a_strides ?b_strides ?out_strides (Bf16)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 16)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelAdd-Bf16-bytes"
)
(rule
((= ?kind (KernelAdd ?shape ?a_strides ?b_strides ?out_strides (Int)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 32)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelAdd-Int-bytes"
)
(rule
((= ?kind (KernelAdd ?shape ?a_strides ?b_strides ?out_strides (Bool)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 8)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelAdd-Bool-bytes"
)
(rule
((= ?kind (KernelAdd ?shape ?a_strides ?b_strides ?out_strides (I4)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 4)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelAdd-I4-bytes"
)
(rule
((= ?kind (KernelAdd ?shape ?a_strides ?b_strides ?out_strides (TF32)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 19)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelAdd-TF32-bytes"
)
(rule
((= ?kind (KernelBatchMatMul ?out_shape ?k_dim ?a_stride ?a_k_stride ?b_stride ?b_k_stride ?out_strides (F32)))
(= ?__cuda_elems (n_elements ?out_shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 32)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelBatchMatMul-F32-bytes"
)
(rule
((= ?kind (KernelBatchMatMul ?out_shape ?k_dim ?a_stride ?a_k_stride ?b_stride ?b_k_stride ?out_strides (F16)))
(= ?__cuda_elems (n_elements ?out_shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 16)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelBatchMatMul-F16-bytes"
)
(rule
((= ?kind (KernelBatchMatMul ?out_shape ?k_dim ?a_stride ?a_k_stride ?b_stride ?b_k_stride ?out_strides (Bf16)))
(= ?__cuda_elems (n_elements ?out_shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 16)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelBatchMatMul-Bf16-bytes"
)
(rule
((= ?kind (KernelBatchMatMul ?out_shape ?k_dim ?a_stride ?a_k_stride ?b_stride ?b_k_stride ?out_strides (Int)))
(= ?__cuda_elems (n_elements ?out_shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 32)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelBatchMatMul-Int-bytes"
)
(rule
((= ?kind (KernelBatchMatMul ?out_shape ?k_dim ?a_stride ?a_k_stride ?b_stride ?b_k_stride ?out_strides (Bool)))
(= ?__cuda_elems (n_elements ?out_shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 8)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelBatchMatMul-Bool-bytes"
)
(rule
((= ?kind (KernelBatchMatMul ?out_shape ?k_dim ?a_stride ?a_k_stride ?b_stride ?b_k_stride ?out_strides (I4)))
(= ?__cuda_elems (n_elements ?out_shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 4)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelBatchMatMul-I4-bytes"
)
(rule
((= ?kind (KernelBatchMatMul ?out_shape ?k_dim ?a_stride ?a_k_stride ?b_stride ?b_k_stride ?out_strides (TF32)))
(= ?__cuda_elems (n_elements ?out_shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 19)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelBatchMatMul-TF32-bytes"
)
(rule
((= ?kind (KernelBatchMatVec ?out_shape ?k_dim ?a_stride ?a_k_stride ?b_stride ?b_k_stride ?out_strides (F32)))
(= ?__cuda_elems (n_elements ?out_shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 32)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelBatchMatVec-F32-bytes"
)
(rule
((= ?kind (KernelBatchMatVec ?out_shape ?k_dim ?a_stride ?a_k_stride ?b_stride ?b_k_stride ?out_strides (F16)))
(= ?__cuda_elems (n_elements ?out_shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 16)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelBatchMatVec-F16-bytes"
)
(rule
((= ?kind (KernelBatchMatVec ?out_shape ?k_dim ?a_stride ?a_k_stride ?b_stride ?b_k_stride ?out_strides (Bf16)))
(= ?__cuda_elems (n_elements ?out_shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 16)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelBatchMatVec-Bf16-bytes"
)
(rule
((= ?kind (KernelBatchMatVec ?out_shape ?k_dim ?a_stride ?a_k_stride ?b_stride ?b_k_stride ?out_strides (Int)))
(= ?__cuda_elems (n_elements ?out_shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 32)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelBatchMatVec-Int-bytes"
)
(rule
((= ?kind (KernelBatchMatVec ?out_shape ?k_dim ?a_stride ?a_k_stride ?b_stride ?b_k_stride ?out_strides (Bool)))
(= ?__cuda_elems (n_elements ?out_shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 8)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelBatchMatVec-Bool-bytes"
)
(rule
((= ?kind (KernelBatchMatVec ?out_shape ?k_dim ?a_stride ?a_k_stride ?b_stride ?b_k_stride ?out_strides (I4)))
(= ?__cuda_elems (n_elements ?out_shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 4)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelBatchMatVec-I4-bytes"
)
(rule
((= ?kind (KernelBatchMatVec ?out_shape ?k_dim ?a_stride ?a_k_stride ?b_stride ?b_k_stride ?out_strides (TF32)))
(= ?__cuda_elems (n_elements ?out_shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 19)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelBatchMatVec-TF32-bytes"
)
(rule
((= ?kind (KernelCast ?size (F32) ?src_dtype)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?size (MNum 32)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelCast-F32-bytes"
)
(rule
((= ?kind (KernelCast ?size (F16) ?src_dtype)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?size (MNum 16)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelCast-F16-bytes"
)
(rule
((= ?kind (KernelCast ?size (Bf16) ?src_dtype)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?size (MNum 16)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelCast-Bf16-bytes"
)
(rule
((= ?kind (KernelCast ?size (Int) ?src_dtype)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?size (MNum 32)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelCast-Int-bytes"
)
(rule
((= ?kind (KernelCast ?size (Bool) ?src_dtype)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?size (MNum 8)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelCast-Bool-bytes"
)
(rule
((= ?kind (KernelCast ?size (I4) ?src_dtype)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?size (MNum 4)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelCast-I4-bytes"
)
(rule
((= ?kind (KernelCast ?size (TF32) ?src_dtype)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?size (MNum 19)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelCast-TF32-bytes"
)
(rule
((= ?kind (KernelConstant ?value)))
((cuda_output_bytes ?kind (MNum 4)))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelConstant-f32-scalar"
)
(rule
((= ?kind (KernelEmbed ?batch_shape ?token_stride ?out_stride ?embed_dim))
(= ?__cuda_elems (n_elements ?batch_shape)))
((cuda_output_bytes ?kind (MMul (MMul ?__cuda_elems ?embed_dim) (MNum 4))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelEmbed-f32-embed"
)
(rule
((= ?kind (KernelExp ?shape ?strides ?out_strides (F32)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 32)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelExp-F32-bytes"
)
(rule
((= ?kind (KernelExp ?shape ?strides ?out_strides (F16)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 16)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelExp-F16-bytes"
)
(rule
((= ?kind (KernelExp ?shape ?strides ?out_strides (Bf16)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 16)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelExp-Bf16-bytes"
)
(rule
((= ?kind (KernelExp ?shape ?strides ?out_strides (Int)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 32)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelExp-Int-bytes"
)
(rule
((= ?kind (KernelExp ?shape ?strides ?out_strides (Bool)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 8)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelExp-Bool-bytes"
)
(rule
((= ?kind (KernelExp ?shape ?strides ?out_strides (I4)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 4)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelExp-I4-bytes"
)
(rule
((= ?kind (KernelExp ?shape ?strides ?out_strides (TF32)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 19)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelExp-TF32-bytes"
)
(rule
((= ?kind (KernelExp2 ?shape ?strides ?out_strides (F32)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 32)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelExp2-F32-bytes"
)
(rule
((= ?kind (KernelExp2 ?shape ?strides ?out_strides (F16)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 16)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelExp2-F16-bytes"
)
(rule
((= ?kind (KernelExp2 ?shape ?strides ?out_strides (Bf16)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 16)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelExp2-Bf16-bytes"
)
(rule
((= ?kind (KernelExp2 ?shape ?strides ?out_strides (Int)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 32)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelExp2-Int-bytes"
)
(rule
((= ?kind (KernelExp2 ?shape ?strides ?out_strides (Bool)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 8)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelExp2-Bool-bytes"
)
(rule
((= ?kind (KernelExp2 ?shape ?strides ?out_strides (I4)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 4)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelExp2-I4-bytes"
)
(rule
((= ?kind (KernelExp2 ?shape ?strides ?out_strides (TF32)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 19)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelExp2-TF32-bytes"
)
(rule
((= ?kind (KernelGather ?out_shape ?index_strides ?data_shape ?data_strides ?out_strides (F32)))
(= ?__cuda_elems (n_elements ?out_shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 32)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelGather-F32-bytes"
)
(rule
((= ?kind (KernelGather ?out_shape ?index_strides ?data_shape ?data_strides ?out_strides (F16)))
(= ?__cuda_elems (n_elements ?out_shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 16)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelGather-F16-bytes"
)
(rule
((= ?kind (KernelGather ?out_shape ?index_strides ?data_shape ?data_strides ?out_strides (Bf16)))
(= ?__cuda_elems (n_elements ?out_shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 16)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelGather-Bf16-bytes"
)
(rule
((= ?kind (KernelGather ?out_shape ?index_strides ?data_shape ?data_strides ?out_strides (Int)))
(= ?__cuda_elems (n_elements ?out_shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 32)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelGather-Int-bytes"
)
(rule
((= ?kind (KernelGather ?out_shape ?index_strides ?data_shape ?data_strides ?out_strides (Bool)))
(= ?__cuda_elems (n_elements ?out_shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 8)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelGather-Bool-bytes"
)
(rule
((= ?kind (KernelGather ?out_shape ?index_strides ?data_shape ?data_strides ?out_strides (I4)))
(= ?__cuda_elems (n_elements ?out_shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 4)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelGather-I4-bytes"
)
(rule
((= ?kind (KernelGather ?out_shape ?index_strides ?data_shape ?data_strides ?out_strides (TF32)))
(= ?__cuda_elems (n_elements ?out_shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 19)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelGather-TF32-bytes"
)
(rule
((= ?kind (KernelIota ?expr ?range)))
((cuda_output_bytes ?kind (MMul ?range (MNum 4))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelIota-int-range"
)
(rule
((= ?kind (KernelLessThan ?shape ?a_strides ?b_strides ?out_strides ?dtype))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind ?__cuda_elems))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelLessThan-bool-shape"
)
(rule
((= ?kind (KernelLog2 ?shape ?strides ?out_strides (F32)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 32)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelLog2-F32-bytes"
)
(rule
((= ?kind (KernelLog2 ?shape ?strides ?out_strides (F16)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 16)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelLog2-F16-bytes"
)
(rule
((= ?kind (KernelLog2 ?shape ?strides ?out_strides (Bf16)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 16)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelLog2-Bf16-bytes"
)
(rule
((= ?kind (KernelLog2 ?shape ?strides ?out_strides (Int)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 32)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelLog2-Int-bytes"
)
(rule
((= ?kind (KernelLog2 ?shape ?strides ?out_strides (Bool)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 8)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelLog2-Bool-bytes"
)
(rule
((= ?kind (KernelLog2 ?shape ?strides ?out_strides (I4)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 4)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelLog2-I4-bytes"
)
(rule
((= ?kind (KernelLog2 ?shape ?strides ?out_strides (TF32)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 19)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelLog2-TF32-bytes"
)
(rule
((= ?kind (KernelMax ?shape ?iters ?strides ?iter_stride ?out_strides (F32)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 32)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelMax-F32-bytes"
)
(rule
((= ?kind (KernelMax ?shape ?iters ?strides ?iter_stride ?out_strides (F16)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 16)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelMax-F16-bytes"
)
(rule
((= ?kind (KernelMax ?shape ?iters ?strides ?iter_stride ?out_strides (Bf16)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 16)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelMax-Bf16-bytes"
)
(rule
((= ?kind (KernelMax ?shape ?iters ?strides ?iter_stride ?out_strides (Int)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 32)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelMax-Int-bytes"
)
(rule
((= ?kind (KernelMax ?shape ?iters ?strides ?iter_stride ?out_strides (Bool)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 8)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelMax-Bool-bytes"
)
(rule
((= ?kind (KernelMax ?shape ?iters ?strides ?iter_stride ?out_strides (I4)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 4)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelMax-I4-bytes"
)
(rule
((= ?kind (KernelMax ?shape ?iters ?strides ?iter_stride ?out_strides (TF32)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 19)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelMax-TF32-bytes"
)
(rule
((= ?kind (KernelMean ?shape ?iters ?strides ?iter_stride ?out_strides (F32)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 32)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelMean-F32-bytes"
)
(rule
((= ?kind (KernelMean ?shape ?iters ?strides ?iter_stride ?out_strides (F16)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 16)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelMean-F16-bytes"
)
(rule
((= ?kind (KernelMean ?shape ?iters ?strides ?iter_stride ?out_strides (Bf16)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 16)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelMean-Bf16-bytes"
)
(rule
((= ?kind (KernelMean ?shape ?iters ?strides ?iter_stride ?out_strides (Int)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 32)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelMean-Int-bytes"
)
(rule
((= ?kind (KernelMean ?shape ?iters ?strides ?iter_stride ?out_strides (Bool)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 8)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelMean-Bool-bytes"
)
(rule
((= ?kind (KernelMean ?shape ?iters ?strides ?iter_stride ?out_strides (I4)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 4)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelMean-I4-bytes"
)
(rule
((= ?kind (KernelMean ?shape ?iters ?strides ?iter_stride ?out_strides (TF32)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 19)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelMean-TF32-bytes"
)
(rule
((= ?kind (KernelMod ?shape ?a_strides ?b_strides ?out_strides (F32)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 32)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelMod-F32-bytes"
)
(rule
((= ?kind (KernelMod ?shape ?a_strides ?b_strides ?out_strides (F16)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 16)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelMod-F16-bytes"
)
(rule
((= ?kind (KernelMod ?shape ?a_strides ?b_strides ?out_strides (Bf16)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 16)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelMod-Bf16-bytes"
)
(rule
((= ?kind (KernelMod ?shape ?a_strides ?b_strides ?out_strides (Int)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 32)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelMod-Int-bytes"
)
(rule
((= ?kind (KernelMod ?shape ?a_strides ?b_strides ?out_strides (Bool)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 8)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelMod-Bool-bytes"
)
(rule
((= ?kind (KernelMod ?shape ?a_strides ?b_strides ?out_strides (I4)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 4)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelMod-I4-bytes"
)
(rule
((= ?kind (KernelMod ?shape ?a_strides ?b_strides ?out_strides (TF32)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 19)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelMod-TF32-bytes"
)
(rule
((= ?kind (KernelMul ?shape ?a_strides ?b_strides ?out_strides (F32)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 32)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelMul-F32-bytes"
)
(rule
((= ?kind (KernelMul ?shape ?a_strides ?b_strides ?out_strides (F16)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 16)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelMul-F16-bytes"
)
(rule
((= ?kind (KernelMul ?shape ?a_strides ?b_strides ?out_strides (Bf16)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 16)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelMul-Bf16-bytes"
)
(rule
((= ?kind (KernelMul ?shape ?a_strides ?b_strides ?out_strides (Int)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 32)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelMul-Int-bytes"
)
(rule
((= ?kind (KernelMul ?shape ?a_strides ?b_strides ?out_strides (Bool)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 8)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelMul-Bool-bytes"
)
(rule
((= ?kind (KernelMul ?shape ?a_strides ?b_strides ?out_strides (I4)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 4)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelMul-I4-bytes"
)
(rule
((= ?kind (KernelMul ?shape ?a_strides ?b_strides ?out_strides (TF32)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 19)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelMul-TF32-bytes"
)
(rule
((= ?kind (KernelRecip ?shape ?strides ?out_strides (F32)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 32)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelRecip-F32-bytes"
)
(rule
((= ?kind (KernelRecip ?shape ?strides ?out_strides (F16)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 16)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelRecip-F16-bytes"
)
(rule
((= ?kind (KernelRecip ?shape ?strides ?out_strides (Bf16)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 16)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelRecip-Bf16-bytes"
)
(rule
((= ?kind (KernelRecip ?shape ?strides ?out_strides (Int)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 32)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelRecip-Int-bytes"
)
(rule
((= ?kind (KernelRecip ?shape ?strides ?out_strides (Bool)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 8)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelRecip-Bool-bytes"
)
(rule
((= ?kind (KernelRecip ?shape ?strides ?out_strides (I4)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 4)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelRecip-I4-bytes"
)
(rule
((= ?kind (KernelRecip ?shape ?strides ?out_strides (TF32)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 19)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelRecip-TF32-bytes"
)
(rule
((= ?kind (KernelScatter ?dest_shape ?dest_strides ?index_shape ?index_strides ?src_strides ?out_strides (F32)))
(= ?__cuda_elems (n_elements ?dest_shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 32)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelScatter-F32-bytes"
)
(rule
((= ?kind (KernelScatter ?dest_shape ?dest_strides ?index_shape ?index_strides ?src_strides ?out_strides (F16)))
(= ?__cuda_elems (n_elements ?dest_shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 16)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelScatter-F16-bytes"
)
(rule
((= ?kind (KernelScatter ?dest_shape ?dest_strides ?index_shape ?index_strides ?src_strides ?out_strides (Bf16)))
(= ?__cuda_elems (n_elements ?dest_shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 16)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelScatter-Bf16-bytes"
)
(rule
((= ?kind (KernelScatter ?dest_shape ?dest_strides ?index_shape ?index_strides ?src_strides ?out_strides (Int)))
(= ?__cuda_elems (n_elements ?dest_shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 32)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelScatter-Int-bytes"
)
(rule
((= ?kind (KernelScatter ?dest_shape ?dest_strides ?index_shape ?index_strides ?src_strides ?out_strides (Bool)))
(= ?__cuda_elems (n_elements ?dest_shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 8)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelScatter-Bool-bytes"
)
(rule
((= ?kind (KernelScatter ?dest_shape ?dest_strides ?index_shape ?index_strides ?src_strides ?out_strides (I4)))
(= ?__cuda_elems (n_elements ?dest_shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 4)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelScatter-I4-bytes"
)
(rule
((= ?kind (KernelScatter ?dest_shape ?dest_strides ?index_shape ?index_strides ?src_strides ?out_strides (TF32)))
(= ?__cuda_elems (n_elements ?dest_shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 19)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelScatter-TF32-bytes"
)
(rule
((= ?kind (KernelScatterNoCopy ?dest_shape ?dest_strides ?index_shape ?index_strides ?src_strides ?out_strides (F32)))
(= ?__cuda_elems (n_elements ?dest_shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 32)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelScatterNoCopy-F32-bytes"
)
(rule
((= ?kind (KernelScatterNoCopy ?dest_shape ?dest_strides ?index_shape ?index_strides ?src_strides ?out_strides (F16)))
(= ?__cuda_elems (n_elements ?dest_shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 16)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelScatterNoCopy-F16-bytes"
)
(rule
((= ?kind (KernelScatterNoCopy ?dest_shape ?dest_strides ?index_shape ?index_strides ?src_strides ?out_strides (Bf16)))
(= ?__cuda_elems (n_elements ?dest_shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 16)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelScatterNoCopy-Bf16-bytes"
)
(rule
((= ?kind (KernelScatterNoCopy ?dest_shape ?dest_strides ?index_shape ?index_strides ?src_strides ?out_strides (Int)))
(= ?__cuda_elems (n_elements ?dest_shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 32)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelScatterNoCopy-Int-bytes"
)
(rule
((= ?kind (KernelScatterNoCopy ?dest_shape ?dest_strides ?index_shape ?index_strides ?src_strides ?out_strides (Bool)))
(= ?__cuda_elems (n_elements ?dest_shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 8)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelScatterNoCopy-Bool-bytes"
)
(rule
((= ?kind (KernelScatterNoCopy ?dest_shape ?dest_strides ?index_shape ?index_strides ?src_strides ?out_strides (I4)))
(= ?__cuda_elems (n_elements ?dest_shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 4)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelScatterNoCopy-I4-bytes"
)
(rule
((= ?kind (KernelScatterNoCopy ?dest_shape ?dest_strides ?index_shape ?index_strides ?src_strides ?out_strides (TF32)))
(= ?__cuda_elems (n_elements ?dest_shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 19)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelScatterNoCopy-TF32-bytes"
)
(rule
((= ?kind (KernelSigmoid ?shape ?strides ?out_strides (F32)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 32)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelSigmoid-F32-bytes"
)
(rule
((= ?kind (KernelSigmoid ?shape ?strides ?out_strides (F16)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 16)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelSigmoid-F16-bytes"
)
(rule
((= ?kind (KernelSigmoid ?shape ?strides ?out_strides (Bf16)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 16)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelSigmoid-Bf16-bytes"
)
(rule
((= ?kind (KernelSigmoid ?shape ?strides ?out_strides (Int)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 32)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelSigmoid-Int-bytes"
)
(rule
((= ?kind (KernelSigmoid ?shape ?strides ?out_strides (Bool)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 8)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelSigmoid-Bool-bytes"
)
(rule
((= ?kind (KernelSigmoid ?shape ?strides ?out_strides (I4)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 4)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelSigmoid-I4-bytes"
)
(rule
((= ?kind (KernelSigmoid ?shape ?strides ?out_strides (TF32)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 19)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelSigmoid-TF32-bytes"
)
(rule
((= ?kind (KernelSin ?shape ?strides ?out_strides (F32)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 32)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelSin-F32-bytes"
)
(rule
((= ?kind (KernelSin ?shape ?strides ?out_strides (F16)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 16)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelSin-F16-bytes"
)
(rule
((= ?kind (KernelSin ?shape ?strides ?out_strides (Bf16)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 16)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelSin-Bf16-bytes"
)
(rule
((= ?kind (KernelSin ?shape ?strides ?out_strides (Int)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 32)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelSin-Int-bytes"
)
(rule
((= ?kind (KernelSin ?shape ?strides ?out_strides (Bool)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 8)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelSin-Bool-bytes"
)
(rule
((= ?kind (KernelSin ?shape ?strides ?out_strides (I4)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 4)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelSin-I4-bytes"
)
(rule
((= ?kind (KernelSin ?shape ?strides ?out_strides (TF32)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 19)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelSin-TF32-bytes"
)
(rule
((= ?kind (KernelSoftmax ?shape ?in_strides ?out_strides ?reduce_dim ?reduce_stride ?dtype))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MMul ?__cuda_elems (MNum 4))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelSoftmax-f32-shape"
)
(rule
((= ?kind (KernelSqrt ?shape ?strides ?out_strides (F32)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 32)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelSqrt-F32-bytes"
)
(rule
((= ?kind (KernelSqrt ?shape ?strides ?out_strides (F16)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 16)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelSqrt-F16-bytes"
)
(rule
((= ?kind (KernelSqrt ?shape ?strides ?out_strides (Bf16)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 16)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelSqrt-Bf16-bytes"
)
(rule
((= ?kind (KernelSqrt ?shape ?strides ?out_strides (Int)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 32)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelSqrt-Int-bytes"
)
(rule
((= ?kind (KernelSqrt ?shape ?strides ?out_strides (Bool)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 8)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelSqrt-Bool-bytes"
)
(rule
((= ?kind (KernelSqrt ?shape ?strides ?out_strides (I4)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 4)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelSqrt-I4-bytes"
)
(rule
((= ?kind (KernelSqrt ?shape ?strides ?out_strides (TF32)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 19)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelSqrt-TF32-bytes"
)
(rule
((= ?kind (KernelSum ?shape ?iters ?strides ?iter_stride ?out_strides (F32)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 32)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelSum-F32-bytes"
)
(rule
((= ?kind (KernelSum ?shape ?iters ?strides ?iter_stride ?out_strides (F16)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 16)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelSum-F16-bytes"
)
(rule
((= ?kind (KernelSum ?shape ?iters ?strides ?iter_stride ?out_strides (Bf16)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 16)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelSum-Bf16-bytes"
)
(rule
((= ?kind (KernelSum ?shape ?iters ?strides ?iter_stride ?out_strides (Int)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 32)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelSum-Int-bytes"
)
(rule
((= ?kind (KernelSum ?shape ?iters ?strides ?iter_stride ?out_strides (Bool)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 8)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelSum-Bool-bytes"
)
(rule
((= ?kind (KernelSum ?shape ?iters ?strides ?iter_stride ?out_strides (I4)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 4)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelSum-I4-bytes"
)
(rule
((= ?kind (KernelSum ?shape ?iters ?strides ?iter_stride ?out_strides (TF32)))
(= ?__cuda_elems (n_elements ?shape)))
((cuda_output_bytes ?kind (MCeilDiv (MMul ?__cuda_elems (MNum 19)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-KernelSum-TF32-bytes"
)
(rule
((= ?kind (LoopInput ?loop_id ?stream_id ?dtype)))
((cuda_output_bytes ?kind (MNum 0)))
:ruleset cuda_memory_analysis
:name "cuda-memory-LoopInput-zero"
)
(rule
((= ?kind (LoopInputStatic ?loop_id ?stream_id ?dtype)))
((cuda_output_bytes ?kind (MNum 0)))
:ruleset cuda_memory_analysis
:name "cuda-memory-LoopInputStatic-zero"
)
(rule
((= ?kind (LoopOutput ?loop_id ?stream_id ?dtype)))
((cuda_output_bytes ?kind (MNum 0)))
:ruleset cuda_memory_analysis
:name "cuda-memory-LoopOutput-zero"
)
(rule
((= ?kind (LoopOutputSelect ?loop_id ?stream_id ?iter ?dtype)))
((cuda_output_bytes ?kind (MNum 0)))
:ruleset cuda_memory_analysis
:name "cuda-memory-LoopOutputSelect-zero"
)
(rule
((= ?kind (cublaslt ?m ?n ?k ?a_layout ?b_layout ?a_order ?b_order ?c_order ?d_order ?lda ?ldb ?ldc ?ldd ?batch_count ?stride_a ?stride_b ?stride_c ?stride_d ?a_dtype ?b_dtype ?c_dtype (F32) ?compute_type ?scale_dtype ?alpha ?beta ?epilogue)))
((cuda_output_bytes ?kind (MCeilDiv (MMul (MMul (MMul ?batch_count ?m) ?n) (MNum 32)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-cublaslt-F32-bytes"
)
(rule
((= ?kind (cublaslt ?m ?n ?k ?a_layout ?b_layout ?a_order ?b_order ?c_order ?d_order ?lda ?ldb ?ldc ?ldd ?batch_count ?stride_a ?stride_b ?stride_c ?stride_d ?a_dtype ?b_dtype ?c_dtype (F16) ?compute_type ?scale_dtype ?alpha ?beta ?epilogue)))
((cuda_output_bytes ?kind (MCeilDiv (MMul (MMul (MMul ?batch_count ?m) ?n) (MNum 16)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-cublaslt-F16-bytes"
)
(rule
((= ?kind (cublaslt ?m ?n ?k ?a_layout ?b_layout ?a_order ?b_order ?c_order ?d_order ?lda ?ldb ?ldc ?ldd ?batch_count ?stride_a ?stride_b ?stride_c ?stride_d ?a_dtype ?b_dtype ?c_dtype (Bf16) ?compute_type ?scale_dtype ?alpha ?beta ?epilogue)))
((cuda_output_bytes ?kind (MCeilDiv (MMul (MMul (MMul ?batch_count ?m) ?n) (MNum 16)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-cublaslt-Bf16-bytes"
)
(rule
((= ?kind (cublaslt ?m ?n ?k ?a_layout ?b_layout ?a_order ?b_order ?c_order ?d_order ?lda ?ldb ?ldc ?ldd ?batch_count ?stride_a ?stride_b ?stride_c ?stride_d ?a_dtype ?b_dtype ?c_dtype (Int) ?compute_type ?scale_dtype ?alpha ?beta ?epilogue)))
((cuda_output_bytes ?kind (MCeilDiv (MMul (MMul (MMul ?batch_count ?m) ?n) (MNum 32)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-cublaslt-Int-bytes"
)
(rule
((= ?kind (cublaslt ?m ?n ?k ?a_layout ?b_layout ?a_order ?b_order ?c_order ?d_order ?lda ?ldb ?ldc ?ldd ?batch_count ?stride_a ?stride_b ?stride_c ?stride_d ?a_dtype ?b_dtype ?c_dtype (Bool) ?compute_type ?scale_dtype ?alpha ?beta ?epilogue)))
((cuda_output_bytes ?kind (MCeilDiv (MMul (MMul (MMul ?batch_count ?m) ?n) (MNum 8)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-cublaslt-Bool-bytes"
)
(rule
((= ?kind (cublaslt ?m ?n ?k ?a_layout ?b_layout ?a_order ?b_order ?c_order ?d_order ?lda ?ldb ?ldc ?ldd ?batch_count ?stride_a ?stride_b ?stride_c ?stride_d ?a_dtype ?b_dtype ?c_dtype (I4) ?compute_type ?scale_dtype ?alpha ?beta ?epilogue)))
((cuda_output_bytes ?kind (MCeilDiv (MMul (MMul (MMul ?batch_count ?m) ?n) (MNum 4)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-cublaslt-I4-bytes"
)
(rule
((= ?kind (cublaslt ?m ?n ?k ?a_layout ?b_layout ?a_order ?b_order ?c_order ?d_order ?lda ?ldb ?ldc ?ldd ?batch_count ?stride_a ?stride_b ?stride_c ?stride_d ?a_dtype ?b_dtype ?c_dtype (TF32) ?compute_type ?scale_dtype ?alpha ?beta ?epilogue)))
((cuda_output_bytes ?kind (MCeilDiv (MMul (MMul (MMul ?batch_count ?m) ?n) (MNum 19)) (MNum 8))))
:ruleset cuda_memory_analysis
:name "cuda-memory-cublaslt-TF32-bytes"
)