Skip to content

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

\[ \ln_q(x) = \frac{x^{1-q} - 1}{1 - q}, \qquad \exp_q(x) = \big[1 + (1 - q)\,x\big]_+^{\frac{1}{1-q}}, \]

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:

\[ \ln_q(x) = \log x \cdot \frac{e^{u} - 1}{u},\quad u = (1-q)\log x, \qquad \exp_q(x) = \exp\!\Big(x \cdot \frac{\log(1 + a)}{a}\Big),\quad a = (1-q)x. \]

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 q-logarithm, same shape as x.

Source code in qjax/core/functions.py
def q_log(x: Array, q: Scalar) -> jax.Array:
    r"""``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``.

    Args:
        x: Non-negative input, any shape.
        q: Entropic index (scalar).

    Returns:
        Element-wise ``q``-logarithm, same shape as ``x``.
    """
    x = jnp.asarray(x, dtype=jnp.result_type(float))
    q = as_scalar_q(q)
    one_minus_q = 1.0 - q

    # Evaluate log on a sanitized argument so that x <= 0 contributes no NaN to
    # the gradient of the in-domain branch.
    positive = x > 0.0
    log_x = jnp.log(jnp.where(positive, x, 1.0))
    deformed = log_x * expm1_over_t(one_minus_q * log_x)

    safe_denom = jnp.where(one_minus_q == 0.0, 1.0, one_minus_q)
    at_zero = jnp.where(one_minus_q > 0.0, -1.0 / safe_denom, -jnp.inf)
    return jnp.where(positive, deformed, jnp.where(x == 0.0, at_zero, jnp.nan))

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 q-exponential, same shape as x.

Source code in qjax/core/functions.py
def q_exp(x: Array, q: Scalar) -> jax.Array:
    r"""``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.

    Args:
        x: Input, any shape.
        q: Entropic index (scalar).

    Returns:
        Element-wise ``q``-exponential, same shape as ``x``.
    """
    x = jnp.asarray(x, dtype=jnp.result_type(float))
    q = as_scalar_q(q)
    one_minus_q = 1.0 - q
    a = one_minus_q * x

    # Sanitize before log1p so the clipped region contributes no NaN gradient.
    in_support = a > -1.0
    safe_a = jnp.where(in_support, a, 0.0)
    finite = jnp.exp(x * log1p_over_t(safe_a))

    cut_off = jnp.where(one_minus_q > 0.0, 0.0, jnp.inf)
    return jnp.where(in_support, finite, cut_off)

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 q-sum.

Source code in qjax/core/functions.py
def q_add(a: Array, b: Array, q: Scalar) -> jax.Array:
    r"""``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``.

    Args:
        a: First operand.
        b: Second operand.
        q: Entropic index (scalar).

    Returns:
        Element-wise ``q``-sum.
    """
    a = jnp.asarray(a, dtype=jnp.result_type(float))
    b = jnp.asarray(b, dtype=jnp.result_type(float))
    one_minus_q = 1.0 - as_scalar_q(q)
    return a + b + one_minus_q * a * b

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 q-difference.

Source code in qjax/core/functions.py
def q_diff(a: Array, b: Array, q: Scalar) -> jax.Array:
    r"""``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``.

    Args:
        a: First operand.
        b: Second operand.
        q: Entropic index (scalar).

    Returns:
        Element-wise ``q``-difference.
    """
    a = jnp.asarray(a, dtype=jnp.result_type(float))
    b = jnp.asarray(b, dtype=jnp.result_type(float))
    one_minus_q = 1.0 - as_scalar_q(q)
    return (a - b) / (1.0 + one_minus_q * b)

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 q-product.

Source code in qjax/core/functions.py
def q_prod(a: Array, b: Array, q: Scalar) -> jax.Array:
    r"""``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``.

    Args:
        a: First operand (non-negative).
        b: Second operand (non-negative).
        q: Entropic index (scalar).

    Returns:
        Element-wise ``q``-product.
    """
    a = jnp.asarray(a, dtype=jnp.result_type(float))
    b = jnp.asarray(b, dtype=jnp.result_type(float))
    q = as_scalar_q(q)
    one_minus_q = 1.0 - q
    safe_exp = jnp.where(one_minus_q == 0.0, 1.0, one_minus_q)
    base = _safe_power(a, safe_exp) + _safe_power(b, safe_exp) - 1.0
    return jnp.where(near_one(q), a * b, _clipped_power(base, safe_exp))

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 q-quotient.

Source code in qjax/core/functions.py
def q_div(a: Array, b: Array, q: Scalar) -> jax.Array:
    r"""``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``.

    Args:
        a: Numerator (non-negative).
        b: Denominator (non-negative).
        q: Entropic index (scalar).

    Returns:
        Element-wise ``q``-quotient.
    """
    a = jnp.asarray(a, dtype=jnp.result_type(float))
    b = jnp.asarray(b, dtype=jnp.result_type(float))
    q = as_scalar_q(q)
    one_minus_q = 1.0 - q
    safe_exp = jnp.where(one_minus_q == 0.0, 1.0, one_minus_q)
    base = _safe_power(a, safe_exp) - _safe_power(b, safe_exp) + 1.0
    # Guard the classical quotient too: a bare ``a / b`` back-propagates NaN from
    # a zero denominator even when this branch is unselected (e.g. q = 2). The
    # divide-by-zero limit is spelled with constants rather than as ``a * inf``,
    # whose derivative is ``inf`` and would still poison the zero cotangent.
    nonzero_b = b != 0.0
    signed_inf = jnp.where(a > 0.0, jnp.inf, jnp.where(a < 0.0, -jnp.inf, jnp.nan))
    classical = jnp.where(nonzero_b, a / jnp.where(nonzero_b, b, 1.0), signed_inf)
    return jnp.where(near_one(q), classical, _clipped_power(base, safe_exp))

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(
    p: Array, q: Scalar, axis: int = -1
) -> Array

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 axis: non-negative and normalized.

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 axis.

Source code in qjax/core/entropy.py
def tsallis_entropy(p: Array, q: Scalar, axis: int = -1) -> jax.Array:
    r"""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``.

    Args:
        p: Probability mass values along ``axis``: non-negative and normalized.
        q: Entropic index (scalar).
        axis: Axis over which the distribution is defined.

    Returns:
        Tsallis entropy reduced over ``axis``.
    """
    p = jnp.asarray(p, dtype=jnp.result_type(float))
    q = as_scalar_q(q)

    # Evaluated in the equivalent form S_q(p) = sum_i p_i ln_q(1/p_i), written
    # through the entire function (e^t - 1)/t. This is exact for a normalized p
    # and, unlike a hard `|q - 1| < eps` branch onto the Shannon expression,
    # keeps both the value and the derivative *with respect to q* correct
    # through q = 1 (the classical branch is q-independent, so its q-derivative
    # would be an identically zero cliff).
    #
    # The 0*log(0) = 0 convention is applied by masking on a sanitized p: a bare
    # p * log(p) is masked to 0 in value but still back-propagates
    # 0 * log(0) = 0 * -inf = NaN at a zero coordinate.
    support = p > 0.0
    safe_p = jnp.where(support, p, 1.0)
    log_p = jnp.log(safe_p)
    terms = jnp.where(support, -safe_p * log_p * expm1_over_t((q - 1.0) * log_p), 0.0)
    return jnp.sum(terms, axis=axis)

tsallis_cross_entropy

tsallis_cross_entropy(
    p: Array, y: Array, q: Scalar, axis: int = -1
) -> Array

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 axis.

required
y Array

Target distribution along axis (e.g. one-hot labels).

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 axis.

Source code in qjax/core/entropy.py
def tsallis_cross_entropy(p: Array, y: Array, q: Scalar, axis: int = -1) -> jax.Array:
    r"""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``.

    Args:
        p: Predicted probabilities along ``axis``.
        y: Target distribution along ``axis`` (e.g. one-hot labels).
        q: Entropic index (scalar).
        axis: Axis over which the distributions are defined.

    Returns:
        Tsallis cross-entropy reduced over ``axis``.
    """
    p = jnp.asarray(p, dtype=jnp.result_type(float))
    y = jnp.asarray(y, dtype=jnp.result_type(float))
    # Apply the 0 * ln_q(0) = 0 convention. Where the target mass is zero the
    # term is dropped, and ``q_log`` is evaluated on a safe argument so a zero
    # prediction at an *unused* class (common for sparse ``tsallis_entmax``
    # outputs) yields neither a NaN value nor a NaN gradient. A zero prediction
    # *on* a positive-target class is a genuine +inf loss and is left intact.
    safe_p = jnp.where(y > 0.0, p, 1.0)
    contrib = jnp.where(y > 0.0, y * q_log(safe_p, q), 0.0)
    return -jnp.sum(contrib, axis=axis)

tsallis_divergence

tsallis_divergence(
    p: Array, r: Array, q: Scalar, axis: int = -1
) -> Array

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 axis, non-negative and normalized.

required
r Array

Second (reference) distribution along axis.

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 axis.

Source code in qjax/core/entropy.py
def tsallis_divergence(p: Array, r: Array, q: Scalar, axis: int = -1) -> jax.Array:
    r"""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``.

    Args:
        p: First distribution along ``axis``, non-negative and normalized.
        r: Second (reference) distribution along ``axis``.
        q: Entropic index (scalar).
        axis: Axis over which the distributions are defined.

    Returns:
        Tsallis divergence reduced over ``axis``.
    """
    p = jnp.asarray(p, dtype=jnp.result_type(float))
    r = jnp.asarray(r, dtype=jnp.result_type(float))
    q = as_scalar_q(q)
    q_minus_one = q - 1.0

    # Evaluated as D_q(p||r) = -sum_i p_i ln_q(r_i / p_i) — the equivalent form
    # already named in the docstring — written through the entire function
    # (e^t - 1)/t so that the Kullback-Leibler limit, and its q-derivative, come
    # out of the same expression instead of a separate q-independent branch.
    #
    # One mask serves both regimes. An earlier version guarded only the KL
    # branch, which still leaked NaN: the deformed branch evaluates
    # r ** (1-q) = 0 ** negative = inf at a zero reference and back-propagates
    # NaN even when unselected.
    mask = (p > 0.0) & (r > 0.0)
    safe_p = jnp.where(mask, p, 1.0)
    safe_r = jnp.where(mask, r, 1.0)
    log_ratio = jnp.log(safe_p / safe_r)
    supported = safe_p * log_ratio * expm1_over_t(q_minus_one * log_ratio)

    # Where p > 0 but r == 0 the reference assigns no mass to an outcome p deems
    # possible: ln_q(0) is -inf for q >= 1, giving a genuinely divergent term,
    # and -1/(1-q) for q < 1, giving the finite p/(1-q). Where p == 0 the term
    # vanishes under the 0 * ln_q(0) = 0 convention.
    safe_gap = jnp.where(q_minus_one < 0.0, -q_minus_one, 1.0)
    unsupported = jnp.where(q_minus_one < 0.0, p / safe_gap, jnp.inf)
    terms = jnp.where(mask, supported, jnp.where(p > 0.0, unsupported, 0.0))
    return jnp.sum(terms, axis=axis)

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

\[ p(x) = \frac{\sqrt{\beta}}{C_q}\,\exp_q\!\big(-\beta x^2\big), \]

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 q < 3.

required

Returns:

Type Description
Array

The scalar normalization constant C_q.

Source code in qjax/core/distributions.py
def normalization(q: Scalar) -> jax.Array:
    r"""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))}$

    Args:
        q: Entropic index (scalar), required to satisfy ``q < 3``.

    Returns:
        The scalar normalization constant ``C_q``.
    """
    q = as_scalar_q(q)
    sqrt_pi = jnp.sqrt(jnp.pi)

    def below(qb):  # q < 1
        a = 1.0 - qb
        log_c = (
            jnp.log(2.0)
            + 0.5 * jnp.log(jnp.pi)
            + gammaln(1.0 / a)
            - jnp.log(3.0 - qb)
            - 0.5 * jnp.log(a)
            - gammaln((3.0 - qb) / (2.0 * a))
        )
        return jnp.exp(log_c)

    def above(qa):  # 1 < q < 3
        b = qa - 1.0
        log_c = (
            0.5 * jnp.log(jnp.pi)
            + gammaln((3.0 - qa) / (2.0 * b))
            - 0.5 * jnp.log(b)
            - gammaln(1.0 / b)
        )
        return jnp.exp(log_c)

    # Each branch is evaluated with a sanitized argument that stays strictly
    # inside its own domain, so the unused branch produces neither a NaN value
    # nor a NaN gradient before the final selection by the sign of (q - 1).
    q_safe_below = jnp.where(q < 1.0, q, 0.0)
    q_safe_above = jnp.where(q > 1.0, q, 2.0)
    c_below = below(q_safe_below)
    c_above = above(q_safe_above)
    return jnp.where(q < 1.0, c_below, jnp.where(q > 1.0, c_above, sqrt_pi))

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), q < 3.

1.0
beta Scalar

Inverse-width parameter beta > 0.

1.0

Returns:

Type Description
Array

Density values, same shape as x.

Source code in qjax/core/distributions.py
def q_gaussian_pdf(x: Array, q: Scalar = 1.0, beta: Scalar = 1.0) -> jax.Array:
    r"""Density of the ``q``-Gaussian, $\sqrt{\beta}/C_q\,\exp_q(-\beta x^2)$.

    Args:
        x: Evaluation points, any shape.
        q: Entropic index (scalar), ``q < 3``.
        beta: Inverse-width parameter ``beta > 0``.

    Returns:
        Density values, same shape as ``x``.
    """
    x = jnp.asarray(x, dtype=jnp.result_type(float))
    beta = jnp.asarray(beta, dtype=jnp.result_type(float))
    return jnp.sqrt(beta) / normalization(q) * q_exp(-beta * x**2, q)

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), q < 3.

1.0
beta Scalar

Inverse-width parameter beta > 0.

1.0

Returns:

Type Description
Array

Log-density values, same shape as x.

Source code in qjax/core/distributions.py
def q_gaussian_logpdf(x: Array, q: Scalar = 1.0, beta: Scalar = 1.0) -> jax.Array:
    r"""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.

    Args:
        x: Evaluation points, any shape.
        q: Entropic index (scalar), ``q < 3``.
        beta: Inverse-width parameter ``beta > 0``.

    Returns:
        Log-density values, same shape as ``x``.
    """
    x = jnp.asarray(x, dtype=jnp.result_type(float))
    beta = jnp.asarray(beta, dtype=jnp.result_type(float))
    log_prefactor = 0.5 * jnp.log(beta) - jnp.log(normalization(q))
    return log_prefactor + jnp.log(q_exp(-beta * x**2, q))

