fix(cleanup): ruff format
This commit is contained in:
parent
71349af5b9
commit
bf6434ea6c
1 changed files with 10 additions and 12 deletions
|
|
@ -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,
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
|
||||||
Reference in a new issue