MLIR 24.0.0git
MathToXeVM.cpp
Go to the documentation of this file.
1//===-- MathToXeVM.cpp - conversion from Math to XeVM ---------------------===//
2//
3// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.
4// See https://llvm.org/LICENSE.txt for license information.
5// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
6//
7//===----------------------------------------------------------------------===//
8
19#include "mlir/Pass/Pass.h"
21#include "llvm/Support/FormatVariadic.h"
22
25
26namespace mlir {
27#define GEN_PASS_DEF_CONVERTMATHTOXEVM
28#include "mlir/Conversion/Passes.h.inc"
29} // namespace mlir
30
31using namespace mlir;
32
33#define DEBUG_TYPE "math-to-xevm"
34
35static bool isSizeOneVector(Type type) {
36 auto vecType = dyn_cast<VectorType>(type);
37 return vecType && vecType.getShape().size() == 1 &&
38 vecType.getShape()[0] == 1 && vecType.getElementType().isFloat();
39}
40
42 if (type.isFloat())
43 return true;
44 if (auto vecType = dyn_cast<VectorType>(type)) {
45 if (!vecType.getElementType().isFloat())
46 return false;
47 // SPIRV distinguishes between vectors and matrices: OpenCL native math
48 // intrsinics are not compatible with matrices.
49 ArrayRef<int64_t> shape = vecType.getShape();
50 if (shape.size() != 1)
51 return false;
52 // SPIRV has no size-1 vector type; such degenerate vectors are handled
53 // by unwrapping to the scalar intrinsic (see matchAndRewrite).
54 if (shape[0] == 1)
55 return true;
56 // SPIRV only allows vectors of size 2, 3, 4, 8, 16.
57 if (shape[0] == 2 || shape[0] == 3 || shape[0] == 4 || shape[0] == 8 ||
58 shape[0] == 16)
59 return true;
60 }
61 return false;
62}
63
64/// Convert math ops marked with `fast` (`afn`) to native OpenCL intrinsics.
65template <typename Op>
66struct ConvertNativeFuncPattern final : public OpConversionPattern<Op> {
67
69 PatternBenefit benefit = 1)
70 : OpConversionPattern<Op>(context, benefit), nativeFunc(nativeFunc) {}
71
72 inline std::string
73 getMangledNativeFuncName(const ArrayRef<Type> operandTypes) const {
74 std::string mangledFuncName =
75 "_Z" + std::to_string(nativeFunc.size()) + nativeFunc.str();
76
77 auto appendFloatToMangledFunc = [&mangledFuncName](Type type) {
78 if (type.isF32())
79 mangledFuncName += "f";
80 else if (type.isF16())
81 mangledFuncName += "Dh";
82 else if (type.isF64())
83 mangledFuncName += "d";
84 };
85
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());
90 } else
91 appendFloatToMangledFunc(type);
92 }
93
94 return mangledFuncName;
95 }
96
97 LogicalResult
98 matchAndRewrite(Op op, typename Op::Adaptor adaptor,
99 ConversionPatternRewriter &rewriter) const override {
100 if (!isSPIRVCompatibleFloatOrVec(op.getType()))
101 return failure();
102
103 arith::FastMathFlags fastFlags = op.getFastmath();
104 if (!arith::bitEnumContainsAll(fastFlags, arith::FastMathFlags::afn))
105 return rewriter.notifyMatchFailure(op, "not a fastmath `afn` operation");
106
107 Location loc = op.getLoc();
108
109 // SPIRV has no size-1 vector type: such vectors are the degenerate result
110 // of distributing/linearizing larger vectors down to a single element (e.g.
111 // by the XeGPU lowering pipeline). They have no OpenCL vector intrinsic, so
112 // unwrap them to the scalar element type and use the scalar intrinsic.
113 SmallVector<Value, 1> operands(adaptor.getOperands());
114 SmallVector<Type, 1> operandTypes;
115 bool unwrapSizeOneVec = isSizeOneVector(op.getType());
116 for (Value &operand : operands) {
117 Type opTy = operand.getType();
118 // This pass only supports operations on vectors that are already in SPIRV
119 // supported vector sizes: Distributing unsupported vector sizes to SPIRV
120 // supported vector sizes are done in other blocking optimization passes.
122 return rewriter.notifyMatchFailure(
123 op, llvm::formatv("incompatible operand type: '{0}'", opTy));
124 if (unwrapSizeOneVec) {
125 assert(isSizeOneVector(opTy) &&
126 "expected all operands to be size-1 vectors");
127 opTy = cast<VectorType>(opTy).getElementType();
128 operand = vector::ExtractOp::create(rewriter, loc, operand,
130 }
131 operandTypes.push_back(opTy);
132 }
133
134 Type resultType = unwrapSizeOneVec
135 ? cast<VectorType>(op.getType()).getElementType()
136 : op.getType();
137
138 auto moduleOp = op->template getParentWithTrait<OpTrait::SymbolTable>();
139 auto funcOpRes = LLVM::lookupOrCreateFn(
140 rewriter, moduleOp, getMangledNativeFuncName(operandTypes),
141 operandTypes, resultType);
142 assert(!failed(funcOpRes));
143 LLVM::LLVMFuncOp funcOp = funcOpRes.value();
144
145 auto callOp = LLVM::CallOp::create(rewriter, loc, funcOp, operands);
146 // Preserve fastmath flags in our MLIR op when converting to llvm function
147 // calls, in order to allow further fastmath optimizations: We thus need to
148 // convert arith fastmath attrs into attrs recognized by llvm.
150 callOp.setFastmathFlagsAttr(
151 fastAttrConverter.getProperties().getFastmathFlags());
152 callOp->setDiscardableAttrs(fastAttrConverter.getDiscardableAttrs());
153
154 if (unwrapSizeOneVec) {
155 // Re-wrap the scalar result back into a size-1 vector to preserve types.
156 rewriter.replaceOpWithNewOp<vector::BroadcastOp>(op, op.getType(),
157 callOp.getResult());
158 } else {
159 rewriter.replaceOp(op, callOp);
160 }
161 return success();
162 }
163
164 const StringRef nativeFunc;
165};
166
167template <typename OpTy>
169 RewritePatternSet &patterns,
170 PatternBenefit benefit,
171 StringRef opName) {
172 std::string prefix = "__spirv_ocl_";
173 std::string mangledName = "_Z" +
174 std::to_string(prefix.size() + opName.size()) +
175 prefix + opName.str();
176
177 patterns.add<ScalarizeVectorOpLowering<OpTy>>(converter, benefit);
179 converter, mangledName + "f", mangledName + "d",
180 /*f32ApproxFunc=*/"", /*f16Func=*/"",
181 /*i32Func=*/"", benefit, LLVM::cconv::CConv::SPIR_FUNC);
182}
183
185 const LLVMTypeConverter &converter, RewritePatternSet &patterns,
186 PatternBenefit benefit) {
187 populateOCLExtSetOpPatterns<math::AcosOp>(converter, patterns, benefit,
188 "acos");
189 populateOCLExtSetOpPatterns<math::AcoshOp>(converter, patterns, benefit,
190 "acosh");
191 populateOCLExtSetOpPatterns<math::AsinOp>(converter, patterns, benefit,
192 "asin");
193 populateOCLExtSetOpPatterns<math::AsinhOp>(converter, patterns, benefit,
194 "asinh");
195 populateOCLExtSetOpPatterns<math::AtanOp>(converter, patterns, benefit,
196 "atan");
197 populateOCLExtSetOpPatterns<math::Atan2Op>(converter, patterns, benefit,
198 "atan2");
199 populateOCLExtSetOpPatterns<math::AtanhOp>(converter, patterns, benefit,
200 "atanh");
201 populateOCLExtSetOpPatterns<math::CbrtOp>(converter, patterns, benefit,
202 "cbrt");
203 populateOCLExtSetOpPatterns<math::CopySignOp>(converter, patterns, benefit,
204 "copysign");
205 populateOCLExtSetOpPatterns<math::CosOp>(converter, patterns, benefit, "cos");
206 populateOCLExtSetOpPatterns<math::CoshOp>(converter, patterns, benefit,
207 "cosh");
208 populateOCLExtSetOpPatterns<math::ErfOp>(converter, patterns, benefit, "erf");
209 populateOCLExtSetOpPatterns<math::ErfcOp>(converter, patterns, benefit,
210 "erfc");
211 populateOCLExtSetOpPatterns<math::ExpOp>(converter, patterns, benefit, "exp");
212 populateOCLExtSetOpPatterns<math::Exp2Op>(converter, patterns, benefit,
213 "exp2");
214 populateOCLExtSetOpPatterns<math::ExpM1Op>(converter, patterns, benefit,
215 "expm1");
216 populateOCLExtSetOpPatterns<math::LogOp>(converter, patterns, benefit, "log");
217 populateOCLExtSetOpPatterns<math::Log10Op>(converter, patterns, benefit,
218 "log10");
219 populateOCLExtSetOpPatterns<math::Log1pOp>(converter, patterns, benefit,
220 "log1p");
221 populateOCLExtSetOpPatterns<math::Log2Op>(converter, patterns, benefit,
222 "log2");
223 populateOCLExtSetOpPatterns<math::PowFOp>(converter, patterns, benefit,
224 "pow");
225 populateOCLExtSetOpPatterns<math::RsqrtOp>(converter, patterns, benefit,
226 "rsqrt");
227 populateOCLExtSetOpPatterns<math::SinOp>(converter, patterns, benefit, "sin");
228 populateOCLExtSetOpPatterns<math::SinhOp>(converter, patterns, benefit,
229 "sinh");
230 populateOCLExtSetOpPatterns<math::SqrtOp>(converter, patterns, benefit,
231 "sqrt");
232 populateOCLExtSetOpPatterns<math::TanOp>(converter, patterns, benefit, "tan");
233 populateOCLExtSetOpPatterns<math::TanhOp>(converter, patterns, benefit,
234 "tanh");
235}
236
238 bool convertArith,
239 PatternBenefit benefit) {
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);
262 if (convertArith)
264 patterns.getContext(), "__spirv_ocl_native_divide", benefit);
265}
266
267namespace {
268struct ConvertMathToXeVMPass
269 : public impl::ConvertMathToXeVMBase<ConvertMathToXeVMPass> {
270 using Base::Base;
271 void runOnOperation() override;
272};
273} // namespace
274
275void ConvertMathToXeVMPass::runOnOperation() {
276 Operation *op = getOperation();
277 MLIRContext *ctx = op->getContext();
278
279 // Simplify first, so the cheaper form is what the conversion below lowers. A
280 // simplification rewrites a whole expression, so it must run before the
281 // lowering turns the parts of that expression into calls. Only the ops these
282 // patterns can match are given to the driver, and folding is off, so nothing
283 // else is touched.
284 {
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);
292 });
293 GreedyRewriteConfig config;
294 config.enableFolding(false);
295 if (failed(applyOpPatternsGreedily(candidates, frozen, config)))
296 return signalPassFailure();
297 }
298
299 const auto &dl = getAnalysis<DataLayoutAnalysis>();
300
301 RewritePatternSet patterns(&getContext());
302 LowerToLLVMOptions options(ctx, dl.getAtOrAbove(op));
303 LLVMTypeConverter converter(ctx, options);
304 ConversionTarget target(getContext());
305
306 // The native (`afn`) patterns must outrank the precise OCL patterns: an op
307 // marked `afn` gets the native intrinsic, and every other op falls through to
308 // the precise OCL intrinsic.
309 constexpr unsigned oclBenefit = 1;
310 populateMathToXeVMConversionPatterns(patterns, convertArith, oclBenefit + 1);
311 if (convertToOCL) {
313 oclBenefit);
314 target
315 .addIllegalOp<LLVM::CosOp, LLVM::ExpOp, LLVM::Exp2Op, LLVM::LogOp,
316 LLVM::Log10Op, LLVM::Log2Op, LLVM::SinOp, LLVM::SqrtOp>();
317 }
318 target.addLegalDialect<BuiltinDialect, LLVM::LLVMDialect>();
319 // The size-1-vector patterns unwrap to the scalar intrinsic via
320 // vector.extract / vector.broadcast; these must be legal for the partial
321 // conversion to succeed.
322 target.addLegalOp<vector::ExtractOp, vector::BroadcastOp>();
323 if (failed(
324 applyPartialConversion(getOperation(), target, std::move(patterns))))
325 signalPassFailure();
326}
return success()
b getContext())
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...
Definition Location.h:76
MLIRContext is the top-level object for a collection of MLIR operations.
Definition MLIRContext.h:63
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.
Definition Operation.h:115
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),...
Definition Operation.h:849
MLIRContext * getContext()
Return the context this operation is associated with.
Definition Operation.h:233
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...
Definition Types.h:74
bool isFloat() const
Return true if this is an float type (with the specified width).
Definition Types.cpp:47
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
Definition Value.h:96
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`.
detail::InFlightRemark failed(Location loc, RemarkOpts opts)
Report an optimization remark that failed.
Definition Remarks.h:734
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.