sample

sample(
    key: Array,
    q: Scalar = 1.0,
    beta: Scalar = 1.0,
    shape: tuple[int, ...] = (),
) -> Array

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,

\[ X = \frac{T_\nu}{\sqrt{(3-q)\,\beta}}, \qquad T_\nu \sim \mathrm{t}(\nu), \]

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 jax.random.PRNGKey.

required
q Scalar

Entropic index (scalar), 1 <= q < 3.

1.0
beta Scalar

Inverse-width parameter beta > 0.

1.0
shape tuple[int, ...]

Output shape.

()

Returns:

Type Description
Array

Samples of the requested shape.

Source code in qjax/core/distributions.py
def sample(
    key: jax.Array,
    q: Scalar = 1.0,
    beta: Scalar = 1.0,
    shape: tuple[int, ...] = (),
) -> jax.Array:
    r"""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,

    $$
    X = \frac{T_\nu}{\sqrt{(3-q)\,\beta}}, \qquad T_\nu \sim \mathrm{t}(\nu),
    $$

    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.

    Args:
        key: A `jax.random.PRNGKey`.
        q: Entropic index (scalar), ``1 <= q < 3``.
        beta: Inverse-width parameter ``beta > 0``.
        shape: Output shape.

    Returns:
        Samples of the requested ``shape``.
    """
    q = as_scalar_q(q)
    beta = jnp.asarray(beta, dtype=jnp.result_type(float))
    k_z, k_w = jax.random.split(key)
    z = jax.random.normal(k_z, shape)

    def gaussian(_):
        return z / jnp.sqrt(2.0 * beta)

    def student_t(_):
        nu = (3.0 - q) / (q - 1.0)
        w = 2.0 * jax.random.gamma(k_w, nu / 2.0, shape)  # chi-square with nu dof
        t = z / jnp.sqrt(w / nu)
        return t / jnp.sqrt((3.0 - q) * beta)

    return jax.lax.cond(near_one(q), gaussian, student_t, operand=None)

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:

\[ \mathrm{entmax}_q(z) = \arg\max_{p \in \Delta}\; \langle p, z \rangle + S_q^{T}(p), \]

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)

\[ p_i = \big[(q - 1)\,z_i - \tau\big]_+^{\,1/(q-1)}, \]

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:

\[ \dot p = T\!\Big(s \odot \dot z + \big(h + s \odot z\big)\tfrac{\dot q}{q - 1}\Big), \]

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(
    z: Array,
    q: Scalar = 2.0,
    axis: int = -1,
    num_iters: int = 50,
) -> Array

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 axis.

required
q Scalar

Entropic index (scalar), q > 0. q = 1 -> softmax, q = 2 -> sparsemax. Values above 1 give sparse outputs; values below 1 give outputs denser than softmax.

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 z that sum to one along axis.

Raises:

Type Description
ValueError

If q is a concrete value and q <= 0. A traced, non-positive q yields NaN instead.

Source code in qjax/core/activations.py
def tsallis_entmax(
    z: Array,
    q: Scalar = 2.0,
    axis: int = -1,
    num_iters: int = 50,
) -> jax.Array:
    r"""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)``.

    Args:
        z: Scores / logits; the distribution is formed along ``axis``.
        q: Entropic index (scalar), ``q > 0``. ``q = 1`` -> softmax,
            ``q = 2`` -> sparsemax. Values above ``1`` give sparse outputs;
            values below ``1`` give outputs denser than softmax.
        axis: Axis over which to normalize.
        num_iters: Number of bisection steps for the threshold search. A Newton
            polish step follows, so the result is accurate even for small values.

    Returns:
        Probabilities with the same shape as ``z`` that sum to one along ``axis``.

    Raises:
        ValueError: If ``q`` is a concrete value and ``q <= 0``. A traced,
            non-positive ``q`` yields ``NaN`` instead.
    """
    z = jnp.asarray(z, dtype=jnp.result_type(float))
    q = positive_q_or_nan(as_scalar_q(q))

    # Move the working axis to the end for a uniform reduction layout.
    z = jnp.moveaxis(z, axis, -1)

    near = near_one(q)
    # Under vmap over a batched q, lax.cond lowers to a select and *both*
    # branches execute, so keep the entmax branch clear of the 1/(q-1) pole.
    # The clamped region is exactly the region the select discards.
    q_eff = jnp.where(near, jnp.where(q < 1.0, 1.0 - Q_EPS, 1.0 + Q_EPS), q)

    p = jax.lax.cond(
        near,
        lambda _: _softmax_core(z, q),
        lambda _: _entmax_core(z, q_eff, num_iters),
        operand=None,
    )
    return jnp.moveaxis(p, -1, axis)

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:

\[ q = q_{\min} + (q_{\max} - q_{\min})\,\sigma(q_{\mathrm{raw}}). \]

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 q_raw.

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
def bounded_q(q_raw: Scalar, lo: Scalar = 1.0, hi: Scalar = 3.0) -> jax.Array:
    r"""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.

    Args:
        q_raw: Unconstrained real parameter, any shape.
        lo: Lower bound of the open interval (exclusive).
        hi: Upper bound of the open interval (exclusive).

    Returns:
        The entropic index, same shape as ``q_raw``.

    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
    """
    lo = jnp.asarray(lo, dtype=jnp.result_type(float))
    hi = jnp.asarray(hi, dtype=jnp.result_type(float))
    return lo + (hi - lo) * jax.nn.sigmoid(jnp.asarray(q_raw, dtype=jnp.result_type(float)))

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 (lo, hi).

required
lo Scalar

Lower bound used by bounded_q.

1.0
hi Scalar

Upper bound used by bounded_q.

3.0

Returns:

Type Description
Array

The q_raw for which bounded_q(q_raw, lo, hi) == q.

