From 5329ae18ba809d1d25b0ec4f97cfb67f1334eb05 Mon Sep 17 00:00:00 2001 From: Jona Reynaert Date: Thu, 2 Apr 2026 10:36:18 +0200 Subject: [PATCH] fix: removed deterministic flag --- experiments/simulate.py | 39 ++++----------------------------------- 1 file changed, 4 insertions(+), 35 deletions(-) diff --git a/experiments/simulate.py b/experiments/simulate.py index 3b79082..3dca83c 100644 --- a/experiments/simulate.py +++ b/experiments/simulate.py @@ -41,8 +41,6 @@ class CleanRLPPOPolicy: network_params: Any, actor_params: Any, action_dim: int, - deterministic: bool = True, - seed: int = 0, ) -> None: from brittle_star_project.rl import Actor, Network @@ -54,16 +52,12 @@ class CleanRLPPOPolicy: "network_params": network_params, "actor_params": actor_params, } - self._deterministic = deterministic - self._rng = jax.random.PRNGKey(int(seed)) @staticmethod def load( path: Path, *, action_dim: int, - deterministic: bool, - seed: int, ) -> "CleanRLPPOPolicy": def _get_index(container: Any, idx: int) -> Any: if isinstance(container, (list, tuple)): @@ -129,26 +123,16 @@ class CleanRLPPOPolicy: network_params=network_params, actor_params=actor_params, action_dim=action_dim, - deterministic=deterministic, - seed=seed, ) - def reset(self, seed: int) -> None: - self._rng = jax.random.PRNGKey(int(seed)) - def act(self, *, observations: dict[str, Any]) -> np.ndarray: obs = _flatten_obs_dict(observations) hidden = self._network_apply(self._params["network_params"], obs) - mean, log_std = self._actor_apply(self._params["actor_params"], hidden) + mean, _log_std = self._actor_apply(self._params["actor_params"], hidden) - if self._deterministic: - action = mean - else: - self._rng, sub = jax.random.split(self._rng) - noise = jax.random.normal(sub, shape=mean.shape) - action = mean + noise * jnp.exp(log_std) - - return np.asarray(action, dtype=np.float32).ravel() + # Always evaluate with the actor mean. + # (Sampling adds exploration noise, which is useful for training but not for evaluation.) + return np.asarray(mean, dtype=np.float32).ravel() def _get_observations(state: Any) -> dict[str, Any]: @@ -174,8 +158,6 @@ def _rollout_one_episode_headless( Returns (return, length, reached_target, final_xy_dist). """ - - policy.reset(seed) state = env.reset(seed=seed) ep_return = 0.0 @@ -294,8 +276,6 @@ def _run_one_episode_viewer( def parse_args() -> argparse.Namespace: - from brittle_star_project.environment import Task - p = argparse.ArgumentParser( description="Run a trained policy for exactly one episode (viewer or headless)." ) @@ -305,12 +285,6 @@ def parse_args() -> argparse.Namespace: required=True, help=("Path to a CleanRL/Flax '.cleanrl_model' checkpoint (saved by src/train.py)."), ) - p.add_argument( - "--deterministic", - action=argparse.BooleanOptionalAction, - default=True, - help="Use mean action (deterministic) or sample actions (stochastic).", - ) p.add_argument( "--headless", action="store_true", @@ -369,14 +343,9 @@ def main() -> None: policy = CleanRLPPOPolicy.load( model_path, action_dim=nu, - deterministic=bool(args.deterministic), - seed=seed_for_env, ) - # Reset policy RNG if a seed was provided. default_seed = seed_for_env - if args.seed is not None: - policy.reset(int(args.seed)) # ======= SIMULATION =======