1
Fork 0

Morphology/brittle star 2 arms (#17)

Provide the features to read in Environment (Morphology, arena, ...) config files in the JSON format. The example JSON file contains the config for a brittle star with 2 arms
This commit is contained in:
RobinMeersman 2026-04-01 19:17:42 +02:00 committed by GitHub
parent ecdbe74df4
commit f594cafead
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
13 changed files with 113 additions and 100 deletions

8
configs/example.json Normal file
View file

@ -0,0 +1,8 @@
{
"morphology": {
"num_arms": 2,
"num_segments_per_arm": 4,
"use_p_control": true,
"use_torque_control": false
}
}

View file

@ -4,17 +4,15 @@ import argparse
from pathlib import Path from pathlib import Path
from brittle_star_project import ( from brittle_star_project import (
ArenaConfig,
Backend, Backend,
BrittleStarEnv, BrittleStarEnv,
BrittleStarEnvFactory, BrittleStarEnvFactory,
EnvConfig, SimulationConfig,
MorphologyConfig, simulate_policy,
Task,
) )
from brittle_star_project.environment import from_json
from brittle_star_project.rl import RLModel # imports concrete models via rl.__init__ from brittle_star_project.rl import RLModel # imports concrete models via rl.__init__
from brittle_star_project.rl.base import get_rl_model_registry from brittle_star_project.rl.base import get_rl_model_registry
from brittle_star_project.renderer import SimulationConfig, simulate_policy
MODEL_BY_NAME = get_rl_model_registry() MODEL_BY_NAME = get_rl_model_registry()
MODEL_OPTIONS = sorted(MODEL_BY_NAME) MODEL_OPTIONS = sorted(MODEL_BY_NAME)
@ -35,9 +33,9 @@ def parse_args() -> argparse.Namespace:
help="Which model class to instantiate when --model is omitted.", help="Which model class to instantiate when --model is omitted.",
) )
p.add_argument( p.add_argument(
"--task", "--backend",
choices=[t.value for t in Task], choices=[b for b in Backend],
default=Task.DIRECTED_LOCOMOTION.value, default=Backend.MJX,
) )
p.add_argument("--seed", type=int, default=None) p.add_argument("--seed", type=int, default=None)
return p.parse_args() return p.parse_args()
@ -46,14 +44,11 @@ def parse_args() -> argparse.Namespace:
def main() -> None: def main() -> None:
args = parse_args() args = parse_args()
morphology_cfg, arena_cfg, env_cfg = from_json("../configs/test.json")
# ======= ENVIRONMENT SETUP ======= # ======= ENVIRONMENT SETUP =======
backend = Backend.MJC backend = args.backend
task = Task(args.task)
morphology_cfg = MorphologyConfig()
arena_cfg = ArenaConfig(attach_target=(task == Task.DIRECTED_LOCOMOTION))
env_cfg = EnvConfig(task=task)
factory = BrittleStarEnvFactory() factory = BrittleStarEnvFactory()
raw_env = factory.create_environment(backend, morphology_cfg, arena_cfg, env_cfg) raw_env = factory.create_environment(backend, morphology_cfg, arena_cfg, env_cfg)

View file

@ -393,4 +393,3 @@ extend-ignore = [
"PLW0603", # global-statement "PLW0603", # global-statement
# "PLW1404", # implicit-str-concat # "PLW1404", # implicit-str-concat
] ]

View file

@ -0,0 +1,25 @@
from .brittle_star_project import (
ArenaConfig,
Backend,
BrittleStarEnv,
BrittleStarEnvFactory,
EnvConfig,
MorphologyConfig,
Task,
simulate_policy,
SimulationConfig,
ControlPolicy,
)
__all__ = [
"ArenaConfig",
"Backend",
"BrittleStarEnv",
"BrittleStarEnvFactory",
"EnvConfig",
"MorphologyConfig",
"Task",
"simulate_policy",
"SimulationConfig",
"ControlPolicy",
]

View file

