# 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 {py:mod}`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: ```python 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: ```python 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 [](../tutorials/tutorials/custom_models.md) for details and a full example. In prfmodel, the {py:mod}`prfmodel._backend._compile` module exports a backend-specific `compile_fun` that enables optimization via the above-mentioned functions. Both {py:class}`prfmodel.fitters.GridFitter` and {py:class}`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.md). Additional backend-specific functions are imported via {py:mod}`prfmodel._backend._external`. ## Further Reading - TensorFlow `tf.function` [guide](https://www.tensorflow.org/guide/function). - PyTorch `torch.compile` [tutorial](https://docs.pytorch.org/tutorials/intermediate/torch_compile_tutorial.html). - JAX `jax.jit` [tutorial](https://docs.jax.dev/en/latest/jit-compilation.html).