ekzhang/jax-js

JAX in JavaScript – ML library for the web, running on WebGPU & Wasm

해결하는 문제

jax-js는 브라우저에서 고성능 수치 계산과 JAX 스타일 연산을 가능하게 하는 머신러닝 프레임워크입니다. 백엔드 서버 없이도 CPU 및 GPU 가속을 활용하여 복잡한 수학적 애플리케이션, 신경망, 시뮬레이션을 클라이언트 측에서 직접 실행할 수 있습니다.

작동 방식

이 프레임워크는 배열 연산을 컴파일러 표현으로 변환한 후, WebAssembly (Wasm) 및 WebGPU용 최적화된 커널을 생성합니다. 여러 디바이스 백엔드를 지원합니다:

  • WebGPU: 고성능과 신경망 처리를 위해 주로 추천되는 선택지.
  • Wasm: 멀티스레드 CPU 백엔드.
  • WebGL: 오래된 브라우저용 백업 옵션.
  • CPU: 디버깅을 위한 인터프리터 기반 JS 백엔드.

성능 최적화를 위해 jit() 함수를 제공하여 여러 연산을 하나의 GPU 디스패치로 통합하는 커널 퓨전을 지원합니다. 이는 메모리 대역폭 병목을 줄이는 데 기여합니다. 또한 JavaScript의 가비지 컬렉션 환경에서 큰 배열을 관리하기 위해 수동 참조 카운팅 메모리 모델(.ref.dispose())을 구현했습니다.

대상 사용자

NumPy 및 JAX와 호환되는 API를 유지하면서, 인터랙티브 시각화, 음성 어시스턴트, 브라우저 내 LLM 추론과 같은 포터블하고 고성능의 ML 애플리케이션을 브라우저에서 개발하고자 하는 개발자에게 적합합니다.

주요 특징

  • JAX 스타일 변환: grad()를 통한 자동 미분, vmap()를 통한 자동 벡터화, jit()를 통한 커널 퓨전을 지원.
  • 고성능: CPU에서는 OpenBLAS와 유사한 행렬 곱셈 속도를 달성하고, 고성능 Apple Silicon에서는 WebGPU를 통해 7000 GFLOP/s 이상의 성능을 구현.
  • 광범위한 호환성: Chrome, Firefox, Safari, Node.js/Deno에서 동작하며 Float16, Float32, Float64를 지원.
  • 의존성 없음: 외부 의존성이 없는 순수 자체 구현.
  • 에코시스템: Safetensors 로딩, ONNX 모델 임포트, Adam 및 SGD와 같은 최적화 알고리즘 구현을 위한 보조 라이브러리 포함.

관련

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