1
Fork 0

feat: More low lvl logs, reward, advantage, returns

This commit is contained in:
JibrilExe 2026-04-11 13:42:52 +02:00
parent aad086cb7d
commit c0e7569773
4 changed files with 54 additions and 16 deletions

View file

@ -1,5 +1,5 @@
# Configuration for debug session # Configuration for debug session
exp_name: "debug-experiment-10042026" # started on april 10 exp_name: "debug-experiment"
seed: 42 seed: 42
track: true track: true
wandb_project_name: "Let's-find-that-bug" wandb_project_name: "Let's-find-that-bug"

View file

@ -50,7 +50,7 @@ if __name__ == "__main__":
run_name = f"{args.exp_name}__seed_{args.seed}__{git_hash}__{int(time.time())}" run_name = f"{args.exp_name}__seed_{args.seed}__{git_hash}__{int(time.time())}"
if args.run_dir is None: if args.run_dir is None:
run_dir = f"runs/{run_name}" run_dir = f"/data/gent/465/vsc46589/runs/{run_name}"
else: else:
run_dir = args.run_dir run_dir = args.run_dir
@ -69,9 +69,6 @@ if __name__ == "__main__":
env = make_env(args.env_config_path, args.num_envs) env = make_env(args.env_config_path, args.num_envs)
raw_env = env.raw raw_env = env.raw
logger.log(
{"run_dir": run_dir}
)
torch.backends.cudnn.deterministic = args.torch_deterministic torch.backends.cudnn.deterministic = args.torch_deterministic

View file

@ -58,6 +58,10 @@ class Storage:
returns: jnp.array returns: jnp.array
rewards: 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": def replace(self, **kwargs) -> "Storage":
fs = fields(self) fs = fields(self)
return Storage(**{f.name: kwargs.get(f.name, getattr(self, f.name)) for f in fs}) return Storage(**{f.name: kwargs.get(f.name, getattr(self, f.name)) for f in fs})

View file

@ -11,8 +11,6 @@ import numpy as np
import optax import optax
from flax.training.train_state import TrainState 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.dataclasses import EpisodeStatistics, PPOArgs
from brittle_star_project.environment.BrittleStarJaxEnvWrapper import BrittleStarJaxEnvWrapper from brittle_star_project.environment.BrittleStarJaxEnvWrapper import BrittleStarJaxEnvWrapper
from brittle_star_project.MLPs.mlps import ( from brittle_star_project.MLPs.mlps import (
@ -23,6 +21,7 @@ from brittle_star_project.MLPs.mlps import (
Storage, Storage,
) )
from brittle_star_project.ppo import PPO from brittle_star_project.ppo import PPO
from experiment_logger import get_logger
@jax.jit @jax.jit
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:
@ -67,15 +66,17 @@ def _get_action_and_value_noise(
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)
std = jnp.exp(log_std) std = jnp.exp(log_std)
action = mean + noise * std raw_action = mean + noise * std
clipped_action = _clip_action( clipped_action = _clip_action(
action, raw_action,
action_low, action_low,
action_high 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) 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( def _step_once(
@ -90,21 +91,24 @@ def _step_once(
action_high action_high
): ):
agent_state, episode_stats, obs, done, key, env_state = carry 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 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, (next_obs, reward, next_done) = env_step_fn(
episode_stats, env_state, action episode_stats, env_state, clipped_action
) )
storage = Storage( storage = Storage(
obs=obs, obs=obs,
actions=action, actions=clipped_action,
raw_actions=raw_action,
logprobs=logprob, logprobs=logprob,
dones=done, dones=done,
values=value, values=value,
rewards=reward, rewards=reward,
means=mean,
stds=std,
returns=jnp.zeros_like(reward), returns=jnp.zeros_like(reward),
advantages=jnp.zeros_like(reward), advantages=jnp.zeros_like(reward),
) )
@ -371,7 +375,37 @@ class PPOTrainer:
start_time, start_time,
iteration_time_start, iteration_time_start,
training_measurements, 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 = { metrics = {
"charts/avg_episodic_return": training_measurements.avg_episodic_return, "charts/avg_episodic_return": training_measurements.avg_episodic_return,
"charts/avg_episodic_length": np.mean( "charts/avg_episodic_length": np.mean(
@ -394,6 +428,7 @@ class PPOTrainer:
"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)
), ),
**storage_metrics
} }
self.logger.log(metrics, step=global_step) self.logger.log(metrics, step=global_step)
@ -464,6 +499,7 @@ class PPOTrainer:
avg_terminated_length=avg_terminated_length, avg_terminated_length=avg_terminated_length,
avg_truncated_length=avg_truncated_length, avg_truncated_length=avg_truncated_length,
), ),
storage
) )
def _close(self): def _close(self):
@ -507,7 +543,7 @@ 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, training_measurements = self._step( env_state, next_obs, next_done, training_measurements, storage = self._step(
env_state, next_obs, next_done, iteration=iteration env_state, next_obs, next_done, iteration=iteration
) )
@ -518,6 +554,7 @@ class PPOTrainer:
start_time, start_time,
iteration_time_start, iteration_time_start,
training_measurements, training_measurements,
storage
) )
sps = int(global_step / (time.time() - start_time)) sps = int(global_step / (time.time() - start_time))