1
Fork 0
This commit is contained in:
JibrilExe 2026-05-02 11:12:01 +02:00
commit 9398592e34
7 changed files with 103 additions and 56 deletions

View file

@ -2,9 +2,9 @@
# Lower timestep count for quick iterations/testing. # Lower timestep count for quick iterations/testing.
learning_rate: 0.0005 learning_rate: 0.0005
total_timesteps: 65536 total_timesteps: 1024
num_envs: 512 num_envs: 32
num_steps: 128 num_steps: 32
anneal_lr: true anneal_lr: true
gamma: 0.99 gamma: 0.99
gae_lambda: 0.95 gae_lambda: 0.95

View file

@ -8,6 +8,7 @@ from brittle_star_project.configs.register_configs import register_configs
from brittle_star_project.trainers.PPOTrainer import PPOTrainer from brittle_star_project.trainers.PPOTrainer import PPOTrainer
from brittle_star_project.environment.BrittleStarJaxEnvWrapper import BrittleStarJaxEnvWrapper from brittle_star_project.environment.BrittleStarJaxEnvWrapper import BrittleStarJaxEnvWrapper
from experiment_logger import init_logger, get_logger from experiment_logger import init_logger, get_logger
import logging
def make_env(cfg: BrittleStarConfig) -> BrittleStarJaxEnvWrapper: def make_env(cfg: BrittleStarConfig) -> BrittleStarJaxEnvWrapper:
@ -41,6 +42,7 @@ def main(dict_cfg: DictConfig):
base_dir=os.path.dirname(run_dir), base_dir=os.path.dirname(run_dir),
) )
logger = get_logger() logger = get_logger()
logger.set_level(logging.DEBUG)
logger.info(f"Hydra-initialized run: {run_name}") logger.info(f"Hydra-initialized run: {run_name}")
logger.info(f"Output directory: {run_dir}") logger.info(f"Output directory: {run_dir}")

View file

@ -4,7 +4,7 @@ import jax.numpy as jnp
@flax.struct.dataclass @flax.struct.dataclass
class EpisodeStatistics: class EpisodeStatistics:
episode_returns: jnp.array episode_returns: jnp.ndarray
episode_lengths: jnp.array episode_lengths: jnp.ndarray
returned_episode_returns: jnp.array returned_episode_returns: jnp.ndarray
returned_episode_lengths: jnp.array returned_episode_lengths: jnp.ndarray

View file

