google-deepmind / google-deepmind/rlax

Stop action gradient in policy gradient loss

Open
#109 6 comments 1 reaction 0 assignees View on GitHub
Dominant language
Python
Stars
1.4k
Forks
109
Avg merge
2h 15m
Merged PRs (30d)
1

Description

The current implementation of `policy_gradient_loss` is:
```python
log_pi_a_t = distributions.softmax().logprob(a_t, logits_t)
adv_t = jax.lax.select(use_stop_gradient, jax.lax.stop_gradient(adv_t), adv_t)
loss_per_timestep = -log_pi_a_t * adv_t
```
It's good that the gradients are already stopped around the advantages, but they should also be stopped around the actions to ensure an unbiased gradient estimator.

This is important when the actions are sampled as part of the training graph (MPO-style algos, imagination training with world models) rather than coming from the replay buffer, and the actor distribution implements a gradient for `sample()` (e.g. gaussian, or straight-through categoricals).

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.