JaxDirichletLevelSetCredalSet¶
- class probly.representation.credal_set.jax.JaxDirichletLevelSetCredalSet(alphas: Array, threshold: Array)[source]¶
Bases:
JaxAxisProtected[Any],JaxCategoricalCredalSet,DirichletLevelSetCredalSetDirichlet density level set credal set.
Contains all distributions whose Dirichlet likelihood is at least
thresholdtimes the peak density.- alphas: jax.Array¶
- astype(dtype: DTypeLike | None, copy: bool = False, device: jax.Device | Sharding | None = None) Self[source]¶
Return a copy of the array cast to the given data type.
- Parameters:
dtype – The target data type.
copy – Whether to always return a copy.
device – The device the result should live on.
- Returns:
A copy with every array-valued field cast to the given data type.
- property at: JaxAxisProtectedIndexUpdateHelper[J]¶
Return the out-of-place indexed update helper, mirroring
jax.Array.at.
- property barycenter: JaxCategoricalDistribution¶
Return the Dirichlet mean as the barycenter.
- block_until_ready() Self[source]¶
Block until the asynchronous computation of every array-valued field has finished.
- Returns:
The object itself.
- property device: Device¶
Device of the array.
- flatten(order: str = 'C') Self[source]¶
Return a flattened version of the array.
- Parameters:
order – The read/write element order.
- Returns:
The flattened array.
- classmethod from_jax_sample(sample: JaxSample[JaxCategoricalDistribution]) Self[source]¶
Create from a jax sample (not supported).
- Raises:
NotImplementedError – Always, as this credal set type cannot be created from samples.
- classmethod from_sample(sample: Sample[JaxCategoricalDistribution]) Self[source]¶
Create a credal set from a finite sample.
- Parameters:
sample – The sample to create the credal set from.
- Returns:
The created credal set.
- item(*args: int) bool | int | float | complex[source]¶
Return the array as a Python scalar.
- Parameters:
*args – Optional index of the element to return.
- Returns:
The selected element as a Python scalar.
- lower(key: Array | None = None) Array[source]¶
Compute per-class lower bounds via Monte Carlo sampling.
- Parameters:
key – PRNG key used for sampling. Defaults to a fixed key for determinism.
- Returns:
Lower probability bounds, shape (…, K).
- permitted_functions = {}¶
- protected_values(func: Callable | None = None) dict[str, JaxProtectedValue] | None[source]¶
Return all protected field values as-is.
Optionally takes the jax function that triggered the call for context. This can be used to conditionally modify the returned values or prevent them from being accessed.
- ravel(order: str = 'C') Self[source]¶
Return a flattened version of the array.
- Parameters:
order – The read/write element order.
- Returns:
The flattened array.
- reshape(*args: int | Sequence[int], order: str = 'C') Self[source]¶
Return a copy with reshaped protected values.
- squeeze(axis: int | Sequence[int] | None = None) Self[source]¶
Return the array with axes of length one removed.
- Parameters:
axis – The axis or axes to remove. If omitted, all length-one axes are removed.
- Returns:
The squeezed array.
- stop_gradient() Self[source]¶
Return a copy detached from the autodiff graph.
Subclasses holding NumPy sidecar fields should override this to leave those fields untouched.
- Returns:
A copy through which gradients do not propagate.
- take_along_axis(indices: Array, axis: int = -1) Self[source]¶
Return a copy with gathered protected values along a batch dimension.
- threshold: jax.Array¶
- to_device(device: Any, /, *, stream: int | Any | None = None) Self[source]¶
Move the array to the given device.
- Parameters:
device – The target device.
stream – Unsupported by JAX, must be None.
- Returns:
A copy with every array-valued field on the given device.
- Raises:
NotImplementedError – If
streamis not None.
- transpose(*args: int | Sequence[int] | None) Self[source]¶
Return a transposed version of the array.
- Parameters:
*args – The permutation of the axes, either as a single sequence or as separate integers. If omitted, the axes are reversed.
- Returns:
The transposed array.
- tree_flatten() tuple[tuple[Any, ...], Any][source]¶
Split the object into pytree children and static auxiliary data.
Fields holding arrays (anything exposing
shape, plusNone) become children, all remaining fields become auxiliary data. Subclasses may override this together withtree_unflatten().- Returns:
The children and the auxiliary data.
- classmethod tree_unflatten(aux_data: Any, children: tuple[Any, ...]) Self[source]¶
Rebuild an object from pytree children and auxiliary data.
The object is built without calling
__init__, because JAX passes tracer objects here during tracing and those would fail the usual validation.- Parameters:
aux_data – The auxiliary data produced by
tree_flatten().children – The children produced by
tree_flatten().
- Returns:
The reconstructed object.
- type = 'categorical'¶
- upper(key: Array | None = None) Array[source]¶
Compute per-class upper bounds via Monte Carlo sampling.
- Parameters:
key – PRNG key used for sampling. Defaults to a fixed key for determinism.
- Returns:
Upper probability bounds, shape (…, K).
- view(dtype: DTypeLike | None = None, type: None = None) Self[source]¶
Return a bit-cast view of the array.
- Parameters:
dtype – The data type to reinterpret the underlying bytes as.
type – Unsupported by JAX, must be
None.
- Returns:
A copy with every array-valued field bit-cast to the given data type.
- Raises:
NotImplementedError – If
typeis not None.