Source code in qjax/nn/reparam.py
def inverse_bounded_q(q: Scalar, lo: Scalar = 1.0, hi: Scalar = 3.0) -> jax.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.

    Args:
        q: Target entropic index, strictly inside ``(lo, hi)``.
        lo: Lower bound used by `bounded_q`.
        hi: Upper bound used by `bounded_q`.

    Returns:
        The ``q_raw`` for which ``bounded_q(q_raw, lo, hi) == q``.
    """
    q = jnp.asarray(q, dtype=jnp.result_type(float))
    lo = jnp.asarray(lo, dtype=jnp.result_type(float))
    hi = jnp.asarray(hi, dtype=jnp.result_type(float))
    unit = (q - lo) / (hi - lo)
    return jnp.log(unit) - jnp.log1p(-unit)

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 (..., n_queries, d_key). A single query per batch element may be passed as (..., d_key).

required
keys Array

Key vectors, shape (..., n_keys, d_key).

required
values Array

Value vectors, shape (..., n_keys, d_value).

required
q Scalar

Entropic index (scalar), q > 0. 1 is softmax attention, 2 is sparsemax attention.

2.0
mask Array | None

Optional boolean array broadcastable to the score shape (..., n_queries, n_keys). False positions are excluded.

None
scale Scalar | None

Divisor applied to the scores. Defaults to sqrt(d_key).

None
num_iters int

Bisection steps for the entmax threshold search.

50

Returns:

Type Description
Array

A (context, attention) pair. context has shape

Array

(..., n_queries, d_value) and attention has shape

tuple[Array, Array]

(..., n_queries, n_keys) and sums to one over the last axis.

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
def 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[jax.Array, jax.Array]:
    r"""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.

    Args:
        queries: Query vectors, shape ``(..., n_queries, d_key)``. A single
            query per batch element may be passed as ``(..., d_key)``.
        keys: Key vectors, shape ``(..., n_keys, d_key)``.
        values: Value vectors, shape ``(..., n_keys, d_value)``.
        q: Entropic index (scalar), ``q > 0``. ``1`` is softmax attention,
            ``2`` is sparsemax attention.
        mask: Optional boolean array broadcastable to the score shape
            ``(..., n_queries, n_keys)``. ``False`` positions are excluded.
        scale: Divisor applied to the scores. Defaults to ``sqrt(d_key)``.
        num_iters: Bisection steps for the ``entmax`` threshold search.

    Returns:
        A ``(context, attention)`` pair. ``context`` has shape
        ``(..., n_queries, d_value)`` and ``attention`` has shape
        ``(..., n_queries, n_keys)`` and sums to one over the last axis.

    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))
    """
    queries = jnp.asarray(queries, dtype=jnp.result_type(float))
    keys = jnp.asarray(keys, dtype=jnp.result_type(float))
    values = jnp.asarray(values, dtype=jnp.result_type(float))

    # A single query per batch element is the common case in attention pooling;
    # add the query axis, then drop it again on the way out.
    squeeze_query = queries.ndim == keys.ndim - 1
    if squeeze_query:
        queries = queries[..., None, :]

    if scale is None:
        scale = jnp.sqrt(jnp.asarray(keys.shape[-1], dtype=queries.dtype))
    scores = jnp.einsum("...qd,...kd->...qk", queries, keys) / scale

    if mask is not None:
        # -inf scores are driven to exactly zero weight by entmax for every q,
        # and keep the masked positions out of the threshold search.
        scores = jnp.where(jnp.asarray(mask, dtype=bool), scores, -jnp.inf)

    attention = tsallis_entmax(scores, q=q, axis=-1, num_iters=num_iters)
    context = jnp.einsum("...qk,...kd->...qd", attention, values)

    if squeeze_query:
        context = context[..., 0, :]
        attention = attention[..., 0, :]
    return context, attention

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,

\[ \ln_q(0) = \frac{-1}{1 - q} \quad (q < 1), \]

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 from_logits (the default), otherwise probabilities that already sum to one along axis.

required
targets Array

Target distribution along axis. Either one-hot or soft; integer class indices are not accepted, so encode them first.

required
q Scalar

Entropic index of the loss. 1 is the standard cross-entropy; values below 1 bound the penalty on confidently wrong predictions (robust to label noise); values above 1 sharpen it.

1.0
from_logits bool

Whether to normalize the input first.

True
normalizer_q Scalar | None

Entropic index of the tsallis_entmax used to normalize logits. Defaults to q, coupling the two; pass 1.0 to keep an ordinary softmax under a deformed loss. Ignored when from_logits is False.

None
axis int

Axis holding the class distribution.

-1
reduction str

"mean", "sum", or "none".

'mean'

Returns:

Type Description
Array

The reduced loss, or the per-example losses when reduction="none".

Raises:

Type Description
ValueError

If reduction is not one of the three accepted values.

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
def 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",
) -> jax.Array:
    r"""``q``-deformed cross-entropy loss.

    Args:
        logits_or_probs: Unnormalized scores if ``from_logits`` (the default),
            otherwise probabilities that already sum to one along ``axis``.
        targets: Target distribution along ``axis``. Either one-hot or soft;
            integer class indices are *not* accepted, so encode them first.
        q: Entropic index of the *loss*. ``1`` is the standard cross-entropy;
            values below ``1`` bound the penalty on confidently wrong
            predictions (robust to label noise); values above ``1`` sharpen it.
        from_logits: Whether to normalize the input first.
        normalizer_q: Entropic index of the `tsallis_entmax` used to
            normalize logits. Defaults to ``q``, coupling the two; pass ``1.0``
            to keep an ordinary softmax under a deformed loss. Ignored when
            ``from_logits`` is ``False``.
        axis: Axis holding the class distribution.
        reduction: ``"mean"``, ``"sum"``, or ``"none"``.

    Returns:
        The reduced loss, or the per-example losses when ``reduction="none"``.

    Raises:
        ValueError: If ``reduction`` is not one of the three accepted values.

    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.
    """
    if reduction not in ("mean", "sum", "none"):
        raise ValueError(f"reduction must be 'mean', 'sum', or 'none'; got {reduction!r}.")

    logits_or_probs = jnp.asarray(logits_or_probs, dtype=jnp.result_type(float))
    targets = jnp.asarray(targets, dtype=jnp.result_type(float))

    if from_logits:
        norm_q = q if normalizer_q is None else normalizer_q
        probs = tsallis_entmax(logits_or_probs, q=norm_q, axis=axis)
    else:
        probs = logits_or_probs

    per_example = tsallis_cross_entropy(probs, targets, q, axis=axis)
    if reduction == "mean":
        return jnp.mean(per_example)
    if reduction == "sum":
        return jnp.sum(per_example)
    return per_example

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

\[p_\theta(s) = \prod_{i=1}^{N} p_\theta(s_i \mid s_1, \dots, s_{i-1})\]

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

made_masks(
    num_spins: int, hidden: Sequence[int]
) -> list[Array]

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 N.

required
hidden Sequence[int]

Widths of the hidden layers, at least one.

required

Returns:

Type Description
list[Array]

A list of len(hidden) + 1 masks aligned with the weight matrices of

list[Array]

made_init; mask k has shape (in_k, out_k).

Raises:

Type Description
ValueError

If hidden is empty or num_spins is below 2.

Source code in qjax/nn/autoregressive.py
def made_masks(num_spins: int, hidden: Sequence[int]) -> list[jax.Array]:
    r"""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.

    Args:
        num_spins: Number of spins ``N``.
        hidden: Widths of the hidden layers, at least one.

    Returns:
        A list of ``len(hidden) + 1`` masks aligned with the weight matrices of
        `made_init`; mask ``k`` has shape ``(in_k, out_k)``.

    Raises:
        ValueError: If ``hidden`` is empty or ``num_spins`` is below 2.
    """
    if not hidden:
        raise ValueError("MADE needs at least one hidden layer.")
    if num_spins < 2:
        raise ValueError(f"MADE needs at least 2 spins; got {num_spins}.")

    input_degrees = jnp.arange(1, num_spins + 1)
    # Hidden degrees cycle through 1..N-1: degree N would let a unit see every
    # input, and no output could then use it.
    hidden_degrees = [1 + jnp.arange(width) % (num_spins - 1) for width in hidden]

    masks = []
    previous = input_degrees
    for degrees in hidden_degrees:
        masks.append((previous[:, None] <= degrees[None, :]).astype(jnp.result_type(float)))
        previous = degrees
    masks.append((previous[:, None] < input_degrees[None, :]).astype(jnp.result_type(float)))
    return masks

made_init

made_init(
    key: Array, num_spins: int, hidden: Sequence[int]
) -> Params

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 N.

required
hidden Sequence[int]

Widths of the hidden layers.

required

Returns:

Type Description
Params

A pytree {"weights": [...], "biases": [...]} whose weight shapes

Params

match the masks from made_masks.

Source code in qjax/nn/autoregressive.py
def made_init(key: jax.Array, num_spins: int, hidden: Sequence[int]) -> Params:
    """Initialize MADE parameters with Glorot-scaled weights and zero biases.

    Args:
        key: PRNG key.
        num_spins: Number of spins ``N``.
        hidden: Widths of the hidden layers.

    Returns:
        A pytree ``{"weights": [...], "biases": [...]}`` whose weight shapes
        match the masks from `made_masks`.
    """
    widths = [num_spins, *hidden, num_spins]
    keys = jax.random.split(key, len(widths) - 1)
    weights, biases = [], []
    for layer_key, fan_in, fan_out in zip(keys, widths[:-1], widths[1:], strict=True):
        scale = jnp.sqrt(2.0 / (fan_in + fan_out))
        weights.append(jax.random.normal(layer_key, (fan_in, fan_out)) * scale)
        biases.append(jnp.zeros((fan_out,)))
    return {"weights": weights, "biases": biases}

made_conditionals

made_conditionals(
    params: Params, masks: Sequence[Array], spins: Array
) -> Array

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 made_init.

required
masks Sequence[Array]

Masks from made_masks.

required
spins Array

Configurations of shape (B, N) with entries in {-1, +1}. Entries at or after position i cannot influence logit i, so partially filled configurations are safe to pass during sampling.

required

Returns:

Type Description
Array

Logits of shape (B, N).

Source code in qjax/nn/autoregressive.py
def made_conditionals(params: Params, masks: Sequence[jax.Array], spins: Array) -> jax.Array:
    r"""Logits of $p(s_i = +1 \mid s_{<i})$ for every site, in one forward pass.

    Args:
        params: Parameters from `made_init`.
        masks: Masks from `made_masks`.
        spins: Configurations of shape ``(B, N)`` with entries in ``{-1, +1}``.
            Entries at or after position ``i`` cannot influence logit ``i``, so
            partially filled configurations are safe to pass during sampling.

    Returns:
        Logits of shape ``(B, N)``.
    """
    weights, biases = params["weights"], params["biases"]
    activation = jnp.asarray(spins, dtype=jnp.result_type(float))
    for weight, bias, mask in zip(weights[:-1], biases[:-1], masks[:-1], strict=True):
        activation = jnp.tanh(activation @ (weight * mask) + bias)
    return activation @ (weights[-1] * masks[-1]) + biases[-1]

made_log_prob

made_log_prob(
    params: Params, masks: Sequence[Array], spins: Array
) -> Array

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 made_init.

required
masks Sequence[Array]

Masks from made_masks.

required
spins Array

Configurations of shape (B, N) with entries in {-1, +1}.

required

Returns:

Type Description
Array

Log-probabilities of shape (B,). Summed over all 2**N

Array

configurations these exponentiate to exactly 1.

Source code in qjax/nn/autoregressive.py
def made_log_prob(params: Params, masks: Sequence[jax.Array], spins: Array) -> jax.Array:
    r"""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.

    Args:
        params: Parameters from `made_init`.
        masks: Masks from `made_masks`.
        spins: Configurations of shape ``(B, N)`` with entries in ``{-1, +1}``.

    Returns:
        Log-probabilities of shape ``(B,)``. Summed over all ``2**N``
        configurations these exponentiate to exactly ``1``.
    """
    spins = jnp.asarray(spins, dtype=jnp.result_type(float))
    logits = made_conditionals(params, masks, spins)
    return -jnp.sum(jax.nn.softplus(-spins * logits), axis=-1)

made_sample

made_sample(
    key: Array,
    params: Params,
    masks: Sequence[Array],
    num_samples: int,
) -> Array

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 made_init.

required
masks Sequence[Array]

Masks from made_masks.

required
num_samples int

Number of configurations to draw.

required

Returns:

Type Description
Array

Configurations of shape (num_samples, N) with entries in

Array

{-1.0, +1.0}, distributed exactly as made_log_prob says.

Source code in qjax/nn/autoregressive.py
def made_sample(
    key: jax.Array, params: Params, masks: Sequence[jax.Array], num_samples: int
) -> jax.Array:
    """Draw exact samples by filling in one spin at a time.

    Args:
        key: PRNG key.
        params: Parameters from `made_init`.
        masks: Masks from `made_masks`.
        num_samples: Number of configurations to draw.

    Returns:
        Configurations of shape ``(num_samples, N)`` with entries in
        ``{-1.0, +1.0}``, distributed exactly as `made_log_prob` says.
    """
    num_spins = masks[0].shape[0]
    Carry = tuple[jax.Array, jax.Array]

    def step(carry: Carry, site: jax.Array) -> tuple[Carry, None]:
        chain_key, state = carry
        chain_key, subkey = jax.random.split(chain_key)
        logits = made_conditionals(params, masks, state)
        probability = jax.nn.sigmoid(logits[:, site])
        draw = jax.random.uniform(subkey, (num_samples,)) < probability
        return (chain_key, state.at[:, site].set(jnp.where(draw, 1.0, -1.0))), None

    initial: Carry = (key, jnp.zeros((num_samples, num_spins), dtype=jnp.result_type(float)))
    (_, spins), _ = jax.lax.scan(step, initial, jnp.arange(num_spins))
    return spins

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,

\[H(s) = -J \sum_{\langle i j \rangle} s_i s_j, \qquad s_i \in \{-1, +1\},\]

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:

  1. ising_exact_observables -- exhaustive enumeration of all \(2^{L^2}\) states, for \(L \le 4\).
  2. ising_transfer_matrix_log_z -- the \(2^L \times 2^L\) transfer matrix, exact for the finite periodic lattice, for \(L \le 10\).
  3. 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 (..., L, L).

required

Returns:

Type Description
Array

An array of the same shape whose entry [..., i, j] is the sum of the

Array

spins at the four sites adjacent to (i, j).

Source code in qjax/physics/lattice.py
def neighbour_sum(spins: Array) -> jax.Array:
    """Sum of the four nearest neighbours of every site, with periodic wrap.

    Args:
        spins: Spin configuration(s) of shape ``(..., L, L)``.

    Returns:
        An array of the same shape whose entry ``[..., i, j]`` is the sum of the
        spins at the four sites adjacent to ``(i, j)``.
    """
    spins = jnp.asarray(spins, dtype=jnp.result_type(float))
    return (
        jnp.roll(spins, 1, axis=-1)
        + jnp.roll(spins, -1, axis=-1)
        + jnp.roll(spins, 1, axis=-2)
        + jnp.roll(spins, -1, axis=-2)
    )

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 (..., L, L).

required
coupling Scalar

Exchange coupling J. Positive is ferromagnetic.

1.0

Returns:

Type Description
Array

Total energy per configuration, shape (...). For the all-aligned

Array

state this is -2 J L**2.

Source code in qjax/physics/lattice.py
def ising_energy(spins: Array, coupling: Scalar = 1.0) -> jax.Array:
    r"""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$.

    Args:
        spins: Spin configuration(s) of shape ``(..., L, L)``.
        coupling: Exchange coupling ``J``. Positive is ferromagnetic.

    Returns:
        Total energy per configuration, shape ``(...)``. For the all-aligned
        state this is ``-2 J L**2``.
    """
    spins = jnp.asarray(spins, dtype=jnp.result_type(float))
    pairs = jnp.sum(spins * neighbour_sum(spins), axis=(-2, -1))
    return -0.5 * jnp.asarray(coupling, dtype=jnp.result_type(float)) * pairs

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 (..., L, L).

required

Returns:

Type Description
Array

Magnetization per site, shape (...), in [-1, 1].

Source code in qjax/physics/lattice.py
def ising_magnetization(spins: Array) -> jax.Array:
    """Signed magnetization per site, ``mean(s)``.

    Args:
        spins: Spin configuration(s) of shape ``(..., L, L)``.

    Returns:
        Magnetization per site, shape ``(...)``, in ``[-1, 1]``.
    """
    spins = jnp.asarray(spins, dtype=jnp.result_type(float))
    return jnp.mean(spins, axis=(-2, -1))

checkerboard_sweep

checkerboard_sweep(
    key: Array,
    spins: Array,
    beta: Scalar,
    coupling: Scalar = 1.0,
) -> Array

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 (L, L).

required
beta Scalar

Inverse temperature 1 / T (a scalar; vmap for a batch).

required
coupling Scalar

Exchange coupling J.

1.0

Returns:

Type Description
Array

The configuration after one sweep, shape (L, L).

Source code in qjax/physics/lattice.py
def checkerboard_sweep(
    key: jax.Array, spins: jax.Array, beta: Scalar, coupling: Scalar = 1.0
) -> jax.Array:
    """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.

    Args:
        key: PRNG key.
        spins: A single configuration of shape ``(L, L)``.
        beta: Inverse temperature ``1 / T`` (a scalar; ``vmap`` for a batch).
        coupling: Exchange coupling ``J``.

    Returns:
        The configuration after one sweep, shape ``(L, L)``.
    """
    beta = jnp.asarray(beta, dtype=jnp.result_type(float))
    coupling = jnp.asarray(coupling, dtype=jnp.result_type(float))
    size = spins.shape[-1]
    index = jnp.arange(size)
    parity = (index[:, None] + index[None, :]) % 2

    for colour in (0, 1):
        key, subkey = jax.random.split(key)
        # Flipping s_i costs Delta E = 2 J s_i n_i.
        delta = 2.0 * coupling * spins * neighbour_sum(spins)
        # min(1, exp(-beta dE)); the clamp keeps exp() from overflowing to inf
        # for strongly uphill moves, which is a no-op for the comparison.
        accept_prob = jnp.exp(jnp.minimum(-beta * delta, 0.0))
        accepted = jax.random.uniform(subkey, spins.shape) < accept_prob
        spins = jnp.where(accepted & (parity == colour), -spins, spins)
    return spins

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 (L, L).

required
beta Scalar

Inverse temperature 1 / T (a scalar; vmap for a batch).

required
sweeps int

Number of full sweeps (a Python int; it sets the scan length).

required
coupling Scalar

Exchange coupling J.

1.0

Returns:

Type Description
Array

The configuration after sweeps sweeps, shape (L, L).

Source code in qjax/physics/lattice.py
def metropolis_chain(
    key: jax.Array, spins: jax.Array, beta: Scalar, sweeps: int, coupling: Scalar = 1.0
) -> jax.Array:
    """Run ``sweeps`` Metropolis sweeps and return the final configuration.

    Args:
        key: PRNG key.
        spins: Initial configuration of shape ``(L, L)``.
        beta: Inverse temperature ``1 / T`` (a scalar; ``vmap`` for a batch).
        sweeps: Number of full sweeps (a Python int; it sets the scan length).
        coupling: Exchange coupling ``J``.

    Returns:
        The configuration after ``sweeps`` sweeps, shape ``(L, L)``.
    """
    Carry = tuple[jax.Array, jax.Array]

    def step(carry: Carry, _: None) -> tuple[Carry, None]:
        chain_key, state = carry
        chain_key, subkey = jax.random.split(chain_key)
        return (chain_key, checkerboard_sweep(subkey, state, beta, coupling)), None

    (_, final), _ = jax.lax.scan(step, (key, spins), None, length=sweeps)
    return final

wolff_update

wolff_update(
    key: Array,
    spins: Array,
    beta: Scalar,
    coupling: Scalar = 1.0,
) -> Array

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 (L, L).

required
beta Scalar

Inverse temperature 1 / T (a scalar; vmap for a batch).

required
coupling Scalar

Exchange coupling J, positive (ferromagnetic).

1.0

Returns:

Type Description
Array

The configuration after one cluster flip, shape (L, L).

Source code in qjax/physics/lattice.py
def wolff_update(
    key: jax.Array, spins: jax.Array, beta: Scalar, coupling: Scalar = 1.0
) -> jax.Array:
    r"""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.

    Args:
        key: PRNG key.
        spins: A single configuration of shape ``(L, L)``.
        beta: Inverse temperature ``1 / T`` (a scalar; ``vmap`` for a batch).
        coupling: Exchange coupling ``J``, positive (ferromagnetic).

    Returns:
        The configuration after one cluster flip, shape ``(L, L)``.
    """
    beta = jnp.asarray(beta, dtype=jnp.result_type(float))
    coupling = jnp.asarray(coupling, dtype=jnp.result_type(float))
    size = spins.shape[-1]
    add_probability = -jnp.expm1(-2.0 * beta * coupling)

    seed_key, grow_key = jax.random.split(key)
    seed = jax.random.randint(seed_key, (2,), 0, size)
    start = jnp.zeros((size, size), dtype=bool).at[seed[0], seed[1]].set(True)
    aligned = spins == spins[seed[0], seed[1]]

    Carry = tuple[jax.Array, jax.Array, jax.Array]

    def growing(carry: Carry) -> jax.Array:
        _, _, frontier = carry
        return jnp.any(frontier)

    def grow(carry: Carry) -> Carry:
        grow_key, cluster, frontier = carry
        grow_key, right_key, down_key = jax.random.split(grow_key, 3)
        # One uniform per bond per iteration: ``right[i, j]`` is the bond from
        # (i, j) to (i, j+1) and ``down[i, j]`` the bond to (i+1, j).
        right = jax.random.uniform(right_key, (size, size)) < add_probability
        down = jax.random.uniform(down_key, (size, size)) < add_probability
        reached = (
            jnp.roll(frontier & right, 1, axis=-1)  # bond crossed rightward
            | (jnp.roll(frontier, -1, axis=-1) & right)  # ... and leftward
            | jnp.roll(frontier & down, 1, axis=-2)  # downward
            | (jnp.roll(frontier, -1, axis=-2) & down)  # upward
        )
        accepted = reached & aligned & ~cluster
        return grow_key, cluster | accepted, accepted

    _, cluster, _ = jax.lax.while_loop(growing, grow, (grow_key, start, start))
    return jnp.where(cluster, -spins, spins)

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 (L, L).

required
beta Scalar

Inverse temperature 1 / T (a scalar; vmap for a batch).

required
updates int

Number of cluster updates (a Python int; it sets the scan length).

required
coupling Scalar

Exchange coupling J.

1.0

Returns:

Type Description
Array

The configuration after updates cluster flips, shape (L, L).

Source code in qjax/physics/lattice.py
def wolff_chain(
    key: jax.Array, spins: jax.Array, beta: Scalar, updates: int, coupling: Scalar = 1.0
) -> jax.Array:
    """Run ``updates`` Wolff cluster updates and return the final configuration.

    Args:
        key: PRNG key.
        spins: Initial configuration of shape ``(L, L)``.
        beta: Inverse temperature ``1 / T`` (a scalar; ``vmap`` for a batch).
        updates: Number of cluster updates (a Python int; it sets the scan
            length).
        coupling: Exchange coupling ``J``.

    Returns:
        The configuration after ``updates`` cluster flips, shape ``(L, L)``.
    """
    Carry = tuple[jax.Array, jax.Array]

    def step(carry: Carry, _: None) -> tuple[Carry, None]:
        chain_key, state = carry
        chain_key, subkey = jax.random.split(chain_key)
        return (chain_key, wolff_update(subkey, state, beta, coupling)), None

    (_, final), _ = jax.lax.scan(step, (key, spins), None, length=updates)
    return final

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 L.

required
temperatures Array

Temperatures to sample at, shape (T,).

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 algorithm="wolff".

required
coupling Scalar

Exchange coupling J.

1.0
algorithm str

"metropolis" for checkerboard_sweep, "wolff" for wolff_update.

'metropolis'

Returns:

Type Description
Array

Configurations of shape (T, num_samples, L, L) with entries in

Array

{-1.0, +1.0}.

Raises:

Type Description
ValueError

If algorithm is neither "metropolis" nor "wolff".

Source code in qjax/physics/lattice.py
def sample_ising(
    key: jax.Array,
    size: int,
    temperatures: Array,
    num_samples: int,
    sweeps: int,
    coupling: Scalar = 1.0,
    algorithm: str = "metropolis",
) -> jax.Array:
    r"""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.

    Args:
        key: PRNG key.
        size: Linear lattice size ``L``.
        temperatures: Temperatures to sample at, shape ``(T,)``.
        num_samples: Independent configurations per temperature.
        sweeps: Equilibration steps per chain (a Python int) -- Metropolis
            sweeps, or Wolff cluster updates when ``algorithm="wolff"``.
        coupling: Exchange coupling ``J``.
        algorithm: ``"metropolis"`` for `checkerboard_sweep`, ``"wolff"`` for
            `wolff_update`.

    Returns:
        Configurations of shape ``(T, num_samples, L, L)`` with entries in
        ``{-1.0, +1.0}``.

    Raises:
        ValueError: If ``algorithm`` is neither ``"metropolis"`` nor
            ``"wolff"``.
    """
    if algorithm not in ("metropolis", "wolff"):
        raise ValueError(f"algorithm must be 'metropolis' or 'wolff'; got {algorithm!r}.")
    chain = metropolis_chain if algorithm == "metropolis" else wolff_chain
    temperatures = jnp.asarray(temperatures, dtype=jnp.result_type(float))
    num_temperatures = temperatures.shape[0]
    num_chains = num_temperatures * num_samples

    keys = jax.random.split(key, num_chains + 1)
    start = jax.random.bernoulli(keys[0], 0.5, (num_chains, size, size))
    initial = jnp.where(start, 1.0, -1.0)
    betas = jnp.repeat(1.0 / temperatures, num_samples)

    run = jax.vmap(lambda k, s, b: chain(k, s, b, sweeps, coupling))
    final = run(keys[1:], initial, betas)
    return final.reshape(num_temperatures, num_samples, size, size)

ising_all_configurations

ising_all_configurations(size: int) -> Array

Enumerate every spin configuration of an L x L lattice.

Parameters:

Name Type Description Default
size int

Linear lattice size L; L**2 must not exceed MAX_ENUMERATED_SPINS.

required

Returns:

Type Description
Array

All configurations, shape (2**(L*L), L, L), entries in

Array

{-1.0, +1.0}.

Raises:

Type Description
ValueError

If L**2 exceeds MAX_ENUMERATED_SPINS.

Source code in qjax/physics/lattice.py
def ising_all_configurations(size: int) -> jax.Array:
    """Enumerate every spin configuration of an ``L x L`` lattice.

    Args:
        size: Linear lattice size ``L``; ``L**2`` must not exceed
            `MAX_ENUMERATED_SPINS`.

    Returns:
        All configurations, shape ``(2**(L*L), L, L)``, entries in
        ``{-1.0, +1.0}``.

    Raises:
        ValueError: If ``L**2`` exceeds `MAX_ENUMERATED_SPINS`.
    """
    num_spins = size * size
    if num_spins > MAX_ENUMERATED_SPINS:
        raise ValueError(
            f"enumerating {num_spins} spins needs 2**{num_spins} states; "
            f"the limit is {MAX_ENUMERATED_SPINS}."
        )
    states = jnp.arange(2**num_spins, dtype=jnp.int32)
    bits = (states[:, None] >> jnp.arange(num_spins, dtype=jnp.int32)[None, :]) & 1
    spins = 1.0 - 2.0 * bits.astype(jnp.result_type(float))
    return spins.reshape(-1, size, size)

ising_boltzmann_probabilities

ising_boltzmann_probabilities(
    size: int, temperature: Scalar, coupling: Scalar = 1.0
) -> Array

Exact Boltzmann weights over the full state space, in enumeration order.

Parameters:

Name Type Description Default
size int

Linear lattice size L.

required
temperature Scalar

Temperature T.

required
coupling Scalar

Exchange coupling J.

1.0

Returns:

Type Description
Array

Normalized probabilities of shape (2**(L*L),), ordered to match

Array

ising_all_configurations.

Source code in qjax/physics/lattice.py
def ising_boltzmann_probabilities(
    size: int, temperature: Scalar, coupling: Scalar = 1.0
) -> jax.Array:
    """Exact Boltzmann weights over the full state space, in enumeration order.

    Args:
        size: Linear lattice size ``L``.
        temperature: Temperature ``T``.
        coupling: Exchange coupling ``J``.

    Returns:
        Normalized probabilities of shape ``(2**(L*L),)``, ordered to match
        `ising_all_configurations`.
    """
    energies = ising_energy(ising_all_configurations(size), coupling)
    log_weights = -energies / jnp.asarray(temperature, dtype=jnp.result_type(float))
    return jnp.exp(log_weights - logsumexp(log_weights))

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 L; L**2 <= MAX_ENUMERATED_SPINS.

required
temperature Scalar

Temperature T.

required
coupling Scalar

Exchange coupling J.

1.0

Returns:

Type Description
dict[str, Array]

A dict with log_z, free_energy_per_site, energy_per_site,

dict[str, Array]

abs_magnetization, magnetization_squared and

dict[str, Array]

heat_capacity (per site).

Source code in qjax/physics/lattice.py
def ising_exact_observables(
    size: int, temperature: Scalar, coupling: Scalar = 1.0
) -> dict[str, jax.Array]:
    """Exact thermodynamics by exhaustive enumeration of the state space.

    Args:
        size: Linear lattice size ``L``; ``L**2 <=`` `MAX_ENUMERATED_SPINS`.
        temperature: Temperature ``T``.
        coupling: Exchange coupling ``J``.

    Returns:
        A dict with ``log_z``, ``free_energy_per_site``, ``energy_per_site``,
        ``abs_magnetization``, ``magnetization_squared`` and
        ``heat_capacity`` (per site).
    """
    temperature = jnp.asarray(temperature, dtype=jnp.result_type(float))
    beta = 1.0 / temperature
    num_spins = size * size

    configurations = ising_all_configurations(size)
    energies = ising_energy(configurations, coupling)
    magnetizations = ising_magnetization(configurations)

    log_weights = -beta * energies
    log_z = logsumexp(log_weights)
    weights = jnp.exp(log_weights - log_z)

    mean_energy = jnp.sum(weights * energies)
    mean_energy_squared = jnp.sum(weights * energies**2)
    variance = mean_energy_squared - mean_energy**2

    return {
        "log_z": log_z,
        "free_energy_per_site": -log_z / (beta * num_spins),
        "energy_per_site": mean_energy / num_spins,
        "abs_magnetization": jnp.sum(weights * jnp.abs(magnetizations)),
        "magnetization_squared": jnp.sum(weights * magnetizations**2),
        "heat_capacity": beta**2 * variance / num_spins,
    }

ising_transfer_matrix_log_z

ising_transfer_matrix_log_z(
    size: int, temperature: Scalar, coupling: Scalar = 1.0
) -> Array

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 L; at most MAX_TRANSFER_SIZE.

required
temperature Scalar

Temperature T.

required
coupling Scalar

Exchange coupling J.

1.0

Returns:

Type Description
Array

A 0-d array holding log Z.

Raises:

Type Description
ValueError

If size exceeds MAX_TRANSFER_SIZE.

Source code in qjax/physics/lattice.py
def ising_transfer_matrix_log_z(
    size: int, temperature: Scalar, coupling: Scalar = 1.0
) -> jax.Array:
    r"""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.

    Args:
        size: Linear lattice size ``L``; at most `MAX_TRANSFER_SIZE`.
        temperature: Temperature ``T``.
        coupling: Exchange coupling ``J``.

    Returns:
        A 0-d array holding ``log Z``.

    Raises:
        ValueError: If ``size`` exceeds `MAX_TRANSFER_SIZE`.
    """
    if size > MAX_TRANSFER_SIZE:
        raise ValueError(
            f"the transfer matrix is {2**size} x {2**size} at L = {size}; "
            f"the limit is L = {MAX_TRANSFER_SIZE}."
        )
    beta = 1.0 / jnp.asarray(temperature, dtype=jnp.result_type(float))
    coupling = jnp.asarray(coupling, dtype=jnp.result_type(float))

    states = jnp.arange(2**size, dtype=jnp.int32)
    bits = (states[:, None] >> jnp.arange(size, dtype=jnp.int32)[None, :]) & 1
    columns = 1.0 - 2.0 * bits.astype(jnp.result_type(float))

    # Bonds inside a column (periodic along it) and between adjacent columns.
    intra = jnp.sum(columns * jnp.roll(columns, 1, axis=-1), axis=-1)
    inter = columns @ columns.T
    log_transfer = beta * coupling * (inter + 0.5 * (intra[:, None] + intra[None, :]))

    shift = jnp.max(log_transfer)
    eigenvalues = jnp.linalg.eigvalsh(jnp.exp(log_transfer - shift))
    largest = jnp.max(jnp.abs(eigenvalues))
    ratios = eigenvalues / largest

    # Tr T^L = e^{L shift} lambda_max^L sum_i (lambda_i / lambda_max)^L. Take the
    # power through |ratio| so an odd L keeps the sign of a negative eigenvalue.
    powered = jnp.abs(ratios) ** size
    if size % 2:
        powered = jnp.sign(ratios) * powered
    return size * (shift + jnp.log(largest)) + jnp.log(jnp.sum(powered))

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.

\[-\beta f = \ln(2 \cosh 2\beta J) + \frac{1}{\pi} \int_0^{\pi/2} \ln\!\Big[\tfrac12\big(1 + \sqrt{1 - \kappa^2 \sin^2\phi}\,\big)\Big] \, d\phi, \qquad \kappa = \frac{2 \sinh 2\beta J}{\cosh^2 2\beta J}.\]

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) T, any shape.

required
coupling Scalar

Exchange coupling J.

1.0
num_quad int

Midpoint quadrature nodes on (0, pi/2).

4096

Returns:

Type Description
Array

Free energy per site, same shape as temperature. Tends to

Array

-T ln 2 as T -> inf and to -2 J as T -> 0.

Source code in qjax/physics/lattice.py
def onsager_free_energy_per_site(
    temperature: Array, coupling: Scalar = 1.0, num_quad: int = 4096
) -> jax.Array:
    r"""Onsager's exact free energy per site in the thermodynamic limit.

    $$-\beta f = \ln(2 \cosh 2\beta J)
      + \frac{1}{\pi} \int_0^{\pi/2}
        \ln\!\Big[\tfrac12\big(1 + \sqrt{1 - \kappa^2 \sin^2\phi}\,\big)\Big]
        \, d\phi,
      \qquad \kappa = \frac{2 \sinh 2\beta J}{\cosh^2 2\beta J}.$$

    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.

    Args:
        temperature: Temperature(s) ``T``, any shape.
        coupling: Exchange coupling ``J``.
        num_quad: Midpoint quadrature nodes on ``(0, pi/2)``.

    Returns:
        Free energy per site, same shape as ``temperature``. Tends to
        ``-T ln 2`` as ``T -> inf`` and to ``-2 J`` as ``T -> 0``.
    """
    temperature = jnp.asarray(temperature, dtype=jnp.result_type(float))
    beta = 1.0 / temperature
    argument = 2.0 * beta * jnp.asarray(coupling, dtype=jnp.result_type(float))

    magnitude = jnp.abs(argument)
    decay = jnp.exp(-2.0 * magnitude)
    log_two_cosh = magnitude + jnp.log1p(decay)
    sech = 2.0 * jnp.exp(-magnitude) / (1.0 + decay)
    kappa = 2.0 * jnp.tanh(argument) * sech

    phi = (jnp.arange(num_quad, dtype=jnp.result_type(float)) + 0.5) * (0.5 * jnp.pi / num_quad)
    radicand = 1.0 - (kappa[..., None] * jnp.sin(phi)) ** 2
    integrand = jnp.log(0.5 * (1.0 + jnp.sqrt(jnp.maximum(radicand, 0.0))))
    # mean * (pi/2) is the integral; dividing by pi leaves mean / 2.
    return -(log_two_cosh + 0.5 * jnp.mean(integrand, axis=-1)) / beta

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) T, any shape.

required
coupling Scalar

Exchange coupling J.

1.0

Returns:

Type Description
Array

Spontaneous magnetization in [0, 1], same shape as temperature.

Source code in qjax/physics/lattice.py
def onsager_magnetization(temperature: Array, coupling: Scalar = 1.0) -> jax.Array:
    r"""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$.

    Args:
        temperature: Temperature(s) ``T``, any shape.
        coupling: Exchange coupling ``J``.

    Returns:
        Spontaneous magnetization in ``[0, 1]``, same shape as ``temperature``.
    """
    temperature = jnp.asarray(temperature, dtype=jnp.result_type(float))
    beta = 1.0 / temperature
    sinh = jnp.sinh(2.0 * beta * jnp.asarray(coupling, dtype=jnp.result_type(float)))
    inner = 1.0 - sinh ** (-4.0)
    # Double-where: the fractional power of a negative base would return NaN and
    # back-propagate NaN even from the unselected branch.
    ordered = inner > 0.0
    return jnp.where(ordered, jnp.where(ordered, inner, 1.0) ** 0.125, 0.0)

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) T, any shape.

required
coupling Scalar

Exchange coupling J.

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 temperature. Equals

Array

-sqrt(2) J at T_c.

Source code in qjax/physics/lattice.py
def onsager_energy_per_site(
    temperature: Array, coupling: Scalar = 1.0, num_quad: int = 4096
) -> jax.Array:
    r"""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.

    Args:
        temperature: Temperature(s) ``T``, any shape.
        coupling: Exchange coupling ``J``.
        num_quad: Quadrature nodes passed through to the free energy.

    Returns:
        Internal energy per site, same shape as ``temperature``. Equals
        ``-sqrt(2) J`` at ``T_c``.
    """
    temperature = jnp.asarray(temperature, dtype=jnp.result_type(float))

    def beta_free_energy(beta: jax.Array) -> jax.Array:
        return beta * onsager_free_energy_per_site(1.0 / beta, coupling, num_quad)

    derivative = jax.grad(beta_free_energy)
    for _ in range(temperature.ndim):
        derivative = jax.vmap(derivative)
    return derivative(1.0 / temperature)

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_cumulant(
    magnetization: Array, axis: int = -1
) -> Array

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 axis reduced.

Source code in qjax/physics/observables.py
def binder_cumulant(magnetization: Array, axis: int = -1) -> jax.Array:
    r"""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).

    Args:
        magnetization: Per-configuration magnetizations.
        axis: Axis to average over.

    Returns:
        The cumulant, with ``axis`` reduced.
    """
    magnetization = jnp.asarray(magnetization, dtype=jnp.result_type(float))
    second = jnp.mean(magnetization**2, axis=axis)
    fourth = jnp.mean(magnetization**4, axis=axis)
    safe = jnp.where(second == 0.0, 1.0, second)
    return jnp.where(second == 0.0, jnp.nan, 1.0 - fourth / (3.0 * safe**2))

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 (T,).

required
curve Array

Values on that grid, shape (T,).

required
level Scalar

The level to cross.

0.5

Returns:

Type Description
Array

A 0-d array with the crossing temperature, or NaN if the curve never

Array

crosses level on the grid.

Source code in qjax/physics/observables.py
def crossing_temperature(temperatures: Array, curve: Array, level: Scalar = 0.5) -> jax.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.

    Args:
        temperatures: Strictly ordered temperature grid, shape ``(T,)``.
        curve: Values on that grid, shape ``(T,)``.
        level: The level to cross.

    Returns:
        A 0-d array with the crossing temperature, or ``NaN`` if the curve never
        crosses ``level`` on the grid.
    """
    x = jnp.asarray(temperatures, dtype=jnp.result_type(float))
    y = jnp.asarray(curve, dtype=jnp.result_type(float))
    above = y >= level
    changes = above[:-1] != above[1:]
    index = jnp.argmax(changes)
    crossed = jnp.any(changes)
    return jnp.where(crossed, _interpolate(x, y, index, index + 1, level), jnp.nan)

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 (T,) with T >= 3.

