! python --version
import sys
!{sys.executable} -m pip list | grep inferactively-pymdpPython 3.10.12
inferactively-pymdp 1.0.4
[notice] A new release of pip is available: 23.0.1 -> 26.2.1
[notice] To update, run: pip install --upgrade pip
PyMDP agents learns to avoid a randomly placed obstacle
Kobus Esterhuysen
September 29, 2026
September 30, 2026
In this problem an active inference agent needs to learn to avoid a randomly placed obstacle by either moving left or right in order to reach one of the end locations that each has a reward. It moves through a 1D grid and the agent does not know the location of the obstacle. The obstacle is placed either on the left or the right in a probabilistic way at the start of an epoch. The obstacle is more likely to be placed on the left, than on the right.
More detail:
'LEFT_REWARD', 'LEFT', 'CENTER', 'RIGHT', 'RIGHT_REWARD'CENTER) at the beginning of an epoch. The agent has an accurate belief about this starting location.LEFT_REWARD and RIGHT_REWARD.LEFT or location RIGHTCOLLISION is observed.Location observation modality, o_0Reward observation modality, o_1Move observation modality, o_2o_1 == 'REWARD'LEFT_REWARD and RIGHT_REWARD end locations are different from the environment’s reality.A few coding conventions:
Python 3.10.12
inferactively-pymdp 1.0.4
[notice] A new release of pip is available: 23.0.1 -> 26.2.1
[notice] To update, run: pip install --upgrade pip
def plot_likelihood(
matrix,
title_str="Likelihood distribution (A)",
yticklabels=None,
xticklabels=None
):
"""
Plots a 2-D likelihood matrix as a heatmap
"""
matrix = np.asarray(matrix)
if not np.isclose(matrix.sum(axis=0), 1.0).all():
raise ValueError("Distribution not column-normalized! Please normalize (ensure matrix.sum(axis=0) == 1.0 for all columns)")
fig = plt.figure(figsize = (5,5))
ax = sns.heatmap(
matrix,
yticklabels = yticklabels,
xticklabels = xticklabels,
## cmap = 'gray',
cmap = "OrRd",
linewidths=1,
cbar = True,
square=True,
vmin = 0.0,
vmax = 1.0)
plt.title(title_str)
plt.show()
def plot_beliefs(belief_dist, title_str=""):
"""
Plot a categorical distribution or belief distribution, stored in the 1-D numpy vector `belief_dist`
"""
belief_dist = np.asarray(belief_dist)
if not np.isclose(belief_dist.sum(), 1.0):
raise ValueError("Distribution not normalized! Please normalize")
plt.grid(zorder=0)
plt.bar(range(belief_dist.shape[0]), belief_dist, color='r', zorder=3)
plt.xticks(range(belief_dist.shape[0]))
plt.title(title_str)
plt.show()
## --- replacements for pymdp.legacy.utils helpers that no longer exist in the JAX API ---
def np_zeros_list(shapes):
"""Replaces np_zeros_list(): a plain list of (mutable) numpy arrays.
The model is built in numpy and converted to jnp when the Agent is created."""
return [np.zeros(s) for s in shapes]
def np_uniform_list(dims):
"""Replaces np_uniform_list(): a list of uniform categorical vectors."""
return [np.ones(d)/d for d in dims]
def onehot(idx, dim):
"""Replaces onehot()."""
v = np.zeros(dim); v[idx] = 1.0
return v
def is_normalized(dist):
"""Replaces is_normalized(): True if every column (axis 0) sums to 1."""
if isinstance(dist, (list, tuple)):
return all(is_normalized(d) for d in dist)
return bool(np.allclose(np.asarray(dist).sum(axis=0), 1.0))There is no pre-existing data to be analyzed.
There is no pre-existing data to be prepared.
Please review the narrative in section 1.
This section attempts to answer three important questions:
For this problem, the only metric we are interested in is the probability of observing a REWARD. The only source of uncertainty is the location of the obstacle.
We assume that the reality of the generative process is given by the following implementation:
class ObstacleEnvir():
def __init__(self, pILeftI=0.7):
self.sˣ0_0 = 'CENTER' ## initial location
self.sˣ_0 = self.sˣ0_0 ## current location
self.sˣ_1 = None ## obstacle location
self.pILeftI = pILeftI ## probability of obstacle being on the Left
def step(self, a_lab):
if a_lab == "MOVE_LEFT": ## action label
if self.sˣ_0 == 'LEFT_REWARD':
self.sˣ_0 = 'LEFT_REWARD'
o_0 = 'LEFT_REWARD' ## location obs
o_1 = 'NULL' ## reward obs
o_2 = 'STAYED' ## move obs
elif self.sˣ_0 == 'LEFT':
self.sˣ_0 = 'LEFT_REWARD'
o_0 = 'LEFT_REWARD'
o_1 = 'REWARD'
o_2 = 'MOVED_LEFT'
elif self.sˣ_0 == 'CENTER':
if self.sˣ_1=='LEFT':
self.sˣ_0 = 'CENTER'
o_0 = 'CENTER'
o_1 = 'COLLISION'
o_2 = 'STAYED'
else: ## self.s̆_1=='RIGHT'
self.sˣ_0 = 'LEFT'
o_0 = 'LEFT'
o_1 = 'NULL'
o_2 = 'MOVED_LEFT'
elif self.sˣ_0 == 'RIGHT':
self.sˣ_0 = 'CENTER'
o_0 = 'CENTER'
o_1 = 'NULL'
o_2 = 'MOVED_LEFT'
elif self.sˣ_0 == 'RIGHT_REWARD':
self.sˣ_0 = 'RIGHT'
o_0 = 'RIGHT'
o_1 = 'NULL'
o_2 = 'MOVED_LEFT'
else:
print(f'ERROR: Invalid action_lab: {a_lab}')
elif a_lab == "MOVE_RIGHT":
if self.sˣ_0 == 'RIGHT_REWARD':
self.sˣ_0 = 'RIGHT_REWARD'
o_0 = 'RIGHT_REWARD'
o_1 = 'NULL'
o_2 = 'STAYED'
elif self.sˣ_0 == 'RIGHT':
self.sˣ_0 = 'RIGHT_REWARD'
o_0 = 'RIGHT_REWARD'
o_1 = 'REWARD'
o_2 = 'MOVED_RIGHT'
elif self.sˣ_0 == 'CENTER':
if self.sˣ_1=='RIGHT':
self.sˣ_0 = 'CENTER'
o_0 = 'CENTER'
o_1 = 'COLLISION'
o_2 = 'STAYED'
else: ## self.sˣ_1=='Left'
self.sˣ_0 = 'RIGHT'
o_0 = 'RIGHT'
o_1 = 'NULL'
o_2 = 'MOVED_RIGHT'
elif self.sˣ_0 == 'LEFT':
self.sˣ_0 = 'CENTER'
o_0 = 'CENTER'
o_1 = 'NULL'
o_2 = 'MOVED_RIGHT'
elif self.sˣ_0 == 'LEFT_REWARD':
self.sˣ_0 = 'LEFT'
o_0 = 'LEFT'
o_1 = 'NULL'
o_2 = 'MOVED_RIGHT'
else:
print(f'ERROR: Invalid action_lab: {a_lab}')
obs = [o_0, o_1, o_2]
return obs
def reset(self):
self.sˣ_0 = self.sˣ0_0
self.sˣ_1 = 'LEFT' if random.random() < self.pILeftI else 'RIGHT'
print(f'{self.sˣ_1=}')
print(f'Re-initialized location to {self.sˣ0_0}')
o_0 = self.sˣ_0
o_2 = 'STAYED'
o_1 = 'NULL'
print(f'{o_0=}, {o_1=}, {o_2=}')
return o_0, o_1, o_2Give this environment class a quick test drive:
self.sˣ_1='LEFT'
Re-initialized location to CENTER
o_0='CENTER', o_1='NULL', o_2='STAYED'
('CENTER', 'NULL', 'STAYED')
The uncertainty of the location of the Obstacle Location is provided for in the reset() function of class ObstacleEnvir()
The following lookup dictionary allows for lookup between indexes & labels:
lup = { ## lookup between indexes & labels
## control/action factors
'a_0': [ ## Move
'MOVE_LEFT', 'MOVE_RIGHT'],
## state factors
's_0': [ ## Location
'LEFT_REWARD', 'LEFT', 'CENTER', 'RIGHT', 'RIGHT_REWARD'],
's_1': [ ## Obstacle Location ## uncontrollable
'LEFT', 'RIGHT'],
## observation modalities
'o_0': [ ## Location observation
'LEFT_REWARD', 'LEFT', 'CENTER', 'RIGHT', 'RIGHT_REWARD'],
'o_1': [ ## Reward observation
'NULL', 'REWARD', 'COLLISION'],
'o_2': [ ## Move observation
'STAYED', 'MOVED_LEFT', 'MOVED_RIGHT'],
}[array([[[0., 0.],
[0., 0.],
[0., 0.],
[0., 0.],
[0., 0.]],
[[0., 0.],
[0., 0.],
[0., 0.],
[0., 0.],
[0., 0.]],
[[0., 0.],
[0., 0.],
[0., 0.],
[0., 0.],
[0., 0.]],
[[0., 0.],
[0., 0.],
[0., 0.],
[0., 0.],
[0., 0.]],
[[0., 0.],
[0., 0.],
[0., 0.],
[0., 0.],
[0., 0.]]]),
array([[[0., 0.],
[0., 0.],
[0., 0.],
[0., 0.],
[0., 0.]],
[[0., 0.],
[0., 0.],
[0., 0.],
[0., 0.],
[0., 0.]],
[[0., 0.],
[0., 0.],
[0., 0.],
[0., 0.],
[0., 0.]]]),
array([[[0., 0.],
[0., 0.],
[0., 0.],
[0., 0.],
[0., 0.]],
[[0., 0.],
[0., 0.],
[0., 0.],
[0., 0.],
[0., 0.]],
[[0., 0.],
[0., 0.],
[0., 0.],
[0., 0.],
[0., 0.]]])]
Assume there is no noise in the observations. This means the Location of the agent is reported accurately as \(o_0\) (proprioceptive observation modality). Note that the agent can not occupy the location where the obstacle is.
## o_0 s_0 s_1
_A[0][:, :, 0] = np.eye(_s_dims[0])
_A[0][:, :, 1] = np.eye(_s_dims[0])
print(_A[0].shape)
_A[0](5, 5, 2)
array([[[1., 1.],
[0., 0.],
[0., 0.],
[0., 0.],
[0., 0.]],
[[0., 0.],
[1., 1.],
[0., 0.],
[0., 0.],
[0., 0.]],
[[0., 0.],
[0., 0.],
[1., 1.],
[0., 0.],
[0., 0.]],
[[0., 0.],
[0., 0.],
[0., 0.],
[1., 1.],
[0., 0.]],
[[0., 0.],
[0., 0.],
[0., 0.],
[0., 0.],
[1., 1.]]])
_A[1][
lup['o_1'].index('REWARD'),
lup['s_0'].index('RIGHT_REWARD'),
lup['s_1'].index('LEFT')
] = 1.0
_A[1][
lup['o_1'].index('REWARD'),
lup['s_0'].index('LEFT_REWARD'),
lup['s_1'].index('RIGHT')
] = 1.0
_A[1][
lup['o_1'].index('NULL'),
lup['s_0'].index('CENTER'),
lup['s_1'].index('RIGHT')
] = 0.5
_A[1][
lup['o_1'].index('COLLISION'),
lup['s_0'].index('CENTER'),
lup['s_1'].index('RIGHT')
] = 0.5
_A[1][
lup['o_1'].index('NULL'),
lup['s_0'].index('CENTER'),
lup['s_1'].index('LEFT')
] = 0.5
_A[1][
lup['o_1'].index('COLLISION'),
lup['s_0'].index('CENTER'),
lup['s_1'].index('LEFT')
] = 0.5
_A[1][
lup['o_1'].index('NULL'),
lup['s_0'].index('RIGHT'),
lup['s_1'].index('LEFT')
] = 1.0
_A[1][
lup['o_1'].index('NULL'),
lup['s_0'].index('LEFT'),
lup['s_1'].index('RIGHT')
] = 1.0
_A[1][
lup['o_1'].index('NULL'),
lup['s_0'].index('LEFT_REWARD'),
lup['s_1'].index('LEFT')
] = 1.0
_A[1][
lup['o_1'].index('NULL'),
lup['s_0'].index('RIGHT_REWARD'),
lup['s_1'].index('RIGHT')
] = 1.0
_A[1][
lup['o_1'].index('NULL'),
lup['s_0'].index('LEFT'),
lup['s_1'].index('LEFT')
] = 1.0
_A[1][
lup['o_1'].index('NULL'),
lup['s_0'].index('RIGHT'),
lup['s_1'].index('RIGHT')
] = 1.0
print(_A[1].shape)
_A[1](3, 5, 2)
array([[[1. , 0. ],
[1. , 1. ],
[0.5, 0.5],
[1. , 1. ],
[0. , 1. ]],
[[0. , 1. ],
[0. , 0. ],
[0. , 0. ],
[0. , 0. ],
[1. , 0. ]],
[[0. , 0. ],
[0. , 0. ],
[0.5, 0.5],
[0. , 0. ],
[0. , 0. ]]])
plot_likelihood(
_A[1][:,:,0],
title_str="""
Modality $o_1$ (Reward observation)
vs
Factor $s_0$ (Location)
($s_1=\mathrm{LEFT}$)
""",
yticklabels=lup['o_1'],
xticklabels=lup['s_0']
)
plot_likelihood(
_A[1][:,:,1],
title_str="""
Modality $o_1$ (Reward observation)
vs
Factor $s_0$ (Location)
($s_1=\mathrm{RIGHT}$)
""",
yticklabels=lup['o_1'],
xticklabels=lup['s_0']
)
plot_likelihood(
_A[1][:,0,:],
title_str="""
Modality $o_1$ (Reward observation)
vs
Factor $s_1$ (Obstacle Location)
($s_0=\mathrm{LEFT\_REWARD}$)
""",
yticklabels=lup['o_1'],
xticklabels=lup['s_1']
)
plot_likelihood(
_A[1][:,1,:],
title_str="""
Modality $o_1$ (Reward observation)
vs
Factor $s_1$ (Obstacle Location)
($s_0=\mathrm{LEFT}$)
""",
yticklabels=lup['o_1'],
xticklabels=lup['s_1']
)
plot_likelihood(
_A[1][:,2,:],
title_str="""
Modality $o_1$ (Reward observation)
vs
Factor $s_1$ (Obstacle Location)
($s_0=\mathrm{CENTER}$)
""",
yticklabels=lup['o_1'],
xticklabels=lup['s_1']
)
_A[2][
lup['o_2'].index('STAYED'),
lup['s_0'].index('CENTER'),
lup['s_1'].index('RIGHT')
] = 0.5
_A[2][
lup['o_2'].index('MOVED_RIGHT'),
lup['s_0'].index('CENTER'),
lup['s_1'].index('RIGHT')
] = 0.5
_A[2][
lup['o_2'].index('STAYED'),
lup['s_0'].index('CENTER'),
lup['s_1'].index('LEFT')
] = 0.5
_A[2][
lup['o_2'].index('MOVED_LEFT'),
lup['s_0'].index('CENTER'),
lup['s_1'].index('LEFT')
] = 0.5
_A[2][
lup['o_2'].index('STAYED'),
lup['s_0'].index('LEFT_REWARD'),
lup['s_1'].index('LEFT')
] = 1.0
_A[2][
lup['o_2'].index('MOVED_LEFT'),
lup['s_0'].index('LEFT_REWARD'),
lup['s_1'].index('RIGHT')
] = 0.6
_A[2][
lup['o_2'].index('STAYED'),
lup['s_0'].index('LEFT_REWARD'),
lup['s_1'].index('RIGHT')
] = 0.4
_A[2][
lup['o_2'].index('STAYED'),
lup['s_0'].index('LEFT'),
lup['s_1'].index('LEFT')
] = 1.0
_A[2][
lup['o_2'].index('MOVED_LEFT'),
lup['s_0'].index('LEFT'),
lup['s_1'].index('RIGHT')
] = 0.9
_A[2][
lup['o_2'].index('MOVED_RIGHT'),
lup['s_0'].index('LEFT'),
lup['s_1'].index('RIGHT')
] = 0.1
_A[2][
lup['o_2'].index('MOVED_LEFT'),
lup['s_0'].index('RIGHT'),
lup['s_1'].index('LEFT')
] = 0.1
_A[2][
lup['o_2'].index('MOVED_RIGHT'),
lup['s_0'].index('RIGHT'),
lup['s_1'].index('LEFT')
] = 0.9
_A[2][
lup['o_2'].index('STAYED'),
lup['s_0'].index('RIGHT'),
lup['s_1'].index('RIGHT')
] = 1.0
_A[2][
lup['o_2'].index('STAYED'),
lup['s_0'].index('LEFT'),
lup['s_1'].index('LEFT')
] = 1.0
_A[2][
lup['o_2'].index('MOVED_RIGHT'),
lup['s_0'].index('RIGHT_REWARD'),
lup['s_1'].index('LEFT')
] = 0.6
_A[2][
lup['o_2'].index('STAYED'),
lup['s_0'].index('RIGHT_REWARD'),
lup['s_1'].index('LEFT')
] = 0.4
_A[2][
lup['o_2'].index('STAYED'),
lup['s_0'].index('RIGHT_REWARD'),
lup['s_1'].index('RIGHT')
] = 1.0
print(_A[2].shape)
_A[2](3, 5, 2)
array([[[1. , 0.4],
[1. , 0. ],
[0.5, 0.5],
[0. , 1. ],
[0.4, 1. ]],
[[0. , 0.6],
[0. , 0.9],
[0.5, 0. ],
[0.1, 0. ],
[0. , 0. ]],
[[0. , 0. ],
[0. , 0.1],
[0. , 0.5],
[0.9, 0. ],
[0.6, 0. ]]])
plot_likelihood(
_A[2][:,:,0],
title_str="""
Modality $o_2$ (Move observation)
vs
Factor $s_0$ (Location)
($s_1=\mathrm{LEFT}$)
""",
yticklabels=lup['o_2'],
xticklabels=lup['s_0']
)
plot_likelihood(
_A[2][:,:,1],
title_str="""
Modality $o_2$ (Move observation)
vs
Factor $s_0$ (Location)
($s_1=\mathrm{RIGHT}$)
""",
yticklabels=lup['o_2'],
xticklabels=lup['s_0']
)
plot_likelihood(
_A[2][:,0,:],
title_str="""
Modality $o_2$ (Move observation)
vs
Factor $s_1$ (Obstacle Location)
($s_0=\mathrm{LEFT\_REWARD}$)
""",
yticklabels=lup['o_2'],
xticklabels=lup['s_1']
)
plot_likelihood(
_A[2][:,1,:],
title_str="""
Modality $o_2$ (Move observation)
vs
Factor $s_1$ (Obstacle Location)
($s_0=\mathrm{LEFT}$)
""",
yticklabels=lup['o_2'],
xticklabels=lup['s_1']
)
plot_likelihood(
_A[2][:,2,:],
title_str="""
Modality $o_2$ (Move observation)
vs
Factor $s_1$ (Obstacle Location)
($s_0=\mathrm{CENTER}$)
""",
yticklabels=lup['o_2'],
xticklabels=lup['s_1']
)
plot_likelihood(
_A[2][:,3,:],
title_str="""
Modality $o_2$ (Move observation)
vs
Factor $s_1$ (Obstacle Location)
($s_0=\mathrm{RIGHT}$)
""",
yticklabels=lup['o_2'],
xticklabels=lup['s_1']
)
plot_likelihood(
_A[2][:,4,:],
title_str="""
Modality $o_2$ (Move observation)
vs
Factor $s_1$ (Obstacle Location)
($s_0=\mathrm{RIGHT\_REWARD}$)
""",
yticklabels=lup['o_2'],
xticklabels=lup['s_1']
)
=== _s_dims:
[5, 2]
=== _o_dims:
[5, 3, 3]
[array([[[1., 1.],
[0., 0.],
[0., 0.],
[0., 0.],
[0., 0.]],
[[0., 0.],
[1., 1.],
[0., 0.],
[0., 0.],
[0., 0.]],
[[0., 0.],
[0., 0.],
[1., 1.],
[0., 0.],
[0., 0.]],
[[0., 0.],
[0., 0.],
[0., 0.],
[1., 1.],
[0., 0.]],
[[0., 0.],
[0., 0.],
[0., 0.],
[0., 0.],
[1., 1.]]]),
array([[[1. , 0. ],
[1. , 1. ],
[0.5, 0.5],
[1. , 1. ],
[0. , 1. ]],
[[0. , 1. ],
[0. , 0. ],
[0. , 0. ],
[0. , 0. ],
[1. , 0. ]],
[[0. , 0. ],
[0. , 0. ],
[0.5, 0.5],
[0. , 0. ],
[0. , 0. ]]]),
array([[[1. , 0.4],
[1. , 0. ],
[0.5, 0.5],
[0. , 1. ],
[0.4, 1. ]],
[[0. , 0.6],
[0. , 0.9],
[0.5, 0. ],
[0.1, 0. ],
[0. , 0. ]],
[[0. , 0. ],
[0. , 0.1],
[0. , 0.5],
[0.9, 0. ],
[0.6, 0. ]]])]
[[5, 5, 2], [2, 2, 1]]
[array([[[0., 0.],
[0., 0.],
[0., 0.],
[0., 0.],
[0., 0.]],
[[0., 0.],
[0., 0.],
[0., 0.],
[0., 0.],
[0., 0.]],
[[0., 0.],
[0., 0.],
[0., 0.],
[0., 0.],
[0., 0.]],
[[0., 0.],
[0., 0.],
[0., 0.],
[0., 0.],
[0., 0.]],
[[0., 0.],
[0., 0.],
[0., 0.],
[0., 0.],
[0., 0.]]]),
array([[[0.],
[0.]],
[[0.],
[0.]]])]
_B[0][
lup['s_0'].index('LEFT_REWARD'),
lup['s_0'].index('LEFT_REWARD'),
lup['a_0'].index('MOVE_LEFT')
] = 1.0
_B[0][
lup['s_0'].index('LEFT'),
lup['s_0'].index('LEFT_REWARD'),
lup['a_0'].index('MOVE_RIGHT')
] = 1.0
_B[0][
lup['s_0'].index('LEFT_REWARD'),
lup['s_0'].index('LEFT'),
lup['a_0'].index('MOVE_LEFT')
] = 1.0
_B[0][
lup['s_0'].index('CENTER'),
lup['s_0'].index('LEFT'),
lup['a_0'].index('MOVE_RIGHT')
] = 1.0
_B[0][ ## obstacle in LEFT
lup['s_0'].index('CENTER'),
lup['s_0'].index('CENTER'),
lup['a_0'].index('MOVE_LEFT')
] = 0.5
_B[0][ ## obstacle in RIGHT
lup['s_0'].index('LEFT'),
lup['s_0'].index('CENTER'),
lup['a_0'].index('MOVE_LEFT')
] = 0.5
_B[0][ ## obstacle in LEFT
lup['s_0'].index('RIGHT'),
lup['s_0'].index('CENTER'),
lup['a_0'].index('MOVE_RIGHT')
] = 0.5
_B[0][ ## obstacle in RIGHT
lup['s_0'].index('CENTER'),
lup['s_0'].index('CENTER'),
lup['a_0'].index('MOVE_RIGHT')
] = 0.5
_B[0][
lup['s_0'].index('CENTER'),
lup['s_0'].index('RIGHT'),
lup['a_0'].index('MOVE_LEFT')
] = 1.0
_B[0][
lup['s_0'].index('RIGHT_REWARD'),
lup['s_0'].index('RIGHT'),
lup['a_0'].index('MOVE_RIGHT')
] = 1.0
_B[0][
lup['s_0'].index('RIGHT'),
lup['s_0'].index('RIGHT_REWARD'),
lup['a_0'].index('MOVE_LEFT')
] = 1.0
_B[0][
lup['s_0'].index('RIGHT_REWARD'),
lup['s_0'].index('RIGHT_REWARD'),
lup['a_0'].index('MOVE_RIGHT')
] = 1.0
print(_B[0].shape)
_B[0](5, 5, 2)
array([[[1. , 0. ],
[1. , 0. ],
[0. , 0. ],
[0. , 0. ],
[0. , 0. ]],
[[0. , 1. ],
[0. , 0. ],
[0.5, 0. ],
[0. , 0. ],
[0. , 0. ]],
[[0. , 0. ],
[0. , 1. ],
[0.5, 0.5],
[1. , 0. ],
[0. , 0. ]],
[[0. , 0. ],
[0. , 0. ],
[0. , 0.5],
[0. , 0. ],
[1. , 0. ]],
[[0. , 0. ],
[0. , 0. ],
[0. , 0. ],
[0. , 1. ],
[0. , 1. ]]])
_B[1][
:,
:,
0 ## for uncontrollable, 0 for the only value: 'null/no-action'
] = np.eye(_s_dims[1]) ##.
_B[1]array([[[1.],
[0.]],
[[0.],
[1.]]])
plot_likelihood(
_B[1][:,:,lup['a_0'].index('MOVE_LEFT')],
title_str="""
Factor $s_{1,t}$ (Obstacle Location)
vs
Factor $s_{1,t-1}$ (Obstacle Location)
($a_{0,t-1}=\mathrm{MOVE\_LEFT}$)""",
yticklabels=lup['s_1'],
xticklabels=lup['s_1'],
)
=== _a_dims:
[2, 1]
=== _s_dims:
[5, 2]
[array([[[1. , 0. ],
[1. , 0. ],
[0. , 0. ],
[0. , 0. ],
[0. , 0. ]],
[[0. , 1. ],
[0. , 0. ],
[0.5, 0. ],
[0. , 0. ],
[0. , 0. ]],
[[0. , 0. ],
[0. , 1. ],
[0.5, 0.5],
[1. , 0. ],
[0. , 0. ]],
[[0. , 0. ],
[0. , 0. ],
[0. , 0.5],
[0. , 0. ],
[1. , 0. ]],
[[0. , 0. ],
[0. , 0. ],
[0. , 0. ],
[0. , 1. ],
[0. , 1. ]]]),
array([[[1.],
[0.]],
[[0.],
[1.]]])]
[array([0., 0., 0., 0., 0.]), array([0., 0., 0.]), array([0., 0., 0.])]
_C[1][1] = 2.0 ## make the agent want to encounter the 'Reward' observation level
_C[1][2] = -4.0 ## make the agent NOT want to encounter the 'Collision' observation level
## we don't care about the other observation modalities for preferences
_C[array([0., 0., 0., 0., 0.]), array([ 0., 2., -4.]), array([0., 0., 0.])]
The agent believes it always starts in \(s_0=\mathrm{'Center'}\). This belief aligns with reality as defined by the generative process.
[array([0., 0., 1., 0., 0.]), array([0.5, 0.5])]
## list of the indices of the hidden state factors that are controllable
## s_0 (Location): controllable; s_1 (Obstacle Location): not controllable
controllable_indices = [0]
agent = Agent(
A=[jnp.asarray(a) for a in _A], ## JAX: convert numpy model to jnp arrays
B=[jnp.asarray(b) for b in _B],
C=[jnp.asarray(c) for c in _C],
D=[jnp.asarray(d) for d in _D],
policy_len = 6,
control_fac_idx=controllable_indices,
action_selection="deterministic", ## same as legacy default
sampling_mode="marginal", ## legacy default; JAX default is "full"
)
agentAgent(
A=[f32[1,5,5,2], f32[1,3,5,2], f32[1,3,5,2]],
B=[f32[1,5,5,2], f32[1,2,2,1]],
C=[f32[1,5], f32[1,3], f32[1,3]],
D=[f32[1,5], f32[1,2]],
E=f32[1,64],
pA=None,
pB=None,
gamma=weak_f32[1],
alpha=weak_f32[1],
policies=Policies(
_policy_tup=(
((0, 0), (0, 0), (0, 0), (0, 0), (0, 0), (0, 0)),
((0, 0), (0, 0), (0, 0), (0, 0), (0, 0), (1, 0)),
((0, 0), (0, 0), (0, 0), (0, 0), (1, 0), (0, 0)),
((0, 0), (0, 0), (0, 0), (0, 0), (1, 0), (1, 0)),
((0, 0), (0, 0), (0, 0), (1, 0), (0, 0), (0, 0)),
((0, 0), (0, 0), (0, 0), (1, 0), (0, 0), (1, 0)),
((0, 0), (0, 0), (0, 0), (1, 0), (1, 0), (0, 0)),
((0, 0), (0, 0), (0, 0), (1, 0), (1, 0), (1, 0)),
((0, 0), (0, 0), (1, 0), (0, 0), (0, 0), (0, 0)),
((0, 0), (0, 0), (1, 0), (0, 0), (0, 0), (1, 0)),
((0, 0), (0, 0), (1, 0), (0, 0), (1, 0), (0, 0)),
((0, 0), (0, 0), (1, 0), (0, 0), (1, 0), (1, 0)),
((0, 0), (0, 0), (1, 0), (1, 0), (0, 0), (0, 0)),
((0, 0), (0, 0), (1, 0), (1, 0), (0, 0), (1, 0)),
((0, 0), (0, 0), (1, 0), (1, 0), (1, 0), (0, 0)),
((0, 0), (0, 0), (1, 0), (1, 0), (1, 0), (1, 0)),
((0, 0), (1, 0), (0, 0), (0, 0), (0, 0), (0, 0)),
((0, 0), (1, 0), (0, 0), (0, 0), (0, 0), (1, 0)),
((0, 0), (1, 0), (0, 0), (0, 0), (1, 0), (0, 0)),
((0, 0), (1, 0), (0, 0), (0, 0), (1, 0), (1, 0)),
((0, 0), (1, 0), (0, 0), (1, 0), (0, 0), (0, 0)),
((0, 0), (1, 0), (0, 0), (1, 0), (0, 0), (1, 0)),
((0, 0), (1, 0), (0, 0), (1, 0), (1, 0), (0, 0)),
((0, 0), (1, 0), (0, 0), (1, 0), (1, 0), (1, 0)),
((0, 0), (1, 0), (1, 0), (0, 0), (0, 0), (0, 0)),
((0, 0), (1, 0), (1, 0), (0, 0), (0, 0), (1, 0)),
((0, 0), (1, 0), (1, 0), (0, 0), (1, 0), (0, 0)),
((0, 0), (1, 0), (1, 0), (0, 0), (1, 0), (1, 0)),
((0, 0), (1, 0), (1, 0), (1, 0), (0, 0), (0, 0)),
((0, 0), (1, 0), (1, 0), (1, 0), (0, 0), (1, 0)),
((0, 0), (1, 0), (1, 0), (1, 0), (1, 0), (0, 0)),
((0, 0), (1, 0), (1, 0), (1, 0), (1, 0), (1, 0)),
((1, 0), (0, 0), (0, 0), (0, 0), (0, 0), (0, 0)),
((1, 0), (0, 0), (0, 0), (0, 0), (0, 0), (1, 0)),
((1, 0), (0, 0), (0, 0), (0, 0), (1, 0), (0, 0)),
((1, 0), (0, 0), (0, 0), (0, 0), (1, 0), (1, 0)),
((1, 0), (0, 0), (0, 0), (1, 0), (0, 0), (0, 0)),
((1, 0), (0, 0), (0, 0), (1, 0), (0, 0), (1, 0)),
((1, 0), (0, 0), (0, 0), (1, 0), (1, 0), (0, 0)),
((1, 0), (0, 0), (0, 0), (1, 0), (1, 0), (1, 0)),
((1, 0), (0, 0), (1, 0), (0, 0), (0, 0), (0, 0)),
((1, 0), (0, 0), (1, 0), (0, 0), (0, 0), (1, 0)),
((1, 0), (0, 0), (1, 0), (0, 0), (1, 0), (0, 0)),
((1, 0), (0, 0), (1, 0), (0, 0), (1, 0), (1, 0)),
((1, 0), (0, 0), (1, 0), (1, 0), (0, 0), (0, 0)),
((1, 0), (0, 0), (1, 0), (1, 0), (0, 0), (1, 0)),
((1, 0), (0, 0), (1, 0), (1, 0), (1, 0), (0, 0)),
((1, 0), (0, 0), (1, 0), (1, 0), (1, 0), (1, 0)),
((1, 0), (1, 0), (0, 0), (0, 0), (0, 0), (0, 0)),
((1, 0), (1, 0), (0, 0), (0, 0), (0, 0), (1, 0)),
((1, 0), (1, 0), (0, 0), (0, 0), (1, 0), (0, 0)),
((1, 0), (1, 0), (0, 0), (0, 0), (1, 0), (1, 0)),
((1, 0), (1, 0), (0, 0), (1, 0), (0, 0), (0, 0)),
((1, 0), (1, 0), (0, 0), (1, 0), (0, 0), (1, 0)),
((1, 0), (1, 0), (0, 0), (1, 0), (1, 0), (0, 0)),
((1, 0), (1, 0), (0, 0), (1, 0), (1, 0), (1, 0)),
((1, 0), (1, 0), (1, 0), (0, 0), (0, 0), (0, 0)),
((1, 0), (1, 0), (1, 0), (0, 0), (0, 0), (1, 0)),
((1, 0), (1, 0), (1, 0), (0, 0), (1, 0), (0, 0)),
((1, 0), (1, 0), (1, 0), (0, 0), (1, 0), (1, 0)),
((1, 0), (1, 0), (1, 0), (1, 0), (0, 0), (0, 0)),
((1, 0), (1, 0), (1, 0), (1, 0), (0, 0), (1, 0)),
((1, 0), (1, 0), (1, 0), (1, 0), (1, 0), (0, 0)),
((1, 0), (1, 0), (1, 0), (1, 0), (1, 0), (1, 0))
),
_dtype=dtype('int32'),
horizon=6,
num_policies=64
),
inductive_threshold=weak_f32[1],
inductive_epsilon=weak_f32[1],
H=None,
I=[f32[1,1,5], f32[1,1,2]],
A_dependencies=[[0, 1], [0, 1], [0, 1]],
B_dependencies=[[0], [1]],
B_action_dependencies=[[0], [1]],
action_maps=None,
batch_size=1,
num_iter=16,
num_obs=[5, 3, 3],
num_modalities=3,
num_states=[5, 2],
num_factors=2,
num_controls=[2, 1],
num_controls_multi=None,
control_fac_idx=[0],
policy_len=6,
inference_horizon=None,
inductive_depth=1,
use_utility=True,
use_states_info_gain=True,
use_param_info_gain=False,
use_inductive=False,
categorical_obs=False,
preprocess_fn=None,
action_selection='deterministic',
sampling_mode='marginal',
inference_algo='fpi',
learn_A=False,
learn_B=False,
learn_C=False,
learn_D=False,
learn_E=False
)
T = 20 ## number of timesteps
obs_lab = envir.reset() ## reset the environment and get an initial observation
obs_idx = [
lup['o_0'].index(obs_lab[0]),
lup['o_1'].index(obs_lab[1]),
lup['o_2'].index(obs_lab[2])]
## JAX: the agent no longer stores beliefs internally; we carry the empirical prior ourselves
prior = agent.D ## at t=0 the prior over states is D
##. with a
for t in range(T): ##. slide
print(f'{t=}')
print(f'{obs_idx=}')
obs = [jnp.array([o]) for o in obs_idx] ## leading batch dim (batch_size=1)
qs = agent.infer_states(obs, empirical_prior=prior) ##. ##. infer.
qIsI = [np.asarray(q[0, -1]) for q in qs] ## drop batch & time dims -> one vector per factor
plot_beliefs(qIsI[0], title_str = f"Beliefs about the Location at time {t}")
plot_beliefs(qIsI[1], title_str = f"Beliefs about the Obstacle Location at time {t}")
qIpiI, neg_efe = agent.infer_policies(qs) ##. (neg_efe = -EFE)
rng_key, key_action = jr.split(rng_key)
act = agent.sample_action(qIpiI, rng_key=jr.split(key_action, agent.batch_size)) ##. Act future &
act_idx = np.asarray(act[0]) ## drop batch dim
print(f'{act_idx=}')
act_idx_controllable = int(act_idx[0]); print(f'{act_idx_controllable=}')
act_lab_controllable = lup['a_0'][act_idx_controllable] ##.
## JAX: propagate beliefs through B under the chosen action -> prior for the next step
prior = agent.update_empirical_prior(act, qs)
obs_lab = envir.step(act_lab_controllable) ##. next observe so
obs_idx = [
lup['o_0'].index(obs_lab[0]),
lup['o_1'].index(obs_lab[1]),
lup['o_2'].index(obs_lab[2])]
print(f'location_obs: {obs_idx[0]}, reward_obs: {obs_idx[1]}, move_obs: {obs_idx[2]}')self.sˣ_1='RIGHT'
Re-initialized location to CENTER
o_0='CENTER', o_1='NULL', o_2='STAYED'
t=0
obs_idx=[2, 0, 0]
act_idx=array([0, 0], dtype=int32)
act_idx_controllable=0
location_obs: 1, reward_obs: 0, move_obs: 1
t=1
obs_idx=[1, 0, 1]
act_idx=array([0, 0], dtype=int32)
act_idx_controllable=0
location_obs: 0, reward_obs: 1, move_obs: 1
t=2
obs_idx=[0, 1, 1]
act_idx=array([0, 0], dtype=int32)
act_idx_controllable=0
location_obs: 0, reward_obs: 0, move_obs: 0
t=3
obs_idx=[0, 0, 0]
act_idx=array([0, 0], dtype=int32)
act_idx_controllable=0
location_obs: 0, reward_obs: 0, move_obs: 0
t=4
obs_idx=[0, 0, 0]
act_idx=array([1, 0], dtype=int32)
act_idx_controllable=1
location_obs: 1, reward_obs: 0, move_obs: 2
t=5
obs_idx=[1, 0, 2]
act_idx=array([1, 0], dtype=int32)
act_idx_controllable=1
location_obs: 2, reward_obs: 0, move_obs: 2
t=6
obs_idx=[2, 0, 2]
act_idx=array([0, 0], dtype=int32)
act_idx_controllable=0
location_obs: 1, reward_obs: 0, move_obs: 1
t=7
obs_idx=[1, 0, 1]
act_idx=array([0, 0], dtype=int32)
act_idx_controllable=0
location_obs: 0, reward_obs: 1, move_obs: 1
t=8
obs_idx=[0, 1, 1]
act_idx=array([0, 0], dtype=int32)
act_idx_controllable=0
location_obs: 0, reward_obs: 0, move_obs: 0
t=9
obs_idx=[0, 0, 0]
act_idx=array([0, 0], dtype=int32)
act_idx_controllable=0
location_obs: 0, reward_obs: 0, move_obs: 0
t=10
obs_idx=[0, 0, 0]
act_idx=array([1, 0], dtype=int32)
act_idx_controllable=1
location_obs: 1, reward_obs: 0, move_obs: 2
t=11
obs_idx=[1, 0, 2]
act_idx=array([1, 0], dtype=int32)
act_idx_controllable=1
location_obs: 2, reward_obs: 0, move_obs: 2
t=12
obs_idx=[2, 0, 2]
act_idx=array([0, 0], dtype=int32)
act_idx_controllable=0
location_obs: 1, reward_obs: 0, move_obs: 1
t=13
obs_idx=[1, 0, 1]
act_idx=array([0, 0], dtype=int32)
act_idx_controllable=0
location_obs: 0, reward_obs: 1, move_obs: 1
t=14
obs_idx=[0, 1, 1]
act_idx=array([0, 0], dtype=int32)
act_idx_controllable=0
location_obs: 0, reward_obs: 0, move_obs: 0
t=15
obs_idx=[0, 0, 0]
act_idx=array([0, 0], dtype=int32)
act_idx_controllable=0
location_obs: 0, reward_obs: 0, move_obs: 0
t=16
obs_idx=[0, 0, 0]
act_idx=array([1, 0], dtype=int32)
act_idx_controllable=1
location_obs: 1, reward_obs: 0, move_obs: 2
t=17
obs_idx=[1, 0, 2]
act_idx=array([1, 0], dtype=int32)
act_idx_controllable=1
location_obs: 2, reward_obs: 0, move_obs: 2
t=18
obs_idx=[2, 0, 2]
act_idx=array([0, 0], dtype=int32)
act_idx_controllable=0
location_obs: 1, reward_obs: 0, move_obs: 1
t=19
obs_idx=[1, 0, 1]
act_idx=array([0, 0], dtype=int32)
act_idx_controllable=0
location_obs: 0, reward_obs: 1, move_obs: 1








































