[Mlir-commits] [mlir] [mlir-c] Add ConversionTarget dynamic legality C API (PR #206161)
Maksim Levental
llvmlistbot at llvm.org
Sun Jun 28 18:31:37 PDT 2026
https://github.com/makslevental updated https://github.com/llvm/llvm-project/pull/206161
>From 006f5fa83d35d161fdd8ed7834268352fc390cf6 Mon Sep 17 00:00:00 2001
From: makslevental <maksim.levental at gmail.com>
Date: Fri, 26 Jun 2026 12:20:59 -0700
Subject: [PATCH] [mlir-c] Add ConversionTarget dynamic legality C API
Add mlirConversionTargetAddDynamicallyLegalOp,
mlirConversionTargetAddDynamicallyLegalDialect,
mlirConversionTargetMarkOpRecursivelyLegal, and
mlirConversionTargetMarkUnknownOpDynamicallyLegal to enable
per-instance legality callbacks from C.
---
mlir/include/mlir-c/Rewrite.h | 28 ++++++++++++++
mlir/lib/CAPI/Transforms/Rewrite.cpp | 44 +++++++++++++++++++++
mlir/test/CAPI/rewrite.c | 58 ++++++++++++++++++++++++++++
3 files changed, 130 insertions(+)
diff --git a/mlir/include/mlir-c/Rewrite.h b/mlir/include/mlir-c/Rewrite.h
index ac243a4c9d8f9..eed23a15264fb 100644
--- a/mlir/include/mlir-c/Rewrite.h
+++ b/mlir/include/mlir-c/Rewrite.h
@@ -533,6 +533,34 @@ MLIR_CAPI_EXPORTED void
mlirConversionTargetAddIllegalDialect(MlirConversionTarget target,
MlirStringRef dialectName);
+/// Callback for dynamic legality checks. Return true if the operation is
+/// legal, false if illegal.
+typedef bool (*MlirConversionTargetDynamicLegalityCallback)(MlirOperation op,
+ void *userData);
+
+/// Register the given operation as dynamically legal, with a callback to
+/// determine per-instance legality.
+MLIR_CAPI_EXPORTED void mlirConversionTargetAddDynamicallyLegalOp(
+ MlirConversionTarget target, MlirStringRef opName,
+ MlirConversionTargetDynamicLegalityCallback callback, void *userData);
+
+/// Register the given dialect as dynamically legal, with a callback to
+/// determine per-instance legality for all operations in the dialect.
+MLIR_CAPI_EXPORTED void mlirConversionTargetAddDynamicallyLegalDialect(
+ MlirConversionTarget target, MlirStringRef dialectName,
+ MlirConversionTargetDynamicLegalityCallback callback, void *userData);
+
+/// Mark the given operation as recursively legal. The optional callback (may
+/// be NULL) determines whether a specific instance is recursively legal.
+MLIR_CAPI_EXPORTED void mlirConversionTargetMarkOpRecursivelyLegal(
+ MlirConversionTarget target, MlirStringRef opName,
+ MlirConversionTargetDynamicLegalityCallback callback, void *userData);
+
+/// Mark unknown operations as dynamically legal, with a callback.
+MLIR_CAPI_EXPORTED void mlirConversionTargetMarkUnknownOpDynamicallyLegal(
+ MlirConversionTarget target,
+ MlirConversionTargetDynamicLegalityCallback callback, void *userData);
+
//===----------------------------------------------------------------------===//
/// TypeConverter API
//===----------------------------------------------------------------------===//
diff --git a/mlir/lib/CAPI/Transforms/Rewrite.cpp b/mlir/lib/CAPI/Transforms/Rewrite.cpp
index 56ce9212f4811..12f56b397161d 100644
--- a/mlir/lib/CAPI/Transforms/Rewrite.cpp
+++ b/mlir/lib/CAPI/Transforms/Rewrite.cpp
@@ -575,6 +575,50 @@ void mlirConversionTargetAddIllegalDialect(MlirConversionTarget target,
unwrap(target)->addIllegalDialect(unwrap(dialectName));
}
+void mlirConversionTargetAddDynamicallyLegalOp(
+ MlirConversionTarget target, MlirStringRef opName,
+ MlirConversionTargetDynamicLegalityCallback callback, void *userData) {
+ MLIRContext *ctx = &unwrap(target)->getContext();
+ OperationName name(unwrap(opName), ctx);
+ unwrap(target)->addDynamicallyLegalOp(
+ name, [callback, userData](Operation *op) -> std::optional<bool> {
+ return callback(wrap(op), userData);
+ });
+}
+
+void mlirConversionTargetAddDynamicallyLegalDialect(
+ MlirConversionTarget target, MlirStringRef dialectName,
+ MlirConversionTargetDynamicLegalityCallback callback, void *userData) {
+ unwrap(target)->addDynamicallyLegalDialect(
+ [callback, userData](Operation *op) -> std::optional<bool> {
+ return callback(wrap(op), userData);
+ },
+ unwrap(dialectName));
+}
+
+void mlirConversionTargetMarkOpRecursivelyLegal(
+ MlirConversionTarget target, MlirStringRef opName,
+ MlirConversionTargetDynamicLegalityCallback callback, void *userData) {
+ MLIRContext *ctx = &unwrap(target)->getContext();
+ OperationName name(unwrap(opName), ctx);
+ ConversionTarget::DynamicLegalityCallbackFn fn;
+ if (callback) {
+ fn = [callback, userData](Operation *op) -> std::optional<bool> {
+ return callback(wrap(op), userData);
+ };
+ }
+ unwrap(target)->markOpRecursivelyLegal(name, fn);
+}
+
+void mlirConversionTargetMarkUnknownOpDynamicallyLegal(
+ MlirConversionTarget target,
+ MlirConversionTargetDynamicLegalityCallback callback, void *userData) {
+ unwrap(target)->markUnknownOpDynamicallyLegal(
+ [callback, userData](Operation *op) -> std::optional<bool> {
+ return callback(wrap(op), userData);
+ });
+}
+
//===----------------------------------------------------------------------===//
/// TypeConverter API
//===----------------------------------------------------------------------===//
diff --git a/mlir/test/CAPI/rewrite.c b/mlir/test/CAPI/rewrite.c
index 3809dd6a7843f..774ff8e2d1a4b 100644
--- a/mlir/test/CAPI/rewrite.c
+++ b/mlir/test/CAPI/rewrite.c
@@ -623,6 +623,63 @@ void testCloneWithMapping(MlirContext ctx) {
fprintf(stderr, "testCloneWithMapping: PASSED\n");
}
+static bool dynamicLegalityAlwaysLegal(MlirOperation op, void *userData) {
+ (void)op;
+ intptr_t *counter = (intptr_t *)userData;
+ (*counter)++;
+ return true;
+}
+
+static bool dynamicLegalityAlwaysIllegal(MlirOperation op, void *userData) {
+ (void)op;
+ intptr_t *counter = (intptr_t *)userData;
+ (*counter)++;
+ return false;
+}
+
+void testConversionTargetDynamicLegality(MlirContext ctx) {
+ // CHECK-LABEL: @testConversionTargetDynamicLegality
+ fprintf(stderr, "@testConversionTargetDynamicLegality\n");
+
+ MlirConversionTarget target = mlirConversionTargetCreate(ctx);
+
+ // Test addDynamicallyLegalOp
+ intptr_t legalCounter = 0;
+ mlirConversionTargetAddDynamicallyLegalOp(
+ target, mlirStringRefCreateFromCString("dialect.op1"),
+ dynamicLegalityAlwaysLegal, &legalCounter);
+
+ // Test addDynamicallyLegalDialect
+ intptr_t dialectCounter = 0;
+ mlirConversionTargetAddDynamicallyLegalDialect(
+ target, mlirStringRefCreateFromCString("dialect"),
+ dynamicLegalityAlwaysIllegal, &dialectCounter);
+
+ // Test markOpRecursivelyLegal (with callback) - op must be legal first
+ intptr_t recursiveCounter = 0;
+ mlirConversionTargetAddLegalOp(
+ target, mlirStringRefCreateFromCString("builtin.module"));
+ mlirConversionTargetMarkOpRecursivelyLegal(
+ target, mlirStringRefCreateFromCString("builtin.module"),
+ dynamicLegalityAlwaysLegal, &recursiveCounter);
+
+ // Test markOpRecursivelyLegal (without callback - NULL) - op must be legal
+ mlirConversionTargetAddLegalOp(target,
+ mlirStringRefCreateFromCString("func.func"));
+ mlirConversionTargetMarkOpRecursivelyLegal(
+ target, mlirStringRefCreateFromCString("func.func"), NULL, NULL);
+
+ // Test markUnknownOpDynamicallyLegal
+ intptr_t unknownCounter = 0;
+ mlirConversionTargetMarkUnknownOpDynamicallyLegal(
+ target, dynamicLegalityAlwaysLegal, &unknownCounter);
+
+ mlirConversionTargetDestroy(target);
+
+ // CHECK: testConversionTargetDynamicLegality: PASSED
+ fprintf(stderr, "testConversionTargetDynamicLegality: PASSED\n");
+}
+
int main(void) {
MlirContext ctx = mlirContextCreate();
mlirContextSetAllowUnregisteredDialects(ctx, true);
@@ -638,6 +695,7 @@ int main(void) {
testReplaceUses(ctx);
testGreedyRewriteDriverConfig(ctx);
testCloneWithMapping(ctx);
+ testConversionTargetDynamicLegality(ctx);
mlirContextDestroy(ctx);
return 0;
More information about the Mlir-commits
mailing list