From 1ad940ce024fd74d66bc20fb08be0cecdadc5faa Mon Sep 17 00:00:00 2001 From: cmekeirl Date: Fri, 20 Mar 2026 10:52:33 +0100 Subject: [PATCH 01/11] Initial proposal for structure --- src/MLPs/centralized.py | 0 src/MLPs/fully_arm.py | 0 src/MLPs/helpers.py | 0 src/MLPs/ring_arm.py | 0 src/MLPs/segment.py | 0 5 files changed, 0 insertions(+), 0 deletions(-) create mode 100644 src/MLPs/centralized.py create mode 100644 src/MLPs/fully_arm.py create mode 100644 src/MLPs/helpers.py create mode 100644 src/MLPs/ring_arm.py create mode 100644 src/MLPs/segment.py diff --git a/src/MLPs/centralized.py b/src/MLPs/centralized.py new file mode 100644 index 0000000..e69de29 diff --git a/src/MLPs/fully_arm.py b/src/MLPs/fully_arm.py new file mode 100644 index 0000000..e69de29 diff --git a/src/MLPs/helpers.py b/src/MLPs/helpers.py new file mode 100644 index 0000000..e69de29 diff --git a/src/MLPs/ring_arm.py b/src/MLPs/ring_arm.py new file mode 100644 index 0000000..e69de29 diff --git a/src/MLPs/segment.py b/src/MLPs/segment.py new file mode 100644 index 0000000..e69de29 From 4e3eee8ac6c8a5c0fa2eeb7ab6ca56557caff4c5 Mon Sep 17 00:00:00 2001 From: cmekeirl Date: Fri, 20 Mar 2026 11:29:35 +0100 Subject: [PATCH 02/11] Basic centralized model --- src/MLPs/centralized.py | 30 ++++++++++++++++++++++++++++++ 1 file changed, 30 insertions(+) diff --git a/src/MLPs/centralized.py b/src/MLPs/centralized.py index e69de29..70abbec 100644 --- a/src/MLPs/centralized.py +++ b/src/MLPs/centralized.py @@ -0,0 +1,30 @@ +from typing import Sequence +import flax.linen as nn +import jax.numpy as jnp +from flax.linen.initializers import constant, orthogonal + + +class Network(nn.Module): + hidden_size: int = 256 + @nn.compact + def __call__(self, x): + # x shape: (batch, obs_dim) + x = nn.Dense(self.hidden_size, kernel_init=orthogonal(jnp.sqrt(2)))(x) + x = nn.tanh(x) + x = nn.Dense(self.hidden_size, kernel_init=orthogonal(jnp.sqrt(2)))(x) + x = nn.tanh(x) + return x + + +class Critic(nn.Module): + @nn.compact + def __call__(self, x): + return nn.Dense(1, kernel_init=orthogonal(1), bias_init=constant(0.0))(x) + + +class Actor(nn.Module): + action_dim: Sequence[int] + + @nn.compact + def __call__(self, x): + return nn.Dense(self.action_dim, kernel_init=orthogonal(0.01), bias_init=constant(0.0))(x) From 2263c97fe8e143566e3a4694f7a22bb6a0f56a8f Mon Sep 17 00:00:00 2001 From: Tibo De Peuter Date: Wed, 1 Apr 2026 17:13:08 +0000 Subject: [PATCH 03/11] style: ruff format --- src/MLPs/centralized.py | 1 + 1 file changed, 1 insertion(+) diff --git a/src/MLPs/centralized.py b/src/MLPs/centralized.py index 70abbec..14e3a22 100644 --- a/src/MLPs/centralized.py +++ b/src/MLPs/centralized.py @@ -6,6 +6,7 @@ from flax.linen.initializers import constant, orthogonal class Network(nn.Module): hidden_size: int = 256 + @nn.compact def __call__(self, x): # x shape: (batch, obs_dim) From 15ae86796f63921aa8d86bd339bef517d96f636b Mon Sep 17 00:00:00 2001 From: cedric Date: Thu, 2 Apr 2026 17:38:46 +0000 Subject: [PATCH 04/11] feat: generic network --- src/MLPs/fully_arm.py | 0 src/MLPs/helpers.py | 0 src/MLPs/{centralized.py => mlps.py} | 17 +++++++++-------- src/MLPs/ring_arm.py | 0 src/MLPs/segment.py | 0 5 files changed, 9 insertions(+), 8 deletions(-) delete mode 100644 src/MLPs/fully_arm.py delete mode 100644 src/MLPs/helpers.py rename src/MLPs/{centralized.py => mlps.py} (50%) delete mode 100644 src/MLPs/ring_arm.py delete mode 100644 src/MLPs/segment.py diff --git a/src/MLPs/fully_arm.py b/src/MLPs/fully_arm.py deleted file mode 100644 index e69de29..0000000 diff --git a/src/MLPs/helpers.py b/src/MLPs/helpers.py deleted file mode 100644 index e69de29..0000000 diff --git a/src/MLPs/centralized.py b/src/MLPs/mlps.py similarity index 50% rename from src/MLPs/centralized.py rename to src/MLPs/mlps.py index 14e3a22..f97452d 100644 --- a/src/MLPs/centralized.py +++ b/src/MLPs/mlps.py @@ -1,19 +1,20 @@ -from typing import Sequence +from typing import Sequence, Callable import flax.linen as nn import jax.numpy as jnp from flax.linen.initializers import constant, orthogonal -class Network(nn.Module): - hidden_size: int = 256 +# example usage: network = SemiGenericNetwork(layer_sizes=[256, 256], activation=nn.relu) +# semi generic so we can easily make a config for it in experiments +class SemiGenericNetwork(nn.Module): + layer_sizes: Sequence[int] = [64, 64] # default 2 layers of 64 neurons + activation: Callable = nn.tanh # default tanh @nn.compact def __call__(self, x): - # x shape: (batch, obs_dim) - x = nn.Dense(self.hidden_size, kernel_init=orthogonal(jnp.sqrt(2)))(x) - x = nn.tanh(x) - x = nn.Dense(self.hidden_size, kernel_init=orthogonal(jnp.sqrt(2)))(x) - x = nn.tanh(x) + for size in self.layer_sizes: + x = nn.Dense(size, kernel_init=orthogonal(jnp.sqrt(2)))(x) + x = self.activation(x) return x diff --git a/src/MLPs/ring_arm.py b/src/MLPs/ring_arm.py deleted file mode 100644 index e69de29..0000000 diff --git a/src/MLPs/segment.py b/src/MLPs/segment.py deleted file mode 100644 index e69de29..0000000 From 9879fea9e221c0839acb4260f1de1b5be46a8992 Mon Sep 17 00:00:00 2001 From: cedric Date: Thu, 2 Apr 2026 17:50:44 +0000 Subject: [PATCH 05/11] fix: refactored rl folder into mlps --- src/MLPs/mlps.py | 32 +++- src/brittle_star_project/rl/DummyAgent.py | 72 -------- src/brittle_star_project/rl/__init__.py | 25 --- src/brittle_star_project/rl/base.py | 162 ------------------ .../rl/random_policy_model.py | 52 ------ src/train.py | 6 +- 6 files changed, 34 insertions(+), 315 deletions(-) delete mode 100644 src/brittle_star_project/rl/DummyAgent.py delete mode 100644 src/brittle_star_project/rl/__init__.py delete mode 100644 src/brittle_star_project/rl/base.py delete mode 100644 src/brittle_star_project/rl/random_policy_model.py diff --git a/src/MLPs/mlps.py b/src/MLPs/mlps.py index f97452d..e4b5969 100644 --- a/src/MLPs/mlps.py +++ b/src/MLPs/mlps.py @@ -1,6 +1,10 @@ -from typing import Sequence, Callable +from dataclasses import dataclass, fields + +import flax import flax.linen as nn import jax.numpy as jnp +import jax.tree_util +from typing import Sequence, Callable from flax.linen.initializers import constant, orthogonal @@ -30,3 +34,29 @@ class Actor(nn.Module): @nn.compact def __call__(self, x): return nn.Dense(self.action_dim, kernel_init=orthogonal(0.01), bias_init=constant(0.0))(x) + + +@jax.tree_util.register_dataclass +@dataclass +class AgentParams: + network_params: flax.core.FrozenDict + actor_params: flax.core.FrozenDict + critic_params: flax.core.FrozenDict + critic_network_params: flax.core.FrozenDict + + +@jax.tree_util.register_dataclass +@dataclass +class Storage: + obs: jnp.array + actions: jnp.array + logprobs: jnp.array + dones: jnp.array + values: jnp.array + advantages: jnp.array + returns: jnp.array + rewards: jnp.array + + def replace(self, **kwargs) -> "Storage": + fs = fields(self) + return Storage(**{f.name: kwargs.get(f.name, getattr(self, f.name)) for f in fs}) diff --git a/src/brittle_star_project/rl/DummyAgent.py b/src/brittle_star_project/rl/DummyAgent.py deleted file mode 100644 index bc027dc..0000000 --- a/src/brittle_star_project/rl/DummyAgent.py +++ /dev/null @@ -1,72 +0,0 @@ -from dataclasses import dataclass, fields - -import flax -import flax.linen as nn -import jax.numpy as jnp -import jax.tree_util -import numpy as np -from flax.linen.initializers import constant, orthogonal - - -class Network(nn.Module): - """ - Dummy model only used for testing purposes - - inspired by: https://github.com/vwxyzjn/cleanrl/blob/master/cleanrl/ppo_atari_envpool_xla_jax_scan.py - """ - - hidden_dim: int = 195 - - @nn.compact - def __call__(self, x): - x = nn.Dense(self.hidden_dim, kernel_init=orthogonal(np.sqrt(2)), bias_init=constant(0.0))( - x - ) - x = nn.relu(x) - x = nn.Dense(self.hidden_dim, kernel_init=orthogonal(np.sqrt(2)), bias_init=constant(0.0))( - x - ) - x = nn.relu(x) - return x - - -class Critic(nn.Module): - @nn.compact - def __call__(self, x): - return nn.Dense(1, kernel_init=orthogonal(1), bias_init=constant(0.0))(x) - - -class Actor(nn.Module): - action_dim: int - - @nn.compact - def __call__(self, x): - mean = nn.Dense(self.action_dim, kernel_init=orthogonal(0.01), bias_init=constant(0.0))(x) - log_std = self.param("log_std", nn.initializers.zeros, (self.action_dim,)) - return mean, log_std - - -@jax.tree_util.register_dataclass -@dataclass -class AgentParams: - network_params: flax.core.FrozenDict - actor_params: flax.core.FrozenDict - critic_params: flax.core.FrozenDict - critic_network_params: flax.core.FrozenDict - - -@jax.tree_util.register_dataclass -@dataclass -class Storage: - obs: jnp.array - actions: jnp.array - logprobs: jnp.array - dones: jnp.array - values: jnp.array - advantages: jnp.array - returns: jnp.array - rewards: jnp.array - - def replace(self, **kwargs) -> "Storage": - fs = fields(self) - return Storage(**{f.name: kwargs.get(f.name, getattr(self, f.name)) for f in fs}) diff --git a/src/brittle_star_project/rl/__init__.py b/src/brittle_star_project/rl/__init__.py deleted file mode 100644 index e58e1c6..0000000 --- a/src/brittle_star_project/rl/__init__.py +++ /dev/null @@ -1,25 +0,0 @@ -from .DummyAgent import Network, Critic, Actor, AgentParams, Storage -from .base import ( - RLAlgorithm, - RLModel, - Transition, - create_model, - register_rl_model, - registered_model_types, -) -from .random_policy_model import RandomPolicyModel - -__all__ = [ - "RLAlgorithm", - "RLModel", - "RandomPolicyModel", - "Transition", - "create_model", - "register_rl_model", - "registered_model_types", - "Network", - "Critic", - "Actor", - "AgentParams", - "Storage", -] diff --git a/src/brittle_star_project/rl/base.py b/src/brittle_star_project/rl/base.py deleted file mode 100644 index 8e3c439..0000000 --- a/src/brittle_star_project/rl/base.py +++ /dev/null @@ -1,162 +0,0 @@ -from __future__ import annotations - -from abc import ABC, abstractmethod -from dataclasses import dataclass -import json -from pathlib import Path -from typing import Any - - -@dataclass(frozen=True, slots=True) -class Transition: - """A minimal transition container for RL. - - This is intentionally generic because the underlying env state type may be a - JAX pytree, a numpy struct, or something library-specific. - """ - - obs: Any - action: Any - reward: float - next_obs: Any - terminated: bool - truncated: bool - info: dict[str, Any] | None = None - - -class RLAlgorithm(ABC): - """Insertable RL algorithm interface.""" - - @abstractmethod - def select_action(self, *, obs: Any, rng: Any | None = None) -> Any: - raise NotImplementedError - - def observe(self, transition: Transition) -> None: - """Optional hook to store transitions.""" - - def update(self, *, rng: Any | None = None) -> dict[str, float]: - """Optional hook to run one training update.""" - - return {} - - def save(self, path: str) -> None: - raise NotImplementedError("Save not implemented") - - def load(self, path: str) -> None: - raise NotImplementedError("Load not implemented") - - -_RL_MODEL_REGISTRY: dict[str, type["RLModel"]] = {} - - -def registered_model_types() -> list[str]: - return sorted(_RL_MODEL_REGISTRY) - - -def create_model(type_name: str, *, payload: dict[str, Any]) -> "RLModel": - model_cls = _RL_MODEL_REGISTRY.get(type_name) - if model_cls is None: - known = ", ".join(sorted(_RL_MODEL_REGISTRY)) or "" - raise ValueError(f"Unknown RLModel type '{type_name}'. Known: {known}") - return model_cls.from_payload(payload) - - -def get_rl_model_registry() -> dict[str, type["RLModel"]]: - """Return a copy of the current RLModel registry. - - The registry is populated by importing concrete model modules that use the - `@register_rl_model(...)` decorator. - """ - - return dict(_RL_MODEL_REGISTRY) - - -def register_rl_model(*type_names: str): - """Decorator to register an `RLModel` for generic loading. - - Concrete model modules should apply this decorator, so `base.py` never needs - to import concrete models (avoids circular imports). - """ - - if not type_names: - raise TypeError("register_rl_model() requires at least one type name") - - primary = type_names[0] - - def _decorator(cls: type[RLModel]): - for name in type_names: - _RL_MODEL_REGISTRY[name] = cls - cls.type_name = primary - return cls - - return _decorator - - -class RLModel(ABC): - """Serializable policy/model interface. - - This is the artifact that `train.py` writes and `simulate.py` loads. - """ - - # Overwritten by the `@register_rl_model(...)` decorator. - type_name: str = "RLModel" - - def reset(self, seed: int | None = None) -> None: - """Optional hook for RNG/stateful models.""" - - @abstractmethod - def act(self, *, obs: Any | None = None, t: float = 0.0) -> Any: - raise NotImplementedError - - def train(self, *, env: Any, num_epochs: int = 1) -> None: - """Optional training hook. - - Many models won't learn; for those this can be a no-op. - """ - - _ = (env, num_epochs) - - def to_payload(self) -> dict[str, Any]: - """Return JSON-serializable model parameters.""" - - return {} - - @classmethod - def from_payload(cls, payload: dict[str, Any]) -> "RLModel": - """Reconstruct a model from `to_payload()` output.""" - - return cls(**payload) # type: ignore[arg-type] - - def save(self, path: str | Path) -> Path: - out = Path(path) - out.parent.mkdir(parents=True, exist_ok=True) - doc = { - "type": self.type_name, - "version": 1, - "payload": self.to_payload(), - } - out.write_text(json.dumps(doc, indent=2, sort_keys=True) + "\n") - return out - - @classmethod - def load(cls, path: str | Path) -> "RLModel": - p = Path(path) - doc = json.loads(p.read_text()) - - type_name = doc.get("type") - if not isinstance(type_name, str): - raise ValueError("Model artifact missing string field 'type'") - - model_cls = _RL_MODEL_REGISTRY.get(type_name) - if model_cls is None: - known = ", ".join(sorted(_RL_MODEL_REGISTRY)) or "" - raise ValueError(f"Unknown RLModel type '{type_name}'. Known: {known}") - - payload = doc.get("payload") - # Backward compatibility: older artifacts stored fields at top-level. - if payload is None: - payload = {k: v for k, v in doc.items() if k not in ("type", "version")} - if not isinstance(payload, dict): - raise ValueError("Model artifact field 'payload' must be an object") - - return model_cls.from_payload(payload) diff --git a/src/brittle_star_project/rl/random_policy_model.py b/src/brittle_star_project/rl/random_policy_model.py deleted file mode 100644 index 9372d49..0000000 --- a/src/brittle_star_project/rl/random_policy_model.py +++ /dev/null @@ -1,52 +0,0 @@ -from __future__ import annotations - -from dataclasses import dataclass, field - -import numpy as np - -from .base import RLModel, register_rl_model - - -@register_rl_model("random") -@dataclass(slots=True) -class RandomPolicyModel(RLModel): - """A minimal, serializable policy model that outputs random controls. - - This is intentionally *not* a learning algorithm yet. It exists so we can: - - produce a stable model artifact from `train.py` - - load that artifact in `simulate.py` - - drive the MuJoCo viewer with the model's actions - """ - - nu: int = 0 - seed: int = 0 - ctrl_noise_scale: float = 0.5 - - _rng: np.random.RandomState = field(init=False, repr=False) - - def __post_init__(self) -> None: - self.reset(self.seed) - - def reset(self, seed: int | None = None) -> None: - if seed is not None: - self.seed = int(seed) - self._rng = np.random.RandomState(self.seed) - - def act(self, *, obs: np.ndarray | None = None, t: float = 0.0) -> np.ndarray: - if self.nu <= 0: - return np.zeros((0,), dtype=np.float32) - ctrl = self.ctrl_noise_scale * self._rng.randn(self.nu) - return ctrl.astype(np.float32) - - def to_payload(self) -> dict[str, object]: - return { - "seed": int(self.seed), - "ctrl_noise_scale": float(self.ctrl_noise_scale), - } - - @classmethod - def from_payload(cls, payload: dict[str, object]) -> RandomPolicyModel: - return cls( - seed=int(payload.get("seed", 0)), - ctrl_noise_scale=float(payload.get("ctrl_noise_scale", 0.5)), - ) diff --git a/src/train.py b/src/train.py index ac6849e..89b85d5 100644 --- a/src/train.py +++ b/src/train.py @@ -19,7 +19,7 @@ from torch.utils.tensorboard import SummaryWriter from brittle_star_project.dataclasses import PPOArgs from brittle_star_project.dataclasses.EpisodeStatistics import EpisodeStatistics from brittle_star_project.environment.BrittleStarJaxEnvWrapper import BrittleStarJaxEnvWrapper -from brittle_star_project.rl import Actor, AgentParams, Critic, Network, Storage +from MLPs.mlps import SemiGenericNetwork, Actor, Critic, AgentParams, Storage from ppo import PPO @@ -116,8 +116,8 @@ def train(args: PPOArgs): return args.learning_rate * frac print("Initializing the models...") - network = Network() - critic_network = Network() + network = SemiGenericNetwork() + critic_network = SemiGenericNetwork() actor = Actor(action_dim=env.single_action_space.shape[0]) # continuous actions for MJX critic = Critic() From 95840b789b01a7e4ef6ef7484d8a67e99c4fb164 Mon Sep 17 00:00:00 2001 From: cedric Date: Thu, 2 Apr 2026 17:59:48 +0000 Subject: [PATCH 06/11] fix: usage of field because flax wont allow mutable class object --- src/MLPs/mlps.py | 12 +++++++----- 1 file changed, 7 insertions(+), 5 deletions(-) diff --git a/src/MLPs/mlps.py b/src/MLPs/mlps.py index e4b5969..f49ac5c 100644 --- a/src/MLPs/mlps.py +++ b/src/MLPs/mlps.py @@ -1,4 +1,4 @@ -from dataclasses import dataclass, fields +from dataclasses import dataclass, fields, field import flax import flax.linen as nn @@ -11,8 +11,8 @@ from flax.linen.initializers import constant, orthogonal # example usage: network = SemiGenericNetwork(layer_sizes=[256, 256], activation=nn.relu) # semi generic so we can easily make a config for it in experiments class SemiGenericNetwork(nn.Module): - layer_sizes: Sequence[int] = [64, 64] # default 2 layers of 64 neurons - activation: Callable = nn.tanh # default tanh + layer_sizes: Sequence[int] = field(default_factory=lambda: [64, 64]) + activation: Callable = nn.tanh @nn.compact def __call__(self, x): @@ -29,11 +29,13 @@ class Critic(nn.Module): class Actor(nn.Module): - action_dim: Sequence[int] + action_dim: int @nn.compact def __call__(self, x): - return nn.Dense(self.action_dim, kernel_init=orthogonal(0.01), bias_init=constant(0.0))(x) + mean = nn.Dense(self.action_dim, kernel_init=orthogonal(0.01), bias_init=constant(0.0))(x) + log_std = self.param("log_std", nn.initializers.zeros, (self.action_dim,)) + return mean, log_std @jax.tree_util.register_dataclass From 729f122be47580df85d11af99c254f02f4f07211 Mon Sep 17 00:00:00 2001 From: cedric Date: Thu, 2 Apr 2026 18:00:53 +0000 Subject: [PATCH 07/11] fix: removed useless comment --- src/MLPs/mlps.py | 1 - 1 file changed, 1 deletion(-) diff --git a/src/MLPs/mlps.py b/src/MLPs/mlps.py index f49ac5c..6c587d3 100644 --- a/src/MLPs/mlps.py +++ b/src/MLPs/mlps.py @@ -8,7 +8,6 @@ from typing import Sequence, Callable from flax.linen.initializers import constant, orthogonal -# example usage: network = SemiGenericNetwork(layer_sizes=[256, 256], activation=nn.relu) # semi generic so we can easily make a config for it in experiments class SemiGenericNetwork(nn.Module): layer_sizes: Sequence[int] = field(default_factory=lambda: [64, 64]) From 40c7e32344a727e364e2830549f3120826d7353a Mon Sep 17 00:00:00 2001 From: cedric Date: Thu, 2 Apr 2026 18:11:49 +0000 Subject: [PATCH 08/11] feat: more generic naming for critic, since messager will use the same --- src/MLPs/mlps.py | 2 +- src/train.py | 5 +++-- 2 files changed, 4 insertions(+), 3 deletions(-) diff --git a/src/MLPs/mlps.py b/src/MLPs/mlps.py index 6c587d3..1bae0c4 100644 --- a/src/MLPs/mlps.py +++ b/src/MLPs/mlps.py @@ -21,7 +21,7 @@ class SemiGenericNetwork(nn.Module): return x -class Critic(nn.Module): +class OneDenseLayerMLP(nn.Module): @nn.compact def __call__(self, x): return nn.Dense(1, kernel_init=orthogonal(1), bias_init=constant(0.0))(x) diff --git a/src/train.py b/src/train.py index 89b85d5..7528f26 100644 --- a/src/train.py +++ b/src/train.py @@ -19,7 +19,7 @@ from torch.utils.tensorboard import SummaryWriter from brittle_star_project.dataclasses import PPOArgs from brittle_star_project.dataclasses.EpisodeStatistics import EpisodeStatistics from brittle_star_project.environment.BrittleStarJaxEnvWrapper import BrittleStarJaxEnvWrapper -from MLPs.mlps import SemiGenericNetwork, Actor, Critic, AgentParams, Storage +from MLPs.mlps import SemiGenericNetwork, Actor, OneDenseLayerMLP, AgentParams, Storage from ppo import PPO @@ -119,7 +119,8 @@ def train(args: PPOArgs): network = SemiGenericNetwork() critic_network = SemiGenericNetwork() actor = Actor(action_dim=env.single_action_space.shape[0]) # continuous actions for MJX - critic = Critic() + critic = OneDenseLayerMLP() + # messager = OneDenseLayerMLP() sample_obs = jnp.concatenate( [ From 79221372dfca74b10ce74326900a7a963e059013 Mon Sep 17 00:00:00 2001 From: cedric Date: Thu, 2 Apr 2026 18:18:10 +0000 Subject: [PATCH 09/11] fix: uniform generic network naming --- src/MLPs/mlps.py | 2 +- src/train.py | 12 +++++++++--- 2 files changed, 10 insertions(+), 4 deletions(-) diff --git a/src/MLPs/mlps.py b/src/MLPs/mlps.py index 1bae0c4..1fdc6a2 100644 --- a/src/MLPs/mlps.py +++ b/src/MLPs/mlps.py @@ -9,7 +9,7 @@ from flax.linen.initializers import constant, orthogonal # semi generic so we can easily make a config for it in experiments -class SemiGenericNetwork(nn.Module): +class GenericDenseLayersWithActivation(nn.Module): layer_sizes: Sequence[int] = field(default_factory=lambda: [64, 64]) activation: Callable = nn.tanh diff --git a/src/train.py b/src/train.py index 7528f26..10f14bf 100644 --- a/src/train.py +++ b/src/train.py @@ -19,7 +19,13 @@ from torch.utils.tensorboard import SummaryWriter from brittle_star_project.dataclasses import PPOArgs from brittle_star_project.dataclasses.EpisodeStatistics import EpisodeStatistics from brittle_star_project.environment.BrittleStarJaxEnvWrapper import BrittleStarJaxEnvWrapper -from MLPs.mlps import SemiGenericNetwork, Actor, OneDenseLayerMLP, AgentParams, Storage +from MLPs.mlps import ( + GenericDenseLayersWithActivation, + Actor, + OneDenseLayerMLP, + AgentParams, + Storage, +) from ppo import PPO @@ -116,8 +122,8 @@ def train(args: PPOArgs): return args.learning_rate * frac print("Initializing the models...") - network = SemiGenericNetwork() - critic_network = SemiGenericNetwork() + network = GenericDenseLayersWithActivation() + critic_network = GenericDenseLayersWithActivation() actor = Actor(action_dim=env.single_action_space.shape[0]) # continuous actions for MJX critic = OneDenseLayerMLP() # messager = OneDenseLayerMLP() From 1c4e40b5fbbd211fd88678025388301cb58391cd Mon Sep 17 00:00:00 2001 From: cedric Date: Thu, 2 Apr 2026 18:37:37 +0000 Subject: [PATCH 10/11] fix: uniform naming of models accross code --- src/MLPs/mlps.py | 4 ++-- src/ppo.py | 32 +++++++++++++++----------------- src/train.py | 37 ++++++++++++++++++++++--------------- 3 files changed, 39 insertions(+), 34 deletions(-) diff --git a/src/MLPs/mlps.py b/src/MLPs/mlps.py index 1fdc6a2..6abb540 100644 --- a/src/MLPs/mlps.py +++ b/src/MLPs/mlps.py @@ -40,10 +40,10 @@ class Actor(nn.Module): @jax.tree_util.register_dataclass @dataclass class AgentParams: - network_params: flax.core.FrozenDict + sensor_params: flax.core.FrozenDict actor_params: flax.core.FrozenDict critic_params: flax.core.FrozenDict - critic_network_params: flax.core.FrozenDict + feature_extractor_params: flax.core.FrozenDict @jax.tree_util.register_dataclass diff --git a/src/ppo.py b/src/ppo.py index 23c1e45..f2dfb82 100644 --- a/src/ppo.py +++ b/src/ppo.py @@ -8,9 +8,7 @@ import jax.numpy as jnp # Chose to use a class as it seemed the easiest way to integrate the CleanRL code style # with our need to seperate concerns class PPO: - def __init__( - self, args, input_network, action_network, critic, critic_network, message_passer=None - ): + def __init__(self, args, sensor, actor, critic, feature_extractor, message_passer=None): self.args = args if not message_passer: @@ -20,10 +18,10 @@ class PPO: partial( ppo_loss, args=args, - input_network_apply=input_network.apply, - action_network_apply=action_network.apply, + sensor_apply=sensor.apply, + actor_apply=actor.apply, critic_apply=critic.apply, - critic_network_apply=critic_network.apply, + feature_extractor_apply=feature_extractor.apply, message_passer=message_passer, ), has_aux=True, @@ -92,15 +90,15 @@ def get_action_and_value2( action_apply, message_passer, critic_apply, - critic_network_apply, + feature_extractor_apply, params: flax.core.FrozenDict, x: jnp.ndarray, action: jnp.ndarray, ): - hidden_network = input_apply(params["network_params"], x) - hidden_critic = critic_network_apply(params["critic_network_params"], x) - hidden_network = message_passer(hidden_network) - mean, log_std = action_apply(params["actor_params"], hidden_network) + hidden_sensor = input_apply(params["sensor_params"], x) + hidden_critic = feature_extractor_apply(params["feature_extractor_params"], x) + hidden_sensor = message_passer(hidden_sensor) + mean, log_std = action_apply(params["actor_params"], hidden_sensor) std = jnp.exp(log_std) logprob = -0.5 * (((action - mean) / std) ** 2 + 2 * log_std + jnp.log(2 * jnp.pi)).sum(-1) @@ -118,18 +116,18 @@ def ppo_loss( mb_advantages, mb_returns, args, - input_network_apply, - action_network_apply, + sensor_apply, + actor_apply, message_passer, critic_apply, - critic_network_apply, + feature_extractor_apply, ): newlogprob, entropy, newvalue = get_action_and_value2( - input_network_apply, - action_network_apply, + sensor_apply, + actor_apply, message_passer, critic_apply, - critic_network_apply, + feature_extractor_apply, params, x, a, diff --git a/src/train.py b/src/train.py index 10f14bf..1d3bb37 100644 --- a/src/train.py +++ b/src/train.py @@ -72,7 +72,7 @@ def train(args: PPOArgs): random.seed(args.seed) np.random.seed(args.seed) key = jax.random.PRNGKey(args.seed) - key, network_key, actor_key, critic_key, critic_network_key = jax.random.split(key, 5) + key, sensor_key, actor_key, critic_key, feature_extractor_key = jax.random.split(key, 5) torch.backends.cudnn.deterministic = args.torch_deterministic device = torch.device("cuda" if torch.cuda.is_available() and args.cuda else "cpu") @@ -122,8 +122,8 @@ def train(args: PPOArgs): return args.learning_rate * frac print("Initializing the models...") - network = GenericDenseLayersWithActivation() - critic_network = GenericDenseLayersWithActivation() + sensor = GenericDenseLayersWithActivation() + feature_extractor = GenericDenseLayersWithActivation() actor = Actor(action_dim=env.single_action_space.shape[0]) # continuous actions for MJX critic = OneDenseLayerMLP() # messager = OneDenseLayerMLP() @@ -135,15 +135,17 @@ def train(args: PPOArgs): if v.size > 0 ] ) - 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)) - critic_params = critic.init(critic_key, critic_network.apply(critic_network_params, sample_obs)) + sensor_params = sensor.init(sensor_key, sample_obs) + feature_extractor_params = feature_extractor.init(feature_extractor_key, sample_obs) + actor_params = actor.init(actor_key, sensor.apply(sensor_params, sample_obs)) + critic_params = critic.init( + critic_key, feature_extractor.apply(feature_extractor_params, sample_obs) + ) agent_state = TrainState.create( apply_fn=None, params=asdict( - AgentParams(network_params, actor_params, critic_params, critic_network_params) + AgentParams(sensor_params, actor_params, critic_params, feature_extractor_params) ), tx=optax.chain( optax.clip_by_global_norm(args.max_grad_norm), @@ -153,11 +155,11 @@ def train(args: PPOArgs): ), ) - network.apply = jax.jit(network.apply) - critic_network.apply = jax.jit(critic_network.apply) + sensor.apply = jax.jit(sensor.apply) + feature_extractor.apply = jax.jit(feature_extractor.apply) actor.apply = jax.jit(actor.apply) critic.apply = jax.jit(critic.apply) - ppo_instance = PPO(args, network, actor, critic, critic_network) + ppo_instance = PPO(args, sensor, actor, critic, feature_extractor) @jax.jit def get_action_and_value_noise( @@ -165,7 +167,11 @@ def train(args: PPOArgs): next_obs: jnp.ndarray, key: jax.random.PRNGKey, ): - hidden = network.apply(agent_state.params["network_params"], next_obs) + hidden = sensor.apply(agent_state.params["sensor_params"], next_obs) + hidden_critic = feature_extractor.apply( + agent_state.params["feature_extractor_params"], next_obs + ) + # Continuous actions: sample from a Gaussian parameterized by the actor mean, log_std = actor.apply(agent_state.params["actor_params"], hidden) key, subkey = jax.random.split(key) @@ -173,7 +179,7 @@ def train(args: PPOArgs): std = jnp.exp(log_std) action = mean + noise * std logprob = -0.5 * (((action - mean) / std) ** 2 + 2 * log_std + jnp.log(2 * jnp.pi)).sum(-1) - value = critic.apply(agent_state.params["critic_params"], hidden) + value = critic.apply(agent_state.params["critic_params"], hidden_critic) return action, logprob, value.squeeze(-1), key @jax.jit @@ -189,7 +195,7 @@ def train(args: PPOArgs): def compute_gae(agent_state, next_obs, next_done, storage): next_value = critic.apply( agent_state.params["critic_params"], - network.apply(agent_state.params["network_params"], next_obs), + sensor.apply(agent_state.params["sensor_params"], next_obs), ).squeeze(-1) advantages = jnp.zeros((args.num_envs,)) @@ -307,9 +313,10 @@ def train(args: PPOArgs): [ vars(args), [ - agent_state.params["network_params"], + agent_state.params["sensor_params"], agent_state.params["actor_params"], agent_state.params["critic_params"], + agent_state.params["feature_extractor_params"], ], ] ) From 3104678f96f2764567a2c743d6b6da6ba2ca884a Mon Sep 17 00:00:00 2001 From: cedric Date: Fri, 3 Apr 2026 08:11:50 +0000 Subject: [PATCH 11/11] fix: uniform naming --- src/ppo.py | 12 ++++++------ 1 file changed, 6 insertions(+), 6 deletions(-) diff --git a/src/ppo.py b/src/ppo.py index f2dfb82..cf6c69e 100644 --- a/src/ppo.py +++ b/src/ppo.py @@ -85,9 +85,9 @@ that are now not in the same scope @partial(jax.jit, static_argnums=(0, 1, 2, 3, 4)) -def get_action_and_value2( - input_apply, - action_apply, +def get_action_and_value( + sensor_apply, + actor_apply, message_passer, critic_apply, feature_extractor_apply, @@ -95,10 +95,10 @@ def get_action_and_value2( x: jnp.ndarray, action: jnp.ndarray, ): - hidden_sensor = input_apply(params["sensor_params"], x) + hidden_sensor = sensor_apply(params["sensor_params"], x) hidden_critic = feature_extractor_apply(params["feature_extractor_params"], x) hidden_sensor = message_passer(hidden_sensor) - mean, log_std = action_apply(params["actor_params"], hidden_sensor) + mean, log_std = actor_apply(params["actor_params"], hidden_sensor) std = jnp.exp(log_std) logprob = -0.5 * (((action - mean) / std) ** 2 + 2 * log_std + jnp.log(2 * jnp.pi)).sum(-1) @@ -122,7 +122,7 @@ def ppo_loss( critic_apply, feature_extractor_apply, ): - newlogprob, entropy, newvalue = get_action_and_value2( + newlogprob, entropy, newvalue = get_action_and_value( sensor_apply, actor_apply, message_passer,