required
curve Array

Values on that grid, shape (T,).

required

Returns:

Type Description
Array

A 0-d array with the peak temperature.

Source code in qjax/physics/observables.py
def peak_temperature(temperatures: Array, curve: Array) -> jax.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.

    Args:
        temperatures: Ordered temperature grid, shape ``(T,)`` with ``T >= 3``.
        curve: Values on that grid, shape ``(T,)``.

    Returns:
        A 0-d array with the peak temperature.
    """
    x = jnp.asarray(temperatures, dtype=jnp.result_type(float))
    y = jnp.asarray(curve, dtype=jnp.result_type(float))
    centre = jnp.clip(jnp.argmax(y), 1, x.shape[0] - 2)
    x1, x2, x3 = x[centre - 1], x[centre], x[centre + 1]
    y1, y2, y3 = y[centre - 1], y[centre], y[centre + 1]

    d1 = (x1 - x2) * (x1 - x3)
    d2 = (x2 - x1) * (x2 - x3)
    d3 = (x3 - x1) * (x3 - x2)
    quadratic = y1 / d1 + y2 / d2 + y3 / d3
    linear = -(y1 * (x2 + x3) / d1 + y2 * (x1 + x3) / d2 + y3 * (x1 + x2) / d3)

    degenerate = quadratic == 0.0
    safe = jnp.where(degenerate, 1.0, quadratic)
    return jnp.where(degenerate, x2, -0.5 * linear / safe)

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 (T,).

required
curve Array

Values on that grid, shape (T,).

required

Returns:

Type Description
Array

A 0-d array with the width, or NaN if the curve does not fall back

Array

below the half level on both sides of its maximum.

Source code in qjax/physics/observables.py
def half_width(temperatures: Array, curve: Array) -> jax.Array:
    r"""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$.

    Args:
        temperatures: Ordered temperature grid, shape ``(T,)``.
        curve: Values on that grid, shape ``(T,)``.

    Returns:
        A 0-d array with the width, or ``NaN`` if the curve does not fall back
        below the half level on both sides of its maximum.
    """
    x = jnp.asarray(temperatures, dtype=jnp.result_type(float))
    y = jnp.asarray(curve, dtype=jnp.result_type(float))
    size = y.shape[0]
    level = 0.5 * (jnp.max(y) + jnp.min(y))

    peak = jnp.argmax(y)
    index = jnp.arange(size)
    below = y < level
    left = jnp.max(jnp.where(below & (index < peak), index, -1))
    right = jnp.min(jnp.where(below & (index > peak), index, size))

    bracketed = (left >= 0) & (right < size)
    safe_left = jnp.clip(left, 0, size - 2)
    safe_right = jnp.clip(right, 1, size - 1)
    lower = _interpolate(x, y, safe_left, safe_left + 1, level)
    upper = _interpolate(x, y, safe_right - 1, safe_right, level)
    return jnp.where(bracketed, upper - lower, jnp.nan)

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 L, shape (S,).

required
estimates Array

The size-dependent estimates, shape (S,).

required
nu Scalar

Correlation-length exponent used to build the abscissa.

1.0

Returns:

Type Description
Array

(intercept, slope, intercept_stderr). The standard error is the

Array

usual OLS one and is 0 for a perfect fit.

Source code in qjax/physics/observables.py
def finite_size_extrapolation(
    sizes: Array, estimates: Array, nu: Scalar = 1.0
) -> tuple[jax.Array, jax.Array, jax.Array]:
    r"""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$.

    Args:
        sizes: Linear lattice sizes ``L``, shape ``(S,)``.
        estimates: The size-dependent estimates, shape ``(S,)``.
        nu: Correlation-length exponent used to build the abscissa.

    Returns:
        ``(intercept, slope, intercept_stderr)``. The standard error is the
        usual OLS one and is ``0`` for a perfect fit.
    """
    x = jnp.asarray(sizes, dtype=jnp.result_type(float)) ** (-1.0 / nu)
    y = jnp.asarray(estimates, dtype=jnp.result_type(float))
    count = x.shape[0]

    x_mean, y_mean = jnp.mean(x), jnp.mean(y)
    sxx = jnp.sum((x - x_mean) ** 2)
    sxy = jnp.sum((x - x_mean) * (y - y_mean))
    slope = sxy / sxx
    intercept = y_mean - slope * x_mean

    residual = y - (intercept + slope * x)
    dof = max(count - 2, 1)
    variance = jnp.sum(residual**2) / dof
    stderr = jnp.sqrt(variance * (1.0 / count + x_mean**2 / sxx))
    return intercept, slope, stderr

spinglass

The Sherrington-Kirkpatrick spin glass: Hamiltonian and exact thermodynamics.

The SK model is the mean-field spin glass,

\[H(s) = -\tfrac12 \sum_{i \neq j} J_{ij} s_i s_j, \qquad J_{ij} = J_{ji} \sim \mathcal N(0, 1/N),\]

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

sk_couplings(key: Array, num_spins: int) -> Array

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 N.

required

Returns:

Type Description
Array

A symmetric (N, N) matrix with zeros on the diagonal.

Source code in qjax/physics/spinglass.py
def sk_couplings(key: jax.Array, num_spins: int) -> jax.Array:
    """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.

    Args:
        key: PRNG key.
        num_spins: Number of spins ``N``.

    Returns:
        A symmetric ``(N, N)`` matrix with zeros on the diagonal.
    """
    raw = jax.random.normal(key, (num_spins, num_spins)) / jnp.sqrt(num_spins)
    upper = jnp.triu(raw, 1)
    return upper + upper.T

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 (..., N) with entries in {-1, +1}.

required
couplings Array

Symmetric (N, N) coupling matrix with zero diagonal.

required

Returns:

Type Description
Array

Total energy per configuration, shape (...).

Source code in qjax/physics/spinglass.py
def sk_energy(spins: Array, couplings: Array) -> jax.Array:
    r"""SK energy $-\tfrac12 s^{\mathsf T} J s$.

    Args:
        spins: Configuration(s) of shape ``(..., N)`` with entries in
            ``{-1, +1}``.
        couplings: Symmetric ``(N, N)`` coupling matrix with zero diagonal.

    Returns:
        Total energy per configuration, shape ``(...)``.
    """
    spins = jnp.asarray(spins, dtype=jnp.result_type(float))
    couplings = jnp.asarray(couplings, dtype=jnp.result_type(float))
    return -0.5 * jnp.einsum("...i,ij,...j->...", spins, couplings, spins)

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 (N, N) coupling matrix, N <= MAX_ENUMERATED_SPINS.

required
temperature Scalar

Temperature T.

required
chunk int

Configurations per chunk; rounded down to a power of two and capped at 2**N.

4096

Returns:

Type Description
dict[str, Array]

A dict with log_z, free_energy_per_spin, energy_per_spin,

dict[str, Array]

ground_state_energy_per_spin and correlations (an (N, N)

dict[str, Array]

array of \(\langle s_i s_j \rangle\)).

Raises:

Type Description
ValueError

If N exceeds MAX_ENUMERATED_SPINS.

Source code in qjax/physics/spinglass.py
def sk_exact_observables(
    couplings: Array, temperature: Scalar, chunk: int = 4096
) -> dict[str, jax.Array]:
    r"""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.

    Args:
        couplings: Symmetric ``(N, N)`` coupling matrix, ``N <=``
            `MAX_ENUMERATED_SPINS`.
        temperature: Temperature ``T``.
        chunk: Configurations per chunk; rounded down to a power of two and
            capped at ``2**N``.

    Returns:
        A dict with ``log_z``, ``free_energy_per_spin``, ``energy_per_spin``,
        ``ground_state_energy_per_spin`` and ``correlations`` (an ``(N, N)``
        array of $\langle s_i s_j \rangle$).

    Raises:
        ValueError: If ``N`` exceeds `MAX_ENUMERATED_SPINS`.
    """
    couplings = jnp.asarray(couplings, dtype=jnp.result_type(float))
    num_spins = couplings.shape[-1]
    if num_spins > MAX_ENUMERATED_SPINS:
        raise ValueError(
            f"enumerating {num_spins} spins needs 2**{num_spins} states; "
            f"the limit is {MAX_ENUMERATED_SPINS}."
        )
    total = 2**num_spins
    chunk = _largest_power_of_two(chunk, total)
    beta = 1.0 / jnp.asarray(temperature, dtype=jnp.result_type(float))

    Carry = tuple[jax.Array, jax.Array, jax.Array, jax.Array, jax.Array]

    def step(carry: Carry, offset: jax.Array) -> tuple[Carry, None]:
        peak, mass, energy_mass, correlation_mass, minimum = carry
        spins = _chunk_configurations(offset, chunk, num_spins)
        energies = sk_energy(spins, couplings)
        log_weights = -beta * energies

        new_peak = jnp.maximum(peak, jnp.max(log_weights))
        rescale = jnp.exp(peak - new_peak)
        weights = jnp.exp(log_weights - new_peak)
        return (
            new_peak,
            mass * rescale + jnp.sum(weights),
            energy_mass * rescale + jnp.sum(weights * energies),
            correlation_mass * rescale + jnp.einsum("bi,bj,b->ij", spins, spins, weights),
            jnp.minimum(minimum, jnp.min(energies)),
        ), None

    zero = jnp.zeros((), dtype=jnp.result_type(float))
    initial: Carry = (
        jnp.full((), -jnp.inf, dtype=jnp.result_type(float)),
        zero,
        zero,
        jnp.zeros((num_spins, num_spins), dtype=jnp.result_type(float)),
        jnp.full((), jnp.inf, dtype=jnp.result_type(float)),
    )
    offsets = jnp.arange(0, total, chunk, dtype=jnp.int32)
    (peak, mass, energy_mass, correlation_mass, minimum), _ = jax.lax.scan(step, initial, offsets)

    log_z = peak + jnp.log(mass)
    return {
        "log_z": log_z,
        "free_energy_per_spin": -log_z / (beta * num_spins),
        "energy_per_spin": energy_mass / mass / num_spins,
        "ground_state_energy_per_spin": minimum / num_spins,
        "correlations": correlation_mass / mass,
    }

sk_exact_correlations

sk_exact_correlations(
    couplings: Array, temperature: Scalar, chunk: int = 4096
) -> Array

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 (N, N) coupling matrix.

required
temperature Scalar

Temperature T.

required
chunk int

Configurations per enumeration chunk.

4096

Returns:

Type Description
Array

An (N, N) array of two-point correlations.

Source code in qjax/physics/spinglass.py
def sk_exact_correlations(couplings: Array, temperature: Scalar, chunk: int = 4096) -> jax.Array:
    r"""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.

    Args:
        couplings: Symmetric ``(N, N)`` coupling matrix.
        temperature: Temperature ``T``.
        chunk: Configurations per enumeration chunk.

    Returns:
        An ``(N, N)`` array of two-point correlations.
    """
    return sk_exact_observables(couplings, temperature, chunk)["correlations"]

clusters

Lennard-Jones atomic clusters: potential, local quenching, and geometry.

The pair potential

\[V(r) = 4\epsilon\Big[(\sigma/r)^{12} - (\sigma/r)^{6}\Big]\]

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 (..., n, 3).

required
epsilon Scalar

Well depth.

1.0
sigma Scalar

Length scale; the pair minimum sits at 2**(1/6) * sigma.

1.0
softening float

Floor on the squared separation.

1e-12

Returns:

Type Description
Array

Total energy per cluster, shape (...). Equals -1, -3, -6

Array

for a regular simplex of 2, 3, 4 atoms at r = 2**(1/6) sigma.

Source code in qjax/physics/clusters.py
def lj_energy(
    positions: Array,
    epsilon: Scalar = 1.0,
    sigma: Scalar = 1.0,
    softening: float = 1e-12,
) -> jax.Array:
    r"""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.

    Args:
        positions: Atomic coordinates of shape ``(..., n, 3)``.
        epsilon: Well depth.
        sigma: Length scale; the pair minimum sits at ``2**(1/6) * sigma``.
        softening: Floor on the squared separation.

    Returns:
        Total energy per cluster, shape ``(...)``. Equals ``-1``, ``-3``, ``-6``
        for a regular simplex of 2, 3, 4 atoms at ``r = 2**(1/6) sigma``.
    """
    positions = jnp.asarray(positions, dtype=jnp.result_type(float))
    epsilon = jnp.asarray(epsilon, dtype=jnp.result_type(float))
    sigma = jnp.asarray(sigma, dtype=jnp.result_type(float))

    offsets = positions[..., :, None, :] - positions[..., None, :, :]
    squared = jnp.sum(offsets**2, axis=-1)
    off_diagonal = ~jnp.eye(positions.shape[-2], dtype=bool)

    safe = jnp.where(off_diagonal, jnp.maximum(squared, softening), 1.0)
    inverse_sixth = (sigma**2 / safe) ** 3
    pair = 4.0 * epsilon * (inverse_sixth**2 - inverse_sixth)
    return 0.5 * jnp.sum(jnp.where(off_diagonal, pair, 0.0), axis=(-2, -1))

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 (..., n, 3).

required
container_radius Scalar

Radius R beyond which the wall acts.

required
stiffness Scalar

Wall stiffness k.

10.0
epsilon Scalar

Well depth.

1.0
sigma Scalar

Length scale.

1.0

Returns:

Type Description
Array

Confined energy per cluster, shape (...). Differentiable everywhere,

Array

including at the origin.

Source code in qjax/physics/clusters.py
def lj_energy_confined(
    positions: Array,
    container_radius: Scalar,
    stiffness: Scalar = 10.0,
    epsilon: Scalar = 1.0,
    sigma: Scalar = 1.0,
) -> jax.Array:
    r"""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.

    Args:
        positions: Atomic coordinates of shape ``(..., n, 3)``.
        container_radius: Radius ``R`` beyond which the wall acts.
        stiffness: Wall stiffness ``k``.
        epsilon: Well depth.
        sigma: Length scale.

    Returns:
        Confined energy per cluster, shape ``(...)``. Differentiable everywhere,
        including at the origin.
    """
    positions = jnp.asarray(positions, dtype=jnp.result_type(float))
    # ``sqrt`` has an infinite derivative at zero, and the ``maximum`` below then
    # multiplies it by zero -- so an atom sitting exactly at the origin (the
    # centre of a re-centred icosahedron, say) would give a finite energy and a
    # NaN gradient. Sanitize the radicand before the root, as ``lj_energy`` does.
    squared_radius = jnp.sum(positions**2, axis=-1)
    inside = squared_radius > 0.0
    radius = jnp.where(inside, jnp.sqrt(jnp.where(inside, squared_radius, 1.0)), 0.0)
    excess = jnp.maximum(radius - jnp.asarray(container_radius, dtype=radius.dtype), 0.0)
    wall = jnp.asarray(stiffness, dtype=radius.dtype) * jnp.sum(excess**2, axis=-1)
    return lj_energy(positions, epsilon, sigma) + wall

lj_random_cluster

lj_random_cluster(
    key: Array, num_atoms: int, radius: Scalar
) -> Array

Draw atomic positions uniformly inside a ball.

Parameters:

Name Type Description Default
key Array

PRNG key.

required
num_atoms int

Number of atoms n.

required
radius Scalar

Ball radius.

required

Returns:

Type Description
Array

Positions of shape (n, 3).

Source code in qjax/physics/clusters.py
def lj_random_cluster(key: jax.Array, num_atoms: int, radius: Scalar) -> jax.Array:
    """Draw atomic positions uniformly inside a ball.

    Args:
        key: PRNG key.
        num_atoms: Number of atoms ``n``.
        radius: Ball radius.

    Returns:
        Positions of shape ``(n, 3)``.
    """
    direction_key, magnitude_key = jax.random.split(key)
    direction = jax.random.normal(direction_key, (num_atoms, 3))
    direction /= jnp.linalg.norm(direction, axis=-1, keepdims=True)
    # u**(1/3) makes the radial law uniform in volume.
    magnitude = jax.random.uniform(magnitude_key, (num_atoms, 1)) ** (1.0 / 3.0)
    return direction * magnitude * jnp.asarray(radius, dtype=jnp.result_type(float))

equidistant_cluster

equidistant_cluster(
    num_atoms: int, distance: Scalar
) -> Array

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 (num_atoms, 3).

Raises:

Type Description
ValueError

If num_atoms is outside 1 to 4 (a regular simplex with more vertices does not embed in three dimensions).

Source code in qjax/physics/clusters.py
def equidistant_cluster(num_atoms: int, distance: Scalar) -> jax.Array:
    """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.

    Args:
        num_atoms: Number of atoms, 1 to 4.
        distance: Edge length.

    Returns:
        Positions of shape ``(num_atoms, 3)``.

    Raises:
        ValueError: If ``num_atoms`` is outside 1 to 4 (a regular simplex with
            more vertices does not embed in three dimensions).
    """
    if not 1 <= num_atoms <= 4:
        raise ValueError(f"a regular simplex in 3-D has at most 4 vertices; got {num_atoms}.")
    third = jnp.sqrt(jnp.asarray(3.0, dtype=jnp.result_type(float)))
    vertices = jnp.array(
        [
            [0.0, 0.0, 0.0],
            [1.0, 0.0, 0.0],
            [0.5, float(third / 2.0), 0.0],
            [0.5, float(third / 6.0), float(jnp.sqrt(jnp.asarray(2.0 / 3.0)))],
        ],
        dtype=jnp.result_type(float),
    )
    return vertices[:num_atoms] * jnp.asarray(distance, dtype=jnp.result_type(float))

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 (n, 3).

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

(final_positions, energies) where energies has shape

Array

(steps + 1,) and holds the energy before each step and after the

tuple[Array, Array]

last.

Source code in qjax/physics/clusters.py
def lj_quench(
    positions: Array,
    steps: int = 40,
    learning_rate: Scalar = 2e-3,
    epsilon: Scalar = 1.0,
    sigma: Scalar = 1.0,
) -> tuple[jax.Array, jax.Array]:
    r"""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.

    Args:
        positions: Starting coordinates of shape ``(n, 3)``.
        steps: Number of Adam steps.
        learning_rate: Adam step size.
        epsilon: Well depth.
        sigma: Length scale.

    Returns:
        ``(final_positions, energies)`` where ``energies`` has shape
        ``(steps + 1,)`` and holds the energy before each step and after the
        last.
    """
    positions = jnp.asarray(positions, dtype=jnp.result_type(float))
    gradient = jax.grad(lambda x: lj_energy(x, epsilon, sigma))
    beta1, beta2, eps = 0.9, 0.999, 1e-8

    Carry = tuple[jax.Array, jax.Array, jax.Array]

    def step(carry: Carry, count: jax.Array) -> tuple[Carry, jax.Array]:
        state, first, second = carry
        energy = lj_energy(state, epsilon, sigma)
        grads = gradient(state)
        first = beta1 * first + (1.0 - beta1) * grads
        second = beta2 * second + (1.0 - beta2) * grads**2
        scale = count + 1.0
        first_hat = first / (1.0 - beta1**scale)
        second_hat = second / (1.0 - beta2**scale)
        state = state - learning_rate * first_hat / (jnp.sqrt(second_hat) + eps)
        return (state, first, second), energy

    initial: Carry = (positions, jnp.zeros_like(positions), jnp.zeros_like(positions))
    (final, _, _), trace = jax.lax.scan(
        step, initial, jnp.arange(steps, dtype=jnp.result_type(float))
    )
    energies = jnp.concatenate([trace, lj_energy(final, epsilon, sigma)[None]])
    return final, energies

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 (..., n, 3).

required
cutoff Scalar

Neighbour cutoff, in the same units as positions. The default sits between the first and second neighbour shells of a close-packed cluster at sigma = 1.

1.35

Returns:

Type Description
Array

Integer neighbour counts of shape (..., n).

Source code in qjax/physics/clusters.py
def coordination_numbers(positions: Array, cutoff: Scalar = 1.35) -> jax.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.

    Args:
        positions: Atomic coordinates of shape ``(..., n, 3)``.
        cutoff: Neighbour cutoff, in the same units as ``positions``. The
            default sits between the first and second neighbour shells of a
            close-packed cluster at ``sigma = 1``.

    Returns:
        Integer neighbour counts of shape ``(..., n)``.
    """
    positions = jnp.asarray(positions, dtype=jnp.result_type(float))
    offsets = positions[..., :, None, :] - positions[..., None, :, :]
    squared = jnp.sum(offsets**2, axis=-1)
    off_diagonal = ~jnp.eye(positions.shape[-2], dtype=bool)
    within = (squared < jnp.asarray(cutoff, dtype=squared.dtype) ** 2) & off_diagonal
    return jnp.sum(within, axis=-1)

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

