Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
6 changes: 2 additions & 4 deletions src/enzyme_ad/jax/Passes/ArithRaising.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -818,11 +818,9 @@ struct ArithRaisingPass
// TODO: either SI or UI is wrong
RaiseToConvert<arith::ExtUIOp>, RaiseToConvert<arith::ExtSIOp>,
RaiseToConvert<arith::TruncIOp>, RaiseMulAdd<math::FmaOp>,
RaiseMulAdd<enzymexla::FMulAddOp>, RaiseCopySign, RaiseTruncOp,
RaiseAtan, RaiseLog10, RaiseLog2, RaiseExp2, RaiseMaxNumF,
RaiseCopySign, RaiseTruncOp, RaiseAtan, RaiseMaxNumF,
RaiseMinNumF, RaiseIsNaN, RaiseConstant, RaiseFPToSI,
RaiseSIToFP, RaiseUIToFP, RaiseSelect, RaiseCmpI, RaiseSinCos>(
context);
RaiseSIToFP, RaiseUIToFP, RaiseSelect, RaiseCmpI, RaiseSinCos>(context);

walkAndApplyPatterns(getOperation(), std::move(patterns));
}
Expand Down
77 changes: 69 additions & 8 deletions src/enzyme_ad/jax/Passes/ConvertPolygeistToLLVM.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -412,17 +412,68 @@ struct Memref2PointerOpLowering
}
};

struct LGammaOpLowering : public OpRewritePattern<enzymexla::LGammaOp> {
using OpRewritePattern<enzymexla::LGammaOp>::OpRewritePattern;

LogicalResult matchAndRewrite(enzymexla::LGammaOp op,
PatternRewriter &rewriter) const override {

Type ty = op.getResult().getType();
if (!ty.isF32() && !ty.isF64())
return failure();

bool onGPU = op->getParentOfType<gpu::GPUFuncOp>() != nullptr;

StringRef fnname = ty.isF32() ? onGPU ? "__nv_lgammaf" : "lgammaf"
: onGPU ? "__nv_lgamma"
: "lgamma";

auto moduleOp = SymbolTable::getNearestSymbolTable(op);
auto fn = LLVM::lookupOrCreateFn(rewriter, moduleOp, fnname, {ty}, ty);
if (failed(fn))
return failure();

rewriter.replaceOpWithNewOp<LLVM::CallOp>(op, *fn, op->getOperands());
return success();
}
};

struct TGammaOpLowering : public OpRewritePattern<enzymexla::TGammaOp> {
using OpRewritePattern<enzymexla::TGammaOp>::OpRewritePattern;

LogicalResult matchAndRewrite(enzymexla::TGammaOp op,
PatternRewriter &rewriter) const override {

Type ty = op.getResult().getType();
if (!ty.isF32() && !ty.isF64())
return failure();

bool onGPU = op->getParentOfType<gpu::GPUFuncOp>() != nullptr;

StringRef fnname = ty.isF32() ? onGPU ? "__nv_tgammaf" : "tgammaf"
: onGPU ? "__nv_tgamma"
: "tgamma";

auto moduleOp = SymbolTable::getNearestSymbolTable(op);
auto fn = LLVM::lookupOrCreateFn(rewriter, moduleOp, fnname, {ty}, ty);
if (failed(fn))
return failure();

rewriter.replaceOpWithNewOp<LLVM::CallOp>(op, *fn, op->getOperands());
return success();
}
};

// Back to the intrinsic it was raised from. Not llvm.intr.fma: fmuladd only
// permits fusing, while fma requires the single rounding, which a target
// without FMA units honors with a libm call per multiply-add.
struct FMulAddOpLowering : public ConvertOpToLLVMPattern<enzymexla::FMulAddOp> {
using ConvertOpToLLVMPattern<enzymexla::FMulAddOp>::ConvertOpToLLVMPattern;
struct FMulAddOpLowering : public OpRewritePattern<enzymexla::FMulAddOp> {
using OpRewritePattern<enzymexla::FMulAddOp>::OpRewritePattern;

LogicalResult
matchAndRewrite(enzymexla::FMulAddOp op, OpAdaptor transformed,
ConversionPatternRewriter &rewriter) const override {
rewriter.replaceOpWithNewOp<LLVM::FMulAddOp>(
op, transformed.getA(), transformed.getB(), transformed.getC());
LogicalResult matchAndRewrite(enzymexla::FMulAddOp op,
PatternRewriter &rewriter) const override {
rewriter.replaceOpWithNewOp<LLVM::FMulAddOp>(op, op.getA(), op.getB(),
op.getC());
return success();
}
};
Expand Down Expand Up @@ -490,6 +541,16 @@ struct Pointer2MemrefOpLowering
}
};

