MLIR 24.0.0git
MemRefToSPIRV.cpp
Go to the documentation of this file.
1//===- MemRefToSPIRV.cpp - MemRef to SPIR-V Patterns ----------------------===//
2//
3// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.
4// See https://llvm.org/LICENSE.txt for license information.
5// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
6//
7//===----------------------------------------------------------------------===//
8//
9// This file implements patterns to convert MemRef dialect to SPIR-V dialect.
10//
11//===----------------------------------------------------------------------===//
12
22#include "mlir/IR/MLIRContext.h"
23#include "mlir/IR/Visitors.h"
25#include <cassert>
26#include <limits>
27#include <optional>
28
29#define DEBUG_TYPE "memref-to-spirv-pattern"
30
31using namespace mlir;
32
33//===----------------------------------------------------------------------===//
34// Utility functions
35//===----------------------------------------------------------------------===//
36
37/// Returns the offset of the value in `targetBits` representation.
38///
39/// `srcIdx` is an index into a 1-D array with each element having `sourceBits`.
40/// It's assumed to be non-negative.
41///
42/// When accessing an element in the array treating as having elements of
43/// `targetBits`, multiple values are loaded in the same time. The method
44/// returns the offset where the `srcIdx` locates in the value. For example, if
45/// `sourceBits` equals to 8 and `targetBits` equals to 32, the x-th element is
46/// located at (x % 4) * 8. Because there are four elements in one i32, and one
47/// element has 8 bits.
48static Value getOffsetForBitwidth(Location loc, Value srcIdx, int sourceBits,
49 int targetBits, OpBuilder &builder) {
50 assert(targetBits % sourceBits == 0);
51 Type type = srcIdx.getType();
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);
55 auto srcBitsValue =
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);
59}
60
61/// Returns an adjusted spirv::AccessChainOp. Based on the
62/// extension/capabilities, certain integer bitwidths `sourceBits` might not be
63/// supported. During conversion if a memref of an unsupported type is used,
64/// load/stores to this memref need to be modified to use a supported higher
65/// bitwidth `targetBits` and extracting the required bits. For an accessing a
66/// 1D array (spirv.array or spirv.rtarray), the last index is modified to load
67/// the bits needed. The extraction of the actual bits needed are handled
68/// separately. Note that this only works for a 1-D tensor.
69static Value
71 spirv::AccessChainOp op, int sourceBits,
72 int targetBits, OpBuilder &builder) {
73 assert(targetBits % sourceBits == 0);
74 const auto loc = op.getLoc();
75 Value lastDim = op->getOperand(op.getNumOperands() - 1);
76 Type type = lastDim.getType();
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());
80 // There are two elements if this is a 1-D tensor.
81 assert(indices.size() == 2);
82 indices.back() = builder.createOrFold<spirv::SDivOp>(loc, lastDim, idx);
83 Type t = typeConverter.convertType(op.getComponentPtr().getType());
84 return spirv::AccessChainOp::create(builder, loc, t, op.getBasePtr(),
85 indices);
86}
87
88/// Casts the given `srcBool` into an integer of `dstType`.
89static Value castBoolToIntN(Location loc, Value srcBool, Type dstType,
90 OpBuilder &builder) {
91 assert(srcBool.getType().isInteger(1));
92 if (dstType.isInteger(1))
93 return srcBool;
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,
97 zero);
98}
99
100/// Returns the `targetBits`-bit value shifted by the given `offset`, and cast
101/// to the type destination type, and masked.
102static Value shiftValue(Location loc, Value value, Value offset, Value mask,
103 OpBuilder &builder) {
104 IntegerType dstType = cast<IntegerType>(mask.getType());
105 int targetBits = static_cast<int>(dstType.getWidth());
106 int valueBits = value.getType().getIntOrFloatBitWidth();
107 assert(valueBits <= targetBits);
108
109 if (valueBits == 1) {
110 value = castBoolToIntN(loc, value, dstType, builder);
111 } else {
112 if (valueBits < targetBits) {
113 value = spirv::UConvertOp::create(
114 builder, loc, builder.getIntegerType(targetBits), value);
115 }
116
117 value = builder.createOrFold<spirv::BitwiseAndOp>(loc, value, mask);
118 }
119 return builder.createOrFold<spirv::ShiftLeftLogicalOp>(loc, value.getType(),
120 value, offset);
121}
122
123/// Returns true if the allocations of memref `type` generated from `allocOp`
124/// can be lowered to SPIR-V.
125static bool isAllocationSupported(Operation *allocOp, MemRefType type) {
126 // Currently only support static shape
127 if (!type.hasStaticShape())
128 return false;
129
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)
133 return false;
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)
137 return false;
138 // Function allocations of memref-compatible element types may be lowered
139 // to SPIRV pointers/arrays of the corresponding SPIRV element type.
140 if (isa<MemRefElementTypeInterface>(type.getElementType()))
141 return true;
142 } else {
143 return false;
144 }
145
146 // Support memref of int or float types, or their vector/complex types.
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();
152 return elementType.isIntOrFloat();
153}
154
155/// Returns the scope to use for atomic operations use for emulating store
156/// operations of unsupported integer bitwidths, based on the memref
157/// type. Returns std::nullopt on failure.
158static std::optional<spirv::Scope> getAtomicOpScope(MemRefType type) {
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;
165 default:
166 break;
167 }
168 return {};
169}
170
171/// Returns the MemorySemantics storage-class bit corresponding to `sc`.
172/// Per SPIR-V spec section 3.32 (Memory Semantics) this bit must be OR'd
173/// with the ordering bits (Acquire/Release/...) on atomic operations.
174static spirv::MemorySemantics
175getMemorySemanticsForStorageClass(spirv::StorageClass sc) {
176 switch (sc) {
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;
188 default:
189 return spirv::MemorySemantics::None;
190 }
191}
192
193/// Returns the AcquireRelease memory semantics OR'd with the storage-class
194/// bit derived from the memory space of `type`.
195static spirv::MemorySemantics getAtomicAcqRelMemorySemantics(MemRefType type) {
196 auto sc = cast<spirv::StorageClassAttr>(type.getMemorySpace()).getValue();
197 return spirv::MemorySemantics::AcquireRelease |
199}
200
201/// Extracts the element type from a SPIR-V pointer type pointing to storage.
202///
203/// For Kernel capability, the pointer points directly to the element type
204/// (possibly wrapped in an array). For Vulkan, the pointer points to a struct
205/// containing an array or runtime array, and we need to unwrap to get the
206/// element type.
207static Type
209 const SPIRVTypeConverter &typeConverter) {
210 if (typeConverter.allows(spirv::Capability::Kernel)) {
211 if (auto arrayType = dyn_cast<spirv::ArrayType>(pointeeType))
212 return arrayType.getElementType();
213 return pointeeType;
214 }
215 // For Vulkan we need to extract element from wrapping struct and array.
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();
220}
221
222/// Casts the given `srcInt` into a boolean value.
223static Value castIntNToBool(Location loc, Value srcInt, OpBuilder &builder) {
224 if (srcInt.getType().isInteger(1))
225 return srcInt;
226
227 auto one = spirv::ConstantOp::getZero(srcInt.getType(), loc, builder);
228 return builder.createOrFold<spirv::INotEqualOp>(loc, srcInt, one);
229}
230
231//===----------------------------------------------------------------------===//
232// Operation conversion
233//===----------------------------------------------------------------------===//
234
235// Note that DRR cannot be used for the patterns in this file: we may need to
236// convert type along the way, which requires ConversionPattern. DRR generates
237// normal RewritePattern.
238
239namespace {
240
241/// Converts memref.alloca to SPIR-V Function variables.
242class AllocaOpPattern final : public OpConversionPattern<memref::AllocaOp> {
243public:
244 using Base::Base;
245
246 LogicalResult
247 matchAndRewrite(memref::AllocaOp allocaOp, OpAdaptor adaptor,
248 ConversionPatternRewriter &rewriter) const override;
249};
250
251/// Converts an allocation operation to SPIR-V. Currently only supports lowering
252/// to Workgroup memory when the size is constant. Note that this pattern needs
253/// to be applied in a pass that runs at least at spirv.module scope since it
254/// wil ladd global variables into the spirv.module.
255class AllocOpPattern final : public OpConversionPattern<memref::AllocOp> {
256public:
257 using Base::Base;
258
259 LogicalResult
260 matchAndRewrite(memref::AllocOp operation, OpAdaptor adaptor,
261 ConversionPatternRewriter &rewriter) const override;
262};
263
264/// Converts memref.automic_rmw operations to SPIR-V atomic operations.
265class AtomicRMWOpPattern final
266 : public OpConversionPattern<memref::AtomicRMWOp> {
267public:
268 using Base::Base;
269
270 LogicalResult
271 matchAndRewrite(memref::AtomicRMWOp atomicOp, OpAdaptor adaptor,
272 ConversionPatternRewriter &rewriter) const override;
273};
274
275/// Removed a deallocation if it is a supported allocation. Currently only
276/// removes deallocation if the memory space is workgroup memory.
277class DeallocOpPattern final : public OpConversionPattern<memref::DeallocOp> {
278public:
279 using Base::Base;
280
281 LogicalResult
282 matchAndRewrite(memref::DeallocOp operation, OpAdaptor adaptor,
283 ConversionPatternRewriter &rewriter) const override;
284};
285
286/// Converts memref.load to spirv.Load + spirv.AccessChain on integers.
287class IntLoadOpPattern final : public OpConversionPattern<memref::LoadOp> {
288public:
289 using Base::Base;
290
291 LogicalResult
292 matchAndRewrite(memref::LoadOp loadOp, OpAdaptor adaptor,
293 ConversionPatternRewriter &rewriter) const override;
294};
295
296/// Converts memref.load to spirv.Load + spirv.AccessChain.
297class LoadOpPattern final : public OpConversionPattern<memref::LoadOp> {
298public:
299 using Base::Base;
300
301 LogicalResult
302 matchAndRewrite(memref::LoadOp loadOp, OpAdaptor adaptor,
303 ConversionPatternRewriter &rewriter) const override;
304};
305
306/// Converts memref.load to spirv.Image + spirv.ImageFetch
307class ImageLoadOpPattern final : public OpConversionPattern<memref::LoadOp> {
308public:
309 using Base::Base;
310
311 LogicalResult
312 matchAndRewrite(memref::LoadOp loadOp, OpAdaptor adaptor,
313 ConversionPatternRewriter &rewriter) const override;
314};
315
316/// Converts memref.store to spirv.Store on integers.
317class IntStoreOpPattern final : public OpConversionPattern<memref::StoreOp> {
318public:
319 using Base::Base;
320
321 LogicalResult
322 matchAndRewrite(memref::StoreOp storeOp, OpAdaptor adaptor,
323 ConversionPatternRewriter &rewriter) const override;
324};
325
326/// Converts memref.memory_space_cast to the appropriate spirv cast operations.
327class MemorySpaceCastOpPattern final
328 : public OpConversionPattern<memref::MemorySpaceCastOp> {
329public:
330 using Base::Base;
331
332 LogicalResult
333 matchAndRewrite(memref::MemorySpaceCastOp addrCastOp, OpAdaptor adaptor,
334 ConversionPatternRewriter &rewriter) const override;
335};
336
337/// Converts memref.store to spirv.Store.
338class StoreOpPattern final : public OpConversionPattern<memref::StoreOp> {
339public:
340 using Base::Base;
341
342 LogicalResult
343 matchAndRewrite(memref::StoreOp storeOp, OpAdaptor adaptor,
344 ConversionPatternRewriter &rewriter) const override;
345};
346
347/// Converts memref.copy to spirv.CopyMemory.
348class CopyOpPattern final : public OpConversionPattern<memref::CopyOp> {
349public:
350 using Base::Base;
351
352 LogicalResult
353 matchAndRewrite(memref::CopyOp copyOp, OpAdaptor adaptor,
354 ConversionPatternRewriter &rewriter) const override;
355};
356
357class ReinterpretCastPattern final
358 : public OpConversionPattern<memref::ReinterpretCastOp> {
359public:
360 using Base::Base;
361
362 LogicalResult
363 matchAndRewrite(memref::ReinterpretCastOp op, OpAdaptor adaptor,
364 ConversionPatternRewriter &rewriter) const override;
365};
366
367class CastPattern final : public OpConversionPattern<memref::CastOp> {
368public:
369 using Base::Base;
370
371 LogicalResult
372 matchAndRewrite(memref::CastOp op, OpAdaptor adaptor,
373 ConversionPatternRewriter &rewriter) const override {
374 Value src = adaptor.getSource();
375 Type srcType = src.getType();
376
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;
382 });
383
384 rewriter.replaceOp(op, src);
385 return success();
386 }
387};
388
389/// Converts memref.extract_aligned_pointer_as_index to spirv.ConvertPtrToU.
390class ExtractAlignedPointerAsIndexOpPattern final
391 : public OpConversionPattern<memref::ExtractAlignedPointerAsIndexOp> {
392public:
393 using Base::Base;
394
395 LogicalResult
396 matchAndRewrite(memref::ExtractAlignedPointerAsIndexOp extractOp,
397 OpAdaptor adaptor,
398 ConversionPatternRewriter &rewriter) const override;
399};
400} // namespace
401
402//===----------------------------------------------------------------------===//
403// AllocaOp
404//===----------------------------------------------------------------------===//
405
406LogicalResult
407AllocaOpPattern::matchAndRewrite(memref::AllocaOp allocaOp, OpAdaptor adaptor,
408 ConversionPatternRewriter &rewriter) const {
409 MemRefType allocType = allocaOp.getType();
410 if (!isAllocationSupported(allocaOp, allocType))
411 return rewriter.notifyMatchFailure(allocaOp, "unhandled allocation type");
412
413 // Get the SPIR-V type for the allocation.
414 Type spirvType = getTypeConverter()->convertType(allocType);
415 if (!spirvType)
416 return rewriter.notifyMatchFailure(allocaOp, "type conversion failed");
417
418 auto function = allocaOp->getParentOfType<FunctionOpInterface>();
419 if (!function)
420 return rewriter.notifyMatchFailure(allocaOp,
421 "requires a containing function");
422
423 // SPIR-V requires Function variables to be declared in the first block.
424 OpBuilder::InsertionGuard guard(rewriter);
425 Block &entryBlock = function->getRegion(0).front();
426 Block::iterator insertionPoint = entryBlock.begin();
427 // Insert the variable after any existing ones to preserve ordering.
428 while (insertionPoint != entryBlock.end() &&
429 isa<spirv::VariableOp>(*insertionPoint))
430 ++insertionPoint;
431 rewriter.setInsertionPoint(&entryBlock, insertionPoint);
432 Value variable = spirv::VariableOp::create(
433 rewriter, allocaOp.getLoc(), spirvType, spirv::StorageClass::Function,
434 /*initializer=*/nullptr);
435 rewriter.replaceOp(allocaOp, variable);
436 return success();
437}
438
439//===----------------------------------------------------------------------===//
440// AllocOp
441//===----------------------------------------------------------------------===//
442
443LogicalResult
444AllocOpPattern::matchAndRewrite(memref::AllocOp operation, OpAdaptor adaptor,
445 ConversionPatternRewriter &rewriter) const {
446 MemRefType allocType = operation.getType();
447 if (!isAllocationSupported(operation, allocType))
448 return rewriter.notifyMatchFailure(operation, "unhandled allocation type");
449
450 // Get the SPIR-V type for the allocation.
451 Type spirvType = getTypeConverter()->convertType(allocType);
452 if (!spirvType)
453 return rewriter.notifyMatchFailure(operation, "type conversion failed");
454
455 // Insert spirv.GlobalVariable for this allocation.
456 Operation *parent =
457 SymbolTable::getNearestSymbolTable(operation->getParentOp());
458 if (!parent)
459 return failure();
460 Location loc = operation.getLoc();
461 spirv::GlobalVariableOp varOp;
462 {
463 OpBuilder::InsertionGuard guard(rewriter);
464 Block &entryBlock = *parent->getRegion(0).begin();
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,
471 /*initializer=*/nullptr);
472 }
473
474 // Get pointer to global variable at the current scope.
475 rewriter.replaceOpWithNewOp<spirv::AddressOfOp>(operation, varOp);
476 return success();
477}
478
479//===----------------------------------------------------------------------===//
480// AllocOp
481//===----------------------------------------------------------------------===//
482
483LogicalResult
484AtomicRMWOpPattern::matchAndRewrite(memref::AtomicRMWOp atomicOp,
485 OpAdaptor adaptor,
486 ConversionPatternRewriter &rewriter) const {
487 auto memrefType = cast<MemRefType>(atomicOp.getMemref().getType());
488 std::optional<spirv::Scope> scope = getAtomicOpScope(memrefType);
489 if (!scope)
490 return rewriter.notifyMatchFailure(atomicOp,
491 "unsupported memref memory space");
492
493 auto &typeConverter = *getTypeConverter<SPIRVTypeConverter>();
494 Type resultType = typeConverter.convertType(atomicOp.getType());
495 if (!resultType)
496 return rewriter.notifyMatchFailure(atomicOp,
497 "failed to convert result type");
498
499 auto loc = atomicOp.getLoc();
500 Value ptr =
501 spirv::getElementPtr(typeConverter, memrefType, adaptor.getMemref(),
502 adaptor.getIndices(), loc, rewriter);
503
504 if (!ptr)
505 return failure();
506
507 // Determine the source and destination bitwidths. The source is the original
508 // memref element type and the destination is the SPIR-V storage type (e.g.,
509 // i32 for Vulkan).
510 int srcBits = memrefType.getElementType().getIntOrFloatBitWidth();
511 auto pointerType = typeConverter.convertType<spirv::PointerType>(memrefType);
512 if (!pointerType)
513 return rewriter.notifyMatchFailure(atomicOp,
514 "failed to convert memref type");
515
516 Type pointeeType = pointerType.getPointeeType();
517 Type storageElemType =
518 getElementTypeForStoragePointer(pointeeType, typeConverter);
519 if (!storageElemType || !storageElemType.isIntOrFloat())
520 return rewriter.notifyMatchFailure(
521 atomicOp, "failed to determine destination element type");
522
523 int dstBits = static_cast<int>(storageElemType.getIntOrFloatBitWidth());
524 assert(dstBits % srcBits == 0);
525
526 spirv::MemorySemantics memSem = getAtomicAcqRelMemorySemantics(memrefType);
527
528 // When the source and destination bitwidths match, emit the atomic operation
529 // directly.
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()); \
535 break
536
537 switch (atomicOp.getKind()) {
538 ATOMIC_CASE(addf, EXTAtomicFAddOp);
539 ATOMIC_CASE(addi, AtomicIAddOp);
540 ATOMIC_CASE(maxs, AtomicSMaxOp);
541 ATOMIC_CASE(maxu, AtomicUMaxOp);
542 ATOMIC_CASE(mins, AtomicSMinOp);
543 ATOMIC_CASE(minu, AtomicUMinOp);
544 ATOMIC_CASE(ori, AtomicOrOp);
545 ATOMIC_CASE(andi, AtomicAndOp);
546 ATOMIC_CASE(xori, AtomicXorOp);
547 default:
548 return rewriter.notifyMatchFailure(atomicOp, "unimplemented atomic kind");
549 }
550
551#undef ATOMIC_CASE
552
553 return success();
554 }
555
556 // Sub-element-width atomic: the element type (e.g., i8) is narrower than the
557 // storage type (e.g., i32). We need to adjust the index and shift/mask the
558 // value to operate on the correct bits within the wider storage element.
559 //
560 // Only ori and andi can be emulated because they operate bitwise and don't
561 // carry across byte boundaries. Other kinds (addi, max, min) would require
562 // CAS loops.
563 if (atomicOp.getKind() != arith::AtomicRMWKind::ori &&
564 atomicOp.getKind() != arith::AtomicRMWKind::andi) {
565 return rewriter.notifyMatchFailure(
566 atomicOp,
567 "atomic op on sub-element-width types is only supported for ori/andi");
568 }
569
570 // Bitcasting is currently unsupported for Kernel capability /
571 // spirv.PtrAccessChain.
572 if (typeConverter.allows(spirv::Capability::Kernel))
573 return rewriter.notifyMatchFailure(
574 atomicOp,
575 "sub-element-width atomic ops unsupported with Kernel capability");
576
577 auto dstType = cast<IntegerType>(storageElemType);
578
579 auto accessChainOp = ptr.getDefiningOp<spirv::AccessChainOp>();
580 if (!accessChainOp)
581 return failure();
582
583 // Compute the bit offset within the storage element and adjust the pointer
584 // to address the containing storage element.
585 assert(accessChainOp.getIndices().size() == 2);
586 Value lastDim = accessChainOp->getOperand(accessChainOp.getNumOperands() - 1);
587 Value offset = getOffsetForBitwidth(loc, lastDim, srcBits, dstBits, rewriter);
588 Value adjustedPtr = adjustAccessChainForBitwidth(typeConverter, accessChainOp,
589 srcBits, dstBits, rewriter);
590 Value result;
591 switch (atomicOp.getKind()) {
592 case arith::AtomicRMWKind::ori: {
593 // OR only sets bits, so shifting the value to the target position and
594 // ORing with zeros in other positions preserves the unaffected bits.
595 Value elemMask = rewriter.createOrFold<spirv::ConstantOp>(
596 loc, dstType, rewriter.getIntegerAttr(dstType, (1uLL << srcBits) - 1));
597 Value storeVal =
598 shiftValue(loc, adaptor.getValue(), offset, elemMask, rewriter);
599 result = spirv::AtomicOrOp::create(rewriter, loc, dstType, adjustedPtr,
600 *scope, memSem, storeVal);
601 break;
602 }
603 case arith::AtomicRMWKind::andi: {
604 // Build a mask that preserves all bits outside the target element
605 // and applies the operand mask to the target element.
606 // mask = (operand << offset) | ~(elemMask << offset)
607 Value elemMask = rewriter.createOrFold<spirv::ConstantOp>(
608 loc, dstType, rewriter.getIntegerAttr(dstType, (1uLL << srcBits) - 1));
609 Value storeVal =
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,
616 invertedElemMask);
617 result = spirv::AtomicAndOp::create(rewriter, loc, dstType, adjustedPtr,
618 *scope, memSem, mask);
619 break;
620 }
621 default:
622 return rewriter.notifyMatchFailure(atomicOp, "unimplemented atomic kind");
623 }
624
625 // The atomic op returns the old value of the full storage element (e.g.,
626 // i32). Extract the original sub-element value from the correct position.
627 result = rewriter.createOrFold<spirv::ShiftRightLogicalOp>(loc, dstType,
628 result, offset);
629 Value mask = rewriter.createOrFold<spirv::ConstantOp>(
630 loc, dstType, rewriter.getIntegerAttr(dstType, (1uLL << srcBits) - 1));
631 result =
632 rewriter.createOrFold<spirv::BitwiseAndOp>(loc, dstType, result, mask);
633 rewriter.replaceOp(atomicOp, result);
634
635 return success();
636}
637
638//===----------------------------------------------------------------------===//
639// DeallocOp
640//===----------------------------------------------------------------------===//
641
642LogicalResult
643DeallocOpPattern::matchAndRewrite(memref::DeallocOp operation,
644 OpAdaptor adaptor,
645 ConversionPatternRewriter &rewriter) const {
646 MemRefType deallocType = cast<MemRefType>(operation.getMemref().getType());
647 if (!isAllocationSupported(operation, deallocType))
648 return rewriter.notifyMatchFailure(operation, "unhandled allocation type");
649 rewriter.eraseOp(operation);
650 return success();
651}
652
653//===----------------------------------------------------------------------===//
654// LoadOp
655//===----------------------------------------------------------------------===//
656
658 spirv::MemoryAccessAttr memoryAccess;
659 IntegerAttr alignment;
660};
661
662/// Given an accessed SPIR-V pointer, calculates its alignment requirements, if
663/// any.
664static FailureOr<MemoryRequirements>
665calculateMemoryRequirements(Value accessedPtr, bool isNontemporal,
666 uint64_t preferredAlignment) {
667 if (preferredAlignment >= std::numeric_limits<uint32_t>::max()) {
668 return failure();
669 }
670
671 MLIRContext *ctx = accessedPtr.getContext();
672
673 auto memoryAccess = spirv::MemoryAccess::None;
674 if (isNontemporal) {
675 memoryAccess = spirv::MemoryAccess::Nontemporal;
676 }
677
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) {
684 return MemoryRequirements{spirv::MemoryAccessAttr{}, IntegerAttr{}};
685 }
686 return MemoryRequirements{spirv::MemoryAccessAttr::get(ctx, memoryAccess),
687 IntegerAttr{}};
688 }
689
690 // PhysicalStorageBuffers require the `Aligned` attribute.
691 // Other storage types may show an `Aligned` attribute.
692 std::optional<int64_t> sizeInBytes;
693 Type rawPointeeType = ptrType.getPointeeType();
694 if (auto scalarType = dyn_cast<spirv::ScalarType>(rawPointeeType)) {
695 // For scalar types, the alignment is determined by their size.
696 sizeInBytes = scalarType.getSizeInBytes();
697 } else if (auto vecType = dyn_cast<VectorType>(rawPointeeType)) {
698 // For vector element types, the alignment should equal the total size of
699 // the vector.
700 if (auto scalarElem =
701 dyn_cast<spirv::ScalarType>(vecType.getElementType())) {
702 if (auto elemSize = scalarElem.getSizeInBytes())
703 sizeInBytes = *elemSize * vecType.getNumElements();
704 }
705 }
706
707 if (!sizeInBytes.has_value())
708 return failure();
709
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);
714 return MemoryRequirements{memAccessAttr, alignment};
715}
716
717/// Given an accessed SPIR-V pointer and the original memref load/store
718/// `memAccess` op, calculates the alignment requirements, if any. Takes into
719/// account the alignment attributes applied to the load/store op.
720template <class LoadOrStoreOp>
721static FailureOr<MemoryRequirements>
722calculateMemoryRequirements(Value accessedPtr, LoadOrStoreOp loadOrStoreOp) {
723 static_assert(
724 llvm::is_one_of<LoadOrStoreOp, memref::LoadOp, memref::StoreOp>::value,
725 "Must be called on either memref::LoadOp or memref::StoreOp");
726
727 return calculateMemoryRequirements(accessedPtr,
728 loadOrStoreOp.getNontemporal(),
729 loadOrStoreOp.getAlignment().value_or(0));
730}
731
732LogicalResult
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())
738 return failure();
739
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");
745
746 if (memorySpaceAttr.getValue() == spirv::StorageClass::Image)
747 return rewriter.notifyMatchFailure(
748 loadOp,
749 "failed to lower memref in image storage class to storage buffer");
750
751 const auto &typeConverter = *getTypeConverter<SPIRVTypeConverter>();
752 Value accessChain =
753 spirv::getElementPtr(typeConverter, memrefType, adaptor.getMemref(),
754 adaptor.getIndices(), loc, rewriter);
755
756 if (!accessChain)
757 return failure();
758
759 int srcBits = memrefType.getElementType().getIntOrFloatBitWidth();
760 bool isBool = srcBits == 1;
761 if (isBool)
762 srcBits = typeConverter.getOptions().boolNumBits;
763
764 auto pointerType = typeConverter.convertType<spirv::PointerType>(memrefType);
765 if (!pointerType)
766 return rewriter.notifyMatchFailure(loadOp, "failed to convert memref type");
767
768 Type pointeeType = pointerType.getPointeeType();
769 Type dstType = getElementTypeForStoragePointer(pointeeType, typeConverter);
770 int dstBits = dstType.getIntOrFloatBitWidth();
771 assert(dstBits % srcBits == 0);
772
773 // If the rewritten load op has the same bit width, use the loading value
774 // directly.
775 if (srcBits == dstBits) {
776 auto memoryRequirements = calculateMemoryRequirements(accessChain, loadOp);
777 if (failed(memoryRequirements))
778 return rewriter.notifyMatchFailure(
779 loadOp, "failed to determine memory requirements");
780
781 auto [memoryAccess, alignment] = *memoryRequirements;
782 Value loadVal = spirv::LoadOp::create(rewriter, loc, accessChain,
783 memoryAccess, alignment);
784 if (isBool)
785 loadVal = castIntNToBool(loc, loadVal, rewriter);
786 rewriter.replaceOp(loadOp, loadVal);
787 return success();
788 }
789
790 // Bitcasting is currently unsupported for Kernel capability /
791 // spirv.PtrAccessChain.
792 if (typeConverter.allows(spirv::Capability::Kernel))
793 return failure();
794
795 auto accessChainOp = accessChain.getDefiningOp<spirv::AccessChainOp>();
796 if (!accessChainOp)
797 return failure();
798
799 // Assume that getElementPtr() works linearizely. If it's a scalar, the method
800 // still returns a linearized accessing. If the accessing is not linearized,
801 // there will be offset issues.
802 assert(accessChainOp.getIndices().size() == 2);
803 Value adjustedPtr = adjustAccessChainForBitwidth(typeConverter, accessChainOp,
804 srcBits, dstBits, rewriter);
805 auto memoryRequirements = calculateMemoryRequirements(adjustedPtr, loadOp);
806 if (failed(memoryRequirements))
807 return rewriter.notifyMatchFailure(
808 loadOp, "failed to determine memory requirements");
809
810 auto [memoryAccess, alignment] = *memoryRequirements;
811 Value spvLoadOp = spirv::LoadOp::create(rewriter, loc, dstType, adjustedPtr,
812 memoryAccess, alignment);
813
814 // Shift the bits to the rightmost.
815 // ____XXXX________ -> ____________XXXX
816 Value lastDim = accessChainOp->getOperand(accessChainOp.getNumOperands() - 1);
817 Value offset = getOffsetForBitwidth(loc, lastDim, srcBits, dstBits, rewriter);
818 Value result = rewriter.createOrFold<spirv::ShiftRightArithmeticOp>(
819 loc, spvLoadOp.getType(), spvLoadOp, offset);
820
821 // Apply the mask to extract corresponding bits.
822 Value mask = rewriter.createOrFold<spirv::ConstantOp>(
823 loc, dstType, rewriter.getIntegerAttr(dstType, (1 << srcBits) - 1));
824 result =
825 rewriter.createOrFold<spirv::BitwiseAndOp>(loc, dstType, result, mask);
826
827 // Apply sign extension on the loading value unconditionally. The signedness
828 // semantic is carried in the operator itself, we relies other pattern to
829 // handle the casting.
830 IntegerAttr shiftValueAttr =
831 rewriter.getIntegerAttr(dstType, dstBits - srcBits);
832 Value shiftValue =
833 rewriter.createOrFold<spirv::ConstantOp>(loc, dstType, shiftValueAttr);
834 result = rewriter.createOrFold<spirv::ShiftLeftLogicalOp>(loc, dstType,
836 result = rewriter.createOrFold<spirv::ShiftRightArithmeticOp>(
837 loc, dstType, result, shiftValue);
838
839 rewriter.replaceOp(loadOp, result);
840
841 assert(accessChainOp.use_empty());
842 rewriter.eraseOp(accessChainOp);
843
844 return success();
845}
846
847LogicalResult
848LoadOpPattern::matchAndRewrite(memref::LoadOp loadOp, OpAdaptor adaptor,
849 ConversionPatternRewriter &rewriter) const {
850 auto memrefType = cast<MemRefType>(loadOp.getMemref().getType());
851 if (memrefType.getElementType().isSignlessInteger())
852 return failure();
853
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");
859
860 if (memorySpaceAttr.getValue() == spirv::StorageClass::Image)
861 return rewriter.notifyMatchFailure(
862 loadOp,
863 "failed to lower memref in image storage class to storage buffer");
864
865 Value loadPtr = spirv::getElementPtr(
866 *getTypeConverter<SPIRVTypeConverter>(), memrefType, adaptor.getMemref(),
867 adaptor.getIndices(), loadOp.getLoc(), rewriter);
868
869 if (!loadPtr)
870 return failure();
871
872 auto memoryRequirements = calculateMemoryRequirements(loadPtr, loadOp);
873 if (failed(memoryRequirements))
874 return rewriter.notifyMatchFailure(
875 loadOp, "failed to determine memory requirements");
876
877 auto [memoryAccess, alignment] = *memoryRequirements;
878 rewriter.replaceOpWithNewOp<spirv::LoadOp>(loadOp, loadPtr, memoryAccess,
879 alignment);
880 return success();
881}
882
883template <typename OpAdaptor>
884static FailureOr<SmallVector<Value>>
885extractLoadCoordsForComposite(memref::LoadOp loadOp, OpAdaptor adaptor,
886 ConversionPatternRewriter &rewriter) {
887 // At present we only support linear "tiling" as specified in Vulkan, this
888 // means that texels are assumed to be laid out in memory in a row-major
889 // order. This allows us to support any memref layout that is a permutation of
890 // the dimensions. Future work will pass an optional image layout to the
891 // rewrite pattern so that we can support optimized target specific tilings.
892 SmallVector<Value> indices = adaptor.getIndices();
893 AffineMap map = loadOp.getMemRefType().getLayout().getAffineMap();
894 if (!map.isPermutation())
895 return rewriter.notifyMatchFailure(
896 loadOp,
897 "Cannot lower memrefs with memory layout which is not a permutation");
898
899 // The memrefs layout determines the dimension ordering so we need to follow
900 // the map to get the ordering of the dimensions/indices.
901 const unsigned dimCount = map.getNumDims();
902 SmallVector<Value, 3> coords(dimCount);
903 for (unsigned dim = 0; dim < dimCount; ++dim)
904 coords[map.getDimPosition(dim)] = indices[dim];
905
906 // We need to reverse the coordinates because the memref layout is slowest to
907 // fastest moving and the vector coordinates for the image op is fastest to
908 // slowest moving.
909 return llvm::to_vector(llvm::reverse(coords));
910}
911
912LogicalResult
913ImageLoadOpPattern::matchAndRewrite(memref::LoadOp loadOp, OpAdaptor adaptor,
914 ConversionPatternRewriter &rewriter) const {
915 auto memrefType = cast<MemRefType>(loadOp.getMemref().getType());
916
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");
922
923 if (memorySpaceAttr.getValue() != spirv::StorageClass::Image)
924 return rewriter.notifyMatchFailure(
925 loadOp, "failed to lower memref in non-image storage class to image");
926
927 Value loadPtr = adaptor.getMemref();
928 auto memoryRequirements = calculateMemoryRequirements(loadPtr, loadOp);
929 if (failed(memoryRequirements))
930 return rewriter.notifyMatchFailure(
931 loadOp, "failed to determine memory requirements");
932
933 const auto [memoryAccess, alignment] = *memoryRequirements;
934
935 if (!loadOp.getMemRefType().hasRank())
936 return rewriter.notifyMatchFailure(
937 loadOp, "cannot lower unranked memrefs to SPIR-V images");
938
939 // We currently only support lowering of scalar memref elements to texels in
940 // the R[16|32][f|i|ui] formats. Future work will enable lowering of vector
941 // elements to texels in richer formats.
942 if (!isa<spirv::ScalarType>(loadOp.getMemRefType().getElementType()))
943 return rewriter.notifyMatchFailure(
944 loadOp,
945 "cannot lower memrefs who's element type is not a SPIR-V scalar type"
946 "to SPIR-V images");
947
948 // We currently only support sampled images since OpImageFetch does not work
949 // for plain images and the OpImageRead instruction needs to be materialized
950 // instead or texels need to be accessed via atomics through a texel pointer.
951 // Future work will generalize support to plain images.
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");
958
959 // Materialize the lowering.
960 Location loc = loadOp->getLoc();
961 auto imageLoadOp =
962 spirv::LoadOp::create(rewriter, loc, loadPtr, memoryAccess, alignment);
963 // Extract the image from the sampled image.
964 auto imageOp = spirv::ImageOp::create(rewriter, loc, imageLoadOp);
965
966 // Build a vector of coordinates or just a scalar index if we have a 1D image.
967 Value coords;
968 if (memrefType.getRank() == 1) {
969 coords = adaptor.getIndices()[0];
970 } else {
971 FailureOr<SmallVector<Value>> maybeCoords =
972 extractLoadCoordsForComposite(loadOp, adaptor, rewriter);
973 if (failed(maybeCoords))
974 return failure();
975 auto coordVectorType = VectorType::get({loadOp.getMemRefType().getRank()},
976 adaptor.getIndices().getType()[0]);
977 coords = spirv::CompositeConstructOp::create(rewriter, loc, coordVectorType,
978 maybeCoords.value());
979 }
980
981 // Fetch the value out of the image.
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{});
986
987 // Note that because OpImageFetch returns a rank 4 vector we need to extract
988 // the elements corresponding to the load which will since we only support the
989 // R[16|32][f|i|ui] formats will always be the R(red) 0th vector element.
990 auto compositeExtractOp =
991 spirv::CompositeExtractOp::create(rewriter, loc, fetchOp, 0);
992
993 rewriter.replaceOp(loadOp, compositeExtractOp);
994 return success();
995}
996
997LogicalResult
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");
1004
1005 auto loc = storeOp.getLoc();
1006 auto &typeConverter = *getTypeConverter<SPIRVTypeConverter>();
1007 Value accessChain =
1008 spirv::getElementPtr(typeConverter, memrefType, adaptor.getMemref(),
1009 adaptor.getIndices(), loc, rewriter);
1010
1011 if (!accessChain)
1012 return rewriter.notifyMatchFailure(
1013 storeOp, "failed to convert element pointer type");
1014
1015 int srcBits = memrefType.getElementType().getIntOrFloatBitWidth();
1016
1017 bool isBool = srcBits == 1;
1018 if (isBool)
1019 srcBits = typeConverter.getOptions().boolNumBits;
1020
1021 auto pointerType = typeConverter.convertType<spirv::PointerType>(memrefType);
1022 if (!pointerType)
1023 return rewriter.notifyMatchFailure(storeOp,
1024 "failed to convert memref type");
1025
1026 Type pointeeType = pointerType.getPointeeType();
1027 auto dstType = dyn_cast<IntegerType>(
1028 getElementTypeForStoragePointer(pointeeType, typeConverter));
1029 if (!dstType)
1030 return rewriter.notifyMatchFailure(
1031 storeOp, "failed to determine destination element type");
1032
1033 int dstBits = static_cast<int>(dstType.getWidth());
1034 assert(dstBits % srcBits == 0);
1035
1036 if (srcBits == dstBits) {
1037 auto memoryRequirements = calculateMemoryRequirements(accessChain, storeOp);
1038 if (failed(memoryRequirements))
1039 return rewriter.notifyMatchFailure(
1040 storeOp, "failed to determine memory requirements");
1041
1042 auto [memoryAccess, alignment] = *memoryRequirements;
1043 Value storeVal = adaptor.getValue();
1044 if (isBool)
1045 storeVal = castBoolToIntN(loc, storeVal, dstType, rewriter);
1046 rewriter.replaceOpWithNewOp<spirv::StoreOp>(storeOp, accessChain, storeVal,
1047 memoryAccess, alignment);
1048 return success();
1049 }
1050
1051 // Bitcasting is currently unsupported for Kernel capability /
1052 // spirv.PtrAccessChain.
1053 if (typeConverter.allows(spirv::Capability::Kernel))
1054 return failure();
1055
1056 auto accessChainOp = accessChain.getDefiningOp<spirv::AccessChainOp>();
1057 if (!accessChainOp)
1058 return failure();
1059
1060 // Since there are multiple threads in the processing, the emulation will be
1061 // done with atomic operations. E.g., if the stored value is i8, rewrite the
1062 // StoreOp to:
1063 // 1) load a 32-bit integer
1064 // 2) clear 8 bits in the loaded value
1065 // 3) set 8 bits in the loaded value
1066 // 4) store 32-bit value back
1067 //
1068 // Step 2 is done with AtomicAnd, and step 3 is done with AtomicOr (of the
1069 // loaded 32-bit value and the shifted 8-bit store value) as another atomic
1070 // step.
1071 assert(accessChainOp.getIndices().size() == 2);
1072 Value lastDim = accessChainOp->getOperand(accessChainOp.getNumOperands() - 1);
1073 Value offset = getOffsetForBitwidth(loc, lastDim, srcBits, dstBits, rewriter);
1074
1075 // Create a mask to clear the destination. E.g., if it is the second i8 in
1076 // i32, 0xFFFF00FF is created.
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);
1081 clearBitsMask =
1082 rewriter.createOrFold<spirv::NotOp>(loc, dstType, clearBitsMask);
1083
1084 Value storeVal = shiftValue(loc, adaptor.getValue(), offset, mask, rewriter);
1085 Value adjustedPtr = adjustAccessChainForBitwidth(typeConverter, accessChainOp,
1086 srcBits, dstBits, rewriter);
1087 std::optional<spirv::Scope> scope = getAtomicOpScope(memrefType);
1088 if (!scope)
1089 return rewriter.notifyMatchFailure(storeOp, "atomic scope not available");
1090
1091 spirv::MemorySemantics memSem = getAtomicAcqRelMemorySemantics(memrefType);
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);
1096
1097 // The AtomicOrOp has no side effect. Since it is already inserted, we can
1098 // just remove the original StoreOp. Note that rewriter.replaceOp()
1099 // doesn't work because it only accepts that the numbers of result are the
1100 // same.
1101 rewriter.eraseOp(storeOp);
1102
1103 assert(accessChainOp.use_empty());
1104 rewriter.eraseOp(accessChainOp);
1105
1106 return success();
1107}
1108
1109//===----------------------------------------------------------------------===//
1110// MemorySpaceCastOp
1111//===----------------------------------------------------------------------===//
1112
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");
1121
1122 auto sourceType = dyn_cast<MemRefType>(addrCastOp.getSource().getType());
1123 if (!sourceType)
1124 return rewriter.notifyMatchFailure(
1125 loc, "SPIR-V lowering requires ranked memref types");
1126 auto resultType = cast<MemRefType>(addrCastOp.getResult().getType());
1127
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";
1134 });
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";
1141 });
1142
1143 spirv::StorageClass sourceSc = sourceStorageClassAttr.getValue();
1144 spirv::StorageClass resultSc = resultStorageClassAttr.getValue();
1145
1146 Value result = adaptor.getSource();
1147 Type resultPtrType = typeConverter.convertType(resultType);
1148 if (!resultPtrType)
1149 return rewriter.notifyMatchFailure(addrCastOp,
1150 "failed to convert memref type");
1151
1152 Type genericPtrType = resultPtrType;
1153 // SPIR-V doesn't have a general address space cast operation. Instead, it has
1154 // conversions to and from generic pointers. To implement the general case,
1155 // we use specific-to-generic conversions when the source class is not
1156 // generic. Then when the result storage class is not generic, we convert the
1157 // generic pointer (either the input on ar intermediate result) to that
1158 // class. This also means that we'll need the intermediate generic pointer
1159 // type if neither the source or destination have it.
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);
1168 }
1169 if (sourceSc != spirv::StorageClass::Generic) {
1170 result = spirv::PtrCastToGenericOp::create(rewriter, loc, genericPtrType,
1171 result);
1172 }
1173 if (resultSc != spirv::StorageClass::Generic) {
1174 result =
1175 spirv::GenericCastToPtrOp::create(rewriter, loc, resultPtrType, result);
1176 }
1177 rewriter.replaceOp(addrCastOp, result);
1178 return success();
1179}
1180
1181LogicalResult
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");
1187 auto storePtr = spirv::getElementPtr(
1188 *getTypeConverter<SPIRVTypeConverter>(), memrefType, adaptor.getMemref(),
1189 adaptor.getIndices(), storeOp.getLoc(), rewriter);
1190
1191 if (!storePtr)
1192 return rewriter.notifyMatchFailure(storeOp, "type conversion failed");
1193
1194 auto memoryRequirements = calculateMemoryRequirements(storePtr, storeOp);
1195 if (failed(memoryRequirements))
1196 return rewriter.notifyMatchFailure(
1197 storeOp, "failed to determine memory requirements");
1198
1199 auto [memoryAccess, alignment] = *memoryRequirements;
1200 rewriter.replaceOpWithNewOp<spirv::StoreOp>(
1201 storeOp, storePtr, adaptor.getValue(), memoryAccess, alignment);
1202 return success();
1203}
1204
1205//===----------------------------------------------------------------------===//
1206// CopyOp
1207//===----------------------------------------------------------------------===//
1208
1209LogicalResult
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");
1215
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");
1224 }
1225
1226 // The converted operands are SPIR-V pointers to the source and target
1227 // storage. spirv.CopyMemory copies the whole pointed-to object, so it only
1228 // applies when both pointers point to the same fixed-size element type.
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");
1235
1236 if (sourcePtrType.getPointeeType() != targetPtrType.getPointeeType())
1237 return rewriter.notifyMatchFailure(
1238 copyOp, "source and target pointee types do not match");
1239
1240 rewriter.replaceOpWithNewOp<spirv::CopyMemoryOp>(
1241 copyOp, target, source, /*memory_access=*/spirv::MemoryAccessAttr{},
1242 /*alignment=*/IntegerAttr{}, /*source_memory_access=*/
1243 spirv::MemoryAccessAttr{}, /*source_alignment=*/IntegerAttr{});
1244 return success();
1245}
1246
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());
1252
1253 if (!srcType)
1254 return rewriter.notifyMatchFailure(op, [&](Diagnostic &diag) {
1255 diag << "invalid src type " << src.getType();
1256 });
1257
1258 const TypeConverter *converter = getTypeConverter();
1259
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();
1264 });
1265
1266 OpFoldResult offset =
1267 getMixedValues(adaptor.getStaticOffsets(), adaptor.getOffsets(), rewriter)
1268 .front();
1269 if (isZeroInteger(offset)) {
1270 rewriter.replaceOp(op, src);
1271 return success();
1272 }
1273
1274 Type intType = converter->convertType(rewriter.getIndexType());
1275 if (!intType)
1276 return rewriter.notifyMatchFailure(op, "failed to convert index type");
1277
1278 Location loc = op.getLoc();
1279 auto offsetValue = [&]() -> Value {
1280 if (auto val = dyn_cast<Value>(offset))
1281 return val;
1282
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);
1286 }();
1287
1288 rewriter.replaceOpWithNewOp<spirv::InBoundsPtrAccessChainOp>(
1289 op, src, offsetValue, ValueRange());
1290 return success();
1291}
1292
1293//===----------------------------------------------------------------------===//
1294// ExtractAlignedPointerAsIndexOp
1295//===----------------------------------------------------------------------===//
1296
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());
1304 return success();
1305}
1306
1307//===----------------------------------------------------------------------===//
1308// Pattern population
1309//===----------------------------------------------------------------------===//
1310
1311namespace mlir {
1313 RewritePatternSet &patterns) {
1314 patterns.add<AllocaOpPattern, AllocOpPattern, AtomicRMWOpPattern,
1315 CopyOpPattern, DeallocOpPattern, IntLoadOpPattern,
1316 ImageLoadOpPattern, IntStoreOpPattern, LoadOpPattern,
1317 MemorySpaceCastOpPattern, StoreOpPattern, ReinterpretCastPattern,
1318 CastPattern, ExtractAlignedPointerAsIndexOpPattern>(
1319 typeConverter, patterns.getContext());
1320}
1321} // namespace mlir
return success()
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.
Definition AffineMap.h:46
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
Definition Block.h:165
iterator end()
Definition Block.h:169
iterator begin()
Definition Block.h:168
auto getOps()
Return an iterator range over the operations within this block that are of 'OpT'.
Definition Block.h:213
IntegerAttr getIntegerAttr(Type type, int64_t value)
Definition Builders.cpp:237
IntegerType getIntegerType(unsigned width)
Definition Builders.cpp:75
This class defines the main interface for locations in MLIR and acts as a non-nullable wrapper around...
Definition Location.h:76
MLIRContext is the top-level object for a collection of MLIR operations.
Definition MLIRContext.h:63
This class helps build Operations.
Definition Builders.h:210
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...
Definition Builders.h:528
Operation is the basic unit of execution within MLIR.
Definition Operation.h:87
Region & getRegion(unsigned index)
Returns the region held by this operation at position 'index'.
Definition Operation.h:738
iterator begin()
Definition Region.h:55
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...
Definition Types.h:74
bool isInteger() const
Return true if this is an integer type (with the specified width).
Definition Types.cpp:58
bool isIntOrFloat() const
Return true if this is an integer (of any signedness) or a float type.
Definition Types.cpp:118
unsigned getIntOrFloatBitWidth() const
Return the bit width of an integer or a float type, assert failure on other types.
Definition Types.cpp:124
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
Definition Value.h:96
MLIRContext * getContext() const
Utility to get the associated MLIRContext that this value is defined in.
Definition Value.h:108
Type getType() const
Return the type of this value.
Definition Value.h:105
Operation * getDefiningOp() const
If this value is the result of an operation, return the operation that defines it.
Definition Value.cpp:18
detail::InFlightRemark failed(Location loc, RemarkOpts opts)
Report an optimization remark that failed.
Definition Remarks.h:734
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.
Definition Utils.cpp:311
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