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