assert_is_replicated in Analytic policy gradients training
- Dominant language
- Jupyter Notebook
- Stars
- 3.2k
- Forks
- 349
- PR merge metrics
- No merged PRs in 30d
Description
When I try to use a 4-gpus machine to run the Analytic policy gradients training in parallel, it reports an AssertionError in `brax/training/agents/apg/train.py` line 255. Seems that it is because `training_state` becomes different on the devices while it should be replicated.
I only make minimum change according to the example training code.
```
import functools
from datetime import datetime
# from brax.training.agents.apg.train import train as apgtrain
from train import train as apgtrain
from brax import envs
env_name = 'humanoidstandup' # @param ['ant', 'halfcheetah', 'hopper', 'humanoid', 'humanoidstandup', 'inverted_pendulum', 'inverted_double_pendulum', 'pusher', 'reacher', 'walker2d']
backend = 'generalized' # @param ['generalized', 'positional', 'spring']
env = envs.get_environment(env_name=env_name,
backend=backend)
train_fn = {
'humanoidstandup': functools.partial(apgtrain, episode_length=320,
action_repeat=1,
num_envs=16,
num_eval_envs=4,
learning_rate = 1e-4,
seed = 0,
max_gradient_norm = 1e8,
num_evals = 10,
normalize_observations = True,
deterministic_eval = False)
}[env_name]
xdata, ydata = [], []
times = [datetime.now()]
def progress(num_steps, metrics):
times.append(datetime.now())
print(num_steps, metrics['eval/episode_reward'])
print("begin")
make_inference_fn, params, _ = train_fn(environment=env, progress_fn=progress)
print("end")
```
To make the error comes sooner, I add `pmap.assert_is_replicated(training_state)` in the iteration of `brax/training/agents/apg/train.py`.
```
for it in range(num_evals_after_init):
logging.info('starting iteration %s %s', it, time.time() - xt)
# optimization
epoch_key, local_key = jax.random.split(local_key)
epoch_keys = jax.random.split(epoch_key, local_devices_to_use)
(training_state,
training_metrics) = training_epoch_with_timing(training_state, epoch_keys)
######################## I add it here #############################
pmap.assert_is_replicated(training_state)
####################################################################
if process_id == 0:
# Run evals.
metrics = evaluator.run_evaluation(
_unpmap(
(training_state.normalizer_params, training_state.policy_params)),
training_metrics)
logging.info(metrics)
progress_fn(it + 1, metrics)
```
And the full output is:
```
begin
0 2238.8042
1 2367.4116
Traceback (most recent call last):
File "xxxxx.py", line 36, in
make_inference_fn, params, _ = train_fn(environment=env, progress_fn=progress)
File "/home/vipuser/playbrax/train.py", line 227, in train
pmap.assert_is_replicated(training_state)
File "/home/vipuser/playbrax/brax/brax/training/pmap.py", line 70, in assert_is_replicated
assert jax.pmap(f, axis_name='i')(x)[0], debug
AssertionError: None
```
If I use `from brax.training.agents.apg.train import train as apgtrain`, the full output will become:
```
begin
0 2233.8481
1 2273.3516
2 2460.1377
3 2319.5432
4 2250.9502
5 2289.2446
6 nan
7 nan
8 nan
9 nan
Traceback (most recent call last):
File "xxxxx.py", line 36, in
make_inference_fn, params, _ = train_fn(environment=env, progress_fn=progress)
File "/home/vipuser/playbrax/train.py", line 227, in train
pmap.assert_is_replicated(training_state)
File "/home/vipuser/playbrax/brax/brax/training/pmap.py", line 255, in assert_is_replicated
assert jax.pmap(f, axis_name='i')(x)[0], debug
AssertionError: None
```
Contributor guide
Assessment
This issue has not been assessed yet.