"""Backend abstraction for array operations.
Provides a unified :class:`Backend` class that delegates array creation and
manipulation to one of three libraries: **torch**, **jax**, or **numpy**.
A module-level :data:`backend` instance is created on import and serves as the
primary interface for users.
"""
import os
import importlib, importlib.util
from typing import TYPE_CHECKING, TypeVar
import numpy as np
from . import utils
if TYPE_CHECKING:
import torch # type: ignore
import jax.numpy as jnp # type: ignore
ArrayLike = TypeVar("ArrayLike", np.ndarray, "torch.Tensor", "jnp.ndarray")
else:
ArrayLike = TypeVar("ArrayLike")
[docs]
class Backend:
"""Unified interface for array operations across torch, jax, and numpy.
Provides a single API for creating and manipulating arrays regardless of
the underlying library. Methods such as ``make_array``, ``concatenate``,
``to``, ``sigmoid``, and ``logit`` are dynamically bound when the backend
is set, delegating to the appropriate library-specific implementation.
Parameters
----------
backend : str, optional
Backend name: ``"torch"``, ``"jax"``, or ``"numpy"``. If ``None``,
reads from the ``CASKADE_BACKEND`` environment variable, defaulting
to ``"torch"``.
Examples
--------
Use the module-level ``backend`` instance to switch backends::
from caskade import backend
backend.backend = "numpy"
arr = backend.make_array([1.0, 2.0, 3.0])
"""
def __init__(self, backend=None):
"""Initialize the backend.
Parameters
----------
backend : str, optional
Backend name: ``"torch"``, ``"jax"``, or ``"numpy"``. If ``None``,
reads from the ``CASKADE_BACKEND`` environment variable, defaulting
to ``"torch"``.
"""
self.backend = backend
@property
def backend(self):
"""str : Name of the active backend (``"torch"``, ``"jax"``, or ``"numpy"``)."""
return self._backend
@backend.setter
def backend(self, backend):
if backend is None:
backend = os.getenv("CASKADE_BACKEND", "none").lower()
if backend == "none": # Try to find available backend
if importlib.util.find_spec("torch") is not None:
backend = "torch"
elif importlib.util.find_spec("jax") is not None:
backend = "jax"
else:
backend = "numpy"
self.module = self._load_backend(backend)
self._backend = backend
def _load_backend(self, backend):
if backend == "torch":
self.setup_torch()
return importlib.import_module("torch")
elif backend == "jax":
self.setup_jax()
return importlib.import_module("jax.numpy")
elif backend == "numpy":
self.setup_numpy()
return importlib.import_module("numpy")
else:
raise ValueError(f"Unsupported backend: {backend}")
[docs]
def setup_torch(self):
self.make_array = self._make_array_torch
self._array_type = self._array_type_torch
self.concatenate = self._concatenate_torch
self.broadcast_cat = utils.broadcast_cat_torch
self.tolist = self._tolist_torch
self.view = self._view_torch
self.detach = self._detach_torch
self.as_array = self._as_array_torch
self.to = self._to_torch
self.to_numpy = self._to_numpy_torch
self.ste_clip = self._ste_clip_torch
[docs]
def setup_jax(self):
self.jax = importlib.import_module("jax")
self.make_array = self._make_array_jax
self._array_type = self._array_type_jax
self.concatenate = self._concatenate_jax
self.broadcast_cat = utils.broadcast_cat_jax
self.tolist = self._tolist_jax
self.view = self._view_jax
self.detach = self._detach_jax
self.as_array = self._as_array_jax
self.to = self._to_jax
self.to_numpy = self._to_numpy_jax
self.ste_clip = self._ste_clip_jax
[docs]
def setup_numpy(self):
self.make_array = self._make_array_numpy
self._array_type = self._array_type_numpy
self.concatenate = self._concatenate_numpy
self.broadcast_cat = utils.broadcast_cat_numpy
self.tolist = self._tolist_numpy
self.view = self._view_numpy
self.detach = self._detach_numpy
self.as_array = self._as_array_numpy
self.to = self._to_numpy
self.to_numpy = self._to_numpy_numpy
self.ste_clip = self._ste_clip_numpy
@property
def array_type(self):
"""type : The array class for the active backend.
Returns ``torch.Tensor``, ``jax.numpy.ndarray``, or ``numpy.ndarray``
depending on the current backend. Useful for ``isinstance`` checks.
Returns
-------
type
The array class used by the active backend.
Examples
--------
::
isinstance(my_array, backend.array_type)
"""
return self._array_type()
def _make_array_torch(self, array, dtype=None, device=None):
return self.module.tensor(array, dtype=dtype, device=device)
def _make_array_jax(self, array, dtype=None, **kwargs):
return self.module.array(array, dtype=dtype)
def _make_array_numpy(self, array, dtype=None, **kwargs):
return self.module.array(array, dtype=dtype)
def _array_type_torch(self):
return self.module.Tensor
def _array_type_jax(self):
return self.module.ndarray
def _array_type_numpy(self):
return self.module.ndarray
def _concatenate_torch(self, arrays, axis=0):
return self.module.cat(arrays, dim=axis)
def _concatenate_jax(self, arrays, axis=0):
return self.module.concatenate(arrays, axis=axis)
def _concatenate_numpy(self, arrays, axis=0):
return self.module.concatenate(arrays, axis=axis)
def _detach_torch(self, array):
return array.detach()
def _detach_jax(self, array):
return array
def _detach_numpy(self, array):
return array
def _tolist_torch(self, array):
return array.detach().cpu().tolist()
def _tolist_jax(self, array):
return array.block_until_ready().tolist()
def _tolist_numpy(self, array):
return array.tolist()
def _view_torch(self, array, shape):
return array.reshape(shape)
def _view_jax(self, array, shape):
return array.reshape(shape)
def _view_numpy(self, array, shape):
return array.reshape(shape)
def _as_array_torch(self, array, dtype=None, device=None):
return self.module.as_tensor(array, dtype=dtype, device=device)
def _as_array_jax(self, array, dtype=None, **kwargs):
return self.module.asarray(array, dtype=dtype)
def _as_array_numpy(self, array, dtype=None, **kwargs):
return self.module.asarray(array, dtype=dtype)
def _to_torch(self, array, dtype=None, device=None):
return array.to(dtype=dtype, device=device)
def _to_jax(self, array, dtype=None, device=None):
return self.jax.device_put(array.astype(dtype), device=device)
def _to_numpy(self, array, dtype=None, **kwargs):
return array.astype(dtype)
def _to_numpy_torch(self, array):
return array.detach().cpu().numpy()
def _to_numpy_jax(self, array):
return np.array(array.block_until_ready())
def _to_numpy_numpy(self, array):
return array
[docs]
def any(self, array):
"""Test whether any element evaluates to ``True``.
Parameters
----------
array : ArrayLike
Input array.
Returns
-------
ArrayLike
Scalar result; ``True`` if any element is non-zero.
"""
return self.module.any(array)
[docs]
def all(self, array):
"""Test whether all elements evaluate to ``True``.
Parameters
----------
array : ArrayLike
Input array.
Returns
-------
ArrayLike
Scalar result; ``True`` if every element is non-zero.
"""
return self.module.all(array)
[docs]
def log(self, array):
"""Compute the natural logarithm element-wise.
Parameters
----------
array : ArrayLike
Input array.
Returns
-------
ArrayLike
Element-wise natural logarithm of the input.
"""
return self.module.log(array)
[docs]
def exp(self, array):
"""Compute the exponential element-wise.
Parameters
----------
array : ArrayLike
Input array.
Returns
-------
ArrayLike
Element-wise exponential of the input.
"""
return self.module.exp(array)
[docs]
def sum(self, array, axis=None):
"""Sum array elements over a given axis.
Parameters
----------
array : ArrayLike
Input array.
axis : int or None, optional
Axis along which to sum. If ``None``, sums all elements.
Returns
-------
ArrayLike
Sum of elements.
"""
return self.module.sum(array, axis=axis)
def _ste_clip_torch(self, array, min_val, max_val):
"""Clip function such that gradients are preserved at the boundaries."""
clipped = self.module.clamp(array, min=min_val, max=max_val)
return clipped.detach() + (array - array.detach())
def _ste_clip_jax(self, array, min_val, max_val):
"""Clip function such that gradients are preserved at the boundaries."""
clipped = self.module.clip(array, min=min_val, max=max_val)
return self.jax.lax.stop_gradient(clipped) + (array - self.jax.lax.stop_gradient(array))
def _ste_clip_numpy(self, array, min_val, max_val):
"""Standard clip function of numpy."""
return np.clip(array, a_min=min_val, a_max=max_val)
#: Module-level :class:`Backend` instance used as the default entry point.
#: Import and configure this object to switch backends globally::
#:
#: from caskade import backend
#: backend.backend = "numpy"
backend = Backend()