国里愛彦(くにさと よしひこ)

  • 専修大学人間科学部心理学科 教授
  • 専門:計算論的精神医学、認知行動療法
  • COI開示:国里の著書や関連する団体について紹介するが、特定企業との利益相反はない。
  • 本発表のPythonコードの作成、校正、情報検索にあたり生成AI(主にgpt-5.6-sol)を使用している。

本日の内容

  • 計算論的精神医学
  • 能動的推論モデル
  • pymdpによる能動的推論モデルの実装

計算論的精神医学

計算論的精神医学のはじまり

  • 2010年代に、計算論的神経科学の精神医学への応用がはじまり、計算論的精神医学としてまとまる。
  • 2011年くらいから複数の雑誌に計算論的精神医学に関する総説が掲載(Huys, Moutoussis, and Williams 2011; Montague et al. 2012)。 2017年には、『Computational Psychiatry』と題する書籍も出版(Anticevic and Murray 2017)
  • 精神障害に関連した神経・認知的現象の数理モデルの構築を目的とする学際的研究領域

日本での計算論的精神医学

計算論的精神医学コロキウム

  • 計算論的精神医学研究の推進を目的として2025年に社団法人化
  • CPSY TOKYOを毎年開催(今年度は第5回目を2027年2月18日-19日に開催予定)
  • 2026年8月24日に初学者向けに「Claude Codeではじめる計算論的精神医学」を開催した

計算論的とは?

  • 私達や動物の行動の背後にある計算過程を数式で表現したモデルを用いること。

  • ここでの計算は意識的な計算に限定されず、脳内のあらゆる情報処理においてなされる無意識的な計算も含む。

  • 計算論的神経科学や計算論的精神医学は、計算論的アプローチをとる。

計算論的精神医学の基本的枠組み

4つの生成モデル

生成モデル、シミュレーション、パラメータ推定

計算論的精神医学への期待

  • 理解: 症状・行動の水準と生物学的水準との間の説明のギャップを埋める
  • 理論: 自然言語で書かれたモデルの数理モデル化とその統合
  • アセスメント: 計算論的アプローチにより疾患マーカーの洗練化が期待される(計算論表現型)
  • シミュレーション: 効果的な治療の探索や臨床的な経過の予想

能動的推論モデル

ベイズ推論モデルとは

  • ベイズ推論モデルは、ベイズの定理を用いて、エージェントが世界について知る過程をモデル化する。
  • 具体的には、事前信念をデータの情報(尤度)で更新して事後信念を得るプロセスをモデル化する。
  • 計算論的精神医学で用いられるベイズ推論モデル:パラメータ化信念更新モデル、カルマンフィルターモデル、階層ガウシアンフィルター、能動的推論モデル

ベイズ推論モデルの世界観

  • 感覚入力\(o\)の背後には、直接観測できない外部状態 \(s^*\) がある(外部状態 \(s^*\) から感覚入力\(o\)が生成されるプロセスは生成過程\(p(s^*, o)\)と呼ぶ)。
  • 外部状態 \(s^*\) は直接観測できないので、入ってきた感覚入力\(o\)から外部状態 \(s^*\) を推論する。

ベイズ推論モデルの世界観

  • 外部状態 \(s^*\) の推論で、外部状態の信念\(s\)と感覚入力\(o\)の同時確率分布\(p(s,o)\)を用いる(生成モデル)
  • 感覚入力\(o\)から外部状態 \(s^*\) についての信念\(s\)をベイズ更新することが知覚
  • 信念の更新だけでなく、行動選択\(a\)によって外界の状態にも働きかける

生成モデル

  • 生成モデルは、エージェントが外界の生成過程を近似したもの。生成モデルがあれば、エージェントの知覚や行動についてシミュレーションできる。
  • 認知課題などによって行動データがあれば、生成モデルを用いてパラメータ推定を行うこともできる(ただし、能動的推論ではまだ多くはない)。
  • 生成モデルは同時確率分布 \(p(s, o)\) に限定されることもあるが、広くデータを生成するモデルという意味で用いることもある。

ベイズの定理

\[ p(s|o) = \frac{p(o|s)p(s)}{p(o)} \]

  • 【右辺の分子】ある信念\(s\)の事前確率 \(p(s)\) をある信念 \(s\)の下で感覚入力 \(o\)が観測される確率(尤度, \(p(o|s)\) )に掛け合わせることで、観測に伴う世界の状態について信念を更新している。

自由エネルギー原理(free-energy principle)

  • 自由エネルギー原理は、フリストンによって提唱されたベイズ推論モデルをベースにした脳と心の統一的理論(Friston 2010)
  • 知覚・行動・学習を「自由エネルギーの最小化」という一つの原理で統一的に説明しようとする

サプライザル

  • ベイズ推論では、感覚入力\(o\)の得られにくさであるサプライザル(\(-\ln p(o)\))が重要(注意:確率論的な量であり、主観的な「驚き」そのものではない)
  • サプライザルが小さいほど、その生成モデルの予測性能が良い(ただし、その計算は困難または不可能)

変分自由エネルギー\(\textbf{F}\)

  • サプライザル (\(-\ln p(o)\)) は直接計算しにくいため、近似事後分布 (\(q(s)\)) を導入し、サプライザルの上界となる変分自由エネルギー (\(F[q(s)]\)) を最小化する。これにより、\(q(s)\) を真の事後分布 \(p(s\mid o)\) に近づける。
  • \(F\)は以下の式で計算され、近似事後分布\(q(s)\)と生成モデル\(p(s, o)\)があれば計算ができる。

\[ F = \int q(s)\ln \frac{q(s)}{p(s,o)}ds \]

変分自由エネルギー \(\textbf{F}\)

  • 変分自由エネルギーFは、以下のように表現できる。

\[ F = D_{KL}[q(s)\,\Vert\,p(s|o)]- \ln p(o) \]

  • 第1項は、カルバック・ライブラーダイバージェンスであり、確率分布\(q\)\(p\)の差異を表し、ゼロ以上の値を取る(\(q\)\(p\)が同じならゼロ)。第2項はサプライザル。

\(F\)はサプライザルの上界であり、\(F\)を最小にする \(q\) を探索することで、近似事後分布を真の事後分布へ近づける。

階層的神経回路による\(\textbf{F}\)の最小化

  • 上の状態から下の状態へ予測が送られ、感覚入力\(o\)と比較され、予測誤差が生じる。
  • 予測誤差を下から上へ伝達して状態更新(予測誤差最小化)
  • 予測誤差は精度(分散の逆数)で重みづけされ、信頼できる情報ほど推論に影響する(Bogacz 2017)

知覚と行動と学習:\(\textbf{F}\)を下げる方法

  • \(F\)を下げる方法(Friston 2010; Buckley et al. 2017)
    • 知覚:信念\(q(s)\)を更新し、予測を感覚入力に合わせる
    • 行動:外界に働きかけて感覚入力\(o\)自体を変えて、感覚入力を予測に合わせる
    • 学習:生成モデルのパラメータをゆっくり更新する

→知覚・行動・学習が自由エネルギー最小化として統一的に記述できる

期待自由エネルギー\(\textbf{G}\)

  • 行動選択には、未来に生じることを考慮する必要があるが、変分自由エネルギー\(F\)では考慮されていない。

