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()
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
mean_conc_beta-0.0.11.tar.gz
(7.8 kB
view details)
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.0.11.tar.gz.
File metadata
- Download URL: mean_conc_beta-0.0.11.tar.gz
- Upload date:
- Size: 7.8 kB
- Tags: Source
- Uploaded using Trusted Publishing? No
- Uploaded via:
uv/0.8.17
File hashes
| Algorithm | Hash digest | |
|---|---|---|
| SHA256 |
ad8971af14a9865cb4886dc81968dd5be77d867a0abae1076ecc8589965199bf
|
|
| MD5 |
d4df0c643d0270968f347b71850a1913
|
|
| BLAKE2b-256 |
0eb5fd4eaa8908085552244c02dd9bc3c9994302e7f23f808f2aa10473390856
|
File details
Details for the file mean_conc_beta-0.0.11-py3-none-any.whl.
File metadata
- Download URL: mean_conc_beta-0.0.11-py3-none-any.whl
- Upload date:
- Size: 7.2 kB
- Tags: Python 3
- Uploaded using Trusted Publishing? No
- Uploaded via:
uv/0.8.17
File hashes
| Algorithm | Hash digest | |
|---|---|---|
| SHA256 |
7b57bae41283438aacc6ee207ac84f29f412775a72d6829c0801e2aa58285c21
|
|
| MD5 |
014482ee9b3846030cc56958d36aee17
|
|
| BLAKE2b-256 |
2d86ffdd3bc62b817f2483786e2739e8fbdbb7cdc2c5379cd066532b55c5fdec
|