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 位浮點與整數格式。
  • 低精度處理:提供指引與工具,以管理運算過程中的精度損失,例如建議使用更高精度的累加。

相關

  • 專案
  • 專案
  • 專案
  • 專案
  • 專案