1
Fork 0

fix: ruff format

This commit is contained in:
JibrilExe 2026-04-14 21:14:26 +02:00
parent 0a9bf2e0a5
commit 20ba63a303
5 changed files with 57 additions and 52 deletions

View file

@ -7,7 +7,14 @@ wandb_entity: "SEL3-2026-Groep-4"
run_dir: "/data/gent/465/vsc46589" run_dir: "/data/gent/465/vsc46589"
num_envs: 32 num_envs: 32
num_steps: 32 num_steps: 32
num_minibatches: 32
total_timesteps: 409600 total_timesteps: 409600
num_arms: 2 num_arms: 2
cuda: true cuda: true
ent_coef: 0.005
vf_coef: 1.0
clip_coef: 0.2
anneal_lr: true
learning_rate: 0.0003

View file

@ -58,9 +58,9 @@ class Storage:
returns: jnp.array returns: jnp.array
rewards: jnp.array rewards: jnp.array
raw_actions: jnp.ndarray = None # before clipping raw_actions: jnp.ndarray = None # before clipping
means: jnp.ndarray = None # policy mean means: jnp.ndarray = None # policy mean
stds: jnp.ndarray = None # policy std stds: jnp.ndarray = None # policy std
def replace(self, **kwargs) -> "Storage": def replace(self, **kwargs) -> "Storage":
fs = fields(self) fs = fields(self)

View file

@ -31,7 +31,7 @@ class EnvConfig:
task: Task = Task.DIRECTED_LOCOMOTION task: Task = Task.DIRECTED_LOCOMOTION
simulation_time: float = 500.0 simulation_time: float = 10000.0
num_physics_steps_per_control_step: int = 10 num_physics_steps_per_control_step: int = 10
time_scale: int = 2 time_scale: int = 2

View file

