microsoft/microxcaling
PyTorch emulation library for Microscaling (MX)-compatible data formats
解决的问题
该库允许研究人员和数据科学家在 PyTorch 中模拟 MX 兼容的数据格式和 bfloat 量化。它使人们能够在不依赖原生支持这些格式的专用硬件的情况下,探索不同低精度数值格式(如 FP8、FP4 和 INT8)对深度神经网络(DNN)性能和精度的影响。
工作原理
该库通过在更高精度(float32、bfloat16 或 fp16)中执行计算,同时将值限制在目标 MX 或 bfloat 格式的范围和精度内,来模拟低精度格式。它提供了标准 PyTorch 模块和函数(如 torch.matmul、torch.linear 和 torch.nn.LayerNorm)的即插即用替代品。
为了在模拟速度和数值精度上优于原生 PyTorch GPU 操作,该项目包含用于量化的自定义 CUDA 扩展。
适用人群
专为关注 DNN 中量化和数值精度探索的数据科学家和 AI 研究人员设计。
主要亮点
- 广泛格式支持:支持多种 MX 兼容格式,包括 FP8(e4m3、e5m2)、FP4(e2m1)和 INT8。
- 灵活配置:使用
mx_specs字典配置比例位数、权重和激活的元素格式以及块大小。 - 无缝集成:提供两种集成路径:手动替换 PyTorch 模块,或通过
mx_mapping.inject_pyt_ops自动注入操作。 - 高性能:包含自定义 CUDA 内核,以避免已知的 PyTorch GPU 数值不准确问题并提升模拟速度。
- 全面覆盖:涵盖前向和反向传播量化,以及 GELU、Softmax 和 LayerNorm 等元素级运算。
相关
- 项目
- 项目
- 项目
- 项目
- 项目