Python中如何为Numpy数组实现自动形状断言?
问题
我编写了一个包含单个方法的类(代码如下):
import numpy as np import numpy.typing as npt import pyvista as pv from typing import Tuple, TypeVar, Literal class SurfaceTriangleMesh: int_type = TypeVar("int_type", bound=int) flt_type = TypeVar("flt_type", bound=float) int_Shape2D = Tuple[int_type, int_type] flt_Shape2D = Tuple[flt_type, flt_type] int_Shape2DType = TypeVar("int_Shape2DType", bound=int_Shape2D) flt_Shape2DType = TypeVar("flt_Shape2DType", bound=flt_Shape2D) def __init__(self, mesh: pv.PolyData): self.mesh = mesh def add_new_points(self, points_to_add: npt.NDArray[flt_Shape2DType]): """ Add new points to surface mesh from attribute self.mesh :param points_to_add: points-to-add array of shape [N * 3] of type float :return: new surface mesh of type pyvista.PolyData """ pass
在add_new_points方法中,我希望传入形状为[*, 3]的Numpy数组,并在执行后续操作前检查该数组的形状。请问是否存在“自动”的形状断言方式?即能否借助某些类或第三方库,在其内部自动完成Numpy数组的形状与类型断言,无需手动编写断言代码?还是必须手动编写断言?
解决方案
1. 第三方库实现自动断言
有几个成熟的库可以帮你自动处理Numpy数组的形状和类型校验,无需手动编写断言逻辑:
(1) Pydantic(支持Numpy类型验证)
Pydantic可以通过自定义模型约束输入数组的形状和类型,自动完成校验:
import numpy as np import numpy.typing as npt import pyvista as pv from pydantic import BaseModel, ValidationError class PointsInput(BaseModel): points_to_add: npt.NDArray[np.float64] @classmethod def validate_array(cls, v): if v.ndim != 2 or v.shape[1] != 3: raise ValueError(f"输入数组形状必须为[*, 3],当前为{v.shape}") if not np.issubdtype(v.dtype, np.floating): raise ValueError(f"输入数组必须为浮点类型,当前为{v.dtype}") return v model_config = {"arbitrary_types_allowed": True} class SurfaceTriangleMesh: def __init__(self, mesh: pv.PolyData): self.mesh = mesh def add_new_points(self, points_to_add: npt.NDArray): try: validated_input = PointsInput(points_to_add=points_to_add) # 后续操作使用validated_input.points_to_add即可 except ValidationError as e: raise ValueError(f"输入参数不合法: {e}") from e pass
(2) enforce库(轻量级运行时校验)
enforce是专门的运行时类型检查库,支持直接标注Numpy数组的形状和类型,通过装饰器自动校验:
import numpy as np import numpy.typing as npt import pyvista as pv from enforce import runtime_validation, types class SurfaceTriangleMesh: def __init__(self, mesh: pv.PolyData): self.mesh = mesh @runtime_validation def add_new_points(self, points_to_add: types.NumpyArray[np.float64, (*, 3)]): # 方法内部无需额外校验,不符合要求会自动抛出异常 pass
2. 静态类型检查(开发阶段校验)
如果只需要在开发阶段做静态检查,不需要运行时断言,可以使用mypy配合mypy-numpy插件。它会根据你标注的类型(比如npt.NDArray[Tuple[Any, 3]])在代码检查阶段提示形状不匹配的问题,但不会在运行时生效。
3. 手动断言(无依赖方案)
如果不想引入第三方库,手动编写断言也足够简洁直接:
def add_new_points(self, points_to_add: npt.NDArray[np.float64]): assert isinstance(points_to_add, np.ndarray), "输入必须是Numpy数组" assert points_to_add.ndim == 2 and points_to_add.shape[1] == 3, f"输入数组形状必须为[*, 3],当前为{points_to_add.shape}" assert np.issubdtype(points_to_add.dtype, np.floating), f"输入数组必须为浮点类型,当前为{points_to_add.dtype}" # 后续业务逻辑 pass
内容的提问来源于stack exchange,提问作者IzaeDA
相关产品推荐
相关产品推荐