@ -38,16 +38,19 @@ _ALLOWED_OBS_KEYS = {
} }
# TODO: clip scaled reward? # TODO: clip scaled reward?
@jax.jit @jax.jit
def _get_xy_distance_to_target(obs_dict: dict) -> jnp.ndarray: def _get_xy_distance_to_target(obs_dict: dict) -> jnp.ndarray:
"""Extract xy_distance_to_target for all environments.""" """Extract xy_distance_to_target for all environments."""
# obs_dict is a dict of arrays with leading batch dimension (num_envs, ...) # 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,) return obs_dict["xy_distance_to_target"].squeeze(-1) # shape: (num_envs,)
@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:
return jnp.clip(action, low, high) return jnp.clip(action, low, high)
def _compute_explained_variance(values: jnp.ndarray, returns: jnp.ndarray) -> float: def _compute_explained_variance(values: jnp.ndarray, returns: jnp.ndarray) -> float:
var_returns = jnp.var(returns) var_returns = jnp.var(returns)
explained_var = 1.0 - jnp.var(returns - values) / (var_returns + 1e-8) explained_var = 1.0 - jnp.var(returns - values) / (var_returns + 1e-8)
@ -59,17 +62,20 @@ def _linear_schedule(count, minibatch_count, update_epochs, num_iterations, lear
frac = 1.0 - (count // (minibatch_count * update_epochs)) / num_iterations frac = 1.0 - (count // (minibatch_count * update_epochs)) / num_iterations
return learning_rate * frac return learning_rate * frac
@jax.jit @jax.jit
def _normalize_obs(obs, mean, var, eps=1e-8): def _normalize_obs(obs, mean, var, eps=1e-8):
return jnp.clip((obs - mean) / jnp.sqrt(var + eps), -10.0, 10.0) return jnp.clip((obs - mean) / jnp.sqrt(var + eps), -10.0, 10.0)
@jax.jit @jax.jit
def _convert_obs_dict_to_array(obs_dict: dict) -> jnp.ndarray: def _convert_obs_dict_to_array(obs_dict: dict) -> jnp.ndarray:
"""Convert the raw observation dict → flat array, filtering unwanted keys.""" """Convert the raw observation dict → flat array, filtering unwanted keys."""
def _filter_and_flatten(o: dict) -> jnp.ndarray: def _filter_and_flatten(o: dict) -> jnp.ndarray:
values = [] values = []
for key in sorted(o.keys()): 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] v = o[key]
if v.size > 0: if v.size > 0:
values.append(jnp.asarray(v).flatten()) values.append(jnp.asarray(v).flatten())
@ -87,7 +93,7 @@ def _get_action_and_value_noise(
next_obs: jnp.ndarray, next_obs: jnp.ndarray,
key: jax.random.PRNGKey, key: jax.random.PRNGKey,
action_low, action_low,
action_high action_high,
): ):
hidden = sensor.apply(agent_state.params["sensor_params"], next_obs) hidden = sensor.apply(agent_state.params["sensor_params"], next_obs)
hidden_critic = feature_extractor.apply( hidden_critic = feature_extractor.apply(
@ -100,18 +106,13 @@ def _get_action_and_value_noise(
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)
raw_action = mean + noise * std raw_action = mean + noise * std
clipped_action = _clip_action( clipped_action = _clip_action(raw_action, action_low, action_high)
raw_action,
action_low,
action_high
)
logprob = -0.5 * (((raw_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, raw_action, logprob, value.squeeze(-1), mean, std, key return clipped_action, raw_action, logprob, value.squeeze(-1), mean, std, key
def _step_once( def _step_once(
carry, carry,
_, _,
@ -121,7 +122,7 @@ def _step_once(
actor: Actor, actor: Actor,
critic: OneDenseLayerMLP, critic: OneDenseLayerMLP,
action_low, action_low,
action_high action_high,
): ):
agent_state, episode_stats, obs, done, key, env_state = carry agent_state, episode_stats, obs, done, key, env_state = carry
clipped_action, raw_action, logprob, value, mean, std, key = _get_action_and_value_noise( 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) next_env_state = env_step_fn(env_state, action)
reward = next_env_state.reward reward = next_env_state.reward
reward *= 5000 reward *= 20000
reward = jnp.clip(reward, -1, 1) reward = jnp.clip(reward, -10, 10)
terminated = next_env_state.terminated terminated = next_env_state.terminated
truncated = next_env_state.truncated truncated = next_env_state.truncated
done = terminated | truncated done = terminated | truncated
@ -192,7 +193,7 @@ def _rollout_jit(
actor: Actor, actor: Actor,
critic: OneDenseLayerMLP, critic: OneDenseLayerMLP,
action_low, action_low,
action_high action_high,
): ):
(agent_state, episode_stats, next_obs, next_done, key, env_state), storage = jax.lax.scan( (agent_state, episode_stats, next_obs, next_done, key, env_state), storage = jax.lax.scan(
partial( partial(
@ -203,7 +204,7 @@ def _rollout_jit(
critic=critic, critic=critic,
env_step_fn=step_env_fn, env_step_fn=step_env_fn,
action_low=action_low, action_low=action_low,
action_high=action_high action_high=action_high,
), ),
(agent_state, episode_stats, next_obs, next_done, key, env_state), (agent_state, episode_stats, next_obs, next_done, key, env_state),
(), (),
@ -295,7 +296,7 @@ class PPOTrainer:
actor=self.actor, actor=self.actor,
critic=self.critic, critic=self.critic,
action_low=action_low, action_low=action_low,
action_high=action_high action_high=action_high,
) )
) )
self._compute_gae_jit = jax.jit( self._compute_gae_jit = jax.jit(
@ -429,35 +430,32 @@ class PPOTrainer:
training_measurements, training_measurements,
storage, storage,
next_obs, next_obs,
xy_distance xy_distance,
): ):
data = jax.device_get({ data = jax.device_get(
'rewards': storage.rewards[0], {
'values': storage.values[0], "rewards": storage.rewards[0],
'returns': storage.returns[0], "values": storage.values[0],
'advantages': storage.advantages[0], "returns": storage.returns[0],
'actions': storage.actions[0], "advantages": storage.advantages[0],
'raw_actions': storage.raw_actions[0], "actions": storage.actions[0],
'means': storage.means[0], "raw_actions": storage.raw_actions[0],
'stds': storage.stds[0], "means": storage.means[0],
'logprobs': storage.logprobs[0], "stds": storage.stds[0],
}) "logprobs": storage.logprobs[0],
}
)
storage_metrics = { storage_metrics = {
"rollout/env0/return_mean": float(np.mean(data['returns'])), "rollout/env0/return_mean": float(np.mean(data["returns"])),
"rollout/env0/advantage_mean": float(np.mean(data["advantages"])),
"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/value_mean": float(np.mean(data['values'])), "rollout/env0/reward_mean": float(np.mean(data["rewards"])),
"rollout/env0/value_vs_return_diff": float(np.mean(data['values'] - data['returns'])), "rollout/env0/mean_mean": float(np.mean(data["means"])),
"rollout/env0/logprob_mean": float(np.mean(data["logprobs"])),
"rollout/env0/reward_mean": float(np.mean(data['rewards'])), "rollout/env0/action_mean": float(np.mean(data["actions"])),
"rollout/env0/raw_action_mean": float(np.mean(data["raw_actions"])),
"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)): for i in range(len(xy_distance)):
@ -485,7 +483,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 **storage_metrics,
} }
self.logger.log(metrics, step=global_step) self.logger.log(metrics, step=global_step)
@ -556,7 +554,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 storage,
) )
def _close(self): def _close(self):
@ -617,7 +615,7 @@ class PPOTrainer:
training_measurements, training_measurements,
storage, storage,
next_obs, next_obs,
xy_distance xy_distance,
) )
sps = int(global_step / (time.time() - start_time)) sps = int(global_step / (time.time() - start_time))