MLIR 24.0.0git
TranslateToCpp.cpp
Go to the documentation of this file.
1//===- TranslateToCpp.cpp - Translating to C++ calls ----------------------===//
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
13#include "mlir/IR/BuiltinOps.h"
15#include "mlir/IR/Dialect.h"
16#include "mlir/IR/Operation.h"
17#include "mlir/IR/SymbolTable.h"
18#include "mlir/IR/Value.h"
20#include "mlir/Support/LLVM.h"
22#include "llvm/ADT/ScopedHashTable.h"
23#include "llvm/ADT/StringExtras.h"
24#include "llvm/ADT/TypeSwitch.h"
25#include "llvm/Support/Casting.h"
26#include "llvm/Support/Debug.h"
27#include "llvm/Support/FormatVariadic.h"
28#include <stack>
29
30#define DEBUG_TYPE "translate-to-cpp"
31
32using namespace mlir;
33using namespace mlir::emitc;
34using llvm::formatv;
35
36/// Convenience functions to produce interleaved output with functions returning
37/// a LogicalResult. This is different than those in STLExtras as functions used
38/// on each element doesn't return a string.
39template <typename ForwardIterator, typename UnaryFunctor,
40 typename NullaryFunctor>
41static inline LogicalResult
43 UnaryFunctor eachFn, NullaryFunctor betweenFn) {
44 if (begin == end)
45 return success();
46 if (failed(eachFn(*begin)))
47 return failure();
48 ++begin;
49 for (; begin != end; ++begin) {
50 betweenFn();
51 if (failed(eachFn(*begin)))
52 return failure();
53 }
54 return success();
55}
56
57template <typename Container, typename UnaryFunctor>
58static inline LogicalResult interleaveCommaWithError(const Container &c,
59 raw_ostream &os,
60 UnaryFunctor eachFn) {
61 return interleaveWithError(c.begin(), c.end(), eachFn, [&]() { os << ", "; });
62}
63
64/// Return the precedence of a operator as an integer, higher values
65/// imply higher precedence.
66static FailureOr<int> getOperatorPrecedence(Operation *operation) {
68 .Case([&](emitc::AddressOfOp op) { return 15; })
69 .Case([&](emitc::AddOp op) { return 12; })
70 .Case([&](emitc::BitwiseAndOp op) { return 7; })
71 .Case([&](emitc::BitwiseLeftShiftOp op) { return 11; })
72 .Case([&](emitc::BitwiseNotOp op) { return 15; })
73 .Case([&](emitc::BitwiseOrOp op) { return 5; })
74 .Case([&](emitc::BitwiseRightShiftOp op) { return 11; })
75 .Case([&](emitc::BitwiseXorOp op) { return 6; })
76 .Case([&](emitc::CallOp op) { return 16; })
77 .Case([&](emitc::CallOpaqueOp op) { return 16; })
78 .Case([&](emitc::CastOp op) { return 15; })
79 .Case([&](emitc::CmpOp op) -> FailureOr<int> {
80 switch (op.getPredicate()) {
81 case emitc::CmpPredicate::eq:
82 case emitc::CmpPredicate::ne:
83 return 8;
84 case emitc::CmpPredicate::lt:
85 case emitc::CmpPredicate::le:
86 case emitc::CmpPredicate::gt:
87 case emitc::CmpPredicate::ge:
88 return 9;
89 case emitc::CmpPredicate::three_way:
90 return 10;
91 }
92 return op->emitError("unsupported cmp predicate");
93 })
94 .Case([&](emitc::ConditionalOp op) { return 2; })
95 .Case([&](emitc::ConstantOp op) { return 17; })
96 .Case([&](emitc::DereferenceOp op) { return 15; })
97 .Case([&](emitc::DivOp op) { return 13; })
98 .Case([&](emitc::GetGlobalOp op) { return 18; })
99 .Case([&](emitc::GetFieldOp op) { return 18; })
100 .Case([&](emitc::LiteralOp op) { return 18; })
101 .Case([&](emitc::LoadOp op) { return 16; })
102 .Case([&](emitc::LogicalAndOp op) { return 4; })
103 .Case([&](emitc::LogicalNotOp op) { return 15; })
104 .Case([&](emitc::LogicalOrOp op) { return 3; })
105 .Case([&](emitc::MemberCallOpaqueOp op) { return 16; })
106 .Case([&](emitc::MemberOfPtrOp op) { return 17; })
107 .Case([&](emitc::MemberOp op) { return 17; })
108 .Case([&](emitc::MulOp op) { return 13; })
109 .Case([&](emitc::PostDecrementOp op) { return 16; })
110 .Case([&](emitc::PostIncrementOp op) { return 16; })
111 .Case([&](emitc::PreDecrementOp op) { return 15; })
112 .Case([&](emitc::PreIncrementOp op) { return 15; })
113 .Case([&](emitc::RemOp op) { return 13; })
114 .Case([&](emitc::SubOp op) { return 12; })
115 .Case([&](emitc::SubscriptOp op) { return 17; })
116 .Case([&](emitc::UnaryMinusOp op) { return 15; })
117 .Case([&](emitc::UnaryPlusOp op) { return 15; })
118 .Default([](auto op) { return op->emitError("unsupported operation"); });
119}
120
121static bool shouldBeInlined(Operation *op);
122
123namespace {
124/// Emitter that uses dialect specific emitters to emit C++ code.
125struct CppEmitter {
126 explicit CppEmitter(raw_ostream &os, bool declareVariablesAtTop,
127 StringRef fileId);
128
129 /// Emits attribute or returns failure.
130 LogicalResult emitAttribute(Location loc, Attribute attr);
131
132 /// Emits operation 'op' with/without training semicolon or returns failure.
133 ///
134 /// For operations that should never be followed by a semicolon, like ForOp,
135 /// the `trailingSemicolon` argument is ignored and a semicolon is not
136 /// emitted.
137 LogicalResult emitOperation(Operation &op, bool trailingSemicolon);
138
139 /// Emits type 'type' or returns failure.
140 LogicalResult emitType(Location loc, Type type);
141
142 /// Emits array of types as a std::tuple of the emitted types.
143 /// - emits void for an empty array;
144 /// - emits the type of the only element for arrays of size one;
145 /// - emits a std::tuple otherwise;
146 LogicalResult emitTypes(Location loc, ArrayRef<Type> types);
147
148 /// Emits array of types as a std::tuple of the emitted types independently of
149 /// the array size.
150 LogicalResult emitTupleType(Location loc, ArrayRef<Type> types);
151
152 /// Emits an assignment for a variable which has been declared previously.
153 LogicalResult emitVariableAssignment(OpResult result);
154
155 /// Emits a variable declaration for a result of an operation.
156 LogicalResult emitVariableDeclaration(OpResult result,
157 bool trailingSemicolon);
158
159 /// Emits a declaration of a variable with the given type and name.
160 LogicalResult emitVariableDeclaration(Location loc, Type type,
161 StringRef name);
162
163 /// Emits the variable declaration and assignment prefix for 'op'.
164 /// - emits separate variable followed by std::tie for multi-valued operation;
165 /// - emits single type followed by variable for single result;
166 /// - emits nothing if no value produced by op;
167 /// Emits final '=' operator where a type is produced. Returns failure if
168 /// any result type could not be converted.
169 LogicalResult emitAssignPrefix(Operation &op);
170
171 /// Emits a global variable declaration or definition.
172 LogicalResult emitGlobalVariable(GlobalOp op);
173
174 /// Emits a label for the block.
175 LogicalResult emitLabel(Block &block);
176
177 /// Emits the operands and atttributes of the operation. All operands are
178 /// emitted first and then all attributes in alphabetical order.
179 LogicalResult emitOperandsAndAttributes(Operation &op,
180 ArrayRef<StringRef> exclude = {});
181
182 /// Emits the operands of the operation. All operands are emitted in order.
183 LogicalResult emitOperands(Operation &op);
184
185 /// Emits value as an operand of some operation. Unless \p isInBrackets is
186 /// true, operands emitted as sub-expressions will be parenthesized if needed
187 /// in order to enforce correct evaluation based on precedence and
188 /// associativity.
189 LogicalResult emitOperand(Value value, bool isInBrackets = false);
190
191 /// Emit an expression as a C expression.
192 LogicalResult emitExpression(Operation *op);
193
194 /// Return the existing or a new name for a Value.
195 StringRef getOrCreateName(Value val);
196
197 /// Return the existing or a new name for a loop induction variable of an
198 /// emitc::ForOp.
199 StringRef getOrCreateInductionVarName(Value val);
200
201 /// Return the existing or a new label of a Block.
202 StringRef getOrCreateName(Block &block);
203
204 LogicalResult emitInlinedExpression(Value value);
205
206 /// Whether to map an mlir integer to a unsigned integer in C++.
207 bool shouldMapToUnsigned(IntegerType::SignednessSemantics val);
208
209 /// Abstract RAII helper function to manage entering/exiting C++ scopes.
210 struct Scope {
211 ~Scope() { emitter.labelInScopeCount.pop(); }
212
213 private:
214 llvm::ScopedHashTableScope<Value, std::string> valueMapperScope;
215 llvm::ScopedHashTableScope<Block *, std::string> blockMapperScope;
216
217 protected:
218 Scope(CppEmitter &emitter)
219 : valueMapperScope(emitter.valueMapper),
220 blockMapperScope(emitter.blockMapper), emitter(emitter) {
221 emitter.labelInScopeCount.push(emitter.labelInScopeCount.top());
222 }
223 CppEmitter &emitter;
224 };
225
226 /// RAII helper function to manage entering/exiting functions, while re-using
227 /// value names.
228 struct FunctionScope : Scope {
229 FunctionScope(CppEmitter &emitter) : Scope(emitter) {
230 // Re-use value names.
231 emitter.resetValueCounter();
232 }
233 };
234
235 /// RAII helper function to manage entering/exiting emitc::forOp loops and
236 /// handle induction variable naming.
237 struct LoopScope : Scope {
238 LoopScope(CppEmitter &emitter) : Scope(emitter) {
239 emitter.increaseLoopNestingLevel();
240 }
241 ~LoopScope() { emitter.decreaseLoopNestingLevel(); }
242 };
243
244 /// Returns wether the Value is assigned to a C++ variable in the scope.
245 bool hasValueInScope(Value val);
246
247 // Returns whether a label is assigned to the block.
248 bool hasBlockLabel(Block &block);
249
250 /// Returns the output stream.
251 raw_indented_ostream &ostream() { return os; };
252
253 /// Returns if all variables for op results and basic block arguments need to
254 /// be declared at the beginning of a function.
255 bool shouldDeclareVariablesAtTop() { return declareVariablesAtTop; };
256
257 /// Returns whether this file op should be emitted
258 bool shouldEmitFile(FileOp file) {
259 return !fileId.empty() && file.getId() == fileId;
260 }
261
262 /// Is expression currently being emitted.
263 bool isEmittingExpression() { return !emittedExpressionPrecedence.empty(); }
264
265 /// Determine whether given value is part of the expression potentially being
266 /// emitted.
267 bool isPartOfCurrentExpression(Value value) {
268 Operation *def = value.getDefiningOp();
269 return def ? isPartOfCurrentExpression(def) : false;
270 }
271
272 /// Determine whether given operation is part of the expression potentially
273 /// being emitted.
274 bool isPartOfCurrentExpression(Operation *def) {
275 return isEmittingExpression() && shouldBeInlined(def);
276 };
277
278 // Resets the value counter to 0.
279 void resetValueCounter();
280
281 // Increases the loop nesting level by 1.
282 void increaseLoopNestingLevel();
283
284 // Decreases the loop nesting level by 1.
285 void decreaseLoopNestingLevel();
286
287private:
288 using ValueMapper = llvm::ScopedHashTable<Value, std::string>;
289 using BlockMapper = llvm::ScopedHashTable<Block *, std::string>;
290
291 /// Output stream to emit to.
292 raw_indented_ostream os;
293
294 /// Boolean to enforce that all variables for op results and block
295 /// arguments are declared at the beginning of the function. This also
296 /// includes results from ops located in nested regions.
297 bool declareVariablesAtTop;
298
299 /// Only emit file ops whos id matches this value.
300 std::string fileId;
301
302 /// Map from value to name of C++ variable that contain the name.
303 ValueMapper valueMapper;
304
305 /// Map from block to name of C++ label.
306 BlockMapper blockMapper;
307
308 /// Default values representing outermost scope.
309 llvm::ScopedHashTableScope<Value, std::string> defaultValueMapperScope;
310 llvm::ScopedHashTableScope<Block *, std::string> defaultBlockMapperScope;
311
312 std::stack<int64_t> labelInScopeCount;
313
314 /// Keeps track of the amount of nested loops the emitter currently operates
315 /// in.
316 uint64_t loopNestingLevel{0};
317
318 /// Emitter-level count of created values to enable unique identifiers.
319 unsigned int valueCount{0};
320
321 /// State of the current expression being emitted.
322 SmallVector<int> emittedExpressionPrecedence;
323
324 void pushExpressionPrecedence(int precedence) {
325 emittedExpressionPrecedence.push_back(precedence);
326 }
327 void popExpressionPrecedence() { emittedExpressionPrecedence.pop_back(); }
328 static int lowestPrecedence() { return 0; }
329 int getExpressionPrecedence() {
330 if (emittedExpressionPrecedence.empty())
331 return lowestPrecedence();
332 return emittedExpressionPrecedence.back();
333 }
334};
335} // namespace
336
337/// Determine whether operation \p op should be emitted inline, i.e.
338/// as part of its user. This function recommends inlining of any expressions
339/// that can be inlined unless it is used by another expression, under the
340/// assumption that any expression fusion/re-materialization was taken care of
341/// by transformations run by the backend.
342static bool shouldBeInlined(Operation *op) {
343 // CExpression operations are inlined if and only if they are marked as
344 // always-inline or reside in an ExpressionOp.
345 if (auto cExpression = dyn_cast<CExpressionInterface>(op))
346 return cExpression.alwaysInline() || isa<ExpressionOp>(op->getParentOp());
347
348 // Only other inlinable operation is ExpressionOp itself.
349 ExpressionOp expressionOp = dyn_cast<ExpressionOp>(op);
350 if (!expressionOp)
351 return false;
352
353 // Inline if the root operation is an always-inline CExpression.
354 if (cast<CExpressionInterface>(expressionOp.getRootOp()).alwaysInline())
355 return true;
356
357 // Do not inline if expression is marked as such.
358 if (expressionOp.getDoNotInline())
359 return false;
360
361 // Do not inline expressions with multiple uses.
362 Value result = expressionOp.getResult();
363 if (!result.hasOneUse())
364 return false;
365
366 Operation *user = *result.getUsers().begin();
367
368 // Do not inline expressions used by other expressions or by ops with the
369 // CExpressionInterface. If this was intended, the user could have been merged
370 // into the expression op.
371 if (isa<emitc::ExpressionOp, emitc::CExpressionInterface>(*user))
372 return false;
373
374 // Expressions with no side-effects can safely be inlined.
375 if (!expressionOp.hasSideEffects())
376 return true;
377
378 // Expressions with side-effects can be only inlined if side-effect ordering
379 // in the program is provably retained.
380
381 // Require the user to immediately follow the expression.
382 if (++Block::iterator(expressionOp) != Block::iterator(user))
383 return false;
384
385 // These single-operand ops are safe.
386 if (isa<emitc::IfOp, emitc::SwitchOp, emitc::ReturnOp>(user))
387 return true;
388
389 // For assignment look for specific cases to inline as evaluation order of
390 // its lvalue and rvalue is undefined in C.
391 if (auto assignOp = dyn_cast<emitc::AssignOp>(user)) {
392 // Inline if this assignment is of the form `<var> = <expression>`.
393 if (expressionOp.getResult() == assignOp.getValue() &&
394 isa_and_present<VariableOp>(assignOp.getVar().getDefiningOp()))
395 return true;
396 }
397
398 return false;
399}
400
401/// Helper function to check if a value traces back to a const global.
402/// Handles direct GetGlobalOp and GetGlobalOp through one or more SubscriptOps.
403/// Returns the GlobalOp if found and it has const_specifier, nullptr otherwise.
404static emitc::GlobalOp getConstGlobal(Value value, Operation *fromOp) {
405 while (auto subscriptOp = value.getDefiningOp<emitc::SubscriptOp>()) {
406 value = subscriptOp.getValue();
407 }
408
409 auto getGlobalOp = value.getDefiningOp<emitc::GetGlobalOp>();
410 if (!getGlobalOp)
411 return nullptr;
412
413 // Find the nearest symbol table to check whether the global is const.
415 fromOp, getGlobalOp.getNameAttr());
416
417 if (globalOp && globalOp.getConstSpecifier())
418 return globalOp;
419
420 return nullptr;
421}
422
423/// Emit address-of with a cast to strip const qualification.
424/// Produces: (ResultType)(&operand)
425static LogicalResult emitAddressOfWithConstCast(CppEmitter &emitter,
426 Operation &op, Value operand) {
427 raw_ostream &os = emitter.ostream();
428 os << "(";
429 if (failed(emitter.emitType(op.getLoc(), op.getResult(0).getType())))
430 return failure();
431 os << ")(&";
432 if (failed(emitter.emitOperand(operand)))
433 return failure();
434 os << ")";
435 return success();
436}
437
438static LogicalResult printOperation(CppEmitter &emitter,
439 emitc::DereferenceOp dereferenceOp) {
440 raw_ostream &os = emitter.ostream();
441 Operation &op = *dereferenceOp.getOperation();
442
443 if (failed(emitter.emitAssignPrefix(op)))
444 return failure();
445 os << "*";
446 return emitter.emitOperand(dereferenceOp.getPointer());
447}
448
449static LogicalResult printOperation(CppEmitter &emitter,
450 emitc::GetFieldOp getFieldOp) {
451 if (!emitter.isPartOfCurrentExpression(getFieldOp.getOperation()))
452 return success();
453
454 emitter.ostream() << getFieldOp.getFieldName();
455 return success();
456}
457
458static LogicalResult printOperation(CppEmitter &emitter,
459 emitc::GetGlobalOp getGlobalOp) {
460 if (!emitter.isPartOfCurrentExpression(getGlobalOp.getOperation()))
461 return success();
462
463 emitter.ostream() << getGlobalOp.getName();
464 return success();
465}
466
467static LogicalResult printOperation(CppEmitter &emitter,
468 emitc::LiteralOp literalOp) {
469 if (!emitter.isPartOfCurrentExpression(literalOp.getOperation()))
470 return success();
471
472 emitter.ostream() << literalOp.getValue();
473 return success();
474}
475
476static LogicalResult printOperation(CppEmitter &emitter,
477 emitc::MemberOp memberOp) {
478 if (memberOp.alwaysInline()) {
479 if (!emitter.isPartOfCurrentExpression(memberOp.getOperation()))
480 return success();
481 } else {
482 if (failed(emitter.emitAssignPrefix(*memberOp.getOperation())))
483 return failure();
484 }
485 if (failed(emitter.emitOperand(memberOp.getOperand())))
486 return failure();
487 emitter.ostream() << "." << memberOp.getMember();
488 return success();
489}
490
491static LogicalResult printOperation(CppEmitter &emitter,
492 emitc::MemberOfPtrOp memberOfPtrOp) {
493 if (!emitter.isPartOfCurrentExpression(memberOfPtrOp.getOperation()))
494 return success();
495
496 if (failed(emitter.emitOperand(memberOfPtrOp.getOperand())))
497 return failure();
498 emitter.ostream() << "->" << memberOfPtrOp.getMember();
499 return success();
500}
501
502static LogicalResult printOperation(CppEmitter &emitter,
503 emitc::SubscriptOp subscriptOp) {
504 if (!emitter.isPartOfCurrentExpression(subscriptOp.getOperation())) {
505 return success();
506 }
507
508 raw_ostream &os = emitter.ostream();
509 if (failed(emitter.emitOperand(subscriptOp.getValue())))
510 return failure();
511 for (auto index : subscriptOp.getIndices()) {
512 os << "[";
513 if (failed(emitter.emitOperand(index, /*isInBrackets=*/true)))
514 return failure();
515 os << "]";
516 }
517 return success();
518}
519
520static LogicalResult printConstantOp(CppEmitter &emitter, Operation *operation,
521 Attribute value) {
522 OpResult result = operation->getResult(0);
523
524 // Only emit an assignment as the variable was already declared when printing
525 // the FuncOp.
526 if (emitter.shouldDeclareVariablesAtTop()) {
527 // Skip the assignment if the emitc.constant has no value.
528 if (auto oAttr = dyn_cast<emitc::OpaqueAttr>(value)) {
529 if (oAttr.getValue().empty())
530 return success();
531 }
532
533 if (failed(emitter.emitVariableAssignment(result)))
534 return failure();
535 return emitter.emitAttribute(operation->getLoc(), value);
536 }
537
538 // Emit a variable declaration for an emitc.constant op without value.
539 if (auto oAttr = dyn_cast<emitc::OpaqueAttr>(value)) {
540 if (oAttr.getValue().empty())
541 // The semicolon gets printed by the emitOperation function.
542 return emitter.emitVariableDeclaration(result,
543 /*trailingSemicolon=*/false);
544 }
545
546 // Emit a variable declaration.
547 if (failed(emitter.emitAssignPrefix(*operation)))
548 return failure();
549 return emitter.emitAttribute(operation->getLoc(), value);
550}
551
552static LogicalResult printOperation(CppEmitter &emitter,
553 emitc::AddressOfOp addressOfOp) {
554 raw_ostream &os = emitter.ostream();
555 Operation &op = *addressOfOp.getOperation();
556
557 if (failed(emitter.emitAssignPrefix(op)))
558 return failure();
559
560 Value operand = addressOfOp.getReference();
561
562 // Check if we're taking address of a const global.
563 if (getConstGlobal(operand, &op))
564 return emitAddressOfWithConstCast(emitter, op, operand);
565
566 os << "&";
567 return emitter.emitOperand(operand);
568}
569
570static LogicalResult printOperation(CppEmitter &emitter,
571 emitc::ConstantOp constantOp) {
572 Operation *operation = constantOp.getOperation();
573 Attribute value = constantOp.getValue();
574
575 if (emitter.isPartOfCurrentExpression(operation))
576 return emitter.emitAttribute(operation->getLoc(), value);
577
578 return printConstantOp(emitter, operation, value);
579}
580
581static LogicalResult printOperation(CppEmitter &emitter,
582 emitc::VariableOp variableOp) {
583 Operation *operation = variableOp.getOperation();
584 Attribute value = variableOp.getValue();
585
586 return printConstantOp(emitter, operation, value);
587}
588
589static LogicalResult printOperation(CppEmitter &emitter,
590 emitc::GlobalOp globalOp) {
591
592 return emitter.emitGlobalVariable(globalOp);
593}
594
595static LogicalResult printOperation(CppEmitter &emitter,
596 emitc::AssignOp assignOp) {
597 if (failed(emitter.emitOperand(assignOp.getVar())))
598 return failure();
599
600 emitter.ostream() << " = ";
601
602 return emitter.emitOperand(assignOp.getValue());
603}
604
605static LogicalResult
606printCompoundAssignmentOperation(CppEmitter &emitter, Operation *operation,
607 StringRef compoundAssignmentOperator) {
608 if (failed(emitter.emitOperand(operation->getOperand(0))))
609 return failure();
610
611 emitter.ostream() << " " << compoundAssignmentOperator << " ";
612
613 return emitter.emitOperand(operation->getOperand(1));
614}
615
616static LogicalResult printOperation(CppEmitter &emitter,
617 emitc::AddAssignOp addAssignOp) {
618 return printCompoundAssignmentOperation(emitter, addAssignOp, "+=");
619}
620
621static LogicalResult printOperation(CppEmitter &emitter,
622 emitc::SubAssignOp subAssignOp) {
623 return printCompoundAssignmentOperation(emitter, subAssignOp, "-=");
624}
625
626static LogicalResult printOperation(CppEmitter &emitter,
627 emitc::MulAssignOp mulAssignOp) {
628 return printCompoundAssignmentOperation(emitter, mulAssignOp, "*=");
629}
630
631static LogicalResult printOperation(CppEmitter &emitter,
632 emitc::DivAssignOp divAssignOp) {
633 return printCompoundAssignmentOperation(emitter, divAssignOp, "/=");
634}
635
636static LogicalResult printOperation(CppEmitter &emitter,
637 emitc::RemAssignOp remAssignOp) {
638 return printCompoundAssignmentOperation(emitter, remAssignOp, "%=");
639}
640
641static LogicalResult printOperation(CppEmitter &emitter, emitc::LoadOp loadOp) {
642 if (failed(emitter.emitAssignPrefix(*loadOp)))
643 return failure();
644
645 return emitter.emitOperand(loadOp.getOperand());
646}
647
648static LogicalResult printBinaryOperation(CppEmitter &emitter,
649 Operation *operation,
650 StringRef binaryOperator) {
651 raw_ostream &os = emitter.ostream();
652
653 if (failed(emitter.emitAssignPrefix(*operation)))
654 return failure();
655
656 if (failed(emitter.emitOperand(operation->getOperand(0))))
657 return failure();
658
659 os << " " << binaryOperator << " ";
660
661 if (failed(emitter.emitOperand(operation->getOperand(1))))
662 return failure();
663
664 return success();
665}
666
667static LogicalResult printUnaryOperation(CppEmitter &emitter,
668 Operation *operation,
669 StringRef unaryOperator) {
670 raw_ostream &os = emitter.ostream();
671
672 if (failed(emitter.emitAssignPrefix(*operation)))
673 return failure();
674
675 os << unaryOperator;
676
677 if (failed(emitter.emitOperand(operation->getOperand(0))))
678 return failure();
679
680 return success();
681}
682
683static LogicalResult printPostfixUnaryOperation(CppEmitter &emitter,
684 Operation *operation,
685 StringRef unaryOperator) {
686 raw_ostream &os = emitter.ostream();
687
688 if (failed(emitter.emitAssignPrefix(*operation)))
689 return failure();
690
691 if (failed(emitter.emitOperand(operation->getOperand(0))))
692 return failure();
693
694 os << unaryOperator;
695
696 return success();
697}
698
699static LogicalResult printOperation(CppEmitter &emitter, emitc::AddOp addOp) {
700 Operation *operation = addOp.getOperation();
701
702 return printBinaryOperation(emitter, operation, "+");
703}
704
705static LogicalResult printOperation(CppEmitter &emitter, emitc::DivOp divOp) {
706 Operation *operation = divOp.getOperation();
707
708 return printBinaryOperation(emitter, operation, "/");
709}
710
711static LogicalResult printOperation(CppEmitter &emitter, emitc::MulOp mulOp) {
712 Operation *operation = mulOp.getOperation();
713
714 return printBinaryOperation(emitter, operation, "*");
715}
716
717static LogicalResult printOperation(CppEmitter &emitter, emitc::RemOp remOp) {
718 Operation *operation = remOp.getOperation();
719
720 return printBinaryOperation(emitter, operation, "%");
721}
722
723static LogicalResult printOperation(CppEmitter &emitter, emitc::SubOp subOp) {
724 Operation *operation = subOp.getOperation();
725
726 return printBinaryOperation(emitter, operation, "-");
727}
728
729static LogicalResult emitSwitchCase(CppEmitter &emitter,
730 raw_indented_ostream &os, Region &region) {
731 for (Region::OpIterator iteratorOp = region.op_begin(), end = region.op_end();
732 std::next(iteratorOp) != end; ++iteratorOp) {
733 if (failed(emitter.emitOperation(*iteratorOp, /*trailingSemicolon=*/true)))
734 return failure();
735 }
736 os << "break;\n";
737 return success();
738}
739
740static LogicalResult printOperation(CppEmitter &emitter,
741 emitc::SwitchOp switchOp) {
742 raw_indented_ostream &os = emitter.ostream();
743
744 os << "switch (";
745 if (failed(emitter.emitOperand(switchOp.getArg())))
746 return failure();
747 os << ") {";
748
749 for (auto pair : llvm::zip(switchOp.getCases(), switchOp.getCaseRegions())) {
750 os << "\ncase " << std::get<0>(pair) << ": {\n";
751 os.indent();
752
753 if (failed(emitSwitchCase(emitter, os, std::get<1>(pair))))
754 return failure();
755
756 os.unindent() << "}";
757 }
758
759 os << "\ndefault: {\n";
760 os.indent();
761
762 if (failed(emitSwitchCase(emitter, os, switchOp.getDefaultRegion())))
763 return failure();
764
765 os.unindent() << "}\n}";
766 return success();
767}
768
769static LogicalResult printOperation(CppEmitter &emitter, emitc::DoOp doOp) {
770 raw_indented_ostream &os = emitter.ostream();
771
772 os << "do {\n";
773 os.indent();
774
775 Block &bodyBlock = doOp.getBodyRegion().front();
776 for (Operation &op : bodyBlock) {
777 if (failed(emitter.emitOperation(op, /*trailingSemicolon=*/true)))
778 return failure();
779 }
780
781 os.unindent() << "} while (";
782
783 Block &condBlock = doOp.getConditionRegion().front();
784 auto condYield = cast<emitc::YieldOp>(condBlock.back());
785 if (failed(emitter.emitExpression(
786 cast<emitc::ExpressionOp>(condYield.getOperand(0).getDefiningOp()))))
787 return failure();
788
789 os << ");";
790 return success();
791}
792
793static LogicalResult printOperation(CppEmitter &emitter, emitc::CmpOp cmpOp) {
794 Operation *operation = cmpOp.getOperation();
795
796 StringRef binaryOperator;
797
798 switch (cmpOp.getPredicate()) {
799 case emitc::CmpPredicate::eq:
800 binaryOperator = "==";
801 break;
802 case emitc::CmpPredicate::ne:
803 binaryOperator = "!=";
804 break;
805 case emitc::CmpPredicate::lt:
806 binaryOperator = "<";
807 break;
808 case emitc::CmpPredicate::le:
809 binaryOperator = "<=";
810 break;
811 case emitc::CmpPredicate::gt:
812 binaryOperator = ">";
813 break;
814 case emitc::CmpPredicate::ge:
815 binaryOperator = ">=";
816 break;
817 case emitc::CmpPredicate::three_way:
818 binaryOperator = "<=>";
819 break;
820 }
821
822 return printBinaryOperation(emitter, operation, binaryOperator);
823}
824
825static LogicalResult printOperation(CppEmitter &emitter,
826 emitc::ConditionalOp conditionalOp) {
827 raw_ostream &os = emitter.ostream();
828
829 if (failed(emitter.emitAssignPrefix(*conditionalOp)))
830 return failure();
831
832 if (failed(emitter.emitOperand(conditionalOp.getCondition())))
833 return failure();
834
835 os << " ? ";
836
837 if (failed(emitter.emitOperand(conditionalOp.getTrueValue())))
838 return failure();
839
840 os << " : ";
841
842 if (failed(emitter.emitOperand(conditionalOp.getFalseValue())))
843 return failure();
844
845 return success();
846}
847
848static LogicalResult printOperation(CppEmitter &emitter,
849 emitc::VerbatimOp verbatimOp) {
850 raw_ostream &os = emitter.ostream();
851
852 FailureOr<SmallVector<ReplacementItem>> items =
853 verbatimOp.parseFormatString();
854 if (failed(items))
855 return failure();
856
857 auto fmtArg = verbatimOp.getFmtArgs().begin();
858
859 for (ReplacementItem &item : *items) {
860 if (auto *str = std::get_if<StringRef>(&item)) {
861 os << *str;
862 } else {
863 if (failed(emitter.emitOperand(*fmtArg++)))
864 return failure();
865 }
866 }
867
868 return success();
869}
870
871static LogicalResult printOperation(CppEmitter &emitter,
872 cf::BranchOp branchOp) {
873 raw_ostream &os = emitter.ostream();
874 Block &successor = *branchOp.getSuccessor();
875
876 for (auto pair :
877 llvm::zip(branchOp.getOperands(), successor.getArguments())) {
878 Value &operand = std::get<0>(pair);
879 BlockArgument &argument = std::get<1>(pair);
880 os << emitter.getOrCreateName(argument) << " = "
881 << emitter.getOrCreateName(operand) << ";\n";
882 }
883
884 os << "goto ";
885 if (!(emitter.hasBlockLabel(successor)))
886 return branchOp.emitOpError("unable to find label for successor block");
887 os << emitter.getOrCreateName(successor);
888 return success();
889}
890
891static LogicalResult printOperation(CppEmitter &emitter,
892 cf::CondBranchOp condBranchOp) {
893 raw_indented_ostream &os = emitter.ostream();
894 Block &trueSuccessor = *condBranchOp.getTrueDest();
895 Block &falseSuccessor = *condBranchOp.getFalseDest();
896
897 os << "if (";
898 if (failed(emitter.emitOperand(condBranchOp.getCondition())))
899 return failure();
900 os << ") {\n";
901
902 os.indent();
903
904 // If condition is true.
905 for (auto pair : llvm::zip(condBranchOp.getTrueOperands(),
906 trueSuccessor.getArguments())) {
907 Value &operand = std::get<0>(pair);
908 BlockArgument &argument = std::get<1>(pair);
909 os << emitter.getOrCreateName(argument) << " = "
910 << emitter.getOrCreateName(operand) << ";\n";
911 }
912
913 os << "goto ";
914 if (!(emitter.hasBlockLabel(trueSuccessor))) {
915 return condBranchOp.emitOpError("unable to find label for successor block");
916 }
917 os << emitter.getOrCreateName(trueSuccessor) << ";\n";
918 os.unindent() << "} else {\n";
919 os.indent();
920 // If condition is false.
921 for (auto pair : llvm::zip(condBranchOp.getFalseOperands(),
922 falseSuccessor.getArguments())) {
923 Value &operand = std::get<0>(pair);
924 BlockArgument &argument = std::get<1>(pair);
925 os << emitter.getOrCreateName(argument) << " = "
926 << emitter.getOrCreateName(operand) << ";\n";
927 }
928
929 os << "goto ";
930 if (!(emitter.hasBlockLabel(falseSuccessor))) {
931 return condBranchOp.emitOpError()
932 << "unable to find label for successor block";
933 }
934 os << emitter.getOrCreateName(falseSuccessor) << ";\n";
935 os.unindent() << "}";
936 return success();
937}
938
939static LogicalResult printCallOperation(CppEmitter &emitter, Operation *callOp,
940 StringRef callee) {
941 if (failed(emitter.emitAssignPrefix(*callOp)))
942 return failure();
943
944 raw_ostream &os = emitter.ostream();
945 os << callee << "(";
946 if (failed(emitter.emitOperands(*callOp)))
947 return failure();
948 os << ")";
949 return success();
950}
951
952static LogicalResult printOperation(CppEmitter &emitter, func::CallOp callOp) {
953 Operation *operation = callOp.getOperation();
954 StringRef callee = callOp.getCallee();
955
956 return printCallOperation(emitter, operation, callee);
957}
958
959static LogicalResult printOperation(CppEmitter &emitter, emitc::CallOp callOp) {
960 Operation *operation = callOp.getOperation();
961 StringRef callee = callOp.getCallee();
962
963 return printCallOperation(emitter, operation, callee);
964}
965
966template <typename OpTy>
967static LogicalResult
968printOpaqueCallCommon(CppEmitter &emitter, OpTy op, StringRef callee,
969 std::optional<ArrayAttr> templateArgs,
970 std::optional<ArrayAttr> args, bool isMemberCall,
971 Value receiver = nullptr) {
972 raw_ostream &os = emitter.ostream();
973
974 if (failed(emitter.emitAssignPrefix(*op.getOperation())))
975 return failure();
976
977 if (isMemberCall) {
978 assert(receiver && "Expected receiver for member call");
979 if (failed(emitter.emitOperand(receiver)))
980 return failure();
981
982 if (llvm::isa<emitc::PointerType>(receiver.getType()))
983 os << "->";
984 else
985 os << ".";
986 }
987
988 os << callee;
989
990 // Template arguments can't refer to SSA values and as such the template
991 // arguments which are supplied in form of attributes can be emitted as is. We
992 // don't need to handle integer attributes specially like we do for arguments
993 // - see below.
994 auto emitTemplateArgs = [&](Attribute attr) -> LogicalResult {
995 return emitter.emitAttribute(op.getLoc(), attr);
996 };
997
998 if (templateArgs) {
999 os << "<";
1000 if (failed(interleaveCommaWithError(*templateArgs, os, emitTemplateArgs)))
1001 return failure();
1002 os << ">";
1003 }
1004
1005 auto emitArgs = [&](Attribute attr) -> LogicalResult {
1006 if (auto t = dyn_cast<IntegerAttr>(attr)) {
1007 if (t.getType().isIndex()) {
1008 int64_t idx = t.getInt();
1009 Value operand = op.getArgOperands()[idx];
1010 return emitter.emitOperand(operand, /*isInBrackets=*/false);
1011 }
1012 }
1013 if (failed(emitter.emitAttribute(op.getLoc(), attr)))
1014 return failure();
1015
1016 return success();
1017 };
1018
1019 os << "(";
1020
1021 LogicalResult emittedArgs = success();
1022 if (args) {
1023 emittedArgs = interleaveCommaWithError(*args, os, emitArgs);
1024 } else {
1025 emittedArgs =
1026 interleaveCommaWithError(op.getArgOperands(), os, [&](Value operand) {
1027 return emitter.emitOperand(operand, /*isInBrackets=*/true);
1028 });
1029 }
1030 if (failed(emittedArgs))
1031 return failure();
1032 os << ")";
1033 return success();
1034}
1035
1036static LogicalResult printOperation(CppEmitter &emitter,
1037 emitc::CallOpaqueOp callOpaqueOp) {
1038 return printOpaqueCallCommon(emitter, callOpaqueOp, callOpaqueOp.getCallee(),
1039 callOpaqueOp.getTemplateArgs(),
1040 callOpaqueOp.getArgs(),
1041 /*isMemberCall=*/false);
1042}
1043
1044static LogicalResult
1045printOperation(CppEmitter &emitter,
1046 emitc::MemberCallOpaqueOp memberCallOpaqueOp) {
1047 return printOpaqueCallCommon(
1048 emitter, memberCallOpaqueOp, memberCallOpaqueOp.getCallee(),
1049 memberCallOpaqueOp.getTemplateArgs(), memberCallOpaqueOp.getArgs(),
1050 /*isMemberCall=*/true, memberCallOpaqueOp.getReceiver());
1051}
1052
1053static LogicalResult printOperation(CppEmitter &emitter,
1054 emitc::BitwiseAndOp bitwiseAndOp) {
1055 Operation *operation = bitwiseAndOp.getOperation();
1056 return printBinaryOperation(emitter, operation, "&");
1057}
1058
1059static LogicalResult
1060printOperation(CppEmitter &emitter,
1061 emitc::BitwiseLeftShiftOp bitwiseLeftShiftOp) {
1062 Operation *operation = bitwiseLeftShiftOp.getOperation();
1063 return printBinaryOperation(emitter, operation, "<<");
1064}
1065
1066static LogicalResult printOperation(CppEmitter &emitter,
1067 emitc::BitwiseNotOp bitwiseNotOp) {
1068 Operation *operation = bitwiseNotOp.getOperation();
1069 return printUnaryOperation(emitter, operation, "~");
1070}
1071
1072static LogicalResult printOperation(CppEmitter &emitter,
1073 emitc::BitwiseOrOp bitwiseOrOp) {
1074 Operation *operation = bitwiseOrOp.getOperation();
1075 return printBinaryOperation(emitter, operation, "|");
1076}
1077
1078static LogicalResult
1079printOperation(CppEmitter &emitter,
1080 emitc::BitwiseRightShiftOp bitwiseRightShiftOp) {
1081 Operation *operation = bitwiseRightShiftOp.getOperation();
1082 return printBinaryOperation(emitter, operation, ">>");
1083}
1084
1085static LogicalResult printOperation(CppEmitter &emitter,
1086 emitc::BitwiseXorOp bitwiseXorOp) {
1087 Operation *operation = bitwiseXorOp.getOperation();
1088 return printBinaryOperation(emitter, operation, "^");
1089}
1090
1091static LogicalResult printOperation(CppEmitter &emitter,
1092 emitc::PreIncrementOp preIncrementOp) {
1093 Operation *operation = preIncrementOp.getOperation();
1094 return printUnaryOperation(emitter, operation, "++");
1095}
1096
1097static LogicalResult printOperation(CppEmitter &emitter,
1098 emitc::PostIncrementOp postIncrementOp) {
1099 Operation *operation = postIncrementOp.getOperation();
1100 return printPostfixUnaryOperation(emitter, operation, "++");
1101}
1102
1103static LogicalResult printOperation(CppEmitter &emitter,
1104 emitc::PreDecrementOp preDecrementOp) {
1105 Operation *operation = preDecrementOp.getOperation();
1106 return printUnaryOperation(emitter, operation, "--");
1107}
1108
1109static LogicalResult printOperation(CppEmitter &emitter,
1110 emitc::PostDecrementOp postDecrementOp) {
1111 Operation *operation = postDecrementOp.getOperation();
1112 return printPostfixUnaryOperation(emitter, operation, "--");
1113}
1114
1115static LogicalResult printOperation(CppEmitter &emitter,
1116 emitc::UnaryPlusOp unaryPlusOp) {
1117 Operation *operation = unaryPlusOp.getOperation();
1118 return printUnaryOperation(emitter, operation, "+");
1119}
1120
1121static LogicalResult printOperation(CppEmitter &emitter,
1122 emitc::UnaryMinusOp unaryMinusOp) {
1123 Operation *operation = unaryMinusOp.getOperation();
1124 return printUnaryOperation(emitter, operation, "-");
1125}
1126
1127static LogicalResult printOperation(CppEmitter &emitter, emitc::CastOp castOp) {
1128 raw_ostream &os = emitter.ostream();
1129 Operation &op = *castOp.getOperation();
1130
1131 if (failed(emitter.emitAssignPrefix(op)))
1132 return failure();
1133 os << "(";
1134 if (failed(emitter.emitType(op.getLoc(), op.getResult(0).getType())))
1135 return failure();
1136 os << ") ";
1137 return emitter.emitOperand(castOp.getOperand());
1138}
1139
1140static LogicalResult printOperation(CppEmitter &emitter,
1141 emitc::ExpressionOp expressionOp) {
1142 if (shouldBeInlined(expressionOp))
1143 return success();
1144
1145 Operation &op = *expressionOp.getOperation();
1146
1147 if (failed(emitter.emitAssignPrefix(op)))
1148 return failure();
1149
1150 return emitter.emitExpression(expressionOp);
1151}
1152
1153static LogicalResult printOperation(CppEmitter &emitter,
1154 emitc::IncludeOp includeOp) {
1155 raw_ostream &os = emitter.ostream();
1156
1157 os << "#include ";
1158 if (includeOp.getIsStandardInclude())
1159 os << "<" << includeOp.getInclude() << ">";
1160 else
1161 os << "\"" << includeOp.getInclude() << "\"";
1162
1163 return success();
1164}
1165
1166static LogicalResult printOperation(CppEmitter &emitter,
1167 emitc::LogicalAndOp logicalAndOp) {
1168 Operation *operation = logicalAndOp.getOperation();
1169 return printBinaryOperation(emitter, operation, "&&");
1170}
1171
1172static LogicalResult printOperation(CppEmitter &emitter,
1173 emitc::LogicalNotOp logicalNotOp) {
1174 Operation *operation = logicalNotOp.getOperation();
1175 return printUnaryOperation(emitter, operation, "!");
1176}
1177
1178static LogicalResult printOperation(CppEmitter &emitter,
1179 emitc::LogicalOrOp logicalOrOp) {
1180 Operation *operation = logicalOrOp.getOperation();
1181 return printBinaryOperation(emitter, operation, "||");
1182}
1183
1184static LogicalResult printOperation(CppEmitter &emitter, emitc::ForOp forOp) {
1185 raw_indented_ostream &os = emitter.ostream();
1186
1187 // Utility function to determine whether a value is an expression that will be
1188 // inlined, and as such should be wrapped in parentheses in order to guarantee
1189 // its precedence and associativity.
1190 auto requiresParentheses = [&](Value value) {
1191 auto expressionOp = value.getDefiningOp<ExpressionOp>();
1192 if (!expressionOp)
1193 return false;
1194 return shouldBeInlined(expressionOp);
1195 };
1196
1197 os << "for (";
1198 if (failed(
1199 emitter.emitType(forOp.getLoc(), forOp.getInductionVar().getType())))
1200 return failure();
1201 os << " ";
1202 os << emitter.getOrCreateInductionVarName(forOp.getInductionVar());
1203 os << " = ";
1204 if (failed(emitter.emitOperand(forOp.getLowerBound())))
1205 return failure();
1206 os << "; ";
1207 os << emitter.getOrCreateInductionVarName(forOp.getInductionVar());
1208 os << " < ";
1209 Value upperBound = forOp.getUpperBound();
1210 bool upperBoundRequiresParentheses = requiresParentheses(upperBound);
1211 if (upperBoundRequiresParentheses)
1212 os << "(";
1213 if (failed(emitter.emitOperand(upperBound)))
1214 return failure();
1215 if (upperBoundRequiresParentheses)
1216 os << ")";
1217 os << "; ";
1218 os << emitter.getOrCreateInductionVarName(forOp.getInductionVar());
1219 os << " += ";
1220 if (failed(emitter.emitOperand(forOp.getStep())))
1221 return failure();
1222 os << ") {\n";
1223 os.indent();
1224
1225 CppEmitter::LoopScope lScope(emitter);
1226
1227 Region &forRegion = forOp.getRegion();
1228 auto regionOps = forRegion.getOps();
1229
1230 // We skip the trailing yield op.
1231 for (auto it = regionOps.begin(); std::next(it) != regionOps.end(); ++it) {
1232 if (failed(emitter.emitOperation(*it, /*trailingSemicolon=*/true)))
1233 return failure();
1234 }
1235
1236 os.unindent() << "}";
1237
1238 return success();
1239}
1240
1241static LogicalResult printOperation(CppEmitter &emitter, emitc::IfOp ifOp) {
1242 raw_indented_ostream &os = emitter.ostream();
1243
1244 // Helper function to emit all ops except the last one, expected to be
1245 // emitc::yield.
1246 auto emitAllExceptLast = [&emitter](Region &region) {
1247 Region::OpIterator it = region.op_begin(), end = region.op_end();
1248 for (; std::next(it) != end; ++it) {
1249 if (failed(emitter.emitOperation(*it, /*trailingSemicolon=*/true)))
1250 return failure();
1251 }
1252 assert(isa<emitc::YieldOp>(*it) &&
1253 "Expected last operation in the region to be emitc::yield");
1254 return success();
1255 };
1256
1257 os << "if (";
1258 if (failed(emitter.emitOperand(ifOp.getCondition())))
1259 return failure();
1260 os << ") {\n";
1261 os.indent();
1262 if (failed(emitAllExceptLast(ifOp.getThenRegion())))
1263 return failure();
1264 os.unindent() << "}";
1265
1266 Region &elseRegion = ifOp.getElseRegion();
1267 if (!elseRegion.empty()) {
1268 os << " else {\n";
1269 os.indent();
1270 if (failed(emitAllExceptLast(elseRegion)))
1271 return failure();
1272 os.unindent() << "}";
1273 }
1274
1275 return success();
1276}
1277
1278static LogicalResult printOperation(CppEmitter &emitter,
1279 func::ReturnOp returnOp) {
1280 raw_ostream &os = emitter.ostream();
1281 os << "return";
1282 switch (returnOp.getNumOperands()) {
1283 case 0:
1284 return success();
1285 case 1:
1286 os << " ";
1287 if (failed(emitter.emitOperand(returnOp.getOperand(0))))
1288 return failure();
1289 return success();
1290 default:
1291 os << " std::make_tuple(";
1292 if (failed(emitter.emitOperandsAndAttributes(*returnOp.getOperation())))
1293 return failure();
1294 os << ")";
1295 return success();
1296 }
1297}
1298
1299static LogicalResult printOperation(CppEmitter &emitter,
1300 emitc::ReturnOp returnOp) {
1301 raw_ostream &os = emitter.ostream();
1302 os << "return";
1303 if (returnOp.getNumOperands() == 0)
1304 return success();
1305
1306 os << " ";
1307 if (failed(emitter.emitOperand(returnOp.getOperand())))
1308 return failure();
1309 return success();
1310}
1311
1312static LogicalResult printOperation(CppEmitter &emitter, ModuleOp moduleOp) {
1313 for (Operation &op : moduleOp) {
1314 if (failed(emitter.emitOperation(op, /*trailingSemicolon=*/false)))
1315 return failure();
1316 }
1317 return success();
1318}
1319
1320static LogicalResult printOperation(CppEmitter &emitter, ClassOp classOp) {
1321 raw_indented_ostream &os = emitter.ostream();
1322 ClassType classType = classOp.getClassType();
1323 os << stringifyClassType(classType) << " " << classOp.getSymName();
1324 if (classOp.getFinalSpecifier())
1325 os << " final";
1326 os << " {\n";
1327
1328 if (classType == ClassType::class_)
1329 os << " public:\n";
1330
1331 os.indent();
1332
1333 for (Operation &op : classOp) {
1334 if (failed(emitter.emitOperation(op, /*trailingSemicolon=*/false)))
1335 return failure();
1336 }
1337
1338 os.unindent();
1339 os << "};";
1340 return success();
1341}
1342
1343static LogicalResult printOperation(CppEmitter &emitter, FieldOp fieldOp) {
1344 raw_ostream &os = emitter.ostream();
1345 if (failed(emitter.emitVariableDeclaration(
1346 fieldOp->getLoc(), fieldOp.getType(), fieldOp.getSymName())))
1347 return failure();
1348 std::optional<Attribute> initialValue = fieldOp.getInitialValue();
1349 if (initialValue) {
1350 os << " = ";
1351 if (failed(emitter.emitAttribute(fieldOp->getLoc(), *initialValue)))
1352 return failure();
1353 }
1354
1355 os << ";";
1356 return success();
1357}
1358
1359static LogicalResult printOperation(CppEmitter &emitter, FileOp file) {
1360 if (!emitter.shouldEmitFile(file))
1361 return success();
1362
1363 for (Operation &op : file) {
1364 if (failed(emitter.emitOperation(op, /*trailingSemicolon=*/false)))
1365 return failure();
1366 }
1367 return success();
1368}
1369
1370static LogicalResult printFunctionArgs(CppEmitter &emitter,
1371 Operation *functionOp,
1372 ArrayRef<Type> arguments) {
1373 raw_indented_ostream &os = emitter.ostream();
1374
1375 return (
1376 interleaveCommaWithError(arguments, os, [&](Type arg) -> LogicalResult {
1377 return emitter.emitType(functionOp->getLoc(), arg);
1378 }));
1379}
1380
1381static LogicalResult printFunctionArgs(CppEmitter &emitter,
1382 Operation *functionOp,
1383 Region::BlockArgListType arguments) {
1384 raw_indented_ostream &os = emitter.ostream();
1385
1387 arguments, os, [&](BlockArgument arg) -> LogicalResult {
1388 return emitter.emitVariableDeclaration(
1389 functionOp->getLoc(), arg.getType(), emitter.getOrCreateName(arg));
1390 }));
1391}
1392
1393static LogicalResult printFunctionBody(CppEmitter &emitter,
1394 Operation *functionOp,
1395 Region::BlockListType &blocks) {
1396 raw_indented_ostream &os = emitter.ostream();
1397 os.indent();
1398
1399 if (emitter.shouldDeclareVariablesAtTop()) {
1400 // Declare all variables that hold op results including those from nested
1401 // regions.
1403 functionOp->walk<WalkOrder::PreOrder>([&](Operation *op) -> WalkResult {
1404 if (isa<emitc::ExpressionOp>(op->getParentOp()) ||
1405 (isa<emitc::ExpressionOp>(op) &&
1406 shouldBeInlined(cast<emitc::ExpressionOp>(op))))
1407 return WalkResult::skip();
1408 for (OpResult result : op->getResults()) {
1409 if (failed(emitter.emitVariableDeclaration(
1410 result, /*trailingSemicolon=*/true))) {
1411 return WalkResult(
1412 op->emitError("unable to declare result variable for op"));
1413 }
1414 }
1415 return WalkResult::advance();
1416 });
1417 if (result.wasInterrupted())
1418 return failure();
1419 }
1420
1421 // Create label names for basic blocks.
1422 for (Block &block : blocks) {
1423 emitter.getOrCreateName(block);
1424 }
1425
1426 // Declare variables for basic block arguments.
1427 for (Block &block : llvm::drop_begin(blocks)) {
1428 for (BlockArgument &arg : block.getArguments()) {
1429 if (emitter.hasValueInScope(arg))
1430 return functionOp->emitOpError(" block argument #")
1431 << arg.getArgNumber() << " is out of scope";
1432 if (isa<ArrayType, LValueType>(arg.getType()))
1433 return functionOp->emitOpError("cannot emit block argument #")
1434 << arg.getArgNumber() << " with type " << arg.getType();
1435 if (failed(
1436 emitter.emitType(block.getParentOp()->getLoc(), arg.getType()))) {
1437 return failure();
1438 }
1439 os << " " << emitter.getOrCreateName(arg) << ";\n";
1440 }
1441 }
1442
1443 for (Block &block : blocks) {
1444 // Only print a label if the block has predecessors.
1445 if (!block.hasNoPredecessors()) {
1446 if (failed(emitter.emitLabel(block)))
1447 return failure();
1448 }
1449 for (Operation &op : block.getOperations()) {
1450 if (failed(emitter.emitOperation(op, /*trailingSemicolon=*/true)))
1451 return failure();
1452 }
1453 }
1454
1455 os.unindent();
1456
1457 return success();
1458}
1459
1460static LogicalResult printOperation(CppEmitter &emitter,
1461 func::FuncOp functionOp) {
1462 // We need to declare variables at top if the function has multiple blocks.
1463 if (!emitter.shouldDeclareVariablesAtTop() &&
1464 functionOp.getBlocks().size() > 1) {
1465 return functionOp.emitOpError(
1466 "with multiple blocks needs variables declared at top");
1467 }
1468
1469 if (llvm::any_of(functionOp.getArgumentTypes(), llvm::IsaPred<LValueType>)) {
1470 return functionOp.emitOpError()
1471 << "cannot emit lvalue type as argument type";
1472 }
1473
1474 if (llvm::any_of(functionOp.getResultTypes(), llvm::IsaPred<ArrayType>)) {
1475 return functionOp.emitOpError() << "cannot emit array type as result type";
1476 }
1477
1478 CppEmitter::FunctionScope scope(emitter);
1479 raw_indented_ostream &os = emitter.ostream();
1480 if (failed(emitter.emitTypes(functionOp.getLoc(),
1481 functionOp.getFunctionType().getResults())))
1482 return failure();
1483 os << " " << functionOp.getName();
1484
1485 os << "(";
1486 Operation *operation = functionOp.getOperation();
1487 if (failed(printFunctionArgs(emitter, operation, functionOp.getArguments())))
1488 return failure();
1489 os << ") {\n";
1490 if (failed(printFunctionBody(emitter, operation, functionOp.getBlocks())))
1491 return failure();
1492 os << "}";
1493
1494 return success();
1495}
1496
1497static LogicalResult printOperation(CppEmitter &emitter,
1498 emitc::FuncOp functionOp) {
1499 // We need to declare variables at top if the function has multiple blocks.
1500 if (!emitter.shouldDeclareVariablesAtTop() &&
1501 functionOp.getBlocks().size() > 1) {
1502 return functionOp.emitOpError(
1503 "with multiple blocks needs variables declared at top");
1504 }
1505
1506 CppEmitter::FunctionScope scope(emitter);
1507 raw_indented_ostream &os = emitter.ostream();
1508 if (functionOp.getSpecifiers()) {
1509 for (Attribute specifier : functionOp.getSpecifiersAttr()) {
1510 os << cast<StringAttr>(specifier).str() << " ";
1511 }
1512 }
1513
1514 if (failed(emitter.emitTypes(functionOp.getLoc(),
1515 functionOp.getFunctionType().getResults())))
1516 return failure();
1517 os << " " << functionOp.getName();
1518
1519 os << "(";
1520 Operation *operation = functionOp.getOperation();
1521 if (functionOp.isExternal()) {
1522 if (failed(printFunctionArgs(emitter, operation,
1523 functionOp.getArgumentTypes())))
1524 return failure();
1525 os << ");";
1526 return success();
1527 }
1528 if (failed(printFunctionArgs(emitter, operation, functionOp.getArguments())))
1529 return failure();
1530 os << ") {\n";
1531 if (failed(printFunctionBody(emitter, operation, functionOp.getBlocks())))
1532 return failure();
1533 os << "}";
1534
1535 return success();
1536}
1537
1538static LogicalResult printOperation(CppEmitter &emitter,
1539 DeclareFuncOp declareFuncOp) {
1540 raw_indented_ostream &os = emitter.ostream();
1541
1542 CppEmitter::FunctionScope scope(emitter);
1544 declareFuncOp, declareFuncOp.getSymNameAttr());
1545
1546 if (!functionOp)
1547 return failure();
1548
1549 if (functionOp.getSpecifiers()) {
1550 for (Attribute specifier : functionOp.getSpecifiersAttr()) {
1551 os << cast<StringAttr>(specifier).str() << " ";
1552 }
1553 }
1554
1555 if (failed(emitter.emitTypes(functionOp.getLoc(),
1556 functionOp.getFunctionType().getResults())))
1557 return failure();
1558 os << " " << functionOp.getName();
1559
1560 os << "(";
1561 Operation *operation = functionOp.getOperation();
1562 if (failed(printFunctionArgs(emitter, operation, functionOp.getArguments())))
1563 return failure();
1564 os << ");";
1565
1566 return success();
1567}
1568
1569CppEmitter::CppEmitter(raw_ostream &os, bool declareVariablesAtTop,
1570 StringRef fileId)
1571 : os(os), declareVariablesAtTop(declareVariablesAtTop),
1572 fileId(fileId.str()), defaultValueMapperScope(valueMapper),
1573 defaultBlockMapperScope(blockMapper) {
1574 labelInScopeCount.push(0);
1575}
1576
1577/// Return the existing or a new name for a Value.
1578StringRef CppEmitter::getOrCreateName(Value val) {
1579 if (!valueMapper.count(val)) {
1580 valueMapper.insert(val, formatv("v{0}", ++valueCount));
1581 }
1582 return *valueMapper.begin(val);
1583}
1584
1585/// Return the existing or a new name for a loop induction variable Value.
1586/// Loop induction variables follow natural naming: i, j, k, ..., t, uX.
1587StringRef CppEmitter::getOrCreateInductionVarName(Value val) {
1588 if (!valueMapper.count(val)) {
1589
1590 int64_t identifier = 'i' + loopNestingLevel;
1591
1592 if (identifier >= 'i' && identifier <= 't') {
1593 valueMapper.insert(val,
1594 formatv("{0}{1}", (char)identifier, ++valueCount));
1595 } else {
1596 // If running out of letters, continue with uX.
1597 valueMapper.insert(val, formatv("u{0}", ++valueCount));
1598 }
1599 }
1600 return *valueMapper.begin(val);
1601}
1602
1603/// Return the existing or a new label for a Block.
1604StringRef CppEmitter::getOrCreateName(Block &block) {
1605 if (!blockMapper.count(&block))
1606 blockMapper.insert(&block, formatv("label{0}", ++labelInScopeCount.top()));
1607 return *blockMapper.begin(&block);
1608}
1609
1610bool CppEmitter::shouldMapToUnsigned(IntegerType::SignednessSemantics val) {
1611 switch (val) {
1612 case IntegerType::Signless:
1613 return false;
1614 case IntegerType::Signed:
1615 return false;
1616 case IntegerType::Unsigned:
1617 return true;
1618 }
1619 llvm_unreachable("Unexpected IntegerType::SignednessSemantics");
1620}
1621
1622bool CppEmitter::hasValueInScope(Value val) { return valueMapper.count(val); }
1623
1624bool CppEmitter::hasBlockLabel(Block &block) {
1625 return blockMapper.count(&block);
1626}
1627
1628LogicalResult CppEmitter::emitAttribute(Location loc, Attribute attr) {
1629 auto printInt = [&](const APInt &val, bool isUnsigned) {
1630 if (val.getBitWidth() == 1) {
1631 if (val.getBoolValue())
1632 os << "true";
1633 else
1634 os << "false";
1635 } else {
1636 SmallString<128> strValue;
1637 val.toString(strValue, 10, !isUnsigned, false);
1638 os << strValue;
1639 }
1640 };
1641
1642 auto printFloat = [&](const APFloat &val) {
1643 if (val.isFinite()) {
1644 SmallString<128> strValue;
1645 // Use default values of toString except don't truncate zeros.
1646 val.toString(strValue, 0, 0, false);
1647 os << strValue;
1648 switch (llvm::APFloatBase::SemanticsToEnum(val.getSemantics())) {
1649 case llvm::APFloatBase::S_IEEEhalf:
1650 os << "f16";
1651 break;
1652 case llvm::APFloatBase::S_BFloat:
1653 os << "bf16";
1654 break;
1655 case llvm::APFloatBase::S_IEEEsingle:
1656 os << "f";
1657 break;
1658 case llvm::APFloatBase::S_IEEEdouble:
1659 break;
1660 default:
1661 llvm_unreachable("unsupported floating point type");
1662 };
1663 } else if (val.isNaN()) {
1664 os << "NAN";
1665 } else if (val.isInfinity()) {
1666 if (val.isNegative())
1667 os << "-";
1668 os << "INFINITY";
1669 }
1670 };
1671
1672 // Print floating point attributes.
1673 if (auto fAttr = dyn_cast<FloatAttr>(attr)) {
1674 if (!isa<Float16Type, BFloat16Type, Float32Type, Float64Type>(
1675 fAttr.getType())) {
1676 return emitError(
1677 loc, "expected floating point attribute to be f16, bf16, f32 or f64");
1678 }
1679 printFloat(fAttr.getValue());
1680 return success();
1681 }
1682 if (auto dense = dyn_cast<DenseFPElementsAttr>(attr)) {
1683 if (!isa<Float16Type, BFloat16Type, Float32Type, Float64Type>(
1684 dense.getElementType())) {
1685 return emitError(
1686 loc, "expected floating point attribute to be f16, bf16, f32 or f64");
1687 }
1688 os << '{';
1689 interleaveComma(dense, os, [&](const APFloat &val) { printFloat(val); });
1690 os << '}';
1691 return success();
1692 }
1693
1694 // Print integer attributes.
1695 if (auto iAttr = dyn_cast<IntegerAttr>(attr)) {
1696 if (auto iType = dyn_cast<IntegerType>(iAttr.getType())) {
1697 printInt(iAttr.getValue(), shouldMapToUnsigned(iType.getSignedness()));
1698 return success();
1699 }
1700 if (auto iType = dyn_cast<IndexType>(iAttr.getType())) {
1701 printInt(iAttr.getValue(), false);
1702 return success();
1703 }
1704 }
1705 if (auto dense = dyn_cast<DenseIntElementsAttr>(attr)) {
1706 if (auto iType = dyn_cast<IntegerType>(
1707 cast<ShapedType>(dense.getType()).getElementType())) {
1708 os << '{';
1709 interleaveComma(dense, os, [&](const APInt &val) {
1710 printInt(val, shouldMapToUnsigned(iType.getSignedness()));
1711 });
1712 os << '}';
1713 return success();
1714 }
1715 if (auto iType = dyn_cast<IndexType>(
1716 cast<ShapedType>(dense.getType()).getElementType())) {
1717 os << '{';
1718 interleaveComma(dense, os,
1719 [&](const APInt &val) { printInt(val, false); });
1720 os << '}';
1721 return success();
1722 }
1723 }
1724
1725 // Print opaque attributes.
1726 if (auto oAttr = dyn_cast<emitc::OpaqueAttr>(attr)) {
1727 os << oAttr.getValue();
1728 return success();
1729 }
1730
1731 // Print symbolic reference attributes.
1732 if (auto sAttr = dyn_cast<SymbolRefAttr>(attr)) {
1733 if (sAttr.getNestedReferences().size() > 1)
1734 return emitError(loc, "attribute has more than 1 nested reference");
1735 os << sAttr.getRootReference().getValue();
1736 return success();
1737 }
1738
1739 // Print type attributes.
1740 if (auto type = dyn_cast<TypeAttr>(attr))
1741 return emitType(loc, type.getValue());
1742
1743 return emitError(loc, "cannot emit attribute: ") << attr;
1744}
1745
1746LogicalResult CppEmitter::emitExpression(Operation *op) {
1747 assert(emittedExpressionPrecedence.empty() &&
1748 "Expected precedence stack to be empty");
1749 Operation *rootOp = nullptr;
1750
1751 if (auto expressionOp = dyn_cast<ExpressionOp>(op)) {
1752 rootOp = expressionOp.getRootOp();
1753 } else {
1754 assert(cast<CExpressionInterface>(op).alwaysInline() &&
1755 "Expected an always-inline operation");
1756 assert(!isa<ExpressionOp>(op->getParentOp()) &&
1757 "Expected operation to have no containing expression");
1758 rootOp = op;
1759 }
1760 FailureOr<int> precedence = getOperatorPrecedence(rootOp);
1761 if (failed(precedence))
1762 return failure();
1763 pushExpressionPrecedence(precedence.value());
1764
1765 if (failed(emitOperation(*rootOp, /*trailingSemicolon=*/false)))
1766 return failure();
1767
1768 popExpressionPrecedence();
1769 assert(emittedExpressionPrecedence.empty() &&
1770 "Expected precedence stack to be empty");
1771
1772 return success();
1773}
1774
1775LogicalResult CppEmitter::emitOperand(Value value, bool isInBrackets) {
1776 if (isPartOfCurrentExpression(value)) {
1777 Operation *def = value.getDefiningOp();
1778 assert(def && "Expected operand to be defined by an operation");
1779 if (auto expressionOp = dyn_cast<ExpressionOp>(def))
1780 def = expressionOp.getRootOp();
1781 // Within an expression, `emitc.load` emits only its operand, so determine
1782 // parentheses from the operand rather than from the load operation.
1783 if (auto loadOp = dyn_cast<emitc::LoadOp>(def))
1784 return emitOperand(loadOp.getOperand(), isInBrackets);
1785 FailureOr<int> precedence = getOperatorPrecedence(def);
1786 if (failed(precedence))
1787 return failure();
1788
1789 // Unless already in brackets, sub-expressions with equal or lower
1790 // precedence need to be parenthesized as they might be evaluated in the
1791 // wrong order depending on the shape of the expression tree.
1792 bool encloseInParenthesis =
1793 !isInBrackets && precedence.value() <= getExpressionPrecedence();
1794
1795 if (encloseInParenthesis)
1796 os << "(";
1797 pushExpressionPrecedence(precedence.value());
1798
1799 if (failed(emitOperation(*def, /*trailingSemicolon=*/false)))
1800 return failure();
1801
1802 if (encloseInParenthesis)
1803 os << ")";
1804
1805 popExpressionPrecedence();
1806 return success();
1807 }
1808
1809 if (Operation *def = value.getDefiningOp(); def && shouldBeInlined(def))
1810 return emitExpression(def);
1811
1812 if (BlockArgument arg = dyn_cast<BlockArgument>(value)) {
1813 // If this operand is a block argument of an expression, emit instead the
1814 // matching expression parameter.
1815 Operation *argOp = arg.getParentBlock()->getParentOp();
1816 if (auto expressionOp = dyn_cast<ExpressionOp>(argOp))
1817 return emitOperand(expressionOp->getOperand(arg.getArgNumber()));
1818 }
1819
1820 os << getOrCreateName(value);
1821 return success();
1822}
1823
1824LogicalResult CppEmitter::emitOperands(Operation &op) {
1825 return interleaveCommaWithError(op.getOperands(), os, [&](Value operand) {
1826 // Emit operand under guarantee that if it's part of an expression then it
1827 // is being emitted within brackets.
1828 return emitOperand(operand, /*isInBrackets=*/true);
1829 });
1830}
1831
1832LogicalResult
1833CppEmitter::emitOperandsAndAttributes(Operation &op,
1834 ArrayRef<StringRef> exclude) {
1835 if (failed(emitOperands(op)))
1836 return failure();
1837 auto attrs = op.getDiscardableAttrs();
1838 // Insert comma in between operands and non-filtered attributes if needed.
1839 if (op.getNumOperands() > 0) {
1840 for (NamedAttribute attr : attrs) {
1841 if (!llvm::is_contained(exclude, attr.getName().strref())) {
1842 os << ", ";
1843 break;
1844 }
1845 }
1846 }
1847 // Emit attributes.
1848 auto emitNamedAttribute = [&](NamedAttribute attr) -> LogicalResult {
1849 if (llvm::is_contained(exclude, attr.getName().strref()))
1850 return success();
1851 os << "/* " << attr.getName().getValue() << " */";
1852 if (failed(emitAttribute(op.getLoc(), attr.getValue())))
1853 return failure();
1854 return success();
1855 };
1856 return interleaveCommaWithError(attrs, os, emitNamedAttribute);
1857}
1858
1859LogicalResult CppEmitter::emitVariableAssignment(OpResult result) {
1860 if (!hasValueInScope(result)) {
1861 return result.getDefiningOp()->emitOpError(
1862 "result variable for the operation has not been declared");
1863 }
1864 os << getOrCreateName(result) << " = ";
1865 return success();
1866}
1867
1868LogicalResult CppEmitter::emitVariableDeclaration(OpResult result,
1869 bool trailingSemicolon) {
1870 if (auto cExpression =
1871 dyn_cast<CExpressionInterface>(result.getDefiningOp())) {
1872 if (cExpression.alwaysInline())
1873 return success();
1874 }
1875 if (hasValueInScope(result)) {
1876 return result.getDefiningOp()->emitError(
1877 "result variable for the operation already declared");
1878 }
1879 if (failed(emitVariableDeclaration(result.getOwner()->getLoc(),
1880 result.getType(),
1881 getOrCreateName(result))))
1882 return failure();
1883 if (trailingSemicolon)
1884 os << ";\n";
1885 return success();
1886}
1887
1888LogicalResult CppEmitter::emitGlobalVariable(GlobalOp op) {
1889 if (op.getExternSpecifier())
1890 os << "extern ";
1891 else if (op.getStaticSpecifier())
1892 os << "static ";
1893 if (op.getConstSpecifier())
1894 os << "const ";
1895
1896 if (failed(emitVariableDeclaration(op->getLoc(), op.getType(),
1897 op.getSymName()))) {
1898 return failure();
1899 }
1900
1901 std::optional<Attribute> initialValue = op.getInitialValue();
1902 if (initialValue) {
1903 os << " = ";
1904 if (failed(emitAttribute(op->getLoc(), *initialValue)))
1905 return failure();
1906 }
1907
1908 os << ";";
1909 return success();
1910}
1911
1912LogicalResult CppEmitter::emitAssignPrefix(Operation &op) {
1913 // If op is being emitted as part of an expression, bail out.
1914 if (isEmittingExpression())
1915 return success();
1916
1917 switch (op.getNumResults()) {
1918 case 0:
1919 break;
1920 case 1: {
1921 OpResult result = op.getResult(0);
1922 if (shouldDeclareVariablesAtTop()) {
1923 if (failed(emitVariableAssignment(result)))
1924 return failure();
1925 } else {
1926 if (failed(emitVariableDeclaration(result, /*trailingSemicolon=*/false)))
1927 return failure();
1928 os << " = ";
1929 }
1930 break;
1931 }
1932 default:
1933 if (!shouldDeclareVariablesAtTop()) {
1934 for (OpResult result : op.getResults()) {
1935 if (failed(emitVariableDeclaration(result, /*trailingSemicolon=*/true)))
1936 return failure();
1937 }
1938 }
1939 os << "std::tie(";
1940 interleaveComma(op.getResults(), os,
1941 [&](Value result) { os << getOrCreateName(result); });
1942 os << ") = ";
1943 }
1944 return success();
1945}
1946
1947LogicalResult CppEmitter::emitLabel(Block &block) {
1948 if (!hasBlockLabel(block))
1949 return block.getParentOp()->emitError("label for block not found");
1950 // FIXME: Add feature in `raw_indented_ostream` to ignore indent for block
1951 // label instead of using `getOStream`.
1952 os.getOStream() << getOrCreateName(block) << ":\n";
1953 return success();
1954}
1955
1956LogicalResult CppEmitter::emitOperation(Operation &op, bool trailingSemicolon) {
1957 LogicalResult status =
1958 llvm::TypeSwitch<Operation *, LogicalResult>(&op)
1959 // Builtin ops.
1960 .Case([&](ModuleOp op) { return printOperation(*this, op); })
1961 // CF ops.
1962 .Case<cf::BranchOp, cf::CondBranchOp>(
1963 [&](auto op) { return printOperation(*this, op); })
1964 // EmitC ops.
1965 .Case<emitc::AddAssignOp, emitc::AddressOfOp, emitc::AddOp,
1966 emitc::AssignOp, emitc::BitwiseAndOp, emitc::BitwiseLeftShiftOp,
1967 emitc::BitwiseNotOp, emitc::BitwiseOrOp,
1968 emitc::BitwiseRightShiftOp, emitc::BitwiseXorOp, emitc::CallOp,
1969 emitc::CallOpaqueOp, emitc::CastOp, emitc::ClassOp,
1970 emitc::CmpOp, emitc::ConditionalOp, emitc::ConstantOp,
1971 emitc::DeclareFuncOp, emitc::DereferenceOp, emitc::DivAssignOp,
1972 emitc::DivOp, emitc::DoOp, emitc::ExpressionOp, emitc::FieldOp,
1973 emitc::FileOp, emitc::ForOp, emitc::FuncOp, emitc::GetFieldOp,
1974 emitc::GetGlobalOp, emitc::GlobalOp, emitc::IfOp,
1975 emitc::IncludeOp, emitc::LiteralOp, emitc::LoadOp,
1976 emitc::LogicalAndOp, emitc::LogicalNotOp, emitc::LogicalOrOp,
1977 emitc::MemberCallOpaqueOp, emitc::MemberOfPtrOp,
1978 emitc::MemberOp, emitc::MulAssignOp, emitc::MulOp,
1979 emitc::PostDecrementOp, emitc::PostIncrementOp,
1980 emitc::PreDecrementOp, emitc::PreIncrementOp,
1981 emitc::RemAssignOp, emitc::RemOp, emitc::ReturnOp,
1982 emitc::SubAssignOp, emitc::SubscriptOp, emitc::SubOp,
1983 emitc::SwitchOp, emitc::UnaryMinusOp, emitc::UnaryPlusOp,
1984 emitc::VariableOp, emitc::VerbatimOp>(
1985
1986 [&](auto op) { return printOperation(*this, op); })
1987 // Func ops.
1988 .Case<func::CallOp, func::FuncOp, func::ReturnOp>(
1989 [&](auto op) { return printOperation(*this, op); })
1990 .Default([&](Operation *) {
1991 return op.emitOpError("unable to find printer for op");
1992 });
1993
1994 if (failed(status))
1995 return failure();
1996
1997 if (auto cExpression = dyn_cast<CExpressionInterface>(op)) {
1998 if (cExpression.alwaysInline())
1999 return success();
2000 }
2001
2002 if (isEmittingExpression() ||
2003 (isa<emitc::ExpressionOp>(op) &&
2004 shouldBeInlined(cast<emitc::ExpressionOp>(op))))
2005 return success();
2006
2007 // Never emit a semicolon for some operations, especially if endening with
2008 // `}`.
2009 trailingSemicolon &=
2010 !isa<cf::CondBranchOp, emitc::DeclareFuncOp, emitc::DoOp, emitc::FileOp,
2011 emitc::ForOp, emitc::IfOp, emitc::IncludeOp, emitc::SwitchOp,
2012 emitc::VerbatimOp>(op);
2013
2014 os << (trailingSemicolon ? ";\n" : "\n");
2015
2016 return success();
2017}
2018
2019LogicalResult CppEmitter::emitVariableDeclaration(Location loc, Type type,
2020 StringRef name) {
2021 if (auto arrType = dyn_cast<emitc::ArrayType>(type)) {
2022 if (failed(emitType(loc, arrType.getElementType())))
2023 return failure();
2024 os << " " << name;
2025 for (auto dim : arrType.getShape()) {
2026 os << "[" << dim << "]";
2027 }
2028 return success();
2029 }
2030 if (failed(emitType(loc, type)))
2031 return failure();
2032 os << " " << name;
2033 return success();
2034}
2035
2036LogicalResult CppEmitter::emitType(Location loc, Type type) {
2037 if (auto iType = dyn_cast<IntegerType>(type)) {
2038 switch (iType.getWidth()) {
2039 case 1:
2040 return (os << "bool"), success();
2041 case 8:
2042 case 16:
2043 case 32:
2044 case 64:
2045 if (shouldMapToUnsigned(iType.getSignedness()))
2046 return (os << "uint" << iType.getWidth() << "_t"), success();
2047 else
2048 return (os << "int" << iType.getWidth() << "_t"), success();
2049 default:
2050 return emitError(loc, "cannot emit integer type ") << type;
2051 }
2052 }
2053 if (auto fType = dyn_cast<FloatType>(type)) {
2054 switch (fType.getWidth()) {
2055 case 16: {
2056 if (llvm::isa<Float16Type>(type))
2057 return (os << "_Float16"), success();
2058 if (llvm::isa<BFloat16Type>(type))
2059 return (os << "__bf16"), success();
2060 else
2061 return emitError(loc, "cannot emit float type ") << type;
2062 }
2063 case 32:
2064 return (os << "float"), success();
2065 case 64:
2066 return (os << "double"), success();
2067 default:
2068 return emitError(loc, "cannot emit float type ") << type;
2069 }
2070 }
2071 if (auto iType = dyn_cast<IndexType>(type))
2072 return (os << "size_t"), success();
2073 if (auto sType = dyn_cast<emitc::SizeTType>(type))
2074 return (os << "size_t"), success();
2075 if (auto sType = dyn_cast<emitc::SignedSizeTType>(type))
2076 return (os << "ssize_t"), success();
2077 if (auto pType = dyn_cast<emitc::PtrDiffTType>(type))
2078 return (os << "ptrdiff_t"), success();
2079 if (auto tType = dyn_cast<TensorType>(type)) {
2080 if (!tType.hasRank())
2081 return emitError(loc, "cannot emit unranked tensor type");
2082 if (!tType.hasStaticShape())
2083 return emitError(loc, "cannot emit tensor type with non static shape");
2084 os << "Tensor<";
2085 if (isa<ArrayType>(tType.getElementType()))
2086 return emitError(loc, "cannot emit tensor of array type ") << type;
2087 if (failed(emitType(loc, tType.getElementType())))
2088 return failure();
2089 auto shape = tType.getShape();
2090 for (auto dimSize : shape) {
2091 os << ", ";
2092 os << dimSize;
2093 }
2094 os << ">";
2095 return success();
2096 }
2097 if (auto tType = dyn_cast<TupleType>(type))
2098 return emitTupleType(loc, tType.getTypes());
2099 if (auto oType = dyn_cast<emitc::OpaqueType>(type)) {
2100 os << oType.getValue();
2101 return success();
2102 }
2103 if (auto aType = dyn_cast<emitc::ArrayType>(type)) {
2104 if (failed(emitType(loc, aType.getElementType())))
2105 return failure();
2106 for (auto dim : aType.getShape())
2107 os << "[" << dim << "]";
2108 return success();
2109 }
2110 if (auto lType = dyn_cast<emitc::LValueType>(type))
2111 return emitType(loc, lType.getValueType());
2112 if (auto pType = dyn_cast<emitc::PointerType>(type)) {
2113 if (isa<ArrayType>(pType.getPointee()))
2114 return emitError(loc, "cannot emit pointer to array type ") << type;
2115 if (failed(emitType(loc, pType.getPointee())))
2116 return failure();
2117 os << "*";
2118 return success();
2119 }
2120 return emitError(loc, "cannot emit type ") << type;
2121}
2122
2123LogicalResult CppEmitter::emitTypes(Location loc, ArrayRef<Type> types) {
2124 switch (types.size()) {
2125 case 0:
2126 os << "void";
2127 return success();
2128 case 1:
2129 return emitType(loc, types.front());
2130 default:
2131 return emitTupleType(loc, types);
2132 }
2133}
2134
2135LogicalResult CppEmitter::emitTupleType(Location loc, ArrayRef<Type> types) {
2136 if (llvm::any_of(types, llvm::IsaPred<ArrayType>)) {
2137 return emitError(loc, "cannot emit tuple of array type");
2138 }
2139 os << "std::tuple<";
2141 types, os, [&](Type type) { return emitType(loc, type); })))
2142 return failure();
2143 os << ">";
2144 return success();
2145}
2146
2147void CppEmitter::resetValueCounter() { valueCount = 0; }
2148
2149void CppEmitter::increaseLoopNestingLevel() { loopNestingLevel++; }
2150
2151void CppEmitter::decreaseLoopNestingLevel() { loopNestingLevel--; }
2152
2154 bool declareVariablesAtTop,
2155 StringRef fileId) {
2156 CppEmitter emitter(os, declareVariablesAtTop, fileId);
2157 return emitter.emitOperation(*op, /*trailingSemicolon=*/false);
2158}
return success()
false
Parses a map_entries map type from a string format back into its numeric value.
static LogicalResult printCallOperation(CppEmitter &emitter, Operation *callOp, StringRef callee)
static FailureOr< int > getOperatorPrecedence(Operation *operation)
Return the precedence of a operator as an integer, higher values imply higher precedence.
static LogicalResult printFunctionArgs(CppEmitter &emitter, Operation *functionOp, ArrayRef< Type > arguments)
static LogicalResult printCompoundAssignmentOperation(CppEmitter &emitter, Operation *operation, StringRef compoundAssignmentOperator)
static LogicalResult printFunctionBody(CppEmitter &emitter, Operation *functionOp, Region::BlockListType &blocks)
static LogicalResult printConstantOp(CppEmitter &emitter, Operation *operation, Attribute value)
static LogicalResult emitSwitchCase(CppEmitter &emitter, raw_indented_ostream &os, Region &region)
static LogicalResult interleaveCommaWithError(const Container &c, raw_ostream &os, UnaryFunctor eachFn)
static LogicalResult printBinaryOperation(CppEmitter &emitter, Operation *operation, StringRef binaryOperator)
static bool shouldBeInlined(Operation *op)
Determine whether operation op should be emitted inline, i.e.
static LogicalResult printOperation(CppEmitter &emitter, emitc::DereferenceOp dereferenceOp)
static LogicalResult printUnaryOperation(CppEmitter &emitter, Operation *operation, StringRef unaryOperator)
static LogicalResult emitAddressOfWithConstCast(CppEmitter &emitter, Operation &op, Value operand)
Emit address-of with a cast to strip const qualification.
static LogicalResult printPostfixUnaryOperation(CppEmitter &emitter, Operation *operation, StringRef unaryOperator)
static LogicalResult interleaveWithError(ForwardIterator begin, ForwardIterator end, UnaryFunctor eachFn, NullaryFunctor betweenFn)
Convenience functions to produce interleaved output with functions returning a LogicalResult.
static LogicalResult printOpaqueCallCommon(CppEmitter &emitter, OpTy op, StringRef callee, std::optional< ArrayAttr > templateArgs, std::optional< ArrayAttr > args, bool isMemberCall, Value receiver=nullptr)
static emitc::GlobalOp getConstGlobal(Value value, Operation *fromOp)
Helper function to check if a value traces back to a const global.
Attributes are known-constant values of operations.
Definition Attributes.h:25
This class represents an argument of a Block.
Definition Value.h:306
Block represents an ordered list of Operations.
Definition Block.h:34
OpListType::iterator iterator
Definition Block.h:165
Operation & front()
Definition Block.h:178
Operation & back()
Definition Block.h:177
BlockArgListType getArguments()
Definition Block.h:112
Block * getSuccessor(unsigned i)
Definition Block.cpp:274
Operation * getParentOp()
Returns the closest surrounding operation that contains this block.
Definition Block.cpp:31
This is a value defined by a result of an operation.
Definition Value.h:454
Operation is the basic unit of execution within MLIR.
Definition Operation.h:87
Value getOperand(unsigned idx)
Definition Operation.h:375
OpResult getResult(unsigned idx)
Get the 'idx'th result of this operation.
Definition Operation.h:432
Location getLoc()
The source location the operation was defined or derived from.
Definition Operation.h:240
Operation * getParentOp()
Returns the closest surrounding operation that contains this operation or nullptr if this is a top-le...
Definition Operation.h:251
unsigned getNumOperands()
Definition Operation.h:371
InFlightDiagnostic emitError(const Twine &message={})
Emit an error about fatal conditions with this operation, reporting up to any diagnostic handlers tha...
auto getDiscardableAttrs()
Return a range of all of discardable attributes on this operation.
Definition Operation.h:538
operand_range getOperands()
Returns an iterator on the underlying Value's.
Definition Operation.h:403
std::enable_if_t< llvm::function_traits< std::decay_t< FnT > >::num_args==1, RetT > walk(FnT &&callback)
Walk the operation by calling the callback for each nested operation (including this one),...
Definition Operation.h:849
result_range getResults()
Definition Operation.h:440
InFlightDiagnostic emitOpError(const Twine &message={})
Emit an error with the op name prefixed, like "'dim' op " which is convenient for verifiers.
unsigned getNumResults()
Return the number of results held by this operation.
Definition Operation.h:429
This class provides iteration over the held operations of blocks directly within a region.
Definition Region.h:147
This class contains a list of basic blocks and a link to the parent operation it is attached to.
Definition Region.h:26
llvm::iplist< Block > BlockListType
Definition Region.h:44
OpIterator op_begin()
Return iterators that walk the operations nested directly within this region.
Definition Region.h:178
iterator_range< OpIterator > getOps()
Definition Region.h:180
bool empty()
Definition Region.h:60
MutableArrayRef< BlockArgument > BlockArgListType
Definition Region.h:93
OpIterator op_end()
Definition Region.h:179
static Operation * lookupNearestSymbolFrom(Operation *from, StringAttr symbol)
Returns the operation registered with the given symbol name within the closest parent operation of,...
Instances of the Type class are uniqued, have an immutable identifier and an optional mutable compone...
Definition Types.h:74
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
Definition Value.h:96
Type getType() const
Return the type of this value.
Definition Value.h:105
Operation * getDefiningOp() const
If this value is the result of an operation, return the operation that defines it.
Definition Value.cpp:18
A utility result that is used to signal how to proceed with an ongoing walk:
Definition WalkResult.h:29
static WalkResult skip()
Definition WalkResult.h:48
static WalkResult advance()
Definition WalkResult.h:47
raw_ostream subclass that simplifies indention a sequence of code.
raw_indented_ostream & indent()
Increases the indent and returning this raw_indented_ostream.
raw_indented_ostream & unindent()
Decreases the indent and returning this raw_indented_ostream.
LogicalResult translateToCpp(Operation *op, raw_ostream &os, bool declareVariablesAtTop=false, StringRef fileId={})
Translates the given operation to C++ code.
std::variant< StringRef, Placeholder > ReplacementItem
Definition EmitC.h:54
detail::InFlightRemark failed(Location loc, RemarkOpts opts)
Report an optimization remark that failed.
Definition Remarks.h:734
Include the generated interface declarations.
InFlightDiagnostic emitError(Location loc)
Utility method to emit an error message using this location.
This iterator enumerates the elements in "forward" order.
Definition Visitors.h:31