[Mlir-commits] [mlir] 1e0a4c7 - [mlir][emitc] Lower multiple results as a struct (#200659)
llvmlistbot at llvm.org
llvmlistbot at llvm.org
Wed Jun 10 01:46:05 PDT 2026
Author: Gil Rapaport
Date: 2026-06-10T11:46:00+03:00
New Revision: 1e0a4c7a9154e46ef52a7c5b0ddbca69fbdcfacd
URL: https://github.com/llvm/llvm-project/commit/1e0a4c7a9154e46ef52a7c5b0ddbca69fbdcfacd
DIFF: https://github.com/llvm/llvm-project/commit/1e0a4c7a9154e46ef52a7c5b0ddbca69fbdcfacd.diff
LOG: [mlir][emitc] Lower multiple results as a struct (#200659)
Previously, func-to-emitc lowering rejected func.{func,call,return} with
more than one result/operand. Such ops are directly handled by the
translator which emits an `std::tuple` packing ther results, but is only
relevant for C++ users. This patch lifts that restriction by packing
multiple return values into an automatically-generated struct, e.g. for
a function returning (i32, i32):
emitc.class struct @return_i32_i32 {
emitc.field @field0 : i32
emitc.field @field1 : i32
}
On return, the operands are packed into a local struct variable which is
then loaded and returned. On call sites, the struct is stored in a local
variable, and each field is extracted to recreate the individual SSA
values of the original results. As with single-result functions,
`emitc.array` return types are not supported.
If a class with that name already exists, it is verified to have exactly
the expected fields with the correct types and no methods. Two functions
with the same return type tuple share a single class definition.
Backward compatibility is maintained by a new lower-to-cpp option which
defaults to `true` (unlike the same flag in memref-to-emitc), in which case
func-to-emitc continues to bail out on multi-return functions.
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 ®istry);
} // 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 ®istry) {
diff --git a/mlir/lib/Conversion/FuncToEmitC/FuncToEmitC.cpp b/mlir/lib/Conversion/FuncToEmitC/FuncToEmitC.cpp
index 4801f07d82c9f..27fd81ba2eca7 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 ®istry) {
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());
+ 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