MLIR 24.0.0git
MemoryOps.cpp
Go to the documentation of this file.
1//===- MemoryOps.cpp - MLIR SPIR-V Memory Ops ----------------------------===//
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// Defines the memory operations in the SPIR-V dialect.
10//
11//===----------------------------------------------------------------------===//
12
15
16#include "SPIRVOpUtils.h"
17#include "SPIRVParsingUtils.h"
19#include "mlir/IR/Diagnostics.h"
20
21#include "llvm/ADT/StringExtras.h"
22#include "llvm/Support/Casting.h"
23
24using namespace mlir::spirv::AttrNames;
25
26namespace mlir::spirv {
27
28/// Parses optional memory access (a.k.a. memory operand) attributes attached to
29/// a memory access operand/pointer. Specifically, parses the following syntax:
30/// (`[` memory-access `]`)?
31/// where:
32/// memory-access ::= `"None"` | `"Volatile"` | `"Aligned", `
33/// integer-literal | `"NonTemporal"`
34template <typename MemoryOpTy>
36 OperationState &state) {
37 // Parse an optional list of attributes staring with '['
38 if (parser.parseOptionalLSquare()) {
39 // Nothing to do
40 return success();
41 }
42
43 spirv::MemoryAccess memoryAccessAttr;
44 StringAttr memoryAccessAttrName =
45 MemoryOpTy::getMemoryAccessAttrName(state.name);
47 memoryAccessAttr, parser, state, memoryAccessAttrName))
48 return failure();
49
50 if (spirv::bitEnumContainsAll(memoryAccessAttr,
51 spirv::MemoryAccess::Aligned)) {
52 // Parse integer attribute for alignment.
53 Attribute alignmentAttr;
54 StringAttr alignmentAttrName = MemoryOpTy::getAlignmentAttrName(state.name);
55 Type i32Type = parser.getBuilder().getIntegerType(32);
56 if (parser.parseComma() ||
57 parser.parseAttribute(alignmentAttr, i32Type, alignmentAttrName,
58 state.attributes)) {
59 return failure();
60 }
61 }
62 return parser.parseRSquare();
63}
64
65// TODO Make sure to merge this and the previous function into one template
66// parameterized by memory access attribute name and alignment. Doing so now
67// results in VS2017 in producing an internal error (at the call site) that's
68// not detailed enough to understand what is happening.
69template <typename MemoryOpTy>
71 OperationState &state) {
72 // Parse an optional list of attributes staring with '['
73 if (parser.parseOptionalLSquare()) {
74 // Nothing to do
75 return success();
76 }
77
78 spirv::MemoryAccess memoryAccessAttr;
79 StringRef memoryAccessAttrName =
80 MemoryOpTy::getSourceMemoryAccessAttrName(state.name);
82 memoryAccessAttr, parser, state, memoryAccessAttrName))
83 return failure();
84
85 if (spirv::bitEnumContainsAll(memoryAccessAttr,
86 spirv::MemoryAccess::Aligned)) {
87 // Parse integer attribute for alignment.
88 Attribute alignmentAttr;
89 StringAttr alignmentAttrName =
90 MemoryOpTy::getSourceAlignmentAttrName(state.name);
91 Type i32Type = parser.getBuilder().getIntegerType(32);
92 if (parser.parseComma() ||
93 parser.parseAttribute(alignmentAttr, i32Type, alignmentAttrName,
94 state.attributes)) {
95 return failure();
96 }
97 }
98 return parser.parseRSquare();
99}
100
101// TODO Make sure to merge this and the previous function into one template
102// parameterized by memory access attribute name and alignment. Doing so now
103// results in VS2017 in producing an internal error (at the call site) that's
104// not detailed enough to understand what is happening.
105template <typename MemoryOpTy>
107 MemoryOpTy memoryOp, OpAsmPrinter &printer,
108 SmallVectorImpl<StringRef> &elidedAttrs,
109 std::optional<spirv::MemoryAccess> memoryAccessAtrrValue = std::nullopt,
110 std::optional<uint32_t> alignmentAttrValue = std::nullopt) {
111
112 // Print optional memory access attribute.
113 if (auto memAccess =
114 (memoryAccessAtrrValue ? memoryAccessAtrrValue
115 : memoryOp.getSourceMemoryAccess())) {
116 elidedAttrs.push_back(memoryOp.getSourceMemoryAccessAttrName());
117
118 printer << ", [\"" << stringifyMemoryAccess(*memAccess) << "\"";
119
120 if (spirv::bitEnumContainsAll(*memAccess, spirv::MemoryAccess::Aligned)) {
121 // Print integer alignment attribute.
122 if (auto alignment =
123 (alignmentAttrValue ? alignmentAttrValue
124 : memoryOp.getSourceAlignment())) {
125 elidedAttrs.push_back(memoryOp.getSourceAlignmentAttrName());
126 printer << ", " << *alignment;
127 }
128 }
129 printer << "]";
130 }
131 elidedAttrs.push_back(spirv::attributeName<spirv::StorageClass>());
132}
133
134template <typename MemoryOpTy>
136 MemoryOpTy memoryOp, OpAsmPrinter &printer,
137 SmallVectorImpl<StringRef> &elidedAttrs,
138 std::optional<spirv::MemoryAccess> memoryAccessAtrrValue = std::nullopt,
139 std::optional<uint32_t> alignmentAttrValue = std::nullopt) {
140 // Print optional memory access attribute.
141 if (auto memAccess = (memoryAccessAtrrValue ? memoryAccessAtrrValue
142 : memoryOp.getMemoryAccess())) {
143 elidedAttrs.push_back(memoryOp.getMemoryAccessAttrName());
144
145 printer << " [\"" << stringifyMemoryAccess(*memAccess) << "\"";
146
147 if (spirv::bitEnumContainsAll(*memAccess, spirv::MemoryAccess::Aligned)) {
148 // Print integer alignment attribute.
149 if (auto alignment = (alignmentAttrValue ? alignmentAttrValue
150 : memoryOp.getAlignment())) {
151 elidedAttrs.push_back(memoryOp.getAlignmentAttrName());
152 printer << ", " << *alignment;
153 }
154 }
155 printer << "]";
156 }
157 elidedAttrs.push_back(spirv::attributeName<spirv::StorageClass>());
158}
159
160template <typename LoadStoreOpTy>
161static LogicalResult verifyLoadStorePtrAndValTypes(LoadStoreOpTy op, Value ptr,
162 Value val) {
163 // ODS already checks ptr is spirv::PointerType. Just check that the pointee
164 // type of the pointer and the type of the value are the same
165 //
166 // TODO: Check that the value type satisfies restrictions of
167 // SPIR-V OpLoad/OpStore operations
168 if (val.getType() !=
169 cast<spirv::PointerType>(ptr.getType()).getPointeeType()) {
170 return op.emitOpError("mismatch in result type and pointer type");
171 }
172 return success();
173}
174
175namespace {
176/// Whether the pointer a memory operands mask applies to is read through,
177/// written through, or both.
178enum class MemoryAccessKind { Read, Write, ReadWrite };
179} // namespace
180
181/// Verifies the memory operands mask `memAccessAttr` of `op` and its companion
182/// alignment attribute `alignmentAttr`. `kind` tells how the pointer the mask
183/// applies to is accessed.
184static LogicalResult
186 spirv::MemoryAccessAttr memAccessAttr,
187 Attribute alignmentAttr, MemoryAccessKind kind) {
188 // ODS checks for attributes values. Just need to verify that if the
189 // memory-access attribute is Aligned, then the alignment attribute must be
190 // present.
191 if (!memAccessAttr) {
192 // Alignment attribute shouldn't be present if memory access attribute is
193 // not present.
194 if (alignmentAttr) {
195 return op->emitOpError(
196 "invalid alignment specification without aligned memory access "
197 "specification");
198 }
199 return success();
200 }
201
202 spirv::MemoryAccess memAccess = memAccessAttr.getValue();
203
204 // MakePointerAvailable applies to writes through the pointer.
205 if (kind == MemoryAccessKind::Read &&
206 spirv::bitEnumContainsAll(memAccess,
207 spirv::MemoryAccess::MakePointerAvailable)) {
208 return op->emitOpError(
209 "not compatible with memory operand 'MakePointerAvailable'");
210 }
211
212 // MakePointerVisible applies to reads through the pointer.
213 if (kind == MemoryAccessKind::Write &&
214 spirv::bitEnumContainsAll(memAccess,
215 spirv::MemoryAccess::MakePointerVisible)) {
216 return op->emitOpError(
217 "not compatible with memory operand 'MakePointerVisible'");
218 }
219
220 if (spirv::bitEnumContainsAny(memAccess,
221 spirv::MemoryAccess::MakePointerAvailable |
222 spirv::MemoryAccess::MakePointerVisible) &&
223 !spirv::bitEnumContainsAll(memAccess,
224 spirv::MemoryAccess::NonPrivatePointer)) {
225 return op->emitOpError(
226 "memory operand 'MakePointerAvailable' or 'MakePointerVisible' "
227 "requires 'NonPrivatePointer' to also be specified");
228 }
229
230 if (spirv::bitEnumContainsAll(memAccess, spirv::MemoryAccess::Aligned)) {
231 if (!alignmentAttr) {
232 return op->emitOpError("missing alignment value");
233 }
234 } else {
235 if (alignmentAttr) {
236 return op->emitOpError(
237 "invalid alignment specification with non-aligned memory access "
238 "specification");
239 }
240 }
241 return success();
242}
243
244/// Verifies the default (non-Source) memory operands mask of `memoryOp`.
245template <typename MemoryOpTy>
246static LogicalResult verifyMemoryAccessAttribute(MemoryOpTy memoryOp,
247 MemoryAccessKind kind) {
248 return verifyMemoryAccessAttribute(memoryOp.getOperation(),
249 memoryOp.getMemoryAccessAttr(),
250 memoryOp.getAlignmentAttr(), kind);
251}
252
253//===----------------------------------------------------------------------===//
254// spirv.AccessChainOp
255//===----------------------------------------------------------------------===//
256
258 auto ptrType = dyn_cast<spirv::PointerType>(type);
259 if (!ptrType) {
260 emitError(baseLoc, "'spirv.AccessChain' op expected a pointer "
261 "to composite type, but provided ")
262 << type;
263 return nullptr;
264 }
265
266 auto resultType = ptrType.getPointeeType();
267 auto resultStorageClass = ptrType.getStorageClass();
268 int32_t index = 0;
269
270 for (auto indexSSA : indices) {
271 auto cType = dyn_cast<spirv::CompositeType>(resultType);
272 if (!cType) {
273 emitError(
274 baseLoc,
275 "'spirv.AccessChain' op cannot extract from non-composite type ")
276 << resultType << " with index " << index;
277 return nullptr;
278 }
279 index = 0;
280 if (isa<spirv::StructType>(resultType)) {
281 Operation *op = indexSSA.getDefiningOp();
282 if (!op) {
283 emitError(baseLoc, "'spirv.AccessChain' op index must be an "
284 "integer spirv.Constant to access "
285 "element of spirv.struct");
286 return nullptr;
287 }
288
289 // TODO: this should be relaxed to allow
290 // integer literals of other bitwidths.
291 if (failed(spirv::extractValueFromConstOp(op, index))) {
292 emitError(
293 baseLoc,
294 "'spirv.AccessChain' index must be an integer spirv.Constant to "
295 "access element of spirv.struct, but provided ")
296 << op->getName();
297 return nullptr;
298 }
299 if (index < 0 || static_cast<uint64_t>(index) >= cType.getNumElements()) {
300 emitError(baseLoc, "'spirv.AccessChain' op index ")
301 << index << " out of bounds for " << resultType;
302 return nullptr;
303 }
304 }
305 resultType = cType.getElementType(index);
306 }
307 return spirv::PointerType::get(resultType, resultStorageClass);
308}
309
310void AccessChainOp::build(OpBuilder &builder, OperationState &state,
311 Value basePtr, ValueRange indices) {
312 auto type = getElementPtrType(basePtr.getType(), indices, state.location);
313 assert(type && "Unable to deduce return type based on basePtr and indices");
314 build(builder, state, type, basePtr, indices);
315}
316
317template <typename Op>
318static LogicalResult verifyAccessChain(Op accessChainOp, ValueRange indices) {
319 auto resultType = getElementPtrType(accessChainOp.getBasePtr().getType(),
320 indices, accessChainOp.getLoc());
321 if (!resultType)
322 return failure();
323
324 auto providedResultType =
325 dyn_cast<spirv::PointerType>(accessChainOp.getType());
326 if (!providedResultType)
327 return accessChainOp.emitOpError(
328 "result type must be a pointer, but provided")
329 << providedResultType;
330
331 if (resultType != providedResultType)
332 return accessChainOp.emitOpError("invalid result type: expected ")
333 << resultType << ", but provided " << providedResultType;
334
335 return success();
336}
337
338LogicalResult AccessChainOp::verify() {
339 return verifyAccessChain(*this, getIndices());
340}
341
342//===----------------------------------------------------------------------===//
343// spirv.InBoundsAccessChainOp
344//===----------------------------------------------------------------------===//
345
346void InBoundsAccessChainOp::build(OpBuilder &builder, OperationState &state,
347 Value basePtr, ValueRange indices) {
348 Type type = getElementPtrType(basePtr.getType(), indices, state.location);
349 assert(type && "Unable to deduce return type based on basePtr and indices");
350 build(builder, state, type, basePtr, indices);
351}
352
353LogicalResult InBoundsAccessChainOp::verify() {
354 return verifyAccessChain(*this, getIndices());
355}
356
357//===----------------------------------------------------------------------===//
358// spirv.LoadOp
359//===----------------------------------------------------------------------===//
360
361void LoadOp::build(OpBuilder &builder, OperationState &state, Value basePtr,
362 MemoryAccessAttr memoryAccess, IntegerAttr alignment) {
363 auto ptrType = cast<spirv::PointerType>(basePtr.getType());
364 build(builder, state, ptrType.getPointeeType(), basePtr, memoryAccess,
365 alignment);
366}
367
368ParseResult LoadOp::parse(OpAsmParser &parser, OperationState &result) {
369 // Parse the storage class specification
370 spirv::StorageClass storageClass;
371 OpAsmParser::UnresolvedOperand ptrInfo;
372 Type elementType;
373 if (parseEnumStrAttr(storageClass, parser) || parser.parseOperand(ptrInfo) ||
375 parser.parseOptionalAttrDict(result.attributes) || parser.parseColon() ||
376 parser.parseType(elementType)) {
377 return failure();
378 }
379
380 auto ptrType = spirv::PointerType::get(elementType, storageClass);
381 if (parser.resolveOperand(ptrInfo, ptrType, result.operands)) {
382 return failure();
383 }
384
385 result.addTypes(elementType);
386 return success();
387}
388
389void LoadOp::print(OpAsmPrinter &printer) {
390 SmallVector<StringRef, 4> elidedAttrs;
391 StringRef sc = stringifyStorageClass(
392 cast<spirv::PointerType>(getPtr().getType()).getStorageClass());
393 printer << " \"" << sc << "\" " << getPtr();
394
395 printMemoryAccessAttribute(*this, printer, elidedAttrs);
396
397 printer.printOptionalAttrDict(
398 (*this)->getDiscardableAttrDictionary().getValue(), elidedAttrs);
399 printer << " : " << getType();
400}
401
402LogicalResult LoadOp::verify() {
403 // SPIR-V spec : "Result Type is the type of the loaded object. It must be a
404 // type with fixed size; i.e., it cannot be, nor include, any
405 // OpTypeRuntimeArray types."
406 if (failed(verifyLoadStorePtrAndValTypes(*this, getPtr(), getValue()))) {
407 return failure();
408 }
409 return verifyMemoryAccessAttribute(*this, MemoryAccessKind::Read);
410}
411
412//===----------------------------------------------------------------------===//
413// spirv.StoreOp
414//===----------------------------------------------------------------------===//
415
416ParseResult StoreOp::parse(OpAsmParser &parser, OperationState &result) {
417 // Parse the storage class specification
418 spirv::StorageClass storageClass;
419 SmallVector<OpAsmParser::UnresolvedOperand, 2> operandInfo;
420 auto loc = parser.getCurrentLocation();
421 Type elementType;
422 if (parseEnumStrAttr(storageClass, parser) ||
423 parser.parseOperandList(operandInfo, 2) ||
425 parser.parseColon() || parser.parseType(elementType)) {
426 return failure();
427 }
428
429 auto ptrType = spirv::PointerType::get(elementType, storageClass);
430 if (parser.resolveOperands(operandInfo, {ptrType, elementType}, loc,
431 result.operands)) {
432 return failure();
433 }
434 return success();
435}
436
437void StoreOp::print(OpAsmPrinter &printer) {
438 SmallVector<StringRef, 4> elidedAttrs;
439 StringRef sc = stringifyStorageClass(
440 cast<spirv::PointerType>(getPtr().getType()).getStorageClass());
441 printer << " \"" << sc << "\" " << getPtr() << ", " << getValue();
442
443 printMemoryAccessAttribute(*this, printer, elidedAttrs);
444
445 printer << " : " << getValue().getType();
446 printer.printOptionalAttrDict(
447 (*this)->getDiscardableAttrDictionary().getValue(), elidedAttrs);
448}
449
450LogicalResult StoreOp::verify() {
451 // SPIR-V spec : "Pointer is the pointer to store through. Its type must be an
452 // OpTypePointer whose Type operand is the same as the type of Object."
453 if (failed(verifyLoadStorePtrAndValTypes(*this, getPtr(), getValue())))
454 return failure();
455 return verifyMemoryAccessAttribute(*this, MemoryAccessKind::Write);
456}
457
458//===----------------------------------------------------------------------===//
459// spirv.CopyMemory
460//===----------------------------------------------------------------------===//
461
462void CopyMemoryOp::print(OpAsmPrinter &printer) {
463 printer << ' ';
464
465 StringRef targetStorageClass = stringifyStorageClass(
466 cast<spirv::PointerType>(getTarget().getType()).getStorageClass());
467 printer << " \"" << targetStorageClass << "\" " << getTarget() << ", ";
468
469 StringRef sourceStorageClass = stringifyStorageClass(
470 cast<spirv::PointerType>(getSource().getType()).getStorageClass());
471 printer << " \"" << sourceStorageClass << "\" " << getSource();
472
473 SmallVector<StringRef, 4> elidedAttrs;
474 printMemoryAccessAttribute(*this, printer, elidedAttrs);
475 printSourceMemoryAccessAttribute(*this, printer, elidedAttrs,
476 getSourceMemoryAccess(),
477 getSourceAlignment());
478
479 printer.printOptionalAttrDict(
480 (*this)->getDiscardableAttrDictionary().getValue(), elidedAttrs);
481
482 Type pointeeType =
483 cast<spirv::PointerType>(getTarget().getType()).getPointeeType();
484 printer << " : " << pointeeType;
485}
486
487ParseResult CopyMemoryOp::parse(OpAsmParser &parser, OperationState &result) {
488 spirv::StorageClass targetStorageClass;
489 OpAsmParser::UnresolvedOperand targetPtrInfo;
490
491 spirv::StorageClass sourceStorageClass;
492 OpAsmParser::UnresolvedOperand sourcePtrInfo;
493
494 Type elementType;
495
496 if (parseEnumStrAttr(targetStorageClass, parser) ||
497 parser.parseOperand(targetPtrInfo) || parser.parseComma() ||
498 parseEnumStrAttr(sourceStorageClass, parser) ||
499 parser.parseOperand(sourcePtrInfo) ||
501 return failure();
502 }
503
504 if (!parser.parseOptionalComma()) {
505 // Parse 2nd memory access attributes.
507 return failure();
508 }
509 }
510
511 if (parser.parseColon() || parser.parseType(elementType))
512 return failure();
513
514 if (parser.parseOptionalAttrDict(result.attributes))
515 return failure();
516
517 auto targetPtrType = spirv::PointerType::get(elementType, targetStorageClass);
518 auto sourcePtrType = spirv::PointerType::get(elementType, sourceStorageClass);
519
520 if (parser.resolveOperand(targetPtrInfo, targetPtrType, result.operands) ||
521 parser.resolveOperand(sourcePtrInfo, sourcePtrType, result.operands)) {
522 return failure();
523 }
524
525 return success();
526}
527
528LogicalResult CopyMemoryOp::verify() {
529 Type targetType =
530 cast<spirv::PointerType>(getTarget().getType()).getPointeeType();
531
532 Type sourceType =
533 cast<spirv::PointerType>(getSource().getType()).getPointeeType();
534
535 if (targetType != sourceType)
536 return emitOpError("both operands must be pointers to the same type");
537
538 // A lone mask applies to both operands. Only the first of two masks is
539 // Target-only.
540 MemoryAccessKind targetKind = getSourceMemoryAccess()
541 ? MemoryAccessKind::Write
542 : MemoryAccessKind::ReadWrite;
543 if (failed(verifyMemoryAccessAttribute(*this, targetKind)))
544 return failure();
545
547 getOperation(), getSourceMemoryAccessAttr(), getSourceAlignmentAttr(),
548 MemoryAccessKind::Read);
549}
550
551//===----------------------------------------------------------------------===//
552// spirv.InBoundsPtrAccessChainOp
553//===----------------------------------------------------------------------===//
554
555void InBoundsPtrAccessChainOp::build(OpBuilder &builder, OperationState &state,
556 Value basePtr, Value element,
558 auto type = getElementPtrType(basePtr.getType(), indices, state.location);
559 assert(type && "Unable to deduce return type based on basePtr and indices");
560 build(builder, state, type, basePtr, element, indices);
561}
562
563LogicalResult InBoundsPtrAccessChainOp::verify() {
564 return verifyAccessChain(*this, getIndices());
565}
566
567//===----------------------------------------------------------------------===//
568// spirv.PtrAccessChainOp
569//===----------------------------------------------------------------------===//
570
571void PtrAccessChainOp::build(OpBuilder &builder, OperationState &state,
572 Value basePtr, Value element, ValueRange indices) {
573 auto type = getElementPtrType(basePtr.getType(), indices, state.location);
574 assert(type && "Unable to deduce return type based on basePtr and indices");
575 build(builder, state, type, basePtr, element, indices);
576}
577
578LogicalResult PtrAccessChainOp::verify() {
579 return verifyAccessChain(*this, getIndices());
580}
581
582//===----------------------------------------------------------------------===//
583// spirv.Variable
584//===----------------------------------------------------------------------===//
585
586ParseResult VariableOp::parse(OpAsmParser &parser, OperationState &result) {
587 // Parse optional initializer
588 std::optional<OpAsmParser::UnresolvedOperand> initInfo;
589 if (succeeded(parser.parseOptionalKeyword("init"))) {
590 initInfo = OpAsmParser::UnresolvedOperand();
591 if (parser.parseLParen() || parser.parseOperand(*initInfo) ||
592 parser.parseRParen())
593 return failure();
594 }
595
596 if (parseVariableDecorations(parser, result)) {
597 return failure();
598 }
599
600 // Parse result pointer type
601 Type type;
602 if (parser.parseColon())
603 return failure();
604 auto loc = parser.getCurrentLocation();
605 if (parser.parseType(type))
606 return failure();
607
608 auto ptrType = dyn_cast<spirv::PointerType>(type);
609 if (!ptrType)
610 return parser.emitError(loc, "expected spirv.ptr type");
611 result.addTypes(ptrType);
612
613 // Resolve the initializer operand
614 if (initInfo) {
615 if (parser.resolveOperand(*initInfo, ptrType.getPointeeType(),
616 result.operands))
617 return failure();
618 }
619
620 auto attr = parser.getBuilder().getAttr<spirv::StorageClassAttr>(
621 ptrType.getStorageClass());
623
624 return success();
625}
626
627void VariableOp::print(OpAsmPrinter &printer) {
628 SmallVector<StringRef, 4> elidedAttrs{
630 // Print optional initializer
631 if (getNumOperands() != 0)
632 printer << " init(" << getInitializer() << ")";
633
634 printVariableDecorations(*this, printer, elidedAttrs);
635 printer << " : " << getType();
636}
637
638LogicalResult VariableOp::verify() {
639 // SPIR-V spec: "Storage Class is the Storage Class of the memory holding the
640 // object. It cannot be Generic. It must be the same as the Storage Class
641 // operand of the Result Type."
642 if (getStorageClass() != spirv::StorageClass::Function) {
643 return emitOpError(
644 "can only be used to model function-level variables. Use "
645 "spirv.GlobalVariable for module-level variables.");
646 }
647
648 auto pointerType = cast<spirv::PointerType>(getPointer().getType());
649 if (getStorageClass() != pointerType.getStorageClass())
650 return emitOpError(
651 "storage class must match result pointer's storage class");
652
653 if (getNumOperands() != 0) {
654 // SPIR-V spec: "Initializer must be an <id> from a constant instruction or
655 // a global (module scope) OpVariable instruction".
656 auto *initOp = getOperand(0).getDefiningOp();
657 if (!initOp || !isa<spirv::ConstantOp, // for normal constant
658 spirv::ReferenceOfOp, // for spec constant
659 spirv::AddressOfOp>(initOp))
660 return emitOpError("initializer must be the result of a "
661 "constant or spirv.GlobalVariable op");
662 }
663
664 auto getDecorationAttr = [op = getOperation()](spirv::Decoration decoration) {
665 return op->getDiscardableAttr(spirv::getDecorationString(decoration));
666 };
667
668 // TODO: generate these strings using ODS.
669 for (auto decoration :
670 {spirv::Decoration::DescriptorSet, spirv::Decoration::Binding,
671 spirv::Decoration::BuiltIn}) {
672 if (auto attr = getDecorationAttr(decoration))
673 return emitOpError("cannot have '")
674 << spirv::getDecorationString(decoration)
675 << "' attribute (only allowed in spirv.GlobalVariable)";
676 }
677
679 getPointeeType())))
680 return failure();
681
682 return success();
683}
684
685} // namespace mlir::spirv
static Value getPointer(Location loc, Value value, ConversionPatternRewriter &rewriter)
return success()
getNumOperands() - 1))) return failure()
Given a list of lists of parsed operands, populates uniqueOperands with unique operands.
virtual Builder & getBuilder() const =0
Return a builder which provides useful access to MLIRContext, global objects like types and attribute...
virtual ParseResult parseOptionalAttrDict(NamedAttrList &result)=0
Parse a named dictionary into 'result' if it is present.
virtual ParseResult parseOptionalKeyword(StringRef keyword)=0
Parse the given keyword if present.
virtual ParseResult parseRParen()=0
Parse a ) token.
virtual InFlightDiagnostic emitError(SMLoc loc, const Twine &message={})=0
Emit a diagnostic at the specified location and return failure.
virtual ParseResult parseRSquare()=0
Parse a ] token.
virtual SMLoc getCurrentLocation()=0
Get the location of the next token and store it into the argument.
virtual ParseResult parseOptionalComma()=0
Parse a , token if present.
virtual ParseResult parseColon()=0
Parse a : token.
virtual ParseResult parseLParen()=0
Parse a ( token.
virtual ParseResult parseType(Type &result)=0
Parse a type.
virtual ParseResult parseComma()=0
Parse a , token.
virtual ParseResult parseOptionalLSquare()=0
Parse a [ token if present.
virtual ParseResult parseAttribute(Attribute &result, Type type={})=0
Parse an arbitrary attribute of a given type and return it in result.
Attributes are known-constant values of operations.
Definition Attributes.h:25
IntegerType getIntegerType(unsigned width)
Definition Builders.cpp:75
Attr getAttr(Args &&...args)
Get or construct an instance of the attribute Attr with provided arguments.
Definition Builders.h:101
This class defines the main interface for locations in MLIR and acts as a non-nullable wrapper around...
Definition Location.h:76
The OpAsmParser has methods for interacting with the asm parser: parsing things from it,...
virtual ParseResult resolveOperand(const UnresolvedOperand &operand, Type type, SmallVectorImpl< Value > &result)=0
Resolve an operand to an SSA value, emitting an error on failure.
ParseResult resolveOperands(Operands &&operands, Type type, SmallVectorImpl< Value > &result)
Resolve a list of operands to SSA values, emitting an error on failure, or appending the results to t...
virtual ParseResult parseOperand(UnresolvedOperand &result, bool allowResultNumber=true)=0
Parse a single SSA value operand name along with a result number if allowResultNumber is true.
virtual ParseResult parseOperandList(SmallVectorImpl< UnresolvedOperand > &result, Delimiter delimiter=Delimiter::None, bool allowResultNumber=true, int requiredOperandCount=-1)=0
Parse zero or more SSA comma-separated operand references with a specified surrounding delimiter,...
This is a pure-virtual base class that exposes the asmprinter hooks necessary to implement a custom p...
virtual void printOptionalAttrDict(ArrayRef< NamedAttribute > attrs, ArrayRef< StringRef > elidedAttrs={})=0
If the specified operation has attributes, print out an attribute dictionary with their values.
This class helps build Operations.
Definition Builders.h:210
InFlightDiagnostic emitOpError(const Twine &message={})
Emit an error with the op name prefixed, like "'dim' op " which is convenient for verifiers.
Location getLoc()
The source location the operation was defined or derived from.
This provides public APIs that all operations should have.
Operation is the basic unit of execution within MLIR.
Definition Operation.h:87
OperationName getName()
The name of an operation is the key identifier for it.
Definition Operation.h:115
InFlightDiagnostic emitOpError(const Twine &message={})
Emit an error with the op name prefixed, like "'dim' op " which is convenient for verifiers.
Instances of the Type class are uniqued, have an immutable identifier and an optional mutable compone...
Definition Types.h:74
This class provides an abstraction over the different types of ranges over Values.
Definition ValueRange.h:389
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
Definition Value.h:96
Type getType() const
Return the type of this value.
Definition Value.h:105
static PointerType get(Type pointeeType, StorageClass storageClass)
Operation::operand_range getIndices(Operation *op)
Get the indices that the given load/store operation is operating on.
Definition Utils.cpp:18
detail::InFlightRemark failed(Location loc, RemarkOpts opts)
Report an optimization remark that failed.
Definition Remarks.h:734
static ParseResult parseSourceMemoryAccessAttributes(OpAsmParser &parser, OperationState &state)
Definition MemoryOps.cpp:70
ParseResult parseEnumStrAttr(EnumClass &value, OpAsmParser &parser, StringRef attrName=spirv::attributeName< EnumClass >())
Parses the next string attribute in parser as an enumerant of the given EnumClass.
static LogicalResult verifyMemoryAccessAttribute(Operation *op, spirv::MemoryAccessAttr memAccessAttr, Attribute alignmentAttr, MemoryAccessKind kind)
Verifies the memory operands mask memAccessAttr of op and its companion alignment attribute alignment...
static void printSourceMemoryAccessAttribute(MemoryOpTy memoryOp, OpAsmPrinter &printer, SmallVectorImpl< StringRef > &elidedAttrs, std::optional< spirv::MemoryAccess > memoryAccessAtrrValue=std::nullopt, std::optional< uint32_t > alignmentAttrValue=std::nullopt)
ParseResult parseMemoryAccessAttributes(OpAsmParser &parser, OperationState &state)
Parses optional memory access (a.k.a.
Definition MemoryOps.cpp:35
static Type getElementPtrType(Type type, ValueRange indices, Location baseLoc)
void printVariableDecorations(Operation *op, OpAsmPrinter &printer, SmallVectorImpl< StringRef > &elidedAttrs)
Definition SPIRVOps.cpp:133
LogicalResult verifyPhysicalStorageBufferDecorations(Operation *op, Type pointeeType)
Verifies the SPV_KHR_physical_storage_buffer rule that a variable whose pointee is a pointer (or arra...
Definition SPIRVOps.cpp:93
static LogicalResult verifyLoadStorePtrAndValTypes(LoadStoreOpTy op, Value ptr, Value val)
constexpr StringRef attributeName()
static void printMemoryAccessAttribute(MemoryOpTy memoryOp, OpAsmPrinter &printer, SmallVectorImpl< StringRef > &elidedAttrs, std::optional< spirv::MemoryAccess > memoryAccessAtrrValue=std::nullopt, std::optional< uint32_t > alignmentAttrValue=std::nullopt)
LogicalResult extractValueFromConstOp(Operation *op, int32_t &value)
Definition SPIRVOps.cpp:49
std::string getDecorationString(Decoration decoration)
Converts a SPIR-V Decoration enum value to its snake_case string representation for use in MLIR attri...
ParseResult parseVariableDecorations(OpAsmParser &parser, OperationState &state)
static LogicalResult verifyAccessChain(Op accessChainOp, ValueRange indices)
Type getType(OpFoldResult ofr)
Returns the int type of the integer in ofr.
Definition Utils.cpp:311
InFlightDiagnostic emitError(Location loc)
Utility method to emit an error message using this location.
This represents an operation in an abstracted form, suitable for use with the builder APIs.