jax-ml/ml_dtypes

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

해결하는 문제

기계학습에서 흔히 사용되지만 NumPy에서 기본적으로 지원하지 않는 특수한 수치 데이터형(dttype)의 독립형 구현을 제공합니다. 이를 통해 ML 라이브러리는 메모리 사용량을 줄이고 계산 효율을 높이기 위해 저정밀도 형식을 사용할 수 있습니다.

작동 방식

이 라이브러리는 여러 NumPy dtype 확장을 구현하고, 이를 NumPy 배열 내에서 직접 사용할 수 있도록 등록합니다. 다음과 같은 다양한 형식을 지원합니다:

  • Bfloat16: 단정밀도 부동소수점의 단순화된 버전.
  • 8비트 부동소수점: 지수와 가수 비트의 다양한 구성 (예: float8_e4m3, float8_e5m2).
  • 마이크로스케일링(MX) 형식: 바이트 미만의 부동소수점 표현 (4비트 및 6비트).
  • 좁은 정수형: 1, 2, 4비트 정수형 (확장된 바이트로 저장).
  • 복소수형: 16비트 복소 부동소수점 수 (complex32bcomplex32).

대상 사용자

NumPy의 핵심 엔진을 다시 작성하지 않고도 학습 또는 추론 시 저정밀도 수치 형식을 사용하고자 하는 기계학습 라이브러리 개발자 및 연구자.

주요 특징

  • NumPy 통합: dtype을 NumPy에 등록하여 문자열 이름으로 참조 가능하게 합니다.
  • 광범위한 형식 지원: 8비트, 4비트, 6비트 부동소수점 및 정수 형식을 다양하게 구현합니다.
  • 저정밀도 처리: 산술 연산 중 정밀도 손실을 관리하기 위한 가이드라인과 도구를 제공하며, 높은 정밀도 누적을 권장합니다.

관련

  • 프로젝트
  • 프로젝트
  • 프로젝트
  • 프로젝트
  • 프로젝트