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 OpTy>
309 if (
auto properties =
310 dyn_cast_or_null<DictionaryAttr>(op->getPropertiesAsAttribute()))
311 attrs.append(properties.getValue());
317 using SPIRVToLLVMConversion<spirv::AccessChainOp>::SPIRVToLLVMConversion;
320 matchAndRewrite(spirv::AccessChainOp op, OpAdaptor adaptor,
321 ConversionPatternRewriter &rewriter)
const override {
323 getTypeConverter()->convertType(op.getComponentPtr().getType());
325 return rewriter.notifyMatchFailure(op,
"type conversion failed");
327 auto indices = llvm::to_vector<4>(adaptor.getIndices());
328 Type indexType = op.getIndices().front().getType();
329 auto llvmIndexType = getTypeConverter()->convertType(indexType);
331 return rewriter.notifyMatchFailure(op,
"type conversion failed");
333 LLVM::ConstantOp::create(rewriter, op.getLoc(), llvmIndexType,
334 rewriter.getIntegerAttr(indexType, 0));
337 auto elementType = getTypeConverter()->convertType(
338 cast<spirv::PointerType>(op.getBasePtr().getType()).getPointeeType());
340 return rewriter.notifyMatchFailure(op,
"type conversion failed");
341 rewriter.replaceOpWithNewOp<LLVM::GEPOp>(op, dstType, elementType,
342 adaptor.getBasePtr(),
indices);
349 using SPIRVToLLVMConversion<spirv::AddressOfOp>::SPIRVToLLVMConversion;
352 matchAndRewrite(spirv::AddressOfOp op, OpAdaptor adaptor,
353 ConversionPatternRewriter &rewriter)
const override {
354 auto dstType = getTypeConverter()->convertType(op.getPointer().getType());
356 return rewriter.notifyMatchFailure(op,
"type conversion failed");
357 rewriter.replaceOpWithNewOp<LLVM::AddressOfOp>(op, dstType,
363class BitFieldInsertPattern
366 using SPIRVToLLVMConversion<spirv::BitFieldInsertOp>::SPIRVToLLVMConversion;
369 matchAndRewrite(spirv::BitFieldInsertOp op, OpAdaptor adaptor,
370 ConversionPatternRewriter &rewriter)
const override {
371 auto srcType = op.getType();
372 auto dstType = getTypeConverter()->convertType(srcType);
374 return rewriter.notifyMatchFailure(op,
"type conversion failed");
375 Location loc = op.getLoc();
379 *getTypeConverter(), rewriter);
381 *getTypeConverter(), rewriter);
385 Value maskShiftedByCount =
386 LLVM::ShlOp::create(rewriter, loc, dstType, minusOne, count);
387 Value negated = LLVM::XOrOp::create(rewriter, loc, dstType,
388 maskShiftedByCount, minusOne);
389 Value maskShiftedByCountAndOffset =
390 LLVM::ShlOp::create(rewriter, loc, dstType, negated, offset);
391 Value mask = LLVM::XOrOp::create(rewriter, loc, dstType,
392 maskShiftedByCountAndOffset, minusOne);
397 LLVM::AndOp::create(rewriter, loc, dstType, op.getBase(), mask);
398 Value insertShiftedByOffset =
399 LLVM::ShlOp::create(rewriter, loc, dstType, op.getInsert(), offset);
400 rewriter.replaceOpWithNewOp<LLVM::OrOp>(op, dstType, baseAndMask,
401 insertShiftedByOffset);
407class ConstantScalarAndVectorPattern
410 using SPIRVToLLVMConversion<spirv::ConstantOp>::SPIRVToLLVMConversion;
413 matchAndRewrite(spirv::ConstantOp constOp, OpAdaptor adaptor,
414 ConversionPatternRewriter &rewriter)
const override {
415 auto srcType = constOp.getType();
416 if (!isa<VectorType>(srcType) && !srcType.isIntOrFloat())
419 auto dstType = getTypeConverter()->convertType(srcType);
421 return rewriter.notifyMatchFailure(constOp,
"type conversion failed");
430 auto signlessType = rewriter.getIntegerType(
getBitWidth(srcType));
432 if (isa<VectorType>(srcType)) {
433 auto dstElementsAttr = cast<DenseIntElementsAttr>(constOp.getValue());
434 rewriter.replaceOpWithNewOp<LLVM::ConstantOp>(
436 dstElementsAttr.mapValues(
437 signlessType, [&](
const APInt &value) {
return value; }));
440 auto srcAttr = cast<IntegerAttr>(constOp.getValue());
441 auto dstAttr = rewriter.getIntegerAttr(signlessType, srcAttr.getValue());
442 rewriter.replaceOpWithNewOp<LLVM::ConstantOp>(constOp, dstType, dstAttr);
445 rewriter.replaceOpWithNewOp<LLVM::ConstantOp>(
446 constOp, dstType, adaptor.getOperands(),
447 collectAttrsForConversion(constOp));
452class BitFieldSExtractPattern
455 using SPIRVToLLVMConversion<spirv::BitFieldSExtractOp>::SPIRVToLLVMConversion;
458 matchAndRewrite(spirv::BitFieldSExtractOp op, OpAdaptor adaptor,
459 ConversionPatternRewriter &rewriter)
const override {
460 auto srcType = op.getType();
461 auto dstType = getTypeConverter()->convertType(srcType);
463 return rewriter.notifyMatchFailure(op,
"type conversion failed");
464 Location loc = op.getLoc();
468 *getTypeConverter(), rewriter);
470 *getTypeConverter(), rewriter);
473 IntegerType integerType;
474 if (
auto vecType = dyn_cast<VectorType>(srcType))
475 integerType = cast<IntegerType>(vecType.getElementType());
477 integerType = cast<IntegerType>(srcType);
479 auto baseSize = rewriter.getIntegerAttr(integerType,
getBitWidth(srcType));
481 isa<VectorType>(srcType)
482 ? LLVM::ConstantOp::create(
483 rewriter, loc, dstType,
484 SplatElementsAttr::get(cast<ShapedType>(srcType), baseSize))
485 : LLVM::ConstantOp::create(rewriter, loc, dstType, baseSize);
489 Value countPlusOffset =
490 LLVM::AddOp::create(rewriter, loc, dstType, count, offset);
491 Value amountToShiftLeft =
492 LLVM::SubOp::create(rewriter, loc, dstType, size, countPlusOffset);
493 Value baseShiftedLeft = LLVM::ShlOp::create(
494 rewriter, loc, dstType, op.getBase(), amountToShiftLeft);
497 Value amountToShiftRight =
498 LLVM::AddOp::create(rewriter, loc, dstType, offset, amountToShiftLeft);
499 rewriter.replaceOpWithNewOp<LLVM::AShrOp>(op, dstType, baseShiftedLeft,
505class BitFieldUExtractPattern
508 using SPIRVToLLVMConversion<spirv::BitFieldUExtractOp>::SPIRVToLLVMConversion;
511 matchAndRewrite(spirv::BitFieldUExtractOp op, OpAdaptor adaptor,
512 ConversionPatternRewriter &rewriter)
const override {
513 auto srcType = op.getType();
514 auto dstType = getTypeConverter()->convertType(srcType);
516 return rewriter.notifyMatchFailure(op,
"type conversion failed");
517 Location loc = op.getLoc();
521 *getTypeConverter(), rewriter);
523 *getTypeConverter(), rewriter);
527 Value maskShiftedByCount =
528 LLVM::ShlOp::create(rewriter, loc, dstType, minusOne, count);
529 Value mask = LLVM::XOrOp::create(rewriter, loc, dstType, maskShiftedByCount,
534 LLVM::LShrOp::create(rewriter, loc, dstType, op.getBase(), offset);
535 rewriter.replaceOpWithNewOp<LLVM::AndOp>(op, dstType, shiftedBase, mask);
542 using SPIRVToLLVMConversion<spirv::BranchOp>::SPIRVToLLVMConversion;
545 matchAndRewrite(spirv::BranchOp branchOp, OpAdaptor adaptor,
546 ConversionPatternRewriter &rewriter)
const override {
547 rewriter.replaceOpWithNewOp<LLVM::BrOp>(branchOp, adaptor.getOperands(),
548 branchOp.getTarget());
553class BranchConditionalConversionPattern
556 using SPIRVToLLVMConversion<
557 spirv::BranchConditionalOp>::SPIRVToLLVMConversion;
560 matchAndRewrite(spirv::BranchConditionalOp op, OpAdaptor adaptor,
561 ConversionPatternRewriter &rewriter)
const override {
564 if (
auto weights = op.getBranchWeights()) {
565 SmallVector<int32_t> weightValues;
566 for (
auto weight : weights->getAsRange<IntegerAttr>())
567 weightValues.push_back(weight.getInt());
571 rewriter.replaceOpWithNewOp<LLVM::CondBrOp>(
572 op, op.getCondition(), op.getTrueBlockArguments(),
573 op.getFalseBlockArguments(), branchWeights, op.getTrueBlock(),
582class CompositeExtractPattern
585 using SPIRVToLLVMConversion<spirv::CompositeExtractOp>::SPIRVToLLVMConversion;
588 matchAndRewrite(spirv::CompositeExtractOp op, OpAdaptor adaptor,
589 ConversionPatternRewriter &rewriter)
const override {
590 auto dstType = this->getTypeConverter()->convertType(op.getType());
592 return rewriter.notifyMatchFailure(op,
"type conversion failed");
594 Type containerType = op.getComposite().getType();
595 if (isa<VectorType>(containerType)) {
596 Location loc = op.getLoc();
597 IntegerAttr value = cast<IntegerAttr>(op.getIndices()[0]);
599 rewriter.replaceOpWithNewOp<LLVM::ExtractElementOp>(
600 op, dstType, adaptor.getComposite(), index);
604 rewriter.replaceOpWithNewOp<LLVM::ExtractValueOp>(
605 op, adaptor.getComposite(),
606 LLVM::convertArrayToIndices(op.getIndices()));
614class CompositeInsertPattern
617 using SPIRVToLLVMConversion<spirv::CompositeInsertOp>::SPIRVToLLVMConversion;
620 matchAndRewrite(spirv::CompositeInsertOp op, OpAdaptor adaptor,
621 ConversionPatternRewriter &rewriter)
const override {
622 auto dstType = this->getTypeConverter()->convertType(op.getType());
624 return rewriter.notifyMatchFailure(op,
"type conversion failed");
626 Type containerType = op.getComposite().getType();
627 if (isa<VectorType>(containerType)) {
628 Location loc = op.getLoc();
629 IntegerAttr value = cast<IntegerAttr>(op.getIndices()[0]);
631 rewriter.replaceOpWithNewOp<LLVM::InsertElementOp>(
632 op, dstType, adaptor.getComposite(), adaptor.getObject(), index);
636 rewriter.replaceOpWithNewOp<LLVM::InsertValueOp>(
637 op, adaptor.getComposite(), adaptor.getObject(),
638 LLVM::convertArrayToIndices(op.getIndices()));
645template <
typename SPIRVOp,
typename LLVMOp>
648 using SPIRVToLLVMConversion<SPIRVOp>::SPIRVToLLVMConversion;
651 matchAndRewrite(SPIRVOp op,
typename SPIRVOp::Adaptor adaptor,
652 ConversionPatternRewriter &rewriter)
const override {
653 auto dstType = this->getTypeConverter()->convertType(op.getType());
655 return rewriter.notifyMatchFailure(op,
"type conversion failed");
656 rewriter.template replaceOpWithNewOp<LLVMOp>(
657 op, dstType, adaptor.getOperands(), collectAttrsForConversion(op));
667template <
typename SPIRVOp,
typename LLVMOp>
670 using SPIRVToLLVMConversion<SPIRVOp>::SPIRVToLLVMConversion;
673 matchAndRewrite(SPIRVOp op,
typename SPIRVOp::Adaptor adaptor,
674 ConversionPatternRewriter &rewriter)
const override {
675 Type dstType = this->getTypeConverter()->convertType(op.getType());
677 return rewriter.notifyMatchFailure(op,
"type conversion failed");
679 Location loc = op.getLoc();
680 Type operandType = adaptor.getOperand1().getType();
681 Type overflowType = rewriter.getI1Type();
682 if (
auto vecType = dyn_cast<VectorType>(operandType))
683 overflowType = VectorType::get(vecType.getShape(), overflowType);
685 Type intrType = LLVM::LLVMStructType::getLiteral(
686 rewriter.getContext(), {operandType, overflowType});
687 Value intrResult = LLVMOp::create(
688 rewriter, loc, intrType, adaptor.getOperand1(), adaptor.getOperand2());
689 Value lowBits = LLVM::ExtractValueOp::create(rewriter, loc, intrResult, 0);
690 Value overflow = LLVM::ExtractValueOp::create(rewriter, loc, intrResult, 1);
691 overflow = LLVM::ZExtOp::create(rewriter, loc, operandType, overflow);
693 Value
result = LLVM::PoisonOp::create(rewriter, loc, dstType);
694 result = LLVM::InsertValueOp::create(rewriter, loc,
result, lowBits,
695 ArrayRef<int64_t>{0});
696 result = LLVM::InsertValueOp::create(rewriter, loc,
result, overflow,
697 ArrayRef<int64_t>{1});
698 rewriter.replaceOp(op,
result);
705class ExecutionModePattern
708 using SPIRVToLLVMConversion<spirv::ExecutionModeOp>::SPIRVToLLVMConversion;
711 matchAndRewrite(spirv::ExecutionModeOp op, OpAdaptor adaptor,
712 ConversionPatternRewriter &rewriter)
const override {
716 ModuleOp module = op->getParentOfType<ModuleOp>();
717 spirv::ExecutionModeAttr executionModeAttr = op.getExecutionModeAttr();
718 std::string moduleName;
719 if (module.getName().has_value())
720 moduleName =
"_" +
module.getName()->str();
723 std::string executionModeInfoName = llvm::formatv(
724 "__spv_{0}_{1}_execution_mode_info_{2}", moduleName, op.getFn().str(),
725 static_cast<uint32_t
>(executionModeAttr.getValue()));
727 MLIRContext *context = rewriter.getContext();
728 OpBuilder::InsertionGuard guard(rewriter);
729 rewriter.setInsertionPointToStart(module.getBody());
736 auto llvmI32Type = IntegerType::get(context, 32);
737 SmallVector<Type, 2> fields;
738 fields.push_back(llvmI32Type);
740 if (!values.empty()) {
741 auto arrayType = LLVM::LLVMArrayType::get(llvmI32Type, values.size());
742 fields.push_back(arrayType);
744 auto structType = LLVM::LLVMStructType::getLiteral(context, fields);
747 auto global = LLVM::GlobalOp::create(
748 rewriter, UnknownLoc::get(context), structType,
true,
749 LLVM::Linkage::External, executionModeInfoName, Attribute(),
751 Location loc = global.getLoc();
752 Region ®ion = global.getInitializerRegion();
753 Block *block = rewriter.createBlock(®ion);
756 rewriter.setInsertionPointToStart(block);
757 Value structValue = LLVM::PoisonOp::create(rewriter, loc, structType);
758 Value executionMode = LLVM::ConstantOp::create(
759 rewriter, loc, llvmI32Type,
760 rewriter.getI32IntegerAttr(
761 static_cast<uint32_t
>(executionModeAttr.getValue())));
762 SmallVector<int64_t> position{0};
763 structValue = LLVM::InsertValueOp::create(rewriter, loc, structValue,
764 executionMode, position);
767 for (
unsigned i = 0, e = values.size(); i < e; ++i) {
768 auto attr = values.getValue()[i];
769 Value entry = LLVM::ConstantOp::create(rewriter, loc, llvmI32Type, attr);
770 structValue = LLVM::InsertValueOp::create(
771 rewriter, loc, structValue, entry, ArrayRef<int64_t>({1, i}));
773 LLVM::ReturnOp::create(rewriter, loc, ArrayRef<Value>({structValue}));
774 rewriter.eraseOp(op);
783class GlobalVariablePattern
786 template <
typename... Args>
787 GlobalVariablePattern(spirv::ClientAPI clientAPI, Args &&...args)
788 : SPIRVToLLVMConversion<spirv::GlobalVariableOp>(
789 std::forward<Args>(args)...),
790 clientAPI(clientAPI) {}
793 matchAndRewrite(spirv::GlobalVariableOp op, OpAdaptor adaptor,
794 ConversionPatternRewriter &rewriter)
const override {
797 if (op.getInitializer())
800 auto srcType = cast<spirv::PointerType>(op.getType());
801 auto dstType = getTypeConverter()->convertType(srcType.getPointeeType());
803 return rewriter.notifyMatchFailure(op,
"type conversion failed");
808 auto storageClass = srcType.getStorageClass();
809 switch (storageClass) {
810 case spirv::StorageClass::Input:
811 case spirv::StorageClass::Private:
812 case spirv::StorageClass::Output:
813 case spirv::StorageClass::StorageBuffer:
814 case spirv::StorageClass::UniformConstant:
823 bool isConstant = (storageClass == spirv::StorageClass::Input) ||
824 (storageClass == spirv::StorageClass::UniformConstant);
830 auto linkage = storageClass == spirv::StorageClass::Private
831 ? LLVM::Linkage::Private
832 : LLVM::Linkage::External;
833 StringAttr locationAttrName = op.getLocationAttrName();
834 IntegerAttr locationAttr = op.getLocationAttr();
835 auto newGlobalOp = rewriter.replaceOpWithNewOp<LLVM::GlobalOp>(
836 op, dstType, isConstant, linkage, op.getSymName(), Attribute(),
841 newGlobalOp->setDiscardableAttr(locationAttrName, locationAttr);
847 spirv::ClientAPI clientAPI;
852template <
typename SPIRVOp,
typename LLVMExtOp,
typename LLVMTruncOp>
855 using SPIRVToLLVMConversion<SPIRVOp>::SPIRVToLLVMConversion;
858 matchAndRewrite(SPIRVOp op,
typename SPIRVOp::Adaptor adaptor,
859 ConversionPatternRewriter &rewriter)
const override {
861 Type fromType = op.getOperand().getType();
862 Type toType = op.getType();
864 auto dstType = this->getTypeConverter()->convertType(toType);
866 return rewriter.notifyMatchFailure(op,
"type conversion failed");
869 rewriter.template replaceOpWithNewOp<LLVMExtOp>(op, dstType,
870 adaptor.getOperands());
874 rewriter.template replaceOpWithNewOp<LLVMTruncOp>(op, dstType,
875 adaptor.getOperands());
882class FunctionCallPattern
885 using SPIRVToLLVMConversion<spirv::FunctionCallOp>::SPIRVToLLVMConversion;
888 matchAndRewrite(spirv::FunctionCallOp callOp, OpAdaptor adaptor,
889 ConversionPatternRewriter &rewriter)
const override {
890 if (callOp.getNumResults() == 0) {
891 auto newOp = rewriter.replaceOpWithNewOp<LLVM::CallOp>(
892 callOp,
TypeRange(), adaptor.getOperands(),
893 collectAttrsForConversion(callOp));
894 newOp.getProperties().operandSegmentSizes = {
895 static_cast<int32_t
>(adaptor.getOperands().size()), 0};
896 newOp.getProperties().op_bundle_sizes = rewriter.getDenseI32ArrayAttr({});
901 auto dstType = getTypeConverter()->convertType(callOp.getType(0));
903 return rewriter.notifyMatchFailure(callOp,
"type conversion failed");
904 auto newOp = rewriter.replaceOpWithNewOp<LLVM::CallOp>(
905 callOp, dstType, adaptor.getOperands(),
906 collectAttrsForConversion(callOp));
907 newOp.getProperties().operandSegmentSizes = {
908 static_cast<int32_t
>(adaptor.getOperands().size()), 0};
909 newOp.getProperties().op_bundle_sizes = rewriter.getDenseI32ArrayAttr({});
915template <
typename SPIRVOp, LLVM::FCmpPredicate predicate>
918 using SPIRVToLLVMConversion<SPIRVOp>::SPIRVToLLVMConversion;
921 matchAndRewrite(SPIRVOp op,
typename SPIRVOp::Adaptor adaptor,
922 ConversionPatternRewriter &rewriter)
const override {
924 auto dstType = this->getTypeConverter()->convertType(op.getType());
926 return rewriter.notifyMatchFailure(op,
"type conversion failed");
928 rewriter.template replaceOpWithNewOp<LLVM::FCmpOp>(
929 op, dstType, predicate, op.getOperand1(), op.getOperand2());
935template <
typename SPIRVOp, LLVM::ICmpPredicate predicate>
938 using SPIRVToLLVMConversion<SPIRVOp>::SPIRVToLLVMConversion;
941 matchAndRewrite(SPIRVOp op,
typename SPIRVOp::Adaptor adaptor,
942 ConversionPatternRewriter &rewriter)
const override {
944 auto dstType = this->getTypeConverter()->convertType(op.getType());
946 return rewriter.notifyMatchFailure(op,
"type conversion failed");
948 rewriter.template replaceOpWithNewOp<LLVM::ICmpOp>(
949 op, dstType, predicate, op.getOperand1(), op.getOperand2());
954class InverseSqrtPattern
957 using SPIRVToLLVMConversion<spirv::GLInverseSqrtOp>::SPIRVToLLVMConversion;
960 matchAndRewrite(spirv::GLInverseSqrtOp op, OpAdaptor adaptor,
961 ConversionPatternRewriter &rewriter)
const override {
962 auto srcType = op.getType();
963 auto dstType = getTypeConverter()->convertType(srcType);
965 return rewriter.notifyMatchFailure(op,
"type conversion failed");
967 Location loc = op.getLoc();
969 Value sqrt = LLVM::SqrtOp::create(rewriter, loc, dstType, op.getOperand());
970 rewriter.replaceOpWithNewOp<LLVM::FDivOp>(op, dstType, one, sqrt);
977class VectorTimesScalarPattern
980 using SPIRVToLLVMConversion<
981 spirv::VectorTimesScalarOp>::SPIRVToLLVMConversion;
984 matchAndRewrite(spirv::VectorTimesScalarOp op, OpAdaptor adaptor,
985 ConversionPatternRewriter &rewriter)
const override {
986 Type srcType = op.getType();
987 Type dstType = getTypeConverter()->convertType(srcType);
989 return rewriter.notifyMatchFailure(op,
"type conversion failed");
991 unsigned numElements = op.getVector().getType().getNumElements();
992 Value broadcasted =
broadcast(op.getLoc(), adaptor.getScalar(), numElements,
993 *getTypeConverter(), rewriter);
994 rewriter.replaceOpWithNewOp<LLVM::FMulOp>(op, dstType, adaptor.getVector(),
1003 using SPIRVToLLVMConversion<spirv::SNegateOp>::SPIRVToLLVMConversion;
1006 matchAndRewrite(spirv::SNegateOp op, OpAdaptor adaptor,
1007 ConversionPatternRewriter &rewriter)
const override {
1008 Type srcType = op.getType();
1009 Type dstType = getTypeConverter()->convertType(srcType);
1011 return rewriter.notifyMatchFailure(op,
"type conversion failed");
1013 Location loc = op.getLoc();
1014 IntegerAttr zeroAttr = rewriter.getIntegerAttr(
1018 rewriter.replaceOpWithNewOp<LLVM::SubOp>(op, dstType, zero,
1019 adaptor.getOperand());
1026template <
typename SPIRVOp,
typename LLVMMinOp,
typename LLVMMaxOp>
1029 using SPIRVToLLVMConversion<SPIRVOp>::SPIRVToLLVMConversion;
1032 matchAndRewrite(SPIRVOp op,
typename SPIRVOp::Adaptor adaptor,
1033 ConversionPatternRewriter &rewriter)
const override {
1034 Type dstType = this->getTypeConverter()->convertType(op.getType());
1036 return rewriter.notifyMatchFailure(op,
"type conversion failed");
1038 Location loc = op.getLoc();
1039 Value
max = LLVMMaxOp::create(rewriter, loc, dstType, adaptor.getX(),
1041 rewriter.template replaceOpWithNewOp<LLVMMinOp>(op, dstType,
max,
1052 using SPIRVToLLVMConversion<spirv::FModOp>::SPIRVToLLVMConversion;
1055 matchAndRewrite(spirv::FModOp op, OpAdaptor adaptor,
1056 ConversionPatternRewriter &rewriter)
const override {
1057 Type dstType = getTypeConverter()->convertType(op.getType());
1059 return rewriter.notifyMatchFailure(op,
"type conversion failed");
1061 Location loc = op.getLoc();
1062 Value
lhs = adaptor.getOperand1();
1063 Value
rhs = adaptor.getOperand2();
1064 Value
div = LLVM::FDivOp::create(rewriter, loc, dstType,
lhs,
rhs);
1065 Value floored = LLVM::FFloorOp::create(rewriter, loc, dstType,
div);
1066 Value scaled = LLVM::FMulOp::create(rewriter, loc, dstType,
rhs, floored);
1067 rewriter.replaceOpWithNewOp<LLVM::FSubOp>(op, dstType,
lhs, scaled);
1078 using SPIRVToLLVMConversion<spirv::SModOp>::SPIRVToLLVMConversion;
1081 matchAndRewrite(spirv::SModOp op, OpAdaptor adaptor,
1082 ConversionPatternRewriter &rewriter)
const override {
1083 Type srcType = op.getType();
1084 Type dstType = getTypeConverter()->convertType(srcType);
1086 return rewriter.notifyMatchFailure(op,
"type conversion failed");
1088 Location loc = op.getLoc();
1089 Value
lhs = adaptor.getOperand1();
1090 Value
rhs = adaptor.getOperand2();
1091 Type i1Type = rewriter.getI1Type();
1092 auto vecSrcType = dyn_cast<VectorType>(srcType);
1094 vecSrcType ? VectorType::get(vecSrcType.getShape(), i1Type) : i1Type;
1096 Value
rem = LLVM::SRemOp::create(rewriter, loc, dstType,
lhs,
rhs);
1097 IntegerAttr zeroAttr = rewriter.getIntegerAttr(
1102 Value remNonZero = LLVM::ICmpOp::create(rewriter, loc, cmpType,
1103 LLVM::ICmpPredicate::ne,
rem, zero);
1104 Value remNeg = LLVM::ICmpOp::create(rewriter, loc, cmpType,
1105 LLVM::ICmpPredicate::slt,
rem, zero);
1106 Value rhsNeg = LLVM::ICmpOp::create(rewriter, loc, cmpType,
1107 LLVM::ICmpPredicate::slt,
rhs, zero);
1108 Value signMismatch =
1109 LLVM::XOrOp::create(rewriter, loc, cmpType, remNeg, rhsNeg);
1111 LLVM::AndOp::create(rewriter, loc, cmpType, remNonZero, signMismatch);
1113 Value adjusted = LLVM::AddOp::create(rewriter, loc, dstType,
rem,
rhs);
1114 rewriter.replaceOpWithNewOp<LLVM::SelectOp>(op, dstType, needsAdjust,
1121template <
typename SPIRVOp>
1124 using SPIRVToLLVMConversion<SPIRVOp>::SPIRVToLLVMConversion;
1127 matchAndRewrite(SPIRVOp op,
typename SPIRVOp::Adaptor adaptor,
1128 ConversionPatternRewriter &rewriter)
const override {
1129 if (!op.getMemoryAccess()) {
1131 *this->getTypeConverter(), 0,
1135 auto memoryAccess = *op.getMemoryAccess();
1136 switch (memoryAccess) {
1137 case spirv::MemoryAccess::Aligned:
1138 case spirv::MemoryAccess::None:
1139 case spirv::MemoryAccess::Nontemporal:
1140 case spirv::MemoryAccess::Volatile: {
1141 unsigned alignment =
1142 memoryAccess == spirv::MemoryAccess::Aligned ? *op.getAlignment() : 0;
1143 bool isNonTemporal = memoryAccess == spirv::MemoryAccess::Nontemporal;
1144 bool isVolatile = memoryAccess == spirv::MemoryAccess::Volatile;
1146 *this->getTypeConverter(), alignment,
1147 isVolatile, isNonTemporal);
1157template <
typename SPIRVOp>
1160 using SPIRVToLLVMConversion<SPIRVOp>::SPIRVToLLVMConversion;
1163 matchAndRewrite(SPIRVOp notOp,
typename SPIRVOp::Adaptor adaptor,
1164 ConversionPatternRewriter &rewriter)
const override {
1165 auto srcType = notOp.getType();
1166 auto dstType = this->getTypeConverter()->convertType(srcType);
1168 return rewriter.notifyMatchFailure(notOp,
"type conversion failed");
1170 Location loc = notOp.getLoc();
1172 rewriter.template replaceOpWithNewOp<LLVM::XOrOp>(notOp, dstType,
1173 notOp.getOperand(), mask);
1179template <
typename SPIRVOp>
1182 using SPIRVToLLVMConversion<SPIRVOp>::SPIRVToLLVMConversion;
1185 matchAndRewrite(SPIRVOp op,
typename SPIRVOp::Adaptor adaptor,
1186 ConversionPatternRewriter &rewriter)
const override {
1187 rewriter.eraseOp(op);
1194 using SPIRVToLLVMConversion<spirv::ReturnOp>::SPIRVToLLVMConversion;
1197 matchAndRewrite(spirv::ReturnOp returnOp, OpAdaptor adaptor,
1198 ConversionPatternRewriter &rewriter)
const override {
1199 rewriter.replaceOpWithNewOp<LLVM::ReturnOp>(returnOp, ArrayRef<Type>(),
1207 using SPIRVToLLVMConversion<spirv::ReturnValueOp>::SPIRVToLLVMConversion;
1210 matchAndRewrite(spirv::ReturnValueOp returnValueOp, OpAdaptor adaptor,
1211 ConversionPatternRewriter &rewriter)
const override {
1212 rewriter.replaceOpWithNewOp<LLVM::ReturnOp>(returnValueOp, ArrayRef<Type>(),
1213 adaptor.getOperands());
1220 using SPIRVToLLVMConversion<spirv::UnreachableOp>::SPIRVToLLVMConversion;
1223 matchAndRewrite(spirv::UnreachableOp unreachableOp, OpAdaptor adaptor,
1224 ConversionPatternRewriter &rewriter)
const override {
1225 rewriter.replaceOpWithNewOp<LLVM::UnreachableOp>(unreachableOp);
1234 bool convergent =
true) {
1235 auto func = dyn_cast_or_null<LLVM::LLVMFuncOp>(
1241 func = LLVM::LLVMFuncOp::create(
1242 b, symbolTable->
getLoc(), name,
1243 LLVM::LLVMFunctionType::get(resultType, paramTypes));
1244 func.setCConv(LLVM::cconv::CConv::SPIR_FUNC);
1245 func.setConvergent(convergent);
1246 func.setNoUnwind(
true);
1247 func.setWillReturn(
true);
1252 LLVM::LLVMFuncOp
func,
1254 auto call = LLVM::CallOp::create(builder, loc,
func, args);
1255 call.setCConv(
func.getCConv());
1256 call.setConvergentAttr(
func.getConvergentAttr());
1257 call.setNoUnwindAttr(
func.getNoUnwindAttr());
1258 call.setWillReturnAttr(
func.getWillReturnAttr());
1262template <
typename BarrierOpTy>
1265 using OpAdaptor =
typename SPIRVToLLVMConversion<BarrierOpTy>::OpAdaptor;
1267 using SPIRVToLLVMConversion<BarrierOpTy>::SPIRVToLLVMConversion;
1269 static constexpr StringRef getFuncName();
1272 matchAndRewrite(BarrierOpTy controlBarrierOp, OpAdaptor adaptor,
1273 ConversionPatternRewriter &rewriter)
const override {
1274 constexpr StringRef funcName = getFuncName();
1275 Operation *symbolTable =
1276 controlBarrierOp->template getParentWithTrait<OpTrait::SymbolTable>();
1278 Type i32 = rewriter.getI32Type();
1280 Type voidTy = rewriter.getType<LLVM::LLVMVoidType>();
1281 LLVM::LLVMFuncOp func =
1284 Location loc = controlBarrierOp->getLoc();
1285 Value execution = LLVM::ConstantOp::create(
1286 rewriter, loc, i32,
static_cast<int32_t
>(adaptor.getExecutionScope()));
1287 Value memory = LLVM::ConstantOp::create(
1288 rewriter, loc, i32,
static_cast<int32_t
>(adaptor.getMemoryScope()));
1289 Value semantics = LLVM::ConstantOp::create(
1290 rewriter, loc, i32,
static_cast<int32_t
>(adaptor.getMemorySemantics()));
1293 {execution, memory, semantics});
1295 rewriter.replaceOp(controlBarrierOp, call);
1302StringRef getTypeMangling(
Type type,
bool isSigned) {
1304 .Case([](Float16Type) {
return "Dh"; })
1305 .Case([](Float32Type) {
return "f"; })
1306 .Case([](Float64Type) {
return "d"; })
1307 .Case([isSigned](IntegerType intTy) {
1308 switch (intTy.getWidth()) {
1312 return (isSigned) ?
"a" :
"c";
1314 return (isSigned) ?
"s" :
"t";
1316 return (isSigned) ?
"i" :
"j";
1318 return (isSigned) ?
"l" :
"m";
1320 llvm_unreachable(
"Unsupported integer width");
1323 .DefaultUnreachable(
"No mangling defined");
1326template <
typename ReduceOp>
1327constexpr StringLiteral getGroupFuncName();
1330constexpr StringLiteral getGroupFuncName<spirv::GroupIAddOp>() {
1331 return "_Z17__spirv_GroupIAddii";
1334constexpr StringLiteral getGroupFuncName<spirv::GroupFAddOp>() {
1335 return "_Z17__spirv_GroupFAddii";
1338constexpr StringLiteral getGroupFuncName<spirv::GroupSMinOp>() {
1339 return "_Z17__spirv_GroupSMinii";
1342constexpr StringLiteral getGroupFuncName<spirv::GroupUMinOp>() {
1343 return "_Z17__spirv_GroupUMinii";
1346constexpr StringLiteral getGroupFuncName<spirv::GroupFMinOp>() {
1347 return "_Z17__spirv_GroupFMinii";
1350constexpr StringLiteral getGroupFuncName<spirv::GroupSMaxOp>() {
1351 return "_Z17__spirv_GroupSMaxii";
1354constexpr StringLiteral getGroupFuncName<spirv::GroupUMaxOp>() {
1355 return "_Z17__spirv_GroupUMaxii";
1358constexpr StringLiteral getGroupFuncName<spirv::GroupFMaxOp>() {
1359 return "_Z17__spirv_GroupFMaxii";
1362constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformIAddOp>() {
1363 return "_Z27__spirv_GroupNonUniformIAddii";
1366constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformFAddOp>() {
1367 return "_Z27__spirv_GroupNonUniformFAddii";
1370constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformIMulOp>() {
1371 return "_Z27__spirv_GroupNonUniformIMulii";
1374constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformFMulOp>() {
1375 return "_Z27__spirv_GroupNonUniformFMulii";
1378constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformSMinOp>() {
1379 return "_Z27__spirv_GroupNonUniformSMinii";
1382constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformUMinOp>() {
1383 return "_Z27__spirv_GroupNonUniformUMinii";
1386constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformFMinOp>() {
1387 return "_Z27__spirv_GroupNonUniformFMinii";
1390constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformSMaxOp>() {
1391 return "_Z27__spirv_GroupNonUniformSMaxii";
1394constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformUMaxOp>() {
1395 return "_Z27__spirv_GroupNonUniformUMaxii";
1398constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformFMaxOp>() {
1399 return "_Z27__spirv_GroupNonUniformFMaxii";
1402constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformBitwiseAndOp>() {
1403 return "_Z33__spirv_GroupNonUniformBitwiseAndii";
1406constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformBitwiseOrOp>() {
1407 return "_Z32__spirv_GroupNonUniformBitwiseOrii";
1410constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformBitwiseXorOp>() {
1411 return "_Z33__spirv_GroupNonUniformBitwiseXorii";
1414constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformLogicalAndOp>() {
1415 return "_Z33__spirv_GroupNonUniformLogicalAndii";
1418constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformLogicalOrOp>() {
1419 return "_Z32__spirv_GroupNonUniformLogicalOrii";
1422constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformLogicalXorOp>() {
1423 return "_Z33__spirv_GroupNonUniformLogicalXorii";
1427template <
typename ReduceOp,
bool Signed = false,
bool NonUniform = false>
1430 using SPIRVToLLVMConversion<ReduceOp>::SPIRVToLLVMConversion;
1433 matchAndRewrite(ReduceOp op,
typename ReduceOp::Adaptor adaptor,
1434 ConversionPatternRewriter &rewriter)
const override {
1436 Type retTy = op.getResult().getType();
1440 SmallString<36> funcName = getGroupFuncName<ReduceOp>();
1441 funcName += getTypeMangling(retTy,
false);
1443 Type i32Ty = rewriter.getI32Type();
1444 SmallVector<Type> paramTypes{i32Ty, i32Ty, retTy};
1445 if constexpr (NonUniform) {
1446 if (adaptor.getClusterSize()) {
1448 paramTypes.push_back(i32Ty);
1452 Operation *symbolTable =
1453 op->template getParentWithTrait<OpTrait::SymbolTable>();
1455 LLVM::LLVMFuncOp func =
1458 Location loc = op.getLoc();
1459 Value scope = LLVM::ConstantOp::create(
1460 rewriter, loc, i32Ty,
1461 static_cast<int32_t
>(adaptor.getExecutionScope()));
1462 Value groupOp = LLVM::ConstantOp::create(
1463 rewriter, loc, i32Ty,
1464 static_cast<int32_t
>(adaptor.getGroupOperation()));
1465 SmallVector<Value> operands{scope, groupOp};
1466 operands.append(adaptor.getOperands().begin(), adaptor.getOperands().end());
1469 rewriter.replaceOp(op, call);
1476ControlBarrierPattern<spirv::ControlBarrierOp>::getFuncName() {
1477 return "_Z22__spirv_ControlBarrieriii";
1482ControlBarrierPattern<spirv::INTELControlBarrierArriveOp>::getFuncName() {
1483 return "_Z33__spirv_ControlBarrierArriveINTELiii";
1488ControlBarrierPattern<spirv::INTELControlBarrierWaitOp>::getFuncName() {
1489 return "_Z31__spirv_ControlBarrierWaitINTELiii";
1542 using SPIRVToLLVMConversion<spirv::LoopOp>::SPIRVToLLVMConversion;
1545 matchAndRewrite(spirv::LoopOp loopOp, OpAdaptor adaptor,
1546 ConversionPatternRewriter &rewriter)
const override {
1548 if (loopOp.getLoopControl() != spirv::LoopControl::None)
1552 if (loopOp.getBody().empty()) {
1553 rewriter.eraseOp(loopOp);
1557 Location loc = loopOp.getLoc();
1561 Block *currentBlock = rewriter.getBlock();
1563 Block *endBlock = rewriter.splitBlock(currentBlock, position);
1567 Block *entryBlock = loopOp.getEntryBlock();
1569 auto brOp = dyn_cast<spirv::BranchOp>(entryBlock->
getOperations().front());
1572 Block *headerBlock = loopOp.getHeaderBlock();
1573 rewriter.setInsertionPointToEnd(currentBlock);
1574 LLVM::BrOp::create(rewriter, loc, brOp.getBlockArguments(), headerBlock);
1575 rewriter.eraseBlock(entryBlock);
1578 Block *mergeBlock = loopOp.getMergeBlock();
1581 rewriter.setInsertionPointToEnd(mergeBlock);
1582 LLVM::BrOp::create(rewriter, loc, terminatorOperands, endBlock);
1584 rewriter.inlineRegionBefore(loopOp.getBody(), endBlock);
1595 using SPIRVToLLVMConversion<spirv::SelectionOp>::SPIRVToLLVMConversion;
1598 matchAndRewrite(spirv::SelectionOp op, OpAdaptor adaptor,
1599 ConversionPatternRewriter &rewriter)
const override {
1603 if (op.getSelectionControl() != spirv::SelectionControl::None)
1610 if (op.getBody().getBlocks().size() <= 2) {
1611 rewriter.eraseOp(op);
1615 Location loc = op.getLoc();
1619 auto *currentBlock = rewriter.getInsertionBlock();
1620 rewriter.setInsertionPointAfter(op);
1621 auto position = rewriter.getInsertionPoint();
1622 auto *continueBlock = rewriter.splitBlock(currentBlock, position);
1625 for (
auto ty : op.getResultTypes()) {
1626 Type dstTy = getTypeConverter()->convertType(ty);
1628 return rewriter.notifyMatchFailure(op,
"failed to convert type");
1629 continueBlock->addArgument(dstTy, loc);
1636 auto *headerBlock = op.getHeaderBlock();
1638 auto condBrOp = dyn_cast<spirv::BranchConditionalOp>(
1644 auto *mergeBlock = op.getMergeBlock();
1647 rewriter.setInsertionPointToEnd(mergeBlock);
1648 LLVM::BrOp::create(rewriter, loc, terminatorOperands, continueBlock);
1651 Block *trueBlock = condBrOp.getTrueBlock();
1652 Block *falseBlock = condBrOp.getFalseBlock();
1653 rewriter.setInsertionPointToEnd(currentBlock);
1654 LLVM::CondBrOp::create(rewriter, loc, condBrOp.getCondition(), trueBlock,
1655 condBrOp.getTrueTargetOperands(), falseBlock,
1656 condBrOp.getFalseTargetOperands());
1658 rewriter.eraseBlock(headerBlock);
1659 rewriter.inlineRegionBefore(op.getBody(), continueBlock);
1660 rewriter.replaceOp(op, continueBlock->getArguments());
1669template <
typename SPIRVOp,
typename LLVMOp>
1672 using SPIRVToLLVMConversion<SPIRVOp>::SPIRVToLLVMConversion;
1675 matchAndRewrite(SPIRVOp op,
typename SPIRVOp::Adaptor adaptor,
1676 ConversionPatternRewriter &rewriter)
const override {
1678 auto dstType = this->getTypeConverter()->convertType(op.getType());
1680 return rewriter.notifyMatchFailure(op,
"type conversion failed");
1682 Type op1Type = op.getOperand1().getType();
1683 Type op2Type = op.getOperand2().getType();
1685 if (op1Type == op2Type) {
1686 rewriter.template replaceOpWithNewOp<LLVMOp>(op, dstType,
1687 adaptor.getOperands());
1691 std::optional<uint64_t> dstTypeWidth =
1693 std::optional<uint64_t> op2TypeWidth =
1696 if (!dstTypeWidth || !op2TypeWidth)
1699 Location loc = op.getLoc();
1701 if (op2TypeWidth < dstTypeWidth) {
1704 LLVM::ZExtOp::create(rewriter, loc, dstType, adaptor.getOperand2());
1707 LLVM::SExtOp::create(rewriter, loc, dstType, adaptor.getOperand2());
1709 }
else if (op2TypeWidth == dstTypeWidth) {
1710 extended = adaptor.getOperand2();
1716 LLVMOp::create(rewriter, loc, dstType, adaptor.getOperand1(), extended);
1717 rewriter.replaceOp(op,
result);
1727 using SPIRVToLLVMConversion<spirv::GLSAbsOp>::SPIRVToLLVMConversion;
1730 matchAndRewrite(spirv::GLSAbsOp op, OpAdaptor adaptor,
1731 ConversionPatternRewriter &rewriter)
const override {
1732 Type dstType = getTypeConverter()->convertType(op.getType());
1734 return rewriter.notifyMatchFailure(op,
"type conversion failed");
1736 rewriter.replaceOpWithNewOp<LLVM::AbsOp>(op, dstType, adaptor.getOperand(),
1745 using SPIRVToLLVMConversion<spirv::GLFractOp>::SPIRVToLLVMConversion;
1748 matchAndRewrite(spirv::GLFractOp op, OpAdaptor adaptor,
1749 ConversionPatternRewriter &rewriter)
const override {
1750 Type dstType = getTypeConverter()->convertType(op.getType());
1752 return rewriter.notifyMatchFailure(op,
"type conversion failed");
1754 Location loc = op.getLoc();
1755 Value operand = adaptor.getOperand();
1756 Value floored = LLVM::FFloorOp::create(rewriter, loc, dstType, operand);
1757 rewriter.replaceOpWithNewOp<LLVM::FSubOp>(op, dstType, operand, floored);
1766 using SPIRVToLLVMConversion<spirv::GLFMixOp>::SPIRVToLLVMConversion;
1769 matchAndRewrite(spirv::GLFMixOp op, OpAdaptor adaptor,
1770 ConversionPatternRewriter &rewriter)
const override {
1771 Type dstType = getTypeConverter()->convertType(op.getType());
1773 return rewriter.notifyMatchFailure(op,
"type conversion failed");
1775 Location loc = op.getLoc();
1776 Value x = adaptor.getX();
1777 Value y = adaptor.getY();
1778 Value a = adaptor.getA();
1780 Value oneMinusA = LLVM::FSubOp::create(rewriter, loc, dstType, one, a);
1781 Value
lhs = LLVM::FMulOp::create(rewriter, loc, dstType, x, oneMinusA);
1782 Value
rhs = LLVM::FMulOp::create(rewriter, loc, dstType, y, a);
1783 rewriter.replaceOpWithNewOp<LLVM::FAddOp>(op, dstType,
lhs,
rhs);
1792 using SPIRVToLLVMConversion<spirv::CLMixOp>::SPIRVToLLVMConversion;
1795 matchAndRewrite(spirv::CLMixOp op, OpAdaptor adaptor,
1796 ConversionPatternRewriter &rewriter)
const override {
1797 Type dstType = getTypeConverter()->convertType(op.getType());
1799 return rewriter.notifyMatchFailure(op,
"type conversion failed");
1801 Location loc = op.getLoc();
1802 Value x = adaptor.getX();
1803 Value y = adaptor.getY();
1804 Value a = adaptor.getZ();
1805 Value diff = LLVM::FSubOp::create(rewriter, loc, dstType, y, x);
1806 rewriter.replaceOpWithNewOp<LLVM::FMAOp>(op, dstType, a, diff, x);
1813template <
typename SPIRVOp>
1816 template <
typename... Args>
1817 ScalePattern(
double scale, Args &&...args)
1818 : SPIRVToLLVMConversion<SPIRVOp>(std::forward<Args>(args)...),
1822 matchAndRewrite(SPIRVOp op,
typename SPIRVOp::Adaptor adaptor,
1823 ConversionPatternRewriter &rewriter)
const override {
1824 Type srcType = op.getType();
1825 Type dstType = this->getTypeConverter()->convertType(srcType);
1827 return rewriter.notifyMatchFailure(op,
"type conversion failed");
1829 Location loc = op.getLoc();
1831 rewriter.replaceOpWithNewOp<LLVM::FMulOp>(op, dstType, adaptor.getOperand(),
1843template <
typename SPIRVOp,
bool isFloat>
1846 using SPIRVToLLVMConversion<SPIRVOp>::SPIRVToLLVMConversion;
1849 matchAndRewrite(SPIRVOp op,
typename SPIRVOp::Adaptor adaptor,
1850 ConversionPatternRewriter &rewriter)
const override {
1851 Type srcType = op.getType();
1852 Type dstType = this->getTypeConverter()->convertType(srcType);
1854 return rewriter.notifyMatchFailure(op,
"type conversion failed");
1856 Location loc = op.getLoc();
1857 Value operand = adaptor.getOperand();
1858 auto vecSrcType = dyn_cast<VectorType>(srcType);
1859 Type i1Type = rewriter.getI1Type();
1861 vecSrcType ? VectorType::get(vecSrcType.getShape(), i1Type) : i1Type;
1863 Value zero, one, minusOne, gt, lt;
1864 if constexpr (isFloat) {
1868 gt = LLVM::FCmpOp::create(rewriter, loc, cmpType,
1869 LLVM::FCmpPredicate::ogt, operand, zero);
1870 lt = LLVM::FCmpOp::create(rewriter, loc, cmpType,
1871 LLVM::FCmpPredicate::olt, operand, zero);
1875 rewriter.getIntegerAttr(intElemType, 0));
1877 rewriter.getIntegerAttr(intElemType, 1));
1879 gt = LLVM::ICmpOp::create(rewriter, loc, cmpType,
1880 LLVM::ICmpPredicate::sgt, operand, zero);
1881 lt = LLVM::ICmpOp::create(rewriter, loc, cmpType,
1882 LLVM::ICmpPredicate::slt, operand, zero);
1886 LLVM::SelectOp::create(rewriter, loc, dstType, lt, minusOne, zero);
1887 rewriter.replaceOpWithNewOp<LLVM::SelectOp>(op, dstType, gt, one,
1895 using SPIRVToLLVMConversion<spirv::VariableOp>::SPIRVToLLVMConversion;
1898 matchAndRewrite(spirv::VariableOp varOp, OpAdaptor adaptor,
1899 ConversionPatternRewriter &rewriter)
const override {
1900 auto srcType = varOp.getType();
1902 auto pointerTo = cast<spirv::PointerType>(srcType).getPointeeType();
1903 auto init = varOp.getInitializer();
1904 if (init && !pointerTo.isIntOrFloat() && !isa<VectorType>(pointerTo))
1907 auto dstType = getTypeConverter()->convertType(srcType);
1909 return rewriter.notifyMatchFailure(varOp,
"type conversion failed");
1911 Location loc = varOp.getLoc();
1914 auto elementType = getTypeConverter()->convertType(pointerTo);
1916 return rewriter.notifyMatchFailure(varOp,
"type conversion failed");
1917 rewriter.replaceOpWithNewOp<LLVM::AllocaOp>(varOp, dstType, elementType,
1921 auto elementType = getTypeConverter()->convertType(pointerTo);
1923 return rewriter.notifyMatchFailure(varOp,
"type conversion failed");
1925 LLVM::AllocaOp::create(rewriter, loc, dstType, elementType, size);
1926 LLVM::StoreOp::create(rewriter, loc, adaptor.getInitializer(), allocated);
1927 rewriter.replaceOp(varOp, allocated);
1936class BitcastConversionPattern
1939 using SPIRVToLLVMConversion<spirv::BitcastOp>::SPIRVToLLVMConversion;
1942 matchAndRewrite(spirv::BitcastOp bitcastOp, OpAdaptor adaptor,
1943 ConversionPatternRewriter &rewriter)
const override {
1944 auto dstType = getTypeConverter()->convertType(bitcastOp.getType());
1946 return rewriter.notifyMatchFailure(bitcastOp,
"type conversion failed");
1949 if (isa<LLVM::LLVMPointerType>(dstType)) {
1950 rewriter.replaceOp(bitcastOp, adaptor.getOperand());
1954 rewriter.replaceOpWithNewOp<LLVM::BitcastOp>(
1955 bitcastOp, dstType, adaptor.getOperands(),
1956 collectAttrsForConversion(bitcastOp));
1967 using SPIRVToLLVMConversion<spirv::FuncOp>::SPIRVToLLVMConversion;
1970 matchAndRewrite(spirv::FuncOp funcOp, OpAdaptor adaptor,
1971 ConversionPatternRewriter &rewriter)
const override {
1975 auto funcType = funcOp.getFunctionType();
1976 TypeConverter::SignatureConversion signatureConverter(
1977 funcType.getNumInputs());
1978 auto llvmType =
static_cast<const LLVMTypeConverter *
>(getTypeConverter())
1979 ->convertFunctionSignature(
1981 false, signatureConverter);
1986 Location loc = funcOp.getLoc();
1987 StringRef name = funcOp.getName();
1988 auto newFuncOp = LLVM::LLVMFuncOp::create(rewriter, loc, name, llvmType);
1991 MLIRContext *context = funcOp.getContext();
1992 switch (funcOp.getFunctionControl()) {
1993 case spirv::FunctionControl::Inline:
1994 newFuncOp.setAlwaysInline(
true);
1996 case spirv::FunctionControl::DontInline:
1997 newFuncOp.setNoInline(
true);
2000#define DISPATCH(functionControl, llvmAttr) \
2001 case functionControl: \
2002 newFuncOp->setDiscardableAttr("passthrough", \
2003 ArrayAttr::get(context, {llvmAttr})); \
2006 DISPATCH(spirv::FunctionControl::Pure,
2007 StringAttr::get(context,
"readonly"));
2008 DISPATCH(spirv::FunctionControl::Const,
2009 StringAttr::get(context,
"readnone"));
2019 rewriter.inlineRegionBefore(funcOp.getBody(), newFuncOp.getBody(),
2021 if (
failed(rewriter.convertRegionTypes(
2022 &newFuncOp.getBody(), *getTypeConverter(), &signatureConverter))) {
2025 rewriter.eraseOp(funcOp);
2036 using SPIRVToLLVMConversion<spirv::ModuleOp>::SPIRVToLLVMConversion;
2039 matchAndRewrite(spirv::ModuleOp spvModuleOp, OpAdaptor adaptor,
2040 ConversionPatternRewriter &rewriter)
const override {
2043 ModuleOp::create(rewriter, spvModuleOp.getLoc(), spvModuleOp.getName());
2044 rewriter.inlineRegionBefore(spvModuleOp.getRegion(), newModuleOp.getBody());
2047 rewriter.eraseBlock(&newModuleOp.getBodyRegion().back());
2048 rewriter.eraseOp(spvModuleOp);
2057class VectorShufflePattern
2060 using SPIRVToLLVMConversion<spirv::VectorShuffleOp>::SPIRVToLLVMConversion;
2062 matchAndRewrite(spirv::VectorShuffleOp op, OpAdaptor adaptor,
2063 ConversionPatternRewriter &rewriter)
const override {
2064 Location loc = op.getLoc();
2065 auto components = adaptor.getComponents();
2066 auto vector1 = adaptor.getVector1();
2067 auto vector2 = adaptor.getVector2();
2068 int vector1Size = cast<VectorType>(vector1.getType()).getNumElements();
2069 int vector2Size = cast<VectorType>(vector2.getType()).getNumElements();
2070 if (vector1Size == vector2Size) {
2071 rewriter.replaceOpWithNewOp<LLVM::ShuffleVectorOp>(
2072 op, vector1, vector2,
2073 LLVM::convertArrayToIndices<int32_t>(components));
2077 auto dstType = getTypeConverter()->convertType(op.getType());
2079 return rewriter.notifyMatchFailure(op,
"type conversion failed");
2080 auto scalarType = cast<VectorType>(dstType).getElementType();
2081 auto componentsArray = components.getValue();
2082 auto *context = rewriter.getContext();
2083 auto llvmI32Type = IntegerType::get(context, 32);
2084 Value targetOp = LLVM::PoisonOp::create(rewriter, loc, dstType);
2085 for (
unsigned i = 0; i < componentsArray.size(); i++) {
2086 if (!isa<IntegerAttr>(componentsArray[i]))
2087 return op.emitError(
"unable to support non-constant component");
2089 int indexVal = cast<IntegerAttr>(componentsArray[i]).getInt();
2094 Value baseVector = vector1;
2095 if (indexVal >= vector1Size) {
2096 offsetVal = vector1Size;
2097 baseVector = vector2;
2100 Value dstIndex = LLVM::ConstantOp::create(
2101 rewriter, loc, llvmI32Type,
2102 rewriter.getIntegerAttr(rewriter.getI32Type(), i));
2103 Value index = LLVM::ConstantOp::create(
2104 rewriter, loc, llvmI32Type,
2105 rewriter.getIntegerAttr(rewriter.getI32Type(), indexVal - offsetVal));
2107 auto extractOp = LLVM::ExtractElementOp::create(rewriter, loc, scalarType,
2109 targetOp = LLVM::InsertElementOp::create(rewriter, loc, dstType, targetOp,
2110 extractOp, dstIndex);
2112 rewriter.replaceOp(op, targetOp);
2123 spirv::ClientAPI clientAPI) {
2140 spirv::ClientAPI clientAPI) {
2143 DirectConversionPattern<spirv::IAddOp, LLVM::AddOp>,
2144 DirectConversionPattern<spirv::IMulOp, LLVM::MulOp>,
2145 DirectConversionPattern<spirv::ISubOp, LLVM::SubOp>,
2146 DirectConversionPattern<spirv::FAddOp, LLVM::FAddOp>,
2147 DirectConversionPattern<spirv::FDivOp, LLVM::FDivOp>,
2148 DirectConversionPattern<spirv::FMulOp, LLVM::FMulOp>,
2149 DirectConversionPattern<spirv::FNegateOp, LLVM::FNegOp>,
2150 DirectConversionPattern<spirv::FRemOp, LLVM::FRemOp>,
2151 DirectConversionPattern<spirv::FSubOp, LLVM::FSubOp>,
2152 DirectConversionPattern<spirv::SDivOp, LLVM::SDivOp>,
2153 DirectConversionPattern<spirv::SRemOp, LLVM::SRemOp>,
2154 DirectConversionPattern<spirv::UDivOp, LLVM::UDivOp>,
2155 DirectConversionPattern<spirv::UModOp, LLVM::URemOp>, FModPattern,
2156 SModPattern, VectorTimesScalarPattern, SNegatePattern,
2157 ArithmeticWithOverflowPattern<spirv::IAddCarryOp,
2158 LLVM::UAddWithOverflowOp>,
2159 ArithmeticWithOverflowPattern<spirv::ISubBorrowOp,
2160 LLVM::USubWithOverflowOp>,
2163 BitFieldInsertPattern, BitFieldUExtractPattern, BitFieldSExtractPattern,
2164 DirectConversionPattern<spirv::BitCountOp, LLVM::CtPopOp>,
2165 DirectConversionPattern<spirv::BitReverseOp, LLVM::BitReverseOp>,
2166 DirectConversionPattern<spirv::BitwiseAndOp, LLVM::AndOp>,
2167 DirectConversionPattern<spirv::BitwiseOrOp, LLVM::OrOp>,
2168 DirectConversionPattern<spirv::BitwiseXorOp, LLVM::XOrOp>,
2169 NotPattern<spirv::NotOp>,
2172 BitcastConversionPattern,
2173 DirectConversionPattern<spirv::ConvertFToSOp, LLVM::FPToSIOp>,
2174 DirectConversionPattern<spirv::ConvertFToUOp, LLVM::FPToUIOp>,
2175 DirectConversionPattern<spirv::ConvertSToFOp, LLVM::SIToFPOp>,
2176 DirectConversionPattern<spirv::ConvertUToFOp, LLVM::UIToFPOp>,
2177 IndirectCastPattern<spirv::FConvertOp, LLVM::FPExtOp, LLVM::FPTruncOp>,
2178 IndirectCastPattern<spirv::SConvertOp, LLVM::SExtOp, LLVM::TruncOp>,
2179 IndirectCastPattern<spirv::UConvertOp, LLVM::ZExtOp, LLVM::TruncOp>,
2180 DirectConversionPattern<spirv::ConvertPtrToUOp, LLVM::PtrToIntOp>,
2181 DirectConversionPattern<spirv::ConvertUToPtrOp, LLVM::IntToPtrOp>,
2182 DirectConversionPattern<spirv::PtrCastToGenericOp, LLVM::AddrSpaceCastOp>,
2183 DirectConversionPattern<spirv::GenericCastToPtrOp, LLVM::AddrSpaceCastOp>,
2184 DirectConversionPattern<spirv::GenericCastToPtrExplicitOp,
2185 LLVM::AddrSpaceCastOp>,
2188 IComparePattern<spirv::IEqualOp, LLVM::ICmpPredicate::eq>,
2189 IComparePattern<spirv::INotEqualOp, LLVM::ICmpPredicate::ne>,
2190 FComparePattern<spirv::FOrdEqualOp, LLVM::FCmpPredicate::oeq>,
2191 FComparePattern<spirv::FOrdGreaterThanOp, LLVM::FCmpPredicate::ogt>,
2192 FComparePattern<spirv::FOrdGreaterThanEqualOp, LLVM::FCmpPredicate::oge>,
2193 FComparePattern<spirv::FOrdLessThanEqualOp, LLVM::FCmpPredicate::ole>,
2194 FComparePattern<spirv::FOrdLessThanOp, LLVM::FCmpPredicate::olt>,
2195 FComparePattern<spirv::FOrdNotEqualOp, LLVM::FCmpPredicate::one>,
2196 FComparePattern<spirv::FUnordEqualOp, LLVM::FCmpPredicate::ueq>,
2197 FComparePattern<spirv::FUnordGreaterThanOp, LLVM::FCmpPredicate::ugt>,
2198 FComparePattern<spirv::FUnordGreaterThanEqualOp,
2199 LLVM::FCmpPredicate::uge>,
2200 FComparePattern<spirv::FUnordLessThanEqualOp, LLVM::FCmpPredicate::ule>,
2201 FComparePattern<spirv::FUnordLessThanOp, LLVM::FCmpPredicate::ult>,
2202 FComparePattern<spirv::FUnordNotEqualOp, LLVM::FCmpPredicate::une>,
2203 FComparePattern<spirv::OrderedOp, LLVM::FCmpPredicate::ord>,
2204 FComparePattern<spirv::UnorderedOp, LLVM::FCmpPredicate::uno>,
2205 IComparePattern<spirv::SGreaterThanOp, LLVM::ICmpPredicate::sgt>,
2206 IComparePattern<spirv::SGreaterThanEqualOp, LLVM::ICmpPredicate::sge>,
2207 IComparePattern<spirv::SLessThanEqualOp, LLVM::ICmpPredicate::sle>,
2208 IComparePattern<spirv::SLessThanOp, LLVM::ICmpPredicate::slt>,
2209 IComparePattern<spirv::UGreaterThanOp, LLVM::ICmpPredicate::ugt>,
2210 IComparePattern<spirv::UGreaterThanEqualOp, LLVM::ICmpPredicate::uge>,
2211 IComparePattern<spirv::ULessThanEqualOp, LLVM::ICmpPredicate::ule>,
2212 IComparePattern<spirv::ULessThanOp, LLVM::ICmpPredicate::ult>,
2215 ConstantScalarAndVectorPattern,
2218 BranchConversionPattern, BranchConditionalConversionPattern,
2219 FunctionCallPattern, LoopPattern, SelectionPattern,
2220 ErasePattern<spirv::MergeOp>,
2223 ErasePattern<spirv::EntryPointOp>, ExecutionModePattern,
2226 DirectConversionPattern<spirv::GLCeilOp, LLVM::FCeilOp>,
2227 DirectConversionPattern<spirv::GLCosOp, LLVM::CosOp>,
2228 DirectConversionPattern<spirv::GLExpOp, LLVM::ExpOp>,
2229 DirectConversionPattern<spirv::GLExp2Op, LLVM::Exp2Op>,
2230 DirectConversionPattern<spirv::GLFAbsOp, LLVM::FAbsOp>,
2231 DirectConversionPattern<spirv::GLFloorOp, LLVM::FFloorOp>,
2232 DirectConversionPattern<spirv::GLFmaOp, LLVM::FMAOp>,
2233 ClampPattern<spirv::GLFClampOp, LLVM::MinNumOp, LLVM::MaxNumOp>,
2234 ClampPattern<spirv::GLSClampOp, LLVM::SMinOp, LLVM::SMaxOp>,
2235 ClampPattern<spirv::GLUClampOp, LLVM::UMinOp, LLVM::UMaxOp>,
2236 DirectConversionPattern<spirv::GLFMaxOp, LLVM::MaxNumOp>,
2237 DirectConversionPattern<spirv::GLFMinOp, LLVM::MinNumOp>,
2238 DirectConversionPattern<spirv::GLNMaxOp, LLVM::MaxNumOp>,
2239 DirectConversionPattern<spirv::GLNMinOp, LLVM::MinNumOp>,
2240 DirectConversionPattern<spirv::GLLogOp, LLVM::LogOp>,
2241 DirectConversionPattern<spirv::GLLog2Op, LLVM::Log2Op>,
2242 DirectConversionPattern<spirv::GLPowOp, LLVM::PowOp>,
2243 DirectConversionPattern<spirv::GLRoundOp, LLVM::RoundOp>,
2244 DirectConversionPattern<spirv::GLRoundEvenOp, LLVM::RoundEvenOp>,
2245 DirectConversionPattern<spirv::GLSinOp, LLVM::SinOp>,
2246 DirectConversionPattern<spirv::GLSinhOp, LLVM::SinhOp>,
2247 DirectConversionPattern<spirv::GLCoshOp, LLVM::CoshOp>,
2248 DirectConversionPattern<spirv::GLSMaxOp, LLVM::SMaxOp>,
2249 DirectConversionPattern<spirv::GLSMinOp, LLVM::SMinOp>,
2250 DirectConversionPattern<spirv::GLSqrtOp, LLVM::SqrtOp>,
2251 DirectConversionPattern<spirv::GLUMaxOp, LLVM::UMaxOp>,
2252 DirectConversionPattern<spirv::GLUMinOp, LLVM::UMinOp>,
2253 DirectConversionPattern<spirv::GLTruncOp, LLVM::FTruncOp>,
2254 DirectConversionPattern<spirv::GLAsinOp, LLVM::ASinOp>,
2255 DirectConversionPattern<spirv::GLAcosOp, LLVM::ACosOp>,
2256 DirectConversionPattern<spirv::GLAtanOp, LLVM::ATanOp>,
2257 DirectConversionPattern<spirv::GLTanOp, LLVM::TanOp>,
2258 DirectConversionPattern<spirv::GLTanhOp, LLVM::TanhOp>,
2259 InverseSqrtPattern, SAbsPattern, FractPattern,
2260 SignPattern<spirv::GLFSignOp,
true>,
2261 SignPattern<spirv::GLSSignOp,
false>, GLFMixPattern,
2264 DirectConversionPattern<spirv::CLCeilOp, LLVM::FCeilOp>,
2265 DirectConversionPattern<spirv::CLCosOp, LLVM::CosOp>,
2266 DirectConversionPattern<spirv::CLExpOp, LLVM::ExpOp>,
2267 DirectConversionPattern<spirv::CLExp2Op, LLVM::Exp2Op>,
2268 DirectConversionPattern<spirv::CLExp10Op, LLVM::Exp10Op>,
2269 DirectConversionPattern<spirv::CLFAbsOp, LLVM::FAbsOp>,
2270 DirectConversionPattern<spirv::CLFloorOp, LLVM::FFloorOp>,
2271 DirectConversionPattern<spirv::CLFmaOp, LLVM::FMAOp>,
2272 DirectConversionPattern<spirv::CLFMaxOp, LLVM::MaxNumOp>,
2273 DirectConversionPattern<spirv::CLFMinOp, LLVM::MinNumOp>,
2274 DirectConversionPattern<spirv::CLLogOp, LLVM::LogOp>,
2275 DirectConversionPattern<spirv::CLLog2Op, LLVM::Log2Op>,
2276 DirectConversionPattern<spirv::CLLog10Op, LLVM::Log10Op>,
2277 DirectConversionPattern<spirv::CLPowOp, LLVM::PowOp>,
2278 DirectConversionPattern<spirv::CLRintOp, LLVM::RintOp>,
2279 DirectConversionPattern<spirv::CLRoundOp, LLVM::RoundOp>,
2280 DirectConversionPattern<spirv::CLSinOp, LLVM::SinOp>,
2281 DirectConversionPattern<spirv::CLSinhOp, LLVM::SinhOp>,
2282 DirectConversionPattern<spirv::CLCoshOp, LLVM::CoshOp>,
2283 DirectConversionPattern<spirv::CLTanOp, LLVM::TanOp>,
2284 DirectConversionPattern<spirv::CLTanhOp, LLVM::TanhOp>,
2285 DirectConversionPattern<spirv::CLAsinOp, LLVM::ASinOp>,
2286 DirectConversionPattern<spirv::CLAcosOp, LLVM::ACosOp>,
2287 DirectConversionPattern<spirv::CLAtanOp, LLVM::ATanOp>,
2288 DirectConversionPattern<spirv::CLAtan2Op, LLVM::ATan2Op>,
2289 DirectConversionPattern<spirv::CLSqrtOp, LLVM::SqrtOp>,
2290 DirectConversionPattern<spirv::CLTruncOp, LLVM::FTruncOp>,
2291 DirectConversionPattern<spirv::CLCopysignOp, LLVM::CopySignOp>,
2292 DirectConversionPattern<spirv::CLFmodOp, LLVM::FRemOp>,
2293 DirectConversionPattern<spirv::CLSMaxOp, LLVM::SMaxOp>,
2294 DirectConversionPattern<spirv::CLSMinOp, LLVM::SMinOp>,
2295 DirectConversionPattern<spirv::CLUMaxOp, LLVM::UMaxOp>,
2296 DirectConversionPattern<spirv::CLUMinOp, LLVM::UMinOp>, CLMixPattern,
2299 DirectConversionPattern<spirv::LogicalAndOp, LLVM::AndOp>,
2300 DirectConversionPattern<spirv::LogicalOrOp, LLVM::OrOp>,
2301 IComparePattern<spirv::LogicalEqualOp, LLVM::ICmpPredicate::eq>,
2302 IComparePattern<spirv::LogicalNotEqualOp, LLVM::ICmpPredicate::ne>,
2303 NotPattern<spirv::LogicalNotOp>,
2306 AccessChainPattern, AddressOfPattern, LoadStorePattern<spirv::LoadOp>,
2307 LoadStorePattern<spirv::StoreOp>, VariablePattern,
2310 CompositeExtractPattern, CompositeInsertPattern,
2311 DirectConversionPattern<spirv::SelectOp, LLVM::SelectOp>,
2312 DirectConversionPattern<spirv::UndefOp, LLVM::UndefOp>,
2313 VectorShufflePattern,
2316 ShiftPattern<spirv::ShiftRightArithmeticOp, LLVM::AShrOp>,
2317 ShiftPattern<spirv::ShiftRightLogicalOp, LLVM::LShrOp>,
2318 ShiftPattern<spirv::ShiftLeftLogicalOp, LLVM::ShlOp>,
2321 ReturnPattern, ReturnValuePattern,
2327 ControlBarrierPattern<spirv::ControlBarrierOp>,
2328 ControlBarrierPattern<spirv::INTELControlBarrierArriveOp>,
2329 ControlBarrierPattern<spirv::INTELControlBarrierWaitOp>,
2332 GroupReducePattern<spirv::GroupIAddOp>,
2333 GroupReducePattern<spirv::GroupFAddOp>,
2334 GroupReducePattern<spirv::GroupFMinOp>,
2335 GroupReducePattern<spirv::GroupUMinOp>,
2336 GroupReducePattern<spirv::GroupSMinOp,
true>,
2337 GroupReducePattern<spirv::GroupFMaxOp>,
2338 GroupReducePattern<spirv::GroupUMaxOp>,
2339 GroupReducePattern<spirv::GroupSMaxOp,
true>,
2340 GroupReducePattern<spirv::GroupNonUniformIAddOp,
false,
2342 GroupReducePattern<spirv::GroupNonUniformFAddOp,
false,
2344 GroupReducePattern<spirv::GroupNonUniformIMulOp,
false,
2346 GroupReducePattern<spirv::GroupNonUniformFMulOp,
false,
2348 GroupReducePattern<spirv::GroupNonUniformSMinOp,
true,
2350 GroupReducePattern<spirv::GroupNonUniformUMinOp,
false,
2352 GroupReducePattern<spirv::GroupNonUniformFMinOp,
false,
2354 GroupReducePattern<spirv::GroupNonUniformSMaxOp,
true,
2356 GroupReducePattern<spirv::GroupNonUniformUMaxOp,
false,
2358 GroupReducePattern<spirv::GroupNonUniformFMaxOp,
false,
2360 GroupReducePattern<spirv::GroupNonUniformBitwiseAndOp,
false,
2362 GroupReducePattern<spirv::GroupNonUniformBitwiseOrOp,
false,
2364 GroupReducePattern<spirv::GroupNonUniformBitwiseXorOp,
false,
2366 GroupReducePattern<spirv::GroupNonUniformLogicalAndOp,
false,
2368 GroupReducePattern<spirv::GroupNonUniformLogicalOrOp,
false,
2370 GroupReducePattern<spirv::GroupNonUniformLogicalXorOp,
false,
2374 patterns.
add<GlobalVariablePattern>(clientAPI, patterns.
getContext(),
2377 patterns.
add<ScalePattern<spirv::GLRadiansOp>>(
2378 0.017453292519943295, patterns.
getContext(), typeConverter);
2380 patterns.
add<ScalePattern<spirv::GLDegreesOp>>(
2381 57.29577951308232, patterns.
getContext(), typeConverter);
2386 patterns.
add<FuncConversionPattern>(patterns.
getContext(), typeConverter);
2391 patterns.
add<ModuleConversionPattern>(patterns.
getContext(), typeConverter);
2400 auto spvModules =
module.getOps<spirv::ModuleOp>();
2401 for (
auto spvModule : spvModules) {
2402 spvModule.walk([&](spirv::GlobalVariableOp op) {
2403 IntegerAttr descriptorSet = op.getDescriptorSetAttr();
2404 IntegerAttr binding = op.getBindingAttr();
2407 if (descriptorSet && binding) {
2410 auto moduleAndName =
2411 spvModule.getName().has_value()
2412 ? spvModule.getName()->str() +
"_" + op.getSymName().str()
2413 : op.getSymName().str();
2415 llvm::formatv(
"{0}_descriptor_set{1}_binding{2}", moduleAndName,
2416 std::to_string(descriptorSet.getInt()),
2417 std::to_string(binding.getInt()));
2418 auto nameAttr = StringAttr::get(op->getContext(), name);
2423 op.emitError(
"unable to replace all symbol uses for ") << name;
2425 op.removeDescriptorSetAttr();
2426 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 Type getElementType(Type type)
Determine the element type of type.
static Value max(ImplicitLocOpBuilder &builder, Value value, Value bound)
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...
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.