\(F\)ではなく期待自由エネルギー\(G\)を計算する

  • 期待自由エネルギー\(G\)では、可能性のある行為の系列である方策(\(\pi\))に対して生じるであろう未来のデータを生成モデルから予測して計算に用いる。
  • 期待自由エネルギー\(G\)が小さい方策ほど選ばれやすく、結果として\(G\)の最小化が達成される(能動的推論)

期待自由エネルギー\(\textbf{G}\)

\[ \begin{aligned} G(\pi) = - \mathbb{E}_{q(\widetilde{s},\widetilde{o}|\pi)}[D_{KL}[q(\widetilde{s}|\widetilde{o},\pi) \,\Vert\, q(\widetilde{s}|\pi)]] \\ - \mathbb{E}_{q(\widetilde{o}|\pi)}[\ln p(\widetilde{o}|C)] \end{aligned} \]

  • 第1項は未来の観測による状態の不確実性の減少の程度で、新しい情報を求める認識的価値。第2項の\(C\) は選好であり、好ましい観測を求める実利的価値。情報利得(\(D_{KL}\))が大きいほど\(G\)は小さく、好ましい観測が予測されるほど\(G\)は小さくなる
  • 期待自由エネルギーは認識的価値と実利的価値とのバランスで定まり、探索-利用のバランスを扱っている

期待自由エネルギー\(\textbf{G}\)の別の見方

\[ G_\pi =D_{\mathrm{KL}} \left[ q(o\mid\pi)\Vert p(o\mid C) \right] + \mathbb E_{q(s\mid\pi)} \left[ H[p(o\mid s)] \right] \]

  • \(G\)は「リスク+曖昧さ」とも書ける(Smith, Friston, and Whyte 2022)
  • リスク:予測される観測と好ましい観測分布とのズレ(好ましい結果が得られそうにないほど大きい)
  • 曖昧さ:状態が分かっていても、どの観測が得られるか判別しにくい程度
  • \(G\)の最小化=「好ましい観測が得られやすく、状態がはっきり分かる観測が得られる」方策の選択

POMDP(部分観測マルコフ決定過程)

POMDPは、外界の状態を直接観測できない状況で、得られた観測から状態を推論しながら行動する枠組み。

  • 隠れ状態 \(s_t\):直接は見えない外界の状態
  • 観測 \(o_t\):状態について得られる部分的・曖昧な情報
  • 行動 \(a_t\):次の状態に影響を与える選択
  • 信念 \(q(s_t)\):隠れ状態についての確率分布

→POMDPは、ベイズ的な状態推定と逐次的な意思決定を統合した枠組み

能動的推論とPOMDP

  • 能動的推論では、POMDPを単なる外部環境の記述ではなく、エージェントがもつ世界の生成モデルとして用いる
  • 観測から隠れ状態を推論し、各方策のもとで将来の状態と観測を予測する
  • 状態信念 \(q(s_t)\) は変分自由エネルギー \(F\)、方策 \(q(\pi)\) は期待自由エネルギー \(G\) に基づいて更新する

因子グラフでPOMDPを表す

  • 因子グラフは、POMDPの同時確率分布を、局所的な関係である因子の積として表現する
  • 状態\(s\)、観測\(o\)、方策\(\pi\)などの変数を円で表す
  • 状態と観測、現在と次の状態など、変数間の確率的な関係は四角で表す

生成モデルを記述するAからD

  • A :尤度
  • B: 状態遷移
  • C: 選好
  • D: 試行開始時の状態事前分布

生成モデルを記述するAからD

記号 確率分布 役割
A \(p(o_\tau\mid s_\tau)\) 状態から観測への尤度
B \(p(s_{\tau+1}\mid s_\tau,\pi)\) 状態遷移と行動の効果
C \(\ln p(o_\tau\mid C)\) 観測に対する選好
D \(p(s_1)\) 試行開始時の状態事前分布

\(\tau\)は、エージェントの信念が対象とする内部的な時刻を表す。

推論と行動選択に関わる量

  • F:変分自由エネルギー
  • G:期待自由エネルギー
  • \(\gamma\):期待自由エネルギーの精度
  • E:方策の事前分布

\[ q(\pi)=\operatorname{softmax}\left(\ln E-\gamma G(\pi)\right) \]

各量の説明

記号 説明
F 状態事後分布を求めるときに最小化する変分自由エネルギー
G 将来の選好充足と情報獲得を方策ごとに評価する期待自由エネルギー
\(\gamma\) \(G\)の差を方策確率に反映する強さ。大きいほど低い\(G\)の方策に集中する
E 期待自由エネルギーとは別に方策の選びやすさを定める事前分布\(p(\pi)\)

能動的推論での学習

  • \(\textbf{A}\)\(\textbf{B}\)の各条件付き確率の列と\(\textbf{D}\)は、それぞれカテゴリ分布の確率ベクトル (\(\boldsymbol\theta\)) を表す。学習では、この (\(\boldsymbol\theta\)) を固定値とせず、それぞれにディリクレ分布を置く。

\[ p(\boldsymbol\theta\mid\mathbf c)=\operatorname{Dir}(\boldsymbol\theta\mid\mathbf c) =\frac{\Gamma(\sum_k c_k)}{\prod_k\Gamma(c_k)} \prod_k\theta_k^{c_k-1} \]

\(c_k>0\)は、それぞれの結果をどれほど経験したと考えるかを表す擬似カウント\(\textbf{A}\)\(\textbf{B}\)\(\textbf{D}\)に対するディリクレパラメータ(\(\mathbf c\)) を小文字の\(\mathbf a\)\(\mathbf b\)\(\mathbf d\)で表す(Smith, Friston, and Whyte 2022)

ディリクレパラメータ:期待確率と総集中度

\[ c_0=\sum_k c_k, \qquad \mathbb E[\theta_k]=\frac{c_k}{c_0} \]

擬似カウント\(\mathbf c\) 期待確率 総集中度\(c_0\) 経験1回の影響
\([1,1]\) \(0.5/0.5\) \(2\) 大きい
\([50,50]\) \(0.5/0.5\) \(100\) 小さい
  • \(c_k\):ディリクレパラメータ(擬似カウント)
  • \(c_0\):総集中度であり、確率ベクトルに対する信念の強さを表す。

推定した経験を擬似カウントへ加える

\[ \mathbf c_{\mathrm{new}} =\mathbf c_{\mathrm{old}}+\eta\,\Delta\mathbf c \]

学習対象 擬似カウントとして加える経験\(\Delta\mathbf c\)
\(D\) 初期状態についての事後信念\(q(s_1)\)
\(A\) 観測と状態の同時生起\(\sum_\tau\mathbf o_\tau\otimes\mathbf s_\tau\)
\(B\) 連続する状態と行動の同時生起
  • \(\mathbf s_\tau=q(s_\tau)\)は状態についての事後信念、\(\eta\)は新しい経験を加える重み。
  • \(\mathbf a\)の各列を列の合計で割ると、状態\(s\)から観測\(o\)が生じる期待確率\(P(o\mid s)\)になる(\(B\)も同様、\(D\)はベクトル全体を正規化)

学習するために選ぶ

  • パラメータ学習:結果を観測した後に、\(A\)の擬似カウントを更新する。
  • 能動的な学習:選択する前に、その選択が\(A\)の不確実性をどれだけ減らすかを評価する。
  • 方策は期待される選好充足だけでなく、\(A\)について学べる程度でも評価される。総集中度が増えて\(A\)への確信が強まるほど、この情報価値は小さくなる(Smith, Friston, and Whyte 2022)

