28#include "llvm/ADT/Sequence.h"
29#include "llvm/ADT/StringExtras.h"
30#include "llvm/ADT/TypeSwitch.h"
35#include "mlir/Dialect/SPIRV/IR/SPIRVOpsDialect.cpp.inc"
44 return llvm::any_of(region, [](
Block &block) {
46 return isa<spirv::ReturnOp, spirv::ReturnValueOp>(terminator);
52struct SPIRVInlinerInterface :
public DialectInlinerInterface {
53 using DialectInlinerInterface::DialectInlinerInterface;
57 bool wouldBeCloned)
const final {
64 IRMapping &)
const final {
67 auto *op = dest->getParentOp();
68 return isa<spirv::FuncOp, spirv::SelectionOp, spirv::LoopOp>(op);
75 IRMapping &)
const final {
77 if ((isa<spirv::SelectionOp, spirv::LoopOp>(op)) &&
85 if (isa<spirv::KillOp>(op))
93 void handleTerminator(Operation *op,
Block *newDest)
const final {
94 if (
auto returnOp = dyn_cast<spirv::ReturnOp>(op)) {
95 auto builder = OpBuilder(op);
96 spirv::BranchOp::create(builder, op->getLoc(), newDest);
98 }
else if (
auto retValOp = dyn_cast<spirv::ReturnValueOp>(op)) {
99 auto builder = OpBuilder(op);
100 spirv::BranchOp::create(builder, retValOp->getLoc(), newDest,
101 retValOp->getOperands());
108 void handleTerminator(Operation *op,
ValueRange valuesToRepl)
const final {
110 auto retValOp = dyn_cast<spirv::ReturnValueOp>(op);
115 assert(valuesToRepl.size() == 1 &&
116 "spirv.ReturnValue expected to only handle one result");
117 valuesToRepl.front().replaceAllUsesWith(retValOp.getValue());
126void SPIRVDialect::initialize() {
127 registerAttributes();
130 registerSPIRVDialectOperations(
this);
132 addInterfaces<SPIRVInlinerInterface>();
135 allowUnknownOperations();
136 declarePromisedInterface<gpu::TargetAttrInterface, TargetEnvAttr>();
139std::string SPIRVDialect::getAttributeName(Decoration decoration) {
148template <
typename ValTy>
149static std::optional<ValTy>
parseAndVerify(SPIRVDialect
const &dialect,
171 if (
auto t = dyn_cast<FloatType>(type)) {
174 "only 8/16/32/64-bit float type allowed but found ")
178 }
else if (
auto t = dyn_cast<IntegerType>(type)) {
181 "only 1/8/16/32/64-bit integer type allowed but found ")
185 }
else if (
auto t = dyn_cast<VectorType>(type)) {
186 if (t.getRank() != 1) {
187 parser.
emitError(typeLoc,
"only 1-D vector allowed but found ") << t;
190 if (t.getNumElements() < 2) {
191 parser.
emitError(typeLoc,
"SPIR-V does not allow one-element vectors");
194 if (t.getNumElements() > 4) {
196 typeLoc,
"vector length has to be less than or equal to 4 but found ")
197 << t.getNumElements();
200 if (!isa<ScalarType>(t.getElementType())) {
203 "vector element type must be a SPIR-V scalar type but found ")
204 << t.getElementType();
207 }
else if (
auto t = dyn_cast<TensorArmType>(type)) {
208 if (!isa<ScalarType>(t.getElementType())) {
210 typeLoc,
"only scalar element type allowed in tensor type but found ")
211 << t.getElementType();
216 << type <<
" to compose SPIR-V types";
230 if (
auto t = dyn_cast<VectorType>(type)) {
231 if (t.getRank() != 1) {
232 parser.
emitError(typeLoc,
"only 1-D vector allowed but found ") << t;
235 if (t.getNumElements() > 4 || t.getNumElements() < 2) {
237 "matrix columns size has to be less than or equal "
238 "to 4 and greater than or equal 2, but found ")
239 << t.getNumElements();
243 if (!isa<FloatType>(t.getElementType())) {
244 parser.
emitError(typeLoc,
"matrix columns' elements must be of "
246 << t.getElementType();
250 parser.
emitError(typeLoc,
"matrix must be composed using vector "
266 auto imageType = dyn_cast<ImageType>(type);
269 "sampled image must be composed using image type, got ")
274 if (llvm::is_contained({Dim::SubpassData, Dim::Buffer}, imageType.getDim())) {
276 typeLoc,
"sampled image Dim must not be SubpassData or Buffer, got ")
277 << stringifyDim(imageType.getDim());
303 if (!(stride = *optStride)) {
304 parser.
emitError(strideLoc,
"ArrayStride must be greater than zero");
326 if (countDims.size() != 1) {
328 "expected single integer for array element count");
336 parser.
emitError(countLoc,
"expected array length greater than 0");
366 if (dims.size() != 2) {
367 parser.
emitError(countLoc,
"expected row and column count");
380 CooperativeMatrixUseKHR use;
398 bool unranked =
false;
410 if (!unranked && dims.empty()) {
411 parser.
emitError(countLoc,
"arm.tensors do not support rank zero");
415 if (llvm::is_contained(dims, 0)) {
416 parser.
emitError(countLoc,
"arm.tensors do not support zero dimensions");
420 if (llvm::any_of(dims, [](
int64_t dim) {
return dim < 0; }) &&
421 llvm::any_of(dims, [](
int64_t dim) {
return dim > 0; })) {
422 parser.
emitError(countLoc,
"arm.tensor shape dimensions must be either "
423 "fully dynamic or completed shaped");
455 StringRef storageClassSpec;
460 auto storageClass = symbolizeStorageClass(storageClassSpec);
462 parser.
emitError(storageClassLoc,
"unknown storage class: ")
501 if (countDims.size() != 1) {
502 parser.
emitError(countLoc,
"expected single unsigned "
503 "integer for number of columns");
507 int64_t columnCount = countDims[0];
509 if (columnCount < 2 || columnCount > 4) {
510 parser.
emitError(countLoc,
"matrix is expected to have 2, 3, or 4 "
527template <
typename ValTy>
536 auto val = spirv::symbolizeEnum<ValTy>(enumSpec);
538 parser.
emitError(enumLoc,
"unknown attribute: '") << enumSpec <<
"'";
552template <
typename IntTy>
555 IntTy offsetVal = std::numeric_limits<IntTy>::max();
572template <
typename ParseType,
typename... Args>
573struct ParseCommaSeparatedList {
574 std::optional<std::tuple<ParseType, Args...>>
575 operator()(SPIRVDialect
const &dialect, DialectAsmParser &parser)
const {
580 auto numArgs = std::tuple_size<std::tuple<Args...>>::value;
583 auto remainingValues = ParseCommaSeparatedList<Args...>{}(dialect, parser);
584 if (!remainingValues)
586 return std::tuple_cat(std::tuple<ParseType>(parseVal.value()),
587 remainingValues.value());
593template <
typename ParseType>
594struct ParseCommaSeparatedList<ParseType> {
595 std::optional<std::tuple<ParseType>>
596 operator()(SPIRVDialect
const &dialect, DialectAsmParser &parser)
const {
598 return std::tuple<ParseType>(*value);
625 ParseCommaSeparatedList<
Type, Dim, ImageDepthInfo, ImageArrayedInfo,
626 ImageSamplingInfo, ImageSamplerUseInfo,
627 ImageFormat>{}(dialect, parser);
663 if (failed(*offsetParseResult))
666 if (offsetInfo.size() != memberTypes.size() - 1) {
668 "offset specification must be given for "
671 offsetInfo.push_back(offset);
683 auto parseDecorations = [&]() {
685 if (!memberDecoration)
694 memberDecorationInfo.emplace_back(
695 static_cast<uint32_t
>(memberTypes.size() - 1),
696 memberDecoration.value(), memberDecorationValue);
698 memberDecorationInfo.emplace_back(
699 static_cast<uint32_t
>(memberTypes.size() - 1),
700 memberDecoration.value(), UnitAttr::get(dialect.getContext()));
726 StringRef identifier;
727 FailureOr<DialectAsmParser::CyclicParseReset> cyclicParse;
736 if (succeeded(cyclicParse)) {
739 "recursive struct reference not nested in struct definition");
750 if (failed(cyclicParse)) {
752 "identifier already used for an enclosing struct");
767 if (!identifier.empty())
778 if (!isa<SPIRVType>(memberType)) {
780 "member type must be a valid SPIR-V type");
783 memberTypes.push_back(memberType);
787 memberDecorationInfo))
791 if (!offsetInfo.empty() && memberTypes.size() != offsetInfo.size()) {
793 "offset specification must be given for all members");
802 auto parseStructDecoration = [&]() {
803 std::optional<spirv::Decoration> decoration =
814 structDecorationInfo.emplace_back(decoration.value(), decorationValue);
816 structDecorationInfo.emplace_back(decoration.value(),
817 UnitAttr::get(dialect.getContext()));
823 if (failed(parseStructDecoration()))
829 if (!identifier.empty()) {
830 if (failed(idStructTy.
trySetBody(memberTypes, offsetInfo,
831 memberDecorationInfo,
832 structDecorationInfo)))
838 structDecorationInfo);
853 if (keyword ==
"array")
855 if (keyword ==
"coopmatrix")
857 if (keyword ==
"image")
859 if (keyword ==
"ptr")
861 if (keyword ==
"rtarray")
863 if (keyword ==
"sampled_image")
865 if (keyword ==
"sampler")
867 if (keyword ==
"named_barrier")
869 if (keyword ==
"struct")
871 if (keyword ==
"matrix")
873 if (keyword ==
"arm.tensor")
886 os <<
", stride=" << stride;
893 os <<
", stride=" << stride;
904 <<
", " << stringifyImageDepthInfo(type.
getDepthInfo()) <<
", "
918 os <<
"named_barrier";
922 FailureOr<AsmPrinter::CyclicPrintReset> cyclicPrint;
930 if (failed(cyclicPrint)) {
940 auto printMember = [&](
unsigned i) {
944 if (type.
hasOffset() || !decorations.empty()) {
948 if (!decorations.empty())
952 os << stringifyDecoration(decoration.decoration);
953 if (decoration.hasValue()) {
958 llvm::interleaveComma(decorations, os, eachFn);
962 llvm::interleaveComma(llvm::seq<unsigned>(0, type.
getNumElements()), os,
968 if (!decorations.empty()) {
971 os << stringifyDecoration(decoration.decoration);
972 if (decoration.hasValue()) {
977 llvm::interleaveComma(decorations, os, eachFn);
1000 if (ShapedType::isDynamic(dim))
1017 [&](
auto type) {
print(type, os); })
1018 .DefaultUnreachable(
"Unhandled SPIR-V type");
1028 if (
auto poison = dyn_cast<ub::PoisonAttr>(value))
1029 return ub::PoisonOp::create(builder, loc, type, poison);
1031 if (!spirv::ConstantOp::isBuildableWith(type))
1034 return spirv::ConstantOp::create(builder, loc, type, value);
1041LogicalResult SPIRVDialect::verifyOperationAttribute(
Operation *op,
1043 StringRef symbol = attribute.
getName().strref();
1047 if (!isa<spirv::EntryPointABIAttr>(attr)) {
1049 << symbol <<
"' attribute must be an entry point ABI attribute";
1052 if (!isa<spirv::TargetEnvAttr>(attr))
1053 return op->
emitError(
"'") << symbol <<
"' must be a spirv::TargetEnvAttr";
1055 if (!isa<spirv::LoopControlAttr>(attr))
1057 << symbol <<
"' must be a spirv::LoopControlAttr";
1059 if (!isa<spirv::SelectionControlAttr>(attr))
1061 << symbol <<
"' must be a spirv::SelectionControlAttr";
1063 return op->
emitError(
"found unsupported '")
1064 << symbol <<
"' attribute on operation";
1074 StringRef symbol = attribute.
getName().strref();
1078 auto varABIAttr = dyn_cast<spirv::InterfaceVarABIAttr>(attr);
1081 << symbol <<
"' must be a spirv::InterfaceVarABIAttr";
1085 <<
"' attribute cannot specify storage class "
1086 "when attaching to a non-scalar value";
1089 if (symbol == spirv::DecorationAttr::name) {
1090 if (!isa<spirv::DecorationAttr>(attr))
1092 << symbol <<
"' must be a spirv::DecorationAttr";
1096 return emitError(loc,
"found unsupported '")
1097 << symbol <<
"' attribute on region argument";
1100LogicalResult SPIRVDialect::verifyRegionArgAttribute(
Operation *op,
1101 unsigned regionIndex,
1104 auto funcOp = dyn_cast<FunctionOpInterface>(op);
1107 Type argType = funcOp.getArgumentTypes()[argIndex];
1112LogicalResult SPIRVDialect::verifyRegionResultAttribute(
1113 Operation *op,
unsigned ,
unsigned resultIndex,
1115 if (
auto graphOp = dyn_cast<spirv::GraphARMOp>(op))
1117 op->
getLoc(), graphOp.getResultTypes()[resultIndex], attribute);
1119 "cannot attach SPIR-V attributes to region result which is "
1120 "not part of a spirv::GraphARMOp type");
static bool isLegalToInline(InlinerInterface &interface, Region *src, Region *insertRegion, bool shouldCloneInlinedRegion, IRMapping &valueMapping)
Utility to check that all of the operations within 'src' can be inlined.
std::optional< unsigned > parseAndVerify< unsigned >(SPIRVDialect const &dialect, DialectAsmParser &parser)
static std::optional< IntTy > parseAndVerifyInteger(SPIRVDialect const &dialect, DialectAsmParser &parser)
static LogicalResult parseOptionalArrayStride(const SPIRVDialect &dialect, DialectAsmParser &parser, unsigned &stride)
Parses an optional , stride = N assembly segment.
static LogicalResult verifyRegionAttribute(Location loc, Type valueType, NamedAttribute attribute)
Verifies the given SPIR-V attribute attached to a value of the given valueType is valid.
static Type parseTensorArmType(SPIRVDialect const &dialect, DialectAsmParser &parser)
static void print(ArrayType type, DialectAsmPrinter &os)
static Type parseSampledImageType(SPIRVDialect const &dialect, DialectAsmParser &parser)
static Type parseAndVerifyType(SPIRVDialect const &dialect, DialectAsmParser &parser)
static ParseResult parseStructMemberDecorations(SPIRVDialect const &dialect, DialectAsmParser &parser, ArrayRef< Type > memberTypes, SmallVectorImpl< StructType::OffsetInfo > &offsetInfo, SmallVectorImpl< StructType::MemberDecorationInfo > &memberDecorationInfo)
static Type parseAndVerifySampledImageType(SPIRVDialect const &dialect, DialectAsmParser &parser)
static Type parseCooperativeMatrixType(SPIRVDialect const &dialect, DialectAsmParser &parser)
std::optional< Type > parseAndVerify< Type >(SPIRVDialect const &dialect, DialectAsmParser &parser)
static Type parseAndVerifyMatrixType(SPIRVDialect const &dialect, DialectAsmParser &parser)
static Type parseArrayType(SPIRVDialect const &dialect, DialectAsmParser &parser)
static bool containsReturn(Region ®ion)
Returns true if the given region contains spirv.Return or spirv.ReturnValue ops.
static Type parseStructType(SPIRVDialect const &dialect, DialectAsmParser &parser)
static Type parseRuntimeArrayType(SPIRVDialect const &dialect, DialectAsmParser &parser)
static Type parseMatrixType(SPIRVDialect const &dialect, DialectAsmParser &parser)
static Type parseImageType(SPIRVDialect const &dialect, DialectAsmParser &parser)
static std::optional< ValTy > parseAndVerify(SPIRVDialect const &dialect, DialectAsmParser &parser)
static Type parsePointerType(SPIRVDialect const &dialect, DialectAsmParser &parser)
virtual OptionalParseResult parseOptionalInteger(APInt &result)=0
Parse an optional integer value from the stream.
virtual ParseResult parseCommaSeparatedList(Delimiter delimiter, function_ref< ParseResult()> parseElementFn, StringRef contextMessage=StringRef())=0
Parse a list of comma-separated items with an optional delimiter.
virtual ParseResult parseOptionalEqual()=0
Parse a = token if 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.
ParseResult parseInteger(IntT &result)
Parse an integer value from the stream.
virtual ParseResult parseOptionalRParen()=0
Parse a ) token if present.
virtual ParseResult parseLess()=0
Parse a '<' token.
virtual ParseResult parseDimensionList(SmallVectorImpl< int64_t > &dimensions, bool allowDynamic=true, bool withTrailingX=true)=0
Parse a dimension list of a tensor or memref type.
virtual ParseResult parseOptionalGreater()=0
Parse a '>' token if present.
virtual ParseResult parseEqual()=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 SMLoc getNameLoc() const =0
Return the location of the original name token.
virtual ParseResult parseOptionalStar()=0
Parse a '*' token if present.
FailureOr< CyclicParseReset > tryStartCyclicParse(AttrOrTypeT attrOrType)
Attempts to start a cyclic parsing region for attrOrType.
virtual ParseResult parseOptionalRSquare()=0
Parse a ] token if present.
virtual ParseResult parseGreater()=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.
ParseResult parseKeyword(StringRef keyword)
Parse a given keyword.
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.
virtual ParseResult parseXInDimensionList()=0
Parse an 'x' token in a dimension list, handling the case where the x is juxtaposed with an element t...
virtual void printAttributeWithoutType(Attribute attr)
Print the given attribute without its type.
FailureOr< CyclicPrintReset > tryStartCyclicPrint(AttrOrTypeT attrOrType)
Attempts to start a cyclic printing region for attrOrType.
Attributes are known-constant values of operations.
Block represents an ordered list of Operations.
Operation * getTerminator()
Get the terminator operation of this block.
The DialectAsmParser has methods for interacting with the asm parser when parsing attributes and type...
This is a pure-virtual base class that exposes the asmprinter hooks necessary to implement a custom p...
This class defines the main interface for locations in MLIR and acts as a non-nullable wrapper around...
NamedAttribute represents a combination of a name and an Attribute value.
StringAttr getName() const
Return the name of the attribute.
Attribute getValue() const
Return the value of the attribute.
This class helps build Operations.
Operation is the basic unit of execution within MLIR.
Location getLoc()
The source location the operation was defined or derived from.
InFlightDiagnostic emitError(const Twine &message={})
Emit an error about fatal conditions with this operation, reporting up to any diagnostic handlers tha...
This class implements Optional functionality for ParseResult.
bool has_value() const
Returns true if we contain a valid ParseResult value.
This class contains a list of basic blocks and a link to the parent operation it is attached to.
Instances of the Type class are uniqued, have an immutable identifier and an optional mutable compone...
Dialect & getDialect() const
Get the dialect this type is registered to.
bool isIntOrIndexOrFloat() const
Return true if this is an integer (of any signedness), index, or float type.
Type getElementType() const
unsigned getArrayStride() const
Returns the array stride in bytes.
unsigned getNumElements() const
static ArrayType get(Type elementType, unsigned elementCount)
Scope getScope() const
Returns the scope of the matrix.
uint32_t getRows() const
Returns the number of rows of the matrix.
uint32_t getColumns() const
Returns the number of columns of the matrix.
static CooperativeMatrixType get(Type elementType, uint32_t rows, uint32_t columns, Scope scope, CooperativeMatrixUseKHR use)
Type getElementType() const
CooperativeMatrixUseKHR getUse() const
Returns the use parameter of the cooperative matrix.
static ImageType get(Type elementType, Dim dim, ImageDepthInfo depth=ImageDepthInfo::DepthUnknown, ImageArrayedInfo arrayed=ImageArrayedInfo::NonArrayed, ImageSamplingInfo samplingInfo=ImageSamplingInfo::SingleSampled, ImageSamplerUseInfo samplerUse=ImageSamplerUseInfo::SamplerUnknown, ImageFormat format=ImageFormat::Unknown)
ImageDepthInfo getDepthInfo() const
ImageArrayedInfo getArrayedInfo() const
ImageFormat getImageFormat() const
ImageSamplerUseInfo getSamplerUseInfo() const
Type getElementType() const
ImageSamplingInfo getSamplingInfo() const
static MatrixType get(Type columnType, uint32_t columnCount)
Type getColumnType() const
unsigned getNumColumns() const
Returns the number of columns.
static NamedBarrierType get(MLIRContext *context)
Type getPointeeType() const
StorageClass getStorageClass() const
static PointerType get(Type pointeeType, StorageClass storageClass)
Type getElementType() const
unsigned getArrayStride() const
Returns the array stride in bytes.
static RuntimeArrayType get(Type elementType)
Type getImageType() const
static SampledImageType get(Type imageType)
static SamplerType get(MLIRContext *context)
static bool isValid(FloatType)
Returns true if the given float type is valid for the SPIR-V dialect.
void getStructDecorations(SmallVectorImpl< StructType::StructDecorationInfo > &structDecorations) const
void getMemberDecorations(SmallVectorImpl< StructType::MemberDecorationInfo > &memberDecorations) const
static StructType getIdentified(MLIRContext *context, StringRef identifier)
Construct an identified StructType.
bool isIdentified() const
Returns true if the StructType is identified.
StringRef getIdentifier() const
For literal structs, return an empty string.
static StructType getEmpty(MLIRContext *context, StringRef identifier="")
Construct a (possibly identified) StructType with no members.
unsigned getNumElements() const
Type getElementType(unsigned) const
LogicalResult trySetBody(ArrayRef< Type > memberTypes, ArrayRef< OffsetInfo > offsetInfo={}, ArrayRef< MemberDecorationInfo > memberDecorations={}, ArrayRef< StructDecorationInfo > structDecorations={})
Sets the contents of an incomplete identified StructType.
static StructType get(ArrayRef< Type > memberTypes, ArrayRef< OffsetInfo > offsetInfo={}, ArrayRef< MemberDecorationInfo > memberDecorations={}, ArrayRef< StructDecorationInfo > structDecorations={})
Construct a literal StructType with at least one member.
uint64_t getMemberOffset(unsigned) const
Type getElementType() const
static TensorArmType get(ArrayRef< int64_t > shape, Type elementType)
ArrayRef< int64_t > getShape() const
StringRef getInterfaceVarABIAttrName()
Returns the attribute name for specifying argument ABI information.
StringRef getLoopControlAttrName()
Returns the attribute name for specifying loop control.
ParseResult parseEnumKeywordAttr(EnumClass &value, ParserType &parser, StringRef attrName=spirv::attributeName< EnumClass >())
Parses the next keyword in parser as an enumerant of the given EnumClass.
StringRef getTargetEnvAttrName()
Returns the attribute name for specifying SPIR-V target environment.
std::string getDecorationString(Decoration decoration)
Converts a SPIR-V Decoration enum value to its snake_case string representation for use in MLIR attri...
StringRef getSelectionControlAttrName()
Returns the attribute name for specifying selection control.
StringRef getEntryPointABIAttrName()
Returns the attribute name for specifying entry point information.
Include the generated interface declarations.
InFlightDiagnostic emitError(Location loc)
Utility method to emit an error message using this location.
llvm::TypeSwitch< T, ResultT > TypeSwitch