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());
63 : binary(binary), context(context), unknownLoc(UnknownLoc::
get(context)),
64 module(createModuleOp()), opBuilder(module->getRegion()),
options(
options)
72LogicalResult spirv::Deserializer::deserialize() {
76 <<
"//+++---------- start deserialization ----------+++//\n";
79 if (
failed(processHeader()))
82 spirv::Opcode opcode = spirv::Opcode::OpNop;
83 ArrayRef<uint32_t> operands;
84 auto binarySize = binary.size();
85 while (curOffset < binarySize) {
95 assert(curOffset == binarySize &&
96 "deserializer should never index beyond the binary end");
98 for (
auto &deferred : deferredInstructions) {
104 if (
failed(resolveDeferredIdDecorations()))
109 LLVM_DEBUG(logger.startLine()
110 <<
"//+++-------- completed deserialization --------+++//\n");
114OwningOpRef<spirv::ModuleOp> spirv::Deserializer::collect() {
115 return std::move(module);
122OwningOpRef<spirv::ModuleOp> spirv::Deserializer::createModuleOp() {
123 OpBuilder builder(context);
124 OperationState state(unknownLoc, spirv::ModuleOp::getOperationName());
125 spirv::ModuleOp::build(builder, state);
129LogicalResult spirv::Deserializer::processHeader() {
132 "SPIR-V binary module must have a 5-word header");
135 return emitError(unknownLoc,
"incorrect magic number");
138 uint32_t majorVersion = (binary[1] << 8) >> 24;
139 uint32_t minorVersion = (binary[1] << 16) >> 24;
140 if (majorVersion == 1) {
141 switch (minorVersion) {
142#define MIN_VERSION_CASE(v) \
144 version = spirv::Version::V_1_##v; \
154#undef MIN_VERSION_CASE
156 return emitError(unknownLoc,
"unsupported SPIR-V minor version: ")
160 return emitError(unknownLoc,
"unsupported SPIR-V major version: ")
170spirv::Deserializer::processCapability(ArrayRef<uint32_t> operands) {
171 if (operands.size() != 1)
172 return emitError(unknownLoc,
"OpCapability must have one parameter");
174 auto cap = spirv::symbolizeCapability(operands[0]);
176 return emitError(unknownLoc,
"unknown capability: ") << operands[0];
178 capabilities.insert(*cap);
182LogicalResult spirv::Deserializer::processExtension(ArrayRef<uint32_t> words) {
186 "OpExtension must have a literal string for the extension name");
189 unsigned wordIndex = 0;
191 if (wordIndex != words.size())
193 "unexpected trailing words in OpExtension instruction");
194 auto ext = spirv::symbolizeExtension(extName);
196 return emitError(unknownLoc,
"unknown extension: ") << extName;
198 extensions.insert(*ext);
203spirv::Deserializer::processExtInstImport(ArrayRef<uint32_t> words) {
204 if (words.size() < 2) {
206 "OpExtInstImport must have a result <id> and a literal "
207 "string for the extended instruction set name");
210 unsigned wordIndex = 1;
212 if (wordIndex != words.size()) {
214 "unexpected trailing words in OpExtInstImport");
219void spirv::Deserializer::attachVCETriple() {
220 module->setVceTripleAttr(spirv::VerCapExtAttr::get(
221 version, capabilities.getArrayRef(), extensions.getArrayRef(), context));
225spirv::Deserializer::processMemoryModel(ArrayRef<uint32_t> operands) {
226 if (operands.size() != 2)
227 return emitError(unknownLoc,
"OpMemoryModel must have two operands");
229 module->setAddressingModel(
230 static_cast<spirv::AddressingModel>(operands.front()));
232 module->setMemoryModel(static_cast<spirv::MemoryModel>(operands.back()));
237template <
typename AttrTy,
typename EnumAttrTy,
typename EnumTy>
241 StringAttr symbol, StringRef decorationName, StringRef cacheControlKind) {
242 if (words.size() != 4) {
243 return emitError(loc,
"OpDecorate with ")
244 << decorationName <<
" needs a cache control integer literal and a "
245 << cacheControlKind <<
" cache control literal";
247 unsigned cacheLevel = words[2];
248 auto cacheControlAttr =
static_cast<EnumTy
>(words[3]);
249 auto value = opBuilder.
getAttr<AttrTy>(cacheLevel, cacheControlAttr);
252 dyn_cast_or_null<ArrayAttr>(decorations[words[0]].
get(symbol)))
253 llvm::append_range(attrs, attrList);
254 attrs.push_back(value);
255 decorations[words[0]].set(symbol, opBuilder.
getArrayAttr(attrs));
259LogicalResult spirv::Deserializer::processDecoration(ArrayRef<uint32_t> words) {
263 if (words.size() < 2) {
265 unknownLoc,
"OpDecorate must have at least result <id> and Decoration");
267 auto decorationName =
268 stringifyDecoration(
static_cast<spirv::Decoration
>(words[1]));
269 if (decorationName.empty()) {
270 return emitError(unknownLoc,
"invalid Decoration code : ") << words[1];
272 auto symbol = getSymbolDecoration(decorationName);
273 switch (
static_cast<spirv::Decoration
>(words[1])) {
274 case spirv::Decoration::FPFastMathMode:
275 if (words.size() != 3) {
276 return emitError(unknownLoc,
"OpDecorate with ")
277 << decorationName <<
" needs a single integer literal";
279 decorations[words[0]].set(
280 symbol, FPFastMathModeAttr::get(opBuilder.getContext(),
281 static_cast<FPFastMathMode
>(words[2])));
283 case spirv::Decoration::FPRoundingMode:
284 if (words.size() != 3) {
285 return emitError(unknownLoc,
"OpDecorate with ")
286 << decorationName <<
" needs a single integer literal";
288 decorations[words[0]].set(
289 symbol, FPRoundingModeAttr::get(opBuilder.getContext(),
290 static_cast<FPRoundingMode
>(words[2])));
292 case spirv::Decoration::DescriptorSet:
293 case spirv::Decoration::Binding:
294 case spirv::Decoration::Location:
295 case spirv::Decoration::SpecId:
296 case spirv::Decoration::Index:
297 case spirv::Decoration::Offset:
298 case spirv::Decoration::XfbBuffer:
299 case spirv::Decoration::XfbStride:
300 if (words.size() != 3) {
301 return emitError(unknownLoc,
"OpDecorate with ")
302 << decorationName <<
" needs a single integer literal";
304 decorations[words[0]].set(
305 symbol, opBuilder.getI32IntegerAttr(
static_cast<int32_t
>(words[2])));
307 case spirv::Decoration::BuiltIn:
308 if (words.size() != 3) {
309 return emitError(unknownLoc,
"OpDecorate with ")
310 << decorationName <<
" needs a single integer literal";
312 decorations[words[0]].set(
313 symbol, opBuilder.getStringAttr(
314 stringifyBuiltIn(
static_cast<spirv::BuiltIn
>(words[2]))));
316 case spirv::Decoration::ArrayStride:
317 if (words.size() != 3) {
318 return emitError(unknownLoc,
"OpDecorate with ")
319 << decorationName <<
" needs a single integer literal";
321 typeDecorations[words[0]] = words[2];
323 case spirv::Decoration::LinkageAttributes: {
324 if (words.size() < 4) {
325 return emitError(unknownLoc,
"OpDecorate with ")
327 <<
" needs at least 1 string and 1 integer literal";
335 unsigned wordIndex = 2;
337 auto linkageTypeAttr = opBuilder.getAttr<::mlir::spirv::LinkageTypeAttr>(
338 static_cast<::mlir::spirv::LinkageType
>(words[wordIndex++]));
339 auto linkageAttr = opBuilder.getAttr<::mlir::spirv::LinkageAttributesAttr>(
340 StringAttr::get(context, linkageName), linkageTypeAttr);
341 decorations[words[0]].set(symbol, dyn_cast<Attribute>(linkageAttr));
344 case spirv::Decoration::Aliased:
345 case spirv::Decoration::AliasedPointer:
346 case spirv::Decoration::Block:
347 case spirv::Decoration::BufferBlock:
348 case spirv::Decoration::Flat:
349 case spirv::Decoration::NonReadable:
350 case spirv::Decoration::NonWritable:
351 case spirv::Decoration::NoPerspective:
352 case spirv::Decoration::NoSignedWrap:
353 case spirv::Decoration::NoUnsignedWrap:
354 case spirv::Decoration::RelaxedPrecision:
355 case spirv::Decoration::Restrict:
356 case spirv::Decoration::RestrictPointer:
357 case spirv::Decoration::NoContraction:
358 case spirv::Decoration::Constant:
359 case spirv::Decoration::Invariant:
360 case spirv::Decoration::Patch:
361 case spirv::Decoration::Coherent:
362 case spirv::Decoration::Volatile:
363 if (words.size() != 2) {
364 return emitError(unknownLoc,
"OpDecorate with ")
365 << decorationName <<
" needs a single target <id>";
367 decorations[words[0]].set(symbol, opBuilder.getUnitAttr());
369 case spirv::Decoration::CacheControlLoadINTEL: {
371 CacheControlLoadINTELAttr, LoadCacheControlAttr, LoadCacheControl>(
372 unknownLoc, opBuilder, decorations, words, symbol, decorationName,
378 case spirv::Decoration::CacheControlStoreINTEL: {
380 CacheControlStoreINTELAttr, StoreCacheControlAttr, StoreCacheControl>(
381 unknownLoc, opBuilder, decorations, words, symbol, decorationName,
387 case spirv::Decoration::AlignmentId:
388 case spirv::Decoration::MaxByteOffsetId:
389 case spirv::Decoration::CounterBuffer:
390 if (words.size() != 3) {
391 return emitError(unknownLoc,
"OpDecorateId with ")
392 << decorationName <<
" needs a single <id> operand";
394 pendingIdDecorations.push_back({words[0],
395 static_cast<spirv::Decoration
>(words[1]),
396 words[2], unknownLoc});
399 return emitError(unknownLoc,
"unhandled Decoration : '") << decorationName;
404LogicalResult spirv::Deserializer::resolveDeferredIdDecorations() {
405 for (
const DeferredIdDecoration &entry : pendingIdDecorations) {
406 StringRef decorationName = stringifyDecoration(entry.decoration);
407 StringAttr symbol = getSymbolDecoration(decorationName);
411 StringRef operandSymName;
412 if (spirv::GlobalVariableOp varOp =
413 globalVariableMap.lookup(entry.operandID))
414 operandSymName = varOp.getSymName();
415 else if (spirv::SpecConstantOp specOp =
416 specConstMap.lookup(entry.operandID))
417 operandSymName = specOp.getSymName();
419 return emitError(entry.loc,
"OpDecorateId with ")
420 << decorationName <<
" references <id> " << entry.operandID
421 <<
" which is not a global variable or specialization constant";
428 Operation *targetOp =
nullptr;
429 if (spirv::GlobalVariableOp varOp =
430 globalVariableMap.lookup(entry.targetID))
432 else if (spirv::SpecConstantOp specOp = specConstMap.lookup(entry.targetID))
434 else if (spirv::FuncOp fnOp = funcMap.lookup(entry.targetID))
436 else if (Value v = valueMap.lookup(entry.targetID))
437 targetOp = v.getDefiningOp();
440 return emitError(entry.loc,
"OpDecorateId with ")
441 << decorationName <<
" references unknown target <id> "
450spirv::Deserializer::processMemberDecoration(ArrayRef<uint32_t> words) {
452 if (words.size() < 3) {
454 "OpMemberDecorate must have at least 3 operands");
457 auto decoration =
static_cast<spirv::Decoration
>(words[2]);
458 if (decoration == spirv::Decoration::Offset && words.size() != 4) {
460 " missing offset specification in OpMemberDecorate with "
461 "Offset decoration");
463 ArrayRef<uint32_t> decorationOperands;
464 if (words.size() > 3) {
465 decorationOperands = words.slice(3);
467 memberDecorationMap[words[0]][words[1]][decoration] = decorationOperands;
471LogicalResult spirv::Deserializer::processMemberName(ArrayRef<uint32_t> words) {
472 if (words.size() < 3) {
473 return emitError(unknownLoc,
"OpMemberName must have at least 3 operands");
475 unsigned wordIndex = 2;
477 if (wordIndex != words.size()) {
479 "unexpected trailing words in OpMemberName instruction");
481 memberNameMap[words[0]][words[1]] = name;
487 if (!decorations.contains(argID)) {
488 argAttrs[argIndex] = DictionaryAttr::get(context, {});
492 spirv::DecorationAttr foundDecorationAttr;
494 for (
auto decoration :
495 {spirv::Decoration::Aliased, spirv::Decoration::Restrict,
496 spirv::Decoration::AliasedPointer,
497 spirv::Decoration::RestrictPointer}) {
499 if (decAttr.getName() !=
503 if (foundDecorationAttr)
505 "more than one Aliased/Restrict decorations for "
506 "function argument with result <id> ")
509 foundDecorationAttr = spirv::DecorationAttr::get(context, decoration);
514 spirv::Decoration::RelaxedPrecision))) {
519 if (foundDecorationAttr)
520 return emitError(unknownLoc,
"already found a decoration for function "
521 "argument with result <id> ")
524 foundDecorationAttr = spirv::DecorationAttr::get(
525 context, spirv::Decoration::RelaxedPrecision);
529 if (!foundDecorationAttr)
530 return emitError(unknownLoc,
"unimplemented decoration support for "
531 "function argument with result <id> ")
534 NamedAttribute attr(StringAttr::get(context, spirv::DecorationAttr::name),
535 foundDecorationAttr);
536 argAttrs[argIndex] = DictionaryAttr::get(context, attr);
543 return emitError(unknownLoc,
"found function inside function");
547 if (operands.size() != 4) {
548 return emitError(unknownLoc,
"OpFunction must have 4 parameters");
552 return emitError(unknownLoc,
"undefined result type from <id> ")
556 uint32_t fnID = operands[1];
557 if (funcMap.count(fnID)) {
558 return emitError(unknownLoc,
"duplicate function definition/declaration");
561 auto fnControl = spirv::symbolizeFunctionControl(operands[2]);
563 return emitError(unknownLoc,
"unknown Function Control: ") << operands[2];
567 if (!fnType || !isa<FunctionType>(fnType)) {
568 return emitError(unknownLoc,
"unknown function type from <id> ")
571 auto functionType = cast<FunctionType>(fnType);
573 if ((
isVoidType(resultType) && functionType.getNumResults() != 0) ||
574 (functionType.getNumResults() == 1 &&
575 functionType.getResult(0) != resultType)) {
576 return emitError(unknownLoc,
"mismatch in function type ")
577 << functionType <<
" and return type " << resultType <<
" specified";
581 auto funcOp = spirv::FuncOp::create(opBuilder, unknownLoc, fnName,
582 functionType, fnControl.value());
584 if (decorations.count(fnID)) {
585 for (
auto attr : decorations[fnID].getAttrs()) {
589 curFunction = funcMap[fnID] = funcOp;
590 auto *entryBlock = funcOp.addEntryBlock();
593 <<
"//===-------------------------------------------===//\n";
594 logger.startLine() <<
"[fn] name: " << fnName <<
"\n";
595 logger.startLine() <<
"[fn] type: " << fnType <<
"\n";
596 logger.startLine() <<
"[fn] ID: " << fnID <<
"\n";
597 logger.startLine() <<
"[fn] entry block: " << entryBlock <<
"\n";
602 argAttrs.resize(functionType.getNumInputs());
605 if (functionType.getNumInputs()) {
606 for (
size_t i = 0, e = functionType.getNumInputs(); i != e; ++i) {
607 auto argType = functionType.getInput(i);
608 spirv::Opcode opcode = spirv::Opcode::OpNop;
611 spirv::Opcode::OpFunctionParameter))) {
614 if (opcode != spirv::Opcode::OpFunctionParameter) {
617 "missing OpFunctionParameter instruction for argument ")
620 if (operands.size() != 2) {
623 "expected result type and result <id> for OpFunctionParameter");
625 auto argDefinedType =
getType(operands[0]);
626 if (!argDefinedType || argDefinedType != argType) {
628 "mismatch in argument type between function type "
630 << functionType <<
" and argument type definition "
631 << argDefinedType <<
" at argument " << i;
634 return emitError(unknownLoc,
"duplicate definition of result <id> ")
641 auto argValue = funcOp.getArgument(i);
642 valueMap[operands[1]] = argValue;
646 if (llvm::any_of(argAttrs, [](
Attribute attr) {
647 auto argAttr = cast<DictionaryAttr>(attr);
648 return !argAttr.empty();
650 funcOp.setArgAttrsAttr(ArrayAttr::get(context, argAttrs));
655 auto linkageAttr = funcOp.getLinkageAttributes();
656 auto hasImportLinkage =
657 linkageAttr && (linkageAttr.value().getLinkageType().
getValue() ==
658 spirv::LinkageType::Import);
659 if (hasImportLinkage)
666 spirv::Opcode opcode = spirv::Opcode::OpNop;
675 spirv::Opcode::OpFunctionEnd))) {
678 if (opcode == spirv::Opcode::OpFunctionEnd) {
681 if (opcode != spirv::Opcode::OpLabel) {
682 return emitError(unknownLoc,
"a basic block must start with OpLabel");
684 if (instOperands.size() != 1) {
685 return emitError(unknownLoc,
"OpLabel should only have result <id>");
687 blockMap[instOperands[0]] = entryBlock;
695 spirv::Opcode::OpFunctionEnd)) &&
696 opcode != spirv::Opcode::OpFunctionEnd) {
701 if (opcode != spirv::Opcode::OpFunctionEnd) {
711 if (!operands.empty()) {
712 return emitError(unknownLoc,
"unexpected operands for OpFunctionEnd");
723 curFunction = std::nullopt;
728 <<
"//===-------------------------------------------===//\n";
735 if (operands.size() < 2) {
737 "missing graph defintion in OpGraphEntryPointARM");
740 unsigned wordIndex = 0;
741 uint32_t graphID = operands[wordIndex++];
742 if (!graphMap.contains(graphID)) {
744 "missing graph definition/declaration with id ")
748 spirv::GraphARMOp graphARM = graphMap[graphID];
750 graphARM.setSymName(name);
751 graphARM.setEntryPoint(
true);
754 for (
int64_t size = operands.size(); wordIndex < size; ++wordIndex) {
756 interface.push_back(SymbolRefAttr::get(arg.getOperation()));
758 return emitError(unknownLoc,
"undefined result <id> ")
759 << operands[wordIndex] <<
" while decoding OpGraphEntryPoint";
765 opBuilder.setInsertionPoint(graphARM);
766 spirv::GraphEntryPointARMOp::create(
767 opBuilder, unknownLoc, SymbolRefAttr::get(opBuilder.getContext(), name),
768 opBuilder.getArrayAttr(interface));
776 return emitError(unknownLoc,
"found graph inside graph");
779 if (operands.size() < 2) {
780 return emitError(unknownLoc,
"OpGraphARM must have at least 2 parameters");
784 if (!type || !isa<GraphType>(type)) {
785 return emitError(unknownLoc,
"unknown graph type from <id> ")
788 auto graphType = cast<GraphType>(type);
789 if (graphType.getNumResults() <= 0) {
790 return emitError(unknownLoc,
"expected at least one result");
793 uint32_t graphID = operands[1];
794 if (graphMap.count(graphID)) {
795 return emitError(unknownLoc,
"duplicate graph definition/declaration");
800 spirv::GraphARMOp::create(opBuilder, unknownLoc, graphName, graphType);
801 curGraph = graphMap[graphID] = graphOp;
802 Block *entryBlock = graphOp.addEntryBlock();
805 <<
"//===-------------------------------------------===//\n";
806 logger.startLine() <<
"[graph] name: " << graphName <<
"\n";
807 logger.startLine() <<
"[graph] type: " << graphType <<
"\n";
808 logger.startLine() <<
"[graph] ID: " << graphID <<
"\n";
809 logger.startLine() <<
"[graph] entry block: " << entryBlock <<
"\n";
814 for (
auto [
index, argType] : llvm::enumerate(graphType.getInputs())) {
815 spirv::Opcode opcode;
818 spirv::Opcode::OpGraphInputARM))) {
821 if (operands.size() != 3) {
822 return emitError(unknownLoc,
"expected result type, result <id> and "
823 "input index for OpGraphInputARM");
827 if (!argDefinedType) {
828 return emitError(unknownLoc,
"unknown operand type <id> ") << operands[0];
831 if (argDefinedType != argType) {
833 "mismatch in argument type between graph type "
835 << graphType <<
" and argument type definition " << argDefinedType
836 <<
" at argument " <<
index;
839 return emitError(unknownLoc,
"duplicate definition of result <id> ")
844 if (!inputIndexAttr) {
846 "unable to read inputIndex value from constant op ")
849 BlockArgument argValue = graphOp.getArgument(inputIndexAttr.getInt());
850 valueMap[operands[1]] = argValue;
853 graphOutputs.resize(graphType.getNumResults());
859 blockMap[graphID] = entryBlock;
866 spirv::Opcode opcode;
876 }
while (opcode != spirv::Opcode::OpGraphEndARM);
883 if (operands.size() != 2) {
886 "expected value id and output index for OpGraphSetOutputARM");
889 uint32_t
id = operands[0];
892 return emitError(unknownLoc,
"could not find result <id> ") << id;
896 if (!outputIndexAttr) {
898 "unable to read outputIndex value from constant op ")
901 graphOutputs[outputIndexAttr.getInt()] = value;
908 spirv::GraphOutputsARMOp::create(opBuilder, unknownLoc, graphOutputs);
911 if (!operands.empty()) {
912 return emitError(unknownLoc,
"unexpected operands for OpGraphEndARM");
916 curGraph = std::nullopt;
917 graphOutputs.clear();
922 <<
"//===-------------------------------------------===//\n";
927std::optional<std::pair<Attribute, Type>>
929 auto constIt = constantMap.find(
id);
930 if (constIt != constantMap.end())
931 return constIt->getSecond();
933 auto replicatedConstIt = constantCompositeReplicateMap.find(
id);
934 if (replicatedConstIt == constantCompositeReplicateMap.end())
937 auto [value, type] = replicatedConstIt->getSecond();
938 auto shapedType = dyn_cast<ShapedType>(type);
944std::optional<std::pair<Attribute, Type>>
946 if (
auto it = constantCompositeReplicateMap.find(
id);
947 it != constantCompositeReplicateMap.end())
952std::optional<spirv::SpecConstOperationMaterializationInfo>
954 auto constIt = specConstOperationMap.find(
id);
955 if (constIt == specConstOperationMap.end())
957 return constIt->getSecond();
961 auto funcName = nameMap.lookup(
id).str();
962 if (funcName.empty()) {
963 funcName =
"spirv_fn_" + std::to_string(
id);
969 std::string graphName = nameMap.lookup(
id).str();
970 if (graphName.empty()) {
971 graphName =
"spirv_graph_" + std::to_string(
id);
977 auto constName = nameMap.lookup(
id).str();
978 if (constName.empty()) {
979 constName =
"spirv_spec_const_" + std::to_string(
id);
986 TypedAttr defaultValue) {
988 auto op = spirv::SpecConstantOp::create(opBuilder, unknownLoc, symName,
990 if (decorations.count(resultID)) {
991 for (
auto attr : decorations[resultID].getAttrs())
994 specConstMap[resultID] = op;
998std::optional<spirv::GraphConstantARMOpMaterializationInfo>
1000 auto graphConstIt = graphConstantMap.find(
id);
1001 if (graphConstIt == graphConstantMap.end())
1002 return std::nullopt;
1003 return graphConstIt->getSecond();
1008 unsigned wordIndex = 0;
1009 if (operands.size() < 3) {
1012 "OpVariable needs at least 3 operands, type, <id> and storage class");
1016 auto type =
getType(operands[wordIndex]);
1018 return emitError(unknownLoc,
"unknown result type <id> : ")
1019 << operands[wordIndex];
1021 auto ptrType = dyn_cast<spirv::PointerType>(type);
1024 "expected a result type <id> to be a spirv.ptr, found : ")
1030 auto variableID = operands[wordIndex];
1031 auto variableName = nameMap.lookup(variableID).str();
1032 if (variableName.empty()) {
1033 variableName =
"spirv_var_" + std::to_string(variableID);
1038 auto storageClass =
static_cast<spirv::StorageClass
>(operands[wordIndex]);
1039 if (ptrType.getStorageClass() != storageClass) {
1040 return emitError(unknownLoc,
"mismatch in storage class of pointer type ")
1041 << type <<
" and that specified in OpVariable instruction : "
1042 << stringifyStorageClass(storageClass);
1049 if (wordIndex < operands.size()) {
1059 return emitError(unknownLoc,
"unknown <id> ")
1060 << operands[wordIndex] <<
"used as initializer";
1062 initializer = SymbolRefAttr::get(op);
1065 if (wordIndex != operands.size()) {
1067 "found more operands than expected when deserializing "
1068 "OpVariable instruction, only ")
1069 << wordIndex <<
" of " << operands.size() <<
" processed";
1072 auto varOp = spirv::GlobalVariableOp::create(
1073 opBuilder, loc, TypeAttr::get(type),
1074 opBuilder.getStringAttr(variableName), initializer);
1077 if (decorations.count(variableID)) {
1078 for (
auto attr : decorations[variableID].getAttrs())
1081 globalVariableMap[variableID] = varOp;
1090 return dyn_cast<IntegerAttr>(constInfo->first);
1094 if (operands.size() < 2) {
1095 return emitError(unknownLoc,
"OpName needs at least 2 operands");
1098 unsigned wordIndex = 1;
1100 if (wordIndex != operands.size()) {
1102 "unexpected trailing words in OpName instruction");
1107 nameMap.emplace_or_assign(operands[0], name);
1118 if (operands.empty()) {
1119 return emitError(unknownLoc,
"type instruction with opcode ")
1120 << spirv::stringifyOpcode(opcode) <<
" needs at least one <id>";
1125 if (typeMap.count(operands[0])) {
1126 return emitError(unknownLoc,
"duplicate definition for result <id> ")
1131 case spirv::Opcode::OpTypeVoid:
1132 if (operands.size() != 1)
1133 return emitError(unknownLoc,
"OpTypeVoid must have no parameters");
1134 typeMap[operands[0]] = opBuilder.getNoneType();
1136 case spirv::Opcode::OpTypeBool:
1137 if (operands.size() != 1)
1138 return emitError(unknownLoc,
"OpTypeBool must have no parameters");
1139 typeMap[operands[0]] = opBuilder.getI1Type();
1141 case spirv::Opcode::OpTypeInt: {
1142 if (operands.size() != 3)
1144 unknownLoc,
"OpTypeInt must have bitwidth and signedness parameters");
1153 auto sign = operands[2] == 1 ? IntegerType::SignednessSemantics::Signed
1154 : IntegerType::SignednessSemantics::Signless;
1155 typeMap[operands[0]] = IntegerType::get(context, operands[1], sign);
1157 case spirv::Opcode::OpTypeFloat: {
1158 if (operands.size() != 2 && operands.size() != 3)
1160 "OpTypeFloat expects either 2 operands (type, bitwidth) "
1161 "or 3 operands (type, bitwidth, encoding), but got ")
1163 uint32_t bitWidth = operands[1];
1166 if (operands.size() == 2) {
1169 floatTy = opBuilder.getF16Type();
1172 floatTy = opBuilder.getF32Type();
1175 floatTy = opBuilder.getF64Type();
1178 return emitError(unknownLoc,
"unsupported OpTypeFloat bitwidth: ")
1183 if (operands.size() == 3) {
1184 if (spirv::FPEncoding(operands[2]) == spirv::FPEncoding::BFloat16KHR &&
1186 floatTy = opBuilder.getBF16Type();
1187 else if (spirv::FPEncoding(operands[2]) ==
1188 spirv::FPEncoding::Float8E4M3EXT &&
1190 floatTy = opBuilder.getF8E4M3FNType();
1191 else if (spirv::FPEncoding(operands[2]) ==
1192 spirv::FPEncoding::Float8E5M2EXT &&
1194 floatTy = opBuilder.getF8E5M2Type();
1196 return emitError(unknownLoc,
"unsupported OpTypeFloat FP encoding: ")
1197 << operands[2] <<
" and bitWidth " << bitWidth;
1200 typeMap[operands[0]] = floatTy;
1202 case spirv::Opcode::OpTypeVector: {
1203 if (operands.size() != 3) {
1206 "OpTypeVector must have element type and count parameters");
1210 return emitError(unknownLoc,
"OpTypeVector references undefined <id> ")
1213 typeMap[operands[0]] = VectorType::get({operands[2]}, elementTy);
1215 case spirv::Opcode::OpTypePointer: {
1218 case spirv::Opcode::OpTypeArray:
1220 case spirv::Opcode::OpTypeCooperativeMatrixKHR:
1222 case spirv::Opcode::OpTypeFunction:
1224 case spirv::Opcode::OpTypeImage:
1226 case spirv::Opcode::OpTypeSampler:
1228 case spirv::Opcode::OpTypeNamedBarrier:
1230 case spirv::Opcode::OpTypeSampledImage:
1232 case spirv::Opcode::OpTypeRuntimeArray:
1234 case spirv::Opcode::OpTypeStruct:
1236 case spirv::Opcode::OpTypeMatrix:
1238 case spirv::Opcode::OpTypeTensorARM:
1240 case spirv::Opcode::OpTypeGraphARM:
1243 return emitError(unknownLoc,
"unhandled type instruction");
1250 if (operands.size() != 3)
1251 return emitError(unknownLoc,
"OpTypePointer must have two parameters");
1253 auto pointeeType =
getType(operands[2]);
1255 return emitError(unknownLoc,
"unknown OpTypePointer pointee type <id> ")
1258 uint32_t typePointerID = operands[0];
1259 auto storageClass =
static_cast<spirv::StorageClass
>(operands[1]);
1262 for (
auto *deferredStructIt = std::begin(deferredStructTypesInfos);
1263 deferredStructIt != std::end(deferredStructTypesInfos);) {
1264 for (
auto *unresolvedMemberIt =
1265 std::begin(deferredStructIt->unresolvedMemberTypes);
1266 unresolvedMemberIt !=
1267 std::end(deferredStructIt->unresolvedMemberTypes);) {
1268 if (unresolvedMemberIt->first == typePointerID) {
1272 deferredStructIt->memberTypes[unresolvedMemberIt->second] =
1273 typeMap[typePointerID];
1274 unresolvedMemberIt =
1275 deferredStructIt->unresolvedMemberTypes.erase(unresolvedMemberIt);
1277 ++unresolvedMemberIt;
1281 if (deferredStructIt->unresolvedMemberTypes.empty()) {
1283 auto structType = deferredStructIt->deferredStructType;
1285 assert(structType &&
"expected a spirv::StructType");
1286 assert(structType.isIdentified() &&
"expected an indentified struct");
1288 if (failed(structType.trySetBody(
1289 deferredStructIt->memberTypes, deferredStructIt->offsetInfo,
1290 deferredStructIt->memberDecorationsInfo,
1291 deferredStructIt->structDecorationsInfo)))
1294 deferredStructIt = deferredStructTypesInfos.erase(deferredStructIt);
1305 if (operands.size() != 3) {
1307 "OpTypeArray must have element type and count parameters");
1312 return emitError(unknownLoc,
"OpTypeArray references undefined <id> ")
1320 return emitError(unknownLoc,
"OpTypeArray count <id> ")
1321 << operands[2] <<
"can only come from normal constant right now";
1324 if (
auto intVal = dyn_cast<IntegerAttr>(countInfo->first)) {
1325 count = intVal.getValue().getZExtValue();
1327 return emitError(unknownLoc,
"OpTypeArray count must come from a "
1328 "scalar integer constant instruction");
1332 elementTy, count, typeDecorations.lookup(operands[0]));
1338 assert(!operands.empty() &&
"No operands for processing function type");
1339 if (operands.size() == 1) {
1340 return emitError(unknownLoc,
"missing return type for OpTypeFunction");
1342 auto returnType =
getType(operands[1]);
1344 return emitError(unknownLoc,
"unknown return type in OpTypeFunction");
1347 for (
size_t i = 2, e = operands.size(); i < e; ++i) {
1348 auto ty =
getType(operands[i]);
1350 return emitError(unknownLoc,
"unknown argument type in OpTypeFunction");
1352 argTypes.push_back(ty);
1358 typeMap[operands[0]] = FunctionType::get(context, argTypes, returnTypes);
1364 if (operands.size() != 6) {
1366 "OpTypeCooperativeMatrixKHR must have element type, "
1367 "scope, row and column parameters, and use");
1373 "OpTypeCooperativeMatrixKHR references undefined <id> ")
1377 std::optional<spirv::Scope> scope =
1382 "OpTypeCooperativeMatrixKHR references undefined scope <id> ")
1391 return emitError(unknownLoc,
"OpTypeCooperativeMatrixKHR `Rows` references "
1392 "undefined constant <id> ")
1396 return emitError(unknownLoc,
"OpTypeCooperativeMatrixKHR `Columns` "
1397 "references undefined constant <id> ")
1401 return emitError(unknownLoc,
"OpTypeCooperativeMatrixKHR `Use` references "
1402 "undefined constant <id> ")
1405 unsigned rows = rowsAttr.getInt();
1406 unsigned columns = columnsAttr.getInt();
1408 std::optional<spirv::CooperativeMatrixUseKHR> use =
1409 spirv::symbolizeCooperativeMatrixUseKHR(useAttr.getInt());
1413 "OpTypeCooperativeMatrixKHR references undefined use <id> ")
1417 typeMap[operands[0]] =
1424 if (operands.size() != 2) {
1425 return emitError(unknownLoc,
"OpTypeRuntimeArray must have two operands");
1430 "OpTypeRuntimeArray references undefined <id> ")
1434 memberType, typeDecorations.lookup(operands[0]));
1442 if (operands.empty()) {
1443 return emitError(unknownLoc,
"OpTypeStruct must have at least result <id>");
1446 if (operands.size() == 1) {
1448 typeMap[operands[0]] =
1457 for (
auto op : llvm::drop_begin(operands, 1)) {
1459 bool typeForwardPtr = (typeForwardPointerIDs.count(op) != 0);
1461 if (!memberType && !typeForwardPtr)
1462 return emitError(unknownLoc,
"OpTypeStruct references undefined <id> ")
1466 unresolvedMemberTypes.emplace_back(op, memberTypes.size());
1468 memberTypes.push_back(memberType);
1473 if (memberDecorationMap.count(operands[0])) {
1474 auto &allMemberDecorations = memberDecorationMap[operands[0]];
1475 for (
auto memberIndex : llvm::seq<uint32_t>(0, memberTypes.size())) {
1476 if (allMemberDecorations.count(memberIndex)) {
1477 for (
auto &memberDecoration : allMemberDecorations[memberIndex]) {
1479 if (memberDecoration.first == spirv::Decoration::Offset) {
1481 if (offsetInfo.empty()) {
1482 offsetInfo.resize(memberTypes.size());
1484 offsetInfo[memberIndex] = memberDecoration.second[0];
1486 auto intType = mlir::IntegerType::get(context, 32);
1487 if (!memberDecoration.second.empty()) {
1488 memberDecorationsInfo.emplace_back(
1489 memberIndex, memberDecoration.first,
1490 IntegerAttr::get(intType, memberDecoration.second[0]));
1492 memberDecorationsInfo.emplace_back(
1493 memberIndex, memberDecoration.first, UnitAttr::get(context));
1502 if (decorations.count(operands[0])) {
1505 std::optional<spirv::Decoration> decoration = spirv::symbolizeDecoration(
1506 llvm::convertToCamelFromSnakeCase(decorationAttr.getName(),
true));
1507 assert(decoration.has_value());
1508 structDecorationsInfo.emplace_back(decoration.value(),
1509 decorationAttr.getValue());
1513 uint32_t structID = operands[0];
1514 std::string structIdentifier = nameMap.lookup(structID).str();
1516 if (structIdentifier.empty()) {
1517 assert(unresolvedMemberTypes.empty() &&
1518 "didn't expect unresolved member types");
1520 memberTypes, offsetInfo, memberDecorationsInfo, structDecorationsInfo);
1523 typeMap[structID] = structTy;
1525 if (!unresolvedMemberTypes.empty())
1526 deferredStructTypesInfos.push_back(
1527 {structTy, unresolvedMemberTypes, memberTypes, offsetInfo,
1528 memberDecorationsInfo, structDecorationsInfo});
1529 else if (failed(structTy.trySetBody(memberTypes, offsetInfo,
1530 memberDecorationsInfo,
1531 structDecorationsInfo)))
1542 if (operands.size() != 3) {
1544 return emitError(unknownLoc,
"OpTypeMatrix must have 3 operands"
1545 " (result_id, column_type, and column_count)");
1551 "OpTypeMatrix references undefined column type.")
1555 uint32_t colsCount = operands[2];
1562 unsigned size = operands.size();
1563 if (size < 2 || size > 4)
1564 return emitError(unknownLoc,
"OpTypeTensorARM must have 2-4 operands "
1565 "(result_id, element_type, (rank), (shape)) ")
1571 "OpTypeTensorARM references undefined element type ")
1581 return emitError(unknownLoc,
"OpTypeTensorARM rank must come from a "
1582 "scalar integer constant instruction");
1583 unsigned rank = rankAttr.getValue().getZExtValue();
1590 std::optional<std::pair<Attribute, Type>> shapeInfo =
1593 return emitError(unknownLoc,
"OpTypeTensorARM shape must come from a "
1594 "constant instruction of type OpTypeArray");
1596 ArrayAttr shapeArrayAttr = dyn_cast<ArrayAttr>(shapeInfo->first);
1598 for (
auto dimAttr : shapeArrayAttr.getValue()) {
1599 auto dimIntAttr = dyn_cast<IntegerAttr>(dimAttr);
1601 return emitError(unknownLoc,
"OpTypeTensorARM shape has an invalid "
1603 shape.push_back(dimIntAttr.getValue().getSExtValue());
1611 unsigned size = operands.size();
1613 return emitError(unknownLoc,
"OpTypeGraphARM must have at least 2 operands "
1614 "(result_id, num_inputs, (inout0_type, "
1615 "inout1_type, ...))")
1618 uint32_t numInputs = operands[1];
1621 for (
unsigned i = 2; i < size; ++i) {
1625 "OpTypeGraphARM references undefined element type.")
1628 if (i - 2 >= numInputs) {
1629 returnTypes.push_back(inOutTy);
1631 argTypes.push_back(inOutTy);
1634 typeMap[operands[0]] = GraphType::get(context, argTypes, returnTypes);
1640 if (operands.size() != 2)
1642 "OpTypeForwardPointer instruction must have two operands");
1644 typeForwardPointerIDs.insert(operands[0]);
1654 if (operands.size() != 8)
1657 "OpTypeImage with non-eight operands are not supported yet");
1661 return emitError(unknownLoc,
"OpTypeImage references undefined <id>: ")
1664 auto dim = spirv::symbolizeDim(operands[2]);
1666 return emitError(unknownLoc,
"unknown Dim for OpTypeImage: ")
1669 auto depthInfo = spirv::symbolizeImageDepthInfo(operands[3]);
1671 return emitError(unknownLoc,
"unknown Depth for OpTypeImage: ")
1674 auto arrayedInfo = spirv::symbolizeImageArrayedInfo(operands[4]);
1676 return emitError(unknownLoc,
"unknown Arrayed for OpTypeImage: ")
1679 auto samplingInfo = spirv::symbolizeImageSamplingInfo(operands[5]);
1681 return emitError(unknownLoc,
"unknown MS for OpTypeImage: ") << operands[5];
1683 auto samplerUseInfo = spirv::symbolizeImageSamplerUseInfo(operands[6]);
1684 if (!samplerUseInfo)
1685 return emitError(unknownLoc,
"unknown Sampled for OpTypeImage: ")
1688 auto format = spirv::symbolizeImageFormat(operands[7]);
1690 return emitError(unknownLoc,
"unknown Format for OpTypeImage: ")
1694 elementTy, dim.value(), depthInfo.value(), arrayedInfo.value(),
1695 samplingInfo.value(), samplerUseInfo.value(), format.value());
1701 if (operands.size() != 2)
1702 return emitError(unknownLoc,
"OpTypeSampledImage must have two operands");
1707 "OpTypeSampledImage references undefined <id>: ")
1716 if (operands.size() != 1)
1717 return emitError(unknownLoc,
"OpTypeSampler must have no parameters");
1725 if (operands.size() != 1)
1726 return emitError(unknownLoc,
"OpTypeNamedBarrier must have no parameters");
1738 StringRef opname = isSpec ?
"OpSpecConstant" :
"OpConstant";
1740 if (operands.size() < 2) {
1742 << opname <<
" must have type <id> and result <id>";
1744 if (operands.size() < 3) {
1746 << opname <<
" must have at least 1 more parameter";
1751 return emitError(unknownLoc,
"undefined result type from <id> ")
1755 auto checkOperandSizeForBitwidth = [&](
unsigned bitwidth) -> LogicalResult {
1756 if (bitwidth == 64) {
1757 if (operands.size() == 4) {
1761 << opname <<
" should have 2 parameters for 64-bit values";
1763 if (bitwidth <= 32) {
1764 if (operands.size() == 3) {
1770 <<
" should have 1 parameter for values with no more than 32 bits";
1772 return emitError(unknownLoc,
"unsupported OpConstant bitwidth: ")
1776 auto resultID = operands[1];
1778 if (
auto intType = dyn_cast<IntegerType>(resultType)) {
1779 auto bitwidth = intType.getWidth();
1780 if (failed(checkOperandSizeForBitwidth(bitwidth))) {
1785 if (bitwidth == 64) {
1792 } words = {operands[2], operands[3]};
1793 value = APInt(64, llvm::bit_cast<uint64_t>(words),
true);
1794 }
else if (bitwidth <= 32) {
1795 value = APInt(bitwidth, operands[2],
true,
1799 auto attr = opBuilder.getIntegerAttr(intType, value);
1806 constantMap.try_emplace(resultID, attr, intType);
1812 if (
auto floatType = dyn_cast<FloatType>(resultType)) {
1813 auto bitwidth = floatType.getWidth();
1814 if (failed(checkOperandSizeForBitwidth(bitwidth))) {
1819 if (floatType.isF64()) {
1826 } words = {operands[2], operands[3]};
1827 value = APFloat(llvm::bit_cast<double>(words));
1828 }
else if (floatType.isF32()) {
1829 value = APFloat(llvm::bit_cast<float>(operands[2]));
1830 }
else if (floatType.isF16()) {
1831 APInt data(16, operands[2]);
1832 value = APFloat(APFloat::IEEEhalf(), data);
1833 }
else if (floatType.isBF16()) {
1834 APInt data(16, operands[2]);
1835 value = APFloat(APFloat::BFloat(), data);
1836 }
else if (floatType.isF8E4M3FN()) {
1837 APInt data(8, operands[2]);
1838 value = APFloat(APFloat::Float8E4M3FN(), data);
1839 }
else if (floatType.isF8E5M2()) {
1840 APInt data(8, operands[2]);
1841 value = APFloat(APFloat::Float8E5M2(), data);
1844 auto attr = opBuilder.getFloatAttr(floatType, value);
1850 constantMap.try_emplace(resultID, attr, floatType);
1856 return emitError(unknownLoc,
"OpConstant can only generate values of "
1857 "scalar integer or floating-point type");
1862 if (operands.size() != 2) {
1864 << (isSpec ?
"Spec" :
"") <<
"Constant"
1865 << (isTrue ?
"True" :
"False")
1866 <<
" must have type <id> and result <id>";
1869 auto attr = opBuilder.getBoolAttr(isTrue);
1870 auto resultID = operands[1];
1876 constantMap.try_emplace(resultID, attr, opBuilder.getI1Type());
1884 if (operands.size() < 2) {
1886 "OpConstantComposite must have type <id> and result <id>");
1888 if (operands.size() < 3) {
1890 "OpConstantComposite must have at least 1 parameter");
1895 return emitError(unknownLoc,
"undefined result type from <id> ")
1900 elements.reserve(operands.size() - 2);
1901 for (
unsigned i = 2, e = operands.size(); i < e; ++i) {
1904 return emitError(unknownLoc,
"OpConstantComposite component <id> ")
1905 << operands[i] <<
" must come from a normal constant";
1907 elements.push_back(elementInfo->first);
1910 auto resultID = operands[1];
1911 if (
auto tensorType = dyn_cast<TensorArmType>(resultType)) {
1914 if (
auto denseElemAttr = dyn_cast<DenseElementsAttr>(element)) {
1915 for (
auto value : denseElemAttr.getValues<
Attribute>())
1916 flattenedElems.push_back(value);
1918 flattenedElems.push_back(element);
1922 constantMap.try_emplace(resultID, attr, tensorType);
1923 }
else if (
auto shapedType = dyn_cast<ShapedType>(resultType)) {
1927 constantMap.try_emplace(resultID, attr, shapedType);
1928 }
else if (isa<spirv::ArrayType, spirv::StructType>(resultType)) {
1929 auto attr = opBuilder.getArrayAttr(elements);
1930 constantMap.try_emplace(resultID, attr, resultType);
1932 return emitError(unknownLoc,
"unsupported OpConstantComposite type: ")
1941 if (operands.size() != 3) {
1944 "OpConstantCompositeReplicateEXT expects 3 operands but found ")
1950 return emitError(unknownLoc,
"undefined result type from <id> ")
1954 auto compositeType = dyn_cast<CompositeType>(resultType);
1955 if (!compositeType) {
1957 "result type from <id> is not a composite type")
1961 uint32_t resultID = operands[1];
1962 uint32_t constantID = operands[2];
1964 std::optional<std::pair<Attribute, Type>> replicatedConstantCompositeInfo =
1966 if (replicatedConstantCompositeInfo.has_value()) {
1967 constantCompositeReplicateMap.try_emplace(
1968 resultID, replicatedConstantCompositeInfo.value().first, resultType);
1972 std::optional<std::pair<Attribute, Type>> constantInfo =
1974 if (constantInfo.has_value()) {
1975 constantCompositeReplicateMap.try_emplace(
1976 resultID, constantInfo.value().first, resultType);
1980 return emitError(unknownLoc,
"OpConstantCompositeReplicateEXT operand <id> ")
1982 <<
" must come from a normal constant or a "
1983 "OpConstantCompositeReplicateEXT";
1988 if (operands.size() < 2) {
1991 "OpSpecConstantComposite must have type <id> and result <id>");
1993 if (operands.size() < 3) {
1995 "OpSpecConstantComposite must have at least 1 parameter");
2000 return emitError(unknownLoc,
"undefined result type from <id> ")
2004 auto resultID = operands[1];
2008 elements.reserve(operands.size() - 2);
2009 for (
unsigned i = 2, e = operands.size(); i < e; ++i) {
2011 elements.push_back(SymbolRefAttr::get(elementInfo));
2014 auto op = spirv::SpecConstantCompositeOp::create(
2015 opBuilder, unknownLoc, TypeAttr::get(resultType), symName,
2016 opBuilder.getArrayAttr(elements));
2017 specConstCompositeMap[resultID] = op;
2024 if (operands.size() != 3) {
2025 return emitError(unknownLoc,
"OpSpecConstantCompositeReplicateEXT expects "
2026 "3 operands but found ")
2032 return emitError(unknownLoc,
"undefined result type from <id> ")
2036 auto compositeType = dyn_cast<CompositeType>(resultType);
2037 if (!compositeType) {
2039 "result type from <id> is not a composite type")
2043 uint32_t resultID = operands[1];
2046 spirv::SpecConstantOp constituentSpecConstantOp =
2048 auto op = spirv::EXTSpecConstantCompositeReplicateOp::create(
2049 opBuilder, unknownLoc, TypeAttr::get(resultType), symName,
2050 SymbolRefAttr::get(constituentSpecConstantOp));
2052 specConstCompositeReplicateMap[resultID] = op;
2059 if (operands.size() < 3)
2060 return emitError(unknownLoc,
"OpConstantOperation must have type <id>, "
2061 "result <id>, and operand opcode");
2063 uint32_t resultTypeID = operands[0];
2066 return emitError(unknownLoc,
"undefined result type from <id> ")
2069 uint32_t resultID = operands[1];
2070 spirv::Opcode enclosedOpcode =
static_cast<spirv::Opcode
>(operands[2]);
2071 auto emplaceResult = specConstOperationMap.try_emplace(
2074 enclosedOpcode, resultTypeID,
2077 if (!emplaceResult.second)
2078 return emitError(unknownLoc,
"value with <id>: ")
2079 << resultID <<
" is probably defined before.";
2085 uint32_t resultID, spirv::Opcode enclosedOpcode, uint32_t resultTypeID,
2101 llvm::SaveAndRestore valueMapGuard(valueMap, newValueMap);
2102 constexpr uint32_t fakeID =
static_cast<uint32_t
>(-3);
2105 enclosedOpResultTypeAndOperands.push_back(resultTypeID);
2106 enclosedOpResultTypeAndOperands.push_back(fakeID);
2107 enclosedOpResultTypeAndOperands.append(enclosedOpOperands.begin(),
2108 enclosedOpOperands.end());
2123 auto specConstOperationOp =
2124 spirv::SpecConstantOperationOp::create(opBuilder, loc, resultType);
2126 Region &body = specConstOperationOp.getBody();
2128 body.
getBlocks().splice(body.
end(), curBlock->getParent()->getBlocks(),
2135 opBuilder.setInsertionPointToEnd(&block);
2137 spirv::YieldOp::create(opBuilder, loc, block.
front().
getResult(0));
2138 return specConstOperationOp.getResult();
2143 if (operands.size() != 2) {
2145 "OpConstantNull must only have type <id> and result <id>");
2150 return emitError(unknownLoc,
"undefined result type from <id> ")
2154 auto resultID = operands[1];
2156 if (resultType.
isIntOrFloat() || isa<VectorType>(resultType)) {
2157 attr = opBuilder.getZeroAttr(resultType);
2158 }
else if (
auto tensorType = dyn_cast<TensorArmType>(resultType)) {
2159 if (
auto element = opBuilder.getZeroAttr(tensorType.getElementType()))
2166 constantMap.try_emplace(resultID, attr, resultType);
2170 return emitError(unknownLoc,
"unsupported OpConstantNull type: ")
2176 if (operands.size() < 3) {
2178 <<
"OpGraphConstantARM must have at least 2 operands";
2183 return emitError(unknownLoc,
"undefined result type from <id> ")
2187 uint32_t resultID = operands[1];
2189 if (!dyn_cast<spirv::TensorArmType>(resultType)) {
2190 return emitError(unknownLoc,
"result must be of type OpTypeTensorARM");
2193 APInt graph_constant_id = APInt(32, operands[2],
true);
2194 Type i32Ty = opBuilder.getIntegerType(32);
2195 IntegerAttr attr = opBuilder.getIntegerAttr(i32Ty, graph_constant_id);
2196 graphConstantMap.try_emplace(
2208 LLVM_DEBUG(logger.startLine() <<
"[block] got exiting block for id = " <<
id
2209 <<
" @ " << block <<
"\n");
2216 auto *block = curFunction->addBlock();
2217 LLVM_DEBUG(logger.startLine() <<
"[block] created block for id = " <<
id
2218 <<
" @ " << block <<
"\n");
2219 return blockMap[id] = block;
2224 return emitError(unknownLoc,
"OpBranch must appear inside a block");
2227 if (operands.size() != 1) {
2228 return emitError(unknownLoc,
"OpBranch must take exactly one target label");
2236 spirv::BranchOp::create(opBuilder, loc,
target);
2246 "OpBranchConditional must appear inside a block");
2249 if (operands.size() != 3 && operands.size() != 5) {
2251 "OpBranchConditional must have condition, true label, "
2252 "false label, and optionally two branch weights");
2255 auto condition =
getValue(operands[0]);
2259 std::optional<std::pair<uint32_t, uint32_t>> weights;
2260 if (operands.size() == 5) {
2261 weights = std::make_pair(operands[3], operands[4]);
2267 spirv::BranchConditionalOp::create(
2268 opBuilder, loc, condition, trueBlock,
2278 return emitError(unknownLoc,
"OpLabel must appear inside a function");
2281 if (operands.size() != 1) {
2282 return emitError(unknownLoc,
"OpLabel should only have result <id>");
2285 auto labelID = operands[0];
2288 LLVM_DEBUG(logger.startLine()
2289 <<
"[block] populating block " << block <<
"\n");
2291 assert(block->empty() &&
"re-deserialize the same block!");
2293 opBuilder.setInsertionPointToStart(block);
2294 blockMap[labelID] = curBlock = block;
2301 return emitError(unknownLoc,
"a graph block must appear inside a graph");
2306 LLVM_DEBUG(logger.startLine()
2307 <<
"[block] populating block " << block <<
"\n");
2309 assert(block->
empty() &&
"re-deserialize the same block!");
2311 opBuilder.setInsertionPointToStart(block);
2312 blockMap[graphID] = curBlock = block;
2320 return emitError(unknownLoc,
"OpSelectionMerge must appear in a block");
2323 if (operands.size() < 2) {
2326 "OpSelectionMerge must specify merge target and selection control");
2331 auto selectionControl = operands[1];
2333 if (!blockMergeInfo.try_emplace(curBlock, loc, selectionControl, mergeBlock)
2337 "a block cannot have more than one OpSelectionMerge instruction");
2346 return emitError(unknownLoc,
"OpLoopMerge must appear in a block");
2349 if (operands.size() < 3) {
2350 return emitError(unknownLoc,
"OpLoopMerge must specify merge target, "
2351 "continue target and loop control");
2357 uint32_t loopControl = operands[2];
2360 .try_emplace(curBlock, loc, loopControl, mergeBlock, continueBlock)
2364 "a block cannot have more than one OpLoopMerge instruction");
2372 return emitError(unknownLoc,
"OpPhi must appear in a block");
2375 if (operands.size() < 4) {
2376 return emitError(unknownLoc,
"OpPhi must specify result type, result <id>, "
2377 "and variable-parent pairs");
2382 BlockArgument blockArg = curBlock->addArgument(blockArgType, unknownLoc);
2383 valueMap[operands[1]] = blockArg;
2384 LLVM_DEBUG(logger.startLine()
2385 <<
"[phi] created block argument " << blockArg
2386 <<
" id = " << operands[1] <<
" of type " << blockArgType <<
"\n");
2390 for (
unsigned i = 2, e = operands.size(); i < e; i += 2) {
2391 uint32_t value = operands[i];
2393 std::pair<Block *, Block *> predecessorTargetPair{predecessor, curBlock};
2394 blockPhiInfo[predecessorTargetPair].push_back(value);
2395 LLVM_DEBUG(logger.startLine() <<
"[phi] predecessor @ " << predecessor
2396 <<
" with arg id = " << value <<
"\n");
2404 return emitError(unknownLoc,
"OpSwitch must appear in a block");
2406 if (operands.size() < 2)
2407 return emitError(unknownLoc,
"OpSwitch must at least specify selector and "
2408 "a default target");
2410 if (operands.size() % 2)
2412 "OpSwitch must at have an even number of operands: "
2413 "selector, default target and any number of literal and "
2414 "label <id> pairs");
2422 for (
unsigned i = 2, e = operands.size(); i < e; i += 2) {
2423 literals.push_back(operands[i]);
2428 spirv::SwitchOp::create(opBuilder, loc, selector, defaultBlock,
2437class ControlFlowStructurizer {
2440 ControlFlowStructurizer(
Location loc, uint32_t control,
2443 llvm::ScopedPrinter &logger)
2444 : location(loc), control(control), blockMergeInfo(mergeInfo),
2445 headerBlock(header), mergeBlock(merge), continueBlock(cont),
2448 ControlFlowStructurizer(
Location loc, uint32_t control,
2451 : location(loc), control(control), blockMergeInfo(mergeInfo),
2452 headerBlock(header), mergeBlock(merge), continueBlock(cont) {}
2462 LogicalResult structurize();
2467 spirv::SelectionOp createSelectionOp(uint32_t selectionControl);
2470 spirv::LoopOp createLoopOp(uint32_t loopControl);
2473 void collectBlocksInConstruct();
2482 Block *continueBlock;
2488 llvm::ScopedPrinter &logger;
2494ControlFlowStructurizer::createSelectionOp(uint32_t selectionControl) {
2497 OpBuilder builder(&mergeBlock->front());
2499 auto control =
static_cast<spirv::SelectionControl
>(selectionControl);
2500 auto selectionOp = spirv::SelectionOp::create(builder, location, control);
2501 selectionOp.addMergeBlock(builder);
2506spirv::LoopOp ControlFlowStructurizer::createLoopOp(uint32_t loopControl) {
2509 OpBuilder builder(&mergeBlock->front());
2511 auto control =
static_cast<spirv::LoopControl
>(loopControl);
2512 auto loopOp = spirv::LoopOp::create(builder, location, control);
2513 loopOp.addEntryAndMergeBlock(builder);
2518void ControlFlowStructurizer::collectBlocksInConstruct() {
2519 assert(constructBlocks.empty() &&
"expected empty constructBlocks");
2522 constructBlocks.insert(headerBlock);
2526 for (
unsigned i = 0; i < constructBlocks.size(); ++i) {
2527 for (
auto *successor : constructBlocks[i]->getSuccessors())
2528 if (successor != mergeBlock)
2529 constructBlocks.insert(successor);
2533LogicalResult ControlFlowStructurizer::structurize() {
2534 Operation *op =
nullptr;
2535 bool isLoop = continueBlock !=
nullptr;
2537 if (
auto loopOp = createLoopOp(control))
2538 op = loopOp.getOperation();
2540 if (
auto selectionOp = createSelectionOp(control))
2541 op = selectionOp.getOperation();
2550 mapper.
map(mergeBlock, &body.
back());
2552 collectBlocksInConstruct();
2573 OpBuilder builder(body);
2574 for (
auto *block : constructBlocks) {
2577 auto *newBlock = builder.createBlock(&body.
back());
2578 mapper.
map(block, newBlock);
2579 LLVM_DEBUG(logger.startLine() <<
"[cf] cloned block " << newBlock
2580 <<
" from block " << block <<
"\n");
2582 for (BlockArgument blockArg : block->getArguments()) {
2584 newBlock->addArgument(blockArg.getType(), blockArg.getLoc());
2585 mapper.
map(blockArg, newArg);
2586 LLVM_DEBUG(logger.startLine() <<
"[cf] remapped block argument "
2587 << blockArg <<
" to " << newArg <<
"\n");
2590 LLVM_DEBUG(logger.startLine()
2591 <<
"[cf] block " << block <<
" is a function entry block\n");
2594 for (
auto &op : *block)
2595 newBlock->push_back(op.
clone(mapper));
2599 auto remapOperands = [&](Operation *op) {
2601 if (Value mappedOp = mapper.
lookupOrNull(operand.get()))
2602 operand.set(mappedOp);
2605 succOp.set(mappedOp);
2607 for (
auto &block : body)
2608 block.walk(remapOperands);
2616 headerBlock->replaceAllUsesWith(mergeBlock);
2619 logger.startLine() <<
"[cf] after cloning and fixing references:\n";
2620 headerBlock->getParentOp()->print(logger.getOStream());
2621 logger.startLine() <<
"\n";
2625 if (!mergeBlock->args_empty()) {
2626 return mergeBlock->getParentOp()->emitError(
2627 "OpPhi in loop merge block unsupported");
2633 for (BlockArgument blockArg : headerBlock->getArguments())
2634 mergeBlock->addArgument(blockArg.getType(), blockArg.getLoc());
2638 SmallVector<Value, 4> blockArgs;
2639 if (!headerBlock->args_empty())
2640 blockArgs = {mergeBlock->args_begin(), mergeBlock->args_end()};
2644 builder.setInsertionPointToEnd(&body.front());
2645 spirv::BranchOp::create(builder, location, mapper.
lookupOrNull(headerBlock),
2646 ArrayRef<Value>(blockArgs));
2651 SmallVector<Value> valuesToYield;
2654 SmallVector<Value> outsideUses;
2668 for (BlockArgument blockArg : mergeBlock->getArguments()) {
2673 body.back().addArgument(blockArg.getType(), blockArg.getLoc());
2674 valuesToYield.push_back(body.back().getArguments().back());
2675 outsideUses.push_back(blockArg);
2680 LLVM_DEBUG(logger.startLine() <<
"[cf] cleaning up blocks after clone\n");
2683 for (
auto *block : constructBlocks)
2684 block->dropAllReferences();
2689 for (
Block *block : constructBlocks) {
2690 for (Operation &op : *block) {
2694 outsideUses.push_back(
result);
2697 for (BlockArgument &arg : block->getArguments()) {
2698 if (!arg.use_empty()) {
2700 outsideUses.push_back(arg);
2705 assert(valuesToYield.size() == outsideUses.size());
2709 if (!valuesToYield.empty()) {
2710 LLVM_DEBUG(logger.startLine()
2711 <<
"[cf] yielding values from the selection / loop region\n");
2714 auto mergeOps = body.back().getOps<spirv::MergeOp>();
2715 Operation *merge = llvm::getSingleElement(mergeOps);
2717 merge->setOperands(valuesToYield);
2725 builder.setInsertionPoint(&mergeBlock->front());
2727 Operation *newOp =
nullptr;
2730 newOp = spirv::LoopOp::create(builder, location,
2732 static_cast<spirv::LoopControl
>(control));
2734 newOp = spirv::SelectionOp::create(
2736 static_cast<spirv::SelectionControl
>(control));
2746 for (
unsigned i = 0, e = outsideUses.size(); i != e; ++i)
2747 outsideUses[i].replaceAllUsesWith(op->
getResult(i));
2753 mergeBlock->eraseArguments(0, mergeBlock->getNumArguments());
2760 for (
auto *block : constructBlocks) {
2761 if (!block->use_empty())
2762 return emitError(block->getParent()->getLoc(),
2763 "failed control flow structurization: "
2764 "block has uses outside of the "
2765 "enclosing selection/loop construct");
2766 for (Operation &op : *block)
2768 return op.
emitOpError(
"failed control flow structurization: value has "
2769 "uses outside of the "
2770 "enclosing selection/loop construct");
2771 for (BlockArgument &arg : block->getArguments())
2772 if (!arg.use_empty())
2773 return emitError(arg.getLoc(),
"failed control flow structurization: "
2774 "block argument has uses outside of the "
2775 "enclosing selection/loop construct");
2779 for (
auto *block : constructBlocks) {
2819 auto updateMergeInfo = [&](
Block *block) -> WalkResult {
2820 auto it = blockMergeInfo.find(block);
2821 if (it != blockMergeInfo.end()) {
2823 Location loc = it->second.loc;
2827 return emitError(loc,
"failed control flow structurization: nested "
2828 "loop header block should be remapped!");
2830 Block *newContinue = it->second.continueBlock;
2834 return emitError(loc,
"failed control flow structurization: nested "
2835 "loop continue block should be remapped!");
2838 Block *newMerge = it->second.mergeBlock;
2840 newMerge = mappedTo;
2844 blockMergeInfo.
erase(it);
2845 blockMergeInfo.try_emplace(newHeader, loc, it->second.control, newMerge,
2852 if (block->walk(updateMergeInfo).wasInterrupted())
2860 LLVM_DEBUG(logger.startLine() <<
"[cf] changing entry block " << block
2861 <<
" to only contain a spirv.Branch op\n");
2865 builder.setInsertionPointToEnd(block);
2866 spirv::BranchOp::create(builder, location, mergeBlock);
2868 LLVM_DEBUG(logger.startLine() <<
"[cf] erasing block " << block <<
"\n");
2873 LLVM_DEBUG(logger.startLine()
2874 <<
"[cf] after structurizing construct with header block "
2875 << headerBlock <<
":\n"
2884 <<
"//----- [phi] start wiring up block arguments -----//\n";
2890 for (
const auto &info : blockPhiInfo) {
2891 Block *block = info.first.first;
2895 logger.startLine() <<
"[phi] block " << block <<
"\n";
2896 logger.startLine() <<
"[phi] before creating block argument:\n";
2898 logger.startLine() <<
"\n";
2904 opBuilder.setInsertionPoint(op);
2907 blockArgs.reserve(phiInfo.size());
2908 for (uint32_t valueId : phiInfo) {
2910 blockArgs.push_back(value);
2911 LLVM_DEBUG(logger.startLine() <<
"[phi] block argument " << value
2912 <<
" id = " << valueId <<
"\n");
2914 return emitError(unknownLoc,
"OpPhi references undefined value!");
2918 if (
auto branchOp = dyn_cast<spirv::BranchOp>(op)) {
2920 spirv::BranchOp::create(opBuilder, branchOp.getLoc(),
2921 branchOp.getTarget(), blockArgs);
2923 }
else if (
auto branchCondOp = dyn_cast<spirv::BranchConditionalOp>(op)) {
2924 assert((branchCondOp.getTrueBlock() ==
target ||
2925 branchCondOp.getFalseBlock() ==
target) &&
2926 "expected target to be either the true or false target");
2927 if (
target == branchCondOp.getTrueTarget())
2928 spirv::BranchConditionalOp::create(
2929 opBuilder, branchCondOp.getLoc(), branchCondOp.getCondition(),
2930 blockArgs, branchCondOp.getFalseBlockArguments(),
2931 branchCondOp.getBranchWeightsAttr(), branchCondOp.getTrueTarget(),
2932 branchCondOp.getFalseTarget());
2934 spirv::BranchConditionalOp::create(
2935 opBuilder, branchCondOp.getLoc(), branchCondOp.getCondition(),
2936 branchCondOp.getTrueBlockArguments(), blockArgs,
2937 branchCondOp.getBranchWeightsAttr(), branchCondOp.getTrueBlock(),
2938 branchCondOp.getFalseBlock());
2940 branchCondOp.erase();
2941 }
else if (
auto switchOp = dyn_cast<spirv::SwitchOp>(op)) {
2942 if (
target == switchOp.getDefaultTarget()) {
2946 spirv::SwitchOp::create(
2947 opBuilder, switchOp.getLoc(), switchOp.getSelector(),
2948 switchOp.getDefaultTarget(), blockArgs, literals,
2949 switchOp.getTargets(), targetOperands);
2953 auto it = llvm::find(targets,
target);
2954 assert(it != targets.end());
2955 size_t index = std::distance(targets.begin(), it);
2956 switchOp.getTargetOperandsMutable(
index).assign(blockArgs);
2959 return emitError(unknownLoc,
"unimplemented terminator for Phi creation");
2963 logger.startLine() <<
"[phi] after creating block argument:\n";
2965 logger.startLine() <<
"\n";
2968 blockPhiInfo.clear();
2973 <<
"//--- [phi] completed wiring up block arguments ---//\n";
2981 for (
auto [block, mergeInfo] : blockMergeInfoCopy) {
2983 if (mergeInfo.continueBlock)
2986 if (!block->mightHaveTerminator())
2989 Operation *terminator = block->getTerminator();
2992 if (!isa<spirv::BranchConditionalOp, spirv::SwitchOp>(terminator))
2996 bool splitHeaderMergeBlock =
false;
2997 for (
const auto &[_, mergeInfo] : blockMergeInfo) {
2998 if (mergeInfo.mergeBlock == block)
2999 splitHeaderMergeBlock =
true;
3006 if (!llvm::hasSingleElement(*block) || splitHeaderMergeBlock) {
3009 spirv::BranchOp::create(builder, block->getParent()->getLoc(), newBlock);
3013 blockMergeInfo.erase(block);
3014 blockMergeInfo.try_emplace(newBlock, mergeInfo);
3022 if (!options.enableControlFlowStructurization) {
3026 <<
"//----- [cf] skip structurizing control flow -----//\n";
3034 <<
"//----- [cf] start structurizing control flow -----//\n";
3039 logger.startLine() <<
"[cf] split conditional blocks\n";
3040 logger.startLine() <<
"\n";
3047 while (!blockMergeInfo.empty()) {
3048 Block *headerBlock = blockMergeInfo.
begin()->first;
3052 logger.startLine() <<
"[cf] header block " << headerBlock <<
":\n";
3053 headerBlock->
print(logger.getOStream());
3054 logger.startLine() <<
"\n";
3058 assert(mergeBlock &&
"merge block cannot be nullptr");
3060 return emitError(unknownLoc,
"OpPhi in loop merge block unimplemented");
3062 logger.startLine() <<
"[cf] merge block " << mergeBlock <<
":\n";
3063 mergeBlock->print(logger.getOStream());
3064 logger.startLine() <<
"\n";
3068 LLVM_DEBUG(
if (continueBlock) {
3069 logger.startLine() <<
"[cf] continue block " << continueBlock <<
":\n";
3070 continueBlock->print(logger.getOStream());
3071 logger.startLine() <<
"\n";
3075 blockMergeInfo.
erase(blockMergeInfo.begin());
3076 ControlFlowStructurizer structurizer(mergeInfo.
loc, mergeInfo.
control,
3077 blockMergeInfo, headerBlock,
3078 mergeBlock, continueBlock
3084 if (failed(structurizer.structurize()))
3091 <<
"//--- [cf] completed structurizing control flow ---//\n";
3104 auto fileName = debugInfoMap.lookup(debugLine->fileID).str();
3105 if (fileName.empty())
3106 fileName =
"<unknown>";
3118 if (operands.size() != 3)
3119 return emitError(unknownLoc,
"OpLine must have 3 operands");
3120 debugLine =
DebugLine{operands[0], operands[1], operands[2]};
3128 if (operands.size() < 2)
3129 return emitError(unknownLoc,
"OpString needs at least 2 operands");
3131 if (!debugInfoMap.lookup(operands[0]).empty())
3133 "duplicate debug string found for result <id> ")
3136 unsigned wordIndex = 1;
3138 if (wordIndex != operands.size())
3140 "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 void setInherentOrDiscardableAttr(Operation *op, StringAttr name, Attribute value)
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.
void setInherentAttr(Operation *op, StringAttr name, Attribute value) const
std::optional< Attribute > getInherentAttr(Operation *op, StringRef name) const
Lookup an inherent attribute by name, this method isn't recommended and may be removed in the future.
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.
void setDiscardableAttr(StringAttr name, Attribute value)
Set a discardable attribute by name.
OpResult getResult(unsigned idx)
Get the 'idx'th result of this operation.
MutableArrayRef< OpOperand > getOpOperands()
OperationName getName()
The name of an operation is the key identifier for it.
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)
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.