void mlir::enzyme::populateEnzymeXLAMathToLLVMConversionPatterns(
RewritePatternSet &patterns) {

// clang-format off
patterns.add<FMulAddOpLowering>(patterns.getContext());
patterns.add<TGammaOpLowering>(patterns.getContext());
patterns.add<LGammaOpLowering>(patterns.getContext());
// clang-format on
}

void populatePolygeistToLLVMConversionPatterns(LLVMTypeConverter &converter,
RewritePatternSet &patterns) {
// clang-format off
Expand All @@ -500,7 +561,7 @@ void populatePolygeistToLLVMConversionPatterns(LLVMTypeConverter &converter,
patterns.add<Stream2TokenOpLowering>(converter);
patterns.add<Memref2PointerOpLowering>(converter);
patterns.add<Pointer2MemrefOpLowering>(converter);
patterns.add<FMulAddOpLowering>(converter);
enzyme::populateEnzymeXLAMathToLLVMConversionPatterns(patterns);
// clang-format on
}

Expand Down
35 changes: 28 additions & 7 deletions src/enzyme_ad/jax/Passes/LowerEnzymeXLAMathPatterns.td
Original file line number Diff line number Diff line change
Expand Up @@ -33,8 +33,19 @@ def GeluApproximationNONE : ConstantAttr<EnzymeXLA_GeluApproximationAttr, "::mli
def GeluApproximationTANH : ConstantAttr<EnzymeXLA_GeluApproximationAttr, "::mlir::enzymexla::GeluApproximation::TANH">;
def GeluApproximationSIGMOID : ConstantAttr<EnzymeXLA_GeluApproximationAttr, "::mlir::enzymexla::GeluApproximation::SIGMOID">;

def AllOperandsAreTensors : Constraint<
CPred<"llvm::all_of($0.getDefiningOp()->getOperandTypes(), [](Type ty) { "
"return llvm::isa<mlir::TensorType>(ty); })">,
"all operands are of type Tensor">;

def AllResultsAreTensors : Constraint<
CPred<"llvm::all_of($0.getDefiningOp()->getResultTypes(), [](Type ty) { "
"return llvm::isa<mlir::TensorType>(ty); })">,
"all results are of type Tensor">;

def : Pat<(ReluOp:$op $input),
(StableHLO_MaxOp $input, (CreateConstOp0 $input))>;
(StableHLO_MaxOp $input, (CreateConstOp0 $input)),
[(AllOperandsAreTensors $op), (AllResultsAreTensors $op)]>;

// Gelu Approximation: NONE
def : Pat<(GeluOp:$op $x, GeluApproximationNONE),
Expand All @@ -46,7 +57,8 @@ def : Pat<(GeluOp:$op $x, GeluApproximationNONE),
(CHLO_ErfOp (StableHLO_MulOp $x, (CreateInvSqrt2Op $op)))
)
)
)
),
[(AllOperandsAreTensors $op), (AllResultsAreTensors $op)]
>;

// Gelu Approximation: TANH
Expand All @@ -71,7 +83,8 @@ def : Pat<(GeluOp:$op $x, GeluApproximationTANH),
)
)
)
)
),
[(AllOperandsAreTensors $op), (AllResultsAreTensors $op)]
>;

// Gelu Approximation: SIGMOID
Expand All @@ -93,7 +106,8 @@ def : Pat<(GeluOp:$op $x, GeluApproximationSIGMOID),
),
ConstDefaultResultAccuracyAttr
)
)
),
[(AllOperandsAreTensors $op), (AllResultsAreTensors $op)]
>;

