Skip to main content

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}
}

Release files for mean-conc-beta 0.2.2

For a detailed explanation of source distributions (sdists) and built distributions (wheels), please see the package formats documentation.

Built distribution (wheel)

Table of built distributions (wheels) for mean-conc-beta 0.2.2
File Interpreter ABI Platform
mean_conc_beta-0.2.2-py3-none-any.whl Python 3 none any Details

Release files / mean_conc_beta-0.2.2-py3-none-any.whl

Download URL mean_conc_beta-0.2.2-py3-none-any.whl
Size 9.6 kB
Tags Python 3
SHA-256 checksum
How to use checksums
7692f433ebebed24e89538f25cfabddd70974d121914553e7cfd3b6729542867
BLAKE2b-256 checksum
How to use checksums
24c244ee10c06659d405a8b252dcbfa4dbd6797bea5f9ac1be5f847a2e9a2820
Upload date
Uploaded using Trusted Publishing?
What is trusted publishing?
No
Uploaded via uv/0.8.17

Release history Release notifications | RSS feed

This release

0.2.2 This release

1 release file

0.2.1

2 release files

0.2.0

2 release files

0.1.5

2 release files

0.1.4

2 release files

0.1.2

2 release files

0.1.1

2 release files

0.1.0

2 release files

0.0.12

2 release files

0.0.11

2 release files

0.0.9

2 release files

0.0.8

2 release files

0.0.7

2 release files

0.0.6

2 release files

0.0.5

2 release files

0.0.4

2 release files

0.0.3

2 release files

0.0.2

2 release files

0.0.1

2 release files

Anthropic, PBC Visionary sponsor Bloomberg Visionary sponsor Hudson River Trading Visionary sponsor Meta Visionary sponsor NVIDIA Visionary sponsor Microsoft Sustainability sponsor Depot Continuous Integration AWS Cloud computing and Security Sponsor Datadog Monitoring Fastly CDN Google Download Analytics Sentry Error logging StatusPage Status page