[Mlir-commits] [mlir] 123078c - Reland emitc lower multi return functions (#203026)

llvmlistbot at llvm.org llvmlistbot at llvm.org
Thu Jun 11 01:46:22 PDT 2026


Author: Gil Rapaport
Date: 2026-06-11T11:46:17+03:00
New Revision: 123078c21cfbe4c6abe1052e53739f9e933e8c1d

URL: https://github.com/llvm/llvm-project/commit/123078c21cfbe4c6abe1052e53739f9e933e8c1d
DIFF: https://github.com/llvm/llvm-project/commit/123078c21cfbe4c6abe1052e53739f9e933e8c1d.diff

LOG: Reland emitc lower multi return functions (#203026)

Reland #200659 reverted by #202911.

Fixed GCC 7 func-to-emitc build: Use the adaptor operand types
when creating the multi-return struct type instead of relying on an
implicit conversion from ValueRange to TypeRange.

Failed buildbot:
https://lab.llvm.org/buildbot/#/builders/116/builds/29302

Assisted-by: Copilot

Added: 
    

Modified: 
    mlir/include/mlir/Conversion/ConvertToEmitC/ConvertToEmitCPatternInterface.td
    mlir/include/mlir/Conversion/ConvertToEmitC/ToEmitCInterface.h
    mlir/include/mlir/Conversion/FuncToEmitC/FuncToEmitC.h
    mlir/include/mlir/Conversion/Passes.td
    mlir/lib/Conversion/ArithToEmitC/ArithToEmitC.cpp
    mlir/lib/Conversion/ConvertToEmitC/ConvertToEmitCPass.cpp
    mlir/lib/Conversion/FuncToEmitC/FuncToEmitC.cpp
    mlir/lib/Conversion/FuncToEmitC/FuncToEmitCPass.cpp
    mlir/lib/Conversion/MemRefToEmitC/MemRefToEmitC.cpp
    mlir/lib/Conversion/SCFToEmitC/SCFToEmitC.cpp
    mlir/test/Conversion/FuncToEmitC/func-to-emitc-failed.mlir
    mlir/test/Conversion/FuncToEmitC/func-to-emitc.mlir
    mlir/test/Target/Cpp/func.mlir

Removed: 
    


################################################################################
diff  --git a/mlir/include/mlir/Conversion/ConvertToEmitC/ConvertToEmitCPatternInterface.td b/mlir/include/mlir/Conversion/ConvertToEmitC/ConvertToEmitCPatternInterface.td
index e300826cdc281..9f494097e4251 100644
--- a/mlir/include/mlir/Conversion/ConvertToEmitC/ConvertToEmitCPatternInterface.td
+++ b/mlir/include/mlir/Conversion/ConvertToEmitC/ConvertToEmitCPatternInterface.td
@@ -14,7 +14,8 @@ def ConvertToEmitCPatternInterface : DialectInterface<"ConvertToEmitCPatternInte
       }],
       "void", "populateConvertToEmitCConversionPatterns",
       (ins "::mlir::ConversionTarget &":$target, "::mlir::TypeConverter &":$typeConverter,
-           "::mlir::RewritePatternSet &":$patterns)
+           "::mlir::RewritePatternSet &":$patterns,
+           "::std::optional<bool>":$lowerToCpp)
     >
   ];
 }

diff  --git a/mlir/include/mlir/Conversion/ConvertToEmitC/ToEmitCInterface.h b/mlir/include/mlir/Conversion/ConvertToEmitC/ToEmitCInterface.h
index f6546fc27e0b3..1532f89984ea7 100644
--- a/mlir/include/mlir/Conversion/ConvertToEmitC/ToEmitCInterface.h
+++ b/mlir/include/mlir/Conversion/ConvertToEmitC/ToEmitCInterface.h
@@ -13,6 +13,8 @@
 #include "mlir/IR/MLIRContext.h"
 #include "mlir/IR/OpDefinition.h"
 
+#include <optional>
+
 namespace mlir {
 class ConversionTarget;
 class TypeConverter;

diff  --git a/mlir/include/mlir/Conversion/FuncToEmitC/FuncToEmitC.h b/mlir/include/mlir/Conversion/FuncToEmitC/FuncToEmitC.h
index be6b6cfe5a6db..5e9e09bb14d74 100644
--- a/mlir/include/mlir/Conversion/FuncToEmitC/FuncToEmitC.h
+++ b/mlir/include/mlir/Conversion/FuncToEmitC/FuncToEmitC.h
@@ -15,7 +15,8 @@ class RewritePatternSet;
 class TypeConverter;
 
 void populateFuncToEmitCPatterns(const TypeConverter &typeConverter,
-                                 RewritePatternSet &patterns);
+                                 RewritePatternSet &patterns,
+                                 bool lowerToCpp = true);
 
 void registerConvertFuncToEmitCInterface(DialectRegistry &registry);
 } // namespace mlir

diff  --git a/mlir/include/mlir/Conversion/Passes.td b/mlir/include/mlir/Conversion/Passes.td
index c30dd3b07d028..07e0e0c4e29e8 100644
--- a/mlir/include/mlir/Conversion/Passes.td
+++ b/mlir/include/mlir/Conversion/Passes.td
@@ -26,6 +26,8 @@ def ConvertToEmitC : Pass<"convert-to-emitc"> {
   let options = [
     ListOption<"filterDialects", "filter-dialects", "std::string",
                "Test conversion patterns of only the specified dialects">,
+    Option<"lowerToCpp", "lower-to-cpp", "bool", "",
+           "Target C++ (true) instead of C (false)">,
   ];
 }
 
