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())
308 using SPIRVToLLVMConversion<spirv::AccessChainOp>::SPIRVToLLVMConversion;
311 matchAndRewrite(spirv::AccessChainOp op, OpAdaptor adaptor,
312 ConversionPatternRewriter &rewriter)
const override {
314 getTypeConverter()->convertType(op.getComponentPtr().getType());
316 return rewriter.notifyMatchFailure(op,
"type conversion failed");
318 auto indices = llvm::to_vector<4>(adaptor.getIndices());
319 Type indexType = op.getIndices().front().getType();
320 auto llvmIndexType = getTypeConverter()->convertType(indexType);
322 return rewriter.notifyMatchFailure(op,
"type conversion failed");
324 LLVM::ConstantOp::create(rewriter, op.getLoc(), llvmIndexType,
325 rewriter.getIntegerAttr(indexType, 0));
328 auto elementType = getTypeConverter()->convertType(
329 cast<spirv::PointerType>(op.getBasePtr().getType()).getPointeeType());
331 return rewriter.notifyMatchFailure(op,
"type conversion failed");
332 rewriter.replaceOpWithNewOp<LLVM::GEPOp>(op, dstType, elementType,
333 adaptor.getBasePtr(),
indices);
340 using SPIRVToLLVMConversion<spirv::AddressOfOp>::SPIRVToLLVMConversion;
343 matchAndRewrite(spirv::AddressOfOp op, OpAdaptor adaptor,
344 ConversionPatternRewriter &rewriter)
const override {
345 auto dstType = getTypeConverter()->convertType(op.getPointer().getType());
347 return rewriter.notifyMatchFailure(op,
"type conversion failed");
348 rewriter.replaceOpWithNewOp<LLVM::AddressOfOp>(op, dstType,
354class BitFieldInsertPattern
357 using SPIRVToLLVMConversion<spirv::BitFieldInsertOp>::SPIRVToLLVMConversion;
360 matchAndRewrite(spirv::BitFieldInsertOp op, OpAdaptor adaptor,
361 ConversionPatternRewriter &rewriter)
const override {
362 auto srcType = op.getType();
363 auto dstType = getTypeConverter()->convertType(srcType);
365 return rewriter.notifyMatchFailure(op,
"type conversion failed");
366 Location loc = op.getLoc();
370 *getTypeConverter(), rewriter);
372 *getTypeConverter(), rewriter);
376 Value maskShiftedByCount =
377 LLVM::ShlOp::create(rewriter, loc, dstType, minusOne, count);
378 Value negated = LLVM::XOrOp::create(rewriter, loc, dstType,
379 maskShiftedByCount, minusOne);
380 Value maskShiftedByCountAndOffset =
381 LLVM::ShlOp::create(rewriter, loc, dstType, negated, offset);
382 Value mask = LLVM::XOrOp::create(rewriter, loc, dstType,
383 maskShiftedByCountAndOffset, minusOne);
388 LLVM::AndOp::create(rewriter, loc, dstType, op.getBase(), mask);
389 Value insertShiftedByOffset =
390 LLVM::ShlOp::create(rewriter, loc, dstType, op.getInsert(), offset);
391 rewriter.replaceOpWithNewOp<LLVM::OrOp>(op, dstType, baseAndMask,
392 insertShiftedByOffset);
398class ConstantScalarAndVectorPattern
401 using SPIRVToLLVMConversion<spirv::ConstantOp>::SPIRVToLLVMConversion;
404 matchAndRewrite(spirv::ConstantOp constOp, OpAdaptor adaptor,
405 ConversionPatternRewriter &rewriter)
const override {
406 auto srcType = constOp.getType();
407 if (!isa<VectorType>(srcType) && !srcType.isIntOrFloat())
410 auto dstType = getTypeConverter()->convertType(srcType);
412 return rewriter.notifyMatchFailure(constOp,
"type conversion failed");
421 auto signlessType = rewriter.getIntegerType(
getBitWidth(srcType));
423 if (isa<VectorType>(srcType)) {
424 auto dstElementsAttr = cast<DenseIntElementsAttr>(constOp.getValue());
425 rewriter.replaceOpWithNewOp<LLVM::ConstantOp>(
427 dstElementsAttr.mapValues(
428 signlessType, [&](
const APInt &value) {
return value; }));
431 auto srcAttr = cast<IntegerAttr>(constOp.getValue());
432 auto dstAttr = rewriter.getIntegerAttr(signlessType, srcAttr.getValue());
433 rewriter.replaceOpWithNewOp<LLVM::ConstantOp>(constOp, dstType, dstAttr);
436 rewriter.replaceOpWithNewOp<LLVM::ConstantOp>(
437 constOp, dstType, adaptor.getOperands(), constOp->getAttrs());
442class BitFieldSExtractPattern
445 using SPIRVToLLVMConversion<spirv::BitFieldSExtractOp>::SPIRVToLLVMConversion;
448 matchAndRewrite(spirv::BitFieldSExtractOp op, OpAdaptor adaptor,
449 ConversionPatternRewriter &rewriter)
const override {
450 auto srcType = op.getType();
451 auto dstType = getTypeConverter()->convertType(srcType);
453 return rewriter.notifyMatchFailure(op,
"type conversion failed");
454 Location loc = op.getLoc();
458 *getTypeConverter(), rewriter);
460 *getTypeConverter(), rewriter);
463 IntegerType integerType;
464 if (
auto vecType = dyn_cast<VectorType>(srcType))
465 integerType = cast<IntegerType>(vecType.getElementType());
467 integerType = cast<IntegerType>(srcType);
469 auto baseSize = rewriter.getIntegerAttr(integerType,
getBitWidth(srcType));
471 isa<VectorType>(srcType)
472 ? LLVM::ConstantOp::create(
473 rewriter, loc, dstType,
474 SplatElementsAttr::get(cast<ShapedType>(srcType), baseSize))
475 : LLVM::ConstantOp::create(rewriter, loc, dstType, baseSize);
479 Value countPlusOffset =
480 LLVM::AddOp::create(rewriter, loc, dstType, count, offset);
481 Value amountToShiftLeft =
482 LLVM::SubOp::create(rewriter, loc, dstType, size, countPlusOffset);
483 Value baseShiftedLeft = LLVM::ShlOp::create(
484 rewriter, loc, dstType, op.getBase(), amountToShiftLeft);
487 Value amountToShiftRight =
488 LLVM::AddOp::create(rewriter, loc, dstType, offset, amountToShiftLeft);
489 rewriter.replaceOpWithNewOp<LLVM::AShrOp>(op, dstType, baseShiftedLeft,
495class BitFieldUExtractPattern
498 using SPIRVToLLVMConversion<spirv::BitFieldUExtractOp>::SPIRVToLLVMConversion;
501 matchAndRewrite(spirv::BitFieldUExtractOp op, OpAdaptor adaptor,
502 ConversionPatternRewriter &rewriter)
const override {
503 auto srcType = op.getType();
504 auto dstType = getTypeConverter()->convertType(srcType);
506 return rewriter.notifyMatchFailure(op,
"type conversion failed");
507 Location loc = op.getLoc();
511 *getTypeConverter(), rewriter);
513 *getTypeConverter(), rewriter);
517 Value maskShiftedByCount =
518 LLVM::ShlOp::create(rewriter, loc, dstType, minusOne, count);
519 Value mask = LLVM::XOrOp::create(rewriter, loc, dstType, maskShiftedByCount,
524 LLVM::LShrOp::create(rewriter, loc, dstType, op.getBase(), offset);
525 rewriter.replaceOpWithNewOp<LLVM::AndOp>(op, dstType, shiftedBase, mask);
532 using SPIRVToLLVMConversion<spirv::BranchOp>::SPIRVToLLVMConversion;
535 matchAndRewrite(spirv::BranchOp branchOp, OpAdaptor adaptor,
536 ConversionPatternRewriter &rewriter)
const override {
537 rewriter.replaceOpWithNewOp<LLVM::BrOp>(branchOp, adaptor.getOperands(),
538 branchOp.getTarget());
543class BranchConditionalConversionPattern
546 using SPIRVToLLVMConversion<
547 spirv::BranchConditionalOp>::SPIRVToLLVMConversion;
550 matchAndRewrite(spirv::BranchConditionalOp op, OpAdaptor adaptor,
551 ConversionPatternRewriter &rewriter)
const override {
554 if (
auto weights = op.getBranchWeights()) {
555 SmallVector<int32_t> weightValues;
556 for (
auto weight : weights->getAsRange<IntegerAttr>())
557 weightValues.push_back(weight.getInt());
561 rewriter.replaceOpWithNewOp<LLVM::CondBrOp>(
562 op, op.getCondition(), op.getTrueBlockArguments(),
563 op.getFalseBlockArguments(), branchWeights, op.getTrueBlock(),
572class CompositeExtractPattern
575 using SPIRVToLLVMConversion<spirv::CompositeExtractOp>::SPIRVToLLVMConversion;
578 matchAndRewrite(spirv::CompositeExtractOp op, OpAdaptor adaptor,
579 ConversionPatternRewriter &rewriter)
const override {
580 auto dstType = this->getTypeConverter()->convertType(op.getType());
582 return rewriter.notifyMatchFailure(op,
"type conversion failed");
584 Type containerType = op.getComposite().getType();
585 if (isa<VectorType>(containerType)) {
586 Location loc = op.getLoc();
587 IntegerAttr value = cast<IntegerAttr>(op.getIndices()[0]);
589 rewriter.replaceOpWithNewOp<LLVM::ExtractElementOp>(
590 op, dstType, adaptor.getComposite(), index);
594 rewriter.replaceOpWithNewOp<LLVM::ExtractValueOp>(
595 op, adaptor.getComposite(),
596 LLVM::convertArrayToIndices(op.getIndices()));
604class CompositeInsertPattern
607 using SPIRVToLLVMConversion<spirv::CompositeInsertOp>::SPIRVToLLVMConversion;
610 matchAndRewrite(spirv::CompositeInsertOp op, OpAdaptor adaptor,
611 ConversionPatternRewriter &rewriter)
const override {
612 auto dstType = this->getTypeConverter()->convertType(op.getType());
614 return rewriter.notifyMatchFailure(op,
"type conversion failed");
616 Type containerType = op.getComposite().getType();
617 if (isa<VectorType>(containerType)) {
618 Location loc = op.getLoc();
619 IntegerAttr value = cast<IntegerAttr>(op.getIndices()[0]);
621 rewriter.replaceOpWithNewOp<LLVM::InsertElementOp>(
622 op, dstType, adaptor.getComposite(), adaptor.getObject(), index);
626 rewriter.replaceOpWithNewOp<LLVM::InsertValueOp>(
627 op, adaptor.getComposite(), adaptor.getObject(),
628 LLVM::convertArrayToIndices(op.getIndices()));
635template <
typename SPIRVOp,
typename LLVMOp>
638 using SPIRVToLLVMConversion<SPIRVOp>::SPIRVToLLVMConversion;
641 matchAndRewrite(SPIRVOp op,
typename SPIRVOp::Adaptor adaptor,
642 ConversionPatternRewriter &rewriter)
const override {
643 auto dstType = this->getTypeConverter()->convertType(op.getType());
645 return rewriter.notifyMatchFailure(op,
"type conversion failed");
646 rewriter.template replaceOpWithNewOp<LLVMOp>(
647 op, dstType, adaptor.getOperands(), op->getAttrs());
657template <
typename SPIRVOp,
typename LLVMOp>
660 using SPIRVToLLVMConversion<SPIRVOp>::SPIRVToLLVMConversion;
663 matchAndRewrite(SPIRVOp op,
typename SPIRVOp::Adaptor adaptor,
664 ConversionPatternRewriter &rewriter)
const override {
665 Type dstType = this->getTypeConverter()->convertType(op.getType());
667 return rewriter.notifyMatchFailure(op,
"type conversion failed");
669 Location loc = op.getLoc();
670 Type operandType = adaptor.getOperand1().getType();
671 Type overflowType = rewriter.getI1Type();
672 if (
auto vecType = dyn_cast<VectorType>(operandType))
673 overflowType = VectorType::get(vecType.getShape(), overflowType);
675 Type intrType = LLVM::LLVMStructType::getLiteral(
676 rewriter.getContext(), {operandType, overflowType});
677 Value intrResult = LLVMOp::create(
678 rewriter, loc, intrType, adaptor.getOperand1(), adaptor.getOperand2());
679 Value lowBits = LLVM::ExtractValueOp::create(rewriter, loc, intrResult, 0);
680 Value overflow = LLVM::ExtractValueOp::create(rewriter, loc, intrResult, 1);
681 overflow = LLVM::ZExtOp::create(rewriter, loc, operandType, overflow);
683 Value
result = LLVM::PoisonOp::create(rewriter, loc, dstType);
684 result = LLVM::InsertValueOp::create(rewriter, loc,
result, lowBits,
685 ArrayRef<int64_t>{0});
686 result = LLVM::InsertValueOp::create(rewriter, loc,
result, overflow,
687 ArrayRef<int64_t>{1});
688 rewriter.replaceOp(op,
result);
695class ExecutionModePattern
698 using SPIRVToLLVMConversion<spirv::ExecutionModeOp>::SPIRVToLLVMConversion;
701 matchAndRewrite(spirv::ExecutionModeOp op, OpAdaptor adaptor,
702 ConversionPatternRewriter &rewriter)
const override {
706 ModuleOp module = op->getParentOfType<ModuleOp>();
707 spirv::ExecutionModeAttr executionModeAttr = op.getExecutionModeAttr();
708 std::string moduleName;
709 if (module.getName().has_value())
710 moduleName =
"_" +
module.getName()->str();
713 std::string executionModeInfoName = llvm::formatv(
714 "__spv_{0}_{1}_execution_mode_info_{2}", moduleName, op.getFn().str(),
715 static_cast<uint32_t
>(executionModeAttr.getValue()));
717 MLIRContext *context = rewriter.getContext();
718 OpBuilder::InsertionGuard guard(rewriter);
719 rewriter.setInsertionPointToStart(module.getBody());
726 auto llvmI32Type = IntegerType::get(context, 32);
727 SmallVector<Type, 2> fields;
728 fields.push_back(llvmI32Type);
730 if (!values.empty()) {
731 auto arrayType = LLVM::LLVMArrayType::get(llvmI32Type, values.size());
732 fields.push_back(arrayType);
734 auto structType = LLVM::LLVMStructType::getLiteral(context, fields);
737 auto global = LLVM::GlobalOp::create(
738 rewriter, UnknownLoc::get(context), structType,
true,
739 LLVM::Linkage::External, executionModeInfoName, Attribute(),
741 Location loc = global.getLoc();
742 Region ®ion = global.getInitializerRegion();
743 Block *block = rewriter.createBlock(®ion);
746 rewriter.setInsertionPointToStart(block);
747 Value structValue = LLVM::PoisonOp::create(rewriter, loc, structType);
748 Value executionMode = LLVM::ConstantOp::create(
749 rewriter, loc, llvmI32Type,
750 rewriter.getI32IntegerAttr(
751 static_cast<uint32_t
>(executionModeAttr.getValue())));
752 SmallVector<int64_t> position{0};
753 structValue = LLVM::InsertValueOp::create(rewriter, loc, structValue,
754 executionMode, position);
757 for (
unsigned i = 0, e = values.size(); i < e; ++i) {
758 auto attr = values.getValue()[i];
759 Value entry = LLVM::ConstantOp::create(rewriter, loc, llvmI32Type, attr);
760 structValue = LLVM::InsertValueOp::create(
761 rewriter, loc, structValue, entry, ArrayRef<int64_t>({1, i}));
763 LLVM::ReturnOp::create(rewriter, loc, ArrayRef<Value>({structValue}));
764 rewriter.eraseOp(op);
773class GlobalVariablePattern
776 template <
typename... Args>
777 GlobalVariablePattern(spirv::ClientAPI clientAPI, Args &&...args)
778 : SPIRVToLLVMConversion<spirv::GlobalVariableOp>(
779 std::forward<Args>(args)...),
780 clientAPI(clientAPI) {}
783 matchAndRewrite(spirv::GlobalVariableOp op, OpAdaptor adaptor,
784 ConversionPatternRewriter &rewriter)
const override {
787 if (op.getInitializer())
790 auto srcType = cast<spirv::PointerType>(op.getType());
791 auto dstType = getTypeConverter()->convertType(srcType.getPointeeType());
793 return rewriter.notifyMatchFailure(op,
"type conversion failed");
798 auto storageClass = srcType.getStorageClass();
799 switch (storageClass) {
800 case spirv::StorageClass::Input:
801 case spirv::StorageClass::Private:
802 case spirv::StorageClass::Output:
803 case spirv::StorageClass::StorageBuffer:
804 case spirv::StorageClass::UniformConstant:
813 bool isConstant = (storageClass == spirv::StorageClass::Input) ||
814 (storageClass == spirv::StorageClass::UniformConstant);
820 auto linkage = storageClass == spirv::StorageClass::Private
821 ? LLVM::Linkage::Private
822 : LLVM::Linkage::External;
823 StringAttr locationAttrName = op.getLocationAttrName();
824 IntegerAttr locationAttr = op.getLocationAttr();
825 auto newGlobalOp = rewriter.replaceOpWithNewOp<LLVM::GlobalOp>(
826 op, dstType, isConstant, linkage, op.getSymName(), Attribute(),
831 newGlobalOp->setAttr(locationAttrName, locationAttr);
837 spirv::ClientAPI clientAPI;
842template <
typename SPIRVOp,
typename LLVMExtOp,
typename LLVMTruncOp>
845 using SPIRVToLLVMConversion<SPIRVOp>::SPIRVToLLVMConversion;
848 matchAndRewrite(SPIRVOp op,
typename SPIRVOp::Adaptor adaptor,
849 ConversionPatternRewriter &rewriter)
const override {
851 Type fromType = op.getOperand().getType();
852 Type toType = op.getType();
854 auto dstType = this->getTypeConverter()->convertType(toType);
856 return rewriter.notifyMatchFailure(op,
"type conversion failed");
859 rewriter.template replaceOpWithNewOp<LLVMExtOp>(op, dstType,
860 adaptor.getOperands());
864 rewriter.template replaceOpWithNewOp<LLVMTruncOp>(op, dstType,
865 adaptor.getOperands());
872class FunctionCallPattern
875 using SPIRVToLLVMConversion<spirv::FunctionCallOp>::SPIRVToLLVMConversion;
878 matchAndRewrite(spirv::FunctionCallOp callOp, OpAdaptor adaptor,
879 ConversionPatternRewriter &rewriter)
const override {
880 if (callOp.getNumResults() == 0) {
881 auto newOp = rewriter.replaceOpWithNewOp<LLVM::CallOp>(
882 callOp,
TypeRange(), adaptor.getOperands(), callOp->getAttrs());
883 newOp.getProperties().operandSegmentSizes = {
884 static_cast<int32_t
>(adaptor.getOperands().size()), 0};
885 newOp.getProperties().op_bundle_sizes = rewriter.getDenseI32ArrayAttr({});
890 auto dstType = getTypeConverter()->convertType(callOp.getType(0));
892 return rewriter.notifyMatchFailure(callOp,
"type conversion failed");
893 auto newOp = rewriter.replaceOpWithNewOp<LLVM::CallOp>(
894 callOp, dstType, adaptor.getOperands(), callOp->getAttrs());
895 newOp.getProperties().operandSegmentSizes = {
896 static_cast<int32_t
>(adaptor.getOperands().size()), 0};
897 newOp.getProperties().op_bundle_sizes = rewriter.getDenseI32ArrayAttr({});
903template <
typename SPIRVOp, LLVM::FCmpPredicate predicate>
906 using SPIRVToLLVMConversion<SPIRVOp>::SPIRVToLLVMConversion;
909 matchAndRewrite(SPIRVOp op,
typename SPIRVOp::Adaptor adaptor,
910 ConversionPatternRewriter &rewriter)
const override {
912 auto dstType = this->getTypeConverter()->convertType(op.getType());
914 return rewriter.notifyMatchFailure(op,
"type conversion failed");
916 rewriter.template replaceOpWithNewOp<LLVM::FCmpOp>(
917 op, dstType, predicate, op.getOperand1(), op.getOperand2());
923template <
typename SPIRVOp, LLVM::ICmpPredicate predicate>
926 using SPIRVToLLVMConversion<SPIRVOp>::SPIRVToLLVMConversion;
929 matchAndRewrite(SPIRVOp op,
typename SPIRVOp::Adaptor adaptor,
930 ConversionPatternRewriter &rewriter)
const override {
932 auto dstType = this->getTypeConverter()->convertType(op.getType());
934 return rewriter.notifyMatchFailure(op,
"type conversion failed");
936 rewriter.template replaceOpWithNewOp<LLVM::ICmpOp>(
937 op, dstType, predicate, op.getOperand1(), op.getOperand2());
942class InverseSqrtPattern
945 using SPIRVToLLVMConversion<spirv::GLInverseSqrtOp>::SPIRVToLLVMConversion;
948 matchAndRewrite(spirv::GLInverseSqrtOp op, OpAdaptor adaptor,
949 ConversionPatternRewriter &rewriter)
const override {
950 auto srcType = op.getType();
951 auto dstType = getTypeConverter()->convertType(srcType);
953 return rewriter.notifyMatchFailure(op,
"type conversion failed");
955 Location loc = op.getLoc();
957 Value sqrt = LLVM::SqrtOp::create(rewriter, loc, dstType, op.getOperand());
958 rewriter.replaceOpWithNewOp<LLVM::FDivOp>(op, dstType, one, sqrt);
965class VectorTimesScalarPattern
968 using SPIRVToLLVMConversion<
969 spirv::VectorTimesScalarOp>::SPIRVToLLVMConversion;
972 matchAndRewrite(spirv::VectorTimesScalarOp op, OpAdaptor adaptor,
973 ConversionPatternRewriter &rewriter)
const override {
974 Type srcType = op.getType();
975 Type dstType = getTypeConverter()->convertType(srcType);
977 return rewriter.notifyMatchFailure(op,
"type conversion failed");
979 unsigned numElements = op.getVector().getType().getNumElements();
980 Value broadcasted =
broadcast(op.getLoc(), adaptor.getScalar(), numElements,
981 *getTypeConverter(), rewriter);
982 rewriter.replaceOpWithNewOp<LLVM::FMulOp>(op, dstType, adaptor.getVector(),
991 using SPIRVToLLVMConversion<spirv::SNegateOp>::SPIRVToLLVMConversion;
994 matchAndRewrite(spirv::SNegateOp op, OpAdaptor adaptor,
995 ConversionPatternRewriter &rewriter)
const override {
996 Type srcType = op.getType();
997 Type dstType = getTypeConverter()->convertType(srcType);
999 return rewriter.notifyMatchFailure(op,
"type conversion failed");
1001 Location loc = op.getLoc();
1002 IntegerAttr zeroAttr = rewriter.getIntegerAttr(
1006 rewriter.replaceOpWithNewOp<LLVM::SubOp>(op, dstType, zero,
1007 adaptor.getOperand());
1014template <
typename SPIRVOp,
typename LLVMMinOp,
typename LLVMMaxOp>
1017 using SPIRVToLLVMConversion<SPIRVOp>::SPIRVToLLVMConversion;
1020 matchAndRewrite(SPIRVOp op,
typename SPIRVOp::Adaptor adaptor,
1021 ConversionPatternRewriter &rewriter)
const override {
1022 Type dstType = this->getTypeConverter()->convertType(op.getType());
1024 return rewriter.notifyMatchFailure(op,
"type conversion failed");
1026 Location loc = op.getLoc();
1027 Value
max = LLVMMaxOp::create(rewriter, loc, dstType, adaptor.getX(),
1029 rewriter.template replaceOpWithNewOp<LLVMMinOp>(op, dstType,
max,
1040 using SPIRVToLLVMConversion<spirv::FModOp>::SPIRVToLLVMConversion;
1043 matchAndRewrite(spirv::FModOp op, OpAdaptor adaptor,
1044 ConversionPatternRewriter &rewriter)
const override {
1045 Type dstType = getTypeConverter()->convertType(op.getType());
1047 return rewriter.notifyMatchFailure(op,
"type conversion failed");
1049 Location loc = op.getLoc();
1050 Value
lhs = adaptor.getOperand1();
1051 Value
rhs = adaptor.getOperand2();
1052 Value
div = LLVM::FDivOp::create(rewriter, loc, dstType,
lhs,
rhs);
1053 Value floored = LLVM::FFloorOp::create(rewriter, loc, dstType,
div);
1054 Value scaled = LLVM::FMulOp::create(rewriter, loc, dstType,
rhs, floored);
1055 rewriter.replaceOpWithNewOp<LLVM::FSubOp>(op, dstType,
lhs, scaled);
1066 using SPIRVToLLVMConversion<spirv::SModOp>::SPIRVToLLVMConversion;
1069 matchAndRewrite(spirv::SModOp op, OpAdaptor adaptor,
1070 ConversionPatternRewriter &rewriter)
const override {
1071 Type srcType = op.getType();
1072 Type dstType = getTypeConverter()->convertType(srcType);
1074 return rewriter.notifyMatchFailure(op,
"type conversion failed");
1076 Location loc = op.getLoc();
1077 Value
lhs = adaptor.getOperand1();
1078 Value
rhs = adaptor.getOperand2();
1079 Type i1Type = rewriter.getI1Type();
1080 auto vecSrcType = dyn_cast<VectorType>(srcType);
1082 vecSrcType ? VectorType::get(vecSrcType.getShape(), i1Type) : i1Type;
1084 Value
rem = LLVM::SRemOp::create(rewriter, loc, dstType,
lhs,
rhs);
1085 IntegerAttr zeroAttr = rewriter.getIntegerAttr(
1090 Value remNonZero = LLVM::ICmpOp::create(rewriter, loc, cmpType,
1091 LLVM::ICmpPredicate::ne,
rem, zero);
1092 Value remNeg = LLVM::ICmpOp::create(rewriter, loc, cmpType,
1093 LLVM::ICmpPredicate::slt,
rem, zero);
1094 Value rhsNeg = LLVM::ICmpOp::create(rewriter, loc, cmpType,
1095 LLVM::ICmpPredicate::slt,
rhs, zero);
1096 Value signMismatch =
1097 LLVM::XOrOp::create(rewriter, loc, cmpType, remNeg, rhsNeg);
1099 LLVM::AndOp::create(rewriter, loc, cmpType, remNonZero, signMismatch);
1101 Value adjusted = LLVM::AddOp::create(rewriter, loc, dstType,
rem,
rhs);
1102 rewriter.replaceOpWithNewOp<LLVM::SelectOp>(op, dstType, needsAdjust,
1109template <
typename SPIRVOp>
1112 using SPIRVToLLVMConversion<SPIRVOp>::SPIRVToLLVMConversion;
1115 matchAndRewrite(SPIRVOp op,
typename SPIRVOp::Adaptor adaptor,
1116 ConversionPatternRewriter &rewriter)
const override {
1117 if (!op.getMemoryAccess()) {
1119 *this->getTypeConverter(), 0,
1123 auto memoryAccess = *op.getMemoryAccess();
1124 switch (memoryAccess) {
1125 case spirv::MemoryAccess::Aligned:
1126 case spirv::MemoryAccess::None:
1127 case spirv::MemoryAccess::Nontemporal:
1128 case spirv::MemoryAccess::Volatile: {
1129 unsigned alignment =
1130 memoryAccess == spirv::MemoryAccess::Aligned ? *op.getAlignment() : 0;
1131 bool isNonTemporal = memoryAccess == spirv::MemoryAccess::Nontemporal;
1132 bool isVolatile = memoryAccess == spirv::MemoryAccess::Volatile;
1134 *this->getTypeConverter(), alignment,
1135 isVolatile, isNonTemporal);
1145template <
typename SPIRVOp>
1148 using SPIRVToLLVMConversion<SPIRVOp>::SPIRVToLLVMConversion;
1151 matchAndRewrite(SPIRVOp notOp,
typename SPIRVOp::Adaptor adaptor,
1152 ConversionPatternRewriter &rewriter)
const override {
1153 auto srcType = notOp.getType();
1154 auto dstType = this->getTypeConverter()->convertType(srcType);
1156 return rewriter.notifyMatchFailure(notOp,
"type conversion failed");
1158 Location loc = notOp.getLoc();
1160 rewriter.template replaceOpWithNewOp<LLVM::XOrOp>(notOp, dstType,
1161 notOp.getOperand(), mask);
1167template <
typename SPIRVOp>
1170 using SPIRVToLLVMConversion<SPIRVOp>::SPIRVToLLVMConversion;
1173 matchAndRewrite(SPIRVOp op,
typename SPIRVOp::Adaptor adaptor,
1174 ConversionPatternRewriter &rewriter)
const override {
1175 rewriter.eraseOp(op);
1182 using SPIRVToLLVMConversion<spirv::ReturnOp>::SPIRVToLLVMConversion;
1185 matchAndRewrite(spirv::ReturnOp returnOp, OpAdaptor adaptor,
1186 ConversionPatternRewriter &rewriter)
const override {
1187 rewriter.replaceOpWithNewOp<LLVM::ReturnOp>(returnOp, ArrayRef<Type>(),
1195 using SPIRVToLLVMConversion<spirv::ReturnValueOp>::SPIRVToLLVMConversion;
1198 matchAndRewrite(spirv::ReturnValueOp returnValueOp, OpAdaptor adaptor,
1199 ConversionPatternRewriter &rewriter)
const override {
1200 rewriter.replaceOpWithNewOp<LLVM::ReturnOp>(returnValueOp, ArrayRef<Type>(),
1201 adaptor.getOperands());
1208 using SPIRVToLLVMConversion<spirv::UnreachableOp>::SPIRVToLLVMConversion;
1211 matchAndRewrite(spirv::UnreachableOp unreachableOp, OpAdaptor adaptor,
1212 ConversionPatternRewriter &rewriter)
const override {
1213 rewriter.replaceOpWithNewOp<LLVM::UnreachableOp>(unreachableOp);
1222 bool convergent =
true) {
1223 auto func = dyn_cast_or_null<LLVM::LLVMFuncOp>(
1229 func = LLVM::LLVMFuncOp::create(
1230 b, symbolTable->
getLoc(), name,
1231 LLVM::LLVMFunctionType::get(resultType, paramTypes));
1232 func.setCConv(LLVM::cconv::CConv::SPIR_FUNC);
1233 func.setConvergent(convergent);
1234 func.setNoUnwind(
true);
1235 func.setWillReturn(
true);
1240 LLVM::LLVMFuncOp
func,
1242 auto call = LLVM::CallOp::create(builder, loc,
func, args);
1243 call.setCConv(
func.getCConv());
1244 call.setConvergentAttr(
func.getConvergentAttr());
1245 call.setNoUnwindAttr(
func.getNoUnwindAttr());
1246 call.setWillReturnAttr(
func.getWillReturnAttr());
1250template <
typename BarrierOpTy>
1253 using OpAdaptor =
typename SPIRVToLLVMConversion<BarrierOpTy>::OpAdaptor;
1255 using SPIRVToLLVMConversion<BarrierOpTy>::SPIRVToLLVMConversion;
1257 static constexpr StringRef getFuncName();
1260 matchAndRewrite(BarrierOpTy controlBarrierOp, OpAdaptor adaptor,
1261 ConversionPatternRewriter &rewriter)
const override {
1262 constexpr StringRef funcName = getFuncName();
1263 Operation *symbolTable =
1264 controlBarrierOp->template getParentWithTrait<OpTrait::SymbolTable>();
1266 Type i32 = rewriter.getI32Type();
1268 Type voidTy = rewriter.getType<LLVM::LLVMVoidType>();
1269 LLVM::LLVMFuncOp func =
1272 Location loc = controlBarrierOp->getLoc();
1273 Value execution = LLVM::ConstantOp::create(
1274 rewriter, loc, i32,
static_cast<int32_t
>(adaptor.getExecutionScope()));
1275 Value memory = LLVM::ConstantOp::create(
1276 rewriter, loc, i32,
static_cast<int32_t
>(adaptor.getMemoryScope()));
1277 Value semantics = LLVM::ConstantOp::create(
1278 rewriter, loc, i32,
static_cast<int32_t
>(adaptor.getMemorySemantics()));
1281 {execution, memory, semantics});
1283 rewriter.replaceOp(controlBarrierOp, call);
1290StringRef getTypeMangling(
Type type,
bool isSigned) {
1292 .Case([](Float16Type) {
return "Dh"; })
1293 .Case([](Float32Type) {
return "f"; })
1294 .Case([](Float64Type) {
return "d"; })
1295 .Case([isSigned](IntegerType intTy) {
1296 switch (intTy.getWidth()) {
1300 return (isSigned) ?
"a" :
"c";
1302 return (isSigned) ?
"s" :
"t";
1304 return (isSigned) ?
"i" :
"j";
1306 return (isSigned) ?
"l" :
"m";
1308 llvm_unreachable(
"Unsupported integer width");
1311 .DefaultUnreachable(
"No mangling defined");
1314template <
typename ReduceOp>
1315constexpr StringLiteral getGroupFuncName();
1318constexpr StringLiteral getGroupFuncName<spirv::GroupIAddOp>() {
1319 return "_Z17__spirv_GroupIAddii";
1322constexpr StringLiteral getGroupFuncName<spirv::GroupFAddOp>() {
1323 return "_Z17__spirv_GroupFAddii";
1326constexpr StringLiteral getGroupFuncName<spirv::GroupSMinOp>() {
1327 return "_Z17__spirv_GroupSMinii";
1330constexpr StringLiteral getGroupFuncName<spirv::GroupUMinOp>() {
1331 return "_Z17__spirv_GroupUMinii";
1334constexpr StringLiteral getGroupFuncName<spirv::GroupFMinOp>() {
1335 return "_Z17__spirv_GroupFMinii";
1338constexpr StringLiteral getGroupFuncName<spirv::GroupSMaxOp>() {
1339 return "_Z17__spirv_GroupSMaxii";
1342constexpr StringLiteral getGroupFuncName<spirv::GroupUMaxOp>() {
1343 return "_Z17__spirv_GroupUMaxii";
1346constexpr StringLiteral getGroupFuncName<spirv::GroupFMaxOp>() {
1347 return "_Z17__spirv_GroupFMaxii";
1350constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformIAddOp>() {
1351 return "_Z27__spirv_GroupNonUniformIAddii";
1354constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformFAddOp>() {
1355 return "_Z27__spirv_GroupNonUniformFAddii";
1358constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformIMulOp>() {
1359 return "_Z27__spirv_GroupNonUniformIMulii";
1362constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformFMulOp>() {
1363 return "_Z27__spirv_GroupNonUniformFMulii";
1366constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformSMinOp>() {
1367 return "_Z27__spirv_GroupNonUniformSMinii";
1370constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformUMinOp>() {
1371 return "_Z27__spirv_GroupNonUniformUMinii";
1374constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformFMinOp>() {
1375 return "_Z27__spirv_GroupNonUniformFMinii";
1378constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformSMaxOp>() {
1379 return "_Z27__spirv_GroupNonUniformSMaxii";
1382constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformUMaxOp>() {
1383 return "_Z27__spirv_GroupNonUniformUMaxii";
1386constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformFMaxOp>() {
1387 return "_Z27__spirv_GroupNonUniformFMaxii";
1390constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformBitwiseAndOp>() {
1391 return "_Z33__spirv_GroupNonUniformBitwiseAndii";
1394constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformBitwiseOrOp>() {
1395 return "_Z32__spirv_GroupNonUniformBitwiseOrii";
1398constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformBitwiseXorOp>() {
1399 return "_Z33__spirv_GroupNonUniformBitwiseXorii";
1402constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformLogicalAndOp>() {
1403 return "_Z33__spirv_GroupNonUniformLogicalAndii";
1406constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformLogicalOrOp>() {
1407 return "_Z32__spirv_GroupNonUniformLogicalOrii";
1410constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformLogicalXorOp>() {
1411 return "_Z33__spirv_GroupNonUniformLogicalXorii";
1415template <
typename ReduceOp,
bool Signed = false,
bool NonUniform = false>
1418 using SPIRVToLLVMConversion<ReduceOp>::SPIRVToLLVMConversion;
1421 matchAndRewrite(ReduceOp op,
typename ReduceOp::Adaptor adaptor,
1422 ConversionPatternRewriter &rewriter)
const override {
1424 Type retTy = op.getResult().getType();
1428 SmallString<36> funcName = getGroupFuncName<ReduceOp>();
1429 funcName += getTypeMangling(retTy,
false);
1431 Type i32Ty = rewriter.getI32Type();
1432 SmallVector<Type> paramTypes{i32Ty, i32Ty, retTy};
1433 if constexpr (NonUniform) {
1434 if (adaptor.getClusterSize()) {
1436 paramTypes.push_back(i32Ty);
1440 Operation *symbolTable =
1441 op->template getParentWithTrait<OpTrait::SymbolTable>();
1443 LLVM::LLVMFuncOp func =
1446 Location loc = op.getLoc();
1447 Value scope = LLVM::ConstantOp::create(
1448 rewriter, loc, i32Ty,
1449 static_cast<int32_t
>(adaptor.getExecutionScope()));
1450 Value groupOp = LLVM::ConstantOp::create(
1451 rewriter, loc, i32Ty,
1452 static_cast<int32_t
>(adaptor.getGroupOperation()));
1453 SmallVector<Value> operands{scope, groupOp};
1454 operands.append(adaptor.getOperands().begin(), adaptor.getOperands().end());
1457 rewriter.replaceOp(op, call);
1464ControlBarrierPattern<spirv::ControlBarrierOp>::getFuncName() {
1465 return "_Z22__spirv_ControlBarrieriii";
1470ControlBarrierPattern<spirv::INTELControlBarrierArriveOp>::getFuncName() {
1471 return "_Z33__spirv_ControlBarrierArriveINTELiii";
1476ControlBarrierPattern<spirv::INTELControlBarrierWaitOp>::getFuncName() {
1477 return "_Z31__spirv_ControlBarrierWaitINTELiii";
1530 using SPIRVToLLVMConversion<spirv::LoopOp>::SPIRVToLLVMConversion;
1533 matchAndRewrite(spirv::LoopOp loopOp, OpAdaptor adaptor,
1534 ConversionPatternRewriter &rewriter)
const override {
1536 if (loopOp.getLoopControl() != spirv::LoopControl::None)
1540 if (loopOp.getBody().empty()) {
1541 rewriter.eraseOp(loopOp);
1545 Location loc = loopOp.getLoc();
1549 Block *currentBlock = rewriter.getBlock();
1551 Block *endBlock = rewriter.splitBlock(currentBlock, position);
1555 Block *entryBlock = loopOp.getEntryBlock();
1557 auto brOp = dyn_cast<spirv::BranchOp>(entryBlock->
getOperations().front());
1560 Block *headerBlock = loopOp.getHeaderBlock();
1561 rewriter.setInsertionPointToEnd(currentBlock);
1562 LLVM::BrOp::create(rewriter, loc, brOp.getBlockArguments(), headerBlock);
1563 rewriter.eraseBlock(entryBlock);
1566 Block *mergeBlock = loopOp.getMergeBlock();
1569 rewriter.setInsertionPointToEnd(mergeBlock);
1570 LLVM::BrOp::create(rewriter, loc, terminatorOperands, endBlock);
1572 rewriter.inlineRegionBefore(loopOp.getBody(), endBlock);
1583 using SPIRVToLLVMConversion<spirv::SelectionOp>::SPIRVToLLVMConversion;
1586 matchAndRewrite(spirv::SelectionOp op, OpAdaptor adaptor,
1587 ConversionPatternRewriter &rewriter)
const override {
1591 if (op.getSelectionControl() != spirv::SelectionControl::None)
1598 if (op.getBody().getBlocks().size() <= 2) {
1599 rewriter.eraseOp(op);
1603 Location loc = op.getLoc();
1607 auto *currentBlock = rewriter.getInsertionBlock();
1608 rewriter.setInsertionPointAfter(op);
1609 auto position = rewriter.getInsertionPoint();
1610 auto *continueBlock = rewriter.splitBlock(currentBlock, position);
1613 for (
auto ty : op.getResultTypes()) {
1614 Type dstTy = getTypeConverter()->convertType(ty);
1616 return rewriter.notifyMatchFailure(op,
"failed to convert type");
1617 continueBlock->addArgument(dstTy, loc);
1624 auto *headerBlock = op.getHeaderBlock();
1626 auto condBrOp = dyn_cast<spirv::BranchConditionalOp>(
1632 auto *mergeBlock = op.getMergeBlock();
1635 rewriter.setInsertionPointToEnd(mergeBlock);
1636 LLVM::BrOp::create(rewriter, loc, terminatorOperands, continueBlock);
1639 Block *trueBlock = condBrOp.getTrueBlock();
1640 Block *falseBlock = condBrOp.getFalseBlock();
1641 rewriter.setInsertionPointToEnd(currentBlock);
1642 LLVM::CondBrOp::create(rewriter, loc, condBrOp.getCondition(), trueBlock,
1643 condBrOp.getTrueTargetOperands(), falseBlock,
1644 condBrOp.getFalseTargetOperands());
1646 rewriter.eraseBlock(headerBlock);
1647 rewriter.inlineRegionBefore(op.getBody(), continueBlock);
1648 rewriter.replaceOp(op, continueBlock->getArguments());
1657template <
typename SPIRVOp,
typename LLVMOp>
1660 using SPIRVToLLVMConversion<SPIRVOp>::SPIRVToLLVMConversion;
1663 matchAndRewrite(SPIRVOp op,
typename SPIRVOp::Adaptor adaptor,
1664 ConversionPatternRewriter &rewriter)
const override {
1666 auto dstType = this->getTypeConverter()->convertType(op.getType());
1668 return rewriter.notifyMatchFailure(op,
"type conversion failed");
1670 Type op1Type = op.getOperand1().getType();
1671 Type op2Type = op.getOperand2().getType();
1673 if (op1Type == op2Type) {
1674 rewriter.template replaceOpWithNewOp<LLVMOp>(op, dstType,
1675 adaptor.getOperands());
1679 std::optional<uint64_t> dstTypeWidth =
1681 std::optional<uint64_t> op2TypeWidth =
1684 if (!dstTypeWidth || !op2TypeWidth)
1687 Location loc = op.getLoc();
1689 if (op2TypeWidth < dstTypeWidth) {
1692 LLVM::ZExtOp::create(rewriter, loc, dstType, adaptor.getOperand2());
1695 LLVM::SExtOp::create(rewriter, loc, dstType, adaptor.getOperand2());
1697 }
else if (op2TypeWidth == dstTypeWidth) {
1698 extended = adaptor.getOperand2();
1704 LLVMOp::create(rewriter, loc, dstType, adaptor.getOperand1(), extended);
1705 rewriter.replaceOp(op,
result);
1715 using SPIRVToLLVMConversion<spirv::GLSAbsOp>::SPIRVToLLVMConversion;
1718 matchAndRewrite(spirv::GLSAbsOp op, OpAdaptor adaptor,
1719 ConversionPatternRewriter &rewriter)
const override {
1720 Type dstType = getTypeConverter()->convertType(op.getType());
1722 return rewriter.notifyMatchFailure(op,
"type conversion failed");
1724 rewriter.replaceOpWithNewOp<LLVM::AbsOp>(op, dstType, adaptor.getOperand(),
1733 using SPIRVToLLVMConversion<spirv::GLFractOp>::SPIRVToLLVMConversion;
1736 matchAndRewrite(spirv::GLFractOp op, OpAdaptor adaptor,
1737 ConversionPatternRewriter &rewriter)
const override {
1738 Type dstType = getTypeConverter()->convertType(op.getType());
1740 return rewriter.notifyMatchFailure(op,
"type conversion failed");
1742 Location loc = op.getLoc();
1743 Value operand = adaptor.getOperand();
1744 Value floored = LLVM::FFloorOp::create(rewriter, loc, dstType, operand);
1745 rewriter.replaceOpWithNewOp<LLVM::FSubOp>(op, dstType, operand, floored);
1754 using SPIRVToLLVMConversion<spirv::GLFMixOp>::SPIRVToLLVMConversion;
1757 matchAndRewrite(spirv::GLFMixOp op, OpAdaptor adaptor,
1758 ConversionPatternRewriter &rewriter)
const override {
1759 Type dstType = getTypeConverter()->convertType(op.getType());
1761 return rewriter.notifyMatchFailure(op,
"type conversion failed");
1763 Location loc = op.getLoc();
1764 Value x = adaptor.getX();
1765 Value y = adaptor.getY();
1766 Value a = adaptor.getA();
1768 Value oneMinusA = LLVM::FSubOp::create(rewriter, loc, dstType, one, a);
1769 Value
lhs = LLVM::FMulOp::create(rewriter, loc, dstType, x, oneMinusA);
1770 Value
rhs = LLVM::FMulOp::create(rewriter, loc, dstType, y, a);
1771 rewriter.replaceOpWithNewOp<LLVM::FAddOp>(op, dstType,
lhs,
rhs);
1780 using SPIRVToLLVMConversion<spirv::CLMixOp>::SPIRVToLLVMConversion;
1783 matchAndRewrite(spirv::CLMixOp op, OpAdaptor adaptor,
1784 ConversionPatternRewriter &rewriter)
const override {
1785 Type dstType = getTypeConverter()->convertType(op.getType());
1787 return rewriter.notifyMatchFailure(op,
"type conversion failed");
1789 Location loc = op.getLoc();
1790 Value x = adaptor.getX();
1791 Value y = adaptor.getY();
1792 Value a = adaptor.getZ();
1793 Value diff = LLVM::FSubOp::create(rewriter, loc, dstType, y, x);
1794 rewriter.replaceOpWithNewOp<LLVM::FMAOp>(op, dstType, a, diff, x);
1801template <
typename SPIRVOp>
1804 template <
typename... Args>
1805 ScalePattern(
double scale, Args &&...args)
1806 : SPIRVToLLVMConversion<SPIRVOp>(std::forward<Args>(args)...),
1810 matchAndRewrite(SPIRVOp op,
typename SPIRVOp::Adaptor adaptor,
1811 ConversionPatternRewriter &rewriter)
const override {
1812 Type srcType = op.getType();
1813 Type dstType = this->getTypeConverter()->convertType(srcType);
1815 return rewriter.notifyMatchFailure(op,
"type conversion failed");
1817 Location loc = op.getLoc();
1819 rewriter.replaceOpWithNewOp<LLVM::FMulOp>(op, dstType, adaptor.getOperand(),
1831template <
typename SPIRVOp,
bool isFloat>
1834 using SPIRVToLLVMConversion<SPIRVOp>::SPIRVToLLVMConversion;
1837 matchAndRewrite(SPIRVOp op,
typename SPIRVOp::Adaptor adaptor,
1838 ConversionPatternRewriter &rewriter)
const override {
1839 Type srcType = op.getType();
1840 Type dstType = this->getTypeConverter()->convertType(srcType);
1842 return rewriter.notifyMatchFailure(op,
"type conversion failed");
1844 Location loc = op.getLoc();
1845 Value operand = adaptor.getOperand();
1846 auto vecSrcType = dyn_cast<VectorType>(srcType);
1847 Type i1Type = rewriter.getI1Type();
1849 vecSrcType ? VectorType::get(vecSrcType.getShape(), i1Type) : i1Type;
1851 Value zero, one, minusOne, gt, lt;
1852 if constexpr (isFloat) {
1856 gt = LLVM::FCmpOp::create(rewriter, loc, cmpType,
1857 LLVM::FCmpPredicate::ogt, operand, zero);
1858 lt = LLVM::FCmpOp::create(rewriter, loc, cmpType,
1859 LLVM::FCmpPredicate::olt, operand, zero);
1863 rewriter.getIntegerAttr(intElemType, 0));
1865 rewriter.getIntegerAttr(intElemType, 1));
1867 gt = LLVM::ICmpOp::create(rewriter, loc, cmpType,
1868 LLVM::ICmpPredicate::sgt, operand, zero);
1869 lt = LLVM::ICmpOp::create(rewriter, loc, cmpType,
1870 LLVM::ICmpPredicate::slt, operand, zero);
1874 LLVM::SelectOp::create(rewriter, loc, dstType, lt, minusOne, zero);
1875 rewriter.replaceOpWithNewOp<LLVM::SelectOp>(op, dstType, gt, one,
1883 using SPIRVToLLVMConversion<spirv::VariableOp>::SPIRVToLLVMConversion;
1886 matchAndRewrite(spirv::VariableOp varOp, OpAdaptor adaptor,
1887 ConversionPatternRewriter &rewriter)
const override {
1888 auto srcType = varOp.getType();
1890 auto pointerTo = cast<spirv::PointerType>(srcType).getPointeeType();
1891 auto init = varOp.getInitializer();
1892 if (init && !pointerTo.isIntOrFloat() && !isa<VectorType>(pointerTo))
1895 auto dstType = getTypeConverter()->convertType(srcType);
1897 return rewriter.notifyMatchFailure(varOp,
"type conversion failed");
1899 Location loc = varOp.getLoc();
1902 auto elementType = getTypeConverter()->convertType(pointerTo);
1904 return rewriter.notifyMatchFailure(varOp,
"type conversion failed");
1905 rewriter.replaceOpWithNewOp<LLVM::AllocaOp>(varOp, dstType, elementType,
1909 auto elementType = getTypeConverter()->convertType(pointerTo);
1911 return rewriter.notifyMatchFailure(varOp,
"type conversion failed");
1913 LLVM::AllocaOp::create(rewriter, loc, dstType, elementType, size);
1914 LLVM::StoreOp::create(rewriter, loc, adaptor.getInitializer(), allocated);
1915 rewriter.replaceOp(varOp, allocated);
1924class BitcastConversionPattern
1927 using SPIRVToLLVMConversion<spirv::BitcastOp>::SPIRVToLLVMConversion;
1930 matchAndRewrite(spirv::BitcastOp bitcastOp, OpAdaptor adaptor,
1931 ConversionPatternRewriter &rewriter)
const override {
1932 auto dstType = getTypeConverter()->convertType(bitcastOp.getType());
1934 return rewriter.notifyMatchFailure(bitcastOp,
"type conversion failed");
1937 if (isa<LLVM::LLVMPointerType>(dstType)) {
1938 rewriter.replaceOp(bitcastOp, adaptor.getOperand());
1942 rewriter.replaceOpWithNewOp<LLVM::BitcastOp>(
1943 bitcastOp, dstType, adaptor.getOperands(), bitcastOp->getAttrs());
1954 using SPIRVToLLVMConversion<spirv::FuncOp>::SPIRVToLLVMConversion;
1957 matchAndRewrite(spirv::FuncOp funcOp, OpAdaptor adaptor,
1958 ConversionPatternRewriter &rewriter)
const override {
1962 auto funcType = funcOp.getFunctionType();
1963 TypeConverter::SignatureConversion signatureConverter(
1964 funcType.getNumInputs());
1965 auto llvmType =
static_cast<const LLVMTypeConverter *
>(getTypeConverter())
1966 ->convertFunctionSignature(
1968 false, signatureConverter);
1973 Location loc = funcOp.getLoc();
1974 StringRef name = funcOp.getName();
1975 auto newFuncOp = LLVM::LLVMFuncOp::create(rewriter, loc, name, llvmType);
1978 MLIRContext *context = funcOp.getContext();
1979 switch (funcOp.getFunctionControl()) {
1980 case spirv::FunctionControl::Inline:
1981 newFuncOp.setAlwaysInline(
true);
1983 case spirv::FunctionControl::DontInline:
1984 newFuncOp.setNoInline(
true);
1987#define DISPATCH(functionControl, llvmAttr) \
1988 case functionControl: \
1989 newFuncOp->setAttr("passthrough", ArrayAttr::get(context, {llvmAttr})); \
1992 DISPATCH(spirv::FunctionControl::Pure,
1993 StringAttr::get(context,
"readonly"));
1994 DISPATCH(spirv::FunctionControl::Const,
1995 StringAttr::get(context,
"readnone"));
2005 rewriter.inlineRegionBefore(funcOp.getBody(), newFuncOp.getBody(),
2007 if (
failed(rewriter.convertRegionTypes(
2008 &newFuncOp.getBody(), *getTypeConverter(), &signatureConverter))) {
2011 rewriter.eraseOp(funcOp);
2022 using SPIRVToLLVMConversion<spirv::ModuleOp>::SPIRVToLLVMConversion;
2025 matchAndRewrite(spirv::ModuleOp spvModuleOp, OpAdaptor adaptor,
2026 ConversionPatternRewriter &rewriter)
const override {
2029 ModuleOp::create(rewriter, spvModuleOp.getLoc(), spvModuleOp.getName());
2030 rewriter.inlineRegionBefore(spvModuleOp.getRegion(), newModuleOp.getBody());
2033 rewriter.eraseBlock(&newModuleOp.getBodyRegion().back());
2034 rewriter.eraseOp(spvModuleOp);
2043class VectorShufflePattern
2046 using SPIRVToLLVMConversion<spirv::VectorShuffleOp>::SPIRVToLLVMConversion;
2048 matchAndRewrite(spirv::VectorShuffleOp op, OpAdaptor adaptor,
2049 ConversionPatternRewriter &rewriter)
const override {
2050 Location loc = op.getLoc();
2051 auto components = adaptor.getComponents();
2052 auto vector1 = adaptor.getVector1();
2053 auto vector2 = adaptor.getVector2();
2054 int vector1Size = cast<VectorType>(vector1.getType()).getNumElements();
2055 int vector2Size = cast<VectorType>(vector2.getType()).getNumElements();
2056 if (vector1Size == vector2Size) {
2057 rewriter.replaceOpWithNewOp<LLVM::ShuffleVectorOp>(
2058 op, vector1, vector2,
2059 LLVM::convertArrayToIndices<int32_t>(components));
2063 auto dstType = getTypeConverter()->convertType(op.getType());
2065 return rewriter.notifyMatchFailure(op,
"type conversion failed");
2066 auto scalarType = cast<VectorType>(dstType).getElementType();
2067 auto componentsArray = components.getValue();
2068 auto *context = rewriter.getContext();
2069 auto llvmI32Type = IntegerType::get(context, 32);
2070 Value targetOp = LLVM::PoisonOp::create(rewriter, loc, dstType);
2071 for (
unsigned i = 0; i < componentsArray.size(); i++) {
2072 if (!isa<IntegerAttr>(componentsArray[i]))
2073 return op.emitError(
"unable to support non-constant component");
2075 int indexVal = cast<IntegerAttr>(componentsArray[i]).getInt();
2080 Value baseVector = vector1;
2081 if (indexVal >= vector1Size) {
2082 offsetVal = vector1Size;
2083 baseVector = vector2;
2086 Value dstIndex = LLVM::ConstantOp::create(
2087 rewriter, loc, llvmI32Type,
2088 rewriter.getIntegerAttr(rewriter.getI32Type(), i));
2089 Value index = LLVM::ConstantOp::create(
2090 rewriter, loc, llvmI32Type,
2091 rewriter.getIntegerAttr(rewriter.getI32Type(), indexVal - offsetVal));
2093 auto extractOp = LLVM::ExtractElementOp::create(rewriter, loc, scalarType,
2095 targetOp = LLVM::InsertElementOp::create(rewriter, loc, dstType, targetOp,
2096 extractOp, dstIndex);
2098 rewriter.replaceOp(op, targetOp);
2109 spirv::ClientAPI clientAPI) {
2126 spirv::ClientAPI clientAPI) {
2129 DirectConversionPattern<spirv::IAddOp, LLVM::AddOp>,
2130 DirectConversionPattern<spirv::IMulOp, LLVM::MulOp>,
2131 DirectConversionPattern<spirv::ISubOp, LLVM::SubOp>,
2132 DirectConversionPattern<spirv::FAddOp, LLVM::FAddOp>,
2133 DirectConversionPattern<spirv::FDivOp, LLVM::FDivOp>,
2134 DirectConversionPattern<spirv::FMulOp, LLVM::FMulOp>,
2135 DirectConversionPattern<spirv::FNegateOp, LLVM::FNegOp>,
2136 DirectConversionPattern<spirv::FRemOp, LLVM::FRemOp>,
2137 DirectConversionPattern<spirv::FSubOp, LLVM::FSubOp>,
2138 DirectConversionPattern<spirv::SDivOp, LLVM::SDivOp>,
2139 DirectConversionPattern<spirv::SRemOp, LLVM::SRemOp>,
2140 DirectConversionPattern<spirv::UDivOp, LLVM::UDivOp>,
2141 DirectConversionPattern<spirv::UModOp, LLVM::URemOp>, FModPattern,
2142 SModPattern, VectorTimesScalarPattern, SNegatePattern,
2143 ArithmeticWithOverflowPattern<spirv::IAddCarryOp,
2144 LLVM::UAddWithOverflowOp>,
2145 ArithmeticWithOverflowPattern<spirv::ISubBorrowOp,
2146 LLVM::USubWithOverflowOp>,
2149 BitFieldInsertPattern, BitFieldUExtractPattern, BitFieldSExtractPattern,
2150 DirectConversionPattern<spirv::BitCountOp, LLVM::CtPopOp>,
2151 DirectConversionPattern<spirv::BitReverseOp, LLVM::BitReverseOp>,
2152 DirectConversionPattern<spirv::BitwiseAndOp, LLVM::AndOp>,
2153 DirectConversionPattern<spirv::BitwiseOrOp, LLVM::OrOp>,
2154 DirectConversionPattern<spirv::BitwiseXorOp, LLVM::XOrOp>,
2155 NotPattern<spirv::NotOp>,
2158 BitcastConversionPattern,
2159 DirectConversionPattern<spirv::ConvertFToSOp, LLVM::FPToSIOp>,
2160 DirectConversionPattern<spirv::ConvertFToUOp, LLVM::FPToUIOp>,
2161 DirectConversionPattern<spirv::ConvertSToFOp, LLVM::SIToFPOp>,
2162 DirectConversionPattern<spirv::ConvertUToFOp, LLVM::UIToFPOp>,
2163 IndirectCastPattern<spirv::FConvertOp, LLVM::FPExtOp, LLVM::FPTruncOp>,
2164 IndirectCastPattern<spirv::SConvertOp, LLVM::SExtOp, LLVM::TruncOp>,
2165 IndirectCastPattern<spirv::UConvertOp, LLVM::ZExtOp, LLVM::TruncOp>,
2166 DirectConversionPattern<spirv::ConvertPtrToUOp, LLVM::PtrToIntOp>,
2167 DirectConversionPattern<spirv::ConvertUToPtrOp, LLVM::IntToPtrOp>,
2168 DirectConversionPattern<spirv::PtrCastToGenericOp, LLVM::AddrSpaceCastOp>,
2169 DirectConversionPattern<spirv::GenericCastToPtrOp, LLVM::AddrSpaceCastOp>,
2170 DirectConversionPattern<spirv::GenericCastToPtrExplicitOp,
2171 LLVM::AddrSpaceCastOp>,
2174 IComparePattern<spirv::IEqualOp, LLVM::ICmpPredicate::eq>,
2175 IComparePattern<spirv::INotEqualOp, LLVM::ICmpPredicate::ne>,
2176 FComparePattern<spirv::FOrdEqualOp, LLVM::FCmpPredicate::oeq>,
2177 FComparePattern<spirv::FOrdGreaterThanOp, LLVM::FCmpPredicate::ogt>,
2178 FComparePattern<spirv::FOrdGreaterThanEqualOp, LLVM::FCmpPredicate::oge>,
2179 FComparePattern<spirv::FOrdLessThanEqualOp, LLVM::FCmpPredicate::ole>,
2180 FComparePattern<spirv::FOrdLessThanOp, LLVM::FCmpPredicate::olt>,
2181 FComparePattern<spirv::FOrdNotEqualOp, LLVM::FCmpPredicate::one>,
2182 FComparePattern<spirv::FUnordEqualOp, LLVM::FCmpPredicate::ueq>,
2183 FComparePattern<spirv::FUnordGreaterThanOp, LLVM::FCmpPredicate::ugt>,
2184 FComparePattern<spirv::FUnordGreaterThanEqualOp,
2185 LLVM::FCmpPredicate::uge>,
2186 FComparePattern<spirv::FUnordLessThanEqualOp, LLVM::FCmpPredicate::ule>,
2187 FComparePattern<spirv::FUnordLessThanOp, LLVM::FCmpPredicate::ult>,
2188 FComparePattern<spirv::FUnordNotEqualOp, LLVM::FCmpPredicate::une>,
2189 FComparePattern<spirv::OrderedOp, LLVM::FCmpPredicate::ord>,
2190 FComparePattern<spirv::UnorderedOp, LLVM::FCmpPredicate::uno>,
2191 IComparePattern<spirv::SGreaterThanOp, LLVM::ICmpPredicate::sgt>,
2192 IComparePattern<spirv::SGreaterThanEqualOp, LLVM::ICmpPredicate::sge>,
2193 IComparePattern<spirv::SLessThanEqualOp, LLVM::ICmpPredicate::sle>,
2194 IComparePattern<spirv::SLessThanOp, LLVM::ICmpPredicate::slt>,
2195 IComparePattern<spirv::UGreaterThanOp, LLVM::ICmpPredicate::ugt>,
2196 IComparePattern<spirv::UGreaterThanEqualOp, LLVM::ICmpPredicate::uge>,
2197 IComparePattern<spirv::ULessThanEqualOp, LLVM::ICmpPredicate::ule>,
2198 IComparePattern<spirv::ULessThanOp, LLVM::ICmpPredicate::ult>,
2201 ConstantScalarAndVectorPattern,
2204 BranchConversionPattern, BranchConditionalConversionPattern,
2205 FunctionCallPattern, LoopPattern, SelectionPattern,
2206 ErasePattern<spirv::MergeOp>,
2209 ErasePattern<spirv::EntryPointOp>, ExecutionModePattern,
2212 DirectConversionPattern<spirv::GLCeilOp, LLVM::FCeilOp>,
2213 DirectConversionPattern<spirv::GLCosOp, LLVM::CosOp>,
2214 DirectConversionPattern<spirv::GLExpOp, LLVM::ExpOp>,
2215 DirectConversionPattern<spirv::GLExp2Op, LLVM::Exp2Op>,
2216 DirectConversionPattern<spirv::GLFAbsOp, LLVM::FAbsOp>,
2217 DirectConversionPattern<spirv::GLFloorOp, LLVM::FFloorOp>,
2218 DirectConversionPattern<spirv::GLFmaOp, LLVM::FMAOp>,
2219 ClampPattern<spirv::GLFClampOp, LLVM::MinNumOp, LLVM::MaxNumOp>,
2220 ClampPattern<spirv::GLSClampOp, LLVM::SMinOp, LLVM::SMaxOp>,
2221 ClampPattern<spirv::GLUClampOp, LLVM::UMinOp, LLVM::UMaxOp>,
2222 DirectConversionPattern<spirv::GLFMaxOp, LLVM::MaxNumOp>,
2223 DirectConversionPattern<spirv::GLFMinOp, LLVM::MinNumOp>,
2224 DirectConversionPattern<spirv::GLNMaxOp, LLVM::MaxNumOp>,
2225 DirectConversionPattern<spirv::GLNMinOp, LLVM::MinNumOp>,
2226 DirectConversionPattern<spirv::GLLogOp, LLVM::LogOp>,
2227 DirectConversionPattern<spirv::GLLog2Op, LLVM::Log2Op>,
2228 DirectConversionPattern<spirv::GLPowOp, LLVM::PowOp>,
2229 DirectConversionPattern<spirv::GLRoundOp, LLVM::RoundOp>,
2230 DirectConversionPattern<spirv::GLRoundEvenOp, LLVM::RoundEvenOp>,
2231 DirectConversionPattern<spirv::GLSinOp, LLVM::SinOp>,
2232 DirectConversionPattern<spirv::GLSinhOp, LLVM::SinhOp>,
2233 DirectConversionPattern<spirv::GLCoshOp, LLVM::CoshOp>,
2234 DirectConversionPattern<spirv::GLSMaxOp, LLVM::SMaxOp>,
2235 DirectConversionPattern<spirv::GLSMinOp, LLVM::SMinOp>,
2236 DirectConversionPattern<spirv::GLSqrtOp, LLVM::SqrtOp>,
2237 DirectConversionPattern<spirv::GLUMaxOp, LLVM::UMaxOp>,
2238 DirectConversionPattern<spirv::GLUMinOp, LLVM::UMinOp>,
2239 DirectConversionPattern<spirv::GLTruncOp, LLVM::FTruncOp>,
2240 DirectConversionPattern<spirv::GLAsinOp, LLVM::ASinOp>,
2241 DirectConversionPattern<spirv::GLAcosOp, LLVM::ACosOp>,
2242 DirectConversionPattern<spirv::GLAtanOp, LLVM::ATanOp>,
2243 DirectConversionPattern<spirv::GLTanOp, LLVM::TanOp>,
2244 DirectConversionPattern<spirv::GLTanhOp, LLVM::TanhOp>,
2245 InverseSqrtPattern, SAbsPattern, FractPattern,
2246 SignPattern<spirv::GLFSignOp,
true>,
2247 SignPattern<spirv::GLSSignOp,
false>, GLFMixPattern,
2250 DirectConversionPattern<spirv::CLCeilOp, LLVM::FCeilOp>,
2251 DirectConversionPattern<spirv::CLCosOp, LLVM::CosOp>,
2252 DirectConversionPattern<spirv::CLExpOp, LLVM::ExpOp>,
2253 DirectConversionPattern<spirv::CLExp2Op, LLVM::Exp2Op>,
2254 DirectConversionPattern<spirv::CLExp10Op, LLVM::Exp10Op>,
2255 DirectConversionPattern<spirv::CLFAbsOp, LLVM::FAbsOp>,
2256 DirectConversionPattern<spirv::CLFloorOp, LLVM::FFloorOp>,
2257 DirectConversionPattern<spirv::CLFmaOp, LLVM::FMAOp>,
2258 DirectConversionPattern<spirv::CLFMaxOp, LLVM::MaxNumOp>,
2259 DirectConversionPattern<spirv::CLFMinOp, LLVM::MinNumOp>,
2260 DirectConversionPattern<spirv::CLLogOp, LLVM::LogOp>,
2261 DirectConversionPattern<spirv::CLLog2Op, LLVM::Log2Op>,
2262 DirectConversionPattern<spirv::CLLog10Op, LLVM::Log10Op>,
2263 DirectConversionPattern<spirv::CLPowOp, LLVM::PowOp>,
2264 DirectConversionPattern<spirv::CLRintOp, LLVM::RintOp>,
2265 DirectConversionPattern<spirv::CLRoundOp, LLVM::RoundOp>,
2266 DirectConversionPattern<spirv::CLSinOp, LLVM::SinOp>,
2267 DirectConversionPattern<spirv::CLSinhOp, LLVM::SinhOp>,
2268 DirectConversionPattern<spirv::CLCoshOp, LLVM::CoshOp>,
2269 DirectConversionPattern<spirv::CLTanOp, LLVM::TanOp>,
2270 DirectConversionPattern<spirv::CLTanhOp, LLVM::TanhOp>,
2271 DirectConversionPattern<spirv::CLAsinOp, LLVM::ASinOp>,
2272 DirectConversionPattern<spirv::CLAcosOp, LLVM::ACosOp>,
2273 DirectConversionPattern<spirv::CLAtanOp, LLVM::ATanOp>,
2274 DirectConversionPattern<spirv::CLAtan2Op, LLVM::ATan2Op>,
2275 DirectConversionPattern<spirv::CLSqrtOp, LLVM::SqrtOp>,
2276 DirectConversionPattern<spirv::CLTruncOp, LLVM::FTruncOp>,
2277 DirectConversionPattern<spirv::CLCopysignOp, LLVM::CopySignOp>,
2278 DirectConversionPattern<spirv::CLFmodOp, LLVM::FRemOp>,
2279 DirectConversionPattern<spirv::CLSMaxOp, LLVM::SMaxOp>,
2280 DirectConversionPattern<spirv::CLSMinOp, LLVM::SMinOp>,
2281 DirectConversionPattern<spirv::CLUMaxOp, LLVM::UMaxOp>,
2282 DirectConversionPattern<spirv::CLUMinOp, LLVM::UMinOp>, CLMixPattern,
2285 DirectConversionPattern<spirv::LogicalAndOp, LLVM::AndOp>,
2286 DirectConversionPattern<spirv::LogicalOrOp, LLVM::OrOp>,
2287 IComparePattern<spirv::LogicalEqualOp, LLVM::ICmpPredicate::eq>,
2288 IComparePattern<spirv::LogicalNotEqualOp, LLVM::ICmpPredicate::ne>,
2289 NotPattern<spirv::LogicalNotOp>,
2292 AccessChainPattern, AddressOfPattern, LoadStorePattern<spirv::LoadOp>,
2293 LoadStorePattern<spirv::StoreOp>, VariablePattern,
2296 CompositeExtractPattern, CompositeInsertPattern,
2297 DirectConversionPattern<spirv::SelectOp, LLVM::SelectOp>,
2298 DirectConversionPattern<spirv::UndefOp, LLVM::UndefOp>,
2299 VectorShufflePattern,
2302 ShiftPattern<spirv::ShiftRightArithmeticOp, LLVM::AShrOp>,
2303 ShiftPattern<spirv::ShiftRightLogicalOp, LLVM::LShrOp>,
2304 ShiftPattern<spirv::ShiftLeftLogicalOp, LLVM::ShlOp>,
2307 ReturnPattern, ReturnValuePattern,
2313 ControlBarrierPattern<spirv::ControlBarrierOp>,
2314 ControlBarrierPattern<spirv::INTELControlBarrierArriveOp>,
2315 ControlBarrierPattern<spirv::INTELControlBarrierWaitOp>,
2318 GroupReducePattern<spirv::GroupIAddOp>,
2319 GroupReducePattern<spirv::GroupFAddOp>,
2320 GroupReducePattern<spirv::GroupFMinOp>,
2321 GroupReducePattern<spirv::GroupUMinOp>,
2322 GroupReducePattern<spirv::GroupSMinOp,
true>,
2323 GroupReducePattern<spirv::GroupFMaxOp>,
2324 GroupReducePattern<spirv::GroupUMaxOp>,
2325 GroupReducePattern<spirv::GroupSMaxOp,
true>,
2326 GroupReducePattern<spirv::GroupNonUniformIAddOp,
false,
2328 GroupReducePattern<spirv::GroupNonUniformFAddOp,
false,
2330 GroupReducePattern<spirv::GroupNonUniformIMulOp,
false,
2332 GroupReducePattern<spirv::GroupNonUniformFMulOp,
false,
2334 GroupReducePattern<spirv::GroupNonUniformSMinOp,
true,
2336 GroupReducePattern<spirv::GroupNonUniformUMinOp,
false,
2338 GroupReducePattern<spirv::GroupNonUniformFMinOp,
false,
2340 GroupReducePattern<spirv::GroupNonUniformSMaxOp,
true,
2342 GroupReducePattern<spirv::GroupNonUniformUMaxOp,
false,
2344 GroupReducePattern<spirv::GroupNonUniformFMaxOp,
false,
2346 GroupReducePattern<spirv::GroupNonUniformBitwiseAndOp,
false,
2348 GroupReducePattern<spirv::GroupNonUniformBitwiseOrOp,
false,
2350 GroupReducePattern<spirv::GroupNonUniformBitwiseXorOp,
false,
2352 GroupReducePattern<spirv::GroupNonUniformLogicalAndOp,
false,
2354 GroupReducePattern<spirv::GroupNonUniformLogicalOrOp,
false,
2356 GroupReducePattern<spirv::GroupNonUniformLogicalXorOp,
false,
2360 patterns.
add<GlobalVariablePattern>(clientAPI, patterns.
getContext(),
2363 patterns.
add<ScalePattern<spirv::GLRadiansOp>>(
2364 0.017453292519943295, patterns.
getContext(), typeConverter);
2366 patterns.
add<ScalePattern<spirv::GLDegreesOp>>(
2367 57.29577951308232, patterns.
getContext(), typeConverter);
2372 patterns.
add<FuncConversionPattern>(patterns.
getContext(), typeConverter);
2377 patterns.
add<ModuleConversionPattern>(patterns.
getContext(), typeConverter);
2388 auto spvModules =
module.getOps<spirv::ModuleOp>();
2389 for (
auto spvModule : spvModules) {
2390 spvModule.walk([&](spirv::GlobalVariableOp op) {
2391 IntegerAttr descriptorSet =
2393 IntegerAttr binding = op->getAttrOfType<IntegerAttr>(
kBinding);
2396 if (descriptorSet && binding) {
2399 auto moduleAndName =
2400 spvModule.getName().has_value()
2401 ? spvModule.getName()->str() +
"_" + op.getSymName().str()
2402 : op.getSymName().str();
2404 llvm::formatv(
"{0}_descriptor_set{1}_binding{2}", moduleAndName,
2405 std::to_string(descriptorSet.getInt()),
2406 std::to_string(binding.getInt()));
2407 auto nameAttr = StringAttr::get(op->getContext(), name);
2412 op.emitError(
"unable to replace all symbol uses for ") << name;
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 constexpr StringRef kDescriptorSet
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 constexpr StringRef kBinding
Hook for descriptor set and binding number encoding.
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...
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.