2026-09-06

wasm 上で動く Stan ランタイムを作った

#stanwasm #Stan #WebAssembly #Rust

最近、データ解析と WebAssembly にハマっています。
なんか、AI・コンピュータの総合格闘技感があって好きです。
ベイズ推論を Web でできたら面白いだろうなぁという思考停止な理由で stan を WebAssembly 上で動かせる stanwasm を作成しました。
紹介がてらその中での思考過程も言語化できればと思っています。
stanwasm は Stan 言語のサブセットを、パースからコンパイル、サンプリングまでブラウザの中だけで完結させる実装です。
パーサ・評価器・自動微分テープ・コード生成器はこのプロジェクトで Rust でゼロから書き、サンプリングだけは Adrian Seyboldt と PyMC 開発者による nuts-rs を wasm にコンパイルして借りています。

stanwasm でのベイズ推論のデモ。インストールなしでブラウザ内でサンプリングが走る
habakan/stanwasmブラウザだけで動く Stan サブセット。Apache-2.0、npm では stanwasmgithub.com

まだ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 で事後分布を計算するまでの処理には、ツールチェーンとしてコンパイラが関わっていることにあります。

整理すると、こういう構造になっています。

Stan のツールチェーン Stan で書いたモデルは OCaml 製の stanc3 が C++ へ変換し、その C++ を自動微分の Stan Math とサンプラーの stan::services と一緒に手元の C++ コンパイラでビルドして実行ファイルにする。走らせるとサンプルが CSV で出る。CmdStan はこのビルドを回して CLI を付けるもので、CmdStanPy と CmdStanR はその CmdStan を呼ぶ。RStan と PyStan は同じ部品を自前の経路で組み込む。Stan 本体のリポジトリは Stan Math を submodule として抱えており、その Stan Math が Eigen・Boost・TBB・SUNDIALS を抱えている。 Stan のツールチェーン インターフェース CmdStan · CmdStanPy · CmdStanR · RStan · PyStan 開発者が呼ぶ API。いずれも下のビルドを回すので、C++ コンパイラが要る data JSON / Rdump compile() sample(data=...) ビルドの経路 — インターフェースがこれを回す model.stan Stan 言語 stanc3 Stan → C++ model.hpp モデルコード (C++) C++ コンパイラ 実行ファイル モデルごと draws (CSV) コンパイル時に取り込まれるもの STAN 本体 — STAN-DEV/STAN Stan Math — 自動微分 SUBMODULE · STAN-DEV/MATH stan::services · mcmc NUTS / HMC Eigen Boost TBB SUNDIALS 導関数もサンプラーも、コンパイル時に 実行ファイルへ取り込まれる

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 のツールチェーン new StanModel にソースとデータを渡すと、ロード時に一度だけパースとデータの束ね込みが走り、評価しながら自動微分テープに記録する。記録したテープを使う経路が 2 つある。既定の sample は nuts-rs がそのテープを毎回読んで勾配を計算する。もう一方は compileToWasm を呼ぶ経路で、これはテープではなくパース済みのモデルを入力に、ダミーのパラメータで評価しながらテープを記録し直し、その場でモデル専用の wasm を出力する。出力したバイト列は JS 側で instantiate して setAotExports で結び、sampleViaAot を呼ぶと nuts-rs がその wasm を呼ぶ。同じモデルから出た 2 経路。 stanwasm のツールチェーン インターフェース npm: stanwasm (JS)· pystanwasm (Pyodide 上の Python) 開発者が呼ぶ API。どちらも同じ wasm を読み込むだけで、ビルド工程は無い data JSON new StanModel(src, data) ロード時に 1 回だけ model.stan Stan 言語 parser 再帰下降 autodiff tape 評価しながら記録 評価しながら記録する。 記録はここ 1 回だけ。 データはここで束ねられ、テープに入る。 codegen はまだ走らない。 記録したテープを使う 2 経路 sample(...) 既定。上のテープをそのまま使う nuts-rs 上と同じテープ 勾配のたびに op を読んで分岐する。 テープは作り直さない。 draws compileToWasm() → sampleViaAot(...) 入力は上のテープではなく、パース済みのモデル codegen 記録し直す モデル専用 wasm nuts-rs draws 勾配のたびに この wasm を呼ぶ。 分岐も間接参照も無い。

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 になります。

