microsoft/microxcaling

PyTorch emulation library for Microscaling (MX)-compatible data formats

해결하는 문제

이 라이브러리는 PyTorch 내에서 MX 호환 데이터 형식과 bfloat 양자화를 에뮬레이션할 수 있게 해줍니다. 전용 하드웨어를 필요로 하지 않고도, FP8, FP4, INT8과 같은 다양한 저정밀도 수치 형식이 딥 뉴럴 네트워크(DNN)의 성능과 정확도에 미치는 영향을 탐색할 수 있습니다.

작동 방식

이 라이브러리는 계산을 더 높은 정밀도(float32, bfloat16, 또는 fp16)로 수행하면서도, 타겟 MX 또는 bfloat 형식의 범위와 정밀도에 값이 제한되도록 시뮬레이션합니다. torch.matmul, torch.linear, torch.nn.LayerNorm와 같은 표준 PyTorch 모듈과 함수에 대한 드롭인 대체품을 제공합니다.

기존 PyTorch GPU 연산의 알려진 수치 부정확성을 피하고 시뮬레이션 속도를 향상시키기 위해, 프로젝트에는 양자화를 위한 커스텀 CUDA 확장이 포함되어 있습니다.

대상 사용자

DNN에서 양자화 및 수치 정밀도 탐색에 집중하는 데이터 사이언티스트 및 AI 연구자에게 적합합니다.

주요 특징

  • 광범위한 형식 지원: FP8(e4m3, e5m2), FP4(e2m1), INT8 등 다양한 MX 호환 형식을 지원합니다.
  • 유연한 구성: mx_specs 사전을 사용하여 스케일 비트, 가중치 및 활성화의 요소 형식, 블록 크기를 구성할 수 있습니다.
  • 무결성 통합: PyTorch 모듈을 수동으로 교체하거나 mx_mapping.inject_pyt_ops를 통해 연산을 자동으로 삽입하는 두 가지 통합 경로를 제공합니다.
  • 고성능: PyTorch GPU의 알려진 수치 부정확성을 피하고 시뮬레이션 속도를 높이는 커스텀 CUDA 커널을 포함합니다.
  • 포괄적인 커버리지: 전방 및 역방향 전파 양자화, GELU, Softmax, LayerNorm과 같은 요소 연산을 모두 커버합니다.

관련

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