Backends¶
This page contains information about the Keras backends.
prfmodel is built on Keras 3.0 to perform GPU-accelerated operations and gradient tracking. Keras provides a unified
API for several backends that perform the actual computations. prfmodel currently supports three backends: TensorFlow,
PyTorch, and JAX. At least one backend must be installed alongside prfmodel to use the package. Some tasks cannot
be solved by solely relying on Keras and backend-specific implementations are needed. These live in the private
prfmodel._backend module and are mainly related to stochastic gradient descent fitting and compilation.
Importantly, only the backend selected by the user is imported in the package.
The three backends differ in how they optimize operations:
TensorFlow uses graph-based execution by wrapping functions inside
tf.function.PyTorch uses compilation through
torch.compile.JAX uses just-in-time compilation via
jax.jit.
These optimizations can yield substantial increases in speed. However, they also require optimized functions to be
traceable. This means that their inputs and outputs must be tensor objects, operations should be done with,
keras.ops (or backend-specific functions, e.g., jax.numpy), and that the function logic must not
depend on concrete values of the inputs (e.g., no if-else branching on input values). The reason is that, when the
optimization graph is being built, the values of the input tensors are not yet available, leading to an error.
For example, this function could not be optimized in prfmodel:
import numpy as np
import pandas as pd
def my_fun(
design: np.ndarray, # inputs must be tensors
parameters: pd.Dataframe, # inputs must be tensors
) -> np.ndarray: # outputs must be tensors
if design[0] > 0: # control flow must not depend on input values
...
However, the input shapes are already known during compilation, so the function logic can depend on them:
from keras import ops
from prfmodel.typing import Tensor
from prfmodel.utils import TensorFrame
def my_fun(
# use backend-agnostic tensor objects instead of numpy.array
# (e.g., in JAX, this will be jax.Array)
design: Tensor,
# use backend-agnostic tensor frame instead of pandas.DataFrame
parameters: TensorFrame,
) -> Tensor:
if design.shape[0] > 0: # control flow can depend on input shape
# use keras.ops for operations inside function
design_log = ops.log(design)
...
See Creating a custom model for details and a full example.
In prfmodel, the prfmodel._backend._compile module exports a backend-specific compile_fun that enables
optimization via the above-mentioned functions. Both prfmodel.fitters.GridFitter and
prfmodel.fitters.SGDFitter internally use this compile_fun to speed up computations if their
compile_step flag is enabled. For the optimization to work, model classes must follow a specific implementation
logic that is explained in Architecture.
Additional backend-specific functions are imported via prfmodel._backend._external.