val#

gpjax.parameters.val(x)[source]#

Return a parameter’s constrained value.

Call this wherever a parameter meets arithmetic. Models are always held in their wrapped form – including the model fit returns – so val is the single point at which the constraining bijection is applied.

Safe to apply to anything: a parameter is resolved to its constrained value (recursively, so nested wrappers such as paramax.non_trainable are handled), while a plain array or float is returned unchanged.

Parameters:

x – A parameter, or any value that does not need unwrapping.

Returns:

The constrained value of x if it is a parameter, else x itself.

Example

>>> import jax.numpy as jnp
>>> from gpjax.parameters import PositiveReal, val
>>> float(val(PositiveReal(jnp.array(2.0))))
2.0
>>> float(val(jnp.array(2.0)))
2.0

Expand for references to gpjax.parameters.val

val

Backend Module Design

Classification

Kernel Guide

Dual Parameterisation of Sparse GPs (t-SVGP)

Natural Gradients in Practice

The Sharp Bits / How does the parameter system work? / When do I need to call val()?