From ac352ef431d243f04aa4069ce8a030ea581e68a5 Mon Sep 17 00:00:00 2001 From: cedric Date: Fri, 10 Apr 2026 10:17:18 +0000 Subject: [PATCH 01/11] feat: start of debug setup, experiment description,.. --- configs/hpc/debug.yaml | 13 +++ .../used_variables.md | 82 +++++++++++++++++++ scripts/train.py | 14 +++- 3 files changed, 108 insertions(+), 1 deletion(-) create mode 100644 configs/hpc/debug.yaml create mode 100644 experiments/debug-experiment-10042026/used_variables.md diff --git a/configs/hpc/debug.yaml b/configs/hpc/debug.yaml new file mode 100644 index 0000000..d62df3e --- /dev/null +++ b/configs/hpc/debug.yaml @@ -0,0 +1,13 @@ +# Configuration for debug session +exp_name: "debug-experiment-10042026" # started on april 10 +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 + +cuda: true + diff --git a/experiments/debug-experiment-10042026/used_variables.md b/experiments/debug-experiment-10042026/used_variables.md new file mode 100644 index 0000000..ed10279 --- /dev/null +++ b/experiments/debug-experiment-10042026/used_variables.md @@ -0,0 +1,82 @@ +## Default envconfig +task: Task = Task.DIRECTED_LOCOMOTION +simulation_time: float = 5.0 +num_physics_steps_per_control_step: int = 10 +time_scale: int = 2 +camera_ids: list[int] = field(default_factory=lambda: [0, 1]) +render_size: tuple[int, int] = (480, 640) +joint_randomization_noise_scale: float = 0.0 +target_distance: float = 3.0 +light_perlin_noise_scale: int = 0 + + +## Default ppoargs +seed: int = 1 +torch_deterministic: bool = True +cuda: bool = True +track: bool = False +checkpoint_frequency: int = 100 +learning_rate: float = 2.5e-4 +anneal_lr: bool = True +gamma: float = 0.99 +gae_lambda: float = 0.95 +num_minibatches: int = 4 +update_epochs: int = 4 +norm_adv: bool = True +clip_coef: float = 0.1 +clip_vloss: bool = True +ent_coef: float = 0.01 +vf_coef: float = 0.5 +max_grad_norm: float = 0.5 +target_kl: float | None = None +batch_size: int = 0 +minibatch_size: int = 0 +num_iterations: int = 0 + +## Used config file: +num_envs: 32 +num_steps: 32 +total_timesteps: 102400 + +## Arena config: +size: tuple[float, float] = (10.0, 5.0) +sand_ground_color: bool = True +attach_target: bool = True +wall_height: float = 1.5 +wall_thickness: float = 0.1 + +## Morphology: +num_arms: int = 5 +num_segments_per_arm: int = 4 +use_p_control: bool = True +use_torque_control: bool = False + +## MLPs: +### Sensor & Feature_extractor: +class GenericDenseLayersWithActivation(nn.Module): + layer_sizes: Sequence[int] = field(default_factory=lambda: [64, 64]) + activation: Callable = nn.tanh + + @nn.compact + def __call__(self, x): + for size in self.layer_sizes: + x = nn.Dense(size, kernel_init=orthogonal(jnp.sqrt(2)))(x) + x = self.activation(x) + return x + +### Actor: +class Actor(nn.Module): + action_dim: int + @nn.compact + def __call__(self, x): + mean = nn.Dense(self.action_dim, kernel_init=orthogonal(0.01), bias_init=constant(0.0))(x) + log_std = self.param("log_std", nn.initializers.zeros, (self.action_dim,)) + return mean, log_std + +### Critic: +class OneDenseLayerMLP(nn.Module): + @nn.compact + def __call__(self, x): + return nn.Dense(1, kernel_init=orthogonal(1), bias_init=constant(0.0))(x) + +### Observations: diff --git a/scripts/train.py b/scripts/train.py index 492c887..5c646bf 100644 --- a/scripts/train.py +++ b/scripts/train.py @@ -68,7 +68,19 @@ if __name__ == "__main__": print_config(args, title="PPO Training Configuration") env = make_env(args.env_config_path, args.num_envs) - + raw_env = env.raw + print( + "\n\n\n Observation space \n", + raw_env.observation_space, + "Action space \n", + raw_env.action_space, + ) + print( + "\n\n\n Observation space \n", + raw_env.observation_space, + "Action space \n", + raw_env.action_space, + ) torch.backends.cudnn.deterministic = args.torch_deterministic ppo_trainer = PPOTrainer(args, env, run_dir, run_name) From 1ef257544a0e08942500690420dfd31bb706ab7e Mon Sep 17 00:00:00 2001 From: JibrilExe Date: Fri, 10 Apr 2026 15:18:33 +0200 Subject: [PATCH 02/11] feat: first approach to action clipping --- configs/hpc/debug.yaml | 3 ++- .../used_variables.md | 2 +- scripts/hpc/train.pbs | 4 +-- scripts/train.py | 14 +++------- .../trainers/PPOTrainer.py | 27 ++++++++++++++++--- 5 files changed, 32 insertions(+), 18 deletions(-) diff --git a/configs/hpc/debug.yaml b/configs/hpc/debug.yaml index d62df3e..20c65e8 100644 --- a/configs/hpc/debug.yaml +++ b/configs/hpc/debug.yaml @@ -8,6 +8,7 @@ wandb_entity: "SEL3-2026-Groep-4" num_envs: 32 num_steps: 32 total_timesteps: 102400 - +num_arms: 2 +num_segments_per_arm: 1 cuda: true diff --git a/experiments/debug-experiment-10042026/used_variables.md b/experiments/debug-experiment-10042026/used_variables.md index ed10279..0dcba35 100644 --- a/experiments/debug-experiment-10042026/used_variables.md +++ b/experiments/debug-experiment-10042026/used_variables.md @@ -46,7 +46,7 @@ wall_height: float = 1.5 wall_thickness: float = 0.1 ## Morphology: -num_arms: int = 5 +num_arms: int = 2 num_segments_per_arm: int = 4 use_p_control: bool = True use_torque_control: bool = False diff --git a/scripts/hpc/train.pbs b/scripts/hpc/train.pbs index e7d21f0..e4ce355 100644 --- a/scripts/hpc/train.pbs +++ b/scripts/hpc/train.pbs @@ -65,8 +65,8 @@ fi # TODO Once experiments get serious, change the config python scripts/train.py \ - --env-config-path configs/hpc/smoke_test.yaml \ - --hyperparameter-config-path configs/hpc/smoke_test.yaml \ + --env-config-path configs/hpc/debug.yaml \ + --hyperparameter-config-path configs/hpc/debug.yaml \ --run-dir "$SCRATCH_RUNDIR" echo ">>> Staging out results to $DATA_RUNDIR..." diff --git a/scripts/train.py b/scripts/train.py index 5c646bf..1f543f7 100644 --- a/scripts/train.py +++ b/scripts/train.py @@ -69,18 +69,10 @@ if __name__ == "__main__": env = make_env(args.env_config_path, args.num_envs) raw_env = env.raw - print( - "\n\n\n Observation space \n", - raw_env.observation_space, - "Action space \n", - raw_env.action_space, - ) - print( - "\n\n\n Observation space \n", - raw_env.observation_space, - "Action space \n", - raw_env.action_space, + logger.log( + {"run_dir": run_dir} ) + torch.backends.cudnn.deterministic = args.torch_deterministic ppo_trainer = PPOTrainer(args, env, run_dir, run_name) diff --git a/src/brittle_star_project/trainers/PPOTrainer.py b/src/brittle_star_project/trainers/PPOTrainer.py index 8c238f1..432030a 100644 --- a/src/brittle_star_project/trainers/PPOTrainer.py +++ b/src/brittle_star_project/trainers/PPOTrainer.py @@ -24,6 +24,9 @@ from brittle_star_project.MLPs.mlps import ( ) from brittle_star_project.ppo import PPO +@jax.jit +def _clip_action(action: jnp.ndarray, low: jnp.ndarray, high: jnp.ndarray) -> jnp.ndarray: + return jnp.clip(action, low, high) @jax.jit def _linear_schedule(count, minibatch_count, update_epochs, num_iterations, learning_rate): @@ -47,6 +50,8 @@ def _get_action_and_value_noise( agent_state: TrainState, next_obs: jnp.ndarray, key: jax.random.PRNGKey, + action_low, + action_high ): hidden = sensor.apply(agent_state.params["sensor_params"], next_obs) hidden_critic = feature_extractor.apply( @@ -59,9 +64,14 @@ def _get_action_and_value_noise( noise = jax.random.normal(subkey, shape=mean.shape) std = jnp.exp(log_std) action = mean + noise * std - logprob = -0.5 * (((action - mean) / std) ** 2 + 2 * log_std + jnp.log(2 * jnp.pi)).sum(-1) + clipped_action = _clip_action( + action, + action_low, + action_high + ) + logprob = -0.5 * (((clipped_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 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 @@ -73,10 +83,12 @@ def _step_once( feature_extractor: GenericDenseLayersWithActivation, actor: Actor, critic: OneDenseLayerMLP, + action_low, + action_high ): agent_state, episode_stats, obs, done, key, env_state = carry action, logprob, value, key = _get_action_and_value_noise( - sensor, feature_extractor, actor, critic, agent_state, obs, key + 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( @@ -140,6 +152,8 @@ def _rollout_jit( feature_extractor: GenericDenseLayersWithActivation, actor: Actor, critic: OneDenseLayerMLP, + action_low, + action_high ): (agent_state, episode_stats, next_obs, next_done, key, env_state), storage = jax.lax.scan( partial( @@ -149,6 +163,8 @@ def _rollout_jit( actor=actor, critic=critic, env_step_fn=step_env_fn, + action_low=action_low, + action_high=action_high ), (agent_state, episode_stats, next_obs, next_done, key, env_state), (), @@ -215,6 +231,9 @@ class PPOTrainer: self.actor.apply = jax.jit(self.actor.apply) self.critic.apply = jax.jit(self.critic.apply) + action_low = jnp.asarray(self.env.single_action_space.low, dtype=jnp.float32) + action_high = jnp.asarray(self.env.single_action_space.high, dtype=jnp.float32) + self._rollout_jit = jax.jit( partial( _rollout_jit, @@ -224,6 +243,8 @@ class PPOTrainer: feature_extractor=self.feature_extractor, actor=self.actor, critic=self.critic, + action_low=action_low, + action_high=action_high ) ) self._compute_gae_jit = jax.jit( From c0e7569773c17c1ded9e70a7a2ec18b1cebad3e6 Mon Sep 17 00:00:00 2001 From: JibrilExe Date: Sat, 11 Apr 2026 13:42:52 +0200 Subject: [PATCH 03/11] feat: More low lvl logs, reward, advantage, returns --- configs/hpc/debug.yaml | 4 +- scripts/train.py | 5 +- src/brittle_star_project/MLPs/mlps.py | 4 ++ .../trainers/PPOTrainer.py | 57 +++++++++++++++---- 4 files changed, 54 insertions(+), 16 deletions(-) diff --git a/configs/hpc/debug.yaml b/configs/hpc/debug.yaml index 20c65e8..c387dc1 100644 --- a/configs/hpc/debug.yaml +++ b/configs/hpc/debug.yaml @@ -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 diff --git a/scripts/train.py b/scripts/train.py index 1f543f7..886bd54 100644 --- a/scripts/train.py +++ b/scripts/train.py @@ -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 diff --git a/src/brittle_star_project/MLPs/mlps.py b/src/brittle_star_project/MLPs/mlps.py index 6abb540..9568a36 100644 --- a/src/brittle_star_project/MLPs/mlps.py +++ b/src/brittle_star_project/MLPs/mlps.py @@ -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}) diff --git a/src/brittle_star_project/trainers/PPOTrainer.py b/src/brittle_star_project/trainers/PPOTrainer.py index 1e554f8..8a9b341 100644 --- a/src/brittle_star_project/trainers/PPOTrainer.py +++ b/src/brittle_star_project/trainers/PPOTrainer.py @@ -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)) From 1c328bdde1fc56e87db0ab5547a335dc89a9a927 Mon Sep 17 00:00:00 2001 From: JibrilExe Date: Sun, 12 Apr 2026 11:12:46 +0200 Subject: [PATCH 04/11] fix: observation filtering, action clipping, log_std clip, reward discount, bigger MLP --- scripts/train.py | 2 +- src/brittle_star_project/ppo.py | 1 + .../trainers/PPOTrainer.py | 53 +++++++++++++------ 3 files changed, 39 insertions(+), 17 deletions(-) diff --git a/scripts/train.py b/scripts/train.py index 886bd54..e4b5786 100644 --- a/scripts/train.py +++ b/scripts/train.py @@ -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"/data/gent/465/vsc46589/runs/{run_name}" + run_dir = f"runs/{run_name}" else: run_dir = args.run_dir diff --git a/src/brittle_star_project/ppo.py b/src/brittle_star_project/ppo.py index cf6c69e..60d9671 100644 --- a/src/brittle_star_project/ppo.py +++ b/src/brittle_star_project/ppo.py @@ -99,6 +99,7 @@ def get_action_and_value( hidden_critic = feature_extractor_apply(params["feature_extractor_params"], x) hidden_sensor = message_passer(hidden_sensor) mean, log_std = actor_apply(params["actor_params"], hidden_sensor) + log_std = jnp.clip(log_std, -5, 2) std = jnp.exp(log_std) logprob = -0.5 * (((action - mean) / std) ** 2 + 2 * log_std + jnp.log(2 * jnp.pi)).sum(-1) diff --git a/src/brittle_star_project/trainers/PPOTrainer.py b/src/brittle_star_project/trainers/PPOTrainer.py index 8a9b341..2dc02f9 100644 --- a/src/brittle_star_project/trainers/PPOTrainer.py +++ b/src/brittle_star_project/trainers/PPOTrainer.py @@ -23,6 +23,21 @@ from brittle_star_project.MLPs.mlps import ( from brittle_star_project.ppo import PPO from experiment_logger import get_logger +# TODO: move to config +_ALLOWED_OBS_KEYS = { + "joint_position", + "joint_velocity", + "joint_actuator_force", + "actuator_force", + "disk_position", + "disk_rotation", + "disk_linear_velocity", + "disk_angular_velocity", + "unit_xy_direction_to_target", + "xy_distance_to_target", +} +# TODO: clip scaled reward? + @jax.jit def _clip_action(action: jnp.ndarray, low: jnp.ndarray, high: jnp.ndarray) -> jnp.ndarray: return jnp.clip(action, low, high) @@ -41,9 +56,17 @@ def _linear_schedule(count, minibatch_count, update_epochs, num_iterations, lear @jax.jit def _convert_obs_dict_to_array(obs_dict: dict) -> jnp.ndarray: - return jax.vmap(lambda o: jnp.concatenate([v.flatten() for v in o.values() if v.size > 0]))( - obs_dict - ) + """Convert the raw observation dict → flat array, filtering unwanted keys.""" + def _filter_and_flatten(o: dict) -> jnp.ndarray: + values = [] + for key in sorted(o.keys()): + if key in _ALLOWED_OBS_KEYS: #TODO: NORMALIZATION or .. of observations?? + v = o[key] + if v.size > 0: + values.append(jnp.asarray(v).flatten()) + return jnp.concatenate(values) + + return jax.vmap(_filter_and_flatten)(obs_dict) def _get_action_and_value_noise( @@ -63,6 +86,7 @@ def _get_action_and_value_noise( ) mean, log_std = actor.apply(agent_state.params["actor_params"], hidden) + log_std = jnp.clip(log_std, -5, 2) key, subkey = jax.random.split(key) noise = jax.random.normal(subkey, shape=mean.shape) std = jnp.exp(log_std) @@ -75,7 +99,7 @@ def _get_action_and_value_noise( 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, raw_action, logprob, value.squeeze(-1), mean, std, key + return raw_action, raw_action, logprob, value.squeeze(-1), mean, std, key @@ -119,6 +143,7 @@ def _step_env_wrapped(episode_stats, env_state, action, env_step_fn): next_env_state = env_step_fn(env_state, action) reward = next_env_state.reward + reward *= 100 terminated = next_env_state.terminated truncated = next_env_state.truncated done = terminated | truncated @@ -211,7 +236,9 @@ def _compute_gae_jit( (dones[1:], values[1:], values[:-1], storage.rewards), reverse=True, ) - return storage.replace(advantages=advantages, returns=advantages + storage.values) + returns = advantages + storage.values + advantages = (advantages - advantages.mean()) / (advantages.std() + 1e-8) + return storage.replace(advantages=advantages, returns=returns) @dataclass @@ -289,8 +316,8 @@ class PPOTrainer: def _init_agent(self): self.logger.info("[AGENT]: Initializing agent...") - sensor = GenericDenseLayersWithActivation() - feature_extractor = GenericDenseLayersWithActivation() + sensor = GenericDenseLayersWithActivation(layer_sizes=[300, 300, 300]) + feature_extractor = GenericDenseLayersWithActivation(layer_sizes=[300, 300, 300]) actor = Actor(action_dim=self.env.single_action_space.shape[0]) critic = OneDenseLayerMLP() return sensor, feature_extractor, actor, critic @@ -302,15 +329,9 @@ class PPOTrainer: self.key, 5 ) - sample_obs = jnp.concatenate( - [ - v.flatten() - for v in self.env.single_observation_space.sample( - rng=jax.random.PRNGKey(0) - ).values() - if v.size > 0 - ] - ) + dummy_reset = self.env.reset(seed=0) + sample_obs = _convert_obs_dict_to_array(dummy_reset.observations)[0] # take first env + sensor_params = self.sensor.init(sensor_key, sample_obs) feature_extractor_params = self.feature_extractor.init(feature_extractor_key, sample_obs) actor_params = self.actor.init(actor_key, self.sensor.apply(sensor_params, sample_obs)) From c5a8dd2b3aeabd86cd214b0e086958fb50ea72b0 Mon Sep 17 00:00:00 2001 From: JibrilExe Date: Sun, 12 Apr 2026 12:54:48 +0200 Subject: [PATCH 05/11] feat: expanded used var doc, and added distance to target log --- configs/hpc/debug.yaml | 2 +- .../used_variables.md | 16 +++++++++++++- .../environment/env_config.py | 2 +- .../trainers/PPOTrainer.py | 21 ++++++++++++++++--- 4 files changed, 35 insertions(+), 6 deletions(-) diff --git a/configs/hpc/debug.yaml b/configs/hpc/debug.yaml index c387dc1..e2a9f9d 100644 --- a/configs/hpc/debug.yaml +++ b/configs/hpc/debug.yaml @@ -4,7 +4,7 @@ seed: 42 track: true wandb_project_name: "Let's-find-that-bug" wandb_entity: "SEL3-2026-Groep-4" - +run_dir: "/data/gent/465/vsc46589" num_envs: 32 num_steps: 32 total_timesteps: 102400 diff --git a/experiments/debug-experiment-10042026/used_variables.md b/experiments/debug-experiment-10042026/used_variables.md index 0dcba35..e49c9e4 100644 --- a/experiments/debug-experiment-10042026/used_variables.md +++ b/experiments/debug-experiment-10042026/used_variables.md @@ -1,6 +1,6 @@ ## Default envconfig task: Task = Task.DIRECTED_LOCOMOTION -simulation_time: float = 5.0 +simulation_time: float = 500.0 num_physics_steps_per_control_step: int = 10 time_scale: int = 2 camera_ids: list[int] = field(default_factory=lambda: [0, 1]) @@ -53,6 +53,8 @@ use_torque_control: bool = False ## MLPs: ### Sensor & Feature_extractor: +Both with 3 layers of 300 neurons per layer. + class GenericDenseLayersWithActivation(nn.Module): layer_sizes: Sequence[int] = field(default_factory=lambda: [64, 64]) activation: Callable = nn.tanh @@ -80,3 +82,15 @@ class OneDenseLayerMLP(nn.Module): return nn.Dense(1, kernel_init=orthogonal(1), bias_init=constant(0.0))(x) ### Observations: +_ALLOWED_OBS_KEYS = { + "joint_position", + "joint_velocity", + "joint_actuator_force", + "actuator_force", + "disk_position", + "disk_rotation", + "disk_linear_velocity", + "disk_angular_velocity", + "unit_xy_direction_to_target", + "xy_distance_to_target", +} \ No newline at end of file diff --git a/src/brittle_star_project/environment/env_config.py b/src/brittle_star_project/environment/env_config.py index 78083e9..c2e8067 100644 --- a/src/brittle_star_project/environment/env_config.py +++ b/src/brittle_star_project/environment/env_config.py @@ -31,7 +31,7 @@ class EnvConfig: task: Task = Task.DIRECTED_LOCOMOTION - simulation_time: float = 5.0 + simulation_time: float = 500.0 num_physics_steps_per_control_step: int = 10 time_scale: int = 2 diff --git a/src/brittle_star_project/trainers/PPOTrainer.py b/src/brittle_star_project/trainers/PPOTrainer.py index 2dc02f9..6997911 100644 --- a/src/brittle_star_project/trainers/PPOTrainer.py +++ b/src/brittle_star_project/trainers/PPOTrainer.py @@ -38,6 +38,12 @@ _ALLOWED_OBS_KEYS = { } # TODO: clip scaled reward? +@jax.jit +def _get_xy_distance_to_target(obs_dict: dict) -> jnp.ndarray: + """Extract xy_distance_to_target for all environments.""" + # obs_dict is a dict of arrays with leading batch dimension (num_envs, ...) + return obs_dict["xy_distance_to_target"].squeeze(-1) # shape: (num_envs,) + @jax.jit def _clip_action(action: jnp.ndarray, low: jnp.ndarray, high: jnp.ndarray) -> jnp.ndarray: return jnp.clip(action, low, high) @@ -143,7 +149,7 @@ def _step_env_wrapped(episode_stats, env_state, action, env_step_fn): next_env_state = env_step_fn(env_state, action) reward = next_env_state.reward - reward *= 100 + reward *= 1000 terminated = next_env_state.terminated truncated = next_env_state.truncated done = terminated | truncated @@ -396,7 +402,9 @@ class PPOTrainer: start_time, iteration_time_start, training_measurements, - storage + storage, + next_obs, + xy_distance ): data = jax.device_get({ 'rewards': storage.rewards[0], # (num_steps,) @@ -425,6 +433,9 @@ class PPOTrainer: "rollout/env0/action_mean": float(np.mean(data['actions'])), "rollout/env0/raw_action_mean": float(np.mean(data['raw_actions'])), + + "charts/env0_xy_distance_to_target": float(xy_distance[0]), + "charts/env1_xy_distance_to_target": float(xy_distance[1]), } metrics = { @@ -568,6 +579,8 @@ class PPOTrainer: env_state, next_obs, next_done, iteration=iteration ) + xy_distance = _get_xy_distance_to_target(env_state.observations) + global_step += self.args.num_steps * self.args.num_envs self._log( global_step, @@ -575,7 +588,9 @@ class PPOTrainer: start_time, iteration_time_start, training_measurements, - storage + storage, + next_obs, + xy_distance ) sps = int(global_step / (time.time() - start_time)) From 631d01ebf660156881f7d158c0e71be122e23f8d Mon Sep 17 00:00:00 2001 From: JibrilExe Date: Sun, 12 Apr 2026 13:09:53 +0200 Subject: [PATCH 06/11] feat: log distance to target for each env --- configs/hpc/debug.yaml | 3 +-- src/brittle_star_project/trainers/PPOTrainer.py | 12 ++++++------ 2 files changed, 7 insertions(+), 8 deletions(-) diff --git a/configs/hpc/debug.yaml b/configs/hpc/debug.yaml index e2a9f9d..9fb9170 100644 --- a/configs/hpc/debug.yaml +++ b/configs/hpc/debug.yaml @@ -7,8 +7,7 @@ wandb_entity: "SEL3-2026-Groep-4" run_dir: "/data/gent/465/vsc46589" num_envs: 32 num_steps: 32 -total_timesteps: 102400 +total_timesteps: 409600 num_arms: 2 -num_segments_per_arm: 1 cuda: true diff --git a/src/brittle_star_project/trainers/PPOTrainer.py b/src/brittle_star_project/trainers/PPOTrainer.py index 6997911..c924519 100644 --- a/src/brittle_star_project/trainers/PPOTrainer.py +++ b/src/brittle_star_project/trainers/PPOTrainer.py @@ -407,10 +407,10 @@ class PPOTrainer: xy_distance ): data = jax.device_get({ - 'rewards': storage.rewards[0], # (num_steps,) + 'rewards': storage.rewards[0], 'values': storage.values[0], - 'returns': storage.returns[0], # NEW - 'advantages': storage.advantages[0],# NEW + 'returns': storage.returns[0], + 'advantages': storage.advantages[0], 'actions': storage.actions[0], 'raw_actions': storage.raw_actions[0], 'means': storage.means[0], @@ -433,11 +433,11 @@ class PPOTrainer: "rollout/env0/action_mean": float(np.mean(data['actions'])), "rollout/env0/raw_action_mean": float(np.mean(data['raw_actions'])), - - "charts/env0_xy_distance_to_target": float(xy_distance[0]), - "charts/env1_xy_distance_to_target": float(xy_distance[1]), } + for i in range(len(xy_distance)): + storage_metrics[f"env_data/env{i}_xy_dist_target"] = float(xy_distance[i]) + metrics = { "charts/avg_episodic_return": training_measurements.avg_episodic_return, "charts/avg_episodic_length": np.mean( From 1f8fdbdc592609ac52dc40f499b8878c9298141e Mon Sep 17 00:00:00 2001 From: JibrilExe Date: Tue, 14 Apr 2026 19:07:15 +0200 Subject: [PATCH 07/11] feat: observation normalization --- .../trainers/PPOTrainer.py | 30 +++++++++++++++++-- 1 file changed, 28 insertions(+), 2 deletions(-) diff --git a/src/brittle_star_project/trainers/PPOTrainer.py b/src/brittle_star_project/trainers/PPOTrainer.py index c924519..ec19e26 100644 --- a/src/brittle_star_project/trainers/PPOTrainer.py +++ b/src/brittle_star_project/trainers/PPOTrainer.py @@ -59,6 +59,9 @@ def _linear_schedule(count, minibatch_count, update_epochs, num_iterations, lear frac = 1.0 - (count // (minibatch_count * update_epochs)) / num_iterations return learning_rate * frac +@jax.jit +def _normalize_obs(obs, mean, var, eps=1e-8): + return jnp.clip((obs - mean) / jnp.sqrt(var + eps), -10.0, 10.0) @jax.jit def _convert_obs_dict_to_array(obs_dict: dict) -> jnp.ndarray: @@ -105,7 +108,7 @@ def _get_action_and_value_noise( 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 raw_action, raw_action, logprob, value.squeeze(-1), mean, std, key + return clipped_action, raw_action, logprob, value.squeeze(-1), mean, std, key @@ -337,7 +340,9 @@ class PPOTrainer: dummy_reset = self.env.reset(seed=0) sample_obs = _convert_obs_dict_to_array(dummy_reset.observations)[0] # take first env - + self.obs_mean = jnp.zeros((len(sample_obs),)) + self.obs_var = jnp.ones((len(sample_obs),)) + self.obs_count = 1e-4 sensor_params = self.sensor.init(sensor_key, sample_obs) feature_extractor_params = self.feature_extractor.init(feature_extractor_key, sample_obs) actor_params = self.actor.init(actor_key, self.sensor.apply(sensor_params, sample_obs)) @@ -376,6 +381,25 @@ class PPOTrainer: returned_episode_returns=jnp.zeros(self.args.num_envs, jnp.float32), returned_episode_lengths=jnp.zeros(self.args.num_envs, dtype=jnp.int32), ) + + def _update_obs_stats(self, obs: jnp.ndarray): + batch_mean = jnp.mean(obs, axis=0) + batch_var = jnp.var(obs, axis=0) + batch_count = obs.shape[0] + + delta = batch_mean - self.obs_mean + total_count = self.obs_count + batch_count + + new_mean = self.obs_mean + delta * batch_count / total_count + + m_a = self.obs_var * self.obs_count + m_b = batch_var * batch_count + M2 = m_a + m_b + delta**2 * self.obs_count * batch_count / total_count + new_var = M2 / total_count + + self.obs_mean = new_mean + self.obs_var = new_var + self.obs_count = total_count def _rollout(self, env_state, next_obs, next_done) -> tuple[Any, ...]: return self._rollout_jit( @@ -578,6 +602,8 @@ class PPOTrainer: env_state, next_obs, next_done, training_measurements, storage = self._step( env_state, next_obs, next_done, iteration=iteration ) + self._update_obs_stats(next_obs) + next_obs = _normalize_obs(next_obs, self.obs_mean, self.obs_var) xy_distance = _get_xy_distance_to_target(env_state.observations) From 7a87facc8aaa43936ddeb97b4b1dff9fe45c3841 Mon Sep 17 00:00:00 2001 From: JibrilExe Date: Tue, 14 Apr 2026 19:23:20 +0200 Subject: [PATCH 08/11] feat: value loss clipping --- src/brittle_star_project/ppo.py | 14 +++++++++++++- 1 file changed, 13 insertions(+), 1 deletion(-) diff --git a/src/brittle_star_project/ppo.py b/src/brittle_star_project/ppo.py index 60d9671..ffe2064 100644 --- a/src/brittle_star_project/ppo.py +++ b/src/brittle_star_project/ppo.py @@ -143,7 +143,19 @@ def ppo_loss( pg_loss1 = -mb_advantages * ratio pg_loss2 = -mb_advantages * jnp.clip(ratio, 1 - args.clip_coef, 1 + args.clip_coef) pg_loss = jnp.maximum(pg_loss1, pg_loss2).mean() - v_loss = 0.5 * ((newvalue - mb_returns) ** 2).mean() + old_value = mb_returns - mb_advantages + + v_clipped = old_value + jnp.clip( + newvalue - old_value, + -args.clip_coef, + args.clip_coef + ) + + v_loss_unclipped = (newvalue - mb_returns) ** 2 + v_loss_clipped = (v_clipped - mb_returns) ** 2 + + v_loss = 0.5 * jnp.maximum(v_loss_unclipped, v_loss_clipped).mean() + entropy_loss = entropy.mean() loss = pg_loss - args.ent_coef * entropy_loss + v_loss * args.vf_coef return loss, (pg_loss, v_loss, entropy_loss, jax.lax.stop_gradient(approx_kl)) From 0a9bf2e0a55f8d02c35e451cc2812a002c5e6d88 Mon Sep 17 00:00:00 2001 From: JibrilExe Date: Tue, 14 Apr 2026 20:30:21 +0200 Subject: [PATCH 09/11] fix: dont store clipped action, and upgraded reward scale --- src/brittle_star_project/ppo.py | 15 ++------------- src/brittle_star_project/trainers/PPOTrainer.py | 5 +++-- 2 files changed, 5 insertions(+), 15 deletions(-) diff --git a/src/brittle_star_project/ppo.py b/src/brittle_star_project/ppo.py index ffe2064..0156b82 100644 --- a/src/brittle_star_project/ppo.py +++ b/src/brittle_star_project/ppo.py @@ -143,19 +143,8 @@ def ppo_loss( pg_loss1 = -mb_advantages * ratio pg_loss2 = -mb_advantages * jnp.clip(ratio, 1 - args.clip_coef, 1 + args.clip_coef) pg_loss = jnp.maximum(pg_loss1, pg_loss2).mean() - old_value = mb_returns - mb_advantages - - v_clipped = old_value + jnp.clip( - newvalue - old_value, - -args.clip_coef, - args.clip_coef - ) - - v_loss_unclipped = (newvalue - mb_returns) ** 2 - v_loss_clipped = (v_clipped - mb_returns) ** 2 - - v_loss = 0.5 * jnp.maximum(v_loss_unclipped, v_loss_clipped).mean() - + + v_loss = 0.5 * ((newvalue - mb_returns) ** 2).mean() entropy_loss = entropy.mean() loss = pg_loss - args.ent_coef * entropy_loss + v_loss * args.vf_coef return loss, (pg_loss, v_loss, entropy_loss, jax.lax.stop_gradient(approx_kl)) diff --git a/src/brittle_star_project/trainers/PPOTrainer.py b/src/brittle_star_project/trainers/PPOTrainer.py index ec19e26..3dcf6db 100644 --- a/src/brittle_star_project/trainers/PPOTrainer.py +++ b/src/brittle_star_project/trainers/PPOTrainer.py @@ -134,7 +134,7 @@ def _step_once( storage = Storage( obs=obs, - actions=clipped_action, + actions=raw_action, raw_actions=raw_action, logprobs=logprob, dones=done, @@ -152,7 +152,8 @@ def _step_env_wrapped(episode_stats, env_state, action, env_step_fn): next_env_state = env_step_fn(env_state, action) reward = next_env_state.reward - reward *= 1000 + reward *= 5000 + reward = jnp.clip(reward, -1, 1) terminated = next_env_state.terminated truncated = next_env_state.truncated done = terminated | truncated From 20ba63a303c771d4bface9b35bb6217f017a719c Mon Sep 17 00:00:00 2001 From: JibrilExe Date: Tue, 14 Apr 2026 21:14:26 +0200 Subject: [PATCH 10/11] fix: ruff format --- configs/hpc/debug.yaml | 7 ++ src/brittle_star_project/MLPs/mlps.py | 6 +- .../environment/env_config.py | 2 +- src/brittle_star_project/ppo.py | 2 +- .../trainers/PPOTrainer.py | 92 +++++++++---------- 5 files changed, 57 insertions(+), 52 deletions(-) diff --git a/configs/hpc/debug.yaml b/configs/hpc/debug.yaml index 9fb9170..a0f5640 100644 --- a/configs/hpc/debug.yaml +++ b/configs/hpc/debug.yaml @@ -7,7 +7,14 @@ wandb_entity: "SEL3-2026-Groep-4" run_dir: "/data/gent/465/vsc46589" num_envs: 32 num_steps: 32 +num_minibatches: 32 total_timesteps: 409600 num_arms: 2 cuda: true +ent_coef: 0.005 +vf_coef: 1.0 +clip_coef: 0.2 + +anneal_lr: true +learning_rate: 0.0003 \ No newline at end of file diff --git a/src/brittle_star_project/MLPs/mlps.py b/src/brittle_star_project/MLPs/mlps.py index 9568a36..5e2deb5 100644 --- a/src/brittle_star_project/MLPs/mlps.py +++ b/src/brittle_star_project/MLPs/mlps.py @@ -58,9 +58,9 @@ 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 + 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) diff --git a/src/brittle_star_project/environment/env_config.py b/src/brittle_star_project/environment/env_config.py index c2e8067..65ec229 100644 --- a/src/brittle_star_project/environment/env_config.py +++ b/src/brittle_star_project/environment/env_config.py @@ -31,7 +31,7 @@ class EnvConfig: task: Task = Task.DIRECTED_LOCOMOTION - simulation_time: float = 500.0 + simulation_time: float = 10000.0 num_physics_steps_per_control_step: int = 10 time_scale: int = 2 diff --git a/src/brittle_star_project/ppo.py b/src/brittle_star_project/ppo.py index 0156b82..1ea37a8 100644 --- a/src/brittle_star_project/ppo.py +++ b/src/brittle_star_project/ppo.py @@ -143,7 +143,7 @@ def ppo_loss( pg_loss1 = -mb_advantages * ratio pg_loss2 = -mb_advantages * jnp.clip(ratio, 1 - args.clip_coef, 1 + args.clip_coef) pg_loss = jnp.maximum(pg_loss1, pg_loss2).mean() - + v_loss = 0.5 * ((newvalue - mb_returns) ** 2).mean() entropy_loss = entropy.mean() loss = pg_loss - args.ent_coef * entropy_loss + v_loss * args.vf_coef diff --git a/src/brittle_star_project/trainers/PPOTrainer.py b/src/brittle_star_project/trainers/PPOTrainer.py index 3dcf6db..58a477e 100644 --- a/src/brittle_star_project/trainers/PPOTrainer.py +++ b/src/brittle_star_project/trainers/PPOTrainer.py @@ -38,16 +38,19 @@ _ALLOWED_OBS_KEYS = { } # TODO: clip scaled reward? + @jax.jit def _get_xy_distance_to_target(obs_dict: dict) -> jnp.ndarray: """Extract xy_distance_to_target for all environments.""" # obs_dict is a dict of arrays with leading batch dimension (num_envs, ...) return obs_dict["xy_distance_to_target"].squeeze(-1) # shape: (num_envs,) + @jax.jit def _clip_action(action: jnp.ndarray, low: jnp.ndarray, high: jnp.ndarray) -> jnp.ndarray: 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) @@ -59,22 +62,25 @@ def _linear_schedule(count, minibatch_count, update_epochs, num_iterations, lear frac = 1.0 - (count // (minibatch_count * update_epochs)) / num_iterations return learning_rate * frac + @jax.jit def _normalize_obs(obs, mean, var, eps=1e-8): return jnp.clip((obs - mean) / jnp.sqrt(var + eps), -10.0, 10.0) + @jax.jit def _convert_obs_dict_to_array(obs_dict: dict) -> jnp.ndarray: """Convert the raw observation dict → flat array, filtering unwanted keys.""" + def _filter_and_flatten(o: dict) -> jnp.ndarray: values = [] for key in sorted(o.keys()): - if key in _ALLOWED_OBS_KEYS: #TODO: NORMALIZATION or .. of observations?? + if key in _ALLOWED_OBS_KEYS: # TODO: NORMALIZATION or .. of observations?? v = o[key] if v.size > 0: values.append(jnp.asarray(v).flatten()) return jnp.concatenate(values) - + return jax.vmap(_filter_and_flatten)(obs_dict) @@ -87,7 +93,7 @@ def _get_action_and_value_noise( next_obs: jnp.ndarray, key: jax.random.PRNGKey, action_low, - action_high + action_high, ): hidden = sensor.apply(agent_state.params["sensor_params"], next_obs) hidden_critic = feature_extractor.apply( @@ -100,16 +106,11 @@ def _get_action_and_value_noise( noise = jax.random.normal(subkey, shape=mean.shape) std = jnp.exp(log_std) raw_action = mean + noise * std - clipped_action = _clip_action( - raw_action, - action_low, - action_high - ) + clipped_action = _clip_action(raw_action, action_low, action_high) 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, raw_action, logprob, value.squeeze(-1), mean, std, key + return clipped_action, raw_action, logprob, value.squeeze(-1), mean, std, key def _step_once( @@ -121,7 +122,7 @@ def _step_once( actor: Actor, critic: OneDenseLayerMLP, action_low, - action_high + action_high, ): agent_state, episode_stats, obs, done, key, env_state = carry clipped_action, raw_action, logprob, value, mean, std, key = _get_action_and_value_noise( @@ -152,8 +153,8 @@ def _step_env_wrapped(episode_stats, env_state, action, env_step_fn): next_env_state = env_step_fn(env_state, action) reward = next_env_state.reward - reward *= 5000 - reward = jnp.clip(reward, -1, 1) + reward *= 20000 + reward = jnp.clip(reward, -10, 10) terminated = next_env_state.terminated truncated = next_env_state.truncated done = terminated | truncated @@ -192,7 +193,7 @@ def _rollout_jit( actor: Actor, critic: OneDenseLayerMLP, action_low, - action_high + action_high, ): (agent_state, episode_stats, next_obs, next_done, key, env_state), storage = jax.lax.scan( partial( @@ -203,7 +204,7 @@ def _rollout_jit( critic=critic, env_step_fn=step_env_fn, action_low=action_low, - action_high=action_high + action_high=action_high, ), (agent_state, episode_stats, next_obs, next_done, key, env_state), (), @@ -295,7 +296,7 @@ class PPOTrainer: actor=self.actor, critic=self.critic, action_low=action_low, - action_high=action_high + action_high=action_high, ) ) self._compute_gae_jit = jax.jit( @@ -382,7 +383,7 @@ class PPOTrainer: returned_episode_returns=jnp.zeros(self.args.num_envs, jnp.float32), returned_episode_lengths=jnp.zeros(self.args.num_envs, dtype=jnp.int32), ) - + def _update_obs_stats(self, obs: jnp.ndarray): batch_mean = jnp.mean(obs, axis=0) batch_var = jnp.var(obs, axis=0) @@ -429,39 +430,36 @@ class PPOTrainer: training_measurements, storage, next_obs, - xy_distance + xy_distance, ): - data = jax.device_get({ - 'rewards': storage.rewards[0], - 'values': storage.values[0], - 'returns': storage.returns[0], - 'advantages': storage.advantages[0], - 'actions': storage.actions[0], - 'raw_actions': storage.raw_actions[0], - 'means': storage.means[0], - 'stds': storage.stds[0], - 'logprobs': storage.logprobs[0], - }) + data = jax.device_get( + { + "rewards": storage.rewards[0], + "values": storage.values[0], + "returns": storage.returns[0], + "advantages": storage.advantages[0], + "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'])), + "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"])), } for i in range(len(xy_distance)): - storage_metrics[f"env_data/env{i}_xy_dist_target"] = float(xy_distance[i]) + storage_metrics[f"env_data/env{i}_xy_dist_target"] = float(xy_distance[i]) metrics = { "charts/avg_episodic_return": training_measurements.avg_episodic_return, @@ -485,7 +483,7 @@ class PPOTrainer: "charts/SPS_update": int( self.args.num_envs * self.args.num_steps / (time.time() - iteration_time_start) ), - **storage_metrics + **storage_metrics, } self.logger.log(metrics, step=global_step) @@ -556,7 +554,7 @@ class PPOTrainer: avg_terminated_length=avg_terminated_length, avg_truncated_length=avg_truncated_length, ), - storage + storage, ) def _close(self): @@ -617,7 +615,7 @@ class PPOTrainer: training_measurements, storage, next_obs, - xy_distance + xy_distance, ) sps = int(global_step / (time.time() - start_time)) From ae6beb174bb341ad8498ee032f8a11b9aee7b383 Mon Sep 17 00:00:00 2001 From: cedric Date: Wed, 15 Apr 2026 05:39:18 +0000 Subject: [PATCH 11/11] fix: updated used vars --- .../used_variables.md | 25 +++++++++++++------ 1 file changed, 18 insertions(+), 7 deletions(-) diff --git a/experiments/debug-experiment-10042026/used_variables.md b/experiments/debug-experiment-10042026/used_variables.md index e49c9e4..92f23a1 100644 --- a/experiments/debug-experiment-10042026/used_variables.md +++ b/experiments/debug-experiment-10042026/used_variables.md @@ -20,23 +20,35 @@ learning_rate: float = 2.5e-4 anneal_lr: bool = True gamma: float = 0.99 gae_lambda: float = 0.95 -num_minibatches: int = 4 update_epochs: int = 4 norm_adv: bool = True -clip_coef: float = 0.1 clip_vloss: bool = True -ent_coef: float = 0.01 -vf_coef: float = 0.5 max_grad_norm: float = 0.5 target_kl: float | None = None batch_size: int = 0 minibatch_size: int = 0 num_iterations: int = 0 -## Used config file: +## Used config file: (hpc/debug.yaml) +exp_name: "debug-experiment" +seed: 42 +track: true +wandb_project_name: "Let's-find-that-bug" +wandb_entity: "SEL3-2026-Groep-4" +run_dir: "/data/gent/465/vsc46589" num_envs: 32 num_steps: 32 -total_timesteps: 102400 +num_minibatches: 32 +total_timesteps: 409600 +num_arms: 2 +cuda: true + +ent_coef: 0.005 +vf_coef: 1.0 +clip_coef: 0.2 + +anneal_lr: true +learning_rate: 0.0003 ## Arena config: size: tuple[float, float] = (10.0, 5.0) @@ -46,7 +58,6 @@ wall_height: float = 1.5 wall_thickness: float = 0.1 ## Morphology: -num_arms: int = 2 num_segments_per_arm: int = 4 use_p_control: bool = True use_torque_control: bool = False