@ -2,6 +2,7 @@ from .environment.env_types import Backend, Task
from .environment.env_config import ArenaConfig, EnvConfig, MorphologyConfig from .environment.env_config import ArenaConfig, EnvConfig, MorphologyConfig
from .environment.factory import BrittleStarEnvFactory from .environment.factory import BrittleStarEnvFactory
from .environment.env_wrapper import BrittleStarEnv from .environment.env_wrapper import BrittleStarEnv
from .render import simulate_policy, SimulationConfig, ControlPolicy
__all__ = [ __all__ = [
"ArenaConfig", "ArenaConfig",
@ -11,4 +12,7 @@ __all__ = [
"EnvConfig", "EnvConfig",
"MorphologyConfig", "MorphologyConfig",
"Task", "Task",
"simulate_policy",
"SimulationConfig",
"ControlPolicy",
] ]

View file

@ -7,6 +7,9 @@ class PPOArgs:
source: https://github.com/vwxyzjn/cleanrl/blob/master/cleanrl/ppo_atari_envpool_xla_jax_scan.py source: https://github.com/vwxyzjn/cleanrl/blob/master/cleanrl/ppo_atari_envpool_xla_jax_scan.py
""" """
# path to environment config file, if None, use default config
config_path: str | None = None
# the name of this experiment # the name of this experiment
exp_name: str = "brittle_star_ppo" exp_name: str = "brittle_star_ppo"
@ -41,8 +44,6 @@ class PPOArgs:
hf_entity: str = "" hf_entity: str = ""
# ==== Algorithm specific dataclasses ==== # ==== Algorithm specific dataclasses ====
# the id of the environment
env_id: str = "" # todo
# total timesteps of the experiments # total timesteps of the experiments
total_timesteps: int = 10000000 total_timesteps: int = 10000000
@ -51,7 +52,7 @@ class PPOArgs:
learning_rate: float = 2.5e-4 learning_rate: float = 2.5e-4
# the number of parallel game environments # the number of parallel game environments
num_envs: int = 16 num_envs: int = 100
# the number of steps to run in each environment per policy rollout # the number of steps to run in each environment per policy rollout
num_steps: int = 128 num_steps: int = 128

View file

@ -8,6 +8,7 @@ from brittle_star_project import (
ArenaConfig, ArenaConfig,
Backend, Backend,
) )
from brittle_star_project.environment import from_json
class BrittleStarJaxEnvWrapper: class BrittleStarJaxEnvWrapper:
@ -76,3 +77,21 @@ class BrittleStarJaxEnvWrapper:
return BrittleStarJaxEnvWrapper( return BrittleStarJaxEnvWrapper(
morphology, arena, env_config, num_envs=num_envs, backend=backend morphology, arena, env_config, num_envs=num_envs, backend=backend
) )
@staticmethod
def from_config(
config_path: str, num_envs: int, backend: Backend = Backend.MJX
) -> "BrittleStarJaxEnvWrapper":
morphology_cfg, arena_cfg, env_cfg = from_json(config_path)
return BrittleStarJaxEnvWrapper(
morphology_cfg, arena_cfg, env_cfg, num_envs=num_envs, backend=backend
)
def __str__(self):
morphology_str = str(self._morphology)
arena_str = str(self._arena)
env_config_str = str(self._env_config)
return (
f"BrittleStarJaxEnvWrapper(backend={self._backend}, num_envs={self._num_envs}, "
+ f"morphology={morphology_str}, arena={arena_str}, env_config={env_config_str})"
)

View file

@ -1,4 +1,4 @@
from .env_config import ArenaConfig, EnvConfig, MorphologyConfig from .env_config import ArenaConfig, EnvConfig, MorphologyConfig, from_json
from .env_types import Backend, Task from .env_types import Backend, Task
from .env_wrapper import BrittleStarEnv, StepResult from .env_wrapper import BrittleStarEnv, StepResult
from .factory import BrittleStarEnvFactory from .factory import BrittleStarEnvFactory
@ -12,4 +12,5 @@ __all__ = [
"BrittleStarEnv", "BrittleStarEnv",
"StepResult", "StepResult",
"BrittleStarEnvFactory", "BrittleStarEnvFactory",
"from_json",
] ]

View file

