pyro-ppl/numpyro

Probabilistic programming with NumPy powered by JAX for autograd and JIT compilation to GPU/TPU/CPU.

解決する課題

NumPyroは、モデルに対してベイズ推論を実行するために設計された確率的プログラミングライブラリです。JAXを活用して自動微分とJust-In-Time (JIT) コンパイルを行うことで、確率的プログラミングにおける推論速度の遅さという問題を解決し、CPU、GPU、およびTPU上でモデルを効率的に実行できるようにします。

仕組み

NumPyroは、Pyro確率的プログラミング言語に対してNumPyのようなバックエンドを提供します。JAXを使用して推論アルゴリズムの統合ステップ全体をXLA最適化カーネルにコンパイルし、Pythonのオーバーヘッドを削減します。このライブラリには、包括的な分布クラス、制約、および双射変換が含まれており、カスタム推論アルゴリズムを実装するためのエフェクトハンドラーをサポートしています。

対象者

複雑な確率モデリングおよびベイズ推論を実行する必要がある研究者や開発者、特にPyroやPyTorchの分布APIに精通している方を対象としています。

ハイライト

  • JAXによるパフォーマンス: JITコンパイルとautogradを使用して、ハードウェアアクセラレータ上でのMCMCおよび変分推論を加速します。
  • コアMCMCアルゴリズム: No-U-Turn Sampler (NUTS)、Hamiltonian Monte Carlo (HMC)、離散変数用のMixedHMC、および大規模データセット用のHMCECSを実装しています。
  • 変分推論: 離散潜在変数を含むモデル向けの柔軟なガイドを備えた、自動微分変分推論 (ADVI) をサポートしています。
  • 柔軟な分布: 幅広い分布クラスを提供し、TensorFlow Probability (TFP) の分布もサポートしています。
  • PythonでC++のような速度: Iterative NUTSを通じて、NUTSのツリー構築段階におけるPythonのオーバーヘッドを排除します。

関連

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