@@ -453,6 +455,10 @@ def ConvertControlFlowToSPIRVPass : Pass<"convert-cf-to-spirv"> {
 def ConvertFuncToEmitC : Pass<"convert-func-to-emitc", "ModuleOp"> {
   let summary = "Convert Func dialect to EmitC dialect";
   let dependentDialects = ["emitc::EmitCDialect"];
+  let options = [Option<
+      "lowerToCpp", "lower-to-cpp", "bool",
+      /*default=*/"true",
+      /*description=*/"Target C++ (true) instead of C (false)">];
 }
 
 //===----------------------------------------------------------------------===//

diff  --git a/mlir/lib/Conversion/ArithToEmitC/ArithToEmitC.cpp b/mlir/lib/Conversion/ArithToEmitC/ArithToEmitC.cpp
index d003dc7a6dff3..5b074130925a4 100644
--- a/mlir/lib/Conversion/ArithToEmitC/ArithToEmitC.cpp
+++ b/mlir/lib/Conversion/ArithToEmitC/ArithToEmitC.cpp
@@ -33,7 +33,7 @@ struct ArithToEmitCDialectInterface : public ConvertToEmitCPatternInterface {
   /// and mark dialect legal for the conversion target.
   void populateConvertToEmitCConversionPatterns(
       ConversionTarget &target, TypeConverter &typeConverter,
-      RewritePatternSet &patterns) const final {
+      RewritePatternSet &patterns, std::optional<bool> lowerToCpp) const final {
     populateArithToEmitCPatterns(typeConverter, patterns);
   }
 };

diff  --git a/mlir/lib/Conversion/ConvertToEmitC/ConvertToEmitCPass.cpp b/mlir/lib/Conversion/ConvertToEmitC/ConvertToEmitCPass.cpp
index ee6d7d5e3c554..4f060eafa14cc 100644
--- a/mlir/lib/Conversion/ConvertToEmitC/ConvertToEmitCPass.cpp
+++ b/mlir/lib/Conversion/ConvertToEmitC/ConvertToEmitCPass.cpp
@@ -31,7 +31,8 @@ namespace {
 class ConvertToEmitCPassInterface {
 public:
   ConvertToEmitCPassInterface(MLIRContext *context,
-                              ArrayRef<std::string> filterDialects);
+                              ArrayRef<std::string> filterDialects,
+                              std::optional<bool> lowerToCpp);
   virtual ~ConvertToEmitCPassInterface() = default;
 
   /// Get the dependent dialects used by `convert-to-emitc`.
@@ -60,6 +61,7 @@ class ConvertToEmitCPassInterface {
   MLIRContext *context;
   /// List of dialects names to use as filters.
   ArrayRef<std::string> filterDialects;
+  std::optional<bool> lowerToCpp;
 };
 
 /// This DialectExtension can be attached to the context, which will invoke the
@@ -124,7 +126,7 @@ struct StaticConvertToEmitC : public ConvertToEmitCPassInterface {
     // Populate the patterns with the dialect interface.
     if (failed(visitInterfaces([&](ConvertToEmitCPatternInterface *iface) {
           iface->populateConvertToEmitCConversionPatterns(
-              *target, *typeConverter, tempPatterns);
+              *target, *typeConverter, tempPatterns, lowerToCpp);
         })))
       return failure();
     this->patterns =
@@ -160,7 +162,11 @@ class ConvertToEmitC : public impl::ConvertToEmitCBase<ConvertToEmitC> {
 
   LogicalResult initialize(MLIRContext *context) final {
     std::shared_ptr<ConvertToEmitCPassInterface> impl;
-    impl = std::make_shared<StaticConvertToEmitC>(context, filterDialects);
+    std::optional<bool> lowerToCppOverride;
+    if (this->lowerToCpp.hasValue())
+      lowerToCppOverride = this->lowerToCpp;
+    impl = std::make_shared<StaticConvertToEmitC>(context, filterDialects,
+                                                  lowerToCppOverride);
     if (failed(impl->initialize()))
       return failure();
     this->impl = impl;
@@ -180,8 +186,10 @@ class ConvertToEmitC : public impl::ConvertToEmitCBase<ConvertToEmitC> {
 //===----------------------------------------------------------------------===//
 
 ConvertToEmitCPassInterface::ConvertToEmitCPassInterface(
-    MLIRContext *context, ArrayRef<std::string> filterDialects)
-    : context(context), filterDialects(filterDialects) {}
+    MLIRContext *context, ArrayRef<std::string> filterDialects,
+    std::optional<bool> lowerToCpp)
+    : context(context), filterDialects(filterDialects), lowerToCpp(lowerToCpp) {
+}
 
 void ConvertToEmitCPassInterface::getDependentDialects(
     DialectRegistry &registry) {

diff  --git a/mlir/lib/Conversion/FuncToEmitC/FuncToEmitC.cpp b/mlir/lib/Conversion/FuncToEmitC/FuncToEmitC.cpp
index 4801f07d82c9f..14352ea3595e3 100644
--- a/mlir/lib/Conversion/FuncToEmitC/FuncToEmitC.cpp
+++ b/mlir/lib/Conversion/FuncToEmitC/FuncToEmitC.cpp
@@ -16,12 +16,122 @@
 #include "mlir/Conversion/ConvertToEmitC/ToEmitCInterface.h"
 #include "mlir/Dialect/EmitC/IR/EmitC.h"
 #include "mlir/Dialect/Func/IR/FuncOps.h"
+#include "mlir/IR/BuiltinAttributes.h"
+#include "mlir/IR/SymbolTable.h"
 #include "mlir/Transforms/DialectConversion.h"
+#include "llvm/ADT/StringExtras.h"
+#include "llvm/Support/LogicalResult.h"
 
 using namespace mlir;
 
 namespace {
 
+//===----------------------------------------------------------------------===//
+// Multi-return struct helpers
+//===----------------------------------------------------------------------===//
+
+// Looks up or creates an `emitc.class` named after `types` in the nearest
+// enclosing symbol table of `op`, suitable for packing those types as plain
+// struct fields (field0, field1, ...). If the class already exists it is
+// verified to have exactly the right fields and no methods. Returns the
+// corresponding !emitc.opaque<"struct ..."> type on success.
+static FailureOr<emitc::OpaqueType>
+getOrCreateMultiReturnType(ConversionPatternRewriter &rewriter, Location loc,
+                           Operation *op, TypeRange types) {
+  // Build the struct name from the types, e.g. "return_i32_i32". Each type is
+  // printed and non-alphanumeric characters are replaced with '_'.
+  std::string structName = "return";
+  for (Type type : types) {
+    std::string typeName;
+    llvm::raw_string_ostream os(typeName);
+    type.print(os);
+    std::replace_if(
+        typeName.begin(), typeName.end(),
+        [](char c) { return !llvm::isAlnum(c); }, '_');
+    structName += "_" + typeName;
+  }
+
+  // Find the enclosing symbol table and the direct child op within it that
+  // contains `op`; the class will be inserted immediately before that child.
+  Operation *symbolTableOp = SymbolTable::getNearestSymbolTable(op);
+  Operation *insertBefore = op;
+  while (insertBefore->getParentOp() != symbolTableOp)
+    insertBefore = insertBefore->getParentOp();
+
+  if (Operation *sym = SymbolTable::lookupSymbolIn(symbolTableOp, structName)) {
+    auto classOp = dyn_cast<emitc::ClassOp>(sym);
+    if (!classOp)
+      return emitError(loc) << "symbol '" << structName
+                            << "' exists but is not an emitc.class";
+
+    if (classOp.getClassType() != emitc::ClassType::struct_)
+      return emitError(loc)
+             << "existing class '" << structName << "' is not a struct";
+
+    SmallVector<emitc::FieldOp> fields;
+    for (Operation &bodyOp : classOp.getBody().front()) {
+      if (isa<emitc::FuncOp>(bodyOp))
+        return emitError(loc) << "existing class '" << structName
+                              << "' has methods; expected a plain struct";
+      if (auto fieldOp = dyn_cast<emitc::FieldOp>(bodyOp))
+        fields.push_back(fieldOp);
+    }
+    if (fields.size() != types.size())
+      return emitError(loc) << "existing class '" << structName
+                            << "' has wrong number of fields";
+    for (auto [i, fieldOp] : llvm::enumerate(fields)) {
+      if (fieldOp.getSymName() != "field" + std::to_string(i))
+        return emitError(loc) << "existing class '" << structName
+                              << "': unexpected field name at index " << i;
+      if (fieldOp.getTypeAttr().getValue() != types[i])
+        return emitError(loc) << "existing class '" << structName
+                              << "': wrong type for field " << i;
+    }
+  } else {
+    // Create the ClassOp before `insertBefore`, then restore the insertion
+    // point.
+    auto savedIP = rewriter.saveInsertionPoint();
+    rewriter.setInsertionPoint(insertBefore);
+
+    emitc::ClassOp classOp = emitc::ClassOp::create(rewriter, loc, structName,
+                                                    /*final_specifier=*/false,
+                                                    emitc::ClassType::struct_);
+    rewriter.createBlock(&classOp.getBody());
+    rewriter.setInsertionPointToStart(&classOp.getBody().front());
+
+    for (auto [i, type] : llvm::enumerate(types)) {
+      auto fieldName = rewriter.getStringAttr("field" + std::to_string(i));
+      emitc::FieldOp::create(rewriter, loc, fieldName, TypeAttr::get(type),
+                             nullptr);
+    }
+
+    rewriter.restoreInsertionPoint(savedIP);
+  }
+  return emitc::OpaqueType::get(rewriter.getContext(), "struct " + structName);
+}
+
+// Packs multiple SSA values into an emitc.class struct variable and loads the
+// result as a single SSA value of the opaque struct type.
+static Value packValuesIntoStruct(ConversionPatternRewriter &rewriter,
+                                  Location loc, ValueRange values,
+                                  emitc::OpaqueType structType) {
+  MLIRContext *ctx = rewriter.getContext();
+  auto noInit = emitc::OpaqueAttr::get(ctx, "");
+  Value structLv =
+      emitc::VariableOp::create(rewriter, loc,
+                                emitc::LValueType::get(structType), noInit)
+          .getResult();
+  for (auto [i, val] : llvm::enumerate(values)) {
+    Value fieldLv =
+        emitc::MemberOp::create(
+            rewriter, loc, emitc::LValueType::get(val.getType()),
+            rewriter.getStringAttr("field" + std::to_string(i)), structLv)
+            .getResult();
+    emitc::AssignOp::create(rewriter, loc, fieldLv, val);
+  }
+  return emitc::LoadOp::create(rewriter, loc, structType, structLv).getResult();
+}
+
 /// Implement the interface to convert Func to EmitC.
 struct FuncToEmitCDialectInterface : public ConvertToEmitCPatternInterface {
   FuncToEmitCDialectInterface(Dialect *dialect)
@@ -31,8 +141,9 @@ struct FuncToEmitCDialectInterface : public ConvertToEmitCPatternInterface {
   /// and mark dialect legal for the conversion target.
   void populateConvertToEmitCConversionPatterns(
       ConversionTarget &target, TypeConverter &typeConverter,
-      RewritePatternSet &patterns) const final {
-    populateFuncToEmitCPatterns(typeConverter, patterns);
+      RewritePatternSet &patterns, std::optional<bool> lowerToCpp) const final {
+    populateFuncToEmitCPatterns(typeConverter, patterns,
+                                lowerToCpp.value_or(true));
   }
 };
 } // namespace
@@ -50,45 +161,102 @@ void mlir::registerConvertFuncToEmitCInterface(DialectRegistry &registry) {
 namespace {
 class CallOpConversion final : public OpConversionPattern<func::CallOp> {
 public:
-  using OpConversionPattern<func::CallOp>::OpConversionPattern;
+  CallOpConversion(const TypeConverter &typeConverter, MLIRContext *ctx,
+                   bool lowerToCpp)
+      : OpConversionPattern<func::CallOp>(typeConverter, ctx),
+        lowerToCpp(lowerToCpp) {}
 
   LogicalResult
   matchAndRewrite(func::CallOp callOp, OpAdaptor adaptor,
                   ConversionPatternRewriter &rewriter) const override {
-    // Multiple results func cannot be converted to `emitc.func`.
-    if (callOp.getNumResults() > 1)
+    // Do not convert multiple-return functions if lowering target is Cpp.
+    // The translator will emit the return values as an std::tuple.
+    if (callOp.getNumResults() > 1 && lowerToCpp)
       return rewriter.notifyMatchFailure(
           callOp, "only functions with zero or one result can be converted");
 
-    if (callOp.getNumResults() == 1) {
-      Type resultType =
-          getTypeConverter()->convertType(callOp.getResult(0).getType());
+    SmallVector<Type> convertedResultTypes;
+    for (Type t : callOp.getResultTypes()) {
+      Type resultType = getTypeConverter()->convertType(t);
       if (!resultType)
         return rewriter.notifyMatchFailure(callOp,
                                            "result type conversion failed");
       if (isa<emitc::ArrayType>(resultType))
         return rewriter.notifyMatchFailure(
             callOp, "function calls returning arrays are not supported");
+      convertedResultTypes.push_back(resultType);
+    }
+
+    if (callOp.getNumResults() <= 1) {
+      rewriter.replaceOpWithNewOp<emitc::CallOp>(
+          callOp, callOp.getResultTypes(), adaptor.getOperands(),
+          callOp->getAttrs());
+      return success();
     }
 
-    rewriter.replaceOpWithNewOp<emitc::CallOp>(callOp, callOp.getResultTypes(),
-                                               adaptor.getOperands(),
-                                               callOp->getAttrs());
+    // Multi-result call: determine the struct type.
+    Location loc = callOp.getLoc();
+
+    auto structType =
+        getOrCreateMultiReturnType(rewriter, loc, callOp, convertedResultTypes);
+    if (failed(structType))
+      return rewriter.notifyMatchFailure(callOp,
+                                         "incompatible multi-return struct");
+
+    // Emit a call returning the packed struct.
+    Value structVal =
+        emitc::CallOp::create(rewriter, loc, callOp.getCalleeAttr(),
+                              TypeRange{*structType}, adaptor.getOperands())
+            .getResult(0);
 
+    // Unpack struct fields to replace the original multiple results.
+    MLIRContext *ctx = rewriter.getContext();
+    auto noInit = emitc::OpaqueAttr::get(ctx, "");
+    Value structLv =
+        emitc::VariableOp::create(rewriter, loc,
+                                  emitc::LValueType::get(*structType), noInit)
+            .getResult();
+    emitc::AssignOp::create(rewriter, loc, structLv, structVal);
+    SmallVector<Value> results;
+    for (auto [i, result] : llvm::enumerate(callOp.getResults())) {
+      if (result.use_empty()) {
+        results.push_back(Value()); // No replacement needed.
+        continue;
+      }
+      Type fieldType = convertedResultTypes[i];
+      StringAttr fieldName =
+          rewriter.getStringAttr("field" + std::to_string(i));
+      Value fieldLv = emitc::MemberOp::create(rewriter, loc,
+                                              emitc::LValueType::get(fieldType),
+                                              fieldName, structLv)
+                          .getResult();
+      results.push_back(
+          emitc::LoadOp::create(rewriter, loc, fieldType, fieldLv).getResult());
+    }
+
+    rewriter.replaceOp(callOp, results);
     return success();
   }
+
+private:
+  bool lowerToCpp;
 };
 
 class FuncOpConversion final : public OpConversionPattern<func::FuncOp> {
 public:
-  using OpConversionPattern<func::FuncOp>::OpConversionPattern;
+  FuncOpConversion(const TypeConverter &typeConverter, MLIRContext *ctx,
+                   bool lowerToCpp)
+      : OpConversionPattern<func::FuncOp>(typeConverter, ctx),
+        lowerToCpp(lowerToCpp) {}
 
   LogicalResult
   matchAndRewrite(func::FuncOp funcOp, OpAdaptor adaptor,
                   ConversionPatternRewriter &rewriter) const override {
     FunctionType fnType = funcOp.getFunctionType();
 
-    if (fnType.getNumResults() > 1)
+    // Do not convert multiple-return functions if lowering target is Cpp.
+    // The translator will emit the return values as an std::tuple.
+    if (fnType.getNumResults() > 1 && lowerToCpp)
       return rewriter.notifyMatchFailure(
           funcOp, "only functions with zero or one result can be converted");
 
@@ -102,15 +270,28 @@ class FuncOpConversion final : public OpConversionPattern<func::FuncOp> {
       signatureConverter.addInputs(argType.index(), convertedType);
     }
 
-    Type resultType;
-    if (fnType.getNumResults() == 1) {
-      resultType = getTypeConverter()->convertType(fnType.getResult(0));
+    SmallVector<Type> convertedResultTypes;
+    for (Type t : fnType.getResults()) {
+      Type resultType = getTypeConverter()->convertType(t);
       if (!resultType)
         return rewriter.notifyMatchFailure(funcOp,
                                            "result type conversion failed");
       if (isa<emitc::ArrayType>(resultType))
         return rewriter.notifyMatchFailure(
             funcOp, "functions returning arrays are not supported");
+      convertedResultTypes.push_back(resultType);
+    }
+
+    Type resultType;
+    if (fnType.getNumResults() == 1) {
+      resultType = convertedResultTypes[0];
+    } else if (fnType.getNumResults() > 1) {
+      auto structTypeOrErr = getOrCreateMultiReturnType(
+          rewriter, funcOp.getLoc(), funcOp, convertedResultTypes);
+      if (failed(structTypeOrErr))
+        return rewriter.notifyMatchFailure(funcOp,
+                                           "incompatible multi-return struct");
+      resultType = *structTypeOrErr;
     }
 
     // Create the converted `emitc.func` op.
@@ -151,28 +332,57 @@ class FuncOpConversion final : public OpConversionPattern<func::FuncOp> {
 
     return success();
   }
+
+private:
+  bool lowerToCpp;
 };
 
 class ReturnOpConversion final : public OpConversionPattern<func::ReturnOp> {
 public:
-  using OpConversionPattern<func::ReturnOp>::OpConversionPattern;
+  ReturnOpConversion(const TypeConverter &typeConverter, MLIRContext *ctx,
+                     bool lowerToCpp)
+      : OpConversionPattern<func::ReturnOp>(typeConverter, ctx),
+        lowerToCpp(lowerToCpp) {}
 
   LogicalResult
   matchAndRewrite(func::ReturnOp returnOp, OpAdaptor adaptor,
                   ConversionPatternRewriter &rewriter) const override {
-    if (returnOp.getNumOperands() > 1)
+    // Do not convert multiple-return functions if lowering target is Cpp.
+    // The translator will emit the return values as an std::tuple.
+    if (returnOp.getNumOperands() > 1 && lowerToCpp)
       return rewriter.notifyMatchFailure(
           returnOp, "only zero or one operand is supported");
-    if (returnOp.getNumOperands() == 1 &&
-        isa<emitc::ArrayType>(adaptor.getOperands()[0].getType()))
+
+    if (llvm::any_of(adaptor.getOperands(), [](Value operand) {
+          return isa<emitc::ArrayType>(operand.getType());
+        }))
       return rewriter.notifyMatchFailure(returnOp,
                                          "returning arrays is not supported");
 
-    rewriter.replaceOpWithNewOp<emitc::ReturnOp>(
-        returnOp,
-        returnOp.getNumOperands() ? adaptor.getOperands()[0] : nullptr);
+    if (returnOp.getNumOperands() <= 1) {
+      rewriter.replaceOpWithNewOp<emitc::ReturnOp>(
+          returnOp,
+          returnOp.getNumOperands() ? adaptor.getOperands()[0] : nullptr);
+      return success();
+    }
+
+    // Multi-operand return: pack values into a struct.
+    Location loc = returnOp.getLoc();
+
+    auto structType = getOrCreateMultiReturnType(
+        rewriter, loc, returnOp, adaptor.getOperands().getTypes());
+    if (failed(structType))
+      return rewriter.notifyMatchFailure(returnOp,
+                                         "incompatible multi-return struct");
+
+    Value structVal =
+        packValuesIntoStruct(rewriter, loc, adaptor.getOperands(), *structType);
+    rewriter.replaceOpWithNewOp<emitc::ReturnOp>(returnOp, structVal);
     return success();
   }
+
+private:
+  bool lowerToCpp;
 };
 } // namespace
 
@@ -181,9 +391,10 @@ class ReturnOpConversion final : public OpConversionPattern<func::ReturnOp> {
 //===----------------------------------------------------------------------===//
 
 void mlir::populateFuncToEmitCPatterns(const TypeConverter &typeConverter,
-                                       RewritePatternSet &patterns) {
+                                       RewritePatternSet &patterns,
+                                       bool lowerToCpp) {
   MLIRContext *ctx = patterns.getContext();
 
   patterns.add<CallOpConversion, FuncOpConversion, ReturnOpConversion>(
-      typeConverter, ctx);
+      typeConverter, ctx, lowerToCpp);
 }

diff  --git a/mlir/lib/Conversion/FuncToEmitC/FuncToEmitCPass.cpp b/mlir/lib/Conversion/FuncToEmitC/FuncToEmitCPass.cpp
index b82a7266dc95f..01129f3e4c5cd 100644
--- a/mlir/lib/Conversion/FuncToEmitC/FuncToEmitCPass.cpp
+++ b/mlir/lib/Conversion/FuncToEmitC/FuncToEmitCPass.cpp
@@ -28,6 +28,7 @@ using namespace mlir;
 namespace {
 struct ConvertFuncToEmitC
     : public impl::ConvertFuncToEmitCBase<ConvertFuncToEmitC> {
+  using Base::Base;
   void runOnOperation() override;
 };
 } // namespace
@@ -48,7 +49,7 @@ void ConvertFuncToEmitC::runOnOperation() {
     return type;
   });
 
-  populateFuncToEmitCPatterns(typeConverter, patterns);
+  populateFuncToEmitCPatterns(typeConverter, patterns, this->lowerToCpp);
 
   if (failed(
           applyPartialConversion(getOperation(), target, std::move(patterns))))

diff  --git a/mlir/lib/Conversion/MemRefToEmitC/MemRefToEmitC.cpp b/mlir/lib/Conversion/MemRefToEmitC/MemRefToEmitC.cpp
index 87144aac9f6f9..693ebc7bc3bd0 100644
--- a/mlir/lib/Conversion/MemRefToEmitC/MemRefToEmitC.cpp
+++ b/mlir/lib/Conversion/MemRefToEmitC/MemRefToEmitC.cpp
@@ -44,7 +44,7 @@ struct MemRefToEmitCDialectInterface : public ConvertToEmitCPatternInterface {
   /// and mark dialect legal for the conversion target.
   void populateConvertToEmitCConversionPatterns(
       ConversionTarget &target, TypeConverter &typeConverter,
-      RewritePatternSet &patterns) const final {
+      RewritePatternSet &patterns, std::optional<bool> lowerToCpp) const final {
     populateMemRefToEmitCTypeConversion(typeConverter);
     populateMemRefToEmitCConversionPatterns(patterns, typeConverter);
   }

diff  --git a/mlir/lib/Conversion/SCFToEmitC/SCFToEmitC.cpp b/mlir/lib/Conversion/SCFToEmitC/SCFToEmitC.cpp
index d7c943ef7c4f1..b4616e23a7066 100644
--- a/mlir/lib/Conversion/SCFToEmitC/SCFToEmitC.cpp
+++ b/mlir/lib/Conversion/SCFToEmitC/SCFToEmitC.cpp
@@ -42,7 +42,7 @@ struct SCFToEmitCDialectInterface : public ConvertToEmitCPatternInterface {
   /// and mark dialect legal for the conversion target.
   void populateConvertToEmitCConversionPatterns(
       ConversionTarget &target, TypeConverter &typeConverter,
-      RewritePatternSet &patterns) const final {
+      RewritePatternSet &patterns, std::optional<bool> lowerToCpp) const final {
     populateEmitCSizeTTypeConversions(typeConverter);
     populateSCFToEmitCConversionPatterns(patterns, typeConverter);
   }

diff  --git a/mlir/test/Conversion/FuncToEmitC/func-to-emitc-failed.mlir b/mlir/test/Conversion/FuncToEmitC/func-to-emitc-failed.mlir
index d85069371e691..46e5319d7d17c 100644
--- a/mlir/test/Conversion/FuncToEmitC/func-to-emitc-failed.mlir
+++ b/mlir/test/Conversion/FuncToEmitC/func-to-emitc-failed.mlir
@@ -1,4 +1,4 @@
-// RUN: mlir-opt -convert-func-to-emitc %s -split-input-file -verify-diagnostics
+// RUN: mlir-opt -convert-func-to-emitc="lower-to-cpp=false" %s -split-input-file -verify-diagnostics
 
 // expected-error at +1 {{failed to legalize operation 'func.func'}}
 func.func @unsuppoted_emitc_type(%arg0: i4) -> i4 {
@@ -97,3 +97,89 @@ func.func private @caller(%arg0: memref<1xi64>, %arg1: i64) -> memref<1xi64> {
   %0 = call @callee(%arg1) : (i64) -> i64
   return %arg0 : memref<1xi64>
 }
+
+// -----
+
+// A symbol with the auto-generated struct name already exists but is not an
+// emitc.class (here it is an emitc.func).
+emitc.func @return_i32_i32() { emitc.return }
+// expected-error at +2 {{symbol 'return_i32_i32' exists but is not an emitc.class}}
+// expected-error at +1 {{failed to legalize operation 'func.func'}}
+func.func @symbol_not_a_class(%arg0: i32) -> (i32, i32) {
+  return %arg0, %arg0 : i32, i32
+}
+
+// -----
+
+// The existing emitc.class is not a struct (class_type != struct).
+emitc.class @return_i32_i32 {
+  emitc.field @field0 : i32
+  emitc.field @field1 : i32
+}
+
+// expected-error at +2 {{existing class 'return_i32_i32' is not a struct}}
+// expected-error at +1 {{failed to legalize operation 'func.func'}}
+func.func @class_not_a_struct(%arg0: i32) -> (i32, i32) {
+  return %arg0, %arg0 : i32, i32
+}
+
+// -----
+
+// The existing emitc.class has a method, so it cannot be used as a plain
+// struct.
+emitc.class struct @return_i32_i32 {
+  emitc.func @method() { emitc.return }
+}
+
+// expected-error at +2 {{existing class 'return_i32_i32' has methods; expected a plain struct}}
+// expected-error at +1 {{failed to legalize operation 'func.func'}}
+func.func @class_has_methods(%arg0: i32) -> (i32, i32) {
+  return %arg0, %arg0 : i32, i32
+}
+
+// -----
+
+// The existing emitc.class has fewer fields than the return types require.
+emitc.class struct @return_i32_i32 {
+  emitc.field @field0 : i32
+}
+
+// expected-error at +2 {{existing class 'return_i32_i32' has wrong number of fields}}
+// expected-error at +1 {{failed to legalize operation 'func.func'}}
+func.func @class_wrong_field_count(%arg0: i32) -> (i32, i32) {
+  return %arg0, %arg0 : i32, i32
+}
+
+// -----
+
+// The existing emitc.class has fields with unexpected names.
+emitc.class struct @return_i32_i32 {
+  emitc.field @a : i32
+  emitc.field @b : i32
+}
+
+// expected-error at +2 {{existing class 'return_i32_i32': unexpected field name at index 0}}
+// expected-error at +1 {{failed to legalize operation 'func.func'}}
+func.func @class_wrong_field_names(%arg0: i32) -> (i32, i32) {
+  return %arg0, %arg0 : i32, i32
+}
+
+// -----
+
+// The existing emitc.class has fields with the wrong types.
+emitc.class struct @return_i32_i32 {
+  emitc.field @field0 : i64
+  emitc.field @field1 : i32
+}
+
+// expected-error at +2 {{existing class 'return_i32_i32': wrong type for field 0}}
+// expected-error at +1 {{failed to legalize operation 'func.func'}}
+func.func @class_wrong_field_types(%arg0: i32) -> (i32, i32) {
+  return %arg0, %arg0 : i32, i32
+}
+
+// -----
+
+// Multi-result function where one result is an array type.
+// expected-error at +1 {{failed to legalize operation 'func.func'}}
+func.func private @multi_result_with_array() -> (i32, !emitc.array<10xi32>)

diff  --git a/mlir/test/Conversion/FuncToEmitC/func-to-emitc.mlir b/mlir/test/Conversion/FuncToEmitC/func-to-emitc.mlir
index 6824a64dda3ef..1a2a8e764e22d 100644
--- a/mlir/test/Conversion/FuncToEmitC/func-to-emitc.mlir
+++ b/mlir/test/Conversion/FuncToEmitC/func-to-emitc.mlir
@@ -1,5 +1,5 @@
-// RUN: mlir-opt -split-input-file -convert-func-to-emitc %s | FileCheck %s
-// RUN: mlir-opt -split-input-file -convert-to-emitc="filter-dialects=func" %s | FileCheck %s
+// RUN: mlir-opt -split-input-file -convert-func-to-emitc="lower-to-cpp=false" %s | FileCheck %s
+// RUN: mlir-opt -split-input-file -convert-to-emitc="filter-dialects=func lower-to-cpp=false" %s | FileCheck %s
 
 // CHECK-LABEL: emitc.func @foo()
 // CHECK-NEXT: return
@@ -75,3 +75,97 @@ func.func @call() {
   call @return_void() : () -> ()
   return
 }
+
+// -----
+
+// Multi-result function: check that an emitc.class struct is created and the
+// function returns the packed struct.
+// CHECK-LABEL:   emitc.class struct @return_i32_i32 {
+// CHECK:           emitc.field @field0 : i32
+// CHECK:           emitc.field @field1 : i32
+// CHECK:         }
+// CHECK-LABEL:   emitc.func @return_two(
+// CHECK-SAME:      %[[ARG0:.*]]: i32,
+// CHECK-SAME:      %[[ARG1:.*]]: i32) -> !emitc.opaque<"struct return_i32_i32"> {
+// CHECK:           %[[VAL_0:.*]] = "emitc.variable"() <{value = #emitc.opaque<"">}> : () -> !emitc.lvalue<!emitc.opaque<"struct return_i32_i32">>
+// CHECK:           %[[VAL_1:.*]] = "emitc.member"(%[[VAL_0]]) <{member = "field0"}> : (!emitc.lvalue<!emitc.opaque<"struct return_i32_i32">>) -> !emitc.lvalue<i32>
+// CHECK:           assign %[[ARG0]] : i32 to %[[VAL_1]] : <i32>
+// CHECK:           %[[VAL_2:.*]] = "emitc.member"(%[[VAL_0]]) <{member = "field1"}> : (!emitc.lvalue<!emitc.opaque<"struct return_i32_i32">>) -> !emitc.lvalue<i32>
+// CHECK:           assign %[[ARG1]] : i32 to %[[VAL_2]] : <i32>
+// CHECK:           %[[VAL_3:.*]] = load %[[VAL_0]] : <!emitc.opaque<"struct return_i32_i32">>
+// CHECK:           return %[[VAL_3]] : !emitc.opaque<"struct return_i32_i32">
+// CHECK:         }
+func.func @return_two(%arg0: i32, %arg1: i32) -> (i32, i32) {
+  return %arg0, %arg1 : i32, i32
+}
+
+// -----
+
+// Call to a multi-result function: check that the call returns the struct and
+// that only the field actually used is extracted.
+// CHECK-LABEL:   emitc.class struct @return_i32_i32 {
+// CHECK:           emitc.field @field0 : i32
+// CHECK:           emitc.field @field1 : i32
+// CHECK:         }
+// CHECK-LABEL:   emitc.func @return_two(
+// CHECK-SAME:      %[[ARG0:.*]]: i32,
+// CHECK-SAME:      %[[ARG1:.*]]: i32) -> !emitc.opaque<"struct return_i32_i32"> {
+// CHECK-NEXT:      %[[VAL_0:.*]] = "emitc.variable"() <{value = #emitc.opaque<"">}> : () -> !emitc.lvalue<!emitc.opaque<"struct return_i32_i32">>
+// CHECK-NEXT:      %[[VAL_1:.*]] = "emitc.member"(%[[VAL_0]]) <{member = "field0"}> : (!emitc.lvalue<!emitc.opaque<"struct return_i32_i32">>) -> !emitc.lvalue<i32>
+// CHECK-NEXT:      assign %[[ARG0]] : i32 to %[[VAL_1]] : <i32>
+// CHECK-NEXT:      %[[VAL_2:.*]] = "emitc.member"(%[[VAL_0]]) <{member = "field1"}> : (!emitc.lvalue<!emitc.opaque<"struct return_i32_i32">>) -> !emitc.lvalue<i32>
+// CHECK-NEXT:      assign %[[ARG1]] : i32 to %[[VAL_2]] : <i32>
+// CHECK-NEXT:      %[[VAL_3:.*]] = load %[[VAL_0]] : <!emitc.opaque<"struct return_i32_i32">>
+// CHECK-NEXT:      return %[[VAL_3]] : !emitc.opaque<"struct return_i32_i32">
+// CHECK-NEXT:    }
+// CHECK-LABEL:   emitc.func @caller(
+// CHECK-SAME:      %[[ARG0:.*]]: i32) -> i32 {
+// CHECK-NEXT:      %[[VAL_0:.*]] = call @return_two(%[[ARG0]], %[[ARG0]]) : (i32, i32) -> !emitc.opaque<"struct return_i32_i32">
+// CHECK-NEXT:      %[[VAL_1:.*]] = "emitc.variable"() <{value = #emitc.opaque<"">}> : () -> !emitc.lvalue<!emitc.opaque<"struct return_i32_i32">>
+// CHECK-NEXT:      assign %[[VAL_0]] : !emitc.opaque<"struct return_i32_i32"> to %[[VAL_1]] : <!emitc.opaque<"struct return_i32_i32">>
+// CHECK-NEXT:      %[[VAL_2:.*]] = "emitc.member"(%[[VAL_1]]) <{member = "field1"}> : (!emitc.lvalue<!emitc.opaque<"struct return_i32_i32">>) -> !emitc.lvalue<i32>
+// CHECK-NEXT:      %[[VAL_3:.*]] = load %[[VAL_2]] : <i32>
+// CHECK-NEXT:      return %[[VAL_3]] : i32
+// CHECK-NEXT:    }
+func.func @return_two(%arg0: i32, %arg1: i32) -> (i32, i32) {
+  return %arg0, %arg1 : i32, i32
+}
+func.func @caller(%arg0: i32) -> i32 {
+  %0, %1 = call @return_two(%arg0, %arg0) : (i32, i32) -> (i32, i32)
+  return %1 : i32
+}
+
+// -----
+
+// Two functions returning the same type tuple share one emitc.class.
+// CHECK-LABEL:   emitc.class struct @return_i32_i32 {
+// CHECK:           emitc.field @field0 : i32
+// CHECK:           emitc.field @field1 : i32
+// CHECK:         }
+// CHECK-LABEL:   emitc.func @first(
+// CHECK-SAME:                      %[[ARG0:.*]]: i32) -> !emitc.opaque<"struct return_i32_i32"> {
+// CHECK-NEXT:      %[[VAL_0:.*]] = "emitc.variable"() <{value = #emitc.opaque<"">}> : () -> !emitc.lvalue<!emitc.opaque<"struct return_i32_i32">>
+// CHECK-NEXT:      %[[VAL_1:.*]] = "emitc.member"(%[[VAL_0]]) <{member = "field0"}> : (!emitc.lvalue<!emitc.opaque<"struct return_i32_i32">>) -> !emitc.lvalue<i32>
+// CHECK-NEXT:      assign %[[ARG0]] : i32 to %[[VAL_1]] : <i32>
+// CHECK-NEXT:      %[[VAL_2:.*]] = "emitc.member"(%[[VAL_0]]) <{member = "field1"}> : (!emitc.lvalue<!emitc.opaque<"struct return_i32_i32">>) -> !emitc.lvalue<i32>
+// CHECK-NEXT:      assign %[[ARG0]] : i32 to %[[VAL_2]] : <i32>
+// CHECK-NEXT:      %[[VAL_3:.*]] = load %[[VAL_0]] : <!emitc.opaque<"struct return_i32_i32">>
+// CHECK-NEXT:      return %[[VAL_3]] : !emitc.opaque<"struct return_i32_i32">
+// CHECK-NEXT:    }
+// CHECK-LABEL:   emitc.func @second(
+// CHECK-SAME:      %[[ARG0:.*]]: i32) -> !emitc.opaque<"struct return_i32_i32"> {
+// CHECK-NEXT:      %[[VAL_0:.*]] = "emitc.variable"() <{value = #emitc.opaque<"">}> : () -> !emitc.lvalue<!emitc.opaque<"struct return_i32_i32">>
+// CHECK-NEXT:      %[[VAL_1:.*]] = "emitc.member"(%[[VAL_0]]) <{member = "field0"}> : (!emitc.lvalue<!emitc.opaque<"struct return_i32_i32">>) -> !emitc.lvalue<i32>
+// CHECK-NEXT:      assign %[[ARG0]] : i32 to %[[VAL_1]] : <i32>
+// CHECK-NEXT:      %[[VAL_2:.*]] = "emitc.member"(%[[VAL_0]]) <{member = "field1"}> : (!emitc.lvalue<!emitc.opaque<"struct return_i32_i32">>) -> !emitc.lvalue<i32>
+// CHECK-NEXT:      assign %[[ARG0]] : i32 to %[[VAL_2]] : <i32>
+// CHECK-NEXT:      %[[VAL_3:.*]] = load %[[VAL_0]] : <!emitc.opaque<"struct return_i32_i32">>
+// CHECK-NEXT:      return %[[VAL_3]] : !emitc.opaque<"struct return_i32_i32">
+// CHECK-NEXT:    }
+// CHECK-NOT:     emitc.class
+func.func @first(%arg0: i32) -> (i32, i32) {
+  return %arg0, %arg0 : i32, i32
+}
+func.func @second(%arg0: i32) -> (i32, i32) {
+  return %arg0, %arg0 : i32, i32
+}

diff  --git a/mlir/test/Target/Cpp/func.mlir b/mlir/test/Target/Cpp/func.mlir
index 9c9ea55bfc4e1..82f1ee9f6ec2b 100644
--- a/mlir/test/Target/Cpp/func.mlir
+++ b/mlir/test/Target/Cpp/func.mlir
@@ -43,3 +43,66 @@ emitc.func private @extern_func(i32) attributes {specifiers = ["extern"]}
 
 emitc.func private @array_arg(!emitc.array<3xi32>) attributes {specifiers = ["extern"]}
 // CPP-DEFAULT: extern void array_arg(int32_t[3]);
+
+emitc.class struct @return_i32_i32 {
+  emitc.field @field0 : i32
+  emitc.field @field1 : i32
+}
+
+emitc.func @return_two(%arg0: i32, %arg1: i32) -> !emitc.opaque<"struct return_i32_i32"> {
+  %0 = "emitc.variable"() <{value = #emitc.opaque<"">}> : () -> !emitc.lvalue<!emitc.opaque<"struct return_i32_i32">>
+  %1 = "emitc.member"(%0) <{member = "field0"}> : (!emitc.lvalue<!emitc.opaque<"struct return_i32_i32">>) -> !emitc.lvalue<i32>
+  assign %arg0 : i32 to %1 : <i32>
+  %2 = "emitc.member"(%0) <{member = "field1"}> : (!emitc.lvalue<!emitc.opaque<"struct return_i32_i32">>) -> !emitc.lvalue<i32>
+  assign %arg1 : i32 to %2 : <i32>
+  %3 = load %0 : <!emitc.opaque<"struct return_i32_i32">>
+  return %3 : !emitc.opaque<"struct return_i32_i32">
+}
+
+emitc.func @call_two(%arg0: i32) -> i32 {
+  %0 = call @return_two(%arg0, %arg0) : (i32, i32) -> !emitc.opaque<"struct return_i32_i32">
+  %1 = "emitc.variable"() <{value = #emitc.opaque<"">}> : () -> !emitc.lvalue<!emitc.opaque<"struct return_i32_i32">>
+  assign %0 : !emitc.opaque<"struct return_i32_i32"> to %1 : <!emitc.opaque<"struct return_i32_i32">>
+  %2 = "emitc.member"(%1) <{member = "field1"}> : (!emitc.lvalue<!emitc.opaque<"struct return_i32_i32">>) -> !emitc.lvalue<i32>
+  %3 = load %2 : <i32>
+  return %3 : i32
+}
+
+// CPP-DEFAULT: struct return_i32_i32 {
+// CPP-DEFAULT-NEXT:   int32_t field0;
+// CPP-DEFAULT-NEXT:   int32_t field1;
+// CPP-DEFAULT-NEXT: };
+// CPP-DEFAULT-NEXT: struct return_i32_i32 return_two(int32_t [[V1:[^ ]*]], int32_t [[V2:[^ ]*]]) {
+// CPP-DEFAULT-NEXT:   struct return_i32_i32 [[V3:[^ ]*]];
+// CPP-DEFAULT-NEXT:   [[V3]].field0 = [[V1]];
+// CPP-DEFAULT-NEXT:   [[V3]].field1 = [[V2]];
+// CPP-DEFAULT-NEXT:   struct return_i32_i32 [[V4:[^ ]*]] = [[V3]];
+// CPP-DEFAULT-NEXT:   return [[V4]];
+// CPP-DEFAULT-NEXT: }
+// CPP-DEFAULT-NEXT: int32_t call_two(int32_t [[V1:[^ ]*]]) {
+// CPP-DEFAULT-NEXT:   struct return_i32_i32 [[V2:[^ ]*]] = return_two([[V1]], [[V1]]);
+// CPP-DEFAULT-NEXT:   struct return_i32_i32 [[V3:[^ ]*]];
+// CPP-DEFAULT-NEXT:   [[V3]] = [[V2]];
+// CPP-DEFAULT-NEXT:   int32_t [[V4:[^ ]*]] = [[V3]].field1;
+// CPP-DEFAULT-NEXT:   return [[V4]];
+
+// CPP-DECLTOP: struct return_i32_i32 {
+// CPP-DECLTOP-NEXT:   int32_t field0;
+// CPP-DECLTOP-NEXT:   int32_t field1;
+// CPP-DECLTOP-NEXT: };
+// CPP-DECLTOP-NEXT: struct return_i32_i32 return_two(int32_t [[V1:[^ ]*]], int32_t [[V2:[^ ]*]]) {
+// CPP-DECLTOP-NEXT:   struct return_i32_i32 [[V3:[^ ]*]];
+// CPP-DECLTOP-NEXT:   struct return_i32_i32 [[V4:[^ ]*]];
+// CPP-DECLTOP:        [[V3]].field0 = [[V1]];
+// CPP-DECLTOP-NEXT:   [[V3]].field1 = [[V2]];
+// CPP-DECLTOP-NEXT:   [[V4]] = [[V3]];
+// CPP-DECLTOP-NEXT:   return [[V4]];
+// CPP-DECLTOP-NEXT: }
+// CPP-DECLTOP-NEXT: int32_t call_two(int32_t [[V1:[^ ]*]]) {
+// CPP-DECLTOP-NEXT:   struct return_i32_i32 [[V2:[^ ]*]];
+// CPP-DECLTOP-NEXT:   struct return_i32_i32 [[V3:[^ ]*]];
+// CPP-DECLTOP-NEXT:   int32_t [[V4:[^ ]*]];
+// CPP-DECLTOP-NEXT:   [[V2]] = return_two([[V1]], [[V1]]);
+// CPP-DECLTOP:        [[V3]] = [[V2]];
+// CPP-DECLTOP-NEXT:   [[V4]] = [[V3]].field1;
+// CPP-DECLTOP-NEXT:   return [[V4]];


        


More information about the Mlir-commits mailing list