[Mlir-commits] [mlir] Reland emitc lower multi return functions (PR #203026)
llvmlistbot at llvm.org
llvmlistbot at llvm.org
Wed Jun 10 09:20:09 PDT 2026
llvmorg-github-actions[bot] wrote:
<!--LLVM PR SUMMARY COMMENT-->
@llvm/pr-subscribers-mlir
Author: Gil Rapaport (aniragil)
<details>
<summary>Changes</summary>
[mlir][emitc] Fix GCC 7 func-to-emitc build
Use the return 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
---
Patch is 38.07 KiB, truncated to 20.00 KiB below, full version: https://github.com/llvm/llvm-project/pull/203026.diff
13 Files Affected:
- (modified) mlir/include/mlir/Conversion/ConvertToEmitC/ConvertToEmitCPatternInterface.td (+2-1)
- (modified) mlir/include/mlir/Conversion/ConvertToEmitC/ToEmitCInterface.h (+2)
- (modified) mlir/include/mlir/Conversion/FuncToEmitC/FuncToEmitC.h (+2-1)
- (modified) mlir/include/mlir/Conversion/Passes.td (+6)
- (modified) mlir/lib/Conversion/ArithToEmitC/ArithToEmitC.cpp (+1-1)
- (modified) mlir/lib/Conversion/ConvertToEmitC/ConvertToEmitCPass.cpp (+13-5)
- (modified) mlir/lib/Conversion/FuncToEmitC/FuncToEmitC.cpp (+236-25)
- (modified) mlir/lib/Conversion/FuncToEmitC/FuncToEmitCPass.cpp (+2-1)
- (modified) mlir/lib/Conversion/MemRefToEmitC/MemRefToEmitC.cpp (+1-1)
- (modified) mlir/lib/Conversion/SCFToEmitC/SCFToEmitC.cpp (+1-1)
- (modified) mlir/test/Conversion/FuncToEmitC/func-to-emitc-failed.mlir (+87-1)
- (modified) mlir/test/Conversion/FuncToEmitC/func-to-emitc.mlir (+96-2)
- (modified) mlir/test/Target/Cpp/func.mlir (+63)
``````````diff
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..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 ®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 operan...
[truncated]
``````````
</details>
https://github.com/llvm/llvm-project/pull/203026
More information about the Mlir-commits
mailing list