feat: More low lvl logs, reward, advantage, returns
This commit is contained in:
parent
aad086cb7d
commit
c0e7569773
4 changed files with 54 additions and 16 deletions
|
|
@ -1,10 +1,10 @@
|
|||
# Configuration for debug session
|
||||
exp_name: "debug-experiment-10042026" # started on april 10
|
||||
exp_name: "debug-experiment"
|
||||
seed: 42
|
||||
track: true
|
||||
wandb_project_name: "Let's-find-that-bug"
|
||||
wandb_entity: "SEL3-2026-Groep-4"
|
||||
|
||||
|
||||
num_envs: 32
|
||||
num_steps: 32
|
||||
total_timesteps: 102400
|
||||
|
|
|
|||
|
|
@ -50,7 +50,7 @@ if __name__ == "__main__":
|
|||
run_name = f"{args.exp_name}__seed_{args.seed}__{git_hash}__{int(time.time())}"
|
||||
|
||||
if args.run_dir is None:
|
||||
run_dir = f"runs/{run_name}"
|
||||
run_dir = f"/data/gent/465/vsc46589/runs/{run_name}"
|
||||
else:
|
||||
run_dir = args.run_dir
|
||||
|
||||
|
|
@ -69,9 +69,6 @@ if __name__ == "__main__":
|
|||
|
||||
env = make_env(args.env_config_path, args.num_envs)
|
||||
raw_env = env.raw
|
||||
logger.log(
|
||||
{"run_dir": run_dir}
|
||||
)
|
||||
|
||||
torch.backends.cudnn.deterministic = args.torch_deterministic
|
||||
|
||||
|
|
|
|||
|
|
@ -58,6 +58,10 @@ class Storage:
|
|||
returns: jnp.array
|
||||
rewards: jnp.array
|
||||
|
||||
raw_actions: jnp.ndarray = None # before clipping
|
||||
means: jnp.ndarray = None # policy mean
|
||||
stds: jnp.ndarray = None # policy std
|
||||
|
||||
def replace(self, **kwargs) -> "Storage":
|
||||
fs = fields(self)
|
||||
return Storage(**{f.name: kwargs.get(f.name, getattr(self, f.name)) for f in fs})
|
||||
|
|
|
|||
|
|
@ -11,8 +11,6 @@ import numpy as np
|
|||
import optax
|
||||
from flax.training.train_state import TrainState
|
||||
|
||||
from experiment_logger import get_logger
|
||||
|
||||
from brittle_star_project.dataclasses import EpisodeStatistics, PPOArgs
|
||||
from brittle_star_project.environment.BrittleStarJaxEnvWrapper import BrittleStarJaxEnvWrapper
|
||||
from brittle_star_project.MLPs.mlps import (
|
||||
|
|
@ -23,6 +21,7 @@ from brittle_star_project.MLPs.mlps import (
|
|||
Storage,
|
||||
)
|
||||
from brittle_star_project.ppo import PPO
|
||||
from experiment_logger import get_logger
|
||||
|
||||
@jax.jit
|
||||
def _clip_action(action: jnp.ndarray, low: jnp.ndarray, high: jnp.ndarray) -> jnp.ndarray:
|
||||
|
|
@ -67,15 +66,17 @@ def _get_action_and_value_noise(
|
|||
key, subkey = jax.random.split(key)
|
||||
noise = jax.random.normal(subkey, shape=mean.shape)
|
||||
std = jnp.exp(log_std)
|
||||
action = mean + noise * std
|
||||
raw_action = mean + noise * std
|
||||
clipped_action = _clip_action(
|
||||
action,
|
||||
raw_action,
|
||||
action_low,
|
||||
action_high
|
||||
)
|
||||
logprob = -0.5 * (((clipped_action - mean) / std) ** 2 + 2 * log_std + jnp.log(2 * jnp.pi)).sum(-1)
|
||||
logprob = -0.5 * (((raw_action - mean) / std) ** 2 + 2 * log_std + jnp.log(2 * jnp.pi)).sum(-1)
|
||||
value = critic.apply(agent_state.params["critic_params"], hidden_critic)
|
||||
return clipped_action, logprob, value.squeeze(-1), key
|
||||
|
||||
return clipped_action, raw_action, logprob, value.squeeze(-1), mean, std, key
|
||||
|
||||
|
||||
|
||||
def _step_once(
|
||||
|
|
@ -90,21 +91,24 @@ def _step_once(
|
|||
action_high
|
||||
):
|
||||
agent_state, episode_stats, obs, done, key, env_state = carry
|
||||
action, logprob, value, key = _get_action_and_value_noise(
|
||||
clipped_action, raw_action, logprob, value, mean, std, key = _get_action_and_value_noise(
|
||||
sensor, feature_extractor, actor, critic, agent_state, obs, key, action_low, action_high
|
||||
)
|
||||
|
||||
episode_stats, env_state, (next_obs, reward, next_done) = env_step_fn(
|
||||
episode_stats, env_state, action
|
||||
episode_stats, env_state, clipped_action
|
||||
)
|
||||
|
||||
storage = Storage(
|
||||
obs=obs,
|
||||
actions=action,
|
||||
actions=clipped_action,
|
||||
raw_actions=raw_action,
|
||||
logprobs=logprob,
|
||||
dones=done,
|
||||
values=value,
|
||||
rewards=reward,
|
||||
means=mean,
|
||||
stds=std,
|
||||
returns=jnp.zeros_like(reward),
|
||||
advantages=jnp.zeros_like(reward),
|
||||
)
|
||||
|
|
@ -371,7 +375,37 @@ class PPOTrainer:
|
|||
start_time,
|
||||
iteration_time_start,
|
||||
training_measurements,
|
||||
storage
|
||||
):
|
||||
data = jax.device_get({
|
||||
'rewards': storage.rewards[0], # (num_steps,)
|
||||
'values': storage.values[0],
|
||||
'returns': storage.returns[0], # NEW
|
||||
'advantages': storage.advantages[0],# NEW
|
||||
'actions': storage.actions[0],
|
||||
'raw_actions': storage.raw_actions[0],
|
||||
'means': storage.means[0],
|
||||
'stds': storage.stds[0],
|
||||
'logprobs': storage.logprobs[0],
|
||||
})
|
||||
|
||||
storage_metrics = {
|
||||
"rollout/env0/return_mean": float(np.mean(data['returns'])),
|
||||
|
||||
"rollout/env0/advantage_mean": float(np.mean(data['advantages'])),
|
||||
|
||||
"rollout/env0/value_mean": float(np.mean(data['values'])),
|
||||
"rollout/env0/value_vs_return_diff": float(np.mean(data['values'] - data['returns'])),
|
||||
|
||||
"rollout/env0/reward_mean": float(np.mean(data['rewards'])),
|
||||
|
||||
"rollout/env0/mean_mean": float(np.mean(data['means'])),
|
||||
"rollout/env0/logprob_mean": float(np.mean(data['logprobs'])),
|
||||
|
||||
"rollout/env0/action_mean": float(np.mean(data['actions'])),
|
||||
"rollout/env0/raw_action_mean": float(np.mean(data['raw_actions'])),
|
||||
}
|
||||
|
||||
metrics = {
|
||||
"charts/avg_episodic_return": training_measurements.avg_episodic_return,
|
||||
"charts/avg_episodic_length": np.mean(
|
||||
|
|
@ -394,6 +428,7 @@ class PPOTrainer:
|
|||
"charts/SPS_update": int(
|
||||
self.args.num_envs * self.args.num_steps / (time.time() - iteration_time_start)
|
||||
),
|
||||
**storage_metrics
|
||||
}
|
||||
self.logger.log(metrics, step=global_step)
|
||||
|
||||
|
|
@ -464,6 +499,7 @@ class PPOTrainer:
|
|||
avg_terminated_length=avg_terminated_length,
|
||||
avg_truncated_length=avg_truncated_length,
|
||||
),
|
||||
storage
|
||||
)
|
||||
|
||||
def _close(self):
|
||||
|
|
@ -507,7 +543,7 @@ class PPOTrainer:
|
|||
for iteration in iter_bar:
|
||||
iteration_time_start = time.time()
|
||||
|
||||
env_state, next_obs, next_done, training_measurements = self._step(
|
||||
env_state, next_obs, next_done, training_measurements, storage = self._step(
|
||||
env_state, next_obs, next_done, iteration=iteration
|
||||
)
|
||||
|
||||
|
|
@ -518,6 +554,7 @@ class PPOTrainer:
|
|||
start_time,
|
||||
iteration_time_start,
|
||||
training_measurements,
|
||||
storage
|
||||
)
|
||||
|
||||
sps = int(global_step / (time.time() - start_time))
|
||||
|
|
|
|||
Reference in a new issue