flowchart TB
    A["Aへの現在の信念"] --> I["期待される<br/>情報利得"]
    I --> C["選択"]
    C --> O["結果を観測"]
    O --> U["Aの擬似カウントを更新"]
    U --> A

階層的推論(深層生成モデル)

  • Sandved-Smith et al. (2021) のメタアウェアネスモデルの生成モデル
  • 上の層は時間的にゆっくり変化
  • 層が増えるほど複雑な現象のモデル化ができる

pymdpによる能動的推論モデルの実装

能動的推論モデルの実装

  • 尤度(\(\textbf{A}\))、遷移確率 (\(\textbf{B}\))、観測に対する事前選好 (\(\textbf{C}\))、初期状態に関する事前信念 (\(\textbf{D}\))を設定すれば、期待自由エネルギーを最小化するように振る舞う能動的推論エージェントを作ることができる。
  • MATLAB上で動作する脳機能画像解析用ソフトであるSPM12 に同封されているspm_MDP_VB_X.mを使う際もこれら\(\textbf{ABCD}\)を設定する。
  • 能動的推論に関しては、MATLABで動作するSPM、Pythonパッケージのpymdp、JuliaパッケージのActiveInference.jlなどがある。

pymdpは、離散状態空間の能動的推論モデルをPythonで構築するオープンソースパッケージ

特徴 できること
生成モデル 複数の状態因子・観測モダリティを、A・B・C・Dなどで定義する
推論と行動 状態推論、方策評価、期待自由エネルギーに基づく行動選択を行う
パラメータ学習 エージェントが経験に応じて、観測モデルAや遷移モデルBなどを更新する
パラメータ推定 行動データにモデルを当てはめ、\(\gamma\)などを推定する(NumPyro・pybefitと連携)
シミュレーション Envは外部環境を表し、rollout()は複数時点の観測・推論・行動・状態遷移を反復する

pymdp v1.0.0以降の特徴

v1.0.0以降はJAXベースになった

JAXでできること

  • jitで反復処理をコンパイルし、CPU・GPU・TPUで実行
  • batch_sizevmapで複数エージェントを並列化
  • 自動微分とNumPyroを用いたパラメータ推定

コード上の違い

  • jax.numpy配列と先頭のバッチ次元を使用
  • 乱数を jax.random.PRNGKey で管理
  • 更新後の値やAgentを明示的に受け取る

Google Colabの使い方

  • Google Colabを使う場合は、このチュートリアル用のリンクにアクセスし、「ファイル」→「ドライブにコピーを保存」をクリックする。

Google Colabの使い方

セットアップ

  • 以下のパッケージ(バージョン)を用いる。
パッケージ バージョン
inferactively-pymdp 1.0.4
jax / jaxlib 0.6.2
equinox 0.13.8
numpyro 0.19.0

バージョンを固定しないと、JAX・Equinox・NumPyroの互換性が崩れる場合がある。初回のインストール後は、ランタイムを再起動してから以降のセルを実行する。

セットアップ

# 最初の行は1度だけ実行
!pip install -q --upgrade inferactively-pymdp==1.0.4 jax==0.6.2 jaxlib==0.6.2 equinox==0.13.8 numpyro==0.19.0

from importlib.metadata import version
import equinox as eqx
import jax
import jax.numpy as jnp
from jax import random as jr
import matplotlib.pyplot as plt
import numpy as np
import pandas as pd
import numpyro
import numpyro.distributions as dist
from numpyro.diagnostics import effective_sample_size, gelman_rubin
from numpyro.infer import MCMC, NUTS

from pymdp.agent import Agent

SEED = 7
plt.style.use("seaborn-v0_8-whitegrid")
print("pymdp:", version("inferactively-pymdp"))
print("jax:", jax.__version__, "numpyro:", numpyro.__version__)

チュートリアルの内容

  1. ランダムドット運動識別: 最初の一歩
  2. 逆転学習: 隠れた文脈の推論
  3. パラメータ推定: 逆転学習課題のデータからパラメータ推定
  4. 変動性のある逆転学習: パラメータの学習
  5. 変動性のある逆転学習: 階層モデル

ランダムドット運動識別課題

多数の点のうち一部だけが左右いずれかへ動く刺激を見て、全体として左向きか右向きかを2肢強制選択する。

課題の流れ

  1. 試行開始
    • 隠れ状態:運動方向は左/右、反応状態は未反応
  2. 感覚証拠を観測し、運動方向の信念を更新
  3. 各反応の正誤を予測し、左/右の行動を選択
    • 状態遷移:反応状態は左/右へ、運動方向は試行内で変化しない

※反応後に正誤フィードバックは提示しない。

エージェントをABCDで表現

