[mlir python] Add nanobind support for standalone dialects. (#117922)
This PR allows out-of-tree dialects to write Python dialect modules using nanobind instead of pybind11. It may make sense to migrate in-tree dialects and some of the ODS Python infrastructure to nanobind, but that is a topic for a future change. This PR makes the following changes: * adds nanobind to the CMake and Bazel build systems. We also add robin_map to the Bazel build, which is a dependency of nanobind. * adds a PYTHON_BINDING_LIBRARY option to various CMake functions, such as declare_mlir_python_extension, allowing users to select a Python binding library. * creates a fork of mlir/include/mlir/Bindings/Python/PybindAdaptors.h named NanobindAdaptors.h. This plays the same role, using nanobind instead of pybind11. * splits CollectDiagnosticsToStringScope out of PybindAdaptors.h and into a new header mlir/include/mlir/Bindings/Python/Diagnostics.h, since it is code that is no way related to pybind11 or for that matter, Python. * changed the standalone Python extension example to have both pybind11 and nanobind variants. * changed mlir/python/mlir/dialects/python_test.py to have both pybind11 and nanobind variants. Notes: * A slightly unfortunate thing that I needed to do in the CMake integration was to use FindPython in addition to FindPython3, since nanobind's CMake integration expects the Python_ names for variables. Perhaps there's a better way to do this.
This commit is contained in:
@@ -1,12 +1,33 @@
|
||||
# RUN: %PYTHON %s | FileCheck %s
|
||||
# RUN: %PYTHON %s pybind11 | FileCheck %s
|
||||
# RUN: %PYTHON %s nanobind | FileCheck %s
|
||||
|
||||
import sys
|
||||
from mlir.ir import *
|
||||
import mlir.dialects.func as func
|
||||
import mlir.dialects.python_test as test
|
||||
import mlir.dialects.tensor as tensor
|
||||
import mlir.dialects.arith as arith
|
||||
|
||||
test.register_python_test_dialect(get_dialect_registry())
|
||||
if sys.argv[1] == "pybind11":
|
||||
from mlir._mlir_libs._mlirPythonTestPybind11 import (
|
||||
TestAttr,
|
||||
TestType,
|
||||
TestTensorValue,
|
||||
TestIntegerRankedTensorType,
|
||||
)
|
||||
|
||||
test.register_python_test_dialect(get_dialect_registry(), use_nanobind=False)
|
||||
elif sys.argv[1] == "nanobind":
|
||||
from mlir._mlir_libs._mlirPythonTestNanobind import (
|
||||
TestAttr,
|
||||
TestType,
|
||||
TestTensorValue,
|
||||
TestIntegerRankedTensorType,
|
||||
)
|
||||
|
||||
test.register_python_test_dialect(get_dialect_registry(), use_nanobind=True)
|
||||
else:
|
||||
raise ValueError("Expected pybind11 or nanobind as argument")
|
||||
|
||||
|
||||
def run(f):
|
||||
@@ -308,7 +329,7 @@ def testOptionalOperandOp():
|
||||
@run
|
||||
def testCustomAttribute():
|
||||
with Context() as ctx, Location.unknown():
|
||||
a = test.TestAttr.get()
|
||||
a = TestAttr.get()
|
||||
# CHECK: #python_test.test_attr
|
||||
print(a)
|
||||
|
||||
@@ -325,11 +346,11 @@ def testCustomAttribute():
|
||||
print(repr(op2.test_attr))
|
||||
|
||||
# The following cast must not assert.
|
||||
b = test.TestAttr(a)
|
||||
b = TestAttr(a)
|
||||
|
||||
unit = UnitAttr.get()
|
||||
try:
|
||||
test.TestAttr(unit)
|
||||
TestAttr(unit)
|
||||
except ValueError as e:
|
||||
assert "Cannot cast attribute to TestAttr" in str(e)
|
||||
else:
|
||||
@@ -338,7 +359,7 @@ def testCustomAttribute():
|
||||
# The following must trigger a TypeError from our adaptors and must not
|
||||
# crash.
|
||||
try:
|
||||
test.TestAttr(42)
|
||||
TestAttr(42)
|
||||
except TypeError as e:
|
||||
assert "Expected an MLIR object" in str(e)
|
||||
else:
|
||||
@@ -347,7 +368,7 @@ def testCustomAttribute():
|
||||
# The following must trigger a TypeError from pybind (therefore, not
|
||||
# checking its message) and must not crash.
|
||||
try:
|
||||
test.TestAttr(42, 56)
|
||||
TestAttr(42, 56)
|
||||
except TypeError:
|
||||
pass
|
||||
else:
|
||||
@@ -357,12 +378,12 @@ def testCustomAttribute():
|
||||
@run
|
||||
def testCustomType():
|
||||
with Context() as ctx:
|
||||
a = test.TestType.get()
|
||||
a = TestType.get()
|
||||
# CHECK: !python_test.test_type
|
||||
print(a)
|
||||
|
||||
# The following cast must not assert.
|
||||
b = test.TestType(a)
|
||||
b = TestType(a)
|
||||
# Instance custom types should have typeids
|
||||
assert isinstance(b.typeid, TypeID)
|
||||
# Subclasses of ir.Type should not have a static_typeid
|
||||
@@ -374,7 +395,7 @@ def testCustomType():
|
||||
|
||||
i8 = IntegerType.get_signless(8)
|
||||
try:
|
||||
test.TestType(i8)
|
||||
TestType(i8)
|
||||
except ValueError as e:
|
||||
assert "Cannot cast type to TestType" in str(e)
|
||||
else:
|
||||
@@ -383,7 +404,7 @@ def testCustomType():
|
||||
# The following must trigger a TypeError from our adaptors and must not
|
||||
# crash.
|
||||
try:
|
||||
test.TestType(42)
|
||||
TestType(42)
|
||||
except TypeError as e:
|
||||
assert "Expected an MLIR object" in str(e)
|
||||
else:
|
||||
@@ -392,7 +413,7 @@ def testCustomType():
|
||||
# The following must trigger a TypeError from pybind (therefore, not
|
||||
# checking its message) and must not crash.
|
||||
try:
|
||||
test.TestType(42, 56)
|
||||
TestType(42, 56)
|
||||
except TypeError:
|
||||
pass
|
||||
else:
|
||||
@@ -405,7 +426,7 @@ def testTensorValue():
|
||||
with Context() as ctx, Location.unknown():
|
||||
i8 = IntegerType.get_signless(8)
|
||||
|
||||
class Tensor(test.TestTensorValue):
|
||||
class Tensor(TestTensorValue):
|
||||
def __str__(self):
|
||||
return super().__str__().replace("Value", "Tensor")
|
||||
|
||||
@@ -425,9 +446,9 @@ def testTensorValue():
|
||||
|
||||
# Classes of custom types that inherit from concrete types should have
|
||||
# static_typeid
|
||||
assert isinstance(test.TestIntegerRankedTensorType.static_typeid, TypeID)
|
||||
assert isinstance(TestIntegerRankedTensorType.static_typeid, TypeID)
|
||||
# And it should be equal to the in-tree concrete type
|
||||
assert test.TestIntegerRankedTensorType.static_typeid == t.type.typeid
|
||||
assert TestIntegerRankedTensorType.static_typeid == t.type.typeid
|
||||
|
||||
d = tensor.EmptyOp([1, 2, 3], IntegerType.get_signless(5)).result
|
||||
# CHECK: Value(%{{.*}} = tensor.empty() : tensor<1x2x3xi5>)
|
||||
@@ -491,7 +512,7 @@ def inferReturnTypeComponents():
|
||||
@run
|
||||
def testCustomTypeTypeCaster():
|
||||
with Context() as ctx, Location.unknown():
|
||||
a = test.TestType.get()
|
||||
a = TestType.get()
|
||||
assert a.typeid is not None
|
||||
|
||||
b = Type.parse("!python_test.test_type")
|
||||
@@ -500,7 +521,7 @@ def testCustomTypeTypeCaster():
|
||||
# CHECK: TestType(!python_test.test_type)
|
||||
print(repr(b))
|
||||
|
||||
c = test.TestIntegerRankedTensorType.get([10, 10], 5)
|
||||
c = TestIntegerRankedTensorType.get([10, 10], 5)
|
||||
# CHECK: tensor<10x10xi5>
|
||||
print(c)
|
||||
# CHECK: TestIntegerRankedTensorType(tensor<10x10xi5>)
|
||||
@@ -511,7 +532,7 @@ def testCustomTypeTypeCaster():
|
||||
|
||||
@register_type_caster(c.typeid)
|
||||
def type_caster(pytype):
|
||||
return test.TestIntegerRankedTensorType(pytype)
|
||||
return TestIntegerRankedTensorType(pytype)
|
||||
|
||||
except RuntimeError as e:
|
||||
print(e)
|
||||
@@ -530,7 +551,7 @@ def testCustomTypeTypeCaster():
|
||||
|
||||
@register_type_caster(c.typeid, replace=True)
|
||||
def type_caster(pytype):
|
||||
return test.TestIntegerRankedTensorType(pytype)
|
||||
return TestIntegerRankedTensorType(pytype)
|
||||
|
||||
d = tensor.EmptyOp([10, 10], IntegerType.get_signless(5)).result
|
||||
# CHECK: tensor<10x10xi5>
|
||||
|
||||
Reference in New Issue
Block a user