Skip to content

About

No description, website, or topics provided.

Resources

Stars

8 stars

Watchers

0 watching

Forks

Latest commit

 

History

7 Commits

Folders and files

NameName
Last commit message
Last commit date
 
 
 
 
 
 
 
 

Repository files navigation

catnat-torch

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.

What is CatNat?

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 $s_i \in \mathbb{R}$ is transformed by an activation function $a(\cdot)$ into the Bernoulli probability of taking one branch at the corresponding internal node. The probability of a leaf is then the product of the branch probabilities along the path from the root to that leaf. See the figure below for a visualisation of the binary tree structure.

CatNat binary tree
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) — maps K-1 raw scores, such as neural-network outputs, to a categorical distribution over K outcomes. The activation maps each score (in $\mathbb{R}$) to a Bernoulli split probability (in $(0, 1)$). With use_logprobs=True, the tree is evaluated in log-space and returns log-probabilities (better for numerical stability); with False, it is evaluated in probability-space and returns probabilities.
  • CatNatCategorical — a torch.distributions.Categorical-compatible distribution wrapping the above. Since the advantage of CatNat is in the diagonal FIM, it is recommended to use CatNatCategorical as it does not re-normalise the probabilities.

Installation

From git:

pip install git+https://github.com/allemanenti/catnat-torch.git

Usage

1) catnat -- given K-1 scores, get categorical log-probabilities or probabilities

import 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 tree

You 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)

Probability-space output

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)

2) CatNatCategorical -- torch.distributions-compatible distribution

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()

Notes

  • CatNatCategorical does 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 what catnat() produces. Skipping the softmax / log-softmax step is what preserves the diagonal Fisher Information Matrix.

Citation

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

License

MIT License.

About

No description, website, or topics provided.

Resources

Stars

8 stars

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages