patrick-kidger/jaxtyping

Type annotations and runtime checking for shape and dtype of JAX/NumPy/PyTorch/etc. arrays. https://docs.kidger.site/jaxtyping/

解决的问题

解决在不同机器学习框架之间追踪和验证数组与张量形状和数据类型的问题。在复杂模型中,维度不匹配常常导致难以调试的运行时错误。

工作原理

该库提供专门的类型注解,允许开发者指定张量的预期形状和数据类型(例如 Float[Tensor, "dim1 dim2"])。这些注解与 typeguardbeartype 等运行时类型检查工具兼容,可在执行时强制实施这些约束。

适用人群

使用 JAX、PyTorch、NumPy、TensorFlow 或 MLX 的研究人员和工程师,希望提升代码的可读性、可维护性,并减少与形状相关的错误。

主要亮点

  • 支持 JAX、PyTorch、NumPy、TensorFlow 和 MLX 等多个框架。
  • 与兼容库配合使用时具备运行时类型检查能力。
  • 对数组维度和轴提供清晰、描述性的注解。

相关

  • 项目
  • 项目
  • 项目
  • 项目
  • 项目