23#include "llvm/ADT/STLExtras.h"
24#include "llvm/ADT/Sequence.h"
25#include "llvm/ADT/SmallVector.h"
26#include "llvm/ADT/StringExtras.h"
27#include "llvm/ADT/bit.h"
28#include "llvm/Support/Debug.h"
29#include "llvm/Support/SaveAndRestore.h"
30#include "llvm/Support/raw_ostream.h"
35#define DEBUG_TYPE "spirv-deserialization"
44 isa_and_nonnull<spirv::FuncOp>(block->
getParentOp());
54 : binary(binary), context(context), unknownLoc(UnknownLoc::
get(context)),
55 module(createModuleOp()), opBuilder(module->getRegion()),
options(
options)
63LogicalResult spirv::Deserializer::deserialize() {
67 <<
"//+++---------- start deserialization ----------+++//\n";
70 if (
failed(processHeader()))
73 spirv::Opcode opcode = spirv::Opcode::OpNop;
74 ArrayRef<uint32_t> operands;
75 auto binarySize = binary.size();
76 while (curOffset < binarySize) {
86 assert(curOffset == binarySize &&
87 "deserializer should never index beyond the binary end");
89 for (
auto &deferred : deferredInstructions) {
95 if (
failed(resolveDeferredIdDecorations()))
100 LLVM_DEBUG(logger.startLine()
101 <<
"//+++-------- completed deserialization --------+++//\n");
105OwningOpRef<spirv::ModuleOp> spirv::Deserializer::collect() {
106 return std::move(module);
113OwningOpRef<spirv::ModuleOp> spirv::Deserializer::createModuleOp() {
114 OpBuilder builder(context);
115 OperationState state(unknownLoc, spirv::ModuleOp::getOperationName());
116 spirv::ModuleOp::build(builder, state);
120LogicalResult spirv::Deserializer::processHeader() {
123 "SPIR-V binary module must have a 5-word header");
126 return emitError(unknownLoc,
"incorrect magic number");
129 uint32_t majorVersion = (binary[1] << 8) >> 24;
130 uint32_t minorVersion = (binary[1] << 16) >> 24;
131 if (majorVersion == 1) {
132 switch (minorVersion) {
133#define MIN_VERSION_CASE(v) \
135 version = spirv::Version::V_1_##v; \
145#undef MIN_VERSION_CASE
147 return emitError(unknownLoc,
"unsupported SPIR-V minor version: ")
151 return emitError(unknownLoc,
"unsupported SPIR-V major version: ")
161spirv::Deserializer::processCapability(ArrayRef<uint32_t> operands) {
162 if (operands.size() != 1)
163 return emitError(unknownLoc,
"OpCapability must have one parameter");
165 auto cap = spirv::symbolizeCapability(operands[0]);
167 return emitError(unknownLoc,
"unknown capability: ") << operands[0];
169 capabilities.insert(*cap);
173LogicalResult spirv::Deserializer::processExtension(ArrayRef<uint32_t> words) {
177 "OpExtension must have a literal string for the extension name");
180 unsigned wordIndex = 0;
182 if (wordIndex != words.size())
184 "unexpected trailing words in OpExtension instruction");
185 auto ext = spirv::symbolizeExtension(extName);
187 return emitError(unknownLoc,
"unknown extension: ") << extName;
189 extensions.insert(*ext);
194spirv::Deserializer::processExtInstImport(ArrayRef<uint32_t> words) {
195 if (words.size() < 2) {
197 "OpExtInstImport must have a result <id> and a literal "
198 "string for the extended instruction set name");
201 unsigned wordIndex = 1;
203 if (wordIndex != words.size()) {
205 "unexpected trailing words in OpExtInstImport");
210void spirv::Deserializer::attachVCETriple() {
212 spirv::ModuleOp::getVCETripleAttrName(),
214 extensions.getArrayRef(), context));
218spirv::Deserializer::processMemoryModel(ArrayRef<uint32_t> operands) {
219 if (operands.size() != 2)
220 return emitError(unknownLoc,
"OpMemoryModel must have two operands");
223 module->getAddressingModelAttrName(),
224 opBuilder.getAttr<spirv::AddressingModelAttr>(
225 static_cast<spirv::AddressingModel
>(operands.front())));
227 (*module)->setAttr(module->getMemoryModelAttrName(),
228 opBuilder.getAttr<spirv::MemoryModelAttr>(
229 static_cast<spirv::MemoryModel
>(operands.back())));
234template <
typename AttrTy,
typename EnumAttrTy,
typename EnumTy>
238 StringAttr symbol, StringRef decorationName, StringRef cacheControlKind) {
239 if (words.size() != 4) {
240 return emitError(loc,
"OpDecorate with ")
241 << decorationName <<
" needs a cache control integer literal and a "
242 << cacheControlKind <<
" cache control literal";
244 unsigned cacheLevel = words[2];
245 auto cacheControlAttr =
static_cast<EnumTy
>(words[3]);
246 auto value = opBuilder.
getAttr<AttrTy>(cacheLevel, cacheControlAttr);
249 dyn_cast_or_null<ArrayAttr>(decorations[words[0]].
get(symbol)))
250 llvm::append_range(attrs, attrList);
251 attrs.push_back(value);
252 decorations[words[0]].set(symbol, opBuilder.
getArrayAttr(attrs));
256LogicalResult spirv::Deserializer::processDecoration(ArrayRef<uint32_t> words) {
260 if (words.size() < 2) {
262 unknownLoc,
"OpDecorate must have at least result <id> and Decoration");
264 auto decorationName =
265 stringifyDecoration(
static_cast<spirv::Decoration
>(words[1]));
266 if (decorationName.empty()) {
267 return emitError(unknownLoc,
"invalid Decoration code : ") << words[1];
269 auto symbol = getSymbolDecoration(decorationName);
270 switch (
static_cast<spirv::Decoration
>(words[1])) {
271 case spirv::Decoration::FPFastMathMode:
272 if (words.size() != 3) {
273 return emitError(unknownLoc,
"OpDecorate with ")
274 << decorationName <<
" needs a single integer literal";
276 decorations[words[0]].set(
277 symbol, FPFastMathModeAttr::get(opBuilder.getContext(),
278 static_cast<FPFastMathMode
>(words[2])));
280 case spirv::Decoration::FPRoundingMode:
281 if (words.size() != 3) {
282 return emitError(unknownLoc,
"OpDecorate with ")
283 << decorationName <<
" needs a single integer literal";
285 decorations[words[0]].set(
286 symbol, FPRoundingModeAttr::get(opBuilder.getContext(),
287 static_cast<FPRoundingMode
>(words[2])));
289 case spirv::Decoration::DescriptorSet:
290 case spirv::Decoration::Binding:
291 case spirv::Decoration::Location:
292 case spirv::Decoration::SpecId:
293 case spirv::Decoration::Index:
294 case spirv::Decoration::Offset:
295 case spirv::Decoration::XfbBuffer:
296 case spirv::Decoration::XfbStride:
297 if (words.size() != 3) {
298 return emitError(unknownLoc,
"OpDecorate with ")
299 << decorationName <<
" needs a single integer literal";
301 decorations[words[0]].set(
302 symbol, opBuilder.getI32IntegerAttr(
static_cast<int32_t
>(words[2])));
304 case spirv::Decoration::BuiltIn:
305 if (words.size() != 3) {
306 return emitError(unknownLoc,
"OpDecorate with ")
307 << decorationName <<
" needs a single integer literal";
309 decorations[words[0]].set(
310 symbol, opBuilder.getStringAttr(
311 stringifyBuiltIn(
static_cast<spirv::BuiltIn
>(words[2]))));
313 case spirv::Decoration::ArrayStride:
314 if (words.size() != 3) {
315 return emitError(unknownLoc,
"OpDecorate with ")
316 << decorationName <<
" needs a single integer literal";
318 typeDecorations[words[0]] = words[2];
320 case spirv::Decoration::LinkageAttributes: {
321 if (words.size() < 4) {
322 return emitError(unknownLoc,
"OpDecorate with ")
324 <<
" needs at least 1 string and 1 integer literal";
332 unsigned wordIndex = 2;
334 auto linkageTypeAttr = opBuilder.getAttr<::mlir::spirv::LinkageTypeAttr>(
335 static_cast<::mlir::spirv::LinkageType
>(words[wordIndex++]));
336 auto linkageAttr = opBuilder.getAttr<::mlir::spirv::LinkageAttributesAttr>(
337 StringAttr::get(context, linkageName), linkageTypeAttr);
338 decorations[words[0]].set(symbol, dyn_cast<Attribute>(linkageAttr));
341 case spirv::Decoration::Aliased:
342 case spirv::Decoration::AliasedPointer:
343 case spirv::Decoration::Block:
344 case spirv::Decoration::BufferBlock:
345 case spirv::Decoration::Flat:
346 case spirv::Decoration::NonReadable:
347 case spirv::Decoration::NonWritable:
348 case spirv::Decoration::NoPerspective:
349 case spirv::Decoration::NoSignedWrap:
350 case spirv::Decoration::NoUnsignedWrap:
351 case spirv::Decoration::RelaxedPrecision:
352 case spirv::Decoration::Restrict:
353 case spirv::Decoration::RestrictPointer:
354 case spirv::Decoration::NoContraction:
355 case spirv::Decoration::Constant:
356 case spirv::Decoration::Invariant:
357 case spirv::Decoration::Patch:
358 case spirv::Decoration::Coherent:
359 case spirv::Decoration::Volatile:
360 if (words.size() != 2) {
361 return emitError(unknownLoc,
"OpDecorate with ")
362 << decorationName <<
" needs a single target <id>";
364 decorations[words[0]].set(symbol, opBuilder.getUnitAttr());
366 case spirv::Decoration::CacheControlLoadINTEL: {
368 CacheControlLoadINTELAttr, LoadCacheControlAttr, LoadCacheControl>(
369 unknownLoc, opBuilder, decorations, words, symbol, decorationName,
375 case spirv::Decoration::CacheControlStoreINTEL: {
377 CacheControlStoreINTELAttr, StoreCacheControlAttr, StoreCacheControl>(
378 unknownLoc, opBuilder, decorations, words, symbol, decorationName,
384 case spirv::Decoration::AlignmentId:
385 case spirv::Decoration::MaxByteOffsetId:
386 case spirv::Decoration::CounterBuffer:
387 if (words.size() != 3) {
388 return emitError(unknownLoc,
"OpDecorateId with ")
389 << decorationName <<
" needs a single <id> operand";
391 pendingIdDecorations.push_back({words[0],
392 static_cast<spirv::Decoration
>(words[1]),
393 words[2], unknownLoc});
396 return emitError(unknownLoc,
"unhandled Decoration : '") << decorationName;
401LogicalResult spirv::Deserializer::resolveDeferredIdDecorations() {
402 for (
const DeferredIdDecoration &entry : pendingIdDecorations) {
403 StringRef decorationName = stringifyDecoration(entry.decoration);
404 StringAttr symbol = getSymbolDecoration(decorationName);
408 StringRef operandSymName;
409 if (spirv::GlobalVariableOp varOp =
410 globalVariableMap.lookup(entry.operandID))
411 operandSymName = varOp.getSymName();
412 else if (spirv::SpecConstantOp specOp =
413 specConstMap.lookup(entry.operandID))
414 operandSymName = specOp.getSymName();
416 return emitError(entry.loc,
"OpDecorateId with ")
417 << decorationName <<
" references <id> " << entry.operandID
418 <<
" which is not a global variable or specialization constant";
425 Operation *targetOp =
nullptr;
426 if (spirv::GlobalVariableOp varOp =
427 globalVariableMap.lookup(entry.targetID))
429 else if (spirv::SpecConstantOp specOp = specConstMap.lookup(entry.targetID))
431 else if (spirv::FuncOp fnOp = funcMap.lookup(entry.targetID))
433 else if (Value v = valueMap.lookup(entry.targetID))
434 targetOp = v.getDefiningOp();
437 return emitError(entry.loc,
"OpDecorateId with ")
438 << decorationName <<
" references unknown target <id> "
441 targetOp->
setAttr(symbol, symRef);
447spirv::Deserializer::processMemberDecoration(ArrayRef<uint32_t> words) {
449 if (words.size() < 3) {
451 "OpMemberDecorate must have at least 3 operands");
454 auto decoration =
static_cast<spirv::Decoration
>(words[2]);
455 if (decoration == spirv::Decoration::Offset && words.size() != 4) {
457 " missing offset specification in OpMemberDecorate with "
458 "Offset decoration");
460 ArrayRef<uint32_t> decorationOperands;
461 if (words.size() > 3) {
462 decorationOperands = words.slice(3);
464 memberDecorationMap[words[0]][words[1]][decoration] = decorationOperands;
468LogicalResult spirv::Deserializer::processMemberName(ArrayRef<uint32_t> words) {
469 if (words.size() < 3) {
470 return emitError(unknownLoc,
"OpMemberName must have at least 3 operands");
472 unsigned wordIndex = 2;
474 if (wordIndex != words.size()) {
476 "unexpected trailing words in OpMemberName instruction");
478 memberNameMap[words[0]][words[1]] = name;
484 if (!decorations.contains(argID)) {
485 argAttrs[argIndex] = DictionaryAttr::get(context, {});
489 spirv::DecorationAttr foundDecorationAttr;
491 for (
auto decoration :
492 {spirv::Decoration::Aliased, spirv::Decoration::Restrict,
493 spirv::Decoration::AliasedPointer,
494 spirv::Decoration::RestrictPointer}) {
496 if (decAttr.getName() !=
500 if (foundDecorationAttr)
502 "more than one Aliased/Restrict decorations for "
503 "function argument with result <id> ")
506 foundDecorationAttr = spirv::DecorationAttr::get(context, decoration);
511 spirv::Decoration::RelaxedPrecision))) {
516 if (foundDecorationAttr)
517 return emitError(unknownLoc,
"already found a decoration for function "
518 "argument with result <id> ")
521 foundDecorationAttr = spirv::DecorationAttr::get(
522 context, spirv::Decoration::RelaxedPrecision);
526 if (!foundDecorationAttr)
527 return emitError(unknownLoc,
"unimplemented decoration support for "
528 "function argument with result <id> ")
531 NamedAttribute attr(StringAttr::get(context, spirv::DecorationAttr::name),
532 foundDecorationAttr);
533 argAttrs[argIndex] = DictionaryAttr::get(context, attr);
540 return emitError(unknownLoc,
"found function inside function");
544 if (operands.size() != 4) {
545 return emitError(unknownLoc,
"OpFunction must have 4 parameters");
549 return emitError(unknownLoc,
"undefined result type from <id> ")
553 uint32_t fnID = operands[1];
554 if (funcMap.count(fnID)) {
555 return emitError(unknownLoc,
"duplicate function definition/declaration");
558 auto fnControl = spirv::symbolizeFunctionControl(operands[2]);
560 return emitError(unknownLoc,
"unknown Function Control: ") << operands[2];
564 if (!fnType || !isa<FunctionType>(fnType)) {
565 return emitError(unknownLoc,
"unknown function type from <id> ")
568 auto functionType = cast<FunctionType>(fnType);
570 if ((
isVoidType(resultType) && functionType.getNumResults() != 0) ||
571 (functionType.getNumResults() == 1 &&
572 functionType.getResult(0) != resultType)) {
573 return emitError(unknownLoc,
"mismatch in function type ")
574 << functionType <<
" and return type " << resultType <<
" specified";
578 auto funcOp = spirv::FuncOp::create(opBuilder, unknownLoc, fnName,
579 functionType, fnControl.value());
581 if (decorations.count(fnID)) {
582 for (
auto attr : decorations[fnID].getAttrs()) {
583 funcOp->setAttr(attr.getName(), attr.getValue());
586 curFunction = funcMap[fnID] = funcOp;
587 auto *entryBlock = funcOp.addEntryBlock();
590 <<
"//===-------------------------------------------===//\n";
591 logger.startLine() <<
"[fn] name: " << fnName <<
"\n";
592 logger.startLine() <<
"[fn] type: " << fnType <<
"\n";
593 logger.startLine() <<
"[fn] ID: " << fnID <<
"\n";
594 logger.startLine() <<
"[fn] entry block: " << entryBlock <<
"\n";
599 argAttrs.resize(functionType.getNumInputs());
602 if (functionType.getNumInputs()) {
603 for (
size_t i = 0, e = functionType.getNumInputs(); i != e; ++i) {
604 auto argType = functionType.getInput(i);
605 spirv::Opcode opcode = spirv::Opcode::OpNop;
608 spirv::Opcode::OpFunctionParameter))) {
611 if (opcode != spirv::Opcode::OpFunctionParameter) {
614 "missing OpFunctionParameter instruction for argument ")
617 if (operands.size() != 2) {
620 "expected result type and result <id> for OpFunctionParameter");
622 auto argDefinedType =
getType(operands[0]);
623 if (!argDefinedType || argDefinedType != argType) {
625 "mismatch in argument type between function type "
627 << functionType <<
" and argument type definition "
628 << argDefinedType <<
" at argument " << i;
631 return emitError(unknownLoc,
"duplicate definition of result <id> ")
638 auto argValue = funcOp.getArgument(i);
639 valueMap[operands[1]] = argValue;
643 if (llvm::any_of(argAttrs, [](
Attribute attr) {
644 auto argAttr = cast<DictionaryAttr>(attr);
645 return !argAttr.empty();
647 funcOp.setArgAttrsAttr(ArrayAttr::get(context, argAttrs));
652 auto linkageAttr = funcOp.getLinkageAttributes();
653 auto hasImportLinkage =
654 linkageAttr && (linkageAttr.value().getLinkageType().
getValue() ==
655 spirv::LinkageType::Import);
656 if (hasImportLinkage)
663 spirv::Opcode opcode = spirv::Opcode::OpNop;
672 spirv::Opcode::OpFunctionEnd))) {
675 if (opcode == spirv::Opcode::OpFunctionEnd) {
678 if (opcode != spirv::Opcode::OpLabel) {
679 return emitError(unknownLoc,
"a basic block must start with OpLabel");
681 if (instOperands.size() != 1) {
682 return emitError(unknownLoc,
"OpLabel should only have result <id>");
684 blockMap[instOperands[0]] = entryBlock;
692 spirv::Opcode::OpFunctionEnd)) &&
693 opcode != spirv::Opcode::OpFunctionEnd) {
698 if (opcode != spirv::Opcode::OpFunctionEnd) {
708 if (!operands.empty()) {
709 return emitError(unknownLoc,
"unexpected operands for OpFunctionEnd");
720 curFunction = std::nullopt;
725 <<
"//===-------------------------------------------===//\n";
732 if (operands.size() < 2) {
734 "missing graph defintion in OpGraphEntryPointARM");
737 unsigned wordIndex = 0;
738 uint32_t graphID = operands[wordIndex++];
739 if (!graphMap.contains(graphID)) {
741 "missing graph definition/declaration with id ")
745 spirv::GraphARMOp graphARM = graphMap[graphID];
747 graphARM.setSymName(name);
748 graphARM.setEntryPoint(
true);
751 for (
int64_t size = operands.size(); wordIndex < size; ++wordIndex) {
753 interface.push_back(SymbolRefAttr::get(arg.getOperation()));
755 return emitError(unknownLoc,
"undefined result <id> ")
756 << operands[wordIndex] <<
" while decoding OpGraphEntryPoint";
762 opBuilder.setInsertionPoint(graphARM);
763 spirv::GraphEntryPointARMOp::create(
764 opBuilder, unknownLoc, SymbolRefAttr::get(opBuilder.getContext(), name),
765 opBuilder.getArrayAttr(interface));
773 return emitError(unknownLoc,
"found graph inside graph");
776 if (operands.size() < 2) {
777 return emitError(unknownLoc,
"OpGraphARM must have at least 2 parameters");
781 if (!type || !isa<GraphType>(type)) {
782 return emitError(unknownLoc,
"unknown graph type from <id> ")
785 auto graphType = cast<GraphType>(type);
786 if (graphType.getNumResults() <= 0) {
787 return emitError(unknownLoc,
"expected at least one result");
790 uint32_t graphID = operands[1];
791 if (graphMap.count(graphID)) {
792 return emitError(unknownLoc,
"duplicate graph definition/declaration");
797 spirv::GraphARMOp::create(opBuilder, unknownLoc, graphName, graphType);
798 curGraph = graphMap[graphID] = graphOp;
799 Block *entryBlock = graphOp.addEntryBlock();
802 <<
"//===-------------------------------------------===//\n";
803 logger.startLine() <<
"[graph] name: " << graphName <<
"\n";
804 logger.startLine() <<
"[graph] type: " << graphType <<
"\n";
805 logger.startLine() <<
"[graph] ID: " << graphID <<
"\n";
806 logger.startLine() <<
"[graph] entry block: " << entryBlock <<
"\n";
811 for (
auto [
index, argType] : llvm::enumerate(graphType.getInputs())) {
812 spirv::Opcode opcode;
815 spirv::Opcode::OpGraphInputARM))) {
818 if (operands.size() != 3) {
819 return emitError(unknownLoc,
"expected result type, result <id> and "
820 "input index for OpGraphInputARM");
824 if (!argDefinedType) {
825 return emitError(unknownLoc,
"unknown operand type <id> ") << operands[0];
828 if (argDefinedType != argType) {
830 "mismatch in argument type between graph type "
832 << graphType <<
" and argument type definition " << argDefinedType
833 <<
" at argument " <<
index;
836 return emitError(unknownLoc,
"duplicate definition of result <id> ")
841 if (!inputIndexAttr) {
843 "unable to read inputIndex value from constant op ")
846 BlockArgument argValue = graphOp.getArgument(inputIndexAttr.getInt());
847 valueMap[operands[1]] = argValue;
850 graphOutputs.resize(graphType.getNumResults());
856 blockMap[graphID] = entryBlock;
863 spirv::Opcode opcode;
873 }
while (opcode != spirv::Opcode::OpGraphEndARM);
880 if (operands.size() != 2) {
883 "expected value id and output index for OpGraphSetOutputARM");
886 uint32_t
id = operands[0];
889 return emitError(unknownLoc,
"could not find result <id> ") << id;
893 if (!outputIndexAttr) {
895 "unable to read outputIndex value from constant op ")
898 graphOutputs[outputIndexAttr.getInt()] = value;
905 spirv::GraphOutputsARMOp::create(opBuilder, unknownLoc, graphOutputs);
908 if (!operands.empty()) {
909 return emitError(unknownLoc,
"unexpected operands for OpGraphEndARM");
913 curGraph = std::nullopt;
914 graphOutputs.clear();
919 <<
"//===-------------------------------------------===//\n";
924std::optional<std::pair<Attribute, Type>>
926 auto constIt = constantMap.find(
id);
927 if (constIt == constantMap.end())
929 return constIt->getSecond();
932std::optional<std::pair<Attribute, Type>>
934 if (
auto it = constantCompositeReplicateMap.find(
id);
935 it != constantCompositeReplicateMap.end())
940std::optional<spirv::SpecConstOperationMaterializationInfo>
942 auto constIt = specConstOperationMap.find(
id);
943 if (constIt == specConstOperationMap.end())
945 return constIt->getSecond();
949 auto funcName = nameMap.lookup(
id).str();
950 if (funcName.empty()) {
951 funcName =
"spirv_fn_" + std::to_string(
id);
957 std::string graphName = nameMap.lookup(
id).str();
958 if (graphName.empty()) {
959 graphName =
"spirv_graph_" + std::to_string(
id);
965 auto constName = nameMap.lookup(
id).str();
966 if (constName.empty()) {
967 constName =
"spirv_spec_const_" + std::to_string(
id);
974 TypedAttr defaultValue) {
976 auto op = spirv::SpecConstantOp::create(opBuilder, unknownLoc, symName,
978 if (decorations.count(resultID)) {
979 for (
auto attr : decorations[resultID].getAttrs())
980 op->setAttr(attr.getName(), attr.getValue());
982 specConstMap[resultID] = op;
986std::optional<spirv::GraphConstantARMOpMaterializationInfo>
988 auto graphConstIt = graphConstantMap.find(
id);
989 if (graphConstIt == graphConstantMap.end())
991 return graphConstIt->getSecond();
996 unsigned wordIndex = 0;
997 if (operands.size() < 3) {
1000 "OpVariable needs at least 3 operands, type, <id> and storage class");
1004 auto type =
getType(operands[wordIndex]);
1006 return emitError(unknownLoc,
"unknown result type <id> : ")
1007 << operands[wordIndex];
1009 auto ptrType = dyn_cast<spirv::PointerType>(type);
1012 "expected a result type <id> to be a spirv.ptr, found : ")
1018 auto variableID = operands[wordIndex];
1019 auto variableName = nameMap.lookup(variableID).str();
1020 if (variableName.empty()) {
1021 variableName =
"spirv_var_" + std::to_string(variableID);
1026 auto storageClass =
static_cast<spirv::StorageClass
>(operands[wordIndex]);
1027 if (ptrType.getStorageClass() != storageClass) {
1028 return emitError(unknownLoc,
"mismatch in storage class of pointer type ")
1029 << type <<
" and that specified in OpVariable instruction : "
1030 << stringifyStorageClass(storageClass);
1037 if (wordIndex < operands.size()) {
1047 return emitError(unknownLoc,
"unknown <id> ")
1048 << operands[wordIndex] <<
"used as initializer";
1050 initializer = SymbolRefAttr::get(op);
1053 if (wordIndex != operands.size()) {
1055 "found more operands than expected when deserializing "
1056 "OpVariable instruction, only ")
1057 << wordIndex <<
" of " << operands.size() <<
" processed";
1060 auto varOp = spirv::GlobalVariableOp::create(
1061 opBuilder, loc, TypeAttr::get(type),
1062 opBuilder.getStringAttr(variableName), initializer);
1065 if (decorations.count(variableID)) {
1066 for (
auto attr : decorations[variableID].getAttrs())
1067 varOp->setAttr(attr.getName(), attr.getValue());
1069 globalVariableMap[variableID] = varOp;
1078 return dyn_cast<IntegerAttr>(constInfo->first);
1082 if (operands.size() < 2) {
1083 return emitError(unknownLoc,
"OpName needs at least 2 operands");
1086 unsigned wordIndex = 1;
1088 if (wordIndex != operands.size()) {
1090 "unexpected trailing words in OpName instruction");
1095 nameMap.emplace_or_assign(operands[0], name);
1106 if (operands.empty()) {
1107 return emitError(unknownLoc,
"type instruction with opcode ")
1108 << spirv::stringifyOpcode(opcode) <<
" needs at least one <id>";
1113 if (typeMap.count(operands[0])) {
1114 return emitError(unknownLoc,
"duplicate definition for result <id> ")
1119 case spirv::Opcode::OpTypeVoid:
1120 if (operands.size() != 1)
1121 return emitError(unknownLoc,
"OpTypeVoid must have no parameters");
1122 typeMap[operands[0]] = opBuilder.getNoneType();
1124 case spirv::Opcode::OpTypeBool:
1125 if (operands.size() != 1)
1126 return emitError(unknownLoc,
"OpTypeBool must have no parameters");
1127 typeMap[operands[0]] = opBuilder.getI1Type();
1129 case spirv::Opcode::OpTypeInt: {
1130 if (operands.size() != 3)
1132 unknownLoc,
"OpTypeInt must have bitwidth and signedness parameters");
1141 auto sign = operands[2] == 1 ? IntegerType::SignednessSemantics::Signed
1142 : IntegerType::SignednessSemantics::Signless;
1143 typeMap[operands[0]] = IntegerType::get(context, operands[1], sign);
1145 case spirv::Opcode::OpTypeFloat: {
1146 if (operands.size() != 2 && operands.size() != 3)
1148 "OpTypeFloat expects either 2 operands (type, bitwidth) "
1149 "or 3 operands (type, bitwidth, encoding), but got ")
1151 uint32_t bitWidth = operands[1];
1154 if (operands.size() == 2) {
1157 floatTy = opBuilder.getF16Type();
1160 floatTy = opBuilder.getF32Type();
1163 floatTy = opBuilder.getF64Type();
1166 return emitError(unknownLoc,
"unsupported OpTypeFloat bitwidth: ")
1171 if (operands.size() == 3) {
1172 if (spirv::FPEncoding(operands[2]) == spirv::FPEncoding::BFloat16KHR &&
1174 floatTy = opBuilder.getBF16Type();
1175 else if (spirv::FPEncoding(operands[2]) ==
1176 spirv::FPEncoding::Float8E4M3EXT &&
1178 floatTy = opBuilder.getF8E4M3FNType();
1179 else if (spirv::FPEncoding(operands[2]) ==
1180 spirv::FPEncoding::Float8E5M2EXT &&
1182 floatTy = opBuilder.getF8E5M2Type();
1184 return emitError(unknownLoc,
"unsupported OpTypeFloat FP encoding: ")
1185 << operands[2] <<
" and bitWidth " << bitWidth;
1188 typeMap[operands[0]] = floatTy;
1190 case spirv::Opcode::OpTypeVector: {
1191 if (operands.size() != 3) {
1194 "OpTypeVector must have element type and count parameters");
1198 return emitError(unknownLoc,
"OpTypeVector references undefined <id> ")
1201 typeMap[operands[0]] = VectorType::get({operands[2]}, elementTy);
1203 case spirv::Opcode::OpTypePointer: {
1206 case spirv::Opcode::OpTypeArray:
1208 case spirv::Opcode::OpTypeCooperativeMatrixKHR:
1210 case spirv::Opcode::OpTypeFunction:
1212 case spirv::Opcode::OpTypeImage:
1214 case spirv::Opcode::OpTypeSampler:
1216 case spirv::Opcode::OpTypeNamedBarrier:
1218 case spirv::Opcode::OpTypeSampledImage:
1220 case spirv::Opcode::OpTypeRuntimeArray:
1222 case spirv::Opcode::OpTypeStruct:
1224 case spirv::Opcode::OpTypeMatrix:
1226 case spirv::Opcode::OpTypeTensorARM:
1228 case spirv::Opcode::OpTypeGraphARM:
1231 return emitError(unknownLoc,
"unhandled type instruction");
1238 if (operands.size() != 3)
1239 return emitError(unknownLoc,
"OpTypePointer must have two parameters");
1241 auto pointeeType =
getType(operands[2]);
1243 return emitError(unknownLoc,
"unknown OpTypePointer pointee type <id> ")
1246 uint32_t typePointerID = operands[0];
1247 auto storageClass =
static_cast<spirv::StorageClass
>(operands[1]);
1250 for (
auto *deferredStructIt = std::begin(deferredStructTypesInfos);
1251 deferredStructIt != std::end(deferredStructTypesInfos);) {
1252 for (
auto *unresolvedMemberIt =
1253 std::begin(deferredStructIt->unresolvedMemberTypes);
1254 unresolvedMemberIt !=
1255 std::end(deferredStructIt->unresolvedMemberTypes);) {
1256 if (unresolvedMemberIt->first == typePointerID) {
1260 deferredStructIt->memberTypes[unresolvedMemberIt->second] =
1261 typeMap[typePointerID];
1262 unresolvedMemberIt =
1263 deferredStructIt->unresolvedMemberTypes.erase(unresolvedMemberIt);
1265 ++unresolvedMemberIt;
1269 if (deferredStructIt->unresolvedMemberTypes.empty()) {
1271 auto structType = deferredStructIt->deferredStructType;
1273 assert(structType &&
"expected a spirv::StructType");
1274 assert(structType.isIdentified() &&
"expected an indentified struct");
1276 if (failed(structType.trySetBody(
1277 deferredStructIt->memberTypes, deferredStructIt->offsetInfo,
1278 deferredStructIt->memberDecorationsInfo,
1279 deferredStructIt->structDecorationsInfo)))
1282 deferredStructIt = deferredStructTypesInfos.erase(deferredStructIt);
1293 if (operands.size() != 3) {
1295 "OpTypeArray must have element type and count parameters");
1300 return emitError(unknownLoc,
"OpTypeArray references undefined <id> ")
1308 return emitError(unknownLoc,
"OpTypeArray count <id> ")
1309 << operands[2] <<
"can only come from normal constant right now";
1312 if (
auto intVal = dyn_cast<IntegerAttr>(countInfo->first)) {
1313 count = intVal.getValue().getZExtValue();
1315 return emitError(unknownLoc,
"OpTypeArray count must come from a "
1316 "scalar integer constant instruction");
1320 elementTy, count, typeDecorations.lookup(operands[0]));
1326 assert(!operands.empty() &&
"No operands for processing function type");
1327 if (operands.size() == 1) {
1328 return emitError(unknownLoc,
"missing return type for OpTypeFunction");
1330 auto returnType =
getType(operands[1]);
1332 return emitError(unknownLoc,
"unknown return type in OpTypeFunction");
1335 for (
size_t i = 2, e = operands.size(); i < e; ++i) {
1336 auto ty =
getType(operands[i]);
1338 return emitError(unknownLoc,
"unknown argument type in OpTypeFunction");
1340 argTypes.push_back(ty);
1346 typeMap[operands[0]] = FunctionType::get(context, argTypes, returnTypes);
1352 if (operands.size() != 6) {
1354 "OpTypeCooperativeMatrixKHR must have element type, "
1355 "scope, row and column parameters, and use");
1361 "OpTypeCooperativeMatrixKHR references undefined <id> ")
1365 std::optional<spirv::Scope> scope =
1370 "OpTypeCooperativeMatrixKHR references undefined scope <id> ")
1379 return emitError(unknownLoc,
"OpTypeCooperativeMatrixKHR `Rows` references "
1380 "undefined constant <id> ")
1384 return emitError(unknownLoc,
"OpTypeCooperativeMatrixKHR `Columns` "
1385 "references undefined constant <id> ")
1389 return emitError(unknownLoc,
"OpTypeCooperativeMatrixKHR `Use` references "
1390 "undefined constant <id> ")
1393 unsigned rows = rowsAttr.getInt();
1394 unsigned columns = columnsAttr.getInt();
1396 std::optional<spirv::CooperativeMatrixUseKHR> use =
1397 spirv::symbolizeCooperativeMatrixUseKHR(useAttr.getInt());
1401 "OpTypeCooperativeMatrixKHR references undefined use <id> ")
1405 typeMap[operands[0]] =
1412 if (operands.size() != 2) {
1413 return emitError(unknownLoc,
"OpTypeRuntimeArray must have two operands");
1418 "OpTypeRuntimeArray references undefined <id> ")
1422 memberType, typeDecorations.lookup(operands[0]));
1430 if (operands.empty()) {
1431 return emitError(unknownLoc,
"OpTypeStruct must have at least result <id>");
1434 if (operands.size() == 1) {
1436 typeMap[operands[0]] =
1445 for (
auto op : llvm::drop_begin(operands, 1)) {
1447 bool typeForwardPtr = (typeForwardPointerIDs.count(op) != 0);
1449 if (!memberType && !typeForwardPtr)
1450 return emitError(unknownLoc,
"OpTypeStruct references undefined <id> ")
1454 unresolvedMemberTypes.emplace_back(op, memberTypes.size());
1456 memberTypes.push_back(memberType);
1461 if (memberDecorationMap.count(operands[0])) {
1462 auto &allMemberDecorations = memberDecorationMap[operands[0]];
1463 for (
auto memberIndex : llvm::seq<uint32_t>(0, memberTypes.size())) {
1464 if (allMemberDecorations.count(memberIndex)) {
1465 for (
auto &memberDecoration : allMemberDecorations[memberIndex]) {
1467 if (memberDecoration.first == spirv::Decoration::Offset) {
1469 if (offsetInfo.empty()) {
1470 offsetInfo.resize(memberTypes.size());
1472 offsetInfo[memberIndex] = memberDecoration.second[0];
1474 auto intType = mlir::IntegerType::get(context, 32);
1475 if (!memberDecoration.second.empty()) {
1476 memberDecorationsInfo.emplace_back(
1477 memberIndex, memberDecoration.first,
1478 IntegerAttr::get(intType, memberDecoration.second[0]));
1480 memberDecorationsInfo.emplace_back(
1481 memberIndex, memberDecoration.first, UnitAttr::get(context));
1490 if (decorations.count(operands[0])) {
1493 std::optional<spirv::Decoration> decoration = spirv::symbolizeDecoration(
1494 llvm::convertToCamelFromSnakeCase(decorationAttr.getName(),
true));
1495 assert(decoration.has_value());
1496 structDecorationsInfo.emplace_back(decoration.value(),
1497 decorationAttr.getValue());
1501 uint32_t structID = operands[0];
1502 std::string structIdentifier = nameMap.lookup(structID).str();
1504 if (structIdentifier.empty()) {
1505 assert(unresolvedMemberTypes.empty() &&
1506 "didn't expect unresolved member types");
1508 memberTypes, offsetInfo, memberDecorationsInfo, structDecorationsInfo);
1511 typeMap[structID] = structTy;
1513 if (!unresolvedMemberTypes.empty())
1514 deferredStructTypesInfos.push_back(
1515 {structTy, unresolvedMemberTypes, memberTypes, offsetInfo,
1516 memberDecorationsInfo, structDecorationsInfo});
1517 else if (failed(structTy.trySetBody(memberTypes, offsetInfo,
1518 memberDecorationsInfo,
1519 structDecorationsInfo)))
1530 if (operands.size() != 3) {
1532 return emitError(unknownLoc,
"OpTypeMatrix must have 3 operands"
1533 " (result_id, column_type, and column_count)");
1539 "OpTypeMatrix references undefined column type.")
1543 uint32_t colsCount = operands[2];
1550 unsigned size = operands.size();
1551 if (size < 2 || size > 4)
1552 return emitError(unknownLoc,
"OpTypeTensorARM must have 2-4 operands "
1553 "(result_id, element_type, (rank), (shape)) ")
1559 "OpTypeTensorARM references undefined element type ")
1569 return emitError(unknownLoc,
"OpTypeTensorARM rank must come from a "
1570 "scalar integer constant instruction");
1571 unsigned rank = rankAttr.getValue().getZExtValue();
1578 std::optional<std::pair<Attribute, Type>> shapeInfo =
1581 return emitError(unknownLoc,
"OpTypeTensorARM shape must come from a "
1582 "constant instruction of type OpTypeArray");
1584 ArrayAttr shapeArrayAttr = dyn_cast<ArrayAttr>(shapeInfo->first);
1586 for (
auto dimAttr : shapeArrayAttr.getValue()) {
1587 auto dimIntAttr = dyn_cast<IntegerAttr>(dimAttr);
1589 return emitError(unknownLoc,
"OpTypeTensorARM shape has an invalid "
1591 shape.push_back(dimIntAttr.getValue().getSExtValue());
1599 unsigned size = operands.size();
1601 return emitError(unknownLoc,
"OpTypeGraphARM must have at least 2 operands "
1602 "(result_id, num_inputs, (inout0_type, "
1603 "inout1_type, ...))")
1606 uint32_t numInputs = operands[1];
1609 for (
unsigned i = 2; i < size; ++i) {
1613 "OpTypeGraphARM references undefined element type.")
1616 if (i - 2 >= numInputs) {
1617 returnTypes.push_back(inOutTy);
1619 argTypes.push_back(inOutTy);
1622 typeMap[operands[0]] = GraphType::get(context, argTypes, returnTypes);
1628 if (operands.size() != 2)
1630 "OpTypeForwardPointer instruction must have two operands");
1632 typeForwardPointerIDs.insert(operands[0]);
1642 if (operands.size() != 8)
1645 "OpTypeImage with non-eight operands are not supported yet");
1649 return emitError(unknownLoc,
"OpTypeImage references undefined <id>: ")
1652 auto dim = spirv::symbolizeDim(operands[2]);
1654 return emitError(unknownLoc,
"unknown Dim for OpTypeImage: ")
1657 auto depthInfo = spirv::symbolizeImageDepthInfo(operands[3]);
1659 return emitError(unknownLoc,
"unknown Depth for OpTypeImage: ")
1662 auto arrayedInfo = spirv::symbolizeImageArrayedInfo(operands[4]);
1664 return emitError(unknownLoc,
"unknown Arrayed for OpTypeImage: ")
1667 auto samplingInfo = spirv::symbolizeImageSamplingInfo(operands[5]);
1669 return emitError(unknownLoc,
"unknown MS for OpTypeImage: ") << operands[5];
1671 auto samplerUseInfo = spirv::symbolizeImageSamplerUseInfo(operands[6]);
1672 if (!samplerUseInfo)
1673 return emitError(unknownLoc,
"unknown Sampled for OpTypeImage: ")
1676 auto format = spirv::symbolizeImageFormat(operands[7]);
1678 return emitError(unknownLoc,
"unknown Format for OpTypeImage: ")
1682 elementTy, dim.value(), depthInfo.value(), arrayedInfo.value(),
1683 samplingInfo.value(), samplerUseInfo.value(), format.value());
1689 if (operands.size() != 2)
1690 return emitError(unknownLoc,
"OpTypeSampledImage must have two operands");
1695 "OpTypeSampledImage references undefined <id>: ")
1704 if (operands.size() != 1)
1705 return emitError(unknownLoc,
"OpTypeSampler must have no parameters");
1713 if (operands.size() != 1)
1714 return emitError(unknownLoc,
"OpTypeNamedBarrier must have no parameters");
1726 StringRef opname = isSpec ?
"OpSpecConstant" :
"OpConstant";
1728 if (operands.size() < 2) {
1730 << opname <<
" must have type <id> and result <id>";
1732 if (operands.size() < 3) {
1734 << opname <<
" must have at least 1 more parameter";
1739 return emitError(unknownLoc,
"undefined result type from <id> ")
1743 auto checkOperandSizeForBitwidth = [&](
unsigned bitwidth) -> LogicalResult {
1744 if (bitwidth == 64) {
1745 if (operands.size() == 4) {
1749 << opname <<
" should have 2 parameters for 64-bit values";
1751 if (bitwidth <= 32) {
1752 if (operands.size() == 3) {
1758 <<
" should have 1 parameter for values with no more than 32 bits";
1760 return emitError(unknownLoc,
"unsupported OpConstant bitwidth: ")
1764 auto resultID = operands[1];
1766 if (
auto intType = dyn_cast<IntegerType>(resultType)) {
1767 auto bitwidth = intType.getWidth();
1768 if (failed(checkOperandSizeForBitwidth(bitwidth))) {
1773 if (bitwidth == 64) {
1780 } words = {operands[2], operands[3]};
1781 value = APInt(64, llvm::bit_cast<uint64_t>(words),
true);
1782 }
else if (bitwidth <= 32) {
1783 value = APInt(bitwidth, operands[2],
true,
1787 auto attr = opBuilder.getIntegerAttr(intType, value);
1794 constantMap.try_emplace(resultID, attr, intType);
1800 if (
auto floatType = dyn_cast<FloatType>(resultType)) {
1801 auto bitwidth = floatType.getWidth();
1802 if (failed(checkOperandSizeForBitwidth(bitwidth))) {
1807 if (floatType.isF64()) {
1814 } words = {operands[2], operands[3]};
1815 value = APFloat(llvm::bit_cast<double>(words));
1816 }
else if (floatType.isF32()) {
1817 value = APFloat(llvm::bit_cast<float>(operands[2]));
1818 }
else if (floatType.isF16()) {
1819 APInt data(16, operands[2]);
1820 value = APFloat(APFloat::IEEEhalf(), data);
1821 }
else if (floatType.isBF16()) {
1822 APInt data(16, operands[2]);
1823 value = APFloat(APFloat::BFloat(), data);
1824 }
else if (floatType.isF8E4M3FN()) {
1825 APInt data(8, operands[2]);
1826 value = APFloat(APFloat::Float8E4M3FN(), data);
1827 }
else if (floatType.isF8E5M2()) {
1828 APInt data(8, operands[2]);
1829 value = APFloat(APFloat::Float8E5M2(), data);
1832 auto attr = opBuilder.getFloatAttr(floatType, value);
1838 constantMap.try_emplace(resultID, attr, floatType);
1844 return emitError(unknownLoc,
"OpConstant can only generate values of "
1845 "scalar integer or floating-point type");
1850 if (operands.size() != 2) {
1852 << (isSpec ?
"Spec" :
"") <<
"Constant"
1853 << (isTrue ?
"True" :
"False")
1854 <<
" must have type <id> and result <id>";
1857 auto attr = opBuilder.getBoolAttr(isTrue);
1858 auto resultID = operands[1];
1864 constantMap.try_emplace(resultID, attr, opBuilder.getI1Type());
1872 if (operands.size() < 2) {
1874 "OpConstantComposite must have type <id> and result <id>");
1876 if (operands.size() < 3) {
1878 "OpConstantComposite must have at least 1 parameter");
1883 return emitError(unknownLoc,
"undefined result type from <id> ")
1888 elements.reserve(operands.size() - 2);
1889 for (
unsigned i = 2, e = operands.size(); i < e; ++i) {
1892 return emitError(unknownLoc,
"OpConstantComposite component <id> ")
1893 << operands[i] <<
" must come from a normal constant";
1895 elements.push_back(elementInfo->first);
1898 auto resultID = operands[1];
1899 if (
auto tensorType = dyn_cast<TensorArmType>(resultType)) {
1902 if (
auto denseElemAttr = dyn_cast<DenseElementsAttr>(element)) {
1903 for (
auto value : denseElemAttr.getValues<
Attribute>())
1904 flattenedElems.push_back(value);
1906 flattenedElems.push_back(element);
1910 constantMap.try_emplace(resultID, attr, tensorType);
1911 }
else if (
auto shapedType = dyn_cast<ShapedType>(resultType)) {
1915 constantMap.try_emplace(resultID, attr, shapedType);
1916 }
else if (isa<spirv::ArrayType, spirv::StructType>(resultType)) {
1917 auto attr = opBuilder.getArrayAttr(elements);
1918 constantMap.try_emplace(resultID, attr, resultType);
1920 return emitError(unknownLoc,
"unsupported OpConstantComposite type: ")
1929 if (operands.size() != 3) {
1932 "OpConstantCompositeReplicateEXT expects 3 operands but found ")
1938 return emitError(unknownLoc,
"undefined result type from <id> ")
1942 auto compositeType = dyn_cast<CompositeType>(resultType);
1943 if (!compositeType) {
1945 "result type from <id> is not a composite type")
1949 uint32_t resultID = operands[1];
1950 uint32_t constantID = operands[2];
1952 std::optional<std::pair<Attribute, Type>> constantInfo =
1954 if (constantInfo.has_value()) {
1955 constantCompositeReplicateMap.try_emplace(
1956 resultID, constantInfo.value().first, resultType);
1960 std::optional<std::pair<Attribute, Type>> replicatedConstantCompositeInfo =
1962 if (replicatedConstantCompositeInfo.has_value()) {
1963 constantCompositeReplicateMap.try_emplace(
1964 resultID, replicatedConstantCompositeInfo.value().first, resultType);
1968 return emitError(unknownLoc,
"OpConstantCompositeReplicateEXT operand <id> ")
1970 <<
" must come from a normal constant or a "
1971 "OpConstantCompositeReplicateEXT";
1976 if (operands.size() < 2) {
1979 "OpSpecConstantComposite must have type <id> and result <id>");
1981 if (operands.size() < 3) {
1983 "OpSpecConstantComposite must have at least 1 parameter");
1988 return emitError(unknownLoc,
"undefined result type from <id> ")
1992 auto resultID = operands[1];
1996 elements.reserve(operands.size() - 2);
1997 for (
unsigned i = 2, e = operands.size(); i < e; ++i) {
1999 elements.push_back(SymbolRefAttr::get(elementInfo));
2002 auto op = spirv::SpecConstantCompositeOp::create(
2003 opBuilder, unknownLoc, TypeAttr::get(resultType), symName,
2004 opBuilder.getArrayAttr(elements));
2005 specConstCompositeMap[resultID] = op;
2012 if (operands.size() != 3) {
2013 return emitError(unknownLoc,
"OpSpecConstantCompositeReplicateEXT expects "
2014 "3 operands but found ")
2020 return emitError(unknownLoc,
"undefined result type from <id> ")
2024 auto compositeType = dyn_cast<CompositeType>(resultType);
2025 if (!compositeType) {
2027 "result type from <id> is not a composite type")
2031 uint32_t resultID = operands[1];
2034 spirv::SpecConstantOp constituentSpecConstantOp =
2036 auto op = spirv::EXTSpecConstantCompositeReplicateOp::create(
2037 opBuilder, unknownLoc, TypeAttr::get(resultType), symName,
2038 SymbolRefAttr::get(constituentSpecConstantOp));
2040 specConstCompositeReplicateMap[resultID] = op;
2047 if (operands.size() < 3)
2048 return emitError(unknownLoc,
"OpConstantOperation must have type <id>, "
2049 "result <id>, and operand opcode");
2051 uint32_t resultTypeID = operands[0];
2054 return emitError(unknownLoc,
"undefined result type from <id> ")
2057 uint32_t resultID = operands[1];
2058 spirv::Opcode enclosedOpcode =
static_cast<spirv::Opcode
>(operands[2]);
2059 auto emplaceResult = specConstOperationMap.try_emplace(
2062 enclosedOpcode, resultTypeID,
2065 if (!emplaceResult.second)
2066 return emitError(unknownLoc,
"value with <id>: ")
2067 << resultID <<
" is probably defined before.";
2073 uint32_t resultID, spirv::Opcode enclosedOpcode, uint32_t resultTypeID,
2089 llvm::SaveAndRestore valueMapGuard(valueMap, newValueMap);
2090 constexpr uint32_t fakeID =
static_cast<uint32_t
>(-3);
2093 enclosedOpResultTypeAndOperands.push_back(resultTypeID);
2094 enclosedOpResultTypeAndOperands.push_back(fakeID);
2095 enclosedOpResultTypeAndOperands.append(enclosedOpOperands.begin(),
2096 enclosedOpOperands.end());
2111 auto specConstOperationOp =
2112 spirv::SpecConstantOperationOp::create(opBuilder, loc, resultType);
2114 Region &body = specConstOperationOp.getBody();
2116 body.
getBlocks().splice(body.
end(), curBlock->getParent()->getBlocks(),
2123 opBuilder.setInsertionPointToEnd(&block);
2125 spirv::YieldOp::create(opBuilder, loc, block.
front().
getResult(0));
2126 return specConstOperationOp.getResult();
2131 if (operands.size() != 2) {
2133 "OpConstantNull must only have type <id> and result <id>");
2138 return emitError(unknownLoc,
"undefined result type from <id> ")
2142 auto resultID = operands[1];
2144 if (resultType.
isIntOrFloat() || isa<VectorType>(resultType)) {
2145 attr = opBuilder.getZeroAttr(resultType);
2146 }
else if (
auto tensorType = dyn_cast<TensorArmType>(resultType)) {
2147 if (
auto element = opBuilder.getZeroAttr(tensorType.getElementType()))
2154 constantMap.try_emplace(resultID, attr, resultType);
2158 return emitError(unknownLoc,
"unsupported OpConstantNull type: ")
2164 if (operands.size() < 3) {
2166 <<
"OpGraphConstantARM must have at least 2 operands";
2171 return emitError(unknownLoc,
"undefined result type from <id> ")
2175 uint32_t resultID = operands[1];
2177 if (!dyn_cast<spirv::TensorArmType>(resultType)) {
2178 return emitError(unknownLoc,
"result must be of type OpTypeTensorARM");
2181 APInt graph_constant_id = APInt(32, operands[2],
true);
2182 Type i32Ty = opBuilder.getIntegerType(32);
2183 IntegerAttr attr = opBuilder.getIntegerAttr(i32Ty, graph_constant_id);
2184 graphConstantMap.try_emplace(
2196 LLVM_DEBUG(logger.startLine() <<
"[block] got exiting block for id = " <<
id
2197 <<
" @ " << block <<
"\n");
2204 auto *block = curFunction->addBlock();
2205 LLVM_DEBUG(logger.startLine() <<
"[block] created block for id = " <<
id
2206 <<
" @ " << block <<
"\n");
2207 return blockMap[id] = block;
2212 return emitError(unknownLoc,
"OpBranch must appear inside a block");
2215 if (operands.size() != 1) {
2216 return emitError(unknownLoc,
"OpBranch must take exactly one target label");
2224 spirv::BranchOp::create(opBuilder, loc,
target);
2234 "OpBranchConditional must appear inside a block");
2237 if (operands.size() != 3 && operands.size() != 5) {
2239 "OpBranchConditional must have condition, true label, "
2240 "false label, and optionally two branch weights");
2243 auto condition =
getValue(operands[0]);
2247 std::optional<std::pair<uint32_t, uint32_t>> weights;
2248 if (operands.size() == 5) {
2249 weights = std::make_pair(operands[3], operands[4]);
2255 spirv::BranchConditionalOp::create(
2256 opBuilder, loc, condition, trueBlock,
2266 return emitError(unknownLoc,
"OpLabel must appear inside a function");
2269 if (operands.size() != 1) {
2270 return emitError(unknownLoc,
"OpLabel should only have result <id>");
2273 auto labelID = operands[0];
2276 LLVM_DEBUG(logger.startLine()
2277 <<
"[block] populating block " << block <<
"\n");
2279 assert(block->empty() &&
"re-deserialize the same block!");
2281 opBuilder.setInsertionPointToStart(block);
2282 blockMap[labelID] = curBlock = block;
2289 return emitError(unknownLoc,
"a graph block must appear inside a graph");
2294 LLVM_DEBUG(logger.startLine()
2295 <<
"[block] populating block " << block <<
"\n");
2297 assert(block->
empty() &&
"re-deserialize the same block!");
2299 opBuilder.setInsertionPointToStart(block);
2300 blockMap[graphID] = curBlock = block;
2308 return emitError(unknownLoc,
"OpSelectionMerge must appear in a block");
2311 if (operands.size() < 2) {
2314 "OpSelectionMerge must specify merge target and selection control");
2319 auto selectionControl = operands[1];
2321 if (!blockMergeInfo.try_emplace(curBlock, loc, selectionControl, mergeBlock)
2325 "a block cannot have more than one OpSelectionMerge instruction");
2334 return emitError(unknownLoc,
"OpLoopMerge must appear in a block");
2337 if (operands.size() < 3) {
2338 return emitError(unknownLoc,
"OpLoopMerge must specify merge target, "
2339 "continue target and loop control");
2345 uint32_t loopControl = operands[2];
2348 .try_emplace(curBlock, loc, loopControl, mergeBlock, continueBlock)
2352 "a block cannot have more than one OpLoopMerge instruction");
2360 return emitError(unknownLoc,
"OpPhi must appear in a block");
2363 if (operands.size() < 4) {
2364 return emitError(unknownLoc,
"OpPhi must specify result type, result <id>, "
2365 "and variable-parent pairs");
2370 BlockArgument blockArg = curBlock->addArgument(blockArgType, unknownLoc);
2371 valueMap[operands[1]] = blockArg;
2372 LLVM_DEBUG(logger.startLine()
2373 <<
"[phi] created block argument " << blockArg
2374 <<
" id = " << operands[1] <<
" of type " << blockArgType <<
"\n");
2378 for (
unsigned i = 2, e = operands.size(); i < e; i += 2) {
2379 uint32_t value = operands[i];
2381 std::pair<Block *, Block *> predecessorTargetPair{predecessor, curBlock};
2382 blockPhiInfo[predecessorTargetPair].push_back(value);
2383 LLVM_DEBUG(logger.startLine() <<
"[phi] predecessor @ " << predecessor
2384 <<
" with arg id = " << value <<
"\n");
2392 return emitError(unknownLoc,
"OpSwitch must appear in a block");
2394 if (operands.size() < 2)
2395 return emitError(unknownLoc,
"OpSwitch must at least specify selector and "
2396 "a default target");
2398 if (operands.size() % 2)
2400 "OpSwitch must at have an even number of operands: "
2401 "selector, default target and any number of literal and "
2402 "label <id> pairs");
2410 for (
unsigned i = 2, e = operands.size(); i < e; i += 2) {
2411 literals.push_back(operands[i]);
2416 spirv::SwitchOp::create(opBuilder, loc, selector, defaultBlock,
2425class ControlFlowStructurizer {
2428 ControlFlowStructurizer(
Location loc, uint32_t control,
2431 llvm::ScopedPrinter &logger)
2432 : location(loc), control(control), blockMergeInfo(mergeInfo),
2433 headerBlock(header), mergeBlock(merge), continueBlock(cont),
2436 ControlFlowStructurizer(
Location loc, uint32_t control,
2439 : location(loc), control(control), blockMergeInfo(mergeInfo),
2440 headerBlock(header), mergeBlock(merge), continueBlock(cont) {}
2450 LogicalResult structurize();
2455 spirv::SelectionOp createSelectionOp(uint32_t selectionControl);
2458 spirv::LoopOp createLoopOp(uint32_t loopControl);
2461 void collectBlocksInConstruct();
2470 Block *continueBlock;
2476 llvm::ScopedPrinter &logger;
2482ControlFlowStructurizer::createSelectionOp(uint32_t selectionControl) {
2485 OpBuilder builder(&mergeBlock->front());
2487 auto control =
static_cast<spirv::SelectionControl
>(selectionControl);
2488 auto selectionOp = spirv::SelectionOp::create(builder, location, control);
2489 selectionOp.addMergeBlock(builder);
2494spirv::LoopOp ControlFlowStructurizer::createLoopOp(uint32_t loopControl) {
2497 OpBuilder builder(&mergeBlock->front());
2499 auto control =
static_cast<spirv::LoopControl
>(loopControl);
2500 auto loopOp = spirv::LoopOp::create(builder, location, control);
2501 loopOp.addEntryAndMergeBlock(builder);
2506void ControlFlowStructurizer::collectBlocksInConstruct() {
2507 assert(constructBlocks.empty() &&
"expected empty constructBlocks");
2510 constructBlocks.insert(headerBlock);
2514 for (
unsigned i = 0; i < constructBlocks.size(); ++i) {
2515 for (
auto *successor : constructBlocks[i]->getSuccessors())
2516 if (successor != mergeBlock)
2517 constructBlocks.insert(successor);
2521LogicalResult ControlFlowStructurizer::structurize() {
2522 Operation *op =
nullptr;
2523 bool isLoop = continueBlock !=
nullptr;
2525 if (
auto loopOp = createLoopOp(control))
2526 op = loopOp.getOperation();
2528 if (
auto selectionOp = createSelectionOp(control))
2529 op = selectionOp.getOperation();
2538 mapper.
map(mergeBlock, &body.
back());
2540 collectBlocksInConstruct();
2561 OpBuilder builder(body);
2562 for (
auto *block : constructBlocks) {
2565 auto *newBlock = builder.createBlock(&body.
back());
2566 mapper.
map(block, newBlock);
2567 LLVM_DEBUG(logger.startLine() <<
"[cf] cloned block " << newBlock
2568 <<
" from block " << block <<
"\n");
2570 for (BlockArgument blockArg : block->getArguments()) {
2572 newBlock->addArgument(blockArg.getType(), blockArg.getLoc());
2573 mapper.
map(blockArg, newArg);
2574 LLVM_DEBUG(logger.startLine() <<
"[cf] remapped block argument "
2575 << blockArg <<
" to " << newArg <<
"\n");
2578 LLVM_DEBUG(logger.startLine()
2579 <<
"[cf] block " << block <<
" is a function entry block\n");
2582 for (
auto &op : *block)
2583 newBlock->push_back(op.
clone(mapper));
2587 auto remapOperands = [&](Operation *op) {
2589 if (Value mappedOp = mapper.
lookupOrNull(operand.get()))
2590 operand.set(mappedOp);
2593 succOp.set(mappedOp);
2595 for (
auto &block : body)
2596 block.walk(remapOperands);
2604 headerBlock->replaceAllUsesWith(mergeBlock);
2607 logger.startLine() <<
"[cf] after cloning and fixing references:\n";
2608 headerBlock->getParentOp()->print(logger.getOStream());
2609 logger.startLine() <<
"\n";
2613 if (!mergeBlock->args_empty()) {
2614 return mergeBlock->getParentOp()->emitError(
2615 "OpPhi in loop merge block unsupported");
2621 for (BlockArgument blockArg : headerBlock->getArguments())
2622 mergeBlock->addArgument(blockArg.getType(), blockArg.getLoc());
2626 SmallVector<Value, 4> blockArgs;
2627 if (!headerBlock->args_empty())
2628 blockArgs = {mergeBlock->args_begin(), mergeBlock->args_end()};
2632 builder.setInsertionPointToEnd(&body.front());
2633 spirv::BranchOp::create(builder, location, mapper.
lookupOrNull(headerBlock),
2634 ArrayRef<Value>(blockArgs));
2639 SmallVector<Value> valuesToYield;
2642 SmallVector<Value> outsideUses;
2656 for (BlockArgument blockArg : mergeBlock->getArguments()) {
2661 body.back().addArgument(blockArg.getType(), blockArg.getLoc());
2662 valuesToYield.push_back(body.back().getArguments().back());
2663 outsideUses.push_back(blockArg);
2668 LLVM_DEBUG(logger.startLine() <<
"[cf] cleaning up blocks after clone\n");
2671 for (
auto *block : constructBlocks)
2672 block->dropAllReferences();
2677 for (
Block *block : constructBlocks) {
2678 for (Operation &op : *block) {
2682 outsideUses.push_back(
result);
2685 for (BlockArgument &arg : block->getArguments()) {
2686 if (!arg.use_empty()) {
2688 outsideUses.push_back(arg);
2693 assert(valuesToYield.size() == outsideUses.size());
2697 if (!valuesToYield.empty()) {
2698 LLVM_DEBUG(logger.startLine()
2699 <<
"[cf] yielding values from the selection / loop region\n");
2702 auto mergeOps = body.back().getOps<spirv::MergeOp>();
2703 Operation *merge = llvm::getSingleElement(mergeOps);
2705 merge->setOperands(valuesToYield);
2713 builder.setInsertionPoint(&mergeBlock->front());
2715 Operation *newOp =
nullptr;
2718 newOp = spirv::LoopOp::create(builder, location,
2720 static_cast<spirv::LoopControl
>(control));
2722 newOp = spirv::SelectionOp::create(
2724 static_cast<spirv::SelectionControl
>(control));
2734 for (
unsigned i = 0, e = outsideUses.size(); i != e; ++i)
2735 outsideUses[i].replaceAllUsesWith(op->
getResult(i));
2741 mergeBlock->eraseArguments(0, mergeBlock->getNumArguments());
2748 for (
auto *block : constructBlocks) {
2749 if (!block->use_empty())
2750 return emitError(block->getParent()->getLoc(),
2751 "failed control flow structurization: "
2752 "block has uses outside of the "
2753 "enclosing selection/loop construct");
2754 for (Operation &op : *block)
2756 return op.
emitOpError(
"failed control flow structurization: value has "
2757 "uses outside of the "
2758 "enclosing selection/loop construct");
2759 for (BlockArgument &arg : block->getArguments())
2760 if (!arg.use_empty())
2761 return emitError(arg.getLoc(),
"failed control flow structurization: "
2762 "block argument has uses outside of the "
2763 "enclosing selection/loop construct");
2767 for (
auto *block : constructBlocks) {
2807 auto updateMergeInfo = [&](
Block *block) -> WalkResult {
2808 auto it = blockMergeInfo.find(block);
2809 if (it != blockMergeInfo.end()) {
2811 Location loc = it->second.loc;
2815 return emitError(loc,
"failed control flow structurization: nested "
2816 "loop header block should be remapped!");
2818 Block *newContinue = it->second.continueBlock;
2822 return emitError(loc,
"failed control flow structurization: nested "
2823 "loop continue block should be remapped!");
2826 Block *newMerge = it->second.mergeBlock;
2828 newMerge = mappedTo;
2832 blockMergeInfo.
erase(it);
2833 blockMergeInfo.try_emplace(newHeader, loc, it->second.control, newMerge,
2840 if (block->walk(updateMergeInfo).wasInterrupted())
2848 LLVM_DEBUG(logger.startLine() <<
"[cf] changing entry block " << block
2849 <<
" to only contain a spirv.Branch op\n");
2853 builder.setInsertionPointToEnd(block);
2854 spirv::BranchOp::create(builder, location, mergeBlock);
2856 LLVM_DEBUG(logger.startLine() <<
"[cf] erasing block " << block <<
"\n");
2861 LLVM_DEBUG(logger.startLine()
2862 <<
"[cf] after structurizing construct with header block "
2863 << headerBlock <<
":\n"
2872 <<
"//----- [phi] start wiring up block arguments -----//\n";
2878 for (
const auto &info : blockPhiInfo) {
2879 Block *block = info.first.first;
2883 logger.startLine() <<
"[phi] block " << block <<
"\n";
2884 logger.startLine() <<
"[phi] before creating block argument:\n";
2886 logger.startLine() <<
"\n";
2892 opBuilder.setInsertionPoint(op);
2895 blockArgs.reserve(phiInfo.size());
2896 for (uint32_t valueId : phiInfo) {
2898 blockArgs.push_back(value);
2899 LLVM_DEBUG(logger.startLine() <<
"[phi] block argument " << value
2900 <<
" id = " << valueId <<
"\n");
2902 return emitError(unknownLoc,
"OpPhi references undefined value!");
2906 if (
auto branchOp = dyn_cast<spirv::BranchOp>(op)) {
2908 spirv::BranchOp::create(opBuilder, branchOp.getLoc(),
2909 branchOp.getTarget(), blockArgs);
2911 }
else if (
auto branchCondOp = dyn_cast<spirv::BranchConditionalOp>(op)) {
2912 assert((branchCondOp.getTrueBlock() ==
target ||
2913 branchCondOp.getFalseBlock() ==
target) &&
2914 "expected target to be either the true or false target");
2915 if (
target == branchCondOp.getTrueTarget())
2916 spirv::BranchConditionalOp::create(
2917 opBuilder, branchCondOp.getLoc(), branchCondOp.getCondition(),
2918 blockArgs, branchCondOp.getFalseBlockArguments(),
2919 branchCondOp.getBranchWeightsAttr(), branchCondOp.getTrueTarget(),
2920 branchCondOp.getFalseTarget());
2922 spirv::BranchConditionalOp::create(
2923 opBuilder, branchCondOp.getLoc(), branchCondOp.getCondition(),
2924 branchCondOp.getTrueBlockArguments(), blockArgs,
2925 branchCondOp.getBranchWeightsAttr(), branchCondOp.getTrueBlock(),
2926 branchCondOp.getFalseBlock());
2928 branchCondOp.erase();
2929 }
else if (
auto switchOp = dyn_cast<spirv::SwitchOp>(op)) {
2930 if (
target == switchOp.getDefaultTarget()) {
2934 spirv::SwitchOp::create(
2935 opBuilder, switchOp.getLoc(), switchOp.getSelector(),
2936 switchOp.getDefaultTarget(), blockArgs, literals,
2937 switchOp.getTargets(), targetOperands);
2941 auto it = llvm::find(targets,
target);
2942 assert(it != targets.end());
2943 size_t index = std::distance(targets.begin(), it);
2944 switchOp.getTargetOperandsMutable(
index).assign(blockArgs);
2947 return emitError(unknownLoc,
"unimplemented terminator for Phi creation");
2951 logger.startLine() <<
"[phi] after creating block argument:\n";
2953 logger.startLine() <<
"\n";
2956 blockPhiInfo.clear();
2961 <<
"//--- [phi] completed wiring up block arguments ---//\n";
2969 for (
auto [block, mergeInfo] : blockMergeInfoCopy) {
2971 if (mergeInfo.continueBlock)
2974 if (!block->mightHaveTerminator())
2977 Operation *terminator = block->getTerminator();
2980 if (!isa<spirv::BranchConditionalOp, spirv::SwitchOp>(terminator))
2984 bool splitHeaderMergeBlock =
false;
2985 for (
const auto &[_, mergeInfo] : blockMergeInfo) {
2986 if (mergeInfo.mergeBlock == block)
2987 splitHeaderMergeBlock =
true;
2994 if (!llvm::hasSingleElement(*block) || splitHeaderMergeBlock) {
2997 spirv::BranchOp::create(builder, block->getParent()->getLoc(), newBlock);
3001 blockMergeInfo.erase(block);
3002 blockMergeInfo.try_emplace(newBlock, mergeInfo);
3010 if (!options.enableControlFlowStructurization) {
3014 <<
"//----- [cf] skip structurizing control flow -----//\n";
3022 <<
"//----- [cf] start structurizing control flow -----//\n";
3027 logger.startLine() <<
"[cf] split conditional blocks\n";
3028 logger.startLine() <<
"\n";
3035 while (!blockMergeInfo.empty()) {
3036 Block *headerBlock = blockMergeInfo.
begin()->first;
3040 logger.startLine() <<
"[cf] header block " << headerBlock <<
":\n";
3041 headerBlock->
print(logger.getOStream());
3042 logger.startLine() <<
"\n";
3046 assert(mergeBlock &&
"merge block cannot be nullptr");
3048 return emitError(unknownLoc,
"OpPhi in loop merge block unimplemented");
3050 logger.startLine() <<
"[cf] merge block " << mergeBlock <<
":\n";
3051 mergeBlock->print(logger.getOStream());
3052 logger.startLine() <<
"\n";
3056 LLVM_DEBUG(
if (continueBlock) {
3057 logger.startLine() <<
"[cf] continue block " << continueBlock <<
":\n";
3058 continueBlock->print(logger.getOStream());
3059 logger.startLine() <<
"\n";
3063 blockMergeInfo.
erase(blockMergeInfo.begin());
3064 ControlFlowStructurizer structurizer(mergeInfo.
loc, mergeInfo.
control,
3065 blockMergeInfo, headerBlock,
3066 mergeBlock, continueBlock
3072 if (failed(structurizer.structurize()))
3079 <<
"//--- [cf] completed structurizing control flow ---//\n";
3092 auto fileName = debugInfoMap.lookup(debugLine->fileID).str();
3093 if (fileName.empty())
3094 fileName =
"<unknown>";
3106 if (operands.size() != 3)
3107 return emitError(unknownLoc,
"OpLine must have 3 operands");
3108 debugLine =
DebugLine{operands[0], operands[1], operands[2]};
3116 if (operands.size() < 2)
3117 return emitError(unknownLoc,
"OpString needs at least 2 operands");
3119 if (!debugInfoMap.lookup(operands[0]).empty())
3121 "duplicate debug string found for result <id> ")
3124 unsigned wordIndex = 1;
3126 if (wordIndex != operands.size())
3128 "unexpected trailing words in OpString instruction");
static bool isLoop(Operation *op)
Returns true if the given operation represents a loop by testing whether it implements the LoopLikeOp...
static bool isFnEntryBlock(Block *block)
Returns true if the given block is a function entry block.
#define MIN_VERSION_CASE(v)
static LogicalResult deserializeCacheControlDecoration(Location loc, OpBuilder &opBuilder, DenseMap< uint32_t, NamedAttrList > &decorations, ArrayRef< uint32_t > words, StringAttr symbol, StringRef decorationName, StringRef cacheControlKind)
static llvm::ManagedStatic< PassManagerOptions > options
Attributes are known-constant values of operations.
This class represents an argument of a Block.
Block represents an ordered list of Operations.
void erase()
Unlink this Block from its parent region and delete it.
Block * splitBlock(iterator splitBefore)
Split the block into two blocks before the specified operation or iterator.
Operation * getTerminator()
Get the terminator operation of this block.
void print(raw_ostream &os)
bool isEntryBlock()
Return if this block is the entry block in the parent region.
Operation * getParentOp()
Returns the closest surrounding operation that contains this block.
ArrayAttr getArrayAttr(ArrayRef< Attribute > value)
Attr getAttr(Args &&...args)
Get or construct an instance of the attribute Attr with provided arguments.
static DenseElementsAttr get(ShapedType type, ArrayRef< Attribute > values)
Constructs a dense elements attribute from an array of element values.
An attribute that represents a reference to a dense integer vector or tensor object.
static FileLineColLoc get(StringAttr filename, unsigned line, unsigned column)
A symbol reference with a reference path containing a single element.
static FlatSymbolRefAttr get(StringAttr value)
Construct a symbol reference for the given value name.
void map(Value from, Value to)
Inserts a new mapping for 'from' to 'to'.
auto lookupOrNull(T from) const
Lookup a mapped value within the map.
This class defines the main interface for locations in MLIR and acts as a non-nullable wrapper around...
MLIRContext is the top-level object for a collection of MLIR operations.
NamedAttrList is array of NamedAttributes that tracks whether it is sorted and does some basic work t...
NamedAttribute represents a combination of a name and an Attribute value.
RAII guard to reset the insertion point of the builder when destroyed.
This class helps build Operations.
Operation is the basic unit of execution within MLIR.
MutableArrayRef< BlockOperand > getBlockOperands()
Region & getRegion(unsigned index)
Returns the region held by this operation at position 'index'.
bool use_empty()
Returns true if this operation has no uses.
OpResult getResult(unsigned idx)
Get the 'idx'th result of this operation.
MutableArrayRef< OpOperand > getOpOperands()
void setAttr(StringAttr name, Attribute value)
If the an attribute exists with the specified name, change it to the new value.
void print(raw_ostream &os, const OpPrintingFlags &flags={})
static Operation * create(Location location, OperationName name, TypeRange resultTypes, ValueRange operands, NamedAttrList &&attributes, PropertyRef properties, BlockRange successors, unsigned numRegions)
Create a new Operation with the specific fields.
result_range getResults()
Operation * clone(IRMapping &mapper, const CloneOptions &options=CloneOptions::all())
Create a deep copy of this operation, remapping any operands that use values outside of the operation...
InFlightDiagnostic emitOpError(const Twine &message={})
Emit an error with the op name prefixed, like "'dim' op " which is convenient for verifiers.
void erase()
Remove this operation from its parent block and delete it.
This class contains a list of basic blocks and a link to the parent operation it is attached to.
BlockListType & getBlocks()
BlockListType::iterator iterator
void takeBody(Region &other)
Takes body of another region (that region will have no body after this operation completes).
This class implements the successor iterators for Block.
Instances of the Type class are uniqued, have an immutable identifier and an optional mutable compone...
bool isIntOrFloat() const
Return true if this is an integer (of any signedness) or a float type.
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
static WalkResult advance()
static ArrayType get(Type elementType, unsigned elementCount)
static CooperativeMatrixType get(Type elementType, uint32_t rows, uint32_t columns, Scope scope, CooperativeMatrixUseKHR use)
LogicalResult wireUpBlockArgument()
Creates block arguments on predecessors previously recorded when handling OpPhi instructions.
Value materializeSpecConstantOperation(uint32_t resultID, spirv::Opcode enclosedOpcode, uint32_t resultTypeID, ArrayRef< uint32_t > enclosedOpOperands)
Materializes/emits an OpSpecConstantOp instruction.
LogicalResult processOpTypePointer(ArrayRef< uint32_t > operands)
Value getValue(uint32_t id)
Get the Value associated with a result <id>.
LogicalResult processMatrixType(ArrayRef< uint32_t > operands)
LogicalResult processGlobalVariable(ArrayRef< uint32_t > operands)
Processes the OpVariable instructions at current offset into binary.
std::optional< SpecConstOperationMaterializationInfo > getSpecConstantOperation(uint32_t id)
Gets the info needed to materialize the spec constant operation op associated with the given <id>.
LogicalResult processConstantNull(ArrayRef< uint32_t > operands)
Processes a SPIR-V OpConstantNull instruction with the given operands.
LogicalResult processSpecConstantComposite(ArrayRef< uint32_t > operands)
Processes a SPIR-V OpSpecConstantComposite instruction with the given operands.
LogicalResult processInstruction(spirv::Opcode opcode, ArrayRef< uint32_t > operands, bool deferInstructions=true)
Processes a SPIR-V instruction with the given opcode and operands.
LogicalResult processBranchConditional(ArrayRef< uint32_t > operands)
spirv::GlobalVariableOp getGlobalVariable(uint32_t id)
Gets the global variable associated with a result <id> of OpVariable.
LogicalResult createGraphBlock(uint32_t graphID)
Creates a block for graph with the given graphID.
LogicalResult processStructType(ArrayRef< uint32_t > operands)
LogicalResult processGraphARM(ArrayRef< uint32_t > operands)
LogicalResult processSamplerType(ArrayRef< uint32_t > operands)
LogicalResult setFunctionArgAttrs(uint32_t argID, SmallVectorImpl< Attribute > &argAttrs, size_t argIndex)
Sets the function argument's attributes.
LogicalResult structurizeControlFlow()
Extracts blocks belonging to a structured selection/loop into a spirv.mlir.selection/spirv....
LogicalResult processLabel(ArrayRef< uint32_t > operands)
Processes a SPIR-V OpLabel instruction with the given operands.
LogicalResult processSampledImageType(ArrayRef< uint32_t > operands)
LogicalResult processTensorARMType(ArrayRef< uint32_t > operands)
std::optional< spirv::GraphConstantARMOpMaterializationInfo > getGraphConstantARM(uint32_t id)
Gets the GraphConstantARM ID attribute and result type with the given result <id>.
std::optional< std::pair< Attribute, Type > > getConstant(uint32_t id)
Gets the constant's attribute and type associated with the given <id>.
LogicalResult processType(spirv::Opcode opcode, ArrayRef< uint32_t > operands)
Processes a SPIR-V type instruction with given opcode and operands and registers the type into module...
LogicalResult processLoopMerge(ArrayRef< uint32_t > operands)
Processes a SPIR-V OpLoopMerge instruction with the given operands.
LogicalResult processArrayType(ArrayRef< uint32_t > operands)
LogicalResult sliceInstruction(spirv::Opcode &opcode, ArrayRef< uint32_t > &operands, std::optional< spirv::Opcode > expectedOpcode=std::nullopt)
Slices the first instruction out of binary and returns its opcode and operands via opcode and operand...
spirv::SpecConstantCompositeOp getSpecConstantComposite(uint32_t id)
Gets the composite specialization constant with the given result <id>.
LogicalResult processNamedBarrierType(ArrayRef< uint32_t > operands)
SmallVector< uint32_t, 2 > BlockPhiInfo
For OpPhi instructions, we use block arguments to represent them.
LogicalResult processSpecConstantCompositeReplicateEXT(ArrayRef< uint32_t > operands)
Processes a SPIR-V OpSpecConstantCompositeReplicateEXT instruction with the given operands.
LogicalResult processCooperativeMatrixTypeKHR(ArrayRef< uint32_t > operands)
LogicalResult processGraphEntryPointARM(ArrayRef< uint32_t > operands)
LogicalResult processFunction(ArrayRef< uint32_t > operands)
Creates a deserializer for the given SPIR-V binary module.
StringAttr getSymbolDecoration(StringRef decorationName)
Gets the symbol name from the name of decoration.
Block * getOrCreateBlock(uint32_t id)
Gets or creates the block corresponding to the given label <id>.
bool isVoidType(Type type) const
Returns true if the given type is for SPIR-V void type.
std::string getSpecConstantSymbol(uint32_t id)
Returns a symbol to be used for the specialization constant with the given result <id>.
LogicalResult processDebugString(ArrayRef< uint32_t > operands)
Processes a SPIR-V OpString instruction with the given operands.
LogicalResult processPhi(ArrayRef< uint32_t > operands)
Processes a SPIR-V OpPhi instruction with the given operands.
std::string getFunctionSymbol(uint32_t id)
Returns a symbol to be used for the function name with the given result <id>.
void clearDebugLine()
Discontinues any source-level location information that might be active from a previous OpLine instru...
LogicalResult processFunctionType(ArrayRef< uint32_t > operands)
IntegerAttr getConstantInt(uint32_t id)
Gets the constant's integer attribute with the given <id>.
LogicalResult processTypeForwardPointer(ArrayRef< uint32_t > operands)
LogicalResult processSwitch(ArrayRef< uint32_t > operands)
Processes a SPIR-V OpSwitch instruction with the given operands.
LogicalResult processGraphEndARM(ArrayRef< uint32_t > operands)
LogicalResult processImageType(ArrayRef< uint32_t > operands)
LogicalResult processConstantComposite(ArrayRef< uint32_t > operands)
Processes a SPIR-V OpConstantComposite instruction with the given operands.
spirv::SpecConstantOp createSpecConstant(Location loc, uint32_t resultID, TypedAttr defaultValue)
Creates a spirv::SpecConstantOp.
Block * getBlock(uint32_t id) const
Returns the block for the given label <id>.
LogicalResult processGraphTypeARM(ArrayRef< uint32_t > operands)
LogicalResult processBranch(ArrayRef< uint32_t > operands)
std::optional< std::pair< Attribute, Type > > getConstantCompositeReplicate(uint32_t id)
Gets the replicated composite constant's attribute and type associated with the given <id>.
LogicalResult processFunctionEnd(ArrayRef< uint32_t > operands)
Processes OpFunctionEnd and finalizes function.
LogicalResult processRuntimeArrayType(ArrayRef< uint32_t > operands)
LogicalResult processSpecConstantOperation(ArrayRef< uint32_t > operands)
Processes a SPIR-V OpSpecConstantOp instruction with the given operands.
LogicalResult processConstant(ArrayRef< uint32_t > operands, bool isSpec)
Processes a SPIR-V Op{|Spec}Constant instruction with the given operands.
Location createFileLineColLoc(OpBuilder opBuilder)
Creates a FileLineColLoc with the OpLine location information.
LogicalResult processGraphConstantARM(ArrayRef< uint32_t > operands)
Processes a SPIR-V OpGraphConstantARM instruction with the given operands.
LogicalResult processConstantBool(bool isTrue, ArrayRef< uint32_t > operands, bool isSpec)
Processes a SPIR-V Op{|Spec}Constant{True|False} instruction with the given operands.
spirv::SpecConstantOp getSpecConstant(uint32_t id)
Gets the specialization constant with the given result <id>.
LogicalResult processConstantCompositeReplicateEXT(ArrayRef< uint32_t > operands)
Processes a SPIR-V OpConstantCompositeReplicateEXT instruction with the given operands.
LogicalResult processSelectionMerge(ArrayRef< uint32_t > operands)
Processes a SPIR-V OpSelectionMerge instruction with the given operands.
LogicalResult processOpGraphSetOutputARM(ArrayRef< uint32_t > operands)
LogicalResult processDebugLine(ArrayRef< uint32_t > operands)
Processes a SPIR-V OpLine instruction with the given operands.
LogicalResult splitSelectionHeader()
Move a conditional branch or a switch into a separate basic block to avoid unnecessary sinking of def...
std::string getGraphSymbol(uint32_t id)
Returns a symbol to be used for the graph name with the given result <id>.
static ImageType get(Type elementType, Dim dim, ImageDepthInfo depth=ImageDepthInfo::DepthUnknown, ImageArrayedInfo arrayed=ImageArrayedInfo::NonArrayed, ImageSamplingInfo samplingInfo=ImageSamplingInfo::SingleSampled, ImageSamplerUseInfo samplerUse=ImageSamplerUseInfo::SamplerUnknown, ImageFormat format=ImageFormat::Unknown)
static MatrixType get(Type columnType, uint32_t columnCount)
static NamedBarrierType get(MLIRContext *context)
static PointerType get(Type pointeeType, StorageClass storageClass)
static RuntimeArrayType get(Type elementType)
static SampledImageType get(Type imageType)
static SamplerType get(MLIRContext *context)
static StructType getIdentified(MLIRContext *context, StringRef identifier)
Construct an identified StructType.
static StructType getEmpty(MLIRContext *context, StringRef identifier="")
Construct a (possibly identified) StructType with no members.
static StructType get(ArrayRef< Type > memberTypes, ArrayRef< OffsetInfo > offsetInfo={}, ArrayRef< MemberDecorationInfo > memberDecorations={}, ArrayRef< StructDecorationInfo > structDecorations={})
Construct a literal StructType with at least one member.
static TensorArmType get(ArrayRef< int64_t > shape, Type elementType)
static VerCapExtAttr get(Version version, ArrayRef< Capability > capabilities, ArrayRef< Extension > extensions, MLIRContext *context)
Gets a VerCapExtAttr instance.
The OpAsmOpInterface, see OpAsmInterface.td for more details.
SmallVector< Operation * > mergeOps
Computation function returning, for the op currently being tiled or fused, the per-iteration-domain-d...
constexpr uint32_t kMagicNumber
SPIR-V magic number.
llvm::MapVector< Block *, BlockMergeInfo > BlockMergeInfoMap
Map from a selection/loop's header block to its merge (and continue) target.
StringRef decodeStringLiteral(ArrayRef< uint32_t > words, unsigned &wordIndex)
Decodes a string literal in words starting at wordIndex.
constexpr unsigned kHeaderWordCount
SPIR-V binary header word count.
Include the generated interface declarations.
Type getType(OpFoldResult ofr)
Returns the int type of the integer in ofr.
InFlightDiagnostic emitError(Location loc)
Utility method to emit an error message using this location.
static std::string debugString(T &&op)
llvm::SetVector< T, Vector, Set, N > SetVector
auto get(MLIRContext *context, Ts &&...params)
Helper method that injects context only if needed, this helps unify some of the attribute constructio...
llvm::DenseMap< KeyT, ValueT, KeyInfoT, BucketT > DenseMap
A struct for containing a header block's merge and continue targets.
A struct for containing OpLine instruction information.
A struct that collects the info needed to materialize/emit a GraphConstantARMOp.
A struct that collects the info needed to materialize/emit a SpecConstantOperation op.