From 9879fea9e221c0839acb4260f1de1b5be46a8992 Mon Sep 17 00:00:00 2001 From: cedric Date: Thu, 2 Apr 2026 17:50:44 +0000 Subject: [PATCH] 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()