求pyO3中判断输入是否为scipy.sparse.csr_matrix类型的完整示例
完整pyO3示例:判断输入是否为scipy.sparse.csr_matrix类型
下面是一个可直接运行的pyO3扩展函数示例,核心是动态获取scipy.sparse.csr_matrix的类型对象,再用is_instance完成类型检查:
第一步:配置Cargo.toml
确保你的Cargo.toml包含pyO3依赖并开启扩展模块特性:
[package] name = "sparse_checker" version = "0.1.0" edition = "2021" [lib] name = "sparse_checker" crate-type = ["cdylib"] [dependencies] pyo3 = { version = "0.20", features = ["extension-module"] }
第二步:编写src/lib.rs代码
use pyo3::prelude::*; #[pyfunction] fn is_csr_matrix(py: Python<'_>, obj: &PyAny) -> PyResult<bool> { // 动态获取scipy.sparse.csr_matrix的类型对象 let csr_matrix_type = py.get_type::<PyAny>("scipy.sparse", "csr_matrix")?; // 检查输入对象是否属于该类型 Ok(obj.is_instance(csr_matrix_type)?) } #[pymodule] fn sparse_checker(_py: Python<'_>, m: &PyModule) -> PyResult<()> { m.add_function(wrap_pyfunction!(is_csr_matrix, m)?)?; Ok(()) }
编译与测试
- 用maturin编译扩展:
maturin build --release
- 找到编译生成的
.so/.pyd文件,在Python环境中测试:
import scipy.sparse as sp import sparse_checker # 创建测试对象 csr_mat = sp.csr_matrix([[1,2],[3,4]]) dense_mat = [[1,2],[3,4]] print(sparse_checker.is_csr_matrix(csr_mat)) # 输出True print(sparse_checker.is_csr_matrix(dense_mat)) # 输出False
关键逻辑说明
py.get_type::<PyAny>("scipy.sparse", "csr_matrix"):通过模块名和类名动态获取Python类型对象,无需提前绑定类型,解决你提到的类型对象获取困惑。obj.is_instance(csr_matrix_type):用获取到的类型对象直接检查输入对象的归属,返回布尔值。
内容的提问来源于stack exchange,提问作者GHPR
相关产品推荐
相关产品推荐

