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"
|
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
|
||||||
|
|
@ -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
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -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))
|
||||||
|
|
|
||||||
Reference in a new issue