\[T_q(t) = T_q(1)\,\frac{2^{q-1} - 1}{(1 + t)^{q-1} - 1}, \qquad t \ge 1.\]

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

\[\ln_{2-q}(x) = \frac{x^{q-1} - 1}{q - 1},\]

both \((q-1)\) factors cancel and the whole schedule is a ratio of two qjax.q_log calls:

\[T_q(t) = T_q(1)\,\frac{\ln_{2-q} 2}{\ln_{2-q}(1 + t)}.\]

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 t, counted from 1; any shape. At t = 1 the schedule returns initial for every q. At t = 0 the denominator vanishes and the result is +inf, which is the correct limit of the schedule rather than an error.

required
initial Scalar

Temperature at t = 1.

required
q Scalar

Cooling index. q = 1 is the Geman-Geman logarithmic schedule, q = 2 the Cauchy schedule initial / t, and q > 2 cools faster still.

required

Returns:

Type Description
Array

The temperature at each step, same shape as step.

Source code in qjax/physics/annealing.py
def tsallis_schedule(step: Array, initial: Scalar, q: Scalar) -> jax.Array:
    r"""The Tsallis cooling law $T_q(1)\,\ln_{2-q} 2 / \ln_{2-q}(1+t)$.

    Args:
        step: Annealing step ``t``, counted from ``1``; any shape. At ``t = 1``
            the schedule returns ``initial`` for every ``q``. At ``t = 0`` the
            denominator vanishes and the result is ``+inf``, which is the
            correct limit of the schedule rather than an error.
        initial: Temperature at ``t = 1``.
        q: Cooling index. ``q = 1`` is the Geman-Geman logarithmic schedule,
            ``q = 2`` the Cauchy schedule ``initial / t``, and ``q > 2`` cools
            faster still.

    Returns:
        The temperature at each ``step``, same shape as ``step``.
    """
    step = jnp.asarray(step, dtype=jnp.result_type(float))
    index = 2.0 - jnp.asarray(q, dtype=jnp.result_type(float))
    return jnp.asarray(initial, dtype=jnp.result_type(float)) * (
        q_log(2.0, index) / q_log(1.0 + step, index)
    )

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 t, counted from 1.