In the run above, the environment places the obstacle on the right (RIGHT). The agent reaches LEFT_REWARD four times in 20 steps (at \(t = 1, 7, 13, 19\)) and never collides with the obstacle.
Its beliefs about the Obstacle Location are only partly correct. The belief plots show a cycle that repeats every six steps:
CENTER and believes both obstacle locations are equally likely. It moves left. Arriving in LEFT with the observation MOVED_LEFT tells it, correctly, that the obstacle must be on the right.LEFT_REWARD and observes REWARD, which confirms that belief.LEFT_REWARD, it observes NULL and STAYED. Under the agent’s model (A[1] and A[2]), these observations are more likely when the obstacle is on the left, so its belief flips to the wrong side (0.71 at \(t = 3\), 1.0 at \(t = 4\)).LEFT_REWARD, it moves back toward CENTER. There the move observation restores the correct belief, and the cycle starts again.The agent was deliberately designed to NOT have perfect knowledge of the workings of the environment, and these mismatches explain what we see:
LEFT_REWARD and RIGHT_REWARD end locations are different from the environment’s reality. In the environment, REWARD is observed only on the step the agent arrives at a reward location; staying there gives NULL and STAYED. The agent’s observation model depends only on the current states, so it cannot represent “just arrived”. Instead, it links these observations to the obstacle location, which causes the belief flip.COLLISION at CENTER is equally likely for both obstacle locations, and its movement model (B[0]) does not depend on the obstacle. A collision would therefore not tell the agent where the obstacle is. In this run the agent’s first move happened to be toward the free side, so no collision occurred.Despite these mismatches, the agent collects the reward repeatedly. A model in which the agent keeps a correct belief at the reward location and learns from collisions would need the obstacle to affect movement in the model itself, for example with B_dependencies (new in pymdp 1.0) or by combining location and obstacle into a single state factor. That is a topic for a future post.