@ -1,15 +1,27 @@
from functools import partial from functools import partial
import flax
import jax import jax
import jax.numpy as jnp import jax.numpy as jnp
from jax import debug
from flax.core import FrozenDict
from experiment_logger import get_logger from experiment_logger import get_logger
from brittle_star_project.utils import logged_jit
logger = get_logger() logger = get_logger()
# Chose to use a class as it seemed the easiest way to integrate the CleanRL code style # Chose to use a class as it seemed the easiest way to integrate the CleanRL code style
# with our need to seperate concerns # with our need to seperate concerns
class PPO: class PPO:
def __init__(self, args, sensor_apply, actor_apply, critic_apply, feature_extractor_apply, message_passer=None): def __init__(
self,
args,
sensor_apply,
actor_apply,
critic_apply,
feature_extractor_apply,
message_passer=None,
):
self.args = args self.args = args
if not message_passer: if not message_passer:
@ -30,13 +42,13 @@ class PPO:
# This PPO class should be initialized only once, # This PPO class should be initialized only once,
# or this function will need to recompile # or this function will need to recompile
@partial(jax.jit, static_argnums=0) @partial(logged_jit, static_argnums=0)
def update_ppo(self, agent_state, storage, key): def update_ppo(self, agent_state, storage, key):
logger.info(f"[update_ppo] storage.obs shape: {getattr(storage, 'obs', None).shape}") debug.callback(logger.debug, f"[PPO] storage.obs shape: {storage.obs.shape}")
logger.info(f"[update_ppo] storage.actions shape: {storage.actions.shape}") debug.callback(logger.debug, f"[PPO] storage.actions shape: {storage.actions.shape}")
logger.info(f"[update_ppo] storage.logprobs shape: {storage.logprobs.shape}") debug.callback(logger.debug, f"[PPO] storage.logprobs shape: {storage.logprobs.shape}")
logger.info(f"[update_ppo] storage.advantages shape: {storage.advantages.shape}") debug.callback(logger.debug, f"[PPO] storage.advantages shape: {storage.advantages.shape}")
logger.info(f"[update_ppo] storage.returns shape: {storage.returns.shape}") debug.callback(logger.debug, f"[PPO] storage.returns shape: {storage.returns.shape}")
args = self.args args = self.args
ppo_loss_grad_fn = self.ppo_loss_grad_fn ppo_loss_grad_fn = self.ppo_loss_grad_fn
@ -56,11 +68,16 @@ class PPO:
shuffled_storage = jax.tree.map(convert_data, flatten_storage) shuffled_storage = jax.tree.map(convert_data, flatten_storage)
def update_minibatch(agent_state, minibatch): def update_minibatch(agent_state, minibatch):
logger.info(f"[update_ppo] minibatch.obs: {minibatch.obs.shape}") debug.callback(logger.debug, f"[PPO] minibatch.obs: {minibatch.obs.shape}")
logger.info(f"[update_ppo] minibatch.actions: {minibatch.actions.shape}") debug.callback(logger.debug, f"[PPO] minibatch.actions: {minibatch.actions.shape}")
logger.info(f"[update_ppo] minibatch.logprobs: {minibatch.logprobs.shape}") debug.callback(
logger.info(f"[update_ppo] minibatch.advantages: {minibatch.advantages.shape}") logger.debug, f"[PPO] minibatch.logprobs: {minibatch.logprobs.shape}"
logger.info(f"[update_ppo] minibatch.returns: {minibatch.returns.shape}") )
debug.callback(
logger.debug, f"[PPO] minibatch.advantages: {minibatch.advantages.shape}"
)
debug.callback(logger.debug, f"[PPO] minibatch.returns: {minibatch.returns.shape}")
(loss, (pg_loss, v_loss, entropy_loss, approx_kl)), grads = ppo_loss_grad_fn( (loss, (pg_loss, v_loss, entropy_loss, approx_kl)), grads = ppo_loss_grad_fn(
agent_state.params, agent_state.params,
minibatch.obs, minibatch.obs,
@ -70,13 +87,7 @@ class PPO:
minibatch.returns, minibatch.returns,
) )
agent_state = agent_state.apply_gradients(grads=grads) agent_state = agent_state.apply_gradients(grads=grads)
return agent_state, ( return agent_state, (loss, pg_loss, v_loss, entropy_loss, approx_kl)
loss,
pg_loss,
v_loss,
entropy_loss,
approx_kl
)
agent_state, metrics = jax.lax.scan(update_minibatch, agent_state, shuffled_storage) agent_state, metrics = jax.lax.scan(update_minibatch, agent_state, shuffled_storage)
return (agent_state, key), metrics return (agent_state, key), metrics
@ -95,14 +106,14 @@ that are now not in the same scope
""" """
@partial(jax.jit, static_argnums=(0, 1, 2, 3, 4)) @partial(logged_jit, static_argnums=(0, 1, 2, 3, 4))
def get_action_and_value( def get_action_and_value(
sensor_apply, sensor_apply,
actor_apply, actor_apply,
message_passer, message_passer,
critic_apply, critic_apply,
feature_extractor_apply, feature_extractor_apply,
params: flax.core.FrozenDict, params: FrozenDict,
x: jnp.ndarray, x: jnp.ndarray,
action: jnp.ndarray, action: jnp.ndarray,
): ):
@ -110,25 +121,28 @@ def get_action_and_value(
hidden_critic = feature_extractor_apply(params["feature_extractor_params"], x) hidden_critic = feature_extractor_apply(params["feature_extractor_params"], x)
hidden_sensor = message_passer(hidden_sensor) hidden_sensor = message_passer(hidden_sensor)
logger.info(f"[get_action_and_value] hidden_sensor: {hidden_sensor.shape}") debug.callback(logger.debug, f"[SHAPE] hidden_sensor: {hidden_sensor.shape}")
logger.info(f"[get_action_and_value] hidden_critic: {hidden_critic.shape}") debug.callback(logger.debug, f"[SHAPE] hidden_critic: {hidden_critic.shape}")
mean, log_std = actor_apply(params["actor_params"], hidden_sensor) mean, log_std = actor_apply(params["actor_params"], hidden_sensor)
logger.info(f"[get_action_and_value] mean: {mean.shape}") debug.callback(logger.debug, f"[SHAPE] mean: {mean.shape}")
logger.info(f"[get_action_and_value] log_std: {log_std.shape}") debug.callback(logger.debug, f"[SHAPE] log_std: {log_std.shape}")
logger.info(f"[get_action_and_value] action: {action.shape}") debug.callback(logger.debug, f"[SHAPE] action: {action.shape}")
log_std = jnp.clip(log_std, -5, 2) log_std = jnp.clip(log_std, -5, 2)
std = jnp.exp(log_std) std = jnp.exp(log_std)
logprob = -0.5 * (((action - mean) / std) ** 2 + 2 * log_std + jnp.log(2 * jnp.pi)) logprob = -0.5 * (((action - mean) / std) ** 2 + 2 * log_std + jnp.log(2 * jnp.pi))
logger.info(f"[get_action_and_value] logprob pre-sum: {logprob.shape}") debug.callback(logger.debug, f"[SHAPE] logprob pre-sum: {logprob.shape}")
logprob = logprob.sum(axis=(-2, -1)) logprob = logprob.sum(axis=(-2, -1))
logger.info(f"[get_action_and_value] logprob final: {logprob.shape}") debug.callback(logger.debug, f"[SHAPE] logprob final: {logprob.shape}")
entropy = (0.5 + 0.5 * jnp.log(2 * jnp.pi) + log_std).sum(axis=(-2, -1)) entropy = (0.5 + 0.5 * jnp.log(2 * jnp.pi) + log_std).sum(axis=(-2, -1))
value = critic_apply(params["critic_params"], hidden_critic).squeeze(-1) value = critic_apply(params["critic_params"], hidden_critic).squeeze(-1)
logger.info(f"[get_action_and_value] value: {value.shape}") debug.callback(logger.debug, f"[SHAPE] value: {value.shape}")
return logprob, entropy, value return logprob, entropy, value

View file

@ -26,6 +26,7 @@ from brittle_star_project.MLPs.mlps import (
) )
from brittle_star_project.ppo import PPO from brittle_star_project.ppo import PPO
from brittle_star_project.environment import MorphMode from brittle_star_project.environment import MorphMode
from brittle_star_project.utils import logged_jit
logger11 = get_logger() logger11 = get_logger()
# TODO: move to config # TODO: move to config
@ -109,7 +110,7 @@ def build_adjacency(segments_per_arm, mode: MorphMode):
return adj return adj
@jax.jit @logged_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)
@ -120,13 +121,13 @@ def _compute_explained_variance(values: jnp.ndarray, returns: jnp.ndarray) -> fl
return float(explained_var) return float(explained_var)
@jax.jit @logged_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
return learning_rate * frac return learning_rate * frac
@jax.jit @logged_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)
@ -135,7 +136,7 @@ def _convert_obs_dict_to_array_morphology(obs_dict, morph_mode, segments_per_arm
num_segments = segments_per_arm.sum() num_segments = segments_per_arm.sum()
num_arms = jnp.where(segments_per_arm > 0, 1, 0).sum() num_arms = jnp.where(segments_per_arm > 0, 1, 0).sum()
@jax.jit @logged_jit
def _filter_and_flatten(o) -> jnp.ndarray: def _filter_and_flatten(o) -> jnp.ndarray:
# vmap feeds one env at a time — v has NO batch dim here # vmap feeds one env at a time — v has NO batch dim here
# shapes are e.g. (n_features,) or (n_nodes, feat) # shapes are e.g. (n_features,) or (n_nodes, feat)
@ -232,7 +233,6 @@ def _get_action_and_value_noise(
flat_action = raw_action.reshape(raw_action.shape[0], -1) # concat the per agent, keep the envs dim flat_action = raw_action.reshape(raw_action.shape[0], -1) # concat the per agent, keep the envs dim
flat_clipped_action = _clip_action(flat_action, action_low, action_high) flat_clipped_action = _clip_action(flat_action, action_low, action_high)
logprob = -0.5 * (((raw_action - mean) / std) ** 2 + 2 * log_std + jnp.log(2 * jnp.pi)).sum(axis=(-2, -1)) logprob = -0.5 * (((raw_action - mean) / std) ** 2 + 2 * log_std + jnp.log(2 * jnp.pi)).sum(axis=(-2, -1))
value = apply_shared(critic, agent_state.params["critic_params"], hidden_critic) value = apply_shared(critic, agent_state.params["critic_params"], hidden_critic)
@ -351,6 +351,7 @@ def _step_env_wrapped(episode_stats, env_state, action, env_step_fn, morph_mode,
), ),
) )
def apply_per_node(net, params, x): def apply_per_node(net, params, x):
# params: (nodes, ...) # params: (nodes, ...)
# x: (batch, nodes, feat) # x: (batch, nodes, feat)
@ -361,6 +362,7 @@ def apply_per_node(net, params, x):
return jax.vmap(apply_single_node, in_axes=(0, 1), out_axes=1)(params, x) return jax.vmap(apply_single_node, in_axes=(0, 1), out_axes=1)(params, x)
def apply_shared(net, params, x): def apply_shared(net, params, x):
# x: (batch, nodes, feat) # x: (batch, nodes, feat)
# If the critic expects a single vector per environment: # If the critic expects a single vector per environment:
@ -368,6 +370,7 @@ def apply_shared(net, params, x):
x_flattened = x.reshape(batch_size, -1) x_flattened = x.reshape(batch_size, -1)
return jax.vmap(lambda xi: net.apply(params, xi))(x_flattened) return jax.vmap(lambda xi: net.apply(params, xi))(x_flattened)
# TODO: update to work with extra dimension + message passing # TODO: update to work with extra dimension + message passing
def _rollout_jit( def _rollout_jit(
agent_state, agent_state,
@ -428,9 +431,10 @@ def _compute_gae_jit(
critic, critic,
adj_matrix: jnp.ndarray, adj_matrix: jnp.ndarray,
): ):
next_value = apply_shared(critic, next_value = apply_shared(
critic,
agent_state.params["critic_params"], agent_state.params["critic_params"],
apply_shared(feature_extractor,agent_state.params["feature_extractor_params"], next_obs), apply_shared(feature_extractor, agent_state.params["feature_extractor_params"], next_obs),
).squeeze(-1) ).squeeze(-1)
advantages = jnp.zeros((num_envs,)) advantages = jnp.zeros((num_envs,))
@ -500,15 +504,15 @@ class PPOTrainer:
self.needed_copies, self.needed_copies,
) = self._init_agent() ) = self._init_agent()
self.sensor.apply = jax.jit(self.sensor.apply) self.sensor.apply = logged_jit(self.sensor.apply)
self.feature_extractor.apply = jax.jit(self.feature_extractor.apply) self.feature_extractor.apply = logged_jit(self.feature_extractor.apply)
self.actor.apply = jax.jit(self.actor.apply) self.actor.apply = logged_jit(self.actor.apply)
self.critic.apply = jax.jit(self.critic.apply) self.critic.apply = logged_jit(self.critic.apply)
action_low = jnp.asarray(self.env.single_action_space.low, dtype=jnp.float32) action_low = jnp.asarray(self.env.single_action_space.low, dtype=jnp.float32)
action_high = jnp.asarray(self.env.single_action_space.high, dtype=jnp.float32) action_high = jnp.asarray(self.env.single_action_space.high, dtype=jnp.float32)
self._rollout_jit = jax.jit( self._rollout_jit = logged_jit(
partial( partial(
_rollout_jit, _rollout_jit,
max_steps=self.ppo.num_steps, max_steps=self.ppo.num_steps,
@ -528,7 +532,7 @@ class PPOTrainer:
adj_matrix=self.adj, adj_matrix=self.adj,
) )
) )
self._compute_gae_jit = jax.jit( self._compute_gae_jit = logged_jit(
partial( partial(
_compute_gae_jit, _compute_gae_jit,
num_envs=self.ppo.num_envs, num_envs=self.ppo.num_envs,
@ -540,10 +544,17 @@ class PPOTrainer:
) )
) )
apply_sensor = lambda p, x: apply_per_node(self.sensor, p, x) def apply_sensor(p, x):
apply_actor = lambda p, x: apply_per_node(self.actor, p, x) return apply_per_node(self.sensor, p, x)
apply_critic = lambda p, x: apply_shared(self.critic, p, x)
apply_feature = lambda p, x: apply_shared(self.feature_extractor, p, x) def apply_actor(p, x):
return apply_per_node(self.actor, p, x)
def apply_critic(p, x):
return apply_shared(self.critic, p, x)
def apply_feature(p, x):
return apply_shared(self.feature_extractor, p, x)
self._ppo = PPO(self.ppo, apply_sensor, apply_actor, apply_critic, apply_feature) self._ppo = PPO(self.ppo, apply_sensor, apply_actor, apply_critic, apply_feature)
@ -566,11 +577,11 @@ class PPOTrainer:
case MorphMode.CENTRALIZED: case MorphMode.CENTRALIZED:
needed_copies = 1 needed_copies = 1
case MorphMode.FULLY_CONNECTED | MorphMode.RING: case MorphMode.FULLY_CONNECTED | MorphMode.RING:
needed_copies = jnp.where(self.segments_per_arm > 0, 1, 0).sum() needed_copies = jnp.where(self.segments_per_arm > 0, 1, 0).sum().item()
case MorphMode.SEGMENT: case MorphMode.SEGMENT:
needed_copies = ( needed_copies = (
self.segments_per_arm.sum() + jnp.where(self.segments_per_arm > 0, 1, 0).sum() self.segments_per_arm.sum() + jnp.where(self.segments_per_arm > 0, 1, 0).sum()
) ).item()
actor = Actor(action_dim=self.env.single_action_space.shape[0]) actor = Actor(action_dim=self.env.single_action_space.shape[0])
sensor = GenericDenseLayersWithActivation(layer_sizes=[300, 300, 300]) sensor = GenericDenseLayersWithActivation(layer_sizes=[300, 300, 300])

View file

@ -0,0 +1,3 @@
from .logged_jit import logged_jit
__all__ = ["logged_jit"]

View file

@ -0,0 +1,17 @@
import jax
from experiment_logger import get_logger
def logged_jit(fn, **jit_kwargs):
logger = get_logger()
name = getattr(fn, "__name__", getattr(fn, "__qualname__", repr(fn)))
def decorator(func):
def traced_func(*args, **kwargs):
logger.debug(f"[JIT] Compiling {name}...")
return func(*args, **kwargs)
jitted = jax.jit(traced_func, **jit_kwargs)
return jitted
return decorator(fn)