required
initial Scalar

Visiting temperature at t = 1.

required
q_visit Scalar

Visiting index q_V. 1 recovers classical (Boltzmann) annealing, 2 the Cauchy machine.

required

Returns:

Type Description
Array

The visiting temperature at each step.

Source code in qjax/physics/annealing.py
def visiting_temperature(step: Array, initial: Scalar, q_visit: Scalar) -> jax.Array:
    r"""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.

    Args:
        step: Annealing step ``t``, counted from ``1``.
        initial: Visiting temperature at ``t = 1``.
        q_visit: Visiting index ``q_V``. ``1`` recovers classical (Boltzmann)
            annealing, ``2`` the Cauchy machine.

    Returns:
        The visiting temperature at each ``step``.
    """
    return tsallis_schedule(step, initial, q_visit)

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 t, counted from 1.

required
initial Scalar

Acceptance temperature at t = 1.

required
q_accept Scalar

Cooling index for the acceptance temperature.

required

Returns:

Type Description
Array

The acceptance temperature at each step.

Source code in qjax/physics/annealing.py
def acceptance_temperature(step: Array, initial: Scalar, q_accept: Scalar) -> jax.Array:
    r"""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.

    Args:
        step: Annealing step ``t``, counted from ``1``.
        initial: Acceptance temperature at ``t = 1``.
        q_accept: Cooling index for the acceptance temperature.

    Returns:
        The acceptance temperature at each ``step``.
    """
    return tsallis_schedule(step, initial, q_accept)

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

\[\frac{\partial p}{\partial t} = D\,\frac{\partial^2 p^{\,2-q}}{\partial x^2}\]

has the self-similar \(q\)-Gaussian solution of Tsallis & Bukman (1996), whose width obeys \(\langle x^2 \rangle \propto t^{\alpha}\) with

\[\alpha = \frac{2}{3 - q}.\]

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,

\[dp = -\frac{\alpha\,p}{1 + (p/p_c)^2}\,dt + \sqrt{2 D_0}\,dW,\]

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

\[q = 1 + \frac{2 D_0}{\alpha\,p_c^2}, \qquad \beta = \frac{\alpha}{2 D_0}.\]

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, q < 3.

required

Returns:

Type Description
Array

The exponent alpha in <x^2> ~ t**alpha.

Source code in qjax/physics/diffusion.py
def nlfp_exponent(q: Scalar) -> jax.Array:
    r"""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.

    Args:
        q: Entropic index, ``q < 3``.

    Returns:
        The exponent ``alpha`` in ``<x^2> ~ t**alpha``.
    """
    return 2.0 / (3.0 - jnp.asarray(q, dtype=jnp.result_type(float)))

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 alpha.

required

Returns:

Type Description
Array

The entropic index implied by alpha.

Source code in qjax/physics/diffusion.py
def nlfp_index(exponent: Scalar) -> jax.Array:
    r"""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``.

    Args:
        exponent: Measured anomalous exponent ``alpha``.

    Returns:
        The entropic index implied by ``alpha``.
    """
    exponent = jnp.asarray(exponent, dtype=jnp.result_type(float))
    return 3.0 - 2.0 / exponent

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 reference_time.

required
reference_time Scalar

Time at which beta_initial is quoted.

1.0

Returns:

Type Description
Array

beta(t), same shape as time.

