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() 関数を提供し、複数の操作を1つの GPU ディスパッチに統合するカーネル結合を実現します。これによりメモリ帯域幅のボトルネックを低減します。また、JavaScript のガベージコレクション環境で大きな配列を管理するために、手動の参照カウンティングメモリモデル(.ref.dispose())を実装しています。

対象ユーザー

NumPy および JAX と互換性のある API を維持しながら、ブラウザ上でポータブルで高性能な ML アプリケーション(インタラクティブな可視化、音声アシスタント、ブラウザ内 LLM 推論など)を構築したい開発者向けです。

特徴

  • 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 などの最適化手法の実装を支援するヘルパーライブラリを含む。

関連

  • プロジェクト
  • プロジェクト
  • プロジェクト
  • プロジェクト
  • プロジェクト