23#include "llvm/ADT/TypeSwitch.h"
24#include "llvm/Support/FormatVariadic.h"
26#define DEBUG_TYPE "spirv-to-llvm-pattern"
38 if (
auto vecType = dyn_cast<VectorType>(type))
39 return vecType.getElementType().isSignedInteger();
47 if (
auto vecType = dyn_cast<VectorType>(type))
48 return vecType.getElementType().isUnsignedInteger();
55 if (
auto intType = dyn_cast<IntegerType>(type))
56 return intType.getWidth();
57 if (
auto vecType = dyn_cast<VectorType>(type))
58 if (
auto intType = dyn_cast<IntegerType>(vecType.getElementType()))
59 return intType.getWidth();
66 "bitwidth is not supported for this type");
69 auto vecType = dyn_cast<VectorType>(type);
70 auto elementType = vecType.getElementType();
71 assert(elementType.isIntOrFloat() &&
72 "only integers and floats have a bitwidth");
73 return elementType.getIntOrFloatBitWidth();
78 if (
auto vecTy = dyn_cast<VectorType>(type))
79 type = vecTy.getElementType();
80 return cast<IntegerType>(type).getWidth();
87 IntegerAttr scalarAttr) {
88 if (
auto vecType = dyn_cast<VectorType>(srcType))
89 return LLVM::ConstantOp::create(
91 return LLVM::ConstantOp::create(rewriter, loc, dstType, scalarAttr);
97 auto integerType = cast<IntegerType>(
98 isa<VectorType>(srcType) ? cast<VectorType>(srcType).
getElementType()
107 if (
auto vecType = dyn_cast<VectorType>(srcType)) {
108 auto floatType = cast<FloatType>(vecType.getElementType());
109 return LLVM::ConstantOp::create(
110 rewriter, loc, dstType,
114 auto floatType = cast<FloatType>(srcType);
115 return LLVM::ConstantOp::create(rewriter, loc, dstType,
128 auto srcType = value.
getType();
134 if (valueBitWidth < targetBitWidth)
135 return LLVM::ZExtOp::create(rewriter, loc, llvmType, value);
140 if (valueBitWidth > targetBitWidth)
141 return LLVM::TruncOp::create(rewriter, loc, llvmType, value);
148 ConversionPatternRewriter &rewriter) {
149 auto vectorType = VectorType::get(numElements, toBroadcast.
getType());
150 auto llvmVectorType = typeConverter.convertType(vectorType);
151 auto llvmI32Type = typeConverter.convertType(rewriter.getIntegerType(32));
152 Value broadcasted = LLVM::PoisonOp::create(rewriter, loc, llvmVectorType);
153 for (
unsigned i = 0; i < numElements; ++i) {
154 auto index = LLVM::ConstantOp::create(rewriter, loc, llvmI32Type,
155 rewriter.getI32IntegerAttr(i));
156 broadcasted = LLVM::InsertElementOp::create(
157 rewriter, loc, llvmVectorType, broadcasted, toBroadcast,
index);
165 ConversionPatternRewriter &rewriter) {
166 if (
auto vectorType = dyn_cast<VectorType>(srcType)) {
167 unsigned numElements = vectorType.getNumElements();
168 return broadcast(loc, value, numElements, typeConverter, rewriter);
185 ConversionPatternRewriter &rewriter) {
199 if (failed(converter.convertTypes(type.
getElementTypes(), elementsVector)))
201 return LLVM::LLVMStructType::getLiteral(type.getContext(), elementsVector,
209 if (failed(converter.convertTypes(type.
getElementTypes(), elementsVector)))
211 return LLVM::LLVMStructType::getLiteral(type.getContext(), elementsVector,
218 return LLVM::ConstantOp::create(
219 rewriter, loc, IntegerType::get(rewriter.
getContext(), 32),
225 ConversionPatternRewriter &rewriter,
227 unsigned alignment,
bool isVolatile,
228 bool isNonTemporal) {
229 if (
auto loadOp = dyn_cast<spirv::LoadOp>(op)) {
230 auto dstType = typeConverter.convertType(loadOp.getType());
232 return rewriter.notifyMatchFailure(op,
"type conversion failed");
233 rewriter.replaceOpWithNewOp<LLVM::LoadOp>(
234 loadOp, dstType, spirv::LoadOpAdaptor(operands).getPtr(), alignment,
235 isVolatile, isNonTemporal);
238 auto storeOp = cast<spirv::StoreOp>(op);
239 spirv::StoreOpAdaptor adaptor(operands);
240 rewriter.replaceOpWithNewOp<LLVM::StoreOp>(storeOp, adaptor.getValue(),
241 adaptor.getPtr(), alignment,
242 isVolatile, isNonTemporal);
257 auto sizeInBytes = cast<spirv::SPIRVType>(elementType).getSizeInBytes();
258 if (stride != 0 && (!sizeInBytes || *sizeInBytes != stride))
261 auto llvmElementType = converter.convertType(elementType);
263 return LLVM::LLVMArrayType::get(llvmElementType, numElements);
270 spirv::ClientAPI clientAPI) {
271 unsigned addressSpace =
273 return LLVM::LLVMPointerType::get(type.getContext(), addressSpace);
284 return LLVM::LLVMArrayType::get(elementType, 0);
293 if (!memberDecorations.empty())
306template <
typename TargetOp,
typename SourceOp>
308collectAttrsForConversion(SourceOp op,
309 typename TargetOp::Properties &properties,
312 if (
auto sourceProperties =
313 dyn_cast_or_null<DictionaryAttr>(op->getPropertiesAsAttribute()))
314 attrs.append(sourceProperties.getValue());
316 if constexpr (std::is_same_v<
typename TargetOp::Properties,
318 llvm::append_range(discardableAttrs, attrs);
322 TargetOp::populateDefaultProperties(
323 OperationName(TargetOp::getOperationName(), op->getContext()),
325 if (
failed(TargetOp::setPropertiesFromAttr(
326 properties, attrs.getDictionary(op->getContext()),
327 [&]() { return op.emitError(
"failed to convert properties"); })))
330 auto convertedProperties = dyn_cast_or_null<DictionaryAttr>(
331 TargetOp::getPropertiesAsAttr(op->getContext(), properties));
333 StringRef name = attr.getName().getValue();
334 bool isProperty = convertedProperties && convertedProperties.contains(name);
335 if (name ==
"operand_segment_sizes")
336 isProperty |= convertedProperties &&
337 convertedProperties.contains(
"operandSegmentSizes");
338 if (name ==
"result_segment_sizes")
339 isProperty |= convertedProperties &&
340 convertedProperties.contains(
"resultSegmentSizes");
342 discardableAttrs.push_back(attr);
349 using SPIRVToLLVMConversion<spirv::AccessChainOp>::SPIRVToLLVMConversion;
352 matchAndRewrite(spirv::AccessChainOp op, OpAdaptor adaptor,
353 ConversionPatternRewriter &rewriter)
const override {
355 getTypeConverter()->convertType(op.getComponentPtr().getType());
357 return rewriter.notifyMatchFailure(op,
"type conversion failed");
359 auto indices = llvm::to_vector<4>(adaptor.getIndices());
360 Type indexType = op.getIndices().front().getType();
361 auto llvmIndexType = getTypeConverter()->convertType(indexType);
363 return rewriter.notifyMatchFailure(op,
"type conversion failed");
365 LLVM::ConstantOp::create(rewriter, op.getLoc(), llvmIndexType,
366 rewriter.getIntegerAttr(llvmIndexType, 0));
369 auto elementType = getTypeConverter()->convertType(
370 cast<spirv::PointerType>(op.getBasePtr().getType()).getPointeeType());
372 return rewriter.notifyMatchFailure(op,
"type conversion failed");
373 rewriter.replaceOpWithNewOp<LLVM::GEPOp>(op, dstType, elementType,
374 adaptor.getBasePtr(),
indices);
381 using SPIRVToLLVMConversion<spirv::AddressOfOp>::SPIRVToLLVMConversion;
384 matchAndRewrite(spirv::AddressOfOp op, OpAdaptor adaptor,
385 ConversionPatternRewriter &rewriter)
const override {
386 auto dstType = getTypeConverter()->convertType(op.getPointer().getType());
388 return rewriter.notifyMatchFailure(op,
"type conversion failed");
389 rewriter.replaceOpWithNewOp<LLVM::AddressOfOp>(op, dstType,
395class BitFieldInsertPattern
398 using SPIRVToLLVMConversion<spirv::BitFieldInsertOp>::SPIRVToLLVMConversion;
401 matchAndRewrite(spirv::BitFieldInsertOp op, OpAdaptor adaptor,
402 ConversionPatternRewriter &rewriter)
const override {
403 auto srcType = op.getType();
404 auto dstType = getTypeConverter()->convertType(srcType);
406 return rewriter.notifyMatchFailure(op,
"type conversion failed");
407 Location loc = op.getLoc();
411 *getTypeConverter(), rewriter);
413 *getTypeConverter(), rewriter);
417 Value maskShiftedByCount =
418 LLVM::ShlOp::create(rewriter, loc, dstType, minusOne, count);
419 Value negated = LLVM::XOrOp::create(rewriter, loc, dstType,
420 maskShiftedByCount, minusOne);
421 Value maskShiftedByCountAndOffset =
422 LLVM::ShlOp::create(rewriter, loc, dstType, negated, offset);
423 Value mask = LLVM::XOrOp::create(rewriter, loc, dstType,
424 maskShiftedByCountAndOffset, minusOne);
429 LLVM::AndOp::create(rewriter, loc, dstType, op.getBase(), mask);
430 Value insertShiftedByOffset =
431 LLVM::ShlOp::create(rewriter, loc, dstType, op.getInsert(), offset);
432 rewriter.replaceOpWithNewOp<LLVM::OrOp>(op, dstType, baseAndMask,
433 insertShiftedByOffset);
439class ConstantScalarAndVectorPattern
442 using SPIRVToLLVMConversion<spirv::ConstantOp>::SPIRVToLLVMConversion;
445 matchAndRewrite(spirv::ConstantOp constOp, OpAdaptor adaptor,
446 ConversionPatternRewriter &rewriter)
const override {
447 auto srcType = constOp.getType();
448 if (!isa<VectorType>(srcType) && !srcType.isIntOrFloat())
451 auto dstType = getTypeConverter()->convertType(srcType);
453 return rewriter.notifyMatchFailure(constOp,
"type conversion failed");
462 auto signlessType = rewriter.getIntegerType(
getBitWidth(srcType));
464 if (isa<VectorType>(srcType)) {
465 auto dstElementsAttr = cast<DenseIntElementsAttr>(constOp.getValue());
466 rewriter.replaceOpWithNewOp<LLVM::ConstantOp>(
468 dstElementsAttr.mapValues(
469 signlessType, [&](
const APInt &value) {
return value; }));
472 auto srcAttr = cast<IntegerAttr>(constOp.getValue());
473 auto dstAttr = rewriter.getIntegerAttr(signlessType, srcAttr.getValue());
474 rewriter.replaceOpWithNewOp<LLVM::ConstantOp>(constOp, dstType, dstAttr);
477 LLVM::ConstantOp::Properties properties{};
478 SmallVector<NamedAttribute> discardableAttrs;
479 if (
failed(collectAttrsForConversion<LLVM::ConstantOp>(constOp, properties,
482 rewriter.replaceOpWithNewOp<LLVM::ConstantOp>(
483 constOp, dstType, adaptor.getOperands(), properties, discardableAttrs);
488class BitFieldSExtractPattern
491 using SPIRVToLLVMConversion<spirv::BitFieldSExtractOp>::SPIRVToLLVMConversion;
494 matchAndRewrite(spirv::BitFieldSExtractOp op, OpAdaptor adaptor,
495 ConversionPatternRewriter &rewriter)
const override {
496 auto srcType = op.getType();
497 auto dstType = getTypeConverter()->convertType(srcType);
499 return rewriter.notifyMatchFailure(op,
"type conversion failed");
500 Location loc = op.getLoc();
504 *getTypeConverter(), rewriter);
506 *getTypeConverter(), rewriter);
509 IntegerType integerType;
510 if (
auto vecType = dyn_cast<VectorType>(srcType))
511 integerType = cast<IntegerType>(vecType.getElementType());
513 integerType = cast<IntegerType>(srcType);
515 auto baseSize = rewriter.getIntegerAttr(integerType,
getBitWidth(srcType));
517 isa<VectorType>(srcType)
518 ? LLVM::ConstantOp::create(
519 rewriter, loc, dstType,
520 SplatElementsAttr::get(cast<ShapedType>(srcType), baseSize))
521 : LLVM::ConstantOp::create(rewriter, loc, dstType, baseSize);
525 Value countPlusOffset =
526 LLVM::AddOp::create(rewriter, loc, dstType, count, offset);
527 Value amountToShiftLeft =
528 LLVM::SubOp::create(rewriter, loc, dstType, size, countPlusOffset);
529 Value baseShiftedLeft = LLVM::ShlOp::create(
530 rewriter, loc, dstType, op.getBase(), amountToShiftLeft);
533 Value amountToShiftRight =
534 LLVM::AddOp::create(rewriter, loc, dstType, offset, amountToShiftLeft);
535 rewriter.replaceOpWithNewOp<LLVM::AShrOp>(op, dstType, baseShiftedLeft,
541class BitFieldUExtractPattern
544 using SPIRVToLLVMConversion<spirv::BitFieldUExtractOp>::SPIRVToLLVMConversion;
547 matchAndRewrite(spirv::BitFieldUExtractOp op, OpAdaptor adaptor,
548 ConversionPatternRewriter &rewriter)
const override {
549 auto srcType = op.getType();
550 auto dstType = getTypeConverter()->convertType(srcType);
552 return rewriter.notifyMatchFailure(op,
"type conversion failed");
553 Location loc = op.getLoc();
557 *getTypeConverter(), rewriter);
559 *getTypeConverter(), rewriter);
563 Value maskShiftedByCount =
564 LLVM::ShlOp::create(rewriter, loc, dstType, minusOne, count);
565 Value mask = LLVM::XOrOp::create(rewriter, loc, dstType, maskShiftedByCount,
570 LLVM::LShrOp::create(rewriter, loc, dstType, op.getBase(), offset);
571 rewriter.replaceOpWithNewOp<LLVM::AndOp>(op, dstType, shiftedBase, mask);
578 using SPIRVToLLVMConversion<spirv::BranchOp>::SPIRVToLLVMConversion;
581 matchAndRewrite(spirv::BranchOp branchOp, OpAdaptor adaptor,
582 ConversionPatternRewriter &rewriter)
const override {
583 rewriter.replaceOpWithNewOp<LLVM::BrOp>(branchOp, adaptor.getOperands(),
584 branchOp.getTarget());
589class BranchConditionalConversionPattern
592 using SPIRVToLLVMConversion<
593 spirv::BranchConditionalOp>::SPIRVToLLVMConversion;
596 matchAndRewrite(spirv::BranchConditionalOp op, OpAdaptor adaptor,
597 ConversionPatternRewriter &rewriter)
const override {
600 if (
auto weights = op.getBranchWeights()) {
601 SmallVector<int32_t> weightValues;
602 for (
auto weight : weights->getAsRange<IntegerAttr>())
603 weightValues.push_back(weight.getInt());
607 rewriter.replaceOpWithNewOp<LLVM::CondBrOp>(
608 op, op.getCondition(), op.getTrueBlockArguments(),
609 op.getFalseBlockArguments(), branchWeights, op.getTrueBlock(),
618class CompositeExtractPattern
621 using SPIRVToLLVMConversion<spirv::CompositeExtractOp>::SPIRVToLLVMConversion;
624 matchAndRewrite(spirv::CompositeExtractOp op, OpAdaptor adaptor,
625 ConversionPatternRewriter &rewriter)
const override {
626 auto dstType = this->getTypeConverter()->convertType(op.getType());
628 return rewriter.notifyMatchFailure(op,
"type conversion failed");
630 Type containerType = op.getComposite().getType();
631 if (isa<VectorType>(containerType)) {
632 Location loc = op.getLoc();
633 IntegerAttr value = cast<IntegerAttr>(op.getIndices()[0]);
635 rewriter.replaceOpWithNewOp<LLVM::ExtractElementOp>(
636 op, dstType, adaptor.getComposite(), index);
640 rewriter.replaceOpWithNewOp<LLVM::ExtractValueOp>(
641 op, adaptor.getComposite(),
642 LLVM::convertArrayToIndices(op.getIndices()));
650class CompositeInsertPattern
653 using SPIRVToLLVMConversion<spirv::CompositeInsertOp>::SPIRVToLLVMConversion;
656 matchAndRewrite(spirv::CompositeInsertOp op, OpAdaptor adaptor,
657 ConversionPatternRewriter &rewriter)
const override {
658 auto dstType = this->getTypeConverter()->convertType(op.getType());
660 return rewriter.notifyMatchFailure(op,
"type conversion failed");
662 Type containerType = op.getComposite().getType();
663 if (isa<VectorType>(containerType)) {
664 Location loc = op.getLoc();
665 IntegerAttr value = cast<IntegerAttr>(op.getIndices()[0]);
667 rewriter.replaceOpWithNewOp<LLVM::InsertElementOp>(
668 op, dstType, adaptor.getComposite(), adaptor.getObject(), index);
672 rewriter.replaceOpWithNewOp<LLVM::InsertValueOp>(
673 op, adaptor.getComposite(), adaptor.getObject(),
674 LLVM::convertArrayToIndices(op.getIndices()));
681template <
typename SPIRVOp,
typename LLVMOp>
684 using SPIRVToLLVMConversion<SPIRVOp>::SPIRVToLLVMConversion;
687 matchAndRewrite(SPIRVOp op,
typename SPIRVOp::Adaptor adaptor,
688 ConversionPatternRewriter &rewriter)
const override {
689 auto dstType = this->getTypeConverter()->convertType(op.getType());
691 return rewriter.notifyMatchFailure(op,
"type conversion failed");
692 typename LLVMOp::Properties properties{};
693 SmallVector<NamedAttribute> discardableAttrs;
694 if (
failed(collectAttrsForConversion<LLVMOp>(op, properties,
697 rewriter.template replaceOpWithNewOp<LLVMOp>(
698 op, dstType, adaptor.getOperands(), properties, discardableAttrs);
708template <
typename SPIRVOp,
typename LLVMOp>
711 using SPIRVToLLVMConversion<SPIRVOp>::SPIRVToLLVMConversion;
714 matchAndRewrite(SPIRVOp op,
typename SPIRVOp::Adaptor adaptor,
715 ConversionPatternRewriter &rewriter)
const override {
716 Type dstType = this->getTypeConverter()->convertType(op.getType());
718 return rewriter.notifyMatchFailure(op,
"type conversion failed");
720 Location loc = op.getLoc();
721 Type operandType = adaptor.getOperand1().getType();
722 Type overflowType = rewriter.getI1Type();
723 if (
auto vecType = dyn_cast<VectorType>(operandType))
724 overflowType = VectorType::get(vecType.getShape(), overflowType);
726 Type intrType = LLVM::LLVMStructType::getLiteral(
727 rewriter.getContext(), {operandType, overflowType});
728 Value intrResult = LLVMOp::create(
729 rewriter, loc, intrType, adaptor.getOperand1(), adaptor.getOperand2());
730 Value lowBits = LLVM::ExtractValueOp::create(rewriter, loc, intrResult, 0);
731 Value overflow = LLVM::ExtractValueOp::create(rewriter, loc, intrResult, 1);
732 overflow = LLVM::ZExtOp::create(rewriter, loc, operandType, overflow);
734 Value
result = LLVM::PoisonOp::create(rewriter, loc, dstType);
735 result = LLVM::InsertValueOp::create(rewriter, loc,
result, lowBits,
736 ArrayRef<int64_t>{0});
737 result = LLVM::InsertValueOp::create(rewriter, loc,
result, overflow,
738 ArrayRef<int64_t>{1});
739 rewriter.replaceOp(op,
result);
746class ExecutionModePattern
749 using SPIRVToLLVMConversion<spirv::ExecutionModeOp>::SPIRVToLLVMConversion;
752 matchAndRewrite(spirv::ExecutionModeOp op, OpAdaptor adaptor,
753 ConversionPatternRewriter &rewriter)
const override {
757 ModuleOp module = op->getParentOfType<ModuleOp>();
758 spirv::ExecutionModeAttr executionModeAttr = op.getExecutionModeAttr();
759 std::string moduleName;
760 if (module.getName().has_value())
761 moduleName =
"_" +
module.getName()->str();
764 std::string executionModeInfoName = llvm::formatv(
765 "__spv_{0}_{1}_execution_mode_info_{2}", moduleName, op.getFn().str(),
766 static_cast<uint32_t
>(executionModeAttr.getValue()));
768 MLIRContext *context = rewriter.getContext();
769 OpBuilder::InsertionGuard guard(rewriter);
770 rewriter.setInsertionPointToStart(module.getBody());
777 auto llvmI32Type = IntegerType::get(context, 32);
778 SmallVector<Type, 2> fields;
779 fields.push_back(llvmI32Type);
781 if (!values.empty()) {
782 auto arrayType = LLVM::LLVMArrayType::get(llvmI32Type, values.size());
783 fields.push_back(arrayType);
785 auto structType = LLVM::LLVMStructType::getLiteral(context, fields);
788 auto global = LLVM::GlobalOp::create(
789 rewriter, UnknownLoc::get(context), structType,
true,
790 LLVM::Linkage::External, executionModeInfoName, Attribute(),
792 Location loc = global.getLoc();
793 Region ®ion = global.getInitializerRegion();
794 Block *block = rewriter.createBlock(®ion);
797 rewriter.setInsertionPointToStart(block);
798 Value structValue = LLVM::PoisonOp::create(rewriter, loc, structType);
799 Value executionMode = LLVM::ConstantOp::create(
800 rewriter, loc, llvmI32Type,
801 rewriter.getI32IntegerAttr(
802 static_cast<uint32_t
>(executionModeAttr.getValue())));
803 SmallVector<int64_t> position{0};
804 structValue = LLVM::InsertValueOp::create(rewriter, loc, structValue,
805 executionMode, position);
808 for (
unsigned i = 0, e = values.size(); i < e; ++i) {
809 auto attr = values.getValue()[i];
810 Value entry = LLVM::ConstantOp::create(rewriter, loc, llvmI32Type, attr);
811 structValue = LLVM::InsertValueOp::create(
812 rewriter, loc, structValue, entry, ArrayRef<int64_t>({1, i}));
814 LLVM::ReturnOp::create(rewriter, loc, ArrayRef<Value>({structValue}));
815 rewriter.eraseOp(op);
824class GlobalVariablePattern
827 template <
typename... Args>
828 GlobalVariablePattern(spirv::ClientAPI clientAPI, Args &&...args)
829 : SPIRVToLLVMConversion<spirv::GlobalVariableOp>(
830 std::forward<Args>(args)...),
831 clientAPI(clientAPI) {}
834 matchAndRewrite(spirv::GlobalVariableOp op, OpAdaptor adaptor,
835 ConversionPatternRewriter &rewriter)
const override {
838 if (op.getInitializer())
841 auto srcType = cast<spirv::PointerType>(op.getType());
842 auto dstType = getTypeConverter()->convertType(srcType.getPointeeType());
844 return rewriter.notifyMatchFailure(op,
"type conversion failed");
849 auto storageClass = srcType.getStorageClass();
850 switch (storageClass) {
851 case spirv::StorageClass::Input:
852 case spirv::StorageClass::Private:
853 case spirv::StorageClass::Output:
854 case spirv::StorageClass::StorageBuffer:
855 case spirv::StorageClass::UniformConstant:
864 bool isConstant = (storageClass == spirv::StorageClass::Input) ||
865 (storageClass == spirv::StorageClass::UniformConstant);
871 auto linkage = storageClass == spirv::StorageClass::Private
872 ? LLVM::Linkage::Private
873 : LLVM::Linkage::External;
874 StringAttr locationAttrName = op.getLocationAttrName();
875 IntegerAttr locationAttr = op.getLocationAttr();
876 auto newGlobalOp = rewriter.replaceOpWithNewOp<LLVM::GlobalOp>(
877 op, dstType, isConstant, linkage, op.getSymName(), Attribute(),
882 newGlobalOp->setDiscardableAttr(locationAttrName, locationAttr);
888 spirv::ClientAPI clientAPI;
893template <
typename SPIRVOp,
typename LLVMExtOp,
typename LLVMTruncOp>
896 using SPIRVToLLVMConversion<SPIRVOp>::SPIRVToLLVMConversion;
899 matchAndRewrite(SPIRVOp op,
typename SPIRVOp::Adaptor adaptor,
900 ConversionPatternRewriter &rewriter)
const override {
902 Type fromType = op.getOperand().getType();
903 Type toType = op.getType();
905 auto dstType = this->getTypeConverter()->convertType(toType);
907 return rewriter.notifyMatchFailure(op,
"type conversion failed");
910 rewriter.template replaceOpWithNewOp<LLVMExtOp>(op, dstType,
911 adaptor.getOperands());
915 rewriter.template replaceOpWithNewOp<LLVMTruncOp>(op, dstType,
916 adaptor.getOperands());
923class FunctionCallPattern
926 using SPIRVToLLVMConversion<spirv::FunctionCallOp>::SPIRVToLLVMConversion;
929 matchAndRewrite(spirv::FunctionCallOp callOp, OpAdaptor adaptor,
930 ConversionPatternRewriter &rewriter)
const override {
931 LLVM::CallOp::Properties properties{};
932 SmallVector<NamedAttribute> discardableAttrs;
933 if (
failed(collectAttrsForConversion<LLVM::CallOp>(callOp, properties,
936 properties.operandSegmentSizes = {
937 static_cast<int32_t
>(adaptor.getOperands().size()), 0};
938 properties.op_bundle_sizes = rewriter.getDenseI32ArrayAttr({});
940 if (callOp.getNumResults() == 0) {
941 rewriter.replaceOpWithNewOp<LLVM::CallOp>(callOp,
TypeRange(),
942 adaptor.getOperands(),
943 properties, discardableAttrs);
948 auto dstType = getTypeConverter()->convertType(callOp.getType(0));
950 return rewriter.notifyMatchFailure(callOp,
"type conversion failed");
951 rewriter.replaceOpWithNewOp<LLVM::CallOp>(
952 callOp, dstType, adaptor.getOperands(), properties, discardableAttrs);
958template <
typename SPIRVOp, LLVM::FCmpPredicate predicate>
961 using SPIRVToLLVMConversion<SPIRVOp>::SPIRVToLLVMConversion;
964 matchAndRewrite(SPIRVOp op,
typename SPIRVOp::Adaptor adaptor,
965 ConversionPatternRewriter &rewriter)
const override {
967 auto dstType = this->getTypeConverter()->convertType(op.getType());
969 return rewriter.notifyMatchFailure(op,
"type conversion failed");
971 rewriter.template replaceOpWithNewOp<LLVM::FCmpOp>(
972 op, dstType, predicate, op.getOperand1(), op.getOperand2());
978template <
typename SPIRVOp, LLVM::ICmpPredicate predicate>
981 using SPIRVToLLVMConversion<SPIRVOp>::SPIRVToLLVMConversion;
984 matchAndRewrite(SPIRVOp op,
typename SPIRVOp::Adaptor adaptor,
985 ConversionPatternRewriter &rewriter)
const override {
987 auto dstType = this->getTypeConverter()->convertType(op.getType());
989 return rewriter.notifyMatchFailure(op,
"type conversion failed");
991 rewriter.template replaceOpWithNewOp<LLVM::ICmpOp>(
992 op, dstType, predicate, op.getOperand1(), op.getOperand2());
997class InverseSqrtPattern
1000 using SPIRVToLLVMConversion<spirv::GLInverseSqrtOp>::SPIRVToLLVMConversion;
1003 matchAndRewrite(spirv::GLInverseSqrtOp op, OpAdaptor adaptor,
1004 ConversionPatternRewriter &rewriter)
const override {
1005 auto srcType = op.getType();
1006 auto dstType = getTypeConverter()->convertType(srcType);
1008 return rewriter.notifyMatchFailure(op,
"type conversion failed");
1010 Location loc = op.getLoc();
1012 Value sqrt = LLVM::SqrtOp::create(rewriter, loc, dstType, op.getOperand());
1013 rewriter.replaceOpWithNewOp<LLVM::FDivOp>(op, dstType, one, sqrt);
1020class VectorTimesScalarPattern
1023 using SPIRVToLLVMConversion<
1024 spirv::VectorTimesScalarOp>::SPIRVToLLVMConversion;
1027 matchAndRewrite(spirv::VectorTimesScalarOp op, OpAdaptor adaptor,
1028 ConversionPatternRewriter &rewriter)
const override {
1029 Type srcType = op.getType();
1030 Type dstType = getTypeConverter()->convertType(srcType);
1032 return rewriter.notifyMatchFailure(op,
"type conversion failed");
1034 unsigned numElements = op.getVector().getType().getNumElements();
1035 Value broadcasted =
broadcast(op.getLoc(), adaptor.getScalar(), numElements,
1036 *getTypeConverter(), rewriter);
1037 rewriter.replaceOpWithNewOp<LLVM::FMulOp>(op, dstType, adaptor.getVector(),
1046 using SPIRVToLLVMConversion<spirv::SNegateOp>::SPIRVToLLVMConversion;
1049 matchAndRewrite(spirv::SNegateOp op, OpAdaptor adaptor,
1050 ConversionPatternRewriter &rewriter)
const override {
1051 Type srcType = op.getType();
1052 Type dstType = getTypeConverter()->convertType(srcType);
1054 return rewriter.notifyMatchFailure(op,
"type conversion failed");
1056 Location loc = op.getLoc();
1057 IntegerAttr zeroAttr = rewriter.getIntegerAttr(
1061 rewriter.replaceOpWithNewOp<LLVM::SubOp>(op, dstType, zero,
1062 adaptor.getOperand());
1069template <
typename SPIRVOp,
typename LLVMMinOp,
typename LLVMMaxOp>
1072 using SPIRVToLLVMConversion<SPIRVOp>::SPIRVToLLVMConversion;
1075 matchAndRewrite(SPIRVOp op,
typename SPIRVOp::Adaptor adaptor,
1076 ConversionPatternRewriter &rewriter)
const override {
1077 Type dstType = this->getTypeConverter()->convertType(op.getType());
1079 return rewriter.notifyMatchFailure(op,
"type conversion failed");
1081 Location loc = op.getLoc();
1082 Value
max = LLVMMaxOp::create(rewriter, loc, dstType, adaptor.getX(),
1084 rewriter.template replaceOpWithNewOp<LLVMMinOp>(op, dstType,
max,
1095 using SPIRVToLLVMConversion<spirv::FModOp>::SPIRVToLLVMConversion;
1098 matchAndRewrite(spirv::FModOp op, OpAdaptor adaptor,
1099 ConversionPatternRewriter &rewriter)
const override {
1100 Type dstType = getTypeConverter()->convertType(op.getType());
1102 return rewriter.notifyMatchFailure(op,
"type conversion failed");
1104 Location loc = op.getLoc();
1105 Value
lhs = adaptor.getOperand1();
1106 Value
rhs = adaptor.getOperand2();
1107 Value
div = LLVM::FDivOp::create(rewriter, loc, dstType,
lhs,
rhs);
1108 Value floored = LLVM::FFloorOp::create(rewriter, loc, dstType,
div);
1109 Value scaled = LLVM::FMulOp::create(rewriter, loc, dstType,
rhs, floored);
1110 rewriter.replaceOpWithNewOp<LLVM::FSubOp>(op, dstType,
lhs, scaled);
1121 using SPIRVToLLVMConversion<spirv::SModOp>::SPIRVToLLVMConversion;
1124 matchAndRewrite(spirv::SModOp op, OpAdaptor adaptor,
1125 ConversionPatternRewriter &rewriter)
const override {
1126 Type srcType = op.getType();
1127 Type dstType = getTypeConverter()->convertType(srcType);
1129 return rewriter.notifyMatchFailure(op,
"type conversion failed");
1131 Location loc = op.getLoc();
1132 Value
lhs = adaptor.getOperand1();
1133 Value
rhs = adaptor.getOperand2();
1134 Type i1Type = rewriter.getI1Type();
1135 auto vecSrcType = dyn_cast<VectorType>(srcType);
1137 vecSrcType ? VectorType::get(vecSrcType.getShape(), i1Type) : i1Type;
1139 Value
rem = LLVM::SRemOp::create(rewriter, loc, dstType,
lhs,
rhs);
1140 IntegerAttr zeroAttr = rewriter.getIntegerAttr(
1145 Value remNonZero = LLVM::ICmpOp::create(rewriter, loc, cmpType,
1146 LLVM::ICmpPredicate::ne,
rem, zero);
1147 Value remNeg = LLVM::ICmpOp::create(rewriter, loc, cmpType,
1148 LLVM::ICmpPredicate::slt,
rem, zero);
1149 Value rhsNeg = LLVM::ICmpOp::create(rewriter, loc, cmpType,
1150 LLVM::ICmpPredicate::slt,
rhs, zero);
1151 Value signMismatch =
1152 LLVM::XOrOp::create(rewriter, loc, cmpType, remNeg, rhsNeg);
1154 LLVM::AndOp::create(rewriter, loc, cmpType, remNonZero, signMismatch);
1156 Value adjusted = LLVM::AddOp::create(rewriter, loc, dstType,
rem,
rhs);
1157 rewriter.replaceOpWithNewOp<LLVM::SelectOp>(op, dstType, needsAdjust,
1164template <
typename SPIRVOp>
1167 using SPIRVToLLVMConversion<SPIRVOp>::SPIRVToLLVMConversion;
1170 matchAndRewrite(SPIRVOp op,
typename SPIRVOp::Adaptor adaptor,
1171 ConversionPatternRewriter &rewriter)
const override {
1172 if (!op.getMemoryAccess()) {
1174 *this->getTypeConverter(), 0,
1178 auto memoryAccess = *op.getMemoryAccess();
1179 switch (memoryAccess) {
1180 case spirv::MemoryAccess::Aligned:
1181 case spirv::MemoryAccess::None:
1182 case spirv::MemoryAccess::Nontemporal:
1183 case spirv::MemoryAccess::Volatile: {
1184 unsigned alignment =
1185 memoryAccess == spirv::MemoryAccess::Aligned ? *op.getAlignment() : 0;
1186 bool isNonTemporal = memoryAccess == spirv::MemoryAccess::Nontemporal;
1187 bool isVolatile = memoryAccess == spirv::MemoryAccess::Volatile;
1189 *this->getTypeConverter(), alignment,
1190 isVolatile, isNonTemporal);
1200template <
typename SPIRVOp>
1203 using SPIRVToLLVMConversion<SPIRVOp>::SPIRVToLLVMConversion;
1206 matchAndRewrite(SPIRVOp notOp,
typename SPIRVOp::Adaptor adaptor,
1207 ConversionPatternRewriter &rewriter)
const override {
1208 auto srcType = notOp.getType();
1209 auto dstType = this->getTypeConverter()->convertType(srcType);
1211 return rewriter.notifyMatchFailure(notOp,
"type conversion failed");
1213 Location loc = notOp.getLoc();
1215 rewriter.template replaceOpWithNewOp<LLVM::XOrOp>(notOp, dstType,
1216 notOp.getOperand(), mask);
1222template <
typename SPIRVOp>
1225 using SPIRVToLLVMConversion<SPIRVOp>::SPIRVToLLVMConversion;
1228 matchAndRewrite(SPIRVOp op,
typename SPIRVOp::Adaptor adaptor,
1229 ConversionPatternRewriter &rewriter)
const override {
1230 rewriter.eraseOp(op);
1237 using SPIRVToLLVMConversion<spirv::ReturnOp>::SPIRVToLLVMConversion;
1240 matchAndRewrite(spirv::ReturnOp returnOp, OpAdaptor adaptor,
1241 ConversionPatternRewriter &rewriter)
const override {
1242 rewriter.replaceOpWithNewOp<LLVM::ReturnOp>(returnOp, ArrayRef<Type>(),
1250 using SPIRVToLLVMConversion<spirv::ReturnValueOp>::SPIRVToLLVMConversion;
1253 matchAndRewrite(spirv::ReturnValueOp returnValueOp, OpAdaptor adaptor,
1254 ConversionPatternRewriter &rewriter)
const override {
1255 rewriter.replaceOpWithNewOp<LLVM::ReturnOp>(returnValueOp, ArrayRef<Type>(),
1256 adaptor.getOperands());
1263 using SPIRVToLLVMConversion<spirv::UnreachableOp>::SPIRVToLLVMConversion;
1266 matchAndRewrite(spirv::UnreachableOp unreachableOp, OpAdaptor adaptor,
1267 ConversionPatternRewriter &rewriter)
const override {
1268 rewriter.replaceOpWithNewOp<LLVM::UnreachableOp>(unreachableOp);
1277 bool convergent =
true) {
1278 auto func = dyn_cast_or_null<LLVM::LLVMFuncOp>(
1284 func = LLVM::LLVMFuncOp::create(
1285 b, symbolTable->
getLoc(), name,
1286 LLVM::LLVMFunctionType::get(resultType, paramTypes));
1287 func.setCConv(LLVM::cconv::CConv::SPIR_FUNC);
1288 func.setConvergent(convergent);
1289 func.setNoUnwind(
true);
1290 func.setWillReturn(
true);
1295 LLVM::LLVMFuncOp
func,
1297 auto call = LLVM::CallOp::create(builder, loc,
func, args);
1298 call.setCConv(
func.getCConv());
1299 call.setConvergentAttr(
func.getConvergentAttr());
1300 call.setNoUnwindAttr(
func.getNoUnwindAttr());
1301 call.setWillReturnAttr(
func.getWillReturnAttr());
1305template <
typename BarrierOpTy>
1308 using OpAdaptor =
typename SPIRVToLLVMConversion<BarrierOpTy>::OpAdaptor;
1310 using SPIRVToLLVMConversion<BarrierOpTy>::SPIRVToLLVMConversion;
1312 static constexpr StringRef getFuncName();
1315 matchAndRewrite(BarrierOpTy controlBarrierOp, OpAdaptor adaptor,
1316 ConversionPatternRewriter &rewriter)
const override {
1317 constexpr StringRef funcName = getFuncName();
1318 Operation *symbolTable =
1319 controlBarrierOp->template getParentWithTrait<OpTrait::SymbolTable>();
1321 Type i32 = rewriter.getI32Type();
1323 Type voidTy = rewriter.getType<LLVM::LLVMVoidType>();
1324 LLVM::LLVMFuncOp func =
1327 Location loc = controlBarrierOp->getLoc();
1328 Value execution = LLVM::ConstantOp::create(
1329 rewriter, loc, i32,
static_cast<int32_t
>(adaptor.getExecutionScope()));
1330 Value memory = LLVM::ConstantOp::create(
1331 rewriter, loc, i32,
static_cast<int32_t
>(adaptor.getMemoryScope()));
1332 Value semantics = LLVM::ConstantOp::create(
1333 rewriter, loc, i32,
static_cast<int32_t
>(adaptor.getMemorySemantics()));
1336 {execution, memory, semantics});
1338 rewriter.replaceOp(controlBarrierOp, call);
1345StringRef getTypeMangling(
Type type,
bool isSigned) {
1347 .Case([](Float16Type) {
return "Dh"; })
1348 .Case([](Float32Type) {
return "f"; })
1349 .Case([](Float64Type) {
return "d"; })
1350 .Case([isSigned](IntegerType intTy) {
1351 switch (intTy.getWidth()) {
1355 return (isSigned) ?
"a" :
"c";
1357 return (isSigned) ?
"s" :
"t";
1359 return (isSigned) ?
"i" :
"j";
1361 return (isSigned) ?
"l" :
"m";
1363 llvm_unreachable(
"Unsupported integer width");
1366 .DefaultUnreachable(
"No mangling defined");
1369template <
typename ReduceOp>
1370constexpr StringLiteral getGroupFuncName();
1373constexpr StringLiteral getGroupFuncName<spirv::GroupIAddOp>() {
1374 return "_Z17__spirv_GroupIAddii";
1377constexpr StringLiteral getGroupFuncName<spirv::GroupFAddOp>() {
1378 return "_Z17__spirv_GroupFAddii";
1381constexpr StringLiteral getGroupFuncName<spirv::GroupSMinOp>() {
1382 return "_Z17__spirv_GroupSMinii";
1385constexpr StringLiteral getGroupFuncName<spirv::GroupUMinOp>() {
1386 return "_Z17__spirv_GroupUMinii";
1389constexpr StringLiteral getGroupFuncName<spirv::GroupFMinOp>() {
1390 return "_Z17__spirv_GroupFMinii";
1393constexpr StringLiteral getGroupFuncName<spirv::GroupSMaxOp>() {
1394 return "_Z17__spirv_GroupSMaxii";
1397constexpr StringLiteral getGroupFuncName<spirv::GroupUMaxOp>() {
1398 return "_Z17__spirv_GroupUMaxii";
1401constexpr StringLiteral getGroupFuncName<spirv::GroupFMaxOp>() {
1402 return "_Z17__spirv_GroupFMaxii";
1405constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformIAddOp>() {
1406 return "_Z27__spirv_GroupNonUniformIAddii";
1409constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformFAddOp>() {
1410 return "_Z27__spirv_GroupNonUniformFAddii";
1413constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformIMulOp>() {
1414 return "_Z27__spirv_GroupNonUniformIMulii";
1417constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformFMulOp>() {
1418 return "_Z27__spirv_GroupNonUniformFMulii";
1421constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformSMinOp>() {
1422 return "_Z27__spirv_GroupNonUniformSMinii";
1425constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformUMinOp>() {
1426 return "_Z27__spirv_GroupNonUniformUMinii";
1429constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformFMinOp>() {
1430 return "_Z27__spirv_GroupNonUniformFMinii";
1433constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformSMaxOp>() {
1434 return "_Z27__spirv_GroupNonUniformSMaxii";
1437constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformUMaxOp>() {
1438 return "_Z27__spirv_GroupNonUniformUMaxii";
1441constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformFMaxOp>() {
1442 return "_Z27__spirv_GroupNonUniformFMaxii";
1445constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformBitwiseAndOp>() {
1446 return "_Z33__spirv_GroupNonUniformBitwiseAndii";
1449constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformBitwiseOrOp>() {
1450 return "_Z32__spirv_GroupNonUniformBitwiseOrii";
1453constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformBitwiseXorOp>() {
1454 return "_Z33__spirv_GroupNonUniformBitwiseXorii";
1457constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformLogicalAndOp>() {
1458 return "_Z33__spirv_GroupNonUniformLogicalAndii";
1461constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformLogicalOrOp>() {
1462 return "_Z32__spirv_GroupNonUniformLogicalOrii";
1465constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformLogicalXorOp>() {
1466 return "_Z33__spirv_GroupNonUniformLogicalXorii";
1470template <
typename ReduceOp,
bool Signed = false,
bool NonUniform = false>
1473 using SPIRVToLLVMConversion<ReduceOp>::SPIRVToLLVMConversion;
1476 matchAndRewrite(ReduceOp op,
typename ReduceOp::Adaptor adaptor,
1477 ConversionPatternRewriter &rewriter)
const override {
1479 Type retTy = op.getResult().getType();
1483 SmallString<36> funcName = getGroupFuncName<ReduceOp>();
1484 funcName += getTypeMangling(retTy,
false);
1486 Type i32Ty = rewriter.getI32Type();
1487 SmallVector<Type> paramTypes{i32Ty, i32Ty, retTy};
1488 if constexpr (NonUniform) {
1489 if (adaptor.getClusterSize()) {
1491 paramTypes.push_back(i32Ty);
1495 Operation *symbolTable =
1496 op->template getParentWithTrait<OpTrait::SymbolTable>();
1498 LLVM::LLVMFuncOp func =
1501 Location loc = op.getLoc();
1502 Value scope = LLVM::ConstantOp::create(
1503 rewriter, loc, i32Ty,
1504 static_cast<int32_t
>(adaptor.getExecutionScope()));
1505 Value groupOp = LLVM::ConstantOp::create(
1506 rewriter, loc, i32Ty,
1507 static_cast<int32_t
>(adaptor.getGroupOperation()));
1508 SmallVector<Value> operands{scope, groupOp};
1509 operands.append(adaptor.getOperands().begin(), adaptor.getOperands().end());
1512 rewriter.replaceOp(op, call);
1519ControlBarrierPattern<spirv::ControlBarrierOp>::getFuncName() {
1520 return "_Z22__spirv_ControlBarrieriii";
1525ControlBarrierPattern<spirv::INTELControlBarrierArriveOp>::getFuncName() {
1526 return "_Z33__spirv_ControlBarrierArriveINTELiii";
1531ControlBarrierPattern<spirv::INTELControlBarrierWaitOp>::getFuncName() {
1532 return "_Z31__spirv_ControlBarrierWaitINTELiii";
1585 using SPIRVToLLVMConversion<spirv::LoopOp>::SPIRVToLLVMConversion;
1588 matchAndRewrite(spirv::LoopOp loopOp, OpAdaptor adaptor,
1589 ConversionPatternRewriter &rewriter)
const override {
1591 if (loopOp.getLoopControl() != spirv::LoopControl::None)
1595 if (loopOp.getBody().empty()) {
1596 rewriter.eraseOp(loopOp);
1600 Location loc = loopOp.getLoc();
1604 Block *currentBlock = rewriter.getBlock();
1606 Block *endBlock = rewriter.splitBlock(currentBlock, position);
1610 Block *entryBlock = loopOp.getEntryBlock();
1612 auto brOp = dyn_cast<spirv::BranchOp>(entryBlock->
getOperations().front());
1615 Block *headerBlock = loopOp.getHeaderBlock();
1616 rewriter.setInsertionPointToEnd(currentBlock);
1617 LLVM::BrOp::create(rewriter, loc, brOp.getBlockArguments(), headerBlock);
1618 rewriter.eraseBlock(entryBlock);
1621 Block *mergeBlock = loopOp.getMergeBlock();
1624 rewriter.setInsertionPointToEnd(mergeBlock);
1625 LLVM::BrOp::create(rewriter, loc, terminatorOperands, endBlock);
1627 rewriter.inlineRegionBefore(loopOp.getBody(), endBlock);
1638 using SPIRVToLLVMConversion<spirv::SelectionOp>::SPIRVToLLVMConversion;
1641 matchAndRewrite(spirv::SelectionOp op, OpAdaptor adaptor,
1642 ConversionPatternRewriter &rewriter)
const override {
1646 if (op.getSelectionControl() != spirv::SelectionControl::None)
1653 if (op.getBody().getBlocks().size() <= 2) {
1654 rewriter.eraseOp(op);
1658 Location loc = op.getLoc();
1662 auto *currentBlock = rewriter.getInsertionBlock();
1663 rewriter.setInsertionPointAfter(op);
1664 auto position = rewriter.getInsertionPoint();
1665 auto *continueBlock = rewriter.splitBlock(currentBlock, position);
1668 for (
auto ty : op.getResultTypes()) {
1669 Type dstTy = getTypeConverter()->convertType(ty);
1671 return rewriter.notifyMatchFailure(op,
"failed to convert type");
1672 continueBlock->addArgument(dstTy, loc);
1679 auto *headerBlock = op.getHeaderBlock();
1681 auto condBrOp = dyn_cast<spirv::BranchConditionalOp>(
1687 auto *mergeBlock = op.getMergeBlock();
1690 rewriter.setInsertionPointToEnd(mergeBlock);
1691 LLVM::BrOp::create(rewriter, loc, terminatorOperands, continueBlock);
1694 Block *trueBlock = condBrOp.getTrueBlock();
1695 Block *falseBlock = condBrOp.getFalseBlock();
1696 rewriter.setInsertionPointToEnd(currentBlock);
1697 LLVM::CondBrOp::create(rewriter, loc, condBrOp.getCondition(), trueBlock,
1698 condBrOp.getTrueTargetOperands(), falseBlock,
1699 condBrOp.getFalseTargetOperands());
1701 rewriter.eraseBlock(headerBlock);
1702 rewriter.inlineRegionBefore(op.getBody(), continueBlock);
1703 rewriter.replaceOp(op, continueBlock->getArguments());
1712template <
typename SPIRVOp,
typename LLVMOp>
1715 using SPIRVToLLVMConversion<SPIRVOp>::SPIRVToLLVMConversion;
1718 matchAndRewrite(SPIRVOp op,
typename SPIRVOp::Adaptor adaptor,
1719 ConversionPatternRewriter &rewriter)
const override {
1721 auto dstType = this->getTypeConverter()->convertType(op.getType());
1723 return rewriter.notifyMatchFailure(op,
"type conversion failed");
1725 Type op1Type = op.getOperand1().getType();
1726 Type op2Type = op.getOperand2().getType();
1728 if (op1Type == op2Type) {
1729 rewriter.template replaceOpWithNewOp<LLVMOp>(op, dstType,
1730 adaptor.getOperands());
1734 std::optional<uint64_t> dstTypeWidth =
1736 std::optional<uint64_t> op2TypeWidth =
1739 if (!dstTypeWidth || !op2TypeWidth)
1742 Location loc = op.getLoc();
1744 if (op2TypeWidth < dstTypeWidth) {
1747 LLVM::ZExtOp::create(rewriter, loc, dstType, adaptor.getOperand2());
1750 LLVM::SExtOp::create(rewriter, loc, dstType, adaptor.getOperand2());
1752 }
else if (op2TypeWidth == dstTypeWidth) {
1753 extended = adaptor.getOperand2();
1759 LLVMOp::create(rewriter, loc, dstType, adaptor.getOperand1(), extended);
1760 rewriter.replaceOp(op,
result);
1770 using SPIRVToLLVMConversion<spirv::GLSAbsOp>::SPIRVToLLVMConversion;
1773 matchAndRewrite(spirv::GLSAbsOp op, OpAdaptor adaptor,
1774 ConversionPatternRewriter &rewriter)
const override {
1775 Type dstType = getTypeConverter()->convertType(op.getType());
1777 return rewriter.notifyMatchFailure(op,
"type conversion failed");
1779 rewriter.replaceOpWithNewOp<LLVM::AbsOp>(op, dstType, adaptor.getOperand(),
1788 using SPIRVToLLVMConversion<spirv::GLFractOp>::SPIRVToLLVMConversion;
1791 matchAndRewrite(spirv::GLFractOp op, OpAdaptor adaptor,
1792 ConversionPatternRewriter &rewriter)
const override {
1793 Type dstType = getTypeConverter()->convertType(op.getType());
1795 return rewriter.notifyMatchFailure(op,
"type conversion failed");
1797 Location loc = op.getLoc();
1798 Value operand = adaptor.getOperand();
1799 Value floored = LLVM::FFloorOp::create(rewriter, loc, dstType, operand);
1800 rewriter.replaceOpWithNewOp<LLVM::FSubOp>(op, dstType, operand, floored);
1809 using SPIRVToLLVMConversion<spirv::GLFMixOp>::SPIRVToLLVMConversion;
1812 matchAndRewrite(spirv::GLFMixOp op, OpAdaptor adaptor,
1813 ConversionPatternRewriter &rewriter)
const override {
1814 Type dstType = getTypeConverter()->convertType(op.getType());
1816 return rewriter.notifyMatchFailure(op,
"type conversion failed");
1818 Location loc = op.getLoc();
1819 Value x = adaptor.getX();
1820 Value y = adaptor.getY();
1821 Value a = adaptor.getA();
1823 Value oneMinusA = LLVM::FSubOp::create(rewriter, loc, dstType, one, a);
1824 Value
lhs = LLVM::FMulOp::create(rewriter, loc, dstType, x, oneMinusA);
1825 Value
rhs = LLVM::FMulOp::create(rewriter, loc, dstType, y, a);
1826 rewriter.replaceOpWithNewOp<LLVM::FAddOp>(op, dstType,
lhs,
rhs);
1835 using SPIRVToLLVMConversion<spirv::CLMixOp>::SPIRVToLLVMConversion;
1838 matchAndRewrite(spirv::CLMixOp op, OpAdaptor adaptor,
1839 ConversionPatternRewriter &rewriter)
const override {
1840 Type dstType = getTypeConverter()->convertType(op.getType());
1842 return rewriter.notifyMatchFailure(op,
"type conversion failed");
1844 Location loc = op.getLoc();
1845 Value x = adaptor.getX();
1846 Value y = adaptor.getY();
1847 Value a = adaptor.getZ();
1848 Value diff = LLVM::FSubOp::create(rewriter, loc, dstType, y, x);
1849 rewriter.replaceOpWithNewOp<LLVM::FMAOp>(op, dstType, a, diff, x);
1856template <
typename SPIRVOp>
1859 template <
typename... Args>
1860 ScalePattern(
double scale, Args &&...args)
1861 : SPIRVToLLVMConversion<SPIRVOp>(std::forward<Args>(args)...),
1865 matchAndRewrite(SPIRVOp op,
typename SPIRVOp::Adaptor adaptor,
1866 ConversionPatternRewriter &rewriter)
const override {
1867 Type srcType = op.getType();
1868 Type dstType = this->getTypeConverter()->convertType(srcType);
1870 return rewriter.notifyMatchFailure(op,
"type conversion failed");
1872 Location loc = op.getLoc();
1874 rewriter.replaceOpWithNewOp<LLVM::FMulOp>(op, dstType, adaptor.getOperand(),
1886template <
typename SPIRVOp,
bool isFloat>
1889 using SPIRVToLLVMConversion<SPIRVOp>::SPIRVToLLVMConversion;
1892 matchAndRewrite(SPIRVOp op,
typename SPIRVOp::Adaptor adaptor,
1893 ConversionPatternRewriter &rewriter)
const override {
1894 Type srcType = op.getType();
1895 Type dstType = this->getTypeConverter()->convertType(srcType);
1897 return rewriter.notifyMatchFailure(op,
"type conversion failed");
1899 Location loc = op.getLoc();
1900 Value operand = adaptor.getOperand();
1901 auto vecSrcType = dyn_cast<VectorType>(srcType);
1902 Type i1Type = rewriter.getI1Type();
1904 vecSrcType ? VectorType::get(vecSrcType.getShape(), i1Type) : i1Type;
1906 Value zero, one, minusOne, gt, lt;
1907 if constexpr (isFloat) {
1911 gt = LLVM::FCmpOp::create(rewriter, loc, cmpType,
1912 LLVM::FCmpPredicate::ogt, operand, zero);
1913 lt = LLVM::FCmpOp::create(rewriter, loc, cmpType,
1914 LLVM::FCmpPredicate::olt, operand, zero);
1918 rewriter.getIntegerAttr(intElemType, 0));
1920 rewriter.getIntegerAttr(intElemType, 1));
1922 gt = LLVM::ICmpOp::create(rewriter, loc, cmpType,
1923 LLVM::ICmpPredicate::sgt, operand, zero);
1924 lt = LLVM::ICmpOp::create(rewriter, loc, cmpType,
1925 LLVM::ICmpPredicate::slt, operand, zero);
1929 LLVM::SelectOp::create(rewriter, loc, dstType, lt, minusOne, zero);
1930 rewriter.replaceOpWithNewOp<LLVM::SelectOp>(op, dstType, gt, one,
1938 using SPIRVToLLVMConversion<spirv::VariableOp>::SPIRVToLLVMConversion;
1941 matchAndRewrite(spirv::VariableOp varOp, OpAdaptor adaptor,
1942 ConversionPatternRewriter &rewriter)
const override {
1943 auto srcType = varOp.getType();
1945 auto pointerTo = cast<spirv::PointerType>(srcType).getPointeeType();
1946 auto init = varOp.getInitializer();
1947 if (init && !pointerTo.isIntOrFloat() && !isa<VectorType>(pointerTo))
1950 auto dstType = getTypeConverter()->convertType(srcType);
1952 return rewriter.notifyMatchFailure(varOp,
"type conversion failed");
1954 Location loc = varOp.getLoc();
1957 auto elementType = getTypeConverter()->convertType(pointerTo);
1959 return rewriter.notifyMatchFailure(varOp,
"type conversion failed");
1960 rewriter.replaceOpWithNewOp<LLVM::AllocaOp>(varOp, dstType, elementType,
1964 auto elementType = getTypeConverter()->convertType(pointerTo);
1966 return rewriter.notifyMatchFailure(varOp,
"type conversion failed");
1968 LLVM::AllocaOp::create(rewriter, loc, dstType, elementType, size);
1969 LLVM::StoreOp::create(rewriter, loc, adaptor.getInitializer(), allocated);
1970 rewriter.replaceOp(varOp, allocated);
1979class BitcastConversionPattern
1982 using SPIRVToLLVMConversion<spirv::BitcastOp>::SPIRVToLLVMConversion;
1985 matchAndRewrite(spirv::BitcastOp bitcastOp, OpAdaptor adaptor,
1986 ConversionPatternRewriter &rewriter)
const override {
1987 auto dstType = getTypeConverter()->convertType(bitcastOp.getType());
1989 return rewriter.notifyMatchFailure(bitcastOp,
"type conversion failed");
1992 if (isa<LLVM::LLVMPointerType>(dstType)) {
1993 rewriter.replaceOp(bitcastOp, adaptor.getOperand());
1997 LLVM::BitcastOp::Properties properties{};
1998 SmallVector<NamedAttribute> discardableAttrs;
1999 if (
failed(collectAttrsForConversion<LLVM::BitcastOp>(bitcastOp, properties,
2002 rewriter.replaceOpWithNewOp<LLVM::BitcastOp>(bitcastOp, dstType,
2003 adaptor.getOperands(),
2004 properties, discardableAttrs);
2015 using SPIRVToLLVMConversion<spirv::FuncOp>::SPIRVToLLVMConversion;
2018 matchAndRewrite(spirv::FuncOp funcOp, OpAdaptor adaptor,
2019 ConversionPatternRewriter &rewriter)
const override {
2023 auto funcType = funcOp.getFunctionType();
2024 TypeConverter::SignatureConversion signatureConverter(
2025 funcType.getNumInputs());
2026 auto llvmType =
static_cast<const LLVMTypeConverter *
>(getTypeConverter())
2027 ->convertFunctionSignature(
2029 false, signatureConverter);
2034 Location loc = funcOp.getLoc();
2035 StringRef name = funcOp.getName();
2036 auto newFuncOp = LLVM::LLVMFuncOp::create(rewriter, loc, name, llvmType);
2039 MLIRContext *context = funcOp.getContext();
2040 switch (funcOp.getFunctionControl()) {
2041 case spirv::FunctionControl::Inline:
2042 newFuncOp.setAlwaysInline(
true);
2044 case spirv::FunctionControl::DontInline:
2045 newFuncOp.setNoInline(
true);
2048#define DISPATCH(functionControl, llvmAttr) \
2049 case functionControl: \
2050 newFuncOp->setDiscardableAttr("passthrough", \
2051 ArrayAttr::get(context, {llvmAttr})); \
2054 DISPATCH(spirv::FunctionControl::Pure,
2055 StringAttr::get(context,
"readonly"));
2056 DISPATCH(spirv::FunctionControl::Const,
2057 StringAttr::get(context,
"readnone"));
2067 rewriter.inlineRegionBefore(funcOp.getBody(), newFuncOp.getBody(),
2069 if (
failed(rewriter.convertRegionTypes(
2070 &newFuncOp.getBody(), *getTypeConverter(), &signatureConverter))) {
2073 rewriter.eraseOp(funcOp);
2084 using SPIRVToLLVMConversion<spirv::ModuleOp>::SPIRVToLLVMConversion;
2087 matchAndRewrite(spirv::ModuleOp spvModuleOp, OpAdaptor adaptor,
2088 ConversionPatternRewriter &rewriter)
const override {
2091 ModuleOp::create(rewriter, spvModuleOp.getLoc(), spvModuleOp.getName());
2092 rewriter.inlineRegionBefore(spvModuleOp.getRegion(), newModuleOp.getBody());
2095 rewriter.eraseBlock(&newModuleOp.getBodyRegion().back());
2096 rewriter.eraseOp(spvModuleOp);
2105class VectorShufflePattern
2108 using SPIRVToLLVMConversion<spirv::VectorShuffleOp>::SPIRVToLLVMConversion;
2110 matchAndRewrite(spirv::VectorShuffleOp op, OpAdaptor adaptor,
2111 ConversionPatternRewriter &rewriter)
const override {
2112 Location loc = op.getLoc();
2113 auto components = adaptor.getComponents();
2114 auto vector1 = adaptor.getVector1();
2115 auto vector2 = adaptor.getVector2();
2116 int vector1Size = cast<VectorType>(vector1.getType()).getNumElements();
2117 int vector2Size = cast<VectorType>(vector2.getType()).getNumElements();
2118 if (vector1Size == vector2Size) {
2119 rewriter.replaceOpWithNewOp<LLVM::ShuffleVectorOp>(
2120 op, vector1, vector2,
2121 LLVM::convertArrayToIndices<int32_t>(components));
2125 auto dstType = getTypeConverter()->convertType(op.getType());
2127 return rewriter.notifyMatchFailure(op,
"type conversion failed");
2128 auto scalarType = cast<VectorType>(dstType).getElementType();
2129 auto componentsArray = components.getValue();
2130 auto *context = rewriter.getContext();
2131 auto llvmI32Type = IntegerType::get(context, 32);
2132 Value targetOp = LLVM::PoisonOp::create(rewriter, loc, dstType);
2133 for (
unsigned i = 0; i < componentsArray.size(); i++) {
2134 if (!isa<IntegerAttr>(componentsArray[i]))
2135 return op.emitError(
"unable to support non-constant component");
2137 int indexVal = cast<IntegerAttr>(componentsArray[i]).getInt();
2142 Value baseVector = vector1;
2143 if (indexVal >= vector1Size) {
2144 offsetVal = vector1Size;
2145 baseVector = vector2;
2148 Value dstIndex = LLVM::ConstantOp::create(
2149 rewriter, loc, llvmI32Type,
2150 rewriter.getIntegerAttr(rewriter.getI32Type(), i));
2151 Value index = LLVM::ConstantOp::create(
2152 rewriter, loc, llvmI32Type,
2153 rewriter.getIntegerAttr(rewriter.getI32Type(), indexVal - offsetVal));
2155 auto extractOp = LLVM::ExtractElementOp::create(rewriter, loc, scalarType,
2157 targetOp = LLVM::InsertElementOp::create(rewriter, loc, dstType, targetOp,
2158 extractOp, dstIndex);
2160 rewriter.replaceOp(op, targetOp);
2171 spirv::ClientAPI clientAPI) {
2188 spirv::ClientAPI clientAPI) {
2191 DirectConversionPattern<spirv::IAddOp, LLVM::AddOp>,
2192 DirectConversionPattern<spirv::IMulOp, LLVM::MulOp>,
2193 DirectConversionPattern<spirv::ISubOp, LLVM::SubOp>,
2194 DirectConversionPattern<spirv::FAddOp, LLVM::FAddOp>,
2195 DirectConversionPattern<spirv::FDivOp, LLVM::FDivOp>,
2196 DirectConversionPattern<spirv::FMulOp, LLVM::FMulOp>,
2197 DirectConversionPattern<spirv::FNegateOp, LLVM::FNegOp>,
2198 DirectConversionPattern<spirv::FRemOp, LLVM::FRemOp>,
2199 DirectConversionPattern<spirv::FSubOp, LLVM::FSubOp>,
2200 DirectConversionPattern<spirv::SDivOp, LLVM::SDivOp>,
2201 DirectConversionPattern<spirv::SRemOp, LLVM::SRemOp>,
2202 DirectConversionPattern<spirv::UDivOp, LLVM::UDivOp>,
2203 DirectConversionPattern<spirv::UModOp, LLVM::URemOp>, FModPattern,
2204 SModPattern, VectorTimesScalarPattern, SNegatePattern,
2205 ArithmeticWithOverflowPattern<spirv::IAddCarryOp,
2206 LLVM::UAddWithOverflowOp>,
2207 ArithmeticWithOverflowPattern<spirv::ISubBorrowOp,
2208 LLVM::USubWithOverflowOp>,
2211 BitFieldInsertPattern, BitFieldUExtractPattern, BitFieldSExtractPattern,
2212 DirectConversionPattern<spirv::BitCountOp, LLVM::CtPopOp>,
2213 DirectConversionPattern<spirv::BitReverseOp, LLVM::BitReverseOp>,
2214 DirectConversionPattern<spirv::BitwiseAndOp, LLVM::AndOp>,
2215 DirectConversionPattern<spirv::BitwiseOrOp, LLVM::OrOp>,
2216 DirectConversionPattern<spirv::BitwiseXorOp, LLVM::XOrOp>,
2217 NotPattern<spirv::NotOp>,
2220 BitcastConversionPattern,
2221 DirectConversionPattern<spirv::ConvertFToSOp, LLVM::FPToSIOp>,
2222 DirectConversionPattern<spirv::ConvertFToUOp, LLVM::FPToUIOp>,
2223 DirectConversionPattern<spirv::ConvertSToFOp, LLVM::SIToFPOp>,
2224 DirectConversionPattern<spirv::ConvertUToFOp, LLVM::UIToFPOp>,
2225 IndirectCastPattern<spirv::FConvertOp, LLVM::FPExtOp, LLVM::FPTruncOp>,
2226 IndirectCastPattern<spirv::SConvertOp, LLVM::SExtOp, LLVM::TruncOp>,
2227 IndirectCastPattern<spirv::UConvertOp, LLVM::ZExtOp, LLVM::TruncOp>,
2228 DirectConversionPattern<spirv::ConvertPtrToUOp, LLVM::PtrToIntOp>,
2229 DirectConversionPattern<spirv::ConvertUToPtrOp, LLVM::IntToPtrOp>,
2230 DirectConversionPattern<spirv::PtrCastToGenericOp, LLVM::AddrSpaceCastOp>,
2231 DirectConversionPattern<spirv::GenericCastToPtrOp, LLVM::AddrSpaceCastOp>,
2232 DirectConversionPattern<spirv::GenericCastToPtrExplicitOp,
2233 LLVM::AddrSpaceCastOp>,
2236 IComparePattern<spirv::IEqualOp, LLVM::ICmpPredicate::eq>,
2237 IComparePattern<spirv::INotEqualOp, LLVM::ICmpPredicate::ne>,
2238 FComparePattern<spirv::FOrdEqualOp, LLVM::FCmpPredicate::oeq>,
2239 FComparePattern<spirv::FOrdGreaterThanOp, LLVM::FCmpPredicate::ogt>,
2240 FComparePattern<spirv::FOrdGreaterThanEqualOp, LLVM::FCmpPredicate::oge>,
2241 FComparePattern<spirv::FOrdLessThanEqualOp, LLVM::FCmpPredicate::ole>,
2242 FComparePattern<spirv::FOrdLessThanOp, LLVM::FCmpPredicate::olt>,
2243 FComparePattern<spirv::FOrdNotEqualOp, LLVM::FCmpPredicate::one>,
2244 FComparePattern<spirv::FUnordEqualOp, LLVM::FCmpPredicate::ueq>,
2245 FComparePattern<spirv::FUnordGreaterThanOp, LLVM::FCmpPredicate::ugt>,
2246 FComparePattern<spirv::FUnordGreaterThanEqualOp,
2247 LLVM::FCmpPredicate::uge>,
2248 FComparePattern<spirv::FUnordLessThanEqualOp, LLVM::FCmpPredicate::ule>,
2249 FComparePattern<spirv::FUnordLessThanOp, LLVM::FCmpPredicate::ult>,
2250 FComparePattern<spirv::FUnordNotEqualOp, LLVM::FCmpPredicate::une>,
2251 FComparePattern<spirv::OrderedOp, LLVM::FCmpPredicate::ord>,
2252 FComparePattern<spirv::UnorderedOp, LLVM::FCmpPredicate::uno>,
2253 IComparePattern<spirv::SGreaterThanOp, LLVM::ICmpPredicate::sgt>,
2254 IComparePattern<spirv::SGreaterThanEqualOp, LLVM::ICmpPredicate::sge>,
2255 IComparePattern<spirv::SLessThanEqualOp, LLVM::ICmpPredicate::sle>,
2256 IComparePattern<spirv::SLessThanOp, LLVM::ICmpPredicate::slt>,
2257 IComparePattern<spirv::UGreaterThanOp, LLVM::ICmpPredicate::ugt>,
2258 IComparePattern<spirv::UGreaterThanEqualOp, LLVM::ICmpPredicate::uge>,
2259 IComparePattern<spirv::ULessThanEqualOp, LLVM::ICmpPredicate::ule>,
2260 IComparePattern<spirv::ULessThanOp, LLVM::ICmpPredicate::ult>,
2263 ConstantScalarAndVectorPattern,
2266 BranchConversionPattern, BranchConditionalConversionPattern,
2267 FunctionCallPattern, LoopPattern, SelectionPattern,
2268 ErasePattern<spirv::MergeOp>,
2271 ErasePattern<spirv::EntryPointOp>, ExecutionModePattern,
2274 DirectConversionPattern<spirv::GLCeilOp, LLVM::FCeilOp>,
2275 DirectConversionPattern<spirv::GLCosOp, LLVM::CosOp>,
2276 DirectConversionPattern<spirv::GLExpOp, LLVM::ExpOp>,
2277 DirectConversionPattern<spirv::GLExp2Op, LLVM::Exp2Op>,
2278 DirectConversionPattern<spirv::GLFAbsOp, LLVM::FAbsOp>,
2279 DirectConversionPattern<spirv::GLFloorOp, LLVM::FFloorOp>,
2280 DirectConversionPattern<spirv::GLFmaOp, LLVM::FMAOp>,
2281 ClampPattern<spirv::GLFClampOp, LLVM::MinNumOp, LLVM::MaxNumOp>,
2282 ClampPattern<spirv::GLSClampOp, LLVM::SMinOp, LLVM::SMaxOp>,
2283 ClampPattern<spirv::GLUClampOp, LLVM::UMinOp, LLVM::UMaxOp>,
2284 DirectConversionPattern<spirv::GLFMaxOp, LLVM::MaxNumOp>,
2285 DirectConversionPattern<spirv::GLFMinOp, LLVM::MinNumOp>,
2286 DirectConversionPattern<spirv::GLNMaxOp, LLVM::MaxNumOp>,
2287 DirectConversionPattern<spirv::GLNMinOp, LLVM::MinNumOp>,
2288 DirectConversionPattern<spirv::GLLogOp, LLVM::LogOp>,
2289 DirectConversionPattern<spirv::GLLog2Op, LLVM::Log2Op>,
2290 DirectConversionPattern<spirv::GLPowOp, LLVM::PowOp>,
2291 DirectConversionPattern<spirv::GLRoundOp, LLVM::RoundOp>,
2292 DirectConversionPattern<spirv::GLRoundEvenOp, LLVM::RoundEvenOp>,
2293 DirectConversionPattern<spirv::GLSinOp, LLVM::SinOp>,
2294 DirectConversionPattern<spirv::GLSinhOp, LLVM::SinhOp>,
2295 DirectConversionPattern<spirv::GLCoshOp, LLVM::CoshOp>,
2296 DirectConversionPattern<spirv::GLSMaxOp, LLVM::SMaxOp>,
2297 DirectConversionPattern<spirv::GLSMinOp, LLVM::SMinOp>,
2298 DirectConversionPattern<spirv::GLSqrtOp, LLVM::SqrtOp>,
2299 DirectConversionPattern<spirv::GLUMaxOp, LLVM::UMaxOp>,
2300 DirectConversionPattern<spirv::GLUMinOp, LLVM::UMinOp>,
2301 DirectConversionPattern<spirv::GLTruncOp, LLVM::FTruncOp>,
2302 DirectConversionPattern<spirv::GLAsinOp, LLVM::ASinOp>,
2303 DirectConversionPattern<spirv::GLAcosOp, LLVM::ACosOp>,
2304 DirectConversionPattern<spirv::GLAtanOp, LLVM::ATanOp>,
2305 DirectConversionPattern<spirv::GLTanOp, LLVM::TanOp>,
2306 DirectConversionPattern<spirv::GLTanhOp, LLVM::TanhOp>,
2307 InverseSqrtPattern, SAbsPattern, FractPattern,
2308 SignPattern<spirv::GLFSignOp,
true>,
2309 SignPattern<spirv::GLSSignOp,
false>, GLFMixPattern,
2312 DirectConversionPattern<spirv::CLCeilOp, LLVM::FCeilOp>,
2313 DirectConversionPattern<spirv::CLCosOp, LLVM::CosOp>,
2314 DirectConversionPattern<spirv::CLExpOp, LLVM::ExpOp>,
2315 DirectConversionPattern<spirv::CLExp2Op, LLVM::Exp2Op>,
2316 DirectConversionPattern<spirv::CLExp10Op, LLVM::Exp10Op>,
2317 DirectConversionPattern<spirv::CLFAbsOp, LLVM::FAbsOp>,
2318 DirectConversionPattern<spirv::CLFloorOp, LLVM::FFloorOp>,
2319 DirectConversionPattern<spirv::CLFmaOp, LLVM::FMAOp>,
2320 DirectConversionPattern<spirv::CLFMaxOp, LLVM::MaxNumOp>,
2321 DirectConversionPattern<spirv::CLFMinOp, LLVM::MinNumOp>,
2322 DirectConversionPattern<spirv::CLLogOp, LLVM::LogOp>,
2323 DirectConversionPattern<spirv::CLLog2Op, LLVM::Log2Op>,
2324 DirectConversionPattern<spirv::CLLog10Op, LLVM::Log10Op>,
2325 DirectConversionPattern<spirv::CLPowOp, LLVM::PowOp>,
2326 DirectConversionPattern<spirv::CLRintOp, LLVM::RintOp>,
2327 DirectConversionPattern<spirv::CLRoundOp, LLVM::RoundOp>,
2328 DirectConversionPattern<spirv::CLSinOp, LLVM::SinOp>,
2329 DirectConversionPattern<spirv::CLSinhOp, LLVM::SinhOp>,
2330 DirectConversionPattern<spirv::CLCoshOp, LLVM::CoshOp>,
2331 DirectConversionPattern<spirv::CLTanOp, LLVM::TanOp>,
2332 DirectConversionPattern<spirv::CLTanhOp, LLVM::TanhOp>,
2333 DirectConversionPattern<spirv::CLAsinOp, LLVM::ASinOp>,
2334 DirectConversionPattern<spirv::CLAcosOp, LLVM::ACosOp>,
2335 DirectConversionPattern<spirv::CLAtanOp, LLVM::ATanOp>,
2336 DirectConversionPattern<spirv::CLAtan2Op, LLVM::ATan2Op>,
2337 DirectConversionPattern<spirv::CLSqrtOp, LLVM::SqrtOp>,
2338 DirectConversionPattern<spirv::CLTruncOp, LLVM::FTruncOp>,
2339 DirectConversionPattern<spirv::CLCopysignOp, LLVM::CopySignOp>,
2340 DirectConversionPattern<spirv::CLFmodOp, LLVM::FRemOp>,
2341 DirectConversionPattern<spirv::CLSMaxOp, LLVM::SMaxOp>,
2342 DirectConversionPattern<spirv::CLSMinOp, LLVM::SMinOp>,
2343 DirectConversionPattern<spirv::CLUMaxOp, LLVM::UMaxOp>,
2344 DirectConversionPattern<spirv::CLUMinOp, LLVM::UMinOp>, CLMixPattern,
2347 DirectConversionPattern<spirv::LogicalAndOp, LLVM::AndOp>,
2348 DirectConversionPattern<spirv::LogicalOrOp, LLVM::OrOp>,
2349 IComparePattern<spirv::LogicalEqualOp, LLVM::ICmpPredicate::eq>,
2350 IComparePattern<spirv::LogicalNotEqualOp, LLVM::ICmpPredicate::ne>,
2351 NotPattern<spirv::LogicalNotOp>,
2354 AccessChainPattern, AddressOfPattern, LoadStorePattern<spirv::LoadOp>,
2355 LoadStorePattern<spirv::StoreOp>, VariablePattern,
2358 CompositeExtractPattern, CompositeInsertPattern,
2359 DirectConversionPattern<spirv::SelectOp, LLVM::SelectOp>,
2360 DirectConversionPattern<spirv::UndefOp, LLVM::UndefOp>,
2361 VectorShufflePattern,
2364 ShiftPattern<spirv::ShiftRightArithmeticOp, LLVM::AShrOp>,
2365 ShiftPattern<spirv::ShiftRightLogicalOp, LLVM::LShrOp>,
2366 ShiftPattern<spirv::ShiftLeftLogicalOp, LLVM::ShlOp>,
2369 ReturnPattern, ReturnValuePattern,
2375 ControlBarrierPattern<spirv::ControlBarrierOp>,
2376 ControlBarrierPattern<spirv::INTELControlBarrierArriveOp>,
2377 ControlBarrierPattern<spirv::INTELControlBarrierWaitOp>,
2380 GroupReducePattern<spirv::GroupIAddOp>,
2381 GroupReducePattern<spirv::GroupFAddOp>,
2382 GroupReducePattern<spirv::GroupFMinOp>,
2383 GroupReducePattern<spirv::GroupUMinOp>,
2384 GroupReducePattern<spirv::GroupSMinOp,
true>,
2385 GroupReducePattern<spirv::GroupFMaxOp>,
2386 GroupReducePattern<spirv::GroupUMaxOp>,
2387 GroupReducePattern<spirv::GroupSMaxOp,
true>,
2388 GroupReducePattern<spirv::GroupNonUniformIAddOp,
false,
2390 GroupReducePattern<spirv::GroupNonUniformFAddOp,
false,
2392 GroupReducePattern<spirv::GroupNonUniformIMulOp,
false,
2394 GroupReducePattern<spirv::GroupNonUniformFMulOp,
false,
2396 GroupReducePattern<spirv::GroupNonUniformSMinOp,
true,
2398 GroupReducePattern<spirv::GroupNonUniformUMinOp,
false,
2400 GroupReducePattern<spirv::GroupNonUniformFMinOp,
false,
2402 GroupReducePattern<spirv::GroupNonUniformSMaxOp,
true,
2404 GroupReducePattern<spirv::GroupNonUniformUMaxOp,
false,
2406 GroupReducePattern<spirv::GroupNonUniformFMaxOp,
false,
2408 GroupReducePattern<spirv::GroupNonUniformBitwiseAndOp,
false,
2410 GroupReducePattern<spirv::GroupNonUniformBitwiseOrOp,
false,
2412 GroupReducePattern<spirv::GroupNonUniformBitwiseXorOp,
false,
2414 GroupReducePattern<spirv::GroupNonUniformLogicalAndOp,
false,
2416 GroupReducePattern<spirv::GroupNonUniformLogicalOrOp,
false,
2418 GroupReducePattern<spirv::GroupNonUniformLogicalXorOp,
false,
2422 patterns.
add<GlobalVariablePattern>(clientAPI, patterns.
getContext(),
2425 patterns.
add<ScalePattern<spirv::GLRadiansOp>>(
2426 0.017453292519943295, patterns.
getContext(), typeConverter);
2428 patterns.
add<ScalePattern<spirv::GLDegreesOp>>(
2429 57.29577951308232, patterns.
getContext(), typeConverter);
2434 patterns.
add<FuncConversionPattern>(patterns.
getContext(), typeConverter);
2439 patterns.
add<ModuleConversionPattern>(patterns.
getContext(), typeConverter);
2448 auto spvModules =
module.getOps<spirv::ModuleOp>();
2449 for (
auto spvModule : spvModules) {
2450 spvModule.walk([&](spirv::GlobalVariableOp op) {
2451 IntegerAttr descriptorSet = op.getDescriptorSetAttr();
2452 IntegerAttr binding = op.getBindingAttr();
2455 if (descriptorSet && binding) {
2458 auto moduleAndName =
2459 spvModule.getName().has_value()
2460 ? spvModule.getName()->str() +
"_" + op.getSymName().str()
2461 : op.getSymName().str();
2463 llvm::formatv(
"{0}_descriptor_set{1}_binding{2}", moduleAndName,
2464 std::to_string(descriptorSet.getInt()),
2465 std::to_string(binding.getInt()));
2466 auto nameAttr = StringAttr::get(op->getContext(), name);
2471 op.emitError(
"unable to replace all symbol uses for ") << name;
2473 op.removeDescriptorSetAttr();
2474 op.removeBindingAttr();
static LLVM::CallOp createSPIRVBuiltinCall(Location loc, ConversionPatternRewriter &rewriter, LLVM::LLVMFuncOp func, ValueRange args)
static LLVM::LLVMFuncOp lookupOrCreateSPIRVFn(Operation *symbolTable, StringRef name, ArrayRef< Type > paramTypes, Type resultType, bool isMemNone, bool isConvergent)
static Value max(ImplicitLocOpBuilder &builder, Value value, Value bound)
static Type getElementType(Type type, ArrayRef< int32_t > indices, function_ref< InFlightDiagnostic(StringRef)> emitErrorFn)
Walks the given type hierarchy with the given indices, potentially down to component granularity,...
static Value optionallyTruncateOrExtend(Location loc, Value value, Type llvmType, PatternRewriter &rewriter)
Utility function for bitfield ops:
static Value createFPConstant(Location loc, Type srcType, Type dstType, PatternRewriter &rewriter, double value)
Creates llvm.mlir.constant with a floating-point scalar or vector value.
static Value createI32ConstantOf(Location loc, PatternRewriter &rewriter, unsigned value)
Creates LLVM dialect constant with the given value.
static Type convertPointerType(spirv::PointerType type, const TypeConverter &converter, spirv::ClientAPI clientAPI)
Converts SPIR-V pointer type to LLVM pointer.
static Value processCountOrOffset(Location loc, Value value, Type srcType, Type dstType, const TypeConverter &converter, ConversionPatternRewriter &rewriter)
Utility function for bitfield ops: BitFieldInsert, BitFieldSExtract and BitFieldUExtract.
static unsigned getBitWidth(Type type)
Returns the bit width of integer, float or vector of float or integer values.
static LogicalResult replaceWithLoadOrStore(Operation *op, ValueRange operands, ConversionPatternRewriter &rewriter, const TypeConverter &typeConverter, unsigned alignment, bool isVolatile, bool isNonTemporal)
Utility for spirv.Load and spirv.Store conversion.
static Type convertStructTypePacked(spirv::StructType type, const TypeConverter &converter)
Converts SPIR-V struct with no offset to packed LLVM struct.
static bool isSignedIntegerOrVector(Type type)
Returns true if the given type is a signed integer or vector type.
static bool isUnsignedIntegerOrVector(Type type)
Returns true if the given type is an unsigned integer or vector type.
static std::optional< Type > convertRuntimeArrayType(spirv::RuntimeArrayType type, TypeConverter &converter)
Converts SPIR-V runtime array to LLVM array.
static Value optionallyBroadcast(Location loc, Value value, Type srcType, const TypeConverter &typeConverter, ConversionPatternRewriter &rewriter)
Broadcasts the value. If srcType is a scalar, the value remains unchanged.
static Value createConstantAllBitsSet(Location loc, Type srcType, Type dstType, PatternRewriter &rewriter)
Creates llvm.mlir.constant with all bits set for the given type.
static unsigned getLLVMTypeBitWidth(Type type)
Returns the bit width of LLVMType integer or vector.
static std::optional< uint64_t > getIntegerOrVectorElementWidth(Type type)
Returns the width of an integer or of the element type of an integer vector, if applicable.
#define DISPATCH(functionControl, llvmAttr)
static Type convertStructTypeWithOffset(spirv::StructType type, const TypeConverter &converter)
Converts SPIR-V struct with a regular (according to VulkanLayoutUtils) offset to LLVM struct.
static Type convertStructType(spirv::StructType type, const TypeConverter &converter)
Converts SPIR-V struct to LLVM struct.
static Value broadcast(Location loc, Value toBroadcast, unsigned numElements, const TypeConverter &typeConverter, ConversionPatternRewriter &rewriter)
Broadcasts the value to vector with numElements number of elements.
static Value createIntegerConstant(Location loc, Type srcType, Type dstType, PatternRewriter &rewriter, IntegerAttr scalarAttr)
Creates llvm.mlir.constant with a scalar or vector integer value, broadcasting scalarAttr across the ...
static std::optional< Type > convertArrayType(spirv::ArrayType type, TypeConverter &converter)
Converts SPIR-V array type to LLVM array.
OpListType::iterator iterator
OpListType & getOperations()
Operation * getTerminator()
Get the terminator operation of this block.
BlockArgListType getArguments()
IntegerAttr getIntegerAttr(Type type, int64_t value)
FloatAttr getFloatAttr(Type type, double value)
MLIRContext * getContext() const
static DenseElementsAttr get(ShapedType type, ArrayRef< Attribute > values)
Constructs a dense elements attribute from an array of element values.
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...
NamedAttrList is array of NamedAttributes that tracks whether it is sorted and does some basic work t...
NamedAttribute represents a combination of a name and an Attribute value.
This class helps build Operations.
Operation is the basic unit of execution within MLIR.
Region & getRegion(unsigned index)
Returns the region held by this operation at position 'index'.
Location getLoc()
The source location the operation was defined or derived from.
operand_range getOperands()
Returns an iterator on the underlying Value's.
A special type of RewriterBase that coordinates the application of a rewrite pattern on the current I...
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.
static LogicalResult replaceAllSymbolUses(StringAttr oldSymbol, StringAttr newSymbol, Operation *from)
Attempt to replace all uses of the given symbol 'oldSymbol' with the provided symbol 'newSymbol' that...
static Operation * lookupSymbolIn(Operation *op, StringAttr symbol)
Returns the operation registered with the given symbol name with the regions of 'symbolTableOp'.
static void setSymbolName(Operation *symbol, StringAttr name)
Sets the name of the given symbol operation.
Instances of the Type class are uniqued, have an immutable identifier and an optional mutable compone...
bool isSignedInteger() const
Return true if this is a signed integer type (with the specified width).
bool isUnsignedInteger() const
Return true if this is an unsigned integer type (with the specified width).
bool isIntOrFloat() const
Return true if this is an integer (of any signedness) or a float type.
unsigned getIntOrFloatBitWidth() const
Return the bit width of an integer or a float type, assert failure on other types.
This class provides an abstraction over the different types of ranges over Values.
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
Type getType() const
Return the type of this value.
static spirv::StructType decorateType(spirv::StructType structType)
Returns a new StructType with layout decoration.
static DenseArrayAttrImpl get(MLIRContext *context, ArrayRef< int32_t > content)
Type getElementType() const
unsigned getArrayStride() const
Returns the array stride in bytes.
unsigned getNumElements() const
StorageClass getStorageClass() const
Type getElementType() const
unsigned getArrayStride() const
Returns the array stride in bytes.
void getMemberDecorations(SmallVectorImpl< StructType::MemberDecorationInfo > &memberDecorations) const
TypeRange getElementTypes() const
bool isCompatibleType(Type type)
Returns true if the given type is compatible with the LLVM dialect.
Include the generated interface declarations.
unsigned storageClassToAddressSpace(spirv::ClientAPI clientAPI, spirv::StorageClass storageClass)
void populateSPIRVToLLVMTypeConversion(LLVMTypeConverter &typeConverter, spirv::ClientAPI clientAPIForAddressSpaceMapping=spirv::ClientAPI::Unknown)
Populates type conversions with additional SPIR-V types.
void populateSPIRVToLLVMFunctionConversionPatterns(const LLVMTypeConverter &typeConverter, RewritePatternSet &patterns)
Populates the given list with patterns for function conversion from SPIR-V to LLVM.
Type getElementTypeOrSelf(Type type)
Return the element type or return the type itself.
detail::DenseArrayAttrImpl< int32_t > DenseI32ArrayAttr
void populateSPIRVToLLVMConversionPatterns(const LLVMTypeConverter &typeConverter, RewritePatternSet &patterns, spirv::ClientAPI clientAPIForAddressSpaceMapping=spirv::ClientAPI::Unknown)
Populates the given list with patterns that convert from SPIR-V to LLVM.
void encodeBindAttribute(ModuleOp module)
Encodes global variable's descriptor set and binding into its name if they both exist.
void populateSPIRVToLLVMModuleConversionPatterns(const LLVMTypeConverter &typeConverter, RewritePatternSet &patterns)
Populates the given patterns for module conversion from SPIR-V to LLVM.
Structure used by default as a "marker" when no "Properties" are set on an Operation.