[Mlir-commits] [mlir] [mlir-c] Add TypeConverter source and target materialization (PR #208934)
llvmlistbot at llvm.org
llvmlistbot at llvm.org
Sat Jul 11 12:05:44 PDT 2026
llvmorg-github-actions[bot] wrote:
<!--LLVM PR SUMMARY COMMENT-->
@llvm/pr-subscribers-mlir
Author: Maksim Levental (makslevental)
<details>
<summary>Changes</summary>
Exposes `TypeConverter::addSourceMaterialization` and `addTargetMaterialization` through the MLIR C API, continuing the buildout of the dialect-conversion C bindings (follows #<!-- -->206146 and #<!-- -->206161).
Assisted by: Claude
---
Full diff: https://github.com/llvm/llvm-project/pull/208934.diff
3 Files Affected:
- (modified) mlir/include/mlir-c/Rewrite.h (+23)
- (modified) mlir/lib/CAPI/Transforms/Rewrite.cpp (+41)
- (modified) mlir/test/CAPI/rewrite.c (+219)
``````````diff
diff --git a/mlir/include/mlir-c/Rewrite.h b/mlir/include/mlir-c/Rewrite.h
index 3356e6f445e47..e074fc96f2aeb 100644
--- a/mlir/include/mlir-c/Rewrite.h
+++ b/mlir/include/mlir-c/Rewrite.h
@@ -605,6 +605,29 @@ mlirTypeConverterAddConversion(MlirTypeConverter typeConverter,
MLIR_CAPI_EXPORTED MlirType
mlirTypeConverterConvertType(MlirTypeConverter typeConverter, MlirType type);
+/// Callback type for type materializations. Given a builder (passed as a
+/// rewriter), the desired output type, the input values, and a location, the
+/// callback must build a cast-like operation that produces a single value of
+/// `outputType` and return it. Returning a null MlirValue indicates failure, in
+/// which case another registered materialization may be attempted.
+typedef MlirValue (*MlirTypeConverterMaterializationCallback)(
+ MlirRewriterBase rewriter, MlirType outputType, intptr_t nInputs,
+ MlirValue *inputs, MlirLocation loc, void *userData);
+
+/// Register a source materialization with the given TypeConverter. This is
+/// invoked when a replacement value must be converted back to its original
+/// source type because some uses persist beyond the main conversion.
+MLIR_CAPI_EXPORTED void mlirTypeConverterAddSourceMaterialization(
+ MlirTypeConverter typeConverter,
+ MlirTypeConverterMaterializationCallback callback, void *userData);
+
+/// Register a target materialization with the given TypeConverter. This is
+/// invoked when a value must be converted to a target type according to a
+/// pattern's type converter.
+MLIR_CAPI_EXPORTED void mlirTypeConverterAddTargetMaterialization(
+ MlirTypeConverter typeConverter,
+ MlirTypeConverterMaterializationCallback callback, void *userData);
+
//===----------------------------------------------------------------------===//
/// ConversionPattern API
//===----------------------------------------------------------------------===//
diff --git a/mlir/lib/CAPI/Transforms/Rewrite.cpp b/mlir/lib/CAPI/Transforms/Rewrite.cpp
index 083ed6f999ae3..92456d6d8a435 100644
--- a/mlir/lib/CAPI/Transforms/Rewrite.cpp
+++ b/mlir/lib/CAPI/Transforms/Rewrite.cpp
@@ -669,6 +669,47 @@ MlirType mlirTypeConverterConvertType(MlirTypeConverter typeConverter,
return wrap(unwrap(typeConverter)->convertType(unwrap(type)));
}
+namespace {
+/// Wraps a C materialization callback as a C++ materialization callback of the
+/// form `Value(OpBuilder &, Type, ValueRange, Location)`, shared by both source
+/// and target materializations. The builder is always a RewriterBase in the
+/// conversion driver, so it is safe to expose it as an MlirRewriterBase.
+std::function<Value(OpBuilder &, Type, ValueRange, Location)>
+wrapMaterializationCallback(MlirTypeConverterMaterializationCallback callback,
+ void *userData) {
+ return [callback, userData](OpBuilder &builder, Type type, ValueRange inputs,
+ Location loc) -> Value {
+ SmallVector<MlirValue> wrappedInputs;
+ wrappedInputs.reserve(inputs.size());
+ for (Value v : inputs)
+ wrappedInputs.push_back(wrap(v));
+ MlirValue result =
+ callback(wrap(static_cast<RewriterBase *>(&builder)), wrap(type),
+ static_cast<intptr_t>(wrappedInputs.size()),
+ wrappedInputs.data(), wrap(loc), userData);
+ return mlirValueIsNull(result) ? Value() : unwrap(result);
+ };
+}
+} // namespace
+
+void mlirTypeConverterAddSourceMaterialization(
+ MlirTypeConverter typeConverter,
+ MlirTypeConverterMaterializationCallback callback, void *userData) {
+ assert(callback && "expected non-null materialization callback");
+ unwrap(typeConverter)
+ ->addSourceMaterialization(
+ wrapMaterializationCallback(callback, userData));
+}
+
+void mlirTypeConverterAddTargetMaterialization(
+ MlirTypeConverter typeConverter,
+ MlirTypeConverterMaterializationCallback callback, void *userData) {
+ assert(callback && "expected non-null materialization callback");
+ unwrap(typeConverter)
+ ->addTargetMaterialization(
+ wrapMaterializationCallback(callback, userData));
+}
+
//===----------------------------------------------------------------------===//
/// ConversionPattern API
//===----------------------------------------------------------------------===//
diff --git a/mlir/test/CAPI/rewrite.c b/mlir/test/CAPI/rewrite.c
index 439d1355af822..6d87d0236061b 100644
--- a/mlir/test/CAPI/rewrite.c
+++ b/mlir/test/CAPI/rewrite.c
@@ -802,6 +802,223 @@ void testConversionTargetDynamicLegality(MlirContext ctx) {
fprintf(stderr, "testConversionTargetDynamicLegality: PASSED\n");
}
+// Type conversion callback: maps i32 -> i64 and leaves every other type
+// unchanged (identity). Used by the materialization tests below.
+static MlirLogicalResult widenI32ToI64(MlirType type, MlirType *result,
+ void *userData) {
+ (void)userData;
+ if (mlirTypeIsAInteger(type) && mlirIntegerTypeGetWidth(type) == 32)
+ *result = mlirIntegerTypeGet(mlirTypeGetContext(type), 64);
+ else
+ *result = type;
+ return mlirLogicalResultSuccess();
+}
+
+// Materialization callback: builds a `test.cast` op that produces a single
+// value of `outputType` from the given inputs, and records that it ran by
+// bumping the counter passed as userData.
+static MlirValue buildCastMaterialization(MlirRewriterBase rewriter,
+ MlirType outputType, intptr_t nInputs,
+ MlirValue *inputs, MlirLocation loc,
+ void *userData) {
+ intptr_t *counter = (intptr_t *)userData;
+ if (counter)
+ (*counter)++;
+ MlirOperationState state =
+ mlirOperationStateGet(mlirStringRefCreateFromCString("test.cast"), loc);
+ mlirOperationStateAddOperands(&state, nInputs, inputs);
+ mlirOperationStateAddResults(&state, 1, &outputType);
+ MlirOperation castOp = mlirOperationCreate(&state);
+ mlirRewriterBaseInsert(rewriter, castOp);
+ return mlirOperationGetResult(castOp, 0);
+}
+
+// Conversion pattern for `test.source`: replaces it with a `test.source_i64`
+// op whose result has the widened (i64) type. Because the original result type
+// (i32) differs from the replacement type (i64), persisting uses force the
+// framework to insert a source materialization.
+static MlirLogicalResult convertSource(MlirConversionPattern pattern,
+ MlirOperation op, intptr_t nOperands,
+ MlirValue *operands,
+ MlirConversionPatternRewriter rewriter,
+ void *userData) {
+ (void)pattern;
+ (void)nOperands;
+ (void)operands;
+ (void)userData;
+ MlirContext ctx = mlirOperationGetContext(op);
+ MlirLocation loc = mlirOperationGetLocation(op);
+ MlirType i64 = mlirIntegerTypeGet(ctx, 64);
+ MlirOperationState state = mlirOperationStateGet(
+ mlirStringRefCreateFromCString("test.source_i64"), loc);
+ mlirOperationStateAddResults(&state, 1, &i64);
+ MlirOperation newOp = mlirOperationCreate(&state);
+
+ MlirRewriterBase base = mlirPatternRewriterAsBase(
+ mlirConversionPatternRewriterAsPatternRewriter(rewriter));
+ mlirRewriterBaseInsert(base, newOp);
+ MlirValue newVal = mlirOperationGetResult(newOp, 0);
+ mlirRewriterBaseReplaceOpWithValues(base, op, 1, &newVal);
+ return mlirLogicalResultSuccess();
+}
+
+void testTypeConverterSourceMaterialization(MlirContext ctx) {
+ // CHECK-LABEL: @testTypeConverterSourceMaterialization
+ fprintf(stderr, "@testTypeConverterSourceMaterialization\n");
+
+ // `test.source` produces an i32 that is consumed by the (legal) `test.user`.
+ // Converting `test.source` to an i64-producing op leaves `test.user` wanting
+ // the original i32, which triggers a source materialization back to i32.
+ const char *moduleString = "%0 = \"test.source\"() : () -> i32\n"
+ "\"test.user\"(%0) : (i32) -> ()\n";
+ MlirModule module =
+ mlirModuleCreateParse(ctx, mlirStringRefCreateFromCString(moduleString));
+ MlirOperation moduleOp = mlirModuleGetOperation(module);
+
+ MlirTypeConverter converter = mlirTypeConverterCreate();
+ mlirTypeConverterAddConversion(converter, widenI32ToI64, NULL);
+ intptr_t materializationCounter = 0;
+ mlirTypeConverterAddSourceMaterialization(converter, buildCastMaterialization,
+ &materializationCounter);
+
+ MlirRewritePatternSet patterns = mlirRewritePatternSetCreate(ctx);
+ MlirConversionPatternCallbacks callbacks = {NULL, NULL, convertSource};
+ MlirConversionPattern pattern = mlirOpConversionPatternCreate(
+ mlirStringRefCreateFromCString("test.source"), 1, ctx, converter,
+ callbacks, NULL, 0, NULL);
+ mlirRewritePatternSetAdd(patterns,
+ mlirConversionPatternAsRewritePattern(pattern));
+ MlirFrozenRewritePatternSet frozen = mlirFreezeRewritePattern(patterns);
+
+ MlirConversionTarget target = mlirConversionTargetCreate(ctx);
+ mlirConversionTargetAddIllegalOp(
+ target, mlirStringRefCreateFromCString("test.source"));
+ mlirConversionTargetAddLegalOp(
+ target, mlirStringRefCreateFromCString("test.source_i64"));
+ mlirConversionTargetAddLegalOp(target,
+ mlirStringRefCreateFromCString("test.cast"));
+ mlirConversionTargetAddLegalOp(target,
+ mlirStringRefCreateFromCString("test.user"));
+ mlirConversionTargetAddLegalOp(
+ target, mlirStringRefCreateFromCString("builtin.module"));
+
+ MlirConversionConfig config = mlirConversionConfigCreate();
+ MlirLogicalResult result =
+ mlirApplyPartialConversion(moduleOp, target, frozen, config);
+ assert(mlirLogicalResultIsSuccess(result));
+ assert(materializationCounter > 0 &&
+ "source materialization callback must be invoked");
+
+ mlirOperationDump(moduleOp);
+ // clang-format off
+ // CHECK: module {
+ // CHECK-NEXT: %[[v:.*]] = "test.source_i64"() : () -> i64
+ // CHECK-NEXT: %[[c:.*]] = "test.cast"(%[[v]]) : (i64) -> i32
+ // CHECK-NEXT: "test.user"(%[[c]]) : (i32) -> ()
+ // CHECK-NEXT: }
+ // clang-format on
+
+ mlirConversionConfigDestroy(config);
+ mlirConversionTargetDestroy(target);
+ mlirFrozenRewritePatternSetDestroy(frozen);
+ mlirTypeConverterDestroy(converter);
+ mlirModuleDestroy(module);
+
+ // CHECK: testTypeConverterSourceMaterialization: PASSED
+ fprintf(stderr, "testTypeConverterSourceMaterialization: PASSED\n");
+}
+
+// Conversion pattern for `test.consumer`: replaces it with a
+// `test.consumer_legal` op that consumes the (already remapped) operands. The
+// operand of the original op has type i32 but its producer is not converted, so
+// the framework inserts a target materialization to i64 before invoking this
+// pattern -- the remapped `operands` are therefore the i64 cast results.
+static MlirLogicalResult convertConsumer(MlirConversionPattern pattern,
+ MlirOperation op, intptr_t nOperands,
+ MlirValue *operands,
+ MlirConversionPatternRewriter rewriter,
+ void *userData) {
+ (void)pattern;
+ (void)userData;
+ MlirLocation loc = mlirOperationGetLocation(op);
+ MlirOperationState state = mlirOperationStateGet(
+ mlirStringRefCreateFromCString("test.consumer_legal"), loc);
+ mlirOperationStateAddOperands(&state, nOperands, operands);
+ MlirOperation newOp = mlirOperationCreate(&state);
+
+ MlirRewriterBase base = mlirPatternRewriterAsBase(
+ mlirConversionPatternRewriterAsPatternRewriter(rewriter));
+ mlirRewriterBaseInsert(base, newOp);
+ mlirRewriterBaseEraseOp(base, op);
+ return mlirLogicalResultSuccess();
+}
+
+void testTypeConverterTargetMaterialization(MlirContext ctx) {
+ // CHECK-LABEL: @testTypeConverterTargetMaterialization
+ fprintf(stderr, "@testTypeConverterTargetMaterialization\n");
+
+ // `test.consumer` takes an i32 from the (legal, unconverted) `test.producer`.
+ // Converting `test.consumer` requires its operand as i64, which triggers a
+ // target materialization from i32 to i64.
+ const char *moduleString = "%0 = \"test.producer\"() : () -> i32\n"
+ "\"test.consumer\"(%0) : (i32) -> ()\n";
+ MlirModule module =
+ mlirModuleCreateParse(ctx, mlirStringRefCreateFromCString(moduleString));
+ MlirOperation moduleOp = mlirModuleGetOperation(module);
+
+ MlirTypeConverter converter = mlirTypeConverterCreate();
+ mlirTypeConverterAddConversion(converter, widenI32ToI64, NULL);
+ intptr_t materializationCounter = 0;
+ mlirTypeConverterAddTargetMaterialization(converter, buildCastMaterialization,
+ &materializationCounter);
+
+ MlirRewritePatternSet patterns = mlirRewritePatternSetCreate(ctx);
+ MlirConversionPatternCallbacks callbacks = {NULL, NULL, convertConsumer};
+ MlirConversionPattern pattern = mlirOpConversionPatternCreate(
+ mlirStringRefCreateFromCString("test.consumer"), 1, ctx, converter,
+ callbacks, NULL, 0, NULL);
+ mlirRewritePatternSetAdd(patterns,
+ mlirConversionPatternAsRewritePattern(pattern));
+ MlirFrozenRewritePatternSet frozen = mlirFreezeRewritePattern(patterns);
+
+ MlirConversionTarget target = mlirConversionTargetCreate(ctx);
+ mlirConversionTargetAddIllegalOp(
+ target, mlirStringRefCreateFromCString("test.consumer"));
+ mlirConversionTargetAddLegalOp(
+ target, mlirStringRefCreateFromCString("test.producer"));
+ mlirConversionTargetAddLegalOp(
+ target, mlirStringRefCreateFromCString("test.consumer_legal"));
+ mlirConversionTargetAddLegalOp(target,
+ mlirStringRefCreateFromCString("test.cast"));
+ mlirConversionTargetAddLegalOp(
+ target, mlirStringRefCreateFromCString("builtin.module"));
+
+ MlirConversionConfig config = mlirConversionConfigCreate();
+ MlirLogicalResult result =
+ mlirApplyPartialConversion(moduleOp, target, frozen, config);
+ assert(mlirLogicalResultIsSuccess(result));
+ assert(materializationCounter > 0 &&
+ "target materialization callback must be invoked");
+
+ mlirOperationDump(moduleOp);
+ // clang-format off
+ // CHECK: module {
+ // CHECK-NEXT: %[[v:.*]] = "test.producer"() : () -> i32
+ // CHECK-NEXT: %[[c:.*]] = "test.cast"(%[[v]]) : (i32) -> i64
+ // CHECK-NEXT: "test.consumer_legal"(%[[c]]) : (i64) -> ()
+ // CHECK-NEXT: }
+ // clang-format on
+
+ mlirConversionConfigDestroy(config);
+ mlirConversionTargetDestroy(target);
+ mlirFrozenRewritePatternSetDestroy(frozen);
+ mlirTypeConverterDestroy(converter);
+ mlirModuleDestroy(module);
+
+ // CHECK: testTypeConverterTargetMaterialization: PASSED
+ fprintf(stderr, "testTypeConverterTargetMaterialization: PASSED\n");
+}
+
int main(void) {
MlirContext ctx = mlirContextCreate();
mlirContextSetAllowUnregisteredDialects(ctx, true);
@@ -818,6 +1035,8 @@ int main(void) {
testGreedyRewriteDriverConfig(ctx);
testCloneWithMapping(ctx);
testConversionTargetDynamicLegality(ctx);
+ testTypeConverterSourceMaterialization(ctx);
+ testTypeConverterTargetMaterialization(ctx);
mlirContextDestroy(ctx);
return 0;
``````````
</details>
https://github.com/llvm/llvm-project/pull/208934
More information about the Mlir-commits
mailing list