29#define DEBUG_TYPE "memref-to-spirv-pattern"
50 assert(targetBits % sourceBits == 0);
52 IntegerAttr idxAttr = builder.
getIntegerAttr(type, targetBits / sourceBits);
53 auto idx = builder.
createOrFold<spirv::ConstantOp>(loc, type, idxAttr);
54 IntegerAttr srcBitsAttr = builder.
getIntegerAttr(type, sourceBits);
56 builder.
createOrFold<spirv::ConstantOp>(loc, type, srcBitsAttr);
57 auto m = builder.
createOrFold<spirv::UModOp>(loc, srcIdx, idx);
58 return builder.
createOrFold<spirv::IMulOp>(loc, type, m, srcBitsValue);
71 spirv::AccessChainOp op,
int sourceBits,
73 assert(targetBits % sourceBits == 0);
74 const auto loc = op.getLoc();
75 Value lastDim = op->getOperand(op.getNumOperands() - 1);
77 IntegerAttr attr = builder.
getIntegerAttr(type, targetBits / sourceBits);
78 auto idx = builder.
createOrFold<spirv::ConstantOp>(loc, type, attr);
79 auto indices = llvm::to_vector<4>(op.getIndices());
83 Type t = typeConverter.convertType(op.getComponentPtr().getType());
84 return spirv::AccessChainOp::create(builder, loc, t, op.getBasePtr(),
94 Value zero = spirv::ConstantOp::getZero(dstType, loc, builder);
95 Value one = spirv::ConstantOp::getOne(dstType, loc, builder);
96 return builder.
createOrFold<spirv::SelectOp>(loc, dstType, srcBool, one,
104 IntegerType dstType = cast<IntegerType>(mask.getType());
105 int targetBits =
static_cast<int>(dstType.getWidth());
107 assert(valueBits <= targetBits);
109 if (valueBits == 1) {
112 if (valueBits < targetBits) {
113 value = spirv::UConvertOp::create(
117 value = builder.
createOrFold<spirv::BitwiseAndOp>(loc, value, mask);
127 if (!type.hasStaticShape())
130 if (isa<memref::AllocOp, memref::DeallocOp>(allocOp)) {
131 auto sc = dyn_cast_or_null<spirv::StorageClassAttr>(type.getMemorySpace());
132 if (!sc || sc.getValue() != spirv::StorageClass::Workgroup)
134 }
else if (isa<memref::AllocaOp>(allocOp)) {
135 auto sc = dyn_cast_or_null<spirv::StorageClassAttr>(type.getMemorySpace());
136 if (!sc || sc.getValue() != spirv::StorageClass::Function)
140 if (isa<MemRefElementTypeInterface>(type.getElementType()))
147 Type elementType = type.getElementType();
148 if (
auto vecType = dyn_cast<VectorType>(elementType))
149 elementType = vecType.getElementType();
150 if (
auto compType = dyn_cast<ComplexType>(elementType))
151 elementType = compType.getElementType();
159 auto sc = dyn_cast_or_null<spirv::StorageClassAttr>(type.getMemorySpace());
160 switch (sc.getValue()) {
161 case spirv::StorageClass::StorageBuffer:
162 return spirv::Scope::Device;
163 case spirv::StorageClass::Workgroup:
164 return spirv::Scope::Workgroup;
174static spirv::MemorySemantics
177 case spirv::StorageClass::StorageBuffer:
178 case spirv::StorageClass::Uniform:
179 return spirv::MemorySemantics::UniformMemory;
180 case spirv::StorageClass::Workgroup:
181 return spirv::MemorySemantics::WorkgroupMemory;
182 case spirv::StorageClass::CrossWorkgroup:
183 return spirv::MemorySemantics::CrossWorkgroupMemory;
184 case spirv::StorageClass::AtomicCounter:
185 return spirv::MemorySemantics::AtomicCounterMemory;
186 case spirv::StorageClass::Image:
187 return spirv::MemorySemantics::ImageMemory;
189 return spirv::MemorySemantics::None;
196 auto sc = cast<spirv::StorageClassAttr>(type.getMemorySpace()).getValue();
197 return spirv::MemorySemantics::AcquireRelease |
210 if (typeConverter.
allows(spirv::Capability::Kernel)) {
211 if (
auto arrayType = dyn_cast<spirv::ArrayType>(pointeeType))
212 return arrayType.getElementType();
216 Type structElemType = cast<spirv::StructType>(pointeeType).getElementType(0);
217 if (
auto arrayType = dyn_cast<spirv::ArrayType>(structElemType))
218 return arrayType.getElementType();
219 return cast<spirv::RuntimeArrayType>(structElemType).getElementType();
227 auto one = spirv::ConstantOp::getZero(srcInt.
getType(), loc, builder);
228 return builder.
createOrFold<spirv::INotEqualOp>(loc, srcInt, one);
242class AllocaOpPattern final :
public OpConversionPattern<memref::AllocaOp> {
247 matchAndRewrite(memref::AllocaOp allocaOp, OpAdaptor adaptor,
248 ConversionPatternRewriter &rewriter)
const override;
255class AllocOpPattern final :
public OpConversionPattern<memref::AllocOp> {
260 matchAndRewrite(memref::AllocOp operation, OpAdaptor adaptor,
261 ConversionPatternRewriter &rewriter)
const override;
265class AtomicRMWOpPattern final
266 :
public OpConversionPattern<memref::AtomicRMWOp> {
271 matchAndRewrite(memref::AtomicRMWOp atomicOp, OpAdaptor adaptor,
272 ConversionPatternRewriter &rewriter)
const override;
277class DeallocOpPattern final :
public OpConversionPattern<memref::DeallocOp> {
282 matchAndRewrite(memref::DeallocOp operation, OpAdaptor adaptor,
283 ConversionPatternRewriter &rewriter)
const override;
287class IntLoadOpPattern final :
public OpConversionPattern<memref::LoadOp> {
292 matchAndRewrite(memref::LoadOp loadOp, OpAdaptor adaptor,
293 ConversionPatternRewriter &rewriter)
const override;
297class LoadOpPattern final :
public OpConversionPattern<memref::LoadOp> {
302 matchAndRewrite(memref::LoadOp loadOp, OpAdaptor adaptor,
303 ConversionPatternRewriter &rewriter)
const override;
307class ImageLoadOpPattern final :
public OpConversionPattern<memref::LoadOp> {
312 matchAndRewrite(memref::LoadOp loadOp, OpAdaptor adaptor,
313 ConversionPatternRewriter &rewriter)
const override;
317class IntStoreOpPattern final :
public OpConversionPattern<memref::StoreOp> {
322 matchAndRewrite(memref::StoreOp storeOp, OpAdaptor adaptor,
323 ConversionPatternRewriter &rewriter)
const override;
327class MemorySpaceCastOpPattern final
328 :
public OpConversionPattern<memref::MemorySpaceCastOp> {
333 matchAndRewrite(memref::MemorySpaceCastOp addrCastOp, OpAdaptor adaptor,
334 ConversionPatternRewriter &rewriter)
const override;
338class StoreOpPattern final :
public OpConversionPattern<memref::StoreOp> {
343 matchAndRewrite(memref::StoreOp storeOp, OpAdaptor adaptor,
344 ConversionPatternRewriter &rewriter)
const override;
348class CopyOpPattern final :
public OpConversionPattern<memref::CopyOp> {
353 matchAndRewrite(memref::CopyOp copyOp, OpAdaptor adaptor,
354 ConversionPatternRewriter &rewriter)
const override;
357class ReinterpretCastPattern final
358 :
public OpConversionPattern<memref::ReinterpretCastOp> {
363 matchAndRewrite(memref::ReinterpretCastOp op, OpAdaptor adaptor,
364 ConversionPatternRewriter &rewriter)
const override;
367class CastPattern final :
public OpConversionPattern<memref::CastOp> {
372 matchAndRewrite(memref::CastOp op, OpAdaptor adaptor,
373 ConversionPatternRewriter &rewriter)
const override {
374 Value src = adaptor.getSource();
377 const TypeConverter *converter = getTypeConverter();
378 Type dstType = converter->convertType(op.getType());
379 if (srcType != dstType)
380 return rewriter.notifyMatchFailure(op, [&](Diagnostic &diag) {
381 diag <<
"types doesn't match: " << srcType <<
" and " << dstType;
384 rewriter.replaceOp(op, src);
390class ExtractAlignedPointerAsIndexOpPattern final
391 :
public OpConversionPattern<memref::ExtractAlignedPointerAsIndexOp> {
396 matchAndRewrite(memref::ExtractAlignedPointerAsIndexOp extractOp,
398 ConversionPatternRewriter &rewriter)
const override;
407AllocaOpPattern::matchAndRewrite(memref::AllocaOp allocaOp, OpAdaptor adaptor,
408 ConversionPatternRewriter &rewriter)
const {
409 MemRefType allocType = allocaOp.getType();
411 return rewriter.notifyMatchFailure(allocaOp,
"unhandled allocation type");
414 Type spirvType = getTypeConverter()->convertType(allocType);
416 return rewriter.notifyMatchFailure(allocaOp,
"type conversion failed");
418 auto function = allocaOp->getParentOfType<FunctionOpInterface>();
420 return rewriter.notifyMatchFailure(allocaOp,
421 "requires a containing function");
424 OpBuilder::InsertionGuard guard(rewriter);
425 Block &entryBlock = function->getRegion(0).front();
428 while (insertionPoint != entryBlock.
end() &&
429 isa<spirv::VariableOp>(*insertionPoint))
431 rewriter.setInsertionPoint(&entryBlock, insertionPoint);
432 Value variable = spirv::VariableOp::create(
433 rewriter, allocaOp.getLoc(), spirvType, spirv::StorageClass::Function,
435 rewriter.replaceOp(allocaOp, variable);
444AllocOpPattern::matchAndRewrite(memref::AllocOp operation, OpAdaptor adaptor,
445 ConversionPatternRewriter &rewriter)
const {
446 MemRefType allocType = operation.getType();
448 return rewriter.notifyMatchFailure(operation,
"unhandled allocation type");
451 Type spirvType = getTypeConverter()->convertType(allocType);
453 return rewriter.notifyMatchFailure(operation,
"type conversion failed");
460 Location loc = operation.getLoc();
461 spirv::GlobalVariableOp varOp;
463 OpBuilder::InsertionGuard guard(rewriter);
465 rewriter.setInsertionPointToStart(&entryBlock);
466 auto varOps = entryBlock.
getOps<spirv::GlobalVariableOp>();
467 std::string varName =
468 std::string(
"__workgroup_mem__") +
469 std::to_string(std::distance(varOps.begin(), varOps.end()));
470 varOp = spirv::GlobalVariableOp::create(rewriter, loc, spirvType, varName,
475 rewriter.replaceOpWithNewOp<spirv::AddressOfOp>(operation, varOp);
484AtomicRMWOpPattern::matchAndRewrite(memref::AtomicRMWOp atomicOp,
486 ConversionPatternRewriter &rewriter)
const {
487 auto memrefType = cast<MemRefType>(atomicOp.getMemref().getType());
490 return rewriter.notifyMatchFailure(atomicOp,
491 "unsupported memref memory space");
493 auto &typeConverter = *getTypeConverter<SPIRVTypeConverter>();
494 Type resultType = typeConverter.convertType(atomicOp.getType());
496 return rewriter.notifyMatchFailure(atomicOp,
497 "failed to convert result type");
499 auto loc = atomicOp.getLoc();
502 adaptor.getIndices(), loc, rewriter);
510 int srcBits = memrefType.getElementType().getIntOrFloatBitWidth();
511 auto pointerType = typeConverter.convertType<spirv::PointerType>(memrefType);
513 return rewriter.notifyMatchFailure(atomicOp,
514 "failed to convert memref type");
516 Type pointeeType = pointerType.getPointeeType();
517 Type storageElemType =
519 if (!storageElemType || !storageElemType.
isIntOrFloat())
520 return rewriter.notifyMatchFailure(
521 atomicOp,
"failed to determine destination element type");
524 assert(dstBits % srcBits == 0);
530 if (srcBits == dstBits) {
531#define ATOMIC_CASE(kind, spirvOp) \
532 case arith::AtomicRMWKind::kind: \
533 rewriter.replaceOpWithNewOp<spirv::spirvOp>( \
534 atomicOp, resultType, ptr, *scope, memSem, adaptor.getValue()); \
537 switch (atomicOp.getKind()) {
548 return rewriter.notifyMatchFailure(atomicOp,
"unimplemented atomic kind");
563 if (atomicOp.getKind() != arith::AtomicRMWKind::ori &&
564 atomicOp.getKind() != arith::AtomicRMWKind::andi) {
565 return rewriter.notifyMatchFailure(
567 "atomic op on sub-element-width types is only supported for ori/andi");
572 if (typeConverter.allows(spirv::Capability::Kernel))
573 return rewriter.notifyMatchFailure(
575 "sub-element-width atomic ops unsupported with Kernel capability");
577 auto dstType = cast<IntegerType>(storageElemType);
579 auto accessChainOp = ptr.
getDefiningOp<spirv::AccessChainOp>();
585 assert(accessChainOp.getIndices().size() == 2);
586 Value lastDim = accessChainOp->getOperand(accessChainOp.getNumOperands() - 1);
589 srcBits, dstBits, rewriter);
591 switch (atomicOp.getKind()) {
592 case arith::AtomicRMWKind::ori: {
595 Value elemMask = rewriter.createOrFold<spirv::ConstantOp>(
596 loc, dstType, rewriter.getIntegerAttr(dstType, (1uLL << srcBits) - 1));
598 shiftValue(loc, adaptor.getValue(), offset, elemMask, rewriter);
599 result = spirv::AtomicOrOp::create(rewriter, loc, dstType, adjustedPtr,
600 *scope, memSem, storeVal);
603 case arith::AtomicRMWKind::andi: {
607 Value elemMask = rewriter.createOrFold<spirv::ConstantOp>(
608 loc, dstType, rewriter.getIntegerAttr(dstType, (1uLL << srcBits) - 1));
610 shiftValue(loc, adaptor.getValue(), offset, elemMask, rewriter);
611 Value shiftedElemMask = rewriter.createOrFold<spirv::ShiftLeftLogicalOp>(
612 loc, dstType, elemMask, offset);
613 Value invertedElemMask =
614 rewriter.createOrFold<spirv::NotOp>(loc, dstType, shiftedElemMask);
615 Value mask = rewriter.createOrFold<spirv::BitwiseOrOp>(loc, storeVal,
617 result = spirv::AtomicAndOp::create(rewriter, loc, dstType, adjustedPtr,
618 *scope, memSem, mask);
622 return rewriter.notifyMatchFailure(atomicOp,
"unimplemented atomic kind");
627 result = rewriter.createOrFold<spirv::ShiftRightLogicalOp>(loc, dstType,
629 Value mask = rewriter.createOrFold<spirv::ConstantOp>(
630 loc, dstType, rewriter.getIntegerAttr(dstType, (1uLL << srcBits) - 1));
632 rewriter.createOrFold<spirv::BitwiseAndOp>(loc, dstType,
result, mask);
633 rewriter.replaceOp(atomicOp,
result);
643DeallocOpPattern::matchAndRewrite(memref::DeallocOp operation,
645 ConversionPatternRewriter &rewriter)
const {
646 MemRefType deallocType = cast<MemRefType>(operation.getMemref().getType());
648 return rewriter.notifyMatchFailure(operation,
"unhandled allocation type");
649 rewriter.eraseOp(operation);
664static FailureOr<MemoryRequirements>
666 uint64_t preferredAlignment) {
667 if (preferredAlignment >= std::numeric_limits<uint32_t>::max()) {
673 auto memoryAccess = spirv::MemoryAccess::None;
675 memoryAccess = spirv::MemoryAccess::Nontemporal;
678 auto ptrType = cast<spirv::PointerType>(accessedPtr.
getType());
679 bool mayOmitAlignment =
680 !preferredAlignment &&
681 ptrType.getStorageClass() != spirv::StorageClass::PhysicalStorageBuffer;
682 if (mayOmitAlignment) {
683 if (memoryAccess == spirv::MemoryAccess::None) {
692 std::optional<int64_t> sizeInBytes;
693 Type rawPointeeType = ptrType.getPointeeType();
694 if (
auto scalarType = dyn_cast<spirv::ScalarType>(rawPointeeType)) {
696 sizeInBytes = scalarType.getSizeInBytes();
697 }
else if (
auto vecType = dyn_cast<VectorType>(rawPointeeType)) {
700 if (
auto scalarElem =
701 dyn_cast<spirv::ScalarType>(vecType.getElementType())) {
702 if (
auto elemSize = scalarElem.getSizeInBytes())
703 sizeInBytes = *elemSize * vecType.getNumElements();
707 if (!sizeInBytes.has_value())
710 memoryAccess |= spirv::MemoryAccess::Aligned;
711 auto memAccessAttr = spirv::MemoryAccessAttr::get(ctx, memoryAccess);
712 auto alignmentValue = preferredAlignment ? preferredAlignment : *sizeInBytes;
713 auto alignment = IntegerAttr::get(IntegerType::get(ctx, 32), alignmentValue);
720template <
class LoadOrStoreOp>
721static FailureOr<MemoryRequirements>
724 llvm::is_one_of<LoadOrStoreOp, memref::LoadOp, memref::StoreOp>::value,
725 "Must be called on either memref::LoadOp or memref::StoreOp");
728 loadOrStoreOp.getNontemporal(),
729 loadOrStoreOp.getAlignment().value_or(0));
733IntLoadOpPattern::matchAndRewrite(memref::LoadOp loadOp, OpAdaptor adaptor,
734 ConversionPatternRewriter &rewriter)
const {
735 auto loc = loadOp.getLoc();
736 auto memrefType = cast<MemRefType>(loadOp.getMemref().getType());
737 if (!memrefType.getElementType().isSignlessInteger())
740 auto memorySpaceAttr =
741 dyn_cast_if_present<spirv::StorageClassAttr>(memrefType.getMemorySpace());
742 if (!memorySpaceAttr)
743 return rewriter.notifyMatchFailure(
744 loadOp,
"missing memory space SPIR-V storage class attribute");
746 if (memorySpaceAttr.getValue() == spirv::StorageClass::Image)
747 return rewriter.notifyMatchFailure(
749 "failed to lower memref in image storage class to storage buffer");
751 const auto &typeConverter = *getTypeConverter<SPIRVTypeConverter>();
754 adaptor.getIndices(), loc, rewriter);
759 int srcBits = memrefType.getElementType().getIntOrFloatBitWidth();
760 bool isBool = srcBits == 1;
762 srcBits = typeConverter.getOptions().boolNumBits;
764 auto pointerType = typeConverter.convertType<spirv::PointerType>(memrefType);
766 return rewriter.notifyMatchFailure(loadOp,
"failed to convert memref type");
768 Type pointeeType = pointerType.getPointeeType();
771 assert(dstBits % srcBits == 0);
775 if (srcBits == dstBits) {
777 if (
failed(memoryRequirements))
778 return rewriter.notifyMatchFailure(
779 loadOp,
"failed to determine memory requirements");
781 auto [memoryAccess, alignment] = *memoryRequirements;
782 Value loadVal = spirv::LoadOp::create(rewriter, loc, accessChain,
783 memoryAccess, alignment);
786 rewriter.replaceOp(loadOp, loadVal);
792 if (typeConverter.allows(spirv::Capability::Kernel))
795 auto accessChainOp = accessChain.
getDefiningOp<spirv::AccessChainOp>();
802 assert(accessChainOp.getIndices().size() == 2);
804 srcBits, dstBits, rewriter);
806 if (
failed(memoryRequirements))
807 return rewriter.notifyMatchFailure(
808 loadOp,
"failed to determine memory requirements");
810 auto [memoryAccess, alignment] = *memoryRequirements;
811 Value spvLoadOp = spirv::LoadOp::create(rewriter, loc, dstType, adjustedPtr,
812 memoryAccess, alignment);
816 Value lastDim = accessChainOp->getOperand(accessChainOp.getNumOperands() - 1);
818 Value
result = rewriter.createOrFold<spirv::ShiftRightArithmeticOp>(
819 loc, spvLoadOp.
getType(), spvLoadOp, offset);
822 Value mask = rewriter.createOrFold<spirv::ConstantOp>(
823 loc, dstType, rewriter.getIntegerAttr(dstType, (1 << srcBits) - 1));
825 rewriter.createOrFold<spirv::BitwiseAndOp>(loc, dstType,
result, mask);
830 IntegerAttr shiftValueAttr =
831 rewriter.getIntegerAttr(dstType, dstBits - srcBits);
833 rewriter.createOrFold<spirv::ConstantOp>(loc, dstType, shiftValueAttr);
834 result = rewriter.createOrFold<spirv::ShiftLeftLogicalOp>(loc, dstType,
836 result = rewriter.createOrFold<spirv::ShiftRightArithmeticOp>(
839 rewriter.replaceOp(loadOp,
result);
841 assert(accessChainOp.use_empty());
842 rewriter.eraseOp(accessChainOp);
848LoadOpPattern::matchAndRewrite(memref::LoadOp loadOp, OpAdaptor adaptor,
849 ConversionPatternRewriter &rewriter)
const {
850 auto memrefType = cast<MemRefType>(loadOp.getMemref().getType());
851 if (memrefType.getElementType().isSignlessInteger())
854 auto memorySpaceAttr =
855 dyn_cast_if_present<spirv::StorageClassAttr>(memrefType.getMemorySpace());
856 if (!memorySpaceAttr)
857 return rewriter.notifyMatchFailure(
858 loadOp,
"missing memory space SPIR-V storage class attribute");
860 if (memorySpaceAttr.getValue() == spirv::StorageClass::Image)
861 return rewriter.notifyMatchFailure(
863 "failed to lower memref in image storage class to storage buffer");
866 *getTypeConverter<SPIRVTypeConverter>(), memrefType, adaptor.getMemref(),
867 adaptor.getIndices(), loadOp.getLoc(), rewriter);
873 if (
failed(memoryRequirements))
874 return rewriter.notifyMatchFailure(
875 loadOp,
"failed to determine memory requirements");
877 auto [memoryAccess, alignment] = *memoryRequirements;
878 rewriter.replaceOpWithNewOp<spirv::LoadOp>(loadOp, loadPtr, memoryAccess,
883template <
typename OpAdaptor>
884static FailureOr<SmallVector<Value>>
886 ConversionPatternRewriter &rewriter) {
893 AffineMap map = loadOp.getMemRefType().getLayout().getAffineMap();
895 return rewriter.notifyMatchFailure(
897 "Cannot lower memrefs with memory layout which is not a permutation");
903 for (
unsigned dim = 0; dim < dimCount; ++dim)
909 return llvm::to_vector(llvm::reverse(coords));
913ImageLoadOpPattern::matchAndRewrite(memref::LoadOp loadOp, OpAdaptor adaptor,
914 ConversionPatternRewriter &rewriter)
const {
915 auto memrefType = cast<MemRefType>(loadOp.getMemref().getType());
917 auto memorySpaceAttr =
918 dyn_cast_if_present<spirv::StorageClassAttr>(memrefType.getMemorySpace());
919 if (!memorySpaceAttr)
920 return rewriter.notifyMatchFailure(
921 loadOp,
"missing memory space SPIR-V storage class attribute");
923 if (memorySpaceAttr.getValue() != spirv::StorageClass::Image)
924 return rewriter.notifyMatchFailure(
925 loadOp,
"failed to lower memref in non-image storage class to image");
927 Value loadPtr = adaptor.getMemref();
929 if (
failed(memoryRequirements))
930 return rewriter.notifyMatchFailure(
931 loadOp,
"failed to determine memory requirements");
933 const auto [memoryAccess, alignment] = *memoryRequirements;
935 if (!loadOp.getMemRefType().hasRank())
936 return rewriter.notifyMatchFailure(
937 loadOp,
"cannot lower unranked memrefs to SPIR-V images");
942 if (!isa<spirv::ScalarType>(loadOp.getMemRefType().getElementType()))
943 return rewriter.notifyMatchFailure(
945 "cannot lower memrefs who's element type is not a SPIR-V scalar type"
952 auto convertedPointeeType = cast<spirv::PointerType>(
953 getTypeConverter()->convertType(loadOp.getMemRefType()));
954 if (!isa<spirv::SampledImageType>(convertedPointeeType.getPointeeType()))
955 return rewriter.notifyMatchFailure(loadOp,
956 "cannot lower memrefs which do not "
957 "convert to SPIR-V sampled images");
960 Location loc = loadOp->getLoc();
962 spirv::LoadOp::create(rewriter, loc, loadPtr, memoryAccess, alignment);
964 auto imageOp = spirv::ImageOp::create(rewriter, loc, imageLoadOp);
968 if (memrefType.getRank() == 1) {
969 coords = adaptor.getIndices()[0];
971 FailureOr<SmallVector<Value>> maybeCoords =
975 auto coordVectorType = VectorType::get({loadOp.getMemRefType().getRank()},
976 adaptor.getIndices().
getType()[0]);
977 coords = spirv::CompositeConstructOp::create(rewriter, loc, coordVectorType,
978 maybeCoords.value());
982 auto resultVectorType = VectorType::get({4}, loadOp.getType());
983 auto fetchOp = spirv::ImageFetchOp::create(
984 rewriter, loc, resultVectorType, imageOp, coords,
985 mlir::spirv::ImageOperandsAttr{},
ValueRange{});
990 auto compositeExtractOp =
991 spirv::CompositeExtractOp::create(rewriter, loc, fetchOp, 0);
993 rewriter.replaceOp(loadOp, compositeExtractOp);
998IntStoreOpPattern::matchAndRewrite(memref::StoreOp storeOp, OpAdaptor adaptor,
999 ConversionPatternRewriter &rewriter)
const {
1000 auto memrefType = cast<MemRefType>(storeOp.getMemref().getType());
1001 if (!memrefType.getElementType().isSignlessInteger())
1002 return rewriter.notifyMatchFailure(storeOp,
1003 "element type is not a signless int");
1005 auto loc = storeOp.getLoc();
1006 auto &typeConverter = *getTypeConverter<SPIRVTypeConverter>();
1009 adaptor.getIndices(), loc, rewriter);
1012 return rewriter.notifyMatchFailure(
1013 storeOp,
"failed to convert element pointer type");
1015 int srcBits = memrefType.getElementType().getIntOrFloatBitWidth();
1017 bool isBool = srcBits == 1;
1019 srcBits = typeConverter.getOptions().boolNumBits;
1021 auto pointerType = typeConverter.convertType<spirv::PointerType>(memrefType);
1023 return rewriter.notifyMatchFailure(storeOp,
1024 "failed to convert memref type");
1026 Type pointeeType = pointerType.getPointeeType();
1027 auto dstType = dyn_cast<IntegerType>(
1030 return rewriter.notifyMatchFailure(
1031 storeOp,
"failed to determine destination element type");
1033 int dstBits =
static_cast<int>(dstType.getWidth());
1034 assert(dstBits % srcBits == 0);
1036 if (srcBits == dstBits) {
1038 if (
failed(memoryRequirements))
1039 return rewriter.notifyMatchFailure(
1040 storeOp,
"failed to determine memory requirements");
1042 auto [memoryAccess, alignment] = *memoryRequirements;
1043 Value storeVal = adaptor.getValue();
1046 rewriter.replaceOpWithNewOp<spirv::StoreOp>(storeOp, accessChain, storeVal,
1047 memoryAccess, alignment);
1053 if (typeConverter.allows(spirv::Capability::Kernel))
1056 auto accessChainOp = accessChain.
getDefiningOp<spirv::AccessChainOp>();
1071 assert(accessChainOp.getIndices().size() == 2);
1072 Value lastDim = accessChainOp->getOperand(accessChainOp.getNumOperands() - 1);
1077 Value mask = rewriter.createOrFold<spirv::ConstantOp>(
1078 loc, dstType, rewriter.getIntegerAttr(dstType, (1 << srcBits) - 1));
1079 Value clearBitsMask = rewriter.createOrFold<spirv::ShiftLeftLogicalOp>(
1080 loc, dstType, mask, offset);
1082 rewriter.createOrFold<spirv::NotOp>(loc, dstType, clearBitsMask);
1084 Value storeVal =
shiftValue(loc, adaptor.getValue(), offset, mask, rewriter);
1086 srcBits, dstBits, rewriter);
1089 return rewriter.notifyMatchFailure(storeOp,
"atomic scope not available");
1092 Value
result = spirv::AtomicAndOp::create(rewriter, loc, dstType, adjustedPtr,
1093 *scope, memSem, clearBitsMask);
1094 result = spirv::AtomicOrOp::create(rewriter, loc, dstType, adjustedPtr,
1095 *scope, memSem, storeVal);
1101 rewriter.eraseOp(storeOp);
1103 assert(accessChainOp.use_empty());
1104 rewriter.eraseOp(accessChainOp);
1113LogicalResult MemorySpaceCastOpPattern::matchAndRewrite(
1114 memref::MemorySpaceCastOp addrCastOp, OpAdaptor adaptor,
1115 ConversionPatternRewriter &rewriter)
const {
1116 Location loc = addrCastOp.getLoc();
1117 auto &typeConverter = *getTypeConverter<SPIRVTypeConverter>();
1118 if (!typeConverter.allows(spirv::Capability::Kernel))
1119 return rewriter.notifyMatchFailure(
1120 loc,
"address space casts require kernel capability");
1122 auto sourceType = dyn_cast<MemRefType>(addrCastOp.getSource().getType());
1124 return rewriter.notifyMatchFailure(
1125 loc,
"SPIR-V lowering requires ranked memref types");
1126 auto resultType = cast<MemRefType>(addrCastOp.getResult().getType());
1128 auto sourceStorageClassAttr =
1129 dyn_cast_or_null<spirv::StorageClassAttr>(sourceType.getMemorySpace());
1130 if (!sourceStorageClassAttr)
1131 return rewriter.notifyMatchFailure(loc, [sourceType](Diagnostic &diag) {
1132 diag <<
"source address space " << sourceType.getMemorySpace()
1133 <<
" must be a SPIR-V storage class";
1135 auto resultStorageClassAttr =
1136 dyn_cast_or_null<spirv::StorageClassAttr>(resultType.getMemorySpace());
1137 if (!resultStorageClassAttr)
1138 return rewriter.notifyMatchFailure(loc, [resultType](Diagnostic &diag) {
1139 diag <<
"result address space " << resultType.getMemorySpace()
1140 <<
" must be a SPIR-V storage class";
1143 spirv::StorageClass sourceSc = sourceStorageClassAttr.getValue();
1144 spirv::StorageClass resultSc = resultStorageClassAttr.getValue();
1146 Value
result = adaptor.getSource();
1147 Type resultPtrType = typeConverter.convertType(resultType);
1149 return rewriter.notifyMatchFailure(addrCastOp,
1150 "failed to convert memref type");
1152 Type genericPtrType = resultPtrType;
1160 if (sourceSc != spirv::StorageClass::Generic &&
1161 resultSc != spirv::StorageClass::Generic) {
1162 Type intermediateType =
1163 MemRefType::get(sourceType.getShape(), sourceType.getElementType(),
1164 sourceType.getLayout(),
1165 rewriter.getAttr<spirv::StorageClassAttr>(
1166 spirv::StorageClass::Generic));
1167 genericPtrType = typeConverter.convertType(intermediateType);
1169 if (sourceSc != spirv::StorageClass::Generic) {
1170 result = spirv::PtrCastToGenericOp::create(rewriter, loc, genericPtrType,
1173 if (resultSc != spirv::StorageClass::Generic) {
1175 spirv::GenericCastToPtrOp::create(rewriter, loc, resultPtrType,
result);
1177 rewriter.replaceOp(addrCastOp,
result);
1182StoreOpPattern::matchAndRewrite(memref::StoreOp storeOp, OpAdaptor adaptor,
1183 ConversionPatternRewriter &rewriter)
const {
1184 auto memrefType = cast<MemRefType>(storeOp.getMemref().getType());
1185 if (memrefType.getElementType().isSignlessInteger())
1186 return rewriter.notifyMatchFailure(storeOp,
"signless int");
1188 *getTypeConverter<SPIRVTypeConverter>(), memrefType, adaptor.getMemref(),
1189 adaptor.getIndices(), storeOp.getLoc(), rewriter);
1192 return rewriter.notifyMatchFailure(storeOp,
"type conversion failed");
1195 if (
failed(memoryRequirements))
1196 return rewriter.notifyMatchFailure(
1197 storeOp,
"failed to determine memory requirements");
1199 auto [memoryAccess, alignment] = *memoryRequirements;
1200 rewriter.replaceOpWithNewOp<spirv::StoreOp>(
1201 storeOp, storePtr, adaptor.getValue(), memoryAccess, alignment);
1210CopyOpPattern::matchAndRewrite(memref::CopyOp copyOp, OpAdaptor adaptor,
1211 ConversionPatternRewriter &rewriter)
const {
1212 auto memrefType = cast<MemRefType>(copyOp.getSource().getType());
1213 if (!memrefType.hasStaticShape())
1214 return rewriter.notifyMatchFailure(copyOp,
"unsupported dynamic shape");
1216 for (MemRefType type :
1217 {memrefType, cast<MemRefType>(copyOp.getTarget().getType())}) {
1218 auto memorySpaceAttr =
1219 dyn_cast_if_present<spirv::StorageClassAttr>(type.getMemorySpace());
1220 if (memorySpaceAttr &&
1221 memorySpaceAttr.getValue() == spirv::StorageClass::Image)
1222 return rewriter.notifyMatchFailure(
1223 copyOp,
"cannot lower memref.copy in image storage class");
1229 Value source = adaptor.getSource();
1230 Value
target = adaptor.getTarget();
1231 auto sourcePtrType = dyn_cast<spirv::PointerType>(source.
getType());
1232 auto targetPtrType = dyn_cast<spirv::PointerType>(
target.getType());
1233 if (!sourcePtrType || !targetPtrType)
1234 return rewriter.notifyMatchFailure(copyOp,
"failed to convert memref type");
1236 if (sourcePtrType.getPointeeType() != targetPtrType.getPointeeType())
1237 return rewriter.notifyMatchFailure(
1238 copyOp,
"source and target pointee types do not match");
1240 rewriter.replaceOpWithNewOp<spirv::CopyMemoryOp>(
1241 copyOp,
target, source, spirv::MemoryAccessAttr{},
1243 spirv::MemoryAccessAttr{}, IntegerAttr{});
1247LogicalResult ReinterpretCastPattern::matchAndRewrite(
1248 memref::ReinterpretCastOp op, OpAdaptor adaptor,
1249 ConversionPatternRewriter &rewriter)
const {
1250 Value src = adaptor.getSource();
1251 auto srcType = dyn_cast<spirv::PointerType>(src.
getType());
1254 return rewriter.notifyMatchFailure(op, [&](Diagnostic &diag) {
1255 diag <<
"invalid src type " << src.
getType();
1258 const TypeConverter *converter = getTypeConverter();
1260 auto dstType = converter->convertType<spirv::PointerType>(op.getType());
1261 if (dstType != srcType)
1262 return rewriter.notifyMatchFailure(op, [&](Diagnostic &diag) {
1263 diag <<
"invalid dst type " << op.getType();
1266 OpFoldResult offset =
1267 getMixedValues(adaptor.getStaticOffsets(), adaptor.getOffsets(), rewriter)
1270 rewriter.replaceOp(op, src);
1274 Type intType = converter->convertType(rewriter.getIndexType());
1276 return rewriter.notifyMatchFailure(op,
"failed to convert index type");
1278 Location loc = op.getLoc();
1279 auto offsetValue = [&]() -> Value {
1280 if (
auto val = dyn_cast<Value>(offset))
1283 int64_t attrVal = cast<IntegerAttr>(cast<Attribute>(offset)).getInt();
1284 Attribute attr = rewriter.getIntegerAttr(intType, attrVal);
1285 return rewriter.createOrFold<spirv::ConstantOp>(loc, intType, attr);
1288 rewriter.replaceOpWithNewOp<spirv::InBoundsPtrAccessChainOp>(
1297LogicalResult ExtractAlignedPointerAsIndexOpPattern::matchAndRewrite(
1298 memref::ExtractAlignedPointerAsIndexOp extractOp, OpAdaptor adaptor,
1299 ConversionPatternRewriter &rewriter)
const {
1300 auto &typeConverter = *getTypeConverter<SPIRVTypeConverter>();
1301 Type indexType = typeConverter.getIndexType();
1302 rewriter.replaceOpWithNewOp<spirv::ConvertPtrToUOp>(extractOp, indexType,
1303 adaptor.getSource());
1314 patterns.
add<AllocaOpPattern, AllocOpPattern, AtomicRMWOpPattern,
1315 CopyOpPattern, DeallocOpPattern, IntLoadOpPattern,
1316 ImageLoadOpPattern, IntStoreOpPattern, LoadOpPattern,
1317 MemorySpaceCastOpPattern, StoreOpPattern, ReinterpretCastPattern,
1318 CastPattern, ExtractAlignedPointerAsIndexOpPattern>(
static spirv::MemorySemantics getMemorySemanticsForStorageClass(spirv::StorageClass sc)
Returns the MemorySemantics storage-class bit corresponding to sc.
static Value castIntNToBool(Location loc, Value srcInt, OpBuilder &builder)
Casts the given srcInt into a boolean value.
static Type getElementTypeForStoragePointer(Type pointeeType, const SPIRVTypeConverter &typeConverter)
Extracts the element type from a SPIR-V pointer type pointing to storage.
static std::optional< spirv::Scope > getAtomicOpScope(MemRefType type)
Returns the scope to use for atomic operations use for emulating store operations of unsupported inte...
static Value shiftValue(Location loc, Value value, Value offset, Value mask, OpBuilder &builder)
Returns the targetBits-bit value shifted by the given offset, and cast to the type destination type,...
static FailureOr< SmallVector< Value > > extractLoadCoordsForComposite(memref::LoadOp loadOp, OpAdaptor adaptor, ConversionPatternRewriter &rewriter)
static Value adjustAccessChainForBitwidth(const SPIRVTypeConverter &typeConverter, spirv::AccessChainOp op, int sourceBits, int targetBits, OpBuilder &builder)
Returns an adjusted spirv::AccessChainOp.
static bool isAllocationSupported(Operation *allocOp, MemRefType type)
Returns true if the allocations of memref type generated from allocOp can be lowered to SPIR-V.
static Value getOffsetForBitwidth(Location loc, Value srcIdx, int sourceBits, int targetBits, OpBuilder &builder)
Returns the offset of the value in targetBits representation.
static spirv::MemorySemantics getAtomicAcqRelMemorySemantics(MemRefType type)
Returns the AcquireRelease memory semantics OR'd with the storage-class bit derived from the memory s...
#define ATOMIC_CASE(kind, spirvOp)
static FailureOr< MemoryRequirements > calculateMemoryRequirements(Value accessedPtr, bool isNontemporal, uint64_t preferredAlignment)
Given an accessed SPIR-V pointer, calculates its alignment requirements, if any.
static Value castBoolToIntN(Location loc, Value srcBool, Type dstType, OpBuilder &builder)
Casts the given srcBool into an integer of dstType.
A multi-dimensional affine map Affine map's are immutable like Type's, and they are uniqued.
unsigned getDimPosition(unsigned idx) const
Extracts the position of the dimensional expression at the given result, when the caller knows it is ...
unsigned getNumDims() const
bool isPermutation() const
Returns true if the AffineMap represents a symbol-less permutation map.
OpListType::iterator iterator
auto getOps()
Return an iterator range over the operations within this block that are of 'OpT'.
IntegerAttr getIntegerAttr(Type type, int64_t value)
IntegerType getIntegerType(unsigned width)
This class defines the main interface for locations in MLIR and acts as a non-nullable wrapper around...
MLIRContext is the top-level object for a collection of MLIR operations.
This class helps build Operations.
void createOrFold(SmallVectorImpl< Value > &results, Location location, Args &&...args)
Create an operation of specific op type at the current insertion point, and immediately try to fold i...
Operation is the basic unit of execution within MLIR.
Region & getRegion(unsigned index)
Returns the region held by this operation at position 'index'.
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.
Type conversion from builtin types to SPIR-V types for shader interface.
bool allows(spirv::Capability capability) const
Checks if the SPIR-V capability inquired is supported.
static Operation * getNearestSymbolTable(Operation *from)
Returns the nearest symbol table from a given operation from.
Instances of the Type class are uniqued, have an immutable identifier and an optional mutable compone...
bool isInteger() const
Return true if this is an 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 represents an instance of an SSA value in the MLIR system, representing a computable value...
MLIRContext * getContext() const
Utility to get the associated MLIRContext that this value is defined in.
Type getType() const
Return the type of this value.
Operation * getDefiningOp() const
If this value is the result of an operation, return the operation that defines it.
Value getElementPtr(const SPIRVTypeConverter &typeConverter, MemRefType baseType, Value basePtr, ValueRange indices, Location loc, OpBuilder &builder)
Performs the index computation to get to the element at indices of the memory pointed to by basePtr,...
Include the generated interface declarations.
SmallVector< OpFoldResult > getMixedValues(ArrayRef< int64_t > staticValues, ValueRange dynamicValues, MLIRContext *context)
Return a vector of OpFoldResults with the same size a staticValues, but all elements for which Shaped...
Type getType(OpFoldResult ofr)
Returns the int type of the integer in ofr.
bool isZeroInteger(OpFoldResult v)
Return "true" if v is an integer value/attribute with constant value 0.
void populateMemRefToSPIRVPatterns(const SPIRVTypeConverter &typeConverter, RewritePatternSet &patterns)
Appends to a pattern list additional patterns for translating MemRef ops to SPIR-V ops.
spirv::MemoryAccessAttr memoryAccess