diff --git a/src/enzyme_ad/jax/BUILD b/src/enzyme_ad/jax/BUILD index 1746580071..262be58786 100644 --- a/src/enzyme_ad/jax/BUILD +++ b/src/enzyme_ad/jax/BUILD @@ -540,6 +540,7 @@ td_library( deps = [ "@llvm-project//mlir:BuiltinDialectTdFiles", "@llvm-project//mlir:OpBaseTdFiles", + "@stablehlo//:chlo_ops_td_files", "@stablehlo//:stablehlo_ops_td_files", ], ) diff --git a/src/enzyme_ad/jax/CheckedRewrite.h b/src/enzyme_ad/jax/CheckedRewrite.h index 9c5d84e48e..81ecef2cc4 100644 --- a/src/enzyme_ad/jax/CheckedRewrite.h +++ b/src/enzyme_ad/jax/CheckedRewrite.h @@ -1,5 +1,7 @@ #pragma once +#include + #include "mlir/IR/PatternMatch.h" #include "mlir/Interfaces/FunctionInterfaces.h" @@ -39,6 +41,24 @@ static LogicalResult failIfFuncOpInterfaceHasAttr(Operation *op, return success(); } +static LogicalResult checkPreconditions(Operation *op, + PatternRewriter &rewriter, + bool supportsDynamicShapes) { + if (op->hasAttr(kDisablePatternAttrName)) + return rewriter.notifyMatchFailure(op, "disabled by attribute."); + + if (failIfFuncOpInterfaceHasAttr(op, kDisablePatternAttrName, rewriter) + .failed()) + return failure(); + + if (!supportsDynamicShapes) { + if (failIfDynamicShape(op, rewriter).failed()) + return failure(); + } + + return success(); +} + template struct CheckedOpRewritePattern : public OpRewritePattern { using Base = OpRewritePattern; @@ -46,16 +66,10 @@ struct CheckedOpRewritePattern : public OpRewritePattern { LogicalResult matchAndRewrite(OpTy op, PatternRewriter &rewriter) const override final { - LogicalResult res = - failIfFuncOpInterfaceHasAttr(op, kDisablePatternAttrName, rewriter); - if (res.failed()) - return res; - - if (!((Child *)this)->supportsDynamicShapes()) { - LogicalResult res = failIfDynamicShape(op, rewriter); - if (res.failed()) - return res; - } + if (checkPreconditions(op, rewriter, + ((Child *)this)->supportsDynamicShapes()) + .failed()) + return failure(); return ((Child *)this)->matchAndRewriteImpl(op, rewriter); } @@ -71,16 +85,10 @@ struct CheckedOpTraitRewritePattern : public OpTraitRewritePattern { LogicalResult matchAndRewrite(Operation *op, PatternRewriter &rewriter) const override final { - LogicalResult res = - failIfFuncOpInterfaceHasAttr(op, kDisablePatternAttrName, rewriter); - if (res.failed()) - return res; - - if (!((Child *)this)->supportsDynamicShapes()) { - auto res = failIfDynamicShape(op, rewriter); - if (res.failed()) - return res; - } + if (checkPreconditions(op, rewriter, + ((Child *)this)->supportsDynamicShapes()) + .failed()) + return failure(); return ((Child *)this)->matchAndRewriteImpl(op, rewriter); } @@ -88,5 +96,30 @@ struct CheckedOpTraitRewritePattern : public OpTraitRewritePattern { bool supportsDynamicShapes() const { return false; } }; +template +struct has_supports_dynamic_shapes : std::false_type {}; + +template +struct has_supports_dynamic_shapes< + T, std::void_t().supportsDynamicShapes())>> + : std::true_type {}; + +template struct CheckedPattern : public PatternTy { + using PatternTy::PatternTy; + + LogicalResult matchAndRewrite(Operation *op, + PatternRewriter &rewriter) const override { + bool supportsDynamic = false; + if constexpr (has_supports_dynamic_shapes::value) { + supportsDynamic = this->supportsDynamicShapes(); + } + + if (checkPreconditions(op, rewriter, supportsDynamic).failed()) + return failure(); + + return PatternTy::matchAndRewrite(op, rewriter); + } +}; + } // namespace enzyme } // namespace mlir diff --git a/src/enzyme_ad/jax/Passes/EnzymeHLOOpt.cpp b/src/enzyme_ad/jax/Passes/EnzymeHLOOpt.cpp index 25d2c1cbb1..859caff81f 100644 --- a/src/enzyme_ad/jax/Passes/EnzymeHLOOpt.cpp +++ b/src/enzyme_ad/jax/Passes/EnzymeHLOOpt.cpp @@ -4194,62 +4194,6 @@ struct ConvertConcat final } }; -struct ConvertConvertFloat final - : CheckedOpRewritePattern { - using CheckedOpRewritePattern::CheckedOpRewritePattern; - - LogicalResult matchAndRewriteImpl(stablehlo::ConvertOp op, - PatternRewriter &rewriter) const { - auto conv0 = op.getOperand().getDefiningOp(); - if (!conv0) - return failure(); - - auto prev = conv0.getOperand(); - if (isa(prev.getType().getElementType()) && - isa(op.getType().getElementType()) && - isa(conv0.getType().getElementType())) { - if (prev.getType() == op.getType()) { - rewriter.replaceOp(op, prev); - return success(); - } - rewriter.replaceOpWithNewOp(op, op.getType(), prev); - return success(); - } - return failure(); - } -}; - -struct ConvertConvertInt final - : CheckedOpRewritePattern { - using CheckedOpRewritePattern::CheckedOpRewritePattern; - - LogicalResult matchAndRewriteImpl(stablehlo::ConvertOp op, - PatternRewriter &rewriter) const { - auto conv0 = op.getOperand().getDefiningOp(); - if (!conv0) - return failure(); - - auto prev = conv0.getOperand(); - if (isa(prev.getType().getElementType()) && - isa(op.getType().getElementType()) && - isa(conv0.getType().getElementType())) { - // we only do the elimination if we go from low bitwidth to high bitwidth - auto prevwidth = prev.getType().getElementType().getIntOrFloatBitWidth(); - auto midwidth = conv0.getType().getElementType().getIntOrFloatBitWidth(); - if (prevwidth > midwidth) - return rewriter.notifyMatchFailure(op, "prevwidth > midwidth"); - - if (prev.getType() == op.getType()) { - rewriter.replaceOp(op, prev); - return success(); - } - rewriter.replaceOpWithNewOp(op, op.getType(), prev); - return success(); - } - return failure(); - } -}; - struct ReduceConcat final : CheckedOpRewritePattern { using CheckedOpRewritePattern::CheckedOpRewritePattern; @@ -7529,193 +7473,6 @@ struct TransposeElementwiseTransposeSimplify } }; -struct AddSimplify - : public CheckedOpRewritePattern { - using CheckedOpRewritePattern::CheckedOpRewritePattern; - - LogicalResult matchAndRewriteImpl(stablehlo::AddOp op, - PatternRewriter &rewriter) const { - Attribute lhsAttr, rhsAttr; - bool lhsConst = matchPattern(op.getLhs(), m_Constant(&lhsAttr)); - bool rhsConst = matchPattern(op.getRhs(), m_Constant(&rhsAttr)); - - if (lhsConst && (matchPattern(lhsAttr, m_AnyZeroFloat()) || - matchPattern(lhsAttr, m_Zero()) || - matchPattern(lhsAttr, m_AnyZeroComplex()))) { - rewriter.replaceOp(op, op.getRhs()); - return success(); - } - - if (rhsConst && (matchPattern(rhsAttr, m_AnyZeroFloat()) || - matchPattern(rhsAttr, m_Zero()) || - matchPattern(rhsAttr, m_AnyZeroComplex()))) { - rewriter.replaceOp(op, op.getLhs()); - return success(); - } - - return failure(); - } -}; - -struct ReplaceNegAddWithSubtract - : public CheckedOpRewritePattern { - using CheckedOpRewritePattern::CheckedOpRewritePattern; - - LogicalResult matchAndRewriteImpl(stablehlo::AddOp op, - PatternRewriter &rewriter) const { - if (auto rhsNegateOp = op.getRhs().getDefiningOp()) { - if (llvm::hasSingleElement(rhsNegateOp->getUsers())) { - rewriter.replaceOpWithNewOp( - op, op.getLhs(), rhsNegateOp.getOperand()); - return success(); - } - } - - if (auto lhsNegateOp = op.getLhs().getDefiningOp()) { - if (llvm::hasSingleElement(lhsNegateOp->getUsers())) { - rewriter.replaceOpWithNewOp( - op, op.getRhs(), lhsNegateOp.getOperand()); - return success(); - } - } - - return failure(); - } -}; - -struct ReplaceSubtractNegWithAdd - : CheckedOpRewritePattern { - using CheckedOpRewritePattern::CheckedOpRewritePattern; - - LogicalResult matchAndRewriteImpl(stablehlo::SubtractOp op, - PatternRewriter &rewriter) const { - if (auto rhsNegateOp = op.getRhs().getDefiningOp()) { - if (llvm::hasSingleElement(rhsNegateOp->getUsers())) { - rewriter.replaceOpWithNewOp(op, op.getLhs(), - rhsNegateOp.getOperand()); - return success(); - } - } - - return failure(); - } -}; - -struct SubSimplify - : public CheckedOpRewritePattern { - using CheckedOpRewritePattern::CheckedOpRewritePattern; - - LogicalResult matchAndRewriteImpl(stablehlo::SubtractOp op, - PatternRewriter &rewriter) const { - Attribute lhsAttr, rhsAttr; - bool lhsConst = matchPattern(op.getLhs(), m_Constant(&lhsAttr)); - bool rhsConst = matchPattern(op.getRhs(), m_Constant(&rhsAttr)); - - if (rhsConst && (matchPattern(rhsAttr, m_AnyZeroFloat()) || - matchPattern(rhsAttr, m_Zero()) || - matchPattern(rhsAttr, m_AnyZeroComplex()))) { - rewriter.replaceOp(op, op.getLhs()); - return success(); - } - - if (lhsConst && (matchPattern(lhsAttr, m_AnyZeroFloat()) || - matchPattern(lhsAttr, m_Zero()) || - matchPattern(lhsAttr, m_AnyZeroComplex()))) { - rewriter.replaceOpWithNewOp(op, op.getRhs()); - return success(); - } - - if (isa(op.getType().getElementType()) && - op.getLhs() == op.getRhs()) { - rewriter.replaceOpWithNewOp( - op, rewriter.getZeroAttr(op.getType())); - return success(); - } - - return failure(); - } -}; - -struct ExponentialMinusOneFuse - : public CheckedOpRewritePattern { - using CheckedOpRewritePattern< - stablehlo::SubtractOp, ExponentialMinusOneFuse>::CheckedOpRewritePattern; - - LogicalResult matchAndRewriteImpl(stablehlo::SubtractOp op, - PatternRewriter &rewriter) const { - auto lhs = op.getLhs(); - auto rhs = op.getRhs(); - - { // exp(x) - 1 -> expm1(x) - auto defOp = lhs.getDefiningOp(); - if (defOp && llvm::hasSingleElement(defOp->getUsers()) && - (matchPattern(rhs, m_One()) || matchPattern(rhs, m_OneFloat()))) { - rewriter.replaceOpWithNewOp(op, defOp.getOperand()); - return success(); - } - } - - { // 1 - exp(x) -> -expm1(x) - auto defOp = rhs.getDefiningOp(); - if (defOp && llvm::hasSingleElement(defOp->getUsers()) && - (matchPattern(lhs, m_One()) || matchPattern(lhs, m_OneFloat()))) { - auto expm1 = stablehlo::Expm1Op::create( - rewriter, op.getLoc(), op.getType(), defOp.getOperand()); - rewriter.replaceOpWithNewOp(op, expm1); - return success(); - } - } - - return failure(); - } -}; - -struct ExponentialMinusOneAddFuse - : public CheckedOpRewritePattern { - using CheckedOpRewritePattern< - stablehlo::AddOp, ExponentialMinusOneAddFuse>::CheckedOpRewritePattern; - - LogicalResult matchAndRewriteImpl(stablehlo::AddOp op, - PatternRewriter &rewriter) const { - auto lhs = op.getLhs(); - auto rhs = op.getRhs(); - - auto isMinusOne = [](Value val) { - SplatElementsAttr attr; - if (!matchPattern(val, m_Constant(&attr))) - return false; - auto doubleVal = getDoubleFromAttr(attr.getSplatValue()); - return doubleVal && *doubleVal == -1.0; - }; - - { // exp(x) + -1 -> expm1(x) - auto defOp = lhs.getDefiningOp(); - if (defOp && llvm::hasSingleElement(defOp->getUsers()) && - isMinusOne(rhs)) { - rewriter.replaceOpWithNewOp(op, defOp.getOperand()); - return success(); - } - } - - { // -1 + exp(x) -> expm1(x) - auto defOp = rhs.getDefiningOp(); - if (defOp && llvm::hasSingleElement(defOp->getUsers()) && - isMinusOne(lhs)) { - rewriter.replaceOpWithNewOp(op, defOp.getOperand()); - return success(); - } - } - - return failure(); - } -}; - struct TransposeSymmetricSimplify : public CheckedOpRewritePattern { @@ -7801,206 +7558,6 @@ struct NoNanSelfSubSimplify } }; -struct AndSimplify - : public CheckedOpRewritePattern { - using CheckedOpRewritePattern::CheckedOpRewritePattern; - - LogicalResult matchAndRewriteImpl(stablehlo::AndOp op, - PatternRewriter &rewriter) const { - - if (op.getLhs() == op.getRhs()) { - rewriter.replaceOp(op, op.getLhs()); - return success(); - } - - for (int i = 0; i < 2; i++) { - Attribute attr; - if (!matchPattern(op.getOperand(i), m_Constant(&attr))) - continue; - // false & x -> false - if (matchPattern(attr, m_Zero())) { - rewriter.replaceOp(op, op.getOperand(i)); - return success(); - } - // true & x -> x - if (matchPattern(attr, m_AllOnes())) { - rewriter.replaceOp(op, op.getOperand(1 - i)); - return success(); - } - } - - return failure(); - } -}; - -struct OrSimplify - : public CheckedOpRewritePattern { - using CheckedOpRewritePattern::CheckedOpRewritePattern; - - LogicalResult matchAndRewriteImpl(stablehlo::OrOp op, - PatternRewriter &rewriter) const { - - if (op.getLhs() == op.getRhs()) { - rewriter.replaceOp(op, op.getLhs()); - return success(); - } - - for (int i = 0; i < 2; i++) { - Attribute attr; - if (!matchPattern(op.getOperand(i), m_Constant(&attr))) - continue; - // true | x -> true - if (matchPattern(attr, m_AllOnes())) { - rewriter.replaceOp(op, op.getOperand(i)); - return success(); - } - // false | x -> x - if (matchPattern(attr, m_Zero())) { - rewriter.replaceOp(op, op.getOperand(1 - i)); - return success(); - } - } - - return failure(); - } -}; - -struct XorSimplify - : public CheckedOpRewritePattern { - using CheckedOpRewritePattern::CheckedOpRewritePattern; - - LogicalResult matchAndRewriteImpl(stablehlo::XorOp op, - PatternRewriter &rewriter) const { - - for (int i = 0; i < 2; i++) { - Attribute attr; - if (!matchPattern(op.getOperand(i), m_Constant(&attr))) - continue; - // false ^ x -> x - if (matchPattern(attr, m_Zero())) { - rewriter.replaceOp(op, op.getOperand(1 - i)); - return success(); - } - // true ^ x -> not x - if (matchPattern(attr, m_AllOnes())) { - rewriter.replaceOpWithNewOp(op, op.getOperand(1 - i)); - return success(); - } - } - - return failure(); - } -}; - -struct MulSimplify - : public CheckedOpRewritePattern { - using CheckedOpRewritePattern::CheckedOpRewritePattern; - - LogicalResult matchAndRewriteImpl(stablehlo::MulOp op, - PatternRewriter &rewriter) const { - Attribute lhsAttr, rhsAttr; - bool lhsConst = matchPattern(op.getLhs(), m_Constant(&lhsAttr)); - bool rhsConst = matchPattern(op.getRhs(), m_Constant(&rhsAttr)); - - // 1 * x -> x - if (lhsConst && (matchPattern(lhsAttr, m_One()) || - matchPattern(lhsAttr, m_OneFloat()))) { - rewriter.replaceOp(op, op.getRhs()); - return success(); - } - - // -1 * x -> negate x - if (lhsConst && (matchPattern(lhsAttr, m_NegOne()) || - matchPattern(lhsAttr, m_NegOneFloat()))) { - rewriter.replaceOpWithNewOp(op, op.getRhs()); - return success(); - } - - if (!rhsConst) - return failure(); - - // x * 1 -> x - if (rhsConst && (matchPattern(rhsAttr, m_One()) || - matchPattern(rhsAttr, m_OneFloat()))) { - rewriter.replaceOp(op, op.getLhs()); - return success(); - } - - // x * -1 -> negate x - if (rhsConst && (matchPattern(rhsAttr, m_NegOne()) || - matchPattern(rhsAttr, m_NegOneFloat()))) { - rewriter.replaceOpWithNewOp(op, op.getLhs()); - return success(); - } - - return failure(); - } -}; - -struct DivSimplify - : public CheckedOpRewritePattern { - using CheckedOpRewritePattern::CheckedOpRewritePattern; - - LogicalResult matchAndRewriteImpl(stablehlo::DivOp op, - PatternRewriter &rewriter) const { - Attribute rhsAttr; - bool rhsConst = matchPattern(op.getRhs(), m_Constant(&rhsAttr)); - - // x / 1 -> x - if (rhsConst && (matchPattern(rhsAttr, m_OneFloat()) || - matchPattern(rhsAttr, m_One()))) { - rewriter.replaceOp(op, op.getLhs()); - return success(); - } - - // x / -1 -> negate x - if (rhsConst && (matchPattern(rhsAttr, m_NegOneFloat()) || - matchPattern(rhsAttr, m_NegOne()))) { - rewriter.replaceOpWithNewOp(op, op.getLhs()); - return success(); - } - - // x / const -> x * (1 / const) - if (isa(op.getType().getElementType())) { - if (auto rhsDenseAttr = dyn_cast_or_null(rhsAttr)) { - { - DenseElementsAttr lhsAttr; - if (matchPattern(op.getLhs(), m_Constant(&lhsAttr))) - return failure(); // const prop will evaluate this - } - - auto ty = op.getType(); - if (rhsDenseAttr.isSplat()) { - ty = RankedTensorType::get( - {}, cast(op->getResultTypes()[0]).getElementType()); - rhsDenseAttr = rhsDenseAttr.resizeSplat(ty); - } - - auto rhsTen = stablehlo::constantOp(rhsDenseAttr); - auto oneTen = - stablehlo::constantOp(cast(makeAttr(ty, 1))); - auto out = fromTensor(stablehlo::divideOp(oneTen, rhsTen, ty)); - - if (ty != op.getType()) { - out = out.resizeSplat(op.getType()); - } - - rewriter.replaceOpWithNewOp( - op, op.getLhs(), - stablehlo::ConstantOp::create(rewriter, op.getLoc(), op.getType(), - out)); - return success(); - } - } - - return failure(); - } -}; - struct NoNanDivSimplify final : public NoNanCheckedOpRewritePattern { using NoNanCheckedOpRewritePattern< @@ -8034,98 +7591,6 @@ struct NoNanDivSimplify final } }; -struct RemSimplify - : public CheckedOpRewritePattern { - using CheckedOpRewritePattern::CheckedOpRewritePattern; - - LogicalResult matchAndRewriteImpl(stablehlo::RemOp op, - PatternRewriter &rewriter) const { - - if (matchPattern(op.getRhs(), m_One())) { - rewriter.replaceOpWithNewOp( - op, cast(makeAttr(op.getType(), 0))); - return success(); - } - - return failure(); - } -}; - -struct PowSimplify - : public CheckedOpRewritePattern { - using CheckedOpRewritePattern::CheckedOpRewritePattern; - - LogicalResult matchAndRewriteImpl(stablehlo::PowOp op, - PatternRewriter &rewriter) const { - Attribute lhsAttr, rhsAttr; - matchPattern(op.getLhs(), m_Constant(&lhsAttr)); - matchPattern(op.getRhs(), m_Constant(&rhsAttr)); - - // pow(x, 1) -> x - if (rhsAttr && (matchPattern(rhsAttr, m_One()) || - matchPattern(rhsAttr, m_OneFloat()))) { - rewriter.replaceAllUsesWith(op, op.getLhs()); - return success(); - } - - // pow(x, 0) -> 1 || pow(1, x) -> 1 - if ((rhsAttr && (matchPattern(rhsAttr, m_Zero()) || - matchPattern(rhsAttr, m_AnyZeroFloat()))) || - (lhsAttr && (matchPattern(lhsAttr, m_One()) || - matchPattern(lhsAttr, m_OneFloat())))) { - rewriter.replaceOpWithNewOp( - op, op.getType(), cast(makeAttr(op.getType(), 1))); - return success(); - } - - if (isa(op.getType().getElementType())) { - auto rhs = dyn_cast_or_null(rhsAttr); - if (rhs && rhs.isSplat()) { - bool allHalf = true, allNegOne = true, allNegHalf = true, allTwo = true; - auto v = rhs.getSplatValue(); - allHalf &= v.isExactlyValue(0.5); - allNegOne &= v.isExactlyValue(-1.0); - allNegHalf &= v.isExactlyValue(-0.5); - allTwo &= v.isExactlyValue(2.0); - - // pow(X, -1) -> 1 / X - if (allNegOne) { - rewriter.replaceOpWithNewOp( - op, - stablehlo::ConstantOp::create( - rewriter, op.getLoc(), op.getType(), - cast(makeAttr(op.getType(), 1))), - op.getLhs()); - return success(); - } - - // pow(X, -0.5) -> rsqrt(X) - if (allNegHalf) { - rewriter.replaceOpWithNewOp(op, op.getLhs()); - return success(); - } - - // pow(X, 0.5) -> sqrt(X) - if (allHalf) { - rewriter.replaceOpWithNewOp(op, op.getLhs()); - return success(); - } - - // pow(X, 2) -> X * X - if (allTwo) { - rewriter.replaceOpWithNewOp(op, op.getLhs(), - op.getLhs()); - return success(); - } - } - } - - return failure(); - } -}; - struct NoNanZeroBasePowSimplify final : public NoNanCheckedOpRewritePattern { @@ -8515,29 +7980,6 @@ struct ConvertSimplify } }; -struct ConvertIotaSimplify - : public CheckedOpRewritePattern { - using CheckedOpRewritePattern::CheckedOpRewritePattern; - - LogicalResult matchAndRewriteImpl(stablehlo::ConvertOp convertOp, - PatternRewriter &rewriter) const { - auto operand = convertOp.getOperand(); - auto iota = operand.getDefiningOp(); - if (!iota) - return failure(); - - auto targetType = convertOp.getType(); - if (!targetType.getElementType().isInteger()) - return failure(); - - rewriter.replaceOpWithNewOp(convertOp, targetType, - iota.getIotaDimension()); - return success(); - } -}; - struct SliceSimplify : public CheckedOpRewritePattern { using CheckedOpRewritePattern { - using CheckedOpRewritePattern::CheckedOpRewritePattern; - - LogicalResult matchAndRewriteImpl(stablehlo::MaxOp op, - PatternRewriter &rewriter) const { - if (op.getOperand(0) == op.getOperand(1)) { - rewriter.replaceOp(op, op.getOperand(0)); - return success(); - } - - return failure(); - } -}; - -struct MinSimplify - : public CheckedOpRewritePattern { - using CheckedOpRewritePattern::CheckedOpRewritePattern; - - LogicalResult matchAndRewriteImpl(stablehlo::MinOp op, - PatternRewriter &rewriter) const { - if (op.getOperand(0) == op.getOperand(1)) { - rewriter.replaceOp(op, op.getOperand(0)); - return success(); - } - - return failure(); - } -}; - template struct BinBroadcastSplat final : CheckedOpRewritePattern> { @@ -14194,38 +13604,6 @@ struct SelectPadToDUS final } }; -struct SelectSelectSameCond final - : CheckedOpRewritePattern { - using CheckedOpRewritePattern::CheckedOpRewritePattern; - - LogicalResult matchAndRewriteImpl(stablehlo::SelectOp op, - PatternRewriter &rewriter) const { - Value cond = op.getPred(); - - // Case 1: false branch is another select with the same condition - // select(cond, c, select(cond, a, b)) -> select(cond, c, b) - if (auto inner = op.getOnFalse().getDefiningOp()) { - if (inner.getPred() == cond) { - rewriter.modifyOpInPlace( - op, [&]() { op.getOnFalseMutable().assign(inner.getOnFalse()); }); - return success(); - } - } - - // Case 2: true branch is another select with the same condition - // select(cond, select(cond, a, b), c) -> select(cond, a, c) - if (auto inner = op.getOnTrue().getDefiningOp()) { - if (inner.getPred() == cond) { - rewriter.modifyOpInPlace( - op, [&]() { op.getOnTrueMutable().assign(inner.getOnTrue()); }); - return success(); - } - } - - return failure(); - } -}; - struct SelectSelectNegCond final : CheckedOpRewritePattern { using CheckedOpRewritePattern::CheckedOpRewritePattern; @@ -15016,24 +14394,6 @@ struct EmptyReduceOpCanon final } }; -struct DynamicReshapeOpCanon final - : CheckedOpRewritePattern { - using CheckedOpRewritePattern::CheckedOpRewritePattern; - - LogicalResult matchAndRewriteImpl(stablehlo::DynamicReshapeOp op, - PatternRewriter &rewriter) const { - // This is a noop when the output type is already a static shape. - RankedTensorType type = op.getType(); - if (!type.hasStaticShape()) - return failure(); - - rewriter.replaceOpWithNewOp(op, type, - op.getOperand()); - return success(); - } -}; - struct GetTupleElementOpCanon final : CheckedOpRewritePattern { @@ -15095,53 +14455,11 @@ struct ImagOpCanon final guaranteedPurelyRealResult(op.getOperand(), rewriter)) { rewriter.replaceOp( op, stablehlo::ConstantOp::create(rewriter, op->getLoc(), - makeAttr(op.getType(), 0))); - return success(); - } - - return failure(); - } -}; - -// (conj (complex a, (neg b))) -> (complex a b) -struct ConjComplexNegate final - : CheckedOpRewritePattern { - using CheckedOpRewritePattern::CheckedOpRewritePattern; - - LogicalResult matchAndRewriteImpl(chlo::ConjOp op, - PatternRewriter &rewriter) const { - auto complex = op.getOperand().getDefiningOp(); - if (!complex) - return failure(); - - auto neg = complex.getRhs().getDefiningOp(); - if (!neg) - return failure(); - - rewriter.replaceOpWithNewOp( - op, op.getType(), complex.getLhs(), neg.getOperand()); - return success(); - } -}; - -// (neg (imag (conj x))) -> (imag x) -struct NegateImagConj final - : CheckedOpRewritePattern { - using CheckedOpRewritePattern::CheckedOpRewritePattern; - - LogicalResult matchAndRewriteImpl(stablehlo::NegOp op, - PatternRewriter &rewriter) const { - auto imag = op.getOperand().getDefiningOp(); - if (!imag) - return failure(); - - auto conj = imag.getOperand().getDefiningOp(); - if (!conj) - return failure(); + makeAttr(op.getType(), 0))); + return success(); + } - rewriter.replaceOpWithNewOp(op, op.getType(), - conj.getOperand()); - return success(); + return failure(); } }; @@ -21371,24 +20689,6 @@ struct CompareSelectSimplify } }; -// select(!op, lhs, rhs) --> select(op, rhs, lhs) -struct NotSelectSimplify - : public CheckedOpRewritePattern { - using CheckedOpRewritePattern::CheckedOpRewritePattern; - - LogicalResult matchAndRewriteImpl(stablehlo::SelectOp op, - PatternRewriter &rewriter) const { - auto notOp = op.getPred().getDefiningOp(); - if (!notOp) - return failure(); - - rewriter.replaceOpWithNewOp( - op, notOp.getOperand(), op.getOnFalse(), op.getOnTrue()); - return success(); - } -}; - struct NotCompare : public CheckedOpRewritePattern { using CheckedOpRewritePattern::CheckedOpRewritePattern; @@ -22184,68 +21484,6 @@ struct ReduceTransposeSimplify } }; -// (mul (sign x) (abs x)) -> x -// (mul (abs x) (sign x)) -> x -struct SignAbsSimplify - : public CheckedOpRewritePattern { - using CheckedOpRewritePattern::CheckedOpRewritePattern; - - LogicalResult matchAndRewriteImpl(stablehlo::MulOp op, - PatternRewriter &rewriter) const { - auto lhs = op.getOperand(0); - auto rhs = op.getOperand(1); - - auto lhsSignOp = lhs.getDefiningOp(); - if (lhsSignOp) { - auto rhsAbsOp = rhs.getDefiningOp(); - if (!rhsAbsOp) - return failure(); - - if (lhsSignOp.getOperand() != rhsAbsOp.getOperand()) - return failure(); - - rewriter.replaceOp(op, lhsSignOp.getOperand()); - return success(); - } - - auto rhsSignOp = rhs.getDefiningOp(); - if (rhsSignOp) { - auto lhsAbsOp = lhs.getDefiningOp(); - if (!lhsAbsOp) - return failure(); - - if (rhsSignOp.getOperand() != lhsAbsOp.getOperand()) - return failure(); - - rewriter.replaceOp(op, rhsSignOp.getOperand()); - return success(); - } - - return failure(); - } -}; - -struct AbsPositiveSimplify - : public CheckedOpRewritePattern { - using CheckedOpRewritePattern::CheckedOpRewritePattern; - - LogicalResult matchAndRewriteImpl(stablehlo::AbsOp op, - PatternRewriter &rewriter) const { - - auto operand = op.getOperand(); - if (isa(operand.getType().getElementType())) - return failure(); - - if (guaranteedNonNegativeResult(operand.getDefiningOp(), rewriter)) { - rewriter.replaceOp(op, op.getOperand()); - return success(); - } - return failure(); - } -}; - struct TransposeReshapeToBroadcast final : CheckedOpRewritePattern { @@ -25671,47 +24909,6 @@ struct SliceRotate final } }; -struct SquareAbsSimplify - : public CheckedOpRewritePattern { - using CheckedOpRewritePattern::CheckedOpRewritePattern; - - LogicalResult matchAndRewriteImpl(stablehlo::MulOp op, - PatternRewriter &rewriter) const { - auto lhs = op.getLhs(); - auto rhs = op.getRhs(); - - if (lhs != rhs) - return failure(); - - auto absOp = lhs.getDefiningOp(); - if (!absOp) - return failure(); - - auto operand = absOp.getOperand(); - auto operandType = dyn_cast(operand.getType()); - if (!operandType) - return failure(); - - if (isa(operandType.getElementType())) { - // abs(z)^2 = real(z)^2 + imag(z)^2 -- only applied if abs(z) is used in - // this operation - if (!isOnlyUsedInOperation(absOp, op)) { - return failure(); - } - auto real = stablehlo::RealOp::create(rewriter, op.getLoc(), operand); - auto imag = stablehlo::ImagOp::create(rewriter, op.getLoc(), operand); - auto realSq = stablehlo::MulOp::create(rewriter, op.getLoc(), real, real); - auto imagSq = stablehlo::MulOp::create(rewriter, op.getLoc(), imag, imag); - rewriter.replaceOpWithNewOp(op, realSq, imagSq); - return success(); - } else { - // abs(x)^2 = x * x - rewriter.replaceOpWithNewOp(op, operand, operand); - return success(); - } - } -}; - struct ConcatBroadcastSlice : public CheckedOpRewritePattern { @@ -26156,24 +25353,6 @@ struct ReduceReduce final } }; -struct ConjReal final : public CheckedOpRewritePattern { - using CheckedOpRewritePattern::CheckedOpRewritePattern; - - bool supportsDynamicShapes() { return true; } - - LogicalResult matchAndRewriteImpl(chlo::ConjOp op, - PatternRewriter &rewriter) const { - auto input = op.getOperand(); - - if (guaranteedPurelyRealResult(input, rewriter)) { - rewriter.replaceOp(op, input); - return success(); - } - - return failure(); - } -}; - struct ConcatReshapeElementwise final : public CheckedOpRewritePattern { @@ -26509,78 +25688,6 @@ struct IfOpLiftCommonOps final } }; -// used for ops that dont define the Involution trait -template -struct InvolutionSimplify - : public CheckedOpRewritePattern> { - using CheckedOpRewritePattern< - OpTy, InvolutionSimplify>::CheckedOpRewritePattern; - - LogicalResult matchAndRewriteImpl(OpTy op, PatternRewriter &rewriter) const { - auto operandOp = op.getOperand().template getDefiningOp(); - if (!operandOp) - return failure(); - - rewriter.replaceOp(op, operandOp.getOperand()); - return success(); - } -}; - -struct RealConjSimplify final - : public CheckedOpRewritePattern { - using CheckedOpRewritePattern::CheckedOpRewritePattern; - - LogicalResult matchAndRewriteImpl(stablehlo::RealOp op, - PatternRewriter &rewriter) const { - auto operandOp = op.getOperand().getDefiningOp(); - if (!operandOp) - return failure(); - - rewriter.replaceOpWithNewOp(op, operandOp.getOperand()); - return success(); - } -}; - -struct RealConvertSimplify final - : public CheckedOpRewritePattern { - using CheckedOpRewritePattern::CheckedOpRewritePattern; - - LogicalResult matchAndRewriteImpl(stablehlo::RealOp op, - PatternRewriter &rewriter) const { - auto operandOp = op.getOperand().getDefiningOp(); - if (!operandOp || isa(cast( - operandOp.getOperand().getType()) - .getElementType())) { - return failure(); - } - - rewriter.replaceOp(op, operandOp.getOperand()); - return success(); - } -}; - -struct ConjComplexSimplify final - : public CheckedOpRewritePattern { - using CheckedOpRewritePattern::CheckedOpRewritePattern; - - LogicalResult matchAndRewriteImpl(chlo::ConjOp op, - PatternRewriter &rewriter) const { - auto operandOp = op.getOperand().getDefiningOp(); - if (!operandOp) - return failure(); - - auto rhs = operandOp.getRhs(); - if (!matchPattern(rhs, m_Constant())) { - return failure(); - } - - auto negateRhs = stablehlo::NegOp::create(rewriter, op.getLoc(), rhs); - rewriter.replaceOpWithNewOp(op, operandOp.getLhs(), - negateRhs); - return success(); - } -}; - // elementwise_op(all operands have zero imag) -> do the op in real domain and // convert to complex struct ElementwiseComplexSimplify final @@ -27743,234 +26850,6 @@ struct CommonAssociativeCommutativeOpReorder final } }; -struct LogSimplify final - : public CheckedOpRewritePattern { - using CheckedOpRewritePattern::CheckedOpRewritePattern; - - LogicalResult matchAndRewriteImpl(stablehlo::LogOp op, - PatternRewriter &rewriter) const { - { // log(exp(x)) -> x - auto defOp = op.getOperand().getDefiningOp(); - if (defOp) { - rewriter.replaceAllUsesWith(op.getResult(), defOp.getOperand()); - return success(); - } - } - - { // log(pow(x, y)) -> y * log(x) - auto defOp = op.getOperand().getDefiningOp(); - if (defOp) { - rewriter.replaceOpWithNewOp( - op, defOp.getRhs(), - stablehlo::LogOp::create(rewriter, op.getLoc(), defOp.getLhs())); - return success(); - } - } - - { - auto defOp = op.getOperand().getDefiningOp(); - if (defOp) { - auto lhs = defOp.getLhs(); - auto rhs = defOp.getRhs(); - if (lhs == rhs) { // log(mul(a, a)) -> 2 * log(a) - rewriter.replaceOpWithNewOp( - op, - stablehlo::ConstantOp::create( - rewriter, op.getLoc(), lhs.getType(), - cast(makeAttr(lhs.getType(), 2))), - stablehlo::LogOp::create(rewriter, op.getLoc(), lhs)); - return success(); - } - - if (anyOperandIsConstant(defOp) && - !allOperandsAreConstant(defOp)) { // log(mul(a, b)) -> log(a) + - // log(b) if a or b is constant - rewriter.replaceOpWithNewOp( - op, stablehlo::LogOp::create(rewriter, op.getLoc(), lhs), - stablehlo::LogOp::create(rewriter, op.getLoc(), rhs)); - return success(); - } - } - } - - { - auto defOp = op.getOperand().getDefiningOp(); - if (defOp) { - auto lhs = defOp.getLhs(); - auto rhs = defOp.getRhs(); - if (lhs == rhs) { // log(add(a, a)) -> log(2) + log(a) - rewriter.replaceOpWithNewOp( - op, - stablehlo::LogOp::create( - rewriter, op.getLoc(), - stablehlo::ConstantOp::create( - rewriter, op.getLoc(), lhs.getType(), - cast(makeAttr(lhs.getType(), 2)))), - stablehlo::LogOp::create(rewriter, op.getLoc(), lhs)); - return success(); - } - - Attribute rhsAttr, lhsAttr; - matchPattern(rhs, m_Constant(&rhsAttr)); - matchPattern(lhs, m_Constant(&lhsAttr)); - if (rhsAttr && - (matchPattern(rhsAttr, m_One()) || - matchPattern(rhsAttr, m_OneFloat()))) { // log(x + 1) -> log1p(x) - rewriter.replaceOpWithNewOp(op, lhs); - return success(); - } - - if (lhsAttr && - (matchPattern(lhsAttr, m_One()) || - matchPattern(lhsAttr, m_OneFloat()))) { // log(1 + x) -> log1p(x) - rewriter.replaceOpWithNewOp(op, rhs); - return success(); - } - } - } - - { - auto defOp = op.getOperand().getDefiningOp(); - if (defOp) { - auto lhs = defOp.getLhs(); - auto rhs = defOp.getRhs(); - - if (anyOperandIsConstant(defOp) && - !allOperandsAreConstant(defOp)) { // log(div(a, b)) -> log(a) - - // log(b) if a or b is constant - rewriter.replaceOpWithNewOp( - op, stablehlo::LogOp::create(rewriter, op.getLoc(), lhs), - stablehlo::LogOp::create(rewriter, op.getLoc(), rhs)); - return success(); - } - } - } - - { - auto defOp = op.getOperand().getDefiningOp(); - if (defOp && - isOnlyUsedInOperation(defOp, op)) { // log(sqrt(x)) -> log(x) / 2 - rewriter.replaceOpWithNewOp( - op, - stablehlo::LogOp::create(rewriter, op.getLoc(), defOp.getOperand()), - stablehlo::ConstantOp::create( - rewriter, op.getLoc(), defOp.getType(), - cast(makeAttr(defOp.getType(), 2)))); - return success(); - } - } - - { - auto defOp = op.getOperand().getDefiningOp(); - if (defOp && - isOnlyUsedInOperation(defOp, op)) { // log(cbrt(x)) -> log(x) / 3 - rewriter.replaceOpWithNewOp( - op, - stablehlo::LogOp::create(rewriter, op.getLoc(), defOp.getOperand()), - stablehlo::ConstantOp::create( - rewriter, op.getLoc(), defOp.getType(), - cast(makeAttr(defOp.getType(), 3)))); - return success(); - } - } - - { - auto defOp = op.getOperand().getDefiningOp(); - if (defOp && - isOnlyUsedInOperation(defOp, op)) { // log(rsqrt(x)) -> -log(x) / 2 - rewriter.replaceOpWithNewOp( - op, - stablehlo::LogOp::create(rewriter, op.getLoc(), defOp.getOperand()), - stablehlo::ConstantOp::create( - rewriter, op.getLoc(), defOp.getType(), - cast(makeAttr(defOp.getType(), -2)))); - return success(); - } - } - - return failure(); - } -}; - -struct NegMulConstSimplify final - : public CheckedOpRewritePattern { - using CheckedOpRewritePattern::CheckedOpRewritePattern; - - LogicalResult matchAndRewriteImpl(stablehlo::NegOp op, - PatternRewriter &rewriter) const { - auto mulOp = op.getOperand().getDefiningOp(); - if (!mulOp) - return failure(); - - auto lhs = mulOp.getLhs(); - auto rhs = mulOp.getRhs(); - - DenseElementsAttr lhsAttr; - bool lhsIsConst = matchPattern(lhs, m_Constant(&lhsAttr)); - - DenseElementsAttr rhsAttr; - bool rhsIsConst = matchPattern(rhs, m_Constant(&rhsAttr)); - - if (lhsIsConst && rhsIsConst) - return failure(); // const prop will evaluate this - - if (lhsIsConst) { - rewriter.replaceOpWithNewOp( - op, stablehlo::NegOp::create(rewriter, op.getLoc(), lhs), rhs); - return success(); - } - - if (rhsIsConst) { - rewriter.replaceOpWithNewOp( - op, lhs, stablehlo::NegOp::create(rewriter, op.getLoc(), rhs)); - return success(); - } - - return failure(); - } -}; - -struct NegDivConstSimplify final - : public CheckedOpRewritePattern { - using CheckedOpRewritePattern::CheckedOpRewritePattern; - - LogicalResult matchAndRewriteImpl(stablehlo::NegOp op, - PatternRewriter &rewriter) const { - auto divOp = op.getOperand().getDefiningOp(); - if (!divOp) - return failure(); - - auto lhs = divOp.getLhs(); - auto rhs = divOp.getRhs(); - - DenseElementsAttr lhsAttr; - bool lhsIsConst = matchPattern(lhs, m_Constant(&lhsAttr)); - - DenseElementsAttr rhsAttr; - bool rhsIsConst = matchPattern(rhs, m_Constant(&rhsAttr)); - - if (lhsIsConst && rhsIsConst) - return failure(); // const prop will evaluate this - - if (lhsIsConst) { - rewriter.replaceOpWithNewOp( - op, stablehlo::NegOp::create(rewriter, op.getLoc(), lhs), rhs); - return success(); - } - - if (rhsIsConst) { - rewriter.replaceOpWithNewOp( - op, lhs, stablehlo::NegOp::create(rewriter, op.getLoc(), rhs)); - return success(); - } - - return failure(); - } -}; - struct ReshapeDeletionsBroadcastInDimSimplify final : public CheckedOpRewritePattern { @@ -35751,6 +34630,31 @@ static bool divMulReassociable(mlir::Value res, mlir::Attribute aAttr, return true; } +static mlir::Value createReciprocalConstantOp(mlir::PatternRewriter &rewriter, + mlir::Location loc, + mlir::Attribute attr, + mlir::Type type) { + auto denseAttr = llvm::dyn_cast_or_null(attr); + assert(denseAttr && "expected DenseElementsAttr"); + + auto ty = llvm::cast(type); + if (denseAttr.isSplat()) { + ty = mlir::RankedTensorType::get({}, ty.getElementType()); + denseAttr = denseAttr.resizeSplat(ty); + } + + auto rhsTen = stablehlo::constantOp(denseAttr); + auto oneTen = + stablehlo::constantOp(llvm::cast(makeAttr(ty, 1))); + auto out = fromTensor(stablehlo::divideOp(oneTen, rhsTen, ty)); + + if (ty != type) { + out = out.resizeSplat(llvm::cast(type)); + } + + return stablehlo::ConstantOp::create(rewriter, loc, type, out); +} + // clang-format off #include "src/enzyme_ad/jax/Passes/StablehloOptPatterns.cpp.inc" @@ -36158,34 +35062,40 @@ struct EnzymeHLOOptPass patterns.add(context); patterns.add(context); + patterns.add(context); + patterns.add< - AddSimplify, SubSimplify, AndSimplify, MaxSimplify, MinSimplify, - OrSimplify, XorSimplify, MulSimplify, DivSimplify, RemSimplify, - PowSimplify, NoopSlice, NoopReverse, SliceReverse, SliceSlice, + NoopSlice, NoopReverse, SliceReverse, SliceSlice, DynamicSliceDynamicSlice, DynamicSliceSlice, SliceDynamicSlice, - LogSimplify, ShiftRightLogicalSimplify, NegativePadToSlice, - SliceSimplify, ConvertSimplify, TransposeSimplify, DotGeneralSimplify, + ShiftRightLogicalSimplify, NegativePadToSlice, SliceSimplify, + ConvertSimplify, TransposeSimplify, DotGeneralSimplify, DotGeneralReshape, DiagonalTensorDotGeneralRewrite, DynamicSliceToStatic, DynamicUpdateSliceElim, ReduceToReshape, BroadcastToReshape, ReshapeEmptyBroadcast, ReshapeBroadcast, - BroadcastReshape, ConstPropThroughBarrier, ReplaceNegAddWithSubtract, - ReplaceSubtractNegWithAdd, SignAbsSimplify, AbsPositiveSimplify, + BroadcastReshape, ConstPropThroughBarrier, SimplifyBoundary, SimplifyBoundary, SimplifyBoundary, TransposeReshapeToBroadcast, ReshapeTransposeToBroadcast, SelectBroadcastInDim, PowerMultiplyToPower, - NegMulConstSimplify, NegDivConstSimplify, NegatedConstantMulFactoring, NegatedConstantMulFactoring, ReshapeDeletionsBroadcastInDimSimplify, ReshapeInsertionsBroadcastInDimSimplify, CompareIotaConstSimplify, - ConvertIotaSimplify, MinMaxIotaConstSimplify, + MinMaxIotaConstSimplify, MinMaxIotaConstSimplify, ClampIotaConstSimplify, CompareAbs, CompareMul, CompareConvert, AddSelects, CompareNegateConstSimplify, CompareSubtractConstSimplify, SelectSimplify, DynamicSliceReshapeDynamicSlice, - DynamicSliceReshapeSlice, SliceReshapeDynamicSlice, SliceReshapeSlice, - ExponentialMinusOneFuse, ExponentialMinusOneAddFuse>( + DynamicSliceReshapeSlice, SliceReshapeDynamicSlice, SliceReshapeSlice>( context, PatternBenefit(65000)); patterns.add, BinBroadcastSplat, BinBroadcastSplat, - BinBroadcastSplat, RotatePad, ConjReal, - ConvertMulConvert, ConvertBinopConvert, + BinBroadcastSplat, RotatePad, ConvertMulConvert, + ConvertBinopConvert, ConvertBinopConvert, NegateReduceWindowSub>(context); // Unary constant propagation patterns @@ -36271,12 +35181,13 @@ struct EnzymeHLOOptPass if (passses & 512) { patterns.add(context); + ConcatToPad, ConcatAppendingReshape, ReshapeIota, DUSDUS, + DUSDUSConcat, DUSConcat, DUSPad, DUSDUSSubsuming, + SliceDUSToConcat, ConcatConcatToDUS>(context); patterns.add, LICM>(false, context); + patterns.add(context); } if (passses & 1024) @@ -36467,14 +35378,11 @@ struct EnzymeHLOOptPass ChainedDynamicBroadcastInDimCanonicalization, CompareOpCanon, CompareExt, - ConjComplexNegate, - NegateImagConj, ConvertOpCanon, DivideSqrtToMultiplyRsqrt, DynamicBroadcastInDimAllDimsNonExpanding, DynamicBroadcastInDimOpNotActuallyDynamic, DynamicGatherOpIsNotDynamic, - DynamicReshapeOpCanon, EmptyReduceOpCanon, GatherOpCanon, ScatterOpCanon, @@ -36493,7 +35401,6 @@ struct EnzymeHLOOptPass SelectCompIotaConstToDUS, SelectCompIotaConstSimplify, SelectPadToDUS, - SelectSelectSameCond, SelectSelectNegCond, AndPadPad, SelectOpUsedWithinIf, @@ -36505,7 +35412,6 @@ struct EnzymeHLOOptPass WhileDeadResults, ZeroExtentTensorCanon, CompareSelectSimplify, - NotSelectSimplify, CommonCompareExpressionRewrite, ScatterUpdateComputationConstProp, ScatterIndicesAreUnique, @@ -36516,6 +35422,7 @@ struct EnzymeHLOOptPass NotCompare, SliceInternal, SquareAbsSimplify, + SquareAbsSimplifyComplex, DivideDivideSimplify, ConcatReshapeSlice, ConcatBroadcastSlice, @@ -36525,12 +35432,6 @@ struct EnzymeHLOOptPass TransposeAllUsersSlice, ReduceReduce, IfOpLiftCommonOps, - InvolutionSimplify, - InvolutionSimplify, - InvolutionSimplify, - RealConjSimplify, - RealConvertSimplify, - ConjComplexSimplify, ElementwiseComplexSimplify, SplitConvolutionIntoReverseConvolution, ScatterMultiplySimplify, @@ -36613,6 +35514,9 @@ struct EnzymeHLOOptPass PatternBenefit(65000)); patterns.add(max_constant_expansion, context, PatternBenefit(65000)); + patterns.add(context); if (enable_auto_batching_passes) { mlir::enzyme::AutoBatchingPassPipelineOptions options{ @@ -36625,8 +35529,21 @@ struct EnzymeHLOOptPass config.setMaxIterations(max_iterations); config.setUseTopDownTraversal(top_down); config.enableFolding(); - if (failed(applyPatternsGreedily(getOperation(), std::move(patterns), - config))) { + + FrozenRewritePatternSet frozenPatterns(std::move(patterns)); + + auto walkResult = + getOperation()->walk([&](FunctionOpInterface func) -> WalkResult { + if (func->hasAttr(kDisablePatternAttrName)) { + return WalkResult::advance(); + } + if (failed(applyPatternsGreedily(func, frozenPatterns, config))) { + return WalkResult::interrupt(); + } + return WalkResult::advance(); + }); + + if (walkResult.wasInterrupted()) { signalPassFailure(); } } diff --git a/src/enzyme_ad/jax/Passes/LowerEnzymeXLAMath.cpp b/src/enzyme_ad/jax/Passes/LowerEnzymeXLAMath.cpp index 8065d73608..b4015b1c3c 100644 --- a/src/enzyme_ad/jax/Passes/LowerEnzymeXLAMath.cpp +++ b/src/enzyme_ad/jax/Passes/LowerEnzymeXLAMath.cpp @@ -30,15 +30,6 @@ using namespace mlir; using namespace mlir::enzyme; using namespace mlir::stablehlo; -template -static stablehlo::ConstantOp -createConstantOpFromScalar(PatternRewriter &rewriter, Location loc, Type type, - T value) { - return stablehlo::ConstantOp::create( - rewriter, loc, type, - cast(mlir::enzyme::makeAttr(type, value))); -} - namespace { #include "src/enzyme_ad/jax/Passes/LowerEnzymeXLAMathPatterns.cpp.inc" diff --git a/src/enzyme_ad/jax/Passes/StablehloOptPatterns.td b/src/enzyme_ad/jax/Passes/StablehloOptPatterns.td index 61394e843d..7389096354 100644 --- a/src/enzyme_ad/jax/Passes/StablehloOptPatterns.td +++ b/src/enzyme_ad/jax/Passes/StablehloOptPatterns.td @@ -8,6 +8,7 @@ include "mlir/IR/OpBase.td" include "stablehlo/dialect/StablehloOps.td" +include "stablehlo/dialect/ChloOps.td" //////// // Mul/Div || Add/Sub - constant computation lifting @@ -79,6 +80,7 @@ def SubSubConst : BinOpLiftConstantComputation; def HasSameType : Constraint>; +def HasDifferentType : Constraint>; def UseOperand : NativeCodeCall<"$0">; def BitcastConvertCancellation : Pat< @@ -86,3 +88,590 @@ def BitcastConvertCancellation : Pat< (UseOperand $x), [(HasSameType $outer, $x)] >; + +// ConvertConvertFloat +def ConvertConvertFloatIdentity : Pat< + (StableHLO_ConvertOp:$outer (StableHLO_ConvertOp:$inner $x)), + (UseOperand $x), + [(IsFloatTensor $outer), (IsFloatTensor $inner), (IsFloatTensor $x), (HasSameType $outer, $x)]>; + +def ConvertConvertFloat : Pat< + (StableHLO_ConvertOp:$outer (StableHLO_ConvertOp:$inner $x)), + (StableHLO_ConvertOp $x), + [(IsFloatTensor $outer), (IsFloatTensor $inner), (IsFloatTensor $x), (HasDifferentType $outer, $x)]>; + +// ConvertConvertInt +def IsIntegerTensor : Constraint(cast($0.getType()).getElementType())">, + "tensor has an integer element type">; + +def IsWideningConvert : Constraint($0.getType()).getElementType().getIntOrFloatBitWidth() <= " + "cast($1.getType()).getElementType().getIntOrFloatBitWidth()" +>, "widening conversion">; + +def ConvertConvertIntIdentity : Pat< + (StableHLO_ConvertOp:$outer (StableHLO_ConvertOp:$inner $x)), + (UseOperand $x), + [(IsIntegerTensor $outer), (IsIntegerTensor $inner), (IsIntegerTensor $x), + (IsWideningConvert $x, $inner), (HasSameType $outer, $x)]>; + +def ConvertConvertInt : Pat< + (StableHLO_ConvertOp:$outer (StableHLO_ConvertOp:$inner $x)), + (StableHLO_ConvertOp $x), + [(IsIntegerTensor $outer), (IsIntegerTensor $inner), (IsIntegerTensor $x), + (IsWideningConvert $x, $inner), (HasDifferentType $outer, $x)]>; + +// NotSelectSimplify +def NotSelectSimplify : Pat< + (StableHLO_SelectOp (StableHLO_NotOp $cond), $lhs, $rhs), + (StableHLO_SelectOp $cond, $rhs, $lhs)>; + +// SelectSelectSameCond +def SelectSelectSameCondFalse : Pat< + (StableHLO_SelectOp $cond, $c, (StableHLO_SelectOp $cond, $a, $b)), + (StableHLO_SelectOp $cond, $c, $b)>; + +def SelectSelectSameCondTrue : Pat< + (StableHLO_SelectOp $cond, (StableHLO_SelectOp $cond, $a, $b), $c), + (StableHLO_SelectOp $cond, $a, $c)>; + +// RealConjSimplify +def RealConjSimplify : Pat< + (StableHLO_RealOp (CHLO_ConjOp $x)), + (StableHLO_RealOp $x)>; + +// RealConvertSimplify +def IsNotComplexTensor : Constraint(cast($0.getType()).getElementType())">, + "tensor has a non-complex element type">; + +def RealConvertSimplify : Pat< + (StableHLO_RealOp (StableHLO_ConvertOp $x)), + (UseOperand $x), + [(IsNotComplexTensor $x)]>; + +// ConjComplexNegate +def ConjComplexNegate : Pat< + (CHLO_ConjOp (StableHLO_ComplexOp $a, (StableHLO_NegOp $b))), + (StableHLO_ComplexOp $a, $b)>; + +// NegateImagConj +def NegateImagConj : Pat< + (StableHLO_NegOp (StableHLO_ImagOp (CHLO_ConjOp $x))), + (StableHLO_ImagOp $x)>; + +def IsAnyZero : Constraint>; + +def IsOne : Constraint>; + +def IsNegOne : Constraint>; + +def IsAllOnes : Constraint>; + +def IsMinusOneValue : Constraint()); " + " return doubleVal && *doubleVal == -1.0; " + "}($0)" +>>; + +def HasOneUse : Constraint>; + + +def GuaranteedPurelyReal : Constraint, + "result is guaranteed purely real">; + +def GuaranteedNonNegative : Constraint, + "result is guaranteed non-negative">; + +def IsConstantOp : Constraint>; + +def IsStaticShapeTensor : Constraint($0.getType()).hasStaticShape()">, + "tensor has a static shape">; + +def CreateConstOp0 : NativeCodeCall<"createConstantOpFromScalar($_builder, $_loc, $0.getType(), 0.0)">; +def CreateConstOp1 : NativeCodeCall<"createConstantOpFromScalar($_builder, $_loc, $0.getType(), 1.0)">; + +// InvolutionSimplify +def NegNegSimplify : Pat< + (StableHLO_NegOp (StableHLO_NegOp $x)), + (UseOperand $x)>; + +def NotNotSimplify : Pat< + (StableHLO_NotOp (StableHLO_NotOp $x)), + (UseOperand $x)>; + +def ConjConjSimplify : Pat< + (CHLO_ConjOp (CHLO_ConjOp $x)), + (UseOperand $x)>; + +// SignAbsSimplify +def SignAbsSimplify_1 : Pat< + (StableHLO_MulOp (StableHLO_SignOp $x), (StableHLO_AbsOp $x)), + (UseOperand $x)>; + +def SignAbsSimplify_2 : Pat< + (StableHLO_MulOp (StableHLO_AbsOp $x), (StableHLO_SignOp $x)), + (UseOperand $x)>; + +// ReplaceNegAddWithSubtract +def ReplaceNegAddWithSubtractRHS : Pat< + (StableHLO_AddOp $lhs, (StableHLO_NegOp:$neg $rhs)), + (StableHLO_SubtractOp $lhs, $rhs), + [(HasOneUse $neg)]>; + +def ReplaceNegAddWithSubtractLHS : Pat< + (StableHLO_AddOp (StableHLO_NegOp:$neg $lhs), $rhs), + (StableHLO_SubtractOp $rhs, $lhs), + [(HasOneUse $neg)]>; + +// ReplaceSubtractNegWithAdd +def ReplaceSubtractNegWithAdd : Pat< + (StableHLO_SubtractOp $lhs, (StableHLO_NegOp:$neg $rhs)), + (StableHLO_AddOp $lhs, $rhs), + [(HasOneUse $neg)]>; + +// ExponentialMinusOneFuse +def ExponentialMinusOneFuse1 : Pat< + (StableHLO_SubtractOp (StableHLO_ExpOp:$exp $x, $accuracy), $one), + (StableHLO_Expm1Op $x, $accuracy), + [(HasOneUse $exp), (IsOne $one)]>; + +def ExponentialMinusOneFuse2 : Pat< + (StableHLO_SubtractOp $one, (StableHLO_ExpOp:$exp $x, $accuracy)), + (StableHLO_NegOp (StableHLO_Expm1Op $x, $accuracy)), + [(HasOneUse $exp), (IsOne $one)]>; + +// ExponentialMinusOneAddFuse +def ExponentialMinusOneAddFuse1 : Pat< + (StableHLO_AddOp (StableHLO_ExpOp:$exp $x, $accuracy), $minus_one), + (StableHLO_Expm1Op $x, $accuracy), + [(HasOneUse $exp), (IsMinusOneValue $minus_one)]>; + +def ExponentialMinusOneAddFuse2 : Pat< + (StableHLO_AddOp $minus_one, (StableHLO_ExpOp:$exp $x, $accuracy)), + (StableHLO_Expm1Op $x, $accuracy), + [(HasOneUse $exp), (IsMinusOneValue $minus_one)]>; + +// MinSimplify / MaxSimplify +def MinSimplify : Pat< + (StableHLO_MinOp $x, $x), + (UseOperand $x)>; + +def MaxSimplify : Pat< + (StableHLO_MaxOp $x, $x), + (UseOperand $x)>; + +// AddSimplify +def AddSimplifyLHS : Pat< + (StableHLO_AddOp $zero, $x), + (UseOperand $x), + [(IsAnyZero $zero)]>; + +def AddSimplifyRHS : Pat< + (StableHLO_AddOp $x, $zero), + (UseOperand $x), + [(IsAnyZero $zero)]>; + +// SubSimplify +def SubSimplifyRHSZero : Pat< + (StableHLO_SubtractOp $x, $zero), + (UseOperand $x), + [(IsAnyZero $zero)]>; + +def SubSimplifyLHSZero : Pat< + (StableHLO_SubtractOp $zero, $x), + (StableHLO_NegOp $x), + [(IsAnyZero $zero)]>; + +def SubSimplifyIdentical : Pat< + (StableHLO_SubtractOp:$res $x, $x), + (CreateConstOp0 $res), + [(IsIntegerTensor $res)]>; + +// MulSimplify +def MulSimplifyOneLHS : Pat< + (StableHLO_MulOp $one, $x), + (UseOperand $x), + [(IsOne $one)]>; + +def MulSimplifyOneRHS : Pat< + (StableHLO_MulOp $x, $one), + (UseOperand $x), + [(IsOne $one)]>; + +def MulSimplifyNegOneLHS : Pat< + (StableHLO_MulOp $negone, $x), + (StableHLO_NegOp $x), + [(IsNegOne $negone)]>; + +def MulSimplifyNegOneRHS : Pat< + (StableHLO_MulOp $x, $negone), + (StableHLO_NegOp $x), + [(IsNegOne $negone)]>; + +// AndSimplify +def AndSimplifyIdentical : Pat< + (StableHLO_AndOp $x, $x), + (UseOperand $x)>; + +def AndSimplifyFalseLHS : Pat< + (StableHLO_AndOp $zero, $x), + (UseOperand $zero), + [(IsAnyZero $zero)]>; + +def AndSimplifyFalseRHS : Pat< + (StableHLO_AndOp $x, $zero), + (UseOperand $zero), + [(IsAnyZero $zero)]>; + +def AndSimplifyTrueLHS : Pat< + (StableHLO_AndOp $ones, $x), + (UseOperand $x), + [(IsAllOnes $ones)]>; + +def AndSimplifyTrueRHS : Pat< + (StableHLO_AndOp $x, $ones), + (UseOperand $x), + [(IsAllOnes $ones)]>; + +// OrSimplify +def OrSimplifyIdentical : Pat< + (StableHLO_OrOp $x, $x), + (UseOperand $x)>; + +def OrSimplifyTrueLHS : Pat< + (StableHLO_OrOp $ones, $x), + (UseOperand $ones), + [(IsAllOnes $ones)]>; + +def OrSimplifyTrueRHS : Pat< + (StableHLO_OrOp $x, $ones), + (UseOperand $ones), + [(IsAllOnes $ones)]>; + +def OrSimplifyFalseLHS : Pat< + (StableHLO_OrOp $zero, $x), + (UseOperand $x), + [(IsAnyZero $zero)]>; + +def OrSimplifyFalseRHS : Pat< + (StableHLO_OrOp $x, $zero), + (UseOperand $x), + [(IsAnyZero $zero)]>; + +// XorSimplify +def XorSimplifyFalseLHS : Pat< + (StableHLO_XorOp $zero, $x), + (UseOperand $x), + [(IsAnyZero $zero)]>; + +def XorSimplifyFalseRHS : Pat< + (StableHLO_XorOp $x, $zero), + (UseOperand $x), + [(IsAnyZero $zero)]>; + +def XorSimplifyTrueLHS : Pat< + (StableHLO_XorOp $ones, $x), + (StableHLO_NotOp $x), + [(IsAllOnes $ones)]>; + +def XorSimplifyTrueRHS : Pat< + (StableHLO_XorOp $x, $ones), + (StableHLO_NotOp $x), + [(IsAllOnes $ones)]>; + +// ConvertIotaSimplify +def ConvertIotaSimplify : Pat< + (StableHLO_ConvertOp:$res (StableHLO_IotaOp $dim)), + (StableHLO_IotaOp $dim), + [(IsIntegerTensor $res)]>; + +// ConjReal +def ConjRealSimplify : Pat< + (CHLO_ConjOp $x), + (UseOperand $x), + [(GuaranteedPurelyReal $x)], + [], + (addBenefit 10)>; + +// ConjComplexSimplify +def ConjComplexSimplify : Pat< + (CHLO_ConjOp (StableHLO_ComplexOp $lhs, $rhs)), + (StableHLO_ComplexOp $lhs, (StableHLO_NegOp $rhs)), + [(IsConstantOp $rhs)]>; + +// DynamicReshapeOpCanon +def DynamicReshapeOpCanon : Pat< + (StableHLO_DynamicReshapeOp:$res $operand, $shape), + (StableHLO_ReshapeOp $operand), + [(IsStaticShapeTensor $res)]>; + +// AbsPositiveSimplify +def AbsPositiveSimplify : Pat< + (StableHLO_AbsOp:$res $x), + (UseOperand $x), + [(IsNotComplexTensor $res), (GuaranteedNonNegative $x)]>; + +// RemSimplify +def RemSimplify : Pat< + (StableHLO_RemOp:$res $x, $one), + (CreateConstOp0 $res), + [(IsOne $one)]>; + +// PowSimplify +def PowSimplifyOneRHS : Pat< + (StableHLO_PowOp $x, $one), + (UseOperand $x), + [(IsOne $one)], + [], + (addBenefit 65000)>; + +def PowSimplifyZeroRHS : Pat< + (StableHLO_PowOp:$res $x, $zero), + (CreateConstOp1 $res), + [(IsAnyZero $zero)], + [], + (addBenefit 65000)>; + +def PowSimplifyOneLHS : Pat< + (StableHLO_PowOp:$res $one, $x), + (CreateConstOp1 $res), + [(IsOne $one)], + [], + (addBenefit 65000)>; + +// Constraints and helper functions for migrated patterns +def ConstDefaultResultAccuracyAttr : ConstantAttr; + +def IsComplexTensor : Constraint(cast($0.getType()).getElementType())">, + "tensor has a complex element type">; + +def IsNotConstantOp : Constraint>; + +def IsOnlyUsedInOperation : Constraint>; + +def IsNegHalf : Constraint()); " + " return doubleVal && *doubleVal == -0.5; " + "}($0)" +>>; + +def IsHalf : Constraint()); " + " return doubleVal && *doubleVal == 0.5; " + "}($0)" +>>; + +def IsTwo : Constraint()); " + " return doubleVal && *doubleVal == 2.0; " + "}($0)" +>>; + +def CreateConstOp2 : NativeCodeCall<"createConstantOpFromScalar($_builder, $_loc, $0.getType(), 2.0)">; +def CreateConstOp3 : NativeCodeCall<"createConstantOpFromScalar($_builder, $_loc, $0.getType(), 3.0)">; +def CreateConstOpNeg2 : NativeCodeCall<"createConstantOpFromScalar($_builder, $_loc, $0.getType(), -2.0)">; + +def AnyOperandIsConstant : Constraint>; +def NotAllOperandsAreConstant : Constraint>; + +def IsEqual : Constraint>; + +// SquareAbsSimplify +def SquareAbsSimplify : Pat< + (StableHLO_MulOp (StableHLO_AbsOp:$absLhs $x), (StableHLO_AbsOp:$absRhs $y)), + (StableHLO_MulOp $x, $x), + [(IsNotComplexTensor $x), (IsEqual $absLhs, $absRhs)]>; + +def SquareAbsSimplifyComplex : Pat< + (StableHLO_MulOp:$res (StableHLO_AbsOp:$absLhs $x), (StableHLO_AbsOp:$absRhs $y)), + (StableHLO_AddOp (StableHLO_MulOp (StableHLO_RealOp $x), (StableHLO_RealOp $x)), (StableHLO_MulOp (StableHLO_ImagOp $x), (StableHLO_ImagOp $x))), + [(IsComplexTensor $x), (IsEqual $absLhs, $absRhs), (IsOnlyUsedInOperation $absLhs, $res), (IsOnlyUsedInOperation $absRhs, $res)]>; + +// DivSimplify simple patterns +// DivSimplify simple patterns +def DivSimplifyOneRHS : Pat< + (StableHLO_DivOp $x, $one), + (UseOperand $x), + [(IsOne $one)], + [], + (addBenefit 65000)>; + +def DivSimplifyNegOneRHS : Pat< + (StableHLO_DivOp $x, $negone), + (StableHLO_NegOp $x), + [(IsNegOne $negone)], + [], + (addBenefit 65000)>; + +// PowSimplify extra patterns +def PowSimplifyNegOneRHS : Pat< + (StableHLO_PowOp:$res $x, $negone), + (StableHLO_DivOp (CreateConstOp1 $res), $x), + [(IsNegOne $negone), (IsFloatTensor $res)], + [], + (addBenefit 65000)>; + +def PowSimplifyNegHalfRHS : Pat< + (StableHLO_PowOp:$res $x, $neghalf), + (StableHLO_RsqrtOp $x, ConstDefaultResultAccuracyAttr), + [(IsNegHalf $neghalf), (IsFloatTensor $res)], + [], + (addBenefit 65000)>; + +def PowSimplifyHalfRHS : Pat< + (StableHLO_PowOp:$res $x, $half), + (StableHLO_SqrtOp $x, ConstDefaultResultAccuracyAttr), + [(IsHalf $half), (IsFloatTensor $res)], + [], + (addBenefit 65000)>; + +def PowSimplifyTwoRHS : Pat< + (StableHLO_PowOp:$res $x, $two), + (StableHLO_MulOp $x, $x), + [(IsTwo $two), (IsFloatTensor $res)], + [], + (addBenefit 65000)>; + +// LogSimplify patterns +def LogExpSimplify : Pat< + (StableHLO_LogOp (StableHLO_ExpOp $x, $exp_acc), $log_acc), + (UseOperand $x), + [], + [], + (addBenefit 65000)>; + +def LogPowSimplify : Pat< + (StableHLO_LogOp (StableHLO_PowOp $x, $y), $log_acc), + (StableHLO_MulOp $y, (StableHLO_LogOp $x, $log_acc)), + [], + [], + (addBenefit 65000)>; + +def LogMulIdenticalSimplify : Pat< + (StableHLO_LogOp:$res (StableHLO_MulOp $x, $x), $log_acc), + (StableHLO_MulOp (CreateConstOp2 $res), (StableHLO_LogOp $x, $log_acc)), + [], + [], + (addBenefit 65000)>; + +def LogMulConstSimplify : Pat< + (StableHLO_LogOp (StableHLO_MulOp:$mul $a, $b), $log_acc), + (StableHLO_AddOp (StableHLO_LogOp $a, $log_acc), (StableHLO_LogOp $b, $log_acc)), + [(AnyOperandIsConstant $mul), (NotAllOperandsAreConstant $mul)], + [], + (addBenefit 65000)>; + +def LogAddIdenticalSimplify : Pat< + (StableHLO_LogOp:$res (StableHLO_AddOp $x, $x), $log_acc), + (StableHLO_AddOp (StableHLO_LogOp (CreateConstOp2 $res), $log_acc), (StableHLO_LogOp $x, $log_acc)), + [], + [], + (addBenefit 65000)>; + +def LogAddOneRHSSimplify : Pat< + (StableHLO_LogOp (StableHLO_AddOp $x, $one), $log_acc), + (StableHLO_Log1pOp $x, $log_acc), + [(IsOne $one)], + [], + (addBenefit 65000)>; + +def LogAddOneLHSSimplify : Pat< + (StableHLO_LogOp (StableHLO_AddOp $one, $x), $log_acc), + (StableHLO_Log1pOp $x, $log_acc), + [(IsOne $one)], + [], + (addBenefit 65000)>; + +def LogDivConstSimplify : Pat< + (StableHLO_LogOp (StableHLO_DivOp:$div $a, $b), $log_acc), + (StableHLO_SubtractOp (StableHLO_LogOp $a, $log_acc), (StableHLO_LogOp $b, $log_acc)), + [(AnyOperandIsConstant $div), (NotAllOperandsAreConstant $div)], + [], + (addBenefit 65000)>; + +def LogSqrtSimplify : Pat< + (StableHLO_LogOp:$res (StableHLO_SqrtOp:$sqrt $x, $sqrt_acc), $log_acc), + (StableHLO_DivOp (StableHLO_LogOp $x, $log_acc), (CreateConstOp2 $res)), + [(IsOnlyUsedInOperation $sqrt, $res)], + [], + (addBenefit 65000)>; + +def LogCbrtSimplify : Pat< + (StableHLO_LogOp:$res (StableHLO_CbrtOp:$cbrt $x, $cbrt_acc), $log_acc), + (StableHLO_DivOp (StableHLO_LogOp $x, $log_acc), (CreateConstOp3 $res)), + [(IsOnlyUsedInOperation $cbrt, $res)], + [], + (addBenefit 65000)>; + +def LogRsqrtSimplify : Pat< + (StableHLO_LogOp:$res (StableHLO_RsqrtOp:$rsqrt $x, $rsqrt_acc), $log_acc), + (StableHLO_DivOp (StableHLO_LogOp $x, $log_acc), (CreateConstOpNeg2 $res)), + [(IsOnlyUsedInOperation $rsqrt, $res)], + [], + (addBenefit 65000)>; + +// NegMulConstSimplify +def NegMulConstLHSSimplify : Pat< + (StableHLO_NegOp (StableHLO_MulOp $lhs, $rhs)), + (StableHLO_MulOp (StableHLO_NegOp $lhs), $rhs), + [(IsConstantOp $lhs), (IsNotConstantOp $rhs)], + [], + (addBenefit 65000)>; + +def NegMulConstRHSSimplify : Pat< + (StableHLO_NegOp (StableHLO_MulOp $lhs, $rhs)), + (StableHLO_MulOp $lhs, (StableHLO_NegOp $rhs)), + [(IsNotConstantOp $lhs), (IsConstantOp $rhs)], + [], + (addBenefit 65000)>; + +// NegDivConstSimplify +def NegDivConstLHSSimplify : Pat< + (StableHLO_NegOp (StableHLO_DivOp $lhs, $rhs)), + (StableHLO_DivOp (StableHLO_NegOp $lhs), $rhs), + [(IsConstantOp $lhs), (IsNotConstantOp $rhs)], + [], + (addBenefit 65000)>; + +def NegDivConstRHSSimplify : Pat< + (StableHLO_NegOp (StableHLO_DivOp $lhs, $rhs)), + (StableHLO_DivOp $lhs, (StableHLO_NegOp $rhs)), + [(IsNotConstantOp $lhs), (IsConstantOp $rhs)], + [], + (addBenefit 65000)>; + +// DivConstToMulReciprocal +def IsDenseElementsAttr : Constraint($0)">>; + +def CreateReciprocalConstOp : NativeCodeCall< + "::mlir::enzyme::createReciprocalConstantOp($_builder, $_loc, $0, $1.getType())" +>; + +def DivConstToMulReciprocal : Pat< + (StableHLO_DivOp:$op $x, (StableHLO_ConstantOp $value)), + (StableHLO_MulOp $x, (CreateReciprocalConstOp $value, $op)), + [(IsFloatTensor $op), (IsNotConstantOp $x), (IsDenseElementsAttr $value)]>; + diff --git a/src/enzyme_ad/jax/TransformOps/TransformOps.td b/src/enzyme_ad/jax/TransformOps/TransformOps.td index 757838648b..262fa26e05 100644 --- a/src/enzyme_ad/jax/TransformOps/TransformOps.td +++ b/src/enzyme_ad/jax/TransformOps/TransformOps.td @@ -43,11 +43,11 @@ class EnzymeHLOParameterizedPatternOp traits = []> // benefit 65k def ApplyAddSimplifyPatterns : EnzymeHLOPatternOp< "add_simplify"> { - let patterns = ["AddSimplify"]; + let patterns = ["AddSimplifyLHS", "AddSimplifyRHS"]; } def ApplyReplaceNegAddWithSubtract: EnzymeHLOPatternOp< "replace_neg_add_with_subtract"> { - let patterns = ["ReplaceNegAddWithSubtract"]; + let patterns = ["ReplaceNegAddWithSubtractRHS", "ReplaceNegAddWithSubtractLHS"]; } def ApplyReplaceSubtractNegWithAdd: EnzymeHLOPatternOp< "replace_subtract_neg_with_add"> { @@ -55,11 +55,11 @@ def ApplyReplaceSubtractNegWithAdd: EnzymeHLOPatternOp< } def ApplySubSimplifyPatterns : EnzymeHLOPatternOp< "sub_simplify"> { - let patterns = ["SubSimplify"]; + let patterns = ["SubSimplifyRHSZero", "SubSimplifyLHSZero", "SubSimplifyIdentical"]; } def ApplyAndSimplifyPatterns : EnzymeHLOPatternOp< "and_simplify"> { - let patterns = ["AndSimplify"]; + let patterns = ["AndSimplifyIdentical", "AndSimplifyFalseLHS", "AndSimplifyFalseRHS", "AndSimplifyTrueLHS", "AndSimplifyTrueRHS"]; } def ApplyMaxSimplifyPatterns : EnzymeHLOPatternOp< "max_simplify"> { @@ -71,19 +71,19 @@ def ApplyMinSimplifyPatterns : EnzymeHLOPatternOp< } def ApplyOrSimplifyPatterns : EnzymeHLOPatternOp< "or_simplify"> { - let patterns = ["OrSimplify"]; + let patterns = ["OrSimplifyIdentical", "OrSimplifyTrueLHS", "OrSimplifyTrueRHS", "OrSimplifyFalseLHS", "OrSimplifyFalseRHS"]; } def ApplyMulSimplifyPatterns : EnzymeHLOPatternOp< "mul_simplify"> { - let patterns = ["MulSimplify"]; + let patterns = ["MulSimplifyOneLHS", "MulSimplifyOneRHS", "MulSimplifyNegOneLHS", "MulSimplifyNegOneRHS"]; } def ApplyDivSimplifyPatterns : EnzymeHLOPatternOp< "div_simplify"> { - let patterns = ["DivSimplify"]; + let patterns = ["DivSimplifyOneRHS", "DivSimplifyNegOneRHS", "DivConstToMulReciprocal"]; } def ApplyPowSimplifyPatterns : EnzymeHLOPatternOp< "pow_simplify"> { - let patterns = ["PowSimplify"]; + let patterns = ["PowSimplifyOneRHS", "PowSimplifyZeroRHS", "PowSimplifyOneLHS", "PowSimplifyNegOneRHS", "PowSimplifyNegHalfRHS", "PowSimplifyHalfRHS", "PowSimplifyTwoRHS"]; } def ApplyNoopSlicePatterns : EnzymeHLOPatternOp< "noop_slice"> { @@ -275,7 +275,7 @@ def ApplySoftplusConstProp : EnzymeHLOPatternOp< } def SignAbsSimplifyPatterns : EnzymeHLOPatternOp< "sign_abs_simplify"> { - let patterns = ["SignAbsSimplify"]; + let patterns = ["SignAbsSimplify_1", "SignAbsSimplify_2"]; } def AbsPositiveSimplifyPatterns : EnzymeHLOPatternOp< "abs_positive_simplify"> { @@ -284,7 +284,7 @@ def AbsPositiveSimplifyPatterns : EnzymeHLOPatternOp< def ApplySquareAbsSimplifyPatterns : EnzymeHLOPatternOp< "square_abs_simplify"> { - let patterns = ["SquareAbsSimplify"]; + let patterns = ["SquareAbsSimplify", "SquareAbsSimplifyComplex"]; } def ApplyDivideDivideSimplifyPatterns : EnzymeHLOPatternOp< @@ -445,11 +445,11 @@ def ApplyDotTransposePatterns : EnzymeHLOPatternOp< } def ApplyConvertConvertFloatPatterns : EnzymeHLOPatternOp< "convert_convert_float"> { - let patterns = ["ConvertConvertFloat"]; + let patterns = ["ConvertConvertFloatIdentity", "ConvertConvertFloat"]; } def ApplyConvertConvertIntPatterns : EnzymeHLOPatternOp< "convert_convert_int"> { - let patterns = ["ConvertConvertInt"]; + let patterns = ["ConvertConvertIntIdentity", "ConvertConvertInt"]; } def ApplyConvertMulConvertPatterns : EnzymeHLOPatternOp< "convert_mul_convert"> { @@ -1502,7 +1502,7 @@ def SelectPadToDUS : EnzymeHLOPatternOp< def SelectSelectSameCond : EnzymeHLOPatternOp< "select_select_same_cond"> { - let patterns = ["SelectSelectSameCond"]; + let patterns = ["SelectSelectSameCondFalse", "SelectSelectSameCondTrue"]; } def SelectSelectNegCond : EnzymeHLOPatternOp< @@ -2040,7 +2040,7 @@ def ApplySumToConvPatterns : EnzymeHLOParameterizedPatternOp< def ApplyXOrSimplifyPatterns : EnzymeHLOPatternOp< "xor_simplify"> { - let patterns = ["XorSimplify"]; + let patterns = ["XorSimplifyFalseLHS", "XorSimplifyFalseRHS", "XorSimplifyTrueLHS", "XorSimplifyTrueRHS"]; } def SumToReduceWindow : EnzymeHLOPatternOp< @@ -2503,7 +2503,7 @@ def ReduceReduce : EnzymeHLOPatternOp< def ConjReal : EnzymeHLOPatternOp< "conj_real"> { - let patterns = ["ConjReal"]; + let patterns = ["ConjRealSimplify"]; } def ApplyTransposeBatchNormTrainingPatterns : EnzymeHLOPatternOp< @@ -2538,17 +2538,17 @@ def ApplyIfOpLiftCommonOpsPatterns : EnzymeHLOPatternOp< def ApplyInvolutionNegSimplifyPatterns : EnzymeHLOPatternOp< "involution_neg_simplify"> { - let patterns = ["InvolutionSimplify"]; + let patterns = ["NegNegSimplify"]; } def ApplyInvolutionConjSimplifyPatterns : EnzymeHLOPatternOp< "involution_conj_simplify"> { - let patterns = ["InvolutionSimplify"]; + let patterns = ["ConjConjSimplify"]; } def ApplyInvolutionNotSimplifyPatterns : EnzymeHLOPatternOp< "involution_not_simplify"> { - let patterns = ["InvolutionSimplify"]; + let patterns = ["NotNotSimplify"]; } def ApplyRealConjSimplifyPatterns : EnzymeHLOPatternOp< @@ -2570,7 +2570,7 @@ def ApplyConjComplexSimplifyPatterns : EnzymeHLOPatternOp< // ConjReal now performs a detailed analysis that covers this case. def ApplyConjConvertSimplifyPatterns : EnzymeHLOPatternOp< "conj_convert_simplify"> { - let patterns = ["ConjReal"]; + let patterns = ["ConjRealSimplify"]; } def ApplyElementwiseComplexSimplifyPatterns : EnzymeHLOPatternOp< @@ -2690,19 +2690,19 @@ def ApplyPowerMultiplyToPower : EnzymeHLOPatternOp<"power_multiply_to_power"> { } def ApplyLogSimplify : EnzymeHLOPatternOp<"log_simplify"> { - let patterns = ["LogSimplify"]; + let patterns = ["LogExpSimplify", "LogPowSimplify", "LogMulIdenticalSimplify", "LogMulConstSimplify", "LogAddIdenticalSimplify", "LogAddOneRHSSimplify", "LogAddOneLHSSimplify", "LogDivConstSimplify", "LogSqrtSimplify", "LogCbrtSimplify", "LogRsqrtSimplify"]; } def ApplyExponentialMinusOneFuse : EnzymeHLOPatternOp<"exponential_minus_one_fuse"> { - let patterns = ["ExponentialMinusOneFuse", "ExponentialMinusOneAddFuse"]; + let patterns = ["ExponentialMinusOneFuse1", "ExponentialMinusOneFuse2", "ExponentialMinusOneAddFuse1", "ExponentialMinusOneAddFuse2"]; } def ApplyNegMulConstSimplify : EnzymeHLOPatternOp<"neg_mul_const_simplify"> { - let patterns = ["NegMulConstSimplify"]; + let patterns = ["NegMulConstLHSSimplify", "NegMulConstRHSSimplify"]; } def ApplyNegDivConstSimplify : EnzymeHLOPatternOp<"neg_div_const_simplify"> { - let patterns = ["NegDivConstSimplify"]; + let patterns = ["NegDivConstLHSSimplify", "NegDivConstRHSSimplify"]; } def ApplyNegatedConstantMulFactoring : EnzymeHLOPatternOp<"negated_constant_mul_factoring"> { diff --git a/src/enzyme_ad/jax/Utils.h b/src/enzyme_ad/jax/Utils.h index a804e82e86..376670697c 100644 --- a/src/enzyme_ad/jax/Utils.h +++ b/src/enzyme_ad/jax/Utils.h @@ -1280,11 +1280,18 @@ SmallVector computeGatherSliceSizes(stablehlo::ScatterOp &scatterOp); template stablehlo::ConstantOp createConstantOpFromScalar(PatternRewriter &rewriter, - Operation *op, T value) { + Location loc, Type type, + T value) { return stablehlo::ConstantOp::create( - rewriter, op->getLoc(), op->getResult(0).getType(), - cast( - mlir::enzyme::makeAttr(op->getResult(0).getType(), value))); + rewriter, loc, type, + cast(mlir::enzyme::makeAttr(type, value))); +} + +template +stablehlo::ConstantOp createConstantOpFromScalar(PatternRewriter &rewriter, + Operation *op, T value) { + return createConstantOpFromScalar(rewriter, op->getLoc(), + op->getResult(0).getType(), value); } stablehlo::ComparisonDirection diff --git a/test/lit_tests/squareabs.mlir b/test/lit_tests/squareabs.mlir index 177e31df45..6acdc77b2d 100644 --- a/test/lit_tests/squareabs.mlir +++ b/test/lit_tests/squareabs.mlir @@ -18,12 +18,12 @@ func.func @squareabscomplex(%arg0: tensor<2x2xcomplex>) -> tensor<2x2xf32> } // CHECK: func.func @squareabscomplex(%arg0: tensor<2x2xcomplex>) -> tensor<2x2xf32> -// CHECK-NEXT: %0 = stablehlo.real %arg0 : (tensor<2x2xcomplex>) -> tensor<2x2xf32> -// CHECK-NEXT: %1 = stablehlo.imag %arg0 : (tensor<2x2xcomplex>) -> tensor<2x2xf32> -// CHECK-NEXT: %2 = stablehlo.multiply %0, %0 : tensor<2x2xf32> -// CHECK-NEXT: %3 = stablehlo.multiply %1, %1 : tensor<2x2xf32> -// CHECK-NEXT: %4 = stablehlo.add %2, %3 : tensor<2x2xf32> -// CHECK-NEXT: return %4 : tensor<2x2xf32> +// CHECK-DAG: [[REAL:%[0-9]+]] = stablehlo.real %arg0 : (tensor<2x2xcomplex>) -> tensor<2x2xf32> +// CHECK-DAG: [[IMAG:%[0-9]+]] = stablehlo.imag %arg0 : (tensor<2x2xcomplex>) -> tensor<2x2xf32> +// CHECK-DAG: [[REAL2:%[0-9]+]] = stablehlo.multiply [[REAL]], [[REAL]] : tensor<2x2xf32> +// CHECK-DAG: [[IMAG2:%[0-9]+]] = stablehlo.multiply [[IMAG]], [[IMAG]] : tensor<2x2xf32> +// CHECK: [[ADD:%[0-9]+]] = stablehlo.add [[REAL2]], [[IMAG2]] : tensor<2x2xf32> +// CHECK-NEXT: return [[ADD]] : tensor<2x2xf32> // CHECK-NEXT: } // doesn't apply here