パラメータ 内容
尤度行列(\(\textbf{A}\) 運動方向から感覚証拠を生成し、各反応から予測される正誤を表す
遷移行列(\(\textbf{B}\) 運動方向を維持し、行動によって反応状態を左/右へ変える
選好(\(\textbf{C}\) 予測される正答を好み、誤答を避ける
事前信念(\(\textbf{D}\) 運動方向は左右等確率、反応状態は未反応から始める

尤度行列\(A\)の設定

  • 生成モデルには感覚証拠と結果の二つの観測モダリティを置く。感覚証拠は実際に入力し、correct/errorは各方策の結果として内部的に予測する(反応後の正誤フィードバックは入力しない)
観測モデル 条件付き確率 この課題での役割
\(A_{evidence}\) \(P(o_{evidence}\mid s_{motion})\) 運動方向から左右の感覚証拠を生成。方向と一致する証拠の確率は\(0.8\)
\(A_{outcome}\) \(P(o_{outcome}\mid s_{motion},s_{response})\) 各反応から予測されるneutralerrorcorrectを表す

尤度行列\(A\)の設定

A_outcome[outcome(3), motion(2), response(3)]は、方策評価に使う内部的な結果予測

予測される結果(outcome 運動方向(motion 反応状態(response
neutral 左/右 未反応
correct 左反応
error 右反応
error 左反応
correct 右反応

遷移行列\(B\)の設定

B[next_state, current_state, action]は、現在の状態と行動から次の状態が生じる確率を表す。

状態因子 コードでの定義 状態遷移
運動方向 B_motion = eye(2)[motion(2), motion(2), None] 試行の正解方向を左/右のまま維持する。行動では変化しない
反応状態 B_response[response(3), response(3), action(2)] 行動0なら左反応、行動1なら右反応へ移る
  • 未反応状態は\(D\)で初期化され、行動後は左または右の反応状態へ移る

選好\(C\)と事前信念\(D\)の設定

選好\(C\)

  • 感覚証拠:[0, 0]
  • 結果:[0, -2, 2]

結果はneutralerrorcorrectであり、正答を予測する方策が選ばれやすくなる。

\(C\)予測される結果の相対的な望ましさ

開始時の状態信念\(D\)

  • 運動方向:[0.5, 0.5]
  • 反応状態:[1, 0, 0]

試行開始時には左右の運動方向を等確率とし、反応状態はnoneleftrightのうちnoneから始まる。

モデルを定義

\(A\)\(D\)を定義し、Agent()でエージェントを作成する。

A_dependenciesは各観測が依存する状態因子、B_dependenciesは各次状態が依存する現在の状態因子を指定する。

MOTION_LEFT, MOTION_RIGHT = 0, 1
RESPONSE_NONE, RESPONSE_LEFT, RESPONSE_RIGHT = 0, 1, 2
EVIDENCE_LEFT, EVIDENCE_RIGHT = 0, 1
NEUTRAL, ERROR, CORRECT = 0, 1, 2

A_evidence = jnp.array([[0.8, 0.2], [0.2, 0.8]])
A_outcome = np.zeros((3, 2, 3))
A_outcome[NEUTRAL, :, RESPONSE_NONE] = 1.0
for motion in (MOTION_LEFT, MOTION_RIGHT):
    for response in (RESPONSE_LEFT, RESPONSE_RIGHT):
        result = CORRECT if response - 1 == motion else ERROR
        A_outcome[result, motion, response] = 1.0

B_motion = jnp.eye(2)[:, :, None]
B_response = np.zeros((3, 3, 2))
B_response[RESPONSE_LEFT, :, 0] = 1.0
B_response[RESPONSE_RIGHT, :, 1] = 1.0

# A:evidence←motion、outcome←motion+response
A_dependencies = [[0], [0, 1]]
# B:次のmotion←現在のmotion、次のresponse←現在のresponse
B_dependencies = [[0], [1]]

agent = Agent(
    A=[A_evidence, jnp.asarray(A_outcome)],
    B=[B_motion, jnp.asarray(B_response)],
    C=[jnp.zeros(2), jnp.array([0.0, -2.0, 2.0])],
    D=[jnp.array([0.5, 0.5]), jnp.array([1.0, 0.0, 0.0])],
    A_dependencies=A_dependencies, B_dependencies=B_dependencies,
    # 状態因子1(反応状態)だけを行動で制御する
    control_fac_idx=[1],
    policy_len=1, batch_size=1,
    use_states_info_gain=False, action_selection="deterministic",
)

観測から状態を推論し、行動選択

  • 左向きの感覚証拠と、反応前を表すneutralを観測し、状態推論infer_states()→ 方策推論infer_policies()→ 行動選択sample_action()を順に実行する。
observation = [jnp.array([EVIDENCE_LEFT]), jnp.array([NEUTRAL])]
qs = agent.infer_states(observation, empirical_prior=agent.D)
q_pi, neg_efe = agent.infer_policies(qs)
action = agent.sample_action(q_pi)

print("q(motion) =", np.asarray(qs[0][0, -1]))
print("q(response policy) =", np.asarray(q_pi[0]))
print("selected action =", np.asarray(action))

推論結果を可視化

左向きの証拠を観測した後の運動方向の信念\(q(s)\)と、行動の方策確率\(q(\pi)\)を確認する。

posterior_motion = np.asarray(qs[0][0, -1])
response_policy = np.asarray(q_pi[0])

fig, axes = plt.subplots(1, 2, figsize=(8, 3.5))
axes[0].bar(["motion left", "motion right"], posterior_motion)
axes[0].set_title("Posterior belief about motion")
axes[1].bar(["respond left", "respond right"], response_policy)
axes[1].set_title("Response policy probability")
for ax in axes: ax.set_ylim(0, 1)
fig.tight_layout()
plt.show()

最小モデルの挙動

左向きの証拠により左運動の事後信念が高まり、\(A_{outcome}\)\(C\)による正誤の予測を通じて左反応が選ばれる。

逆転学習課題

逆転学習モデル

  • 隠れ状態を直接観測する代わりに、選択行動後の結果から「現在どちらが有利か」を推論する。
種類 因子・モダリティ
隠れ状態 context A優位/B優位
choice 未選択/A/B
観測 outcome 未提示/無報酬/報酬
choice 未選択/A/B

尤度行列\(A\)の設定

このモデルでは、\(A\)を二つの観測モダリティに分ける。

観測モデル 条件付き確率 この課題での役割
\(A_{outcome}\) \(P(o_{outcome}\mid s_{context},s_{choice})\) 文脈と選択から、未提示・無報酬・報酬を予測する
\(A_{choice}\) \(P(o_{choice}\mid s_{choice})\) 現在の選択状態をそのまま観測する単位行列
文脈 A選択時の報酬確率 B選択時の報酬確率
A優位 \(0.8\) \(0.2\)
B優位 \(0.2\) \(0.8\)

遷移行列\(B\)の設定

B[next_state, current_state, action]は、現在の状態と行動から次の状態が生じる確率を表す。

状態因子 コードでの設定 状態遷移
文脈 B_context hazard=0.04とし、前試行の文脈を\(0.96\)で維持(\(0.04\)で反転)
選択状態 B_choice 行動0でA、行動1でBの選択状態へ移る
  • hazardはエージェントの主観的な変化確率
  • 行動で制御できるのはchoiceだけであり、contextは直接変更できない

選好\(C\)と事前信念\(D\)の設定

選好\(C\)

  • 結果:[0, -1, 3]
  • 選択観測:[0, 0, 0]

結果は、未提示・無報酬・報酬の順であり、報酬に正、無報酬に負の値をおく。AとBそのものには選好はない。

開始時の状態信念\(D\)

  • 文脈:[0.5, 0.5]
  • 選択状態:[1, 0, 0]

開始時にはA優位・B優位を等確率とし、まだどちらも選んでいない状態から始める。

逆転学習のAgent

use_states_info_gain=Trueにより、報酬の獲得だけでなく、どちらが有利かという文脈を見分ける情報価値も方策評価に含める。

NULL, LOSS, REWARD = 0, 1, 2
UNDECIDED, ARM_A, ARM_B = 0, 1, 2
hazard = 0.04

A_outcome = np.zeros((3, 2, 3))
A_outcome[NULL, :, UNDECIDED] = 1.0
reward_prob = np.array([[0.8, 0.2], [0.2, 0.8]])
for context in range(2):
    for arm in range(2):
        p = reward_prob[context, arm]
        A_outcome[REWARD, context, arm + 1] = p
        A_outcome[LOSS, context, arm + 1] = 1 - p

B_context = np.array([[1-hazard, hazard],
                      [hazard, 1-hazard]])[:, :, None]
B_choice = np.zeros((3, 3, 2))
B_choice[ARM_A, :, 0], B_choice[ARM_B, :, 1] = 1.0, 1.0

agent = Agent(
    A=[jnp.asarray(A_outcome), jnp.eye(3)],
    B=[jnp.asarray(B_context), jnp.asarray(B_choice)],
    C=[jnp.array([0., -1., 3.]), jnp.zeros(3)],
    D=[jnp.array([.5, .5]), jnp.array([1., 0., 0.])],
    A_dependencies=[[0, 1], [1]], B_dependencies=[[0], [1]],
    control_fac_idx=[1], policy_len=1, batch_size=1,
    use_states_info_gain=True, action_selection="deterministic",
)

80試行のシミュレーション

各試行で方策を評価して選択し、環境から得た結果を使って隠れた文脈の信念を更新する。

def true_reward_probability(trial, arm):
    good_arm = 0 if trial < 40 else 1
    return 0.8 if arm == good_arm else 0.2

rng = np.random.default_rng(SEED)
qs = agent.infer_states([jnp.array([NULL]), jnp.array([UNDECIDED])],
                        empirical_prior=agent.D)
rows = []
for trial in range(80):
    q_pi, neg_efe = agent.infer_policies(qs)
    action = agent.sample_action(q_pi)
    arm = int(action[0, 1])
    prior = agent.update_empirical_prior(action, qs)
    reward = rng.random() < true_reward_probability(trial, arm)
    outcome = REWARD if reward else LOSS
    qs = agent.infer_states(
        [jnp.array([outcome]), jnp.array([arm + 1])], prior
    )
    rows.append((trial + 1, arm, int(reward),
                 true_reward_probability(trial, arm),
                 float(qs[0][0, -1, 0]), float(q_pi[0, 0]),
                 float(neg_efe[0, 0]), float(neg_efe[0, 1])))
reversal_df = pd.DataFrame(rows, columns=[
    "trial", "arm", "reward", "p_reward_true",
    "q_A_good", "q_choose_A", "neg_efe_A", "neg_efe_B"
])

信念・選択・報酬を可視化する

上段に文脈の事後信念、中段に実際の選択、下段は観測報酬と報酬確率をプロットする

true_prob_A = np.array([true_reward_probability(t, 0) for t in range(80)])
fig, axes = plt.subplots(3, 1, figsize=(10, 8), sharex=True)
axes[0].plot(reversal_df["trial"], reversal_df["q_A_good"])
axes[0].axhline(0.5, color="0.5", lw=1)
axes[0].set_ylabel("q(A good)")
axes[1].step(reversal_df["trial"], reversal_df["arm"], where="mid")
axes[1].set_yticks([0, 1], ["A", "B"])
axes[1].invert_yaxis()  # Aを上、Bを下に表示する
axes[1].set_ylabel("choice")
axes[2].scatter(reversal_df["trial"], reversal_df["reward"],
                c=reversal_df["reward"], cmap="RdYlGn", vmin=0, vmax=1,
                label="observed reward")
axes[2].step(reversal_df["trial"], true_prob_A, where="mid",
             color="black", ls="--", lw=2,
             label="environment P(reward | A)")
axes[2].set(xlabel="trial", ylabel="reward / probability")
axes[2].legend()
fig.tight_layout()
plt.show()

逆転後に信念と選択が切り替わる

パラメータ推定

  • \(\gamma\)は、方策間の期待自由エネルギーの差が選択にどれだけ強く反映されるかを表すパラメータ。
  • 以下の式でAの選択確率を計算(\(\gamma\)が大きいほど\(G\)の小さい選択が一貫して選ばれ、小さいとランダムに選ばれる)

\[ P(A)=\frac{\exp(-\gamma G_A)}{\exp(-\gamma G_A)+\exp(-\gamma G_B)} \]

  • 2択では\(P(A)=\operatorname{sigmoid}\{\gamma(G_B-G_A)\}\)と等価である。

異なる \(\gamma\) のデータを生成する

  • 先程のシミュレーションで得られた負の期待自由エネルギー(reversal_df内のneg_efe_Aneg_efe_B)をもとに、 \(\gamma\) だけ変えた行動を生成する(ここでは5名分のTRUE_GAMMASを設定)
TRUE_GAMMAS = np.array([0.35, 0.60, 1.00, 1.80, 3.00])
delta_neg_efe_A = (reversal_df["neg_efe_A"]
                   - reversal_df["neg_efe_B"]).to_numpy()
rng = np.random.default_rng(20260906)

choices_A = []
for gamma in TRUE_GAMMAS:
    p_choose_A = jax.nn.sigmoid(
        gamma * jnp.asarray(delta_neg_efe_A)
    )
    choices_A.append(rng.binomial(1, np.asarray(p_choose_A)))
choices_A = np.asarray(choices_A)

5名分の可視化

  • 全員で共通した状態信念\(q\)(点線)、 \(\gamma\) で異なる選択確率(実線)と選択(点)をプロットする
q_A_before = np.r_[0.5, reversal_df["q_A_good"].to_numpy()[:-1]]
colors = plt.cm.viridis(np.linspace(.12, .88, len(TRUE_GAMMAS)))
fig, axes = plt.subplots(5, 1, figsize=(11, 10), sharex=True, sharey=True)
for person, (gamma, color, ax) in enumerate(
    zip(TRUE_GAMMAS, colors, axes), start=1
):
    p_choose_A = np.asarray(jax.nn.sigmoid(
        gamma * jnp.asarray(delta_neg_efe_A)
    ))
    chosen_A = choices_A[person - 1]
    ax.plot(np.arange(1, 81), q_A_before, "--", color=".35",
            label="shared q(A good)")
    ax.plot(np.arange(1, 81), p_choose_A, color=color,
            label="P(choice=A)")
    ax.scatter(np.arange(1, 81), chosen_A, color=color, s=10, alpha=.45)
    ax.set(ylabel=f"P{person}", ylim=(-.08, 1.08))
    ax.set_title(f"true gamma = {gamma:.2f}", loc="left", fontsize=10)
axes[0].legend(ncol=2, fontsize=8)
axes[-1].set_xlabel("trial")
fig.tight_layout()
plt.show()

5名の選択プロファイル

  • 期待自由エネルギーは共通でも、\(\gamma\)が大きくなるほど(下にいくほど)、期待自由エネルギーの低い方を確実に選択するようになる

NumPyroで\(\gamma\)を推定する

  • 確率的プログラミングライブラリのNumPyroでMCMC法でパラメータ推定する。
  • \(\gamma\)は非負なのでLogNormal事前分布を置き、4本のNUTSチェーンを実行する(推定するのは\(\gamma\)のみ)
def gamma_choice_model(delta_neg_efe_A, choices_A):
    n = choices_A.shape[0]
    gamma = numpyro.sample(
        "gamma",
        dist.LogNormal(jnp.log(1.0), 1.0).expand([n]).to_event(1),
    )
    logits = gamma[:, None] * delta_neg_efe_A[None, :]
    numpyro.sample(
        "choice_A", dist.Bernoulli(logits=logits).to_event(2),
        obs=choices_A,
    )

mcmc = MCMC(
    NUTS(gamma_choice_model, target_accept_prob=0.90),
    num_warmup=400, num_samples=800, num_chains=4,
    chain_method="sequential",
)
mcmc.run(jr.PRNGKey(20260906), jnp.asarray(delta_neg_efe_A),
         jnp.asarray(choices_A))
gamma_chains = mcmc.get_samples(group_by_chain=True)["gamma"]
gamma_samples = np.asarray(gamma_chains).reshape(-1, len(TRUE_GAMMAS))
gamma_rhat = np.asarray(gelman_rubin(gamma_chains))
gamma_ess = np.asarray(effective_sample_size(gamma_chains))
divergences = int(np.asarray(
    mcmc.get_extra_fields(group_by_chain=True)["diverging"]
).sum())

MCMCの収束診断

  • トレースプロットに真値を点線で重ねる。\(\widehat R\)、有効サンプルサイズ、発散数も数値で確認する。
fig, axes = plt.subplots(5, 1, figsize=(11, 9), sharex=True)
for person, ax in enumerate(axes):
    for chain in range(gamma_chains.shape[0]):
        ax.plot(gamma_chains[chain, :, person], lw=.55, alpha=.75,
                label=f"chain {chain+1}" if person == 0 else None)
    ax.axhline(TRUE_GAMMAS[person], color="black", ls="--", lw=1)
    ax.set_ylabel(f"P{person+1}\ngamma")
    ax.set_title(f"R-hat={gamma_rhat[person]:.3f}, "
                 f"ESS={gamma_ess[person]:.0f}", loc="right", fontsize=9)
axes[0].legend(ncol=4, fontsize=8)
axes[-1].set_xlabel("post-warmup draw")
fig.tight_layout()
plt.show()
print("divergences:", divergences)

収束を確認する

  • チェーンの重なりと定常性に加え、\(\widehat R\)が1に近いか、有効サンプルサイズが十分か、発散がないかを確認する。

\(\gamma\)の簡易パラメータリカバリー

  • 事後平均と90%信用区間を求め、真値と近い値を推定できているかを確認する。相関だけでは尺度のずれを検出できないため、RMSEと区間被覆率も計算する。
gamma_mean = gamma_samples.mean(axis=0)
gamma_low = np.quantile(gamma_samples, .05, axis=0)
gamma_high = np.quantile(gamma_samples, .95, axis=0)
correlation = np.corrcoef(TRUE_GAMMAS, gamma_mean)[0, 1]
rmse = np.sqrt(np.mean((TRUE_GAMMAS - gamma_mean)**2))
coverage = np.mean((TRUE_GAMMAS >= gamma_low) &
                   (TRUE_GAMMAS <= gamma_high))

fig, ax = plt.subplots(figsize=(6, 5))
yerr = np.vstack([gamma_mean-gamma_low, gamma_high-gamma_mean])
ax.errorbar(TRUE_GAMMAS, gamma_mean, yerr=yerr, fmt="o", capsize=4)
limit = max(TRUE_GAMMAS.max(), gamma_high.max()) * 1.08
ax.plot([0, limit], [0, limit], "--", color="black")
ax.set(xlim=(0, limit), ylim=(0, limit), xlabel="true gamma",
       ylabel="estimated gamma")
ax.set_aspect("equal", adjustable="box")
plt.show()
print({"correlation": correlation, "RMSE": rmse,
       "coverage_90": coverage, "max_Rhat": gamma_rhat.max(),
       "min_ESS": gamma_ess.min(), "divergences": divergences})

\(\gamma\)の簡易パラメータリカバリー

  • 散布図で確認(相関、RMSEなども)

パラメータリカバリーの補足

  • 1回の逆転学習シミュレーションから得た\(\Delta_t=(-G_{A,t})-(-G_{B,t})\)を、5名に共通する既知の説明変数として固定している。
  • そのうえで\(P(a_t=A)=\operatorname{sigmoid}(\gamma\Delta_t)\)から選択を生成し、同じ選択モデルで\(\gamma\)を推定している。
  • ここで確認するのは固定した方策価値系列のもとでの\(\gamma\)の条件付き回復であり、簡易的な検討である。

参加者ごとの選択・結果から状態信念と期待自由エネルギーを再構成する検討ではない。

変動性のある逆転学習課題の流れ

\(A\)を学習するモデルに

  • 環境は、40試行ごとの安定した反転から10試行ごとの頻繁な反転へ移り、最後は再び40試行安定する。
  • これまで逆転学習課題は文脈についての信念の更新で解いたが、今回は\(A\)の学習によって解く。
  • 隠れ状態はchoiceのみにし(これは行動によって制御されほぼ直接観測)、状態に対する結果の確率を表す\(A\)を学習する。

\(A\)の学習

  • 結果を\(\mathbf o_t\)、状態因子choiceの事後信念を\(\mathbf s_t=q(s_t)\)と表す。結果を観測するたびに、両者の外積を擬似カウントへ加える。

\[ \mathrm{pA}_{t+1} =\mathrm{pA}_t+\eta\,\mathbf o_t\otimes\mathbf s_t, \qquad \eta=0.5 \]

  • 以下の式で更新した擬似カウントから\(A\)を計算する。

\[ A_{t+1}(o\mid s) =\frac{\mathrm{pA}_{t+1}(o,s)}{\sum_{o'}\mathrm{pA}_{t+1}(o',s)} \]

pymdpで\(A\)を学習する手順

  1. pAと、その期待値としての初期\(A\)を設定する
  2. Agent()learn_A=Trueとし、\(A\)の学習を有効にする
  3. 観測後にagent = agent.infer_parameters()で、更新されたAgentを受け取る
  • 元のAgentを直接変更せず、更新されたpA\(A\)を持つAgentを返す。そのため、agent = ...という再代入が必要になる

\(pA_{outcome}\)の行と列

pA_outcome[outcome, choice]は、結果3カテゴリと選択状態3カテゴリの組み合わせに対する擬似カウントである。

選択状態 未提示 無報酬 報酬 事前信念
未選択 20.0 0.1 0.1 反応前には結果が提示されない
A 0.1 1.0 1.0 無報酬と報酬について弱く対称
B 0.1 1.0 1.0 無報酬と報酬について弱く対称
  • pA_choice[observed_choice, choice]は自分の選択状態についての観測モデルであり、ここでは既知に近い強い事前を置く。
  • これらは実際の観測回数ではなく、学習開始前の信念を表す擬似カウントである。

\(\gamma\)\(\alpha\)の役割を分ける

\[ q(\pi)=\operatorname{softmax}\{\ln E-\gamma G(\pi)\}, \qquad P(\mathrm{select}\ \pi)\propto q(\pi)^\alpha \]

  • \(\gamma\):期待自由エネルギー\(G\)を方策確率\(q(\pi)\)へ反映する強さ
  • \(\alpha\)\(q(\pi)\)から実際の行動を確率的に抽出する際の追加的な精度
  • 両方を大きくすると選択が二重に確定的になるため、ここでは\(\gamma=4\)\(\alpha=1\)とする。これにより、\(q(\pi)\)を追加で尖らせずに確率的に行動を選ぶ。

pymdpで\(pA\)を設定

  • pA_outcomepA_choiceを初期\(A\)へ正規化し、擬似カウントとともにAgentへ渡す。
def normalize_columns(x):
    return x / x.sum(axis=0, keepdims=True)

def build_learning_agent():
    # outcome | choice:A/Bについては弱い対称事前
    pA_outcome = np.ones((3, 3))
    pA_outcome[:, UNDECIDED] = [20.0, 0.1, 0.1]
    pA_outcome[:, ARM_A:] = [[0.1, 0.1],
                                [1.0, 1.0],
                                [1.0, 1.0]]
    # choice観測は既知に近い強い事前
    pA_choice = np.full((3, 3), 0.1)
    np.fill_diagonal(pA_choice, 20.0)
    B_choice = np.zeros((3, 3, 2))
    B_choice[ARM_A, :, 0] = 1.0
    B_choice[ARM_B, :, 1] = 1.0

    agent = Agent(
        A=[jnp.asarray(normalize_columns(pA_outcome)),
           jnp.asarray(normalize_columns(pA_choice))],
        B=[jnp.asarray(B_choice)],
        C=[jnp.array([0., -1., 3.]), jnp.zeros(3)],
        D=[jnp.array([1., 0., 0.])],
        pA=[jnp.asarray(pA_outcome), jnp.asarray(pA_choice)],
        A_dependencies=[[0], [0]], B_dependencies=[[0]],
        control_fac_idx=[0], policy_len=1, batch_size=1,
        learn_A=True, use_param_info_gain=True,
        # alpha=1なら、q(pi)を追加で尖らせずに確率的に行動を選ぶ
        action_selection="stochastic", gamma=4.0, alpha=1.0,
    )
    base_pA = [x.copy() for x in agent.pA]
    return agent, base_pA

learning_agent, base_pA = build_learning_agent()
print("initial P(reward | A/B):",
      np.asarray(learning_agent.A[0][0, REWARD, ARM_A:]))

古い擬似カウントの減衰

  • \(pA_0\)を初期擬似カウントとし、新しい結果を加える前に、それ以降に蓄積した経験を保持率\(\rho\)で減衰させるモデルも考える。

\[ pA_t^{-}=pA_0+\rho\left(pA_t-pA_0\right) \]

  • \(\rho=1\)なら経験を累積し続け、\(\rho<1\)なら古い経験の影響が毎試行小さくなる。

古い擬似カウントの減衰

  • 古い擬似カウントの減衰は、pymdpが自動で行う処理ではないため、更新前に明示的に実装する。
def apply_dirichlet_forgetting(agent, base_pA, retention):
    new_pA = [base + retention * (current - base)
              for base, current in zip(base_pA, agent.pA)]
    new_A = [p / p.sum(axis=1, keepdims=True) for p in new_pA]
    return eqx.tree_at(lambda x: (x.pA, x.A), agent, (new_pA, new_A))

def good_arm(trial):
    if trial < 40: return 0
    if trial < 80: return 1
    if trial < 120: return ((trial - 80) // 10) % 2
    return 0

def reward_probability(trial, arm):
    return 0.8 if arm == good_arm(trial) else 0.2

160試行を通してAを学習する

  • 各試行では、①行動選択、②環境から結果を得る、③状態推論、④古い擬似カウントの減衰、⑤新しい経験の追加、の順に処理する。
  • 古い擬似カウントの減衰をしない場合(累積)とする場合(減衰)で実施
def run_volatile_learning(retention, seed=22, n_trials=160):
    agent, base_pA = build_learning_agent()
    rng, key = np.random.default_rng(seed), jr.PRNGKey(seed)
    qs = agent.infer_states(
        [jnp.array([NULL]), jnp.array([UNDECIDED])], agent.D
    )
    rows = []
    for trial in range(n_trials):
        q_pi, _ = agent.infer_policies(qs)
        key, sample_key = jr.split(key)
        action_keys = jr.split(sample_key, agent.batch_size + 1)[1:]
        action = agent.sample_action(q_pi, rng_key=action_keys)
        arm = int(action[0, 0])
        predicted_reward = float(agent.A[0][0, REWARD, arm + 1])
        prior = agent.update_empirical_prior(action, qs)

        p_reward = reward_probability(trial, arm)
        reward = int(rng.random() < p_reward)
        outcome = REWARD if reward else LOSS
        obs = [jnp.array([outcome]), jnp.array([arm + 1])]
        qs = agent.infer_states(obs, empirical_prior=prior)

        if retention < 1.0:
            agent = apply_dirichlet_forgetting(agent, base_pA, retention)
        agent = agent.infer_parameters(
            beliefs_A=qs,
            observations=[jnp.array([[outcome]]),
                          jnp.array([[arm + 1]])],
            actions=None, lr_pA=0.5,
        )
        rows.append({
            "trial": trial + 1, "arm": arm,
            "good_arm": good_arm(trial), "reward": reward,
            "correct_choice": int(arm == good_arm(trial)),
            "predicted_reward": predicted_reward,
            "est_reward_A": float(agent.A[0][0, REWARD, ARM_A]),
            "est_reward_B": float(agent.A[0][0, REWARD, ARM_B]),
        })
    return pd.DataFrame(rows), agent

cumulative_df, _ = run_volatile_learning(retention=1.00)
leaky_df, _ = run_volatile_learning(retention=0.90)

累積・忘却学習を可視化

上段はエージェントが学習したA・Bの報酬確率、下段はAの選択率の8試行移動平均である。下段には環境で設定した\(P(\mathrm{reward}\mid A)\)も示す。乱数seedと報酬系列の生成方法を揃え、保持率\(\rho\)だけを変える。

true_prob_A_160 = np.array([reward_probability(t, 0) for t in range(160)])
fig, axes = plt.subplots(2, 2, figsize=(13, 7), sharex=True, sharey="row")
for col, (label, df) in enumerate([
    ("Cumulative counts", cumulative_df),
    ("Forgetting: rho=0.90", leaky_df),
]):
    axes[0, col].plot(df["trial"], df["est_reward_A"], label="estimated A")
    axes[0, col].plot(df["trial"], df["est_reward_B"], label="estimated B")
    axes[0, col].set(title=label, ylabel="P(reward)")
    choice_A = (df["arm"] == 0).rolling(8, min_periods=1).mean()
    axes[1, col].plot(df["trial"], choice_A, color="C3",
                      label="choice A (rolling 8)")
    axes[1, col].step(df["trial"], true_prob_A_160, where="mid",
                      color="black", ls="--", lw=1.8,
                      label="environment P(reward | A)")
    axes[1, col].set(xlabel="trial", ylabel="P(choice = A)",
                     ylim=(-0.05, 1.05))
    for ax in axes[:, col]: ax.axvspan(80.5, 120.5, color="C4", alpha=.06)
axes[0, 0].legend(ncol=2, fontsize=8)
axes[1, 0].legend(ncol=2, fontsize=8)
fig.tight_layout()
plt.show()

累積学習と忘却学習を比較

  • 忘却で直近の変化に敏感に(左:累積、右:忘却)

階層モデル化:上位層の追加

  • \(A\)を学習するモデルでは保持率\(\rho\)は固定だったが、変動性に合わせて調整できると良い。

→下位層では選択と結果を学び、上位層では環境が安定/変動的のどちらかを推論する階層モデルへ

  • 下位の予測から外れた結果を上位層の「驚き」観測へ変換し、上位層の\(q_t(\mathrm{volatile})\)を下位層の保持率\(\rho\)へ戻す

→環境が変動していると思うほど忘却を強め、最近の結果を重視する

1試行の流れ

  1. 下位層が選択肢A/Bを選択する
  2. 結果を観測し、更新前\(A_{low}\)から予測外かどうかを判定する(\(0.5\)未満かどうかで二値化)
  3. 判定した「通常/驚き」を上位層の観測として入力する
  4. 上位層が\(q_t(\mathrm{volatile})\)を推論する
  5. 上位信念から下位層の保持率を計算する
  6. \(\rho_t\)で下位層の\(pA\)を減衰させ、新しい結果の擬似カウントを加える

上位層の\(A\)\(B\)を定義する

  • \(A_{high}\)は安定/変動状態から通常/驚き観測を予測
  • \(B_{high}\)は高次状態のゆっくりした時間変化を表す
ORDINARY, SURPRISING = 0, 1
STABLE, VOLATILE = 0, 1

def build_meta_agent():
    # P(surprising | stable)=.12, P(surprising | volatile)=.55
    A_high = [jnp.array([[0.88, 0.45],
                         [0.12, 0.55]])]
    # 上位状態は高い確率で持続する
    B_high = [jnp.array([[[0.995], [0.020]],
                         [[0.005], [0.980]]])]
    return Agent(
        A=A_high, B=B_high, C=[jnp.zeros(2)],
        D=[jnp.array([0.9, 0.1])],
        policy_len=1, batch_size=1,
        action_selection="deterministic",
    )

higher = build_meta_agent()
print("A_high:", np.asarray(higher.A[0][0]), sep="\n")
print("B_high:", np.asarray(higher.B[0][0, :, :, 0]), sep="\n")

上位信念が下位層の記憶を調整する

  • 上位層は下位層の\(pA\)保持率を調整するが、下位層の\(A\)や選択を直接決めるわけではない

\[ \rho_t=0.995-0.15q_t(\mathrm{volatile}) \]

  • 上位層が\(q_t(\mathrm{volatile})\)を高く見積もれば、忘却を強め、最近の結果を重視する。

上位・下位Agentを接続

  • 2つのAgentを明示的に接続する。pymdpが階層間の結合を自動生成するわけではない。
def run_hierarchical_learning(seed=22, n_trials=160):
    lower, base_pA = build_learning_agent()
    higher = build_meta_agent()
    rng, key = np.random.default_rng(seed), jr.PRNGKey(seed)
    q_lower = lower.infer_states(
        [jnp.array([NULL]), jnp.array([UNDECIDED])], lower.D
    )
    prior_high, rows = higher.D, []

    for trial in range(n_trials):
        q_pi, _ = lower.infer_policies(q_lower)
        key, sample_key = jr.split(key)
        action_keys = jr.split(sample_key, lower.batch_size + 1)[1:]
        action = lower.sample_action(q_pi, rng_key=action_keys)
        arm = int(action[0, 0])
        prior_lower = lower.update_empirical_prior(action, q_lower)

        reward = int(rng.random() < reward_probability(trial, arm))
        outcome = REWARD if reward else LOSS
        predicted_p = float(lower.A[0][0, outcome, arm + 1])
        high_obs = SURPRISING if predicted_p < 0.50 else ORDINARY
        q_high = higher.infer_states(
            [jnp.array([high_obs])], empirical_prior=prior_high
        )
        q_volatile = float(q_high[0][0, -1, VOLATILE])
        prior_high = higher.update_empirical_prior(jnp.array([[0]]), q_high)

        retention = 0.995 - 0.15 * q_volatile
        lower = apply_dirichlet_forgetting(lower, base_pA, retention)
        obs = [jnp.array([outcome]), jnp.array([arm + 1])]
        q_lower = lower.infer_states(obs, empirical_prior=prior_lower)
        lower = lower.infer_parameters(
            beliefs_A=q_lower,
            observations=[jnp.array([[outcome]]),
                          jnp.array([[arm + 1]])],
            actions=None, lr_pA=0.5,
        )
        rows.append({
            "trial": trial + 1, "arm": arm, "reward": reward,
            "correct_choice": int(arm == good_arm(trial)),
            "q_volatile": q_volatile, "retention": retention,
            "est_reward_A": float(lower.A[0][0, REWARD, ARM_A]),
            "est_reward_B": float(lower.A[0][0, REWARD, ARM_B]),
        })
    return pd.DataFrame(rows), lower, higher

hierarchical_df, lower_final, higher_final = run_hierarchical_learning()

結果を可視化

  • 上段に推定した変動性、中段に保持率、下段に学習した報酬確率と選択を示す。高変動区間は薄紫で示す。
fig, axes = plt.subplots(3, 1, figsize=(11, 8), sharex=True)
task_volatility = np.where(
    (hierarchical_df["trial"] > 80) &
    (hierarchical_df["trial"] <= 120), 1, 0
)
axes[0].plot(hierarchical_df["trial"], hierarchical_df["q_volatile"],
             label="inferred q(volatile)")
axes[0].step(hierarchical_df["trial"], task_volatility, where="mid",
             color="black", ls=":", label="task volatility")
axes[0].set(ylabel="volatility", ylim=(-.03, 1.03))
axes[0].legend()
axes[1].plot(hierarchical_df["trial"], hierarchical_df["retention"])
axes[1].set_ylabel("retention rho")
axes[2].plot(hierarchical_df["trial"], hierarchical_df["est_reward_A"],
             label="estimated A")
axes[2].plot(hierarchical_df["trial"], hierarchical_df["est_reward_B"],
             label="estimated B")
axes[2].step(hierarchical_df["trial"], hierarchical_df["arm"],
             where="mid", color=".3", alpha=.35, label="choice")
axes[2].step(hierarchical_df["trial"], true_prob_A_160, where="mid",
             color="black", ls="--", lw=1.8,
             label="environment P(reward | A)")
axes[2].set(xlabel="trial", ylabel="probability / choice")
axes[2].legend()
for ax in axes: ax.axvspan(80.5, 120.5, color="C4", alpha=.06)
fig.tight_layout()
plt.show()

上位層が保持率を調整する

  • 高変動区間では\(q(volatile)\)が高まり、保持率\(\rho\)が低下

全体のまとめ

  • 計算論的精神医学:行動を生む潜在的な計算過程を生成モデルとして明示し、シミュレーションとデータへの適用を通して検討
  • 能動的推論:離散POMDPの\(\textbf{ABCD}\)で世界と選好を表し、\(F\)に基づく状態推論と\(G\)に基づく方策評価を結びつける
  • pymdpによる実装:ランダムドット課題と逆転学習課題を通して、状態推論・方策評価・行動選択をAgentで実装できる
  • パラメータ推定:NumPyroなどを組み合わせて、選択データからパラメータ推定できる
  • 学習と階層化\(pA\)の更新に明示的な忘却を加えると変動環境に追従でき、2つのAgentを接続すると上位の変動性信念から下位の保持率を調整できる

引用文献

Anticevic, Alan, and John D Murray. 2017. Computational Psychiatry: Mathematical Modeling of Mental Illness. Academic Press.
Bogacz, Rafal. 2017. “A Tutorial on the Free-Energy Framework for Modelling Perception and Learning.” Journal of Mathematical Psychology 76 (Pt B): 198–211.
Buckley, Christopher L, Chang Sub Kim, Simon McGregor, and Anil K Seth. 2017. “The Free Energy Principle for Action and Perception: A Mathematical Review.” Journal of Mathematical Psychology 81 (December): 55–79.
Friston, Karl. 2010. “The Free-Energy Principle: A Unified Brain Theory?” Nature Reviews. Neuroscience 11 (2): 127–38.
Huys, Quentin J M, Michael Moutoussis, and Jonathan Williams. 2011. “Are Computational Models of Any Use to Psychiatry?” Neural Networks: The Official Journal of the International Neural Network Society 24 (6): 544–51.
Montague, P Read, Raymond J Dolan, Karl J Friston, and Peter Dayan. 2012. “Computational Psychiatry.” Trends in Cognitive Sciences 16 (1): 72–80.
Sandved-Smith, Lars, Casper Hesp, Jérémie Mattout, Karl Friston, Antoine Lutz, and Maxwell J D Ramstead. 2021. “Towards a Computational Phenomenology of Mental Action: Modelling Meta-Awareness and Attentional Control with Deep Parametric Active Inference.” Neuroscience of Consciousness 2021 (1): niab018. https://doi.org/10.1093/nc/niab018.
Smith, Ryan, Karl J Friston, and Christopher J Whyte. 2022. “A Step-by-Step Tutorial on Active Inference and Its Application to Empirical Data.” Journal of Mathematical Psychology 107 (April): 102632.
国里愛彦. 2013. “うつとストレスに対する計算論的アプローチ : 計算論的精神医学入門.” ストレス科学 = The Japanese Journal of Stress Sciences : 日本ストレス学会誌 28 (2): 101–7.
国里愛彦, 片平健太郎, 沖村宰, and 山下祐一. 2019. 計算論的精神医学:情報処理過程から読み解く精神障害. 勁草書房.
宗田卓史, 遠山朝子, 国里愛彦, 沖村宰, 片平健太郎, and 山下祐一. 2025. R/Pythonではじめる計算論的精神医学. 京都: 金芳堂.
山下祐一, 松岡洋夫, and 谷淳. 2013. “計算論的精神医学の可能性: 適応行動の代償としての統合失調症.” 精神医学 55 (9): 885–95.