21#include "llvm/ADT/StringExtras.h"
22#include "llvm/Support/Casting.h"
34template <
typename MemoryOpTy>
43 spirv::MemoryAccess memoryAccessAttr;
44 StringAttr memoryAccessAttrName =
45 MemoryOpTy::getMemoryAccessAttrName(state.
name);
47 memoryAccessAttr, parser, state, memoryAccessAttrName))
50 if (spirv::bitEnumContainsAll(memoryAccessAttr,
51 spirv::MemoryAccess::Aligned)) {
54 StringAttr alignmentAttrName = MemoryOpTy::getAlignmentAttrName(state.
name);
69template <
typename MemoryOpTy>
78 spirv::MemoryAccess memoryAccessAttr;
79 StringRef memoryAccessAttrName =
80 MemoryOpTy::getSourceMemoryAccessAttrName(state.
name);
82 memoryAccessAttr, parser, state, memoryAccessAttrName))
85 if (spirv::bitEnumContainsAll(memoryAccessAttr,
86 spirv::MemoryAccess::Aligned)) {
89 StringAttr alignmentAttrName =
90 MemoryOpTy::getSourceAlignmentAttrName(state.
name);
105template <
typename MemoryOpTy>
109 std::optional<spirv::MemoryAccess> memoryAccessAtrrValue = std::nullopt,
110 std::optional<uint32_t> alignmentAttrValue = std::nullopt) {
114 (memoryAccessAtrrValue ? memoryAccessAtrrValue
115 : memoryOp.getSourceMemoryAccess())) {
116 elidedAttrs.push_back(memoryOp.getSourceMemoryAccessAttrName());
118 printer <<
", [\"" << stringifyMemoryAccess(*memAccess) <<
"\"";
120 if (spirv::bitEnumContainsAll(*memAccess, spirv::MemoryAccess::Aligned)) {
123 (alignmentAttrValue ? alignmentAttrValue
124 : memoryOp.getSourceAlignment())) {
125 elidedAttrs.push_back(memoryOp.getSourceAlignmentAttrName());
126 printer <<
", " << *alignment;
134template <
typename MemoryOpTy>
138 std::optional<spirv::MemoryAccess> memoryAccessAtrrValue = std::nullopt,
139 std::optional<uint32_t> alignmentAttrValue = std::nullopt) {
141 if (
auto memAccess = (memoryAccessAtrrValue ? memoryAccessAtrrValue
142 : memoryOp.getMemoryAccess())) {
143 elidedAttrs.push_back(memoryOp.getMemoryAccessAttrName());
145 printer <<
" [\"" << stringifyMemoryAccess(*memAccess) <<
"\"";
147 if (spirv::bitEnumContainsAll(*memAccess, spirv::MemoryAccess::Aligned)) {
149 if (
auto alignment = (alignmentAttrValue ? alignmentAttrValue
150 : memoryOp.getAlignment())) {
151 elidedAttrs.push_back(memoryOp.getAlignmentAttrName());
152 printer <<
", " << *alignment;
160template <
typename LoadStoreOpTy>
169 cast<spirv::PointerType>(
ptr.getType()).getPointeeType()) {
170 return op.emitOpError(
"mismatch in result type and pointer type");
178enum class MemoryAccessKind { Read, Write, ReadWrite };
186 spirv::MemoryAccessAttr memAccessAttr,
187 Attribute alignmentAttr, MemoryAccessKind kind) {
191 if (!memAccessAttr) {
196 "invalid alignment specification without aligned memory access "
202 spirv::MemoryAccess memAccess = memAccessAttr.getValue();
205 if (kind == MemoryAccessKind::Read &&
206 spirv::bitEnumContainsAll(memAccess,
207 spirv::MemoryAccess::MakePointerAvailable)) {
209 "not compatible with memory operand 'MakePointerAvailable'");
213 if (kind == MemoryAccessKind::Write &&
214 spirv::bitEnumContainsAll(memAccess,
215 spirv::MemoryAccess::MakePointerVisible)) {
217 "not compatible with memory operand 'MakePointerVisible'");
220 if (spirv::bitEnumContainsAny(memAccess,
221 spirv::MemoryAccess::MakePointerAvailable |
222 spirv::MemoryAccess::MakePointerVisible) &&
223 !spirv::bitEnumContainsAll(memAccess,
224 spirv::MemoryAccess::NonPrivatePointer)) {
226 "memory operand 'MakePointerAvailable' or 'MakePointerVisible' "
227 "requires 'NonPrivatePointer' to also be specified");
230 if (spirv::bitEnumContainsAll(memAccess, spirv::MemoryAccess::Aligned)) {
231 if (!alignmentAttr) {
237 "invalid alignment specification with non-aligned memory access "
245template <
typename MemoryOpTy>
247 MemoryAccessKind kind) {
249 memoryOp.getMemoryAccessAttr(),
250 memoryOp.getAlignmentAttr(), kind);
258 auto ptrType = dyn_cast<spirv::PointerType>(type);
260 emitError(baseLoc,
"'spirv.AccessChain' op expected a pointer "
261 "to composite type, but provided ")
266 auto resultType = ptrType.getPointeeType();
267 auto resultStorageClass = ptrType.getStorageClass();
270 for (
auto indexSSA :
indices) {
271 auto cType = dyn_cast<spirv::CompositeType>(resultType);
275 "'spirv.AccessChain' op cannot extract from non-composite type ")
276 << resultType <<
" with index " <<
index;
280 if (isa<spirv::StructType>(resultType)) {
281 Operation *op = indexSSA.getDefiningOp();
283 emitError(baseLoc,
"'spirv.AccessChain' op index must be an "
284 "integer spirv.Constant to access "
285 "element of spirv.struct");
294 "'spirv.AccessChain' index must be an integer spirv.Constant to "
295 "access element of spirv.struct, but provided ")
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;
305 resultType = cType.getElementType(
index);
313 assert(type &&
"Unable to deduce return type based on basePtr and indices");
314 build(builder, state, type, basePtr,
indices);
317template <
typename Op>
324 auto providedResultType =
325 dyn_cast<spirv::PointerType>(accessChainOp.getType());
326 if (!providedResultType)
328 "result type must be a pointer, but provided")
329 << providedResultType;
331 if (resultType != providedResultType)
332 return accessChainOp.
emitOpError(
"invalid result type: expected ")
333 << resultType <<
", but provided " << providedResultType;
338LogicalResult AccessChainOp::verify() {
349 assert(type &&
"Unable to deduce return type based on basePtr and indices");
350 build(builder, state, type, basePtr,
indices);
353LogicalResult InBoundsAccessChainOp::verify() {
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,
368ParseResult LoadOp::parse(OpAsmParser &parser, OperationState &
result) {
370 spirv::StorageClass storageClass;
371 OpAsmParser::UnresolvedOperand ptrInfo;
385 result.addTypes(elementType);
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();
398 (*this)->getDiscardableAttrDictionary().getValue(), elidedAttrs);
402LogicalResult LoadOp::verify() {
416ParseResult StoreOp::parse(OpAsmParser &parser, OperationState &
result) {
418 spirv::StorageClass storageClass;
419 SmallVector<OpAsmParser::UnresolvedOperand, 2> operandInfo;
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();
445 printer <<
" : " << getValue().getType();
447 (*this)->getDiscardableAttrDictionary().getValue(), elidedAttrs);
450LogicalResult StoreOp::verify() {
462void CopyMemoryOp::print(OpAsmPrinter &printer) {
465 StringRef targetStorageClass = stringifyStorageClass(
466 cast<spirv::PointerType>(getTarget().
getType()).getStorageClass());
467 printer <<
" \"" << targetStorageClass <<
"\" " << getTarget() <<
", ";
469 StringRef sourceStorageClass = stringifyStorageClass(
470 cast<spirv::PointerType>(getSource().
getType()).getStorageClass());
471 printer <<
" \"" << sourceStorageClass <<
"\" " << getSource();
473 SmallVector<StringRef, 4> elidedAttrs;
476 getSourceMemoryAccess(),
477 getSourceAlignment());
480 (*this)->getDiscardableAttrDictionary().getValue(), elidedAttrs);
483 cast<spirv::PointerType>(getTarget().
getType()).getPointeeType();
484 printer <<
" : " << pointeeType;
487ParseResult CopyMemoryOp::parse(OpAsmParser &parser, OperationState &
result) {
488 spirv::StorageClass targetStorageClass;
489 OpAsmParser::UnresolvedOperand targetPtrInfo;
491 spirv::StorageClass sourceStorageClass;
492 OpAsmParser::UnresolvedOperand sourcePtrInfo;
528LogicalResult CopyMemoryOp::verify() {
530 cast<spirv::PointerType>(getTarget().
getType()).getPointeeType();
533 cast<spirv::PointerType>(getSource().
getType()).getPointeeType();
535 if (targetType != sourceType)
536 return emitOpError(
"both operands must be pointers to the same type");
540 MemoryAccessKind targetKind = getSourceMemoryAccess()
541 ? MemoryAccessKind::Write
542 : MemoryAccessKind::ReadWrite;
547 getOperation(), getSourceMemoryAccessAttr(), getSourceAlignmentAttr(),
548 MemoryAccessKind::Read);
555void InBoundsPtrAccessChainOp::build(OpBuilder &builder, OperationState &state,
556 Value basePtr, Value element,
559 assert(type &&
"Unable to deduce return type based on basePtr and indices");
560 build(builder, state, type, basePtr, element,
indices);
563LogicalResult InBoundsPtrAccessChainOp::verify() {
571void PtrAccessChainOp::build(OpBuilder &builder, OperationState &state,
574 assert(type &&
"Unable to deduce return type based on basePtr and indices");
575 build(builder, state, type, basePtr, element,
indices);
578LogicalResult PtrAccessChainOp::verify() {
586ParseResult VariableOp::parse(OpAsmParser &parser, OperationState &
result) {
588 std::optional<OpAsmParser::UnresolvedOperand> initInfo;
590 initInfo = OpAsmParser::UnresolvedOperand();
608 auto ptrType = dyn_cast<spirv::PointerType>(type);
610 return parser.
emitError(loc,
"expected spirv.ptr type");
621 ptrType.getStorageClass());
627void VariableOp::print(OpAsmPrinter &printer) {
628 SmallVector<StringRef, 4> elidedAttrs{
632 printer <<
" init(" << getInitializer() <<
")";
638LogicalResult VariableOp::verify() {
642 if (getStorageClass() != spirv::StorageClass::Function) {
644 "can only be used to model function-level variables. Use "
645 "spirv.GlobalVariable for module-level variables.");
649 if (getStorageClass() != pointerType.getStorageClass())
651 "storage class must match result pointer's storage class");
656 auto *initOp = getOperand(0).getDefiningOp();
657 if (!initOp || !isa<spirv::ConstantOp,
658 spirv::ReferenceOfOp,
659 spirv::AddressOfOp>(initOp))
660 return emitOpError(
"initializer must be the result of a "
661 "constant or spirv.GlobalVariable op");
664 auto getDecorationAttr = [op = getOperation()](spirv::Decoration decoration) {
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 '")
675 <<
"' attribute (only allowed in spirv.GlobalVariable)";
static Value getPointer(Location loc, Value value, ConversionPatternRewriter &rewriter)
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.
IntegerType getIntegerType(unsigned width)
Attr getAttr(Args &&...args)
Get or construct an instance of the attribute Attr with provided arguments.
This class defines the main interface for locations in MLIR and acts as a non-nullable wrapper around...
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.
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.
OperationName getName()
The name of an operation is the key identifier for it.
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...
This class provides an abstraction over the different types of ranges over Values.
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
Type getType() const
Return the type of this value.
static PointerType get(Type pointeeType, StorageClass storageClass)
Operation::operand_range getIndices(Operation *op)
Get the indices that the given load/store operation is operating on.
static ParseResult parseSourceMemoryAccessAttributes(OpAsmParser &parser, OperationState &state)
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.
static Type getElementPtrType(Type type, ValueRange indices, Location baseLoc)
void printVariableDecorations(Operation *op, OpAsmPrinter &printer, SmallVectorImpl< StringRef > &elidedAttrs)
LogicalResult verifyPhysicalStorageBufferDecorations(Operation *op, Type pointeeType)
Verifies the SPV_KHR_physical_storage_buffer rule that a variable whose pointee is a pointer (or arra...
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)
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.
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.