MLIR 24.0.0git
IRDLToCpp.cpp
Go to the documentation of this file.
1//===- IRDLToCpp.cpp - Converts IRDL definitions to C++ -------------------===//
2//
3// Part of the LLVM Project, under the A0ache 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
11#include "mlir/Support/LLVM.h"
12#include "llvm/ADT/STLExtras.h"
13#include "llvm/ADT/SmallString.h"
14#include "llvm/ADT/SmallVector.h"
15#include "llvm/ADT/StringExtras.h"
16#include "llvm/ADT/StringRef.h"
17#include "llvm/ADT/TypeSwitch.h"
18#include "llvm/Support/FormatVariadic.h"
19#include "llvm/Support/raw_ostream.h"
20
21#include "TemplatingUtils.h"
22
23using namespace mlir;
24
25constexpr char headerTemplateText[] =
26#include "Templates/Header.txt"
27 ;
28
29constexpr char declarationMacroFlag[] = "GEN_DIALECT_DECL_HEADER";
30constexpr char definitionMacroFlag[] = "GEN_DIALECT_DEF";
31
32namespace {
33
34/// The set of strings that can be generated from a Dialect declaraiton
35struct DialectStrings {
36 std::string dialectName;
37 std::string dialectCppName;
38 std::string dialectCppShortName;
39 std::string dialectBaseTypeName;
40
41 std::string namespaceOpen;
42 std::string namespaceClose;
43 std::string namespacePath;
44};
45
46/// The set of strings that can be generated from a Type declaraiton
47struct TypeStrings {
48 StringRef typeName;
49 std::string typeCppName;
50};
51
52/// The set of strings that can be generated from an Operation declaraiton
53struct OpStrings {
54 StringRef opName;
55 std::string opCppName;
56 std::string opScopedCppName;
57 SmallVector<std::string> opNameSpaces;
58 SmallVector<std::string> opResultNames;
59 SmallVector<std::string> opOperandNames;
60 SmallVector<std::string> opRegionNames;
61};
62
63static std::string joinNameList(llvm::ArrayRef<std::string> names) {
64 std::string nameArray;
65 llvm::raw_string_ostream nameArrayStream(nameArray);
66 nameArrayStream << "{\"" << llvm::join(names, "\", \"") << "\"}";
67
68 return nameArray;
69}
70
71/// Prefix identifiers that start with a digit to make them valid C++ names.
72static std::string legalizeCppName(StringRef name) {
73 if (!name.empty() && llvm::isDigit(name.front()))
74 return ("_" + name).str();
75 return name.str();
76}
77
78/// Generates the C++ type name for a TypeOp
79static std::string typeToCppName(irdl::TypeOp type) {
80 return llvm::formatv("{0}Type", legalizeCppName(convertToCamelFromSnakeCase(
81 type.getSymName(), true)));
82}
83
84/// Generates the C++ class name for an OperationOp
85static std::string opToCppName(irdl::OperationOp op) {
86 const auto opName = op.getSymName();
87 const auto periodIndex = opName.find_last_of(".");
88 const auto nameSubstr = periodIndex == std::string::npos
89 ? opName
90 : opName.substr(periodIndex + 1);
91 return llvm::formatv(
92 "{0}Op", legalizeCppName(convertToCamelFromSnakeCase(nameSubstr, true)));
93}
94
95/// Generates the C++ namespace components for an OperationOp.
96static SmallVector<std::string> opToCppNamespaces(irdl::OperationOp op) {
97 auto parts = SmallVector<StringRef>(llvm::split(op.getSymName(), "."));
98 parts.pop_back();
99 return llvm::map_to_vector(parts, legalizeCppName);
100}
101
102// Generates the C++ class name for an OperationOp, scoped to the namespace
103static std::string opToScopedCppName(irdl::OperationOp op) {
104 auto names = opToCppNamespaces(op);
105 names.push_back(opToCppName(op));
106 return llvm::join(names, "::");
107}
108
109/// Generates TypeStrings from a TypeOp
110static TypeStrings getStrings(irdl::TypeOp type) {
111 TypeStrings strings;
112 strings.typeName = type.getSymName();
113 strings.typeCppName = typeToCppName(type);
114 return strings;
115}
116
117/// Generates OpStrings from an OperatioOp
118static OpStrings getStrings(irdl::OperationOp op) {
119 auto operandOp = op.getOp<irdl::OperandsOp>();
120 auto resultOp = op.getOp<irdl::ResultsOp>();
121 auto regionsOp = op.getOp<irdl::RegionsOp>();
122
123 OpStrings strings;
124 strings.opName = op.getSymName();
125 strings.opNameSpaces = opToCppNamespaces(op);
126 strings.opCppName = opToCppName(op);
127 strings.opScopedCppName = opToScopedCppName(op);
128
129 if (operandOp) {
130 strings.opOperandNames = SmallVector<std::string>(
131 llvm::map_range(operandOp->getNames(), [](Attribute attr) {
132 return llvm::formatv("{0}", cast<StringAttr>(attr));
133 }));
134 }
135
136 if (resultOp) {
137 strings.opResultNames = SmallVector<std::string>(
138 llvm::map_range(resultOp->getNames(), [](Attribute attr) {
139 return llvm::formatv("{0}", cast<StringAttr>(attr));
140 }));
141 }
142
143 if (regionsOp) {
144 strings.opRegionNames = SmallVector<std::string>(
145 llvm::map_range(regionsOp->getNames(), [](Attribute attr) {
146 return llvm::formatv("{0}", cast<StringAttr>(attr));
147 }));
148 }
149
150 return strings;
151}
152
153/// Fills a dictionary with values from TypeStrings
154static void fillDict(irdl::detail::dictionary &dict,
155 const TypeStrings &strings) {
156 dict["TYPE_NAME"] = strings.typeName;
157 dict["TYPE_CPP_NAME"] = strings.typeCppName;
158}
159
160/// Fills a dictionary with values from OpStrings
161static void fillDict(irdl::detail::dictionary &dict, const OpStrings &strings) {
162 const auto operandCount = strings.opOperandNames.size();
163 const auto resultCount = strings.opResultNames.size();
164 const auto regionCount = strings.opRegionNames.size();
165
166 dict["OP_NAME"] = strings.opName;
167 dict["OP_CPP_NAME"] = strings.opCppName;
168 dict["OP_SCOPED_CPP_NAME"] = strings.opScopedCppName;
169 dict["OP_OPERAND_COUNT"] = std::to_string(strings.opOperandNames.size());
170 dict["OP_RESULT_COUNT"] = std::to_string(strings.opResultNames.size());
171 dict["OP_OPERAND_INITIALIZER_LIST"] =
172 operandCount ? joinNameList(strings.opOperandNames) : "{\"\"}";
173 dict["OP_RESULT_INITIALIZER_LIST"] =
174 resultCount ? joinNameList(strings.opResultNames) : "{\"\"}";
175 dict["OP_REGION_COUNT"] = std::to_string(regionCount);
176 dict["NAMESPACE_OPEN"] =
177 (dict["NAMESPACE_OPEN"] +
178 llvm::join(llvm::map_range(strings.opNameSpaces,
179 [](llvm::StringRef ref) -> std::string {
180 return llvm::formatv("namespace {0} {{",
181 ref);
182 }),
183 "\n"))
184 .str();
185 dict["NAMESPACE_PATH"] =
186 (dict["NAMESPACE_PATH"] +
187 llvm::join(llvm::map_range(strings.opNameSpaces,
188 [](llvm::StringRef ref) -> std::string {
189 return llvm::formatv("::{0}", ref);
190 }),
191 ""))
192 .str();
193 dict["NAMESPACE_CLOSE"] =
194 (llvm::join(llvm::map_range(llvm::reverse(strings.opNameSpaces),
195 [](llvm::StringRef ref) -> std::string {
196 return llvm::formatv("} // namespace {0}\n",
197 ref);
198 }),
199 "") +
200 dict["NAMESPACE_CLOSE"])
201 .str();
202}
203
204/// Fills a dictionary with values from DialectStrings
205static void fillDict(irdl::detail::dictionary &dict,
206 const DialectStrings &strings) {
207 dict["DIALECT_NAME"] = strings.dialectName;
208 dict["DIALECT_BASE_TYPE_NAME"] = strings.dialectBaseTypeName;
209 dict["DIALECT_CPP_NAME"] = strings.dialectCppName;
210 dict["DIALECT_CPP_SHORT_NAME"] = strings.dialectCppShortName;
211 dict["NAMESPACE_OPEN"] = strings.namespaceOpen;
212 dict["NAMESPACE_CLOSE"] = strings.namespaceClose;
213 dict["NAMESPACE_PATH"] = strings.namespacePath;
214}
215
216static LogicalResult generateTypedefList(irdl::DialectOp &dialect,
217 SmallVector<std::string> &typeNames) {
218 auto typeOps = dialect.getOps<irdl::TypeOp>();
219 auto range = llvm::map_range(typeOps, typeToCppName);
220 typeNames = SmallVector<std::string>(range);
221 return success();
222}
223
224static LogicalResult generateOpList(irdl::DialectOp &dialect,
225 SmallVector<std::string> &opNames) {
226 auto operationOps = dialect.getOps<irdl::OperationOp>();
227 auto range = llvm::map_range(operationOps, opToScopedCppName);
228 opNames = SmallVector<std::string>(range);
229 return success();
230}
231
232} // namespace
233
234static LogicalResult generateTypeInclude(irdl::TypeOp type, raw_ostream &output,
236 static const auto typeDeclTemplate = irdl::detail::Template(
237#include "Templates/TypeDecl.txt"
238 );
239
240 fillDict(dict, getStrings(type));
241 typeDeclTemplate.render(output, dict);
242
243 return success();
244}
245
247 const OpStrings &opStrings) {
248 auto opGetters = std::string{};
249 auto resGetters = std::string{};
250 auto regionGetters = std::string{};
251 auto regionAdaptorGetters = std::string{};
252
253 for (size_t i = 0, end = opStrings.opOperandNames.size(); i < end; ++i) {
254 const auto op =
255 llvm::convertToCamelFromSnakeCase(opStrings.opOperandNames[i], true);
256 opGetters += llvm::formatv("::mlir::Value get{0}() { return "
257 "getStructuredOperands({1}).front(); }\n ",
258 op, i);
259 }
260 for (size_t i = 0, end = opStrings.opResultNames.size(); i < end; ++i) {
261 const auto op =
262 llvm::convertToCamelFromSnakeCase(opStrings.opResultNames[i], true);
263 resGetters += llvm::formatv(
264 R"(::mlir::Value get{0}() { return ::llvm::cast<::mlir::Value>(getStructuredResults({1}).front()); }
265 )",
266 op, i);
267 }
268
269 for (size_t i = 0, end = opStrings.opRegionNames.size(); i < end; ++i) {
270 const auto op =
271 llvm::convertToCamelFromSnakeCase(opStrings.opRegionNames[i], true);
272 regionAdaptorGetters += llvm::formatv(
273 R"(::mlir::Region &get{0}() { return *getRegions()[{1}]; }
274 )",
275 op, i);
276 regionGetters += llvm::formatv(
277 R"(::mlir::Region &get{0}() { return (*this)->getRegion({1}); }
278 )",
279 op, i);
280 }
281
282 dict["OP_OPERAND_GETTER_DECLS"] = opGetters;
283 dict["OP_RESULT_GETTER_DECLS"] = resGetters;
284 dict["OP_REGION_ADAPTER_GETTER_DECLS"] = regionAdaptorGetters;
285 dict["OP_REGION_GETTER_DECLS"] = regionGetters;
286}
287
289 const OpStrings &opStrings) {
290 std::string buildDecls;
291 llvm::raw_string_ostream stream{buildDecls};
292
293 auto resultParams =
294 llvm::join(llvm::map_range(opStrings.opResultNames,
295 [](StringRef name) -> std::string {
296 return llvm::formatv(
297 "::mlir::Type {0}, ",
298 llvm::convertToCamelFromSnakeCase(name));
299 }),
300 "");
301
302 auto operandParams =
303 llvm::join(llvm::map_range(opStrings.opOperandNames,
304 [](StringRef name) -> std::string {
305 return llvm::formatv(
306 "::mlir::Value {0}, ",
307 llvm::convertToCamelFromSnakeCase(name));
308 }),
309 "");
310
311 stream << llvm::formatv(
312 R"(static void build(::mlir::OpBuilder &opBuilder, ::mlir::OperationState &opState, {0} {1} ::llvm::ArrayRef<::mlir::NamedAttribute> attributes = {{});)",
313 resultParams, operandParams);
314 stream << "\n";
315 stream << llvm::formatv(
316 R"(static {0} create(::mlir::OpBuilder &opBuilder, ::mlir::Location location, {1} {2} ::llvm::ArrayRef<::mlir::NamedAttribute> attributes = {{});)",
317 opStrings.opCppName, resultParams, operandParams);
318 stream << "\n";
319 stream << llvm::formatv(
320 R"(static {0} create(::mlir::ImplicitLocOpBuilder &opBuilder, {1} {2} ::llvm::ArrayRef<::mlir::NamedAttribute> attributes = {{});)",
321 opStrings.opCppName, resultParams, operandParams);
322 stream << "\n";
323 dict["OP_BUILD_DECLS"] = buildDecls;
325
326// add traits to the dictionary, return true if any were added
327static SmallVector<std::string> generateTraits(irdl::OperationOp op,
328 const OpStrings &strings) {
329 SmallVector<std::string> cppTraitNames;
330 if (!strings.opRegionNames.empty()) {
331 cppTraitNames.push_back(
332 llvm::formatv("::mlir::OpTrait::NRegions<{0}>::Impl",
333 strings.opRegionNames.size())
334 .str());
335
336 // Requires verifyInvariantsImpl is implemented on the op
337 cppTraitNames.emplace_back("::mlir::OpTrait::OpInvariants");
338 }
339 return cppTraitNames;
341
342static LogicalResult
343generateOperationInclude(irdl::OperationOp op, raw_ostream &output,
344 const irdl::detail::dictionary &dict) {
345 static const auto perOpDeclTemplate = irdl::detail::Template(
346#include "Templates/PerOperationDecl.txt"
347 );
348 const auto opStrings = getStrings(op);
349 auto opDict = dict;
350 fillDict(opDict, opStrings);
351
352 SmallVector<std::string> traitNames = generateTraits(op, opStrings);
353 if (traitNames.empty())
354 opDict["OP_TEMPLATE_ARGS"] = opStrings.opCppName;
355 else
356 opDict["OP_TEMPLATE_ARGS"] = llvm::formatv("{0}, {1}", opStrings.opCppName,
357 llvm::join(traitNames, ", "));
358
359 generateOpGetterDeclarations(opDict, opStrings);
360 generateOpBuilderDeclarations(opDict, opStrings);
361
362 perOpDeclTemplate.render(output, opDict);
363 return success();
364}
365
366static LogicalResult generateInclude(irdl::DialectOp dialect,
367 raw_ostream &output,
368 DialectStrings &dialectStrings) {
369 static const auto dialectDeclTemplate = irdl::detail::Template(
370#include "Templates/DialectDecl.txt"
371 );
372 static const auto typeHeaderDeclTemplate = irdl::detail::Template(
373#include "Templates/TypeHeaderDecl.txt"
374 );
375
377 fillDict(dict, dialectStrings);
378
379 dialectDeclTemplate.render(output, dict);
380 typeHeaderDeclTemplate.render(output, dict);
381
382 auto typeOps = dialect.getOps<irdl::TypeOp>();
383 auto operationOps = dialect.getOps<irdl::OperationOp>();
384
385 for (auto &&typeOp : typeOps) {
386 if (failed(generateTypeInclude(typeOp, output, dict)))
387 return failure();
388 }
389
391 if (failed(generateOpList(dialect, opNames)))
392 return failure();
393
394 auto classDeclarations =
395 llvm::join(llvm::map_range(
396 opNames,
397 [](llvm::StringRef name) -> std::string {
398 if (name.contains("::")) {
399 auto [scope, className] = name.rsplit("::");
400 return llvm::formatv("namespace {0} {{\nclass {1};\n}",
401 scope, className);
402 }
403 return llvm::formatv("class {0};", name);
404 }),
405 "\n");
406 const auto forwardDeclarations = llvm::formatv(
407 "{1}\n{0}\n{2}", std::move(classDeclarations),
408 dialectStrings.namespaceOpen, dialectStrings.namespaceClose);
409
410 output << forwardDeclarations;
411 for (auto &&operationOp : operationOps) {
412 if (failed(generateOperationInclude(operationOp, output, dict)))
413 return failure();
414 }
415
416 return success();
417}
418
420 irdl::detail::dictionary &dict, irdl::OperationOp op,
421 const OpStrings &strings, SmallVectorImpl<std::string> &verifierHelpers,
422 SmallVectorImpl<std::string> &verifierCalls) {
423 auto regionsOp = op.getOp<irdl::RegionsOp>();
424 if (strings.opRegionNames.empty() || !regionsOp)
425 return;
426
427 for (size_t i = 0; i < strings.opRegionNames.size(); ++i) {
428 std::string regionName = strings.opRegionNames[i];
429 std::string helperFnName =
430 llvm::formatv("__mlir_irdl_local_region_constraint_{0}_{1}",
431 strings.opCppName, regionName)
432 .str();
433
434 // Extract the actual region constraint from the IRDL RegionOp
435 std::string condition = "true";
436 std::string textualConditionName = "any region";
437
438 if (auto regionDefOp =
439 regionsOp->getArgs()[i].getDefiningOp<irdl::RegionOp>()) {
440 // Generate constraint condition based on RegionOp attributes
441 SmallVector<std::string> conditionParts;
442 SmallVector<std::string> descriptionParts;
443
444 // Check number of blocks constraint
445 if (auto blockCount = regionDefOp.getNumberOfBlocks()) {
446 conditionParts.push_back(
447 llvm::formatv("region.getBlocks().size() == {0}",
448 blockCount.value())
449 .str());
450 descriptionParts.push_back(
451 llvm::formatv("exactly {0} block(s)", blockCount.value()).str());
452 }
453
454 // Check entry block arguments constraint
455 if (regionDefOp.getConstrainedArguments()) {
456 size_t expectedArgCount = regionDefOp.getEntryBlockArgs().size();
457 conditionParts.push_back(
458 llvm::formatv("region.getNumArguments() == {0}", expectedArgCount)
459 .str());
460 descriptionParts.push_back(
461 llvm::formatv("{0} entry block argument(s)", expectedArgCount)
462 .str());
463 }
464
465 // Combine conditions
466 if (!conditionParts.empty()) {
467 condition = llvm::join(conditionParts, " && ");
468 }
469
470 // Generate descriptive error message
471 if (!descriptionParts.empty()) {
472 textualConditionName =
473 llvm::formatv("region with {0}",
474 llvm::join(descriptionParts, " and "))
475 .str();
476 }
477 }
478
479 verifierHelpers.push_back(llvm::formatv(
480 R"(static ::llvm::LogicalResult {0}(::mlir::Operation *op, ::mlir::Region &region, ::llvm::StringRef regionName, unsigned regionIndex) {{
481 if (!({1})) {{
482 return op->emitOpError("region #") << regionIndex
483 << (regionName.empty() ? " " : " ('" + regionName + "') ")
484 << "failed to verify constraint: {2}";
485 }
486 return ::mlir::success();
487})",
488 helperFnName, condition, textualConditionName));
489
490 verifierCalls.push_back(llvm::formatv(R"(
491 if (::mlir::failed({0}(*this, (*this)->getRegion({1}), "{2}", {1})))
492 return ::mlir::failure();)",
493 helperFnName, i, regionName)
494 .str());
495 }
496}
497
499 irdl::OperationOp op, const OpStrings &strings) {
500 SmallVector<std::string> verifierHelpers;
501 SmallVector<std::string> verifierCalls;
502
503 generateRegionConstraintVerifiers(dict, op, strings, verifierHelpers,
504 verifierCalls);
505
506 // Add an overall verifier that sequences the helper calls
507 std::string verifierDef =
508 llvm::formatv(R"(
509::llvm::LogicalResult {0}::verifyInvariantsImpl() {{
510 if(::mlir::failed(verify()))
511 return ::mlir::failure();
512
513 {1}
514
515 return ::mlir::success();
516})",
517 strings.opCppName, llvm::join(verifierCalls, "\n"));
518
519 dict["OP_VERIFIER_HELPERS"] = llvm::join(verifierHelpers, "\n");
520 dict["OP_VERIFIER"] = verifierDef;
521}
522
523static std::string generateOpDefinition(irdl::detail::dictionary &dict,
524 irdl::OperationOp op) {
525 static const auto perOpDefTemplate = mlir::irdl::detail::Template{
526#include "Templates/PerOperationDef.txt"
527 };
528
529 auto opStrings = getStrings(op);
530 auto opDict = dict;
531 fillDict(opDict, opStrings);
532
533 auto resultTypes = llvm::join(
534 llvm::map_range(opStrings.opResultNames,
535 [](StringRef attr) -> std::string {
536 return llvm::formatv("::mlir::Type {0}, ", attr);
537 }),
538 "");
539 auto operandTypes = llvm::join(
540 llvm::map_range(opStrings.opOperandNames,
541 [](StringRef attr) -> std::string {
542 return llvm::formatv("::mlir::Value {0}, ", attr);
543 }),
544 "");
545 auto operandAdder =
546 llvm::join(llvm::map_range(opStrings.opOperandNames,
547 [](StringRef attr) -> std::string {
548 return llvm::formatv(
549 " opState.addOperands({0});", attr);
550 }),
551 "\n");
552 auto resultAdder = llvm::join(
553 llvm::map_range(opStrings.opResultNames,
554 [](StringRef attr) -> std::string {
555 return llvm::formatv(" opState.addTypes({0});", attr);
556 }),
557 "\n");
558
559 const auto buildDefinition = llvm::formatv(
560 R"(
561void {0}::build(::mlir::OpBuilder &opBuilder, ::mlir::OperationState &opState, {1} {2} ::llvm::ArrayRef<::mlir::NamedAttribute> attributes) {{
562{3}
563{4}
564}
566{0} {0}::create(::mlir::OpBuilder &opBuilder, ::mlir::Location location, {1} {2} ::llvm::ArrayRef<::mlir::NamedAttribute> attributes) {{
567 ::mlir::OperationState __state__(location, getOperationName());
568 build(opBuilder, __state__, {5} {6} attributes);
569 auto __res__ = opBuilder.create(__state__);
570 assert((::llvm::isa<{0}>(__res__)) && "builder didn't return the right type");
571 return ::llvm::cast<{0}>(__res__);
572}
573
574{0} {0}::create(::mlir::ImplicitLocOpBuilder &opBuilder, {1} {2} ::llvm::ArrayRef<::mlir::NamedAttribute> attributes) {{
575 return create(opBuilder, opBuilder.getLoc(), {5} {6} attributes);
576}
577)",
578 opStrings.opCppName, std::move(resultTypes), std::move(operandTypes),
579 std::move(operandAdder), std::move(resultAdder),
580 llvm::join(opStrings.opResultNames, ",") +
581 (!opStrings.opResultNames.empty() ? "," : ""),
582 llvm::join(opStrings.opOperandNames, ",") +
583 (!opStrings.opOperandNames.empty() ? "," : ""));
584
585 opDict["OP_BUILD_DEFS"] = buildDefinition;
586
587 generateVerifiers(opDict, op, opStrings);
588
589 std::string str;
590 llvm::raw_string_ostream stream{str};
591 perOpDefTemplate.render(stream, opDict);
592 return str;
593}
594
595static std::string
596generateTypeVerifierCase(StringRef name, const DialectStrings &dialectStrings) {
597 return llvm::formatv(
598 R"(.Case({1}::{0}::getMnemonic(), [&](llvm::StringRef, llvm::SMLoc) {
599value = {1}::{0}::get(parser.getContext());
600return ::mlir::success(!!value);
601}))",
602 name, dialectStrings.namespacePath);
603}
604
605static LogicalResult generateLib(irdl::DialectOp dialect, raw_ostream &output,
606 DialectStrings &dialectStrings) {
607
608 static const auto typeHeaderDefTemplate = mlir::irdl::detail::Template{
609#include "Templates/TypeHeaderDef.txt"
610 };
611 static const auto typeDefTemplate = mlir::irdl::detail::Template{
612#include "Templates/TypeDef.txt"
613 };
614 static const auto dialectDefTemplate = mlir::irdl::detail::Template{
615#include "Templates/DialectDef.txt"
616 };
617
619 fillDict(dict, dialectStrings);
620
621 typeHeaderDefTemplate.render(output, dict);
622
623 SmallVector<std::string> typeNames;
624 if (failed(generateTypedefList(dialect, typeNames)))
625 return failure();
626
627 dict["TYPE_LIST"] = llvm::join(
628 llvm::map_range(typeNames,
629 [&dialectStrings](llvm::StringRef name) -> std::string {
630 return llvm::formatv(
631 "{0}::{1}", dialectStrings.namespacePath, name);
632 }),
633 ",\n");
634
635 auto typeVerifierGenerator =
636 [&dialectStrings](llvm::StringRef name) -> std::string {
637 return generateTypeVerifierCase(name, dialectStrings);
638 };
639
640 auto typeCase =
641 llvm::join(llvm::map_range(typeNames, typeVerifierGenerator), "\n");
642
643 dict["TYPE_PARSER"] = llvm::formatv(
644 R"(static ::mlir::OptionalParseResult generatedTypeParser(::mlir::AsmParser &parser, ::llvm::StringRef *mnemonic, ::mlir::Type &value) {
645 return ::mlir::AsmParser::KeywordSwitch<::mlir::OptionalParseResult>(parser)
646 {0}
647 .Default([&](llvm::StringRef keyword, llvm::SMLoc) {{
648 *mnemonic = keyword;
649 return std::nullopt;
650 });
651})",
652 std::move(typeCase));
653
654 auto typePrintCase =
655 llvm::join(llvm::map_range(typeNames,
656 [&](llvm::StringRef name) -> std::string {
657 return llvm::formatv(
658 R"(.Case<{1}::{0}>([&](auto t) {
659 printer << {1}::{0}::getMnemonic();
660 return ::mlir::success();
661 }))",
662 name, dialectStrings.namespacePath);
663 }),
664 "\n");
665 dict["TYPE_PRINTER"] = llvm::formatv(
666 R"(static ::llvm::LogicalResult generatedTypePrinter(::mlir::Type def, ::mlir::AsmPrinter &printer) {
667 return ::llvm::TypeSwitch<::mlir::Type, ::llvm::LogicalResult>(def)
668 {0}
669 .Default([](auto) {{ return ::mlir::failure(); });
670})",
671 std::move(typePrintCase));
672
673 dict["TYPE_DEFINES"] =
674 join(map_range(typeNames,
675 [&](StringRef name) -> std::string {
676 return formatv("MLIR_DEFINE_EXPLICIT_TYPE_ID({1}::{0})",
677 name, dialectStrings.namespacePath);
678 }),
679 "\n");
680
681 typeDefTemplate.render(output, dict);
682
683 auto operations = dialect.getOps<irdl::OperationOp>();
685 if (failed(generateOpList(dialect, opNames)))
686 return failure();
687
688 const auto commaSeparatedOpList = llvm::join(
689 map_range(opNames,
690 [&dialectStrings](llvm::StringRef name) -> std::string {
691 return llvm::formatv("{0}::{1}", dialectStrings.namespacePath,
692 name);
693 }),
694 ",\n");
695
696 const auto opDefinitionGenerator = [&dict](irdl::OperationOp op) {
697 return generateOpDefinition(dict, op);
698 };
699
700 const auto perOpDefinitions =
701 llvm::join(llvm::map_range(operations, opDefinitionGenerator), "\n");
703 dict["OP_LIST"] = commaSeparatedOpList;
704 dict["OP_CLASSES"] = perOpDefinitions;
705 output << perOpDefinitions;
706 dialectDefTemplate.render(output, dict);
707
708 return success();
709}
710
711static LogicalResult verifySupported(irdl::DialectOp dialect) {
712 LogicalResult res = success();
713 dialect.walk([&](mlir::Operation *op) {
714 res =
716 .Case(([](irdl::DialectOp) { return success(); }))
717 .Case(([](irdl::OperationOp) { return success(); }))
718 .Case(([](irdl::TypeOp) { return success(); }))
719 .Case(([](irdl::OperandsOp op) -> LogicalResult {
720 if (llvm::all_of(
721 op.getVariadicity(), [](irdl::VariadicityAttr attr) {
722 return attr.getValue() == irdl::Variadicity::single;
723 }))
724 return success();
725 return op.emitError("IRDL C++ translation does not yet support "
726 "variadic operations");
727 }))
728 .Case(([](irdl::ResultsOp op) -> LogicalResult {
729 if (llvm::all_of(
730 op.getVariadicity(), [](irdl::VariadicityAttr attr) {
731 return attr.getValue() == irdl::Variadicity::single;
732 }))
733 return success();
734 return op.emitError(
735 "IRDL C++ translation does not yet support variadic results");
736 }))
737 .Case(([](irdl::AnyOp) { return success(); }))
738 .Case(([](irdl::RegionOp) { return success(); }))
739 .Case(([](irdl::RegionsOp) { return success(); }))
740 .Default([](mlir::Operation *op) -> LogicalResult {
741 return op->emitError("IRDL C++ translation does not yet support "
742 "translation of ")
743 << op->getName() << " operation";
744 });
745
746 if (failed(res))
747 return WalkResult::interrupt();
748
749 return WalkResult::advance();
750 });
751
752 return res;
753}
754
755LogicalResult
757 raw_ostream &output) {
758 static const auto typeDefTempl = detail::Template(
759#include "Templates/TypeDef.txt"
760 );
761
762 llvm::SmallMapVector<DialectOp, DialectStrings, 2> dialectStringTable;
763
764 for (auto dialect : dialects) {
765 if (failed(verifySupported(dialect)))
766 return failure();
767
768 StringRef dialectName = dialect.getSymName();
769
770 SmallVector<SmallString<8>> namespaceAbsolutePath{{"mlir"}, dialectName};
771 std::string namespaceOpen;
772 std::string namespaceClose;
773 std::string namespacePath;
774 llvm::raw_string_ostream namespaceOpenStream(namespaceOpen);
775 llvm::raw_string_ostream namespaceCloseStream(namespaceClose);
776 llvm::raw_string_ostream namespacePathStream(namespacePath);
777 for (auto &pathElement : namespaceAbsolutePath) {
778 namespaceOpenStream << "namespace " << pathElement << " {\n";
779 namespacePathStream << "::" << pathElement;
780 }
781
782 for (auto &pathElement : llvm::reverse(namespaceAbsolutePath))
783 namespaceCloseStream << "} // namespace " << pathElement << "\n";
784
785 std::string cppShortName =
786 llvm::convertToCamelFromSnakeCase(dialectName, true);
787 std::string dialectBaseTypeName = llvm::formatv("{0}Type", cppShortName);
788 std::string cppName = llvm::formatv("{0}Dialect", cppShortName);
789
790 DialectStrings dialectStrings;
791 dialectStrings.dialectName = dialectName;
792 dialectStrings.dialectBaseTypeName = std::move(dialectBaseTypeName);
793 dialectStrings.dialectCppName = std::move(cppName);
794 dialectStrings.dialectCppShortName = std::move(cppShortName);
795 dialectStrings.namespaceOpen = std::move(namespaceOpen);
796 dialectStrings.namespaceClose = std::move(namespaceClose);
797 dialectStrings.namespacePath = std::move(namespacePath);
798
799 dialectStringTable[dialect] = std::move(dialectStrings);
800 }
801
802 // generate the actual header
803 output << headerTemplateText;
804
805 output << llvm::formatv("#ifdef {0}\n#undef {0}\n", declarationMacroFlag);
806 for (auto dialect : dialects) {
807
808 auto &dialectStrings = dialectStringTable[dialect];
809 auto &dialectName = dialectStrings.dialectName;
810
811 if (failed(generateInclude(dialect, output, dialectStrings)))
812 return dialect->emitError("Error in Dialect " + dialectName +
813 " while generating headers");
814 }
815 output << llvm::formatv("#endif // #ifdef {}\n", declarationMacroFlag);
816
817 output << llvm::formatv("#ifdef {0}\n#undef {0}\n ", definitionMacroFlag);
818 for (auto &dialect : dialects) {
819 auto &dialectStrings = dialectStringTable[dialect];
820 auto &dialectName = dialectStrings.dialectName;
821
822 if (failed(generateLib(dialect, output, dialectStrings)))
823 return dialect->emitError("Error in Dialect " + dialectName +
824 " while generating library");
825 }
826 output << llvm::formatv("#endif // #ifdef {}\n", definitionMacroFlag);
827
828 return success();
829}
return success()
std::string join(const Ts &...args)
Helper function to concatenate arguments into a std::string.
static LogicalResult verifySupported(irdl::DialectOp dialect)
static LogicalResult generateInclude(irdl::DialectOp dialect, raw_ostream &output, DialectStrings &dialectStrings)
static void generateOpGetterDeclarations(irdl::detail::dictionary &dict, const OpStrings &opStrings)
static SmallVector< std::string > generateTraits(irdl::OperationOp op, const OpStrings &strings)
static void generateOpBuilderDeclarations(irdl::detail::dictionary &dict, const OpStrings &opStrings)
constexpr char declarationMacroFlag[]
Definition IRDLToCpp.cpp:29
constexpr char headerTemplateText[]
Definition IRDLToCpp.cpp:25
static std::string generateTypeVerifierCase(StringRef name, const DialectStrings &dialectStrings)
static LogicalResult generateTypeInclude(irdl::TypeOp type, raw_ostream &output, irdl::detail::dictionary &dict)
static LogicalResult generateOperationInclude(irdl::OperationOp op, raw_ostream &output, const irdl::detail::dictionary &dict)
static void generateRegionConstraintVerifiers(irdl::detail::dictionary &dict, irdl::OperationOp op, const OpStrings &strings, SmallVectorImpl< std::string > &verifierHelpers, SmallVectorImpl< std::string > &verifierCalls)
static LogicalResult generateLib(irdl::DialectOp dialect, raw_ostream &output, DialectStrings &dialectStrings)
constexpr char definitionMacroFlag[]
Definition IRDLToCpp.cpp:30
static std::string generateOpDefinition(irdl::detail::dictionary &dict, irdl::OperationOp op)
static void generateVerifiers(irdl::detail::dictionary &dict, irdl::OperationOp op, const OpStrings &strings)
Attributes are known-constant values of operations.
Definition Attributes.h:25
Operation is the basic unit of execution within MLIR.
Definition Operation.h:87
InFlightDiagnostic emitError(const Twine &message={})
Emit an error about fatal conditions with this operation, reporting up to any diagnostic handlers tha...
OperationName getName()
The name of an operation is the key identifier for it.
Definition Operation.h:115
static WalkResult advance()
Definition WalkResult.h:47
static WalkResult interrupt()
Definition WalkResult.h:46
Template Code as used by IRDL-to-Cpp.
llvm::StringMap< llvm::SmallString< 8 > > dictionary
A dictionary stores a mapping of template variable names to their assigned string values.
LogicalResult translateIRDLDialectToCpp(llvm::ArrayRef< irdl::DialectOp > dialects, raw_ostream &output)
Translates an IRDL dialect definition to a C++ definition that can be used with MLIR.
detail::InFlightRemark failed(Location loc, RemarkOpts opts)
Report an optimization remark that failed.
Definition Remarks.h:734
Include the generated interface declarations.