fix: ruff format
This commit is contained in:
parent
0a9bf2e0a5
commit
20ba63a303
5 changed files with 57 additions and 52 deletions
|
|
@ -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
|
||||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
|
|
|
|||
Reference in a new issue