merged
This commit is contained in:
commit
9398592e34
7 changed files with 103 additions and 56 deletions
|
|
@ -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
|
||||||
|
|
|
||||||
|
|
@ -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}")
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -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])
|
||||||
|
|
|
||||||
3
src/brittle_star_project/utils/__init__.py
Normal file
3
src/brittle_star_project/utils/__init__.py
Normal file
|
|
@ -0,0 +1,3 @@
|
||||||
|
from .logged_jit import logged_jit
|
||||||
|
|
||||||
|
__all__ = ["logged_jit"]
|
||||||
17
src/brittle_star_project/utils/logged_jit.py
Normal file
17
src/brittle_star_project/utils/logged_jit.py
Normal 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)
|
||||||
Reference in a new issue