@ -1,6 +1,7 @@
from __future__ import annotations from __future__ import annotations
from dataclasses import dataclass, field from dataclasses import dataclass, field
import json
from .env_types import Task from .env_types import Task
@ -48,6 +49,11 @@ class EnvConfig:
# Per docs in upstream env config: integer factors of 200. # Per docs in upstream env config: integer factors of 200.
light_perlin_noise_scale: int = 0 light_perlin_noise_scale: int = 0
@staticmethod
def from_json(path: str) -> EnvConfig: def from_json(path: str) -> tuple[MorphologyConfig, ArenaConfig, EnvConfig]:
pass with open(path, "r") as f:
config_json = json.load(f)
morphology = MorphologyConfig(**config_json.get("morphology", {}))
arena = ArenaConfig(**config_json.get("arena", {}))
env = EnvConfig(**config_json.get("env", {}))
return morphology, arena, env

View file

@ -0,0 +1,3 @@
from .renderer import simulate_policy, SimulationConfig, ControlPolicy
__all__ = ["simulate_policy", "SimulationConfig", "ControlPolicy"]

View file

@ -52,6 +52,7 @@ class AgentParams:
network_params: flax.core.FrozenDict network_params: flax.core.FrozenDict
actor_params: flax.core.FrozenDict actor_params: flax.core.FrozenDict
critic_params: flax.core.FrozenDict critic_params: flax.core.FrozenDict
critic_network_params: flax.core.FrozenDict
@jax.tree_util.register_dataclass @jax.tree_util.register_dataclass

View file

