[Mlir-commits] [mlir] [MLIR][CAPI][Python] Add support for querying memory effect instances (PR #213459)
llvmlistbot at llvm.org
llvmlistbot at llvm.org
Mon Aug 3 06:58:11 PDT 2026
https://github.com/PragmaTwice updated https://github.com/llvm/llvm-project/pull/213459
>From 9d32f07b1e1c304181a6ef965a2c40ccbbca9567 Mon Sep 17 00:00:00 2001
From: PragmaTwice <twice at apache.org>
Date: Sat, 1 Aug 2026 19:05:05 +0800
Subject: [PATCH 1/4] [MLIR][CAPI][Python] Add support for querying memory
effect instances
---
mlir/include/mlir-c/Dialect/Transform.h | 15 +-
mlir/include/mlir-c/Interfaces.h | 71 ++++++-
.../mlir/Bindings/Python/IRInterfaces.h | 26 +--
mlir/include/mlir/CAPI/Interfaces.h | 3 -
mlir/lib/Bindings/Python/DialectTransform.cpp | 62 ++++---
mlir/lib/Bindings/Python/IRInterfaces.cpp | 175 +++++++++++++++---
mlir/lib/CAPI/Dialect/Transform.cpp | 54 ++++--
mlir/lib/CAPI/Interfaces/Interfaces.cpp | 79 +++++++-
mlir/python/mlir/dialects/ext.py | 4 +-
mlir/test/CAPI/CMakeLists.txt | 1 +
mlir/test/CAPI/ir.c | 91 ++++++++-
mlir/test/CAPI/transform.c | 46 +++++
mlir/test/python/dialects/ext.py | 4 +-
.../dialects/memory_effects_op_interface.py | 131 +++++++++----
.../python/dialects/transform_op_interface.py | 20 +-
15 files changed, 634 insertions(+), 148 deletions(-)
diff --git a/mlir/include/mlir-c/Dialect/Transform.h b/mlir/include/mlir-c/Dialect/Transform.h
index cbda09cdbc37b..83f1ed7dc37e9 100644
--- a/mlir/include/mlir-c/Dialect/Transform.h
+++ b/mlir/include/mlir-c/Dialect/Transform.h
@@ -246,25 +246,30 @@ MLIR_CAPI_EXPORTED void mlirPatternDescriptorOpInterfaceAttachFallbackModel(
/// Helper to mark operands as only reading handles.
MLIR_CAPI_EXPORTED void
mlirTransformOnlyReadsHandle(MlirOpOperand *operands, intptr_t numOperands,
- MlirMemoryEffectInstancesList effects);
+ MlirMemoryEffectInstancesCallback callback,
+ void *userData);
/// Helper to mark operands as consuming handles.
MLIR_CAPI_EXPORTED void
mlirTransformConsumesHandle(MlirOpOperand *operands, intptr_t numOperands,
- MlirMemoryEffectInstancesList effects);
+ MlirMemoryEffectInstancesCallback callback,
+ void *userData);
/// Helper to mark results as producing handles.
MLIR_CAPI_EXPORTED void
mlirTransformProducesHandle(MlirValue *results, intptr_t numResults,
- MlirMemoryEffectInstancesList effects);
+ MlirMemoryEffectInstancesCallback callback,
+ void *userData);
/// Helper to mark potential modifications to the payload IR.
MLIR_CAPI_EXPORTED void
-mlirTransformModifiesPayload(MlirMemoryEffectInstancesList effects);
+mlirTransformModifiesPayload(MlirMemoryEffectInstancesCallback callback,
+ void *userData);
/// Helper to mark potential reads from the payload IR.
MLIR_CAPI_EXPORTED void
-mlirTransformOnlyReadsPayload(MlirMemoryEffectInstancesList effects);
+mlirTransformOnlyReadsPayload(MlirMemoryEffectInstancesCallback callback,
+ void *userData);
#ifdef __cplusplus
}
diff --git a/mlir/include/mlir-c/Interfaces.h b/mlir/include/mlir-c/Interfaces.h
index a416a6ab76f87..aa03b7f9e7b66 100644
--- a/mlir/include/mlir-c/Interfaces.h
+++ b/mlir/include/mlir-c/Interfaces.h
@@ -30,7 +30,6 @@ extern "C" {
DEFINE_C_API_STRUCT(MlirMemoryEffect, void);
DEFINE_C_API_STRUCT(MlirMemoryEffectInstance, void);
-DEFINE_C_API_STRUCT(MlirMemoryEffectInstancesList, void);
DEFINE_C_API_STRUCT(MlirSideEffectResource, void);
#undef DEFINE_C_API_STRUCT
@@ -159,6 +158,10 @@ MLIR_CAPI_EXPORTED MlirMemoryEffect mlirMemoryEffectsReadGet(void);
/// Returns the borrowed singleton instance of the write memory effect.
MLIR_CAPI_EXPORTED MlirMemoryEffect mlirMemoryEffectsWriteGet(void);
+/// Returns the TypeID identifying the concrete type of the given memory effect.
+MLIR_CAPI_EXPORTED MlirTypeID
+mlirMemoryEffectGetEffectID(MlirMemoryEffect effect);
+
/// Returns the borrowed singleton instance of the default side effect
/// resource.
MLIR_CAPI_EXPORTED MlirSideEffectResource
@@ -212,15 +215,53 @@ mlirMemoryEffectInstanceCreateForSymbol(MlirMemoryEffect effect,
bool effectOnFullRegion,
MlirSideEffectResource resource);
-/// Destroys a memory effect instance created by one of the functions above.
+/// Destroys an owned memory effect instance created or cloned by this API.
MLIR_CAPI_EXPORTED void
mlirMemoryEffectInstanceDestroy(MlirMemoryEffectInstance instance);
-/// Appends a copy of `instance` to the given list. This does not take ownership
-/// of `instance`; the caller remains responsible for destroying it.
-MLIR_CAPI_EXPORTED void
-mlirMemoryEffectInstancesListAppend(MlirMemoryEffectInstancesList list,
- MlirMemoryEffectInstance instance);
+/// Creates an owned copy of a memory effect instance. The caller must destroy
+/// the returned instance with `mlirMemoryEffectInstanceDestroy`.
+MLIR_CAPI_EXPORTED MlirMemoryEffectInstance
+mlirMemoryEffectInstanceClone(MlirMemoryEffectInstance instance);
+
+/// Returns the memory effect of the given instance.
+MLIR_CAPI_EXPORTED MlirMemoryEffect
+mlirMemoryEffectInstanceGetEffect(MlirMemoryEffectInstance instance);
+
+/// Returns the side effect resource of the given instance.
+MLIR_CAPI_EXPORTED MlirSideEffectResource
+mlirMemoryEffectInstanceGetResource(MlirMemoryEffectInstance instance);
+
+/// Returns the stage of the given instance.
+MLIR_CAPI_EXPORTED int
+mlirMemoryEffectInstanceGetStage(MlirMemoryEffectInstance instance);
+
+/// Returns true if the given instance has effect on every single value of
+/// the resource.
+MLIR_CAPI_EXPORTED bool mlirMemoryEffectInstanceGetEffectOnFullRegion(
+ MlirMemoryEffectInstance instance);
+
+/// Returns the parameters of the given instance, or a null attribute if there
+/// are no parameters.
+MLIR_CAPI_EXPORTED MlirAttribute
+mlirMemoryEffectInstanceGetParameters(MlirMemoryEffectInstance instance);
+
+/// Returns the value (OpOperand, OpResult, or BlockArgument) of the given
+/// instance, or a null value if there is no associated value.
+MLIR_CAPI_EXPORTED MlirValue
+mlirMemoryEffectInstanceGetValue(MlirMemoryEffectInstance instance);
+
+/// Returns the symbol reference of the given instance, or a null attribute if
+/// there is no associated symbol.
+MLIR_CAPI_EXPORTED MlirAttribute
+mlirMemoryEffectInstanceGetSymbolRef(MlirMemoryEffectInstance instance);
+
+/// Callback used to return a list of memory effect instances. `effects` points
+/// to `numEffects` consecutive borrowed instances that are only valid for the
+/// duration of the callback. The caller-provided `userData` is forwarded to
+/// the callback.
+typedef void (*MlirMemoryEffectInstancesCallback)(
+ intptr_t numEffects, MlirMemoryEffectInstance *effects, void *userData);
/// Returns the interface TypeID of the MemoryEffectsOpInterface.
MLIR_CAPI_EXPORTED MlirTypeID mlirMemoryEffectsOpInterfaceTypeID(void);
@@ -231,9 +272,12 @@ typedef struct {
void (*construct)(void *userData);
/// Optional destructor for user data. Set to nullptr to disable it.
void (*destruct)(void *userData);
- /// Get memory effects callback.
- void (*getEffects)(MlirOperation op, MlirMemoryEffectInstancesList effects,
- void *userData);
+ /// Get memory effects callback. Implementations return effects by invoking
+ /// `callback` with an array of memory effect instances. The callback copies
+ /// the instances synchronously, so implementations retain ownership of them.
+ void (*getEffects)(MlirOperation op,
+ MlirMemoryEffectInstancesCallback callback,
+ void *callbackUserData, void *userData);
void *userData;
} MlirMemoryEffectsOpInterfaceCallbacks;
@@ -243,6 +287,13 @@ MLIR_CAPI_EXPORTED void mlirMemoryEffectsOpInterfaceAttachFallbackModel(
MlirContext ctx, MlirStringRef opName,
MlirMemoryEffectsOpInterfaceCallbacks callbacks);
+/// Gets the memory effects of the given operation. The operation must
+/// implement the MemoryEffectsOpInterface. Invokes `callback` once with all
+/// effects; the instances are borrowed and only valid during the callback.
+MLIR_CAPI_EXPORTED void mlirMemoryEffectsOpInterfaceGetEffects(
+ MlirOperation operation, MlirMemoryEffectInstancesCallback callback,
+ void *userData);
+
#ifdef __cplusplus
}
#endif
diff --git a/mlir/include/mlir/Bindings/Python/IRInterfaces.h b/mlir/include/mlir/Bindings/Python/IRInterfaces.h
index c1047c7e20ab7..ab8c70c4a4c82 100644
--- a/mlir/include/mlir/Bindings/Python/IRInterfaces.h
+++ b/mlir/include/mlir/Bindings/Python/IRInterfaces.h
@@ -163,6 +163,12 @@ class PyMemoryEffectInstance {
public:
explicit PyMemoryEffectInstance(MlirMemoryEffectInstance instance)
: instance(instance) {}
+ PyMemoryEffectInstance(const PyMemoryEffect &effect,
+ const nanobind::object &target,
+ const nanobind::object ¶meters, int stage,
+ bool effectOnFullRegion,
+ const PySideEffectResource &resource);
+ PyMemoryEffectInstance(const PyMemoryEffectInstance &) = delete;
PyMemoryEffectInstance(PyMemoryEffectInstance &&other) noexcept
: instance(other.instance) {
other.instance.ptr = nullptr;
@@ -173,24 +179,18 @@ class PyMemoryEffectInstance {
}
MlirMemoryEffectInstance get() const { return instance; }
+ PyMemoryEffect getEffect() const;
+ PySideEffectResource getResource() const;
+ int getStage() const;
+ bool getEffectOnFullRegion() const;
+ nanobind::object getParameters() const;
+ nanobind::object getValue() const;
+ nanobind::object getSymbolRef() const;
private:
MlirMemoryEffectInstance instance;
};
-/// A callback-scoped view of a list of memory effect instances.
-class PyMemoryEffectsInstanceList {
-public:
- explicit PyMemoryEffectsInstanceList(MlirMemoryEffectInstancesList effects)
- : effects(effects) {}
-
- MlirMemoryEffectInstancesList get() const { return effects; }
- operator MlirMemoryEffectInstancesList() const { return effects; }
-
-private:
- MlirMemoryEffectInstancesList effects;
-};
-
} // namespace MLIR_BINDINGS_PYTHON_DOMAIN
} // namespace python
} // namespace mlir
diff --git a/mlir/include/mlir/CAPI/Interfaces.h b/mlir/include/mlir/CAPI/Interfaces.h
index 55d850ccd01eb..1d18f6f7d64a8 100644
--- a/mlir/include/mlir/CAPI/Interfaces.h
+++ b/mlir/include/mlir/CAPI/Interfaces.h
@@ -19,9 +19,6 @@
#include "mlir/CAPI/Wrap.h"
#include "mlir/Interfaces/SideEffectInterfaces.h"
-DEFINE_C_API_PTR_METHODS(
- MlirMemoryEffectInstancesList,
- llvm::SmallVectorImpl<mlir::MemoryEffects::EffectInstance>)
DEFINE_C_API_PTR_METHODS(MlirMemoryEffect, mlir::MemoryEffects::Effect)
DEFINE_C_API_PTR_METHODS(MlirMemoryEffectInstance,
mlir::MemoryEffects::EffectInstance)
diff --git a/mlir/lib/Bindings/Python/DialectTransform.cpp b/mlir/lib/Bindings/Python/DialectTransform.cpp
index a4a91b28522cf..dd8c2d711edb2 100644
--- a/mlir/lib/Bindings/Python/DialectTransform.cpp
+++ b/mlir/lib/Bindings/Python/DialectTransform.cpp
@@ -493,36 +493,55 @@ struct ParamType : PyConcreteType<ParamType> {
//===----------------------------------------------------------------------===//
namespace {
-void onlyReadsHandle(nb::iterable &operands,
- const PyMemoryEffectsInstanceList &effects) {
+void collectMemoryEffectInstances(intptr_t numEffects,
+ MlirMemoryEffectInstance *effects,
+ void *userData) {
+ auto *result = static_cast<std::vector<PyMemoryEffectInstance> *>(userData);
+ result->reserve(result->size() + numEffects);
+ for (intptr_t i = 0; i < numEffects; ++i)
+ result->emplace_back(mlirMemoryEffectInstanceClone(effects[i]));
+}
+
+std::vector<PyMemoryEffectInstance> onlyReadsHandle(nb::iterable &operands) {
std::vector<MlirOpOperand> operandsVec;
for (auto operand : operands)
operandsVec.push_back(nb::cast<PyOpOperand>(operand));
- mlirTransformOnlyReadsHandle(operandsVec.data(), operandsVec.size(), effects);
+ std::vector<PyMemoryEffectInstance> effects;
+ mlirTransformOnlyReadsHandle(operandsVec.data(), operandsVec.size(),
+ collectMemoryEffectInstances, &effects);
+ return effects;
};
-void consumesHandle(nb::iterable &operands,
- const PyMemoryEffectsInstanceList &effects) {
+std::vector<PyMemoryEffectInstance> consumesHandle(nb::iterable &operands) {
std::vector<MlirOpOperand> operandsVec;
for (auto operand : operands)
operandsVec.push_back(nb::cast<PyOpOperand>(operand));
- mlirTransformConsumesHandle(operandsVec.data(), operandsVec.size(), effects);
+ std::vector<PyMemoryEffectInstance> effects;
+ mlirTransformConsumesHandle(operandsVec.data(), operandsVec.size(),
+ collectMemoryEffectInstances, &effects);
+ return effects;
};
-void producesHandle(nb::iterable &results,
- const PyMemoryEffectsInstanceList &effects) {
+std::vector<PyMemoryEffectInstance> producesHandle(nb::iterable &results) {
std::vector<MlirValue> resultsVec;
for (auto result : results)
resultsVec.push_back(nb::cast<PyOpResult>(result).get());
- mlirTransformProducesHandle(resultsVec.data(), resultsVec.size(), effects);
+ std::vector<PyMemoryEffectInstance> effects;
+ mlirTransformProducesHandle(resultsVec.data(), resultsVec.size(),
+ collectMemoryEffectInstances, &effects);
+ return effects;
};
-void modifiesPayload(const PyMemoryEffectsInstanceList &effects) {
- mlirTransformModifiesPayload(effects);
+std::vector<PyMemoryEffectInstance> modifiesPayload() {
+ std::vector<PyMemoryEffectInstance> effects;
+ mlirTransformModifiesPayload(collectMemoryEffectInstances, &effects);
+ return effects;
}
-void onlyReadsPayload(const PyMemoryEffectsInstanceList &effects) {
- mlirTransformOnlyReadsPayload(effects);
+std::vector<PyMemoryEffectInstance> onlyReadsPayload() {
+ std::vector<PyMemoryEffectInstance> effects;
+ mlirTransformOnlyReadsPayload(collectMemoryEffectInstances, &effects);
+ return effects;
}
} // namespace
@@ -546,21 +565,22 @@ static void populateDialectTransformSubmodule(nb::module_ &m) {
PyPatternDescriptorOpInterface::bind(m);
m.def("only_reads_handle", onlyReadsHandle,
- "Mark operands as only reading handles.", nb::arg("operands"),
- nb::arg("effects"));
+ "Returns effects marking operands as only reading handles.",
+ nb::arg("operands"));
m.def("consumes_handle", consumesHandle,
- "Mark operands as consuming handles.", nb::arg("operands"),
- nb::arg("effects"));
+ "Returns effects marking operands as consuming handles.",
+ nb::arg("operands"));
- m.def("produces_handle", producesHandle, "Mark results as producing handles.",
- nb::arg("results"), nb::arg("effects"));
+ m.def("produces_handle", producesHandle,
+ "Returns effects marking results as producing handles.",
+ nb::arg("results"));
m.def("modifies_payload", modifiesPayload,
- "Mark the transform as modifying the payload.", nb::arg("effects"));
+ "Returns effects marking potential payload modifications.");
m.def("only_reads_payload", onlyReadsPayload,
- "Mark the transform as only reading the payload.", nb::arg("effects"));
+ "Returns effects marking payload reads.");
}
} // namespace transform
} // namespace MLIR_BINDINGS_PYTHON_DOMAIN
diff --git a/mlir/lib/Bindings/Python/IRInterfaces.cpp b/mlir/lib/Bindings/Python/IRInterfaces.cpp
index d16015acaf3dc..762dffc827043 100644
--- a/mlir/lib/Bindings/Python/IRInterfaces.cpp
+++ b/mlir/lib/Bindings/Python/IRInterfaces.cpp
@@ -44,13 +44,10 @@ MlirAttribute unwrapOptionalAttribute(const nb::object &attribute) {
return pyAttribute->get();
}
-void appendMemoryEffectInstance(PyMemoryEffectsInstanceList &effects,
- const PyMemoryEffect &effect,
- const nb::object &target,
- const nb::object ¶meters, int stage,
- bool effectOnFullRegion,
- const PySideEffectResource &resource) {
- MlirMemoryEffectInstancesList list = effects.get();
+MlirMemoryEffectInstance createMemoryEffectInstance(
+ const PyMemoryEffect &effect, const nb::object &target,
+ const nb::object ¶meters, int stage, bool effectOnFullRegion,
+ const PySideEffectResource &resource) {
MlirAttribute unwrappedParameters = unwrapOptionalAttribute(parameters);
MlirMemoryEffectInstance rawInstance{nullptr};
@@ -93,9 +90,7 @@ void appendMemoryEffectInstance(PyMemoryEffectsInstanceList &effects,
"SymbolRefAttr, or None");
}
}
-
- PyMemoryEffectInstance instance(rawInstance);
- mlirMemoryEffectInstancesListAppend(list, instance.get());
+ return rawInstance;
}
/// Takes in an optional ist of operands and converts them into a std::vector
@@ -176,6 +171,61 @@ wrapRegions(std::optional<std::vector<PyRegion>> regions) {
} // namespace
+PyMemoryEffectInstance::PyMemoryEffectInstance(
+ const PyMemoryEffect &effect, const nb::object &target,
+ const nb::object ¶meters, int stage, bool effectOnFullRegion,
+ const PySideEffectResource &resource)
+ : PyMemoryEffectInstance(createMemoryEffectInstance(
+ effect, target, parameters, stage, effectOnFullRegion, resource)) {}
+
+PyMemoryEffect PyMemoryEffectInstance::getEffect() const {
+ return PyMemoryEffect(mlirMemoryEffectInstanceGetEffect(instance));
+}
+
+PySideEffectResource PyMemoryEffectInstance::getResource() const {
+ return PySideEffectResource(mlirMemoryEffectInstanceGetResource(instance));
+}
+
+int PyMemoryEffectInstance::getStage() const {
+ return mlirMemoryEffectInstanceGetStage(instance);
+}
+
+bool PyMemoryEffectInstance::getEffectOnFullRegion() const {
+ return mlirMemoryEffectInstanceGetEffectOnFullRegion(instance);
+}
+
+nb::object PyMemoryEffectInstance::getParameters() const {
+ MlirAttribute parameters = mlirMemoryEffectInstanceGetParameters(instance);
+ if (mlirAttributeIsNull(parameters))
+ return nb::none();
+ PyMlirContextRef context =
+ PyMlirContext::forContext(mlirAttributeGetContext(parameters));
+ return PyAttribute(context, parameters).maybeDownCast();
+}
+
+nb::object PyMemoryEffectInstance::getValue() const {
+ MlirValue value = mlirMemoryEffectInstanceGetValue(instance);
+ if (mlirValueIsNull(value))
+ return nb::none();
+ MlirOperation owner =
+ mlirValueIsAOpResult(value)
+ ? mlirOpResultGetOwner(value)
+ : mlirBlockGetParentOperation(mlirBlockArgumentGetOwner(value));
+ PyMlirContextRef context =
+ PyMlirContext::forContext(mlirOperationGetContext(owner));
+ return PyValue(PyOperation::forOperation(context, owner), value)
+ .maybeDownCast();
+}
+
+nb::object PyMemoryEffectInstance::getSymbolRef() const {
+ MlirAttribute symbol = mlirMemoryEffectInstanceGetSymbolRef(instance);
+ if (mlirAttributeIsNull(symbol))
+ return nb::none();
+ PyMlirContextRef context =
+ PyMlirContext::forContext(mlirAttributeGetContext(symbol));
+ return PyAttribute(context, symbol).maybeDownCast();
+}
+
/// Python wrapper for InferTypeOpInterface. This interface has only static
/// methods.
class PyInferTypeOpInterface
@@ -494,22 +544,36 @@ class PyMemoryEffectsOpInterface
nb::handle(static_cast<PyObject *>(userData)).dec_ref();
};
callbacks.getEffects = [](MlirOperation op,
- MlirMemoryEffectInstancesList effects,
- void *userData) {
+ MlirMemoryEffectInstancesCallback callback,
+ void *callbackUserData, void *userData) {
nb::handle pyClass(static_cast<PyObject *>(userData));
// Get the 'get_effects' method from the Python class.
auto pyGetEffects =
nb::cast<nb::callable>(nb::getattr(pyClass, "get_effects"));
- PyMemoryEffectsInstanceList effectsWrapper{effects};
-
PyMlirContextRef context =
PyMlirContext::forContext(mlirOperationGetContext(op));
auto opview = PyOperation::forOperation(context, op)->createOpView();
- // Invoke `pyClass.get_effects(op, effects)`.
- pyGetEffects(opview, effectsWrapper);
+ // Invoke `pyClass.get_effects(op)` and pass the resulting instances back
+ // to the C++ interface as a borrowed array.
+ nb::object result = pyGetEffects(opview);
+ nb::iterable iterable;
+ if (!nb::try_cast<nb::iterable>(result, iterable))
+ throw nb::type_error("get_effects must return an iterable");
+
+ std::vector<nb::object> effectObjects;
+ std::vector<MlirMemoryEffectInstance> effects;
+ for (nb::handle object : iterable) {
+ PyMemoryEffectInstance *effect = nullptr;
+ if (!nb::try_cast<PyMemoryEffectInstance *>(object, effect) || !effect)
+ throw nb::type_error(
+ "get_effects must return MemoryEffectInstance objects");
+ effectObjects.push_back(nb::borrow<nb::object>(object));
+ effects.push_back(effect->get());
+ }
+ callback(effects.size(), effects.data(), callbackUserData);
};
mlirMemoryEffectsOpInterfaceAttachFallbackModel(
@@ -517,7 +581,33 @@ class PyMemoryEffectsOpInterface
callbacks);
}
+ std::vector<PyMemoryEffectInstance> getEffects() {
+ if (isStatic())
+ throw nb::type_error("Cannot query effects on a static interface");
+
+ auto operationObject = getOperationObject();
+ auto *operation = nb::cast<PyOperation *>(operationObject);
+ std::vector<PyMemoryEffectInstance> effects;
+
+ mlirMemoryEffectsOpInterfaceGetEffects(
+ operation->get(),
+ [](intptr_t numEffects, MlirMemoryEffectInstance *effects,
+ void *userData) {
+ auto *result =
+ static_cast<std::vector<PyMemoryEffectInstance> *>(userData);
+ result->reserve(result->size() + numEffects);
+ for (intptr_t i = 0; i < numEffects; ++i) {
+ result->emplace_back(mlirMemoryEffectInstanceClone(effects[i]));
+ }
+ },
+ &effects);
+ return effects;
+ }
+
static void bindDerived(ClassTy &cls) {
+ cls.def("get_effects", &PyMemoryEffectsOpInterface::getEffects,
+ nb::sig("def get_effects(self) -> list[MemoryEffectInstance]"),
+ "Returns the memory effects of the operation.");
cls.attr("attach") = classmethod(
[](const nb::object &cls, const nb::object &opName, nb::object target,
DefaultingPyMlirContext context) {
@@ -539,6 +629,13 @@ void populateIRInterfaces(nb::module_ &m) {
.value("RecursivelySpeculatable",
MlirSpeculatabilityRecursivelySpeculatable);
nb::class_<PyMemoryEffect>(m, "MemoryEffect", "A memory effect.")
+ .def(
+ "__eq__",
+ [](const PyMemoryEffect &self, const PyMemoryEffect &other) {
+ return mlirTypeIDEqual(mlirMemoryEffectGetEffectID(self.get()),
+ mlirMemoryEffectGetEffectID(other.get()));
+ },
+ nb::is_operator(), "Compares two memory effects for equality.")
.def_prop_ro_static("Allocate",
[](nb::object & /*class*/) {
return PyMemoryEffect(
@@ -562,22 +659,44 @@ void populateIRInterfaces(nb::module_ &m) {
return PySideEffectResource(mlirSideEffectsDefaultResourceGet());
});
- nb::class_<PyMemoryEffectsInstanceList>(
- m, "MemoryEffectInstancesList",
- "A memory effect list that is valid only during get_effects.")
- .def("append", &appendMemoryEffectInstance, nb::arg("effect"),
- nb::arg("target").none() = nb::none(), nb::kw_only(),
- nb::arg("parameters").none() = nb::none(), nb::arg("stage") = 0,
- nb::arg("effect_on_full_region") = false,
+ nb::class_<PyMemoryEffectInstance>(m, "MemoryEffectInstance",
+ "A concrete instance of a memory effect.")
+ .def(nb::init<const PyMemoryEffect &, const nb::object &,
+ const nb::object &, int, bool,
+ const PySideEffectResource &>(),
+ nb::arg("effect"), nb::arg("target").none() = nb::none(),
+ nb::kw_only(), nb::arg("parameters").none() = nb::none(),
+ nb::arg("stage") = 0, nb::arg("effect_on_full_region") = false,
nb::arg("resource") =
PySideEffectResource(mlirSideEffectsDefaultResourceGet()),
- nb::sig("def append(self, effect: MemoryEffect, target: OpOperand | "
- "OpResult | BlockArgument | SymbolRefAttr | None = None, *, "
- "parameters: Attribute | None = None, stage: int = 0, "
+ nb::sig("def __init__(self, effect: MemoryEffect, target: "
+ "OpOperand | OpResult | BlockArgument | SymbolRefAttr | "
+ "FlatSymbolRefAttr | None = None, *, parameters: Attribute "
+ "| None = None, stage: int = 0, "
"effect_on_full_region: bool = False, resource: "
"SideEffectResource = ...) -> None"),
- "Append a memory effect instance. The target may be an OpOperand, "
- "OpResult, BlockArgument, SymbolRefAttr, or None.");
+ "Creates a memory effect instance. The target may be an OpOperand, "
+ "OpResult, BlockArgument, SymbolRefAttr, or None.")
+ .def_prop_ro("effect", &PyMemoryEffectInstance::getEffect,
+ "Returns the kind of memory effect.")
+ .def_prop_ro("resource", &PyMemoryEffectInstance::getResource,
+ "Returns the affected side effect resource.")
+ .def_prop_ro("stage", &PyMemoryEffectInstance::getStage,
+ "Returns the stage at which the effect occurs.")
+ .def_prop_ro("effect_on_full_region",
+ &PyMemoryEffectInstance::getEffectOnFullRegion,
+ "Returns whether the effect applies to the full resource.")
+ .def_prop_ro("parameters", &PyMemoryEffectInstance::getParameters,
+ nb::sig("def parameters(self) -> Attribute | None"),
+ "Returns the effect parameters, if any.")
+ .def_prop_ro(
+ "value", &PyMemoryEffectInstance::getValue,
+ nb::sig("def value(self) -> OpResult | BlockArgument | None"),
+ "Returns the affected value, if any.")
+ .def_prop_ro("symbol_ref", &PyMemoryEffectInstance::getSymbolRef,
+ nb::sig("def symbol_ref(self) -> SymbolRefAttr | "
+ "FlatSymbolRefAttr | None"),
+ "Returns the affected symbol reference, if any.");
PyConditionallySpeculatableOpInterface::bind(m);
PyInferShapedTypeOpInterface::bind(m);
diff --git a/mlir/lib/CAPI/Dialect/Transform.cpp b/mlir/lib/CAPI/Dialect/Transform.cpp
index 1ed14255bf5e0..ea8c2eaef8212 100644
--- a/mlir/lib/CAPI/Dialect/Transform.cpp
+++ b/mlir/lib/CAPI/Dialect/Transform.cpp
@@ -385,39 +385,67 @@ void mlirPatternDescriptorOpInterfaceAttachFallbackModel(
// MemoryEffectsOpInterface helpers
//===---------------------------------------------------------------------===//
+static void invokeMemoryEffectInstancesCallback(
+ SmallVectorImpl<MemoryEffects::EffectInstance> &effects,
+ MlirMemoryEffectInstancesCallback callback, void *userData) {
+ SmallVector<MlirMemoryEffectInstance> wrappedEffects;
+ wrappedEffects.reserve(effects.size());
+ for (MemoryEffects::EffectInstance &effect : effects)
+ wrappedEffects.push_back(wrap(&effect));
+ callback(wrappedEffects.size(), wrappedEffects.data(), userData);
+}
+
/// Set the effect for the operands to only read the transform handles.
void mlirTransformOnlyReadsHandle(MlirOpOperand *operands, intptr_t numOperands,
- MlirMemoryEffectInstancesList effects) {
- MutableArrayRef<OpOperand> operandArray(unwrap(*operands), numOperands);
- transform::onlyReadsHandle(operandArray, *unwrap(effects));
+ MlirMemoryEffectInstancesCallback callback,
+ void *userData) {
+ MutableArrayRef<OpOperand> operandArray;
+ if (numOperands != 0)
+ operandArray = MutableArrayRef<OpOperand>(unwrap(*operands), numOperands);
+ SmallVector<MemoryEffects::EffectInstance> effects;
+ transform::onlyReadsHandle(operandArray, effects);
+ invokeMemoryEffectInstancesCallback(effects, callback, userData);
}
/// Set the effect for the operands to consuming the transform handles.
void mlirTransformConsumesHandle(MlirOpOperand *operands, intptr_t numOperands,
- MlirMemoryEffectInstancesList effects) {
- MutableArrayRef<OpOperand> operandArray(unwrap(*operands), numOperands);
- transform::consumesHandle(operandArray, *unwrap(effects));
+ MlirMemoryEffectInstancesCallback callback,
+ void *userData) {
+ MutableArrayRef<OpOperand> operandArray;
+ if (numOperands != 0)
+ operandArray = MutableArrayRef<OpOperand>(unwrap(*operands), numOperands);
+ SmallVector<MemoryEffects::EffectInstance> effects;
+ transform::consumesHandle(operandArray, effects);
+ invokeMemoryEffectInstancesCallback(effects, callback, userData);
}
/// Set the effect for the results to that they produce transform handles.
void mlirTransformProducesHandle(MlirValue *results, intptr_t numResults,
- MlirMemoryEffectInstancesList effects) {
+ MlirMemoryEffectInstancesCallback callback,
+ void *userData) {
// NB: calling `producesHandle()` `numResults` as we cannot cast array of
// `OpResult`s to a single `ResultRange` (and neither is `ResultRange` exposed
// to Python). `producesHandle` iterates over the given `ResultRange` anyway.
- SmallVectorImpl<MemoryEffects::EffectInstance> &effectList = *unwrap(effects);
+ SmallVector<MemoryEffects::EffectInstance> effects;
for (intptr_t i = 0; i < numResults; ++i) {
auto opResult = cast<OpResult>(unwrap(results[i]));
- transform::producesHandle(ResultRange(opResult), effectList);
+ transform::producesHandle(ResultRange(opResult), effects);
}
+ invokeMemoryEffectInstancesCallback(effects, callback, userData);
}
/// Set the effect of potentially modifying payload IR.
-void mlirTransformModifiesPayload(MlirMemoryEffectInstancesList effects) {
- transform::modifiesPayload(*unwrap(effects));
+void mlirTransformModifiesPayload(MlirMemoryEffectInstancesCallback callback,
+ void *userData) {
+ SmallVector<MemoryEffects::EffectInstance> effects;
+ transform::modifiesPayload(effects);
+ invokeMemoryEffectInstancesCallback(effects, callback, userData);
}
/// Set the effect of potentially reading payload IR.
-void mlirTransformOnlyReadsPayload(MlirMemoryEffectInstancesList effects) {
- transform::onlyReadsPayload(*unwrap(effects));
+void mlirTransformOnlyReadsPayload(MlirMemoryEffectInstancesCallback callback,
+ void *userData) {
+ SmallVector<MemoryEffects::EffectInstance> effects;
+ transform::onlyReadsPayload(effects);
+ invokeMemoryEffectInstancesCallback(effects, callback, userData);
}
diff --git a/mlir/lib/CAPI/Interfaces/Interfaces.cpp b/mlir/lib/CAPI/Interfaces/Interfaces.cpp
index 1bf8dffa9c431..65bf80fdcb476 100644
--- a/mlir/lib/CAPI/Interfaces/Interfaces.cpp
+++ b/mlir/lib/CAPI/Interfaces/Interfaces.cpp
@@ -16,6 +16,7 @@
#include "mlir/CAPI/Wrap.h"
#include "mlir/IR/ValueRange.h"
#include "mlir/Interfaces/InferTypeOpInterface.h"
+#include "mlir/Interfaces/SideEffectInterfaces.h"
#include "llvm/ADT/ScopeExit.h"
#include <optional>
@@ -68,6 +69,26 @@ SmallVector<std::unique_ptr<Region>> unwrapRegions(intptr_t nRegions,
return unwrappedRegions;
}
+void invokeMemoryEffectInstancesCallback(
+ SmallVectorImpl<MemoryEffects::EffectInstance> &effects,
+ MlirMemoryEffectInstancesCallback callback, void *userData) {
+ SmallVector<MlirMemoryEffectInstance> wrappedEffects;
+ wrappedEffects.reserve(effects.size());
+ for (MemoryEffects::EffectInstance &effect : effects)
+ wrappedEffects.push_back(wrap(&effect));
+ callback(wrappedEffects.size(), wrappedEffects.data(), userData);
+}
+
+void appendMemoryEffectInstances(intptr_t numEffects,
+ MlirMemoryEffectInstance *effects,
+ void *userData) {
+ auto *unwrappedEffects =
+ static_cast<SmallVectorImpl<MemoryEffects::EffectInstance> *>(userData);
+ unwrappedEffects->reserve(unwrappedEffects->size() + numEffects);
+ for (intptr_t i = 0; i < numEffects; ++i)
+ unwrappedEffects->push_back(*unwrap(effects[i]));
+}
+
} // namespace
bool mlirOperationImplementsInterface(MlirOperation operation,
@@ -298,6 +319,10 @@ MlirMemoryEffect mlirMemoryEffectsWriteGet() {
static_cast<MemoryEffects::Effect *>(MemoryEffects::Write::get()));
}
+MlirTypeID mlirMemoryEffectGetEffectID(MlirMemoryEffect effect) {
+ return wrap(unwrap(effect)->getEffectID());
+}
+
MlirSideEffectResource mlirSideEffectsDefaultResourceGet() {
return wrap(static_cast<SideEffects::Resource *>(
SideEffects::DefaultResource::get()));
@@ -347,9 +372,42 @@ void mlirMemoryEffectInstanceDestroy(MlirMemoryEffectInstance instance) {
delete unwrap(instance);
}
-void mlirMemoryEffectInstancesListAppend(MlirMemoryEffectInstancesList list,
- MlirMemoryEffectInstance instance) {
- unwrap(list)->push_back(*unwrap(instance));
+MlirMemoryEffectInstance
+mlirMemoryEffectInstanceClone(MlirMemoryEffectInstance instance) {
+ return wrap(new MemoryEffects::EffectInstance(*unwrap(instance)));
+}
+
+MlirMemoryEffect
+mlirMemoryEffectInstanceGetEffect(MlirMemoryEffectInstance instance) {
+ return wrap(unwrap(instance)->getEffect());
+}
+
+MlirSideEffectResource
+mlirMemoryEffectInstanceGetResource(MlirMemoryEffectInstance instance) {
+ return wrap(unwrap(instance)->getResource());
+}
+
+int mlirMemoryEffectInstanceGetStage(MlirMemoryEffectInstance instance) {
+ return unwrap(instance)->getStage();
+}
+
+bool mlirMemoryEffectInstanceGetEffectOnFullRegion(
+ MlirMemoryEffectInstance instance) {
+ return unwrap(instance)->getEffectOnFullRegion();
+}
+
+MlirAttribute
+mlirMemoryEffectInstanceGetParameters(MlirMemoryEffectInstance instance) {
+ return wrap(unwrap(instance)->getParameters());
+}
+
+MlirValue mlirMemoryEffectInstanceGetValue(MlirMemoryEffectInstance instance) {
+ return wrap(unwrap(instance)->getValue());
+}
+
+MlirAttribute
+mlirMemoryEffectInstanceGetSymbolRef(MlirMemoryEffectInstance instance) {
+ return wrap(unwrap(instance)->getSymbolRef());
}
MlirTypeID mlirMemoryEffectsOpInterfaceTypeID() {
@@ -389,8 +447,8 @@ class MemoryEffectOpInterfaceFallbackModel
getEffects(Operation *op,
SmallVectorImpl<MemoryEffects::EffectInstance> &effects) const {
assert(callbacks.getEffects && "getEffects callback not set");
- MlirMemoryEffectInstancesList cEffects = wrap(&effects);
- callbacks.getEffects(wrap(op), cEffects, callbacks.userData);
+ callbacks.getEffects(wrap(op), appendMemoryEffectInstances, &effects,
+ callbacks.userData);
}
private:
@@ -417,3 +475,14 @@ void mlirMemoryEffectsOpInterfaceAttachFallbackModel(
assert(model && "Failed to get MemoryEffectOpInterfaceFallbackModel");
model->setCallbacks(callbacks);
}
+
+void mlirMemoryEffectsOpInterfaceGetEffects(
+ MlirOperation operation, MlirMemoryEffectInstancesCallback callback,
+ void *userData) {
+ auto iface = dyn_cast<MemoryEffectOpInterface>(unwrap(operation));
+ assert(iface && "operation does not implement MemoryEffectOpInterface");
+
+ SmallVector<MemoryEffects::EffectInstance> effects;
+ iface.getEffects(effects);
+ invokeMemoryEffectInstancesCallback(effects, callback, userData);
+}
diff --git a/mlir/python/mlir/dialects/ext.py b/mlir/python/mlir/dialects/ext.py
index 4e8d30d82d5f8..bd59772d8a9e1 100644
--- a/mlir/python/mlir/dialects/ext.py
+++ b/mlir/python/mlir/dialects/ext.py
@@ -1000,8 +1000,8 @@ class Pure:
class NoMemoryEffect(ir.MemoryEffectsOpInterface):
@staticmethod
- def get_effects(op, effects):
- pass
+ def get_effects(op):
+ return []
class AlwaysSpeculatable(ir.ConditionallySpeculatable):
@staticmethod
diff --git a/mlir/test/CAPI/CMakeLists.txt b/mlir/test/CAPI/CMakeLists.txt
index 3b3f15a825b38..211e321a9b763 100644
--- a/mlir/test/CAPI/CMakeLists.txt
+++ b/mlir/test/CAPI/CMakeLists.txt
@@ -110,6 +110,7 @@ _add_capi_test_executable(mlir-capi-transform-test
transform.c
LINK_LIBS PRIVATE
MLIRCAPIIR
+ MLIRCAPIInterfaces
MLIRCAPIRegisterEverything
MLIRCAPITransformDialect
)
diff --git a/mlir/test/CAPI/ir.c b/mlir/test/CAPI/ir.c
index f73153bc54129..3b00201349038 100644
--- a/mlir/test/CAPI/ir.c
+++ b/mlir/test/CAPI/ir.c
@@ -2543,6 +2543,28 @@ static MlirSpeculatability conditionallySpeculatableCallback(MlirOperation op,
return MlirSpeculatabilityRecursivelySpeculatable;
}
+typedef struct {
+ intptr_t callbackCount;
+ intptr_t effectCount;
+} MemoryEffectsCallbackData;
+
+static void memoryEffectsCallback(intptr_t numEffects,
+ MlirMemoryEffectInstance *effects,
+ void *userData) {
+ MemoryEffectsCallbackData *data = (MemoryEffectsCallbackData *)userData;
+ ++data->callbackCount;
+ data->effectCount += numEffects;
+ for (intptr_t i = 0; i < numEffects; ++i) {
+ MlirMemoryEffectInstance clone = mlirMemoryEffectInstanceClone(effects[i]);
+ assert(clone.ptr && "expected a cloned memory effect instance");
+ assert(mlirMemoryEffectInstanceGetEffect(clone).ptr &&
+ "expected a memory effect");
+ assert(mlirMemoryEffectInstanceGetResource(clone).ptr &&
+ "expected a side effect resource");
+ mlirMemoryEffectInstanceDestroy(clone);
+ }
+}
+
int testInterfaces(MlirContext ctx) {
// CHECK-LABEL: @testInterfaces
fprintf(stderr, "@testInterfaces\n");
@@ -2590,7 +2612,8 @@ int testInterfaces(MlirContext ctx) {
MlirOperationState storeState = mlirOperationStateGet(storeName, loc);
MlirValue constantResult = mlirOperationGetResult(constantOp, 0);
- mlirOperationStateAddOperands(&storeState, 1, &constantResult);
+ MlirValue storeOperands[] = {constantResult, constantResult};
+ mlirOperationStateAddOperands(&storeState, 2, storeOperands);
MlirOperation storeOp = mlirOperationCreate(&storeState);
if (mlirOperationImplementsInterface(storeOp, condSpecTypeID)) {
fprintf(stderr, "ERROR: Expected memref.store instance to not implement "
@@ -2621,6 +2644,23 @@ int testInterfaces(MlirContext ctx) {
// CHECK: memref.store speculatability: 2
// CHECK: callback count: 1
+ MlirTypeID memoryEffectsTypeID = mlirMemoryEffectsOpInterfaceTypeID();
+ if (!mlirOperationImplementsInterface(storeOp, memoryEffectsTypeID)) {
+ fprintf(
+ stderr,
+ "ERROR: Expected memref.store to implement MemoryEffectsOpInterface\n");
+ return 6;
+ }
+ MemoryEffectsCallbackData memoryEffectsData = {0};
+ mlirMemoryEffectsOpInterfaceGetEffects(storeOp, memoryEffectsCallback,
+ &memoryEffectsData);
+ fprintf(stderr, "memory effects callback count: %" PRIdPTR "\n",
+ memoryEffectsData.callbackCount);
+ fprintf(stderr, "memory effects count: %" PRIdPTR "\n",
+ memoryEffectsData.effectCount);
+ // CHECK: memory effects callback count: 1
+ // CHECK: memory effects count: 1
+
MlirMemoryEffect allocate = mlirMemoryEffectsAllocateGet();
MlirMemoryEffect free = mlirMemoryEffectsFreeGet();
MlirMemoryEffect read = mlirMemoryEffectsReadGet();
@@ -2632,6 +2672,27 @@ int testInterfaces(MlirContext ctx) {
return 6;
}
+ MlirTypeID effectIDs[] = {
+ mlirMemoryEffectGetEffectID(allocate),
+ mlirMemoryEffectGetEffectID(free),
+ mlirMemoryEffectGetEffectID(read),
+ mlirMemoryEffectGetEffectID(write),
+ };
+ for (intptr_t i = 0; i < 4; ++i) {
+ if (mlirTypeIDIsNull(effectIDs[i])) {
+ fprintf(stderr, "ERROR: Expected a non-null memory effect ID\n");
+ return 6;
+ }
+ for (intptr_t j = 0; j < i; ++j) {
+ if (mlirTypeIDEqual(effectIDs[i], effectIDs[j])) {
+ fprintf(stderr, "ERROR: Expected distinct memory effect IDs\n");
+ return 6;
+ }
+ }
+ }
+ fprintf(stderr, "memory effect IDs are distinct\n");
+ // CHECK: memory effect IDs are distinct
+
MlirAttribute nullParameters = {NULL};
MlirOpOperand opOperand = mlirOperationGetOpOperand(storeOp, 0);
MlirBlock block = mlirBlockCreate(1, &i32, &loc);
@@ -2651,16 +2712,38 @@ int testInterfaces(MlirContext ctx) {
mlirMemoryEffectInstanceCreateForSymbol(read, symbol, zero, 4, true,
defaultResource),
};
+ MlirMemoryEffect expectedEffects[] = {allocate, read, write, free, read};
+ MlirValue expectedValues[] = {
+ {NULL}, constantResult, constantResult, blockArgument, {NULL}};
for (intptr_t i = 0; i < 5; ++i) {
if (!instances[i].ptr) {
fprintf(stderr, "ERROR: Expected memory effect instance\n");
return 7;
}
- mlirMemoryEffectInstanceDestroy(instances[i]);
+ if (!mlirTypeIDEqual(mlirMemoryEffectGetEffectID(
+ mlirMemoryEffectInstanceGetEffect(instances[i])),
+ mlirMemoryEffectGetEffectID(expectedEffects[i])) ||
+ mlirMemoryEffectInstanceGetStage(instances[i]) != i ||
+ mlirMemoryEffectInstanceGetEffectOnFullRegion(instances[i]) !=
+ (i == 4) ||
+ !mlirValueEqual(mlirMemoryEffectInstanceGetValue(instances[i]),
+ expectedValues[i])) {
+ fprintf(stderr, "ERROR: Unexpected memory effect instance properties\n");
+ return 7;
+ }
}
+ if (!mlirAttributeEqual(mlirMemoryEffectInstanceGetSymbolRef(instances[4]),
+ symbol) ||
+ !mlirAttributeEqual(mlirMemoryEffectInstanceGetParameters(instances[4]),
+ zero)) {
+ fprintf(stderr, "ERROR: Unexpected symbol memory effect properties\n");
+ return 7;
+ }
+ for (intptr_t i = 0; i < 5; ++i)
+ mlirMemoryEffectInstanceDestroy(instances[i]);
mlirBlockDestroy(block);
- fprintf(stderr, "memory effect instances constructed\n");
- // CHECK: memory effect instances constructed
+ fprintf(stderr, "memory effect instance properties verified\n");
+ // CHECK: memory effect instance properties verified
mlirOperationDestroy(storeOp);
mlirOperationDestroy(constantOp);
diff --git a/mlir/test/CAPI/transform.c b/mlir/test/CAPI/transform.c
index 24d31cc590a10..860fad16b246c 100644
--- a/mlir/test/CAPI/transform.c
+++ b/mlir/test/CAPI/transform.c
@@ -14,6 +14,7 @@
#include "mlir-c/Support.h"
#include <assert.h>
+#include <inttypes.h>
#include <stdio.h>
#include <stdlib.h>
@@ -79,11 +80,56 @@ void testOperationType(MlirContext ctx) {
fprintf(stderr, "\n\n");
}
+typedef struct {
+ intptr_t numEffects;
+ intptr_t numReads;
+ intptr_t numWrites;
+} MemoryEffectCallbackData;
+
+static void collectMemoryEffects(intptr_t numEffects,
+ MlirMemoryEffectInstance *effects,
+ void *userData) {
+ MemoryEffectCallbackData *data = (MemoryEffectCallbackData *)userData;
+ data->numEffects += numEffects;
+ MlirTypeID readID = mlirMemoryEffectGetEffectID(mlirMemoryEffectsReadGet());
+ MlirTypeID writeID = mlirMemoryEffectGetEffectID(mlirMemoryEffectsWriteGet());
+ for (intptr_t i = 0; i < numEffects; ++i) {
+ MlirTypeID effectID = mlirMemoryEffectGetEffectID(
+ mlirMemoryEffectInstanceGetEffect(effects[i]));
+ data->numReads += mlirTypeIDEqual(effectID, readID);
+ data->numWrites += mlirTypeIDEqual(effectID, writeID);
+ }
+}
+
+// CHECK-LABEL: testMemoryEffectHelpers
+void testMemoryEffectHelpers(void) {
+ fprintf(stderr, "testMemoryEffectHelpers\n");
+
+ MemoryEffectCallbackData modifies = {0};
+ mlirTransformModifiesPayload(collectMemoryEffects, &modifies);
+ // CHECK: modifies payload: 2 effects, 1 read, 1 write
+ fprintf(stderr,
+ "modifies payload: %" PRIdPTR " effects, %" PRIdPTR " read, %" PRIdPTR
+ " write\n",
+ modifies.numEffects, modifies.numReads, modifies.numWrites);
+
+ MemoryEffectCallbackData reads = {0};
+ mlirTransformOnlyReadsPayload(collectMemoryEffects, &reads);
+ // CHECK: only reads payload: 1 effects, 1 read, 0 write
+ fprintf(stderr,
+ "only reads payload: %" PRIdPTR " effects, %" PRIdPTR
+ " read, %" PRIdPTR " write\n",
+ reads.numEffects, reads.numReads, reads.numWrites);
+
+ fprintf(stderr, "\n\n");
+}
+
int main(void) {
MlirContext ctx = mlirContextCreate();
mlirDialectHandleRegisterDialect(mlirGetDialectHandle__transform__(), ctx);
testAnyOpType(ctx);
testOperationType(ctx);
+ testMemoryEffectHelpers();
mlirContextDestroy(ctx);
return EXIT_SUCCESS;
}
diff --git a/mlir/test/python/dialects/ext.py b/mlir/test/python/dialects/ext.py
index 6e83da4a4a78a..b98b36ed07e0e 100644
--- a/mlir/test/python/dialects/ext.py
+++ b/mlir/test/python/dialects/ext.py
@@ -987,8 +987,8 @@ class TestIface(Dialect, name="ext_iface"):
class NoMemoryEffectModel(ir.MemoryEffectsOpInterface):
@staticmethod
- def get_effects(op, effects):
- pass
+ def get_effects(op):
+ return []
class AlwaysSpeculatableModel(ir.ConditionallySpeculatable):
@staticmethod
diff --git a/mlir/test/python/dialects/memory_effects_op_interface.py b/mlir/test/python/dialects/memory_effects_op_interface.py
index dc3b744aae2d7..d5370ea5358cb 100644
--- a/mlir/test/python/dialects/memory_effects_op_interface.py
+++ b/mlir/test/python/dialects/memory_effects_op_interface.py
@@ -13,81 +13,91 @@ class MemoryEffectsTest(ext.Dialect, name="memory_effects_test"):
class NoEffectModel(ir.MemoryEffectsOpInterface):
@staticmethod
- def get_effects(op, effects):
- pass
+ def get_effects(op):
+ return []
class ReadModel(ir.MemoryEffectsOpInterface):
@staticmethod
- def get_effects(op, effects):
- effects.append(
- ir.MemoryEffect.Read,
- op.op_operands[0],
- parameters=ir.StringAttr.get("read parameter"),
- stage=1,
- effect_on_full_region=True,
- resource=ir.SideEffectResource.Default,
- )
+ def get_effects(op):
+ return [
+ ir.MemoryEffectInstance(
+ ir.MemoryEffect.Read,
+ op.op_operands[0],
+ parameters=ir.StringAttr.get("read parameter"),
+ stage=1,
+ effect_on_full_region=True,
+ resource=ir.SideEffectResource.Default,
+ )
+ ]
class ReadDeadModel(ir.MemoryEffectsOpInterface):
@staticmethod
- def get_effects(op, effects):
- effects.append(ir.MemoryEffect.Read)
+ def get_effects(op):
+ return [ir.MemoryEffectInstance(ir.MemoryEffect.Read)]
class WriteModel(ir.MemoryEffectsOpInterface):
@staticmethod
- def get_effects(op, effects):
- effects.append(ir.MemoryEffect.Write)
+ def get_effects(op):
+ return [ir.MemoryEffectInstance(ir.MemoryEffect.Write)]
class FreeModel(ir.MemoryEffectsOpInterface):
@staticmethod
- def get_effects(op, effects):
- effects.append(ir.MemoryEffect.Free)
+ def get_effects(op):
+ return [ir.MemoryEffectInstance(ir.MemoryEffect.Free)]
class AllocateModel(ir.MemoryEffectsOpInterface):
@staticmethod
- def get_effects(op, effects):
- effects.append(ir.MemoryEffect.Allocate)
+ def get_effects(op):
+ return [ir.MemoryEffectInstance(ir.MemoryEffect.Allocate)]
class AllocateResultModel(ir.MemoryEffectsOpInterface):
@staticmethod
- def get_effects(op, effects):
- effects.append(ir.MemoryEffect.Allocate, op.results[0])
+ def get_effects(op):
+ return [ir.MemoryEffectInstance(ir.MemoryEffect.Allocate, op.results[0])]
class BlockArgumentTargetModel(ir.MemoryEffectsOpInterface):
@staticmethod
- def get_effects(op, effects):
- effects.append(ir.MemoryEffect.Read, op.regions[0].blocks[0].arguments[0])
+ def get_effects(op):
+ return [
+ ir.MemoryEffectInstance(
+ ir.MemoryEffect.Read, op.regions[0].blocks[0].arguments[0]
+ )
+ ]
class SymbolTargetModel(ir.MemoryEffectsOpInterface):
@staticmethod
- def get_effects(op, effects):
+ def get_effects(op):
try:
- effects.append(ir.MemoryEffect.Read, ir.StringAttr.get("not a symbol"))
+ ir.MemoryEffectInstance(
+ ir.MemoryEffect.Read, ir.StringAttr.get("not a symbol")
+ )
except TypeError as error:
print("invalid symbol target:", error)
try:
- effects.append(ir.MemoryEffect.Read, parameters=42)
+ ir.MemoryEffectInstance(ir.MemoryEffect.Read, parameters=42)
except TypeError as error:
print("invalid parameters:", error)
try:
- effects.append(ir.MemoryEffect.Read, 42)
+ ir.MemoryEffectInstance(ir.MemoryEffect.Read, 42)
except TypeError as error:
print("invalid target:", error)
- effects.append(
- ir.MemoryEffect.Read,
- ir.FlatSymbolRefAttr.get("global"),
- parameters=ir.StringAttr.get("symbol parameter"),
- stage=2,
- effect_on_full_region=True,
- )
+ return [
+ ir.MemoryEffectInstance(
+ ir.MemoryEffect.Read,
+ ir.FlatSymbolRefAttr.get("global"),
+ parameters=ir.StringAttr.get("symbol parameter"),
+ stage=2,
+ effect_on_full_region=True,
+ )
+ ]
class ReadOp(MemoryEffectsTest.Operation, name="read", traits=[ReadModel]):
@@ -167,12 +177,65 @@ def run_pass(source, pipeline):
isinstance(ir.MemoryEffect.Read, ir.MemoryEffect),
isinstance(ir.MemoryEffect.Write, ir.MemoryEffect),
)
+ # CHECK: memory effect equality: True True True True False False
+ print(
+ "memory effect equality:",
+ ir.MemoryEffect.Allocate == ir.MemoryEffect.Allocate,
+ ir.MemoryEffect.Free == ir.MemoryEffect.Free,
+ ir.MemoryEffect.Read == ir.MemoryEffect.Read,
+ ir.MemoryEffect.Write == ir.MemoryEffect.Write,
+ ir.MemoryEffect.Read == ir.MemoryEffect.Write,
+ ir.MemoryEffect.Read == 42,
+ )
# CHECK: default resource property: True
print(
"default resource property:",
isinstance(ir.SideEffectResource.Default, ir.SideEffectResource),
)
+ query_module = ir.Module.parse(
+ """
+ module {
+ func.func @test(%arg0: i32) -> i32 {
+ %0 = "memory_effects_test.read"(%arg0) : (i32) -> i32
+ return %0 : i32
+ }
+ }
+ """
+ )
+ read_op = query_module.body.operations[0].regions[0].blocks[0].operations[0]
+ read_effects = ir.MemoryEffectsOpInterface(read_op).get_effects()
+ read_effect = read_effects[0]
+ # CHECK: queried effects: True 1 True True 1 True True True
+ print(
+ "queried effects:",
+ isinstance(read_effects, list),
+ len(read_effects),
+ isinstance(read_effect, ir.MemoryEffectInstance),
+ read_effect.effect == ir.MemoryEffect.Read,
+ read_effect.stage,
+ read_effect.effect_on_full_region,
+ isinstance(read_effect.resource, ir.SideEffectResource),
+ read_effect.value == read_op.operands[0],
+ )
+ # CHECK: queried optional properties: "read parameter" True
+ print(
+ "queried optional properties:",
+ read_effect.parameters,
+ read_effect.symbol_ref is None,
+ )
+
+ symbol_effect = ir.MemoryEffectInstance(
+ ir.MemoryEffect.Read, ir.FlatSymbolRefAttr.get("global")
+ )
+ # CHECK: symbol effect properties: True True True
+ print(
+ "symbol effect properties:",
+ isinstance(symbol_effect.symbol_ref, ir.FlatSymbolRefAttr),
+ symbol_effect.value is None,
+ symbol_effect.parameters is None,
+ )
+
read_cse = run_pass(
"""
module {
diff --git a/mlir/test/python/dialects/transform_op_interface.py b/mlir/test/python/dialects/transform_op_interface.py
index 811dd0a9149fb..01c80826ce8b9 100644
--- a/mlir/test/python/dialects/transform_op_interface.py
+++ b/mlir/test/python/dialects/transform_op_interface.py
@@ -77,10 +77,12 @@ def schedule_boilerplate():
# Used by most ops defined below.
class MemoryEffectsOpInterfaceFallbackModel(ir.MemoryEffectsOpInterface):
@staticmethod
- def get_effects(op: ir.Operation, effects):
- transform.only_reads_handle(op.op_operands, effects)
- transform.produces_handle(op.results, effects)
- transform.only_reads_payload(effects)
+ def get_effects(op: ir.Operation):
+ return (
+ transform.only_reads_handle(op.op_operands)
+ + transform.produces_handle(op.results)
+ + transform.only_reads_payload()
+ )
# Demonstration of a TransformOpInterface-implementing op that gets named attributes
@@ -235,10 +237,12 @@ def allow_repeated_handle_operands(_op: OneOpInOneOpOut) -> bool:
# TransformOpInterface-implementing ops are also required to implement MemoryEffectsOpInterface. The above defined fallback model works for this op.
class MemoryEffectsOpInterfaceFallbackModel(ir.MemoryEffectsOpInterface):
@staticmethod
- def get_effects(op: ir.Operation, effects):
- transform.consumes_handle(op.op_operands, effects)
- transform.produces_handle(op.results, effects)
- transform.modifies_payload(effects)
+ def get_effects(op: ir.Operation):
+ return (
+ transform.consumes_handle(op.op_operands)
+ + transform.produces_handle(op.results)
+ + transform.modifies_payload()
+ )
MemoryEffectsOpInterfaceFallbackModel.attach(OneOpInOneOpOut.OPERATION_NAME)
>From 1220ed4306ad2743b0b9ca676a5e725ea2b0ec59 Mon Sep 17 00:00:00 2001
From: PragmaTwice <twice at apache.org>
Date: Sun, 2 Aug 2026 15:57:21 +0800
Subject: [PATCH 2/4] avoid owned and borrowed in comments
---
mlir/include/mlir-c/Interfaces.h | 25 ++++++++++++-------------
1 file changed, 12 insertions(+), 13 deletions(-)
diff --git a/mlir/include/mlir-c/Interfaces.h b/mlir/include/mlir-c/Interfaces.h
index aa03b7f9e7b66..09128766ee40a 100644
--- a/mlir/include/mlir-c/Interfaces.h
+++ b/mlir/include/mlir-c/Interfaces.h
@@ -146,24 +146,23 @@ mlirConditionallySpeculatableOpInterfaceGetSpeculatability(
// MemoryEffectsOpInterface
//===---------------------------------------------------------------------===//
-/// Returns the borrowed singleton instance of the allocate memory effect.
+/// Returns the singleton instance of the allocate memory effect.
MLIR_CAPI_EXPORTED MlirMemoryEffect mlirMemoryEffectsAllocateGet(void);
-/// Returns the borrowed singleton instance of the free memory effect.
+/// Returns the singleton instance of the free memory effect.
MLIR_CAPI_EXPORTED MlirMemoryEffect mlirMemoryEffectsFreeGet(void);
-/// Returns the borrowed singleton instance of the read memory effect.
+/// Returns the singleton instance of the read memory effect.
MLIR_CAPI_EXPORTED MlirMemoryEffect mlirMemoryEffectsReadGet(void);
-/// Returns the borrowed singleton instance of the write memory effect.
+/// Returns the singleton instance of the write memory effect.
MLIR_CAPI_EXPORTED MlirMemoryEffect mlirMemoryEffectsWriteGet(void);
/// Returns the TypeID identifying the concrete type of the given memory effect.
MLIR_CAPI_EXPORTED MlirTypeID
mlirMemoryEffectGetEffectID(MlirMemoryEffect effect);
-/// Returns the borrowed singleton instance of the default side effect
-/// resource.
+/// Returns the singleton instance of the default side effect resource.
MLIR_CAPI_EXPORTED MlirSideEffectResource
mlirSideEffectsDefaultResourceGet(void);
@@ -215,15 +214,15 @@ mlirMemoryEffectInstanceCreateForSymbol(MlirMemoryEffect effect,
bool effectOnFullRegion,
MlirSideEffectResource resource);
-/// Destroys an owned memory effect instance created or cloned by this API.
-MLIR_CAPI_EXPORTED void
-mlirMemoryEffectInstanceDestroy(MlirMemoryEffectInstance instance);
-
-/// Creates an owned copy of a memory effect instance. The caller must destroy
+/// Creates a copy of a memory effect instance. The caller must destroy
/// the returned instance with `mlirMemoryEffectInstanceDestroy`.
MLIR_CAPI_EXPORTED MlirMemoryEffectInstance
mlirMemoryEffectInstanceClone(MlirMemoryEffectInstance instance);
+/// Destroys a memory effect instance created or cloned by APIs above.
+MLIR_CAPI_EXPORTED void
+mlirMemoryEffectInstanceDestroy(MlirMemoryEffectInstance instance);
+
/// Returns the memory effect of the given instance.
MLIR_CAPI_EXPORTED MlirMemoryEffect
mlirMemoryEffectInstanceGetEffect(MlirMemoryEffectInstance instance);
@@ -257,7 +256,7 @@ MLIR_CAPI_EXPORTED MlirAttribute
mlirMemoryEffectInstanceGetSymbolRef(MlirMemoryEffectInstance instance);
/// Callback used to return a list of memory effect instances. `effects` points
-/// to `numEffects` consecutive borrowed instances that are only valid for the
+/// to `numEffects` consecutive instances that are only valid for the
/// duration of the callback. The caller-provided `userData` is forwarded to
/// the callback.
typedef void (*MlirMemoryEffectInstancesCallback)(
@@ -289,7 +288,7 @@ MLIR_CAPI_EXPORTED void mlirMemoryEffectsOpInterfaceAttachFallbackModel(
/// Gets the memory effects of the given operation. The operation must
/// implement the MemoryEffectsOpInterface. Invokes `callback` once with all
-/// effects; the instances are borrowed and only valid during the callback.
+/// effects; the lifetime of the instances are only valid during the callback.
MLIR_CAPI_EXPORTED void mlirMemoryEffectsOpInterfaceGetEffects(
MlirOperation operation, MlirMemoryEffectInstancesCallback callback,
void *userData);
>From 8416903249be5016f46017daa7c139c6edfbda7c Mon Sep 17 00:00:00 2001
From: PragmaTwice <twice at apache.org>
Date: Sun, 2 Aug 2026 16:12:26 +0800
Subject: [PATCH 3/4] refine comments
---
mlir/include/mlir-c/Interfaces.h | 19 +++++++++++--------
1 file changed, 11 insertions(+), 8 deletions(-)
diff --git a/mlir/include/mlir-c/Interfaces.h b/mlir/include/mlir-c/Interfaces.h
index 09128766ee40a..e35aae955da92 100644
--- a/mlir/include/mlir-c/Interfaces.h
+++ b/mlir/include/mlir-c/Interfaces.h
@@ -255,10 +255,10 @@ mlirMemoryEffectInstanceGetValue(MlirMemoryEffectInstance instance);
MLIR_CAPI_EXPORTED MlirAttribute
mlirMemoryEffectInstanceGetSymbolRef(MlirMemoryEffectInstance instance);
-/// Callback used to return a list of memory effect instances. `effects` points
-/// to `numEffects` consecutive instances that are only valid for the
-/// duration of the callback. The caller-provided `userData` is forwarded to
-/// the callback.
+/// Callback for receiving a batch of memory effect instances. `effects` points
+/// to `numEffects` consecutive instances. Ownership is not transferred, and
+/// the instances are valid only while `callback` is executing. The
+/// caller-provided `userData` is forwarded to the callback.
typedef void (*MlirMemoryEffectInstancesCallback)(
intptr_t numEffects, MlirMemoryEffectInstance *effects, void *userData);
@@ -271,9 +271,10 @@ typedef struct {
void (*construct)(void *userData);
/// Optional destructor for user data. Set to nullptr to disable it.
void (*destruct)(void *userData);
- /// Get memory effects callback. Implementations return effects by invoking
- /// `callback` with an array of memory effect instances. The callback copies
- /// the instances synchronously, so implementations retain ownership of them.
+ /// Get memory effects callback. Implementations report effects by invoking
+ /// `callback` before returning. The supplied callback copies the instances,
+ /// so implementations retain ownership of the instances and only need to
+ /// keep them valid until `callback` returns.
void (*getEffects)(MlirOperation op,
MlirMemoryEffectInstancesCallback callback,
void *callbackUserData, void *userData);
@@ -288,7 +289,9 @@ MLIR_CAPI_EXPORTED void mlirMemoryEffectsOpInterfaceAttachFallbackModel(
/// Gets the memory effects of the given operation. The operation must
/// implement the MemoryEffectsOpInterface. Invokes `callback` once with all
-/// effects; the lifetime of the instances are only valid during the callback.
+/// effects. Ownership is not transferred; call
+/// `mlirMemoryEffectInstanceClone` from the callback to keep a copy after the
+/// callback returns.
MLIR_CAPI_EXPORTED void mlirMemoryEffectsOpInterfaceGetEffects(
MlirOperation operation, MlirMemoryEffectInstancesCallback callback,
void *userData);
>From 0d1b082af727c545adbbab70d05040af50fc5c56 Mon Sep 17 00:00:00 2001
From: PragmaTwice <twice at apache.org>
Date: Mon, 3 Aug 2026 21:57:45 +0800
Subject: [PATCH 4/4] refine handle comments
---
mlir/include/mlir-c/Dialect/Transform.h | 13 ++++++++-----
mlir/lib/Bindings/Python/DialectTransform.cpp | 13 ++++++++-----
2 files changed, 16 insertions(+), 10 deletions(-)
diff --git a/mlir/include/mlir-c/Dialect/Transform.h b/mlir/include/mlir-c/Dialect/Transform.h
index 83f1ed7dc37e9..fc8ede36ad046 100644
--- a/mlir/include/mlir-c/Dialect/Transform.h
+++ b/mlir/include/mlir-c/Dialect/Transform.h
@@ -243,30 +243,33 @@ MLIR_CAPI_EXPORTED void mlirPatternDescriptorOpInterfaceAttachFallbackModel(
// Transform-specifc MemoryEffectsOpInterface helpers
//===---------------------------------------------------------------------===//
-/// Helper to mark operands as only reading handles.
+/// Invokes `callback` with `OnlyReadsHandle` effects corresponding to operands
+/// which have been marked as having those effects.
MLIR_CAPI_EXPORTED void
mlirTransformOnlyReadsHandle(MlirOpOperand *operands, intptr_t numOperands,
MlirMemoryEffectInstancesCallback callback,
void *userData);
-/// Helper to mark operands as consuming handles.
+/// Invokes `callback` with `ConsumesHandle` effects corresponding to operands
+/// which have been marked as having those effects.
MLIR_CAPI_EXPORTED void
mlirTransformConsumesHandle(MlirOpOperand *operands, intptr_t numOperands,
MlirMemoryEffectInstancesCallback callback,
void *userData);
-/// Helper to mark results as producing handles.
+/// Invokes `callback` with `ProducesHandle` effects corresponding to results
+/// which have been marked as having those effects.
MLIR_CAPI_EXPORTED void
mlirTransformProducesHandle(MlirValue *results, intptr_t numResults,
MlirMemoryEffectInstancesCallback callback,
void *userData);
-/// Helper to mark potential modifications to the payload IR.
+/// Invokes `callback` with `ModifiesPayload` effects.
MLIR_CAPI_EXPORTED void
mlirTransformModifiesPayload(MlirMemoryEffectInstancesCallback callback,
void *userData);
-/// Helper to mark potential reads from the payload IR.
+/// Invokes `callback` with `OnlyReadsPayload` effects.
MLIR_CAPI_EXPORTED void
mlirTransformOnlyReadsPayload(MlirMemoryEffectInstancesCallback callback,
void *userData);
diff --git a/mlir/lib/Bindings/Python/DialectTransform.cpp b/mlir/lib/Bindings/Python/DialectTransform.cpp
index dd8c2d711edb2..06e01ad16f675 100644
--- a/mlir/lib/Bindings/Python/DialectTransform.cpp
+++ b/mlir/lib/Bindings/Python/DialectTransform.cpp
@@ -565,22 +565,25 @@ static void populateDialectTransformSubmodule(nb::module_ &m) {
PyPatternDescriptorOpInterface::bind(m);
m.def("only_reads_handle", onlyReadsHandle,
- "Returns effects marking operands as only reading handles.",
+ "Returns `OnlyReadsHandle` effects corresponding to operands which "
+ "have been marked as having those effects.",
nb::arg("operands"));
m.def("consumes_handle", consumesHandle,
- "Returns effects marking operands as consuming handles.",
+ "Returns `ConsumesHandle` effects corresponding to operands which "
+ "have been marked as having those effects.",
nb::arg("operands"));
m.def("produces_handle", producesHandle,
- "Returns effects marking results as producing handles.",
+ "Returns `ProducesHandle` effects corresponding to results which have "
+ "been marked as having those effects.",
nb::arg("results"));
m.def("modifies_payload", modifiesPayload,
- "Returns effects marking potential payload modifications.");
+ "Returns `ModifiesPayload` effects.");
m.def("only_reads_payload", onlyReadsPayload,
- "Returns effects marking payload reads.");
+ "Returns `OnlyReadsPayload` effects.");
}
} // namespace transform
} // namespace MLIR_BINDINGS_PYTHON_DOMAIN
More information about the Mlir-commits
mailing list