API reference¶
The curated public API is re-exported at the top level (e.g. qjax.q_log); the
canonical definitions live in the submodules documented below. Every entry on
this page is rendered from the source docstrings by
mkdocstrings, so it cannot disagree with the
installed version.
How to read these pages
Signatures show the annotations from the source. Arguments, returns, and
raises come from the Google-style docstring sections — the same text help()
prints in a REPL. Click Source on any entry to see the implementation.
Core — deformed functions¶
functions
¶
q-deformed elementary functions and the Tsallis q-algebra.
This module implements the two foundational maps of non-extensive statistics,
the q-logarithm and q-exponential, together with the q-deformed
arithmetic they induce. Each function is a pure, differentiable JAX expression
that recovers its Boltzmann–Gibbs counterpart as q -> 1.
Definitions¶
The q-logarithm and its inverse, the q-exponential, are
where \([\cdot]_+ = \max(\cdot, 0)\). Both reduce to \(\ln\) and \(\exp\) as \(q \to 1\).
Numerical form¶
Evaluating \((x^{1-q} - 1)/(1-q)\) directly loses catastrophically to
cancellation as \(q \to 1\) — in float32 the relative error peaks near
4e-3 around q = 1.00001. Both functions are therefore written through
jax.numpy.expm1 / jax.numpy.log1p, which are exact in that
regime:
The shared factor is the entire function \((e^t - 1)/t \to 1\) as
\(t \to 0\), so no q = 1 special case is needed at all: the classical
limit falls out of the same expression. Near \(t = 0\) a short Taylor series
replaces the ratio, which keeps not only the value but also the derivative with
respect to q correct — a hard where(q == 1, ...) branch would return a
q-independent expression and hence a spurious zero q-gradient.
q_log
¶
q_log(x: Array, q: Scalar) -> Array
q-logarithm \(\ln_q(x) = (x^{1-q} - 1)/(1-q)\).
Recovers the natural logarithm as q -> 1, continuously and with the
correct derivative in q (see the module docstring).
At x = 0 the limit is -1/(1-q) for q < 1 and -inf for
q >= 1. Negative x is outside the domain and yields NaN.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
x
|
Array
|
Non-negative input, any shape. |
required |
q
|
Scalar
|
Entropic index (scalar). |
required |
Returns:
| Type | Description |
|---|---|
Array
|
Element-wise |
Source code in qjax/core/functions.py
q_exp
¶
q_exp(x: Array, q: Scalar) -> Array
q-exponential \(\exp_q(x) = [1 + (1-q)x]_+^{1/(1-q)}\).
Inverse of q_log and the limit of math.exp as q -> 1.
Past the Tsallis cut-off — that is, wherever 1 + (1-q)x <= 0 — the
exponent 1/(1-q) decides the value: the result is 0 for q < 1
(positive exponent) and +inf for q > 1 (negative exponent), matching
the genuine divergence of the q-exponential there.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
x
|
Array
|
Input, any shape. |
required |
q
|
Scalar
|
Entropic index (scalar). |
required |
Returns:
| Type | Description |
|---|---|
Array
|
Element-wise |
Source code in qjax/core/functions.py
q_add
¶
q_add(a: Array, b: Array, q: Scalar) -> Array
q-addition \(a \oplus_q b = a + b + (1-q)\,a\,b\).
The deformed sum satisfies q_log(x*y) = q_add(q_log(x), q_log(y)) and
reduces to ordinary addition as q -> 1.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
a
|
Array
|
First operand. |
required |
b
|
Array
|
Second operand. |
required |
q
|
Scalar
|
Entropic index (scalar). |
required |
Returns:
| Type | Description |
|---|---|
Array
|
Element-wise |
Source code in qjax/core/functions.py
q_diff
¶
q_diff(a: Array, b: Array, q: Scalar) -> Array
q-subtraction \(a \ominus_q b = (a - b)/(1 + (1-q)b)\).
Inverse of q_add in its first argument: q_add(q_diff(a, b), b) == a.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
a
|
Array
|
First operand. |
required |
b
|
Array
|
Second operand. |
required |
q
|
Scalar
|
Entropic index (scalar). |
required |
Returns:
| Type | Description |
|---|---|
Array
|
Element-wise |
Source code in qjax/core/functions.py
q_prod
¶
q_prod(a: Array, b: Array, q: Scalar) -> Array
q-product \(a \otimes_q b = [a^{1-q} + b^{1-q} - 1]_+^{1/(1-q)}\).
Dual to q_add: it satisfies q_exp(x+y) = q_prod(q_exp(x), q_exp(y))
and reduces to ordinary multiplication as q -> 1. Defined for a, b >= 0.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
a
|
Array
|
First operand (non-negative). |
required |
b
|
Array
|
Second operand (non-negative). |
required |
q
|
Scalar
|
Entropic index (scalar). |
required |
Returns:
| Type | Description |
|---|---|
Array
|
Element-wise |
Source code in qjax/core/functions.py
q_div
¶
q_div(a: Array, b: Array, q: Scalar) -> Array
q-division \(a \oslash_q b = [a^{1-q} - b^{1-q} + 1]_+^{1/(1-q)}\).
Inverse of q_prod in its first argument and the limit of ordinary
division as q -> 1. Defined for a, b >= 0.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
a
|
Array
|
Numerator (non-negative). |
required |
b
|
Array
|
Denominator (non-negative). |
required |
q
|
Scalar
|
Entropic index (scalar). |
required |
Returns:
| Type | Description |
|---|---|
Array
|
Element-wise |
Source code in qjax/core/functions.py
Core — entropy and divergences¶
entropy
¶
Tsallis entropy, cross-entropy, and relative entropy (q-divergence).
These information measures generalize Shannon's entropy and the
Kullback–Leibler divergence through the entropic index q, recovering them in
the limit q -> 1. They are the natural objective functions for
non-extensive learning.
tsallis_entropy
¶
Tsallis entropy \(S_q(p) = (1 - \sum_i p_i^q)/(q - 1)\).
Recovers the Shannon entropy \(-\sum_i p_i \log p_i\) as q -> 1.
The entropy is concave in p and non-negative for probability vectors.
Computed in the equivalent form \(\sum_i p_i \ln_q(1/p_i)\), which
agrees with the definition above whenever p sums to one and, unlike it,
stays finite and correctly differentiable in q at q = 1. The two
forms differ on an unnormalized p, for which the definition above is
genuinely singular at q = 1.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
p
|
Array
|
Probability mass values along |
required |
q
|
Scalar
|
Entropic index (scalar). |
required |
axis
|
int
|
Axis over which the distribution is defined. |
-1
|
Returns:
| Type | Description |
|---|---|
Array
|
Tsallis entropy reduced over |
Source code in qjax/core/entropy.py
tsallis_cross_entropy
¶
Tsallis cross-entropy \(H_q(y, p) = -\sum_i y_i \ln_q p_i\).
A drop-in q-deformed classification loss. With one-hot y it reduces
to -q_log(p_correct, q), and to the standard cross-entropy as q -> 1.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
p
|
Array
|
Predicted probabilities along |
required |
y
|
Array
|
Target distribution along |
required |
q
|
Scalar
|
Entropic index (scalar). |
required |
axis
|
int
|
Axis over which the distributions are defined. |
-1
|
Returns:
| Type | Description |
|---|---|
Array
|
Tsallis cross-entropy reduced over |
Source code in qjax/core/entropy.py
tsallis_divergence
¶
Tsallis relative entropy \(D_q(p\,\|\,r)\).
Defined as \(D_q(p\|r) = \big(\sum_i p_i^q r_i^{1-q} - 1\big)/(q - 1)\),
equivalently \(-\sum_i p_i \ln_q(r_i / p_i)\). Recovers the
Kullback–Leibler divergence as q -> 1 and is non-negative.
The second form is the one evaluated: it agrees with the first whenever
p sums to one and, unlike it, stays correctly differentiable in q at
q = 1.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
p
|
Array
|
First distribution along |
required |
r
|
Array
|
Second (reference) distribution along |
required |
q
|
Scalar
|
Entropic index (scalar). |
required |
axis
|
int
|
Axis over which the distributions are defined. |
-1
|
Returns:
| Type | Description |
|---|---|
Array
|
Tsallis divergence reduced over |
Source code in qjax/core/entropy.py
Core — the q-Gaussian distribution¶
distributions
¶
The q-Gaussian distribution.
The q-Gaussian maximizes Tsallis entropy under a fixed second moment, just
as the Gaussian maximizes Shannon entropy. Its density is
with normalization constant \(C_q\). It interpolates between heavy-tailed
distributions (1 < q < 3; Student-t like) and compactly supported ones
(q < 1), recovering the Gaussian as q -> 1.
normalization
¶
normalization(q: Scalar) -> Array
Normalization constant \(C_q\) of the unit-\beta q-Gaussian.
Defined piecewise (see Tsallis, 2009) so that \(\int p(x)\,dx = 1\):
q < 1: \(C_q = \frac{2\sqrt{\pi}\,\Gamma(1/(1-q))} {(3-q)\sqrt{1-q}\,\Gamma(\tfrac{3-q}{2(1-q)})}\)q = 1: \(C_q = \sqrt{\pi}\)1<q<3: \(C_q = \frac{\sqrt{\pi}\,\Gamma(\tfrac{3-q}{2(q-1)})} {\sqrt{q-1}\,\Gamma(1/(q-1))}\)
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
q
|
Scalar
|
Entropic index (scalar), required to satisfy |
required |
Returns:
| Type | Description |
|---|---|
Array
|
The scalar normalization constant |
Source code in qjax/core/distributions.py
q_gaussian_pdf
¶
q_gaussian_pdf(
x: Array, q: Scalar = 1.0, beta: Scalar = 1.0
) -> Array
Density of the q-Gaussian, \(\sqrt{\beta}/C_q\,\exp_q(-\beta x^2)\).
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
x
|
Array
|
Evaluation points, any shape. |
required |
q
|
Scalar
|
Entropic index (scalar), |
1.0
|
beta
|
Scalar
|
Inverse-width parameter |
1.0
|
Returns:
| Type | Description |
|---|---|
Array
|
Density values, same shape as |
Source code in qjax/core/distributions.py
q_gaussian_logpdf
¶
q_gaussian_logpdf(
x: Array, q: Scalar = 1.0, beta: Scalar = 1.0
) -> Array
Log-density of the q-Gaussian.
Computed as \(\tfrac12\log\beta - \log C_q + \log\!\exp_q(-\beta x^2)\).
Returns -inf outside the (compact) support when q < 1, where the
density vanishes.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
x
|
Array
|
Evaluation points, any shape. |
required |
q
|
Scalar
|
Entropic index (scalar), |
1.0
|
beta
|
Scalar
|
Inverse-width parameter |
1.0
|
Returns:
| Type | Description |
|---|---|
Array
|
Log-density values, same shape as |
Source code in qjax/core/distributions.py
sample
¶
Draw q-Gaussian samples for 1 <= q < 3.
For 1 < q < 3 the q-Gaussian is a rescaled Student-t: with
\(\nu = (3-q)/(q-1)\) degrees of freedom,
where \(T_\nu = Z/\sqrt{W/\nu}\) with \(Z \sim \mathcal N(0,1)\) and
\(W \sim \chi^2_\nu\). This yields exactly the family variance
\(1/((5-3q)\beta)\) for q < 5/3. At q = 1 the Gaussian
\(Z/\sqrt{2\beta}\) is returned.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
key
|
Array
|
A |
required |
q
|
Scalar
|
Entropic index (scalar), |
1.0
|
beta
|
Scalar
|
Inverse-width parameter |
1.0
|
shape
|
tuple[int, ...]
|
Output shape. |
()
|
Returns:
| Type | Description |
|---|---|
Array
|
Samples of the requested |
Source code in qjax/core/distributions.py
Core — activations (entmax)¶
activations
¶
Tsallis entmax: the q-deformed softmax / sparsemax family.
entmax is the probability mapping obtained by regularizing the
maximum-score problem with Tsallis entropy:
with Tsallis entropy \(S_q^{T}(p) = \tfrac{1}{q(q-1)}(1 - \sum_i p_i^q)\). Its solution has the closed form (Peters, Niculae & Martins, 2019)
where the threshold \(\tau\) enforces \(\sum_i p_i = 1\). For q = 1
it is the ordinary softmax; q = 2 is sparsemax. Larger q gives
sparser distributions, q < 1 gives distributions denser (higher entropy)
than softmax, and q -> 0^+ approaches the uniform distribution.
Differentiation¶
The threshold is located by bisection, but the mapping is not differentiated
through that loop — doing so yields a wrong (and non-symmetric) Jacobian. The
solve is wrapped in jax.lax.stop_gradient and the derivative is supplied
by a jax.custom_jvp rule built from the implicit function theorem.
Writing \(s_i = p_i^{2-q}\) and \(h_i = -p_i \log p_i\) (both zero off the support), and letting \(T(v) = v - s\,\sum_j v_j / \sum_j s_j\), both tangents share a single form:
so the Jacobian with respect to z is
\(J = \mathrm{diag}(s) - s s^\top / \sum_i s_i\) — symmetric, positive
semi-definite, and annihilating the all-ones vector (entmax is invariant to
a constant shift of z). Because custom_jvp is used rather than
custom_vjp, forward mode, reverse mode, and higher-order derivatives are all
exact.
tsallis_entmax
¶
Tsallis entmax over a simplex axis.
Solves for the threshold tau such that sum([(q-1)z - tau]_+^{1/(q-1)})
equals one. q = 1 short-circuits to a numerically stable softmax.
Gradients are exact: the threshold search is not differentiated through, and
a jax.custom_jvp rule supplies the implicit-function derivative with
respect to both z and q. The Jacobian in z is
diag(s) - s s^T / sum(s) with s = p ** (2 - q).
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
z
|
Array
|
Scores / logits; the distribution is formed along |
required |
q
|
Scalar
|
Entropic index (scalar), |
2.0
|
axis
|
int
|
Axis over which to normalize. |
-1
|
num_iters
|
int
|
Number of bisection steps for the threshold search. A Newton polish step follows, so the result is accurate even for small values. |
50
|
Returns:
| Type | Description |
|---|---|
Array
|
Probabilities with the same shape as |
Raises:
| Type | Description |
|---|---|
ValueError
|
If |
Source code in qjax/core/activations.py
Neural-network building blocks¶
Framework-agnostic pieces for Tsallis models: everything here operates on plain arrays and pytrees, so it composes with Flax, Equinox, Haiku, or hand-rolled JAX without adding a dependency on any of them.
reparam
¶
Keeping a learnable entropic index inside its valid range.
Treating q as a free parameter and optimizing it directly does not work: the
Tsallis primitives are undefined at q <= 0, the q-Gaussian requires
q < 3, and gradient descent will happily step outside either bound. The
standard remedy — used verbatim in six of this repository's examples before it
lived here — is to optimize an unconstrained real q_raw and squash it:
bounded_q
¶
bounded_q(
q_raw: Scalar, lo: Scalar = 1.0, hi: Scalar = 3.0
) -> Array
Map an unconstrained parameter to an entropic index in (lo, hi).
The map is smooth and strictly monotone, so gradients flow to q_raw
everywhere and the optimizer can never leave the interval. q_raw = 0
corresponds to the midpoint.
The interval is open in exact arithmetic but closed in floating point: the
sigmoid saturates to exactly 0 or 1 once |q_raw| exceeds roughly
37 (float64) or 17 (float32), so q can reach lo or hi
exactly. Choose lo strictly inside the valid domain — q > 0 for the
Tsallis primitives, q < 3 for the q-Gaussian — rather than relying
on strict inequality here.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
q_raw
|
Scalar
|
Unconstrained real parameter, any shape. |
required |
lo
|
Scalar
|
Lower bound of the open interval (exclusive). |
1.0
|
hi
|
Scalar
|
Upper bound of the open interval (exclusive). |
3.0
|
Returns:
| Type | Description |
|---|---|
Array
|
The entropic index, same shape as |
Example
import jax.numpy as jnp from qjax.nn import bounded_q float(bounded_q(jnp.asarray(0.0), 1.0, 3.0)) 2.0
Source code in qjax/nn/reparam.py
inverse_bounded_q
¶
inverse_bounded_q(
q: Scalar, lo: Scalar = 1.0, hi: Scalar = 3.0
) -> Array
Invert bounded_q to initialize q_raw at a chosen q.
Useful for starting training from a meaningful index — q = 1 for a
softmax-like attention map, say — rather than from the interval midpoint.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
q
|
Scalar
|
Target entropic index, strictly inside |
required |
lo
|
Scalar
|
Lower bound used by |
1.0
|
hi
|
Scalar
|
Upper bound used by |
3.0
|
Returns:
| Type | Description |
|---|---|
Array
|
The |
Source code in qjax/nn/reparam.py
attention
¶
Attention whose normalizer is tsallis_entmax.
Replacing the softmax in an attention block with entmax makes the sparsity
of the attention map a property of the entropic index rather than a fixed choice:
q = 1 reproduces ordinary softmax attention, q = 2 gives sparsemax
(most positions receive exactly zero weight), and q < 1 spreads the weight
more evenly than softmax. Because q is differentiable, it can be learned
jointly with the rest of the network — see bounded_q.
entmax_attention
¶
entmax_attention(
queries: Array,
keys: Array,
values: Array,
q: Scalar = 2.0,
mask: Array | None = None,
scale: Scalar | None = None,
num_iters: int = 50,
) -> tuple[Array, Array]
Scaled dot-product attention normalized by tsallis_entmax.
Computes entmax_q(Q K^T / sqrt(d)) V over the last (key) axis. Leading
axes broadcast, so the same call serves single-head, multi-head, and batched
inputs.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
queries
|
Array
|
Query vectors, shape |
required |
keys
|
Array
|
Key vectors, shape |
required |
values
|
Array
|
Value vectors, shape |
required |
q
|
Scalar
|
Entropic index (scalar), |
2.0
|
mask
|
Array | None
|
Optional boolean array broadcastable to the score shape
|
None
|
scale
|
Scalar | None
|
Divisor applied to the scores. Defaults to |
None
|
num_iters
|
int
|
Bisection steps for the |
50
|
Returns:
| Type | Description |
|---|---|
Array
|
A |
Array
|
|
tuple[Array, Array]
|
|
Example
import jax.numpy as jnp from qjax.nn import entmax_attention q_vec = jnp.ones((2, 4)) k = jnp.ones((2, 5, 4)) v = jnp.ones((2, 5, 3)) context, attn = entmax_attention(q_vec, k, v, q=2.0) context.shape, attn.shape ((2, 3), (2, 5))
Source code in qjax/nn/attention.py
losses
¶
Loss functions built on the Tsallis information measures.
The Tsallis cross-entropy \(H_q(y, p) = -\sum_i y_i \ln_q p_i\) is a drop-in
replacement for the usual cross-entropy that recovers it at q = 1.
Its practical appeal is robustness to label noise, and that lives at q < 1.
There the q-logarithm is bounded below,
so a confidently wrong prediction — the signature of a mislabelled example --
contributes at most \(1/(1-q)\) instead of diverging, and its gradient is
capped with it. At q = 1 the penalty is unbounded, and for q > 1 it grows
faster than the logarithm (like \(p^{1-q}\)), which sharpens the model on
clean data at the cost of amplifying bad labels.
tsallis_cross_entropy_loss
¶
tsallis_cross_entropy_loss(
logits_or_probs: Array,
targets: Array,
q: Scalar = 1.0,
from_logits: bool = True,
normalizer_q: Scalar | None = None,
axis: int = -1,
reduction: str = "mean",
) -> Array
q-deformed cross-entropy loss.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
logits_or_probs
|
Array
|
Unnormalized scores if |
required |
targets
|
Array
|
Target distribution along |
required |
q
|
Scalar
|
Entropic index of the loss. |
1.0
|
from_logits
|
bool
|
Whether to normalize the input first. |
True
|
normalizer_q
|
Scalar | None
|
Entropic index of the |
None
|
axis
|
int
|
Axis holding the class distribution. |
-1
|
reduction
|
str
|
|
'mean'
|
Returns:
| Type | Description |
|---|---|
Array
|
The reduced loss, or the per-example losses when |
Raises:
| Type | Description |
|---|---|
ValueError
|
If |
Note
A sparse normalizer (normalizer_q > 1) can assign exactly zero
probability to the true class, for which the loss is a genuine +inf.
That is mathematically correct but hostile to optimization; either pair a
deformed loss with normalizer_q=1.0 or keep q close to 1
early in training.
Source code in qjax/nn/losses.py
autoregressive
¶
A masked autoregressive network over binary spins (MADE).
Variational methods in statistical mechanics need a distribution over \(2^N\) spin configurations that can be both sampled and evaluated exactly -- the second is what a mean-field ansatz gives up and what makes a variational free energy computable without a nested Monte Carlo estimate. An autoregressive factorization
gives both: sampling is \(N\) sequential passes, but the log-probability of a given configuration is a single pass, because every conditional is read off the same forward computation.
MADE (Germain et al., 2015) enforces the factorization by masking the weights of an ordinary MLP: each unit carries a degree, and a connection is kept only when it cannot leak information from \(s_i\) into the conditional for \(s_i\) itself. The result is an exactly normalized distribution with no architectural machinery beyond element-wise masks.
Kept here rather than in qjax.physics because nothing about it is
q-deformed or physical: it is a neural-network building block, framework-
agnostic like the rest of qjax.nn, and the same masked MLP serves any
autoregressive model over binary variables.
References
Germain, M., Gregor, K., Murray, I. & Larochelle, H. (2015). MADE: Masked Autoencoder for Distribution Estimation. ICML. Wu, D., Wang, L. & Zhang, P. (2019). Solving statistical mechanics using variational autoregressive networks. Phys. Rev. Lett. 122, 080602.
made_masks
¶
Binary masks enforcing the autoregressive property, one per weight matrix.
Input \(s_i\) is given degree \(i+1\) and output \(i\) the same. A connection into a hidden unit of degree \(d\) is kept when the incoming degree is \(\le d\); a connection from a hidden unit into output \(i\) is kept when \(d < i+1\). Composing the two, output \(i\) can depend on input \(j\) only if \(j < i\) -- which is exactly the autoregressive condition, and is asserted directly in the tests via a Jacobian.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
num_spins
|
int
|
Number of spins |
required |
hidden
|
Sequence[int]
|
Widths of the hidden layers, at least one. |
required |
Returns:
| Type | Description |
|---|---|
list[Array]
|
A list of |
list[Array]
|
|
Raises:
| Type | Description |
|---|---|
ValueError
|
If |
Source code in qjax/nn/autoregressive.py
made_init
¶
Initialize MADE parameters with Glorot-scaled weights and zero biases.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
key
|
Array
|
PRNG key. |
required |
num_spins
|
int
|
Number of spins |
required |
hidden
|
Sequence[int]
|
Widths of the hidden layers. |
required |
Returns:
| Type | Description |
|---|---|
Params
|
A pytree |
Params
|
match the masks from |
Source code in qjax/nn/autoregressive.py
made_conditionals
¶
Logits of \(p(s_i = +1 \mid s_{<i})\) for every site, in one forward pass.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
params
|
Params
|
Parameters from |
required |
masks
|
Sequence[Array]
|
Masks from |
required |
spins
|
Array
|
Configurations of shape |
required |
Returns:
| Type | Description |
|---|---|
Array
|
Logits of shape |
Source code in qjax/nn/autoregressive.py
made_log_prob
¶
Exact log-probability \(\log p_\theta(s)\) of each configuration.
With \(p(s_i = +1) = \sigma(z_i)\) the two cases collapse into one: \(\log p(s_i) = \log \sigma(s_i z_i) = -\mathrm{softplus}(-s_i z_i)\), which is also the numerically stable form.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
params
|
Params
|
Parameters from |
required |
masks
|
Sequence[Array]
|
Masks from |
required |
spins
|
Array
|
Configurations of shape |
required |
Returns:
| Type | Description |
|---|---|
Array
|
Log-probabilities of shape |
Array
|
configurations these exponentiate to exactly |
Source code in qjax/nn/autoregressive.py
made_sample
¶
Draw exact samples by filling in one spin at a time.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
key
|
Array
|
PRNG key. |
required |
params
|
Params
|
Parameters from |
required |
masks
|
Sequence[Array]
|
Masks from |
required |
num_samples
|
int
|
Number of configurations to draw. |
required |
Returns:
| Type | Description |
|---|---|
Array
|
Configurations of shape |
Array
|
|
Source code in qjax/nn/autoregressive.py
Physics — systems with exact reference values¶
qjax.physics is not re-exported at the top level: import it explicitly
(import qjax.physics as qp) so the flat qjax.* namespace stays the
q-primitives. Its scope rule is deliberate — pure, cheap, exactly-testable
kernels live here; the long runs, controlled comparisons and figures live in
examples/.
reference
¶
Exact and independently-established reference values for the physics examples.
Every number here comes from a closed form, an exhaustive enumeration, or a
published high-precision calculation, and carries its citation. They exist so
that the examples in examples/ can validate themselves rather than merely
illustrate: a curve is only evidence if there is something exact to compare it
against.
References
Onsager, L. (1944). Crystal statistics I. Phys. Rev. 65, 117. Parisi, G. (1980). A sequence of approximated solutions to the SK model. J. Phys. A 13, L115. Wales, D. J. & Doye, J. P. K. (1997). Global optimization by basin-hopping. J. Phys. Chem. A 101, 5111. See also the Cambridge Cluster Database, https://www-wales.ch.cam.ac.uk/CCD.html. Doye, J. P. K., Miller, M. A. & Wales, D. J. (1999). The double-funnel energy landscape of the 38-atom Lennard-Jones cluster. J. Chem. Phys. 110, 6896.
lattice
¶
The 2-D Ising model: Hamiltonian, Metropolis sampler, and exact references.
This module supplies the physical system that the Tsallis machine-learning examples are tested against. It is deliberately small and exact-first: every sampled quantity has at least one independent closed-form or exhaustive counterpart in the same file, so a sampler bug shows up as a disagreement rather than as a plausible-looking curve.
Two Monte Carlo updates are provided: a local checkerboard Metropolis sweep, and the Wolff single-cluster update, which has no critical slowing down and is what makes a finite-size-scaling study at \(T_c\) trustworthy.
The Hamiltonian is the nearest-neighbour Ising model on an \(L \times L\) square lattice with periodic boundaries,
whose critical point is known exactly: \(T_c = 2 J / \ln(1 + \sqrt 2)\) (Onsager, 1944), with \(\nu = 1\), \(\beta = 1/8\) and \(u(T_c) = -\sqrt 2 J\).
Three mutually independent routes to the exact free energy are provided, which is what makes the validation credible:
ising_exact_observables-- exhaustive enumeration of all \(2^{L^2}\) states, for \(L \le 4\).ising_transfer_matrix_log_z-- the \(2^L \times 2^L\) transfer matrix, exact for the finite periodic lattice, for \(L \le 10\).onsager_free_energy_per_site-- Onsager's thermodynamic-limit solution.
References
Onsager, L. (1944). Phys. Rev. 65, 117. Metropolis, N. et al. (1953). J. Chem. Phys. 21, 1087. Wolff, U. (1989). Phys. Rev. Lett. 62, 361.
neighbour_sum
¶
neighbour_sum(spins: Array) -> Array
Sum of the four nearest neighbours of every site, with periodic wrap.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
spins
|
Array
|
Spin configuration(s) of shape |
required |
Returns:
| Type | Description |
|---|---|
Array
|
An array of the same shape whose entry |
Array
|
spins at the four sites adjacent to |
Source code in qjax/physics/lattice.py
ising_energy
¶
ising_energy(spins: Array, coupling: Scalar = 1.0) -> Array
Total Ising energy \(-J \sum_{\langle ij \rangle} s_i s_j\).
Each bond is counted once: the site-wise sum \(\sum_i s_i n_i\) visits every bond twice, hence the factor \(1/2\).
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
spins
|
Array
|
Spin configuration(s) of shape |
required |
coupling
|
Scalar
|
Exchange coupling |
1.0
|
Returns:
| Type | Description |
|---|---|
Array
|
Total energy per configuration, shape |
Array
|
state this is |
Source code in qjax/physics/lattice.py
ising_magnetization
¶
ising_magnetization(spins: Array) -> Array
Signed magnetization per site, mean(s).
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
spins
|
Array
|
Spin configuration(s) of shape |
required |
Returns:
| Type | Description |
|---|---|
Array
|
Magnetization per site, shape |
Source code in qjax/physics/lattice.py
checkerboard_sweep
¶
One Metropolis sweep, updating the two lattice sublattices in turn.
On a bipartite lattice the neighbours of every site lie entirely in the other sublattice, so all sites of one colour can be proposed in parallel without breaking detailed balance. Two half-updates make one full sweep.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
key
|
Array
|
PRNG key. |
required |
spins
|
Array
|
A single configuration of shape |
required |
beta
|
Scalar
|
Inverse temperature |
required |
coupling
|
Scalar
|
Exchange coupling |
1.0
|
Returns:
| Type | Description |
|---|---|
Array
|
The configuration after one sweep, shape |
Source code in qjax/physics/lattice.py
metropolis_chain
¶
metropolis_chain(
key: Array,
spins: Array,
beta: Scalar,
sweeps: int,
coupling: Scalar = 1.0,
) -> Array
Run sweeps Metropolis sweeps and return the final configuration.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
key
|
Array
|
PRNG key. |
required |
spins
|
Array
|
Initial configuration of shape |
required |
beta
|
Scalar
|
Inverse temperature |
required |
sweeps
|
int
|
Number of full sweeps (a Python int; it sets the scan length). |
required |
coupling
|
Scalar
|
Exchange coupling |
1.0
|
Returns:
| Type | Description |
|---|---|
Array
|
The configuration after |
Source code in qjax/physics/lattice.py
wolff_update
¶
One Wolff single-cluster update.
A seed site is chosen uniformly, a cluster is grown outward through bonds between equal spins with probability \(p = 1 - e^{-2\beta J}\) each, and the whole cluster is flipped. The move is accepted unconditionally: \(p\) is exactly the value that makes the cluster-construction and Boltzmann factors cancel.
Why it is here rather than only checkerboard_sweep: a local update's
autocorrelation time grows as \(L^{z}\) with \(z \approx 2.17\) at \(T_c\),
while a cluster update has no such critical slowing down. Note that this
buys decorrelation, not a shortcut to equilibrium from a cold start: one
update flips one cluster, so a chain started from a random configuration
still needs enough updates for the cluster to have swept the lattice --
measured at \(T_c\), L = 32 is nowhere near equilibrium after 40 updates
and settled by about 120. Local sweeps relax a random start more evenly;
cluster updates decorrelate an equilibrated one far better.
The cluster is grown as a boolean mask in a jax.lax.while_loop: every
iteration draws one fresh uniform per bond, so two frontier sites adjacent
to the same candidate test their bonds independently, as the algorithm
requires. The loop is data-dependent, so this is jittable but not
reverse-differentiable -- which a Monte Carlo update never needs to be.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
key
|
Array
|
PRNG key. |
required |
spins
|
Array
|
A single configuration of shape |
required |
beta
|
Scalar
|
Inverse temperature |
required |
coupling
|
Scalar
|
Exchange coupling |
1.0
|
Returns:
| Type | Description |
|---|---|
Array
|
The configuration after one cluster flip, shape |
Source code in qjax/physics/lattice.py
wolff_chain
¶
wolff_chain(
key: Array,
spins: Array,
beta: Scalar,
updates: int,
coupling: Scalar = 1.0,
) -> Array
Run updates Wolff cluster updates and return the final configuration.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
key
|
Array
|
PRNG key. |
required |
spins
|
Array
|
Initial configuration of shape |
required |
beta
|
Scalar
|
Inverse temperature |
required |
updates
|
int
|
Number of cluster updates (a Python int; it sets the scan length). |
required |
coupling
|
Scalar
|
Exchange coupling |
1.0
|
Returns:
| Type | Description |
|---|---|
Array
|
The configuration after |
Source code in qjax/physics/lattice.py
sample_ising
¶
sample_ising(
key: Array,
size: int,
temperatures: Array,
num_samples: int,
sweeps: int,
coupling: Scalar = 1.0,
algorithm: str = "metropolis",
) -> Array
Draw equilibrium configurations at each of several temperatures.
Every sample gets its own independent chain, started from a random configuration. Critical slowing down therefore affects only how long each chain must run to equilibrate, never the independence of the samples -- so no decorrelation sweeps between samples are needed.
How long "long enough" is depends on the update, and the two available here fail in opposite directions: Metropolis sweeps relax a random start evenly but decorrelate slowly at \(T_c\) (\(\tau \sim L^{2.17}\)), while Wolff cluster updates decorrelate without critical slowing down but need enough updates to have touched the whole lattice first. Both are checked against exhaustive enumeration in the test suite.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
key
|
Array
|
PRNG key. |
required |
size
|
int
|
Linear lattice size |
required |
temperatures
|
Array
|
Temperatures to sample at, shape |
required |
num_samples
|
int
|
Independent configurations per temperature. |
required |
sweeps
|
int
|
Equilibration steps per chain (a Python int) -- Metropolis
sweeps, or Wolff cluster updates when |
required |
coupling
|
Scalar
|
Exchange coupling |
1.0
|
algorithm
|
str
|
|
'metropolis'
|
Returns:
| Type | Description |
|---|---|
Array
|
Configurations of shape |
Array
|
|
Raises:
| Type | Description |
|---|---|
ValueError
|
If |
Source code in qjax/physics/lattice.py
ising_all_configurations
¶
Enumerate every spin configuration of an L x L lattice.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
size
|
int
|
Linear lattice size |
required |
Returns:
| Type | Description |
|---|---|
Array
|
All configurations, shape |
Array
|
|
Raises:
| Type | Description |
|---|---|
ValueError
|
If |
Source code in qjax/physics/lattice.py
ising_boltzmann_probabilities
¶
Exact Boltzmann weights over the full state space, in enumeration order.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
size
|
int
|
Linear lattice size |
required |
temperature
|
Scalar
|
Temperature |
required |
coupling
|
Scalar
|
Exchange coupling |
1.0
|
Returns:
| Type | Description |
|---|---|
Array
|
Normalized probabilities of shape |
Array
|
|
Source code in qjax/physics/lattice.py
ising_exact_observables
¶
ising_exact_observables(
size: int, temperature: Scalar, coupling: Scalar = 1.0
) -> dict[str, Array]
Exact thermodynamics by exhaustive enumeration of the state space.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
size
|
int
|
Linear lattice size |
required |
temperature
|
Scalar
|
Temperature |
required |
coupling
|
Scalar
|
Exchange coupling |
1.0
|
Returns:
| Type | Description |
|---|---|
dict[str, Array]
|
A dict with |
dict[str, Array]
|
|
dict[str, Array]
|
|
Source code in qjax/physics/lattice.py
ising_transfer_matrix_log_z
¶
Exact \(\log Z\) of the finite periodic \(L \times L\) lattice.
Builds the \(2^L \times 2^L\) column-to-column transfer matrix and
evaluates \(Z = \operatorname{Tr} T^L\) from its eigenvalues. The matrix
entries reach \(e^{2 \beta J L}\), so the largest element is factored out
before exponentiating and the trace is accumulated relative to
\(\lambda_{\max}\) -- without that, float64 overflows already at
\(L = 10\) and low temperature.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
size
|
int
|
Linear lattice size |
required |
temperature
|
Scalar
|
Temperature |
required |
coupling
|
Scalar
|
Exchange coupling |
1.0
|
Returns:
| Type | Description |
|---|---|
Array
|
A 0-d array holding |
Raises:
| Type | Description |
|---|---|
ValueError
|
If |
Source code in qjax/physics/lattice.py
onsager_free_energy_per_site
¶
onsager_free_energy_per_site(
temperature: Array,
coupling: Scalar = 1.0,
num_quad: int = 4096,
) -> Array
Onsager's exact free energy per site in the thermodynamic limit.
The integrand is bounded everywhere, including at \(T_c\) where
\(\kappa = 1\) (the singularity is in the second derivative), so a plain
midpoint rule converges. cosh and sech are evaluated in log space so
the expression stays finite down to very low temperature.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
temperature
|
Array
|
Temperature(s) |
required |
coupling
|
Scalar
|
Exchange coupling |
1.0
|
num_quad
|
int
|
Midpoint quadrature nodes on |
4096
|
Returns:
| Type | Description |
|---|---|
Array
|
Free energy per site, same shape as |
Array
|
|
Source code in qjax/physics/lattice.py
onsager_magnetization
¶
onsager_magnetization(
temperature: Array, coupling: Scalar = 1.0
) -> Array
Onsager's exact spontaneous magnetization \(m = (1 - \sinh^{-4} 2\beta J)^{1/8}\).
Zero for \(T \ge T_c\) and rising to 1 as \(T \to 0\), with the exact
critical exponent \(\beta = 1/8\).
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
temperature
|
Array
|
Temperature(s) |
required |
coupling
|
Scalar
|
Exchange coupling |
1.0
|
Returns:
| Type | Description |
|---|---|
Array
|
Spontaneous magnetization in |
Source code in qjax/physics/lattice.py
onsager_energy_per_site
¶
onsager_energy_per_site(
temperature: Array,
coupling: Scalar = 1.0,
num_quad: int = 4096,
) -> Array
Onsager's exact internal energy per site, as \(\partial (\beta f)/\partial \beta\).
Taken by automatic differentiation of onsager_free_energy_per_site
rather than from the closed form, which involves a complete elliptic
integral \(K(\kappa)\) that diverges logarithmically at \(T_c\) and so is
hard to quadrature there. \(\beta f\) is smooth at \(T_c\), so its derivative
is well conditioned.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
temperature
|
Array
|
Temperature(s) |
required |
coupling
|
Scalar
|
Exchange coupling |
1.0
|
num_quad
|
int
|
Quadrature nodes passed through to the free energy. |
4096
|
Returns:
| Type | Description |
|---|---|
Array
|
Internal energy per site, same shape as |
Array
|
|
Source code in qjax/physics/lattice.py
observables
¶
Estimators for locating and characterizing a phase transition.
These are the finite-size-scaling tools that turn a family of curves measured at several lattice sizes into a number that can be compared against an exact critical temperature or exponent. They are deliberately model-agnostic: the input is a curve over a temperature grid, whether it came from a Monte Carlo observable or from a neural network's output.
All of them are jittable and take static shapes; the index searches are written
with masks rather than Python control flow so they work under jax.jit.
binder_cumulant
¶
Binder fourth-order cumulant \(U_4 = 1 - \langle m^4\rangle / (3\langle m^2\rangle^2)\).
\(U_4\) is dimensionless at the critical point, so curves measured at different lattice sizes cross there -- the standard way to locate \(T_c\) without knowing it. It equals \(2/3\) deep in the ordered phase (where \(m\) is a two-delta distribution at \(\pm m_0\)) and \(0\) in the disordered phase (where \(m\) is Gaussian).
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
magnetization
|
Array
|
Per-configuration magnetizations. |
required |
axis
|
int
|
Axis to average over. |
-1
|
Returns:
| Type | Description |
|---|---|
Array
|
The cumulant, with |
Source code in qjax/physics/observables.py
crossing_temperature
¶
crossing_temperature(
temperatures: Array, curve: Array, level: Scalar = 0.5
) -> Array
Temperature of the first crossing of level, by linear interpolation.
Used to read a transition temperature off a monotone indicator, such as the probability a classifier assigns to the ordered phase.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
temperatures
|
Array
|
Strictly ordered temperature grid, shape |
required |
curve
|
Array
|
Values on that grid, shape |
required |
level
|
Scalar
|
The level to cross. |
0.5
|
Returns:
| Type | Description |
|---|---|
Array
|
A 0-d array with the crossing temperature, or |
Array
|
crosses |
Source code in qjax/physics/observables.py
peak_temperature
¶
peak_temperature(
temperatures: Array, curve: Array
) -> Array
Temperature of a curve's maximum, refined by a three-point parabola.
Reads off a pseudo-critical temperature from a peaked indicator (a susceptibility, a heat capacity, or the Tsallis entropy of a classifier's output). The parabolic vertex through the grid maximum and its two neighbours recovers sub-grid resolution and handles a non-uniform grid.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
temperatures
|
Array
|
Ordered temperature grid, shape |
required |
curve
|
Array
|
Values on that grid, shape |
required |
Returns:
| Type | Description |
|---|---|
Array
|
A 0-d array with the peak temperature. |
Source code in qjax/physics/observables.py
half_width
¶
half_width(temperatures: Array, curve: Array) -> Array
Full width at half maximum of a peaked curve.
The half level is taken relative to the curve's own minimum, \(\tfrac12(\max + \min)\), so a peak sitting on a non-zero background is measured correctly. For a critical indicator the width scales as \(w(L) \sim L^{-1/\nu}\), which is how the examples recover \(\nu\).
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
temperatures
|
Array
|
Ordered temperature grid, shape |
required |
curve
|
Array
|
Values on that grid, shape |
required |
Returns:
| Type | Description |
|---|---|
Array
|
A 0-d array with the width, or |
Array
|
below the half level on both sides of its maximum. |
Source code in qjax/physics/observables.py
finite_size_extrapolation
¶
finite_size_extrapolation(
sizes: Array, estimates: Array, nu: Scalar = 1.0
) -> tuple[Array, Array, Array]
Extrapolate a size-dependent estimate to the thermodynamic limit.
Fits \(T_c(L) = T_c(\infty) + a L^{-1/\nu}\) by ordinary least squares in \(x = L^{-1/\nu}\) -- the leading finite-size correction for a pseudo-critical temperature. With the exact \(\nu\) supplied, the intercept is the quantity to compare against the exact \(T_c\).
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
sizes
|
Array
|
Linear lattice sizes |
required |
estimates
|
Array
|
The size-dependent estimates, shape |
required |
nu
|
Scalar
|
Correlation-length exponent used to build the abscissa. |
1.0
|
Returns:
| Type | Description |
|---|---|
Array
|
|
Array
|
usual OLS one and is |
Source code in qjax/physics/observables.py
spinglass
¶
The Sherrington-Kirkpatrick spin glass: Hamiltonian and exact thermodynamics.
The SK model is the mean-field spin glass,
with a rugged, frustrated landscape whose ground-state energy per spin tends to Parisi's \(-0.7633\) as \(N \to \infty\). It is the standard hard case for variational methods: a distribution over \(2^N\) states that concentrates on a few of them scores well on the naive objective while getting the physics wrong, which is why the free energy here is always reported against an exact value.
For \(N \le 22\) the full state space is enumerable, so the free energy,
internal energy and two-point correlations are available exactly. The
enumeration is streamed in chunks with a running logsumexp, so the memory
cost is set by the chunk size rather than by \(2^N\).
References
Sherrington, D. & Kirkpatrick, S. (1975). Phys. Rev. Lett. 35, 1792. Parisi, G. (1980). J. Phys. A 13, L115.
sk_couplings
¶
Draw a symmetric SK coupling matrix with zero diagonal.
Off-diagonal entries are Gaussian with variance 1 / num_spins, the
scaling that makes the energy per spin intensive.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
key
|
Array
|
PRNG key. |
required |
num_spins
|
int
|
Number of spins |
required |
Returns:
| Type | Description |
|---|---|
Array
|
A symmetric |
Source code in qjax/physics/spinglass.py
sk_energy
¶
sk_energy(spins: Array, couplings: Array) -> Array
SK energy \(-\tfrac12 s^{\mathsf T} J s\).
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
spins
|
Array
|
Configuration(s) of shape |
required |
couplings
|
Array
|
Symmetric |
required |
Returns:
| Type | Description |
|---|---|
Array
|
Total energy per configuration, shape |
Source code in qjax/physics/spinglass.py
sk_exact_observables
¶
sk_exact_observables(
couplings: Array, temperature: Scalar, chunk: int = 4096
) -> dict[str, Array]
Exact SK thermodynamics by streamed exhaustive enumeration.
Accumulates \(Z\), \(\langle E \rangle\) and \(\langle s_i s_j \rangle\)
over all \(2^N\) configurations in chunks, rescaling the running sums
whenever a new maximum log-weight appears. This is a numerically exact
logsumexp at O(chunk) memory.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
couplings
|
Array
|
Symmetric |
required |
temperature
|
Scalar
|
Temperature |
required |
chunk
|
int
|
Configurations per chunk; rounded down to a power of two and
capped at |
4096
|
Returns:
| Type | Description |
|---|---|
dict[str, Array]
|
A dict with |
dict[str, Array]
|
|
dict[str, Array]
|
array of \(\langle s_i s_j \rangle\)). |
Raises:
| Type | Description |
|---|---|
ValueError
|
If |
Source code in qjax/physics/spinglass.py
sk_exact_correlations
¶
Exact two-point correlations \(\langle s_i s_j \rangle\) by enumeration.
A thin wrapper around sk_exact_observables. The correlation matrix is
the sharpest diagnostic of variational mode collapse: a distribution that
has collapsed onto one configuration reports +/-1 everywhere, however
good its free energy looks.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
couplings
|
Array
|
Symmetric |
required |
temperature
|
Scalar
|
Temperature |
required |
chunk
|
int
|
Configurations per enumeration chunk. |
4096
|
Returns:
| Type | Description |
|---|---|
Array
|
An |
Source code in qjax/physics/spinglass.py
clusters
¶
Lennard-Jones atomic clusters: potential, local quenching, and geometry.
The pair potential
is the standard benchmark landscape for global optimization: an \(n\)-atom cluster has a number of local minima growing roughly exponentially in \(n\) (around \(10^8\) already at \(n = 20\)), so the number of basins, not the dimension, is what makes it hard. Two properties make it a verifiable benchmark rather than a demo:
- For \(n \le 4\) the global minimum is a closed form. A regular simplex with every edge at \(r = 2^{1/6}\sigma\) puts every pair exactly at the potential minimum, contributing \(-\epsilon\) each: \(-1\), \(-3\), \(-6\).
- For larger \(n\) the global minima are tabulated to six decimals in the
Cambridge Cluster Database (see
qjax.physics.reference), including the double-funnel case \(n = 38\).
References
Wales, D. J. & Doye, J. P. K. (1997). J. Phys. Chem. A 101, 5111. Doye, J. P. K., Miller, M. A. & Wales, D. J. (1999). J. Chem. Phys. 110, 6896.
lj_energy
¶
lj_energy(
positions: Array,
epsilon: Scalar = 1.0,
sigma: Scalar = 1.0,
softening: float = 1e-12,
) -> Array
Total Lennard-Jones energy of a cluster, summed over unordered pairs.
The \(r^{-12}\) term diverges as two atoms coincide, so the squared
separations are floored at softening and the diagonal is replaced
before the power rather than after: a bare 1 / 0 would back-propagate
NaN into every coordinate even though the diagonal is masked out of the
sum.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
positions
|
Array
|
Atomic coordinates of shape |
required |
epsilon
|
Scalar
|
Well depth. |
1.0
|
sigma
|
Scalar
|
Length scale; the pair minimum sits at |
1.0
|
softening
|
float
|
Floor on the squared separation. |
1e-12
|
Returns:
| Type | Description |
|---|---|
Array
|
Total energy per cluster, shape |
Array
|
for a regular simplex of 2, 3, 4 atoms at |
Source code in qjax/physics/clusters.py
lj_energy_confined
¶
lj_energy_confined(
positions: Array,
container_radius: Scalar,
stiffness: Scalar = 10.0,
epsilon: Scalar = 1.0,
sigma: Scalar = 1.0,
) -> Array
Lennard-Jones energy plus a soft spherical wall.
A cluster in free space evaporates: an atom kicked far enough away feels no restoring force, and its energy contribution goes to zero rather than to something unphysical, so a global search will happily "solve" the problem by losing atoms. The wall \(k \sum_i [\,|x_i| - R\,]_+^2\) is zero inside the container, so a minimum found strictly inside is a minimum of the bare potential -- which the examples assert rather than assume.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
positions
|
Array
|
Atomic coordinates of shape |
required |
container_radius
|
Scalar
|
Radius |
required |
stiffness
|
Scalar
|
Wall stiffness |
10.0
|
epsilon
|
Scalar
|
Well depth. |
1.0
|
sigma
|
Scalar
|
Length scale. |
1.0
|
Returns:
| Type | Description |
|---|---|
Array
|
Confined energy per cluster, shape |
Array
|
including at the origin. |
Source code in qjax/physics/clusters.py
lj_random_cluster
¶
Draw atomic positions uniformly inside a ball.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
key
|
Array
|
PRNG key. |
required |
num_atoms
|
int
|
Number of atoms |
required |
radius
|
Scalar
|
Ball radius. |
required |
Returns:
| Type | Description |
|---|---|
Array
|
Positions of shape |
Source code in qjax/physics/clusters.py
equidistant_cluster
¶
A regular simplex of 1 to 4 atoms with every pair at distance.
These are the closed-form Lennard-Jones global minima: at
distance = 2**(1/6) sigma every pair sits exactly at the potential
minimum, so the energy is -epsilon times the number of pairs.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
num_atoms
|
int
|
Number of atoms, 1 to 4. |
required |
distance
|
Scalar
|
Edge length. |
required |
Returns:
| Type | Description |
|---|---|
Array
|
Positions of shape |
Raises:
| Type | Description |
|---|---|
ValueError
|
If |
Source code in qjax/physics/clusters.py
lj_quench
¶
lj_quench(
positions: Array,
steps: int = 40,
learning_rate: Scalar = 0.002,
epsilon: Scalar = 1.0,
sigma: Scalar = 1.0,
) -> tuple[Array, Array]
Locally minimize the bare Lennard-Jones energy with Adam.
This is the "quench" of Wales-Doye basin-hopping: it maps a proposed configuration onto the bottom of the basin it fell into, so a Monte Carlo walk explores the graph of local minima rather than the raw landscape. Adam rather than plain gradient descent because the \(r^{-12}\) core makes the gradient scale vary by orders of magnitude across the cluster; the price is that the energy trace is not guaranteed monotone.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
positions
|
Array
|
Starting coordinates of shape |
required |
steps
|
int
|
Number of Adam steps. |
40
|
learning_rate
|
Scalar
|
Adam step size. |
0.002
|
epsilon
|
Scalar
|
Well depth. |
1.0
|
sigma
|
Scalar
|
Length scale. |
1.0
|
Returns:
| Type | Description |
|---|---|
Array
|
|
Array
|
|
tuple[Array, Array]
|
last. |
Source code in qjax/physics/clusters.py
coordination_numbers
¶
coordination_numbers(
positions: Array, cutoff: Scalar = 1.35
) -> Array
Count neighbours within cutoff of each atom.
Distinguishes surface atoms from core atoms, which is what makes an fcc truncated octahedron visually distinguishable from an icosahedron.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
positions
|
Array
|
Atomic coordinates of shape |
required |
cutoff
|
Scalar
|
Neighbour cutoff, in the same units as |
1.35
|
Returns:
| Type | Description |
|---|---|
Array
|
Integer neighbour counts of shape |
Source code in qjax/physics/clusters.py
annealing
¶
Tsallis-Stariolo generalized-simulated-annealing temperature schedules.
Generalized simulated annealing (Tsallis & Stariolo, Physica A 233, 395, 1996) replaces the two Boltzmann ingredients of classical annealing -- the Gaussian proposal and the exponential acceptance -- by their \(q\)-deformed counterparts, and cools with the matching \(q\)-deformed schedule
At \(q = 2\) this is the Cauchy machine of Szu & Hartley; at \(q = 1\) it must
reduce to the Geman-Geman logarithmic schedule \(T_1(t) = T_1(1)\ln 2 /
\ln(1+t)\) -- but the expression above is \(0/0\) there, exactly the pathology
that qjax.shared.series exists to defeat.
The fix needs no new code at all. Since
both \((q-1)\) factors cancel and the whole schedule is a ratio of two
qjax.q_log calls:
Written this way the classical limit falls out of the same expression with no
branch on \(q\), and -- because qjax.q_log carries the limit through the
entire function \((e^t-1)/t\) rather than switching to a \(q\)-independent
formula -- the derivative with respect to \(q\) stays correct and non-zero at
\(q = 1\). A learnable cooling index is therefore just another parameter.
References
Tsallis, C. & Stariolo, D. A. (1996). Physica A 233, 395. Szu, H. & Hartley, R. (1987). Phys. Lett. A 122, 157. Geman, S. & Geman, D. (1984). IEEE TPAMI 6, 721.
tsallis_schedule
¶
tsallis_schedule(
step: Array, initial: Scalar, q: Scalar
) -> Array
The Tsallis cooling law \(T_q(1)\,\ln_{2-q} 2 / \ln_{2-q}(1+t)\).
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
step
|
Array
|
Annealing step |
required |
initial
|
Scalar
|
Temperature at |
required |
q
|
Scalar
|
Cooling index. |
required |
Returns:
| Type | Description |
|---|---|
Array
|
The temperature at each |
Source code in qjax/physics/annealing.py
visiting_temperature
¶
visiting_temperature(
step: Array, initial: Scalar, q_visit: Scalar
) -> Array
Visiting temperature \(T_V(t)\) controlling the proposal step length.
Sets the width of the \(q\)-Gaussian from which trial moves are drawn. With
q_visit > 1 the proposal has power-law tails, so the walk mixes
occasional long Levy-like flights with local moves and escapes a metastable
basin without waiting for a thermally activated crossing.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
step
|
Array
|
Annealing step |
required |
initial
|
Scalar
|
Visiting temperature at |
required |
q_visit
|
Scalar
|
Visiting index |
required |
Returns:
| Type | Description |
|---|---|
Array
|
The visiting temperature at each |
Source code in qjax/physics/annealing.py
acceptance_temperature
¶
acceptance_temperature(
step: Array, initial: Scalar, q_accept: Scalar
) -> Array
Acceptance temperature \(T_A(t)\) entering \(\exp_{q_A}(-\Delta E / T_A)\).
Scheduled by the same Tsallis law as the visiting temperature but with its own index, so the two deformations -- how far the walk proposes and how readily it accepts an uphill move -- can be varied independently.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
step
|
Array
|
Annealing step |
required |
initial
|
Scalar
|
Acceptance temperature at |
required |
q_accept
|
Scalar
|
Cooling index for the acceptance temperature. |
required |
Returns:
| Type | Description |
|---|---|
Array
|
The acceptance temperature at each |
Source code in qjax/physics/annealing.py
diffusion
¶
Anomalous diffusion: scaling relations and estimators for the entropic index.
Nonextensive statistics makes a falsifiable prediction about anomalous diffusion. The nonlinear (porous-medium) Fokker-Planck equation
has the self-similar \(q\)-Gaussian solution of Tsallis & Bukman (1996), whose width obeys \(\langle x^2 \rangle \propto t^{\alpha}\) with
So the entropic index and the anomalous diffusion exponent are not two independent fitted parameters: measuring the shape of the distribution predicts the growth of its width, and vice versa. That is what makes \(q\) a measured physical quantity here rather than a hyperparameter -- and it is checkable, in this module, by two independent estimators.
A second, experimentally realized mechanism yields the same distribution from a linear-noise Langevin equation with saturating (Sisyphus) friction,
whose exact stationary solution is \(P(p) \propto [1 + p^2/p_c^2]^{-\alpha p_c^2 / (2 D_0)}\), a \(q\)-Gaussian with
This is the cold-atom case: Lutz (2003) evaluated the three coefficients semiclassically for atoms in a dissipative optical lattice and obtained \(q = 1 + 44 E_R / U_0\), confirmed experimentally by Douglas, Bergamini & Renzoni (2006).
References
Plastino, A. R. & Plastino, A. (1995). Physica A 222, 347. Tsallis, C. & Bukman, D. J. (1996). Phys. Rev. E 54, R2197. Lutz, E. (2003). Phys. Rev. A 67, 051402(R). Douglas, P., Bergamini, S. & Renzoni, F. (2006). Phys. Rev. Lett. 96, 110601.
nlfp_exponent
¶
nlfp_exponent(q: Scalar) -> Array
Anomalous diffusion exponent \(\alpha = 2/(3-q)\) of the nonlinear FP equation.
q = 1 gives normal diffusion (alpha = 1), q > 1 superdiffusion
and q < 1 subdiffusion.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
q
|
Scalar
|
Entropic index, |
required |
Returns:
| Type | Description |
|---|---|
Array
|
The exponent |
Source code in qjax/physics/diffusion.py
nlfp_index
¶
nlfp_index(exponent: Scalar) -> Array
Invert nlfp_exponent: \(q = 3 - 2/\alpha\).
The second, independent route to the entropic index. It is the only route
for a subdiffusive (\(q < 1\)) process, because the \(q\)-Gaussian then has
compact support: the log-likelihood is -inf outside it, so a
gradient-based fit initialized at q > 1 can never cross into q < 1.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
exponent
|
Scalar
|
Measured anomalous exponent |
required |
Returns:
| Type | Description |
|---|---|
Array
|
The entropic index implied by |
Source code in qjax/physics/diffusion.py
nlfp_scaling_beta
¶
nlfp_scaling_beta(
time: Array,
q: Scalar,
beta_initial: Scalar,
reference_time: Scalar = 1.0,
) -> Array
Self-similar width parameter \(\beta(t)\) of the Tsallis-Bukman solution.
The solution keeps its \(q\)-Gaussian shape for all time and only
rescales, with \(\beta(t) \propto t^{-\alpha}\) and the same
\(\alpha = 2/(3-q)\) that governs the mean-squared displacement. The
diffusivity \(D\) enters only through the pair
(beta_initial, reference_time), which the initial condition fixes, so it
is not an argument here.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
time
|
Array
|
Time(s) at which to evaluate the width, any shape. |
required |
q
|
Scalar
|
Entropic index. |
required |
beta_initial
|
Scalar
|
Width parameter at |
required |
reference_time
|
Scalar
|
Time at which |
1.0
|
Returns:
| Type | Description |
|---|---|
Array
|
|
Source code in qjax/physics/diffusion.py
nlfp_rate
¶
nlfp_rate(q: Scalar, diffusivity: Scalar) -> Array
Rate constant \(K\) of the width ODE \(\dot\beta = -K\,\beta^{(5-q)/2}\).
Substituting the normalized \(q\)-Gaussian \(p = (\sqrt\beta / C_q)\exp_q(-\beta x^2)\) into \(\partial_t p = D\,\partial_{xx} p^{\,2-q}\) makes every \(x\)-dependent factor cancel identically, leaving a scalar ordinary differential equation for the width alone with
where \(C_q\) is qjax.normalization. At \(q = 1\) this is \(K = 4D\), and
nlfp_width then reduces to the heat kernel exactly -- which is the check
that pins the constant.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
q
|
Scalar
|
Entropic index, strictly below |
required |
diffusivity
|
Scalar
|
The coefficient |
required |
Returns:
| Type | Description |
|---|---|
Array
|
The scalar rate constant |
Raises:
| Type | Description |
|---|---|
ValueError
|
If |
Source code in qjax/physics/diffusion.py
nlfp_offset
¶
nlfp_offset(
q: Scalar, diffusivity: Scalar, beta_initial: Scalar
) -> Array
Time offset \(t_\star\) placing width \(\beta_0\) at \(t = 0\).
The self-similar solution is a power law in \(t + t_\star\); the offset is what
replaces the singular point-source initial condition by a \(q\)-Gaussian of
finite width, so that t = 0 is an ordinary regular point.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
q
|
Scalar
|
Entropic index, strictly below |
required |
diffusivity
|
Scalar
|
The coefficient |
required |
beta_initial
|
Scalar
|
Width parameter at |
required |
Returns:
| Type | Description |
|---|---|
Array
|
The scalar offset |
Source code in qjax/physics/diffusion.py
nlfp_width
¶
nlfp_width(
time: Array,
q: Scalar,
diffusivity: Scalar,
beta_initial: Scalar,
initial_time: Scalar = 0.0,
) -> Array
Exact width parameter \(\beta(t)\) of the Tsallis-Bukman solution.
Solving \(\dot\beta = -K\beta^{(5-q)/2}\) from nlfp_rate gives
with \(t_\star\) from nlfp_offset chosen so that
\(\beta(t_0) = \beta_0\).
This is the full solution, diffusivity included. nlfp_scaling_beta is the
weaker statement -- the same power law with the prefactor left to the initial
condition -- and is what to use when D is unknown.
Two properties worth knowing, both pinned by the test suite: at \(q = 1\)
this is exactly the heat kernel width \(1/(4D(t + t_\star))\), and for every
\(q\) it decays as \(t^{-\alpha}\) with the same
\(\alpha = 2/(3-q)\) that nlfp_exponent returns.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
time
|
Array
|
Time(s) at which to evaluate the width, any shape. |
required |
q
|
Scalar
|
Entropic index, strictly below |
required |
diffusivity
|
Scalar
|
The coefficient |
required |
beta_initial
|
Scalar
|
Width parameter at |
required |
initial_time
|
Scalar
|
Time at which |
0.0
|
Returns:
| Type | Description |
|---|---|
Array
|
|
Source code in qjax/physics/diffusion.py
nlfp_density
¶
nlfp_density(
x: Array,
time: Array,
q: Scalar,
diffusivity: Scalar,
beta_initial: Scalar,
initial_time: Scalar = 0.0,
) -> Array
The exact solution of the nonlinear Fokker-Planck equation.
A normalized \(q\)-Gaussian whose width follows nlfp_width:
The shape never changes -- only the width -- which is what "self-similar" means here, and is why one scalar ODE captures the whole solution of a nonlinear partial differential equation.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
x
|
Array
|
Position(s), broadcast against |
required |
time
|
Array
|
Time(s), broadcast against |
required |
q
|
Scalar
|
Entropic index, strictly below |
required |
diffusivity
|
Scalar
|
The coefficient |
required |
beta_initial
|
Scalar
|
Width parameter at |
required |
initial_time
|
Scalar
|
Time at which |
0.0
|
Returns:
| Type | Description |
|---|---|
Array
|
The density, broadcast over |
Array
|
|
Source code in qjax/physics/diffusion.py
nlfp_front
¶
nlfp_front(
time: Array,
q: Scalar,
diffusivity: Scalar,
beta_initial: Scalar,
initial_time: Scalar = 0.0,
) -> Array
Edge of the support, \(x_f(t) = 1/\sqrt{(1-q)\beta(t)}\), for \(q < 1\).
Below \(q = 1\) the \(q\)-Gaussian is compactly supported, so the solution has a genuine moving free boundary: the density is exactly zero beyond \(x_f\), not merely small. That front is the sharpest thing to measure a numerical solution against, and it is the feature a strictly positive parameterization cannot represent at all.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
time
|
Array
|
Time(s) at which to locate the front, any shape. |
required |
q
|
Scalar
|
Entropic index, strictly below |
required |
diffusivity
|
Scalar
|
The coefficient |
required |
beta_initial
|
Scalar
|
Width parameter at |
required |
initial_time
|
Scalar
|
Time at which |
0.0
|
Returns:
| Type | Description |
|---|---|
Array
|
The front position, same shape as |
Array
|
where the support is the whole line. |
Source code in qjax/physics/diffusion.py
nlfp_residual
¶
nlfp_residual(
density_fn: Callable[[Array, Array], Array],
x: Scalar,
time: Scalar,
q: Scalar,
diffusivity: Scalar,
) -> Array
Residual \(\partial_t p - D\,\partial_{xx} p^{\,2-q}\) of a candidate solution.
Takes a callable mapping a scalar (x, t) to a scalar density and
differentiates it, so one operator serves two purposes: it validates the
exact solution (its residual must vanish, which is what gates the derivation
of nlfp_rate) and it trains a neural network (the residual is the loss).
Vectorize over collocation points with jax.vmap.
The second derivative is taken of the pressure variable \(v = p^{\,2-q}\) directly, matching the equation as written. Near a \(q < 1\) front, \(p \sim (x_f - x)^{1/(1-q)}\) gives \(v \sim (x_f - x)^{(2-q)/(1-q)}\) -- cubic at \(q = 1/2\) -- so \(v''\) is continuous and vanishing there rather than singular.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
density_fn
|
Callable[[Array, Array], Array]
|
A callable |
required |
x
|
Scalar
|
Position at which to evaluate the residual. |
required |
time
|
Scalar
|
Time at which to evaluate the residual. |
required |
q
|
Scalar
|
Entropic index, strictly below |
required |
diffusivity
|
Scalar
|
The coefficient |
required |
Returns:
| Type | Description |
|---|---|
Array
|
The scalar residual. Zero for an exact solution. |
Raises:
| Type | Description |
|---|---|
ValueError
|
If |
Source code in qjax/physics/diffusion.py
saturating_langevin_q
¶
saturating_langevin_q(
diffusion: Scalar,
friction: Scalar,
momentum_scale: Scalar,
) -> Array
Entropic index \(q = 1 + 2 D_0 / (\alpha p_c^2)\) of the Sisyphus Langevin process.
Exact for the stationary state of dp = -alpha p / (1 + (p/p_c)**2) dt +
sqrt(2 D_0) dW. Because the three coefficients are chosen by the caller,
this is a rigorous internal reference for a fit: the true q is known.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
diffusion
|
Scalar
|
Momentum diffusion coefficient |
required |
friction
|
Scalar
|
Friction coefficient |
required |
momentum_scale
|
Scalar
|
Saturation momentum |
required |
Returns:
| Type | Description |
|---|---|
Array
|
The exact stationary entropic index. |
Source code in qjax/physics/diffusion.py
saturating_langevin_beta
¶
saturating_langevin_beta(
diffusion: Scalar, friction: Scalar
) -> Array
Width parameter \(\beta = \alpha / (2 D_0)\) of the Sisyphus stationary state.
The companion of saturating_langevin_q: together they specify the exact
stationary \(q\)-Gaussian, so both fitted parameters have a known target.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
diffusion
|
Scalar
|
Momentum diffusion coefficient |
required |
friction
|
Scalar
|
Friction coefficient |
required |
Returns:
| Type | Description |
|---|---|
Array
|
The exact stationary width parameter. |
Source code in qjax/physics/diffusion.py
lutz_q
¶
lutz_q(recoil_over_depth: Array) -> Array
Lutz's cold-atom prediction \(q = 1 + 44 E_R / U_0\).
Obtained by evaluating the three Sisyphus coefficients semiclassically for
atoms in a dissipative optical lattice of depth \(U_0\) and recoil energy
\(E_R\); confirmed experimentally by Douglas, Bergamini & Renzoni (2006).
Unlike saturating_langevin_q, this is a prediction about a real
experiment, not about a simulation.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
recoil_over_depth
|
Array
|
The ratio |
required |
Returns:
| Type | Description |
|---|---|
Array
|
The predicted entropic index, same shape as the input. |
Source code in qjax/physics/diffusion.py
mean_squared_displacement
¶
mean_squared_displacement(
snapshots: Array, origin: Array | None = None
) -> Array
Ensemble mean-squared displacement from a set of snapshots.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
snapshots
|
Array
|
Positions of shape |
required |
origin
|
Array | None
|
Starting positions, shape matching one snapshot. Defaults to
|
None
|
Returns:
| Type | Description |
|---|---|
Array
|
The mean-squared displacement at each snapshot time, shape |
Source code in qjax/physics/diffusion.py
fit_power_law
¶
fit_power_law(
x: Array,
y: Array,
low: int = 0,
high: int | None = None,
) -> tuple[Array, Array, Array]
Least-squares fit of \(y = c\,x^{a}\) on a log-log scale.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
x
|
Array
|
Abscissa, strictly positive, shape |
required |
y
|
Array
|
Ordinate, strictly positive, shape |
required |
low
|
int
|
First index to include; use it to drop an early transient. |
0
|
high
|
int | None
|
One past the last index to include. Defaults to |
None
|
Returns:
| Type | Description |
|---|---|
Array
|
|
Array
|
usual OLS one and is |
Source code in qjax/physics/diffusion.py
histogram_density
¶
histogram_density(samples: Array, edges: Array) -> Array
Normalized histogram density over the given bin edges.
Written with a scatter-add rather than jax.numpy.histogram so it is
cheap to call inside a jax.lax.scan -- the particle simulations need the
density at every step, since the nonlinear Fokker-Planck drift depends on it.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
samples
|
Array
|
Values to bin, any shape (flattened). |
required |
edges
|
Array
|
Monotone bin edges, shape |
required |
Returns:
| Type | Description |
|---|---|
Array
|
Density in each bin, shape |
Array
|
bin widths (or all zeros if no sample falls inside). |
Source code in qjax/physics/diffusion.py
interpolate_density
¶
interpolate_density(
x: Array, edges: Array, density: Array
) -> Array
Evaluate a binned density at arbitrary points by linear interpolation.
Interpolation is on bin centres, so the two outer half-bins have no
bracketing centre: there the value is held flat at the end bin rather than
faded to zero, which would report an empty edge for a bin that is not empty.
Outside edges the density is zero.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
x
|
Array
|
Points at which to evaluate, any shape. |
required |
edges
|
Array
|
Bin edges the density was built on, shape |
required |
density
|
Array
|
Density per bin, shape |
required |
Returns:
| Type | Description |
|---|---|
Array
|
Interpolated density, same shape as |
Array
|
|
Source code in qjax/physics/diffusion.py
Shared — types, validation, and series¶
types
¶
Type aliases shared across qjax.
These aliases keep signatures readable without imposing a hard dependency on a
specific array backend. At runtime every value is a jax.Array, but the
aliases also accept Python scalars and NumPy arrays, which JAX promotes
automatically.
Array and Scalar were previously the same alias, so the distinction
between "an array of any shape" and "a single real number" was documentation
only and a type checker could not act on it. Scalar is now the narrower of
the two: it excludes the nested sequences that Array accepts, which is the
real constraint on an entropic index.
validation
¶
Validation and broadcasting helpers for the entropic index q.
Tsallis primitives are parameterized by a single real number q (the entropic
index). These helpers normalize q to a JAX scalar and provide a numerically
robust mask for the q -> 1 limit, where most closed-form expressions become
indeterminate (0 / 0).
as_scalar_q
¶
as_scalar_q(q: Scalar) -> Array
Coerce an entropic index to a floating-point JAX scalar.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
q
|
Scalar
|
The entropic index, as a Python number or array-like. |
required |
Returns:
| Type | Description |
|---|---|
Array
|
A 0-d |
Source code in qjax/shared/validation.py
near_one
¶
Boolean mask for indices that should use the q -> 1 (classical) limit.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
q
|
Scalar
|
The entropic index. |
required |
eps
|
float
|
Half-width of the neighborhood around |
Q_EPS
|
Returns:
| Type | Description |
|---|---|
Array
|
A boolean array, broadcastable against |
Array
|
|
Source code in qjax/shared/validation.py
positive_q_or_nan
¶
positive_q_or_nan(q: Scalar) -> Array
Reject a non-positive entropic index.
The Tsallis entropy normalizer \(1/(q(q-1))\) is singular at q = 0,
so the entmax family is undefined for q <= 0. A Python-level
raise is impossible under jax.jit, so the check is split: a
statically known q fails loudly at trace time, while a traced q
(e.g. a learnable parameter that wandered out of range) is mapped to NaN
so the failure is visible downstream instead of silently plausible.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
q
|
Scalar
|
The entropic index. |
required |
Returns:
| Type | Description |
|---|---|
Array
|
|
Array
|
is non-positive. |
Raises:
| Type | Description |
|---|---|
ValueError
|
If |
Source code in qjax/shared/validation.py
series
¶
Entire-function ratios used to take the q -> 1 limit stably.
Every Tsallis closed form is a 0/0 indeterminate at q = 1, and the
textbook way to evaluate it — select a separate Boltzmann-Gibbs expression when
|q - 1| is small — has two defects:
- Between the switch-over point and the region where the deformed form is
accurate there is a band in which neither is: subtracting
1fromx^{1-q} ~ 1loses most of the mantissa. In float32 the relative error of(x^{1-q} - 1)/(1-q)peaks near4e-3aroundq = 1.00001. - The classical branch does not depend on
q, so its derivative with respect toqis identically zero. A learnable entropic index that wanders into the switch-over window sees a hard zero gradient.
Both vanish if the indeterminate quotient is written through the entire
functions below, which are analytic at the origin and equal 1 there. The
classical limit then falls out of the same expression, with no branch on q
and no loss of accuracy.
Near the origin each ratio is evaluated by its Taylor series rather than by the direct form, which keeps the value and the derivative correct.
Plots¶
Plotting requires the optional plots extra: pip install "qjax[plots]".
style
¶
Publication-grade plotting style for qjax, themed on the brand ramp.
use_qjax_style configures Matplotlib for research-grade, vector output:
serif text with Computer-Modern math, embedded fonts, thin in-pointing ticks,
and a color cycle drawn from the qjax ramp. qcolors samples a discrete
sequence from that ramp so a family of curves indexed by q shares a coherent
identity, and save_figure writes a tight, font-embedded PDF.
The ramp¶
QJAX_RAMP is the ten-step green-blue scale the logo and documentation are
built from, ordered light to dark. It is registered with Matplotlib as
"qjax" (plus "qjax_r"), so CMAP works anywhere a colormap name does.
It is a sequential scale: it encodes magnitude, which is exactly what a family
of curves indexed by q needs. Two consequences worth knowing:
qcolorswindows the ramp to[0.40, 1.0]by default. The three lightest steps sit between 1.3:1 and 1.7:1 against white — invisible as thin lines. The window starts where the ramp first clears the 2:1 floor for a sequential light end.- The ramp cannot supply a categorical palette. Exhaustive search over all 1820 four-colour subsets (including interpolated mid-steps) found none that passes the categorical checks: every subset with usable separation (normal-vision OKLab ΔE >= 15) buys it from the extremes that fall outside the lightness band and below 3:1 contrast. Where a figure distinguishes methods rather than magnitudes, carry identity with linestyle and markers and let colour be the secondary cue.
qcolors
¶
Sample n evenly spaced colors from the qjax ramp.
Intended for curves indexed by an ordered parameter — a family of q
values, a noise sweep — where the reader should see the ordering in the
color. For unordered categories (competing methods, class labels) the ramp
cannot give reliable separation; see the module docstring.
The default [lo, hi] window starts partway down the ramp so the lightest
curve still reads against a white page.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
n
|
int
|
Number of colors to return. |
required |
lo
|
float
|
Lower bound of the colormap window, in |
QCOLORS_LO
|
hi
|
float
|
Upper bound of the colormap window, in |
QCOLORS_HI
|
Returns:
| Type | Description |
|---|---|
list
|
A list of |
Source code in qjax/plots/style.py
qlinestyles
¶
Return n distinguishable Matplotlib linestyles.
The brand ramp is sequential, so colour alone cannot separate more than about
three unordered categories: sampling it for four competing methods puts two
dark blues side by side that measure well under the readability floor (OKLab
ΔE ~6 against a floor of 15). Pairing qcolors with these dash patterns
supplies the second, non-colour channel, which also keeps figures readable in
grayscale print and for colour-vision deficiencies.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
n
|
int
|
Number of linestyles to return. |
required |
Returns:
| Type | Description |
|---|---|
list[LineStyle]
|
A list of |
list[LineStyle]
|
number of defined patterns. |
Source code in qjax/plots/style.py
use_qjax_style
¶
Apply the qjax publication style (serif math, vector PDF, brand-ramp cycle).
Source code in qjax/plots/style.py
save_figure
¶
Save fig as a tight, font-embedded vector PDF.
The extension is forced to .pdf and the parent directory is created if
needed, so callers can pass a bare stem like figures/q_gaussian.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
fig
|
Figure
|
The figure to write. |
required |
path
|
str | Path
|
Destination path; any extension is replaced with |
required |
transparent
|
bool
|
If |
False
|
Returns:
| Type | Description |
|---|---|
Path
|
The resolved output path. |
Source code in qjax/plots/style.py
functions
¶
Plots of the q-deformed elementary functions across a range of q.
plot_q_log
¶
plot_q_log(
q_values: Sequence[float] = (0.5, 0.8, 1.0, 1.5, 2.0),
x_range: tuple[float, float] = (0.05, 4.0),
num: int = 400,
ax: Axes | None = None,
) -> Axes
Plot the q-logarithm for several entropic indices.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
q_values
|
Sequence[float]
|
Entropic indices to draw, one curve each. |
(0.5, 0.8, 1.0, 1.5, 2.0)
|
x_range
|
tuple[float, float]
|
|
(0.05, 4.0)
|
num
|
int
|
Number of sample points. |
400
|
ax
|
Axes | None
|
Existing axis to draw on; a new one is created if |
None
|
Returns:
| Type | Description |
|---|---|
Axes
|
The axis containing the plot. |
Source code in qjax/plots/functions.py
plot_q_exp
¶
plot_q_exp(
q_values: Sequence[float] = (0.5, 0.8, 1.0, 1.5, 2.0),
x_range: tuple[float, float] = (-3.0, 2.0),
num: int = 400,
ax: Axes | None = None,
) -> Axes
Plot the q-exponential for several entropic indices.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
q_values
|
Sequence[float]
|
Entropic indices to draw, one curve each. |
(0.5, 0.8, 1.0, 1.5, 2.0)
|
x_range
|
tuple[float, float]
|
|
(-3.0, 2.0)
|
num
|
int
|
Number of sample points. |
400
|
ax
|
Axes | None
|
Existing axis to draw on; a new one is created if |
None
|
Returns:
| Type | Description |
|---|---|
Axes
|
The axis containing the plot. |
Source code in qjax/plots/functions.py
distributions
¶
Plots of the q-Gaussian density across a range of q.
plot_q_gaussian
¶
plot_q_gaussian(
q_values: Sequence[float] = (0.5, 1.0, 1.5, 2.0, 2.5),
beta: float = 1.0,
x_range: tuple[float, float] = (-5.0, 5.0),
num: int = 500,
ax: Axes | None = None,
) -> Axes
Plot the q-Gaussian density for several entropic indices.
Lower q gives compact support; q -> 1 is the Gaussian; higher q
(up to 3) gives progressively heavier tails.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
q_values
|
Sequence[float]
|
Entropic indices to draw, one curve each ( |
(0.5, 1.0, 1.5, 2.0, 2.5)
|
beta
|
float
|
Shared inverse-width parameter. |
1.0
|
x_range
|
tuple[float, float]
|
|
(-5.0, 5.0)
|
num
|
int
|
Number of sample points. |
500
|
ax
|
Axes | None
|
Existing axis to draw on; a new one is created if |
None
|
Returns:
| Type | Description |
|---|---|
Axes
|
The axis containing the plot. |