fix: Merge with origin/dev
This commit is contained in:
commit
aad086cb7d
2 changed files with 89 additions and 39 deletions
10
configs/hpc/wandb_expand.yaml
Normal file
10
configs/hpc/wandb_expand.yaml
Normal file
|
|
@ -0,0 +1,10 @@
|
||||||
|
exp_name: "explained_var_fun_more_steps"
|
||||||
|
seed: 42
|
||||||
|
track: true
|
||||||
|
wandb_project_name: "LET-THERE-BE-MORE-LOGGING"
|
||||||
|
wandb_entity: "SEL3-2026-Groep-4"
|
||||||
|
|
||||||
|
num_envs: 16
|
||||||
|
num_steps: 256
|
||||||
|
total_timesteps: 50000
|
||||||
|
cuda: true
|
||||||
|
|
@ -28,6 +28,12 @@ from brittle_star_project.ppo import PPO
|
||||||
def _clip_action(action: jnp.ndarray, low: jnp.ndarray, high: jnp.ndarray) -> jnp.ndarray:
|
def _clip_action(action: jnp.ndarray, low: jnp.ndarray, high: jnp.ndarray) -> jnp.ndarray:
|
||||||
return jnp.clip(action, low, high)
|
return jnp.clip(action, low, high)
|
||||||
|
|
||||||
|
def _compute_explained_variance(values: jnp.ndarray, returns: jnp.ndarray) -> float:
|
||||||
|
var_returns = jnp.var(returns)
|
||||||
|
explained_var = 1.0 - jnp.var(returns - values) / (var_returns + 1e-8)
|
||||||
|
return float(explained_var)
|
||||||
|
|
||||||
|
|
||||||
@jax.jit
|
@jax.jit
|
||||||
def _linear_schedule(count, minibatch_count, update_epochs, num_iterations, learning_rate):
|
def _linear_schedule(count, minibatch_count, update_epochs, num_iterations, learning_rate):
|
||||||
frac = 1.0 - (count // (minibatch_count * update_epochs)) / num_iterations
|
frac = 1.0 - (count // (minibatch_count * update_epochs)) / num_iterations
|
||||||
|
|
@ -41,7 +47,6 @@ def _convert_obs_dict_to_array(obs_dict: dict) -> jnp.ndarray:
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
# removed jit: used in _rollout_jit, so will be compiled with _rollout_jit
|
|
||||||
def _get_action_and_value_noise(
|
def _get_action_and_value_noise(
|
||||||
sensor: GenericDenseLayersWithActivation,
|
sensor: GenericDenseLayersWithActivation,
|
||||||
feature_extractor: GenericDenseLayersWithActivation,
|
feature_extractor: GenericDenseLayersWithActivation,
|
||||||
|
|
@ -58,7 +63,6 @@ def _get_action_and_value_noise(
|
||||||
agent_state.params["feature_extractor_params"], next_obs
|
agent_state.params["feature_extractor_params"], next_obs
|
||||||
)
|
)
|
||||||
|
|
||||||
# Continuous actions: sample from a Gaussian parameterized by the actor
|
|
||||||
mean, log_std = actor.apply(agent_state.params["actor_params"], hidden)
|
mean, log_std = actor.apply(agent_state.params["actor_params"], hidden)
|
||||||
key, subkey = jax.random.split(key)
|
key, subkey = jax.random.split(key)
|
||||||
noise = jax.random.normal(subkey, shape=mean.shape)
|
noise = jax.random.normal(subkey, shape=mean.shape)
|
||||||
|
|
@ -74,7 +78,6 @@ def _get_action_and_value_noise(
|
||||||
return clipped_action, logprob, value.squeeze(-1), key
|
return clipped_action, logprob, value.squeeze(-1), key
|
||||||
|
|
||||||
|
|
||||||
# removed jit: used in _rollout_jit, so will be compiled with _rollout_jit
|
|
||||||
def _step_once(
|
def _step_once(
|
||||||
carry,
|
carry,
|
||||||
_,
|
_,
|
||||||
|
|
@ -108,15 +111,13 @@ def _step_once(
|
||||||
return (agent_state, episode_stats, next_obs, next_done, key, env_state), storage
|
return (agent_state, episode_stats, next_obs, next_done, key, env_state), storage
|
||||||
|
|
||||||
|
|
||||||
# removed jit: used in _rollout_jit, so will be compiled with _rollout_jit
|
|
||||||
def _step_env_wrapped(episode_stats, env_state, action, env_step_fn):
|
def _step_env_wrapped(episode_stats, env_state, action, env_step_fn):
|
||||||
next_env_state = env_step_fn(env_state, action)
|
next_env_state = env_step_fn(env_state, action)
|
||||||
|
|
||||||
# Extract per-environment signals from the state object
|
reward = next_env_state.reward
|
||||||
reward = next_env_state.reward # (num_envs,)
|
terminated = next_env_state.terminated
|
||||||
terminated = next_env_state.terminated # (num_envs,)
|
truncated = next_env_state.truncated
|
||||||
truncated = next_env_state.truncated # (num_envs,)
|
done = terminated | truncated
|
||||||
done = terminated | truncated # (num_envs,)
|
|
||||||
|
|
||||||
new_episode_return = episode_stats.episode_returns + reward
|
new_episode_return = episode_stats.episode_returns + reward
|
||||||
new_episode_length = episode_stats.episode_lengths + 1
|
new_episode_length = episode_stats.episode_lengths + 1
|
||||||
|
|
@ -138,7 +139,6 @@ def _step_env_wrapped(episode_stats, env_state, action, env_step_fn):
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
# jit applied in wrapper method self._rollout_jit using partial
|
|
||||||
def _rollout_jit(
|
def _rollout_jit(
|
||||||
agent_state,
|
agent_state,
|
||||||
episode_stats,
|
episode_stats,
|
||||||
|
|
@ -173,7 +173,6 @@ def _rollout_jit(
|
||||||
return agent_state, episode_stats, next_obs, next_done, storage, key, env_state
|
return agent_state, episode_stats, next_obs, next_done, storage, key, env_state
|
||||||
|
|
||||||
|
|
||||||
# removed jit: used in _compute_gae_jit, so will be compiled with _compute_gae_jit
|
|
||||||
def _compute_gae_once(carry, inp, gamma, gae_lambda):
|
def _compute_gae_once(carry, inp, gamma, gae_lambda):
|
||||||
advantages = carry
|
advantages = carry
|
||||||
nextdone, nextvalues, curvalues, reward = inp
|
nextdone, nextvalues, curvalues, reward = inp
|
||||||
|
|
@ -183,13 +182,20 @@ def _compute_gae_once(carry, inp, gamma, gae_lambda):
|
||||||
return advantages, advantages
|
return advantages, advantages
|
||||||
|
|
||||||
|
|
||||||
# jit applied on partial-wrapped wrapper method self._compute_gae_jit
|
|
||||||
def _compute_gae_jit(
|
def _compute_gae_jit(
|
||||||
agent_state, storage, next_obs, next_done, gamma, gae_lambda, num_envs, sensor, critic
|
agent_state,
|
||||||
|
storage,
|
||||||
|
next_obs,
|
||||||
|
next_done,
|
||||||
|
gamma,
|
||||||
|
gae_lambda,
|
||||||
|
num_envs,
|
||||||
|
feature_extractor,
|
||||||
|
critic,
|
||||||
):
|
):
|
||||||
next_value = critic.apply(
|
next_value = critic.apply(
|
||||||
agent_state.params["critic_params"],
|
agent_state.params["critic_params"],
|
||||||
sensor.apply(agent_state.params["sensor_params"], next_obs),
|
feature_extractor.apply(agent_state.params["feature_extractor_params"], next_obs),
|
||||||
).squeeze(-1)
|
).squeeze(-1)
|
||||||
|
|
||||||
advantages = jnp.zeros((num_envs,))
|
advantages = jnp.zeros((num_envs,))
|
||||||
|
|
@ -205,14 +211,18 @@ def _compute_gae_jit(
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class LossInfo:
|
class TrainingMeasurements:
|
||||||
# todo: better typing
|
loss: jnp.ndarray
|
||||||
loss: Any
|
pg_loss: jnp.ndarray
|
||||||
pg_loss: Any
|
v_loss: jnp.ndarray
|
||||||
v_loss: Any
|
entropy_loss: jnp.ndarray
|
||||||
entropy_loss: Any
|
approx_kl: jnp.ndarray
|
||||||
approx_kl: Any
|
avg_episodic_return: float
|
||||||
avg_episodic_return: Any
|
explained_variance: float
|
||||||
|
num_terminated: int
|
||||||
|
num_truncated: int
|
||||||
|
avg_terminated_length: Any
|
||||||
|
avg_truncated_length: Any
|
||||||
|
|
||||||
|
|
||||||
class PPOTrainer:
|
class PPOTrainer:
|
||||||
|
|
@ -253,7 +263,7 @@ class PPOTrainer:
|
||||||
num_envs=self.args.num_envs,
|
num_envs=self.args.num_envs,
|
||||||
gamma=self.args.gamma,
|
gamma=self.args.gamma,
|
||||||
gae_lambda=self.args.gae_lambda,
|
gae_lambda=self.args.gae_lambda,
|
||||||
sensor=self.sensor,
|
feature_extractor=self.feature_extractor,
|
||||||
critic=self.critic,
|
critic=self.critic,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
@ -277,11 +287,8 @@ class PPOTrainer:
|
||||||
|
|
||||||
sensor = GenericDenseLayersWithActivation()
|
sensor = GenericDenseLayersWithActivation()
|
||||||
feature_extractor = GenericDenseLayersWithActivation()
|
feature_extractor = GenericDenseLayersWithActivation()
|
||||||
actor = Actor(
|
actor = Actor(action_dim=self.env.single_action_space.shape[0])
|
||||||
action_dim=self.env.single_action_space.shape[0]
|
|
||||||
) # continuous actions for MJX
|
|
||||||
critic = OneDenseLayerMLP()
|
critic = OneDenseLayerMLP()
|
||||||
# messenger = OneDenseLayerMLP()
|
|
||||||
return sensor, feature_extractor, actor, critic
|
return sensor, feature_extractor, actor, critic
|
||||||
|
|
||||||
def _init_agent_state(self) -> TrainState:
|
def _init_agent_state(self) -> TrainState:
|
||||||
|
|
@ -332,7 +339,7 @@ class PPOTrainer:
|
||||||
def _init_episode_stats(self) -> EpisodeStatistics:
|
def _init_episode_stats(self) -> EpisodeStatistics:
|
||||||
self.logger.info("[EPISODE STATS]: Initializing episode stats...")
|
self.logger.info("[EPISODE STATS]: Initializing episode stats...")
|
||||||
|
|
||||||
return EpisodeStatistics( # type: ignore[call-arg]
|
return EpisodeStatistics(
|
||||||
episode_returns=jnp.zeros(self.args.num_envs, dtype=jnp.float32),
|
episode_returns=jnp.zeros(self.args.num_envs, dtype=jnp.float32),
|
||||||
episode_lengths=jnp.zeros(self.args.num_envs, dtype=jnp.int32),
|
episode_lengths=jnp.zeros(self.args.num_envs, dtype=jnp.int32),
|
||||||
returned_episode_returns=jnp.zeros(self.args.num_envs, jnp.float32),
|
returned_episode_returns=jnp.zeros(self.args.num_envs, jnp.float32),
|
||||||
|
|
@ -363,21 +370,26 @@ class PPOTrainer:
|
||||||
episode_stats,
|
episode_stats,
|
||||||
start_time,
|
start_time,
|
||||||
iteration_time_start,
|
iteration_time_start,
|
||||||
loss_info,
|
training_measurements,
|
||||||
):
|
):
|
||||||
metrics = {
|
metrics = {
|
||||||
"charts/avg_episodic_return": loss_info.avg_episodic_return,
|
"charts/avg_episodic_return": training_measurements.avg_episodic_return,
|
||||||
"charts/avg_episodic_length": np.mean(
|
"charts/avg_episodic_length": np.mean(
|
||||||
jax.device_get(episode_stats.returned_episode_lengths)
|
jax.device_get(episode_stats.returned_episode_lengths)
|
||||||
),
|
),
|
||||||
"charts/learning_rate": self.agent_state.opt_state[1]
|
"charts/learning_rate": self.agent_state.opt_state[1]
|
||||||
.hyperparams["learning_rate"]
|
.hyperparams["learning_rate"]
|
||||||
.item(),
|
.item(),
|
||||||
"losses/value_loss": loss_info.v_loss[-1, -1].item(),
|
"charts/explained_variance": training_measurements.explained_variance,
|
||||||
"losses/policy_loss": loss_info.pg_loss[-1, -1].item(),
|
"charts/num_terminated": training_measurements.num_terminated,
|
||||||
"losses/entropy": loss_info.entropy_loss[-1, -1].item(),
|
"charts/num_truncated": training_measurements.num_truncated,
|
||||||
"losses/approx_kl": loss_info.approx_kl[-1, -1].item(),
|
"charts/avg_terminated_ep_length": training_measurements.avg_terminated_length,
|
||||||
"losses/loss": loss_info.loss[-1, -1].item(),
|
"charts/avg_truncated_ep_length": training_measurements.avg_truncated_length,
|
||||||
|
"losses/value_loss": training_measurements.v_loss[-1, -1].item(),
|
||||||
|
"losses/policy_loss": training_measurements.pg_loss[-1, -1].item(),
|
||||||
|
"losses/entropy": training_measurements.entropy_loss[-1, -1].item(),
|
||||||
|
"losses/approx_kl": training_measurements.approx_kl[-1, -1].item(),
|
||||||
|
"losses/loss": training_measurements.loss[-1, -1].item(),
|
||||||
"charts/SPS": int(global_step / (time.time() - start_time)),
|
"charts/SPS": int(global_step / (time.time() - start_time)),
|
||||||
"charts/SPS_update": int(
|
"charts/SPS_update": int(
|
||||||
self.args.num_envs * self.args.num_steps / (time.time() - iteration_time_start)
|
self.args.num_envs * self.args.num_steps / (time.time() - iteration_time_start)
|
||||||
|
|
@ -418,17 +430,39 @@ class PPOTrainer:
|
||||||
jnp.mean(jax.device_get(self.episode_stats.returned_episode_returns)).item()
|
jnp.mean(jax.device_get(self.episode_stats.returned_episode_returns)).item()
|
||||||
)
|
)
|
||||||
|
|
||||||
|
explained_var = _compute_explained_variance(storage.values, storage.returns)
|
||||||
|
|
||||||
|
terminated = next_env_state.terminated
|
||||||
|
truncated = next_env_state.truncated
|
||||||
|
episode_lengths = self.episode_stats.returned_episode_lengths
|
||||||
|
|
||||||
|
num_terminated = int(jnp.sum(terminated).item())
|
||||||
|
num_truncated = int(jnp.sum(truncated).item())
|
||||||
|
|
||||||
|
avg_terminated_length = jnp.sum(episode_lengths * terminated) / jnp.maximum(
|
||||||
|
jnp.sum(terminated), 1
|
||||||
|
)
|
||||||
|
|
||||||
|
avg_truncated_length = jnp.sum(episode_lengths * truncated) / jnp.maximum(
|
||||||
|
jnp.sum(truncated), 1
|
||||||
|
)
|
||||||
|
|
||||||
return (
|
return (
|
||||||
next_env_state,
|
next_env_state,
|
||||||
next_obs,
|
next_obs,
|
||||||
next_done,
|
next_done,
|
||||||
LossInfo(
|
TrainingMeasurements(
|
||||||
loss=loss,
|
loss=loss,
|
||||||
pg_loss=pg_loss,
|
pg_loss=pg_loss,
|
||||||
v_loss=v_loss,
|
v_loss=v_loss,
|
||||||
entropy_loss=entropy_loss,
|
entropy_loss=entropy_loss,
|
||||||
approx_kl=approx_kl,
|
approx_kl=approx_kl,
|
||||||
avg_episodic_return=avg_episodic_return,
|
avg_episodic_return=avg_episodic_return,
|
||||||
|
explained_variance=explained_var,
|
||||||
|
num_terminated=num_terminated,
|
||||||
|
num_truncated=num_truncated,
|
||||||
|
avg_terminated_length=avg_terminated_length,
|
||||||
|
avg_truncated_length=avg_truncated_length,
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
@ -473,12 +507,18 @@ class PPOTrainer:
|
||||||
for iteration in iter_bar:
|
for iteration in iter_bar:
|
||||||
iteration_time_start = time.time()
|
iteration_time_start = time.time()
|
||||||
|
|
||||||
env_state, next_obs, next_done, loss_info = self._step(
|
env_state, next_obs, next_done, training_measurements = self._step(
|
||||||
env_state, next_obs, next_done, iteration=iteration
|
env_state, next_obs, next_done, iteration=iteration
|
||||||
)
|
)
|
||||||
|
|
||||||
global_step += self.args.num_steps * self.args.num_envs
|
global_step += self.args.num_steps * self.args.num_envs
|
||||||
self._log(global_step, self.episode_stats, start_time, iteration_time_start, loss_info)
|
self._log(
|
||||||
|
global_step,
|
||||||
|
self.episode_stats,
|
||||||
|
start_time,
|
||||||
|
iteration_time_start,
|
||||||
|
training_measurements,
|
||||||
|
)
|
||||||
|
|
||||||
sps = int(global_step / (time.time() - start_time))
|
sps = int(global_step / (time.time() - start_time))
|
||||||
remaining_steps = self.args.total_timesteps - global_step
|
remaining_steps = self.args.total_timesteps - global_step
|
||||||
|
|
@ -489,7 +529,7 @@ class PPOTrainer:
|
||||||
f"Iteration {iteration}/{self.args.num_iterations} | "
|
f"Iteration {iteration}/{self.args.num_iterations} | "
|
||||||
f"Step {global_step}/{self.args.total_timesteps} | "
|
f"Step {global_step}/{self.args.total_timesteps} | "
|
||||||
f"SPS {sps} | "
|
f"SPS {sps} | "
|
||||||
f"Return {loss_info.avg_episodic_return:.4f} | "
|
f"Return {training_measurements.avg_episodic_return:.4f} | "
|
||||||
f"ETA {eta_str}"
|
f"ETA {eta_str}"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
|
||||||
Reference in a new issue