diff --git a/configs/environment/directed_locomotion.yaml b/configs/environment/directed_locomotion.yaml index 63ad520..465d80e 100644 --- a/configs/environment/directed_locomotion.yaml +++ b/configs/environment/directed_locomotion.yaml @@ -2,11 +2,11 @@ # Baseline task setting. task: DIRECTED_LOCOMOTION -simulation_time: 5.0 +simulation_time: 5000.0 num_physics_steps_per_control_step: 10 time_scale: 2 camera_ids: [0, 1] render_size: [480, 640] joint_randomization_noise_scale: 0.0 -target_distance: 3.0 +target_distance: 0.6 light_perlin_noise_scale: 0 diff --git a/configs/logging/wandb_enabled.yaml b/configs/logging/wandb_enabled.yaml index 2e82781..40c52bf 100644 --- a/configs/logging/wandb_enabled.yaml +++ b/configs/logging/wandb_enabled.yaml @@ -2,7 +2,7 @@ # For production/cloud experiments with weights synced. track: true -wandb_project_name: "PPO-Modularity" +wandb_project_name: "PPO-Modularity - reward engineering" wandb_entity: "SEL3-2026-Groep-4" capture_video: false save_model: true diff --git a/configs/morphology/2_arms.yaml b/configs/morphology/2_arms.yaml new file mode 100644 index 0000000..dc6f159 --- /dev/null +++ b/configs/morphology/2_arms.yaml @@ -0,0 +1,5 @@ +# 2 Arms Morphology Configuration + +segments_per_arm: [4, 4] +use_p_control: true +use_torque_control: false \ No newline at end of file diff --git a/configs/ppo/debug.yaml b/configs/ppo/debug.yaml index 58dbd3a..7732fd3 100644 --- a/configs/ppo/debug.yaml +++ b/configs/ppo/debug.yaml @@ -1,5 +1,5 @@ learning_rate: 0.0003 -total_timesteps: 409600 +total_timesteps: 409600 num_envs: 32 num_steps: 32 anneal_lr: true diff --git a/configs/ppo/fast.yaml b/configs/ppo/fast.yaml index 5001ef2..ecd65fa 100644 --- a/configs/ppo/fast.yaml +++ b/configs/ppo/fast.yaml @@ -3,7 +3,7 @@ learning_rate: 0.0005 total_timesteps: 500000 -num_envs: 8 +num_envs: 32 num_steps: 128 anneal_lr: true gamma: 0.99 diff --git a/configs/ppo/fast2.yaml b/configs/ppo/fast2.yaml new file mode 100644 index 0000000..b2f629a --- /dev/null +++ b/configs/ppo/fast2.yaml @@ -0,0 +1,16 @@ +learning_rate: 0.0003 +total_timesteps: 1228800 +num_envs: 32 +num_steps: 64 +anneal_lr: true +gamma: 0.99 +gae_lambda: 0.95 +num_minibatches: 32 +update_epochs: 4 +norm_adv: true +clip_coef: 0.2 +clip_vloss: true +ent_coef: 0.005 +vf_coef: 1.0 +max_grad_norm: 0.5 +target_kl: null \ No newline at end of file diff --git a/src/brittle_star_project/environment/BrittleStarJaxEnvWrapper.py b/src/brittle_star_project/environment/BrittleStarJaxEnvWrapper.py index a5175a7..3122696 100644 --- a/src/brittle_star_project/environment/BrittleStarJaxEnvWrapper.py +++ b/src/brittle_star_project/environment/BrittleStarJaxEnvWrapper.py @@ -27,7 +27,7 @@ class BrittleStarJaxEnvWrapper: ) # Pre-compute masks for observation padding - self._padding_masks = compute_padding_masks(self._morphology.segments_per_arm) + self._padding_masks = compute_padding_masks(self._morphology.segments_per_arm, (4, 4)) self._vectorized_reset = jax.jit(jax.vmap(self._env.reset)) self._vectorized_step = jax.jit(jax.vmap(self._env.step)) diff --git a/src/brittle_star_project/environment/padded_obs_wrapper.py b/src/brittle_star_project/environment/padded_obs_wrapper.py index 3f22038..4886284 100644 --- a/src/brittle_star_project/environment/padded_obs_wrapper.py +++ b/src/brittle_star_project/environment/padded_obs_wrapper.py @@ -8,7 +8,7 @@ flattened observation maintains the correct physical mapping to the neural netwo from __future__ import annotations -from typing import Any +from typing import Any, Sequence import jax.numpy as jnp # Observation keys whose size scales with the number of joints (2 per segment). @@ -30,8 +30,8 @@ _SEGMENT_SCALED_KEYS = frozenset( def compute_padding_masks( - segments_per_arm: tuple[int, ...], - reference_segments_per_arm: tuple[int, ...] = (4, 4, 4, 4, 4), + segments_per_arm: Sequence[int], + reference_segments_per_arm: Sequence[int] = (4, 4, 4, 4, 4), ) -> dict[str, Any]: """Pre-compute boolean masks for spatial insertion of observations. diff --git a/src/brittle_star_project/trainers/PPOTrainer.py b/src/brittle_star_project/trainers/PPOTrainer.py index e5ecca6..3b102d4 100644 --- a/src/brittle_star_project/trainers/PPOTrainer.py +++ b/src/brittle_star_project/trainers/PPOTrainer.py @@ -151,12 +151,27 @@ def _step_once( return (agent_state, episode_stats, next_obs, next_done, key, env_state), storage +def _reward_fn(env_state, next_env_state): + # if delta distance positive ==> brittle star walking away from target + delta_distance = ( + next_env_state.observations["xy_distance_to_target"] + - env_state.observations["xy_distance_to_target"] + ).squeeze(-1) + + env_reward = next_env_state.reward + clipped_env_reward = jnp.clip(100 * env_reward, -10, 10) + + time_penalty = 0.1 + distance_penalty = jnp.clip(0.5 * delta_distance, -0.5, 0.5) + penalty = time_penalty + distance_penalty + + return jnp.where(next_env_state.terminated, 50.0, clipped_env_reward - penalty) + + 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 *= 20000 - reward = jnp.clip(reward, -10, 10) + reward = _reward_fn(env_state, next_env_state) terminated = next_env_state.terminated truncated = next_env_state.truncated done = terminated | truncated @@ -440,62 +455,49 @@ class PPOTrainer: iteration_time_start, training_measurements, storage, - next_obs, - 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], + "rewards": storage.rewards, + "values": storage.values, + "returns": storage.returns, + "advantages": storage.advantages, } ) - 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_metrics = { + "rollout/reward_mean": float(np.mean(data["rewards"])), + "rollout/return_mean": float(np.mean(data["returns"])), + "rollout/value_mean": float(np.mean(data["values"])), + "rollout/advantage_mean": float(np.mean(data["advantages"])), + "rollout/advantage_std": float(np.std(data["advantages"])), + "rollout/value_vs_return_mse": float(np.mean((data["values"] - data["returns"]) ** 2)), } - 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( - jax.device_get(episode_stats.returned_episode_lengths) + "charts/episodic_return": training_measurements.avg_episodic_return, + "charts/episodic_length": float( + np.mean(jax.device_get(episode_stats.returned_episode_lengths)) ), - "charts/learning_rate": self.agent_state.opt_state[1] - .hyperparams["learning_rate"] - .item(), "charts/explained_variance": training_measurements.explained_variance, - "charts/num_terminated": training_measurements.num_terminated, - "charts/num_truncated": training_measurements.num_truncated, - "charts/avg_terminated_ep_length": training_measurements.avg_terminated_length, - "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/learning_rate": self.agent_state.opt_state[1] + .hyperparams["learning_rate"] + .item(), "charts/SPS": int(global_step / (time.time() - start_time)), "charts/SPS_update": int( self.ppo.num_envs * self.ppo.num_steps / (time.time() - iteration_time_start) ), - **storage_metrics, + "termi_trunci/num_terminated": training_measurements.num_terminated, + "termi_trunci/num_truncated": training_measurements.num_truncated, + "termi_trunci/avg_terminated_ep_length": training_measurements.avg_terminated_length, + "termi_trunci/avg_truncated_ep_length": training_measurements.avg_truncated_length, + **rollout_metrics, } + self.logger.log(metrics, step=global_step) def _step(self, env_state, next_obs, next_done, iteration: int) -> tuple: @@ -630,8 +632,6 @@ class PPOTrainer: iteration_time_start, training_measurements, storage, - next_obs, - xy_distance, ) sps = int(global_step / (time.time() - start_time))