google / google/brax

checkpoint.load_config should ignore None kernel initializer

Open
#664 0 comments 1 reaction 0 assignees View on GitHub
Dominant language
Jupyter Notebook
Stars
3.2k
Forks
349
PR merge metrics
No merged PRs in 30d

Description

Kernel initializers can be set to and saved as `None`. In those cases, we should not retrieve them from `networks.KERNEL_INITIALIZER`.
This issue is probably similar to what's encountered here https://github.com/google/brax/commit/770e83711261ad1b309374c86150403d86b5a250.

For example, `mean_kernel_init_fn` could be `None`.

https://github.com/google/brax/blob/4597b9c22677efdbeed2bba00a517cbe8eb757c6/brax/training/agents/ppo/networks.py#L97

When loaded from the saved checkpoint, we will try to access `networks.KERNEL_INITIALIZER[None]`, resulting in an error:

```
Traceback (most recent call last):
File "/Users/linfeng/workspace/brax/scripts/reproduce.py", line 47, in
inference_fn = ppo_checkpoint.load_policy(
checkpoints[0],
network_factory=network_factory,
deterministic=True,
)
File "/Users/linfeng/workspace/brax/brax/training/agents/ppo/checkpoint.py", line 83, in load_policy
config = load_config(path)
File "/Users/linfeng/workspace/brax/brax/training/agents/ppo/checkpoint.py", line 71, in load_config
return checkpoint.load_config(config_path)
~~~~~~~~~~~~~~~~~~~~~~^^^^^^^^^^^^^
File "/Users/linfeng/workspace/brax/brax/training/checkpoint.py", line 229, in load_config
networks.KERNEL_INITIALIZER[init_fn_name_]
~~~~~~~~~~~~~~~~~~~~~~~~~~~^^^^^^^^^^^^^^^
KeyError: None
```

A reproducible example, with `python==3.13.0`, `brax==0.14.1` and `playground==0.1.0`, run on macOS 15.5:

```python
#!/usr/bin/env python3
import glob
import functools
from pathlib import Path
from brax.training.agents.ppo import train
from brax.training.agents.ppo import networks as ppo_networks
from brax.training.agents.ppo import checkpoint as ppo_checkpoint
from mujoco_playground import wrapper

from mujoco_playground.config import manipulation_params
from mujoco_playground import registry

env_name = "AeroCubeRotateZAxis"
env_cfg = registry.get_default_config(env_name)
env_cfg.episode_length = 10
env = registry.load(env_name, env_cfg)

ppo_params = manipulation_params.brax_ppo_config(env_name)
ppo_params.num_timesteps = 1
ppo_params.num_evals = 1
ppo_params.num_envs = 1
ppo_params.num_minibatches = 1
ppo_params.num_updates_per_batch = 1
ppo_params.batch_size = 1
ppo_training_params = dict(ppo_params)

network_factory = functools.partial(
ppo_networks.make_ppo_networks,
policy_obs_key="privileged_state",
value_obs_key="privileged_state",
policy_hidden_layer_sizes=(128,),
value_hidden_layer_sizes=(128,),
)
del ppo_training_params["network_factory"]
ppo_training_params["network_factory"] = network_factory

checkpoints_path = Path("./checkpoints").expanduser().resolve().as_posix()
make_inference_fn, params, metrics = train.train(
environment=env,
**dict(ppo_training_params),
save_checkpoint_path=checkpoints_path,
wrap_env_fn=wrapper.wrap_for_brax_training,
)

checkpoints = glob.glob(f"{checkpoints_path}/*/")

inference_fn = ppo_checkpoint.load_policy(
checkpoints[0],
network_factory=network_factory,
deterministic=True,
)
```

Contributor guide

Open the contributing guide

Assessment

This issue has not been assessed yet.

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.