MLIR 24.0.0git
MathToSPIRV.cpp
Go to the documentation of this file.
1//===- MathToSPIRV.cpp - Math to SPIR-V Patterns --------------------------===//
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//
9// This file implements patterns to convert Math dialect to SPIR-V dialect.
10//
11//===----------------------------------------------------------------------===//
12
18#include "mlir/IR/Matchers.h"
21#include "llvm/ADT/STLExtras.h"
22#include "llvm/ADT/TypeSwitch.h"
23#include "llvm/Support/FormatVariadic.h"
24
25#define DEBUG_TYPE "math-to-spirv-pattern"
26
27using namespace mlir;
28
29//===----------------------------------------------------------------------===//
30// Utility functions
31//===----------------------------------------------------------------------===//
32
33/// Creates a 32-bit scalar/vector integer constant. Returns nullptr if the
34/// given type is not a 32-bit scalar/vector type.
36 OpBuilder &builder, Location loc) {
37 if (auto vectorType = dyn_cast<VectorType>(type)) {
38 if (!vectorType.getElementType().isInteger(32))
39 return nullptr;
40 SmallVector<int> values(vectorType.getNumElements(), value);
41 return spirv::ConstantOp::create(builder, loc, type,
42 builder.getI32VectorAttr(values));
43 }
44 if (type.isInteger(32))
45 return spirv::ConstantOp::create(builder, loc, type,
46 builder.getI32IntegerAttr(value));
47
48 return nullptr;
49}
50
51/// Check if the type is supported by math-to-spirv conversion. We expect to
52/// only see scalars and vectors at this point, with higher-level types already
53/// lowered.
54static bool isSupportedSourceType(Type originalType) {
55 if (originalType.isIntOrIndexOrFloat())
56 return true;
57
58 if (auto vecTy = dyn_cast<VectorType>(originalType)) {
59 if (!vecTy.getElementType().isIntOrIndexOrFloat())
60 return false;
61 if (vecTy.isScalable())
62 return false;
63 if (vecTy.getRank() > 1)
64 return false;
65
66 return true;
67 }
68
69 return false;
70}
71
72/// Check if all `sourceOp` types are supported by math-to-spirv conversion.
73/// Notify of a match failure othwerise and return a `failure` result.
74/// This is intended to simplify type checks in `OpConversionPattern`s.
75static LogicalResult checkSourceOpTypes(ConversionPatternRewriter &rewriter,
76 Operation *sourceOp) {
77 auto allTypes = llvm::to_vector(sourceOp->getOperandTypes());
78 llvm::append_range(allTypes, sourceOp->getResultTypes());
79
80 for (Type ty : allTypes) {
81 if (!isSupportedSourceType(ty)) {
82 return rewriter.notifyMatchFailure(
83 sourceOp,
84 llvm::formatv(
85 "unsupported source type for Math to SPIR-V conversion: {0}",
86 ty));
87 }
88 }
89
90 return success();
91}
92
93//===----------------------------------------------------------------------===//
94// Operation conversion
95//===----------------------------------------------------------------------===//
96
97// Note that DRR cannot be used for the patterns in this file: we may need to
98// convert type along the way, which requires ConversionPattern. DRR generates
99// normal RewritePattern.
100
101namespace {
102/// Converts elementwise unary, binary, and ternary standard operations to
103/// SPIR-V operations. Checks that source `Op` types are supported.
104template <typename Op, typename SPIRVOp>
105struct CheckedElementwiseOpPattern final
106 : public spirv::ElementwiseOpPattern<Op, SPIRVOp> {
107 using BasePattern = typename spirv::ElementwiseOpPattern<Op, SPIRVOp>;
108 using BasePattern::BasePattern;
109
110 LogicalResult
111 matchAndRewrite(Op op, typename Op::Adaptor adaptor,
112 ConversionPatternRewriter &rewriter) const override {
113 if (LogicalResult res = checkSourceOpTypes(rewriter, op); failed(res))
114 return res;
115
116 return BasePattern::matchAndRewrite(op, adaptor, rewriter);
117 }
118};
119
120/// Converts math.copysign to SPIR-V ops.
121struct CopySignPattern final : public OpConversionPattern<math::CopySignOp> {
122 using Base::Base;
123
124 LogicalResult
125 matchAndRewrite(math::CopySignOp copySignOp, OpAdaptor adaptor,
126 ConversionPatternRewriter &rewriter) const override {
127 if (LogicalResult res = checkSourceOpTypes(rewriter, copySignOp);
128 failed(res))
129 return res;
130
131 // Defer to the CL copysign op on Kernel targets.
132 auto &typeConverter = *getTypeConverter<SPIRVTypeConverter>();
133 if (typeConverter.getTargetEnv().allows(spirv::Capability::Kernel))
134 return rewriter.notifyMatchFailure(copySignOp,
135 "Kernel target has native CL op");
136
137 Type type = getTypeConverter()->convertType(copySignOp.getType());
138 if (!type)
139 return failure();
140
141 FloatType floatType;
142 if (auto scalarType = dyn_cast<FloatType>(copySignOp.getType())) {
143 floatType = scalarType;
144 } else if (auto vectorType = dyn_cast<VectorType>(copySignOp.getType())) {
145 floatType = cast<FloatType>(vectorType.getElementType());
146 } else {
147 return failure();
148 }
149
150 Location loc = copySignOp.getLoc();
151 int bitwidth = floatType.getWidth();
152 Type intType = rewriter.getIntegerType(bitwidth);
153 uint64_t intValue = uint64_t(1) << (bitwidth - 1);
154
155 Value signMask = spirv::ConstantOp::create(
156 rewriter, loc, intType, rewriter.getIntegerAttr(intType, intValue));
157 Value valueMask = spirv::ConstantOp::create(
158 rewriter, loc, intType,
159 rewriter.getIntegerAttr(intType, intValue - 1u));
160
161 if (auto vectorType = dyn_cast<VectorType>(type)) {
162 assert(vectorType.getRank() == 1);
163 int count = vectorType.getNumElements();
164 intType = VectorType::get(count, intType);
165
166 Repeated<Value> signSplat(count, signMask);
167 signMask = spirv::CompositeConstructOp::create(rewriter, loc, intType,
168 signSplat);
169
170 Repeated<Value> valueSplat(count, valueMask);
171 valueMask = spirv::CompositeConstructOp::create(rewriter, loc, intType,
172 valueSplat);
173 }
174
175 Value lhsCast =
176 spirv::BitcastOp::create(rewriter, loc, intType, adaptor.getLhs());
177 Value rhsCast =
178 spirv::BitcastOp::create(rewriter, loc, intType, adaptor.getRhs());
179
180 Value value = spirv::BitwiseAndOp::create(rewriter, loc, intType,
181 ValueRange{lhsCast, valueMask});
182 Value sign = spirv::BitwiseAndOp::create(rewriter, loc, intType,
183 ValueRange{rhsCast, signMask});
184
185 Value result = spirv::BitwiseOrOp::create(rewriter, loc, intType,
186 ValueRange{value, sign});
187 rewriter.replaceOpWithNewOp<spirv::BitcastOp>(copySignOp, type, result);
188 return success();
189 }
190};
191
192/// Converts math.ctlz to SPIR-V ops.
193///
194/// OpenCL targets lower math.ctlz directly to OpenCL.std clz via the generic
195/// elementwise pattern. This pattern handles the shader fallback.
196///
197/// SPIR-V does not have a direct operations for counting leading zeros for
198/// glsl. If Shader capability is supported, we can leverage GL FindUMsb to
199/// calculate it.
200struct CountLeadingZerosPattern final
201 : public OpConversionPattern<math::CountLeadingZerosOp> {
202 using Base::Base;
203
204 LogicalResult
205 matchAndRewrite(math::CountLeadingZerosOp countOp, OpAdaptor adaptor,
206 ConversionPatternRewriter &rewriter) const override {
207 if (LogicalResult res = checkSourceOpTypes(rewriter, countOp); failed(res))
208 return res;
209
210 Type type = getTypeConverter()->convertType(countOp.getType());
211 if (!type)
212 return failure();
213
214 auto &typeConverter = *getTypeConverter<SPIRVTypeConverter>();
215 if (!typeConverter.getTargetEnv().allows(spirv::Capability::Shader))
216 return rewriter.notifyMatchFailure(countOp, "requires Shader capability");
217
218 // The GL FindUMsb fallback only supports 32-bit integer types for now.
219 unsigned bitwidth = 0;
220 if (isa<IntegerType>(type))
221 bitwidth = type.getIntOrFloatBitWidth();
222 if (auto vectorType = dyn_cast<VectorType>(type))
223 bitwidth = vectorType.getElementTypeBitWidth();
224 if (bitwidth != 32)
225 return failure();
226
227 Location loc = countOp.getLoc();
228 Value input = adaptor.getOperand();
229 Value val1 = getScalarOrVectorI32Constant(type, 1, rewriter, loc);
230 Value val31 = getScalarOrVectorI32Constant(type, 31, rewriter, loc);
231 Value val32 = getScalarOrVectorI32Constant(type, 32, rewriter, loc);
232
233 Value msb = spirv::GLFindUMsbOp::create(rewriter, loc, input);
234 // We need to subtract from 31 given that the index returned by GLSL
235 // FindUMsb is counted from the least significant bit. Theoretically this
236 // also gives the correct result even if the integer has all zero bits, in
237 // which case GL FindUMsb would return -1.
238 Value subMsb = spirv::ISubOp::create(rewriter, loc, val31, msb);
239 // However, certain Vulkan implementations have driver bugs for the corner
240 // case where the input is zero. And.. it can be smart to optimize a select
241 // only involving the corner case. So separately compute the result when the
242 // input is either zero or one.
243 Value subInput = spirv::ISubOp::create(rewriter, loc, val32, input);
244 Value cmp = spirv::ULessThanEqualOp::create(rewriter, loc, input, val1);
245 rewriter.replaceOpWithNewOp<spirv::SelectOp>(countOp, cmp, subInput,
246 subMsb);
247 return success();
248 }
249};
250
251/// Converts math.cttz to GL FindILsb. GL FindILsb returns -1 for a zero
252/// input while math.cttz must return the bitwidth, so the zero case is
253/// patched up with a select.
254struct CountTrailingZerosPattern final
255 : public OpConversionPattern<math::CountTrailingZerosOp> {
256 using Base::Base;
257
258 LogicalResult
259 matchAndRewrite(math::CountTrailingZerosOp countOp, OpAdaptor adaptor,
260 ConversionPatternRewriter &rewriter) const override {
261 if (LogicalResult res = checkSourceOpTypes(rewriter, countOp); failed(res))
262 return res;
263
264 Type type = getTypeConverter()->convertType(countOp.getType());
265 if (!type)
266 return failure();
267
268 auto &typeConverter = *getTypeConverter<SPIRVTypeConverter>();
269 if (!typeConverter.getTargetEnv().allows(spirv::Capability::Shader))
270 return rewriter.notifyMatchFailure(countOp, "requires Shader capability");
271
272 unsigned bitwidth = 0;
273 if (isa<IntegerType>(type))
274 bitwidth = type.getIntOrFloatBitWidth();
275 else if (auto vectorType = dyn_cast<VectorType>(type))
276 bitwidth = vectorType.getElementTypeBitWidth();
277
278 Location loc = countOp.getLoc();
279 Value input = adaptor.getOperand();
280 Value val0 = spirv::ConstantOp::getZero(type, loc, rewriter);
281 Type elemType = getElementTypeOrSelf(type);
282 Attribute bwAttr = IntegerAttr::get(elemType, bitwidth);
283 if (auto vecType = dyn_cast<VectorType>(type))
284 bwAttr = SplatElementsAttr::get(vecType, bwAttr);
285 Value valBitwidth = spirv::ConstantOp::create(rewriter, loc, type, bwAttr);
286
287 Value lsb = spirv::GLFindILsbOp::create(rewriter, loc, input);
288 Value isZero = spirv::IEqualOp::create(rewriter, loc, input, val0);
289 rewriter.replaceOpWithNewOp<spirv::SelectOp>(countOp, isZero, valBitwidth,
290 lsb);
291 return success();
292 }
293};
294
295/// Converts math.expm1 to SPIR-V ops.
296///
297/// SPIR-V does not have a direct operations for exp(x)-1. Explicitly lower to
298/// these operations.
299template <typename ExpOp>
300struct ExpM1OpPattern final : public OpConversionPattern<math::ExpM1Op> {
301 using Base::Base;
302
303 LogicalResult
304 matchAndRewrite(math::ExpM1Op operation, OpAdaptor adaptor,
305 ConversionPatternRewriter &rewriter) const override {
306 assert(adaptor.getOperands().size() == 1);
307 if (LogicalResult res = checkSourceOpTypes(rewriter, operation);
308 failed(res))
309 return res;
310
311 Location loc = operation.getLoc();
312 Type type = this->getTypeConverter()->convertType(operation.getType());
313 if (!type)
314 return failure();
315
316 Value exp = ExpOp::create(rewriter, loc, type, adaptor.getOperand());
317 auto one = spirv::ConstantOp::getOne(type, loc, rewriter);
318 rewriter.replaceOpWithNewOp<spirv::FSubOp>(operation, exp, one);
319 return success();
320 }
321};
322
323/// Converts math.log1p to SPIR-V ops.
324///
325/// SPIR-V does not have a direct operations for log(1+x). Explicitly lower to
326/// these operations.
327template <typename LogOp>
328struct Log1pOpPattern final : public OpConversionPattern<math::Log1pOp> {
329 using Base::Base;
330
331 LogicalResult
332 matchAndRewrite(math::Log1pOp operation, OpAdaptor adaptor,
333 ConversionPatternRewriter &rewriter) const override {
334 assert(adaptor.getOperands().size() == 1);
335 if (LogicalResult res = checkSourceOpTypes(rewriter, operation);
336 failed(res))
337 return res;
338
339 Location loc = operation.getLoc();
340 Type type = this->getTypeConverter()->convertType(operation.getType());
341 if (!type)
342 return failure();
343
344 auto one = spirv::ConstantOp::getOne(type, operation.getLoc(), rewriter);
345 Value onePlus =
346 spirv::FAddOp::create(rewriter, loc, one, adaptor.getOperand());
347 rewriter.replaceOpWithNewOp<LogOp>(operation, type, onePlus);
348 return success();
349 }
350};
351
352/// Converts math.log10 to GLSL SPIR-V ops.
353///
354/// GLSL.std.450 has no Log10 instruction. Lower it as:
355/// log10(x) = log(x) * 1/log(10)
356struct Log10OpPattern final : public OpConversionPattern<math::Log10Op> {
357 using Base::Base;
358
359 static constexpr double log10Reciprocal =
360 0.4342944819032518276511289189166050822943970058036665661144537832;
361
362 LogicalResult
363 matchAndRewrite(math::Log10Op operation, OpAdaptor adaptor,
364 ConversionPatternRewriter &rewriter) const override {
365 assert(adaptor.getOperands().size() == 1);
366 if (LogicalResult res = checkSourceOpTypes(rewriter, operation);
367 failed(res))
368 return res;
369
370 Location loc = operation.getLoc();
371 Type type = this->getTypeConverter()->convertType(operation.getType());
372 if (!type)
373 return rewriter.notifyMatchFailure(operation, "type conversion failed");
374
375 auto getConstantValue = [&](double value) {
376 if (auto floatType = dyn_cast<FloatType>(type)) {
377 return spirv::ConstantOp::create(
378 rewriter, loc, type, rewriter.getFloatAttr(floatType, value));
379 }
380 if (auto vectorType = dyn_cast<VectorType>(type)) {
381 Type elemType = vectorType.getElementType();
382
383 if (isa<FloatType>(elemType)) {
384 return spirv::ConstantOp::create(
385 rewriter, loc, type,
387 vectorType, FloatAttr::get(elemType, value).getValue()));
388 }
389 }
390 llvm_unreachable("unimplemented type for log10");
391 };
392
393 Value constantValue = getConstantValue(log10Reciprocal);
394 Value log = spirv::GLLogOp::create(rewriter, loc, adaptor.getOperand());
395 rewriter.replaceOpWithNewOp<spirv::FMulOp>(operation, type, log,
396 constantValue);
397 return success();
398 }
399};
400
401/// Converts math.powf to SPIRV-Ops.
402struct PowFOpPattern final : public OpConversionPattern<math::PowFOp> {
403 using Base::Base;
404
405 LogicalResult
406 matchAndRewrite(math::PowFOp powfOp, OpAdaptor adaptor,
407 ConversionPatternRewriter &rewriter) const override {
408 if (LogicalResult res = checkSourceOpTypes(rewriter, powfOp); failed(res))
409 return res;
410
411 Type dstType = getTypeConverter()->convertType(powfOp.getType());
412 if (!dstType)
413 return failure();
414
415 Location loc = powfOp.getLoc();
416 Type operandType = adaptor.getRhs().getType();
417
418 // Parity-based lowering requires an integer-valued constant exponent.
419 // Otherwise fall back to exp(y*log(x)), which yields NaN for x<0 (matches
420 // C).
421 auto isOdd = [](const APFloat &v) {
422 APSInt i(/*BitWidth=*/64, /*isUnsigned=*/false);
423 bool ignored;
424 v.convertToInteger(i, APFloat::rmTowardZero, &ignored);
425 return i[0];
426 };
427
428 SmallVector<bool> oddMask;
429 Attribute rhsAttr;
430 if (matchPattern(adaptor.getRhs(), m_Constant(&rhsAttr))) {
431 TypeSwitch<Attribute>(rhsAttr)
432 .Case([&](FloatAttr a) {
433 if (a.getValue().isInteger())
434 oddMask.push_back(isOdd(a.getValue()));
435 })
436 .Case([&](SplatElementsAttr a) {
437 APFloat splat = a.getSplatValue<APFloat>();
438 if (splat.isInteger())
439 oddMask.push_back(isOdd(splat));
440 })
441 .Case([&](DenseElementsAttr a) {
442 SmallVector<bool> mask;
443 for (const APFloat &elt : a.getValues<APFloat>()) {
444 if (!elt.isInteger())
445 return;
446 mask.push_back(isOdd(elt));
447 }
448 oddMask = std::move(mask);
449 });
450 }
451
452 if (oddMask.empty()) {
453 Value log = spirv::GLLogOp::create(rewriter, loc, adaptor.getLhs());
454 Value mul = spirv::FMulOp::create(rewriter, loc, adaptor.getRhs(), log);
455 rewriter.replaceOpWithNewOp<spirv::GLExpOp>(powfOp, mul);
456 return success();
457 }
458
459 // GL.Pow is undefined for x < 0; take abs and conditionally negate the
460 // result for lanes whose exponent is odd.
461 Value abs = spirv::GLFAbsOp::create(rewriter, loc, adaptor.getLhs());
462 Value pow = spirv::GLPowOp::create(rewriter, loc, abs, adaptor.getRhs());
463
464 // No odd-parity element: result has the same sign as |lhs|^rhs >= 0.
465 if (llvm::none_of(oddMask, [](bool b) { return b; })) {
466 rewriter.replaceOp(powfOp, pow);
467 return success();
468 }
469
470 Value zero = spirv::ConstantOp::getZero(operandType, loc, rewriter);
471 Value lessThan =
472 spirv::FOrdLessThanOp::create(rewriter, loc, adaptor.getLhs(), zero);
473 Value negate = spirv::FNegateOp::create(rewriter, loc, pow);
474
475 Value shouldNegate;
476 if (llvm::all_equal(oddMask)) {
477 // Every lane has odd exponent: negate iff lhs < 0.
478 shouldNegate = lessThan;
479 } else {
480 // Mixed parity (non-splat dense vector): AND lhs<0 with a per-element
481 // constant odd-mask.
482 auto vecType = cast<VectorType>(operandType);
483 auto maskType = VectorType::get(vecType.getShape(), rewriter.getI1Type());
484 Value oddConst = spirv::ConstantOp::create(
485 rewriter, loc, maskType, DenseElementsAttr::get(maskType, oddMask));
486 shouldNegate =
487 spirv::LogicalAndOp::create(rewriter, loc, lessThan, oddConst);
488 }
489
490 rewriter.replaceOpWithNewOp<spirv::SelectOp>(powfOp, shouldNegate, negate,
491 pow);
492 return success();
493 }
494};
495
496/// Converts math.fpowi to spirv.CL.pown.
497struct PowIOpPattern final : public OpConversionPattern<math::FPowIOp> {
498 using Base::Base;
499
500 LogicalResult
501 matchAndRewrite(math::FPowIOp op, OpAdaptor adaptor,
502 ConversionPatternRewriter &rewriter) const override {
503 if (LogicalResult res = checkSourceOpTypes(rewriter, op); failed(res))
504 return res;
505
506 Type dstType = getTypeConverter()->convertType(op.getType());
507 if (!dstType)
508 return failure();
509
510 rewriter.replaceOpWithNewOp<spirv::CLPownOp>(op, dstType, adaptor.getLhs(),
511 adaptor.getRhs());
512 return success();
513 }
514};
515
516/// Converts math.fpowi to GLSL SPIR-V ops. GL has no integer-power op, so the
517/// exponent is converted to float and lowered through spirv.GL.Pow. As GL.Pow
518/// is undefined for a negative base, the base is made positive and the result
519/// is negated when the base is negative and the exponent is odd.
520struct PowIOpGLPattern final : public OpConversionPattern<math::FPowIOp> {
521 using Base::Base;
522
523 LogicalResult
524 matchAndRewrite(math::FPowIOp op, OpAdaptor adaptor,
525 ConversionPatternRewriter &rewriter) const override {
526 if (LogicalResult res = checkSourceOpTypes(rewriter, op); failed(res))
527 return res;
528
529 Type dstType = getTypeConverter()->convertType(op.getType());
530 if (!dstType)
531 return failure();
532
533 Location loc = op.getLoc();
534 Value base = adaptor.getLhs();
535 Value power = adaptor.getRhs();
536
537 Value expFloat =
538 spirv::ConvertSToFOp::create(rewriter, loc, dstType, power);
539 Value abs = spirv::GLFAbsOp::create(rewriter, loc, base);
540 Value pow = spirv::GLPowOp::create(rewriter, loc, abs, expFloat);
541
542 Value zeroF = spirv::ConstantOp::getZero(dstType, loc, rewriter);
543 Value lessThan = spirv::FOrdLessThanOp::create(rewriter, loc, base, zeroF);
544
545 Type powerType = power.getType();
546 Value oneI = spirv::ConstantOp::getOne(powerType, loc, rewriter);
547 Value lowBit = spirv::BitwiseAndOp::create(rewriter, loc, power, oneI);
548 Value isOdd = spirv::IEqualOp::create(rewriter, loc, lowBit, oneI);
549
550 Value shouldNegate =
551 spirv::LogicalAndOp::create(rewriter, loc, lessThan, isOdd);
552 Value negate = spirv::FNegateOp::create(rewriter, loc, pow);
553 rewriter.replaceOpWithNewOp<spirv::SelectOp>(op, shouldNegate, negate, pow);
554 return success();
555 }
556};
557
558/// Converts math.sincos to SPIR-V ops.
559///
560/// SPIR-V has no fused sincos instruction, so emit separate sin and cos ops
561/// sharing the same operand.
562template <typename SinOp, typename CosOp>
563struct SincosOpPattern final : public OpConversionPattern<math::SincosOp> {
564 using Base::Base;
565
566 LogicalResult
567 matchAndRewrite(math::SincosOp operation, OpAdaptor adaptor,
568 ConversionPatternRewriter &rewriter) const override {
569 if (LogicalResult res = checkSourceOpTypes(rewriter, operation);
570 failed(res))
571 return res;
572
573 Type type =
574 getTypeConverter()->convertType(operation.getOperand().getType());
575 if (!type)
576 return failure();
577
578 Location loc = operation.getLoc();
579 Value sin = SinOp::create(rewriter, loc, type, adaptor.getOperand());
580 Value cos = CosOp::create(rewriter, loc, type, adaptor.getOperand());
581 rewriter.replaceOp(operation, {sin, cos});
582 return success();
583 }
584};
585
586/// Converts math.round to GLSL SPIRV extended ops.
587struct RoundOpPattern final : public OpConversionPattern<math::RoundOp> {
588 using Base::Base;
589
590 LogicalResult
591 matchAndRewrite(math::RoundOp roundOp, OpAdaptor adaptor,
592 ConversionPatternRewriter &rewriter) const override {
593 if (LogicalResult res = checkSourceOpTypes(rewriter, roundOp); failed(res))
594 return res;
595
596 Location loc = roundOp.getLoc();
597 auto ty = getTypeConverter()->convertType(adaptor.getOperand().getType());
598 if (!ty) {
599 return rewriter.notifyMatchFailure(
600 roundOp->getLoc(),
601 llvm::formatv("failed to convert type {0} for SPIR-V",
602 roundOp.getType()));
603 }
604
605 Type ety = getElementTypeOrSelf(ty);
606
607 auto zero = spirv::ConstantOp::getZero(ty, loc, rewriter);
608 auto one = spirv::ConstantOp::getOne(ty, loc, rewriter);
609 Value half;
610 if (VectorType vty = dyn_cast<VectorType>(ty)) {
611 half = spirv::ConstantOp::create(
612 rewriter, loc, vty,
614 rewriter.getFloatAttr(ety, 0.5).getValue()));
615 } else {
616 half = spirv::ConstantOp::create(rewriter, loc, ty,
617 rewriter.getFloatAttr(ety, 0.5));
618 }
619
620 auto abs = spirv::GLFAbsOp::create(rewriter, loc, adaptor.getOperand());
621 auto floor = spirv::GLFloorOp::create(rewriter, loc, abs);
622 auto sub = spirv::FSubOp::create(rewriter, loc, abs, floor);
623 auto greater =
624 spirv::FOrdGreaterThanEqualOp::create(rewriter, loc, sub, half);
625 auto select = spirv::SelectOp::create(rewriter, loc, greater, one, zero);
626 auto add = spirv::FAddOp::create(rewriter, loc, floor, select);
627 rewriter.replaceOpWithNewOp<math::CopySignOp>(roundOp, add,
628 adaptor.getOperand());
629 return success();
630 }
631};
632
633} // namespace
634
635//===----------------------------------------------------------------------===//
636// Pattern population
637//===----------------------------------------------------------------------===//
638
639namespace mlir {
641 RewritePatternSet &patterns) {
642 // Core patterns
643 patterns
644 .add<CopySignPattern,
645 CheckedElementwiseOpPattern<math::CtPopOp, spirv::BitCountOp>,
646 CheckedElementwiseOpPattern<math::IsInfOp, spirv::IsInfOp>,
647 CheckedElementwiseOpPattern<math::IsNaNOp, spirv::IsNanOp>,
648 CheckedElementwiseOpPattern<math::IsFiniteOp, spirv::IsFiniteOp>,
649 CheckedElementwiseOpPattern<math::IsNormalOp, spirv::IsNormalOp>>(
650 typeConverter, patterns.getContext());
651
652 // GLSL patterns
653 patterns
654 .add<CountLeadingZerosPattern, CountTrailingZerosPattern,
655 Log1pOpPattern<spirv::GLLogOp>, Log10OpPattern,
656 ExpM1OpPattern<spirv::GLExpOp>, PowFOpPattern, PowIOpGLPattern,
657 RoundOpPattern, SincosOpPattern<spirv::GLSinOp, spirv::GLCosOp>,
658 CheckedElementwiseOpPattern<math::AbsFOp, spirv::GLFAbsOp>,
659 CheckedElementwiseOpPattern<math::AbsIOp, spirv::GLSAbsOp>,
660 CheckedElementwiseOpPattern<math::AtanOp, spirv::GLAtanOp>,
661 CheckedElementwiseOpPattern<math::Atan2Op, spirv::GLAtan2Op>,
662 CheckedElementwiseOpPattern<math::CeilOp, spirv::GLCeilOp>,
663 CheckedElementwiseOpPattern<math::ClampFOp, spirv::GLFClampOp>,
664 CheckedElementwiseOpPattern<math::CosOp, spirv::GLCosOp>,
665 CheckedElementwiseOpPattern<math::ExpOp, spirv::GLExpOp>,
666 CheckedElementwiseOpPattern<math::Exp2Op, spirv::GLExp2Op>,
667 CheckedElementwiseOpPattern<math::FloorOp, spirv::GLFloorOp>,
668 CheckedElementwiseOpPattern<math::FmaOp, spirv::GLFmaOp>,
669 CheckedElementwiseOpPattern<math::LogOp, spirv::GLLogOp>,
670 CheckedElementwiseOpPattern<math::Log2Op, spirv::GLLog2Op>,
671 CheckedElementwiseOpPattern<math::RoundEvenOp, spirv::GLRoundEvenOp>,
672 CheckedElementwiseOpPattern<math::RsqrtOp, spirv::GLInverseSqrtOp>,
673 CheckedElementwiseOpPattern<math::SinOp, spirv::GLSinOp>,
674 CheckedElementwiseOpPattern<math::SqrtOp, spirv::GLSqrtOp>,
675 CheckedElementwiseOpPattern<math::TanhOp, spirv::GLTanhOp>,
676 CheckedElementwiseOpPattern<math::TanOp, spirv::GLTanOp>,
677 CheckedElementwiseOpPattern<math::TruncOp, spirv::GLTruncOp>,
678 CheckedElementwiseOpPattern<math::AsinOp, spirv::GLAsinOp>,
679 CheckedElementwiseOpPattern<math::AcosOp, spirv::GLAcosOp>,
680 CheckedElementwiseOpPattern<math::SinhOp, spirv::GLSinhOp>,
681 CheckedElementwiseOpPattern<math::CoshOp, spirv::GLCoshOp>,
682 CheckedElementwiseOpPattern<math::AsinhOp, spirv::GLAsinhOp>,
683 CheckedElementwiseOpPattern<math::AcoshOp, spirv::GLAcoshOp>,
684 CheckedElementwiseOpPattern<math::AtanhOp, spirv::GLAtanhOp>>(
685 typeConverter, patterns.getContext());
686
687 // OpenCL patterns
688 patterns.add<
689 Log1pOpPattern<spirv::CLLogOp>, ExpM1OpPattern<spirv::CLExpOp>,
690 SincosOpPattern<spirv::CLSinOp, spirv::CLCosOp>,
691 CheckedElementwiseOpPattern<math::AbsFOp, spirv::CLFAbsOp>,
692 CheckedElementwiseOpPattern<math::AbsIOp, spirv::CLSAbsOp>,
693 CheckedElementwiseOpPattern<math::CountLeadingZerosOp, spirv::CLClzOp>,
694 CheckedElementwiseOpPattern<math::AtanOp, spirv::CLAtanOp>,
695 CheckedElementwiseOpPattern<math::Atan2Op, spirv::CLAtan2Op>,
696 CheckedElementwiseOpPattern<math::CbrtOp, spirv::CLCbrtOp>,
697 CheckedElementwiseOpPattern<math::CeilOp, spirv::CLCeilOp>,
698 CheckedElementwiseOpPattern<math::CopySignOp, spirv::CLCopysignOp>,
699 CheckedElementwiseOpPattern<math::CosOp, spirv::CLCosOp>,
700 CheckedElementwiseOpPattern<math::ErfOp, spirv::CLErfOp>,
701 CheckedElementwiseOpPattern<math::ErfcOp, spirv::CLErfcOp>,
702 CheckedElementwiseOpPattern<math::ExpOp, spirv::CLExpOp>,
703 CheckedElementwiseOpPattern<math::Exp2Op, spirv::CLExp2Op>,
704 CheckedElementwiseOpPattern<math::FloorOp, spirv::CLFloorOp>,
705 CheckedElementwiseOpPattern<math::FmaOp, spirv::CLFmaOp>,
706 CheckedElementwiseOpPattern<math::LogOp, spirv::CLLogOp>,
707 CheckedElementwiseOpPattern<math::Log2Op, spirv::CLLog2Op>,
708 CheckedElementwiseOpPattern<math::Log10Op, spirv::CLLog10Op>,
709 CheckedElementwiseOpPattern<math::PowFOp, spirv::CLPowOp>, PowIOpPattern,
710 CheckedElementwiseOpPattern<math::RoundEvenOp, spirv::CLRintOp>,
711 CheckedElementwiseOpPattern<math::RoundOp, spirv::CLRoundOp>,
712 CheckedElementwiseOpPattern<math::RsqrtOp, spirv::CLRsqrtOp>,
713 CheckedElementwiseOpPattern<math::SinOp, spirv::CLSinOp>,
714 CheckedElementwiseOpPattern<math::SqrtOp, spirv::CLSqrtOp>,
715 CheckedElementwiseOpPattern<math::TanhOp, spirv::CLTanhOp>,
716 CheckedElementwiseOpPattern<math::TanOp, spirv::CLTanOp>,
717 CheckedElementwiseOpPattern<math::TruncOp, spirv::CLTruncOp>,
718 CheckedElementwiseOpPattern<math::AsinOp, spirv::CLAsinOp>,
719 CheckedElementwiseOpPattern<math::AcosOp, spirv::CLAcosOp>,
720 CheckedElementwiseOpPattern<math::SinhOp, spirv::CLSinhOp>,
721 CheckedElementwiseOpPattern<math::CoshOp, spirv::CLCoshOp>,
722 CheckedElementwiseOpPattern<math::AsinhOp, spirv::CLAsinhOp>,
723 CheckedElementwiseOpPattern<math::AcoshOp, spirv::CLAcoshOp>,
724 CheckedElementwiseOpPattern<math::AtanhOp, spirv::CLAtanhOp>>(
725 typeConverter, patterns.getContext());
726}
727
728} // namespace mlir
return success()
b
Return true if permutation is a valid permutation of the outer_dims_perm (case OuterOrInnerPerm::Oute...
static LogicalResult checkSourceOpTypes(ConversionPatternRewriter &rewriter, Operation *sourceOp)
Check if all sourceOp types are supported by math-to-spirv conversion.
static bool isSupportedSourceType(Type originalType)
Check if the type is supported by math-to-spirv conversion.
static Value getScalarOrVectorI32Constant(Type type, int value, OpBuilder &builder, Location loc)
Creates a 32-bit scalar/vector integer constant.
#define mul(a, b)
#define add(a, b)
IntegerAttr getI32IntegerAttr(int32_t value)
Definition Builders.cpp:208
DenseIntElementsAttr getI32VectorAttr(ArrayRef< int32_t > values)
Definition Builders.cpp:130
auto getValues() const
Return the held element values as a range of the given type.
std::enable_if_t<!std::is_base_of< Attribute, T >::value||std::is_same< Attribute, T >::value, T > getSplatValue() const
Return the splat value for this attribute.
static DenseElementsAttr get(ShapedType type, ArrayRef< Attribute > values)
Constructs a dense elements attribute from an array of element values.
static DenseFPElementsAttr get(const ShapedType &type, Arg &&arg)
Get an instance of a DenseFPElementsAttr with the given arguments.
This class defines the main interface for locations in MLIR and acts as a non-nullable wrapper around...
Definition Location.h:76
This class helps build Operations.
Definition Builders.h:210
Operation is the basic unit of execution within MLIR.
Definition Operation.h:87
operand_type_range getOperandTypes()
Definition Operation.h:422
result_type_range getResultTypes()
Definition Operation.h:453
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.
Type conversion from builtin types to SPIR-V types for shader interface.
Instances of the Type class are uniqued, have an immutable identifier and an optional mutable compone...
Definition Types.h:74
bool isIntOrIndexOrFloat() const
Return true if this is an integer (of any signedness), index, or float type.
Definition Types.cpp:122
bool isInteger() const
Return true if this is an integer type (with the specified width).
Definition Types.cpp:58
unsigned getIntOrFloatBitWidth() const
Return the bit width of an integer or a float type, assert failure on other types.
Definition Types.cpp:124
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
Definition Value.h:96
Type getType() const
Return the type of this value.
Definition Value.h:105
DynamicAPInt floor(const Fraction &f)
Definition Fraction.h:77
Fraction abs(const Fraction &f)
Definition Fraction.h:107
detail::InFlightRemark failed(Location loc, RemarkOpts opts)
Report an optimization remark that failed.
Definition Remarks.h:717
Include the generated interface declarations.
bool matchPattern(Value value, const Pattern &pattern)
Entry point for matching a pattern over a Value.
Definition Matchers.h:490
void populateMathToSPIRVPatterns(const SPIRVTypeConverter &typeConverter, RewritePatternSet &patterns)
Appends to a pattern list additional patterns for translating Math ops to SPIR-V ops.
Type getElementTypeOrSelf(Type type)
Return the element type or return the type itself.
llvm::TypeSwitch< T, ResultT > TypeSwitch
Definition LLVM.h:139
detail::constant_op_matcher m_Constant()
Matches a constant foldable operation.
Definition Matchers.h:369
Converts elementwise unary, binary and ternary standard operations to SPIR-V operations.
Definition Pattern.h:24