MLIR 24.0.0git
SPIRVOps.cpp
Go to the documentation of this file.
1//===- SPIRVOps.cpp - MLIR SPIR-V operations ------------------------------===//
2//
3// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.
4// See https://llvm.org/LICENSE.txt for license information.
5// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
6//
7//===----------------------------------------------------------------------===//
8//
9// This file defines the operations in the SPIR-V dialect.
10//
11//===----------------------------------------------------------------------===//
12
14
15#include "SPIRVOpUtils.h"
16#include "SPIRVParsingUtils.h"
17
24#include "mlir/IR/Builders.h"
28#include "mlir/IR/Operation.h"
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"
38#include <cassert>
39#include <numeric>
40#include <optional>
41
42using namespace mlir;
43using namespace mlir::spirv::AttrNames;
44
45//===----------------------------------------------------------------------===//
46// Common utility functions
47//===----------------------------------------------------------------------===//
48
49LogicalResult spirv::extractValueFromConstOp(Operation *op, int32_t &value) {
50 auto constOp = dyn_cast_or_null<spirv::ConstantOp>(op);
51 if (!constOp) {
52 return failure();
53 }
54 auto valueAttr = constOp.getValue();
55 auto integerValueAttr = dyn_cast<IntegerAttr>(valueAttr);
56 if (!integerValueAttr) {
57 return failure();
58 }
59
60 if (integerValueAttr.getType().isSignlessInteger())
61 value = integerValueAttr.getInt();
62 else
63 value = integerValueAttr.getSInt();
64
65 return success();
66}
67
68LogicalResult
70 spirv::MemorySemantics memorySemantics) {
71 // According to the SPIR-V specification:
72 // "Despite being a mask and allowing multiple bits to be combined, it is
73 // invalid for more than one of these four bits to be set: Acquire, Release,
74 // AcquireRelease, or SequentiallyConsistent. Requesting both Acquire and
75 // Release semantics is done by setting the AcquireRelease bit, not by setting
76 // two bits."
77 auto atMostOneInSet = spirv::MemorySemantics::Acquire |
78 spirv::MemorySemantics::Release |
79 spirv::MemorySemantics::AcquireRelease |
80 spirv::MemorySemantics::SequentiallyConsistent;
81
82 auto bitCount =
83 llvm::popcount(static_cast<uint32_t>(memorySemantics & atMostOneInSet));
84 if (bitCount > 1) {
85 return op->emitError(
86 "expected at most one of these four memory constraints "
87 "to be set: `Acquire`, `Release`,"
88 "`AcquireRelease` or `SequentiallyConsistent`");
89 }
90 return success();
91}
92
94 Type pointeeType) {
95 // From SPV_KHR_physical_storage_buffer:
96 // > If an OpVariable's pointee type is a pointer (or array of pointers) in
97 // > PhysicalStorageBuffer storage class, then the variable must be decorated
98 // > with exactly one of AliasedPointer or RestrictPointer.
99 auto pointeePtrType = dyn_cast<spirv::PointerType>(pointeeType);
100 if (!pointeePtrType) {
101 if (auto pointeeArrayType = dyn_cast<spirv::ArrayType>(pointeeType)) {
102 pointeePtrType =
103 dyn_cast<spirv::PointerType>(pointeeArrayType.getElementType());
104 }
105 }
106
107 if (!pointeePtrType || pointeePtrType.getStorageClass() !=
108 spirv::StorageClass::PhysicalStorageBuffer)
109 return success();
110
111 auto getDecorationAttr = [op](spirv::Decoration decoration) {
112 return op->getDiscardableAttr(spirv::getDecorationString(decoration));
113 };
114
115 bool hasAliasedPtr =
116 getDecorationAttr(spirv::Decoration::AliasedPointer) != nullptr;
117 bool hasRestrictPtr =
118 getDecorationAttr(spirv::Decoration::RestrictPointer) != nullptr;
119
120 if (!hasAliasedPtr && !hasRestrictPtr)
121 return op->emitOpError()
122 << " with physical buffer pointer must be decorated "
123 "either 'AliasedPointer' or 'RestrictPointer'";
124
125 if (hasAliasedPtr && hasRestrictPtr)
126 return op->emitOpError()
127 << " with physical buffer pointer must have exactly one "
128 "aliasing decoration";
129
130 return success();
131}
132
134 SmallVectorImpl<StringRef> &elidedAttrs) {
135 NamedAttrList attrs(op->getDiscardableAttrDictionary().getValue());
136 op->getName().populateInherentAttrs(op, attrs);
137
138 // Print optional descriptor binding
139 auto descriptorSetName = llvm::convertToSnakeFromCamelCase(
140 stringifyDecoration(spirv::Decoration::DescriptorSet));
141 auto bindingName = llvm::convertToSnakeFromCamelCase(
142 stringifyDecoration(spirv::Decoration::Binding));
143 auto descriptorSet =
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()
150 << ")";
151 }
152
153 // Print BuiltIn attribute if present
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);
159 }
160
161 printer.printOptionalAttrDict(attrs, elidedAttrs);
162}
163
167 Type type;
168 // If the operand list is in-between parentheses, then we have a generic form.
169 // (see the fallback in `printOneResultOp`).
170 SMLoc loc = parser.getCurrentLocation();
171 if (!parser.parseOptionalLParen()) {
172 if (parser.parseOperandList(ops) || parser.parseRParen() ||
173 parser.parseOptionalAttrDict(result.attributes) ||
174 parser.parseColon() || parser.parseType(type))
175 return failure();
176 auto fnType = dyn_cast<FunctionType>(type);
177 if (!fnType) {
178 parser.emitError(loc, "expected function type");
179 return failure();
180 }
181 if (parser.resolveOperands(ops, fnType.getInputs(), loc, result.operands))
182 return failure();
183 result.addTypes(fnType.getResults());
184 return success();
185 }
186 return failure(parser.parseOperandList(ops) ||
187 parser.parseOptionalAttrDict(result.attributes) ||
188 parser.parseColonType(type) ||
189 parser.resolveOperands(ops, type, result.operands) ||
190 parser.addTypeToList(type, result.types));
191}
192
194 assert(op->getNumResults() == 1 && "op should have one result");
195
196 // If not all the operand and result types are the same, just use the
197 // generic assembly form to avoid omitting information in printing.
198 auto resultType = op->getResult(0).getType();
199 if (llvm::any_of(op->getOperandTypes(),
200 [&](Type type) { return type != resultType; })) {
201 p.printGenericOp(op, /*printOpName=*/false);
202 return;
203 }
204
205 p << ' ';
206 p.printOperands(op->getOperands());
208 // Now we can output only one type for all operands and the result.
209 p << " : " << resultType;
210}
211
212template <typename BlockReadWriteOpTy>
213static LogicalResult verifyBlockReadWritePtrAndValTypes(BlockReadWriteOpTy op,
214 Value ptr, Value val) {
215 auto valType = val.getType();
216 if (auto valVecTy = dyn_cast<VectorType>(valType))
217 valType = valVecTy.getElementType();
218
219 if (valType != cast<spirv::PointerType>(ptr.getType()).getPointeeType()) {
220 return op.emitOpError("mismatch in result type and pointer type");
221 }
222 return success();
223}
224
225/// Walks the given type hierarchy with the given indices, potentially down
226/// to component granularity, to select an element type. Returns null type and
227/// emits errors with the given loc on failure.
228static Type
230 function_ref<InFlightDiagnostic(StringRef)> emitErrorFn) {
231 if (indices.empty()) {
232 emitErrorFn("expected at least one index for spirv.CompositeExtract");
233 return nullptr;
234 }
235
236 for (auto index : indices) {
237 if (auto cType = dyn_cast<spirv::CompositeType>(type)) {
238 if (cType.hasCompileTimeKnownNumElements() &&
239 (index < 0 ||
240 static_cast<uint64_t>(index) >= cType.getNumElements())) {
241 emitErrorFn("index ") << index << " out of bounds for " << type;
242 return nullptr;
243 }
244 type = cType.getElementType(index);
245 } else {
246 emitErrorFn("cannot extract from non-composite type ")
247 << type << " with index " << index;
248 return nullptr;
249 }
250 }
251 return type;
252}
253
254static Type
256 function_ref<InFlightDiagnostic(StringRef)> emitErrorFn) {
257 auto indicesArrayAttr = dyn_cast<ArrayAttr>(indices);
258 if (!indicesArrayAttr) {
259 emitErrorFn("expected a 32-bit integer array attribute for 'indices'");
260 return nullptr;
261 }
262 if (indicesArrayAttr.empty()) {
263 emitErrorFn("expected at least one index for spirv.CompositeExtract");
264 return nullptr;
265 }
266
267 SmallVector<int32_t, 2> indexVals;
268 for (auto indexAttr : indicesArrayAttr) {
269 auto indexIntAttr = dyn_cast<IntegerAttr>(indexAttr);
270 if (!indexIntAttr) {
271 emitErrorFn("expected an 32-bit integer for index, but found '")
272 << indexAttr << "'";
273 return nullptr;
274 }
275 indexVals.push_back(indexIntAttr.getInt());
276 }
277 return getElementType(type, indexVals, emitErrorFn);
278}
279
281 auto errorFn = [&](StringRef err) -> InFlightDiagnostic {
282 return ::mlir::emitError(loc, err);
283 };
284 return getElementType(type, indices, errorFn);
285}
286
288 SMLoc loc) {
289 auto errorFn = [&](StringRef err) -> InFlightDiagnostic {
290 return parser.emitError(loc, err);
291 };
292 return getElementType(type, indices, errorFn);
293}
294
295template <typename ExtendedBinaryOp>
296static LogicalResult verifyArithmeticExtendedBinaryOp(ExtendedBinaryOp op) {
297 auto resultType = cast<spirv::StructType>(op.getType());
298 if (resultType.getNumElements() != 2)
299 return op.emitOpError("expected result struct type containing two members");
300
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");
306
307 return success();
308}
309
313 if (parser.parseOptionalAttrDict(result.attributes) ||
314 parser.parseOperandList(operands) || parser.parseColon())
315 return failure();
316
317 Type resultType;
318 SMLoc loc = parser.getCurrentLocation();
319 if (parser.parseType(resultType))
320 return failure();
321
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");
325
326 SmallVector<Type, 2> operandTypes(2, structType.getElementType(0));
327 if (parser.resolveOperands(operands, operandTypes, loc, result.operands))
328 return failure();
329
330 result.addTypes(resultType);
331 return success();
332}
333
335 OpAsmPrinter &printer) {
336 printer << ' ';
337 printer.printOptionalAttrDict(op->getDiscardableAttrDictionary().getValue());
338 printer.printOperands(op->getOperands());
339 printer << " : " << op->getResultTypes().front();
340}
341
342static LogicalResult verifyShiftOp(Operation *op) {
343 if (op->getOperand(0).getType() != op->getResult(0).getType()) {
344 return op->emitError("expected the same type for the first operand and "
345 "result, but provided ")
346 << op->getOperand(0).getType() << " and "
347 << op->getResult(0).getType();
348 }
349 return success();
350}
351
352//===----------------------------------------------------------------------===//
353// spirv.mlir.addressof
354//===----------------------------------------------------------------------===//
355
356void spirv::AddressOfOp::build(OpBuilder &builder, OperationState &state,
357 spirv::GlobalVariableOp var) {
358 build(builder, state, var.getType(), SymbolRefAttr::get(var));
359}
360
361LogicalResult spirv::AddressOfOp::verify() {
362 auto varOp = dyn_cast_or_null<spirv::GlobalVariableOp>(
363 SymbolTable::lookupNearestSymbolFrom((*this)->getParentOp(),
364 getVariableAttr()));
365 if (!varOp) {
366 return emitOpError("expected spirv.GlobalVariable symbol");
367 }
368 if (getPointer().getType() != varOp.getType()) {
369 return emitOpError(
370 "result type mismatch with the referenced global variable's type");
371 }
372 return success();
373}
374
375//===----------------------------------------------------------------------===//
376// spirv.CompositeConstruct
377//===----------------------------------------------------------------------===//
378
379LogicalResult spirv::CompositeConstructOp::verify() {
380 operand_range constituents = this->getConstituents();
381
382 // There are 4 cases with varying verification rules:
383 // 1. Cooperative Matrices (1 constituent)
384 // 2. Structs (1 constituent for each member)
385 // 3. Arrays (1 constituent for each array element)
386 // 4. Vectors (1 constituent (sub-)element for each vector element)
387
388 auto coopElementType = llvm::TypeSwitch<Type, Type>(getType())
389 .Case([](spirv::CooperativeMatrixType coopType) {
390 return coopType.getElementType();
391 })
392 .Default(nullptr);
393
394 // Case 1. -- matrices.
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();
403 return success();
404 }
405
406 // Case 2./3./4. -- number of constituents matches the number of elements.
407 auto cType = cast<spirv::CompositeType>(getType());
408 if (constituents.size() == cType.getNumElements()) {
409 for (auto index : llvm::seq<uint32_t>(0, constituents.size())) {
410 if (constituents[index].getType() != cType.getElementType(index)) {
411 return emitOpError("operand type mismatch: expected operand type ")
412 << cType.getElementType(index) << ", but provided "
413 << constituents[index].getType();
414 }
415 }
416 return success();
417 }
418
419 // Case 4. -- check that all constituents add up tp the expected vector type.
420 auto resultType = dyn_cast<VectorType>(cType);
421 if (!resultType)
422 return emitOpError(
423 "expected to return a vector or cooperative matrix when the number of "
424 "constituents is less than what the result needs");
425
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();
433
434 Type elementType = component.getType();
435 if (auto vectorType = dyn_cast<VectorType>(component.getType())) {
436 sizes.push_back(vectorType.getNumElements());
437 elementType = vectorType.getElementType();
438 } else {
439 sizes.push_back(1);
440 }
441
442 if (elementType != resultType.getElementType())
443 return emitOpError("operand element type mismatch: expected to be ")
444 << resultType.getElementType() << ", but provided " << elementType;
445 }
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;
450 return success();
451}
452
453//===----------------------------------------------------------------------===//
454// spirv.CompositeExtractOp
455//===----------------------------------------------------------------------===//
456
457void spirv::CompositeExtractOp::build(OpBuilder &builder, OperationState &state,
458 Value composite,
460 auto indexAttr = builder.getI32ArrayAttr(indices);
461 auto elementType =
462 getElementType(composite.getType(), indexAttr, state.location);
463 if (!elementType) {
464 return;
465 }
466 build(builder, state, elementType, composite, indexAttr);
467}
468
469ParseResult spirv::CompositeExtractOp::parse(OpAsmParser &parser,
471 OpAsmParser::UnresolvedOperand compositeInfo;
472 Attribute indicesAttr;
473 StringRef indicesAttrName =
474 spirv::CompositeExtractOp::getIndicesAttrName(result.name);
475 Type compositeType;
476 SMLoc attrLocation;
477
478 if (parser.parseOperand(compositeInfo) ||
479 parser.getCurrentLocation(&attrLocation) ||
480 parser.parseAttribute(indicesAttr, indicesAttrName, result.attributes) ||
481 parser.parseColonType(compositeType) ||
482 parser.resolveOperand(compositeInfo, compositeType, result.operands)) {
483 return failure();
484 }
485
486 Type resultType =
487 getElementType(compositeType, indicesAttr, parser, attrLocation);
488 if (!resultType) {
489 return failure();
490 }
491 result.addTypes(resultType);
492 return success();
493}
494
495void spirv::CompositeExtractOp::print(OpAsmPrinter &printer) {
496 printer << ' ' << getComposite() << getIndices() << " : "
497 << getComposite().getType();
498}
499
500LogicalResult spirv::CompositeExtractOp::verify() {
501 auto indicesArrayAttr = dyn_cast<ArrayAttr>(getIndices());
502 auto resultType =
503 getElementType(getComposite().getType(), indicesArrayAttr, getLoc());
504 if (!resultType)
505 return failure();
506
507 if (resultType != getType()) {
508 return emitOpError("invalid result type: expected ")
509 << resultType << " but provided " << getType();
510 }
511
512 return success();
513}
514
515//===----------------------------------------------------------------------===//
516// spirv.CompositeInsert
517//===----------------------------------------------------------------------===//
518
519void spirv::CompositeInsertOp::build(OpBuilder &builder, OperationState &state,
520 Value object, Value composite,
522 auto indexAttr = builder.getI32ArrayAttr(indices);
523 build(builder, state, composite.getType(), object, composite, indexAttr);
524}
525
526ParseResult spirv::CompositeInsertOp::parse(OpAsmParser &parser,
529 Type objectType, compositeType;
530 Attribute indicesAttr;
531 StringRef indicesAttrName =
532 spirv::CompositeInsertOp::getIndicesAttrName(result.name);
533 auto loc = parser.getCurrentLocation();
534
535 return failure(
536 parser.parseOperandList(operands, 2) ||
537 parser.parseAttribute(indicesAttr, indicesAttrName, result.attributes) ||
538 parser.parseColonType(objectType) ||
539 parser.parseKeywordType("into", compositeType) ||
540 parser.resolveOperands(operands, {objectType, compositeType}, loc,
541 result.operands) ||
542 parser.addTypesToList(compositeType, result.types));
543}
544
545LogicalResult spirv::CompositeInsertOp::verify() {
546 auto indicesArrayAttr = dyn_cast<ArrayAttr>(getIndices());
547 auto objectType =
548 getElementType(getComposite().getType(), indicesArrayAttr, getLoc());
549 if (!objectType)
550 return failure();
551
552 if (objectType != getObject().getType()) {
553 return emitOpError("object operand type should be ")
554 << objectType << ", but found " << getObject().getType();
555 }
556
557 if (getComposite().getType() != getType()) {
558 return emitOpError("result type should be the same as "
559 "the composite type, but found ")
560 << getComposite().getType() << " vs " << getType();
561 }
562
563 return success();
564}
565
566void spirv::CompositeInsertOp::print(OpAsmPrinter &printer) {
567 printer << " " << getObject() << ", " << getComposite() << getIndices()
568 << " : " << getObject().getType() << " into "
569 << getComposite().getType();
570}
571
572//===----------------------------------------------------------------------===//
573// spirv.Constant
574//===----------------------------------------------------------------------===//
575
576ParseResult spirv::ConstantOp::parse(OpAsmParser &parser,
578 Attribute value;
579 StringRef valueAttrName = spirv::ConstantOp::getValueAttrName(result.name);
580 if (parser.parseAttribute(value, valueAttrName, result.attributes))
581 return failure();
582
583 Type type = NoneType::get(parser.getContext());
584 if (auto typedAttr = dyn_cast<TypedAttr>(value))
585 type = typedAttr.getType();
586 if (isa<NoneType, TensorType>(type)) {
587 if (parser.parseColonType(type))
588 return failure();
589 }
590
591 if (isa<TensorArmType>(type)) {
592 if (parser.parseOptionalColon().succeeded())
593 if (parser.parseType(type))
594 return failure();
595 }
596
597 return parser.addTypeToList(type, result.types);
598}
599
600void spirv::ConstantOp::print(OpAsmPrinter &printer) {
601 printer << ' ' << getValue();
602 if (isa<spirv::ArrayType, spirv::StructType>(getType()))
603 printer << " : " << getType();
604}
605
606static LogicalResult verifyConstantType(spirv::ConstantOp op, Attribute value,
607 Type opType) {
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 ")
613 << denseAttr;
614 }
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 << ")";
620 return success();
621 }
622 if (isa<DenseTypedElementsAttr, SparseElementsAttr>(value)) {
623 auto valueType = cast<TypedAttr>(value).getType();
624 if (valueType == opType)
625 return success();
626 auto arrayType = dyn_cast<spirv::ArrayType>(opType);
627 auto shapedType = dyn_cast<ShapedType>(valueType);
628 if (!arrayType)
629 return op.emitOpError("result or element type (")
630 << opType << ") does not match value type (" << valueType
631 << "), must be the same or spirv.array";
632
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();
638 }
639 if (!opElemType.isIntOrFloat())
640 return op.emitOpError("only support nested array result type");
641
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 << ")";
647 }
648
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() << ")";
653 }
654 return success();
655 }
656 if (auto arrayAttr = dyn_cast<ArrayAttr>(value)) {
657 if (auto structType = dyn_cast<spirv::StructType>(opType)) {
658 // Identified (possibly recursive) structs are not supported as constants.
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 (")
664 << arrayAttr.size()
665 << ") does not match number of struct members ("
666 << structType.getNumElements() << ")";
667 for (auto [idx, element] : llvm::enumerate(arrayAttr.getValue())) {
668 if (failed(verifyConstantType(op, element,
669 structType.getElementType(idx))))
670 return failure();
671 }
672 return success();
673 }
674 auto arrayType = dyn_cast<spirv::ArrayType>(opType);
675 if (!arrayType)
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()) {
680 // Verify array elements recursively.
681 if (failed(verifyConstantType(op, element, elemType)))
682 return failure();
683 }
684 return success();
685 }
686 return op.emitOpError("cannot have attribute: ") << value;
687}
688
689LogicalResult spirv::ConstantOp::verify() {
690 // ODS already generates checks to make sure the result type is valid. We just
691 // need to additionally check that the value's attribute type is consistent
692 // with the result type.
693 return verifyConstantType(*this, getValueAttr(), getType());
694}
695
696bool spirv::ConstantOp::isBuildableWith(Type type) {
697 // Must be valid SPIR-V type first.
698 if (!isa<spirv::SPIRVType>(type))
699 return false;
700
701 if (isa<SPIRVDialect>(type.getDialect())) {
702 if (auto structType = dyn_cast<spirv::StructType>(type))
703 return !structType.isIdentified();
704 return isa<spirv::ArrayType>(type);
705 }
706
707 return true;
708}
709
710spirv::ConstantOp spirv::ConstantOp::getZero(Type type, Location loc,
711 OpBuilder &builder) {
712 if (auto intType = dyn_cast<IntegerType>(type)) {
713 unsigned width = intType.getWidth();
714 if (width == 1)
715 return spirv::ConstantOp::create(builder, loc, type,
716 builder.getBoolAttr(false));
717 return spirv::ConstantOp::create(
718 builder, loc, type, builder.getIntegerAttr(type, APInt(width, 0)));
719 }
720 if (auto floatType = dyn_cast<FloatType>(type)) {
721 return spirv::ConstantOp::create(builder, loc, type,
722 builder.getFloatAttr(floatType, 0.0));
723 }
724 if (auto vectorType = dyn_cast<VectorType>(type)) {
725 Type elemType = vectorType.getElementType();
726 if (isa<IntegerType>(elemType)) {
727 return spirv::ConstantOp::create(
728 builder, loc, type,
729 DenseElementsAttr::get(vectorType,
730 IntegerAttr::get(elemType, 0).getValue()));
731 }
732 if (isa<FloatType>(elemType)) {
733 return spirv::ConstantOp::create(
734 builder, loc, type,
735 DenseFPElementsAttr::get(vectorType,
736 FloatAttr::get(elemType, 0.0).getValue()));
737 }
738 }
739
740 llvm_unreachable("unimplemented types for ConstantOp::getZero()");
741}
742
743spirv::ConstantOp spirv::ConstantOp::getOne(Type type, Location loc,
744 OpBuilder &builder) {
745 if (auto intType = dyn_cast<IntegerType>(type)) {
746 unsigned width = intType.getWidth();
747 if (width == 1)
748 return spirv::ConstantOp::create(builder, loc, type,
749 builder.getBoolAttr(true));
750 return spirv::ConstantOp::create(
751 builder, loc, type, builder.getIntegerAttr(type, APInt(width, 1)));
752 }
753 if (auto floatType = dyn_cast<FloatType>(type)) {
754 return spirv::ConstantOp::create(builder, loc, type,
755 builder.getFloatAttr(floatType, 1.0));
756 }
757 if (auto vectorType = dyn_cast<VectorType>(type)) {
758 Type elemType = vectorType.getElementType();
759 if (isa<IntegerType>(elemType)) {
760 return spirv::ConstantOp::create(
761 builder, loc, type,
762 DenseElementsAttr::get(vectorType,
763 IntegerAttr::get(elemType, 1).getValue()));
764 }
765 if (isa<FloatType>(elemType)) {
766 return spirv::ConstantOp::create(
767 builder, loc, type,
768 DenseFPElementsAttr::get(vectorType,
769 FloatAttr::get(elemType, 1.0).getValue()));
770 }
771 }
772
773 llvm_unreachable("unimplemented types for ConstantOp::getOne()");
774}
775
776void mlir::spirv::ConstantOp::getAsmResultNames(
777 llvm::function_ref<void(mlir::Value, llvm::StringRef)> setNameFn) {
778 Type type = getType();
779
780 SmallString<32> specialNameBuffer;
781 llvm::raw_svector_ostream specialName(specialNameBuffer);
782 specialName << "cst";
783
784 IntegerType intTy = dyn_cast<IntegerType>(type);
785
786 if (IntegerAttr intCst = dyn_cast<IntegerAttr>(getValue())) {
787 assert(intTy);
788
789 if (intTy.getWidth() == 1) {
790 return setNameFn(getResult(), (intCst.getInt() ? "true" : "false"));
791 }
792
793 if (intTy.isSignless()) {
794 specialName << intCst.getInt();
795 } else if (intTy.isUnsigned()) {
796 specialName << intCst.getUInt();
797 } else {
798 specialName << intCst.getSInt();
799 }
800 }
801
802 if (intTy || isa<FloatType>(type)) {
803 specialName << '_' << type;
804 }
805
806 if (auto vecType = dyn_cast<VectorType>(type)) {
807 specialName << "_vec_";
808 specialName << vecType.getDimSize(0);
809
810 Type elementType = vecType.getElementType();
811
812 if (isa<IntegerType>(elementType) || isa<FloatType>(elementType)) {
813 specialName << "x" << elementType;
814 }
815 }
816
817 setNameFn(getResult(), specialName.str());
818}
819
820void mlir::spirv::AddressOfOp::getAsmResultNames(
821 llvm::function_ref<void(mlir::Value, llvm::StringRef)> setNameFn) {
822 SmallString<32> specialNameBuffer;
823 llvm::raw_svector_ostream specialName(specialNameBuffer);
824 specialName << getVariable() << "_addr";
825 setNameFn(getResult(), specialName.str());
826}
827
828//===----------------------------------------------------------------------===//
829// spirv.EXTConstantCompositeReplicate
830//===----------------------------------------------------------------------===//
831
832// Returns type of attribute. In case of a TypedAttr this will simply return
833// the type. But for an ArrayAttr which is untyped and can be multidimensional
834// it creates the ArrayType recursively.
836 if (auto typedAttr = dyn_cast<TypedAttr>(attr)) {
837 return typedAttr.getType();
838 }
839
840 if (auto arrayAttr = dyn_cast<ArrayAttr>(attr)) {
841 return spirv::ArrayType::get(getValueType(arrayAttr[0]), arrayAttr.size());
842 }
843
844 return nullptr;
845}
846
847LogicalResult spirv::EXTConstantCompositeReplicateOp::verify() {
848 Type valueType = getValueType(getValue());
849 if (!valueType)
850 return emitError("unknown value attribute type");
851
852 auto compositeType = dyn_cast<spirv::CompositeType>(getType());
853 if (!compositeType)
854 return emitError("result type is not a composite type");
855
856 Type compositeElementType = compositeType.getElementType(0);
857
858 SmallVector<Type, 3> possibleTypes = {compositeElementType};
859 while (auto type = dyn_cast<spirv::CompositeType>(compositeElementType)) {
860 compositeElementType = type.getElementType(0);
861 possibleTypes.push_back(compositeElementType);
862 }
863
864 if (!is_contained(possibleTypes, valueType)) {
865 return emitError("expected value attribute type ")
866 << interleaved(possibleTypes, " or ") << ", but got: " << valueType;
867 }
868
869 return success();
870}
871
872//===----------------------------------------------------------------------===//
873// spirv.ControlBarrierOp
874//===----------------------------------------------------------------------===//
875
876LogicalResult spirv::ControlBarrierOp::verify() {
877 return verifyMemorySemantics(getOperation(), getMemorySemantics());
878}
879
880//===----------------------------------------------------------------------===//
881// spirv.EntryPoint
882//===----------------------------------------------------------------------===//
883
884void spirv::EntryPointOp::build(OpBuilder &builder, OperationState &state,
885 spirv::ExecutionModel executionModel,
886 spirv::FuncOp function,
887 ArrayRef<Attribute> interfaceVars) {
888 build(builder, state,
889 spirv::ExecutionModelAttr::get(builder.getContext(), executionModel),
890 SymbolRefAttr::get(function), builder.getArrayAttr(interfaceVars));
891}
892
893ParseResult spirv::EntryPointOp::parse(OpAsmParser &parser,
895 spirv::ExecutionModel execModel;
896 SmallVector<Attribute, 4> interfaceVars;
897
899 if (parseEnumStrAttr<spirv::ExecutionModelAttr>(execModel, parser, result) ||
900 parser.parseAttribute(fn, Type(), kFnNameAttrName, result.attributes)) {
901 return failure();
902 }
903
904 if (!parser.parseOptionalComma()) {
905 // Parse the interface variables
906 if (parser.parseCommaSeparatedList([&]() -> ParseResult {
907 // The name of the interface variable attribute isnt important
908 FlatSymbolRefAttr var;
909 NamedAttrList attrs;
910 if (parser.parseAttribute(var, Type(), "var_symbol", attrs))
911 return failure();
912 interfaceVars.push_back(var);
913 return success();
914 }))
915 return failure();
916 }
917 result.addAttribute(spirv::EntryPointOp::getInterfaceAttrName(result.name),
918 parser.getBuilder().getArrayAttr(interfaceVars));
919 return success();
920}
921
922void spirv::EntryPointOp::print(OpAsmPrinter &printer) {
923 printer << " \"" << stringifyExecutionModel(getExecutionModel()) << "\" ";
924 printer.printSymbolName(getFn());
925 auto interfaceVars = getInterface().getValue();
926 if (!interfaceVars.empty())
927 printer << ", " << llvm::interleaved(interfaceVars);
928}
929
930LogicalResult spirv::EntryPointOp::verify() {
931 // Checks for fn and interface symbol reference are done in spirv::ModuleOp
932 // verification.
933 return success();
934}
935
936//===----------------------------------------------------------------------===//
937// spirv.ExecutionMode / spirv.ExecutionModeId
938//===----------------------------------------------------------------------===//
939
940namespace {
941// Describes the extra operands a SPIR-V ExecutionMode expects: whether they
942// are <id> operands (only valid on spirv.ExecutionModeId) or literal integers
943// (only valid on spirv.ExecutionMode), and how many of them are required.
944struct ExecutionModeOperandSchema {
945 bool isIdOperand;
946 unsigned numOperands;
947};
948
949ExecutionModeOperandSchema
950getExecutionModeOperandSchema(spirv::ExecutionMode mode) {
951 switch (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:
973 return {/*isIdOperand=*/false, /*numOperands=*/1};
974 case spirv::ExecutionMode::LocalSize:
975 case spirv::ExecutionMode::LocalSizeHint:
976 case spirv::ExecutionMode::MaxWorkgroupSizeINTEL:
977 return {/*isIdOperand=*/false, /*numOperands=*/3};
978 case spirv::ExecutionMode::SubgroupsPerWorkgroupId:
979 return {/*isIdOperand=*/true, /*numOperands=*/1};
980 case spirv::ExecutionMode::LocalSizeId:
981 case spirv::ExecutionMode::LocalSizeHintId:
982 return {/*isIdOperand=*/true, /*numOperands=*/3};
983 default:
984 return {/*isIdOperand=*/false, /*numOperands=*/0};
985 }
986}
987} // namespace
988
989//===----------------------------------------------------------------------===//
990// spirv.ExecutionMode
991//===----------------------------------------------------------------------===//
992
993void spirv::ExecutionModeOp::build(OpBuilder &builder, OperationState &state,
994 spirv::FuncOp function,
995 spirv::ExecutionMode executionMode,
996 ArrayRef<int32_t> params) {
997 build(builder, state, SymbolRefAttr::get(function),
998 spirv::ExecutionModeAttr::get(builder.getContext(), executionMode),
999 builder.getI32ArrayAttr(params));
1000}
1001
1002ParseResult spirv::ExecutionModeOp::parse(OpAsmParser &parser,
1004 spirv::ExecutionMode execMode;
1005 Attribute fn;
1006 if (parser.parseAttribute(fn, kFnNameAttrName, result.attributes) ||
1008 return failure();
1009 }
1010
1012 Type i32Type = parser.getBuilder().getIntegerType(32);
1013 while (!parser.parseOptionalComma()) {
1014 NamedAttrList attr;
1015 Attribute value;
1016 if (parser.parseAttribute(value, i32Type, "value", attr)) {
1017 return failure();
1018 }
1019 values.push_back(cast<IntegerAttr>(value).getInt());
1020 }
1021 StringRef valuesAttrName =
1022 spirv::ExecutionModeOp::getValuesAttrName(result.name);
1023 result.addAttribute(valuesAttrName,
1024 parser.getBuilder().getI32ArrayAttr(values));
1025 return success();
1026}
1027
1028void spirv::ExecutionModeOp::print(OpAsmPrinter &printer) {
1029 printer << " ";
1030 printer.printSymbolName(getFn());
1031 printer << " \"" << stringifyExecutionMode(getExecutionMode()) << "\"";
1032 ArrayAttr values = this->getValues();
1033 if (!values.empty())
1034 printer << ", " << llvm::interleaved(values.getAsValueRange<IntegerAttr>());
1035}
1036
1037LogicalResult spirv::ExecutionModeOp::verify() {
1038 ExecutionModeOperandSchema schema =
1039 getExecutionModeOperandSchema(getExecutionMode());
1040
1041 if (schema.isIdOperand)
1042 return emitOpError("expected ExecutionMode that takes extra operands "
1043 "that are not <id> operands, got: ")
1044 << stringifyExecutionMode(getExecutionMode());
1045
1046 if (getValues().size() != schema.numOperands)
1047 return emitOpError("expected ")
1048 << schema.numOperands << " value operand(s), got "
1049 << getValues().size();
1050
1051 return success();
1052}
1053
1054//===----------------------------------------------------------------------===//
1055// spirv.ExecutionModeId
1056//===----------------------------------------------------------------------===//
1057
1058ParseResult spirv::ExecutionModeIdOp::parse(OpAsmParser &parser,
1060 ExecutionMode execMode;
1061 if (Attribute fn;
1062 parser.parseAttribute(fn, kFnNameAttrName, result.attributes) ||
1063 parseEnumStrAttr<ExecutionModeAttr>(execMode, parser, result)) {
1064 return failure();
1065 }
1066
1068 if (parser.parseCommaSeparatedList([&]() -> ParseResult {
1069 FlatSymbolRefAttr attr;
1070 if (parser.parseAttribute(attr))
1071 return failure();
1072 values.push_back(attr);
1073 return success();
1074 })) {
1075 return failure();
1076 }
1077
1078 StringRef valuesAttrName = getValuesAttrName(result.name);
1079 ArrayAttr valuesAttr = parser.getBuilder().getArrayAttr(values);
1080 result.addAttribute(valuesAttrName, valuesAttr);
1081 return success();
1082}
1083
1084void spirv::ExecutionModeIdOp::print(OpAsmPrinter &printer) {
1085 printer << " ";
1086 printer.printSymbolName(getFn());
1087 printer << " \"" << stringifyExecutionMode(getExecutionMode()) << "\" ";
1088
1089 llvm::interleaveComma(
1090 getValues().getAsValueRange<FlatSymbolRefAttr>(), printer,
1091 [&](StringRef value) { printer.printSymbolName(value); });
1092}
1093
1094LogicalResult spirv::ExecutionModeIdOp::verify() {
1095 ExecutionModeOperandSchema schema =
1096 getExecutionModeOperandSchema(getExecutionMode());
1097
1098 if (!schema.isIdOperand)
1099 return emitOpError("expected ExecutionMode that takes extra operands that "
1100 "are <id> operands, got: ")
1101 << stringifyExecutionMode(getExecutionMode());
1102
1103 if (getValues().size() != schema.numOperands)
1104 return emitOpError("expected ")
1105 << schema.numOperands << " value operand(s), got "
1106 << getValues().size();
1107
1108 for (Attribute value : getValues()) {
1109 auto valueSymbol = dyn_cast<FlatSymbolRefAttr>(value);
1110 if (!valueSymbol)
1111 return emitOpError("expected value operands to be symbol reference");
1113 (*this)->getParentOp(), valueSymbol);
1114 if (!valueOp)
1115 return emitOpError("cannot find symbol referenced by value operand: ")
1116 << valueSymbol.getValue();
1117 }
1118
1119 return success();
1120}
1121
1122//===----------------------------------------------------------------------===//
1123// spirv.func
1124//===----------------------------------------------------------------------===//
1125
1126ParseResult spirv::FuncOp::parse(OpAsmParser &parser, OperationState &result) {
1128 SmallVector<DictionaryAttr> resultAttrs;
1129 SmallVector<Type> resultTypes;
1130 auto &builder = parser.getBuilder();
1131
1132 // Parse the name as a symbol.
1133 StringAttr nameAttr;
1134 if (parser.parseSymbolName(nameAttr, SymbolTable::getSymbolAttrName(),
1135 result.attributes))
1136 return failure();
1137
1138 // Parse the function signature.
1139 bool isVariadic = false;
1141 parser, /*allowVariadic=*/false, entryArgs, isVariadic, resultTypes,
1142 resultAttrs))
1143 return failure();
1144
1145 SmallVector<Type> argTypes;
1146 for (auto &arg : entryArgs)
1147 argTypes.push_back(arg.type);
1148 auto fnType = builder.getFunctionType(argTypes, resultTypes);
1149 result.addAttribute(getFunctionTypeAttrName(result.name),
1150 TypeAttr::get(fnType));
1151
1152 // Parse the optional function control keyword.
1153 spirv::FunctionControl fnControl;
1155 return failure();
1156
1157 // If additional attributes are present, parse them.
1158 if (parser.parseOptionalAttrDictWithKeyword(result.attributes))
1159 return failure();
1160
1161 // Add the attributes to the function arguments.
1162 assert(resultAttrs.size() == resultTypes.size());
1164 builder, result, entryArgs, resultAttrs, getArgAttrsAttrName(result.name),
1165 getResAttrsAttrName(result.name));
1166
1167 // Parse the optional function body.
1168 auto *body = result.addRegion();
1169 OptionalParseResult parseResult =
1170 parser.parseOptionalRegion(*body, entryArgs);
1171 return failure(parseResult.has_value() && failed(*parseResult));
1172}
1173
1174void spirv::FuncOp::print(OpAsmPrinter &printer) {
1175 // Print function name, signature, and control.
1176 printer << " ";
1177 printer.printSymbolName(getSymName());
1178 auto fnType = getFunctionType();
1180 printer, *this, fnType.getInputs(),
1181 /*isVariadic=*/false, fnType.getResults());
1182 printer << " \"" << spirv::stringifyFunctionControl(getFunctionControl())
1183 << "\"";
1185 printer, *this,
1186 {spirv::attributeName<spirv::FunctionControl>(),
1187 getFunctionTypeAttrName(), getArgAttrsAttrName(), getResAttrsAttrName(),
1188 getFunctionControlAttrName()});
1189
1190 // Print the body if this is not an external function.
1191 Region &body = this->getBody();
1192 if (!body.empty()) {
1193 printer << ' ';
1194 printer.printRegion(body, /*printEntryBlockArgs=*/false,
1195 /*printBlockTerminators=*/true);
1196 }
1197}
1198
1199LogicalResult spirv::FuncOp::verifyType() {
1200 FunctionType fnType = getFunctionType();
1201 if (fnType.getNumResults() > 1)
1202 return emitOpError("cannot have more than one result");
1203
1204 auto hasDecorationAttr = [&](spirv::Decoration decoration,
1205 unsigned argIndex) {
1206 auto func = cast<FunctionOpInterface>(getOperation());
1207 for (auto argAttr : cast<FunctionOpInterface>(func).getArgAttrs(argIndex)) {
1208 if (argAttr.getName() != spirv::DecorationAttr::name)
1209 continue;
1210 if (auto decAttr = dyn_cast<spirv::DecorationAttr>(argAttr.getValue()))
1211 return decAttr.getValue() == decoration;
1212 }
1213 return false;
1214 };
1215
1216 for (unsigned i = 0, e = this->getNumArguments(); i != e; ++i) {
1217 Type param = fnType.getInputs()[i];
1218 auto inputPtrType = dyn_cast<spirv::PointerType>(param);
1219 if (!inputPtrType)
1220 continue;
1221
1222 auto pointeePtrType =
1223 dyn_cast<spirv::PointerType>(inputPtrType.getPointeeType());
1224 if (pointeePtrType) {
1225 // SPIR-V spec, from SPV_KHR_physical_storage_buffer:
1226 // > If an OpFunctionParameter is a pointer (or contains a pointer)
1227 // > and the type it points to is a pointer in the PhysicalStorageBuffer
1228 // > storage class, the function parameter must be decorated with exactly
1229 // > one of AliasedPointer or RestrictPointer.
1230 if (pointeePtrType.getStorageClass() !=
1231 spirv::StorageClass::PhysicalStorageBuffer)
1232 continue;
1233
1234 bool hasAliasedPtr =
1235 hasDecorationAttr(spirv::Decoration::AliasedPointer, i);
1236 bool hasRestrictPtr =
1237 hasDecorationAttr(spirv::Decoration::RestrictPointer, i);
1238 if (!hasAliasedPtr && !hasRestrictPtr)
1239 return emitOpError()
1240 << "with a pointer points to a physical buffer pointer must "
1241 "be decorated either 'AliasedPointer' or 'RestrictPointer'";
1242 continue;
1243 }
1244 // SPIR-V spec, from SPV_KHR_physical_storage_buffer:
1245 // > If an OpFunctionParameter is a pointer (or contains a pointer) in
1246 // > the PhysicalStorageBuffer storage class, the function parameter must
1247 // > be decorated with exactly one of Aliased or Restrict.
1248 if (auto pointeeArrayType =
1249 dyn_cast<spirv::ArrayType>(inputPtrType.getPointeeType())) {
1250 pointeePtrType =
1251 dyn_cast<spirv::PointerType>(pointeeArrayType.getElementType());
1252 } else {
1253 pointeePtrType = inputPtrType;
1254 }
1255
1256 if (!pointeePtrType || pointeePtrType.getStorageClass() !=
1257 spirv::StorageClass::PhysicalStorageBuffer)
1258 continue;
1259
1260 bool hasAliased = hasDecorationAttr(spirv::Decoration::Aliased, i);
1261 bool hasRestrict = hasDecorationAttr(spirv::Decoration::Restrict, i);
1262 if (!hasAliased && !hasRestrict)
1263 return emitOpError() << "with physical buffer pointer must be decorated "
1264 "either 'Aliased' or 'Restrict'";
1265 }
1266
1267 return success();
1268}
1269
1270LogicalResult spirv::FuncOp::verifyBody() {
1271 FunctionType fnType = getFunctionType();
1272 if (!isExternal()) {
1273 Block &entryBlock = front();
1274
1275 unsigned numArguments = this->getNumArguments();
1276 if (entryBlock.getNumArguments() != numArguments)
1277 return emitOpError("entry block must have ")
1278 << numArguments << " arguments to match function signature";
1279
1280 for (auto [index, fnArgType, blockArgType] :
1281 llvm::enumerate(getArgumentTypes(), entryBlock.getArgumentTypes())) {
1282 if (blockArgType != fnArgType) {
1283 return emitOpError("type of entry block argument #")
1284 << index << '(' << blockArgType
1285 << ") must match the type of the corresponding argument in "
1286 << "function signature(" << fnArgType << ')';
1287 }
1288 }
1289 }
1290
1291 auto walkResult = walk([fnType](Operation *op) -> WalkResult {
1292 if (auto retOp = dyn_cast<spirv::ReturnOp>(op)) {
1293 if (fnType.getNumResults() != 0)
1294 return retOp.emitOpError("cannot be used in functions returning value");
1295 } else if (auto retOp = dyn_cast<spirv::ReturnValueOp>(op)) {
1296 if (fnType.getNumResults() != 1)
1297 return retOp.emitOpError(
1298 "returns 1 value but enclosing function requires ")
1299 << fnType.getNumResults() << " results";
1300
1301 auto retOperandType = retOp.getValue().getType();
1302 auto fnResultType = fnType.getResult(0);
1303 if (retOperandType != fnResultType)
1304 return retOp.emitOpError(" return value's type (")
1305 << retOperandType << ") mismatch with function's result type ("
1306 << fnResultType << ")";
1307 }
1308 return WalkResult::advance();
1309 });
1310
1311 // TODO: verify other bits like linkage type.
1312
1313 return failure(walkResult.wasInterrupted());
1314}
1315
1316void spirv::FuncOp::build(OpBuilder &builder, OperationState &state,
1317 StringRef name, FunctionType type,
1318 spirv::FunctionControl control,
1321 builder.getStringAttr(name));
1322 state.addAttribute(getFunctionTypeAttrName(state.name), TypeAttr::get(type));
1323 state.addAttribute(spirv::attributeName<spirv::FunctionControl>(),
1324 builder.getAttr<spirv::FunctionControlAttr>(control));
1325 state.attributes.append(attrs.begin(), attrs.end());
1326 state.addRegion();
1327}
1328
1329//===----------------------------------------------------------------------===//
1330// spirv.GLFClampOp
1331//===----------------------------------------------------------------------===//
1332
1333ParseResult spirv::GLFClampOp::parse(OpAsmParser &parser,
1336}
1337void spirv::GLFClampOp::print(OpAsmPrinter &p) { printOneResultOp(*this, p); }
1338
1339//===----------------------------------------------------------------------===//
1340// spirv.GLUClampOp
1341//===----------------------------------------------------------------------===//
1342
1343ParseResult spirv::GLUClampOp::parse(OpAsmParser &parser,
1346}
1347void spirv::GLUClampOp::print(OpAsmPrinter &p) { printOneResultOp(*this, p); }
1348
1349//===----------------------------------------------------------------------===//
1350// spirv.GLSClampOp
1351//===----------------------------------------------------------------------===//
1352
1353ParseResult spirv::GLSClampOp::parse(OpAsmParser &parser,
1356}
1357void spirv::GLSClampOp::print(OpAsmPrinter &p) { printOneResultOp(*this, p); }
1358
1359//===----------------------------------------------------------------------===//
1360// spirv.GLNClampOp
1361//===----------------------------------------------------------------------===//
1362
1363ParseResult spirv::GLNClampOp::parse(OpAsmParser &parser,
1366}
1367void spirv::GLNClampOp::print(OpAsmPrinter &p) { printOneResultOp(*this, p); }
1368
1369//===----------------------------------------------------------------------===//
1370// spirv.GLSmoothStepOp
1371//===----------------------------------------------------------------------===//
1372
1373ParseResult spirv::GLSmoothStepOp::parse(OpAsmParser &parser,
1376}
1377void spirv::GLSmoothStepOp::print(OpAsmPrinter &p) {
1378 printOneResultOp(*this, p);
1379}
1380
1381//===----------------------------------------------------------------------===//
1382// spirv.GLFmaOp
1383//===----------------------------------------------------------------------===//
1384
1385ParseResult spirv::GLFmaOp::parse(OpAsmParser &parser, OperationState &result) {
1387}
1388void spirv::GLFmaOp::print(OpAsmPrinter &p) { printOneResultOp(*this, p); }
1389
1390//===----------------------------------------------------------------------===//
1391// spirv.GlobalVariable
1392//===----------------------------------------------------------------------===//
1393
1394void spirv::GlobalVariableOp::build(OpBuilder &builder, OperationState &state,
1395 Type type, StringRef name,
1396 unsigned descriptorSet, unsigned binding) {
1397 build(builder, state, TypeAttr::get(type), builder.getStringAttr(name));
1398 state.addAttribute(
1399 spirv::SPIRVDialect::getAttributeName(spirv::Decoration::DescriptorSet),
1400 builder.getI32IntegerAttr(descriptorSet));
1401 state.addAttribute(
1402 spirv::SPIRVDialect::getAttributeName(spirv::Decoration::Binding),
1403 builder.getI32IntegerAttr(binding));
1404}
1405
1406void spirv::GlobalVariableOp::build(OpBuilder &builder, OperationState &state,
1407 Type type, StringRef name,
1408 spirv::BuiltIn builtin) {
1409 build(builder, state, TypeAttr::get(type), builder.getStringAttr(name));
1410 state.addAttribute(
1411 spirv::SPIRVDialect::getAttributeName(spirv::Decoration::BuiltIn),
1412 builder.getStringAttr(spirv::stringifyBuiltIn(builtin)));
1413}
1414
1415ParseResult spirv::GlobalVariableOp::parse(OpAsmParser &parser,
1417 // Parse variable name.
1418 StringAttr nameAttr;
1419 StringRef initializerAttrName =
1420 spirv::GlobalVariableOp::getInitializerAttrName(result.name);
1421 if (parser.parseSymbolName(nameAttr, SymbolTable::getSymbolAttrName(),
1422 result.attributes)) {
1423 return failure();
1424 }
1425
1426 // Parse optional initializer
1427 if (succeeded(parser.parseOptionalKeyword(initializerAttrName))) {
1428 FlatSymbolRefAttr initSymbol;
1429 if (parser.parseLParen() ||
1430 parser.parseAttribute(initSymbol, Type(), initializerAttrName,
1431 result.attributes) ||
1432 parser.parseRParen())
1433 return failure();
1434 }
1435
1436 if (parseVariableDecorations(parser, result)) {
1437 return failure();
1438 }
1439
1440 Type type;
1441 StringRef typeAttrName =
1442 spirv::GlobalVariableOp::getTypeAttrName(result.name);
1443 auto loc = parser.getCurrentLocation();
1444 if (parser.parseColonType(type)) {
1445 return failure();
1446 }
1447 if (!isa<spirv::PointerType>(type)) {
1448 return parser.emitError(loc, "expected spirv.ptr type");
1449 }
1450 result.addAttribute(typeAttrName, TypeAttr::get(type));
1451
1452 return success();
1453}
1454
1455void spirv::GlobalVariableOp::print(OpAsmPrinter &printer) {
1456 SmallVector<StringRef, 4> elidedAttrs{
1457 spirv::attributeName<spirv::StorageClass>()};
1458
1459 // Print variable name.
1460 printer << ' ';
1461 printer.printSymbolName(getSymName());
1462 elidedAttrs.push_back(SymbolTable::getSymbolAttrName());
1463
1464 StringRef initializerAttrName = this->getInitializerAttrName();
1465 // Print optional initializer
1466 if (auto initializer = this->getInitializer()) {
1467 printer << " " << initializerAttrName << '(';
1468 printer.printSymbolName(*initializer);
1469 printer << ')';
1470 elidedAttrs.push_back(initializerAttrName);
1471 }
1472
1473 StringRef typeAttrName = this->getTypeAttrName();
1474 elidedAttrs.push_back(typeAttrName);
1475 spirv::printVariableDecorations(*this, printer, elidedAttrs);
1476 printer << " : " << getType();
1477}
1478
1479LogicalResult spirv::GlobalVariableOp::verify() {
1480 if (!isa<spirv::PointerType>(getType()))
1481 return emitOpError("result must be of a !spv.ptr type");
1482
1483 // SPIR-V spec: "Storage Class is the Storage Class of the memory holding the
1484 // object. It cannot be Generic. It must be the same as the Storage Class
1485 // operand of the Result Type."
1486 // Also, Function storage class is reserved by spirv.Variable.
1487 auto storageClass = this->storageClass();
1488 if (storageClass == spirv::StorageClass::Generic ||
1489 storageClass == spirv::StorageClass::Function) {
1490 return emitOpError("storage class cannot be '")
1491 << stringifyStorageClass(storageClass) << "'";
1492 }
1493
1494 // SPIR-V spec: "A module-scope OpVariable with an Initializer operand must
1495 // not be decorated with the Import Linkage Type."
1496 if (std::optional<spirv::LinkageAttributesAttr> linkage =
1497 getLinkageAttributes()) {
1498 if (linkage->getLinkageType().getValue() == spirv::LinkageType::Import &&
1499 getInitializer()) {
1500 return emitOpError(
1501 "with Import linkage type must not have an initializer");
1502 }
1503 }
1504
1505 if (FlatSymbolRefAttr init = getInitializerAttr()) {
1507 (*this)->getParentOp(), init.getAttr());
1508 // TODO: Currently only variable initialization with specialization
1509 // constants is supported. There could be normal constants in the module
1510 // scope as well.
1511 //
1512 // In the current setup we also cannot initialize one global variable with
1513 // another. The problem is that if we try to initialize pointer of type X
1514 // with another pointer type, the validator fails because it expects the
1515 // variable to be initialized to be type X, not pointer to X. Now
1516 // `spirv.GlobalVariable` only allows pointer type, so in the current design
1517 // we cannot initialize one `spirv.GlobalVariable` with another.
1518 if (!initOp ||
1519 !isa<spirv::SpecConstantOp, spirv::SpecConstantCompositeOp>(initOp)) {
1520 return emitOpError("initializer must be result of a "
1521 "spirv.SpecConstant or "
1522 "spirv.SpecConstantCompositeOp op");
1523 }
1524 }
1525
1526 Type pointeeType = cast<spirv::PointerType>(getType()).getPointeeType();
1527 if (failed(
1528 verifyPhysicalStorageBufferDecorations(getOperation(), pointeeType)))
1529 return failure();
1530
1531 return success();
1532}
1533
1534//===----------------------------------------------------------------------===//
1535// spirv.INTEL.SubgroupBlockRead
1536//===----------------------------------------------------------------------===//
1537
1538LogicalResult spirv::INTELSubgroupBlockReadOp::verify() {
1539 if (failed(verifyBlockReadWritePtrAndValTypes(*this, getPtr(), getValue())))
1540 return failure();
1541
1542 return success();
1543}
1544
1545//===----------------------------------------------------------------------===//
1546// spirv.INTEL.SubgroupBlockWrite
1547//===----------------------------------------------------------------------===//
1548
1549ParseResult spirv::INTELSubgroupBlockWriteOp::parse(OpAsmParser &parser,
1551 // Parse the storage class specification
1552 spirv::StorageClass storageClass;
1554 auto loc = parser.getCurrentLocation();
1555 Type elementType;
1556 if (parseEnumStrAttr(storageClass, parser) ||
1557 parser.parseOperandList(operandInfo, 2) || parser.parseColon() ||
1558 parser.parseType(elementType)) {
1559 return failure();
1560 }
1561
1562 auto ptrType = spirv::PointerType::get(elementType, storageClass);
1563 if (auto valVecTy = dyn_cast<VectorType>(elementType))
1564 ptrType = spirv::PointerType::get(valVecTy.getElementType(), storageClass);
1565
1566 if (parser.resolveOperands(operandInfo, {ptrType, elementType}, loc,
1567 result.operands)) {
1568 return failure();
1569 }
1570 return success();
1571}
1572
1573void spirv::INTELSubgroupBlockWriteOp::print(OpAsmPrinter &printer) {
1574 printer << " " << getPtr() << ", " << getValue() << " : "
1575 << getValue().getType();
1576}
1577
1578LogicalResult spirv::INTELSubgroupBlockWriteOp::verify() {
1579 if (failed(verifyBlockReadWritePtrAndValTypes(*this, getPtr(), getValue())))
1580 return failure();
1581
1582 return success();
1583}
1584
1585//===----------------------------------------------------------------------===//
1586// spirv.IAddCarryOp
1587//===----------------------------------------------------------------------===//
1588
1589LogicalResult spirv::IAddCarryOp::verify() {
1590 return ::verifyArithmeticExtendedBinaryOp(*this);
1591}
1592
1593ParseResult spirv::IAddCarryOp::parse(OpAsmParser &parser,
1595 return ::parseArithmeticExtendedBinaryOp(parser, result);
1596}
1597
1598void spirv::IAddCarryOp::print(OpAsmPrinter &printer) {
1599 ::printArithmeticExtendedBinaryOp(*this, printer);
1600}
1601
1602//===----------------------------------------------------------------------===//
1603// spirv.ISubBorrowOp
1604//===----------------------------------------------------------------------===//
1605
1606LogicalResult spirv::ISubBorrowOp::verify() {
1607 return ::verifyArithmeticExtendedBinaryOp(*this);
1608}
1609
1610ParseResult spirv::ISubBorrowOp::parse(OpAsmParser &parser,
1612 return ::parseArithmeticExtendedBinaryOp(parser, result);
1613}
1614
1615void spirv::ISubBorrowOp::print(OpAsmPrinter &printer) {
1616 ::printArithmeticExtendedBinaryOp(*this, printer);
1617}
1618
1619//===----------------------------------------------------------------------===//
1620// spirv.SMulExtended
1621//===----------------------------------------------------------------------===//
1622
1623LogicalResult spirv::SMulExtendedOp::verify() {
1624 return ::verifyArithmeticExtendedBinaryOp(*this);
1625}
1626
1627ParseResult spirv::SMulExtendedOp::parse(OpAsmParser &parser,
1629 return ::parseArithmeticExtendedBinaryOp(parser, result);
1630}
1631
1632void spirv::SMulExtendedOp::print(OpAsmPrinter &printer) {
1633 ::printArithmeticExtendedBinaryOp(*this, printer);
1634}
1635
1636//===----------------------------------------------------------------------===//
1637// spirv.UMulExtended
1638//===----------------------------------------------------------------------===//
1639
1640LogicalResult spirv::UMulExtendedOp::verify() {
1641 return ::verifyArithmeticExtendedBinaryOp(*this);
1642}
1643
1644ParseResult spirv::UMulExtendedOp::parse(OpAsmParser &parser,
1646 return ::parseArithmeticExtendedBinaryOp(parser, result);
1647}
1648
1649void spirv::UMulExtendedOp::print(OpAsmPrinter &printer) {
1650 ::printArithmeticExtendedBinaryOp(*this, printer);
1651}
1652
1653//===----------------------------------------------------------------------===//
1654// spirv.MemoryBarrierOp
1655//===----------------------------------------------------------------------===//
1656
1657LogicalResult spirv::MemoryBarrierOp::verify() {
1658 return verifyMemorySemantics(getOperation(), getMemorySemantics());
1659}
1660
1661//===----------------------------------------------------------------------===//
1662// spirv.MemoryNamedBarrierOp
1663//===----------------------------------------------------------------------===//
1664
1665LogicalResult spirv::MemoryNamedBarrierOp::verify() {
1666 return verifyMemorySemantics(getOperation(), getMemorySemantics());
1667}
1668
1669//===----------------------------------------------------------------------===//
1670// spirv.module
1671//===----------------------------------------------------------------------===//
1672
1673void spirv::ModuleOp::build(OpBuilder &builder, OperationState &state,
1674 std::optional<StringRef> name) {
1675 OpBuilder::InsertionGuard guard(builder);
1676 builder.createBlock(state.addRegion());
1677 if (name) {
1679 builder.getStringAttr(*name));
1680 }
1681}
1682
1683void spirv::ModuleOp::build(OpBuilder &builder, OperationState &state,
1684 spirv::AddressingModel addressingModel,
1685 spirv::MemoryModel memoryModel,
1686 std::optional<VerCapExtAttr> vceTriple,
1687 std::optional<StringRef> name) {
1688 state.addAttribute(
1689 "addressing_model",
1690 builder.getAttr<spirv::AddressingModelAttr>(addressingModel));
1691 state.addAttribute("memory_model",
1692 builder.getAttr<spirv::MemoryModelAttr>(memoryModel));
1693 OpBuilder::InsertionGuard guard(builder);
1694 builder.createBlock(state.addRegion());
1695 if (vceTriple)
1696 state.addAttribute(getVCETripleAttrName(), *vceTriple);
1697 if (name)
1699 builder.getStringAttr(*name));
1700}
1701
1702ParseResult spirv::ModuleOp::parse(OpAsmParser &parser,
1704 Region *body = result.addRegion();
1705
1706 // If the name is present, parse it.
1707 StringAttr nameAttr;
1709 nameAttr, mlir::SymbolTable::getSymbolAttrName(), result.attributes);
1710
1711 // Parse attributes
1712 spirv::AddressingModel addrModel;
1713 spirv::MemoryModel memoryModel;
1715 result) ||
1717 result))
1718 return failure();
1719
1720 if (succeeded(parser.parseOptionalKeyword("requires"))) {
1721 spirv::VerCapExtAttr vceTriple;
1722 if (parser.parseAttribute(vceTriple,
1723 spirv::ModuleOp::getVCETripleAttrName(),
1724 result.attributes))
1725 return failure();
1726 }
1727
1728 if (parser.parseOptionalAttrDictWithKeyword(result.attributes) ||
1729 parser.parseRegion(*body, /*arguments=*/{}))
1730 return failure();
1731
1732 // Make sure we have at least one block.
1733 if (body->empty())
1734 body->push_back(new Block());
1735
1736 return success();
1737}
1738
1739void spirv::ModuleOp::print(OpAsmPrinter &printer) {
1740 if (std::optional<StringRef> name = getName()) {
1741 printer << ' ';
1742 printer.printSymbolName(*name);
1743 }
1744
1745 SmallVector<StringRef, 2> elidedAttrs;
1746
1747 printer << " " << spirv::stringifyAddressingModel(getAddressingModel()) << " "
1748 << spirv::stringifyMemoryModel(getMemoryModel());
1749 auto addressingModelAttrName = spirv::attributeName<spirv::AddressingModel>();
1750 auto memoryModelAttrName = spirv::attributeName<spirv::MemoryModel>();
1751 elidedAttrs.assign({addressingModelAttrName, memoryModelAttrName,
1753
1754 if (std::optional<spirv::VerCapExtAttr> triple = getVceTriple()) {
1755 printer << " requires " << *triple;
1756 elidedAttrs.push_back(spirv::ModuleOp::getVCETripleAttrName());
1757 }
1758
1760 (*this)->getDiscardableAttrDictionary().getValue(), elidedAttrs);
1761 printer << ' ';
1762 printer.printRegion(getRegion());
1763}
1764
1765LogicalResult spirv::ModuleOp::verifyRegions() {
1766 Dialect *dialect = (*this)->getDialect();
1768 entryPoints;
1769 mlir::SymbolTable table(*this);
1770
1771 for (auto &op : *getBody()) {
1772 if (op.getDialect() != dialect)
1773 return op.emitError("'spirv.module' can only contain spirv.* ops");
1774
1775 // For EntryPoint op, check that the function and execution model is not
1776 // duplicated in EntryPointOps. Also verify that the interface specified
1777 // comes from globalVariables here to make this check cheaper.
1778 if (auto entryPointOp = dyn_cast<spirv::EntryPointOp>(op)) {
1779 auto funcOp = table.lookup<spirv::FuncOp>(entryPointOp.getFn());
1780 if (!funcOp) {
1781 return entryPointOp.emitError("function '")
1782 << entryPointOp.getFn() << "' not found in 'spirv.module'";
1783 }
1784 if (auto interface = entryPointOp.getInterface()) {
1785 for (Attribute varRef : interface) {
1786 auto varSymRef = dyn_cast<FlatSymbolRefAttr>(varRef);
1787 if (!varSymRef) {
1788 return entryPointOp.emitError(
1789 "expected symbol reference for interface "
1790 "specification instead of '")
1791 << varRef;
1792 }
1793 auto variableOp =
1794 table.lookup<spirv::GlobalVariableOp>(varSymRef.getValue());
1795 if (!variableOp) {
1796 return entryPointOp.emitError("expected spirv.GlobalVariable "
1797 "symbol reference instead of'")
1798 << varSymRef << "'";
1799 }
1800 }
1801 }
1802
1803 auto key = std::pair<spirv::FuncOp, spirv::ExecutionModel>(
1804 funcOp, entryPointOp.getExecutionModel());
1805 if (!entryPoints.try_emplace(key, entryPointOp).second)
1806 return entryPointOp.emitError("duplicate of a previous EntryPointOp");
1807 } else if (auto funcOp = dyn_cast<spirv::FuncOp>(op)) {
1808 // If the function is external and does not have 'Import'
1809 // linkage_attributes(LinkageAttributes), throw an error. 'Import'
1810 // LinkageAttributes is used to import external functions.
1811 auto linkageAttr = funcOp.getLinkageAttributes();
1812 auto hasImportLinkage =
1813 linkageAttr && (linkageAttr.value().getLinkageType().getValue() ==
1814 spirv::LinkageType::Import);
1815 if (funcOp.isExternal() && !hasImportLinkage)
1816 return op.emitError(
1817 "'spirv.module' cannot contain external functions "
1818 "without 'Import' linkage_attributes (LinkageAttributes)");
1819
1820 // TODO: move this check to spirv.func.
1821 for (auto &block : funcOp)
1822 for (auto &op : block) {
1823 if (op.getDialect() != dialect)
1824 return op.emitError(
1825 "functions in 'spirv.module' can only contain spirv.* ops");
1826 }
1827 }
1828 }
1829
1830 return success();
1831}
1832
1833//===----------------------------------------------------------------------===//
1834// spirv.mlir.referenceof
1835//===----------------------------------------------------------------------===//
1836
1837LogicalResult spirv::ReferenceOfOp::verify() {
1838 auto *specConstSym = SymbolTable::lookupNearestSymbolFrom(
1839 (*this)->getParentOp(), getSpecConstAttr());
1840 Type constType;
1841
1842 auto specConstOp = dyn_cast_or_null<spirv::SpecConstantOp>(specConstSym);
1843 if (specConstOp)
1844 constType = specConstOp.getDefaultValue().getType();
1845
1846 auto specConstCompositeOp =
1847 dyn_cast_or_null<spirv::SpecConstantCompositeOp>(specConstSym);
1848 if (specConstCompositeOp)
1849 constType = specConstCompositeOp.getType();
1850
1851 if (!specConstOp && !specConstCompositeOp)
1852 return emitOpError(
1853 "expected spirv.SpecConstant or spirv.SpecConstantComposite symbol");
1854
1855 if (getReference().getType() != constType)
1856 return emitOpError("result type mismatch with the referenced "
1857 "specialization constant's type");
1858
1859 return success();
1860}
1861
1862//===----------------------------------------------------------------------===//
1863// spirv.SpecConstant
1864//===----------------------------------------------------------------------===//
1865
1866ParseResult spirv::SpecConstantOp::parse(OpAsmParser &parser,
1868 StringAttr nameAttr;
1869 Attribute valueAttr;
1870 StringRef defaultValueAttrName =
1871 spirv::SpecConstantOp::getDefaultValueAttrName(result.name);
1872
1873 if (parser.parseSymbolName(nameAttr, SymbolTable::getSymbolAttrName(),
1874 result.attributes))
1875 return failure();
1876
1877 // Parse optional spec_id.
1878 if (succeeded(parser.parseOptionalKeyword(kSpecIdAttrName))) {
1879 IntegerAttr specIdAttr;
1880 if (parser.parseLParen() ||
1881 parser.parseAttribute(specIdAttr, kSpecIdAttrName, result.attributes) ||
1882 parser.parseRParen())
1883 return failure();
1884 }
1885
1886 if (parser.parseEqual() ||
1887 parser.parseAttribute(valueAttr, defaultValueAttrName, result.attributes))
1888 return failure();
1889
1890 return success();
1891}
1892
1893void spirv::SpecConstantOp::print(OpAsmPrinter &printer) {
1894 printer << ' ';
1895 printer.printSymbolName(getSymName());
1896 if (auto specID =
1897 (*this)->getDiscardableAttrOfType<IntegerAttr>(kSpecIdAttrName))
1898 printer << ' ' << kSpecIdAttrName << '(' << specID.getInt() << ')';
1899 printer << " = " << getDefaultValue();
1900}
1901
1902LogicalResult spirv::SpecConstantOp::verify() {
1903 if (auto specID =
1904 (*this)->getDiscardableAttrOfType<IntegerAttr>(kSpecIdAttrName))
1905 if (specID.getValue().isNegative())
1906 return emitOpError("SpecId cannot be negative");
1907
1908 auto value = getDefaultValue();
1909 if (isa<IntegerAttr, FloatAttr>(value)) {
1910 // Make sure bitwidth is allowed.
1911 if (!isa<spirv::SPIRVType>(value.getType()))
1912 return emitOpError("default value bitwidth disallowed");
1913 return success();
1914 }
1915 return emitOpError(
1916 "default value can only be a bool, integer, or float scalar");
1917}
1918
1919//===----------------------------------------------------------------------===//
1920// spirv.VectorShuffle
1921//===----------------------------------------------------------------------===//
1922
1923LogicalResult spirv::VectorShuffleOp::verify() {
1924 VectorType resultType = cast<VectorType>(getType());
1925
1926 size_t numResultElements = resultType.getNumElements();
1927 if (numResultElements != getComponents().size())
1928 return emitOpError("result type element count (")
1929 << numResultElements
1930 << ") mismatch with the number of component selectors ("
1931 << getComponents().size() << ")";
1932
1933 size_t totalSrcElements =
1934 cast<VectorType>(getVector1().getType()).getNumElements() +
1935 cast<VectorType>(getVector2().getType()).getNumElements();
1936
1937 for (const auto &selector : getComponents().getAsValueRange<IntegerAttr>()) {
1938 uint32_t index = selector.getZExtValue();
1939 if (index >= totalSrcElements &&
1940 index != std::numeric_limits<uint32_t>().max())
1941 return emitOpError("component selector ")
1942 << index << " out of range: expected to be in [0, "
1943 << totalSrcElements << ") or 0xffffffff";
1944 }
1945 return success();
1946}
1947
1948//===----------------------------------------------------------------------===//
1949// spirv.SpecConstantComposite
1950//===----------------------------------------------------------------------===//
1951
1952ParseResult spirv::SpecConstantCompositeOp::parse(OpAsmParser &parser,
1954
1955 StringAttr compositeName;
1956 if (parser.parseSymbolName(compositeName, SymbolTable::getSymbolAttrName(),
1957 result.attributes))
1958 return failure();
1959
1960 if (parser.parseLParen())
1961 return failure();
1962
1963 SmallVector<Attribute, 4> constituents;
1964
1965 do {
1966 // The name of the constituent attribute isn't important
1967 const char *attrName = "spec_const";
1968 FlatSymbolRefAttr specConstRef;
1969 NamedAttrList attrs;
1970
1971 if (parser.parseAttribute(specConstRef, Type(), attrName, attrs))
1972 return failure();
1973
1974 constituents.push_back(specConstRef);
1975 } while (!parser.parseOptionalComma());
1976
1977 if (parser.parseRParen())
1978 return failure();
1979
1980 StringAttr compositeSpecConstituentsName =
1981 spirv::SpecConstantCompositeOp::getConstituentsAttrName(result.name);
1982 result.addAttribute(compositeSpecConstituentsName,
1983 parser.getBuilder().getArrayAttr(constituents));
1984
1985 Type type;
1986 if (parser.parseColonType(type))
1987 return failure();
1988
1989 StringAttr typeAttrName =
1990 spirv::SpecConstantCompositeOp::getTypeAttrName(result.name);
1991 result.addAttribute(typeAttrName, TypeAttr::get(type));
1992
1993 return success();
1994}
1995
1996void spirv::SpecConstantCompositeOp::print(OpAsmPrinter &printer) {
1997 printer << " ";
1998 printer.printSymbolName(getSymName());
1999 printer << " (" << llvm::interleaved(this->getConstituents().getValue())
2000 << ") : " << getType();
2001}
2002
2003LogicalResult spirv::SpecConstantCompositeOp::verify() {
2004 auto cType = dyn_cast<spirv::CompositeType>(getType());
2005 auto constituents = this->getConstituents().getValue();
2006
2007 if (!cType)
2008 return emitError("result type must be a composite type, but provided ")
2009 << getType();
2010
2011 if (isa<spirv::CooperativeMatrixType>(cType))
2012 return emitError("unsupported composite type ") << cType;
2013 if (constituents.size() != cType.getNumElements())
2014 return emitError("has incorrect number of operands: expected ")
2015 << cType.getNumElements() << ", but provided "
2016 << constituents.size();
2017
2018 for (auto index : llvm::seq<uint32_t>(0, constituents.size())) {
2019 auto constituent = cast<FlatSymbolRefAttr>(constituents[index]);
2020
2022 (*this)->getParentOp(), constituent.getAttr());
2023
2024 if (!constituentOp)
2025 return emitError("unknown constituent symbol ") << constituent.getAttr();
2026
2027 Type constituentType;
2028 if (auto specConstOp = dyn_cast<spirv::SpecConstantOp>(constituentOp)) {
2029 constituentType = specConstOp.getDefaultValue().getType();
2030 } else if (auto specConstCompositeOp =
2031 dyn_cast<spirv::SpecConstantCompositeOp>(constituentOp)) {
2032 constituentType = specConstCompositeOp.getType();
2033 } else {
2034 return emitError("unsupported constituent ")
2035 << constituent.getAttr()
2036 << ": must reference a spirv.SpecConstant or "
2037 "spirv.SpecConstantComposite";
2038 }
2039
2040 if (constituentType != cType.getElementType(index))
2041 return emitError("has incorrect types of operands: expected ")
2042 << cType.getElementType(index) << ", but provided "
2043 << constituentType;
2044 }
2045
2046 return success();
2047}
2048
2049//===----------------------------------------------------------------------===//
2050// spirv.EXTSpecConstantCompositeReplicateOp
2051//===----------------------------------------------------------------------===//
2052
2053ParseResult
2054spirv::EXTSpecConstantCompositeReplicateOp::parse(OpAsmParser &parser,
2056 StringAttr compositeName;
2057 FlatSymbolRefAttr specConstRef;
2058 const char *attrName = "spec_const";
2059 NamedAttrList attrs;
2060 Type type;
2061
2062 if (parser.parseSymbolName(compositeName, SymbolTable::getSymbolAttrName(),
2063 result.attributes) ||
2064 parser.parseLParen() ||
2065 parser.parseAttribute(specConstRef, Type(), attrName, attrs) ||
2066 parser.parseRParen() || parser.parseColonType(type))
2067 return failure();
2068
2069 StringAttr compositeSpecConstituentName =
2070 spirv::EXTSpecConstantCompositeReplicateOp::getConstituentAttrName(
2071 result.name);
2072 result.addAttribute(compositeSpecConstituentName, specConstRef);
2073
2074 StringAttr typeAttrName =
2075 spirv::EXTSpecConstantCompositeReplicateOp::getTypeAttrName(result.name);
2076 result.addAttribute(typeAttrName, TypeAttr::get(type));
2077
2078 return success();
2079}
2080
2081void spirv::EXTSpecConstantCompositeReplicateOp::print(OpAsmPrinter &printer) {
2082 printer << " ";
2083 printer.printSymbolName(getSymName());
2084 printer << " (" << this->getConstituent() << ") : " << getType();
2085}
2086
2087LogicalResult spirv::EXTSpecConstantCompositeReplicateOp::verify() {
2088 auto compositeType = dyn_cast<spirv::CompositeType>(getType());
2089 if (!compositeType)
2090 return emitError("result type must be a composite type, but provided ")
2091 << getType();
2092
2094 (*this)->getParentOp(), this->getConstituent());
2095 if (!constituentOp)
2096 return emitError(
2097 "splat spec constant reference defining constituent not found");
2098
2099 auto constituentSpecConstOp = dyn_cast<spirv::SpecConstantOp>(constituentOp);
2100 if (!constituentSpecConstOp)
2101 return emitError("constituent is not a spec constant");
2102
2103 Type constituentType = constituentSpecConstOp.getDefaultValue().getType();
2104 Type compositeElementType = compositeType.getElementType(0);
2105 if (constituentType != compositeElementType)
2106 return emitError("constituent has incorrect type: expected ")
2107 << compositeElementType << ", but provided " << constituentType;
2108
2109 return success();
2110}
2111
2112//===----------------------------------------------------------------------===//
2113// spirv.SpecConstantOperation
2114//===----------------------------------------------------------------------===//
2115
2116ParseResult spirv::SpecConstantOperationOp::parse(OpAsmParser &parser,
2118 Region *body = result.addRegion();
2119
2120 if (parser.parseKeyword("wraps"))
2121 return failure();
2122
2123 body->push_back(new Block);
2124 Block &block = body->back();
2125 Operation *wrappedOp = parser.parseGenericOperation(&block, block.begin());
2126
2127 if (!wrappedOp)
2128 return failure();
2129
2130 OpBuilder builder(parser.getContext());
2131 builder.setInsertionPointToEnd(&block);
2132 spirv::YieldOp::create(builder, wrappedOp->getLoc(), wrappedOp->getResult(0));
2133 result.location = wrappedOp->getLoc();
2134
2135 result.addTypes(wrappedOp->getResult(0).getType());
2136
2137 if (parser.parseOptionalAttrDict(result.attributes))
2138 return failure();
2139
2140 return success();
2141}
2142
2143void spirv::SpecConstantOperationOp::print(OpAsmPrinter &printer) {
2144 printer << " wraps ";
2145 printer.printGenericOp(&getBody().front().front());
2146}
2147
2148LogicalResult spirv::SpecConstantOperationOp::verifyRegions() {
2149 Block &block = getRegion().getBlocks().front();
2150
2151 if (block.getOperations().size() != 2)
2152 return emitOpError("expected exactly 2 nested ops");
2153
2154 Operation &enclosedOp = block.getOperations().front();
2155
2157 return emitOpError("invalid enclosed op");
2158
2159 for (auto operand : enclosedOp.getOperands())
2160 if (!isa_and_present<spirv::ConstantOp, spirv::ReferenceOfOp,
2161 spirv::SpecConstantOperationOp>(
2162 operand.getDefiningOp()))
2163 return emitOpError(
2164 "invalid operand, must be defined by a constant operation");
2165
2166 return success();
2167}
2168
2169//===----------------------------------------------------------------------===//
2170// spirv.GL.FrexpStruct
2171//===----------------------------------------------------------------------===//
2172
2173LogicalResult spirv::GLFrexpStructOp::verify() {
2174 spirv::StructType structTy =
2175 dyn_cast<spirv::StructType>(getResult().getType());
2176
2177 if (structTy.getNumElements() != 2)
2178 return emitError("result type must be a struct type with two memebers");
2179
2180 Type significandTy = structTy.getElementType(0);
2181 Type exponentTy = structTy.getElementType(1);
2182 VectorType exponentVecTy = dyn_cast<VectorType>(exponentTy);
2183 IntegerType exponentIntTy = dyn_cast<IntegerType>(exponentTy);
2184
2185 Type operandTy = getOperand().getType();
2186 VectorType operandVecTy = dyn_cast<VectorType>(operandTy);
2187 FloatType operandFTy = dyn_cast<FloatType>(operandTy);
2188
2189 if (significandTy != operandTy)
2190 return emitError("member zero of the resulting struct type must be the "
2191 "same type as the operand");
2192
2193 if (exponentVecTy) {
2194 IntegerType componentIntTy =
2195 dyn_cast<IntegerType>(exponentVecTy.getElementType());
2196 if (!componentIntTy || componentIntTy.getWidth() != 32)
2197 return emitError("member one of the resulting struct type must"
2198 "be a scalar or vector of 32 bit integer type");
2199 } else if (!exponentIntTy || exponentIntTy.getWidth() != 32) {
2200 return emitError("member one of the resulting struct type "
2201 "must be a scalar or vector of 32 bit integer type");
2202 }
2203
2204 // Check that the two member types have the same number of components
2205 if (operandVecTy && exponentVecTy &&
2206 (exponentVecTy.getNumElements() == operandVecTy.getNumElements()))
2207 return success();
2208
2209 if (operandFTy && exponentIntTy)
2210 return success();
2211
2212 return emitError("member one of the resulting struct type must have the same "
2213 "number of components as the operand type");
2214}
2215
2216//===----------------------------------------------------------------------===//
2217// spirv.GL.Ldexp
2218//===----------------------------------------------------------------------===//
2219
2220static LogicalResult verifyFloatIntegerBuiltin(Operation *op, Type floatType,
2221 Type integerType) {
2222 if (isa<FloatType>(floatType) != isa<IntegerType>(integerType))
2223 return op->emitOpError("operands must both be scalars or vectors");
2224
2225 auto getNumElements = [](Type type) -> unsigned {
2226 if (auto vectorType = dyn_cast<VectorType>(type))
2227 return vectorType.getNumElements();
2228 return 1;
2229 };
2230
2231 if (getNumElements(floatType) != getNumElements(integerType))
2232 return op->emitOpError("operands must have the same number of elements");
2233
2234 return success();
2235}
2236
2237LogicalResult spirv::GLLdexpOp::verify() {
2238 return verifyFloatIntegerBuiltin(getOperation(), getX().getType(),
2239 getExp().getType());
2240}
2241
2242//===----------------------------------------------------------------------===//
2243// spirv.CL.ldexp
2244//===----------------------------------------------------------------------===//
2245
2246LogicalResult spirv::CLLdexpOp::verify() {
2247 return verifyFloatIntegerBuiltin(getOperation(), getX().getType(),
2248 getExp().getType());
2249}
2250
2251//===----------------------------------------------------------------------===//
2252// spirv.CL.pown
2253//===----------------------------------------------------------------------===//
2254
2255LogicalResult spirv::CLPownOp::verify() {
2256 return verifyFloatIntegerBuiltin(getOperation(), getX().getType(),
2257 getY().getType());
2258}
2259
2260//===----------------------------------------------------------------------===//
2261// spirv.CL.rootn
2262//===----------------------------------------------------------------------===//
2263
2264LogicalResult spirv::CLRootnOp::verify() {
2265 return verifyFloatIntegerBuiltin(getOperation(), getX().getType(),
2266 getN().getType());
2267}
2268
2269//===----------------------------------------------------------------------===//
2270// spirv.ShiftLeftLogicalOp
2271//===----------------------------------------------------------------------===//
2272
2273LogicalResult spirv::ShiftLeftLogicalOp::verify() {
2274 return verifyShiftOp(*this);
2275}
2276
2277//===----------------------------------------------------------------------===//
2278// spirv.ShiftRightArithmeticOp
2279//===----------------------------------------------------------------------===//
2280
2281LogicalResult spirv::ShiftRightArithmeticOp::verify() {
2282 return verifyShiftOp(*this);
2283}
2284
2285//===----------------------------------------------------------------------===//
2286// spirv.ShiftRightLogicalOp
2287//===----------------------------------------------------------------------===//
2288
2289LogicalResult spirv::ShiftRightLogicalOp::verify() {
2290 return verifyShiftOp(*this);
2291}
2292
2293//===----------------------------------------------------------------------===//
2294// spirv.VectorTimesScalarOp
2295//===----------------------------------------------------------------------===//
2296
2297LogicalResult spirv::VectorTimesScalarOp::verify() {
2298 if (getVector().getType() != getType())
2299 return emitOpError("vector operand and result type mismatch");
2300 auto scalarType = cast<VectorType>(getType()).getElementType();
2301 if (getScalar().getType() != scalarType)
2302 return emitOpError("scalar operand and result element type match");
2303 return success();
2304}
return success()
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 int64_t getNumElements(Type t)
Compute the total number of elements in the given type, also taking into account nested types.
ArrayAttr()
static Value max(ImplicitLocOpBuilder &builder, Value value, Value bound)
static ParseResult parseArithmeticExtendedBinaryOp(OpAsmParser &parser, OperationState &result)
Definition SPIRVOps.cpp:310
static Type getValueType(Attribute attr)
Definition SPIRVOps.cpp:835
static LogicalResult verifyConstantType(spirv::ConstantOp op, Attribute value, Type opType)
Definition SPIRVOps.cpp:606
static ParseResult parseOneResultSameOperandTypeOp(OpAsmParser &parser, OperationState &result)
Definition SPIRVOps.cpp:164
static LogicalResult verifyArithmeticExtendedBinaryOp(ExtendedBinaryOp op)
Definition SPIRVOps.cpp:296
static LogicalResult verifyFloatIntegerBuiltin(Operation *op, Type floatType, Type integerType)
static LogicalResult verifyShiftOp(Operation *op)
Definition SPIRVOps.cpp:342
static LogicalResult verifyBlockReadWritePtrAndValTypes(BlockReadWriteOpTy op, Value ptr, Value val)
Definition SPIRVOps.cpp:213
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,...
Definition SPIRVOps.cpp:229
static void printOneResultOp(Operation *op, OpAsmPrinter &p)
Definition SPIRVOps.cpp:193
static void printArithmeticExtendedBinaryOp(Operation *op, OpAsmPrinter &printer)
Definition SPIRVOps.cpp:334
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.
Definition Attributes.h:25
Block represents an ordered list of Operations.
Definition Block.h:33
ValueTypeRange< BlockArgListType > getArgumentTypes()
Return a range containing the types of the arguments for this block.
Definition Block.cpp:154
unsigned getNumArguments()
Definition Block.h:152
OpListType & getOperations()
Definition Block.h:161
Operation & front()
Definition Block.h:177
iterator begin()
Definition Block.h:167
IntegerAttr getI32IntegerAttr(int32_t value)
Definition Builders.cpp:208
IntegerAttr getIntegerAttr(Type type, int64_t value)
Definition Builders.cpp:237
ArrayAttr getI32ArrayAttr(ArrayRef< int32_t > values)
Definition Builders.cpp:285
FloatAttr getFloatAttr(Type type, double value)
Definition Builders.cpp:263
FunctionType getFunctionType(TypeRange inputs, TypeRange results)
Definition Builders.cpp:84
IntegerType getIntegerType(unsigned width)
Definition Builders.cpp:75
BoolAttr getBoolAttr(bool value)
Definition Builders.cpp:108
StringAttr getStringAttr(const Twine &bytes)
Definition Builders.cpp:271
ArrayAttr getArrayAttr(ArrayRef< Attribute > value)
Definition Builders.cpp:275
MLIRContext * getContext() const
Definition Builders.h:56
Attr getAttr(Args &&...args)
Get or construct an instance of the attribute Attr with provided arguments.
Definition Builders.h:101
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...
Definition Dialect.h:38
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...
Definition Location.h:76
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 &region, 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 &region, 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.
Definition Builders.h:351
This class helps build Operations.
Definition Builders.h:210
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.
Definition Builders.cpp:439
void setInsertionPointToEnd(Block *block)
Sets the insertion point to the end of the specified block.
Definition Builders.h:439
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
Operation is the basic unit of execution within MLIR.
Definition Operation.h:87
Attribute getDiscardableAttr(StringRef name)
Access a discardable attribute by name, returns a null Attribute if the discardable attribute does no...
Definition Operation.h:478
Dialect * getDialect()
Return the dialect this operation is associated with, or nullptr if the associated dialect is not loa...
Definition Operation.h:237
Value getOperand(unsigned idx)
Definition Operation.h:375
bool hasTrait()
Returns true if the operation was registered with a particular trait, e.g.
Definition Operation.h:794
OpResult getResult(unsigned idx)
Get the 'idx'th result of this operation.
Definition Operation.h:432
Location getLoc()
The source location the operation was defined or derived from.
Definition Operation.h:240
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.
Definition Operation.h:115
DictionaryAttr getDiscardableAttrDictionary()
Return all of the discardable attributes on this operation as a DictionaryAttr.
Definition Operation.h:546
operand_type_range getOperandTypes()
Definition Operation.h:422
result_type_range getResultTypes()
Definition Operation.h:453
operand_range getOperands()
Returns an iterator on the underlying Value's.
Definition Operation.h:403
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.
Definition Operation.h:429
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.
Definition Region.h:26
void push_back(Block *block)
Definition Region.h:61
Block & back()
Definition Region.h:64
bool empty()
Definition Region.h:60
This class allows for representing and managing the symbol table used by operations with the 'SymbolT...
Definition SymbolTable.h:24
static StringRef getSymbolAttrName()
Return the name of the attribute used for symbol names.
Definition SymbolTable.h:76
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...
Definition Types.h:74
Dialect & getDialect() const
Get the dialect this type is registered to.
Definition Types.h:107
Type front()
Return first type in the range.
Definition TypeRange.h:164
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
Definition Value.h:96
Type getType() const
Return the type of this value.
Definition Value.h:105
A utility result that is used to signal how to proceed with an ongoing walk:
Definition WalkResult.h:29
static WalkResult advance()
Definition WalkResult.h:47
static ArrayType get(Type elementType, unsigned elementCount)
static PointerType get(Type pointeeType, StorageClass storageClass)
SPIR-V struct type.
Definition SPIRVTypes.h:274
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.
Definition Visitors.h:102
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.
Definition Utils.cpp:18
detail::InFlightRemark failed(Location loc, RemarkOpts opts)
Report an optimization remark that failed.
Definition Remarks.h:717
uint64_t getN(LevelType lt)
Definition Enums.h:442
constexpr char kFnNameAttrName[]
constexpr char kSpecIdAttrName[]
LogicalResult verifyMemorySemantics(Operation *op, spirv::MemorySemantics memorySemantics)
Definition SPIRVOps.cpp:69
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)
Definition SPIRVOps.cpp:133
LogicalResult verifyPhysicalStorageBufferDecorations(Operation *op, Type pointeeType)
Verifies the SPV_KHR_physical_storage_buffer rule that a variable whose pointee is a pointer (or arra...
Definition SPIRVOps.cpp:93
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)
Definition SPIRVOps.cpp:49
std::string getDecorationString(Decoration decoration)
Converts a SPIR-V Decoration enum value to its snake_case string representation for use in MLIR attri...
ParseResult parseVariableDecorations(OpAsmParser &parser, OperationState &state)
Include the generated interface declarations.
Type getType(OpFoldResult ofr)
Returns the int type of the integer in ofr.
Definition Utils.cpp:307
InFlightDiagnostic emitError(Location loc)
Utility method to emit an error message using this location.
llvm::DenseMap< KeyT, ValueT, KeyInfoT, BucketT > DenseMap
Definition LLVM.h:120
llvm::function_ref< Fn > function_ref
Definition LLVM.h:147
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.