1//===-- CIRLoweringEmitter.cpp - Generate CIR lowering patterns -----------===//
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// This TableGen backend emits CIR operation lowering patterns.
10//
11//===----------------------------------------------------------------------===//
12
13#include "TableGenBackends.h"
14#include "llvm/TableGen/Error.h"
15#include "llvm/TableGen/Record.h"
16#include "llvm/TableGen/TableGenBackend.h"
17#include <string>
18#include <utility>
19#include <vector>
20
21using namespace llvm;
22using namespace clang;
23
24namespace {
25std::vector<std::string> CXXABILoweringPatterns;
26std::vector<std::string> CXXABILoweringPatternsList;
27std::vector<std::string> CXXABILoweringAttrAlwaysLegal;
28std::vector<std::string> LLVMLoweringPatterns;
29std::vector<std::string> LLVMLoweringPatternsList;
30std::string CIRAttrToValueVisitFunc;
31std::vector<std::string> CIRAttrToValueVisitorCaseTypes;
32std::vector<std::string> CIRAttrToValueVisitorDecls;
33
34struct CustomLoweringCtor {
35 struct Param {
36 std::string Type;
37 std::string Name;
38 };
39
40 std::vector<Param> Params;
41};
42
43// Adapted from mlir/lib/TableGen/Operator.cpp
44// Returns the C++ class name of the operation, which is the name of the
45// operation with the dialect prefix removed and the first underscore removed.
46// If the operation name starts with an underscore, the underscore is considered
47// part of the class name.
48std::string GetOpCppClassName(const Record *OpRecord) {
49 StringRef Name = OpRecord->getName();
50 StringRef Prefix;
51 StringRef CppClassName;
52 std::tie(args&: Prefix, args&: CppClassName) = Name.split(Separator: '_');
53 if (Prefix.empty()) {
54 // Class name with a leading underscore and without dialect prefix
55 return Name.str();
56 }
57 if (CppClassName.empty()) {
58 // Class name without dialect prefix
59 return Prefix.str();
60 }
61
62 return CppClassName.str();
63}
64
65std::string GetOpABILoweringPatternName(llvm::StringRef OpName) {
66 std::string Name = "CIR";
67 Name += OpName;
68 Name += "ABILowering";
69 return Name;
70}
71
72std::string GetOpLLVMLoweringPatternName(llvm::StringRef OpName) {
73 std::string Name = "CIRToLLVM";
74 Name += OpName;
75 Name += "Lowering";
76 return Name;
77}
78std::optional<CustomLoweringCtor> parseCustomLoweringCtor(const Record *R) {
79 if (!R)
80 return std::nullopt;
81
82 CustomLoweringCtor Ctor;
83 const DagInit *Args = R->getValueAsDag(FieldName: "dagParams");
84
85 for (const auto &[Arg, Name] : Args->getArgAndNames()) {
86 Ctor.Params.push_back(
87 x: {.Type: Arg->getAsUnquotedString(), .Name: Name->getAsUnquotedString()});
88 }
89
90 return Ctor;
91}
92
93void emitCustomParamList(raw_ostream &Code,
94 ArrayRef<CustomLoweringCtor::Param> Params) {
95 for (const CustomLoweringCtor::Param &Param : Params) {
96 Code << ", ";
97 Code << Param.Type << " " << Param.Name;
98 }
99}
100
101void emitCustomInitList(raw_ostream &Code,
102 ArrayRef<CustomLoweringCtor::Param> Params) {
103 for (const CustomLoweringCtor::Param &P : Params)
104 Code << ", " << P.Name << "(" << P.Name << ")";
105}
106
107void GenerateABILoweringPattern(llvm::StringRef OpName,
108 llvm::StringRef PatternName) {
109 std::string CodeBuffer;
110 llvm::raw_string_ostream Code(CodeBuffer);
111
112 Code << "class " << PatternName
113 << " : public mlir::OpConversionPattern<cir::" << OpName << "> {\n";
114 Code << " [[maybe_unused]] mlir::DataLayout *dataLayout;\n";
115 Code << " [[maybe_unused]] cir::LowerModule *lowerModule;\n";
116 Code << "\n";
117
118 Code << "public:\n";
119 Code << " " << PatternName
120 << "(mlir::MLIRContext *context, const mlir::TypeConverter "
121 "&typeConverter, mlir::DataLayout &dataLayout, cir::LowerModule "
122 "&lowerModule)\n";
123 Code << " : OpConversionPattern<cir::" << OpName
124 << ">(typeConverter, context), dataLayout(&dataLayout), "
125 "lowerModule(&lowerModule) {}\n";
126 Code << "\n";
127
128 Code << " mlir::LogicalResult matchAndRewrite(cir::" << OpName
129 << " op, OpAdaptor adaptor, mlir::ConversionPatternRewriter &rewriter) "
130 "const override;\n";
131
132 Code << "};\n";
133
134 CXXABILoweringPatterns.push_back(x: std::move(CodeBuffer));
135}
136
137void GenerateLLVMLoweringPattern(
138 llvm::StringRef OpName, llvm::StringRef PatternName, bool IsRecursive,
139 llvm::StringRef ExtraDecl, const Record *CustomCtorRec,
140 llvm::StringRef LLVMOp, llvm::StringRef ConstrainedLLVMIntrinsic,
141 bool ConstrainedHasRoundingMode, bool HasZeroResult) {
142 std::optional<CustomLoweringCtor> CustomCtor =
143 parseCustomLoweringCtor(R: CustomCtorRec);
144 std::string CodeBuffer;
145 llvm::raw_string_ostream Code(CodeBuffer);
146
147 Code << "class " << PatternName
148 << " : public mlir::OpConversionPattern<cir::" << OpName << "> {\n";
149 Code << " [[maybe_unused]] mlir::DataLayout const &dataLayout;\n";
150 Code << " [[maybe_unused]] mlir::SymbolTableCollection &symbolTables;\n";
151
152 if (CustomCtor) {
153 for (const CustomLoweringCtor::Param &P : CustomCtor->Params)
154 Code << " " << P.Type << " " << P.Name << ";\n";
155 }
156
157 Code << "\n";
158
159 Code << "public:\n";
160 Code << " using mlir::OpConversionPattern<cir::" << OpName
161 << ">::OpConversionPattern;\n";
162
163 // Constructor
164 Code << " " << PatternName
165 << "(const mlir::TypeConverter &typeConverter, "
166 "mlir::MLIRContext *context, const mlir::DataLayout &dataLayout, "
167 "mlir::SymbolTableCollection &symbolTables";
168
169 if (CustomCtor)
170 emitCustomParamList(Code, Params: CustomCtor->Params);
171
172 Code << ")\n";
173
174 Code << " : OpConversionPattern<cir::" << OpName
175 << ">(typeConverter, context), dataLayout(dataLayout), "
176 "symbolTables(symbolTables)";
177
178 if (CustomCtor)
179 emitCustomInitList(Code, Params: CustomCtor->Params);
180
181 Code << " {\n";
182
183 if (IsRecursive)
184 Code << " setHasBoundedRewriteRecursion();\n";
185
186 Code << " }\n\n";
187
188 if (!ConstrainedLLVMIntrinsic.empty()) {
189 // Generate a matchAndRewrite body for a floating-point operation that
190 // carries an optional `fenv` attribute. The shared `lowerConstrainableFPOp`
191 // helper lowers to the plain `llvmOp` when no `fenv` attribute is present
192 // and to a call to the constrained floating-point intrinsic when one is.
193 assert(!LLVMOp.empty() && "constrainedLLVMIntrinsic requires llvmOp");
194 Code
195 << " mlir::LogicalResult matchAndRewrite(cir::" << OpName
196 << " op, OpAdaptor adaptor, mlir::ConversionPatternRewriter &rewriter) "
197 "const override {\n";
198 Code << " return lowerConstrainableFPOp<mlir::LLVM::" << LLVMOp
199 << ">(\n";
200 Code << " op, adaptor.getOperands(), op.getFenvAttr(),\n";
201 Code << " *getTypeConverter(), rewriter, \""
202 << ConstrainedLLVMIntrinsic << "\",\n";
203 Code << " /*hasRoundingMode=*/"
204 << (ConstrainedHasRoundingMode ? "true" : "false") << ");\n";
205 Code << " }\n";
206 } else if (!LLVMOp.empty()) {
207 // Generate the matchAndRewrite body automatically.
208 Code
209 << " mlir::LogicalResult matchAndRewrite(cir::" << OpName
210 << " op, OpAdaptor adaptor, mlir::ConversionPatternRewriter &rewriter) "
211 "const override {\n";
212 if (HasZeroResult) {
213 Code << " rewriter.replaceOpWithNewOp<mlir::LLVM::" << LLVMOp
214 << ">(op, mlir::TypeRange{}, adaptor.getOperands());\n";
215 } else {
216 Code << " mlir::Type resTy = "
217 "typeConverter->convertType(op.getType());\n";
218 Code << " rewriter.replaceOpWithNewOp<mlir::LLVM::" << LLVMOp
219 << ">(op, resTy, adaptor.getOperands());\n";
220 }
221 Code << " return mlir::success();\n";
222 Code << " }\n";
223 } else {
224 Code
225 << " mlir::LogicalResult matchAndRewrite(cir::" << OpName
226 << " op, OpAdaptor adaptor, mlir::ConversionPatternRewriter &rewriter) "
227 "const override;\n";
228 }
229
230 if (!ExtraDecl.empty()) {
231 Code << "\nprivate:\n";
232 Code << ExtraDecl << "\n";
233 }
234
235 Code << "};\n";
236
237 LLVMLoweringPatterns.push_back(x: std::move(CodeBuffer));
238}
239
240void Generate(const Record *OpRecord) {
241 std::string OpName = GetOpCppClassName(OpRecord);
242
243 if (OpRecord->getValueAsBit(FieldName: "hasCXXABILowering")) {
244 std::string PatternName = GetOpABILoweringPatternName(OpName);
245 GenerateABILoweringPattern(OpName, PatternName);
246 CXXABILoweringPatternsList.push_back(x: std::move(PatternName));
247 }
248
249 if (OpRecord->getValueAsBit(FieldName: "hasLLVMLowering")) {
250 std::string PatternName = GetOpLLVMLoweringPatternName(OpName);
251 bool IsRecursive = OpRecord->getValueAsBit(FieldName: "isLLVMLoweringRecursive");
252 const Record *CustomCtor =
253 OpRecord->getValueAsOptionalDef(FieldName: "customLLVMLoweringConstructorDecl");
254 llvm::StringRef ExtraDecl =
255 OpRecord->getValueAsString(FieldName: "extraLLVMLoweringPatternDecl");
256
257 llvm::StringRef LLVMOp = OpRecord->getValueAsString(FieldName: "llvmOp");
258 llvm::StringRef ConstrainedLLVMIntrinsic =
259 OpRecord->getValueAsString(FieldName: "constrainedLLVMIntrinsic");
260 bool ConstrainedHasRoundingMode =
261 OpRecord->getValueAsBit(FieldName: "constrainedLLVMIntrinsicHasRoundingMode");
262
263 if (!LLVMOp.empty() && CustomCtor)
264 PrintFatalError(ErrorLoc: OpRecord->getLoc(),
265 Msg: "op '" + OpName +
266 "' has both llvmOp and a custom lowering "
267 "constructor, which is not supported");
268
269 if (!ConstrainedLLVMIntrinsic.empty() && LLVMOp.empty())
270 PrintFatalError(ErrorLoc: OpRecord->getLoc(),
271 Msg: "op '" + OpName +
272 "' has constrainedLLVMIntrinsic set but no llvmOp, "
273 "which is not supported");
274
275 const DagInit *ResultsDag = OpRecord->getValueAsDag(FieldName: "results");
276 bool IsZeroResult = ResultsDag->getNumArgs() == 0;
277 GenerateLLVMLoweringPattern(OpName, PatternName, IsRecursive, ExtraDecl,
278 CustomCtorRec: CustomCtor, LLVMOp, ConstrainedLLVMIntrinsic,
279 ConstrainedHasRoundingMode, HasZeroResult: IsZeroResult);
280 // Only automatically register patterns that use the default constructor.
281 // Patterns with a custom constructor must be manually registered by the
282 // lowering pass.
283 if (!CustomCtor)
284 LLVMLoweringPatternsList.push_back(x: std::move(PatternName));
285 }
286}
287
288void GenerateCIREnumAttrs(const Record *Record) {
289 std::string OpName = GetOpCppClassName(OpRecord: Record);
290 // EnumAttr is in a separate hierarchy, so we have to set these separately, as
291 // they never have an 'illegal' CXXABI type in them.
292 CXXABILoweringAttrAlwaysLegal.push_back(x: "cir::" + OpName);
293}
294
295void GenerateAttrToValueVisitor(const Record *Rec) {
296 const Record *DialectRec = Rec->getValueAsDef(FieldName: "dialect");
297 llvm::StringRef Ns = DialectRec->getValueAsString(FieldName: "cppNamespace");
298 Ns.consume_front(Prefix: "::");
299 std::string CppClassRef = Ns.str();
300 CppClassRef += "::";
301 CppClassRef += Rec->getValueAsString(FieldName: "cppClassName");
302
303 std::string CodeBuffer;
304 llvm::raw_string_ostream Code(CodeBuffer);
305 Code << " mlir::Value visitCirAttr(" << CppClassRef << " attr);";
306 CIRAttrToValueVisitorDecls.push_back(x: std::move(CodeBuffer));
307 CIRAttrToValueVisitorCaseTypes.push_back(x: std::move(CppClassRef));
308}
309
310void GenerateAttrToValueVisitFunc() {
311 std::string CodeBuffer;
312 llvm::raw_string_ostream Code(CodeBuffer);
313 Code << " mlir::Value visit(mlir::Attribute attr) {\n"
314 << " return llvm::TypeSwitch<mlir::Attribute, mlir::Value>(attr)\n"
315 << " .Case<\n "
316 << llvm::join(R&: CIRAttrToValueVisitorCaseTypes, Separator: ",\n ")
317 << ">(\n"
318 << " [&](auto attrT) { return visitCirAttr(attrT); })\n"
319 << " .Default([this](mlir::Attribute attr) {\n"
320 << " mlir::emitError(parentOp->getLoc(), \"unsupported CIR "
321 "attribute in LLVM constant lowering\")\n"
322 << " << attr;\n"
323 << " return mlir::Value();\n"
324 << " });\n"
325 << " }\n";
326 CIRAttrToValueVisitFunc = std::move(CodeBuffer);
327}
328
329void GenerateCIRAttrs(const Record *Record) {
330 std::string OpName = GetOpCppClassName(OpRecord: Record);
331 if (!Record->getValueAsBit(FieldName: "canHaveIllegalCXXABIType"))
332 CXXABILoweringAttrAlwaysLegal.push_back(x: "cir::" + OpName);
333 if (Record->getValueAsBit(FieldName: "hasAttrToValueLowering"))
334 GenerateAttrToValueVisitor(Rec: Record);
335}
336} // namespace
337
338void clang::EmitCIRLowering(const llvm::RecordKeeper &RK,
339 llvm::raw_ostream &OS) {
340 emitSourceFileHeader(Desc: "Lowering patterns for CIR operations", OS);
341 for (const auto *OpRecord : RK.getAllDerivedDefinitions(ClassName: "CIR_Op"))
342 Generate(OpRecord);
343 for (const auto *OpRecord : RK.getAllDerivedDefinitions(ClassName: "EnumAttr"))
344 GenerateCIREnumAttrs(Record: OpRecord);
345 for (const auto *OpRecord : RK.getAllDerivedDefinitions(ClassName: "CIR_Attr"))
346 GenerateCIRAttrs(Record: OpRecord);
347 GenerateAttrToValueVisitFunc();
348
349 OS << "#ifdef GET_CIR_ATTR_TO_VALUE_VISITOR_DECLS\n"
350 << CIRAttrToValueVisitFunc << "\n"
351 << llvm::join(R&: CIRAttrToValueVisitorDecls, Separator: "\n") << "\n"
352 << "#endif // GET_CIR_ATTR_TO_VALUE_VISITOR_DECLS\n\n";
353
354 OS << "#ifdef GET_ABI_LOWERING_PATTERNS\n"
355 << llvm::join(R&: CXXABILoweringPatterns, Separator: "\n") << "#endif\n\n";
356 OS << "#ifdef GET_ABI_LOWERING_PATTERNS_LIST\n"
357 << llvm::join(R&: CXXABILoweringPatternsList, Separator: ",\n") << "\n#endif\n\n";
358
359 OS << "#ifdef GET_LLVM_LOWERING_PATTERNS\n"
360 << llvm::join(R&: LLVMLoweringPatterns, Separator: "\n") << "#endif\n\n";
361 OS << "#ifdef GET_LLVM_LOWERING_PATTERNS_LIST\n"
362 << llvm::join(R&: LLVMLoweringPatternsList, Separator: ",\n") << "\n#endif\n\n";
363
364 OS << "#ifdef CXX_ABI_ALWAYS_LEGAL_ATTRS\n"
365 << llvm::join(R&: CXXABILoweringAttrAlwaysLegal, Separator: ",\n") << "\n#endif\n\n";
366}
367