1
Fork 0

fix(cleanup): ruff format

This commit is contained in:
cedric 2026-04-09 22:13:07 +00:00
parent 71349af5b9
commit bf6434ea6c

View file

@ -24,11 +24,13 @@ from brittle_star_project.MLPs.mlps import (
) )
from brittle_star_project.ppo import PPO from brittle_star_project.ppo import PPO
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)
return float(explained_var) return float(explained_var)
@jax.jit @jax.jit
def _linear_schedule(count, minibatch_count, update_epochs, num_iterations, learning_rate): def _linear_schedule(count, minibatch_count, update_epochs, num_iterations, learning_rate):
frac = 1.0 - (count // (minibatch_count * update_epochs)) / num_iterations frac = 1.0 - (count // (minibatch_count * update_epochs)) / num_iterations
@ -412,25 +414,21 @@ class PPOTrainer:
jnp.mean(jax.device_get(self.episode_stats.returned_episode_returns)).item() jnp.mean(jax.device_get(self.episode_stats.returned_episode_returns)).item()
) )
explained_var = _compute_explained_variance( explained_var = _compute_explained_variance(storage.values, storage.returns)
storage.values, storage.returns
)
terminated = next_env_state.terminated # (num_envs,) terminated = next_env_state.terminated # (num_envs,)
truncated = next_env_state.truncated # (num_envs,) truncated = next_env_state.truncated # (num_envs,)
episode_lengths = self.episode_stats.returned_episode_lengths episode_lengths = self.episode_stats.returned_episode_lengths
num_terminated = int(jnp.sum(terminated).item()) num_terminated = int(jnp.sum(terminated).item())
num_truncated = int(jnp.sum(truncated).item()) num_truncated = int(jnp.sum(truncated).item())
avg_terminated_length = ( avg_terminated_length = jnp.sum(episode_lengths * terminated) / jnp.maximum(
jnp.sum(episode_lengths * terminated) / jnp.sum(terminated), 1
jnp.maximum(jnp.sum(terminated), 1)
) )
avg_truncated_length = ( avg_truncated_length = jnp.sum(episode_lengths * truncated) / jnp.maximum(
jnp.sum(episode_lengths * truncated) / jnp.sum(truncated), 1
jnp.maximum(jnp.sum(truncated), 1)
) )
return ( return (
@ -448,7 +446,7 @@ class PPOTrainer:
num_terminated=num_terminated, num_terminated=num_terminated,
num_truncated=num_truncated, num_truncated=num_truncated,
avg_terminated_length=avg_terminated_length, avg_terminated_length=avg_terminated_length,
avg_truncated_length=avg_truncated_length avg_truncated_length=avg_truncated_length,
), ),
) )