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());
654class ExecutionModePattern
657 using SPIRVToLLVMConversion<spirv::ExecutionModeOp>::SPIRVToLLVMConversion;
660 matchAndRewrite(spirv::ExecutionModeOp op, OpAdaptor adaptor,
661 ConversionPatternRewriter &rewriter)
const override {
665 ModuleOp module = op->getParentOfType<ModuleOp>();
666 spirv::ExecutionModeAttr executionModeAttr = op.getExecutionModeAttr();
667 std::string moduleName;
668 if (module.getName().has_value())
669 moduleName =
"_" +
module.getName()->str();
672 std::string executionModeInfoName = llvm::formatv(
673 "__spv_{0}_{1}_execution_mode_info_{2}", moduleName, op.getFn().str(),
674 static_cast<uint32_t
>(executionModeAttr.getValue()));
676 MLIRContext *context = rewriter.getContext();
677 OpBuilder::InsertionGuard guard(rewriter);
678 rewriter.setInsertionPointToStart(module.getBody());
685 auto llvmI32Type = IntegerType::get(context, 32);
686 SmallVector<Type, 2> fields;
687 fields.push_back(llvmI32Type);
689 if (!values.empty()) {
690 auto arrayType = LLVM::LLVMArrayType::get(llvmI32Type, values.size());
691 fields.push_back(arrayType);
693 auto structType = LLVM::LLVMStructType::getLiteral(context, fields);
696 auto global = LLVM::GlobalOp::create(
697 rewriter, UnknownLoc::get(context), structType,
true,
698 LLVM::Linkage::External, executionModeInfoName, Attribute(),
700 Location loc = global.getLoc();
701 Region ®ion = global.getInitializerRegion();
702 Block *block = rewriter.createBlock(®ion);
705 rewriter.setInsertionPointToStart(block);
706 Value structValue = LLVM::PoisonOp::create(rewriter, loc, structType);
707 Value executionMode = LLVM::ConstantOp::create(
708 rewriter, loc, llvmI32Type,
709 rewriter.getI32IntegerAttr(
710 static_cast<uint32_t
>(executionModeAttr.getValue())));
711 SmallVector<int64_t> position{0};
712 structValue = LLVM::InsertValueOp::create(rewriter, loc, structValue,
713 executionMode, position);
716 for (
unsigned i = 0, e = values.size(); i < e; ++i) {
717 auto attr = values.getValue()[i];
718 Value entry = LLVM::ConstantOp::create(rewriter, loc, llvmI32Type, attr);
719 structValue = LLVM::InsertValueOp::create(
720 rewriter, loc, structValue, entry, ArrayRef<int64_t>({1, i}));
722 LLVM::ReturnOp::create(rewriter, loc, ArrayRef<Value>({structValue}));
723 rewriter.eraseOp(op);
732class GlobalVariablePattern
735 template <
typename... Args>
736 GlobalVariablePattern(spirv::ClientAPI clientAPI, Args &&...args)
737 : SPIRVToLLVMConversion<spirv::GlobalVariableOp>(
738 std::forward<Args>(args)...),
739 clientAPI(clientAPI) {}
742 matchAndRewrite(spirv::GlobalVariableOp op, OpAdaptor adaptor,
743 ConversionPatternRewriter &rewriter)
const override {
746 if (op.getInitializer())
749 auto srcType = cast<spirv::PointerType>(op.getType());
750 auto dstType = getTypeConverter()->convertType(srcType.getPointeeType());
752 return rewriter.notifyMatchFailure(op,
"type conversion failed");
757 auto storageClass = srcType.getStorageClass();
758 switch (storageClass) {
759 case spirv::StorageClass::Input:
760 case spirv::StorageClass::Private:
761 case spirv::StorageClass::Output:
762 case spirv::StorageClass::StorageBuffer:
763 case spirv::StorageClass::UniformConstant:
772 bool isConstant = (storageClass == spirv::StorageClass::Input) ||
773 (storageClass == spirv::StorageClass::UniformConstant);
779 auto linkage = storageClass == spirv::StorageClass::Private
780 ? LLVM::Linkage::Private
781 : LLVM::Linkage::External;
782 StringAttr locationAttrName = op.getLocationAttrName();
783 IntegerAttr locationAttr = op.getLocationAttr();
784 auto newGlobalOp = rewriter.replaceOpWithNewOp<LLVM::GlobalOp>(
785 op, dstType, isConstant, linkage, op.getSymName(), Attribute(),
790 newGlobalOp->setAttr(locationAttrName, locationAttr);
796 spirv::ClientAPI clientAPI;
801template <
typename SPIRVOp,
typename LLVMExtOp,
typename LLVMTruncOp>
804 using SPIRVToLLVMConversion<SPIRVOp>::SPIRVToLLVMConversion;
807 matchAndRewrite(SPIRVOp op,
typename SPIRVOp::Adaptor adaptor,
808 ConversionPatternRewriter &rewriter)
const override {
810 Type fromType = op.getOperand().getType();
811 Type toType = op.getType();
813 auto dstType = this->getTypeConverter()->convertType(toType);
815 return rewriter.notifyMatchFailure(op,
"type conversion failed");
818 rewriter.template replaceOpWithNewOp<LLVMExtOp>(op, dstType,
819 adaptor.getOperands());
823 rewriter.template replaceOpWithNewOp<LLVMTruncOp>(op, dstType,
824 adaptor.getOperands());
831class FunctionCallPattern
834 using SPIRVToLLVMConversion<spirv::FunctionCallOp>::SPIRVToLLVMConversion;
837 matchAndRewrite(spirv::FunctionCallOp callOp, OpAdaptor adaptor,
838 ConversionPatternRewriter &rewriter)
const override {
839 if (callOp.getNumResults() == 0) {
840 auto newOp = rewriter.replaceOpWithNewOp<LLVM::CallOp>(
841 callOp,
TypeRange(), adaptor.getOperands(), callOp->getAttrs());
842 newOp.getProperties().operandSegmentSizes = {
843 static_cast<int32_t
>(adaptor.getOperands().size()), 0};
844 newOp.getProperties().op_bundle_sizes = rewriter.getDenseI32ArrayAttr({});
849 auto dstType = getTypeConverter()->convertType(callOp.getType(0));
851 return rewriter.notifyMatchFailure(callOp,
"type conversion failed");
852 auto newOp = rewriter.replaceOpWithNewOp<LLVM::CallOp>(
853 callOp, dstType, adaptor.getOperands(), callOp->getAttrs());
854 newOp.getProperties().operandSegmentSizes = {
855 static_cast<int32_t
>(adaptor.getOperands().size()), 0};
856 newOp.getProperties().op_bundle_sizes = rewriter.getDenseI32ArrayAttr({});
862template <
typename SPIRVOp, LLVM::FCmpPredicate predicate>
865 using SPIRVToLLVMConversion<SPIRVOp>::SPIRVToLLVMConversion;
868 matchAndRewrite(SPIRVOp op,
typename SPIRVOp::Adaptor adaptor,
869 ConversionPatternRewriter &rewriter)
const override {
871 auto dstType = this->getTypeConverter()->convertType(op.getType());
873 return rewriter.notifyMatchFailure(op,
"type conversion failed");
875 rewriter.template replaceOpWithNewOp<LLVM::FCmpOp>(
876 op, dstType, predicate, op.getOperand1(), op.getOperand2());
882template <
typename SPIRVOp, LLVM::ICmpPredicate predicate>
885 using SPIRVToLLVMConversion<SPIRVOp>::SPIRVToLLVMConversion;
888 matchAndRewrite(SPIRVOp op,
typename SPIRVOp::Adaptor adaptor,
889 ConversionPatternRewriter &rewriter)
const override {
891 auto dstType = this->getTypeConverter()->convertType(op.getType());
893 return rewriter.notifyMatchFailure(op,
"type conversion failed");
895 rewriter.template replaceOpWithNewOp<LLVM::ICmpOp>(
896 op, dstType, predicate, op.getOperand1(), op.getOperand2());
901class InverseSqrtPattern
904 using SPIRVToLLVMConversion<spirv::GLInverseSqrtOp>::SPIRVToLLVMConversion;
907 matchAndRewrite(spirv::GLInverseSqrtOp op, OpAdaptor adaptor,
908 ConversionPatternRewriter &rewriter)
const override {
909 auto srcType = op.getType();
910 auto dstType = getTypeConverter()->convertType(srcType);
912 return rewriter.notifyMatchFailure(op,
"type conversion failed");
914 Location loc = op.getLoc();
916 Value sqrt = LLVM::SqrtOp::create(rewriter, loc, dstType, op.getOperand());
917 rewriter.replaceOpWithNewOp<LLVM::FDivOp>(op, dstType, one, sqrt);
925 using SPIRVToLLVMConversion<spirv::SNegateOp>::SPIRVToLLVMConversion;
928 matchAndRewrite(spirv::SNegateOp op, OpAdaptor adaptor,
929 ConversionPatternRewriter &rewriter)
const override {
930 Type srcType = op.getType();
931 Type dstType = getTypeConverter()->convertType(srcType);
933 return rewriter.notifyMatchFailure(op,
"type conversion failed");
935 Location loc = op.getLoc();
936 IntegerAttr zeroAttr = rewriter.getIntegerAttr(
940 rewriter.replaceOpWithNewOp<LLVM::SubOp>(op, dstType, zero,
941 adaptor.getOperand());
948template <
typename SPIRVOp,
typename LLVMMinOp,
typename LLVMMaxOp>
951 using SPIRVToLLVMConversion<SPIRVOp>::SPIRVToLLVMConversion;
954 matchAndRewrite(SPIRVOp op,
typename SPIRVOp::Adaptor adaptor,
955 ConversionPatternRewriter &rewriter)
const override {
956 Type dstType = this->getTypeConverter()->convertType(op.getType());
958 return rewriter.notifyMatchFailure(op,
"type conversion failed");
960 Location loc = op.getLoc();
961 Value
max = LLVMMaxOp::create(rewriter, loc, dstType, adaptor.getX(),
963 rewriter.template replaceOpWithNewOp<LLVMMinOp>(op, dstType,
max,
970template <
typename SPIRVOp>
973 using SPIRVToLLVMConversion<SPIRVOp>::SPIRVToLLVMConversion;
976 matchAndRewrite(SPIRVOp op,
typename SPIRVOp::Adaptor adaptor,
977 ConversionPatternRewriter &rewriter)
const override {
978 if (!op.getMemoryAccess()) {
980 *this->getTypeConverter(), 0,
984 auto memoryAccess = *op.getMemoryAccess();
985 switch (memoryAccess) {
986 case spirv::MemoryAccess::Aligned:
987 case spirv::MemoryAccess::None:
988 case spirv::MemoryAccess::Nontemporal:
989 case spirv::MemoryAccess::Volatile: {
991 memoryAccess == spirv::MemoryAccess::Aligned ? *op.getAlignment() : 0;
992 bool isNonTemporal = memoryAccess == spirv::MemoryAccess::Nontemporal;
993 bool isVolatile = memoryAccess == spirv::MemoryAccess::Volatile;
995 *this->getTypeConverter(), alignment,
996 isVolatile, isNonTemporal);
1006template <
typename SPIRVOp>
1009 using SPIRVToLLVMConversion<SPIRVOp>::SPIRVToLLVMConversion;
1012 matchAndRewrite(SPIRVOp notOp,
typename SPIRVOp::Adaptor adaptor,
1013 ConversionPatternRewriter &rewriter)
const override {
1014 auto srcType = notOp.getType();
1015 auto dstType = this->getTypeConverter()->convertType(srcType);
1017 return rewriter.notifyMatchFailure(notOp,
"type conversion failed");
1019 Location loc = notOp.getLoc();
1021 rewriter.template replaceOpWithNewOp<LLVM::XOrOp>(notOp, dstType,
1022 notOp.getOperand(), mask);
1028template <
typename SPIRVOp>
1031 using SPIRVToLLVMConversion<SPIRVOp>::SPIRVToLLVMConversion;
1034 matchAndRewrite(SPIRVOp op,
typename SPIRVOp::Adaptor adaptor,
1035 ConversionPatternRewriter &rewriter)
const override {
1036 rewriter.eraseOp(op);
1043 using SPIRVToLLVMConversion<spirv::ReturnOp>::SPIRVToLLVMConversion;
1046 matchAndRewrite(spirv::ReturnOp returnOp, OpAdaptor adaptor,
1047 ConversionPatternRewriter &rewriter)
const override {
1048 rewriter.replaceOpWithNewOp<LLVM::ReturnOp>(returnOp, ArrayRef<Type>(),
1056 using SPIRVToLLVMConversion<spirv::ReturnValueOp>::SPIRVToLLVMConversion;
1059 matchAndRewrite(spirv::ReturnValueOp returnValueOp, OpAdaptor adaptor,
1060 ConversionPatternRewriter &rewriter)
const override {
1061 rewriter.replaceOpWithNewOp<LLVM::ReturnOp>(returnValueOp, ArrayRef<Type>(),
1062 adaptor.getOperands());
1069 using SPIRVToLLVMConversion<spirv::UnreachableOp>::SPIRVToLLVMConversion;
1072 matchAndRewrite(spirv::UnreachableOp unreachableOp, OpAdaptor adaptor,
1073 ConversionPatternRewriter &rewriter)
const override {
1074 rewriter.replaceOpWithNewOp<LLVM::UnreachableOp>(unreachableOp);
1083 bool convergent =
true) {
1084 auto func = dyn_cast_or_null<LLVM::LLVMFuncOp>(
1090 func = LLVM::LLVMFuncOp::create(
1091 b, symbolTable->
getLoc(), name,
1092 LLVM::LLVMFunctionType::get(resultType, paramTypes));
1093 func.setCConv(LLVM::cconv::CConv::SPIR_FUNC);
1094 func.setConvergent(convergent);
1095 func.setNoUnwind(
true);
1096 func.setWillReturn(
true);
1101 LLVM::LLVMFuncOp
func,
1103 auto call = LLVM::CallOp::create(builder, loc,
func, args);
1104 call.setCConv(
func.getCConv());
1105 call.setConvergentAttr(
func.getConvergentAttr());
1106 call.setNoUnwindAttr(
func.getNoUnwindAttr());
1107 call.setWillReturnAttr(
func.getWillReturnAttr());
1111template <
typename BarrierOpTy>
1114 using OpAdaptor =
typename SPIRVToLLVMConversion<BarrierOpTy>::OpAdaptor;
1116 using SPIRVToLLVMConversion<BarrierOpTy>::SPIRVToLLVMConversion;
1118 static constexpr StringRef getFuncName();
1121 matchAndRewrite(BarrierOpTy controlBarrierOp, OpAdaptor adaptor,
1122 ConversionPatternRewriter &rewriter)
const override {
1123 constexpr StringRef funcName = getFuncName();
1124 Operation *symbolTable =
1125 controlBarrierOp->template getParentWithTrait<OpTrait::SymbolTable>();
1127 Type i32 = rewriter.getI32Type();
1129 Type voidTy = rewriter.getType<LLVM::LLVMVoidType>();
1130 LLVM::LLVMFuncOp func =
1133 Location loc = controlBarrierOp->getLoc();
1134 Value execution = LLVM::ConstantOp::create(
1135 rewriter, loc, i32,
static_cast<int32_t
>(adaptor.getExecutionScope()));
1136 Value memory = LLVM::ConstantOp::create(
1137 rewriter, loc, i32,
static_cast<int32_t
>(adaptor.getMemoryScope()));
1138 Value semantics = LLVM::ConstantOp::create(
1139 rewriter, loc, i32,
static_cast<int32_t
>(adaptor.getMemorySemantics()));
1142 {execution, memory, semantics});
1144 rewriter.replaceOp(controlBarrierOp, call);
1151StringRef getTypeMangling(
Type type,
bool isSigned) {
1153 .Case([](Float16Type) {
return "Dh"; })
1154 .Case([](Float32Type) {
return "f"; })
1155 .Case([](Float64Type) {
return "d"; })
1156 .Case([isSigned](IntegerType intTy) {
1157 switch (intTy.getWidth()) {
1161 return (isSigned) ?
"a" :
"c";
1163 return (isSigned) ?
"s" :
"t";
1165 return (isSigned) ?
"i" :
"j";
1167 return (isSigned) ?
"l" :
"m";
1169 llvm_unreachable(
"Unsupported integer width");
1172 .DefaultUnreachable(
"No mangling defined");
1175template <
typename ReduceOp>
1176constexpr StringLiteral getGroupFuncName();
1179constexpr StringLiteral getGroupFuncName<spirv::GroupIAddOp>() {
1180 return "_Z17__spirv_GroupIAddii";
1183constexpr StringLiteral getGroupFuncName<spirv::GroupFAddOp>() {
1184 return "_Z17__spirv_GroupFAddii";
1187constexpr StringLiteral getGroupFuncName<spirv::GroupSMinOp>() {
1188 return "_Z17__spirv_GroupSMinii";
1191constexpr StringLiteral getGroupFuncName<spirv::GroupUMinOp>() {
1192 return "_Z17__spirv_GroupUMinii";
1195constexpr StringLiteral getGroupFuncName<spirv::GroupFMinOp>() {
1196 return "_Z17__spirv_GroupFMinii";
1199constexpr StringLiteral getGroupFuncName<spirv::GroupSMaxOp>() {
1200 return "_Z17__spirv_GroupSMaxii";
1203constexpr StringLiteral getGroupFuncName<spirv::GroupUMaxOp>() {
1204 return "_Z17__spirv_GroupUMaxii";
1207constexpr StringLiteral getGroupFuncName<spirv::GroupFMaxOp>() {
1208 return "_Z17__spirv_GroupFMaxii";
1211constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformIAddOp>() {
1212 return "_Z27__spirv_GroupNonUniformIAddii";
1215constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformFAddOp>() {
1216 return "_Z27__spirv_GroupNonUniformFAddii";
1219constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformIMulOp>() {
1220 return "_Z27__spirv_GroupNonUniformIMulii";
1223constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformFMulOp>() {
1224 return "_Z27__spirv_GroupNonUniformFMulii";
1227constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformSMinOp>() {
1228 return "_Z27__spirv_GroupNonUniformSMinii";
1231constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformUMinOp>() {
1232 return "_Z27__spirv_GroupNonUniformUMinii";
1235constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformFMinOp>() {
1236 return "_Z27__spirv_GroupNonUniformFMinii";
1239constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformSMaxOp>() {
1240 return "_Z27__spirv_GroupNonUniformSMaxii";
1243constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformUMaxOp>() {
1244 return "_Z27__spirv_GroupNonUniformUMaxii";
1247constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformFMaxOp>() {
1248 return "_Z27__spirv_GroupNonUniformFMaxii";
1251constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformBitwiseAndOp>() {
1252 return "_Z33__spirv_GroupNonUniformBitwiseAndii";
1255constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformBitwiseOrOp>() {
1256 return "_Z32__spirv_GroupNonUniformBitwiseOrii";
1259constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformBitwiseXorOp>() {
1260 return "_Z33__spirv_GroupNonUniformBitwiseXorii";
1263constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformLogicalAndOp>() {
1264 return "_Z33__spirv_GroupNonUniformLogicalAndii";
1267constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformLogicalOrOp>() {
1268 return "_Z32__spirv_GroupNonUniformLogicalOrii";
1271constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformLogicalXorOp>() {
1272 return "_Z33__spirv_GroupNonUniformLogicalXorii";
1276template <
typename ReduceOp,
bool Signed = false,
bool NonUniform = false>
1279 using SPIRVToLLVMConversion<ReduceOp>::SPIRVToLLVMConversion;
1282 matchAndRewrite(ReduceOp op,
typename ReduceOp::Adaptor adaptor,
1283 ConversionPatternRewriter &rewriter)
const override {
1285 Type retTy = op.getResult().getType();
1289 SmallString<36> funcName = getGroupFuncName<ReduceOp>();
1290 funcName += getTypeMangling(retTy,
false);
1292 Type i32Ty = rewriter.getI32Type();
1293 SmallVector<Type> paramTypes{i32Ty, i32Ty, retTy};
1294 if constexpr (NonUniform) {
1295 if (adaptor.getClusterSize()) {
1297 paramTypes.push_back(i32Ty);
1301 Operation *symbolTable =
1302 op->template getParentWithTrait<OpTrait::SymbolTable>();
1304 LLVM::LLVMFuncOp func =
1307 Location loc = op.getLoc();
1308 Value scope = LLVM::ConstantOp::create(
1309 rewriter, loc, i32Ty,
1310 static_cast<int32_t
>(adaptor.getExecutionScope()));
1311 Value groupOp = LLVM::ConstantOp::create(
1312 rewriter, loc, i32Ty,
1313 static_cast<int32_t
>(adaptor.getGroupOperation()));
1314 SmallVector<Value> operands{scope, groupOp};
1315 operands.append(adaptor.getOperands().begin(), adaptor.getOperands().end());
1318 rewriter.replaceOp(op, call);
1325ControlBarrierPattern<spirv::ControlBarrierOp>::getFuncName() {
1326 return "_Z22__spirv_ControlBarrieriii";
1331ControlBarrierPattern<spirv::INTELControlBarrierArriveOp>::getFuncName() {
1332 return "_Z33__spirv_ControlBarrierArriveINTELiii";
1337ControlBarrierPattern<spirv::INTELControlBarrierWaitOp>::getFuncName() {
1338 return "_Z31__spirv_ControlBarrierWaitINTELiii";
1391 using SPIRVToLLVMConversion<spirv::LoopOp>::SPIRVToLLVMConversion;
1394 matchAndRewrite(spirv::LoopOp loopOp, OpAdaptor adaptor,
1395 ConversionPatternRewriter &rewriter)
const override {
1397 if (loopOp.getLoopControl() != spirv::LoopControl::None)
1401 if (loopOp.getBody().empty()) {
1402 rewriter.eraseOp(loopOp);
1406 Location loc = loopOp.getLoc();
1410 Block *currentBlock = rewriter.getBlock();
1412 Block *endBlock = rewriter.splitBlock(currentBlock, position);
1416 Block *entryBlock = loopOp.getEntryBlock();
1418 auto brOp = dyn_cast<spirv::BranchOp>(entryBlock->
getOperations().front());
1421 Block *headerBlock = loopOp.getHeaderBlock();
1422 rewriter.setInsertionPointToEnd(currentBlock);
1423 LLVM::BrOp::create(rewriter, loc, brOp.getBlockArguments(), headerBlock);
1424 rewriter.eraseBlock(entryBlock);
1427 Block *mergeBlock = loopOp.getMergeBlock();
1430 rewriter.setInsertionPointToEnd(mergeBlock);
1431 LLVM::BrOp::create(rewriter, loc, terminatorOperands, endBlock);
1433 rewriter.inlineRegionBefore(loopOp.getBody(), endBlock);
1444 using SPIRVToLLVMConversion<spirv::SelectionOp>::SPIRVToLLVMConversion;
1447 matchAndRewrite(spirv::SelectionOp op, OpAdaptor adaptor,
1448 ConversionPatternRewriter &rewriter)
const override {
1452 if (op.getSelectionControl() != spirv::SelectionControl::None)
1459 if (op.getBody().getBlocks().size() <= 2) {
1460 rewriter.eraseOp(op);
1464 Location loc = op.getLoc();
1468 auto *currentBlock = rewriter.getInsertionBlock();
1469 rewriter.setInsertionPointAfter(op);
1470 auto position = rewriter.getInsertionPoint();
1471 auto *continueBlock = rewriter.splitBlock(currentBlock, position);
1477 auto *headerBlock = op.getHeaderBlock();
1479 auto condBrOp = dyn_cast<spirv::BranchConditionalOp>(
1485 auto *mergeBlock = op.getMergeBlock();
1488 rewriter.setInsertionPointToEnd(mergeBlock);
1489 LLVM::BrOp::create(rewriter, loc, terminatorOperands, continueBlock);
1492 Block *trueBlock = condBrOp.getTrueBlock();
1493 Block *falseBlock = condBrOp.getFalseBlock();
1494 rewriter.setInsertionPointToEnd(currentBlock);
1495 LLVM::CondBrOp::create(rewriter, loc, condBrOp.getCondition(), trueBlock,
1496 condBrOp.getTrueTargetOperands(), falseBlock,
1497 condBrOp.getFalseTargetOperands());
1499 rewriter.eraseBlock(headerBlock);
1500 rewriter.inlineRegionBefore(op.getBody(), continueBlock);
1501 rewriter.replaceOp(op, continueBlock->getArguments());
1510template <
typename SPIRVOp,
typename LLVMOp>
1513 using SPIRVToLLVMConversion<SPIRVOp>::SPIRVToLLVMConversion;
1516 matchAndRewrite(SPIRVOp op,
typename SPIRVOp::Adaptor adaptor,
1517 ConversionPatternRewriter &rewriter)
const override {
1519 auto dstType = this->getTypeConverter()->convertType(op.getType());
1521 return rewriter.notifyMatchFailure(op,
"type conversion failed");
1523 Type op1Type = op.getOperand1().getType();
1524 Type op2Type = op.getOperand2().getType();
1526 if (op1Type == op2Type) {
1527 rewriter.template replaceOpWithNewOp<LLVMOp>(op, dstType,
1528 adaptor.getOperands());
1532 std::optional<uint64_t> dstTypeWidth =
1534 std::optional<uint64_t> op2TypeWidth =
1537 if (!dstTypeWidth || !op2TypeWidth)
1540 Location loc = op.getLoc();
1542 if (op2TypeWidth < dstTypeWidth) {
1545 LLVM::ZExtOp::create(rewriter, loc, dstType, adaptor.getOperand2());
1548 LLVM::SExtOp::create(rewriter, loc, dstType, adaptor.getOperand2());
1550 }
else if (op2TypeWidth == dstTypeWidth) {
1551 extended = adaptor.getOperand2();
1557 LLVMOp::create(rewriter, loc, dstType, adaptor.getOperand1(), extended);
1558 rewriter.replaceOp(op,
result);
1565 using SPIRVToLLVMConversion<spirv::GLTanOp>::SPIRVToLLVMConversion;
1568 matchAndRewrite(spirv::GLTanOp tanOp, OpAdaptor adaptor,
1569 ConversionPatternRewriter &rewriter)
const override {
1570 auto dstType = getTypeConverter()->convertType(tanOp.getType());
1572 return rewriter.notifyMatchFailure(tanOp,
"type conversion failed");
1574 rewriter.replaceOpWithNewOp<LLVM::TanOp>(tanOp, dstType,
1575 adaptor.getOperands());
1582 using SPIRVToLLVMConversion<spirv::GLTanhOp>::SPIRVToLLVMConversion;
1585 matchAndRewrite(spirv::GLTanhOp tanhOp, OpAdaptor adaptor,
1586 ConversionPatternRewriter &rewriter)
const override {
1587 auto srcType = tanhOp.getType();
1588 auto dstType = getTypeConverter()->convertType(srcType);
1590 return rewriter.notifyMatchFailure(tanhOp,
"type conversion failed");
1592 rewriter.replaceOpWithNewOp<LLVM::TanhOp>(tanhOp, dstType,
1593 adaptor.getOperands());
1603 using SPIRVToLLVMConversion<spirv::GLSAbsOp>::SPIRVToLLVMConversion;
1606 matchAndRewrite(spirv::GLSAbsOp op, OpAdaptor adaptor,
1607 ConversionPatternRewriter &rewriter)
const override {
1608 Type dstType = getTypeConverter()->convertType(op.getType());
1610 return rewriter.notifyMatchFailure(op,
"type conversion failed");
1612 rewriter.replaceOpWithNewOp<LLVM::AbsOp>(op, dstType, adaptor.getOperand(),
1621 using SPIRVToLLVMConversion<spirv::GLFractOp>::SPIRVToLLVMConversion;
1624 matchAndRewrite(spirv::GLFractOp op, OpAdaptor adaptor,
1625 ConversionPatternRewriter &rewriter)
const override {
1626 Type dstType = getTypeConverter()->convertType(op.getType());
1628 return rewriter.notifyMatchFailure(op,
"type conversion failed");
1630 Location loc = op.getLoc();
1631 Value operand = adaptor.getOperand();
1632 Value floored = LLVM::FFloorOp::create(rewriter, loc, dstType, operand);
1633 rewriter.replaceOpWithNewOp<LLVM::FSubOp>(op, dstType, operand, floored);
1642 using SPIRVToLLVMConversion<spirv::GLFMixOp>::SPIRVToLLVMConversion;
1645 matchAndRewrite(spirv::GLFMixOp op, OpAdaptor adaptor,
1646 ConversionPatternRewriter &rewriter)
const override {
1647 Type dstType = getTypeConverter()->convertType(op.getType());
1649 return rewriter.notifyMatchFailure(op,
"type conversion failed");
1651 Location loc = op.getLoc();
1652 Value x = adaptor.getX();
1653 Value y = adaptor.getY();
1654 Value a = adaptor.getA();
1656 Value oneMinusA = LLVM::FSubOp::create(rewriter, loc, dstType, one, a);
1657 Value
lhs = LLVM::FMulOp::create(rewriter, loc, dstType, x, oneMinusA);
1658 Value
rhs = LLVM::FMulOp::create(rewriter, loc, dstType, y, a);
1659 rewriter.replaceOpWithNewOp<LLVM::FAddOp>(op, dstType,
lhs,
rhs);
1668 using SPIRVToLLVMConversion<spirv::CLMixOp>::SPIRVToLLVMConversion;
1671 matchAndRewrite(spirv::CLMixOp op, OpAdaptor adaptor,
1672 ConversionPatternRewriter &rewriter)
const override {
1673 Type dstType = getTypeConverter()->convertType(op.getType());
1675 return rewriter.notifyMatchFailure(op,
"type conversion failed");
1677 Location loc = op.getLoc();
1678 Value x = adaptor.getX();
1679 Value y = adaptor.getY();
1680 Value a = adaptor.getZ();
1681 Value diff = LLVM::FSubOp::create(rewriter, loc, dstType, y, x);
1682 rewriter.replaceOpWithNewOp<LLVM::FMAOp>(op, dstType, a, diff, x);
1689template <
typename SPIRVOp>
1692 template <
typename... Args>
1693 ScalePattern(
double scale, Args &&...args)
1694 : SPIRVToLLVMConversion<SPIRVOp>(std::forward<Args>(args)...),
1698 matchAndRewrite(SPIRVOp op,
typename SPIRVOp::Adaptor adaptor,
1699 ConversionPatternRewriter &rewriter)
const override {
1700 Type srcType = op.getType();
1701 Type dstType = this->getTypeConverter()->convertType(srcType);
1703 return rewriter.notifyMatchFailure(op,
"type conversion failed");
1705 Location loc = op.getLoc();
1707 rewriter.replaceOpWithNewOp<LLVM::FMulOp>(op, dstType, adaptor.getOperand(),
1718 using SPIRVToLLVMConversion<spirv::VariableOp>::SPIRVToLLVMConversion;
1721 matchAndRewrite(spirv::VariableOp varOp, OpAdaptor adaptor,
1722 ConversionPatternRewriter &rewriter)
const override {
1723 auto srcType = varOp.getType();
1725 auto pointerTo = cast<spirv::PointerType>(srcType).getPointeeType();
1726 auto init = varOp.getInitializer();
1727 if (init && !pointerTo.isIntOrFloat() && !isa<VectorType>(pointerTo))
1730 auto dstType = getTypeConverter()->convertType(srcType);
1732 return rewriter.notifyMatchFailure(varOp,
"type conversion failed");
1734 Location loc = varOp.getLoc();
1737 auto elementType = getTypeConverter()->convertType(pointerTo);
1739 return rewriter.notifyMatchFailure(varOp,
"type conversion failed");
1740 rewriter.replaceOpWithNewOp<LLVM::AllocaOp>(varOp, dstType, elementType,
1744 auto elementType = getTypeConverter()->convertType(pointerTo);
1746 return rewriter.notifyMatchFailure(varOp,
"type conversion failed");
1748 LLVM::AllocaOp::create(rewriter, loc, dstType, elementType, size);
1749 LLVM::StoreOp::create(rewriter, loc, adaptor.getInitializer(), allocated);
1750 rewriter.replaceOp(varOp, allocated);
1759class BitcastConversionPattern
1762 using SPIRVToLLVMConversion<spirv::BitcastOp>::SPIRVToLLVMConversion;
1765 matchAndRewrite(spirv::BitcastOp bitcastOp, OpAdaptor adaptor,
1766 ConversionPatternRewriter &rewriter)
const override {
1767 auto dstType = getTypeConverter()->convertType(bitcastOp.getType());
1769 return rewriter.notifyMatchFailure(bitcastOp,
"type conversion failed");
1772 if (isa<LLVM::LLVMPointerType>(dstType)) {
1773 rewriter.replaceOp(bitcastOp, adaptor.getOperand());
1777 rewriter.replaceOpWithNewOp<LLVM::BitcastOp>(
1778 bitcastOp, dstType, adaptor.getOperands(), bitcastOp->getAttrs());
1789 using SPIRVToLLVMConversion<spirv::FuncOp>::SPIRVToLLVMConversion;
1792 matchAndRewrite(spirv::FuncOp funcOp, OpAdaptor adaptor,
1793 ConversionPatternRewriter &rewriter)
const override {
1797 auto funcType = funcOp.getFunctionType();
1798 TypeConverter::SignatureConversion signatureConverter(
1799 funcType.getNumInputs());
1800 auto llvmType =
static_cast<const LLVMTypeConverter *
>(getTypeConverter())
1801 ->convertFunctionSignature(
1803 false, signatureConverter);
1808 Location loc = funcOp.getLoc();
1809 StringRef name = funcOp.getName();
1810 auto newFuncOp = LLVM::LLVMFuncOp::create(rewriter, loc, name, llvmType);
1813 MLIRContext *context = funcOp.getContext();
1814 switch (funcOp.getFunctionControl()) {
1815 case spirv::FunctionControl::Inline:
1816 newFuncOp.setAlwaysInline(
true);
1818 case spirv::FunctionControl::DontInline:
1819 newFuncOp.setNoInline(
true);
1822#define DISPATCH(functionControl, llvmAttr) \
1823 case functionControl: \
1824 newFuncOp->setAttr("passthrough", ArrayAttr::get(context, {llvmAttr})); \
1827 DISPATCH(spirv::FunctionControl::Pure,
1828 StringAttr::get(context,
"readonly"));
1829 DISPATCH(spirv::FunctionControl::Const,
1830 StringAttr::get(context,
"readnone"));
1840 rewriter.inlineRegionBefore(funcOp.getBody(), newFuncOp.getBody(),
1842 if (
failed(rewriter.convertRegionTypes(
1843 &newFuncOp.getBody(), *getTypeConverter(), &signatureConverter))) {
1846 rewriter.eraseOp(funcOp);
1857 using SPIRVToLLVMConversion<spirv::ModuleOp>::SPIRVToLLVMConversion;
1860 matchAndRewrite(spirv::ModuleOp spvModuleOp, OpAdaptor adaptor,
1861 ConversionPatternRewriter &rewriter)
const override {
1864 ModuleOp::create(rewriter, spvModuleOp.getLoc(), spvModuleOp.getName());
1865 rewriter.inlineRegionBefore(spvModuleOp.getRegion(), newModuleOp.getBody());
1868 rewriter.eraseBlock(&newModuleOp.getBodyRegion().back());
1869 rewriter.eraseOp(spvModuleOp);
1878class VectorShufflePattern
1881 using SPIRVToLLVMConversion<spirv::VectorShuffleOp>::SPIRVToLLVMConversion;
1883 matchAndRewrite(spirv::VectorShuffleOp op, OpAdaptor adaptor,
1884 ConversionPatternRewriter &rewriter)
const override {
1885 Location loc = op.getLoc();
1886 auto components = adaptor.getComponents();
1887 auto vector1 = adaptor.getVector1();
1888 auto vector2 = adaptor.getVector2();
1889 int vector1Size = cast<VectorType>(vector1.getType()).getNumElements();
1890 int vector2Size = cast<VectorType>(vector2.getType()).getNumElements();
1891 if (vector1Size == vector2Size) {
1892 rewriter.replaceOpWithNewOp<LLVM::ShuffleVectorOp>(
1893 op, vector1, vector2,
1894 LLVM::convertArrayToIndices<int32_t>(components));
1898 auto dstType = getTypeConverter()->convertType(op.getType());
1900 return rewriter.notifyMatchFailure(op,
"type conversion failed");
1901 auto scalarType = cast<VectorType>(dstType).getElementType();
1902 auto componentsArray = components.getValue();
1903 auto *context = rewriter.getContext();
1904 auto llvmI32Type = IntegerType::get(context, 32);
1905 Value targetOp = LLVM::PoisonOp::create(rewriter, loc, dstType);
1906 for (
unsigned i = 0; i < componentsArray.size(); i++) {
1907 if (!isa<IntegerAttr>(componentsArray[i]))
1908 return op.emitError(
"unable to support non-constant component");
1910 int indexVal = cast<IntegerAttr>(componentsArray[i]).getInt();
1915 Value baseVector = vector1;
1916 if (indexVal >= vector1Size) {
1917 offsetVal = vector1Size;
1918 baseVector = vector2;
1921 Value dstIndex = LLVM::ConstantOp::create(
1922 rewriter, loc, llvmI32Type,
1923 rewriter.getIntegerAttr(rewriter.getI32Type(), i));
1924 Value index = LLVM::ConstantOp::create(
1925 rewriter, loc, llvmI32Type,
1926 rewriter.getIntegerAttr(rewriter.getI32Type(), indexVal - offsetVal));
1928 auto extractOp = LLVM::ExtractElementOp::create(rewriter, loc, scalarType,
1930 targetOp = LLVM::InsertElementOp::create(rewriter, loc, dstType, targetOp,
1931 extractOp, dstIndex);
1933 rewriter.replaceOp(op, targetOp);
1944 spirv::ClientAPI clientAPI) {
1961 spirv::ClientAPI clientAPI) {
1964 DirectConversionPattern<spirv::IAddOp, LLVM::AddOp>,
1965 DirectConversionPattern<spirv::IMulOp, LLVM::MulOp>,
1966 DirectConversionPattern<spirv::ISubOp, LLVM::SubOp>,
1967 DirectConversionPattern<spirv::FAddOp, LLVM::FAddOp>,
1968 DirectConversionPattern<spirv::FDivOp, LLVM::FDivOp>,
1969 DirectConversionPattern<spirv::FMulOp, LLVM::FMulOp>,
1970 DirectConversionPattern<spirv::FNegateOp, LLVM::FNegOp>,
1971 DirectConversionPattern<spirv::FRemOp, LLVM::FRemOp>,
1972 DirectConversionPattern<spirv::FSubOp, LLVM::FSubOp>,
1973 DirectConversionPattern<spirv::SDivOp, LLVM::SDivOp>,
1974 DirectConversionPattern<spirv::SRemOp, LLVM::SRemOp>,
1975 DirectConversionPattern<spirv::UDivOp, LLVM::UDivOp>,
1976 DirectConversionPattern<spirv::UModOp, LLVM::URemOp>, SNegatePattern,
1979 BitFieldInsertPattern, BitFieldUExtractPattern, BitFieldSExtractPattern,
1980 DirectConversionPattern<spirv::BitCountOp, LLVM::CtPopOp>,
1981 DirectConversionPattern<spirv::BitReverseOp, LLVM::BitReverseOp>,
1982 DirectConversionPattern<spirv::BitwiseAndOp, LLVM::AndOp>,
1983 DirectConversionPattern<spirv::BitwiseOrOp, LLVM::OrOp>,
1984 DirectConversionPattern<spirv::BitwiseXorOp, LLVM::XOrOp>,
1985 NotPattern<spirv::NotOp>,
1988 BitcastConversionPattern,
1989 DirectConversionPattern<spirv::ConvertFToSOp, LLVM::FPToSIOp>,
1990 DirectConversionPattern<spirv::ConvertFToUOp, LLVM::FPToUIOp>,
1991 DirectConversionPattern<spirv::ConvertSToFOp, LLVM::SIToFPOp>,
1992 DirectConversionPattern<spirv::ConvertUToFOp, LLVM::UIToFPOp>,
1993 IndirectCastPattern<spirv::FConvertOp, LLVM::FPExtOp, LLVM::FPTruncOp>,
1994 IndirectCastPattern<spirv::SConvertOp, LLVM::SExtOp, LLVM::TruncOp>,
1995 IndirectCastPattern<spirv::UConvertOp, LLVM::ZExtOp, LLVM::TruncOp>,
1996 DirectConversionPattern<spirv::ConvertPtrToUOp, LLVM::PtrToIntOp>,
1997 DirectConversionPattern<spirv::ConvertUToPtrOp, LLVM::IntToPtrOp>,
1998 DirectConversionPattern<spirv::PtrCastToGenericOp, LLVM::AddrSpaceCastOp>,
1999 DirectConversionPattern<spirv::GenericCastToPtrOp, LLVM::AddrSpaceCastOp>,
2000 DirectConversionPattern<spirv::GenericCastToPtrExplicitOp,
2001 LLVM::AddrSpaceCastOp>,
2004 IComparePattern<spirv::IEqualOp, LLVM::ICmpPredicate::eq>,
2005 IComparePattern<spirv::INotEqualOp, LLVM::ICmpPredicate::ne>,
2006 FComparePattern<spirv::FOrdEqualOp, LLVM::FCmpPredicate::oeq>,
2007 FComparePattern<spirv::FOrdGreaterThanOp, LLVM::FCmpPredicate::ogt>,
2008 FComparePattern<spirv::FOrdGreaterThanEqualOp, LLVM::FCmpPredicate::oge>,
2009 FComparePattern<spirv::FOrdLessThanEqualOp, LLVM::FCmpPredicate::ole>,
2010 FComparePattern<spirv::FOrdLessThanOp, LLVM::FCmpPredicate::olt>,
2011 FComparePattern<spirv::FOrdNotEqualOp, LLVM::FCmpPredicate::one>,
2012 FComparePattern<spirv::FUnordEqualOp, LLVM::FCmpPredicate::ueq>,
2013 FComparePattern<spirv::FUnordGreaterThanOp, LLVM::FCmpPredicate::ugt>,
2014 FComparePattern<spirv::FUnordGreaterThanEqualOp,
2015 LLVM::FCmpPredicate::uge>,
2016 FComparePattern<spirv::FUnordLessThanEqualOp, LLVM::FCmpPredicate::ule>,
2017 FComparePattern<spirv::FUnordLessThanOp, LLVM::FCmpPredicate::ult>,
2018 FComparePattern<spirv::FUnordNotEqualOp, LLVM::FCmpPredicate::une>,
2019 FComparePattern<spirv::OrderedOp, LLVM::FCmpPredicate::ord>,
2020 FComparePattern<spirv::UnorderedOp, LLVM::FCmpPredicate::uno>,
2021 IComparePattern<spirv::SGreaterThanOp, LLVM::ICmpPredicate::sgt>,
2022 IComparePattern<spirv::SGreaterThanEqualOp, LLVM::ICmpPredicate::sge>,
2023 IComparePattern<spirv::SLessThanEqualOp, LLVM::ICmpPredicate::sle>,
2024 IComparePattern<spirv::SLessThanOp, LLVM::ICmpPredicate::slt>,
2025 IComparePattern<spirv::UGreaterThanOp, LLVM::ICmpPredicate::ugt>,
2026 IComparePattern<spirv::UGreaterThanEqualOp, LLVM::ICmpPredicate::uge>,
2027 IComparePattern<spirv::ULessThanEqualOp, LLVM::ICmpPredicate::ule>,
2028 IComparePattern<spirv::ULessThanOp, LLVM::ICmpPredicate::ult>,
2031 ConstantScalarAndVectorPattern,
2034 BranchConversionPattern, BranchConditionalConversionPattern,
2035 FunctionCallPattern, LoopPattern, SelectionPattern,
2036 ErasePattern<spirv::MergeOp>,
2039 ErasePattern<spirv::EntryPointOp>, ExecutionModePattern,
2042 DirectConversionPattern<spirv::GLCeilOp, LLVM::FCeilOp>,
2043 DirectConversionPattern<spirv::GLCosOp, LLVM::CosOp>,
2044 DirectConversionPattern<spirv::GLExpOp, LLVM::ExpOp>,
2045 DirectConversionPattern<spirv::GLExp2Op, LLVM::Exp2Op>,
2046 DirectConversionPattern<spirv::GLFAbsOp, LLVM::FAbsOp>,
2047 DirectConversionPattern<spirv::GLFloorOp, LLVM::FFloorOp>,
2048 DirectConversionPattern<spirv::GLFmaOp, LLVM::FMAOp>,
2049 ClampPattern<spirv::GLFClampOp, LLVM::MinNumOp, LLVM::MaxNumOp>,
2050 ClampPattern<spirv::GLSClampOp, LLVM::SMinOp, LLVM::SMaxOp>,
2051 ClampPattern<spirv::GLUClampOp, LLVM::UMinOp, LLVM::UMaxOp>,
2052 DirectConversionPattern<spirv::GLFMaxOp, LLVM::MaxNumOp>,
2053 DirectConversionPattern<spirv::GLFMinOp, LLVM::MinNumOp>,
2054 DirectConversionPattern<spirv::GLNMaxOp, LLVM::MaxNumOp>,
2055 DirectConversionPattern<spirv::GLNMinOp, LLVM::MinNumOp>,
2056 DirectConversionPattern<spirv::GLLogOp, LLVM::LogOp>,
2057 DirectConversionPattern<spirv::GLLog2Op, LLVM::Log2Op>,
2058 DirectConversionPattern<spirv::GLPowOp, LLVM::PowOp>,
2059 DirectConversionPattern<spirv::GLRoundOp, LLVM::RoundOp>,
2060 DirectConversionPattern<spirv::GLRoundEvenOp, LLVM::RoundEvenOp>,
2061 DirectConversionPattern<spirv::GLSinOp, LLVM::SinOp>,
2062 DirectConversionPattern<spirv::GLSinhOp, LLVM::SinhOp>,
2063 DirectConversionPattern<spirv::GLCoshOp, LLVM::CoshOp>,
2064 DirectConversionPattern<spirv::GLSMaxOp, LLVM::SMaxOp>,
2065 DirectConversionPattern<spirv::GLSMinOp, LLVM::SMinOp>,
2066 DirectConversionPattern<spirv::GLSqrtOp, LLVM::SqrtOp>,
2067 DirectConversionPattern<spirv::GLUMaxOp, LLVM::UMaxOp>,
2068 DirectConversionPattern<spirv::GLUMinOp, LLVM::UMinOp>,
2069 DirectConversionPattern<spirv::GLTruncOp, LLVM::FTruncOp>,
2070 DirectConversionPattern<spirv::GLAsinOp, LLVM::ASinOp>,
2071 DirectConversionPattern<spirv::GLAcosOp, LLVM::ACosOp>,
2072 DirectConversionPattern<spirv::GLAtanOp, LLVM::ATanOp>,
2073 InverseSqrtPattern, SAbsPattern, TanPattern, TanhPattern, FractPattern,
2077 DirectConversionPattern<spirv::CLCeilOp, LLVM::FCeilOp>,
2078 DirectConversionPattern<spirv::CLCosOp, LLVM::CosOp>,
2079 DirectConversionPattern<spirv::CLExpOp, LLVM::ExpOp>,
2080 DirectConversionPattern<spirv::CLExp2Op, LLVM::Exp2Op>,
2081 DirectConversionPattern<spirv::CLExp10Op, LLVM::Exp10Op>,
2082 DirectConversionPattern<spirv::CLFAbsOp, LLVM::FAbsOp>,
2083 DirectConversionPattern<spirv::CLFloorOp, LLVM::FFloorOp>,
2084 DirectConversionPattern<spirv::CLFmaOp, LLVM::FMAOp>,
2085 DirectConversionPattern<spirv::CLFMaxOp, LLVM::MaxNumOp>,
2086 DirectConversionPattern<spirv::CLFMinOp, LLVM::MinNumOp>,
2087 DirectConversionPattern<spirv::CLLogOp, LLVM::LogOp>,
2088 DirectConversionPattern<spirv::CLLog2Op, LLVM::Log2Op>,
2089 DirectConversionPattern<spirv::CLLog10Op, LLVM::Log10Op>,
2090 DirectConversionPattern<spirv::CLPowOp, LLVM::PowOp>,
2091 DirectConversionPattern<spirv::CLRintOp, LLVM::RintOp>,
2092 DirectConversionPattern<spirv::CLRoundOp, LLVM::RoundOp>,
2093 DirectConversionPattern<spirv::CLSinOp, LLVM::SinOp>,
2094 DirectConversionPattern<spirv::CLSinhOp, LLVM::SinhOp>,
2095 DirectConversionPattern<spirv::CLCoshOp, LLVM::CoshOp>,
2096 DirectConversionPattern<spirv::CLTanOp, LLVM::TanOp>,
2097 DirectConversionPattern<spirv::CLTanhOp, LLVM::TanhOp>,
2098 DirectConversionPattern<spirv::CLAsinOp, LLVM::ASinOp>,
2099 DirectConversionPattern<spirv::CLAcosOp, LLVM::ACosOp>,
2100 DirectConversionPattern<spirv::CLAtanOp, LLVM::ATanOp>,
2101 DirectConversionPattern<spirv::CLAtan2Op, LLVM::ATan2Op>,
2102 DirectConversionPattern<spirv::CLSqrtOp, LLVM::SqrtOp>,
2103 DirectConversionPattern<spirv::CLTruncOp, LLVM::FTruncOp>,
2104 DirectConversionPattern<spirv::CLCopysignOp, LLVM::CopySignOp>,
2105 DirectConversionPattern<spirv::CLFmodOp, LLVM::FRemOp>,
2106 DirectConversionPattern<spirv::CLSMaxOp, LLVM::SMaxOp>,
2107 DirectConversionPattern<spirv::CLSMinOp, LLVM::SMinOp>,
2108 DirectConversionPattern<spirv::CLUMaxOp, LLVM::UMaxOp>,
2109 DirectConversionPattern<spirv::CLUMinOp, LLVM::UMinOp>, CLMixPattern,
2112 DirectConversionPattern<spirv::LogicalAndOp, LLVM::AndOp>,
2113 DirectConversionPattern<spirv::LogicalOrOp, LLVM::OrOp>,
2114 IComparePattern<spirv::LogicalEqualOp, LLVM::ICmpPredicate::eq>,
2115 IComparePattern<spirv::LogicalNotEqualOp, LLVM::ICmpPredicate::ne>,
2116 NotPattern<spirv::LogicalNotOp>,
2119 AccessChainPattern, AddressOfPattern, LoadStorePattern<spirv::LoadOp>,
2120 LoadStorePattern<spirv::StoreOp>, VariablePattern,
2123 CompositeExtractPattern, CompositeInsertPattern,
2124 DirectConversionPattern<spirv::SelectOp, LLVM::SelectOp>,
2125 DirectConversionPattern<spirv::UndefOp, LLVM::UndefOp>,
2126 VectorShufflePattern,
2129 ShiftPattern<spirv::ShiftRightArithmeticOp, LLVM::AShrOp>,
2130 ShiftPattern<spirv::ShiftRightLogicalOp, LLVM::LShrOp>,
2131 ShiftPattern<spirv::ShiftLeftLogicalOp, LLVM::ShlOp>,
2134 ReturnPattern, ReturnValuePattern,
2140 ControlBarrierPattern<spirv::ControlBarrierOp>,
2141 ControlBarrierPattern<spirv::INTELControlBarrierArriveOp>,
2142 ControlBarrierPattern<spirv::INTELControlBarrierWaitOp>,
2145 GroupReducePattern<spirv::GroupIAddOp>,
2146 GroupReducePattern<spirv::GroupFAddOp>,
2147 GroupReducePattern<spirv::GroupFMinOp>,
2148 GroupReducePattern<spirv::GroupUMinOp>,
2149 GroupReducePattern<spirv::GroupSMinOp,
true>,
2150 GroupReducePattern<spirv::GroupFMaxOp>,
2151 GroupReducePattern<spirv::GroupUMaxOp>,
2152 GroupReducePattern<spirv::GroupSMaxOp,
true>,
2153 GroupReducePattern<spirv::GroupNonUniformIAddOp,
false,
2155 GroupReducePattern<spirv::GroupNonUniformFAddOp,
false,
2157 GroupReducePattern<spirv::GroupNonUniformIMulOp,
false,
2159 GroupReducePattern<spirv::GroupNonUniformFMulOp,
false,
2161 GroupReducePattern<spirv::GroupNonUniformSMinOp,
true,
2163 GroupReducePattern<spirv::GroupNonUniformUMinOp,
false,
2165 GroupReducePattern<spirv::GroupNonUniformFMinOp,
false,
2167 GroupReducePattern<spirv::GroupNonUniformSMaxOp,
true,
2169 GroupReducePattern<spirv::GroupNonUniformUMaxOp,
false,
2171 GroupReducePattern<spirv::GroupNonUniformFMaxOp,
false,
2173 GroupReducePattern<spirv::GroupNonUniformBitwiseAndOp,
false,
2175 GroupReducePattern<spirv::GroupNonUniformBitwiseOrOp,
false,
2177 GroupReducePattern<spirv::GroupNonUniformBitwiseXorOp,
false,
2179 GroupReducePattern<spirv::GroupNonUniformLogicalAndOp,
false,
2181 GroupReducePattern<spirv::GroupNonUniformLogicalOrOp,
false,
2183 GroupReducePattern<spirv::GroupNonUniformLogicalXorOp,
false,
2187 patterns.
add<GlobalVariablePattern>(clientAPI, patterns.
getContext(),
2190 patterns.
add<ScalePattern<spirv::GLRadiansOp>>(
2191 0.017453292519943295, patterns.
getContext(), typeConverter);
2193 patterns.
add<ScalePattern<spirv::GLDegreesOp>>(
2194 57.29577951308232, patterns.
getContext(), typeConverter);
2199 patterns.
add<FuncConversionPattern>(patterns.
getContext(), typeConverter);
2204 patterns.
add<ModuleConversionPattern>(patterns.
getContext(), typeConverter);
2215 auto spvModules =
module.getOps<spirv::ModuleOp>();
2216 for (
auto spvModule : spvModules) {
2217 spvModule.walk([&](spirv::GlobalVariableOp op) {
2218 IntegerAttr descriptorSet =
2220 IntegerAttr binding = op->getAttrOfType<IntegerAttr>(
kBinding);
2223 if (descriptorSet && binding) {
2226 auto moduleAndName =
2227 spvModule.getName().has_value()
2228 ? spvModule.getName()->str() +
"_" + op.getSymName().str()
2229 : op.getSymName().str();
2231 llvm::formatv(
"{0}_descriptor_set{1}_binding{2}", moduleAndName,
2232 std::to_string(descriptorSet.getInt()),
2233 std::to_string(binding.getInt()));
2234 auto nameAttr = StringAttr::get(op->getContext(), name);
2239 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.