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.

Further Reading