jax-ml/ml_dtypes

A stand-alone implementation of several NumPy dtype extensions used in machine learning.

解决的问题

提供一种独立的实现,支持机器学习中常用但 NumPy 原生不支持的专用数值数据类型(dtypes)。这使得 ML 库可以使用低精度格式来减少内存占用并提高计算效率。

工作原理

该库实现了多个 NumPy dtype 扩展,并将其注册,以便可在 NumPy 数组中直接使用。支持广泛的格式,包括:

  • Bfloat16:单精度浮点数的截断版本。
  • 8位浮点数:指数和尾数位数的各种配置(例如 float8_e4m3float8_e5m2)。
  • 微缩缩放(MX)格式:子字节浮点表示(4位和6位)。
  • 窄整数:1、2 和 4 位整数类型(以展开的字节形式存储)。
  • 复数类型:16 位复数浮点数(complex32bcomplex32)。

适用人群

需要在训练或推理中使用低精度数值格式,但又无需重写 NumPy 核心引擎的机器学习库开发者和研究人员。

主要亮点

  • NumPy 集成:将 dtype 注册到 NumPy,使其可通过字符串名称引用。
  • 广泛格式支持:实现了多种 8 位、4 位和 6 位浮点及整数格式。
  • 低精度处理:提供指导和工具,以管理算术运算中的精度损失,例如推荐使用更高精度的累加。

相关

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