Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
7 changes: 6 additions & 1 deletion ss2r/algorithms/mbpo/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,7 +2,10 @@

import ss2r.algorithms.mbpo.networks as mbpo_networks
from ss2r.algorithms.sac.data import get_collection_fn
from ss2r.algorithms.sac.q_transforms import get_reward_q_transform
from ss2r.algorithms.sac.q_transforms import (
get_cost_q_transform,
get_reward_q_transform,
)


def get_train_fn(cfg, checkpoint_path, restore_checkpoint_path):
Expand Down Expand Up @@ -54,6 +57,7 @@ def get_train_fn(cfg, checkpoint_path, restore_checkpoint_path):
policy_obs_key=policy_obs_key,
)
reward_q_transform = get_reward_q_transform(cfg)
cost_q_transform = get_cost_q_transform(cfg)
data_collection = get_collection_fn(cfg)
train_fn = functools.partial(
mbpo.train,
Expand All @@ -62,6 +66,7 @@ def get_train_fn(cfg, checkpoint_path, restore_checkpoint_path):
network_factory=network_factory,
checkpoint_logdir=checkpoint_path,
reward_q_transform=reward_q_transform,
cost_q_transform=cost_q_transform,
get_experience_fn=data_collection,
restore_checkpoint_path=restore_checkpoint_path,
)
Expand Down
41 changes: 27 additions & 14 deletions ss2r/algorithms/mbpo/model_env.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,7 +2,6 @@
import jax.numpy as jnp
from brax import envs
from brax.envs import base
from brax.training.acme import running_statistics

from ss2r.algorithms.sac.types import float32