Source code in qjax/physics/diffusion.py
def nlfp_scaling_beta(
    time: Array, q: Scalar, beta_initial: Scalar, reference_time: Scalar = 1.0
) -> jax.Array:
    r"""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.

    Args:
        time: Time(s) at which to evaluate the width, any shape.
        q: Entropic index.
        beta_initial: Width parameter at ``reference_time``.
        reference_time: Time at which ``beta_initial`` is quoted.

    Returns:
        ``beta(t)``, same shape as ``time``.
    """
    time = jnp.asarray(time, dtype=jnp.result_type(float))
    scale = time / jnp.asarray(reference_time, dtype=jnp.result_type(float))
    return jnp.asarray(beta_initial, dtype=jnp.result_type(float)) * scale ** (-nlfp_exponent(q))

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

\[K = \frac{4 D (2-q)}{C_q^{\,1-q}},\]

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 NLFP_MAX_INDEX.

required
diffusivity Scalar

The coefficient D in the equation.

required

Returns:

Type Description
Array

The scalar rate constant K, with the dimensions of D.

Raises:

Type Description
ValueError

If q is a concrete value at or above NLFP_MAX_INDEX, where the equation is no longer of porous-medium type.

Source code in qjax/physics/diffusion.py
def nlfp_rate(q: Scalar, diffusivity: Scalar) -> jax.Array:
    r"""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

    $$K = \frac{4 D (2-q)}{C_q^{\,1-q}},$$

    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.

    Args:
        q: Entropic index, strictly below `NLFP_MAX_INDEX`.
        diffusivity: The coefficient ``D`` in the equation.

    Returns:
        The scalar rate constant ``K``, with the dimensions of ``D``.

    Raises:
        ValueError: If ``q`` is a concrete value at or above
            `NLFP_MAX_INDEX`, where the equation is no longer of
            porous-medium type.
    """
    _check_index(q)
    q = jnp.asarray(q, dtype=jnp.result_type(float))
    diffusivity = jnp.asarray(diffusivity, dtype=jnp.result_type(float))
    return 4.0 * diffusivity * (2.0 - q) / normalization(q) ** (1.0 - q)

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 NLFP_MAX_INDEX.

required
diffusivity Scalar

The coefficient D.

required
beta_initial Scalar

Width parameter at t = 0.

required

Returns:

Type Description
Array

The scalar offset t_star.

Source code in qjax/physics/diffusion.py
def nlfp_offset(q: Scalar, diffusivity: Scalar, beta_initial: Scalar) -> jax.Array:
    r"""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.

    Args:
        q: Entropic index, strictly below `NLFP_MAX_INDEX`.
        diffusivity: The coefficient ``D``.
        beta_initial: Width parameter at ``t = 0``.

    Returns:
        The scalar offset ``t_star``.
    """
    q = jnp.asarray(q, dtype=jnp.result_type(float))
    beta_initial = jnp.asarray(beta_initial, dtype=jnp.result_type(float))
    rate = nlfp_rate(q, diffusivity)
    return 2.0 / ((3.0 - q) * rate) * beta_initial ** (-(3.0 - q) / 2.0)

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

\[\beta(t) = \Big[\tfrac{3-q}{2}\,K\,(t - t_0 + t_\star)\Big]^{-\frac{2}{3-q}},\]

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 NLFP_MAX_INDEX.

required
diffusivity Scalar

The coefficient D.

required
beta_initial Scalar

Width parameter at initial_time.

required
initial_time Scalar

Time at which beta_initial is quoted.

0.0

Returns:

Type Description
Array

beta(t), same shape as time.

Source code in qjax/physics/diffusion.py
def nlfp_width(
    time: Array,
    q: Scalar,
    diffusivity: Scalar,
    beta_initial: Scalar,
    initial_time: Scalar = 0.0,
) -> jax.Array:
    r"""Exact width parameter $\beta(t)$ of the Tsallis-Bukman solution.

    Solving $\dot\beta = -K\beta^{(5-q)/2}$ from `nlfp_rate` gives

    $$\beta(t) = \Big[\tfrac{3-q}{2}\,K\,(t - t_0 + t_\star)\Big]^{-\frac{2}{3-q}},$$

    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.

    Args:
        time: Time(s) at which to evaluate the width, any shape.
        q: Entropic index, strictly below `NLFP_MAX_INDEX`.
        diffusivity: The coefficient ``D``.
        beta_initial: Width parameter at ``initial_time``.
        initial_time: Time at which ``beta_initial`` is quoted.

    Returns:
        ``beta(t)``, same shape as ``time``.
    """
    time = jnp.asarray(time, dtype=jnp.result_type(float))
    q = jnp.asarray(q, dtype=jnp.result_type(float))
    rate = nlfp_rate(q, diffusivity)
    offset = nlfp_offset(q, diffusivity, beta_initial)
    elapsed = time - jnp.asarray(initial_time, dtype=jnp.result_type(float)) + offset
    return (0.5 * (3.0 - q) * rate * elapsed) ** (-2.0 / (3.0 - q))

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:

\[p(x, t) = \frac{\sqrt{\beta(t)}}{C_q}\,\exp_q\!\big(-\beta(t)\,x^2\big).\]

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 time.

required
time Array

Time(s), broadcast against x.

required
q Scalar

Entropic index, strictly below NLFP_MAX_INDEX.

required
diffusivity Scalar

The coefficient D.

required
beta_initial Scalar

Width parameter at initial_time.

required
initial_time Scalar

Time at which beta_initial is quoted.

0.0

Returns:

Type Description
Array

The density, broadcast over x and time. Exactly zero beyond

Array

nlfp_front when q < 1.

Source code in qjax/physics/diffusion.py
def nlfp_density(
    x: Array,
    time: Array,
    q: Scalar,
    diffusivity: Scalar,
    beta_initial: Scalar,
    initial_time: Scalar = 0.0,
) -> jax.Array:
    r"""The exact solution of the nonlinear Fokker-Planck equation.

    A normalized $q$-Gaussian whose width follows `nlfp_width`:

    $$p(x, t) = \frac{\sqrt{\beta(t)}}{C_q}\,\exp_q\!\big(-\beta(t)\,x^2\big).$$

    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.

    Args:
        x: Position(s), broadcast against ``time``.
        time: Time(s), broadcast against ``x``.
        q: Entropic index, strictly below `NLFP_MAX_INDEX`.
        diffusivity: The coefficient ``D``.
        beta_initial: Width parameter at ``initial_time``.
        initial_time: Time at which ``beta_initial`` is quoted.

    Returns:
        The density, broadcast over ``x`` and ``time``. Exactly zero beyond
        `nlfp_front` when ``q < 1``.
    """
    width = nlfp_width(time, q, diffusivity, beta_initial, initial_time)
    return q_gaussian_pdf(x, q, width)

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 NLFP_MAX_INDEX.

required
diffusivity Scalar

The coefficient D.

required
beta_initial Scalar

Width parameter at initial_time.

required
initial_time Scalar

Time at which beta_initial is quoted.

0.0

Returns:

Type Description
Array

The front position, same shape as time; +inf for q >= 1,

Array

where the support is the whole line.

Source code in qjax/physics/diffusion.py
def nlfp_front(
    time: Array,
    q: Scalar,
    diffusivity: Scalar,
    beta_initial: Scalar,
    initial_time: Scalar = 0.0,
) -> jax.Array:
    r"""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.

    Args:
        time: Time(s) at which to locate the front, any shape.
        q: Entropic index, strictly below `NLFP_MAX_INDEX`.
        diffusivity: The coefficient ``D``.
        beta_initial: Width parameter at ``initial_time``.
        initial_time: Time at which ``beta_initial`` is quoted.

    Returns:
        The front position, same shape as ``time``; ``+inf`` for ``q >= 1``,
        where the support is the whole line.
    """
    q = jnp.asarray(q, dtype=jnp.result_type(float))
    width = nlfp_width(time, q, diffusivity, beta_initial, initial_time)
    compact = q < 1.0
    # Sanitize before the reciprocal square root so the unselected branch
    # contributes neither a NaN value nor a NaN gradient.
    safe = jnp.where(compact, (1.0 - q) * width, 1.0)
    return jnp.where(compact, 1.0 / jnp.sqrt(safe), jnp.inf)

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 (x, t) -> p on scalars.

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 NLFP_MAX_INDEX.

required
diffusivity Scalar

The coefficient D.

required

Returns:

Type Description
Array

The scalar residual. Zero for an exact solution.

Raises:

Type Description
ValueError

If q is a concrete value at or above NLFP_MAX_INDEX.

Source code in qjax/physics/diffusion.py
def nlfp_residual(
    density_fn: Callable[[jax.Array, jax.Array], jax.Array],
    x: Scalar,
    time: Scalar,
    q: Scalar,
    diffusivity: Scalar,
) -> jax.Array:
    r"""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.

    Args:
        density_fn: A callable ``(x, t) -> p`` on scalars.
        x: Position at which to evaluate the residual.
        time: Time at which to evaluate the residual.
        q: Entropic index, strictly below `NLFP_MAX_INDEX`.
        diffusivity: The coefficient ``D``.

    Returns:
        The scalar residual. Zero for an exact solution.

    Raises:
        ValueError: If ``q`` is a concrete value at or above `NLFP_MAX_INDEX`.
    """
    _check_index(q)
    x = jnp.asarray(x, dtype=jnp.result_type(float))
    time = jnp.asarray(time, dtype=jnp.result_type(float))
    exponent = 2.0 - jnp.asarray(q, dtype=jnp.result_type(float))

    def pressure(position: jax.Array, instant: jax.Array) -> jax.Array:
        density = density_fn(position, instant)
        # Sanitize before the fractional power: at a q < 1 front the density is
        # exactly zero, and ``0 ** exponent`` would back-propagate NaN even
        # though the value itself is fine.
        positive = density > 0.0
        return jnp.where(positive, jnp.where(positive, density, 1.0) ** exponent, 0.0)

    time_derivative = jax.grad(density_fn, argnums=1)(x, time)
    curvature = jax.grad(jax.grad(pressure, argnums=0), argnums=0)(x, time)
    return time_derivative - jnp.asarray(diffusivity, dtype=jnp.result_type(float)) * curvature

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 D_0.

required
friction Scalar

Friction coefficient alpha at small momentum.

required
momentum_scale Scalar

Saturation momentum p_c.

required

Returns:

Type Description
Array

The exact stationary entropic index.

Source code in qjax/physics/diffusion.py
def saturating_langevin_q(diffusion: Scalar, friction: Scalar, momentum_scale: Scalar) -> jax.Array:
    r"""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.

    Args:
        diffusion: Momentum diffusion coefficient ``D_0``.
        friction: Friction coefficient ``alpha`` at small momentum.
        momentum_scale: Saturation momentum ``p_c``.

    Returns:
        The exact stationary entropic index.
    """
    diffusion = jnp.asarray(diffusion, dtype=jnp.result_type(float))
    friction = jnp.asarray(friction, dtype=jnp.result_type(float))
    momentum_scale = jnp.asarray(momentum_scale, dtype=jnp.result_type(float))
    return 1.0 + 2.0 * diffusion / (friction * momentum_scale**2)

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 D_0.

required
friction Scalar

Friction coefficient alpha at small momentum.

required

Returns:

Type Description
Array

The exact stationary width parameter.

Source code in qjax/physics/diffusion.py
def saturating_langevin_beta(diffusion: Scalar, friction: Scalar) -> jax.Array:
    r"""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.

    Args:
        diffusion: Momentum diffusion coefficient ``D_0``.
        friction: Friction coefficient ``alpha`` at small momentum.

    Returns:
        The exact stationary width parameter.
    """
    friction = jnp.asarray(friction, dtype=jnp.result_type(float))
    diffusion = jnp.asarray(diffusion, dtype=jnp.result_type(float))
    return friction / (2.0 * diffusion)

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 E_R / U_0, any shape.

required

Returns:

Type Description
Array

The predicted entropic index, same shape as the input.

Source code in qjax/physics/diffusion.py
def lutz_q(recoil_over_depth: Array) -> jax.Array:
    r"""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.

    Args:
        recoil_over_depth: The ratio ``E_R / U_0``, any shape.

    Returns:
        The predicted entropic index, same shape as the input.
    """
    return 1.0 + 44.0 * jnp.asarray(recoil_over_depth, dtype=jnp.result_type(float))

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 (T, P) for one dimension or (T, P, D) for D dimensions, where P is the number of independent walkers. A single walker may be passed as (T,).

required
origin Array | None

Starting positions, shape matching one snapshot. Defaults to snapshots[0].

None

Returns:

Type Description
Array

The mean-squared displacement at each snapshot time, shape (T,).

Source code in qjax/physics/diffusion.py
def mean_squared_displacement(snapshots: Array, origin: Array | None = None) -> jax.Array:
    """Ensemble mean-squared displacement from a set of snapshots.

    Args:
        snapshots: Positions of shape ``(T, P)`` for one dimension or
            ``(T, P, D)`` for ``D`` dimensions, where ``P`` is the number of
            independent walkers. A single walker may be passed as ``(T,)``.
        origin: Starting positions, shape matching one snapshot. Defaults to
            ``snapshots[0]``.

    Returns:
        The mean-squared displacement at each snapshot time, shape ``(T,)``.
    """
    snapshots = jnp.asarray(snapshots, dtype=jnp.result_type(float))
    if snapshots.ndim == 1:
        snapshots = snapshots[:, None]
    start = snapshots[0] if origin is None else jnp.asarray(origin, dtype=snapshots.dtype)
    offsets = snapshots - start
    squared = offsets**2 if snapshots.ndim == 2 else jnp.sum(offsets**2, axis=-1)
    return jnp.mean(squared, axis=-1)

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 (n,).

required
y Array

Ordinate, strictly positive, shape (n,).

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 n.

None

Returns:

Type Description
Array

(exponent, prefactor, exponent_stderr). The standard error is the

Array

usual OLS one and is 0 for a perfect power law.

