21#include "llvm/Support/FormatVariadic.h"
27#define GEN_PASS_DEF_CONVERTMATHTOXEVM
28#include "mlir/Conversion/Passes.h.inc"
33#define DEBUG_TYPE "math-to-xevm"
36 auto vecType = dyn_cast<VectorType>(type);
37 return vecType && vecType.getShape().size() == 1 &&
38 vecType.getShape()[0] == 1 && vecType.getElementType().isFloat();
44 if (
auto vecType = dyn_cast<VectorType>(type)) {
45 if (!vecType.getElementType().isFloat())
50 if (
shape.size() != 1)
74 std::string mangledFuncName =
77 auto appendFloatToMangledFunc = [&mangledFuncName](
Type type) {
79 mangledFuncName +=
"f";
80 else if (type.isF16())
81 mangledFuncName +=
"Dh";
82 else if (type.isF64())
83 mangledFuncName +=
"d";
86 for (
auto type : operandTypes) {
87 if (
auto vecType = dyn_cast<VectorType>(type)) {
88 mangledFuncName +=
"Dv" + std::to_string(vecType.getShape()[0]) +
"_";
89 appendFloatToMangledFunc(vecType.getElementType());
91 appendFloatToMangledFunc(type);
94 return mangledFuncName;
99 ConversionPatternRewriter &rewriter)
const override {
103 arith::FastMathFlags fastFlags = op.getFastmath();
104 if (!arith::bitEnumContainsAll(fastFlags, arith::FastMathFlags::afn))
105 return rewriter.notifyMatchFailure(op,
"not a fastmath `afn` operation");
116 for (
Value &operand : operands) {
117 Type opTy = operand.getType();
122 return rewriter.notifyMatchFailure(
123 op, llvm::formatv(
"incompatible operand type: '{0}'", opTy));
124 if (unwrapSizeOneVec) {
126 "expected all operands to be size-1 vectors");
127 opTy = cast<VectorType>(opTy).getElementType();
128 operand = vector::ExtractOp::create(rewriter, loc, operand,
131 operandTypes.push_back(opTy);
134 Type resultType = unwrapSizeOneVec
135 ? cast<VectorType>(op.getType()).getElementType()
138 auto moduleOp = op->template getParentWithTrait<OpTrait::SymbolTable>();
141 operandTypes, resultType);
142 assert(!failed(funcOpRes));
143 LLVM::LLVMFuncOp funcOp = funcOpRes.value();
145 auto callOp = LLVM::CallOp::create(rewriter, loc, funcOp, operands);
150 callOp.setFastmathFlagsAttr(
154 if (unwrapSizeOneVec) {
156 rewriter.replaceOpWithNewOp<vector::BroadcastOp>(op, op.getType(),
159 rewriter.replaceOp(op, callOp);
167template <
typename OpTy>
172 std::string prefix =
"__spirv_ocl_";
173 std::string mangledName =
"_Z" +
174 std::to_string(prefix.size() + opName.size()) +
175 prefix + opName.str();
179 converter, mangledName +
"f", mangledName +
"d",
181 "", benefit, LLVM::cconv::CConv::SPIR_FUNC);
241 patterns.
getContext(),
"__spirv_ocl_native_exp", benefit);
243 patterns.
getContext(),
"__spirv_ocl_native_cos", benefit);
245 patterns.
getContext(),
"__spirv_ocl_native_exp2", benefit);
247 patterns.
getContext(),
"__spirv_ocl_native_log", benefit);
249 patterns.
getContext(),
"__spirv_ocl_native_log2", benefit);
251 patterns.
getContext(),
"__spirv_ocl_native_log10", benefit);
253 patterns.
getContext(),
"__spirv_ocl_native_powr", benefit);
255 patterns.
getContext(),
"__spirv_ocl_native_rsqrt", benefit);
257 patterns.
getContext(),
"__spirv_ocl_native_sin", benefit);
259 patterns.
getContext(),
"__spirv_ocl_native_sqrt", benefit);
261 patterns.
getContext(),
"__spirv_ocl_native_tan", benefit);
264 patterns.
getContext(),
"__spirv_ocl_native_divide", benefit);
268struct ConvertMathToXeVMPass
271 void runOnOperation()
override;
275void ConvertMathToXeVMPass::runOnOperation() {
276 Operation *op = getOperation();
285 RewritePatternSet simplifications(ctx);
287 FrozenRewritePatternSet frozen(std::move(simplifications));
288 SmallVector<Operation *> candidates;
289 op->
walk([&](Operation *nested) {
290 if (frozen.getOpSpecificNativePatterns().contains(nested->
getName()))
291 candidates.push_back(nested);
293 GreedyRewriteConfig config;
296 return signalPassFailure();
299 const auto &dl = getAnalysis<DataLayoutAnalysis>();
302 LowerToLLVMOptions
options(ctx, dl.getAtOrAbove(op));
303 LLVMTypeConverter converter(ctx,
options);
309 constexpr unsigned oclBenefit = 1;
315 .addIllegalOp<LLVM::CosOp, LLVM::ExpOp, LLVM::Exp2Op, LLVM::LogOp,
316 LLVM::Log10Op, LLVM::Log2Op, LLVM::SinOp, LLVM::SqrtOp>();
318 target.addLegalDialect<BuiltinDialect, LLVM::LLVMDialect>();
322 target.addLegalOp<vector::ExtractOp, vector::BroadcastOp>();
324 applyPartialConversion(getOperation(),
target, std::move(patterns))))
static bool isSizeOneVector(Type type)
static bool isSPIRVCompatibleFloatOrVec(Type type)
static void populateOCLExtSetOpPatterns(const LLVMTypeConverter &converter, RewritePatternSet &patterns, PatternBenefit benefit, StringRef opName)
static llvm::ManagedStatic< PassManagerOptions > options
GreedyRewriteConfig & enableFolding(bool enable=true)
Conversion from types to the LLVM IR dialect.
This class defines the main interface for locations in MLIR and acts as a non-nullable wrapper around...
MLIRContext is the top-level object for a collection of MLIR operations.
Location getLoc()
The source location the operation was defined or derived from.
This provides public APIs that all operations should have.
OperationName getName()
The name of an operation is the key identifier for it.
std::enable_if_t< llvm::function_traits< std::decay_t< FnT > >::num_args==1, RetT > walk(FnT &&callback)
Walk the operation by calling the callback for each nested operation (including this one),...
MLIRContext * getContext()
Return the context this operation is associated with.
This class represents the benefit of a pattern match in a unitless scheme that ranges from 0 (very li...
MLIRContext * getContext() const
RewritePatternSet & add(ConstructorArg &&arg, ConstructorArgs &&...args)
Add an instance of each of the pattern types 'Ts' to the pattern list with the given arguments.
Instances of the Type class are uniqued, have an immutable identifier and an optional mutable compone...
bool isFloat() const
Return true if this is an float type (with the specified width).
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
TargetOp::Properties getProperties() const
ArrayRef< NamedAttribute > getDiscardableAttrs() const
FailureOr< LLVM::LLVMFuncOp > lookupOrCreateFn(OpBuilder &b, Operation *moduleOp, StringRef name, ArrayRef< Type > paramTypes={}, Type resultType={}, bool isVarArg=false, bool isReserved=false, SymbolTableCollection *symbolTables=nullptr)
Create a FuncOp with signature resultType(paramTypes) and name name`.
Include the generated interface declarations.
LogicalResult applyOpPatternsGreedily(ArrayRef< Operation * > ops, const FrozenRewritePatternSet &patterns, GreedyRewriteConfig config=GreedyRewriteConfig(), bool *changed=nullptr, bool *allErased=nullptr)
Rewrite the specified ops by repeatedly applying the highest benefit patterns in a greedy worklist dr...
void populateMathAlgebraicSimplificationPatterns(RewritePatternSet &patterns)
void populateMathToScalarOCLExtSetConversionPatterns(const LLVMTypeConverter &converter, RewritePatternSet &patterns, PatternBenefit benefit=1)
Populate the given list with patterns that convert from Math to OCL LLVM-SPV builtin calls.
void populateMathToXeVMConversionPatterns(RewritePatternSet &patterns, bool convertArith, PatternBenefit benefit=1)
Populate the given list with patterns that convert from Math to XeVM calls.
Convert math ops marked with fast (afn) to native OpenCL intrinsics.
const StringRef nativeFunc
ConvertNativeFuncPattern(MLIRContext *context, StringRef nativeFunc, PatternBenefit benefit=1)
std::string getMangledNativeFuncName(const ArrayRef< Type > operandTypes) const
LogicalResult matchAndRewrite(Op op, typename Op::Adaptor adaptor, ConversionPatternRewriter &rewriter) const override
Rewriting that replaces SourceOp with a CallOp to f32Func or f64Func or f32ApproxFunc or f16Func or i...
Unrolls SourceOp to array/vector elements.