Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 2 additions & 2 deletions src/unilab/envs/manager_based_rl_env.py
Original file line number Diff line number Diff line change
Expand Up @@ -629,8 +629,8 @@ def reset(
if self._autoreset_reset_active:
# Autoreset runs at the tail of step(): keep this step's
# per-step log entries (reward/* etc., computed pre-reset) and
# layer the reset extras (Episode_Reward/* etc.) on top, so
# consumers still see the transition's reward breakdown.
# layer manager reset extras on top, so consumers still see the
# transition's metrics.
step_log = self._state.info.get("log")
if step_log:
log = {**step_log, **log}
Expand Down
13 changes: 1 addition & 12 deletions src/unilab/managers/reward_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -39,8 +39,6 @@ class RewardManager(ManagerBase):

When ``scale_by_dt=True`` (default):
- ``reward_buf`` (returned by ``compute()``) = raw_value * weight * dt
- ``_episode_sums`` (cumulative rewards) are scaled by dt
- ``Episode_Reward/*`` logged metrics are scaled by dt

When ``scale_by_dt=False``:
- ``reward_buf`` = raw_value * weight (no dt scaling)
Expand Down Expand Up @@ -71,9 +69,6 @@ def __init__(

self.cfg = deepcopy(cfg)
super().__init__(env=env)
self._episode_sums = dict()
for term_name in self._term_names:
self._episode_sums[term_name] = np.zeros(self.num_envs, dtype=np.float32)
self._reward_buf = np.zeros(self.num_envs, dtype=np.float32)
self._step_reward = np.zeros((self.num_envs, len(self._term_names)), dtype=np.float32)
# Scratch for the weighted term value, reused across terms to avoid a
Expand Down Expand Up @@ -107,14 +102,9 @@ def active_terms(self) -> list[str]:
def reset(self, env_ids: np.ndarray | slice | None = None) -> dict[str, float]:
if env_ids is None:
env_ids = slice(None)
extras = {}
for key in self._episode_sums.keys():
episodic_sum_avg = float(np.mean(self._episode_sums[key][env_ids]))
extras["Episode_Reward/" + key] = episodic_sum_avg / self._env.max_episode_length_s
self._episode_sums[key][env_ids] = 0.0
for term_cfg in self._class_term_cfgs:
term_cfg.func.reset(env_ids=env_ids)
return extras
return {}

def compute(self, dt: float) -> np.ndarray:
if not np.isfinite(dt) or (self._scale_by_dt and dt <= 0.0):
Expand Down Expand Up @@ -142,7 +132,6 @@ def compute(self, dt: float) -> np.ndarray:
np.multiply(value, term_cfg.weight, out=scratch)
scratch *= scale
self._reward_buf += scratch
self._episode_sums[name] += scratch
np.divide(scratch, scale, out=self._step_reward[:, term_idx])
return self._reward_buf

Expand Down
7 changes: 5 additions & 2 deletions tests/managers/test_core_managers.py
Original file line number Diff line number Diff line change
Expand Up @@ -171,8 +171,11 @@ def test_reward_dt_scaling_reset_and_config_immutability(fake_env: FakeEnv) -> N
manager = RewardManager(cfg, fake_env)
np.testing.assert_allclose(manager.compute(dt=0.25), fake_env.value * 0.5)
assert manager.get_active_iterable_terms(2) == [("stateful", [4.0])]
extras = manager.reset(np.array([1, 2]))
assert extras["Episode_Reward/stateful"] == pytest.approx(0.375)
reset_ids = np.array([1, 2])
extras = manager.reset(reset_ids)
assert extras == {}
assert not any(key.startswith("Episode_Reward/") for key in extras)
assert manager.get_term_cfg("stateful").func.reset_ids is reset_ids
assert cfg["stateful"].func is StatefulReward
assert isinstance(manager.get_term_cfg("stateful").func, StatefulReward)

Expand Down
Loading