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";
139 auto descriptorSetName = llvm::convertToSnakeFromCamelCase(
140 stringifyDecoration(spirv::Decoration::DescriptorSet));
141 auto bindingName = llvm::convertToSnakeFromCamelCase(
142 stringifyDecoration(spirv::Decoration::Binding));
144 dyn_cast_or_null<IntegerAttr>(attrs.
get(descriptorSetName));
145 auto binding = dyn_cast_or_null<IntegerAttr>(attrs.
get(bindingName));
146 if (descriptorSet && binding) {
147 elidedAttrs.push_back(descriptorSetName);
148 elidedAttrs.push_back(bindingName);
149 printer <<
" bind(" << descriptorSet.getInt() <<
", " << binding.getInt()
154 auto builtInName = llvm::convertToSnakeFromCamelCase(
155 stringifyDecoration(spirv::Decoration::BuiltIn));
156 if (
auto builtin = dyn_cast_or_null<StringAttr>(attrs.
get(builtInName))) {
157 printer <<
" " << builtInName <<
"(\"" <<
builtin.getValue() <<
"\")";
158 elidedAttrs.push_back(builtInName);
176 auto fnType = dyn_cast<FunctionType>(type);
178 parser.
emitError(loc,
"expected function type");
183 result.addTypes(fnType.getResults());
194 assert(op->
getNumResults() == 1 &&
"op should have one result");
200 [&](
Type type) { return type != resultType; })) {
209 p <<
" : " << resultType;
212template <
typename BlockReadWriteOpTy>
216 if (
auto valVecTy = dyn_cast<VectorType>(valType))
217 valType = valVecTy.getElementType();
219 if (valType != cast<spirv::PointerType>(
ptr.getType()).getPointeeType()) {
220 return op.emitOpError(
"mismatch in result type and pointer type");
232 emitErrorFn(
"expected at least one index for spirv.CompositeExtract");
237 if (
auto cType = dyn_cast<spirv::CompositeType>(type)) {
238 if (cType.hasCompileTimeKnownNumElements() &&
240 static_cast<uint64_t
>(
index) >= cType.getNumElements())) {
241 emitErrorFn(
"index ") <<
index <<
" out of bounds for " << type;
244 type = cType.getElementType(
index);
246 emitErrorFn(
"cannot extract from non-composite type ")
247 << type <<
" with index " <<
index;
257 auto indicesArrayAttr = dyn_cast<ArrayAttr>(
indices);
258 if (!indicesArrayAttr) {
259 emitErrorFn(
"expected a 32-bit integer array attribute for 'indices'");
262 if (indicesArrayAttr.empty()) {
263 emitErrorFn(
"expected at least one index for spirv.CompositeExtract");
268 for (
auto indexAttr : indicesArrayAttr) {
269 auto indexIntAttr = dyn_cast<IntegerAttr>(indexAttr);
271 emitErrorFn(
"expected an 32-bit integer for index, but found '")
275 indexVals.push_back(indexIntAttr.getInt());
282 return ::mlir::emitError(loc, err);
295template <
typename ExtendedBinaryOp>
297 auto resultType = cast<spirv::StructType>(op.getType());
298 if (resultType.getNumElements() != 2)
299 return op.emitOpError(
"expected result struct type containing two members");
301 if (!llvm::all_equal({op.getOperand1().getType(), op.getOperand2().getType(),
302 resultType.getElementType(0),
303 resultType.getElementType(1)}))
304 return op.emitOpError(
305 "expected all operand types and struct member types are the same");
322 auto structType = dyn_cast<spirv::StructType>(resultType);
323 if (!structType || structType.getNumElements() != 2)
324 return parser.
emitError(loc,
"expected spirv.struct type with two members");
330 result.addTypes(resultType);
344 return op->
emitError(
"expected the same type for the first operand and "
345 "result, but provided ")
357 spirv::GlobalVariableOp var) {
358 build(builder, state, var.getType(), SymbolRefAttr::get(var));
361LogicalResult spirv::AddressOfOp::verify() {
362 auto varOp = dyn_cast_or_null<spirv::GlobalVariableOp>(
366 return emitOpError(
"expected spirv.GlobalVariable symbol");
368 if (getPointer().
getType() != varOp.getType()) {
370 "result type mismatch with the referenced global variable's type");
379LogicalResult spirv::CompositeConstructOp::verify() {
380 operand_range constituents = this->getConstituents();
395 if (coopElementType) {
396 if (constituents.size() != 1)
397 return emitOpError(
"has incorrect number of operands: expected ")
398 <<
"1, but provided " << constituents.size();
399 if (coopElementType != constituents.front().getType())
400 return emitOpError(
"operand type mismatch: expected operand type ")
401 << coopElementType <<
", but provided "
402 << constituents.front().getType();
407 auto cType = cast<spirv::CompositeType>(
getType());
408 if (constituents.size() == cType.getNumElements()) {
409 for (
auto index : llvm::seq<uint32_t>(0, constituents.size())) {
411 return emitOpError(
"operand type mismatch: expected operand type ")
412 << cType.getElementType(
index) <<
", but provided "
413 << constituents[
index].getType();
420 auto resultType = dyn_cast<VectorType>(cType);
423 "expected to return a vector or cooperative matrix when the number of "
424 "constituents is less than what the result needs");
427 for (
Value component : constituents) {
428 if (!isa<VectorType>(component.getType()) &&
429 !component.getType().isIntOrFloat())
430 return emitOpError(
"operand type mismatch: expected operand to have "
431 "a scalar or vector type, but provided ")
432 << component.getType();
434 Type elementType = component.getType();
435 if (
auto vectorType = dyn_cast<VectorType>(component.getType())) {
436 sizes.push_back(vectorType.getNumElements());
437 elementType = vectorType.getElementType();
442 if (elementType != resultType.getElementType())
443 return emitOpError(
"operand element type mismatch: expected to be ")
444 << resultType.getElementType() <<
", but provided " << elementType;
446 unsigned totalCount = llvm::sum_of(sizes);
447 if (totalCount != cType.getNumElements())
448 return emitOpError(
"has incorrect number of operands: expected ")
449 << cType.getNumElements() <<
", but provided " << totalCount;
466 build(builder, state, elementType, composite, indexAttr);
469ParseResult spirv::CompositeExtractOp::parse(
OpAsmParser &parser,
473 StringRef indicesAttrName =
474 spirv::CompositeExtractOp::getIndicesAttrName(
result.name);
491 result.addTypes(resultType);
495void spirv::CompositeExtractOp::print(
OpAsmPrinter &printer) {
496 printer <<
' ' << getComposite() <<
getIndices() <<
" : "
500LogicalResult spirv::CompositeExtractOp::verify() {
501 auto indicesArrayAttr = dyn_cast<ArrayAttr>(
getIndices());
508 return emitOpError(
"invalid result type: expected ")
509 << resultType <<
" but provided " <<
getType();
523 build(builder, state, composite.
getType(),
object, composite, indexAttr);
526ParseResult spirv::CompositeInsertOp::parse(
OpAsmParser &parser,
529 Type objectType, compositeType;
531 StringRef indicesAttrName =
532 spirv::CompositeInsertOp::getIndicesAttrName(
result.name);
545LogicalResult spirv::CompositeInsertOp::verify() {
546 auto indicesArrayAttr = dyn_cast<ArrayAttr>(
getIndices());
552 if (objectType != getObject().
getType()) {
553 return emitOpError(
"object operand type should be ")
554 << objectType <<
", but found " << getObject().getType();
558 return emitOpError(
"result type should be the same as "
559 "the composite type, but found ")
560 << getComposite().getType() <<
" vs " <<
getType();
566void spirv::CompositeInsertOp::print(
OpAsmPrinter &printer) {
567 printer <<
" " << getObject() <<
", " << getComposite() <<
getIndices()
568 <<
" : " << getObject().
getType() <<
" into "
569 << getComposite().getType();
576ParseResult spirv::ConstantOp::parse(
OpAsmParser &parser,
579 StringRef valueAttrName = spirv::ConstantOp::getValueAttrName(
result.name);
584 if (
auto typedAttr = dyn_cast<TypedAttr>(value))
585 type = typedAttr.getType();
586 if (isa<NoneType, TensorType>(type)) {
591 if (isa<TensorArmType>(type)) {
601 printer <<
' ' << getValue();
602 if (isa<spirv::ArrayType, spirv::StructType>(
getType()))
608 if (isa<spirv::CooperativeMatrixType>(opType)) {
609 auto denseAttr = dyn_cast<DenseElementsAttr>(value);
610 if (!denseAttr || !denseAttr.isSplat())
611 return op.emitOpError(
"expected a splat dense attribute for cooperative "
612 "matrix constant, but found ")
615 if (isa<IntegerAttr, FloatAttr>(value)) {
616 auto valueType = cast<TypedAttr>(value).getType();
617 if (valueType != opType)
618 return op.emitOpError(
"result type (")
619 << opType <<
") does not match value type (" << valueType <<
")";
622 if (isa<DenseTypedElementsAttr, SparseElementsAttr>(value)) {
623 auto valueType = cast<TypedAttr>(value).getType();
624 if (valueType == opType)
626 auto arrayType = dyn_cast<spirv::ArrayType>(opType);
627 auto shapedType = dyn_cast<ShapedType>(valueType);
629 return op.emitOpError(
"result or element type (")
630 << opType <<
") does not match value type (" << valueType
631 <<
"), must be the same or spirv.array";
633 int numElements = arrayType.getNumElements();
634 auto opElemType = arrayType.getElementType();
635 while (
auto t = dyn_cast<spirv::ArrayType>(opElemType)) {
636 numElements *= t.getNumElements();
637 opElemType = t.getElementType();
639 if (!opElemType.isIntOrFloat())
640 return op.emitOpError(
"only support nested array result type");
642 auto valueElemType = shapedType.getElementType();
643 if (valueElemType != opElemType) {
644 return op.emitOpError(
"result element type (")
645 << opElemType <<
") does not match value element type ("
646 << valueElemType <<
")";
649 if (numElements != shapedType.getNumElements()) {
650 return op.emitOpError(
"result number of elements (")
651 << numElements <<
") does not match value number of elements ("
652 << shapedType.getNumElements() <<
")";
656 if (
auto arrayAttr = dyn_cast<ArrayAttr>(value)) {
657 if (
auto structType = dyn_cast<spirv::StructType>(opType)) {
659 if (structType.isIdentified())
660 return op.emitOpError(
661 "cannot have an identified struct as a constant type");
662 if (arrayAttr.size() != structType.getNumElements())
663 return op.emitOpError(
"number of constituents (")
665 <<
") does not match number of struct members ("
666 << structType.getNumElements() <<
")";
667 for (
auto [idx, element] : llvm::enumerate(arrayAttr.getValue())) {
669 structType.getElementType(idx))))
674 auto arrayType = dyn_cast<spirv::ArrayType>(opType);
676 return op.emitOpError(
677 "must have spirv.array or spirv.struct result type for array value");
678 Type elemType = arrayType.getElementType();
679 for (
Attribute element : arrayAttr.getValue()) {
686 return op.emitOpError(
"cannot have attribute: ") << value;
689LogicalResult spirv::ConstantOp::verify() {
696bool spirv::ConstantOp::isBuildableWith(
Type type) {
698 if (!isa<spirv::SPIRVType>(type))
702 if (
auto structType = dyn_cast<spirv::StructType>(type))
703 return !structType.isIdentified();
704 return isa<spirv::ArrayType>(type);
710spirv::ConstantOp spirv::ConstantOp::getZero(
Type type,
Location loc,
712 if (
auto intType = dyn_cast<IntegerType>(type)) {
713 unsigned width = intType.getWidth();
715 return spirv::ConstantOp::create(builder, loc, type,
717 return spirv::ConstantOp::create(
718 builder, loc, type, builder.
getIntegerAttr(type, APInt(width, 0)));
720 if (
auto floatType = dyn_cast<FloatType>(type)) {
721 return spirv::ConstantOp::create(builder, loc, type,
724 if (
auto vectorType = dyn_cast<VectorType>(type)) {
725 Type elemType = vectorType.getElementType();
726 if (isa<IntegerType>(elemType)) {
727 return spirv::ConstantOp::create(
730 IntegerAttr::get(elemType, 0).getValue()));
732 if (isa<FloatType>(elemType)) {
733 return spirv::ConstantOp::create(
736 FloatAttr::get(elemType, 0.0).getValue()));
740 llvm_unreachable(
"unimplemented types for ConstantOp::getZero()");
743spirv::ConstantOp spirv::ConstantOp::getOne(
Type type,
Location loc,
745 if (
auto intType = dyn_cast<IntegerType>(type)) {
746 unsigned width = intType.getWidth();
748 return spirv::ConstantOp::create(builder, loc, type,
750 return spirv::ConstantOp::create(
751 builder, loc, type, builder.
getIntegerAttr(type, APInt(width, 1)));
753 if (
auto floatType = dyn_cast<FloatType>(type)) {
754 return spirv::ConstantOp::create(builder, loc, type,
757 if (
auto vectorType = dyn_cast<VectorType>(type)) {
758 Type elemType = vectorType.getElementType();
759 if (isa<IntegerType>(elemType)) {
760 return spirv::ConstantOp::create(
763 IntegerAttr::get(elemType, 1).getValue()));
765 if (isa<FloatType>(elemType)) {
766 return spirv::ConstantOp::create(
769 FloatAttr::get(elemType, 1.0).getValue()));
773 llvm_unreachable(
"unimplemented types for ConstantOp::getOne()");
776void mlir::spirv::ConstantOp::getAsmResultNames(
781 llvm::raw_svector_ostream specialName(specialNameBuffer);
782 specialName <<
"cst";
784 IntegerType intTy = dyn_cast<IntegerType>(type);
786 if (IntegerAttr intCst = dyn_cast<IntegerAttr>(getValue())) {
789 if (intTy.getWidth() == 1) {
790 return setNameFn(getResult(), (intCst.getInt() ?
"true" :
"false"));
793 if (intTy.isSignless()) {
794 specialName << intCst.getInt();
795 }
else if (intTy.isUnsigned()) {
796 specialName << intCst.getUInt();
798 specialName << intCst.getSInt();
802 if (intTy || isa<FloatType>(type)) {
803 specialName <<
'_' << type;
806 if (
auto vecType = dyn_cast<VectorType>(type)) {
807 specialName <<
"_vec_";
808 specialName << vecType.getDimSize(0);
810 Type elementType = vecType.getElementType();
812 if (isa<IntegerType>(elementType) || isa<FloatType>(elementType)) {
813 specialName <<
"x" << elementType;
817 setNameFn(getResult(), specialName.str());
820void mlir::spirv::AddressOfOp::getAsmResultNames(
823 llvm::raw_svector_ostream specialName(specialNameBuffer);
824 specialName << getVariable() <<
"_addr";
825 setNameFn(getResult(), specialName.str());
836 if (
auto typedAttr = dyn_cast<TypedAttr>(attr)) {
837 return typedAttr.getType();
840 if (
auto arrayAttr = dyn_cast<ArrayAttr>(attr)) {
847LogicalResult spirv::EXTConstantCompositeReplicateOp::verify() {
850 return emitError(
"unknown value attribute type");
852 auto compositeType = dyn_cast<spirv::CompositeType>(
getType());
854 return emitError(
"result type is not a composite type");
856 Type compositeElementType = compositeType.getElementType(0);
859 while (
auto type = dyn_cast<spirv::CompositeType>(compositeElementType)) {
860 compositeElementType = type.getElementType(0);
861 possibleTypes.push_back(compositeElementType);
864 if (!is_contained(possibleTypes, valueType)) {
865 return emitError(
"expected value attribute type ")
866 << interleaved(possibleTypes,
" or ") <<
", but got: " << valueType;
876LogicalResult spirv::ControlBarrierOp::verify() {
885 spirv::ExecutionModel executionModel,
886 spirv::FuncOp function,
888 build(builder, state,
889 spirv::ExecutionModelAttr::get(builder.
getContext(), executionModel),
890 SymbolRefAttr::get(function), builder.
getArrayAttr(interfaceVars));
893ParseResult spirv::EntryPointOp::parse(
OpAsmParser &parser,
895 spirv::ExecutionModel execModel;
908 FlatSymbolRefAttr var;
910 if (parser.parseAttribute(var, Type(),
"var_symbol", attrs))
912 interfaceVars.push_back(var);
917 result.addAttribute(spirv::EntryPointOp::getInterfaceAttrName(
result.name),
925 auto interfaceVars = getInterface().getValue();
926 if (!interfaceVars.empty())
927 printer <<
", " << llvm::interleaved(interfaceVars);
930LogicalResult spirv::EntryPointOp::verify() {
944struct ExecutionModeOperandSchema {
946 unsigned numOperands;
949ExecutionModeOperandSchema
950getExecutionModeOperandSchema(spirv::ExecutionMode mode) {
952 case spirv::ExecutionMode::Invocations:
953 case spirv::ExecutionMode::OutputVertices:
954 case spirv::ExecutionMode::VecTypeHint:
955 case spirv::ExecutionMode::SubgroupSize:
956 case spirv::ExecutionMode::SubgroupsPerWorkgroup:
957 case spirv::ExecutionMode::DenormPreserve:
958 case spirv::ExecutionMode::DenormFlushToZero:
959 case spirv::ExecutionMode::SignedZeroInfNanPreserve:
960 case spirv::ExecutionMode::RoundingModeRTE:
961 case spirv::ExecutionMode::RoundingModeRTZ:
962 case spirv::ExecutionMode::OutputPrimitivesEXT:
963 case spirv::ExecutionMode::SharedLocalMemorySizeINTEL:
964 case spirv::ExecutionMode::RoundingModeRTPINTEL:
965 case spirv::ExecutionMode::RoundingModeRTNINTEL:
966 case spirv::ExecutionMode::FloatingPointModeALTINTEL:
967 case spirv::ExecutionMode::FloatingPointModeIEEEINTEL:
968 case spirv::ExecutionMode::MaxWorkDimINTEL:
969 case spirv::ExecutionMode::NumSIMDWorkitemsINTEL:
970 case spirv::ExecutionMode::SchedulerTargetFmaxMhzINTEL:
971 case spirv::ExecutionMode::StreamingInterfaceINTEL:
972 case spirv::ExecutionMode::NamedBarrierCountINTEL:
974 case spirv::ExecutionMode::LocalSize:
975 case spirv::ExecutionMode::LocalSizeHint:
976 case spirv::ExecutionMode::MaxWorkgroupSizeINTEL:
978 case spirv::ExecutionMode::SubgroupsPerWorkgroupId:
980 case spirv::ExecutionMode::LocalSizeId:
981 case spirv::ExecutionMode::LocalSizeHintId:
994 spirv::FuncOp function,
995 spirv::ExecutionMode executionMode,
997 build(builder, state, SymbolRefAttr::get(function),
998 spirv::ExecutionModeAttr::get(builder.
getContext(), executionMode),
1002ParseResult spirv::ExecutionModeOp::parse(
OpAsmParser &parser,
1004 spirv::ExecutionMode execMode;
1019 values.push_back(cast<IntegerAttr>(value).getInt());
1021 StringRef valuesAttrName =
1022 spirv::ExecutionModeOp::getValuesAttrName(
result.name);
1023 result.addAttribute(valuesAttrName,
1028void spirv::ExecutionModeOp::print(
OpAsmPrinter &printer) {
1031 printer <<
" \"" << stringifyExecutionMode(getExecutionMode()) <<
"\"";
1033 if (!values.empty())
1034 printer <<
", " << llvm::interleaved(values.getAsValueRange<IntegerAttr>());
1037LogicalResult spirv::ExecutionModeOp::verify() {
1038 ExecutionModeOperandSchema schema =
1039 getExecutionModeOperandSchema(getExecutionMode());
1041 if (schema.isIdOperand)
1042 return emitOpError(
"expected ExecutionMode that takes extra operands "
1043 "that are not <id> operands, got: ")
1044 << stringifyExecutionMode(getExecutionMode());
1046 if (getValues().size() != schema.numOperands)
1047 return emitOpError(
"expected ")
1048 << schema.numOperands <<
" value operand(s), got "
1049 << getValues().size();
1058ParseResult spirv::ExecutionModeIdOp::parse(
OpAsmParser &parser,
1060 ExecutionMode execMode;
1069 FlatSymbolRefAttr attr;
1070 if (parser.parseAttribute(attr))
1072 values.push_back(attr);
1078 StringRef valuesAttrName = getValuesAttrName(
result.name);
1080 result.addAttribute(valuesAttrName, valuesAttr);
1084void spirv::ExecutionModeIdOp::print(
OpAsmPrinter &printer) {
1087 printer <<
" \"" << stringifyExecutionMode(getExecutionMode()) <<
"\" ";
1089 llvm::interleaveComma(
1090 getValues().getAsValueRange<FlatSymbolRefAttr>(), printer,
1094LogicalResult spirv::ExecutionModeIdOp::verify() {
1095 ExecutionModeOperandSchema schema =
1096 getExecutionModeOperandSchema(getExecutionMode());
1098 if (!schema.isIdOperand)
1099 return emitOpError(
"expected ExecutionMode that takes extra operands that "
1100 "are <id> operands, got: ")
1101 << stringifyExecutionMode(getExecutionMode());
1103 if (getValues().size() != schema.numOperands)
1104 return emitOpError(
"expected ")
1105 << schema.numOperands <<
" value operand(s), got "
1106 << getValues().size();
1109 auto valueSymbol = dyn_cast<FlatSymbolRefAttr>(value);
1111 return emitOpError(
"expected value operands to be symbol reference");
1113 (*this)->getParentOp(), valueSymbol);
1115 return emitOpError(
"cannot find symbol referenced by value operand: ")
1116 << valueSymbol.getValue();
1135 StringAttr nameAttr;
1141 bool isVariadic =
false;
1143 parser,
false, entryArgs, isVariadic, resultTypes,
1148 for (
auto &arg : entryArgs)
1149 argTypes.push_back(arg.type);
1151 result.addAttribute(getFunctionTypeAttrName(
result.name),
1152 TypeAttr::get(fnType));
1155 spirv::FunctionControl fnControl;
1164 assert(resultAttrs.size() == resultTypes.size());
1166 builder,
result, entryArgs, resultAttrs, getArgAttrsAttrName(
result.name),
1167 getResAttrsAttrName(
result.name));
1170 auto *body =
result.addRegion();
1179 if (StringAttr visibility = getSymVisibilityAttr())
1180 printer << visibility.getValue() <<
' ';
1182 auto fnType = getFunctionType();
1184 printer, *
this, fnType.getInputs(),
1185 false, fnType.getResults());
1186 printer <<
" \"" << spirv::stringifyFunctionControl(getFunctionControl())
1190 {spirv::attributeName<spirv::FunctionControl>(),
1191 getFunctionTypeAttrName(), getArgAttrsAttrName(), getResAttrsAttrName(),
1192 getFunctionControlAttrName(), getSymVisibilityAttrName()});
1195 Region &body = this->getBody();
1196 if (!body.empty()) {
1203LogicalResult spirv::FuncOp::verifyType() {
1204 FunctionType fnType = getFunctionType();
1205 if (fnType.getNumResults() > 1)
1206 return emitOpError(
"cannot have more than one result");
1208 auto hasDecorationAttr = [&](spirv::Decoration decoration,
1209 unsigned argIndex) {
1210 auto func = cast<FunctionOpInterface>(getOperation());
1211 for (
auto argAttr : cast<FunctionOpInterface>(
func).
getArgAttrs(argIndex)) {
1212 if (argAttr.getName() != spirv::DecorationAttr::name)
1214 if (
auto decAttr = dyn_cast<spirv::DecorationAttr>(argAttr.getValue()))
1215 return decAttr.getValue() == decoration;
1220 for (
unsigned i = 0, e = this->getNumArguments(); i != e; ++i) {
1221 Type param = fnType.getInputs()[i];
1222 auto inputPtrType = dyn_cast<spirv::PointerType>(param);
1226 auto pointeePtrType =
1227 dyn_cast<spirv::PointerType>(inputPtrType.getPointeeType());
1228 if (pointeePtrType) {
1234 if (pointeePtrType.getStorageClass() !=
1235 spirv::StorageClass::PhysicalStorageBuffer)
1238 bool hasAliasedPtr =
1239 hasDecorationAttr(spirv::Decoration::AliasedPointer, i);
1240 bool hasRestrictPtr =
1241 hasDecorationAttr(spirv::Decoration::RestrictPointer, i);
1242 if (!hasAliasedPtr && !hasRestrictPtr)
1243 return emitOpError()
1244 <<
"with a pointer points to a physical buffer pointer must "
1245 "be decorated either 'AliasedPointer' or 'RestrictPointer'";
1252 if (
auto pointeeArrayType =
1253 dyn_cast<spirv::ArrayType>(inputPtrType.getPointeeType())) {
1255 dyn_cast<spirv::PointerType>(pointeeArrayType.getElementType());
1257 pointeePtrType = inputPtrType;
1260 if (!pointeePtrType || pointeePtrType.getStorageClass() !=
1261 spirv::StorageClass::PhysicalStorageBuffer)
1264 bool hasAliased = hasDecorationAttr(spirv::Decoration::Aliased, i);
1265 bool hasRestrict = hasDecorationAttr(spirv::Decoration::Restrict, i);
1266 if (!hasAliased && !hasRestrict)
1267 return emitOpError() <<
"with physical buffer pointer must be decorated "
1268 "either 'Aliased' or 'Restrict'";
1274LogicalResult spirv::FuncOp::verifyBody() {
1275 FunctionType fnType = getFunctionType();
1276 if (!isExternal()) {
1277 Block &entryBlock = front();
1279 unsigned numArguments = this->getNumArguments();
1281 return emitOpError(
"entry block must have ")
1282 << numArguments <<
" arguments to match function signature";
1284 for (
auto [
index, fnArgType, blockArgType] :
1286 if (blockArgType != fnArgType) {
1287 return emitOpError(
"type of entry block argument #")
1288 <<
index <<
'(' << blockArgType
1289 <<
") must match the type of the corresponding argument in "
1290 <<
"function signature(" << fnArgType <<
')';
1296 if (
auto retOp = dyn_cast<spirv::ReturnOp>(op)) {
1297 if (fnType.getNumResults() != 0)
1298 return retOp.emitOpError(
"cannot be used in functions returning value");
1299 }
else if (
auto retOp = dyn_cast<spirv::ReturnValueOp>(op)) {
1300 if (fnType.getNumResults() != 1)
1301 return retOp.emitOpError(
1302 "returns 1 value but enclosing function requires ")
1303 << fnType.getNumResults() <<
" results";
1305 auto retOperandType = retOp.getValue().getType();
1306 auto fnResultType = fnType.getResult(0);
1307 if (retOperandType != fnResultType)
1308 return retOp.emitOpError(
" return value's type (")
1309 << retOperandType <<
") mismatch with function's result type ("
1310 << fnResultType <<
")";
1317 return failure(walkResult.wasInterrupted());
1321 StringRef name, FunctionType type,
1322 spirv::FunctionControl control,
1326 state.
addAttribute(getFunctionTypeAttrName(state.
name), TypeAttr::get(type));
1327 state.
addAttribute(spirv::attributeName<spirv::FunctionControl>(),
1328 builder.
getAttr<spirv::FunctionControlAttr>(control));
1337ParseResult spirv::GLFClampOp::parse(
OpAsmParser &parser,
1347ParseResult spirv::GLUClampOp::parse(
OpAsmParser &parser,
1357ParseResult spirv::GLSClampOp::parse(
OpAsmParser &parser,
1367ParseResult spirv::GLNClampOp::parse(
OpAsmParser &parser,
1377ParseResult spirv::GLSmoothStepOp::parse(
OpAsmParser &parser,
1399 Type type, StringRef name,
1400 unsigned descriptorSet,
unsigned binding) {
1401 build(builder, state, TypeAttr::get(type), builder.
getStringAttr(name));
1403 spirv::SPIRVDialect::getAttributeName(spirv::Decoration::DescriptorSet),
1406 spirv::SPIRVDialect::getAttributeName(spirv::Decoration::Binding),
1411 Type type, StringRef name,
1413 build(builder, state, TypeAttr::get(type), builder.
getStringAttr(name));
1415 spirv::SPIRVDialect::getAttributeName(spirv::Decoration::BuiltIn),
1419ParseResult spirv::GlobalVariableOp::parse(
OpAsmParser &parser,
1424 StringAttr nameAttr;
1425 StringRef initializerAttrName =
1426 spirv::GlobalVariableOp::getInitializerAttrName(
result.name);
1447 StringRef typeAttrName =
1448 spirv::GlobalVariableOp::getTypeAttrName(
result.name);
1453 if (!isa<spirv::PointerType>(type)) {
1454 return parser.
emitError(loc,
"expected spirv.ptr type");
1456 result.addAttribute(typeAttrName, TypeAttr::get(type));
1461void spirv::GlobalVariableOp::print(
OpAsmPrinter &printer) {
1463 spirv::attributeName<spirv::StorageClass>()};
1467 if (StringAttr visibility = getSymVisibilityAttr())
1468 printer << visibility.getValue() <<
' ';
1470 elidedAttrs.push_back(getSymNameAttrName());
1471 elidedAttrs.push_back(getSymVisibilityAttrName());
1473 StringRef initializerAttrName = this->getInitializerAttrName();
1475 if (
auto initializer = this->getInitializer()) {
1476 printer <<
" " << initializerAttrName <<
'(';
1479 elidedAttrs.push_back(initializerAttrName);
1482 StringRef typeAttrName = this->getTypeAttrName();
1483 elidedAttrs.push_back(typeAttrName);
1485 printer <<
" : " <<
getType();
1488LogicalResult spirv::GlobalVariableOp::verify() {
1489 if (!isa<spirv::PointerType>(
getType()))
1490 return emitOpError(
"result must be of a !spv.ptr type");
1496 auto storageClass = this->storageClass();
1497 if (storageClass == spirv::StorageClass::Generic ||
1498 storageClass == spirv::StorageClass::Function) {
1499 return emitOpError(
"storage class cannot be '")
1500 << stringifyStorageClass(storageClass) <<
"'";
1505 if (std::optional<spirv::LinkageAttributesAttr> linkage =
1506 getLinkageAttributes()) {
1507 if (linkage->getLinkageType().getValue() == spirv::LinkageType::Import &&
1510 "with Import linkage type must not have an initializer");
1516 (*this)->getParentOp(), init.getAttr());
1528 !isa<spirv::SpecConstantOp, spirv::SpecConstantCompositeOp>(initOp)) {
1529 return emitOpError(
"initializer must be result of a "
1530 "spirv.SpecConstant or "
1531 "spirv.SpecConstantCompositeOp op");
1535 Type pointeeType = cast<spirv::PointerType>(
getType()).getPointeeType();
1547LogicalResult spirv::INTELSubgroupBlockReadOp::verify() {
1558ParseResult spirv::INTELSubgroupBlockWriteOp::parse(
OpAsmParser &parser,
1561 spirv::StorageClass storageClass;
1572 if (
auto valVecTy = dyn_cast<VectorType>(elementType))
1582void spirv::INTELSubgroupBlockWriteOp::print(
OpAsmPrinter &printer) {
1583 printer <<
" " << getPtr() <<
", " << getValue() <<
" : "
1584 << getValue().getType();
1587LogicalResult spirv::INTELSubgroupBlockWriteOp::verify() {
1598LogicalResult spirv::IAddCarryOp::verify() {
1599 return ::verifyArithmeticExtendedBinaryOp(*
this);
1602ParseResult spirv::IAddCarryOp::parse(
OpAsmParser &parser,
1604 return ::parseArithmeticExtendedBinaryOp(parser,
result);
1615LogicalResult spirv::ISubBorrowOp::verify() {
1616 return ::verifyArithmeticExtendedBinaryOp(*
this);
1619ParseResult spirv::ISubBorrowOp::parse(
OpAsmParser &parser,
1621 return ::parseArithmeticExtendedBinaryOp(parser,
result);
1624void spirv::ISubBorrowOp::print(
OpAsmPrinter &printer) {
1632LogicalResult spirv::SMulExtendedOp::verify() {
1633 return ::verifyArithmeticExtendedBinaryOp(*
this);
1636ParseResult spirv::SMulExtendedOp::parse(
OpAsmParser &parser,
1638 return ::parseArithmeticExtendedBinaryOp(parser,
result);
1641void spirv::SMulExtendedOp::print(
OpAsmPrinter &printer) {
1649LogicalResult spirv::UMulExtendedOp::verify() {
1650 return ::verifyArithmeticExtendedBinaryOp(*
this);
1653ParseResult spirv::UMulExtendedOp::parse(
OpAsmParser &parser,
1655 return ::parseArithmeticExtendedBinaryOp(parser,
result);
1658void spirv::UMulExtendedOp::print(
OpAsmPrinter &printer) {
1666LogicalResult spirv::MemoryBarrierOp::verify() {
1674LogicalResult spirv::MemoryNamedBarrierOp::verify() {
1683 std::optional<StringRef> name) {
1693 spirv::AddressingModel addressingModel,
1694 spirv::MemoryModel memoryModel,
1695 std::optional<VerCapExtAttr> vceTriple,
1696 std::optional<StringRef> name) {
1699 builder.
getAttr<spirv::AddressingModelAttr>(addressingModel));
1701 builder.
getAttr<spirv::MemoryModelAttr>(memoryModel));
1705 state.
addAttribute(getVCETripleAttrName(), *vceTriple);
1711ParseResult spirv::ModuleOp::parse(
OpAsmParser &parser,
1718 StringAttr nameAttr;
1720 nameAttr, getSymNameAttrName(
result.name),
result.attributes);
1723 spirv::AddressingModel addrModel;
1724 spirv::MemoryModel memoryModel;
1734 spirv::ModuleOp::getVCETripleAttrName(),
1751 if (StringAttr visibility = getSymVisibilityAttr())
1752 printer <<
' ' << visibility.getValue();
1753 if (std::optional<StringRef> name = getName()) {
1762 auto addressingModelAttrName = spirv::attributeName<spirv::AddressingModel>();
1763 auto memoryModelAttrName = spirv::attributeName<spirv::MemoryModel>();
1764 elidedAttrs.assign({addressingModelAttrName, memoryModelAttrName,
1765 getSymNameAttrName(), getSymVisibilityAttrName()});
1767 if (std::optional<spirv::VerCapExtAttr> triple = getVceTriple()) {
1768 printer <<
" requires " << *triple;
1769 elidedAttrs.push_back(spirv::ModuleOp::getVCETripleAttrName());
1773 (*this)->getDiscardableAttrDictionary().getValue(), elidedAttrs);
1778LogicalResult spirv::ModuleOp::verifyRegions() {
1779 Dialect *dialect = (*this)->getDialect();
1784 for (
auto &op : *getBody()) {
1786 return op.
emitError(
"'spirv.module' can only contain spirv.* ops");
1791 if (
auto entryPointOp = dyn_cast<spirv::EntryPointOp>(op)) {
1792 auto funcOp = table.lookup<spirv::FuncOp>(entryPointOp.getFn());
1794 return entryPointOp.emitError(
"function '")
1795 << entryPointOp.getFn() <<
"' not found in 'spirv.module'";
1797 if (
auto interface = entryPointOp.getInterface()) {
1799 auto varSymRef = dyn_cast<FlatSymbolRefAttr>(varRef);
1801 return entryPointOp.emitError(
1802 "expected symbol reference for interface "
1803 "specification instead of '")
1807 table.lookup<spirv::GlobalVariableOp>(varSymRef.getValue());
1809 return entryPointOp.emitError(
"expected spirv.GlobalVariable "
1810 "symbol reference instead of'")
1811 << varSymRef <<
"'";
1816 auto key = std::pair<spirv::FuncOp, spirv::ExecutionModel>(
1817 funcOp, entryPointOp.getExecutionModel());
1818 if (!entryPoints.try_emplace(key, entryPointOp).second)
1819 return entryPointOp.emitError(
"duplicate of a previous EntryPointOp");
1820 }
else if (
auto funcOp = dyn_cast<spirv::FuncOp>(op)) {
1824 auto linkageAttr = funcOp.getLinkageAttributes();
1825 auto hasImportLinkage =
1826 linkageAttr && (linkageAttr.value().getLinkageType().getValue() ==
1827 spirv::LinkageType::Import);
1828 if (funcOp.isExternal() && !hasImportLinkage)
1830 "'spirv.module' cannot contain external functions "
1831 "without 'Import' linkage_attributes (LinkageAttributes)");
1834 for (
auto &block : funcOp)
1835 for (
auto &op : block) {
1838 "functions in 'spirv.module' can only contain spirv.* ops");
1850LogicalResult spirv::ReferenceOfOp::verify() {
1852 (*this)->getParentOp(), getSpecConstAttr());
1855 auto specConstOp = dyn_cast_or_null<spirv::SpecConstantOp>(specConstSym);
1857 constType = specConstOp.getDefaultValue().getType();
1859 auto specConstCompositeOp =
1860 dyn_cast_or_null<spirv::SpecConstantCompositeOp>(specConstSym);
1861 if (specConstCompositeOp)
1862 constType = specConstCompositeOp.getType();
1864 if (!specConstOp && !specConstCompositeOp)
1866 "expected spirv.SpecConstant or spirv.SpecConstantComposite symbol");
1868 if (getReference().
getType() != constType)
1869 return emitOpError(
"result type mismatch with the referenced "
1870 "specialization constant's type");
1879ParseResult spirv::SpecConstantOp::parse(
OpAsmParser &parser,
1883 StringAttr nameAttr;
1885 StringRef defaultValueAttrName =
1886 spirv::SpecConstantOp::getDefaultValueAttrName(
result.name);
1894 IntegerAttr specIdAttr;
1908void spirv::SpecConstantOp::print(
OpAsmPrinter &printer) {
1910 if (StringAttr visibility = getSymVisibilityAttr())
1911 printer << visibility.getValue() <<
' ';
1916 printer <<
" = " << getDefaultValue();
1919LogicalResult spirv::SpecConstantOp::verify() {
1922 if (specID.getValue().isNegative())
1923 return emitOpError(
"SpecId cannot be negative");
1925 auto value = getDefaultValue();
1926 if (isa<IntegerAttr, FloatAttr>(value)) {
1928 if (!isa<spirv::SPIRVType>(value.getType()))
1929 return emitOpError(
"default value bitwidth disallowed");
1933 "default value can only be a bool, integer, or float scalar");
1940LogicalResult spirv::VectorShuffleOp::verify() {
1941 VectorType resultType = cast<VectorType>(
getType());
1943 size_t numResultElements = resultType.getNumElements();
1944 if (numResultElements != getComponents().size())
1945 return emitOpError(
"result type element count (")
1946 << numResultElements
1947 <<
") mismatch with the number of component selectors ("
1948 << getComponents().size() <<
")";
1950 size_t totalSrcElements =
1951 cast<VectorType>(getVector1().
getType()).getNumElements() +
1952 cast<VectorType>(getVector2().
getType()).getNumElements();
1954 for (
const auto &selector : getComponents().getAsValueRange<IntegerAttr>()) {
1955 uint32_t
index = selector.getZExtValue();
1956 if (
index >= totalSrcElements &&
1957 index != std::numeric_limits<uint32_t>().
max())
1958 return emitOpError(
"component selector ")
1959 <<
index <<
" out of range: expected to be in [0, "
1960 << totalSrcElements <<
") or 0xffffffff";
1969ParseResult spirv::SpecConstantCompositeOp::parse(
OpAsmParser &parser,
1974 StringAttr compositeName;
1986 const char *attrName =
"spec_const";
1993 constituents.push_back(specConstRef);
1999 StringAttr compositeSpecConstituentsName =
2000 spirv::SpecConstantCompositeOp::getConstituentsAttrName(
result.name);
2001 result.addAttribute(compositeSpecConstituentsName,
2008 StringAttr typeAttrName =
2009 spirv::SpecConstantCompositeOp::getTypeAttrName(
result.name);
2010 result.addAttribute(typeAttrName, TypeAttr::get(type));
2015void spirv::SpecConstantCompositeOp::print(
OpAsmPrinter &printer) {
2017 if (StringAttr visibility = getSymVisibilityAttr())
2018 printer << visibility.getValue() <<
' ';
2020 printer <<
" (" << llvm::interleaved(this->getConstituents().getValue())
2024LogicalResult spirv::SpecConstantCompositeOp::verify() {
2025 auto cType = dyn_cast<spirv::CompositeType>(
getType());
2026 auto constituents = this->getConstituents().getValue();
2029 return emitError(
"result type must be a composite type, but provided ")
2032 if (isa<spirv::CooperativeMatrixType>(cType))
2033 return emitError(
"unsupported composite type ") << cType;
2034 if (constituents.size() != cType.getNumElements())
2035 return emitError(
"has incorrect number of operands: expected ")
2036 << cType.getNumElements() <<
", but provided "
2037 << constituents.size();
2039 for (
auto index : llvm::seq<uint32_t>(0, constituents.size())) {
2040 auto constituent = cast<FlatSymbolRefAttr>(constituents[
index]);
2043 (*this)->getParentOp(), constituent.getAttr());
2046 return emitError(
"unknown constituent symbol ") << constituent.getAttr();
2048 Type constituentType;
2049 if (
auto specConstOp = dyn_cast<spirv::SpecConstantOp>(constituentOp)) {
2050 constituentType = specConstOp.getDefaultValue().getType();
2051 }
else if (
auto specConstCompositeOp =
2052 dyn_cast<spirv::SpecConstantCompositeOp>(constituentOp)) {
2053 constituentType = specConstCompositeOp.getType();
2055 return emitError(
"unsupported constituent ")
2056 << constituent.getAttr()
2057 <<
": must reference a spirv.SpecConstant or "
2058 "spirv.SpecConstantComposite";
2061 if (constituentType != cType.getElementType(
index))
2062 return emitError(
"has incorrect types of operands: expected ")
2063 << cType.getElementType(
index) <<
", but provided "
2075spirv::EXTSpecConstantCompositeReplicateOp::parse(
OpAsmParser &parser,
2079 StringAttr compositeName;
2081 const char *attrName =
"spec_const";
2092 StringAttr compositeSpecConstituentName =
2093 spirv::EXTSpecConstantCompositeReplicateOp::getConstituentAttrName(
2095 result.addAttribute(compositeSpecConstituentName, specConstRef);
2097 StringAttr typeAttrName =
2098 spirv::EXTSpecConstantCompositeReplicateOp::getTypeAttrName(
result.name);
2099 result.addAttribute(typeAttrName, TypeAttr::get(type));
2104void spirv::EXTSpecConstantCompositeReplicateOp::print(
OpAsmPrinter &printer) {
2106 if (StringAttr visibility = getSymVisibilityAttr())
2107 printer << visibility.getValue() <<
' ';
2109 printer <<
" (" << this->getConstituent() <<
") : " <<
getType();
2112LogicalResult spirv::EXTSpecConstantCompositeReplicateOp::verify() {
2113 auto compositeType = dyn_cast<spirv::CompositeType>(
getType());
2115 return emitError(
"result type must be a composite type, but provided ")
2119 (*this)->getParentOp(), this->getConstituent());
2122 "splat spec constant reference defining constituent not found");
2124 auto constituentSpecConstOp = dyn_cast<spirv::SpecConstantOp>(constituentOp);
2125 if (!constituentSpecConstOp)
2126 return emitError(
"constituent is not a spec constant");
2128 Type constituentType = constituentSpecConstOp.getDefaultValue().getType();
2129 Type compositeElementType = compositeType.getElementType(0);
2130 if (constituentType != compositeElementType)
2131 return emitError(
"constituent has incorrect type: expected ")
2132 << compositeElementType <<
", but provided " << constituentType;
2141ParseResult spirv::SpecConstantOperationOp::parse(
OpAsmParser &parser,
2157 spirv::YieldOp::create(builder, wrappedOp->
getLoc(), wrappedOp->
getResult(0));
2168void spirv::SpecConstantOperationOp::print(
OpAsmPrinter &printer) {
2169 printer <<
" wraps ";
2173LogicalResult spirv::SpecConstantOperationOp::verifyRegions() {
2174 Block &block = getRegion().getBlocks().
front();
2177 return emitOpError(
"expected exactly 2 nested ops");
2182 return emitOpError(
"invalid enclosed op");
2185 if (!isa_and_present<spirv::ConstantOp, spirv::ReferenceOfOp,
2186 spirv::SpecConstantOperationOp>(
2187 operand.getDefiningOp()))
2189 "invalid operand, must be defined by a constant operation");
2198LogicalResult spirv::GLFrexpStructOp::verify() {
2200 dyn_cast<spirv::StructType>(getResult().
getType());
2203 return emitError(
"result type must be a struct type with two memebers");
2207 VectorType exponentVecTy = dyn_cast<VectorType>(exponentTy);
2208 IntegerType exponentIntTy = dyn_cast<IntegerType>(exponentTy);
2210 Type operandTy = getOperand().getType();
2211 VectorType operandVecTy = dyn_cast<VectorType>(operandTy);
2212 FloatType operandFTy = dyn_cast<FloatType>(operandTy);
2214 if (significandTy != operandTy)
2215 return emitError(
"member zero of the resulting struct type must be the "
2216 "same type as the operand");
2218 if (exponentVecTy) {
2219 IntegerType componentIntTy =
2220 dyn_cast<IntegerType>(exponentVecTy.getElementType());
2221 if (!componentIntTy || componentIntTy.getWidth() != 32)
2222 return emitError(
"member one of the resulting struct type must"
2223 "be a scalar or vector of 32 bit integer type");
2224 }
else if (!exponentIntTy || exponentIntTy.getWidth() != 32) {
2225 return emitError(
"member one of the resulting struct type "
2226 "must be a scalar or vector of 32 bit integer type");
2230 if (operandVecTy && exponentVecTy &&
2231 (exponentVecTy.getNumElements() == operandVecTy.getNumElements()))
2234 if (operandFTy && exponentIntTy)
2237 return emitError(
"member one of the resulting struct type must have the same "
2238 "number of components as the operand type");
2247 if (isa<FloatType>(floatType) != isa<IntegerType>(integerType))
2248 return op->
emitOpError(
"operands must both be scalars or vectors");
2251 if (
auto vectorType = dyn_cast<VectorType>(type))
2252 return vectorType.getNumElements();
2257 return op->
emitOpError(
"operands must have the same number of elements");
2262LogicalResult spirv::GLLdexpOp::verify() {
2271LogicalResult spirv::CLLdexpOp::verify() {
2280LogicalResult spirv::CLPownOp::verify() {
2289LogicalResult spirv::CLRootnOp::verify() {
2298LogicalResult spirv::ShiftLeftLogicalOp::verify() {
2306LogicalResult spirv::ShiftRightArithmeticOp::verify() {
2314LogicalResult spirv::ShiftRightLogicalOp::verify() {
2322LogicalResult spirv::VectorTimesScalarOp::verify() {
2324 return emitOpError(
"vector operand and result type mismatch");
2325 auto scalarType = cast<VectorType>(
getType()).getElementType();
2326 if (getScalar().
getType() != scalarType)
2327 return emitOpError(
"scalar operand and result element type match");
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...
Attribute get(StringAttr name) const
Return the specified attribute if present, null otherwise.
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
void populateInherentAttrs(Operation *op, NamedAttrList &attrs) const
Append the inherent attributes stored in the properties of op to attrs.
Operation is the basic unit of execution within MLIR.
Attribute getDiscardableAttr(StringRef name)
Access a discardable attribute by name, returns a null Attribute if the discardable attribute does no...
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.
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...
OperationName getName()
The name of an operation is the key identifier for it.
DictionaryAttr getDiscardableAttrDictionary()
Return all of the discardable attributes on this operation as a DictionaryAttr.
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 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.
ParseResult parseOptionalVisibilityKeyword(OpAsmParser &parser, NamedAttrList &attrs)
Parse an optional visibility attribute keyword (i.e., public, private, or nested) without quotes in a...
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.