Merge branch 'dev' into ci/hpc
This commit is contained in:
commit
ef380d073f
6 changed files with 62 additions and 298 deletions
|
|
@ -1,36 +1,27 @@
|
||||||
from dataclasses import dataclass, fields
|
from dataclasses import dataclass, fields, field
|
||||||
|
|
||||||
import flax
|
import flax
|
||||||
import flax.linen as nn
|
import flax.linen as nn
|
||||||
import jax.numpy as jnp
|
import jax.numpy as jnp
|
||||||
import jax.tree_util
|
import jax.tree_util
|
||||||
import numpy as np
|
from typing import Sequence, Callable
|
||||||
from flax.linen.initializers import constant, orthogonal
|
from flax.linen.initializers import constant, orthogonal
|
||||||
|
|
||||||
|
|
||||||
class Network(nn.Module):
|
# semi generic so we can easily make a config for it in experiments
|
||||||
"""
|
class GenericDenseLayersWithActivation(nn.Module):
|
||||||
Dummy model only used for testing purposes
|
layer_sizes: Sequence[int] = field(default_factory=lambda: [64, 64])
|
||||||
|
activation: Callable = nn.tanh
|
||||||
inspired by: https://github.com/vwxyzjn/cleanrl/blob/master/cleanrl/ppo_atari_envpool_xla_jax_scan.py
|
|
||||||
"""
|
|
||||||
|
|
||||||
hidden_dim: int = 195
|
|
||||||
|
|
||||||
@nn.compact
|
@nn.compact
|
||||||
def __call__(self, x):
|
def __call__(self, x):
|
||||||
x = nn.Dense(self.hidden_dim, kernel_init=orthogonal(np.sqrt(2)), bias_init=constant(0.0))(
|
for size in self.layer_sizes:
|
||||||
x
|
x = nn.Dense(size, kernel_init=orthogonal(jnp.sqrt(2)))(x)
|
||||||
)
|
x = self.activation(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
|
return x
|
||||||
|
|
||||||
|
|
||||||
class Critic(nn.Module):
|
class OneDenseLayerMLP(nn.Module):
|
||||||
@nn.compact
|
@nn.compact
|
||||||
def __call__(self, x):
|
def __call__(self, x):
|
||||||
return nn.Dense(1, kernel_init=orthogonal(1), bias_init=constant(0.0))(x)
|
return nn.Dense(1, kernel_init=orthogonal(1), bias_init=constant(0.0))(x)
|
||||||
|
|
@ -49,10 +40,10 @@ class Actor(nn.Module):
|
||||||
@jax.tree_util.register_dataclass
|
@jax.tree_util.register_dataclass
|
||||||
@dataclass
|
@dataclass
|
||||||
class AgentParams:
|
class AgentParams:
|
||||||
network_params: flax.core.FrozenDict
|
sensor_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
|
feature_extractor_params: flax.core.FrozenDict
|
||||||
|
|
||||||
|
|
||||||
@jax.tree_util.register_dataclass
|
@jax.tree_util.register_dataclass
|
||||||
|
|
@ -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",
|
|
||||||
]
|
|
||||||
|
|
@ -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 "<none>"
|
|
||||||
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 "<none>"
|
|
||||||
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)
|
|
||||||
|
|
@ -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)),
|
|
||||||
)
|
|
||||||
40
src/ppo.py
40
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
|
# 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__(
|
def __init__(self, args, sensor, actor, critic, feature_extractor, message_passer=None):
|
||||||
self, args, input_network, action_network, critic, critic_network, message_passer=None
|
|
||||||
):
|
|
||||||
self.args = args
|
self.args = args
|
||||||
|
|
||||||
if not message_passer:
|
if not message_passer:
|
||||||
|
|
@ -20,10 +18,10 @@ class PPO:
|
||||||
partial(
|
partial(
|
||||||
ppo_loss,
|
ppo_loss,
|
||||||
args=args,
|
args=args,
|
||||||
input_network_apply=input_network.apply,
|
sensor_apply=sensor.apply,
|
||||||
action_network_apply=action_network.apply,
|
actor_apply=actor.apply,
|
||||||
critic_apply=critic.apply,
|
critic_apply=critic.apply,
|
||||||
critic_network_apply=critic_network.apply,
|
feature_extractor_apply=feature_extractor.apply,
|
||||||
message_passer=message_passer,
|
message_passer=message_passer,
|
||||||
),
|
),
|
||||||
has_aux=True,
|
has_aux=True,
|
||||||
|
|
@ -87,20 +85,20 @@ that are now not in the same scope
|
||||||
|
|
||||||
|
|
||||||
@partial(jax.jit, static_argnums=(0, 1, 2, 3, 4))
|
@partial(jax.jit, static_argnums=(0, 1, 2, 3, 4))
|
||||||
def get_action_and_value2(
|
def get_action_and_value(
|
||||||
input_apply,
|
sensor_apply,
|
||||||
action_apply,
|
actor_apply,
|
||||||
message_passer,
|
message_passer,
|
||||||
critic_apply,
|
critic_apply,
|
||||||
critic_network_apply,
|
feature_extractor_apply,
|
||||||
params: flax.core.FrozenDict,
|
params: flax.core.FrozenDict,
|
||||||
x: jnp.ndarray,
|
x: jnp.ndarray,
|
||||||
action: jnp.ndarray,
|
action: jnp.ndarray,
|
||||||
):
|
):
|
||||||
hidden_network = input_apply(params["network_params"], x)
|
hidden_sensor = sensor_apply(params["sensor_params"], x)
|
||||||
hidden_critic = critic_network_apply(params["critic_network_params"], x)
|
hidden_critic = feature_extractor_apply(params["feature_extractor_params"], x)
|
||||||
hidden_network = message_passer(hidden_network)
|
hidden_sensor = message_passer(hidden_sensor)
|
||||||
mean, log_std = action_apply(params["actor_params"], hidden_network)
|
mean, log_std = actor_apply(params["actor_params"], hidden_sensor)
|
||||||
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)).sum(-1)
|
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_advantages,
|
||||||
mb_returns,
|
mb_returns,
|
||||||
args,
|
args,
|
||||||
input_network_apply,
|
sensor_apply,
|
||||||
action_network_apply,
|
actor_apply,
|
||||||
message_passer,
|
message_passer,
|
||||||
critic_apply,
|
critic_apply,
|
||||||
critic_network_apply,
|
feature_extractor_apply,
|
||||||
):
|
):
|
||||||
newlogprob, entropy, newvalue = get_action_and_value2(
|
newlogprob, entropy, newvalue = get_action_and_value(
|
||||||
input_network_apply,
|
sensor_apply,
|
||||||
action_network_apply,
|
actor_apply,
|
||||||
message_passer,
|
message_passer,
|
||||||
critic_apply,
|
critic_apply,
|
||||||
critic_network_apply,
|
feature_extractor_apply,
|
||||||
params,
|
params,
|
||||||
x,
|
x,
|
||||||
a,
|
a,
|
||||||
|
|
|
||||||
48
src/train.py
48
src/train.py
|
|
@ -23,7 +23,13 @@ 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 Actor, AgentParams, Critic, Network, Storage
|
from MLPs.mlps import (
|
||||||
|
GenericDenseLayersWithActivation,
|
||||||
|
Actor,
|
||||||
|
OneDenseLayerMLP,
|
||||||
|
AgentParams,
|
||||||
|
Storage,
|
||||||
|
)
|
||||||
from ppo import PPO
|
from ppo import PPO
|
||||||
|
|
||||||
|
|
||||||
|
|
@ -83,7 +89,7 @@ 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, 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
|
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")
|
||||||
|
|
@ -133,10 +139,11 @@ def train(args: PPOArgs):
|
||||||
return args.learning_rate * frac
|
return args.learning_rate * frac
|
||||||
|
|
||||||
print("Initializing the models...")
|
print("Initializing the models...")
|
||||||
network = Network()
|
sensor = GenericDenseLayersWithActivation()
|
||||||
critic_network = Network()
|
feature_extractor = GenericDenseLayersWithActivation()
|
||||||
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 = OneDenseLayerMLP()
|
||||||
|
# messager = OneDenseLayerMLP()
|
||||||
|
|
||||||
sample_obs = jnp.concatenate(
|
sample_obs = jnp.concatenate(
|
||||||
[
|
[
|
||||||
|
|
@ -145,15 +152,17 @@ def train(args: PPOArgs):
|
||||||
if v.size > 0
|
if v.size > 0
|
||||||
]
|
]
|
||||||
)
|
)
|
||||||
network_params = network.init(network_key, sample_obs)
|
sensor_params = sensor.init(sensor_key, sample_obs)
|
||||||
critic_network_params = critic_network.init(critic_network_key, sample_obs)
|
feature_extractor_params = feature_extractor.init(feature_extractor_key, sample_obs)
|
||||||
actor_params = actor.init(actor_key, network.apply(network_params, sample_obs))
|
actor_params = actor.init(actor_key, sensor.apply(sensor_params, sample_obs))
|
||||||
critic_params = critic.init(critic_key, critic_network.apply(critic_network_params, sample_obs))
|
critic_params = critic.init(
|
||||||
|
critic_key, feature_extractor.apply(feature_extractor_params, sample_obs)
|
||||||
|
)
|
||||||
|
|
||||||
agent_state = TrainState.create(
|
agent_state = TrainState.create(
|
||||||
apply_fn=None,
|
apply_fn=None,
|
||||||
params=asdict(
|
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(
|
tx=optax.chain(
|
||||||
optax.clip_by_global_norm(args.max_grad_norm),
|
optax.clip_by_global_norm(args.max_grad_norm),
|
||||||
|
|
@ -163,11 +172,11 @@ def train(args: PPOArgs):
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
network.apply = jax.jit(network.apply)
|
sensor.apply = jax.jit(sensor.apply)
|
||||||
critic_network.apply = jax.jit(critic_network.apply)
|
feature_extractor.apply = jax.jit(feature_extractor.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)
|
ppo_instance = PPO(args, sensor, actor, critic, feature_extractor)
|
||||||
|
|
||||||
@jax.jit
|
@jax.jit
|
||||||
def get_action_and_value_noise(
|
def get_action_and_value_noise(
|
||||||
|
|
@ -175,7 +184,11 @@ def train(args: PPOArgs):
|
||||||
next_obs: jnp.ndarray,
|
next_obs: jnp.ndarray,
|
||||||
key: jax.random.PRNGKey,
|
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
|
# Continuous actions: sample from a Gaussian parameterized by the actor
|
||||||
mean, log_std = actor.apply(agent_state.params["actor_params"], hidden)
|
mean, log_std = actor.apply(agent_state.params["actor_params"], hidden)
|
||||||
key, subkey = jax.random.split(key)
|
key, subkey = jax.random.split(key)
|
||||||
|
|
@ -183,7 +196,7 @@ def train(args: PPOArgs):
|
||||||
std = jnp.exp(log_std)
|
std = jnp.exp(log_std)
|
||||||
action = mean + noise * std
|
action = mean + noise * std
|
||||||
logprob = -0.5 * (((action - mean) / std) ** 2 + 2 * log_std + jnp.log(2 * jnp.pi)).sum(-1)
|
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
|
return action, logprob, value.squeeze(-1), key
|
||||||
|
|
||||||
@jax.jit
|
@jax.jit
|
||||||
|
|
@ -199,7 +212,7 @@ def train(args: PPOArgs):
|
||||||
def compute_gae(agent_state, next_obs, next_done, storage):
|
def compute_gae(agent_state, next_obs, next_done, storage):
|
||||||
next_value = critic.apply(
|
next_value = critic.apply(
|
||||||
agent_state.params["critic_params"],
|
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)
|
).squeeze(-1)
|
||||||
|
|
||||||
advantages = jnp.zeros((args.num_envs,))
|
advantages = jnp.zeros((args.num_envs,))
|
||||||
|
|
@ -353,9 +366,10 @@ def train(args: PPOArgs):
|
||||||
[
|
[
|
||||||
vars(args),
|
vars(args),
|
||||||
[
|
[
|
||||||
agent_state.params["network_params"],
|
agent_state.params["sensor_params"],
|
||||||
agent_state.params["actor_params"],
|
agent_state.params["actor_params"],
|
||||||
agent_state.params["critic_params"],
|
agent_state.params["critic_params"],
|
||||||
|
agent_state.params["feature_extractor_params"],
|
||||||
],
|
],
|
||||||
]
|
]
|
||||||
)
|
)
|
||||||
|
|
|
||||||
Reference in a new issue