diff --git a/src/enzyme_ad/jax/Passes/ArithRaising.cpp b/src/enzyme_ad/jax/Passes/ArithRaising.cpp index 77b6da1f00..c3b2b431c7 100644 --- a/src/enzyme_ad/jax/Passes/ArithRaising.cpp +++ b/src/enzyme_ad/jax/Passes/ArithRaising.cpp @@ -682,6 +682,7 @@ struct ArithRaisingPass RaiseUnary, RaiseUnary, RaiseUnary, + RaiseUnary, RaiseUnary, RaiseUnary, RaiseUnary, diff --git a/src/enzyme_ad/jax/Passes/LibDeviceFuncsRaisingPass.cpp b/src/enzyme_ad/jax/Passes/LibDeviceFuncsRaisingPass.cpp index ff053141c7..27c0da1e5b 100644 --- a/src/enzyme_ad/jax/Passes/LibDeviceFuncsRaisingPass.cpp +++ b/src/enzyme_ad/jax/Passes/LibDeviceFuncsRaisingPass.cpp @@ -700,6 +700,17 @@ using ConvertFMFMathFromLLVMPattern = using AbsFOpLowering = ConvertFMFMathFromLLVMPattern; + +// llvm.intr.abs carries an is_int_min_poison flag arith has no place for; +// drop it and raise to math.absi. +struct AbsIOpRaising : public OpRewritePattern { + using OpRewritePattern::OpRewritePattern; + LogicalResult matchAndRewrite(LLVM::AbsOp op, + PatternRewriter &rewriter) const override { + rewriter.replaceOpWithNewOp(op, op.getIn()); + return success(); + } +}; using CeilOpLowering = ConvertFMFMathFromLLVMPattern; using CopySignOpLowering = @@ -1415,6 +1426,7 @@ void populateLLVMToMathPatterns(MLIRContext *context, // From // https://github.com/llvm/llvm-project/blob/7d8b4eb0ead277f41ff69525ed807f9f6e227f37/mlir/lib/Conversion/MathToLLVM/MathToLLVM.cpp#L306 // patterns.add(converter); + patterns.add(patterns.getContext()); patterns.add i32 { + %res = "llvm.intr.abs"(%arg0) <{is_int_min_poison = false}> : (i32) -> i32 + func.return %res : i32 + } + + // HLO-LABEL: @tensor_absi + // HLO: stablehlo.abs %arg0 : tensor<20xi32> + // HLO-NOT: math.absi + func.func @tensor_absi(%arg0: tensor<20xi32>) -> tensor<20xi32> { + %res = math.absi %arg0 : tensor<20xi32> + func.return %res : tensor<20xi32> + } +}