Expand All @@ -13,17 +12,21 @@ class ModelBasedEnv(envs.Env):
def __init__(
self,
transitions,
observation_size: int,
action_size: int,
observation_size,
action_size,
model_network,
model_params,
normalizer_params: running_statistics.RunningStatisticsState,
ensemble_selection: str = "mean", # "random", "mean", or "pessimistic"
safety_budget: float = float("inf"),
qc_network,
qc_params,
normalizer_params,
ensemble_selection="mean", # "random", "mean", or "pessimistic"
safety_budget=float("inf"),
):
super().__init__()
self.model_network = model_network
self.model_params = model_params
self.qc_network = qc_network
self.qc_params = qc_params
self.normalizer_params = normalizer_params
self.ensemble_selection = ensemble_selection
self.safety_budget = safety_budget
Expand Down Expand Up @@ -79,11 +82,18 @@ def step(self, state: base.State, action: jax.Array) -> base.State:
truncation = jnp.zeros_like(reward, dtype=jnp.float32)
state.info["cost"] = cost
state.info["truncation"] = truncation
if "cumulative_cost" in state.info:
if self.qc_network is not None:
prev_cumulative_cost = state.info["cumulative_cost"]
curr_discount = state.info.get("curr_discount", jnp.ones_like(reward))
accumulated_cost_for_transition = (
prev_cumulative_cost + curr_discount * cost
prev_cumulative_cost
+ curr_discount
* self.qc_network.apply(
self.normalizer_params,
self.qc_params,
state.obs,
action,
).mean(axis=-1)
)
done = jnp.where(
accumulated_cost_for_transition > self.safety_budget,
Expand All @@ -104,7 +114,6 @@ def reset_states(self, done, state, next_obs):
return state, next_obs

state, next_obs = reset_states(self, done, state, next_obs)

state = state.replace(
obs=next_obs,
reward=reward,
Expand All @@ -130,11 +139,13 @@ def create_model_env(
transitions,
model_network,
model_params,
observation_size: int,
action_size: int,
normalizer_params: running_statistics.RunningStatisticsState,
ensemble_selection: str = "random",
safety_budget: float = float("inf"),
qc_network,
qc_params,
observation_size,
action_size,
normalizer_params,
ensemble_selection="random",
safety_budget=float("inf"),
) -> ModelBasedEnv:
"""Factory function to create a model-based environment."""
return ModelBasedEnv(
Expand All @@ -143,6 +154,8 @@ def create_model_env(
observation_size=observation_size,
action_size=action_size,
model_params=model_params,
qc_network=qc_network,
qc_params=qc_params,
normalizer_params=normalizer_params,
ensemble_selection=ensemble_selection,
safety_budget=safety_budget,
Expand Down
23 changes: 23 additions & 0 deletions ss2r/algorithms/mbpo/networks.py
Original file line number Diff line number Diff line change
Expand Up @@ -19,6 +19,7 @@
import brax.training.agents.sac.networks as sac_networks
import flax
import jax
import jax.nn as jnn
import jax.numpy as jnp
from brax.training import distribution, networks, types
from flax import linen
Expand All @@ -41,6 +42,7 @@ def __call__(
*,
n_critics: int = 2,
n_heads: int = 1,
safe: bool = False,
use_bro: bool = True,
) -> NetworkType:
pass
Expand All @@ -50,6 +52,7 @@ def __call__(
class MBPONetworks:
policy_network: networks.FeedForwardNetwork
qr_network: networks.FeedForwardNetwork
qc_network: networks.FeedForwardNetwork | None
model_network: networks.FeedForwardNetwork
parametric_action_distribution: distribution.ParametricDistribution

Expand Down Expand Up @@ -126,6 +129,7 @@ def make_mbpo_networks(
use_bro: bool = True,
n_critics: int = 2,
n_heads: int = 1,
safe: bool = False,
) -> MBPONetworks:
"""Make SAC networks."""
parametric_action_distribution = distribution.NormalTanhDistribution(
Expand All @@ -150,6 +154,24 @@ def make_mbpo_networks(
n_critics=n_critics,
n_heads=n_heads,
)
if safe:
qc_network = make_q_network(
observation_size,
action_size,
preprocess_observations_fn=preprocess_observations_fn,
hidden_layer_sizes=value_hidden_layer_sizes,
activation=activation,
obs_key=value_obs_key,
use_bro=use_bro,
n_critics=n_critics,
n_heads=n_heads,
)
old_apply = qc_network.apply
qc_network.apply = lambda *args, **kwargs: jnn.softplus(
old_apply(*args, **kwargs)
)
else:
qc_network = None
model_network = make_world_model_ensemble(
observation_size,
action_size,
Expand All @@ -161,6 +183,7 @@ def make_mbpo_networks(
return MBPONetworks(
policy_network=policy_network,
qr_network=qr_network,
qc_network=qc_network,
model_network=model_network,
parametric_action_distribution=parametric_action_distribution,
) # type: ignore
34 changes: 32 additions & 2 deletions ss2r/algorithms/mbpo/train.py
Original file line number Diff line number Diff line change
Expand Up @@ -40,7 +40,7 @@
from ss2r.algorithms.ppo.wrappers import TrackOnlineCosts
from ss2r.algorithms.sac import gradients
from ss2r.algorithms.sac.data import collect_single_step
from ss2r.algorithms.sac.q_transforms import QTransformation, SACBase
from ss2r.algorithms.sac.q_transforms import QTransformation, SACBase, SACCost
from ss2r.algorithms.sac.rae import RAEReplayBuffer
from ss2r.algorithms.sac.types import (
CollectDataFn,
Expand All @@ -59,6 +59,7 @@ def _init_training_state(
alpha_optimizer: optax.GradientTransformation,
policy_optimizer: optax.GradientTransformation,
qr_optimizer: optax.GradientTransformation,
qc_optimizer: optax.GradientTransformation,
model_optimizer: optax.GradientTransformation,
model_ensemble_size: int,
) -> TrainingState:
Expand All @@ -74,6 +75,13 @@ def _init_training_state(
model_keys = jax.random.split(key_model, model_ensemble_size)
model_params = init_model_ensemble(model_keys)
model_optimizer_state = model_optimizer.init(model_params)
if mbpo_network.qc_network is not None:
qc_params = mbpo_network.qc_network.init(key_qr)
assert qc_optimizer is not None
qc_optimizer_state = qc_optimizer.init(qc_params)
else:
qc_params = None
qc_optimizer_state = None
if isinstance(obs_size, Mapping):
obs_shape = {
k: specs.Array(v, jnp.dtype("float32")) for k, v in obs_size.items()
Expand All @@ -86,7 +94,10 @@ def _init_training_state(
policy_params=policy_params,
qr_optimizer_state=qr_optimizer_state,
qr_params=qr_params,
qc_optimizer_state=qc_optimizer_state,
qc_params=qc_params,
target_qr_params=qr_params,
target_qc_params=qc_params,
model_params=model_params,
model_optimizer_state=model_optimizer_state,
gradient_steps=jnp.zeros(()),
Expand Down Expand Up @@ -146,12 +157,14 @@ def train(
safe: bool = False,
safety_budget: float = float("inf"),
reward_q_transform: QTransformation = SACBase(),
cost_q_transform: QTransformation = SACCost(),
use_bro: bool = True,
normalize_budget: bool = True,
reset_on_eval: bool = True,
store_buffer: bool = False,
use_rae: bool = False,
optimism: float = 0.0,
pessimism: float = 0.0,
model_propagation: str = "nominal",
):
if min_replay_size >= num_timesteps:
Expand Down Expand Up @@ -203,6 +216,7 @@ def train(
observation_size=obs_size,
action_size=action_size,
preprocess_observations_fn=normalize_fn,
safe=safe,
use_bro=use_bro,
n_critics=n_critics,
n_heads=n_heads,
Expand All @@ -215,6 +229,7 @@ def train(
)
policy_optimizer = make_optimizer(learning_rate, 1.0)
qr_optimizer = make_optimizer(critic_learning_rate, 1.0)
qc_optimizer = make_optimizer(critic_learning_rate, 1.0)
model_optimizer = make_optimizer(model_learning_rate, 1.0)
if isinstance(obs_size, Mapping):
dummy_obs = {k: jnp.zeros(v) for k, v in obs_size.items()}
Expand Down Expand Up @@ -248,6 +263,7 @@ def train(
alpha_optimizer=alpha_optimizer,
policy_optimizer=policy_optimizer,
qr_optimizer=qr_optimizer,
qc_optimizer=qc_optimizer,
model_optimizer=model_optimizer,
model_ensemble_size=model_ensemble_size,
)
Expand All @@ -260,7 +276,8 @@ def train(
training_state = training_state.replace( # type: ignore
normalizer_params=params[0],
policy_params=params[1],
qr_params=params[3],
qr_params=params[2],
qc_params=params[3],
model_params=params[4],
)
if len(params) >= 6 and use_rae:
Expand Down Expand Up @@ -307,6 +324,12 @@ def train(
critic_loss, qr_optimizer, pmap_axis_name=None
)
)
if safe:
cost_critic_update = gradients.gradient_update_fn( # pytype: disable=wrong-arg-types # jax-ndarray
critic_loss, qc_optimizer, pmap_axis_name=None
)
else:
cost_critic_update = None
model_update = (
gradients.gradient_update_fn( # pytype: disable=wrong-arg-types # jax-ndarray
model_loss, model_optimizer, pmap_axis_name=None
Expand All @@ -328,6 +351,7 @@ def train(
make_model_env = functools.partial(
create_model_env,
model_network=mbpo_network.model_network,
qc_network=mbpo_network.qc_network,
action_size=action_size,
observation_size=obs_size,
ensemble_selection=model_propagation,
Expand All @@ -341,10 +365,13 @@ def train(
sac_replay_buffer,
alpha_update,
critic_update,
cost_critic_update,
model_update,
actor_update,
safe,
min_alpha,
reward_q_transform,
cost_q_transform,
model_grad_updates_per_step,
critic_grad_updates_per_step,
extra_fields,
Expand All @@ -355,6 +382,7 @@ def train(
unroll_length,
num_model_rollouts,
optimism,
pessimism,
model_to_real_data_ratio,
)

Expand Down Expand Up @@ -554,6 +582,7 @@ def training_epoch_with_timing(
training_state.normalizer_params,
training_state.policy_params,
training_state.qr_params,
training_state.qc_params,
training_state.model_params,
)
if store_buffer:
Expand All @@ -575,6 +604,7 @@ def training_epoch_with_timing(
training_state.normalizer_params,
training_state.policy_params,
training_state.qr_params,
training_state.qc_params,
training_state.model_params,
)
if store_buffer:
Expand Down
Loading