如何通过PyO3将NumPy数组传入Rust函数?
解决PyO3传递NumPy数组到Rust函数的类型错误
错误原因分析
你遇到的两类错误核心原因如下:
- Python调用时的类型不匹配:NumPy默认创建
float64(双精度浮点)数组,但你的Rust函数参数指定的是f32(单精度浮点),PyO3无法自动跨精度转换NumPy数组,因此抛出TypeError。 - 使用
ndarray::ArrayD的编译错误:ArrayD是Rustndarray库的原生数组类型,并非PyO3包装的Python对象,PyO3的函数参数 trait 不支持直接将Python的ndarray转换为该类型,必须先接收PyO3提供的PyArray类型再做转换。
解决方案
方案1:接受任意维度的只读数组,返回一维数组
使用PyReadonlyArrayDyn<T>可以兼容任意维度的NumPy数组,同时通过ndarray库的视图安全访问元素:
use numpy::{PyReadonlyArrayDyn, PyArray, ndarray::ArrayViewD}; use pyo3::{Python, PyResult, Bound}; #[pyfunction] fn identity<'py>(py: Python<'py>, arr: PyReadonlyArrayDyn<f32>) -> PyResult<Bound<'py, PyArray<f32, numpy::ndarray::Dim<[usize;1]>>>> { // 将PyO3的数组包装转换为ndarray的只读视图 let array_view: ArrayViewD<'_, f32> = arr.as_array(); // 扁平化数组为一维(如果需要处理多维输入) let flat_view = array_view.into_shape((array_view.len(),))?; // 转换为Vec后生成新的PyArray let vec: Vec<f32> = flat_view.to_vec(); Ok(vec.into_pyarray(py).into_bound(py)) }
Python调用时需要传入float32类型的数组(或者将Rust函数中的f32改为f64适配NumPy默认类型):
import numpy as np import temp # 指定dtype为float32匹配Rust函数参数 z = np.array([5.1, 5.3, 4.1, 6.4], dtype=np.float32) print(temp.identity(z))
方案2:指定接受一维数组,简化签名
如果只需要处理一维数组,可以用Ix1(一维维度的别名)简化签名,直接接收Bound<PyArray<T, Ix1>>:
use numpy::{PyArray, ndarray::Ix1}; use pyo3::{Python, PyResult, Bound}; #[pyfunction] fn identity<'py>(py: Python<'py>, arr: Bound<'py, PyArray<f32, Ix1>>) -> PyResult<Bound<'py, PyArray<f32, Ix1>>> { // 获取数组的只读视图 let array_view = arr.readonly().as_array(); // 复制元素到Vec并生成新数组 let vec: Vec<f32> = array_view.to_vec(); Ok(vec.into_pyarray(py).into_bound(py)) }
同样,Python调用时需保证数组类型与Rust参数一致。
内容的提问来源于stack exchange,提问作者jeffhmr
相关产品推荐
相关产品推荐