// Softplus
Expand All @@ -112,10 +126,12 @@ def : Pat<(SoftplusOp:$op $x),
ConstDefaultResultAccuracyAttr
)
)
)
),
[(AllOperandsAreTensors $op), (AllResultsAreTensors $op)]
>;

def : Pat<(LGammaOp $x), (CHLO_LgammaOp $x)>;
def : Pat<(LGammaOp:$op $x), (CHLO_LgammaOp $x),
[(AllOperandsAreTensors $op), (AllResultsAreTensors $op)]>;

def : Pat<(TGammaOp:$op $x),
(StableHLO_SelectOp
Expand All @@ -127,9 +143,14 @@ def : Pat<(TGammaOp:$op $x),
),
(CreateNanOp $op),
(StableHLO_ExpOp (CHLO_LgammaOp $x), ConstDefaultResultAccuracyAttr)
)
),
[(AllOperandsAreTensors $op), (AllResultsAreTensors $op)]
>;

def : Pat<(FMulAddOp:$op $a, $b, $c),
(StableHLO_AddOp (StableHLO_MulOp $a, $b), $c),
[(AllOperandsAreTensors $op), (AllResultsAreTensors $op)]>;

// Hypot
def : Pat<(HypotOp:$op $x, $y),
(StableHLO_SelectOp
Expand Down
17 changes: 10 additions & 7 deletions src/enzyme_ad/jax/Passes/LowerJIT.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -14,6 +14,7 @@

#include "src/enzyme_ad/jax/Dialect/Ops.h"

#include "mlir/Conversion/LLVMCommon/TypeConverter.h"
#include "mlir/Dialect/Func/IR/FuncOps.h"
#include "mlir/Dialect/GPU/IR/GPUDialect.h"
#include "mlir/Dialect/LLVMIR/LLVMDialect.h"
Expand Down Expand Up @@ -978,13 +979,15 @@ CompileCall(SymbolTableCollection &symbolTable, mlir::Location loc,
if (str.size() > 200)
gmod.setName(str.substr(0, 200));
});
submod->walk([](enzymexla::FMulAddOp op) {
OpBuilder builder(op);
auto newOp = LLVM::FMulAddOp::create(builder, op->getLoc(), op.getA(),
op.getB(), op.getC());
op.getResult().replaceAllUsesWith(newOp.getResult());
op->erase();
});

{
RewritePatternSet patterns(submod.getContext());
populateEnzymeXLAMathToLLVMConversionPatterns(patterns);
if (failed(applyPatternsGreedily(submod, std::move(patterns)))) {
submod.erase();
return {};
}
}

std::string legalName;
submod->walk([&](gpu::LaunchFuncOp gmod) {
Expand Down
1 change: 1 addition & 0 deletions src/enzyme_ad/jax/Passes/Passes.h
Original file line number Diff line number Diff line change
Expand Up @@ -40,6 +40,7 @@ void populateInlineNeverLoopingWhilePattern(RewritePatternSet &patterns);
#define GEN_PASS_REGISTRATION
#include "src/enzyme_ad/jax/Passes/Passes.h.inc"

void populateEnzymeXLAMathToLLVMConversionPatterns(RewritePatternSet &patterns);
void populateLibDeviceFuncsToOpsPatterns(MLIRContext *context,
RewritePatternSet &patterns);

Expand Down
2 changes: 1 addition & 1 deletion test/lit_tests/raising/affine_to_stablehlo_fmuladd.mlir
Original file line number Diff line number Diff line change
@@ -1,4 +1,4 @@
// RUN: enzymexlamlir-opt %s --raise-affine-to-stablehlo --canonicalize --arith-raise --enzyme-hlo-opt=max_constant_expansion=0 | FileCheck %s
// RUN: enzymexlamlir-opt %s --raise-affine-to-stablehlo --canonicalize --lower-enzymexla-math --enzyme-hlo-opt=max_constant_expansion=0 | FileCheck %s

// enzymexla.math.fmuladd rides the same route as math.fma: tensorized by the
// affine raising, then split into multiply+add by arith-raise -- stablehlo has
Expand Down
Loading