AbstractLocationFunction#

class gpjax.kernels.location_functions.AbstractLocationFunction(active_dims=<factory>)[source]#

Bases: _SummaryMixin, Module

Base class for location functions.

A subclass implements __call__, which maps one input point to a scalar on the log scale. Use slice_input to read only the columns in active_dims.

Parameters:

active_dims (list[int] | slice)

slice_input(x)[source]#

Select the columns in active_dims from one input point.

Parameters:

x (Num[jaxlib._jax.Array, 'D'] | Num[ndarray, 'D']) – one input point, with all columns of the data.

Returns:

The selected columns.

Return type:

Num[jaxlib._jax.Array, ‘Q’] | Num[ndarray, ‘Q’]