chore: import upstream snapshot with attribution
This commit is contained in:
@@ -0,0 +1,245 @@
|
||||
"""Example showing how to continue training an Algorithm with a changed config.
|
||||
|
||||
Use the setup shown in this script if you want to continue a prior experiment, but
|
||||
would also like to change some of the config values you originally used.
|
||||
|
||||
This example:
|
||||
- runs a single- or multi-agent CartPole experiment (for multi-agent, we use
|
||||
different learning rates) thereby checkpointing the state of the Algorithm every n
|
||||
iterations. The config used is hereafter called "1st config".
|
||||
- stops the experiment due to some episode return being achieved.
|
||||
- just for testing purposes, restores the entire algorithm from the latest
|
||||
checkpoint and checks, whether the state of the restored algo exactly match the
|
||||
state of the previously saved one.
|
||||
- then changes the original config used (learning rate and other settings) and
|
||||
continues training with the restored algorithm and the changed config until a
|
||||
final episode return is reached. The new config is hereafter called "2nd config".
|
||||
|
||||
|
||||
How to run this script
|
||||
----------------------
|
||||
`python [script file name].py --num-agents=[0 or 2]
|
||||
--stop-reward-first-config=[return at which the algo on 1st config should stop training]
|
||||
--stop-reward=[the final return to achieve after restoration from the checkpoint with
|
||||
the 2nd config]
|
||||
`
|
||||
|
||||
For debugging, use the following additional command line options
|
||||
`--no-tune --num-env-runners=0`
|
||||
which should allow you to set breakpoints anywhere in the RLlib code and
|
||||
have the execution stop there for inspection and debugging.
|
||||
|
||||
For logging to your WandB account, use:
|
||||
`--wandb-key=[your WandB API key] --wandb-project=[some project name]
|
||||
--wandb-run-name=[optional: WandB run name (within the defined project)]`
|
||||
|
||||
|
||||
Results to expect
|
||||
-----------------
|
||||
First, you should see the initial tune.Tuner do it's thing:
|
||||
|
||||
Trial status: 1 RUNNING
|
||||
Current time: 2024-06-03 12:03:39. Total running time: 30s
|
||||
Logical resource usage: 3.0/12 CPUs, 0/0 GPUs
|
||||
╭────────────────────────────────────────────────────────────────────────
|
||||
│ Trial name status iter total time (s)
|
||||
├────────────────────────────────────────────────────────────────────────
|
||||
│ PPO_CartPole-v1_7b1eb_00000 RUNNING 6 16.265
|
||||
╰────────────────────────────────────────────────────────────────────────
|
||||
───────────────────────────────────────────────────────────────────────╮
|
||||
..._sampled_lifetime ..._trained_lifetime ...episodes_lifetime │
|
||||
───────────────────────────────────────────────────────────────────────┤
|
||||
24000 24000 340 │
|
||||
───────────────────────────────────────────────────────────────────────╯
|
||||
...
|
||||
|
||||
The experiment stops at an average episode return of `--stop-reward-first-config`.
|
||||
|
||||
After the validation of the last checkpoint, a new experiment is started from
|
||||
scratch, but with the RLlib callback restoring the Algorithm right after
|
||||
initialization using the previous checkpoint. This new experiment then runs
|
||||
until `--stop-reward` is reached.
|
||||
|
||||
Trial status: 1 RUNNING
|
||||
Current time: 2024-06-03 12:05:00. Total running time: 1min 0s
|
||||
Logical resource usage: 3.0/12 CPUs, 0/0 GPUs
|
||||
╭────────────────────────────────────────────────────────────────────────
|
||||
│ Trial name status iter total time (s)
|
||||
├────────────────────────────────────────────────────────────────────────
|
||||
│ PPO_CartPole-v1_7b1eb_00000 RUNNING 23 14.8372
|
||||
╰────────────────────────────────────────────────────────────────────────
|
||||
───────────────────────────────────────────────────────────────────────╮
|
||||
..._sampled_lifetime ..._trained_lifetime ...episodes_lifetime │
|
||||
───────────────────────────────────────────────────────────────────────┤
|
||||
109078 109078 531 │
|
||||
───────────────────────────────────────────────────────────────────────╯
|
||||
|
||||
And if you are using the `--as-test` option, you should see a finel message:
|
||||
|
||||
```
|
||||
`env_runners/episode_return_mean` of 450.0 reached! ok
|
||||
```
|
||||
"""
|
||||
from ray.rllib.algorithms.algorithm_config import AlgorithmConfig
|
||||
from ray.rllib.algorithms.ppo import PPOConfig
|
||||
from ray.rllib.core import DEFAULT_MODULE_ID
|
||||
from ray.rllib.examples.envs.classes.multi_agent import MultiAgentCartPole
|
||||
from ray.rllib.examples.utils import (
|
||||
add_rllib_example_script_args,
|
||||
run_rllib_example_script_experiment,
|
||||
)
|
||||
from ray.rllib.policy.policy import PolicySpec
|
||||
from ray.rllib.utils.metrics import (
|
||||
ENV_RUNNER_RESULTS,
|
||||
EPISODE_RETURN_MEAN,
|
||||
LEARNER_RESULTS,
|
||||
)
|
||||
from ray.rllib.utils.numpy import convert_to_numpy
|
||||
from ray.rllib.utils.test_utils import check
|
||||
from ray.tune.registry import register_env
|
||||
|
||||
parser = add_rllib_example_script_args(
|
||||
default_reward=450.0, default_timesteps=10000000, default_iters=2000
|
||||
)
|
||||
parser.add_argument(
|
||||
"--stop-reward-first-config",
|
||||
type=float,
|
||||
default=150.0,
|
||||
help="Mean episode return after which the Algorithm on the first config should "
|
||||
"stop training.",
|
||||
)
|
||||
# By default, set `args.checkpoint_freq` to 1 and `args.checkpoint_at_end` to True.
|
||||
parser.set_defaults(
|
||||
checkpoint_freq=1,
|
||||
checkpoint_at_end=True,
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
args = parser.parse_args()
|
||||
|
||||
register_env(
|
||||
"ma_cart", lambda cfg: MultiAgentCartPole({"num_agents": args.num_agents})
|
||||
)
|
||||
|
||||
# Simple generic config.
|
||||
base_config = (
|
||||
PPOConfig()
|
||||
.environment("CartPole-v1" if args.num_agents == 0 else "ma_cart")
|
||||
.env_runners(create_env_on_local_worker=True)
|
||||
.training(lr=0.0001)
|
||||
# TODO (sven): Tune throws a weird error inside the "log json" callback
|
||||
# when running with this option. The `perf` key in the result dict contains
|
||||
# binary data (instead of just 2 float values for mem and cpu usage).
|
||||
# .experimental(_use_msgpack_checkpoints=True)
|
||||
)
|
||||
|
||||
# Setup multi-agent, if required.
|
||||
if args.num_agents > 0:
|
||||
base_config.multi_agent(
|
||||
policies={
|
||||
f"p{aid}": PolicySpec(
|
||||
config=AlgorithmConfig.overrides(
|
||||
lr=5e-5
|
||||
* (aid + 1), # agent 1 has double the learning rate as 0.
|
||||
)
|
||||
)
|
||||
for aid in range(args.num_agents)
|
||||
},
|
||||
policy_mapping_fn=lambda aid, *a, **kw: f"p{aid}",
|
||||
)
|
||||
|
||||
# Define some stopping criterion. Note that this criterion is an avg episode return
|
||||
# to be reached.
|
||||
metric = f"{ENV_RUNNER_RESULTS}/{EPISODE_RETURN_MEAN}"
|
||||
stop = {metric: args.stop_reward_first_config}
|
||||
|
||||
tuner_results = run_rllib_example_script_experiment(
|
||||
base_config,
|
||||
args,
|
||||
stop=stop,
|
||||
keep_ray_up=True,
|
||||
)
|
||||
|
||||
# Perform a very quick test to make sure our algo (upon restoration) did not lose
|
||||
# its ability to perform well in the env.
|
||||
# - Extract the best checkpoint.
|
||||
best_result = tuner_results.get_best_result(metric=metric, mode="max")
|
||||
assert (
|
||||
best_result.metrics[ENV_RUNNER_RESULTS][EPISODE_RETURN_MEAN]
|
||||
>= args.stop_reward_first_config
|
||||
)
|
||||
best_checkpoint_path = best_result.checkpoint.path
|
||||
|
||||
# Rebuild the algorithm (just for testing purposes).
|
||||
test_algo = base_config.build()
|
||||
# Load algo's state from the best checkpoint.
|
||||
test_algo.restore_from_path(best_checkpoint_path)
|
||||
# Perform some checks on the restored state.
|
||||
assert test_algo.training_iteration > 0
|
||||
# Evaluate on the restored algorithm.
|
||||
test_eval_results = test_algo.evaluate()
|
||||
assert (
|
||||
test_eval_results[ENV_RUNNER_RESULTS][EPISODE_RETURN_MEAN]
|
||||
>= args.stop_reward_first_config
|
||||
), test_eval_results[ENV_RUNNER_RESULTS][EPISODE_RETURN_MEAN]
|
||||
# Train one iteration to make sure, the performance does not collapse (e.g. due
|
||||
# to the optimizer weights not having been restored properly).
|
||||
test_results = test_algo.train()
|
||||
assert (
|
||||
test_results[ENV_RUNNER_RESULTS][EPISODE_RETURN_MEAN]
|
||||
>= args.stop_reward_first_config
|
||||
), test_results[ENV_RUNNER_RESULTS][EPISODE_RETURN_MEAN]
|
||||
# Stop the test algorithm again.
|
||||
test_algo.stop()
|
||||
|
||||
# Make sure the algorithm gets restored from a checkpoint right after
|
||||
# initialization. Note that this includes all subcomponents of the algorithm,
|
||||
# including the optimizer states in the LearnerGroup/Learner actors.
|
||||
def on_algorithm_init(algorithm, **kwargs):
|
||||
module_p0 = algorithm.get_module("p0")
|
||||
weight_before = convert_to_numpy(next(iter(module_p0.parameters())))
|
||||
|
||||
algorithm.restore_from_path(best_checkpoint_path)
|
||||
|
||||
# Make sure weights were restored (changed).
|
||||
weight_after = convert_to_numpy(next(iter(module_p0.parameters())))
|
||||
check(weight_before, weight_after, false=True)
|
||||
|
||||
# Change the config.
|
||||
(
|
||||
base_config
|
||||
# Make sure the algorithm gets restored upon initialization.
|
||||
.callbacks(on_algorithm_init=on_algorithm_init)
|
||||
# Change training parameters considerably.
|
||||
.training(
|
||||
lr=0.0003,
|
||||
train_batch_size=5000,
|
||||
grad_clip=100.0,
|
||||
gamma=0.996,
|
||||
num_epochs=6,
|
||||
vf_loss_coeff=0.01,
|
||||
)
|
||||
# Make multi-CPU/GPU.
|
||||
.learners(num_learners=2)
|
||||
# Use more env runners and more envs per env runner.
|
||||
.env_runners(num_env_runners=3, num_envs_per_env_runner=5)
|
||||
)
|
||||
|
||||
# Update the stopping criterium to the final target return per episode.
|
||||
stop = {metric: args.stop_reward}
|
||||
|
||||
# Run a new experiment with the (RLlib) callback `on_algorithm_init` restoring
|
||||
# from the best checkpoint.
|
||||
# Note that the new experiment starts again from iteration=0 (unlike when you
|
||||
# use `tune.Tuner.restore()` after a crash or interrupted trial).
|
||||
tuner_results = run_rllib_example_script_experiment(base_config, args, stop=stop)
|
||||
|
||||
# Assert that we have continued training with a different learning rate.
|
||||
assert (
|
||||
tuner_results[0].metrics[LEARNER_RESULTS][DEFAULT_MODULE_ID][
|
||||
"default_optimizer_learning_rate"
|
||||
]
|
||||
== base_config.lr
|
||||
== 0.0003
|
||||
)
|
||||
@@ -0,0 +1,148 @@
|
||||
"""Example extracting a checkpoint from n trials using one or more custom criteria.
|
||||
|
||||
This example:
|
||||
- runs a CartPole experiment with three different learning rates (three tune
|
||||
"trials"). During the experiment, for each trial, we create a checkpoint at each
|
||||
iteration.
|
||||
- at the end of the experiment, we compare the trials and pick the one that
|
||||
performed best, based on the criterion: Lowest episode count per single iteration
|
||||
(for CartPole, a low episode count means the episodes are very long and thus the
|
||||
reward is also very high).
|
||||
- from that best trial (with the lowest episode count), we then pick those
|
||||
checkpoints that a) have the lowest policy loss (good) and b) have the highest value
|
||||
function loss (bad).
|
||||
|
||||
|
||||
How to run this script
|
||||
----------------------
|
||||
`python [script file name].py`
|
||||
|
||||
For debugging, use the following additional command line options
|
||||
`--no-tune --num-env-runners=0`
|
||||
which should allow you to set breakpoints anywhere in the RLlib code and
|
||||
have the execution stop there for inspection and debugging.
|
||||
|
||||
For logging to your WandB account, use:
|
||||
`--wandb-key=[your WandB API key] --wandb-project=[some project name]
|
||||
--wandb-run-name=[optional: WandB run name (within the defined project)]`
|
||||
|
||||
|
||||
Results to expect
|
||||
-----------------
|
||||
In the console output, you can see the performance of the three different learning
|
||||
rates used here:
|
||||
|
||||
+-----------------------------+------------+-----------------+--------+--------+
|
||||
| Trial name | status | loc | lr | iter |
|
||||
|-----------------------------+------------+-----------------+--------+--------+
|
||||
| PPO_CartPole-v1_d7dbe_00000 | TERMINATED | 127.0.0.1:98487 | 0.01 | 17 |
|
||||
| PPO_CartPole-v1_d7dbe_00001 | TERMINATED | 127.0.0.1:98488 | 0.001 | 8 |
|
||||
| PPO_CartPole-v1_d7dbe_00002 | TERMINATED | 127.0.0.1:98489 | 0.0001 | 9 |
|
||||
+-----------------------------+------------+-----------------+--------+--------+
|
||||
|
||||
+------------------+-------+----------+----------------------+----------------------+
|
||||
| total time (s) | ts | reward | episode_reward_max | episode_reward_min |
|
||||
|------------------+-------+----------+----------------------+----------------------+
|
||||
| 28.1068 | 39797 | 151.11 | 500 | 12 |
|
||||
| 13.304 | 18728 | 158.91 | 500 | 15 |
|
||||
| 14.8848 | 21069 | 167.36 | 500 | 13 |
|
||||
+------------------+-------+----------+----------------------+----------------------+
|
||||
|
||||
+--------------------+
|
||||
| episode_len_mean |
|
||||
|--------------------|
|
||||
| 151.11 |
|
||||
| 158.91 |
|
||||
| 167.36 |
|
||||
+--------------------+
|
||||
"""
|
||||
|
||||
from ray import tune
|
||||
from ray.rllib.core import DEFAULT_MODULE_ID
|
||||
from ray.rllib.examples.utils import (
|
||||
add_rllib_example_script_args,
|
||||
run_rllib_example_script_experiment,
|
||||
)
|
||||
from ray.rllib.utils.metrics import (
|
||||
ENV_RUNNER_RESULTS,
|
||||
EPISODE_RETURN_MEAN,
|
||||
LEARNER_RESULTS,
|
||||
)
|
||||
from ray.tune.registry import get_trainable_cls
|
||||
|
||||
parser = add_rllib_example_script_args(
|
||||
default_reward=450.0,
|
||||
default_timesteps=100000,
|
||||
default_iters=200,
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
args = parser.parse_args()
|
||||
|
||||
# Force-set `args.checkpoint_freq` to 1.
|
||||
args.checkpoint_freq = 1
|
||||
|
||||
# Simple generic config.
|
||||
base_config = (
|
||||
get_trainable_cls(args.algo)
|
||||
.get_default_config()
|
||||
.environment("CartPole-v1")
|
||||
# Run 3 trials, each w/ a different learning rate.
|
||||
.training(lr=tune.grid_search([0.01, 0.001, 0.0001]), train_batch_size=2341)
|
||||
)
|
||||
# Run tune for some iterations and generate checkpoints.
|
||||
results = run_rllib_example_script_experiment(base_config, args)
|
||||
|
||||
# Get the best of the 3 trials by using some metric.
|
||||
# NOTE: Choosing the min `episodes_this_iter` automatically picks the trial
|
||||
# with the best performance (over the entire run (scope="all")):
|
||||
# The fewer episodes, the longer each episode lasted, the more reward we
|
||||
# got each episode.
|
||||
# Setting scope to "last", "last-5-avg", or "last-10-avg" will only compare
|
||||
# (using `mode=min|max`) the average values of the last 1, 5, or 10
|
||||
# iterations with each other, respectively.
|
||||
# Setting scope to "avg" will compare (using `mode`=min|max) the average
|
||||
# values over the entire run.
|
||||
metric = "env_runners/num_episodes"
|
||||
# notice here `scope` is `all`, meaning for each trial,
|
||||
# all results (not just the last one) will be examined.
|
||||
best_result = results.get_best_result(metric=metric, mode="min", scope="all")
|
||||
value_best_metric = best_result.metrics_dataframe[metric].min()
|
||||
best_return_best = best_result.metrics_dataframe[
|
||||
f"{ENV_RUNNER_RESULTS}/{EPISODE_RETURN_MEAN}"
|
||||
].max()
|
||||
print(
|
||||
f"Best trial was the one with lr={best_result.metrics['config']['lr']}. "
|
||||
f"Reached lowest episode count ({value_best_metric}) in a single iteration and "
|
||||
f"an average return of {best_return_best}."
|
||||
)
|
||||
|
||||
# Confirm, we picked the right trial.
|
||||
|
||||
assert (
|
||||
value_best_metric
|
||||
== results.get_dataframe(filter_metric=metric, filter_mode="min")[metric].min()
|
||||
)
|
||||
|
||||
# Get the best checkpoints from the trial, based on different metrics.
|
||||
# Checkpoint with the lowest policy loss value:
|
||||
if not args.old_api_stack:
|
||||
policy_loss_key = f"{LEARNER_RESULTS}/{DEFAULT_MODULE_ID}/policy_loss"
|
||||
else:
|
||||
policy_loss_key = "info/learner/default_policy/learner_stats/policy_loss"
|
||||
best_result = results.get_best_result(metric=policy_loss_key, mode="min")
|
||||
ckpt = best_result.checkpoint
|
||||
lowest_policy_loss = best_result.metrics_dataframe[policy_loss_key].min()
|
||||
print(f"Checkpoint w/ lowest policy loss ({lowest_policy_loss}): {ckpt}")
|
||||
|
||||
# Checkpoint with the highest value-function loss:
|
||||
if not args.old_api_stack:
|
||||
vf_loss_key = f"{LEARNER_RESULTS}/{DEFAULT_MODULE_ID}/vf_loss"
|
||||
else:
|
||||
vf_loss_key = "info/learner/default_policy/learner_stats/vf_loss"
|
||||
best_result = results.get_best_result(metric=vf_loss_key, mode="max")
|
||||
ckpt = best_result.checkpoint
|
||||
highest_value_fn_loss = best_result.metrics_dataframe[vf_loss_key].max()
|
||||
print(f"Checkpoint w/ highest value function loss: {ckpt}")
|
||||
print(f"Highest value function loss: {highest_value_fn_loss}")
|
||||
@@ -0,0 +1,264 @@
|
||||
"""Example showing how to restore an Algorithm from a checkpoint and resume training.
|
||||
|
||||
Use the setup shown in this script if your experiments tend to crash after some time,
|
||||
and you would therefore like to make your setup more robust and fault-tolerant.
|
||||
|
||||
This example:
|
||||
- runs a single- or multi-agent CartPole experiment (for multi-agent, we use
|
||||
different learning rates) thereby checkpointing the state of the Algorithm every n
|
||||
iterations.
|
||||
- stops the experiment due to an expected crash in the algorithm's main process
|
||||
after a certain number of iterations.
|
||||
- just for testing purposes, restores the entire algorithm from the latest
|
||||
checkpoint and checks, whether the state of the restored algo exactly match the
|
||||
state of the crashed one.
|
||||
- then continues training with the restored algorithm until the desired final
|
||||
episode return is reached.
|
||||
|
||||
|
||||
How to run this script
|
||||
----------------------
|
||||
`python [script file name].py --num-agents=[0 or 2]
|
||||
--stop-reward-crash=[the episode return after which the algo should crash]
|
||||
--stop-reward=[the final episode return to achieve after(!) restoration from the
|
||||
checkpoint]
|
||||
`
|
||||
|
||||
For debugging, use the following additional command line options
|
||||
`--no-tune --num-env-runners=0`
|
||||
which should allow you to set breakpoints anywhere in the RLlib code and
|
||||
have the execution stop there for inspection and debugging.
|
||||
|
||||
For logging to your WandB account, use:
|
||||
`--wandb-key=[your WandB API key] --wandb-project=[some project name]
|
||||
--wandb-run-name=[optional: WandB run name (within the defined project)]`
|
||||
|
||||
|
||||
Results to expect
|
||||
-----------------
|
||||
First, you should see the initial tune.Tuner do it's thing:
|
||||
|
||||
Trial status: 1 RUNNING
|
||||
Current time: 2024-06-03 12:03:39. Total running time: 30s
|
||||
Logical resource usage: 3.0/12 CPUs, 0/0 GPUs
|
||||
╭────────────────────────────────────────────────────────────────────────
|
||||
│ Trial name status iter total time (s)
|
||||
├────────────────────────────────────────────────────────────────────────
|
||||
│ PPO_CartPole-v1_7b1eb_00000 RUNNING 6 15.362
|
||||
╰────────────────────────────────────────────────────────────────────────
|
||||
───────────────────────────────────────────────────────────────────────╮
|
||||
..._sampled_lifetime ..._trained_lifetime ...episodes_lifetime │
|
||||
───────────────────────────────────────────────────────────────────────┤
|
||||
24000 24000 340 │
|
||||
───────────────────────────────────────────────────────────────────────╯
|
||||
...
|
||||
|
||||
then, you should see the experiment crashing as soon as the `--stop-reward-crash`
|
||||
has been reached:
|
||||
|
||||
```RuntimeError: Intended crash after reaching trigger return.```
|
||||
|
||||
At some point, the experiment should resume exactly where it left off (using
|
||||
the checkpoint and restored Tuner):
|
||||
|
||||
Trial status: 1 RUNNING
|
||||
Current time: 2024-06-03 12:05:00. Total running time: 1min 0s
|
||||
Logical resource usage: 3.0/12 CPUs, 0/0 GPUs
|
||||
╭────────────────────────────────────────────────────────────────────────
|
||||
│ Trial name status iter total time (s)
|
||||
├────────────────────────────────────────────────────────────────────────
|
||||
│ PPO_CartPole-v1_7b1eb_00000 RUNNING 27 66.1451
|
||||
╰────────────────────────────────────────────────────────────────────────
|
||||
───────────────────────────────────────────────────────────────────────╮
|
||||
..._sampled_lifetime ..._trained_lifetime ...episodes_lifetime │
|
||||
───────────────────────────────────────────────────────────────────────┤
|
||||
108000 108000 531 │
|
||||
───────────────────────────────────────────────────────────────────────╯
|
||||
|
||||
And if you are using the `--as-test` option, you should see a finel message:
|
||||
|
||||
```
|
||||
`env_runners/episode_return_mean` of 500.0 reached! ok
|
||||
```
|
||||
"""
|
||||
import re
|
||||
import time
|
||||
|
||||
from ray import tune
|
||||
from ray.air.integrations.wandb import WandbLoggerCallback
|
||||
from ray.rllib.algorithms.algorithm_config import AlgorithmConfig
|
||||
from ray.rllib.callbacks.callbacks import RLlibCallback
|
||||
from ray.rllib.examples.envs.classes.multi_agent import MultiAgentCartPole
|
||||
from ray.rllib.examples.utils import add_rllib_example_script_args
|
||||
from ray.rllib.policy.policy import PolicySpec
|
||||
from ray.rllib.utils.metrics import (
|
||||
ENV_RUNNER_RESULTS,
|
||||
EPISODE_RETURN_MEAN,
|
||||
)
|
||||
from ray.rllib.utils.test_utils import check_learning_achieved
|
||||
from ray.tune.registry import get_trainable_cls, register_env
|
||||
|
||||
parser = add_rllib_example_script_args(
|
||||
default_reward=500.0, default_timesteps=10000000, default_iters=2000
|
||||
)
|
||||
parser.add_argument(
|
||||
"--stop-reward-crash",
|
||||
type=float,
|
||||
default=200.0,
|
||||
help="Mean episode return after which the Algorithm should crash.",
|
||||
)
|
||||
# By default, set `args.checkpoint_freq` to 1 and `args.checkpoint_at_end` to True.
|
||||
parser.set_defaults(
|
||||
checkpoint_freq=1,
|
||||
checkpoint_at_end=True,
|
||||
)
|
||||
|
||||
|
||||
class CrashAfterNIters(RLlibCallback):
|
||||
"""Callback that makes the algo crash after a certain avg. return is reached."""
|
||||
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
# We have to delay crashing by one iteration just so the checkpoint still
|
||||
# gets created by Tune after(!) we have reached the trigger avg. return.
|
||||
self._should_crash = False
|
||||
|
||||
def on_train_result(self, *, algorithm, metrics_logger, result, **kwargs):
|
||||
# We had already reached the mean-return to crash, the last checkpoint written
|
||||
# (the one from the previous iteration) should yield that exact avg. return.
|
||||
if self._should_crash:
|
||||
raise RuntimeError("Intended crash after reaching trigger return.")
|
||||
# Reached crashing criterion, crash on next iteration.
|
||||
elif result[ENV_RUNNER_RESULTS][EPISODE_RETURN_MEAN] >= args.stop_reward_crash:
|
||||
print(
|
||||
"Reached trigger return of "
|
||||
f"{result[ENV_RUNNER_RESULTS][EPISODE_RETURN_MEAN]}"
|
||||
)
|
||||
self._should_crash = True
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
args = parser.parse_args()
|
||||
|
||||
register_env(
|
||||
"ma_cart", lambda cfg: MultiAgentCartPole({"num_agents": args.num_agents})
|
||||
)
|
||||
|
||||
# Simple generic config.
|
||||
config = (
|
||||
get_trainable_cls(args.algo)
|
||||
.get_default_config()
|
||||
.environment("CartPole-v1" if args.num_agents == 0 else "ma_cart")
|
||||
.env_runners(create_env_on_local_worker=True)
|
||||
.training(lr=0.0001)
|
||||
.callbacks(CrashAfterNIters)
|
||||
)
|
||||
|
||||
# Tune config.
|
||||
# Need a WandB callback?
|
||||
tune_callbacks = []
|
||||
if args.wandb_key:
|
||||
project = args.wandb_project or (
|
||||
args.algo.lower() + "-" + re.sub("\\W+", "-", str(config.env).lower())
|
||||
)
|
||||
tune_callbacks.append(
|
||||
WandbLoggerCallback(
|
||||
api_key=args.wandb_key,
|
||||
project=args.wandb_project,
|
||||
upload_checkpoints=False,
|
||||
**({"name": args.wandb_run_name} if args.wandb_run_name else {}),
|
||||
)
|
||||
)
|
||||
|
||||
# Setup multi-agent, if required.
|
||||
if args.num_agents > 0:
|
||||
config.multi_agent(
|
||||
policies={
|
||||
f"p{aid}": PolicySpec(
|
||||
config=AlgorithmConfig.overrides(
|
||||
lr=5e-5
|
||||
* (aid + 1), # agent 1 has double the learning rate as 0.
|
||||
)
|
||||
)
|
||||
for aid in range(args.num_agents)
|
||||
},
|
||||
policy_mapping_fn=lambda aid, *a, **kw: f"p{aid}",
|
||||
)
|
||||
|
||||
# Define some stopping criterion. Note that this criterion is an avg episode return
|
||||
# to be reached. The stop criterion does not consider the built-in crash we are
|
||||
# triggering through our callback.
|
||||
stop = {
|
||||
f"{ENV_RUNNER_RESULTS}/{EPISODE_RETURN_MEAN}": args.stop_reward,
|
||||
}
|
||||
|
||||
# Run tune for some iterations and generate checkpoints.
|
||||
tuner = tune.Tuner(
|
||||
trainable=config.algo_class,
|
||||
param_space=config,
|
||||
run_config=tune.RunConfig(
|
||||
callbacks=tune_callbacks,
|
||||
checkpoint_config=tune.CheckpointConfig(
|
||||
checkpoint_frequency=args.checkpoint_freq,
|
||||
checkpoint_at_end=args.checkpoint_at_end,
|
||||
),
|
||||
stop=stop,
|
||||
),
|
||||
)
|
||||
tuner_results = tuner.fit()
|
||||
|
||||
# Perform a very quick test to make sure our algo (upon restoration) did not lose
|
||||
# its ability to perform well in the env.
|
||||
# - Extract the best checkpoint.
|
||||
metric = f"{ENV_RUNNER_RESULTS}/{EPISODE_RETURN_MEAN}"
|
||||
best_result = tuner_results.get_best_result(metric=metric, mode="max")
|
||||
assert (
|
||||
best_result.metrics[ENV_RUNNER_RESULTS][EPISODE_RETURN_MEAN]
|
||||
>= args.stop_reward_crash
|
||||
)
|
||||
# - Change our config, such that the restored algo will have an env on the local
|
||||
# EnvRunner (to perform evaluation) and won't crash anymore (remove the crashing
|
||||
# callback).
|
||||
config.callbacks(None)
|
||||
# Rebuild the algorithm (just for testing purposes).
|
||||
test_algo = config.build()
|
||||
# Load algo's state from best checkpoint.
|
||||
test_algo.restore(best_result.checkpoint)
|
||||
# Perform some checks on the restored state.
|
||||
assert test_algo.training_iteration > 0
|
||||
# Evaluate on the restored algorithm.
|
||||
test_eval_results = test_algo.evaluate()
|
||||
assert (
|
||||
test_eval_results[ENV_RUNNER_RESULTS][EPISODE_RETURN_MEAN]
|
||||
>= args.stop_reward_crash
|
||||
), test_eval_results[ENV_RUNNER_RESULTS][EPISODE_RETURN_MEAN]
|
||||
# Train one iteration to make sure, the performance does not collapse (e.g. due
|
||||
# to the optimizer weights not having been restored properly).
|
||||
test_results = test_algo.train()
|
||||
assert (
|
||||
test_results[ENV_RUNNER_RESULTS][EPISODE_RETURN_MEAN] >= args.stop_reward_crash
|
||||
), test_results[ENV_RUNNER_RESULTS][EPISODE_RETURN_MEAN]
|
||||
# Stop the test algorithm again.
|
||||
test_algo.stop()
|
||||
|
||||
# Create a new Tuner from the existing experiment path (which contains the tuner's
|
||||
# own checkpoint file). Note that even the WandB logging will be continued without
|
||||
# creating a new WandB run name.
|
||||
restored_tuner = tune.Tuner.restore(
|
||||
path=tuner_results.experiment_path,
|
||||
trainable=config.algo_class,
|
||||
param_space=config,
|
||||
# Important to set this to True b/c the previous trial had failed (due to our
|
||||
# `CrashAfterNIters` callback).
|
||||
resume_errored=True,
|
||||
)
|
||||
# Continue the experiment exactly where we left off.
|
||||
tuner_results = restored_tuner.fit()
|
||||
|
||||
# Not sure, whether this is really necessary, but we have observed the WandB
|
||||
# logger sometimes not logging some of the last iterations. This sleep here might
|
||||
# give it enough time to do so.
|
||||
time.sleep(20)
|
||||
|
||||
if args.as_test:
|
||||
check_learning_achieved(tuner_results, args.stop_reward, metric=metric)
|
||||
@@ -0,0 +1,170 @@
|
||||
"""Example demonstrating how to load module weights for 1 of n agents from a checkpoint.
|
||||
|
||||
This example:
|
||||
- Runs a multi-agent `Pendulum-v1` experiment with >= 2 policies, p0, p1, etc..
|
||||
- Saves a checkpoint of the `MultiRLModule` every `--checkpoint-freq`
|
||||
iterations.
|
||||
- Stops the experiments after the agents reach a combined return of -800.
|
||||
- Picks the best checkpoint by combined return and restores p0 from it.
|
||||
- Runs a second experiment with the restored `RLModule` for p0 and
|
||||
a fresh `RLModule` for the other policies.
|
||||
- Stops the second experiment after the agents reach a combined return of -800.
|
||||
|
||||
|
||||
How to run this script
|
||||
----------------------
|
||||
`python [script file name].py --num-agents=2
|
||||
--checkpoint-freq=20 --checkpoint-at-end`
|
||||
|
||||
Control the number of agents and policies (RLModules) via --num-agents and
|
||||
--num-policies.
|
||||
|
||||
Control the number of checkpoints by setting `--checkpoint-freq` to a value > 0.
|
||||
Note that the checkpoint frequency is per iteration and this example needs at
|
||||
least a single checkpoint to load the RLModule weights for policy 0.
|
||||
If `--checkpoint-at-end` is set, a checkpoint will be saved at the end of the
|
||||
experiment.
|
||||
|
||||
For debugging, use the following additional command line options
|
||||
`--no-tune --num-env-runners=0`
|
||||
which should allow you to set breakpoints anywhere in the RLlib code and
|
||||
have the execution stop there for inspection and debugging.
|
||||
|
||||
For logging to your WandB account, use:
|
||||
`--wandb-key=[your WandB API key] --wandb-project=[some project name]
|
||||
--wandb-run-name=[optional: WandB run name (within the defined project)]`
|
||||
|
||||
|
||||
Results to expect
|
||||
-----------------
|
||||
You should expect a reward of -400.0 eventually being achieved by a simple
|
||||
single PPO policy. In the second run of the experiment, the MultiRLModule weights
|
||||
for policy 0 are restored from the checkpoint of the first run. The reward for a
|
||||
single agent should be -400.0 again, but the training time should be shorter
|
||||
(around 30 iterations instead of 190) due to the fact that one policy is already
|
||||
an expert from the get go.
|
||||
"""
|
||||
|
||||
from pathlib import Path
|
||||
|
||||
from ray.rllib.algorithms.callbacks import DefaultCallbacks
|
||||
from ray.rllib.core import (
|
||||
COMPONENT_LEARNER,
|
||||
COMPONENT_LEARNER_GROUP,
|
||||
COMPONENT_RL_MODULE,
|
||||
)
|
||||
from ray.rllib.core.rl_module.default_model_config import DefaultModelConfig
|
||||
from ray.rllib.examples.envs.classes.multi_agent import MultiAgentPendulum
|
||||
from ray.rllib.examples.utils import (
|
||||
add_rllib_example_script_args,
|
||||
run_rllib_example_script_experiment,
|
||||
)
|
||||
from ray.rllib.utils.metrics import (
|
||||
ENV_RUNNER_RESULTS,
|
||||
EPISODE_RETURN_MEAN,
|
||||
NUM_ENV_STEPS_SAMPLED_LIFETIME,
|
||||
)
|
||||
from ray.rllib.utils.numpy import convert_to_numpy
|
||||
from ray.rllib.utils.test_utils import check
|
||||
from ray.tune.registry import get_trainable_cls, register_env
|
||||
from ray.tune.result import TRAINING_ITERATION
|
||||
|
||||
parser = add_rllib_example_script_args(
|
||||
# Pendulum-v1 sum of 2 agents (each agent reaches -250).
|
||||
default_reward=-500.0,
|
||||
)
|
||||
parser.set_defaults(
|
||||
checkpoint_freq=1,
|
||||
num_agents=2,
|
||||
)
|
||||
# TODO (sven): This arg is currently ignored (hard-set to 2).
|
||||
parser.add_argument("--num-policies", type=int, default=2)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
args = parser.parse_args()
|
||||
|
||||
# Register our environment with tune.
|
||||
if args.num_agents > 1:
|
||||
register_env(
|
||||
"env",
|
||||
lambda _: MultiAgentPendulum(config={"num_agents": args.num_agents}),
|
||||
)
|
||||
else:
|
||||
raise ValueError(
|
||||
f"`num_agents` must be > 1, but is {args.num_agents}."
|
||||
"Read the script docstring for more information."
|
||||
)
|
||||
|
||||
assert args.checkpoint_freq > 0, (
|
||||
"This example requires at least one checkpoint to load the RLModule "
|
||||
"weights for policy 0."
|
||||
)
|
||||
|
||||
base_config = (
|
||||
get_trainable_cls(args.algo)
|
||||
.get_default_config()
|
||||
.environment("env")
|
||||
.training(
|
||||
train_batch_size_per_learner=512,
|
||||
minibatch_size=64,
|
||||
lambda_=0.1,
|
||||
gamma=0.95,
|
||||
lr=0.0003,
|
||||
vf_clip_param=10.0,
|
||||
)
|
||||
.rl_module(
|
||||
model_config=DefaultModelConfig(fcnet_activation="relu"),
|
||||
)
|
||||
)
|
||||
|
||||
# Add a simple multi-agent setup.
|
||||
if args.num_agents > 0:
|
||||
base_config.multi_agent(
|
||||
policies={f"p{i}" for i in range(args.num_agents)},
|
||||
policy_mapping_fn=lambda aid, *a, **kw: f"p{aid}",
|
||||
)
|
||||
|
||||
# Augment the base config with further settings and train the agents.
|
||||
results = run_rllib_example_script_experiment(base_config, args, keep_ray_up=True)
|
||||
|
||||
# Now swap in the RLModule weights for policy 0.
|
||||
chkpt_path = results.get_best_result().checkpoint.path
|
||||
p_0_module_state_path = (
|
||||
Path(chkpt_path) # <- algorithm's checkpoint dir
|
||||
/ COMPONENT_LEARNER_GROUP # <- learner group
|
||||
/ COMPONENT_LEARNER # <- learner
|
||||
/ COMPONENT_RL_MODULE # <- MultiRLModule
|
||||
/ "p0" # <- (single) RLModule
|
||||
)
|
||||
|
||||
class LoadP0OnAlgoInitCallback(DefaultCallbacks):
|
||||
def on_algorithm_init(self, *, algorithm, **kwargs):
|
||||
module_p0 = algorithm.get_module("p0")
|
||||
weight_before = convert_to_numpy(next(iter(module_p0.parameters())))
|
||||
algorithm.restore_from_path(
|
||||
p_0_module_state_path,
|
||||
component=(
|
||||
COMPONENT_LEARNER_GROUP
|
||||
+ "/"
|
||||
+ COMPONENT_LEARNER
|
||||
+ "/"
|
||||
+ COMPONENT_RL_MODULE
|
||||
+ "/p0"
|
||||
),
|
||||
)
|
||||
# Make sure weights were updated.
|
||||
weight_after = convert_to_numpy(next(iter(module_p0.parameters())))
|
||||
check(weight_before, weight_after, false=True)
|
||||
|
||||
base_config.callbacks(LoadP0OnAlgoInitCallback)
|
||||
|
||||
# Define stopping criteria.
|
||||
stop = {
|
||||
f"{ENV_RUNNER_RESULTS}/{EPISODE_RETURN_MEAN}": -800.0,
|
||||
f"{ENV_RUNNER_RESULTS}/{NUM_ENV_STEPS_SAMPLED_LIFETIME}": 100000,
|
||||
TRAINING_ITERATION: 100,
|
||||
}
|
||||
|
||||
# Run the experiment again with the restored MultiRLModule.
|
||||
run_rllib_example_script_experiment(base_config, args, stop=stop)
|
||||
Reference in New Issue
Block a user