PyTorch implementation of the CatNat parameterization for categorical random variables, introduced in “Beyond Softmax: A Natural Parameterization for Categorical Random Variables” (ICML 2026).
JAX implementation available at allemanenti/catnat-jax.
CatNat is a drop-in alternative to softmax for parameterizing categorical distributions over K outcomes. It is designed to improve optimization by using a natural parameterization whose Fisher Information Matrix (FIM) is diagonal.
While softmax maps K unconstrained logits to K probabilities by exponentiating and normalizing them, CatNat uses only K - 1 unconstrained scores. These scores are assigned to the internal nodes of a binary tree whose leaves correspond to the K categorical outcomes.
Each score

Figure 1. Each score s_i is placed at an
internal node. The activation a(s_i) gives the probability of
taking the left branch, while 1 - a(s_i) gives the probability of
taking the right branch. A leaf probability is the product of the branch
probabilities along its path from the root.
This package provides two entry points:
-
catnat(scores, activation="sigmoid", use_logprobs=True)— mapsK-1raw scores, such as neural-network outputs, to a categorical distribution overKoutcomes. Theactivationmaps each score (in$\mathbb{R}$ ) to a Bernoulli split probability (in $(0, 1)$). Withuse_logprobs=True, the tree is evaluated in log-space and returns log-probabilities (better for numerical stability); withFalse, it is evaluated in probability-space and returns probabilities. -
CatNatCategorical— atorch.distributions.Categorical-compatible distribution wrapping the above. Since the advantage ofCatNatis in the diagonal FIM, it is recommended to useCatNatCategoricalas it does not re-normalise the probabilities.
From git:
pip install git+https://github.com/allemanenti/catnat-torch.gitimport torch
from catnat_torch import catnat
# K - 1 = 4 internal scores -> K = 5 categorical log-probabilities
scores = torch.tensor([0.1, -0.2, 0.3, 0.4])
logprobs = catnat(scores) # default: sigmoid, log-space
logprobs_nat = catnat(scores, "natural") # alternative built-in activation
# to get probabilities:
probs = torch.exp(catnat(scores)) # Stable way, uses log-space in the binary tree
# or
probs = catnat(scores, use_logprobs=False) # May be less stable for extreme scores,
# as it uses probabilities in the treeYou can also pass a custom callable instead of a string. In log mode the callable must return
(log_p_left, log_p_right):
import torch.nn.functional as F
def log_sigmoid_split(x):
return F.logsigmoid(x), F.logsigmoid(-x)
logprobs = catnat(scores, log_sigmoid_split)Set use_logprobs=False to get normalised probabilities. The callable signature
becomes x -> p (Bernoulli probability):
probs = catnat(scores, use_logprobs=False) # shape (5,), sums to 1.0
probs = catnat(scores, "natural", use_logprobs=False)
probs = catnat(scores, torch.sigmoid, use_logprobs=False)CatNatCategorical does not re-normalise its inputs, so feeding it
log-probabilities directly preserves the diagonal FIM that motivates CatNat:
import torch
from catnat_torch import catnat, CatNatCategorical
scores = torch.tensor([0.1, -0.2, 0.3, 0.4])
dist = CatNatCategorical(logprobs=catnat(scores)) # default log-prob path
# Get samples
samples = dist.sample((8,)) # sample 8 outcomes
# Get log-probabilities of samples
log_probs = dist.log_prob(samples) # log-prob of each sample
# Get entropy of the distribution
entropy = dist.entropy()CatNatCategoricaldoes not re-normalise its inputs. The constructor expects already-valid probabilities (non-negative, finite, sum to 1) or normalised log-probabilities (logsumexp == 0), which is exactly whatcatnat()produces. Skipping the softmax / log-softmax step is what preserves the diagonal Fisher Information Matrix.
If you use catnat-torch in your research, please cite the original paper:
@inproceedings{manenti2026beyond,
title={Beyond Softmax: A Natural Parameterization for Categorical Random Variables},
author={Manenti, Alessandro and Alippi, Cesare},
booktitle={International Conference on Machine Learning},
year={2026},
organization={PMLR}
}
MIT License.