MLIR 24.0.0git
ArithToLLVM.cpp
Go to the documentation of this file.
1//===- ArithToLLVM.cpp - Arithmetic to LLVM dialect conversion -------===//
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
10
23#include <type_traits>
24
25namespace mlir {
26#define GEN_PASS_DEF_ARITHTOLLVMCONVERSIONPASS
27#include "mlir/Conversion/Passes.h.inc"
28} // namespace mlir
29
30using namespace mlir;
31
32namespace {
33
34/// Lowering pattern that matches only when the source op's rounding mode
35/// presence agrees with `HasRoundingMode`. This allows registering two
36/// instances of the same pattern for one source op: one that handles the
37/// unconstrained case (no rounding mode, lowering to a regular LLVM op) and
38/// one that handles the constrained case (rounding mode present, lowering to
39/// a constrained LLVM intrinsic).
40///
41/// * `HasRoundingMode`: the pattern matches if and only if the source op has
42/// a rounding mode attribute.
43/// * `AttrConvert`: attribute converter to translate source attributes to
44/// target attributes.
45/// * `FailOnUnsupportedFP`: whether to fail if the source op has unsupported
46/// floating point types.
47template <typename SourceOp, typename TargetOp, bool HasRoundingMode,
48 template <typename, typename> typename AttrConvert =
50 bool FailOnUnsupportedFP = false>
51struct ConstrainedVectorConvertToLLVMPattern
52 : public VectorConvertToLLVMPattern<SourceOp, TargetOp, AttrConvert,
53 FailOnUnsupportedFP> {
54 using VectorConvertToLLVMPattern<
55 SourceOp, TargetOp, AttrConvert,
56 FailOnUnsupportedFP>::VectorConvertToLLVMPattern;
57
58 LogicalResult
59 matchAndRewrite(SourceOp op, typename SourceOp::Adaptor adaptor,
60 ConversionPatternRewriter &rewriter) const override {
61 if (HasRoundingMode != static_cast<bool>(op.getRoundingModeAttr()))
62 return failure();
63 return VectorConvertToLLVMPattern<
64 SourceOp, TargetOp, AttrConvert,
65 FailOnUnsupportedFP>::matchAndRewrite(op, adaptor, rewriter);
66 }
67};
68
69/// No-op bitcast. Propagate type input arg if converted source and dest types
70/// are the same.
71struct IdentityBitcastLowering final
72 : public OpConversionPattern<arith::BitcastOp> {
73 using Base::Base;
74
75 LogicalResult
76 matchAndRewrite(arith::BitcastOp op, OpAdaptor adaptor,
77 ConversionPatternRewriter &rewriter) const final {
78 Value src = adaptor.getIn();
79 Type resultType = getTypeConverter()->convertType(op.getType());
80 if (src.getType() != resultType)
81 return rewriter.notifyMatchFailure(op, "Types are different");
82
83 rewriter.replaceOp(op, src);
84 return success();
85 }
86};
87
88//===----------------------------------------------------------------------===//
89// Straightforward Op Lowerings
90//===----------------------------------------------------------------------===//
91
92using AddFOpLowering =
93 ConstrainedVectorConvertToLLVMPattern<arith::AddFOp, LLVM::FAddOp,
94 /*HasRoundingMode=*/false,
96 /*FailOnUnsupportedFP=*/true>;
97using ConstrainedAddFOpLowering = ConstrainedVectorConvertToLLVMPattern<
98 arith::AddFOp, LLVM::ConstrainedFAddIntr, /*HasRoundingMode=*/true,
99 arith::AttrConverterConstrainedFPToLLVM, /*FailOnUnsupportedFP=*/true>;
100using AddIOpLowering =
101 VectorConvertToLLVMPattern<arith::AddIOp, LLVM::AddOp,
104using BitcastOpLowering =
106using DivFOpLowering =
107 ConstrainedVectorConvertToLLVMPattern<arith::DivFOp, LLVM::FDivOp,
108 /*HasRoundingMode=*/false,
110 /*FailOnUnsupportedFP=*/true>;
111using ConstrainedDivFOpLowering = ConstrainedVectorConvertToLLVMPattern<
112 arith::DivFOp, LLVM::ConstrainedFDivIntr, /*HasRoundingMode=*/true,
113 arith::AttrConverterConstrainedFPToLLVM, /*FailOnUnsupportedFP=*/true>;
114using DivSIOpLowering =
116using DivUIOpLowering =
118using ExtFOpLowering =
119 VectorConvertToLLVMPattern<arith::ExtFOp, LLVM::FPExtOp,
121 /*FailOnUnsupportedFP=*/true>;
122using ExtSIOpLowering =
124using ExtUIOpLowering =
125 VectorConvertToLLVMPattern<arith::ExtUIOp, LLVM::ZExtOp,
127using FPToSIOpLowering =
128 VectorConvertToLLVMPattern<arith::FPToSIOp, LLVM::FPToSIOp,
130 /*FailOnUnsupportedFP=*/true>;
131using FPToUIOpLowering =
132 VectorConvertToLLVMPattern<arith::FPToUIOp, LLVM::FPToUIOp,
134 /*FailOnUnsupportedFP=*/true>;
135using MaximumFOpLowering =
136 VectorConvertToLLVMPattern<arith::MaximumFOp, LLVM::MaximumOp,
138 /*FailOnUnsupportedFP=*/true>;
139using MaxNumFOpLowering =
140 VectorConvertToLLVMPattern<arith::MaxNumFOp, LLVM::MaxNumOp,
142 /*FailOnUnsupportedFP=*/true>;
143using MaximumNumFOpLowering =
144 VectorConvertToLLVMPattern<arith::MaximumNumFOp, LLVM::MaximumNumOp,
146 /*FailOnUnsupportedFP=*/true>;
147using MaxSIOpLowering =
149using MaxUIOpLowering =
151using MinimumFOpLowering =
152 VectorConvertToLLVMPattern<arith::MinimumFOp, LLVM::MinimumOp,
154 /*FailOnUnsupportedFP=*/true>;
155using MinNumFOpLowering =
156 VectorConvertToLLVMPattern<arith::MinNumFOp, LLVM::MinNumOp,
158 /*FailOnUnsupportedFP=*/true>;
159using MinimumNumFOpLowering =
160 VectorConvertToLLVMPattern<arith::MinimumNumFOp, LLVM::MinimumNumOp,
162 /*FailOnUnsupportedFP=*/true>;
163using MinSIOpLowering =
165using MinUIOpLowering =
167using MulFOpLowering =
168 ConstrainedVectorConvertToLLVMPattern<arith::MulFOp, LLVM::FMulOp,
169 /*HasRoundingMode=*/false,
171 /*FailOnUnsupportedFP=*/true>;
172using ConstrainedMulFOpLowering = ConstrainedVectorConvertToLLVMPattern<
173 arith::MulFOp, LLVM::ConstrainedFMulIntr, /*HasRoundingMode=*/true,
174 arith::AttrConverterConstrainedFPToLLVM, /*FailOnUnsupportedFP=*/true>;
175using MulIOpLowering =
176 VectorConvertToLLVMPattern<arith::MulIOp, LLVM::MulOp,
178using NegFOpLowering =
179 VectorConvertToLLVMPattern<arith::NegFOp, LLVM::FNegOp,
181 /*FailOnUnsupportedFP=*/true>;
183using RemFOpLowering =
184 VectorConvertToLLVMPattern<arith::RemFOp, LLVM::FRemOp,
186 /*FailOnUnsupportedFP=*/true>;
187using RemSIOpLowering =
189using RemUIOpLowering =
191using SelectOpLowering =
193using ShLIOpLowering =
194 VectorConvertToLLVMPattern<arith::ShLIOp, LLVM::ShlOp,
196using ShRSIOpLowering =
198using ShRUIOpLowering =
200using SIToFPOpLowering =
202using SubFOpLowering =
203 ConstrainedVectorConvertToLLVMPattern<arith::SubFOp, LLVM::FSubOp,
204 /*HasRoundingMode=*/false,
206 /*FailOnUnsupportedFP=*/true>;
207using ConstrainedSubFOpLowering = ConstrainedVectorConvertToLLVMPattern<
208 arith::SubFOp, LLVM::ConstrainedFSubIntr, /*HasRoundingMode=*/true,
209 arith::AttrConverterConstrainedFPToLLVM, /*FailOnUnsupportedFP=*/true>;
210using SubIOpLowering =
211 VectorConvertToLLVMPattern<arith::SubIOp, LLVM::SubOp,
213using TruncFOpLowering =
214 ConstrainedVectorConvertToLLVMPattern<arith::TruncFOp, LLVM::FPTruncOp,
215 /*HasRoundingMode=*/false,
217 /*FailOnUnsupportedFP=*/true>;
218using ConstrainedTruncFOpLowering = ConstrainedVectorConvertToLLVMPattern<
219 arith::TruncFOp, LLVM::ConstrainedFPTruncIntr, /*HasRoundingMode=*/true,
220 arith::AttrConverterConstrainedFPToLLVM, /*FailOnUnsupportedFP=*/true>;
221using TruncIOpLowering =
222 VectorConvertToLLVMPattern<arith::TruncIOp, LLVM::TruncOp,
224using UIToFPOpLowering =
225 VectorConvertToLLVMPattern<arith::UIToFPOp, LLVM::UIToFPOp,
227 /*FailOnUnsupportedFP=*/true>;
229
230//===----------------------------------------------------------------------===//
231// Op Lowering Patterns
232//===----------------------------------------------------------------------===//
233
234/// Directly lower to LLVM op.
235struct ConstantOpLowering : public ConvertOpToLLVMPattern<arith::ConstantOp> {
237
238 LogicalResult
239 matchAndRewrite(arith::ConstantOp op, OpAdaptor adaptor,
240 ConversionPatternRewriter &rewriter) const override;
241};
242
243/// The lowering of index_cast becomes an integer conversion since index
244/// becomes an integer. If the bit width of the source and target integer
245/// types is the same, just erase the cast. If the target type is wider,
246/// sign-extend the value, otherwise truncate it.
247template <typename OpTy, typename ExtCastTy>
248struct IndexCastOpLowering : public ConvertOpToLLVMPattern<OpTy> {
249 using ConvertOpToLLVMPattern<OpTy>::ConvertOpToLLVMPattern;
250
251 LogicalResult
252 matchAndRewrite(OpTy op, typename OpTy::Adaptor adaptor,
253 ConversionPatternRewriter &rewriter) const override;
254};
255
256using IndexCastOpSILowering =
257 IndexCastOpLowering<arith::IndexCastOp, LLVM::SExtOp>;
258using IndexCastOpUILowering =
259 IndexCastOpLowering<arith::IndexCastUIOp, LLVM::ZExtOp>;
260
261struct AddUIExtendedOpLowering
262 : public ConvertOpToLLVMPattern<arith::AddUIExtendedOp> {
264
265 LogicalResult
266 matchAndRewrite(arith::AddUIExtendedOp op, OpAdaptor adaptor,
267 ConversionPatternRewriter &rewriter) const override;
268};
269
270struct SubUIExtendedOpLowering
271 : public ConvertOpToLLVMPattern<arith::SubUIExtendedOp> {
273
274 LogicalResult
275 matchAndRewrite(arith::SubUIExtendedOp op, OpAdaptor adaptor,
276 ConversionPatternRewriter &rewriter) const override;
277};
278
279template <typename ArithMulOp, bool IsSigned>
280struct MulIExtendedOpLowering : public ConvertOpToLLVMPattern<ArithMulOp> {
281 using ConvertOpToLLVMPattern<ArithMulOp>::ConvertOpToLLVMPattern;
282
283 LogicalResult
284 matchAndRewrite(ArithMulOp op, typename ArithMulOp::Adaptor adaptor,
285 ConversionPatternRewriter &rewriter) const override;
286};
287
288using MulSIExtendedOpLowering =
289 MulIExtendedOpLowering<arith::MulSIExtendedOp, true>;
290using MulUIExtendedOpLowering =
291 MulIExtendedOpLowering<arith::MulUIExtendedOp, false>;
292
293struct CmpIOpLowering : public ConvertOpToLLVMPattern<arith::CmpIOp> {
295
296 LogicalResult
297 matchAndRewrite(arith::CmpIOp op, OpAdaptor adaptor,
298 ConversionPatternRewriter &rewriter) const override;
299};
300
301struct CmpFOpLowering : public ConvertOpToLLVMPattern<arith::CmpFOp> {
303
304 LogicalResult
305 matchAndRewrite(arith::CmpFOp op, OpAdaptor adaptor,
306 ConversionPatternRewriter &rewriter) const override;
307};
308
309/// Lower arith.convertf (same-bitwidth FP cast) to LLVM.
310///
311/// Extends to f32 via llvm.fpext, then truncates to the target type via
312/// llvm.fptrunc. This handles bf16 <-> f16, which is the only same-bitwidth
313/// pair of LLVM-supported FP types.
314struct ConvertFOpLowering : public ConvertOpToLLVMPattern<arith::ConvertFOp> {
316
317 LogicalResult
318 matchAndRewrite(arith::ConvertFOp op, OpAdaptor adaptor,
319 ConversionPatternRewriter &rewriter) const override {
321 *getTypeConverter()))
322 return rewriter.notifyMatchFailure(op, "unsupported floating point type");
323
324 // Only bf16 <-> f16 conversions are supported. There is currently no other
325 // pair of FP types that are valid LLVM types.
326 [[maybe_unused]] auto srcType = getElementTypeOrSelf(op.getIn().getType());
327 [[maybe_unused]] auto dstType = getElementTypeOrSelf(op.getType());
328 assert(((srcType.isBF16() && dstType.isF16()) ||
329 (srcType.isF16() && dstType.isBF16())) &&
330 "only bf16 <-> f16 conversions are supported");
331
332 Type convertedType = getTypeConverter()->convertType(op.getType());
333 if (!convertedType)
334 return rewriter.notifyMatchFailure(op, "failed to convert result type");
335
336 Value input = adaptor.getIn();
337 Location loc = op.getLoc();
338
339 if (!isa<LLVM::LLVMArrayType>(input.getType())) {
340 rewriter.replaceOp(op,
341 emitConversion(rewriter, loc, input, convertedType));
342 return success();
343 }
344
345 if (!isa<VectorType>(op.getType()))
346 return rewriter.notifyMatchFailure(op, "expected vector result type");
347
349 op.getOperation(), adaptor.getOperands(), *getTypeConverter(),
350 [&](Type llvm1DVectorTy, ValueRange operands) -> Value {
351 return emitConversion(rewriter, loc, operands.front(),
352 llvm1DVectorTy);
353 },
354 rewriter);
355 }
356
357private:
358 static Value emitConversion(ConversionPatternRewriter &rewriter, Location loc,
359 Value input, Type targetType) {
360 Type f32Scalar = Float32Type::get(rewriter.getContext());
361 Type f32Ty = f32Scalar;
362 if (auto vecTy = dyn_cast<VectorType>(targetType))
363 f32Ty = VectorType::get(vecTy.getShape(), f32Scalar);
364
365 Value ext = LLVM::FPExtOp::create(rewriter, loc, f32Ty, input);
366 return LLVM::FPTruncOp::create(rewriter, loc, targetType, ext);
367 }
368};
369
370struct SelectOpOneToNLowering : public ConvertOpToLLVMPattern<arith::SelectOp> {
373
374 LogicalResult
375 matchAndRewrite(arith::SelectOp op, Adaptor adaptor,
376 ConversionPatternRewriter &rewriter) const override;
377};
378
379} // namespace
380
381//===----------------------------------------------------------------------===//
382// ConstantOpLowering
383//===----------------------------------------------------------------------===//
384
385/// Retypes `attr` for a `llvm.mlir.constant` of `resultType`. `arith.constant`
386/// requires the value attribute and the result to have the same type, but the
387/// type converter may map the element type to a different one, e.g. `index` to
388/// `i32` or `i64` depending on the configured index bitwidth. Build the
389/// attribute from the converted type so that the two agree wherever that is
390/// representable. Returns a null attribute if `attr` cannot be used for
391/// `resultType`.
392static TypedAttr convertConstantValue(TypedAttr attr, Type resultType) {
393 // Compare the element types, but explicitly ignore the non-scalar portions of
394 // types. This relaxation is required to support multi-dimensional vector
395 // constant that result in nested LLVM array values.
396 Type sourceElementType = LLVM::getConstantElementType(attr.getType());
397 Type targetElementType = LLVM::getConstantElementType(resultType);
398 if (sourceElementType == targetElementType)
399 return attr;
400
401 auto targetIntType = dyn_cast<IntegerType>(targetElementType);
402 if (!targetIntType)
403 return {};
404
405 // The converter maps the low-precision float types that have no LLVM
406 // equivalent to an integer of the same width. The attribute stays a float
407 // attribute in that case.
408 if (auto sourceFloatType = dyn_cast<FloatType>(sourceElementType)) {
409 if (sourceFloatType.getWidth() != targetIntType.getWidth())
410 return {};
411 return attr;
412 }
413
414 // Apart from those floats, `index` is the only element type the converter
415 // rewrites, so anything else is a malformed `arith.constant`. Bail out rather
416 // than reinterpret its value, which would silently change the constant (e.g.
417 // sign-extending an `i1` `true` to -1).
418 if (!isa<IndexType>(sourceElementType))
419 return {};
420 // `index` is signless but holds signed values, so narrowing to a smaller
421 // index bitwidth truncates and widening sign-extends.
422 unsigned width = targetIntType.getWidth();
423
424 if (auto intAttr = dyn_cast<IntegerAttr>(attr))
425 return IntegerAttr::get(targetIntType,
426 intAttr.getValue().sextOrTrunc(width));
427
428 auto retypeValues = [&](DenseIntElementsAttr values) {
429 return values.mapValues(targetIntType, [&](const APInt &value) {
430 return value.sextOrTrunc(width);
431 });
432 };
433
434 if (auto denseAttr = dyn_cast<DenseIntElementsAttr>(attr))
435 return retypeValues(denseAttr);
436
437 if (auto sparseAttr = dyn_cast<SparseElementsAttr>(attr))
438 return SparseElementsAttr::get(
439 cast<ShapedType>(attr.getType()).clone(targetIntType),
440 sparseAttr.getIndices(),
441 retypeValues(cast<DenseIntElementsAttr>(sparseAttr.getValues())));
442
443 // A resource-backed elements attribute refers to a blob laid out for its own
444 // element type. The blob cannot be rewritten here, only reinterpreted, which
445 // is correct exactly when the target type has the same width as the storage
446 // `index` uses in a blob.
447 if (auto resourceAttr = dyn_cast<DenseResourceElementsAttr>(attr)) {
448 if (width != IndexType::kInternalStorageBitWidth)
449 return {};
450 return DenseResourceElementsAttr::get(
451 cast<ShapedType>(attr.getType()).clone(targetIntType),
452 resourceAttr.getRawHandle());
453 }
454
455 return {};
456}
457
458LogicalResult
459ConstantOpLowering::matchAndRewrite(arith::ConstantOp op, OpAdaptor adaptor,
460 ConversionPatternRewriter &rewriter) const {
461 Type resultType = getTypeConverter()->convertType(op.getType());
462 if (!resultType)
463 return rewriter.notifyMatchFailure(op, "failed to convert result type");
464
465 TypedAttr value = convertConstantValue(op.getValue(), resultType);
466 if (!value)
467 return rewriter.notifyMatchFailure(
468 op, "failed to convert value attribute to the converted result type");
469
470 // `arith.constant` has no operands and a single result, so there is nothing
471 // for `oneToOneRewrite` to do here beyond converting `resultType` a second
472 // time.
473 DictionaryAttr discardableAttrs = op->getDiscardableAttrDictionary();
474 auto constantOp =
475 LLVM::ConstantOp::create(rewriter, op.getLoc(), resultType, value);
476 constantOp->setDiscardableAttrs(discardableAttrs);
477 rewriter.replaceOp(op, constantOp);
478 return success();
479}
480
481//===----------------------------------------------------------------------===//
482// IndexCastOpLowering
483//===----------------------------------------------------------------------===//
484
485template <typename OpTy, typename ExtCastTy>
486LogicalResult IndexCastOpLowering<OpTy, ExtCastTy>::matchAndRewrite(
487 OpTy op, typename OpTy::Adaptor adaptor,
488 ConversionPatternRewriter &rewriter) const {
489 Type resultType = op.getResult().getType();
490 Type targetElementType =
491 this->typeConverter->convertType(getElementTypeOrSelf(resultType));
492 Type sourceElementType =
493 this->typeConverter->convertType(getElementTypeOrSelf(op.getIn()));
494 unsigned targetBits = targetElementType.getIntOrFloatBitWidth();
495 unsigned sourceBits = sourceElementType.getIntOrFloatBitWidth();
496
497 if (targetBits == sourceBits) {
498 rewriter.replaceOp(op, adaptor.getIn());
499 return success();
500 }
501
502 // Memref index_cast is a no-op at the LLVM level since LLVM uses opaque
503 // pointers and memrefs of different integer/index element types all convert
504 // to the same LLVM struct type.
505 if (isa<MemRefType>(op.getIn().getType())) {
506 rewriter.replaceOp(op, adaptor.getIn());
507 return success();
508 }
509
510 bool isNonNeg = false;
511 if constexpr (std::is_same_v<ExtCastTy, LLVM::ZExtOp>)
512 isNonNeg = op.getNonNeg();
513
514 // Handle the scalar and 1D vector cases.
515 Type operandType = adaptor.getIn().getType();
516 if (!isa<LLVM::LLVMArrayType>(operandType)) {
517 Type targetType = this->typeConverter->convertType(resultType);
518 if (targetBits < sourceBits) {
519 rewriter.replaceOpWithNewOp<LLVM::TruncOp>(op, targetType,
520 adaptor.getIn());
521 } else {
522 auto extOp = rewriter.replaceOpWithNewOp<ExtCastTy>(op, targetType,
523 adaptor.getIn());
524 if constexpr (std::is_same_v<ExtCastTy, LLVM::ZExtOp>)
525 extOp.setNonNeg(isNonNeg);
526 }
527 return success();
528 }
529
530 if (!isa<VectorType>(resultType))
531 return rewriter.notifyMatchFailure(op, "expected vector result type");
532
534 op.getOperation(), adaptor.getOperands(), *(this->getTypeConverter()),
535 [&](Type llvm1DVectorTy, ValueRange operands) -> Value {
536 typename OpTy::Adaptor adaptor(operands);
537 if (targetBits < sourceBits) {
538 return LLVM::TruncOp::create(rewriter, op.getLoc(), llvm1DVectorTy,
539 adaptor.getIn());
540 }
541 auto extOp = ExtCastTy::create(rewriter, op.getLoc(), llvm1DVectorTy,
542 adaptor.getIn());
543 if constexpr (std::is_same_v<ExtCastTy, LLVM::ZExtOp>) {
544 if (isNonNeg)
545 extOp.setNonNeg(true);
546 }
547 return extOp;
548 },
549 rewriter);
550}
551
552//===----------------------------------------------------------------------===//
553// AddUIExtendedOpLowering
554//===----------------------------------------------------------------------===//
555
556LogicalResult AddUIExtendedOpLowering::matchAndRewrite(
557 arith::AddUIExtendedOp op, OpAdaptor adaptor,
558 ConversionPatternRewriter &rewriter) const {
559 Type operandType = adaptor.getLhs().getType();
560 Type sumResultType = op.getSum().getType();
561 Type overflowResultType = op.getOverflow().getType();
562
563 if (!LLVM::isCompatibleType(operandType))
564 return failure();
565
566 MLIRContext *ctx = rewriter.getContext();
567 Location loc = op.getLoc();
568
569 // Handle the scalar and 1D vector cases.
570 if (!isa<LLVM::LLVMArrayType>(operandType)) {
571 Type newOverflowType = typeConverter->convertType(overflowResultType);
572 Type structType =
573 LLVM::LLVMStructType::getLiteral(ctx, {sumResultType, newOverflowType});
574 Value addOverflow = LLVM::UAddWithOverflowOp::create(
575 rewriter, loc, structType, adaptor.getLhs(), adaptor.getRhs());
576 Value sumExtracted =
577 LLVM::ExtractValueOp::create(rewriter, loc, addOverflow, 0);
578 Value overflowExtracted =
579 LLVM::ExtractValueOp::create(rewriter, loc, addOverflow, 1);
580 rewriter.replaceOp(op, {sumExtracted, overflowExtracted});
581 return success();
582 }
583
584 if (!isa<VectorType>(sumResultType))
585 return rewriter.notifyMatchFailure(loc, "expected vector result types");
586
587 return rewriter.notifyMatchFailure(loc,
588 "ND vector types are not supported yet");
589}
590
591//===----------------------------------------------------------------------===//
592// SubUIExtendedOpLowering
593//===----------------------------------------------------------------------===//
594
595LogicalResult SubUIExtendedOpLowering::matchAndRewrite(
596 arith::SubUIExtendedOp op, OpAdaptor adaptor,
597 ConversionPatternRewriter &rewriter) const {
598 Type operandType = adaptor.getLhs().getType();
599 Type diffResultType = op.getDiff().getType();
600 Type borrowResultType = op.getBorrow().getType();
601
602 if (!LLVM::isCompatibleType(operandType))
603 return failure();
604
605 MLIRContext *ctx = rewriter.getContext();
606 Location loc = op.getLoc();
607
608 // Handle the scalar and 1D vector cases.
609 if (!isa<LLVM::LLVMArrayType>(operandType)) {
610 Type newBorrowType = typeConverter->convertType(borrowResultType);
611 Type structType =
612 LLVM::LLVMStructType::getLiteral(ctx, {diffResultType, newBorrowType});
613 Value subOverflow = LLVM::USubWithOverflowOp::create(
614 rewriter, loc, structType, adaptor.getLhs(), adaptor.getRhs());
615 Value diffExtracted =
616 LLVM::ExtractValueOp::create(rewriter, loc, subOverflow, 0);
617 Value borrowExtracted =
618 LLVM::ExtractValueOp::create(rewriter, loc, subOverflow, 1);
619 rewriter.replaceOp(op, {diffExtracted, borrowExtracted});
620 return success();
621 }
622
623 if (!isa<VectorType>(diffResultType))
624 return rewriter.notifyMatchFailure(loc, "expected vector result types");
625
626 return rewriter.notifyMatchFailure(loc,
627 "ND vector types are not supported yet");
628}
629
630//===----------------------------------------------------------------------===//
631// MulIExtendedOpLowering
632//===----------------------------------------------------------------------===//
633
634template <typename ArithMulOp, bool IsSigned>
635LogicalResult MulIExtendedOpLowering<ArithMulOp, IsSigned>::matchAndRewrite(
636 ArithMulOp op, typename ArithMulOp::Adaptor adaptor,
637 ConversionPatternRewriter &rewriter) const {
638 Type resultType = adaptor.getLhs().getType();
639
640 if (!LLVM::isCompatibleType(resultType))
641 return failure();
642
643 Location loc = op.getLoc();
644
645 // Handle the scalar and 1D vector cases. Because LLVM does not have a
646 // matching extended multiplication intrinsic, perform regular multiplication
647 // on operands zero-extended to i(2*N) bits, and truncate the results back to
648 // iN types.
649 if (!isa<LLVM::LLVMArrayType>(resultType)) {
650 // Shift amount necessary to extract the high bits from widened result.
651 TypedAttr shiftValAttr;
652
653 if (auto intTy = dyn_cast<IntegerType>(resultType)) {
654 unsigned resultBitwidth = intTy.getWidth();
655 auto attrTy = rewriter.getIntegerType(resultBitwidth * 2);
656 shiftValAttr = rewriter.getIntegerAttr(attrTy, resultBitwidth);
657 } else {
658 auto vecTy = cast<VectorType>(resultType);
659 unsigned resultBitwidth = vecTy.getElementTypeBitWidth();
660 auto attrTy = VectorType::get(
661 vecTy.getShape(), rewriter.getIntegerType(resultBitwidth * 2));
662 shiftValAttr = SplatElementsAttr::get(
663 attrTy, APInt(resultBitwidth * 2, resultBitwidth));
664 }
665 Type wideType = shiftValAttr.getType();
666 assert(LLVM::isCompatibleType(wideType) &&
667 "LLVM dialect should support all signless integer types");
668
669 using LLVMExtOp = std::conditional_t<IsSigned, LLVM::SExtOp, LLVM::ZExtOp>;
670 Value lhsExt = LLVMExtOp::create(rewriter, loc, wideType, adaptor.getLhs());
671 Value rhsExt = LLVMExtOp::create(rewriter, loc, wideType, adaptor.getRhs());
672 Value mulExt = LLVM::MulOp::create(rewriter, loc, wideType, lhsExt, rhsExt);
673
674 // Split the 2*N-bit wide result into two N-bit values.
675 Value low = LLVM::TruncOp::create(rewriter, loc, resultType, mulExt);
676 Value shiftVal = LLVM::ConstantOp::create(rewriter, loc, shiftValAttr);
677 Value highExt = LLVM::LShrOp::create(rewriter, loc, mulExt, shiftVal);
678 Value high = LLVM::TruncOp::create(rewriter, loc, resultType, highExt);
679
680 rewriter.replaceOp(op, {low, high});
681 return success();
682 }
683
684 if (!isa<VectorType>(resultType))
685 return rewriter.notifyMatchFailure(op, "expected vector result type");
686
687 return rewriter.notifyMatchFailure(op,
688 "ND vector types are not supported yet");
689}
690
691//===----------------------------------------------------------------------===//
692// CmpIOpLowering
693//===----------------------------------------------------------------------===//
694
695// Convert arith.cmp predicate into the LLVM dialect CmpPredicate. The two enums
696// share numerical values so just cast.
697template <typename LLVMPredType, typename PredType>
698static LLVMPredType convertCmpPredicate(PredType pred) {
699 return static_cast<LLVMPredType>(pred);
700}
701
702LogicalResult
703CmpIOpLowering::matchAndRewrite(arith::CmpIOp op, OpAdaptor adaptor,
704 ConversionPatternRewriter &rewriter) const {
705 Type operandType = adaptor.getLhs().getType();
706 Type resultType = op.getResult().getType();
707
708 // Handle the scalar and 1D vector cases.
709 if (!isa<LLVM::LLVMArrayType>(operandType)) {
710 rewriter.replaceOpWithNewOp<LLVM::ICmpOp>(
711 op, typeConverter->convertType(resultType),
713 adaptor.getLhs(), adaptor.getRhs());
714 return success();
715 }
716
717 if (!isa<VectorType>(resultType))
718 return rewriter.notifyMatchFailure(op, "expected vector result type");
719
721 op.getOperation(), adaptor.getOperands(), *getTypeConverter(),
722 [&](Type llvm1DVectorTy, ValueRange operands) {
723 OpAdaptor adaptor(operands);
724 return LLVM::ICmpOp::create(
725 rewriter, op.getLoc(), llvm1DVectorTy,
727 adaptor.getLhs(), adaptor.getRhs());
728 },
729 rewriter);
730}
731
732//===----------------------------------------------------------------------===//
733// CmpFOpLowering
734//===----------------------------------------------------------------------===//
735
736LogicalResult
737CmpFOpLowering::matchAndRewrite(arith::CmpFOp op, OpAdaptor adaptor,
738 ConversionPatternRewriter &rewriter) const {
739 if (LLVM::detail::isUnsupportedFloatingPointType(*this->getTypeConverter(),
740 op.getLhs().getType()))
741 return rewriter.notifyMatchFailure(op, "unsupported floating point type");
742
743 Type operandType = adaptor.getLhs().getType();
744 Type resultType = op.getResult().getType();
745 LLVM::FastmathFlags fmf =
746 arith::convertArithFastMathFlagsToLLVM(op.getFastmath());
747
748 // Handle the scalar and 1D vector cases.
749 if (!isa<LLVM::LLVMArrayType>(operandType)) {
750 rewriter.replaceOpWithNewOp<LLVM::FCmpOp>(
751 op, typeConverter->convertType(resultType),
753 adaptor.getLhs(), adaptor.getRhs(), fmf);
754 return success();
755 }
756
757 if (!isa<VectorType>(resultType))
758 return rewriter.notifyMatchFailure(op, "expected vector result type");
759
761 op.getOperation(), adaptor.getOperands(), *getTypeConverter(),
762 [&](Type llvm1DVectorTy, ValueRange operands) {
763 OpAdaptor adaptor(operands);
764 return LLVM::FCmpOp::create(
765 rewriter, op.getLoc(), llvm1DVectorTy,
767 adaptor.getLhs(), adaptor.getRhs(), fmf);
768 },
769 rewriter);
770}
771
772//===----------------------------------------------------------------------===//
773// SelectOpOneToNLowering
774//===----------------------------------------------------------------------===//
775
776/// Pattern for arith.select where the true/false values lower to multiple
777/// SSA values (1:N conversion). This pattern generates multiple arith.select
778/// than can be lowered by the 1:1 arith.select pattern.
779LogicalResult SelectOpOneToNLowering::matchAndRewrite(
780 arith::SelectOp op, Adaptor adaptor,
781 ConversionPatternRewriter &rewriter) const {
782 // In case of a 1:1 conversion, the 1:1 pattern will match.
783 if (llvm::hasSingleElement(adaptor.getTrueValue()))
784 return rewriter.notifyMatchFailure(
785 op, "not a 1:N conversion, 1:1 pattern will match");
786 if (!op.getCondition().getType().isInteger(1))
787 return rewriter.notifyMatchFailure(op,
788 "non-i1 conditions are not supported");
789 SmallVector<Value> results;
790 for (auto [trueValue, falseValue] :
791 llvm::zip_equal(adaptor.getTrueValue(), adaptor.getFalseValue()))
792 results.push_back(arith::SelectOp::create(
793 rewriter, op.getLoc(), op.getCondition(), trueValue, falseValue));
794 rewriter.replaceOpWithMultiple(op, {results});
795 return success();
796}
797
798//===----------------------------------------------------------------------===//
799// Pass Definition
800//===----------------------------------------------------------------------===//
801
802namespace {
803struct ArithToLLVMConversionPass
804 : public impl::ArithToLLVMConversionPassBase<ArithToLLVMConversionPass> {
805 using Base::Base;
806
807 void runOnOperation() override {
808 LLVMConversionTarget target(getContext());
809 RewritePatternSet patterns(&getContext());
810
811 const auto &dataLayoutAnalysis = getAnalysis<DataLayoutAnalysis>();
812 LowerToLLVMOptions options(&getContext(),
813 dataLayoutAnalysis.getAtOrAbove(getOperation()));
814 if (indexBitwidth != kDeriveIndexBitwidthFromDataLayout)
815 options.overrideIndexBitwidth(indexBitwidth);
816
817 LLVMTypeConverter converter(&getContext(), options, &dataLayoutAnalysis);
818 arith::populateCeilFloorDivExpandOpsPatterns(patterns);
819 arith::populateArithToLLVMConversionPatterns(converter, patterns);
820
821 if (failed(applyPartialConversion(getOperation(), target,
822 std::move(patterns))))
823 signalPassFailure();
824 }
825};
826} // namespace
827
828//===----------------------------------------------------------------------===//
829// ConvertToLLVMPatternInterface implementation
830//===----------------------------------------------------------------------===//
831
832namespace {
833/// Implement the interface to convert MemRef to LLVM.
834struct ArithToLLVMDialectInterface : public ConvertToLLVMPatternInterface {
835 ArithToLLVMDialectInterface(Dialect *dialect)
836 : ConvertToLLVMPatternInterface(dialect) {}
837
838 void loadDependentDialects(MLIRContext *context) const final {
839 context->loadDialect<LLVM::LLVMDialect>();
840 }
841
842 /// Hook for derived dialect interface to provide conversion patterns
843 /// and mark dialect legal for the conversion target.
844 void populateConvertToLLVMConversionPatterns(
845 ConversionTarget &target, LLVMTypeConverter &typeConverter,
846 RewritePatternSet &patterns) const final {
847 arith::populateCeilFloorDivExpandOpsPatterns(patterns);
848 arith::populateArithToLLVMConversionPatterns(typeConverter, patterns);
849 }
850};
851} // namespace
852
854 DialectRegistry &registry) {
855 registry.addExtension(+[](MLIRContext *ctx, arith::ArithDialect *dialect) {
856 dialect->addInterfaces<ArithToLLVMDialectInterface>();
857 });
858}
859
860//===----------------------------------------------------------------------===//
861// Pattern Population
862//===----------------------------------------------------------------------===//
863
865 const LLVMTypeConverter &converter, RewritePatternSet &patterns) {
866
867 // Set a higher pattern benefit for IdentityBitcastLowering so it will run
868 // before BitcastOpLowering.
869 patterns.add<IdentityBitcastLowering>(converter, patterns.getContext(),
870 /*patternBenefit*/ 10);
871
872 // clang-format off
873 patterns.add<
874 AddFOpLowering,
875 ConstrainedAddFOpLowering,
876 AddIOpLowering,
877 AndIOpLowering,
878 AddUIExtendedOpLowering,
879 SubUIExtendedOpLowering,
880 BitcastOpLowering,
881 ConstantOpLowering,
882 CmpFOpLowering,
883 CmpIOpLowering,
884 DivFOpLowering,
885 ConstrainedDivFOpLowering,
886 DivSIOpLowering,
887 DivUIOpLowering,
888 ExtFOpLowering,
889 ExtSIOpLowering,
890 ExtUIOpLowering,
891 ConvertFOpLowering,
892 FPToSIOpLowering,
893 FPToUIOpLowering,
894 IndexCastOpSILowering,
895 IndexCastOpUILowering,
896 MaximumFOpLowering,
897 MaxNumFOpLowering,
898 MaximumNumFOpLowering,
899 MaxSIOpLowering,
900 MaxUIOpLowering,
901 MinimumFOpLowering,
902 MinNumFOpLowering,
903 MinimumNumFOpLowering,
904 MinSIOpLowering,
905 MinUIOpLowering,
906 MulFOpLowering,
907 ConstrainedMulFOpLowering,
908 MulIOpLowering,
909 MulSIExtendedOpLowering,
910 MulUIExtendedOpLowering,
911 NegFOpLowering,
912 OrIOpLowering,
913 RemFOpLowering,
914 RemSIOpLowering,
915 RemUIOpLowering,
916 SelectOpLowering,
917 SelectOpOneToNLowering,
918 ShLIOpLowering,
919 ShRSIOpLowering,
920 ShRUIOpLowering,
921 SIToFPOpLowering,
922 SubFOpLowering,
923 ConstrainedSubFOpLowering,
924 SubIOpLowering,
925 TruncFOpLowering,
926 ConstrainedTruncFOpLowering,
927 TruncIOpLowering,
928 UIToFPOpLowering,
929 XOrIOpLowering
930 >(converter);
931 // clang-format on
932}
return success()
static LLVMPredType convertCmpPredicate(PredType pred)
static TypedAttr convertConstantValue(TypedAttr attr, Type resultType)
Retypes attr for a llvm.mlir.constant of resultType.
b getContext())
static llvm::ManagedStatic< PassManagerOptions > options
Utility class for operation conversions targeting the LLVM dialect that match exactly one source oper...
Definition Pattern.h:233
ConvertOpToLLVMPattern(const LLVMTypeConverter &typeConverter, PatternBenefit benefit=1)
Definition Pattern.h:239
typename SourceOp::template GenericAdaptor< ArrayRef< ValueRange > > OneToNOpAdaptor
Definition Pattern.h:236
An attribute that represents a reference to a dense integer vector or tensor object.
The DialectRegistry maps a dialect namespace to a constructor for the matching dialect.
bool addExtension(TypeID extensionID, std::unique_ptr< DialectExtensionBase > extension)
Add the given extension to the registry.
Conversion from types to the LLVM IR dialect.
MLIRContext is the top-level object for a collection of MLIR operations.
Definition MLIRContext.h:63
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
unsigned getIntOrFloatBitWidth() const
Return the bit width of an integer or a float type, assert failure on other types.
Definition Types.cpp:124
Type getType() const
Return the type of this value.
Definition Value.h:105
Basic lowering implementation to rewrite Ops with just one result to the LLVM Dialect.
LogicalResult handleMultidimensionalVectors(Operation *op, ValueRange operands, const LLVMTypeConverter &typeConverter, std::function< Value(Type, ValueRange)> createOperand, ConversionPatternRewriter &rewriter)
bool isUnsupportedFloatingPointType(const TypeConverter &typeConverter, Type type)
Return "true" if the given type is an unsupported floating point type.
Definition Pattern.cpp:678
bool opHasUnsupportedFloatingPointTypes(Operation *op, const TypeConverter &typeConverter)
Return "true" if the given op has any unsupported floating point types (either operands or results).
Definition Pattern.cpp:689
Type getConstantElementType(Type type)
Determines the element type of type the way the llvm.mlir.constant verifier does, i....
void populateArithToLLVMConversionPatterns(const LLVMTypeConverter &converter, RewritePatternSet &patterns)
void registerConvertArithToLLVMInterface(DialectRegistry &registry)
detail::InFlightRemark failed(Location loc, RemarkOpts opts)
Report an optimization remark that failed.
Definition Remarks.h:734
Include the generated interface declarations.
static constexpr unsigned kDeriveIndexBitwidthFromDataLayout
Value to pass as bitwidth for the index type when the converter is expected to derive the bitwidth fr...
Type getElementTypeOrSelf(Type type)
Return the element type or return the type itself.