From a8d1b5f1b45236976e79c2c533886771094ad4ef Mon Sep 17 00:00:00 2001 From: Jiacheng Huang Date: Thu, 6 Aug 2026 20:20:27 +0800 Subject: [PATCH 1/3] refactor: minimize direct Python C API usage --- src/pybind11_utils.h | 62 ++++++++++++++++-------------------- tests/conftest.py | 1 + tests/test_pybind11_utils.py | 60 ++++++++++++++++++++++++++++++++++ 3 files changed, 89 insertions(+), 34 deletions(-) create mode 100644 tests/test_pybind11_utils.py diff --git a/src/pybind11_utils.h b/src/pybind11_utils.h index 6d34e2d06..3ae7e512f 100644 --- a/src/pybind11_utils.h +++ b/src/pybind11_utils.h @@ -20,66 +20,60 @@ namespace infini::ops { namespace detail { -inline PyObject* InternedName(const char* value) { +inline py::handle InternedName(const char* value) { // Keep one reference for the process lifetime; Python objects must not be // decref'd by a static destructor after interpreter finalization. auto* name{PyUnicode_InternFromString(value)}; if (name == nullptr) throw py::error_already_set(); - return name; + return py::handle{name}; } -inline PyObject* DataPtrName() { - static PyObject* const name{InternedName("data_ptr")}; +inline py::handle DataPtrName() { + static const auto name{InternedName("data_ptr")}; return name; } -inline PyObject* ShapeName() { - static PyObject* const name{InternedName("shape")}; +inline py::handle ShapeName() { + static const auto name{InternedName("shape")}; return name; } -inline PyObject* DTypeName() { - static PyObject* const name{InternedName("dtype")}; +inline py::handle DTypeName() { + static const auto name{InternedName("dtype")}; return name; } -inline PyObject* DeviceName() { - static PyObject* const name{InternedName("device")}; +inline py::handle DeviceName() { + static const auto name{InternedName("device")}; return name; } -inline PyObject* TypeName() { - static PyObject* const name{InternedName("type")}; +inline py::handle TypeName() { + static const auto name{InternedName("type")}; return name; } -inline PyObject* IndexName() { - static PyObject* const name{InternedName("index")}; +inline py::handle IndexName() { + static const auto name{InternedName("index")}; return name; } -inline PyObject* StrideName() { - static PyObject* const name{InternedName("stride")}; +inline py::handle StrideName() { + static const auto name{InternedName("stride")}; return name; } -inline py::object GetAttr(py::handle obj, PyObject* name) { - auto* value{PyObject_GetAttr(obj.ptr(), name)}; - if (value == nullptr) throw py::error_already_set(); - return py::reinterpret_steal(value); -} - -inline py::object CallMethodNoArgs(py::handle obj, PyObject* name) { - auto* value{PyObject_CallMethodNoArgs(obj.ptr(), name)}; +inline py::object CallMethodNoArgs(py::handle obj, py::handle name) { + auto* value{PyObject_CallMethodNoArgs(obj.ptr(), name.ptr())}; if (value == nullptr) throw py::error_already_set(); return py::reinterpret_steal(value); } template -Integer IntegerFromPyObject(PyObject* obj) { +Integer IntegerFromPyObject(py::handle obj) { static_assert(std::is_integral_v); if constexpr (std::is_unsigned_v) { - const auto value{PyLong_AsUnsignedLongLong(obj)}; + const auto value{PyLong_AsUnsignedLongLong(obj.ptr())}; if (value == static_cast(-1) && PyErr_Occurred()) { throw py::error_already_set(); } @@ -90,7 +84,7 @@ Integer IntegerFromPyObject(PyObject* obj) { } return static_cast(value); } else { - const auto value{PyLong_AsLongLong(obj)}; + const auto value{PyLong_AsLongLong(obj.ptr())}; if (value == -1 && PyErr_Occurred()) throw py::error_already_set(); if (value < std::numeric_limits::min() || value > std::numeric_limits::max()) { @@ -113,7 +107,7 @@ Vector VectorFromSequence(py::handle obj) { result.reserve(static_cast(size)); for (Py_ssize_t i = 0; i < size; ++i) { result.push_back(IntegerFromPyObject( - PySequence_Fast_GET_ITEM(sequence.ptr(), i))); + py::handle{PySequence_Fast_GET_ITEM(sequence.ptr(), i)})); } return result; } @@ -247,8 +241,8 @@ inline std::optional TryDeviceTypeFromString( namespace detail { inline Device DeviceFromPybind11HandleImpl(py::handle obj) { - auto device_obj{detail::GetAttr(obj, detail::DeviceName())}; - auto device_type_obj{detail::GetAttr(device_obj, detail::TypeName())}; + auto device_obj{py::getattr(obj, detail::DeviceName())}; + auto device_type_obj{py::getattr(device_obj, detail::TypeName())}; std::string device_type_storage; std::string_view device_type_str; if (PyUnicode_Check(device_type_obj.ptr())) { @@ -262,7 +256,7 @@ inline Device DeviceFromPybind11HandleImpl(py::handle obj) { device_type_storage = device_type_obj.cast(); device_type_str = device_type_storage; } - auto device_index_obj{detail::GetAttr(device_obj, detail::IndexName())}; + auto device_index_obj{py::getattr(device_obj, detail::IndexName())}; auto device_index{device_index_obj.is_none() ? 0 : device_index_obj.cast()}; @@ -275,10 +269,10 @@ inline Tensor TensorFromPybind11HandleImpl(py::handle obj) { .cast())}; auto shape{detail::VectorFromSequence( - detail::GetAttr(obj, detail::ShapeName()))}; + py::getattr(obj, detail::ShapeName()))}; - auto dtype{DataTypeFromPybind11HandleImpl( - detail::GetAttr(obj, detail::DTypeName()))}; + auto dtype{ + DataTypeFromPybind11HandleImpl(py::getattr(obj, detail::DTypeName()))}; auto device{DeviceFromPybind11HandleImpl(obj)}; diff --git a/tests/conftest.py b/tests/conftest.py index 09122d62e..603e8cf6c 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -205,6 +205,7 @@ def _set_random_seed(seed): "tests/test_generate_ninetoothed_ops.py", "tests/test_generate_torch_ops.py", "tests/test_generate_wrappers.py", + "tests/test_pybind11_utils.py", } _SMOKE_TORCH_OPS = {"abs", "clamp", "exp", "relu"} diff --git a/tests/test_pybind11_utils.py b/tests/test_pybind11_utils.py new file mode 100644 index 000000000..7bc102ce4 --- /dev/null +++ b/tests/test_pybind11_utils.py @@ -0,0 +1,60 @@ +import pytest +import torch + +import infini.ops + + +class _IteratingTuple(tuple): + def __new__(cls, stored_values, iterated_values): + instance = super().__new__(cls, stored_values) + instance.iterated_values = tuple(iterated_values) + + return instance + + def __iter__(self): + return iter(self.iterated_values) + + +class _DuckTensor: + def __init__(self, tensor): + self.tensor = tensor + self.shape = _IteratingTuple((2, 2), tensor.shape) + self.dtype = tensor.dtype + self.device = tensor.device + + def data_ptr(self): + return self.tensor.data_ptr() + + def stride(self): + return _IteratingTuple((2, 1), self.tensor.stride()) + + +def test_duck_tensor_conversion_respects_sequence_subclass_iteration(): + implementation_indices = infini.ops.Add.active_implementation_indices("cpu") + + if not implementation_indices: + pytest.skip("CPU Add implementation is not active") + + input = torch.tensor( + ( + (-1.0, -2.0, -3.0), + (-4.0, -5.0, -6.0), + ) + ) + other = torch.tensor( + ( + (1.0, 2.0, 3.0), + (4.0, 5.0, 6.0), + ) + ) + out = torch.full_like(input, -123.0) + + infini.ops.add( + _DuckTensor(input), + _DuckTensor(other), + _DuckTensor(out), + stream=0, + implementation_index=implementation_indices[0], + ) + + assert torch.equal(out, input + other) From 2b0318f185e055287c8259801689487e785993cf Mon Sep 17 00:00:00 2001 From: Jiacheng Huang Date: Fri, 7 Aug 2026 11:11:35 +0800 Subject: [PATCH 2/3] refactor: clarify integer conversion helper name --- src/pybind11_utils.h | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/src/pybind11_utils.h b/src/pybind11_utils.h index 3ae7e512f..2e6e2ceec 100644 --- a/src/pybind11_utils.h +++ b/src/pybind11_utils.h @@ -70,7 +70,7 @@ inline py::object CallMethodNoArgs(py::handle obj, py::handle name) { } template -Integer IntegerFromPyObject(py::handle obj) { +Integer IntegerFromPybind11Handle(py::handle obj) { static_assert(std::is_integral_v); if constexpr (std::is_unsigned_v) { const auto value{PyLong_AsUnsignedLongLong(obj.ptr())}; @@ -106,7 +106,7 @@ Vector VectorFromSequence(py::handle obj) { Vector result; result.reserve(static_cast(size)); for (Py_ssize_t i = 0; i < size; ++i) { - result.push_back(IntegerFromPyObject( + result.push_back(IntegerFromPybind11Handle( py::handle{PySequence_Fast_GET_ITEM(sequence.ptr(), i)})); } return result; From 9ca118261b2da25a0d08325f870e39334641a506 Mon Sep 17 00:00:00 2001 From: Jiacheng Huang Date: Fri, 7 Aug 2026 11:26:22 +0800 Subject: [PATCH 3/3] test: drop unrelated tensor conversion coverage --- tests/conftest.py | 1 - tests/test_pybind11_utils.py | 60 ------------------------------------ 2 files changed, 61 deletions(-) delete mode 100644 tests/test_pybind11_utils.py diff --git a/tests/conftest.py b/tests/conftest.py index 603e8cf6c..09122d62e 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -205,7 +205,6 @@ def _set_random_seed(seed): "tests/test_generate_ninetoothed_ops.py", "tests/test_generate_torch_ops.py", "tests/test_generate_wrappers.py", - "tests/test_pybind11_utils.py", } _SMOKE_TORCH_OPS = {"abs", "clamp", "exp", "relu"} diff --git a/tests/test_pybind11_utils.py b/tests/test_pybind11_utils.py deleted file mode 100644 index 7bc102ce4..000000000 --- a/tests/test_pybind11_utils.py +++ /dev/null @@ -1,60 +0,0 @@ -import pytest -import torch - -import infini.ops - - -class _IteratingTuple(tuple): - def __new__(cls, stored_values, iterated_values): - instance = super().__new__(cls, stored_values) - instance.iterated_values = tuple(iterated_values) - - return instance - - def __iter__(self): - return iter(self.iterated_values) - - -class _DuckTensor: - def __init__(self, tensor): - self.tensor = tensor - self.shape = _IteratingTuple((2, 2), tensor.shape) - self.dtype = tensor.dtype - self.device = tensor.device - - def data_ptr(self): - return self.tensor.data_ptr() - - def stride(self): - return _IteratingTuple((2, 1), self.tensor.stride()) - - -def test_duck_tensor_conversion_respects_sequence_subclass_iteration(): - implementation_indices = infini.ops.Add.active_implementation_indices("cpu") - - if not implementation_indices: - pytest.skip("CPU Add implementation is not active") - - input = torch.tensor( - ( - (-1.0, -2.0, -3.0), - (-4.0, -5.0, -6.0), - ) - ) - other = torch.tensor( - ( - (1.0, 2.0, 3.0), - (4.0, 5.0, 6.0), - ) - ) - out = torch.full_like(input, -123.0) - - infini.ops.add( - _DuckTensor(input), - _DuckTensor(other), - _DuckTensor(out), - stream=0, - implementation_index=implementation_indices[0], - ) - - assert torch.equal(out, input + other)