mean-conc-beta
Beta distribution parameterized by mean and concentration for bounded continuous action spaces in reinforcement learning.
Install
$ pip install mean-conc-beta
Usage
import torch
from mean_conc_beta import Beta
# defaults to (-1., 1.), most continuous action spaces are symmetric centered on 0, sometimes off by a scale
# but you can pass in custom bounds, e.g. Beta((-0.4, 0.4))
beta = Beta()
# network output: (batch, num_actions, 2) for raw mean and concentration
params = torch.randn(16, 4, 2, requires_grad = True)
# distribution on the action bounds
dist = beta(params)
# sample actions
actions = dist.sample()
actions_reparam = dist.rsample()
# joint log prob and entropy over the action dimensions
log_prob = dist.log_prob(actions).sum(dim = -1)
entropy = dist.entropy().sum(dim = -1)
# auto-rescale actions directly into env.step, with optional clipping
env_step = beta.rescale_env_step(env.step, target_range = (-2.1, 2.1), clip = (-2., 2.))
next_obs, reward, term, trunc, info = env_step(actions)
# behavior cloning with mse loss on mean
expert_actions = torch.rand(16, 4) * 2. - 1.
pred_mean = beta.mean(params)
bc_loss = (pred_mean - expert_actions).pow(2).mean()
bc_loss.backward()
Per-action bounds
Bounds must have shape (2,) for a single (low, high) pair, or (num_actions, 2) for a stack of per-action pairs.
beta = Beta(bounds = [(-1., 1.), (0., 1.), (-2., 2.)])
params = torch.randn(16, 3, 2)
actions = beta(params).sample()
Mean Squashing
The raw mean is squashed onto the bounds with LeakyTanh by default - the exact tanh forward, with the backward gradient floored at leak so a policy saturated at a bound can still be pulled back. It can be swapped for plain tanh or any other squash:
from mean_conc_beta import Beta, LeakyTanh
beta = Beta(squash_fn = 'tanh')
beta = Beta(squash_fn = LeakyTanh(leak = 0.1))
beta = Beta(squash_fn = 'softsign')
beta = Beta(squash_fn = lambda x: x / (1. + x ** 2).sqrt())
Citations
@article{Ferrari2004BetaRF,
title = {Beta Regression for Modelling Rates and Proportions},
author = {Silvia L. P. Ferrari and Francisco Cribari-Neto},
journal = {Journal of Applied Statistics},
year = {2004},
volume = {31},
pages = {799 - 815}
}
@inproceedings{Chou2017TheBP,
title = {The Beta Policy for Continuous Reinforcement Learning},
author = {Po-Wei Chou and Daniel Maturana and Sebastian Scherer},
booktitle = {International Conference on Machine Learning},
year = {2017}
}
Download files
Download the file for your platform. If you're not sure which to choose, learn more about installing packages.
Source Distribution
Built Distribution
Filter files by name, interpreter, ABI, and platform.
If you're not sure about the file name format, learn more about wheel file names.
Copy a direct link to the current filters
File details
Details for the file mean_conc_beta-0.1.2.tar.gz.
File metadata
- Download URL: mean_conc_beta-0.1.2.tar.gz
- Upload date:
- Size: 10.0 kB
- Tags: Source
- Uploaded using Trusted Publishing? No
- Uploaded via:
uv/0.8.17
File hashes
| Algorithm | Hash digest | |
|---|---|---|
| SHA256 |
96bf65be19dd51861f013c56829f3e584fb1dee866c471c8df4c96b1945f05df
|
|
| MD5 |
841ca71c75bef78f07389edf0bc269a2
|
|
| BLAKE2b-256 |
6cfac6f2e0f5d22d585c13f68c65f7e37eff84ff8be1f645dad1d65268042901
|
File details
Details for the file mean_conc_beta-0.1.2-py3-none-any.whl.
File metadata
- Download URL: mean_conc_beta-0.1.2-py3-none-any.whl
- Upload date:
- Size: 9.4 kB
- Tags: Python 3
- Uploaded using Trusted Publishing? No
- Uploaded via:
uv/0.8.17
File hashes
| Algorithm | Hash digest | |
|---|---|---|
| SHA256 |
abe98f1f838847e592c58e7f4ccafd5b54cdef4745fb7eef2faffeeac5059f11
|
|
| MD5 |
26ba37ff4877727d671f78d392582a5b
|
|
| BLAKE2b-256 |
a8602b639c4950a1030b27ad44bd0450339e5778cc558a29ec0ce330c89502c5
|