31#include "llvm/ADT/APFloat.h"
32#include "llvm/ADT/APInt.h"
33#include "llvm/ADT/ArrayRef.h"
34#include "llvm/ADT/STLExtras.h"
35#include "llvm/ADT/StringExtras.h"
36#include "llvm/ADT/TypeSwitch.h"
37#include "llvm/Support/InterleavedRange.h"
50 auto constOp = dyn_cast_or_null<spirv::ConstantOp>(op);
54 auto valueAttr = constOp.getValue();
55 auto integerValueAttr = dyn_cast<IntegerAttr>(valueAttr);
56 if (!integerValueAttr) {
60 if (integerValueAttr.getType().isSignlessInteger())
61 value = integerValueAttr.getInt();
63 value = integerValueAttr.getSInt();
70 spirv::MemorySemantics memorySemantics) {
77 auto atMostOneInSet = spirv::MemorySemantics::Acquire |
78 spirv::MemorySemantics::Release |
79 spirv::MemorySemantics::AcquireRelease |
80 spirv::MemorySemantics::SequentiallyConsistent;
83 llvm::popcount(
static_cast<uint32_t
>(memorySemantics & atMostOneInSet));
86 "expected at most one of these four memory constraints "
87 "to be set: `Acquire`, `Release`,"
88 "`AcquireRelease` or `SequentiallyConsistent`");
99 auto pointeePtrType = dyn_cast<spirv::PointerType>(pointeeType);
100 if (!pointeePtrType) {
101 if (
auto pointeeArrayType = dyn_cast<spirv::ArrayType>(pointeeType)) {
103 dyn_cast<spirv::PointerType>(pointeeArrayType.getElementType());
107 if (!pointeePtrType || pointeePtrType.getStorageClass() !=
108 spirv::StorageClass::PhysicalStorageBuffer)
111 auto getDecorationAttr = [op](spirv::Decoration decoration) {
116 getDecorationAttr(spirv::Decoration::AliasedPointer) !=
nullptr;
117 bool hasRestrictPtr =
118 getDecorationAttr(spirv::Decoration::RestrictPointer) !=
nullptr;
120 if (!hasAliasedPtr && !hasRestrictPtr)
122 <<
" with physical buffer pointer must be decorated "
123 "either 'AliasedPointer' or 'RestrictPointer'";
125 if (hasAliasedPtr && hasRestrictPtr)
127 <<
" with physical buffer pointer must have exactly one "
128 "aliasing decoration";
137 stringifyDecoration(spirv::Decoration::DescriptorSet));
138 auto bindingName = llvm::convertToSnakeFromCamelCase(
139 stringifyDecoration(spirv::Decoration::Binding));
142 if (descriptorSet && binding) {
145 printer <<
" bind(" << descriptorSet.getInt() <<
", " << binding.getInt()
150 auto builtInName = llvm::convertToSnakeFromCamelCase(
151 stringifyDecoration(spirv::Decoration::BuiltIn));
153 printer <<
" " << builtInName <<
"(\"" <<
builtin.getValue() <<
"\")";
154 elidedAttrs.push_back(builtInName);
172 auto fnType = dyn_cast<FunctionType>(type);
174 parser.
emitError(loc,
"expected function type");
179 result.addTypes(fnType.getResults());
190 assert(op->
getNumResults() == 1 &&
"op should have one result");
196 [&](
Type type) { return type != resultType; })) {
205 p <<
" : " << resultType;
208template <
typename BlockReadWriteOpTy>
212 if (
auto valVecTy = dyn_cast<VectorType>(valType))
213 valType = valVecTy.getElementType();
215 if (valType != cast<spirv::PointerType>(
ptr.getType()).getPointeeType()) {
216 return op.emitOpError(
"mismatch in result type and pointer type");
228 emitErrorFn(
"expected at least one index for spirv.CompositeExtract");
233 if (
auto cType = dyn_cast<spirv::CompositeType>(type)) {
234 if (cType.hasCompileTimeKnownNumElements() &&
236 static_cast<uint64_t
>(
index) >= cType.getNumElements())) {
237 emitErrorFn(
"index ") <<
index <<
" out of bounds for " << type;
240 type = cType.getElementType(
index);
242 emitErrorFn(
"cannot extract from non-composite type ")
243 << type <<
" with index " <<
index;
253 auto indicesArrayAttr = dyn_cast<ArrayAttr>(
indices);
254 if (!indicesArrayAttr) {
255 emitErrorFn(
"expected a 32-bit integer array attribute for 'indices'");
258 if (indicesArrayAttr.empty()) {
259 emitErrorFn(
"expected at least one index for spirv.CompositeExtract");
264 for (
auto indexAttr : indicesArrayAttr) {
265 auto indexIntAttr = dyn_cast<IntegerAttr>(indexAttr);
267 emitErrorFn(
"expected an 32-bit integer for index, but found '")
271 indexVals.push_back(indexIntAttr.getInt());
278 return ::mlir::emitError(loc, err);
291template <
typename ExtendedBinaryOp>
293 auto resultType = cast<spirv::StructType>(op.getType());
294 if (resultType.getNumElements() != 2)
295 return op.emitOpError(
"expected result struct type containing two members");
297 if (!llvm::all_equal({op.getOperand1().getType(), op.getOperand2().getType(),
298 resultType.getElementType(0),
299 resultType.getElementType(1)}))
300 return op.emitOpError(
301 "expected all operand types and struct member types are the same");
318 auto structType = dyn_cast<spirv::StructType>(resultType);
319 if (!structType || structType.getNumElements() != 2)
320 return parser.
emitError(loc,
"expected spirv.struct type with two members");
326 result.addTypes(resultType);
340 return op->
emitError(
"expected the same type for the first operand and "
341 "result, but provided ")
353 spirv::GlobalVariableOp var) {
354 build(builder, state, var.getType(), SymbolRefAttr::get(var));
357LogicalResult spirv::AddressOfOp::verify() {
358 auto varOp = dyn_cast_or_null<spirv::GlobalVariableOp>(
362 return emitOpError(
"expected spirv.GlobalVariable symbol");
364 if (getPointer().
getType() != varOp.getType()) {
366 "result type mismatch with the referenced global variable's type");
375LogicalResult spirv::CompositeConstructOp::verify() {
376 operand_range constituents = this->getConstituents();
391 if (coopElementType) {
392 if (constituents.size() != 1)
393 return emitOpError(
"has incorrect number of operands: expected ")
394 <<
"1, but provided " << constituents.size();
395 if (coopElementType != constituents.front().getType())
396 return emitOpError(
"operand type mismatch: expected operand type ")
397 << coopElementType <<
", but provided "
398 << constituents.front().getType();
403 auto cType = cast<spirv::CompositeType>(
getType());
404 if (constituents.size() == cType.getNumElements()) {
405 for (
auto index : llvm::seq<uint32_t>(0, constituents.size())) {
407 return emitOpError(
"operand type mismatch: expected operand type ")
408 << cType.getElementType(
index) <<
", but provided "
409 << constituents[
index].getType();
416 auto resultType = dyn_cast<VectorType>(cType);
419 "expected to return a vector or cooperative matrix when the number of "
420 "constituents is less than what the result needs");
423 for (
Value component : constituents) {
424 if (!isa<VectorType>(component.getType()) &&
425 !component.getType().isIntOrFloat())
426 return emitOpError(
"operand type mismatch: expected operand to have "
427 "a scalar or vector type, but provided ")
428 << component.getType();
430 Type elementType = component.getType();
431 if (
auto vectorType = dyn_cast<VectorType>(component.getType())) {
432 sizes.push_back(vectorType.getNumElements());
433 elementType = vectorType.getElementType();
438 if (elementType != resultType.getElementType())
439 return emitOpError(
"operand element type mismatch: expected to be ")
440 << resultType.getElementType() <<
", but provided " << elementType;
442 unsigned totalCount = llvm::sum_of(sizes);
443 if (totalCount != cType.getNumElements())
444 return emitOpError(
"has incorrect number of operands: expected ")
445 << cType.getNumElements() <<
", but provided " << totalCount;
462 build(builder, state, elementType, composite, indexAttr);
465ParseResult spirv::CompositeExtractOp::parse(
OpAsmParser &parser,
469 StringRef indicesAttrName =
470 spirv::CompositeExtractOp::getIndicesAttrName(
result.name);
487 result.addTypes(resultType);
491void spirv::CompositeExtractOp::print(
OpAsmPrinter &printer) {
492 printer <<
' ' << getComposite() <<
getIndices() <<
" : "
496LogicalResult spirv::CompositeExtractOp::verify() {
497 auto indicesArrayAttr = dyn_cast<ArrayAttr>(
getIndices());
504 return emitOpError(
"invalid result type: expected ")
505 << resultType <<
" but provided " <<
getType();
519 build(builder, state, composite.
getType(),
object, composite, indexAttr);
522ParseResult spirv::CompositeInsertOp::parse(
OpAsmParser &parser,
525 Type objectType, compositeType;
527 StringRef indicesAttrName =
528 spirv::CompositeInsertOp::getIndicesAttrName(
result.name);
541LogicalResult spirv::CompositeInsertOp::verify() {
542 auto indicesArrayAttr = dyn_cast<ArrayAttr>(
getIndices());
548 if (objectType != getObject().
getType()) {
549 return emitOpError(
"object operand type should be ")
550 << objectType <<
", but found " << getObject().getType();
554 return emitOpError(
"result type should be the same as "
555 "the composite type, but found ")
556 << getComposite().getType() <<
" vs " <<
getType();
562void spirv::CompositeInsertOp::print(
OpAsmPrinter &printer) {
563 printer <<
" " << getObject() <<
", " << getComposite() <<
getIndices()
564 <<
" : " << getObject().
getType() <<
" into "
565 << getComposite().getType();
572ParseResult spirv::ConstantOp::parse(
OpAsmParser &parser,
575 StringRef valueAttrName = spirv::ConstantOp::getValueAttrName(
result.name);
580 if (
auto typedAttr = dyn_cast<TypedAttr>(value))
581 type = typedAttr.getType();
582 if (isa<NoneType, TensorType>(type)) {
587 if (isa<TensorArmType>(type)) {
597 printer <<
' ' << getValue();
598 if (isa<spirv::ArrayType, spirv::StructType>(
getType()))
604 if (isa<spirv::CooperativeMatrixType>(opType)) {
605 auto denseAttr = dyn_cast<DenseElementsAttr>(value);
606 if (!denseAttr || !denseAttr.isSplat())
607 return op.emitOpError(
"expected a splat dense attribute for cooperative "
608 "matrix constant, but found ")
611 if (isa<IntegerAttr, FloatAttr>(value)) {
612 auto valueType = cast<TypedAttr>(value).getType();
613 if (valueType != opType)
614 return op.emitOpError(
"result type (")
615 << opType <<
") does not match value type (" << valueType <<
")";
618 if (isa<DenseTypedElementsAttr, SparseElementsAttr>(value)) {
619 auto valueType = cast<TypedAttr>(value).getType();
620 if (valueType == opType)
622 auto arrayType = dyn_cast<spirv::ArrayType>(opType);
623 auto shapedType = dyn_cast<ShapedType>(valueType);
625 return op.emitOpError(
"result or element type (")
626 << opType <<
") does not match value type (" << valueType
627 <<
"), must be the same or spirv.array";
629 int numElements = arrayType.getNumElements();
630 auto opElemType = arrayType.getElementType();
631 while (
auto t = dyn_cast<spirv::ArrayType>(opElemType)) {
632 numElements *= t.getNumElements();
633 opElemType = t.getElementType();
635 if (!opElemType.isIntOrFloat())
636 return op.emitOpError(
"only support nested array result type");
638 auto valueElemType = shapedType.getElementType();
639 if (valueElemType != opElemType) {
640 return op.emitOpError(
"result element type (")
641 << opElemType <<
") does not match value element type ("
642 << valueElemType <<
")";
645 if (numElements != shapedType.getNumElements()) {
646 return op.emitOpError(
"result number of elements (")
647 << numElements <<
") does not match value number of elements ("
648 << shapedType.getNumElements() <<
")";
652 if (
auto arrayAttr = dyn_cast<ArrayAttr>(value)) {
653 if (
auto structType = dyn_cast<spirv::StructType>(opType)) {
655 if (structType.isIdentified())
656 return op.emitOpError(
657 "cannot have an identified struct as a constant type");
658 if (arrayAttr.size() != structType.getNumElements())
659 return op.emitOpError(
"number of constituents (")
661 <<
") does not match number of struct members ("
662 << structType.getNumElements() <<
")";
663 for (
auto [idx, element] : llvm::enumerate(arrayAttr.getValue())) {
665 structType.getElementType(idx))))
670 auto arrayType = dyn_cast<spirv::ArrayType>(opType);
672 return op.emitOpError(
673 "must have spirv.array or spirv.struct result type for array value");
674 Type elemType = arrayType.getElementType();
675 for (
Attribute element : arrayAttr.getValue()) {
682 return op.emitOpError(
"cannot have attribute: ") << value;
685LogicalResult spirv::ConstantOp::verify() {
692bool spirv::ConstantOp::isBuildableWith(
Type type) {
694 if (!isa<spirv::SPIRVType>(type))
698 if (
auto structType = dyn_cast<spirv::StructType>(type))
699 return !structType.isIdentified();
700 return isa<spirv::ArrayType>(type);
706spirv::ConstantOp spirv::ConstantOp::getZero(
Type type,
Location loc,
708 if (
auto intType = dyn_cast<IntegerType>(type)) {
709 unsigned width = intType.getWidth();
711 return spirv::ConstantOp::create(builder, loc, type,
713 return spirv::ConstantOp::create(
714 builder, loc, type, builder.
getIntegerAttr(type, APInt(width, 0)));
716 if (
auto floatType = dyn_cast<FloatType>(type)) {
717 return spirv::ConstantOp::create(builder, loc, type,
720 if (
auto vectorType = dyn_cast<VectorType>(type)) {
721 Type elemType = vectorType.getElementType();
722 if (isa<IntegerType>(elemType)) {
723 return spirv::ConstantOp::create(
726 IntegerAttr::get(elemType, 0).getValue()));
728 if (isa<FloatType>(elemType)) {
729 return spirv::ConstantOp::create(
732 FloatAttr::get(elemType, 0.0).getValue()));
736 llvm_unreachable(
"unimplemented types for ConstantOp::getZero()");
739spirv::ConstantOp spirv::ConstantOp::getOne(
Type type,
Location loc,
741 if (
auto intType = dyn_cast<IntegerType>(type)) {
742 unsigned width = intType.getWidth();
744 return spirv::ConstantOp::create(builder, loc, type,
746 return spirv::ConstantOp::create(
747 builder, loc, type, builder.
getIntegerAttr(type, APInt(width, 1)));
749 if (
auto floatType = dyn_cast<FloatType>(type)) {
750 return spirv::ConstantOp::create(builder, loc, type,
753 if (
auto vectorType = dyn_cast<VectorType>(type)) {
754 Type elemType = vectorType.getElementType();
755 if (isa<IntegerType>(elemType)) {
756 return spirv::ConstantOp::create(
759 IntegerAttr::get(elemType, 1).getValue()));
761 if (isa<FloatType>(elemType)) {
762 return spirv::ConstantOp::create(
765 FloatAttr::get(elemType, 1.0).getValue()));
769 llvm_unreachable(
"unimplemented types for ConstantOp::getOne()");
772void mlir::spirv::ConstantOp::getAsmResultNames(
777 llvm::raw_svector_ostream specialName(specialNameBuffer);
778 specialName <<
"cst";
780 IntegerType intTy = dyn_cast<IntegerType>(type);
782 if (IntegerAttr intCst = dyn_cast<IntegerAttr>(getValue())) {
785 if (intTy.getWidth() == 1) {
786 return setNameFn(getResult(), (intCst.getInt() ?
"true" :
"false"));
789 if (intTy.isSignless()) {
790 specialName << intCst.getInt();
791 }
else if (intTy.isUnsigned()) {
792 specialName << intCst.getUInt();
794 specialName << intCst.getSInt();
798 if (intTy || isa<FloatType>(type)) {
799 specialName <<
'_' << type;
802 if (
auto vecType = dyn_cast<VectorType>(type)) {
803 specialName <<
"_vec_";
804 specialName << vecType.getDimSize(0);
806 Type elementType = vecType.getElementType();
808 if (isa<IntegerType>(elementType) || isa<FloatType>(elementType)) {
809 specialName <<
"x" << elementType;
813 setNameFn(getResult(), specialName.str());
816void mlir::spirv::AddressOfOp::getAsmResultNames(
819 llvm::raw_svector_ostream specialName(specialNameBuffer);
820 specialName << getVariable() <<
"_addr";
821 setNameFn(getResult(), specialName.str());
832 if (
auto typedAttr = dyn_cast<TypedAttr>(attr)) {
833 return typedAttr.getType();
836 if (
auto arrayAttr = dyn_cast<ArrayAttr>(attr)) {
843LogicalResult spirv::EXTConstantCompositeReplicateOp::verify() {
846 return emitError(
"unknown value attribute type");
848 auto compositeType = dyn_cast<spirv::CompositeType>(
getType());
850 return emitError(
"result type is not a composite type");
852 Type compositeElementType = compositeType.getElementType(0);
855 while (
auto type = dyn_cast<spirv::CompositeType>(compositeElementType)) {
856 compositeElementType = type.getElementType(0);
857 possibleTypes.push_back(compositeElementType);
860 if (!is_contained(possibleTypes, valueType)) {
861 return emitError(
"expected value attribute type ")
862 << interleaved(possibleTypes,
" or ") <<
", but got: " << valueType;
872LogicalResult spirv::ControlBarrierOp::verify() {
881 spirv::ExecutionModel executionModel,
882 spirv::FuncOp function,
884 build(builder, state,
885 spirv::ExecutionModelAttr::get(builder.
getContext(), executionModel),
886 SymbolRefAttr::get(function), builder.
getArrayAttr(interfaceVars));
889ParseResult spirv::EntryPointOp::parse(
OpAsmParser &parser,
891 spirv::ExecutionModel execModel;
904 FlatSymbolRefAttr var;
906 if (parser.parseAttribute(var, Type(),
"var_symbol", attrs))
908 interfaceVars.push_back(var);
913 result.addAttribute(spirv::EntryPointOp::getInterfaceAttrName(
result.name),
921 auto interfaceVars = getInterface().getValue();
922 if (!interfaceVars.empty())
923 printer <<
", " << llvm::interleaved(interfaceVars);
926LogicalResult spirv::EntryPointOp::verify() {
940struct ExecutionModeOperandSchema {
942 unsigned numOperands;
945ExecutionModeOperandSchema
946getExecutionModeOperandSchema(spirv::ExecutionMode mode) {
948 case spirv::ExecutionMode::Invocations:
949 case spirv::ExecutionMode::OutputVertices:
950 case spirv::ExecutionMode::VecTypeHint:
951 case spirv::ExecutionMode::SubgroupSize:
952 case spirv::ExecutionMode::SubgroupsPerWorkgroup:
953 case spirv::ExecutionMode::DenormPreserve:
954 case spirv::ExecutionMode::DenormFlushToZero:
955 case spirv::ExecutionMode::SignedZeroInfNanPreserve:
956 case spirv::ExecutionMode::RoundingModeRTE:
957 case spirv::ExecutionMode::RoundingModeRTZ:
958 case spirv::ExecutionMode::OutputPrimitivesEXT:
959 case spirv::ExecutionMode::SharedLocalMemorySizeINTEL:
960 case spirv::ExecutionMode::RoundingModeRTPINTEL:
961 case spirv::ExecutionMode::RoundingModeRTNINTEL:
962 case spirv::ExecutionMode::FloatingPointModeALTINTEL:
963 case spirv::ExecutionMode::FloatingPointModeIEEEINTEL:
964 case spirv::ExecutionMode::MaxWorkDimINTEL:
965 case spirv::ExecutionMode::NumSIMDWorkitemsINTEL:
966 case spirv::ExecutionMode::SchedulerTargetFmaxMhzINTEL:
967 case spirv::ExecutionMode::StreamingInterfaceINTEL:
968 case spirv::ExecutionMode::NamedBarrierCountINTEL:
970 case spirv::ExecutionMode::LocalSize:
971 case spirv::ExecutionMode::LocalSizeHint:
972 case spirv::ExecutionMode::MaxWorkgroupSizeINTEL:
974 case spirv::ExecutionMode::SubgroupsPerWorkgroupId:
976 case spirv::ExecutionMode::LocalSizeId:
977 case spirv::ExecutionMode::LocalSizeHintId:
990 spirv::FuncOp function,
991 spirv::ExecutionMode executionMode,
993 build(builder, state, SymbolRefAttr::get(function),
994 spirv::ExecutionModeAttr::get(builder.
getContext(), executionMode),
998ParseResult spirv::ExecutionModeOp::parse(
OpAsmParser &parser,
1000 spirv::ExecutionMode execMode;
1015 values.push_back(cast<IntegerAttr>(value).getInt());
1017 StringRef valuesAttrName =
1018 spirv::ExecutionModeOp::getValuesAttrName(
result.name);
1019 result.addAttribute(valuesAttrName,
1024void spirv::ExecutionModeOp::print(
OpAsmPrinter &printer) {
1027 printer <<
" \"" << stringifyExecutionMode(getExecutionMode()) <<
"\"";
1029 if (!values.empty())
1030 printer <<
", " << llvm::interleaved(values.getAsValueRange<IntegerAttr>());
1033LogicalResult spirv::ExecutionModeOp::verify() {
1034 ExecutionModeOperandSchema schema =
1035 getExecutionModeOperandSchema(getExecutionMode());
1037 if (schema.isIdOperand)
1038 return emitOpError(
"expected ExecutionMode that takes extra operands "
1039 "that are not <id> operands, got: ")
1040 << stringifyExecutionMode(getExecutionMode());
1042 if (getValues().size() != schema.numOperands)
1044 << schema.numOperands <<
" value operand(s), got "
1045 << getValues().size();
1054ParseResult spirv::ExecutionModeIdOp::parse(
OpAsmParser &parser,
1056 ExecutionMode execMode;
1065 FlatSymbolRefAttr attr;
1066 if (parser.parseAttribute(attr))
1068 values.push_back(attr);
1074 StringRef valuesAttrName = getValuesAttrName(
result.name);
1076 result.addAttribute(valuesAttrName, valuesAttr);
1080void spirv::ExecutionModeIdOp::print(
OpAsmPrinter &printer) {
1083 printer <<
" \"" << stringifyExecutionMode(getExecutionMode()) <<
"\" ";
1085 llvm::interleaveComma(
1086 getValues().getAsValueRange<FlatSymbolRefAttr>(), printer,
1090LogicalResult spirv::ExecutionModeIdOp::verify() {
1091 ExecutionModeOperandSchema schema =
1092 getExecutionModeOperandSchema(getExecutionMode());
1094 if (!schema.isIdOperand)
1095 return emitOpError(
"expected ExecutionMode that takes extra operands that "
1096 "are <id> operands, got: ")
1097 << stringifyExecutionMode(getExecutionMode());
1099 if (getValues().size() != schema.numOperands)
1101 << schema.numOperands <<
" value operand(s), got "
1102 << getValues().size();
1105 auto valueSymbol = dyn_cast<FlatSymbolRefAttr>(value);
1107 return emitOpError(
"expected value operands to be symbol reference");
1109 (*this)->getParentOp(), valueSymbol);
1111 return emitOpError(
"cannot find symbol referenced by value operand: ")
1112 << valueSymbol.getValue();
1129 StringAttr nameAttr;
1135 bool isVariadic =
false;
1137 parser,
false, entryArgs, isVariadic, resultTypes,
1142 for (
auto &arg : entryArgs)
1143 argTypes.push_back(arg.type);
1145 result.addAttribute(getFunctionTypeAttrName(
result.name),
1146 TypeAttr::get(fnType));
1149 spirv::FunctionControl fnControl;
1158 assert(resultAttrs.size() == resultTypes.size());
1160 builder,
result, entryArgs, resultAttrs, getArgAttrsAttrName(
result.name),
1161 getResAttrsAttrName(
result.name));
1164 auto *body =
result.addRegion();
1174 auto fnType = getFunctionType();
1176 printer, *
this, fnType.getInputs(),
1177 false, fnType.getResults());
1178 printer <<
" \"" << spirv::stringifyFunctionControl(getFunctionControl())
1182 {spirv::attributeName<spirv::FunctionControl>(),
1183 getFunctionTypeAttrName(), getArgAttrsAttrName(), getResAttrsAttrName(),
1184 getFunctionControlAttrName()});
1187 Region &body = this->getBody();
1188 if (!body.empty()) {
1195LogicalResult spirv::FuncOp::verifyType() {
1196 FunctionType fnType = getFunctionType();
1197 if (fnType.getNumResults() > 1)
1198 return emitOpError(
"cannot have more than one result");
1200 auto hasDecorationAttr = [&](spirv::Decoration decoration,
1201 unsigned argIndex) {
1202 auto func = cast<FunctionOpInterface>(getOperation());
1203 for (
auto argAttr : cast<FunctionOpInterface>(
func).
getArgAttrs(argIndex)) {
1204 if (argAttr.getName() != spirv::DecorationAttr::name)
1206 if (
auto decAttr = dyn_cast<spirv::DecorationAttr>(argAttr.getValue()))
1207 return decAttr.getValue() == decoration;
1212 for (
unsigned i = 0, e = this->getNumArguments(); i != e; ++i) {
1213 Type param = fnType.getInputs()[i];
1214 auto inputPtrType = dyn_cast<spirv::PointerType>(param);
1218 auto pointeePtrType =
1219 dyn_cast<spirv::PointerType>(inputPtrType.getPointeeType());
1220 if (pointeePtrType) {
1226 if (pointeePtrType.getStorageClass() !=
1227 spirv::StorageClass::PhysicalStorageBuffer)
1230 bool hasAliasedPtr =
1231 hasDecorationAttr(spirv::Decoration::AliasedPointer, i);
1232 bool hasRestrictPtr =
1233 hasDecorationAttr(spirv::Decoration::RestrictPointer, i);
1234 if (!hasAliasedPtr && !hasRestrictPtr)
1236 <<
"with a pointer points to a physical buffer pointer must "
1237 "be decorated either 'AliasedPointer' or 'RestrictPointer'";
1244 if (
auto pointeeArrayType =
1245 dyn_cast<spirv::ArrayType>(inputPtrType.getPointeeType())) {
1247 dyn_cast<spirv::PointerType>(pointeeArrayType.getElementType());
1249 pointeePtrType = inputPtrType;
1252 if (!pointeePtrType || pointeePtrType.getStorageClass() !=
1253 spirv::StorageClass::PhysicalStorageBuffer)
1256 bool hasAliased = hasDecorationAttr(spirv::Decoration::Aliased, i);
1257 bool hasRestrict = hasDecorationAttr(spirv::Decoration::Restrict, i);
1258 if (!hasAliased && !hasRestrict)
1259 return emitOpError() <<
"with physical buffer pointer must be decorated "
1260 "either 'Aliased' or 'Restrict'";
1266LogicalResult spirv::FuncOp::verifyBody() {
1267 FunctionType fnType = getFunctionType();
1268 if (!isExternal()) {
1269 Block &entryBlock = front();
1271 unsigned numArguments = this->getNumArguments();
1274 << numArguments <<
" arguments to match function signature";
1276 for (
auto [
index, fnArgType, blockArgType] :
1278 if (blockArgType != fnArgType) {
1279 return emitOpError(
"type of entry block argument #")
1280 <<
index <<
'(' << blockArgType
1281 <<
") must match the type of the corresponding argument in "
1282 <<
"function signature(" << fnArgType <<
')';
1288 if (
auto retOp = dyn_cast<spirv::ReturnOp>(op)) {
1289 if (fnType.getNumResults() != 0)
1290 return retOp.emitOpError(
"cannot be used in functions returning value");
1291 }
else if (
auto retOp = dyn_cast<spirv::ReturnValueOp>(op)) {
1292 if (fnType.getNumResults() != 1)
1293 return retOp.emitOpError(
1294 "returns 1 value but enclosing function requires ")
1295 << fnType.getNumResults() <<
" results";
1297 auto retOperandType = retOp.getValue().getType();
1298 auto fnResultType = fnType.getResult(0);
1299 if (retOperandType != fnResultType)
1300 return retOp.emitOpError(
" return value's type (")
1301 << retOperandType <<
") mismatch with function's result type ("
1302 << fnResultType <<
")";
1309 return failure(walkResult.wasInterrupted());
1313 StringRef name, FunctionType type,
1314 spirv::FunctionControl control,
1318 state.
addAttribute(getFunctionTypeAttrName(state.
name), TypeAttr::get(type));
1319 state.
addAttribute(spirv::attributeName<spirv::FunctionControl>(),
1320 builder.
getAttr<spirv::FunctionControlAttr>(control));
1329ParseResult spirv::GLFClampOp::parse(
OpAsmParser &parser,
1339ParseResult spirv::GLUClampOp::parse(
OpAsmParser &parser,
1349ParseResult spirv::GLSClampOp::parse(
OpAsmParser &parser,
1359ParseResult spirv::GLNClampOp::parse(
OpAsmParser &parser,
1369ParseResult spirv::GLSmoothStepOp::parse(
OpAsmParser &parser,
1391 Type type, StringRef name,
1392 unsigned descriptorSet,
unsigned binding) {
1393 build(builder, state, TypeAttr::get(type), builder.
getStringAttr(name));
1395 spirv::SPIRVDialect::getAttributeName(spirv::Decoration::DescriptorSet),
1398 spirv::SPIRVDialect::getAttributeName(spirv::Decoration::Binding),
1403 Type type, StringRef name,
1405 build(builder, state, TypeAttr::get(type), builder.
getStringAttr(name));
1407 spirv::SPIRVDialect::getAttributeName(spirv::Decoration::BuiltIn),
1411ParseResult spirv::GlobalVariableOp::parse(
OpAsmParser &parser,
1414 StringAttr nameAttr;
1415 StringRef initializerAttrName =
1416 spirv::GlobalVariableOp::getInitializerAttrName(
result.name);
1437 StringRef typeAttrName =
1438 spirv::GlobalVariableOp::getTypeAttrName(
result.name);
1443 if (!isa<spirv::PointerType>(type)) {
1444 return parser.
emitError(loc,
"expected spirv.ptr type");
1446 result.addAttribute(typeAttrName, TypeAttr::get(type));
1451void spirv::GlobalVariableOp::print(
OpAsmPrinter &printer) {
1453 spirv::attributeName<spirv::StorageClass>()};
1460 StringRef initializerAttrName = this->getInitializerAttrName();
1462 if (
auto initializer = this->getInitializer()) {
1463 printer <<
" " << initializerAttrName <<
'(';
1466 elidedAttrs.push_back(initializerAttrName);
1469 StringRef typeAttrName = this->getTypeAttrName();
1470 elidedAttrs.push_back(typeAttrName);
1472 printer <<
" : " <<
getType();
1475LogicalResult spirv::GlobalVariableOp::verify() {
1476 if (!isa<spirv::PointerType>(
getType()))
1477 return emitOpError(
"result must be of a !spv.ptr type");
1483 auto storageClass = this->storageClass();
1484 if (storageClass == spirv::StorageClass::Generic ||
1485 storageClass == spirv::StorageClass::Function) {
1487 << stringifyStorageClass(storageClass) <<
"'";
1492 if (std::optional<spirv::LinkageAttributesAttr> linkage =
1493 getLinkageAttributes()) {
1494 if (linkage->getLinkageType().getValue() == spirv::LinkageType::Import &&
1497 "with Import linkage type must not have an initializer");
1502 this->getInitializerAttrName())) {
1504 (*this)->getParentOp(), init.getAttr());
1516 !isa<spirv::SpecConstantOp, spirv::SpecConstantCompositeOp>(initOp)) {
1517 return emitOpError(
"initializer must be result of a "
1518 "spirv.SpecConstant or "
1519 "spirv.SpecConstantCompositeOp op");
1523 Type pointeeType = cast<spirv::PointerType>(
getType()).getPointeeType();
1535LogicalResult spirv::INTELSubgroupBlockReadOp::verify() {
1546ParseResult spirv::INTELSubgroupBlockWriteOp::parse(
OpAsmParser &parser,
1549 spirv::StorageClass storageClass;
1560 if (
auto valVecTy = dyn_cast<VectorType>(elementType))
1570void spirv::INTELSubgroupBlockWriteOp::print(
OpAsmPrinter &printer) {
1571 printer <<
" " << getPtr() <<
", " << getValue() <<
" : "
1572 << getValue().getType();
1575LogicalResult spirv::INTELSubgroupBlockWriteOp::verify() {
1586LogicalResult spirv::IAddCarryOp::verify() {
1587 return ::verifyArithmeticExtendedBinaryOp(*
this);
1590ParseResult spirv::IAddCarryOp::parse(
OpAsmParser &parser,
1592 return ::parseArithmeticExtendedBinaryOp(parser,
result);
1603LogicalResult spirv::ISubBorrowOp::verify() {
1604 return ::verifyArithmeticExtendedBinaryOp(*
this);
1607ParseResult spirv::ISubBorrowOp::parse(
OpAsmParser &parser,
1609 return ::parseArithmeticExtendedBinaryOp(parser,
result);
1612void spirv::ISubBorrowOp::print(
OpAsmPrinter &printer) {
1620LogicalResult spirv::SMulExtendedOp::verify() {
1621 return ::verifyArithmeticExtendedBinaryOp(*
this);
1624ParseResult spirv::SMulExtendedOp::parse(
OpAsmParser &parser,
1626 return ::parseArithmeticExtendedBinaryOp(parser,
result);
1629void spirv::SMulExtendedOp::print(
OpAsmPrinter &printer) {
1637LogicalResult spirv::UMulExtendedOp::verify() {
1638 return ::verifyArithmeticExtendedBinaryOp(*
this);
1641ParseResult spirv::UMulExtendedOp::parse(
OpAsmParser &parser,
1643 return ::parseArithmeticExtendedBinaryOp(parser,
result);
1646void spirv::UMulExtendedOp::print(
OpAsmPrinter &printer) {
1654LogicalResult spirv::MemoryBarrierOp::verify() {
1662LogicalResult spirv::MemoryNamedBarrierOp::verify() {
1671 std::optional<StringRef> name) {
1681 spirv::AddressingModel addressingModel,
1682 spirv::MemoryModel memoryModel,
1683 std::optional<VerCapExtAttr> vceTriple,
1684 std::optional<StringRef> name) {
1687 builder.
getAttr<spirv::AddressingModelAttr>(addressingModel));
1689 builder.
getAttr<spirv::MemoryModelAttr>(memoryModel));
1693 state.
addAttribute(getVCETripleAttrName(), *vceTriple);
1699ParseResult spirv::ModuleOp::parse(
OpAsmParser &parser,
1704 StringAttr nameAttr;
1709 spirv::AddressingModel addrModel;
1710 spirv::MemoryModel memoryModel;
1720 spirv::ModuleOp::getVCETripleAttrName(),
1737 if (std::optional<StringRef> name = getName()) {
1746 auto addressingModelAttrName = spirv::attributeName<spirv::AddressingModel>();
1747 auto memoryModelAttrName = spirv::attributeName<spirv::MemoryModel>();
1748 elidedAttrs.assign({addressingModelAttrName, memoryModelAttrName,
1751 if (std::optional<spirv::VerCapExtAttr> triple = getVceTriple()) {
1752 printer <<
" requires " << *triple;
1753 elidedAttrs.push_back(spirv::ModuleOp::getVCETripleAttrName());
1761LogicalResult spirv::ModuleOp::verifyRegions() {
1762 Dialect *dialect = (*this)->getDialect();
1767 for (
auto &op : *getBody()) {
1769 return op.
emitError(
"'spirv.module' can only contain spirv.* ops");
1774 if (
auto entryPointOp = dyn_cast<spirv::EntryPointOp>(op)) {
1775 auto funcOp = table.lookup<spirv::FuncOp>(entryPointOp.getFn());
1777 return entryPointOp.emitError(
"function '")
1778 << entryPointOp.getFn() <<
"' not found in 'spirv.module'";
1780 if (
auto interface = entryPointOp.getInterface()) {
1782 auto varSymRef = dyn_cast<FlatSymbolRefAttr>(varRef);
1784 return entryPointOp.emitError(
1785 "expected symbol reference for interface "
1786 "specification instead of '")
1790 table.lookup<spirv::GlobalVariableOp>(varSymRef.getValue());
1792 return entryPointOp.emitError(
"expected spirv.GlobalVariable "
1793 "symbol reference instead of'")
1794 << varSymRef <<
"'";
1799 auto key = std::pair<spirv::FuncOp, spirv::ExecutionModel>(
1800 funcOp, entryPointOp.getExecutionModel());
1801 if (!entryPoints.try_emplace(key, entryPointOp).second)
1802 return entryPointOp.emitError(
"duplicate of a previous EntryPointOp");
1803 }
else if (
auto funcOp = dyn_cast<spirv::FuncOp>(op)) {
1807 auto linkageAttr = funcOp.getLinkageAttributes();
1808 auto hasImportLinkage =
1809 linkageAttr && (linkageAttr.value().getLinkageType().getValue() ==
1810 spirv::LinkageType::Import);
1811 if (funcOp.isExternal() && !hasImportLinkage)
1813 "'spirv.module' cannot contain external functions "
1814 "without 'Import' linkage_attributes (LinkageAttributes)");
1817 for (
auto &block : funcOp)
1818 for (
auto &op : block) {
1821 "functions in 'spirv.module' can only contain spirv.* ops");
1833LogicalResult spirv::ReferenceOfOp::verify() {
1835 (*this)->getParentOp(), getSpecConstAttr());
1838 auto specConstOp = dyn_cast_or_null<spirv::SpecConstantOp>(specConstSym);
1840 constType = specConstOp.getDefaultValue().getType();
1842 auto specConstCompositeOp =
1843 dyn_cast_or_null<spirv::SpecConstantCompositeOp>(specConstSym);
1844 if (specConstCompositeOp)
1845 constType = specConstCompositeOp.getType();
1847 if (!specConstOp && !specConstCompositeOp)
1849 "expected spirv.SpecConstant or spirv.SpecConstantComposite symbol");
1851 if (getReference().
getType() != constType)
1852 return emitOpError(
"result type mismatch with the referenced "
1853 "specialization constant's type");
1862ParseResult spirv::SpecConstantOp::parse(
OpAsmParser &parser,
1864 StringAttr nameAttr;
1866 StringRef defaultValueAttrName =
1867 spirv::SpecConstantOp::getDefaultValueAttrName(
result.name);
1875 IntegerAttr specIdAttr;
1889void spirv::SpecConstantOp::print(
OpAsmPrinter &printer) {
1892 if (
auto specID = (*this)->getAttrOfType<IntegerAttr>(
kSpecIdAttrName))
1894 printer <<
" = " << getDefaultValue();
1897LogicalResult spirv::SpecConstantOp::verify() {
1898 if (
auto specID = (*this)->getAttrOfType<IntegerAttr>(
kSpecIdAttrName))
1899 if (specID.getValue().isNegative())
1902 auto value = getDefaultValue();
1903 if (isa<IntegerAttr, FloatAttr>(value)) {
1905 if (!isa<spirv::SPIRVType>(value.getType()))
1906 return emitOpError(
"default value bitwidth disallowed");
1910 "default value can only be a bool, integer, or float scalar");
1917LogicalResult spirv::VectorShuffleOp::verify() {
1918 VectorType resultType = cast<VectorType>(
getType());
1920 size_t numResultElements = resultType.getNumElements();
1921 if (numResultElements != getComponents().size())
1923 << numResultElements
1924 <<
") mismatch with the number of component selectors ("
1925 << getComponents().size() <<
")";
1927 size_t totalSrcElements =
1928 cast<VectorType>(getVector1().
getType()).getNumElements() +
1929 cast<VectorType>(getVector2().
getType()).getNumElements();
1931 for (
const auto &selector : getComponents().getAsValueRange<IntegerAttr>()) {
1932 uint32_t
index = selector.getZExtValue();
1933 if (
index >= totalSrcElements &&
1934 index != std::numeric_limits<uint32_t>().
max())
1936 <<
index <<
" out of range: expected to be in [0, "
1937 << totalSrcElements <<
") or 0xffffffff";
1946ParseResult spirv::SpecConstantCompositeOp::parse(
OpAsmParser &parser,
1949 StringAttr compositeName;
1961 const char *attrName =
"spec_const";
1968 constituents.push_back(specConstRef);
1974 StringAttr compositeSpecConstituentsName =
1975 spirv::SpecConstantCompositeOp::getConstituentsAttrName(
result.name);
1976 result.addAttribute(compositeSpecConstituentsName,
1983 StringAttr typeAttrName =
1984 spirv::SpecConstantCompositeOp::getTypeAttrName(
result.name);
1985 result.addAttribute(typeAttrName, TypeAttr::get(type));
1990void spirv::SpecConstantCompositeOp::print(
OpAsmPrinter &printer) {
1993 printer <<
" (" << llvm::interleaved(this->getConstituents().getValue())
1997LogicalResult spirv::SpecConstantCompositeOp::verify() {
1998 auto cType = dyn_cast<spirv::CompositeType>(
getType());
1999 auto constituents = this->getConstituents().getValue();
2002 return emitError(
"result type must be a composite type, but provided ")
2005 if (isa<spirv::CooperativeMatrixType>(cType))
2006 return emitError(
"unsupported composite type ") << cType;
2007 if (constituents.size() != cType.getNumElements())
2008 return emitError(
"has incorrect number of operands: expected ")
2009 << cType.getNumElements() <<
", but provided "
2010 << constituents.size();
2012 for (
auto index : llvm::seq<uint32_t>(0, constituents.size())) {
2013 auto constituent = cast<FlatSymbolRefAttr>(constituents[
index]);
2016 (*this)->getParentOp(), constituent.getAttr());
2019 return emitError(
"unknown constituent symbol ") << constituent.getAttr();
2021 Type constituentType;
2022 if (
auto specConstOp = dyn_cast<spirv::SpecConstantOp>(constituentOp)) {
2023 constituentType = specConstOp.getDefaultValue().getType();
2024 }
else if (
auto specConstCompositeOp =
2025 dyn_cast<spirv::SpecConstantCompositeOp>(constituentOp)) {
2026 constituentType = specConstCompositeOp.getType();
2028 return emitError(
"unsupported constituent ")
2029 << constituent.getAttr()
2030 <<
": must reference a spirv.SpecConstant or "
2031 "spirv.SpecConstantComposite";
2034 if (constituentType != cType.getElementType(
index))
2035 return emitError(
"has incorrect types of operands: expected ")
2036 << cType.getElementType(
index) <<
", but provided "
2048spirv::EXTSpecConstantCompositeReplicateOp::parse(
OpAsmParser &parser,
2050 StringAttr compositeName;
2052 const char *attrName =
"spec_const";
2063 StringAttr compositeSpecConstituentName =
2064 spirv::EXTSpecConstantCompositeReplicateOp::getConstituentAttrName(
2066 result.addAttribute(compositeSpecConstituentName, specConstRef);
2068 StringAttr typeAttrName =
2069 spirv::EXTSpecConstantCompositeReplicateOp::getTypeAttrName(
result.name);
2070 result.addAttribute(typeAttrName, TypeAttr::get(type));
2075void spirv::EXTSpecConstantCompositeReplicateOp::print(
OpAsmPrinter &printer) {
2078 printer <<
" (" << this->getConstituent() <<
") : " <<
getType();
2081LogicalResult spirv::EXTSpecConstantCompositeReplicateOp::verify() {
2082 auto compositeType = dyn_cast<spirv::CompositeType>(
getType());
2084 return emitError(
"result type must be a composite type, but provided ")
2088 (*this)->getParentOp(), this->getConstituent());
2091 "splat spec constant reference defining constituent not found");
2093 auto constituentSpecConstOp = dyn_cast<spirv::SpecConstantOp>(constituentOp);
2094 if (!constituentSpecConstOp)
2095 return emitError(
"constituent is not a spec constant");
2097 Type constituentType = constituentSpecConstOp.getDefaultValue().getType();
2098 Type compositeElementType = compositeType.getElementType(0);
2099 if (constituentType != compositeElementType)
2100 return emitError(
"constituent has incorrect type: expected ")
2101 << compositeElementType <<
", but provided " << constituentType;
2110ParseResult spirv::SpecConstantOperationOp::parse(
OpAsmParser &parser,
2126 spirv::YieldOp::create(builder, wrappedOp->
getLoc(), wrappedOp->
getResult(0));
2137void spirv::SpecConstantOperationOp::print(
OpAsmPrinter &printer) {
2138 printer <<
" wraps ";
2142LogicalResult spirv::SpecConstantOperationOp::verifyRegions() {
2143 Block &block = getRegion().getBlocks().
front();
2146 return emitOpError(
"expected exactly 2 nested ops");
2154 if (!isa_and_present<spirv::ConstantOp, spirv::ReferenceOfOp,
2155 spirv::SpecConstantOperationOp>(
2156 operand.getDefiningOp()))
2158 "invalid operand, must be defined by a constant operation");
2167LogicalResult spirv::GLFrexpStructOp::verify() {
2169 dyn_cast<spirv::StructType>(getResult().
getType());
2172 return emitError(
"result type must be a struct type with two memebers");
2176 VectorType exponentVecTy = dyn_cast<VectorType>(exponentTy);
2177 IntegerType exponentIntTy = dyn_cast<IntegerType>(exponentTy);
2179 Type operandTy = getOperand().getType();
2180 VectorType operandVecTy = dyn_cast<VectorType>(operandTy);
2181 FloatType operandFTy = dyn_cast<FloatType>(operandTy);
2183 if (significandTy != operandTy)
2184 return emitError(
"member zero of the resulting struct type must be the "
2185 "same type as the operand");
2187 if (exponentVecTy) {
2188 IntegerType componentIntTy =
2189 dyn_cast<IntegerType>(exponentVecTy.getElementType());
2190 if (!componentIntTy || componentIntTy.getWidth() != 32)
2191 return emitError(
"member one of the resulting struct type must"
2192 "be a scalar or vector of 32 bit integer type");
2193 }
else if (!exponentIntTy || exponentIntTy.getWidth() != 32) {
2194 return emitError(
"member one of the resulting struct type "
2195 "must be a scalar or vector of 32 bit integer type");
2199 if (operandVecTy && exponentVecTy &&
2200 (exponentVecTy.getNumElements() == operandVecTy.getNumElements()))
2203 if (operandFTy && exponentIntTy)
2206 return emitError(
"member one of the resulting struct type must have the same "
2207 "number of components as the operand type");
2216 if (isa<FloatType>(floatType) != isa<IntegerType>(integerType))
2217 return op->
emitOpError(
"operands must both be scalars or vectors");
2220 if (
auto vectorType = dyn_cast<VectorType>(type))
2221 return vectorType.getNumElements();
2226 return op->
emitOpError(
"operands must have the same number of elements");
2231LogicalResult spirv::GLLdexpOp::verify() {
2240LogicalResult spirv::CLLdexpOp::verify() {
2249LogicalResult spirv::CLPownOp::verify() {
2258LogicalResult spirv::CLRootnOp::verify() {
2267LogicalResult spirv::ShiftLeftLogicalOp::verify() {
2275LogicalResult spirv::ShiftRightArithmeticOp::verify() {
2283LogicalResult spirv::ShiftRightLogicalOp::verify() {
2291LogicalResult spirv::VectorTimesScalarOp::verify() {
2293 return emitOpError(
"vector operand and result type mismatch");
2294 auto scalarType = cast<VectorType>(
getType()).getElementType();
2295 if (getScalar().
getType() != scalarType)
2296 return emitOpError(
"scalar operand and result element type match");
p<< " : "<< getMemRefType()<< ", "<< getType();}static LogicalResult verifyVectorMemoryOp(Operation *op, MemRefType memrefType, VectorType vectorType) { if(memrefType.getElementType() !=vectorType.getElementType()) return op-> emitOpError("requires memref and vector types of the same elemental type")
Given a list of lists of parsed operands, populates uniqueOperands with unique operands.
static std::string bindingName()
Returns the string name of the Binding decoration.
static std::string descriptorSetName()
Returns the string name of the DescriptorSet decoration.
static int64_t getNumElements(Type t)
Compute the total number of elements in the given type, also taking into account nested types.
static Value max(ImplicitLocOpBuilder &builder, Value value, Value bound)
static ParseResult parseArithmeticExtendedBinaryOp(OpAsmParser &parser, OperationState &result)
static Type getValueType(Attribute attr)
static LogicalResult verifyConstantType(spirv::ConstantOp op, Attribute value, Type opType)
static ParseResult parseOneResultSameOperandTypeOp(OpAsmParser &parser, OperationState &result)
static LogicalResult verifyArithmeticExtendedBinaryOp(ExtendedBinaryOp op)
static LogicalResult verifyFloatIntegerBuiltin(Operation *op, Type floatType, Type integerType)
static LogicalResult verifyShiftOp(Operation *op)
static LogicalResult verifyBlockReadWritePtrAndValTypes(BlockReadWriteOpTy op, Value ptr, Value val)
static Type getElementType(Type type, ArrayRef< int32_t > indices, function_ref< InFlightDiagnostic(StringRef)> emitErrorFn)
Walks the given type hierarchy with the given indices, potentially down to component granularity,...
static void printOneResultOp(Operation *op, OpAsmPrinter &p)
static void printArithmeticExtendedBinaryOp(Operation *op, OpAsmPrinter &printer)
ParseResult parseSymbolName(StringAttr &result)
Parse an -identifier and store it (without the '@' symbol) in a string attribute.
virtual ParseResult parseOptionalSymbolName(StringAttr &result)=0
Parse an optional -identifier and store it (without the '@' symbol) in a string attribute.
virtual Builder & getBuilder() const =0
Return a builder which provides useful access to MLIRContext, global objects like types and attribute...
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 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.
MLIRContext * getContext() const
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 parseOptionalColon()=0
Parse a : token if present.
ParseResult addTypeToList(Type type, SmallVectorImpl< Type > &result)
Add the specified type to the end of the specified type list and return success.
virtual ParseResult parseEqual()=0
Parse a = token.
virtual ParseResult parseOptionalAttrDictWithKeyword(NamedAttrList &result)=0
Parse a named dictionary into 'result' if the attributes keyword is present.
virtual ParseResult parseColonType(Type &result)=0
Parse a colon followed by a type.
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.
ParseResult addTypesToList(ArrayRef< Type > types, SmallVectorImpl< Type > &result)
Add the specified types to the end of the specified type list and return success.
virtual ParseResult parseLParen()=0
Parse a ( token.
virtual ParseResult parseType(Type &result)=0
Parse a type.
virtual ParseResult parseOptionalLParen()=0
Parse a ( token if present.
ParseResult parseKeywordType(const char *keyword, Type &result)
Parse a keyword followed by a type.
ParseResult parseKeyword(StringRef keyword)
Parse a given keyword.
virtual ParseResult parseAttribute(Attribute &result, Type type={})=0
Parse an arbitrary attribute of a given type and return it in result.
virtual void printSymbolName(StringRef symbolRef)
Print the given string as a symbol reference, i.e.
Attributes are known-constant values of operations.
Block represents an ordered list of Operations.
ValueTypeRange< BlockArgListType > getArgumentTypes()
Return a range containing the types of the arguments for this block.
unsigned getNumArguments()
OpListType & getOperations()
IntegerAttr getI32IntegerAttr(int32_t value)
IntegerAttr getIntegerAttr(Type type, int64_t value)
ArrayAttr getI32ArrayAttr(ArrayRef< int32_t > values)
FloatAttr getFloatAttr(Type type, double value)
FunctionType getFunctionType(TypeRange inputs, TypeRange results)
IntegerType getIntegerType(unsigned width)
BoolAttr getBoolAttr(bool value)
StringAttr getStringAttr(const Twine &bytes)
ArrayAttr getArrayAttr(ArrayRef< Attribute > value)
MLIRContext * getContext() const
Attr getAttr(Args &&...args)
Get or construct an instance of the attribute Attr with provided arguments.
static DenseElementsAttr get(ShapedType type, ArrayRef< Attribute > values)
Constructs a dense elements attribute from an array of element values.
static DenseFPElementsAttr get(const ShapedType &type, Arg &&arg)
Get an instance of a DenseFPElementsAttr with the given arguments.
Dialects are groups of MLIR operations, types and attributes, as well as behavior associated with the...
A symbol reference with a reference path containing a single element.
This class represents a diagnostic that is inflight and set to be reported.
This class defines the main interface for locations in MLIR and acts as a non-nullable wrapper around...
NamedAttrList is array of NamedAttributes that tracks whether it is sorted and does some basic work t...
void append(StringRef name, Attribute attr)
Add an attribute with the specified name.
The OpAsmParser has methods for interacting with the asm parser: parsing things from it,...
virtual ParseResult parseRegion(Region ®ion, ArrayRef< Argument > arguments={}, bool enableNameShadowing=false)=0
Parses a region.
virtual ParseResult resolveOperand(const UnresolvedOperand &operand, Type type, SmallVectorImpl< Value > &result)=0
Resolve an operand to an SSA value, emitting an error on failure.
virtual OptionalParseResult parseOptionalRegion(Region ®ion, ArrayRef< Argument > arguments={}, bool enableNameShadowing=false)=0
Parses a region if present.
virtual Operation * parseGenericOperation(Block *insertBlock, Block::iterator insertPt)=0
Parse an operation in its generic form.
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...
void printOperands(const ContainerType &container)
Print a comma separated list of operands.
virtual void printOptionalAttrDictWithKeyword(ArrayRef< NamedAttribute > attrs, ArrayRef< StringRef > elidedAttrs={})=0
If the specified operation has attributes, print out an attribute dictionary prefixed with 'attribute...
virtual void printOptionalAttrDict(ArrayRef< NamedAttribute > attrs, ArrayRef< StringRef > elidedAttrs={})=0
If the specified operation has attributes, print out an attribute dictionary with their values.
virtual void printGenericOp(Operation *op, bool printOpName=true)=0
Print the entire operation with the default generic assembly form.
virtual void printRegion(Region &blocks, bool printEntryBlockArgs=true, bool printBlockTerminators=true, bool printEmptyBlock=false)=0
Prints a region.
RAII guard to reset the insertion point of the builder when destroyed.
This class helps build Operations.
Block * createBlock(Region *parent, Region::iterator insertPt={}, TypeRange argTypes={}, ArrayRef< Location > locs={})
Add new block with 'argTypes' arguments and set the insertion point to the end of it.
void setInsertionPointToEnd(Block *block)
Sets the insertion point to the end of the specified block.
A trait to mark ops that can be enclosed/wrapped in a SpecConstantOperation op.
type_range getType() const
Operation is the basic unit of execution within MLIR.
Dialect * getDialect()
Return the dialect this operation is associated with, or nullptr if the associated dialect is not loa...
Value getOperand(unsigned idx)
bool hasTrait()
Returns true if the operation was registered with a particular trait, e.g.
AttrClass getAttrOfType(StringAttr name)
Attribute getAttr(StringAttr name)
Return the specified attribute if present, null otherwise.
ArrayRef< NamedAttribute > getAttrs()
Return all of the attributes on this operation.
OpResult getResult(unsigned idx)
Get the 'idx'th result of this operation.
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...
operand_type_range getOperandTypes()
result_type_range getResultTypes()
operand_range getOperands()
Returns an iterator on the underlying Value's.
InFlightDiagnostic emitOpError(const Twine &message={})
Emit an error with the op name prefixed, like "'dim' op " which is convenient for verifiers.
unsigned getNumResults()
Return the number of results held by this operation.
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.
void push_back(Block *block)
This class allows for representing and managing the symbol table used by operations with the 'SymbolT...
static StringRef getSymbolAttrName()
Return the name of the attribute used for symbol names.
static Operation * lookupNearestSymbolFrom(Operation *from, StringAttr symbol)
Returns the operation registered with the given symbol name within the closest parent operation of,...
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.
Type front()
Return first type in the range.
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.
A utility result that is used to signal how to proceed with an ongoing walk:
static WalkResult advance()
static ArrayType get(Type elementType, unsigned elementCount)
Type getElementType() const
static PointerType get(Type pointeeType, StorageClass storageClass)
unsigned getNumElements() const
Type getElementType(unsigned) const
An attribute that specifies the SPIR-V (version, capabilities, extensions) triple.
void addArgAndResultAttrs(Builder &builder, OperationState &result, ArrayRef< DictionaryAttr > argAttrs, ArrayRef< DictionaryAttr > resultAttrs, StringAttr argAttrsName, StringAttr resAttrsName)
Adds argument and result attributes, provided as argAttrs and resultAttrs arguments,...
void walk(Operation *op, function_ref< void(Region *)> callback, WalkOrder order)
Walk all of the regions, blocks, or operations nested under (and including) the given operation.
ArrayRef< NamedAttribute > getArgAttrs(FunctionOpInterface op, unsigned index)
Return all of the attributes for the argument at 'index'.
ParseResult parseFunctionSignatureWithArguments(OpAsmParser &parser, bool allowVariadic, SmallVectorImpl< OpAsmParser::Argument > &arguments, bool &isVariadic, SmallVectorImpl< Type > &resultTypes, SmallVectorImpl< DictionaryAttr > &resultAttrs)
Parses a function signature using parser.
void printFunctionAttributes(OpAsmPrinter &p, Operation *op, ArrayRef< StringRef > elided={})
Prints the list of function prefixed with the "attributes" keyword.
void printFunctionSignature(OpAsmPrinter &p, FunctionOpInterface op, ArrayRef< Type > argTypes, bool isVariadic, ArrayRef< Type > resultTypes)
Prints the signature of the function-like operation op.
Operation::operand_range getIndices(Operation *op)
Get the indices that the given load/store operation is operating on.
uint64_t getN(LevelType lt)
constexpr char kFnNameAttrName[]
constexpr char kSpecIdAttrName[]
LogicalResult verifyMemorySemantics(Operation *op, spirv::MemorySemantics memorySemantics)
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.
ParseResult parseEnumKeywordAttr(EnumClass &value, ParserType &parser, StringRef attrName=spirv::attributeName< EnumClass >())
Parses the next keyword in parser as an enumerant of the given EnumClass.
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...
AddressingModel getAddressingModel(TargetEnvAttr targetAttr, bool use64bitAddress)
Returns addressing model selected based on target environment.
FailureOr< ExecutionModel > getExecutionModel(TargetEnvAttr targetAttr)
Returns execution model selected based on target environment.
FailureOr< MemoryModel > getMemoryModel(TargetEnvAttr targetAttr)
Returns memory model selected based on target environment.
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)
Include the generated interface declarations.
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.
llvm::DenseMap< KeyT, ValueT, KeyInfoT, BucketT > DenseMap
llvm::function_ref< Fn > function_ref
This is the representation of an operand reference.
This represents an operation in an abstracted form, suitable for use with the builder APIs.
void addAttribute(StringRef name, Attribute attr)
Add an attribute with the specified name.
Region * addRegion()
Create a region that should be attached to the operation.