自動微分テープの中身と、その読み方三通り テープの各行は、演算とその入力、そして結果を置く場所からなる。入力はデータの y と、パラメータの mu と sigma。順に読めば対数密度の値が出る。逆順に読めば各中間値の偏微分が積み上がり勾配が出る。コード生成はこの両方を命令として書き出し、対数密度と勾配を 1 回で返す関数 log_prob_grad にまとめる。勾配のたびに読むのではなく、作るときに 1 度読むだけになる。 テープの 1 行は「演算・その入力・結果の置き場」 入力 y DATA mu PARAM sigma PARAM TAPE #1 Sub y, mu → v1 #2 Div #1, sigma → v2 #3 Mul #2, #2 → v3 #4 Log sigma → v4 lp = -0.5 * v3 - v4 - 0.5*log(2pi) 順に読む 各行を順に計算し、 対数密度 lp が出る 逆順に読む 各 v の ∂lp/∂v を積み、 勾配が出る 同じテープを、計算せずに命令として書き出す 読む向きは上と同じ。違うのは、その場で値を出さずに、行ごとに対応する wasm 命令を並べる点。 READ AND EMIT 順に読む → 値 #1 f64x2.sub #2 f64x2.div #3 f64x2.mul #4 call log 逆順に読む → 勾配 #4 f64x2.div #3 f64x2.mul #2 f64x2.div #1 f64x2.add export log_prob_grad 対数密度と勾配を 1 回で返す 切り替えではなく、両方が同じ関数の中に並んでいる NUTS-RS CALLS IT FOR EVERY GRADIENT
同じテープを順に読めば値、逆順に読めば勾配、読みながら命令を出力すれば wasm になる

なぜそんなことをするのかを先に書きます。

NUTS は勾配を数万回呼びますが、そのあいだ計算の構造は一度も変わりません。
変わるのはパラメータの値だけです。
テープを読む経路は、その確定した構造を勾配のたびに読み直しています。
構造が最初から決まっているなら、それを埋め込んだコードを 1 回だけ作って、あとはそれを呼ぶだけで済むはずです。
インタプリタとコンパイル済みコードの差、と言い換えてもいい。

やることは単純です。
逆伝播は op[i] を読んで「この演算の微分はこう」と分岐しますが、その分岐先を「この演算に対応する wasm 命令はこれ」に差し替えます (stanwasm-codegen)。
読んでいるテープの形は同じで、Op::Mul を見たときにその場で掛けるか、掛け算の命令を 1 個並べるかだけが違います。
実装ではロード時のテープを使い回さず、ダミーのパラメータで記録し直したものを読みます。

これを順と逆順の両方について行い、対数密度を計算する命令と勾配を積む命令を log_prob_grad という 1 つの関数にまとめます。
対数密度と勾配を切り替えているのではなく、両方の命令が同じ関数の中に並んでいます。 NUTS が要るのも常にこの 2 つ同時だからです。

Stan Math では同じことができません。
ノードに付いた chain メソッドを呼べば微分は進みますが、「これは何の演算か」を値として取り出せません。 どの chain を呼ぶかは C++ をコンパイルした時点で決まっているからです。
書き出すしくみを足すには Stan Math をビルドし直すことになり、コンパイラが要るという最初の問題に戻ります。

記録したテープを、どう使うか ロード時に作った自動微分テープの使い方が二手に分かれる。左はテープを読む側で、勾配のたびにノードごとに op を読んで分岐し、値をロードして間接参照する。既存の Stan の資産を使う道も、呼び方は違うがノードごとにディスパッチする点では同じ側にいる。Stan Math ならノードに付いた chain メソッドを呼ぶ形で、何の演算かを値として取り出すことはできない。右は出力する側で、作るときにテープを読んで対応する wasm 命令を並べたモジュールにするので、勾配のたびの分岐も間接参照も無くなる。こちらが勾配あたり 1.5 から 12 倍速い。 記録したテープを、どう使うか autodiff tape 上の図で作ったもの 毎回読む — 勾配のたびにテープを逆順に読む sample(...) match self.op[i] { Op::Mul => grad[a1] += g * vb, 勾配のたびに、ノードごとに繰り返すこと: op[i] をロード arg1[i] をロード match で分岐 · grad[a1] を間接参照 既存の Stan の資産を使う道も、同じ側 プリコンパイル済みのランタイムを配る形でも、 ノードごとのディスパッチは残る。 var_stack_[i]->chain() Stan Math ならこう。ノードに付いた chain を 呼ぶと微分が進む。呼べば動くが、 「何の演算か」は値として取り出せない。 呼び方は違っても、毎回ディスパッチする点は同じ。 出力する — 作るときに読むだけで、勾配のたびには読まない compileToWasm() → sampleViaAot(...) match self.op[i] { Op::Mul => f.instruction(F64x2Mul), できた wasm の中身。勾配のたびにこれが走る: local.get 2 · local.get 5 f64x2.mul local.set 7 できた wasm には分岐も間接参照も無い。番号は命令の中の定数。 WAT に近い表記で書いているが、実際の出力はバイナリ。 1.5〜12 倍 同じ match の腕を、値を計算するか 命令を出力するかで差し替えているだけ。 テープが実行されるコードではなく データだから、これができる。

命令が 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 で言語のサブセットも狭いですが、インストールなしで試せるものにはなっています。

ベイズのことを考えていたはずなのに、気づいたらアセンブリやコンパイラの勉強になっていました。