Skip to content

Commit cb5d8f4

Browse files
committed
logging: Sequences Containing Done considers critic mask
1 parent 489236b commit cb5d8f4

2 files changed

Lines changed: 5 additions & 3 deletions

File tree

amago/agent.py

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -608,7 +608,9 @@ def masked_avg(x_, dim=0):
608608
"Max TD Target": td_target[where_mask].max(),
609609
"TD Target (test-time gamma)": masked_avg(td_target, -1),
610610
"Mean Reward (in training sequences)": masked_avg(r),
611-
"Sequences Containing Done": d[:, :, 0, 0, 0].any(1).sum(),
611+
"Sequences Containing Done": (d * mask.all(2, keepdim=True))
612+
.any((1, 2, 3))
613+
.sum(),
612614
"Min Reward (in training sequences)": r[where_mask].min(),
613615
"Max Reward (in training sequences)": r[where_mask].max(),
614616
}

docs/tutorial/customization.rst

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -113,14 +113,14 @@ RLDataset
113113
114114
Actor Output Head
115115
~~~~~~~~~~~~~~~~~~~
116-
**Implement** :py:class:`~amago.nets.actor_critic.BaseActorHead`
116+
**Implement**: :py:class:`~amago.nets.actor_critic.BaseActorHead`
117117

118118
**Configure**: ``Agent.actor_type : MyActor``
119119

120120
|
121121
122122
Critic Output Head
123123
~~~~~~~~~~~~~~~~~~~
124-
**Implement** : :py:class:`~amago.nets.actor_critic.BaseCriticHead`
124+
**Implement**: :py:class:`~amago.nets.actor_critic.BaseCriticHead`
125125

126126
**Configure**: ``Agent.critic_type : MyCritic``

0 commit comments

Comments
 (0)