jax.random.chisquare

Contents

jax.random.chisquare#

jax.random.chisquare(key, df, shape=None, dtype=None, *, method='exact', out_sharding=None)[source]#

Sample Chisquare random values with given shape and float dtype.

The values are distributed according to the probability density function:

\[f(x; \nu) \propto x^{\nu/2 - 1}e^{-x/2}\]

on the domain \(0 < x < \infty\), where \(\nu > 0\) represents the degrees of freedom, given by the parameter df.

Parameters:
  • key (ArrayLike) – a PRNG key used as the random key.

  • df (RealArray) – a float or array of floats broadcast-compatible with shape representing the parameter of the distribution.

  • shape (Shape | None) – optional, a tuple of nonnegative integers specifying the result shape. Must be broadcast-compatible with df. The default (None) produces a result shape equal to df.shape.

  • dtype (DTypeLikeFloat | None) – optional, a float dtype for the returned values (default float64 if jax_enable_x64 is true, otherwise float32).

  • method (str) – optional, the sampling algorithm to use, either 'exact' (the default) or 'approximate'. The 'exact' method is a rejection sampler. The 'approximate' method is loop-free and faster but carries a small bias. The gradient w.r.t. df differs between the two methods because of the ambiguity in defining a gradient for random variates.

  • out_sharding (NamedSharding | P | None) – optional, Specifies how the output array should be sharded across devices in multi-device computation. Can be a NamedSharding, a PartitionSpec (P), or None (default). When specified, the output will be sharded according to the given sharding specification. Primarily used in explicit sharding mode. See the explicit sharding tutorial for more details.

Returns:

A random array with the specified dtype and with shape given by shape if shape is not None, or else by df.shape.

Return type:

Array