flowchart TB
A["Aへの現在の信念"] --> I["期待される<br/>情報利得"]
I --> C["選択"]
C --> O["結果を観測"]
O --> U["Aの擬似カウントを更新"]
U --> A
2026-09-06
私達や動物の行動の背後にある計算過程を数式で表現したモデルを用いること。
ここでの計算は意識的な計算に限定されず、脳内のあらゆる情報処理においてなされる無意識的な計算も含む。
計算論的神経科学や計算論的精神医学は、計算論的アプローチをとる。
\[ p(s|o) = \frac{p(o|s)p(s)}{p(o)} \]
\[ F = \int q(s)\ln \frac{q(s)}{p(s,o)}ds \]
\[ F = D_{KL}[q(s)\,\Vert\,p(s|o)]- \ln p(o) \]
→ \(F\)はサプライザルの上界であり、\(F\)を最小にする \(q\) を探索することで、近似事後分布を真の事後分布へ近づける。
→知覚・行動・学習が自由エネルギー最小化として統一的に記述できる
→\(F\)ではなく期待自由エネルギー\(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} \]
\[ 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] \]
POMDPは、外界の状態を直接観測できない状況で、得られた観測から状態を推論しながら行動する枠組み。
→POMDPは、ベイズ的な状態推定と逐次的な意思決定を統合した枠組み
| 記号 | 確率分布 | 役割 |
|---|---|---|
| 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\)は、エージェントの信念が対象とする内部的な時刻を表す。
\[ q(\pi)=\operatorname{softmax}\left(\ln E-\gamma G(\pi)\right) \]
| 記号 | 説明 |
|---|---|
| F | 状態事後分布を求めるときに最小化する変分自由エネルギー |
| G | 将来の選好充足と情報獲得を方策ごとに評価する期待自由エネルギー |
| \(\gamma\) | \(G\)の差を方策確率に反映する強さ。大きいほど低い\(G\)の方策に集中する |
| E | 期待自由エネルギーとは別に方策の選びやすさを定める事前分布\(p(\pi)\) |
\[ 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\) | 小さい |
\[ \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\) | 連続する状態と行動の同時生起 |
flowchart TB
A["Aへの現在の信念"] --> I["期待される<br/>情報利得"]
I --> C["選択"]
C --> O["結果を観測"]
O --> U["Aの擬似カウントを更新"]
U --> A
pymdpは、離散状態空間の能動的推論モデルをPythonで構築するオープンソースパッケージ
| 特徴 | できること |
|---|---|
| 生成モデル | 複数の状態因子・観測モダリティを、A・B・C・Dなどで定義する |
| 推論と行動 | 状態推論、方策評価、期待自由エネルギーに基づく行動選択を行う |
| パラメータ学習 | エージェントが経験に応じて、観測モデルAや遷移モデルBなどを更新する |
| パラメータ推定 | 行動データにモデルを当てはめ、\(\gamma\)などを推定する(NumPyro・pybefitと連携) |
| シミュレーション | Envは外部環境を表し、rollout()は複数時点の観測・推論・行動・状態遷移を反復する |
v1.0.0以降はJAXベースになった
jitで反復処理をコンパイルし、CPU・GPU・TPUで実行batch_sizeとvmapで複数エージェントを並列化jax.numpy配列と先頭のバッチ次元を使用jax.random.PRNGKey で管理| パッケージ | バージョン |
|---|---|
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__)多数の点のうち一部だけが左右いずれかへ動く刺激を見て、全体として左向きか右向きかを2肢強制選択する。
※反応後に正誤フィードバックは提示しない。
| パラメータ | 内容 |
|---|---|
| 尤度行列(\(\textbf{A}\)) | 運動方向から感覚証拠を生成し、各反応から予測される正誤を表す |
| 遷移行列(\(\textbf{B}\)) | 運動方向を維持し、行動によって反応状態を左/右へ変える |
| 選好(\(\textbf{C}\)) | 予測される正答を好み、誤答を避ける |
| 事前信念(\(\textbf{D}\)) | 運動方向は左右等確率、反応状態は未反応から始める |
| 観測モデル | 条件付き確率 | この課題での役割 |
|---|---|---|
| \(A_{evidence}\) | \(P(o_{evidence}\mid s_{motion})\) | 運動方向から左右の感覚証拠を生成。方向と一致する証拠の確率は\(0.8\) |
| \(A_{outcome}\) | \(P(o_{outcome}\mid s_{motion},s_{response})\) | 各反応から予測されるneutral・error・correctを表す |
A_outcome[outcome(3), motion(2), response(3)]は、方策評価に使う内部的な結果予測
予測される結果(outcome) |
運動方向(motion) |
反応状態(response) |
|---|---|---|
neutral |
左/右 | 未反応 |
correct |
左 | 左反応 |
error |
左 | 右反応 |
error |
右 | 左反応 |
correct |
右 | 右反応 |
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なら右反応へ移る |
[0, 0][0, -2, 2]結果はneutral・error・correctであり、正答を予測する方策が選ばれやすくなる。
\(C\)は予測される結果の相対的な望ましさ。
[0.5, 0.5][1, 0, 0]試行開始時には左右の運動方向を等確率とし、反応状態はnone・left・rightのうち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_{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[next_state, current_state, action]は、現在の状態と行動から次の状態が生じる確率を表す。
| 状態因子 | コードでの設定 | 状態遷移 |
|---|---|---|
| 文脈 | B_context |
hazard=0.04とし、前試行の文脈を\(0.96\)で維持(\(0.04\)で反転) |
| 選択状態 | B_choice |
行動0でA、行動1でBの選択状態へ移る |
hazardはエージェントの主観的な変化確率choiceだけであり、contextは直接変更できない[0, -1, 3][0, 0, 0]結果は、未提示・無報酬・報酬の順であり、報酬に正、無報酬に負の値をおく。AとBそのものには選好はない。
[0.5, 0.5][1, 0, 0]開始時にはA優位・B優位を等確率とし、まだどちらも選んでいない状態から始める。
Agentuse_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",
)各試行で方策を評価して選択し、環境から得た結果を使って隠れた文脈の信念を更新する。
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()\[ P(A)=\frac{\exp(-\gamma G_A)}{\exp(-\gamma G_A)+\exp(-\gamma G_B)} \]
reversal_df内のneg_efe_Aとneg_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)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()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())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)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})参加者ごとの選択・結果から状態信念と期待自由エネルギーを再構成する検討ではない。
choiceのみにし(これは行動によって制御されほぼ直接観測)、状態に対する結果の確率を表す\(A\)を学習する。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_{t+1}(o\mid s) =\frac{\mathrm{pA}_{t+1}(o,s)}{\sum_{o'}\mathrm{pA}_{t+1}(o',s)} \]
pAと、その期待値としての初期\(A\)を設定するAgent()でlearn_A=Trueとし、\(A\)の学習を有効にするagent = agent.infer_parameters()で、更新されたAgentを受け取るpAと\(A\)を持つAgentを返す。そのため、agent = ...という再代入が必要になる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]は自分の選択状態についての観測モデルであり、ここでは既知に近い強い事前を置く。\[ q(\pi)=\operatorname{softmax}\{\ln E-\gamma G(\pi)\}, \qquad P(\mathrm{select}\ \pi)\propto q(\pi)^\alpha \]
pA_outcomeとpA_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_t^{-}=pA_0+\rho\left(pA_t-pA_0\right) \]
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.2def 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()→下位層では選択と結果を学び、上位層では環境が安定/変動的のどちらかを推論する階層モデルへ
→環境が変動していると思うほど忘却を強め、最近の結果を重視する
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")\[ \rho_t=0.995-0.15q_t(\mathrm{volatile}) \]
Agentを接続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()Agentで実装できるAgentを接続すると上位の変動性信念から下位の保持率を調整できる