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와 같은 최적화 알고리즘 구현을 위한 보조 라이브러리 포함.
관련
- 프로젝트
- 프로젝트
- 프로젝트
- 프로젝트
- 프로젝트