fix: removed deterministic flag
This commit is contained in:
parent
6db1129ca7
commit
5329ae18ba
1 changed files with 4 additions and 35 deletions
|
|
@ -41,8 +41,6 @@ class CleanRLPPOPolicy:
|
||||||
network_params: Any,
|
network_params: Any,
|
||||||
actor_params: Any,
|
actor_params: Any,
|
||||||
action_dim: int,
|
action_dim: int,
|
||||||
deterministic: bool = True,
|
|
||||||
seed: int = 0,
|
|
||||||
) -> None:
|
) -> None:
|
||||||
from brittle_star_project.rl import Actor, Network
|
from brittle_star_project.rl import Actor, Network
|
||||||
|
|
||||||
|
|
@ -54,16 +52,12 @@ class CleanRLPPOPolicy:
|
||||||
"network_params": network_params,
|
"network_params": network_params,
|
||||||
"actor_params": actor_params,
|
"actor_params": actor_params,
|
||||||
}
|
}
|
||||||
self._deterministic = deterministic
|
|
||||||
self._rng = jax.random.PRNGKey(int(seed))
|
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def load(
|
def load(
|
||||||
path: Path,
|
path: Path,
|
||||||
*,
|
*,
|
||||||
action_dim: int,
|
action_dim: int,
|
||||||
deterministic: bool,
|
|
||||||
seed: int,
|
|
||||||
) -> "CleanRLPPOPolicy":
|
) -> "CleanRLPPOPolicy":
|
||||||
def _get_index(container: Any, idx: int) -> Any:
|
def _get_index(container: Any, idx: int) -> Any:
|
||||||
if isinstance(container, (list, tuple)):
|
if isinstance(container, (list, tuple)):
|
||||||
|
|
@ -129,26 +123,16 @@ class CleanRLPPOPolicy:
|
||||||
network_params=network_params,
|
network_params=network_params,
|
||||||
actor_params=actor_params,
|
actor_params=actor_params,
|
||||||
action_dim=action_dim,
|
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:
|
def act(self, *, observations: dict[str, Any]) -> np.ndarray:
|
||||||
obs = _flatten_obs_dict(observations)
|
obs = _flatten_obs_dict(observations)
|
||||||
hidden = self._network_apply(self._params["network_params"], obs)
|
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:
|
# Always evaluate with the actor mean.
|
||||||
action = mean
|
# (Sampling adds exploration noise, which is useful for training but not for evaluation.)
|
||||||
else:
|
return np.asarray(mean, dtype=np.float32).ravel()
|
||||||
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()
|
|
||||||
|
|
||||||
|
|
||||||
def _get_observations(state: Any) -> dict[str, Any]:
|
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).
|
Returns (return, length, reached_target, final_xy_dist).
|
||||||
"""
|
"""
|
||||||
|
|
||||||
policy.reset(seed)
|
|
||||||
state = env.reset(seed=seed)
|
state = env.reset(seed=seed)
|
||||||
|
|
||||||
ep_return = 0.0
|
ep_return = 0.0
|
||||||
|
|
@ -294,8 +276,6 @@ def _run_one_episode_viewer(
|
||||||
|
|
||||||
|
|
||||||
def parse_args() -> argparse.Namespace:
|
def parse_args() -> argparse.Namespace:
|
||||||
from brittle_star_project.environment import Task
|
|
||||||
|
|
||||||
p = argparse.ArgumentParser(
|
p = argparse.ArgumentParser(
|
||||||
description="Run a trained policy for exactly one episode (viewer or headless)."
|
description="Run a trained policy for exactly one episode (viewer or headless)."
|
||||||
)
|
)
|
||||||
|
|
@ -305,12 +285,6 @@ def parse_args() -> argparse.Namespace:
|
||||||
required=True,
|
required=True,
|
||||||
help=("Path to a CleanRL/Flax '.cleanrl_model' checkpoint (saved by src/train.py)."),
|
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(
|
p.add_argument(
|
||||||
"--headless",
|
"--headless",
|
||||||
action="store_true",
|
action="store_true",
|
||||||
|
|
@ -369,14 +343,9 @@ def main() -> None:
|
||||||
policy = CleanRLPPOPolicy.load(
|
policy = CleanRLPPOPolicy.load(
|
||||||
model_path,
|
model_path,
|
||||||
action_dim=nu,
|
action_dim=nu,
|
||||||
deterministic=bool(args.deterministic),
|
|
||||||
seed=seed_for_env,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
# Reset policy RNG if a seed was provided.
|
|
||||||
default_seed = seed_for_env
|
default_seed = seed_for_env
|
||||||
if args.seed is not None:
|
|
||||||
policy.reset(int(args.seed))
|
|
||||||
|
|
||||||
# ======= SIMULATION =======
|
# ======= SIMULATION =======
|
||||||
|
|
||||||
|
|
|
||||||
Reference in a new issue