1 //===- SPIRVSerializationGen.cpp - SPIR-V serialization utility generator -===//
2 //
3 // Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.
4 // See https://llvm.org/LICENSE.txt for license information.
5 // SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
6 //
7 //===----------------------------------------------------------------------===//
8 //
9 // SPIRVSerializationGen generates common utility functions for SPIR-V
10 // serialization.
11 //
12 //===----------------------------------------------------------------------===//
13 
14 #include "mlir/TableGen/Attribute.h"
15 #include "mlir/TableGen/CodeGenHelpers.h"
16 #include "mlir/TableGen/Format.h"
17 #include "mlir/TableGen/GenInfo.h"
18 #include "mlir/TableGen/Operator.h"
19 #include "llvm/ADT/Sequence.h"
20 #include "llvm/ADT/SmallVector.h"
21 #include "llvm/ADT/StringExtras.h"
22 #include "llvm/ADT/StringMap.h"
23 #include "llvm/ADT/StringRef.h"
24 #include "llvm/ADT/StringSet.h"
25 #include "llvm/Support/FormatVariadic.h"
26 #include "llvm/Support/raw_ostream.h"
27 #include "llvm/TableGen/Error.h"
28 #include "llvm/TableGen/Record.h"
29 #include "llvm/TableGen/TableGenBackend.h"
30 
31 #include <list>
32 
33 using llvm::ArrayRef;
34 using llvm::formatv;
35 using llvm::raw_ostream;
36 using llvm::raw_string_ostream;
37 using llvm::Record;
38 using llvm::RecordKeeper;
39 using llvm::SmallVector;
40 using llvm::SMLoc;
41 using llvm::StringMap;
42 using llvm::StringRef;
43 using mlir::tblgen::Attribute;
44 using mlir::tblgen::EnumAttr;
45 using mlir::tblgen::EnumAttrCase;
46 using mlir::tblgen::NamedAttribute;
47 using mlir::tblgen::NamedTypeConstraint;
48 using mlir::tblgen::NamespaceEmitter;
49 using mlir::tblgen::Operator;
50 
51 //===----------------------------------------------------------------------===//
52 // Availability Wrapper Class
53 //===----------------------------------------------------------------------===//
54 
55 namespace {
56 // Wrapper class with helper methods for accessing availability defined in
57 // TableGen.
58 class Availability {
59 public:
60   explicit Availability(const Record *def);
61 
62   // Returns the name of the direct TableGen class for this availability
63   // instance.
64   StringRef getClass() const;
65 
66   // Returns the generated C++ interface's class namespace.
67   StringRef getInterfaceClassNamespace() const;
68 
69   // Returns the generated C++ interface's class name.
70   StringRef getInterfaceClassName() const;
71 
72   // Returns the generated C++ interface's description.
73   StringRef getInterfaceDescription() const;
74 
75   // Returns the name of the query function insided the generated C++ interface.
76   StringRef getQueryFnName() const;
77 
78   // Returns the return type of the query function insided the generated C++
79   // interface.
80   StringRef getQueryFnRetType() const;
81 
82   // Returns the code for merging availability requirements.
83   StringRef getMergeActionCode() const;
84 
85   // Returns the initializer expression for initializing the final availability
86   // requirements.
87   StringRef getMergeInitializer() const;
88 
89   // Returns the C++ type for an availability instance.
90   StringRef getMergeInstanceType() const;
91 
92   // Returns the C++ statements for preparing availability instance.
93   StringRef getMergeInstancePreparation() const;
94 
95   // Returns the concrete availability instance carried in this case.
96   StringRef getMergeInstance() const;
97 
98   // Returns the underlying LLVM TableGen Record.
getDef() const99   const llvm::Record *getDef() const { return def; }
100 
101 private:
102   // The TableGen definition of this availability.
103   const llvm::Record *def;
104 };
105 } // namespace
106 
Availability(const llvm::Record * def)107 Availability::Availability(const llvm::Record *def) : def(def) {
108   assert(def->isSubClassOf("Availability") &&
109          "must be subclass of TableGen 'Availability' class");
110 }
111 
getClass() const112 StringRef Availability::getClass() const {
113   SmallVector<Record *, 1> parentClass;
114   def->getDirectSuperClasses(parentClass);
115   if (parentClass.size() != 1) {
116     PrintFatalError(def->getLoc(),
117                     "expected to only have one direct superclass");
118   }
119   return parentClass.front()->getName();
120 }
121 
getInterfaceClassNamespace() const122 StringRef Availability::getInterfaceClassNamespace() const {
123   return def->getValueAsString("cppNamespace");
124 }
125 
getInterfaceClassName() const126 StringRef Availability::getInterfaceClassName() const {
127   return def->getValueAsString("interfaceName");
128 }
129 
getInterfaceDescription() const130 StringRef Availability::getInterfaceDescription() const {
131   return def->getValueAsString("interfaceDescription");
132 }
133 
getQueryFnRetType() const134 StringRef Availability::getQueryFnRetType() const {
135   return def->getValueAsString("queryFnRetType");
136 }
137 
getQueryFnName() const138 StringRef Availability::getQueryFnName() const {
139   return def->getValueAsString("queryFnName");
140 }
141 
getMergeActionCode() const142 StringRef Availability::getMergeActionCode() const {
143   return def->getValueAsString("mergeAction");
144 }
145 
getMergeInitializer() const146 StringRef Availability::getMergeInitializer() const {
147   return def->getValueAsString("initializer");
148 }
149 
getMergeInstanceType() const150 StringRef Availability::getMergeInstanceType() const {
151   return def->getValueAsString("instanceType");
152 }
153 
getMergeInstancePreparation() const154 StringRef Availability::getMergeInstancePreparation() const {
155   return def->getValueAsString("instancePreparation");
156 }
157 
getMergeInstance() const158 StringRef Availability::getMergeInstance() const {
159   return def->getValueAsString("instance");
160 }
161 
162 // Returns the availability spec of the given `def`.
getAvailabilities(const Record & def)163 std::vector<Availability> getAvailabilities(const Record &def) {
164   std::vector<Availability> availabilities;
165 
166   if (def.getValue("availability")) {
167     std::vector<Record *> availDefs = def.getValueAsListOfDefs("availability");
168     availabilities.reserve(availDefs.size());
169     for (const Record *avail : availDefs)
170       availabilities.emplace_back(avail);
171   }
172 
173   return availabilities;
174 }
175 
176 //===----------------------------------------------------------------------===//
177 // Availability Interface Definitions AutoGen
178 //===----------------------------------------------------------------------===//
179 
emitInterfaceDef(const Availability & availability,raw_ostream & os)180 static void emitInterfaceDef(const Availability &availability,
181                              raw_ostream &os) {
182 
183   os << availability.getQueryFnRetType() << " ";
184 
185   StringRef cppNamespace = availability.getInterfaceClassNamespace();
186   cppNamespace.consume_front("::");
187   if (!cppNamespace.empty())
188     os << cppNamespace << "::";
189 
190   StringRef methodName = availability.getQueryFnName();
191   os << availability.getInterfaceClassName() << "::" << methodName << "() {\n"
192      << "  return getImpl()->" << methodName << "(getImpl(), getOperation());\n"
193      << "}\n";
194 }
195 
emitInterfaceDefs(const RecordKeeper & recordKeeper,raw_ostream & os)196 static bool emitInterfaceDefs(const RecordKeeper &recordKeeper,
197                               raw_ostream &os) {
198   llvm::emitSourceFileHeader("Availability Interface Definitions", os);
199 
200   auto defs = recordKeeper.getAllDerivedDefinitions("Availability");
201   SmallVector<const Record *, 1> handledClasses;
202   for (const Record *def : defs) {
203     SmallVector<Record *, 1> parent;
204     def->getDirectSuperClasses(parent);
205     if (parent.size() != 1) {
206       PrintFatalError(def->getLoc(),
207                       "expected to only have one direct superclass");
208     }
209     if (llvm::is_contained(handledClasses, parent.front()))
210       continue;
211 
212     Availability availability(def);
213     emitInterfaceDef(availability, os);
214     handledClasses.push_back(parent.front());
215   }
216   return false;
217 }
218 
219 //===----------------------------------------------------------------------===//
220 // Availability Interface Declarations AutoGen
221 //===----------------------------------------------------------------------===//
222 
emitConceptDecl(const Availability & availability,raw_ostream & os)223 static void emitConceptDecl(const Availability &availability, raw_ostream &os) {
224   os << "  class Concept {\n"
225      << "  public:\n"
226      << "    virtual ~Concept() = default;\n"
227      << "    virtual " << availability.getQueryFnRetType() << " "
228      << availability.getQueryFnName()
229      << "(const Concept *impl, Operation *tblgen_opaque_op) const = 0;\n"
230      << "  };\n";
231 }
232 
emitModelDecl(const Availability & availability,raw_ostream & os)233 static void emitModelDecl(const Availability &availability, raw_ostream &os) {
234   for (const char *modelClass : {"Model", "FallbackModel"}) {
235     os << "  template<typename ConcreteOp>\n";
236     os << "  class " << modelClass << " : public Concept {\n"
237        << "  public:\n"
238        << "    " << availability.getQueryFnRetType() << " "
239        << availability.getQueryFnName()
240        << "(const Concept *impl, Operation *tblgen_opaque_op) const final {\n"
241        << "      auto op = llvm::cast<ConcreteOp>(tblgen_opaque_op);\n"
242        << "      (void)op;\n"
243        // Forward to the method on the concrete operation type.
244        << "      return op." << availability.getQueryFnName() << "();\n"
245        << "    }\n"
246        << "  };\n";
247   }
248   os << "  template<typename ConcreteModel, typename ConcreteOp>\n";
249   os << "  class ExternalModel : public FallbackModel<ConcreteOp> {};\n";
250 }
251 
emitInterfaceDecl(const Availability & availability,raw_ostream & os)252 static void emitInterfaceDecl(const Availability &availability,
253                               raw_ostream &os) {
254   StringRef interfaceName = availability.getInterfaceClassName();
255   std::string interfaceTraitsName =
256       std::string(formatv("{0}Traits", interfaceName));
257 
258   StringRef cppNamespace = availability.getInterfaceClassNamespace();
259   NamespaceEmitter nsEmitter(os, cppNamespace);
260 
261   // Emit the traits struct containing the concept and model declarations.
262   os << "namespace detail {\n"
263      << "struct " << interfaceTraitsName << " {\n";
264   emitConceptDecl(availability, os);
265   os << '\n';
266   emitModelDecl(availability, os);
267   os << "};\n} // namespace detail\n\n";
268 
269   // Emit the main interface class declaration.
270   os << "/*\n" << availability.getInterfaceDescription().trim() << "\n*/\n";
271   os << llvm::formatv("class {0} : public OpInterface<{1}, detail::{2}> {\n"
272                       "public:\n"
273                       "  using OpInterface<{1}, detail::{2}>::OpInterface;\n",
274                       interfaceName, interfaceName, interfaceTraitsName);
275 
276   // Emit query function declaration.
277   os << "  " << availability.getQueryFnRetType() << " "
278      << availability.getQueryFnName() << "();\n";
279   os << "};\n\n";
280 }
281 
emitInterfaceDecls(const RecordKeeper & recordKeeper,raw_ostream & os)282 static bool emitInterfaceDecls(const RecordKeeper &recordKeeper,
283                                raw_ostream &os) {
284   llvm::emitSourceFileHeader("Availability Interface Declarations", os);
285 
286   auto defs = recordKeeper.getAllDerivedDefinitions("Availability");
287   SmallVector<const Record *, 4> handledClasses;
288   for (const Record *def : defs) {
289     SmallVector<Record *, 1> parent;
290     def->getDirectSuperClasses(parent);
291     if (parent.size() != 1) {
292       PrintFatalError(def->getLoc(),
293                       "expected to only have one direct superclass");
294     }
295     if (llvm::is_contained(handledClasses, parent.front()))
296       continue;
297 
298     Availability avail(def);
299     emitInterfaceDecl(avail, os);
300     handledClasses.push_back(parent.front());
301   }
302   return false;
303 }
304 
305 //===----------------------------------------------------------------------===//
306 // Availability Interface Hook Registration
307 //===----------------------------------------------------------------------===//
308 
309 // Registers the operation interface generator to mlir-tblgen.
310 static mlir::GenRegistration
311     genInterfaceDecls("gen-avail-interface-decls",
312                       "Generate availability interface declarations",
__anon19f8bfbe0202(const RecordKeeper &records, raw_ostream &os) 313                       [](const RecordKeeper &records, raw_ostream &os) {
314                         return emitInterfaceDecls(records, os);
315                       });
316 
317 // Registers the operation interface generator to mlir-tblgen.
318 static mlir::GenRegistration
319     genInterfaceDefs("gen-avail-interface-defs",
320                      "Generate op interface definitions",
__anon19f8bfbe0302(const RecordKeeper &records, raw_ostream &os) 321                      [](const RecordKeeper &records, raw_ostream &os) {
322                        return emitInterfaceDefs(records, os);
323                      });
324 
325 //===----------------------------------------------------------------------===//
326 // Enum Availability Query AutoGen
327 //===----------------------------------------------------------------------===//
328 
emitAvailabilityQueryForIntEnum(const Record & enumDef,raw_ostream & os)329 static void emitAvailabilityQueryForIntEnum(const Record &enumDef,
330                                             raw_ostream &os) {
331   EnumAttr enumAttr(enumDef);
332   StringRef enumName = enumAttr.getEnumClassName();
333   std::vector<EnumAttrCase> enumerants = enumAttr.getAllCases();
334 
335   // Mapping from availability class name to (enumerant, availability
336   // specification) pairs.
337   llvm::StringMap<llvm::SmallVector<std::pair<EnumAttrCase, Availability>, 1>>
338       classCaseMap;
339 
340   // Place all availability specifications to their corresponding
341   // availability classes.
342   for (const EnumAttrCase &enumerant : enumerants)
343     for (const Availability &avail : getAvailabilities(enumerant.getDef()))
344       classCaseMap[avail.getClass()].push_back({enumerant, avail});
345 
346   for (const auto &classCasePair : classCaseMap) {
347     Availability avail = classCasePair.getValue().front().second;
348 
349     os << formatv("llvm::Optional<{0}> {1}({2} value) {{\n",
350                   avail.getMergeInstanceType(), avail.getQueryFnName(),
351                   enumName);
352 
353     os << "  switch (value) {\n";
354     for (const auto &caseSpecPair : classCasePair.getValue()) {
355       EnumAttrCase enumerant = caseSpecPair.first;
356       Availability avail = caseSpecPair.second;
357       os << formatv("  case {0}::{1}: { {2} return {3}({4}); }\n", enumName,
358                     enumerant.getSymbol(), avail.getMergeInstancePreparation(),
359                     avail.getMergeInstanceType(), avail.getMergeInstance());
360     }
361     // Only emit default if uncovered cases.
362     if (classCasePair.getValue().size() < enumAttr.getAllCases().size())
363       os << "  default: break;\n";
364     os << "  }\n"
365        << "  return llvm::None;\n"
366        << "}\n";
367   }
368 }
369 
emitAvailabilityQueryForBitEnum(const Record & enumDef,raw_ostream & os)370 static void emitAvailabilityQueryForBitEnum(const Record &enumDef,
371                                             raw_ostream &os) {
372   EnumAttr enumAttr(enumDef);
373   StringRef enumName = enumAttr.getEnumClassName();
374   std::string underlyingType = std::string(enumAttr.getUnderlyingType());
375   std::vector<EnumAttrCase> enumerants = enumAttr.getAllCases();
376 
377   // Mapping from availability class name to (enumerant, availability
378   // specification) pairs.
379   llvm::StringMap<llvm::SmallVector<std::pair<EnumAttrCase, Availability>, 1>>
380       classCaseMap;
381 
382   // Place all availability specifications to their corresponding
383   // availability classes.
384   for (const EnumAttrCase &enumerant : enumerants)
385     for (const Availability &avail : getAvailabilities(enumerant.getDef()))
386       classCaseMap[avail.getClass()].push_back({enumerant, avail});
387 
388   for (const auto &classCasePair : classCaseMap) {
389     Availability avail = classCasePair.getValue().front().second;
390 
391     os << formatv("llvm::Optional<{0}> {1}({2} value) {{\n",
392                   avail.getMergeInstanceType(), avail.getQueryFnName(),
393                   enumName);
394 
395     os << formatv(
396         "  assert(::llvm::countPopulation(static_cast<{0}>(value)) <= 1"
397         " && \"cannot have more than one bit set\");\n",
398         underlyingType);
399 
400     os << "  switch (value) {\n";
401     for (const auto &caseSpecPair : classCasePair.getValue()) {
402       EnumAttrCase enumerant = caseSpecPair.first;
403       Availability avail = caseSpecPair.second;
404       os << formatv("  case {0}::{1}: { {2} return {3}({4}); }\n", enumName,
405                     enumerant.getSymbol(), avail.getMergeInstancePreparation(),
406                     avail.getMergeInstanceType(), avail.getMergeInstance());
407     }
408     os << "  default: break;\n";
409     os << "  }\n"
410        << "  return llvm::None;\n"
411        << "}\n";
412   }
413 }
414 
emitEnumDecl(const Record & enumDef,raw_ostream & os)415 static void emitEnumDecl(const Record &enumDef, raw_ostream &os) {
416   EnumAttr enumAttr(enumDef);
417   StringRef enumName = enumAttr.getEnumClassName();
418   StringRef cppNamespace = enumAttr.getCppNamespace();
419   auto enumerants = enumAttr.getAllCases();
420 
421   llvm::SmallVector<StringRef, 2> namespaces;
422   llvm::SplitString(cppNamespace, namespaces, "::");
423 
424   for (auto ns : namespaces)
425     os << "namespace " << ns << " {\n";
426 
427   llvm::StringSet<> handledClasses;
428 
429   // Place all availability specifications to their corresponding
430   // availability classes.
431   for (const EnumAttrCase &enumerant : enumerants)
432     for (const Availability &avail : getAvailabilities(enumerant.getDef())) {
433       StringRef className = avail.getClass();
434       if (handledClasses.count(className))
435         continue;
436       os << formatv("llvm::Optional<{0}> {1}({2} value);\n",
437                     avail.getMergeInstanceType(), avail.getQueryFnName(),
438                     enumName);
439       handledClasses.insert(className);
440     }
441 
442   for (auto ns : llvm::reverse(namespaces))
443     os << "} // namespace " << ns << "\n";
444 }
445 
emitEnumDecls(const RecordKeeper & recordKeeper,raw_ostream & os)446 static bool emitEnumDecls(const RecordKeeper &recordKeeper, raw_ostream &os) {
447   llvm::emitSourceFileHeader("SPIR-V Enum Availability Declarations", os);
448 
449   auto defs = recordKeeper.getAllDerivedDefinitions("EnumAttrInfo");
450   for (const auto *def : defs)
451     emitEnumDecl(*def, os);
452 
453   return false;
454 }
455 
emitEnumDef(const Record & enumDef,raw_ostream & os)456 static void emitEnumDef(const Record &enumDef, raw_ostream &os) {
457   EnumAttr enumAttr(enumDef);
458   StringRef cppNamespace = enumAttr.getCppNamespace();
459 
460   llvm::SmallVector<StringRef, 2> namespaces;
461   llvm::SplitString(cppNamespace, namespaces, "::");
462 
463   for (auto ns : namespaces)
464     os << "namespace " << ns << " {\n";
465 
466   if (enumAttr.isBitEnum()) {
467     emitAvailabilityQueryForBitEnum(enumDef, os);
468   } else {
469     emitAvailabilityQueryForIntEnum(enumDef, os);
470   }
471 
472   for (auto ns : llvm::reverse(namespaces))
473     os << "} // namespace " << ns << "\n";
474   os << "\n";
475 }
476 
emitEnumDefs(const RecordKeeper & recordKeeper,raw_ostream & os)477 static bool emitEnumDefs(const RecordKeeper &recordKeeper, raw_ostream &os) {
478   llvm::emitSourceFileHeader("SPIR-V Enum Availability Definitions", os);
479 
480   auto defs = recordKeeper.getAllDerivedDefinitions("EnumAttrInfo");
481   for (const auto *def : defs)
482     emitEnumDef(*def, os);
483 
484   return false;
485 }
486 
487 //===----------------------------------------------------------------------===//
488 // Enum Availability Query Hook Registration
489 //===----------------------------------------------------------------------===//
490 
491 // Registers the enum utility generator to mlir-tblgen.
492 static mlir::GenRegistration
493     genEnumDecls("gen-spirv-enum-avail-decls",
494                  "Generate SPIR-V enum availability declarations",
__anon19f8bfbe0402(const RecordKeeper &records, raw_ostream &os) 495                  [](const RecordKeeper &records, raw_ostream &os) {
496                    return emitEnumDecls(records, os);
497                  });
498 
499 // Registers the enum utility generator to mlir-tblgen.
500 static mlir::GenRegistration
501     genEnumDefs("gen-spirv-enum-avail-defs",
502                 "Generate SPIR-V enum availability definitions",
__anon19f8bfbe0502(const RecordKeeper &records, raw_ostream &os) 503                 [](const RecordKeeper &records, raw_ostream &os) {
504                   return emitEnumDefs(records, os);
505                 });
506 
507 //===----------------------------------------------------------------------===//
508 // Serialization AutoGen
509 //===----------------------------------------------------------------------===//
510 
511 /// Generates code to serialize attributes of a SPV_Op `op` into `os`. The
512 /// generates code extracts the attribute with name `attrName` from
513 /// `operandList` of `op`.
emitAttributeSerialization(const Attribute & attr,ArrayRef<SMLoc> loc,StringRef tabs,StringRef opVar,StringRef operandList,StringRef attrName,raw_ostream & os)514 static void emitAttributeSerialization(const Attribute &attr,
515                                        ArrayRef<SMLoc> loc, StringRef tabs,
516                                        StringRef opVar, StringRef operandList,
517                                        StringRef attrName, raw_ostream &os) {
518   os << tabs
519      << formatv("if (auto attr = {0}->getAttr(\"{1}\")) {{\n", opVar, attrName);
520   if (attr.getAttrDefName() == "SPV_ScopeAttr" ||
521       attr.getAttrDefName() == "SPV_MemorySemanticsAttr") {
522     os << tabs
523        << formatv("  {0}.push_back(prepareConstantInt({1}.getLoc(), "
524                   "attr.cast<IntegerAttr>()));\n",
525                   operandList, opVar);
526   } else if (attr.getAttrDefName() == "I32ArrayAttr") {
527     // Serialize all the elements of the array
528     os << tabs << "  for (auto attrElem : attr.cast<ArrayAttr>()) {\n";
529     os << tabs
530        << formatv("    {0}.push_back(static_cast<uint32_t>("
531                   "attrElem.cast<IntegerAttr>().getValue().getZExtValue()));\n",
532                   operandList);
533     os << tabs << "  }\n";
534   } else if (attr.isEnumAttr() || attr.getAttrDefName() == "I32Attr") {
535     os << tabs
536        << formatv("  {0}.push_back(static_cast<uint32_t>("
537                   "attr.cast<IntegerAttr>().getValue().getZExtValue()));\n",
538                   operandList);
539   } else if (attr.isEnumAttr() || attr.getAttrDefName() == "TypeAttr") {
540     os << tabs
541        << formatv("  {0}.push_back(static_cast<uint32_t>("
542                   "getTypeID(attr.cast<TypeAttr>().getValue())));\n",
543                   operandList);
544   } else {
545     PrintFatalError(
546         loc,
547         llvm::Twine(
548             "unhandled attribute type in SPIR-V serialization generation : '") +
549             attr.getAttrDefName() + llvm::Twine("'"));
550   }
551   os << tabs << "}\n";
552 }
553 
554 /// Generates code to serialize the operands of a SPV_Op `op` into `os`. The
555 /// generated queries the SSA-ID if operand is a SSA-Value, or serializes the
556 /// attributes. The `operands` vector is updated appropriately. `elidedAttrs`
557 /// updated as well to include the serialized attributes.
emitArgumentSerialization(const Operator & op,ArrayRef<SMLoc> loc,StringRef tabs,StringRef opVar,StringRef operands,StringRef elidedAttrs,raw_ostream & os)558 static void emitArgumentSerialization(const Operator &op, ArrayRef<SMLoc> loc,
559                                       StringRef tabs, StringRef opVar,
560                                       StringRef operands, StringRef elidedAttrs,
561                                       raw_ostream &os) {
562   using mlir::tblgen::Argument;
563 
564   // SPIR-V ops can mix operands and attributes in the definition. These
565   // operands and attributes are serialized in the exact order of the definition
566   // to match SPIR-V binary format requirements. It can cause excessive
567   // generated code bloat because we are emitting code to handle each
568   // operand/attribute separately. So here we probe first to check whether all
569   // the operands are ahead of attributes. Then we can serialize all operands
570   // together.
571 
572   // Whether all operands are ahead of all attributes in the op's spec.
573   bool areOperandsAheadOfAttrs = true;
574   // Find the first attribute.
575   const Argument *it = llvm::find_if(op.getArgs(), [](const Argument &arg) {
576     return arg.is<NamedAttribute *>();
577   });
578   // Check whether all following arguments are attributes.
579   for (const Argument *ie = op.arg_end(); it != ie; ++it) {
580     if (!it->is<NamedAttribute *>()) {
581       areOperandsAheadOfAttrs = false;
582       break;
583     }
584   }
585 
586   // Serialize all operands together.
587   if (areOperandsAheadOfAttrs) {
588     if (op.getNumOperands() != 0) {
589       os << tabs
590          << formatv("for (Value operand : {0}->getOperands()) {{\n", opVar);
591       os << tabs << "  auto id = getValueID(operand);\n";
592       os << tabs << "  assert(id && \"use before def!\");\n";
593       os << tabs << formatv("  {0}.push_back(id);\n", operands);
594       os << tabs << "}\n";
595     }
596     for (const NamedAttribute &attr : op.getAttributes()) {
597       emitAttributeSerialization(
598           (attr.attr.isOptional() ? attr.attr.getBaseAttr() : attr.attr), loc,
599           tabs, opVar, operands, attr.name, os);
600       os << tabs
601          << formatv("{0}.push_back(\"{1}\");\n", elidedAttrs, attr.name);
602     }
603     return;
604   }
605 
606   // Serialize operands separately.
607   auto operandNum = 0;
608   for (unsigned i = 0, e = op.getNumArgs(); i < e; ++i) {
609     auto argument = op.getArg(i);
610     os << tabs << "{\n";
611     if (argument.is<NamedTypeConstraint *>()) {
612       os << tabs
613          << formatv("  for (auto arg : {0}.getODSOperands({1})) {{\n", opVar,
614                     operandNum);
615       os << tabs << "    auto argID = getValueID(arg);\n";
616       os << tabs << "    if (!argID) {\n";
617       os << tabs
618          << formatv("      return emitError({0}.getLoc(), "
619                     "\"operand #{1} has a use before def\");\n",
620                     opVar, operandNum);
621       os << tabs << "    }\n";
622       os << tabs << formatv("    {0}.push_back(argID);\n", operands);
623       os << "    }\n";
624       operandNum++;
625     } else {
626       NamedAttribute *attr = argument.get<NamedAttribute *>();
627       auto newtabs = tabs.str() + "  ";
628       emitAttributeSerialization(
629           (attr->attr.isOptional() ? attr->attr.getBaseAttr() : attr->attr),
630           loc, newtabs, opVar, operands, attr->name, os);
631       os << newtabs
632          << formatv("{0}.push_back(\"{1}\");\n", elidedAttrs, attr->name);
633     }
634     os << tabs << "}\n";
635   }
636 }
637 
638 /// Generates code to serializes the result of SPV_Op `op` into `os`. The
639 /// generated gets the ID for the type of the result (if any), the SSA-ID of
640 /// the result and updates `resultID` with the SSA-ID.
emitResultSerialization(const Operator & op,ArrayRef<SMLoc> loc,StringRef tabs,StringRef opVar,StringRef operands,StringRef resultID,raw_ostream & os)641 static void emitResultSerialization(const Operator &op, ArrayRef<SMLoc> loc,
642                                     StringRef tabs, StringRef opVar,
643                                     StringRef operands, StringRef resultID,
644                                     raw_ostream &os) {
645   if (op.getNumResults() == 1) {
646     StringRef resultTypeID("resultTypeID");
647     os << tabs << formatv("uint32_t {0} = 0;\n", resultTypeID);
648     os << tabs
649        << formatv(
650               "if (failed(processType({0}.getLoc(), {0}.getType(), {1}))) {{\n",
651               opVar, resultTypeID);
652     os << tabs << "  return failure();\n";
653     os << tabs << "}\n";
654     os << tabs << formatv("{0}.push_back({1});\n", operands, resultTypeID);
655     // Create an SSA result <id> for the op
656     os << tabs << formatv("{0} = getNextID();\n", resultID);
657     os << tabs
658        << formatv("valueIDMap[{0}.getResult()] = {1};\n", opVar, resultID);
659     os << tabs << formatv("{0}.push_back({1});\n", operands, resultID);
660   } else if (op.getNumResults() != 0) {
661     PrintFatalError(loc, "SPIR-V ops can only have zero or one result");
662   }
663 }
664 
665 /// Generates code to serialize attributes of SPV_Op `op` that become
666 /// decorations on the `resultID` of the serialized operation `opVar` in the
667 /// SPIR-V binary.
emitDecorationSerialization(const Operator & op,StringRef tabs,StringRef opVar,StringRef elidedAttrs,StringRef resultID,raw_ostream & os)668 static void emitDecorationSerialization(const Operator &op, StringRef tabs,
669                                         StringRef opVar, StringRef elidedAttrs,
670                                         StringRef resultID, raw_ostream &os) {
671   if (op.getNumResults() == 1) {
672     // All non-argument attributes translated into OpDecorate instruction
673     os << tabs << formatv("for (auto attr : {0}->getAttrs()) {{\n", opVar);
674     os << tabs
675        << formatv("  if (llvm::is_contained({0}, attr.getName())) {{",
676                   elidedAttrs);
677     os << tabs << "    continue;\n";
678     os << tabs << "  }\n";
679     os << tabs
680        << formatv(
681               "  if (failed(processDecoration({0}.getLoc(), {1}, attr))) {{\n",
682               opVar, resultID);
683     os << tabs << "    return failure();\n";
684     os << tabs << "  }\n";
685     os << tabs << "}\n";
686   }
687 }
688 
689 /// Generates code to serialize an SPV_Op `op` into `os`.
emitSerializationFunction(const Record * attrClass,const Record * record,const Operator & op,raw_ostream & os)690 static void emitSerializationFunction(const Record *attrClass,
691                                       const Record *record, const Operator &op,
692                                       raw_ostream &os) {
693   // If the record has 'autogenSerialization' set to 0, nothing to do
694   if (!record->getValueAsBit("autogenSerialization"))
695     return;
696 
697   StringRef opVar("op"), operands("operands"), elidedAttrs("elidedAttrs"),
698       resultID("resultID");
699 
700   os << formatv(
701       "template <> LogicalResult\nSerializer::processOp<{0}>({0} {1}) {{\n",
702       op.getQualCppClassName(), opVar);
703 
704   // Special case for ops without attributes in TableGen definitions
705   if (op.getNumAttributes() == 0 && op.getNumVariableLengthOperands() == 0) {
706     std::string extInstSet;
707     std::string opcode;
708     if (record->isSubClassOf("SPV_ExtInstOp")) {
709       extInstSet =
710           formatv("\"{0}\"", record->getValueAsString("extendedInstSetName"));
711       opcode = std::to_string(record->getValueAsInt("extendedInstOpcode"));
712     } else {
713       extInstSet = "\"\"";
714       opcode = formatv("static_cast<uint32_t>(spirv::Opcode::{0})",
715                        record->getValueAsString("spirvOpName"));
716     }
717 
718     os << formatv("  return processOpWithoutGrammarAttr({0}, {1}, {2});\n}\n\n",
719                   opVar, extInstSet, opcode);
720     return;
721   }
722 
723   os << formatv("  SmallVector<uint32_t, 4> {0};\n", operands);
724   os << formatv("  SmallVector<StringRef, 2> {0};\n", elidedAttrs);
725 
726   // Serialize result information.
727   if (op.getNumResults() == 1) {
728     os << formatv("  uint32_t {0} = 0;\n", resultID);
729     emitResultSerialization(op, record->getLoc(), "  ", opVar, operands,
730                             resultID, os);
731   }
732 
733   // Process arguments.
734   emitArgumentSerialization(op, record->getLoc(), "  ", opVar, operands,
735                             elidedAttrs, os);
736 
737   if (record->isSubClassOf("SPV_ExtInstOp")) {
738     os << formatv(
739         "  (void)encodeExtensionInstruction({0}, \"{1}\", {2}, {3});\n", opVar,
740         record->getValueAsString("extendedInstSetName"),
741         record->getValueAsInt("extendedInstOpcode"), operands);
742   } else {
743     // Emit debug info.
744     os << formatv("  (void)emitDebugLine(functionBody, {0}.getLoc());\n",
745                   opVar);
746     os << formatv("  (void)encodeInstructionInto("
747                   "functionBody, spirv::Opcode::{1}, {2});\n",
748                   op.getQualCppClassName(),
749                   record->getValueAsString("spirvOpName"), operands);
750   }
751 
752   // Process decorations.
753   emitDecorationSerialization(op, "  ", opVar, elidedAttrs, resultID, os);
754 
755   os << "  return success();\n";
756   os << "}\n\n";
757 }
758 
759 /// Generates the prologue for the function that dispatches the serialization of
760 /// the operation `opVar` based on its opcode.
initDispatchSerializationFn(StringRef opVar,raw_ostream & os)761 static void initDispatchSerializationFn(StringRef opVar, raw_ostream &os) {
762   os << formatv(
763       "LogicalResult Serializer::dispatchToAutogenSerialization(Operation "
764       "*{0}) {{\n",
765       opVar);
766 }
767 
768 /// Generates the body of the dispatch function. This function generates the
769 /// check that if satisfied, will call the serialization function generated for
770 /// the `op`.
emitSerializationDispatch(const Operator & op,StringRef tabs,StringRef opVar,raw_ostream & os)771 static void emitSerializationDispatch(const Operator &op, StringRef tabs,
772                                       StringRef opVar, raw_ostream &os) {
773   os << tabs
774      << formatv("if (isa<{0}>({1})) {{\n", op.getQualCppClassName(), opVar);
775   os << tabs
776      << formatv("  return processOp(cast<{0}>({1}));\n",
777                 op.getQualCppClassName(), opVar);
778   os << tabs << "}\n";
779 }
780 
781 /// Generates the epilogue for the function that dispatches the serialization of
782 /// the operation.
finalizeDispatchSerializationFn(StringRef opVar,raw_ostream & os)783 static void finalizeDispatchSerializationFn(StringRef opVar, raw_ostream &os) {
784   os << formatv(
785       "  return {0}->emitError(\"unhandled operation serialization\");\n",
786       opVar);
787   os << "}\n\n";
788 }
789 
790 /// Generates code to deserialize the attribute of a SPV_Op into `os`. The
791 /// generated code reads the `words` of the serialized instruction at
792 /// position `wordIndex` and adds the deserialized attribute into `attrList`.
emitAttributeDeserialization(const Attribute & attr,ArrayRef<SMLoc> loc,StringRef tabs,StringRef attrList,StringRef attrName,StringRef words,StringRef wordIndex,raw_ostream & os)793 static void emitAttributeDeserialization(const Attribute &attr,
794                                          ArrayRef<SMLoc> loc, StringRef tabs,
795                                          StringRef attrList, StringRef attrName,
796                                          StringRef words, StringRef wordIndex,
797                                          raw_ostream &os) {
798   if (attr.getAttrDefName() == "SPV_ScopeAttr" ||
799       attr.getAttrDefName() == "SPV_MemorySemanticsAttr") {
800     os << tabs
801        << formatv("{0}.push_back(opBuilder.getNamedAttr(\"{1}\", "
802                   "getConstantInt({2}[{3}++])));\n",
803                   attrList, attrName, words, wordIndex);
804   } else if (attr.getAttrDefName() == "I32ArrayAttr") {
805     os << tabs << "SmallVector<Attribute, 4> attrListElems;\n";
806     os << tabs << formatv("while ({0} < {1}.size()) {{\n", wordIndex, words);
807     os << tabs
808        << formatv(
809               "  "
810               "attrListElems.push_back(opBuilder.getI32IntegerAttr({0}[{1}++]))"
811               ";\n",
812               words, wordIndex);
813     os << tabs << "}\n";
814     os << tabs
815        << formatv("{0}.push_back(opBuilder.getNamedAttr(\"{1}\", "
816                   "opBuilder.getArrayAttr(attrListElems)));\n",
817                   attrList, attrName);
818   } else if (attr.isEnumAttr() || attr.getAttrDefName() == "I32Attr") {
819     os << tabs
820        << formatv("{0}.push_back(opBuilder.getNamedAttr(\"{1}\", "
821                   "opBuilder.getI32IntegerAttr({2}[{3}++])));\n",
822                   attrList, attrName, words, wordIndex);
823   } else if (attr.isEnumAttr() || attr.getAttrDefName() == "TypeAttr") {
824     os << tabs
825        << formatv("{0}.push_back(opBuilder.getNamedAttr(\"{1}\", "
826                   "TypeAttr::get(getType({2}[{3}++]))));\n",
827                   attrList, attrName, words, wordIndex);
828   } else {
829     PrintFatalError(
830         loc, llvm::Twine(
831                  "unhandled attribute type in deserialization generation : '") +
832                  attr.getAttrDefName() + llvm::Twine("'"));
833   }
834 }
835 
836 /// Generates the code to deserialize the result of an SPV_Op `op` into
837 /// `os`. The generated code gets the type of the result specified at
838 /// `words`[`wordIndex`], the SSA ID for the result at position `wordIndex` + 1
839 /// and updates the `resultType` and `valueID` with the parsed type and SSA ID,
840 /// respectively.
emitResultDeserialization(const Operator & op,ArrayRef<SMLoc> loc,StringRef tabs,StringRef words,StringRef wordIndex,StringRef resultTypes,StringRef valueID,raw_ostream & os)841 static void emitResultDeserialization(const Operator &op, ArrayRef<SMLoc> loc,
842                                       StringRef tabs, StringRef words,
843                                       StringRef wordIndex,
844                                       StringRef resultTypes, StringRef valueID,
845                                       raw_ostream &os) {
846   // Deserialize result information if it exists
847   if (op.getNumResults() == 1) {
848     os << tabs << "{\n";
849     os << tabs << formatv("  if ({0} >= {1}.size()) {{\n", wordIndex, words);
850     os << tabs
851        << formatv(
852               "    return emitError(unknownLoc, \"expected result type <id> "
853               "while deserializing {0}\");\n",
854               op.getQualCppClassName());
855     os << tabs << "  }\n";
856     os << tabs << formatv("  auto ty = getType({0}[{1}]);\n", words, wordIndex);
857     os << tabs << "  if (!ty) {\n";
858     os << tabs
859        << formatv(
860               "    return emitError(unknownLoc, \"unknown type result <id> : "
861               "\") << {0}[{1}];\n",
862               words, wordIndex);
863     os << tabs << "  }\n";
864     os << tabs << formatv("  {0}.push_back(ty);\n", resultTypes);
865     os << tabs << formatv("  {0}++;\n", wordIndex);
866     os << tabs << formatv("  if ({0} >= {1}.size()) {{\n", wordIndex, words);
867     os << tabs
868        << formatv(
869               "    return emitError(unknownLoc, \"expected result <id> while "
870               "deserializing {0}\");\n",
871               op.getQualCppClassName());
872     os << tabs << "  }\n";
873     os << tabs << "}\n";
874     os << tabs << formatv("{0} = {1}[{2}++];\n", valueID, words, wordIndex);
875   } else if (op.getNumResults() != 0) {
876     PrintFatalError(loc, "SPIR-V ops can have only zero or one result");
877   }
878 }
879 
880 /// Generates the code to deserialize the operands of an SPV_Op `op` into
881 /// `os`. The generated code reads the `words` of the binary instruction, from
882 /// position `wordIndex` to the end, and either gets the Value corresponding to
883 /// the ID encoded, or deserializes the attributes encoded. The parsed
884 /// operand(attribute) is added to the `operands` list or `attributes` list.
emitOperandDeserialization(const Operator & op,ArrayRef<SMLoc> loc,StringRef tabs,StringRef words,StringRef wordIndex,StringRef operands,StringRef attributes,raw_ostream & os)885 static void emitOperandDeserialization(const Operator &op, ArrayRef<SMLoc> loc,
886                                        StringRef tabs, StringRef words,
887                                        StringRef wordIndex, StringRef operands,
888                                        StringRef attributes, raw_ostream &os) {
889   // Process operands/attributes
890   for (unsigned i = 0, e = op.getNumArgs(); i < e; ++i) {
891     auto argument = op.getArg(i);
892     if (auto *valueArg = argument.dyn_cast<NamedTypeConstraint *>()) {
893       if (valueArg->isVariableLength()) {
894         if (i != e - 1) {
895           PrintFatalError(loc, "SPIR-V ops can have Variadic<..> or "
896                                "Optional<...> arguments only if "
897                                "it's the last argument");
898         }
899         os << tabs
900            << formatv("for (; {0} < {1}.size(); ++{0})", wordIndex, words);
901       } else {
902         os << tabs << formatv("if ({0} < {1}.size())", wordIndex, words);
903       }
904       os << " {\n";
905       os << tabs
906          << formatv("  auto arg = getValue({0}[{1}]);\n", words, wordIndex);
907       os << tabs << "  if (!arg) {\n";
908       os << tabs
909          << formatv(
910                 "    return emitError(unknownLoc, \"unknown result <id> : \") "
911                 "<< {0}[{1}];\n",
912                 words, wordIndex);
913       os << tabs << "  }\n";
914       os << tabs << formatv("  {0}.push_back(arg);\n", operands);
915       if (!valueArg->isVariableLength()) {
916         os << tabs << formatv("  {0}++;\n", wordIndex);
917       }
918       os << tabs << "}\n";
919     } else {
920       os << tabs << formatv("if ({0} < {1}.size()) {{\n", wordIndex, words);
921       auto *attr = argument.get<NamedAttribute *>();
922       auto newtabs = tabs.str() + "  ";
923       emitAttributeDeserialization(
924           (attr->attr.isOptional() ? attr->attr.getBaseAttr() : attr->attr),
925           loc, newtabs, attributes, attr->name, words, wordIndex, os);
926       os << "  }\n";
927     }
928   }
929 
930   os << tabs << formatv("if ({0} != {1}.size()) {{\n", wordIndex, words);
931   os << tabs
932      << formatv(
933             "  return emitError(unknownLoc, \"found more operands than "
934             "expected when deserializing {0}, only \") << {1} << \" of \" << "
935             "{2}.size() << \" processed\";\n",
936             op.getQualCppClassName(), wordIndex, words);
937   os << tabs << "}\n\n";
938 }
939 
940 /// Generates code to update the `attributes` vector with the attributes
941 /// obtained from parsing the decorations in the SPIR-V binary associated with
942 /// an <id> `valueID`
emitDecorationDeserialization(const Operator & op,StringRef tabs,StringRef valueID,StringRef attributes,raw_ostream & os)943 static void emitDecorationDeserialization(const Operator &op, StringRef tabs,
944                                           StringRef valueID,
945                                           StringRef attributes,
946                                           raw_ostream &os) {
947   // Import decorations parsed
948   if (op.getNumResults() == 1) {
949     os << tabs << formatv("if (decorations.count({0})) {{\n", valueID);
950     os << tabs
951        << formatv("  auto attrs = decorations[{0}].getAttrs();\n", valueID);
952     os << tabs
953        << formatv("  {0}.append(attrs.begin(), attrs.end());\n", attributes);
954     os << tabs << "}\n";
955   }
956 }
957 
958 /// Generates code to deserialize an SPV_Op `op` into `os`.
emitDeserializationFunction(const Record * attrClass,const Record * record,const Operator & op,raw_ostream & os)959 static void emitDeserializationFunction(const Record *attrClass,
960                                         const Record *record,
961                                         const Operator &op, raw_ostream &os) {
962   // If the record has 'autogenSerialization' set to 0, nothing to do
963   if (!record->getValueAsBit("autogenSerialization"))
964     return;
965 
966   StringRef resultTypes("resultTypes"), valueID("valueID"), words("words"),
967       wordIndex("wordIndex"), opVar("op"), operands("operands"),
968       attributes("attributes");
969 
970   // Method declaration
971   os << formatv("template <> "
972                 "LogicalResult\nDeserializer::processOp<{0}>(ArrayRef<"
973                 "uint32_t> {1}) {{\n",
974                 op.getQualCppClassName(), words);
975 
976   // Special case for ops without attributes in TableGen definitions
977   if (op.getNumAttributes() == 0 && op.getNumVariableLengthOperands() == 0) {
978     os << formatv("  return processOpWithoutGrammarAttr("
979                   "{0}, \"{1}\", {2}, {3});\n}\n\n",
980                   words, op.getOperationName(),
981                   op.getNumResults() ? "true" : "false", op.getNumOperands());
982     return;
983   }
984 
985   os << formatv("  SmallVector<Type, 1> {0};\n", resultTypes);
986   os << formatv("  size_t {0} = 0; (void){0};\n", wordIndex);
987   os << formatv("  uint32_t {0} = 0; (void){0};\n", valueID);
988 
989   // Deserialize result information
990   emitResultDeserialization(op, record->getLoc(), "  ", words, wordIndex,
991                             resultTypes, valueID, os);
992 
993   os << formatv("  SmallVector<Value, 4> {0};\n", operands);
994   os << formatv("  SmallVector<NamedAttribute, 4> {0};\n", attributes);
995   // Operand deserialization
996   emitOperandDeserialization(op, record->getLoc(), "  ", words, wordIndex,
997                              operands, attributes, os);
998 
999   // Decorations
1000   emitDecorationDeserialization(op, "  ", valueID, attributes, os);
1001 
1002   os << formatv("  Location loc = createFileLineColLoc(opBuilder);\n");
1003   os << formatv("  auto {1} = opBuilder.create<{0}>(loc, {2}, {3}, {4}); "
1004                 "(void){1};\n",
1005                 op.getQualCppClassName(), opVar, resultTypes, operands,
1006                 attributes);
1007   if (op.getNumResults() == 1) {
1008     os << formatv("  valueMap[{0}] = {1}.getResult();\n\n", valueID, opVar);
1009   }
1010 
1011   // According to SPIR-V spec:
1012   // This location information applies to the instructions physically following
1013   // this instruction, up to the first occurrence of any of the following: the
1014   // next end of block.
1015   os << formatv("  if ({0}.hasTrait<OpTrait::IsTerminator>())\n", opVar);
1016   os << formatv("    (void)clearDebugLine();\n");
1017   os << "  return success();\n";
1018   os << "}\n\n";
1019 }
1020 
1021 /// Generates the prologue for the function that dispatches the deserialization
1022 /// based on the `opcode`.
initDispatchDeserializationFn(StringRef opcode,StringRef words,raw_ostream & os)1023 static void initDispatchDeserializationFn(StringRef opcode, StringRef words,
1024                                           raw_ostream &os) {
1025   os << formatv("LogicalResult spirv::Deserializer::"
1026                 "dispatchToAutogenDeserialization(spirv::Opcode {0},"
1027                 " ArrayRef<uint32_t> {1}) {{\n",
1028                 opcode, words);
1029   os << formatv("  switch ({0}) {{\n", opcode);
1030 }
1031 
1032 /// Generates the body of the dispatch function, by generating the case label
1033 /// for an opcode and the call to the method to perform the deserialization.
emitDeserializationDispatch(const Operator & op,const Record * def,StringRef tabs,StringRef words,raw_ostream & os)1034 static void emitDeserializationDispatch(const Operator &op, const Record *def,
1035                                         StringRef tabs, StringRef words,
1036                                         raw_ostream &os) {
1037   os << tabs
1038      << formatv("case spirv::Opcode::{0}:\n",
1039                 def->getValueAsString("spirvOpName"));
1040   os << tabs
1041      << formatv("  return processOp<{0}>({1});\n", op.getQualCppClassName(),
1042                 words);
1043 }
1044 
1045 /// Generates the epilogue for the function that dispatches the deserialization
1046 /// of the operation.
finalizeDispatchDeserializationFn(StringRef opcode,raw_ostream & os)1047 static void finalizeDispatchDeserializationFn(StringRef opcode,
1048                                               raw_ostream &os) {
1049   os << "  default:\n";
1050   os << "    ;\n";
1051   os << "  }\n";
1052   StringRef opcodeVar("opcodeString");
1053   os << formatv("  auto {0} = spirv::stringifyOpcode({1});\n", opcodeVar,
1054                 opcode);
1055   os << formatv("  if (!{0}.empty()) {{\n", opcodeVar);
1056   os << formatv("    return emitError(unknownLoc, \"unhandled deserialization "
1057                 "of \") << {0};\n",
1058                 opcodeVar);
1059   os << "  } else {\n";
1060   os << formatv("   return emitError(unknownLoc, \"unhandled opcode \") << "
1061                 "static_cast<uint32_t>({0});\n",
1062                 opcode);
1063   os << "  }\n";
1064   os << "}\n";
1065 }
1066 
initExtendedSetDeserializationDispatch(StringRef extensionSetName,StringRef instructionID,StringRef words,raw_ostream & os)1067 static void initExtendedSetDeserializationDispatch(StringRef extensionSetName,
1068                                                    StringRef instructionID,
1069                                                    StringRef words,
1070                                                    raw_ostream &os) {
1071   os << formatv("LogicalResult spirv::Deserializer::"
1072                 "dispatchToExtensionSetAutogenDeserialization("
1073                 "StringRef {0}, uint32_t {1}, ArrayRef<uint32_t> {2}) {{\n",
1074                 extensionSetName, instructionID, words);
1075 }
1076 
1077 static void
emitExtendedSetDeserializationDispatch(const RecordKeeper & recordKeeper,raw_ostream & os)1078 emitExtendedSetDeserializationDispatch(const RecordKeeper &recordKeeper,
1079                                        raw_ostream &os) {
1080   StringRef extensionSetName("extensionSetName"),
1081       instructionID("instructionID"), words("words");
1082 
1083   // First iterate over all ops derived from SPV_ExtensionSetOps to get all
1084   // extensionSets.
1085 
1086   // For each of the extensions a separate raw_string_ostream is used to
1087   // generate code into. These are then concatenated at the end. Since
1088   // raw_string_ostream needs a string&, use a vector to store all the string
1089   // that are captured by reference within raw_string_ostream.
1090   StringMap<raw_string_ostream> extensionSets;
1091   std::list<std::string> extensionSetNames;
1092 
1093   initExtendedSetDeserializationDispatch(extensionSetName, instructionID, words,
1094                                          os);
1095   auto defs = recordKeeper.getAllDerivedDefinitions("SPV_ExtInstOp");
1096   for (const auto *def : defs) {
1097     if (!def->getValueAsBit("autogenSerialization")) {
1098       continue;
1099     }
1100     Operator op(def);
1101     auto setName = def->getValueAsString("extendedInstSetName");
1102     if (!extensionSets.count(setName)) {
1103       extensionSetNames.emplace_back("");
1104       extensionSets.try_emplace(setName, extensionSetNames.back());
1105       auto &setos = extensionSets.find(setName)->second;
1106       setos << formatv("  if ({0} == \"{1}\") {{\n", extensionSetName, setName);
1107       setos << formatv("    switch ({0}) {{\n", instructionID);
1108     }
1109     auto &setos = extensionSets.find(setName)->second;
1110     setos << formatv("    case {0}:\n",
1111                      def->getValueAsInt("extendedInstOpcode"));
1112     setos << formatv("      return processOp<{0}>({1});\n",
1113                      op.getQualCppClassName(), words);
1114   }
1115 
1116   // Append the dispatch code for all the extended sets.
1117   for (auto &extensionSet : extensionSets) {
1118     os << extensionSet.second.str();
1119     os << "    default:\n";
1120     os << formatv(
1121         "      return emitError(unknownLoc, \"unhandled deserializations of "
1122         "\") << {0} << \" from extension set \" << {1};\n",
1123         instructionID, extensionSetName);
1124     os << "    }\n";
1125     os << "  }\n";
1126   }
1127 
1128   os << formatv("  return emitError(unknownLoc, \"unhandled deserialization of "
1129                 "extended instruction set {0}\");\n",
1130                 extensionSetName);
1131   os << "}\n";
1132 }
1133 
1134 /// Emits all the autogenerated serialization/deserializations functions for the
1135 /// SPV_Ops.
emitSerializationFns(const RecordKeeper & recordKeeper,raw_ostream & os)1136 static bool emitSerializationFns(const RecordKeeper &recordKeeper,
1137                                  raw_ostream &os) {
1138   llvm::emitSourceFileHeader("SPIR-V Serialization Utilities/Functions", os);
1139 
1140   std::string dSerFnString, dDesFnString, serFnString, deserFnString,
1141       utilsString;
1142   raw_string_ostream dSerFn(dSerFnString), dDesFn(dDesFnString),
1143       serFn(serFnString), deserFn(deserFnString);
1144   Record *attrClass = recordKeeper.getClass("Attr");
1145 
1146   // Emit the serialization and deserialization functions simultaneously.
1147   StringRef opVar("op");
1148   StringRef opcode("opcode"), words("words");
1149 
1150   // Handle the SPIR-V ops.
1151   initDispatchSerializationFn(opVar, dSerFn);
1152   initDispatchDeserializationFn(opcode, words, dDesFn);
1153   auto defs = recordKeeper.getAllDerivedDefinitions("SPV_Op");
1154   for (const auto *def : defs) {
1155     Operator op(def);
1156     emitSerializationFunction(attrClass, def, op, serFn);
1157     emitDeserializationFunction(attrClass, def, op, deserFn);
1158     if (def->getValueAsBit("hasOpcode") || def->isSubClassOf("SPV_ExtInstOp")) {
1159       emitSerializationDispatch(op, "  ", opVar, dSerFn);
1160     }
1161     if (def->getValueAsBit("hasOpcode")) {
1162       emitDeserializationDispatch(op, def, "  ", words, dDesFn);
1163     }
1164   }
1165   finalizeDispatchSerializationFn(opVar, dSerFn);
1166   finalizeDispatchDeserializationFn(opcode, dDesFn);
1167 
1168   emitExtendedSetDeserializationDispatch(recordKeeper, dDesFn);
1169 
1170   os << "#ifdef GET_SERIALIZATION_FNS\n\n";
1171   os << serFn.str();
1172   os << dSerFn.str();
1173   os << "#endif // GET_SERIALIZATION_FNS\n\n";
1174 
1175   os << "#ifdef GET_DESERIALIZATION_FNS\n\n";
1176   os << deserFn.str();
1177   os << dDesFn.str();
1178   os << "#endif // GET_DESERIALIZATION_FNS\n\n";
1179 
1180   return false;
1181 }
1182 
1183 //===----------------------------------------------------------------------===//
1184 // Serialization Hook Registration
1185 //===----------------------------------------------------------------------===//
1186 
1187 static mlir::GenRegistration genSerialization(
1188     "gen-spirv-serialization",
1189     "Generate SPIR-V (de)serialization utilities and functions",
__anon19f8bfbe0702(const RecordKeeper &records, raw_ostream &os) 1190     [](const RecordKeeper &records, raw_ostream &os) {
1191       return emitSerializationFns(records, os);
1192     });
1193 
1194 //===----------------------------------------------------------------------===//
1195 // Op Utils AutoGen
1196 //===----------------------------------------------------------------------===//
1197 
emitEnumGetAttrNameFnDecl(raw_ostream & os)1198 static void emitEnumGetAttrNameFnDecl(raw_ostream &os) {
1199   os << formatv("template <typename EnumClass> inline constexpr StringRef "
1200                 "attributeName();\n");
1201 }
1202 
emitEnumGetAttrNameFnDefn(const EnumAttr & enumAttr,raw_ostream & os)1203 static void emitEnumGetAttrNameFnDefn(const EnumAttr &enumAttr,
1204                                       raw_ostream &os) {
1205   auto enumName = enumAttr.getEnumClassName();
1206   os << formatv("template <> inline StringRef attributeName<{0}>() {{\n",
1207                 enumName);
1208   os << "  "
1209      << formatv("static constexpr const char attrName[] = \"{0}\";\n",
1210                 llvm::convertToSnakeFromCamelCase(enumName));
1211   os << "  return attrName;\n";
1212   os << "}\n";
1213 }
1214 
emitAttrUtils(const RecordKeeper & recordKeeper,raw_ostream & os)1215 static bool emitAttrUtils(const RecordKeeper &recordKeeper, raw_ostream &os) {
1216   llvm::emitSourceFileHeader("SPIR-V Attribute Utilities", os);
1217 
1218   auto defs = recordKeeper.getAllDerivedDefinitions("EnumAttrInfo");
1219   os << "#ifndef MLIR_DIALECT_SPIRV_IR_ATTR_UTILS_H_\n";
1220   os << "#define MLIR_DIALECT_SPIRV_IR_ATTR_UTILS_H_\n";
1221   emitEnumGetAttrNameFnDecl(os);
1222   for (const auto *def : defs) {
1223     EnumAttr enumAttr(*def);
1224     emitEnumGetAttrNameFnDefn(enumAttr, os);
1225   }
1226   os << "#endif // MLIR_DIALECT_SPIRV_IR_ATTR_UTILS_H\n";
1227   return false;
1228 }
1229 
1230 //===----------------------------------------------------------------------===//
1231 // Op Utils Hook Registration
1232 //===----------------------------------------------------------------------===//
1233 
1234 static mlir::GenRegistration
1235     genOpUtils("gen-spirv-attr-utils",
1236                "Generate SPIR-V attribute utility definitions",
__anon19f8bfbe0802(const RecordKeeper &records, raw_ostream &os) 1237                [](const RecordKeeper &records, raw_ostream &os) {
1238                  return emitAttrUtils(records, os);
1239                });
1240 
1241 //===----------------------------------------------------------------------===//
1242 // SPIR-V Availability Impl AutoGen
1243 //===----------------------------------------------------------------------===//
1244 
emitAvailabilityImpl(const Operator & srcOp,raw_ostream & os)1245 static void emitAvailabilityImpl(const Operator &srcOp, raw_ostream &os) {
1246   mlir::tblgen::FmtContext fctx;
1247   fctx.addSubst("overall", "tblgen_overall");
1248 
1249   std::vector<Availability> opAvailabilities =
1250       getAvailabilities(srcOp.getDef());
1251 
1252   // First collect all availability classes this op should implement.
1253   // All availability instances keep information for the generated interface and
1254   // the instance's specific requirement. Here we remember a random instance so
1255   // we can get the information regarding the generated interface.
1256   llvm::StringMap<Availability> availClasses;
1257   for (const Availability &avail : opAvailabilities)
1258     availClasses.try_emplace(avail.getClass(), avail);
1259   for (const NamedAttribute &namedAttr : srcOp.getAttributes()) {
1260     const auto *enumAttr = llvm::dyn_cast<EnumAttr>(&namedAttr.attr);
1261     if (!enumAttr)
1262       continue;
1263 
1264     for (const EnumAttrCase &enumerant : enumAttr->getAllCases())
1265       for (const Availability &caseAvail :
1266            getAvailabilities(enumerant.getDef()))
1267         availClasses.try_emplace(caseAvail.getClass(), caseAvail);
1268   }
1269 
1270   // Then generate implementation for each availability class.
1271   for (const auto &availClass : availClasses) {
1272     StringRef availClassName = availClass.getKey();
1273     Availability avail = availClass.getValue();
1274 
1275     // Generate the implementation method signature.
1276     os << formatv("{0} {1}::{2}() {{\n", avail.getQueryFnRetType(),
1277                   srcOp.getCppClassName(), avail.getQueryFnName());
1278 
1279     // Create the variable for the final requirement and initialize it.
1280     os << formatv("  {0} tblgen_overall = {1};\n", avail.getQueryFnRetType(),
1281                   avail.getMergeInitializer());
1282 
1283     // Update with the op's specific availability spec.
1284     for (const Availability &avail : opAvailabilities)
1285       if (avail.getClass() == availClassName &&
1286           (!avail.getMergeInstancePreparation().empty() ||
1287            !avail.getMergeActionCode().empty())) {
1288         os << "  {\n    "
1289            // Prepare this instance.
1290            << avail.getMergeInstancePreparation()
1291            << "\n    "
1292            // Merge this instance.
1293            << std::string(
1294                   tgfmt(avail.getMergeActionCode(),
1295                         &fctx.addSubst("instance", avail.getMergeInstance())))
1296            << ";\n  }\n";
1297       }
1298 
1299     // Update with enum attributes' specific availability spec.
1300     for (const NamedAttribute &namedAttr : srcOp.getAttributes()) {
1301       const auto *enumAttr = llvm::dyn_cast<EnumAttr>(&namedAttr.attr);
1302       if (!enumAttr)
1303         continue;
1304 
1305       // (enumerant, availability specification) pairs for this availability
1306       // class.
1307       SmallVector<std::pair<EnumAttrCase, Availability>, 1> caseSpecs;
1308 
1309       // Collect all cases' availability specs.
1310       for (const EnumAttrCase &enumerant : enumAttr->getAllCases())
1311         for (const Availability &caseAvail :
1312              getAvailabilities(enumerant.getDef()))
1313           if (availClassName == caseAvail.getClass())
1314             caseSpecs.push_back({enumerant, caseAvail});
1315 
1316       // If this attribute kind does not have any availability spec from any of
1317       // its cases, no more work to do.
1318       if (caseSpecs.empty())
1319         continue;
1320 
1321       if (enumAttr->isBitEnum()) {
1322         // For BitEnumAttr, we need to iterate over each bit to query its
1323         // availability spec.
1324         os << formatv("  for (unsigned i = 0; "
1325                       "i < std::numeric_limits<{0}>::digits; ++i) {{\n",
1326                       enumAttr->getUnderlyingType());
1327         os << formatv("    {0}::{1} tblgen_attrVal = this->{2}() & "
1328                       "static_cast<{0}::{1}>(1 << i);\n",
1329                       enumAttr->getCppNamespace(), enumAttr->getEnumClassName(),
1330                       namedAttr.name);
1331         os << formatv(
1332             "    if (static_cast<{0}>(tblgen_attrVal) == 0) continue;\n",
1333             enumAttr->getUnderlyingType());
1334       } else {
1335         // For IntEnumAttr, we just need to query the value as a whole.
1336         os << "  {\n";
1337         os << formatv("    auto tblgen_attrVal = this->{0}();\n",
1338                       namedAttr.name);
1339       }
1340       os << formatv("    auto tblgen_instance = {0}::{1}(tblgen_attrVal);\n",
1341                     enumAttr->getCppNamespace(), avail.getQueryFnName());
1342       os << "    if (tblgen_instance) "
1343          // TODO` here once ODS supports
1344          // dialect-specific contents so that we can use not implementing the
1345          // availability interface as indication of no requirements.
1346          << std::string(tgfmt(caseSpecs.front().second.getMergeActionCode(),
1347                               &fctx.addSubst("instance", "*tblgen_instance")))
1348          << ";\n";
1349       os << "  }\n";
1350     }
1351 
1352     os << "  return tblgen_overall;\n";
1353     os << "}\n";
1354   }
1355 }
1356 
emitAvailabilityImpl(const RecordKeeper & recordKeeper,raw_ostream & os)1357 static bool emitAvailabilityImpl(const RecordKeeper &recordKeeper,
1358                                  raw_ostream &os) {
1359   llvm::emitSourceFileHeader("SPIR-V Op Availability Implementations", os);
1360 
1361   auto defs = recordKeeper.getAllDerivedDefinitions("SPV_Op");
1362   for (const auto *def : defs) {
1363     Operator op(def);
1364     emitAvailabilityImpl(op, os);
1365   }
1366   return false;
1367 }
1368 
1369 //===----------------------------------------------------------------------===//
1370 // Op Availability Implementation Hook Registration
1371 //===----------------------------------------------------------------------===//
1372 
1373 static mlir::GenRegistration
1374     genOpAvailabilityImpl("gen-spirv-avail-impls",
1375                           "Generate SPIR-V operation utility definitions",
__anon19f8bfbe0902(const RecordKeeper &records, raw_ostream &os) 1376                           [](const RecordKeeper &records, raw_ostream &os) {
1377                             return emitAvailabilityImpl(records, os);
1378                           });
1379 
1380 //===----------------------------------------------------------------------===//
1381 // SPIR-V Capability Implication AutoGen
1382 //===----------------------------------------------------------------------===//
1383 
emitCapabilityImplication(const RecordKeeper & recordKeeper,raw_ostream & os)1384 static bool emitCapabilityImplication(const RecordKeeper &recordKeeper,
1385                                       raw_ostream &os) {
1386   llvm::emitSourceFileHeader("SPIR-V Capability Implication", os);
1387 
1388   EnumAttr enumAttr(recordKeeper.getDef("SPV_CapabilityAttr"));
1389 
1390   os << "ArrayRef<spirv::Capability> "
1391         "spirv::getDirectImpliedCapabilities(spirv::Capability cap) {\n"
1392      << "  switch (cap) {\n"
1393      << "  default: return {};\n";
1394   for (const EnumAttrCase &enumerant : enumAttr.getAllCases()) {
1395     const Record &def = enumerant.getDef();
1396     if (!def.getValue("implies"))
1397       continue;
1398 
1399     std::vector<Record *> impliedCapsDefs = def.getValueAsListOfDefs("implies");
1400     os << "  case spirv::Capability::" << enumerant.getSymbol()
1401        << ": {static const spirv::Capability implies[" << impliedCapsDefs.size()
1402        << "] = {";
1403     llvm::interleaveComma(impliedCapsDefs, os, [&](const Record *capDef) {
1404       os << "spirv::Capability::" << EnumAttrCase(capDef).getSymbol();
1405     });
1406     os << "}; return ArrayRef<spirv::Capability>(implies, "
1407        << impliedCapsDefs.size() << "); }\n";
1408   }
1409   os << "  }\n";
1410   os << "}\n";
1411 
1412   return false;
1413 }
1414 
1415 //===----------------------------------------------------------------------===//
1416 // SPIR-V Capability Implication Hook Registration
1417 //===----------------------------------------------------------------------===//
1418 
1419 static mlir::GenRegistration
1420     genCapabilityImplication("gen-spirv-capability-implication",
1421                              "Generate utility function to return implied "
1422                              "capabilities for a given capability",
__anon19f8bfbe0b02(const RecordKeeper &records, raw_ostream &os) 1423                              [](const RecordKeeper &records, raw_ostream &os) {
1424                                return emitCapabilityImplication(records, os);
1425                              });
1426