Source code for caskade.backend

"""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()