@ -7,6 +7,7 @@ from typing import Callable
import flax import flax
import jax import jax
import jax.numpy as jnp import jax.numpy as jnp
import matplotlib.pyplot as plt
import numpy as np import numpy as np
import optax import optax
import torch import torch
@ -18,7 +19,8 @@ from torch.utils.tensorboard import SummaryWriter
from brittle_star_project.dataclasses import PPOArgs from brittle_star_project.dataclasses import PPOArgs
from brittle_star_project.dataclasses.EpisodeStatistics import EpisodeStatistics from brittle_star_project.dataclasses.EpisodeStatistics import EpisodeStatistics
from brittle_star_project.environment.BrittleStarJaxEnvWrapper import BrittleStarJaxEnvWrapper from brittle_star_project.environment.BrittleStarJaxEnvWrapper import BrittleStarJaxEnvWrapper
from brittle_star_project.rl import Network, Actor, Critic, AgentParams, Storage from brittle_star_project.rl import Actor, AgentParams, Critic, Network, Storage
from ppo import PPO
def convert_obs_dict_to_array(obs_dict: dict) -> jnp.ndarray: def convert_obs_dict_to_array(obs_dict: dict) -> jnp.ndarray:
@ -27,9 +29,11 @@ def convert_obs_dict_to_array(obs_dict: dict) -> jnp.ndarray:
) )
def make_env(num_envs: int) -> Callable: def make_env(config_path: str | None, num_envs: int) -> Callable:
def thunk(): def thunk():
if config_path is None:
return BrittleStarJaxEnvWrapper.default(num_envs=num_envs) return BrittleStarJaxEnvWrapper.default(num_envs=num_envs)
return BrittleStarJaxEnvWrapper.from_config(config_path, num_envs=num_envs)
return thunk return thunk
@ -62,14 +66,15 @@ def train(args: PPOArgs):
random.seed(args.seed) random.seed(args.seed)
np.random.seed(args.seed) np.random.seed(args.seed)
key = jax.random.PRNGKey(args.seed) key = jax.random.PRNGKey(args.seed)
key, network_key, actor_key, critic_key = jax.random.split(key, 4) key, network_key, actor_key, critic_key, critic_network_key = jax.random.split(key, 5)
torch.backends.cudnn.deterministic = args.torch_deterministic torch.backends.cudnn.deterministic = args.torch_deterministic
device = torch.device("cuda" if torch.cuda.is_available() and args.cuda else "cpu") device = torch.device("cuda" if torch.cuda.is_available() and args.cuda else "cpu")
print(f"Running on device: {device}") print(f"Running on device: {device}")
print("Creating the environment...") print("Creating the environment...")
env = make_env(num_envs=args.num_envs)() env = make_env(config_path=args.config_path, num_envs=args.num_envs)()
print(f"Environment: {env}")
episode_stats = EpisodeStatistics( episode_stats = EpisodeStatistics(
episode_returns=jnp.zeros(args.num_envs, dtype=jnp.float32), episode_returns=jnp.zeros(args.num_envs, dtype=jnp.float32),
@ -112,6 +117,7 @@ def train(args: PPOArgs):
print("Initializing the models...") print("Initializing the models...")
network = Network() network = Network()
critic_network = Network()
actor = Actor(action_dim=env.single_action_space.shape[0]) # continuous actions for MJX actor = Actor(action_dim=env.single_action_space.shape[0]) # continuous actions for MJX
critic = Critic() critic = Critic()
@ -123,12 +129,15 @@ def train(args: PPOArgs):
] ]
) )
network_params = network.init(network_key, sample_obs) network_params = network.init(network_key, sample_obs)
critic_network_params = critic_network.init(critic_network_key, sample_obs)
actor_params = actor.init(actor_key, network.apply(network_params, sample_obs)) actor_params = actor.init(actor_key, network.apply(network_params, sample_obs))
critic_params = critic.init(critic_key, network.apply(network_params, sample_obs)) critic_params = critic.init(critic_key, critic_network.apply(critic_network_params, sample_obs))
agent_state = TrainState.create( agent_state = TrainState.create(
apply_fn=None, apply_fn=None,
params=asdict(AgentParams(network_params, actor_params, critic_params)), params=asdict(
AgentParams(network_params, actor_params, critic_params, critic_network_params)
),
tx=optax.chain( tx=optax.chain(
optax.clip_by_global_norm(args.max_grad_norm), optax.clip_by_global_norm(args.max_grad_norm),
optax.inject_hyperparams(optax.adam)( optax.inject_hyperparams(optax.adam)(
@ -138,8 +147,10 @@ def train(args: PPOArgs):
) )
network.apply = jax.jit(network.apply) network.apply = jax.jit(network.apply)
critic_network.apply = jax.jit(critic_network.apply)
actor.apply = jax.jit(actor.apply) actor.apply = jax.jit(actor.apply)
critic.apply = jax.jit(critic.apply) critic.apply = jax.jit(critic.apply)
ppo_instance = PPO(args, network, actor, critic, critic_network)
@jax.jit @jax.jit
def get_action_and_value_noise( def get_action_and_value_noise(
@ -158,20 +169,6 @@ def train(args: PPOArgs):
value = critic.apply(agent_state.params["critic_params"], hidden) value = critic.apply(agent_state.params["critic_params"], hidden)
return action, logprob, value.squeeze(-1), key return action, logprob, value.squeeze(-1), key
@jax.jit
def get_action_and_value(
params: flax.core.FrozenDict,
x: jnp.ndarray,
action: np.ndarray,
):
hidden = network.apply(params["network_params"], x)
mean, log_std = actor.apply(params["actor_params"], hidden)
std = jnp.exp(log_std)
logprob = -0.5 * (((action - mean) / std) ** 2 + 2 * log_std + jnp.log(2 * jnp.pi)).sum(-1)
entropy = (0.5 + 0.5 * jnp.log(2 * jnp.pi) + log_std).sum(-1)
value = critic.apply(params["critic_params"], hidden).squeeze(-1)
return logprob, entropy, value
@jax.jit @jax.jit
def compute_gae_once(carry, inp, gamma, gae_lambda): def compute_gae_once(carry, inp, gamma, gae_lambda):
advantages = carry advantages = carry
@ -199,61 +196,6 @@ def train(args: PPOArgs):
) )
return storage.replace(advantages=advantages, returns=advantages + storage.values) return storage.replace(advantages=advantages, returns=advantages + storage.values)
def ppo_loss(params, x, a, logp, mb_advantages, mb_returns):
newlogprob, entropy, newvalue = get_action_and_value(params, x, a)
logratio = newlogprob - logp
ratio = jnp.exp(logratio)
approx_kl = ((ratio - 1) - logratio).mean()
if args.norm_adv:
mb_advantages = (mb_advantages - mb_advantages.mean()) / (mb_advantages.std() + 1e-8)
pg_loss1 = -mb_advantages * ratio
pg_loss2 = -mb_advantages * jnp.clip(ratio, 1 - args.clip_coef, 1 + args.clip_coef)
pg_loss = jnp.maximum(pg_loss1, pg_loss2).mean()
v_loss = 0.5 * ((newvalue - mb_returns) ** 2).mean()
entropy_loss = entropy.mean()
loss = pg_loss - args.ent_coef * entropy_loss + v_loss * args.vf_coef
return loss, (pg_loss, v_loss, entropy_loss, jax.lax.stop_gradient(approx_kl))
ppo_loss_grad_fn = jax.value_and_grad(ppo_loss, has_aux=True)
@jax.jit
def update_ppo(agent_state, storage, key):
def update_epoch(carry, _):
agent_state, key = carry
key, subkey = jax.random.split(key)
def flatten(x):
return x.reshape((-1,) + x.shape[2:])
def convert_data(x):
x = jax.random.permutation(subkey, x)
return jnp.reshape(x, (args.num_minibatches, -1) + x.shape[1:])
flatten_storage = jax.tree.map(flatten, storage)
shuffled_storage = jax.tree.map(convert_data, flatten_storage)
def update_minibatch(agent_state, minibatch):
(loss, (pg_loss, v_loss, entropy_loss, approx_kl)), grads = ppo_loss_grad_fn(
agent_state.params,
minibatch.obs,
minibatch.actions,
minibatch.logprobs,
minibatch.advantages,
minibatch.returns,
)
agent_state = agent_state.apply_gradients(grads=grads)
return agent_state, (loss, pg_loss, v_loss, entropy_loss, approx_kl, grads)
agent_state, metrics = jax.lax.scan(update_minibatch, agent_state, shuffled_storage)
return (agent_state, key), metrics
(agent_state, key), (loss, pg_loss, v_loss, entropy_loss, approx_kl, grads) = jax.lax.scan(
update_epoch, (agent_state, key), (), length=args.update_epochs
)
return agent_state, loss, pg_loss, v_loss, entropy_loss, approx_kl, key
# --- Main training loop --- # --- Main training loop ---
global_step = 0 global_step = 0
start_time = time.time() start_time = time.time()
@ -303,6 +245,7 @@ def train(args: PPOArgs):
print("Starting training...") print("Starting training...")
iters_bar = tqdm.tqdm(range(1, args.num_iterations + 1)) iters_bar = tqdm.tqdm(range(1, args.num_iterations + 1))
losses = []
for _ in iters_bar: for _ in iters_bar:
iteration_time_start = time.time() iteration_time_start = time.time()
@ -312,10 +255,12 @@ def train(args: PPOArgs):
global_step += args.num_steps * args.num_envs global_step += args.num_steps * args.num_envs
storage = compute_gae(agent_state, next_obs, next_done, storage) storage = compute_gae(agent_state, next_obs, next_done, storage)
agent_state, loss, pg_loss, v_loss, entropy_loss, approx_kl, key = update_ppo( agent_state, loss, pg_loss, v_loss, entropy_loss, approx_kl, key = ppo_instance.update_ppo(
agent_state, storage, key agent_state, storage, key
) )
losses.append(jnp.mean(loss))
avg_episodic_return = np.mean(jax.device_get(episode_stats.returned_episode_returns)) avg_episodic_return = np.mean(jax.device_get(episode_stats.returned_episode_returns))
iters_bar.set_postfix_str( iters_bar.set_postfix_str(
f"global_step={global_step}, avg_episodic_return={avg_episodic_return}" f"global_step={global_step}, avg_episodic_return={avg_episodic_return}"
@ -367,6 +312,12 @@ def train(args: PPOArgs):
env.close() env.close()
writer.close() writer.close()
print("Saving loss plot...")
plt.plot(losses)
plt.title("PPO Loss, mean over minibatches")
plt.savefig(f"runs/{run_name}/{args.exp_name}_losses.png")
plt.close()
def main() -> None: def main() -> None:
args = tyro.cli(PPOArgs) args = tyro.cli(PPOArgs)