diff --git a/configs/architecture/centralized.yaml b/configs/architecture/centralized.yaml index 3b7783d..bb4bc0f 100644 --- a/configs/architecture/centralized.yaml +++ b/configs/architecture/centralized.yaml @@ -6,7 +6,7 @@ name: "centralized" sensor: - hidden_dims: [64, 64] + hidden_dims: [300, 300, 300] activation: "tanh" motor: @@ -14,7 +14,7 @@ motor: activation: "tanh" feature_extractor: - hidden_dims: [64, 64] + hidden_dims: [300, 300, 300] activation: "tanh" critic: diff --git a/configs/architecture/decentralized.yaml b/configs/architecture/decentralized.yaml index fbf07bc..4b9fffb 100644 --- a/configs/architecture/decentralized.yaml +++ b/configs/architecture/decentralized.yaml @@ -6,11 +6,11 @@ name: "decentralized" sensor: - hidden_dims: [64, 64] + hidden_dims: [300, 300, 300] activation: "tanh" propagator: - hidden_dims: [64, 64] + hidden_dims: [300, 300, 300] activation: "tanh" motor: @@ -18,7 +18,7 @@ motor: activation: "tanh" feature_extractor: - hidden_dims: [64, 64] + hidden_dims: [300, 300, 300] activation: "tanh" critic: diff --git a/configs/ppo/debug.yaml b/configs/ppo/debug.yaml new file mode 100644 index 0000000..58dbd3a --- /dev/null +++ b/configs/ppo/debug.yaml @@ -0,0 +1,16 @@ +learning_rate: 0.0003 +total_timesteps: 409600 +num_envs: 32 +num_steps: 32 +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 diff --git a/experiments/debug-experiment-10042026/used_variables.md b/experiments/debug-experiment-10042026/used_variables.md new file mode 100644 index 0000000..92f23a1 --- /dev/null +++ b/experiments/debug-experiment-10042026/used_variables.md @@ -0,0 +1,107 @@ +## Default envconfig +task: Task = Task.DIRECTED_LOCOMOTION +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]) +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 +update_epochs: int = 4 +norm_adv: bool = True +clip_vloss: bool = True +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: (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 +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) +sand_ground_color: bool = True +attach_target: bool = True +wall_height: float = 1.5 +wall_thickness: float = 0.1 + +## Morphology: +num_segments_per_arm: int = 4 +use_p_control: bool = True +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 + + @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: +_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/MLPs/mlps.py b/src/brittle_star_project/MLPs/mlps.py index 6abb540..5e2deb5 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/environment/env_config.py b/src/brittle_star_project/environment/env_config.py index ee30e87..7cb4c21 100644 --- a/src/brittle_star_project/environment/env_config.py +++ b/src/brittle_star_project/environment/env_config.py @@ -44,7 +44,7 @@ class EnvConfig: task: Task = Task.DIRECTED_LOCOMOTION - simulation_time: float = 5.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 cf6c69e..1ea37a8 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) @@ -142,6 +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 b3321b2..e5ecca6 100644 --- a/src/brittle_star_project/trainers/PPOTrainer.py +++ b/src/brittle_star_project/trainers/PPOTrainer.py @@ -25,6 +25,33 @@ from brittle_star_project.MLPs.mlps import ( ) from brittle_star_project.ppo import PPO +# 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 _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) @@ -38,11 +65,25 @@ def _linear_schedule(count, minibatch_count, update_epochs, num_iterations, lear 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: - 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( @@ -53,6 +94,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( @@ -60,13 +103,16 @@ 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) - action = mean + noise * std - logprob = -0.5 * (((action - mean) / std) ** 2 + 2 * log_std + jnp.log(2 * jnp.pi)).sum(-1) + raw_action = mean + noise * std + 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 action, logprob, value.squeeze(-1), key + + return clipped_action, raw_action, logprob, value.squeeze(-1), mean, std, key def _step_once( @@ -77,23 +123,28 @@ 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 + 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=raw_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), ) @@ -104,6 +155,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 *= 20000 + reward = jnp.clip(reward, -10, 10) terminated = next_env_state.terminated truncated = next_env_state.truncated done = terminated | truncated @@ -141,6 +194,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( @@ -150,6 +205,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), (), @@ -192,7 +249,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 @@ -235,6 +294,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, @@ -244,6 +306,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( @@ -274,8 +338,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 @@ -287,15 +351,11 @@ 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 + 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)) @@ -335,6 +395,25 @@ class PPOTrainer: returned_episode_lengths=jnp.zeros(self.ppo.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( self.agent_state, @@ -360,7 +439,39 @@ class PPOTrainer: start_time, 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], + } + ) + + 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"])), + } + + 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( @@ -383,6 +494,7 @@ class PPOTrainer: "charts/SPS_update": int( self.ppo.num_envs * self.ppo.num_steps / (time.time() - iteration_time_start) ), + **storage_metrics, } self.logger.log(metrics, step=global_step) @@ -453,6 +565,7 @@ class PPOTrainer: avg_terminated_length=avg_terminated_length, avg_truncated_length=avg_truncated_length, ), + storage, ) def _close(self): @@ -501,9 +614,13 @@ 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 ) + 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) global_step += self.ppo.num_steps * self.ppo.num_envs self._log( @@ -512,6 +629,9 @@ class PPOTrainer: start_time, iteration_time_start, training_measurements, + storage, + next_obs, + xy_distance, ) sps = int(global_step / (time.time() - start_time))