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_e4m3、float8_e5m2)。 - 微缩缩放(MX)格式:子字节浮点表示(4位和6位)。
- 窄整数:1、2 和 4 位整数类型(以展开的字节形式存储)。
- 复数类型:16 位复数浮点数(
complex32和bcomplex32)。
适用人群
需要在训练或推理中使用低精度数值格式,但又无需重写 NumPy 核心引擎的机器学习库开发者和研究人员。
主要亮点
- NumPy 集成:将 dtype 注册到 NumPy,使其可通过字符串名称引用。
- 广泛格式支持:实现了多种 8 位、4 位和 6 位浮点及整数格式。
- 低精度处理:提供指导和工具,以管理算术运算中的精度损失,例如推荐使用更高精度的累加。
相关
- 项目
- 项目
- 项目
- 项目
- 项目