wasm 上で動く Stan ランタイムを作った
最近、データ解析と WebAssembly にハマっています。
なんか、AI・コンピュータの総合格闘技感があって好きです。
ベイズ推論を Web でできたら面白いだろうなぁという思考停止な理由で stan を WebAssembly 上で動かせる stanwasm を作成しました。
紹介がてらその中での思考過程も言語化できればと思っています。
stanwasm は Stan 言語のサブセットを、パースからコンパイル、サンプリングまでブラウザの中だけで完結させる実装です。
パーサ・評価器・自動微分テープ・コード生成器はこのプロジェクトで Rust でゼロから書き、サンプリングだけは Adrian Seyboldt と PyMC 開発者による nuts-rs を wasm にコンパイルして借りています。
まだAPIが変わり得るalphaですが、npm install [email protected] で読み込めます。
import init, { StanModel } from "stanwasm";
await init();
const model = new StanModel(stanCode, JSON.stringify(data));
const draws = model.sample(model.randomInit(42n), 1000, 1000, 42n);
// 初期値, warmup, サンプリング数, seed
CmdStan や Stan Playground を置き換えるものではなく、それらが合わない「ブラウザに埋め込む」用途を狙っています。
ただ、色々実装を検討した結果、ブラウザの中だけでシンプルに完結させるために Stan の既存資産をほとんど使わない設計になりました。
なので実態としては、「Stan 言語で記述されたモデルをもとに NUTS で計算できる独自ランタイム」を作ったことになります。
Stan を wasm 化するときの留意点
最初に Stan のツールチェーンについてさらっと整理します。
Stan は C++ をベースに実装されており、
確率グラフィカルモデルを Stan の文法で記述すると、それをもとに、MCMC などで事後分布を計算するためのサンプリング処理がコンパイルされます。
Stan 言語で記述された確率分布はパースされたあと、HMC などでサンプリングできるように導関数へと導かれます。その微分を担うのが、高速化のために最適化された自動微分ライブラリ、Stan Math Library です。
この C++ によって長年メンテナンスされてきたコードにより、Stan は事後分布のサンプリングを高速に処理できております。
C++ で実装されているなら emscripten で wasm 化なんてすぐにできるのではないか?そう思われる方もいらっしゃるでしょうが、実は実現するにはいくつかの課題があります。
それは Stan で事後分布を計算するまでの処理には、ツールチェーンとしてコンパイラが関わっていることにあります。
整理すると、こういう構造になっています。
model.stan は、まず stanc3 が C++ のソース (model.hpp) に変換します。
これを自動微分の Stan Math と、NUTS / HMC を持つ stan::services と一緒に、手元の C++ コンパイラでビルドします。
出てくるのはそのモデル専用の実行ファイルで、走らせるとドローが CSV で出ます。
導関数もサンプラーも、この時点で実行ファイルの中に取り込まれています。
どのインターフェースを使うにせよ、このビルドが走るので手元に C++ コンパイラが要ります。
データだけは実行時に渡すので、同じ実行ファイルを別のデータで使い回せます。
もしインターフェースとして web 上で Stan 言語を記述する場合、当然ランタイムはどの分布や関数が来るかを事前に知ることができません。
つまり選択肢が 2 つに絞られます。
ひとつは、従来の経路をそのままブラウザで再現することです。
stanc3 と Stan Math に加えて、C++ コンパイラ自体をブラウザに持ち込むことになります。
もうひとつは、Stan に出てくる分布や関数をあらかじめ全部コンパイルしておき、その一式をランタイムとして配ることです。
どのモデルが来るか分からない以上、実際には使われないものまで含めて積むことになります。
コンパイラは要りませんが、そのぶん配布物が大きくなります。
なので Stan を wasm で動かすための問いとしては、「Stan言語のツールチェーンをどうやってwasmの世界でコンパクトに表現するか」、「そのためにStan のソースとサンプラーの間にどんな表現が置かれるべきか」、ということです。一旦自分が考えた答えは、自動微分グラフを明示的な中間表現とし、その中間表現から wasm の命令を出力することでした。
stanwasm はどう組み立てたか
その前に、なぜ自動微分とテープというものが要るのかを整理しておきます。
NUTS や HMC は、事後分布の勾配を見て次の点を決めます。
サンプリングを回すには、モデルに出てくる確率分布それぞれが対数密度にどう効くか、その勾配が分かっている必要があります。
モデルの中で分布や演算が繋がると、勾配は連鎖律で合成されていきます。
これを自動で計算するのが自動微分です。
合成は出力側から入力側へ逆向きに進むので、対数尤度をどの演算からどの順序で組み立てたかが分かっていないと計算できません。
そこで計算の過程そのものを記録しておきます。
この記録がテープと呼ばれるものです。
Stan はモデルを実行ファイルへコンパイルする過程で、そのモデルの導関数を自動微分で計算するコードごと実行ファイルに取り込みます。
モデルごとにビルドするので、テープに積むノードの型もそのモデルに合わせて決まります。
こうして高速な事後分布の計算を実現しています。
stanwasm と分かれるのは、そのテープをいつ作るのかと、作ったものをどう使うのかです。
同じ枠組みで stanwasm 側を並べると、こうなります。
stanwasm はテープをロード時に 1 回だけ作ります。new StanModel(src, data) の中でパースし、データを束ね、評価しながら記録するところまでが済みます。
前節で見た経路では、C++ をコンパイルした時点で、導関数を構築して実行する演算が実行ファイルの中に入ります。
stanwasm は、ブラウザで特定のモデルを見た後で、その順伝播と逆伝播を wasm モジュールに変えます。それを可能にしている違いは、stanwasm が微分を明示的にデータとして保持している点です。Stan Math も実行時にテープは作ります。var_stack_ が vari_base* の列で、grad() がそれを逆順に読みます。違いはテープの有無ではなく、ノードが何を持っているかです。
毎回テープを読むか、一度だけ命令を出力するか
stanwasm の Tape は、op arg1 val grad を同じ長さの配列として並べて持つだけの構造体です。op[i] に入るのは「この行が何の演算か」を表す番号で、足し算なら 1、掛け算なら 3 といった具合です。Op::Mul はその 3 に付いた名前にすぎません。
演算の種類が、ただの数値としてそこに置いてあるので、逆伝播はそれを読んで自分で分岐できます。
テープが値の列であることの意味は、読み方を変えられる点にあります。先頭から順に読めば対数密度の値、末尾から逆順に読めば勾配、そして順に読みながら命令を出力すれば、そのモデル専用の wasm になります。
なぜそんなことをするのかを先に書きます。
NUTS は勾配を数万回呼びますが、そのあいだ計算の構造は一度も変わりません。
変わるのはパラメータの値だけです。
テープを読む経路は、その確定した構造を勾配のたびに読み直しています。
構造が最初から決まっているなら、それを埋め込んだコードを 1 回だけ作って、あとはそれを呼ぶだけで済むはずです。
インタプリタとコンパイル済みコードの差、と言い換えてもいい。
やることは単純です。
逆伝播は op[i] を読んで「この演算の微分はこう」と分岐しますが、その分岐先を「この演算に対応する wasm 命令はこれ」に差し替えます (stanwasm-codegen)。
読んでいるテープの形は同じで、Op::Mul を見たときにその場で掛けるか、掛け算の命令を 1 個並べるかだけが違います。
実装ではロード時のテープを使い回さず、ダミーのパラメータで記録し直したものを読みます。
これを順と逆順の両方について行い、対数密度を計算する命令と勾配を積む命令を log_prob_grad という 1 つの関数にまとめます。
対数密度と勾配を切り替えているのではなく、両方の命令が同じ関数の中に並んでいます。 NUTS が要るのも常にこの 2 つ同時だからです。
Stan Math では同じことができません。
ノードに付いた chain メソッドを呼べば微分は進みますが、「これは何の演算か」を値として取り出せません。 どの chain を呼ぶかは C++ をコンパイルした時点で決まっているからです。
書き出すしくみを足すには Stan Math をビルドし直すことになり、コンパイラが要るという最初の問題に戻ります。
命令が f64x2 の SIMD なのも、モデル専用だからできることです。汎用にテープを読む形では次に何の演算が来るか分からないので、倍精度 2 つをまとめる判断ができません。
実装でわかったこと
ざっと実装した結果以下のような形で実現できてそうです。
| サイズ | 664 KB、gzip で 253 KB。npm パッケージに全部入っていて、モデルごとの追加ダウンロードはありません |
| コールドスタート | 空白ページから最初のドローまで、50 Mbps で 0.4 秒、1.6 Mbps で 1.6 秒 |
| 動作環境 | Chromium・Firefox・WebKit。iOS も含みます |
| 実行経路 | テープを読む経路と、モデル専用 wasm を出力する経路。形がパラメータで変わるモデル用に、勾配ごとに取り直す経路もあります |
正しさは posteriordb(Stan のモデルと参照結果を集めたベンチマーク集)で確かめています。
| 確かめたこと | 結果 |
|---|---|
| ロード・勾配評価・コンパイルが通るか | 147 posterior のうち 141 |
| 対数密度と勾配が CmdStan 2.38.0 と一致するか | 16 モデルで最悪 3.3e-13、うち 13 モデルは 1e-14 以下 |
| 事後平均が参照ドローと合うか | 参照付き 45 posterior のうち 42 が 0.2 sd 以内、0.5 sd を超えるものは無し |
| モデル専用 wasm はどれだけ速いか | テープを読む経路より勾配あたり 1.5〜12 倍 |
事後平均のチェックは平均だけで、R-hat・ESS・divergence はまだ比較していません。
1.5〜12 倍も自分の 2 経路の内部比較であって、Stan との比較ではありません。
ただ、自前でスクラッチ書いた部分が多いため、犠牲になっている部分も多々あります。
- 言語は Stan のサブセットであって、全体ではありません
- サンプラーが独立したコードなので、同じ seed でも Stan のドローは再現しません
- Stan Math の修正は流れてきません。特殊関数も分布のエッジケースも、自前で書いてテストしました
そして何より、計算の形がパラメータによって変わるものは、一度だけ記録するやり方と相性が悪いです。
適応的な ODE ソルバーがその例で、パラメータ次第でステップの取り方が変わるため、記録したグラフはトレースした時点の形で固まります。
記録して再生する経路ではこれを拒否し、代わりに sampleFresh() が勾配のたびにテープを取り直します。integrate_ode_rk45 はこの経路で動きますが、再生の 5〜8 倍かかります。
ステップ数を固定する ode_rk4_fixed ならグラフが動かないので、再生にも AOT にも乗ります。integrate_ode_bdf はまだありません。
Stan のソースとサンプラーの間に何を置くかという問いに、妥当な答えは一つではありません。
プリコンパイル済みのカーネルの上に静的グラフを置くのも一つで、互換性・成果物のサイズ・特殊化の度合いのトレードが違うだけだと思っています。
実際、stanli がその道を行っています。
stanwasm を作ってから存在を知りました。
Stan Math などの既存資産をうまく使いながら、モデルごとの自動微分処理をインタプリタとして wasm に載せており、素晴らしいなと思いました。
おわりに
Stan のモデルをブラウザだけで動かすために、自動微分グラフを明示的な中間表現として持ち、そこから wasm の命令を出力する形にしました。
まだ alpha で言語のサブセットも狭いですが、インストールなしで試せるものにはなっています。
ベイズのことを考えていたはずなのに、気づいたらアセンブリやコンパイラの勉強になっていました。