Source code in qjax/physics/diffusion.py
def fit_power_law(
    x: Array, y: Array, low: int = 0, high: int | None = None
) -> tuple[jax.Array, jax.Array, jax.Array]:
    r"""Least-squares fit of $y = c\,x^{a}$ on a log-log scale.

    Args:
        x: Abscissa, strictly positive, shape ``(n,)``.
        y: Ordinate, strictly positive, shape ``(n,)``.
        low: First index to include; use it to drop an early transient.
        high: One past the last index to include. Defaults to ``n``.

    Returns:
        ``(exponent, prefactor, exponent_stderr)``. The standard error is the
        usual OLS one and is ``0`` for a perfect power law.
    """
    log_x = jnp.log(jnp.asarray(x, dtype=jnp.result_type(float))[low:high])
    log_y = jnp.log(jnp.asarray(y, dtype=jnp.result_type(float))[low:high])
    count = log_x.shape[0]

    x_mean, y_mean = jnp.mean(log_x), jnp.mean(log_y)
    sxx = jnp.sum((log_x - x_mean) ** 2)
    exponent = jnp.sum((log_x - x_mean) * (log_y - y_mean)) / sxx
    intercept = y_mean - exponent * x_mean

    residual = log_y - (intercept + exponent * log_x)
    dof = max(count - 2, 1)
    stderr = jnp.sqrt(jnp.sum(residual**2) / dof / sxx)
    return exponent, jnp.exp(intercept), stderr

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 (B + 1,).

required

Returns:

Type Description
Array

Density in each bin, shape (B,), integrating to 1 against the

Array

bin widths (or all zeros if no sample falls inside).

Source code in qjax/physics/diffusion.py
def histogram_density(samples: Array, edges: Array) -> jax.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.

    Args:
        samples: Values to bin, any shape (flattened).
        edges: Monotone bin edges, shape ``(B + 1,)``.

    Returns:
        Density in each bin, shape ``(B,)``, integrating to ``1`` against the
        bin widths (or all zeros if no sample falls inside).
    """
    samples = jnp.asarray(samples, dtype=jnp.result_type(float)).reshape(-1)
    edges = jnp.asarray(edges, dtype=jnp.result_type(float))
    bins = edges.shape[0] - 1
    widths = jnp.diff(edges)

    index = jnp.searchsorted(edges, samples, side="right") - 1
    inside = (index >= 0) & (index < bins)
    counts = (
        jnp.zeros(bins, dtype=samples.dtype)
        .at[jnp.where(inside, index, 0)]
        .add(jnp.where(inside, 1.0, 0.0))
    )
    total = jnp.sum(counts)
    return jnp.where(total > 0.0, counts / jnp.where(total > 0.0, total, 1.0) / widths, 0.0)

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 (B + 1,).

required
density Array

Density per bin, shape (B,).

required

Returns:

Type Description
Array

Interpolated density, same shape as x; exactly 0 outside

Array

[edges[0], edges[-1]].

Source code in qjax/physics/diffusion.py
def interpolate_density(x: Array, edges: Array, density: Array) -> jax.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.

    Args:
        x: Points at which to evaluate, any shape.
        edges: Bin edges the density was built on, shape ``(B + 1,)``.
        density: Density per bin, shape ``(B,)``.

    Returns:
        Interpolated density, same shape as ``x``; exactly ``0`` outside
        ``[edges[0], edges[-1]]``.
    """
    edges = jnp.asarray(edges, dtype=jnp.result_type(float))
    density = jnp.asarray(density, dtype=jnp.result_type(float))
    centres = 0.5 * (edges[:-1] + edges[1:])
    x = jnp.asarray(x, dtype=jnp.result_type(float))
    # ``left``/``right`` hold the end bins flat across their outer half.
    inside = (x >= edges[0]) & (x <= edges[-1])
    interpolated = jnp.interp(x, centres, density, left=density[0], right=density[-1])
    return jnp.where(inside, interpolated, 0.0)

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 jax.Array of floating dtype.

Source code in qjax/shared/validation.py
def as_scalar_q(q: Scalar) -> jax.Array:
    """Coerce an entropic index to a floating-point JAX scalar.

    Args:
        q: The entropic index, as a Python number or array-like.

    Returns:
        A 0-d `jax.Array` of floating dtype.
    """
    return jnp.asarray(q, dtype=jnp.result_type(float))

near_one

near_one(q: Scalar, eps: float = Q_EPS) -> Array

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 = 1.

Q_EPS

Returns:

Type Description
Array

A boolean array, broadcastable against q, that is True where

Array

|q - 1| < eps.

Source code in qjax/shared/validation.py
def near_one(q: Scalar, eps: float = Q_EPS) -> jax.Array:
    """Boolean mask for indices that should use the ``q -> 1`` (classical) limit.

    Args:
        q: The entropic index.
        eps: Half-width of the neighborhood around ``q = 1``.

    Returns:
        A boolean array, broadcastable against ``q``, that is ``True`` where
        ``|q - 1| < eps``.
    """
    return jnp.abs(as_scalar_q(q) - 1.0) < eps

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

q as a floating-point JAX scalar, or NaN where a traced q

Array

is non-positive.

Raises:

Type Description
ValueError

If q is a concrete value and q <= 0.

Source code in qjax/shared/validation.py
def positive_q_or_nan(q: Scalar) -> jax.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.

    Args:
        q: The entropic index.

    Returns:
        ``q`` as a floating-point JAX scalar, or ``NaN`` where a traced ``q``
        is non-positive.

    Raises:
        ValueError: If ``q`` is a concrete value and ``q <= 0``.
    """
    q = as_scalar_q(q)
    try:
        static = float(q)
    except (TypeError, jax.errors.ConcretizationTypeError, jax.errors.TracerArrayConversionError):
        return jnp.where(q > 0.0, q, jnp.nan)
    if static <= 0.0:
        raise ValueError(
            f"q must be positive (Tsallis entropy is singular at q = 0); got {static}."
        )
    return q

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:

  1. Between the switch-over point and the region where the deformed form is accurate there is a band in which neither is: subtracting 1 from x^{1-q} ~ 1 loses most of the mantissa. In float32 the relative error of (x^{1-q} - 1)/(1-q) peaks near 4e-3 around q = 1.00001.
  2. The classical branch does not depend on q, so its derivative with respect to q is 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.

expm1_over_t

expm1_over_t(t: Array) -> Array

Entire function (exp(t) - 1)/t, equal to 1 at t = 0.

Source code in qjax/shared/series.py
def expm1_over_t(t: jax.Array) -> jax.Array:
    """Entire function ``(exp(t) - 1)/t``, equal to ``1`` at ``t = 0``."""
    small = jnp.abs(t) < SERIES_CUTOFF
    safe_t = jnp.where(small, 1.0, t)
    series = 1.0 + t * (0.5 + t * (1.0 / 6.0 + t / 24.0))
    return jnp.where(small, series, jnp.expm1(safe_t) / safe_t)

log1p_over_t

log1p_over_t(t: Array) -> Array

Entire function log(1 + t)/t, equal to 1 at t = 0.

Source code in qjax/shared/series.py
def log1p_over_t(t: jax.Array) -> jax.Array:
    """Entire function ``log(1 + t)/t``, equal to ``1`` at ``t = 0``."""
    small = jnp.abs(t) < SERIES_CUTOFF
    safe_t = jnp.where(small, 1.0, t)
    series = 1.0 - t * (0.5 - t * (1.0 / 3.0 - t * 0.25))
    return jnp.where(small, series, jnp.log1p(safe_t) / safe_t)

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:

  • qcolors windows 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

qcolors(
    n: int, lo: float = QCOLORS_LO, hi: float = QCOLORS_HI
) -> list

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 [0, 1].

QCOLORS_LO
hi float

Upper bound of the colormap window, in [0, 1].

QCOLORS_HI

Returns:

Type Description
list

A list of n RGBA tuples.

Source code in qjax/plots/style.py
def qcolors(n: int, lo: float = QCOLORS_LO, hi: float = QCOLORS_HI) -> list:
    """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.

    Args:
        n: Number of colors to return.
        lo: Lower bound of the colormap window, in ``[0, 1]``.
        hi: Upper bound of the colormap window, in ``[0, 1]``.

    Returns:
        A list of ``n`` RGBA tuples.
    """
    cmap = mpl.colormaps[CMAP]
    if n <= 0:
        return []
    if n == 1:
        # A lone curve should get a mid-ramp color, not the faintest step.
        return [cmap(0.5 * (lo + hi))]
    return [cmap(p) for p in np.linspace(lo, hi, n)]

qlinestyles

qlinestyles(n: int) -> list[LineStyle]

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 n linestyle specifications, cycling if n exceeds the

list[LineStyle]

number of defined patterns.

Source code in qjax/plots/style.py
def qlinestyles(n: int) -> list[LineStyle]:
    """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.

    Args:
        n: Number of linestyles to return.

    Returns:
        A list of ``n`` linestyle specifications, cycling if ``n`` exceeds the
        number of defined patterns.
    """
    if n <= 0:
        return []
    return [QLINESTYLES[i % len(QLINESTYLES)] for i in range(n)]

use_qjax_style

use_qjax_style() -> None

Apply the qjax publication style (serif math, vector PDF, brand-ramp cycle).

Source code in qjax/plots/style.py
def use_qjax_style() -> None:
    """Apply the qjax publication style (serif math, vector PDF, brand-ramp cycle)."""
    plt.rcParams.update(
        {
            # Typography: serif body with Computer-Modern math (no system LaTeX
            # required). pdf/ps fonttype 42 embeds editable TrueType outlines.
            "text.usetex": False,
            "font.family": "serif",
            "font.serif": ["CMU Serif", "Times New Roman", "DejaVu Serif"],
            "mathtext.fontset": "cm",
            "axes.formatter.use_mathtext": True,
            "pdf.fonttype": 42,
            "ps.fonttype": 42,
            "font.size": 11,
            "axes.titlesize": 12,
            "axes.labelsize": 11,
            "xtick.labelsize": 9.5,
            "ytick.labelsize": 9.5,
            "legend.fontsize": 9.5,
            # Color: the qjax ramp and a matching discrete cycle.
            "image.cmap": CMAP,
            "axes.prop_cycle": cycler(color=qcolors(5)),
            # Figure / output: single-column default, high-resolution rasters.
            "figure.figsize": (6.0, 4.0),
            "figure.dpi": 150,
            "savefig.dpi": 600,
            "savefig.format": "pdf",
            "savefig.bbox": "tight",
            "savefig.pad_inches": 0.03,
            "savefig.transparent": False,
            # Axes, lines, and ticks.
            "axes.linewidth": 0.8,
            "axes.grid": True,
            "axes.axisbelow": True,
            "grid.alpha": 0.25,
            "grid.linewidth": 0.5,
            "axes.spines.top": False,
            "axes.spines.right": False,
            "lines.linewidth": 1.8,
            "lines.markersize": 5,
            "legend.frameon": False,
            "legend.handlelength": 1.6,
            "xtick.direction": "in",
            "ytick.direction": "in",
            "xtick.minor.visible": True,
            "ytick.minor.visible": True,
            "xtick.major.size": 4.0,
            "ytick.major.size": 4.0,
            "xtick.minor.size": 2.0,
            "ytick.minor.size": 2.0,
        }
    )

save_figure

save_figure(
    fig: Figure, path: str | Path, transparent: bool = False
) -> Path

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 .pdf.

required
transparent bool

If True, write with a transparent background so the figure blends with whatever it is placed on.

False

Returns:

Type Description
Path

The resolved output path.

Source code in qjax/plots/style.py
def save_figure(fig: plt.Figure, path: str | Path, transparent: bool = False) -> Path:
    """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``.

    Args:
        fig: The figure to write.
        path: Destination path; any extension is replaced with ``.pdf``.
        transparent: If ``True``, write with a transparent background so the
            figure blends with whatever it is placed on.

    Returns:
        The resolved output path.
    """
    out = Path(path).with_suffix(".pdf")
    out.parent.mkdir(parents=True, exist_ok=True)
    fig.savefig(out, bbox_inches="tight", pad_inches=0.03, transparent=transparent)
    return out

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]

(min, max) of the (positive) domain.

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

None

Returns:

Type Description
Axes

The axis containing the plot.

Source code in qjax/plots/functions.py
def 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: plt.Axes | None = None,
) -> plt.Axes:
    """Plot the ``q``-logarithm for several entropic indices.

    Args:
        q_values: Entropic indices to draw, one curve each.
        x_range: ``(min, max)`` of the (positive) domain.
        num: Number of sample points.
        ax: Existing axis to draw on; a new one is created if ``None``.

    Returns:
        The axis containing the plot.
    """
    use_qjax_style()
    if ax is None:
        _, ax = plt.subplots()
    x = jnp.linspace(x_range[0], x_range[1], num)
    for q, color in zip(q_values, qcolors(len(q_values)), strict=False):
        ax.plot(x, q_log(x, q), color=color, label=f"q = {q:g}")
    ax.axhline(0.0, color="0.6", lw=0.8)
    ax.set(xlabel="x", ylabel=r"$\ln_q(x)$", title="q-logarithm")
    ax.legend()
    return ax

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]

(min, max) of the domain.

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

None

Returns:

Type Description
Axes

The axis containing the plot.

Source code in qjax/plots/functions.py
def 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: plt.Axes | None = None,
) -> plt.Axes:
    """Plot the ``q``-exponential for several entropic indices.

    Args:
        q_values: Entropic indices to draw, one curve each.
        x_range: ``(min, max)`` of the domain.
        num: Number of sample points.
        ax: Existing axis to draw on; a new one is created if ``None``.

    Returns:
        The axis containing the plot.
    """
    use_qjax_style()
    if ax is None:
        _, ax = plt.subplots()
    x = jnp.linspace(x_range[0], x_range[1], num)
    for q, color in zip(q_values, qcolors(len(q_values)), strict=False):
        ax.plot(x, q_exp(x, q), color=color, label=f"q = {q:g}")
    ax.set(xlabel="x", ylabel=r"$\exp_q(x)$", title="q-exponential")
    ax.legend()
    return ax

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 (q < 3).

(0.5, 1.0, 1.5, 2.0, 2.5)
beta float

Shared inverse-width parameter.

1.0
x_range tuple[float, float]

(min, max) of the domain.

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

None

Returns:

Type Description
Axes

The axis containing the plot.

Source code in qjax/plots/distributions.py
def 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: plt.Axes | None = None,
) -> plt.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.

    Args:
        q_values: Entropic indices to draw, one curve each (``q < 3``).
        beta: Shared inverse-width parameter.
        x_range: ``(min, max)`` of the domain.
        num: Number of sample points.
        ax: Existing axis to draw on; a new one is created if ``None``.

    Returns:
        The axis containing the plot.
    """
    use_qjax_style()
    if ax is None:
        _, ax = plt.subplots()
    x = jnp.linspace(x_range[0], x_range[1], num)
    for q, color in zip(q_values, qcolors(len(q_values)), strict=False):
        ax.plot(x, q_gaussian_pdf(x, q, beta), color=color, label=f"q = {q:g}")
    ax.set(xlabel="x", ylabel="density", title=f"q-Gaussian (β = {beta:g})")
    ax.legend()
    return ax