Skip to content

Differentiation

These functions construct derivatives directly from symbolic expressions. For the usual modelling interface, including derivatives of a named Function, see Building functions. The derivatives guide explains gradients, Jacobians, and Hessians.

Modes

A Jacobian-vector product, jvp, computes a directional derivative without constructing the full Jacobian. A vector-Jacobian product, vjp, propagates output weights back to the inputs. The *_many forms handle several directions or output weights together.

jvp

jvp(expr: Expr, wrt: Expr, seed: Expr) -> Expr

Forward-mode derivative: J(expr, wrt) @ seed, with seed shaped like wrt.

One pass per seed. For many seeds at once use jvp_many, which shares the expensive work.

Source code in src/scaly/ad/forward.py
def jvp(expr: Expr, wrt: Expr, seed: Expr) -> Expr:
  """Forward-mode derivative: ``J(expr, wrt) @ seed``, with ``seed`` shaped like ``wrt``.

  One pass per seed. For many seeds at once use ``jvp_many``, which shares the expensive work.
  """
  return _jvp(expr, {wrt: seed}, {}, {})

jvp_many

jvp_many(expr: Expr, wrt: Expr, seeds: Expr) -> Expr

Forward mode over several seeds in one pass. seeds has shape (n, *wrt.shape).

Structural rules share work across seeds. One cos serves every column of a sin's derivative, so this is much cheaper than n separate jvp calls. Ops without a multi-seed rule fall back to per-seed evaluation. Set SCALY_STRICT_JVP_MANY=1 to raise instead of falling back. Returns shape (n, *expr.shape).

Source code in src/scaly/ad/forward.py
def jvp_many(expr: Expr, wrt: Expr, seeds: Expr) -> Expr:
  """Forward mode over several seeds in one pass. ``seeds`` has shape ``(n, *wrt.shape)``.

  Structural rules share work across seeds. One ``cos`` serves every column of a ``sin``'s
  derivative, so this is much cheaper than ``n`` separate ``jvp`` calls. Ops without a
  multi-seed rule fall back to per-seed evaluation. Set ``SCALY_STRICT_JVP_MANY=1`` to raise
  instead of falling back. Returns shape ``(n, *expr.shape)``.
  """
  if len(seeds.shape) < 1 or seeds.shape[1:] != wrt.shape:
    raise ValueError(f"multi-seed JVP expects seeds shape (nseed, *{wrt.shape}), got {seeds.shape}")
  if seeds.shape[0] == 0:
    return Expr.const(np.zeros((0, *expr.shape), dtype=np.float64))
  strict = env_bool("SCALY_STRICT_JVP_MANY", False)
  try:
    ret = _jvp_many_structural(expr, wrt, seeds, {}, {})
  except _JVPManyUnsupported as unsupported:
    if strict:
      raise NotImplementedError(
        f"structural jvp_many does not support {unsupported.op!r}; SCALY_STRICT_JVP_MANY=1 forbids the unrolled fallback"
      ) from unsupported
    return _jvp_many_unrolled(expr, wrt, seeds)
  expected = (seeds.shape[0], *expr.shape)
  if ret.shape == expected:
    return ret
  if strict:
    raise NotImplementedError(
      f"structural jvp_many returned shape {ret.shape} for {expr.op!r}, expected {expected}; SCALY_STRICT_JVP_MANY=1 forbids the unrolled fallback"
    )
  return _jvp_many_unrolled(expr, wrt, seeds)

vjp

vjp(outputs: Sequence[Expr], wrts: Sequence[Expr], cotangents: Sequence[Expr]) -> tuple[Expr, ...]

Reverse-mode derivative: one adjoint per entry of wrts, seeded by cotangents.

Each cotangent is shaped like its output. One sweep computes the derivative of a scalar with respect to every input at once, which is why gradients go through here.

Source code in src/scaly/ad/reverse.py
def vjp(outputs: Sequence[Expr], wrts: Sequence[Expr], cotangents: Sequence[Expr]) -> tuple[Expr, ...]:
  """Reverse-mode derivative: one adjoint per entry of ``wrts``, seeded by ``cotangents``.

  Each cotangent is shaped like its output. One sweep computes the derivative of a scalar with
  respect to every input at once, which is why gradients go through here.
  """
  if len(outputs) != len(cotangents):
    raise ValueError(f"expected {len(outputs)} cotangents, got {len(cotangents)}")
  adjoints: dict[int, Expr] = {}
  nodes = topo(outputs)
  expr_ids = {e.id for e in nodes}
  dep_memo: dict[tuple[int, int], bool] = {}

  def needed(expr: Expr) -> bool:
    return any(_depends_on(expr, wrt, dep_memo) for wrt in wrts)

  for out, cot in zip(outputs, cotangents, strict=True):
    if out.shape != cot.shape:
      raise ValueError(f"cotangent for output shape {out.shape} has shape {cot.shape}")
    adjoints[out.id] = cot if out.id not in adjoints else adjoints[out.id] + cot

  for expr in reversed(nodes):
    cot = adjoints.get(expr.id)
    if cot is None or expr.op in {ExprOp.INPUT, ExprOp.CONST} or not needed(expr):
      continue
    if expr.op == ExprOp.VMAP:
      for arg, arg_cot in _vmap_vjp(expr, cot, wrts, dep_memo):
        if arg.id in expr_ids:
          adjoints[arg.id] = arg_cot if arg.id not in adjoints else adjoints[arg.id] + arg_cot
      continue
    for arg, arg_cot in zip(expr.args, _local_vjp(expr, cot), strict=True):
      if arg.id in expr_ids:
        adjoints[arg.id] = arg_cot if arg.id not in adjoints else adjoints[arg.id] + arg_cot

  return tuple(adjoints.get(wrt.id, zeros_like(wrt)) for wrt in wrts)

vjp_many

vjp_many(outputs: Sequence[Expr], wrts: Sequence[Expr], cotangents: Sequence[Expr]) -> tuple[Expr, ...]

Reverse mode over several cotangent seeds at once.

Each cotangent has a leading seed axis, (n, *output.shape), and each returned adjoint carries the same leading axis.

Source code in src/scaly/ad/reverse.py
def vjp_many(outputs: Sequence[Expr], wrts: Sequence[Expr], cotangents: Sequence[Expr]) -> tuple[Expr, ...]:
  """Reverse mode over several cotangent seeds at once.

  Each cotangent has a leading seed axis, ``(n, *output.shape)``, and each returned adjoint
  carries the same leading axis.
  """
  if len(outputs) != len(cotangents):
    raise ValueError(f"expected {len(outputs)} cotangents, got {len(cotangents)}")
  nseed: int | None = None
  for out, cot in zip(outputs, cotangents, strict=True):
    if len(cot.shape) < 1 or cot.shape[1:] != out.shape:
      raise ValueError(f"multi-seed VJP expects cotangent shape (nseed, *{out.shape}), got {cot.shape}")
    if nseed is None:
      nseed = cot.shape[0]
    elif cot.shape[0] != nseed:
      raise ValueError(f"all VJP cotangents must have the same leading seed axis, got {nseed} and {cot.shape[0]}")
  nseed = 0 if nseed is None else nseed
  if nseed == 0:
    return tuple(Expr.const(np.zeros((0, *wrt.shape), dtype=np.float64)) for wrt in wrts)

  per_seed = [vjp(outputs, wrts, tuple(cot[i] for cot in cotangents)) for i in range(nseed)]
  return tuple(stack([seed_grads[i] for seed_grads in per_seed], axis=0) for i in range(len(wrts)))

Whole derivatives

jacobian

jacobian(expr: Expr, wrt: Expr) -> Expr

Dense Jacobian d expr / d wrt, shape (expr.size, wrt.size).

Column j is the derivative with respect to wrt[j]. Computed by pushing the whole identity through forward mode in one batched pass, then simplifying.

Source code in src/scaly/ad/derivatives.py
def jacobian(expr: Expr, wrt: Expr) -> Expr:
  """Dense Jacobian ``d expr / d wrt``, shape ``(expr.size, wrt.size)``.

  Column ``j`` is the derivative with respect to ``wrt[j]``. Computed by pushing the whole
  identity through forward mode in one batched pass, then simplifying.
  """
  if wrt.size == 0:
    return Expr.const(np.zeros((expr.size, 0), dtype=np.float64))
  # Batched forward AD: stack the wrt.size identity columns as a (wrt.size, *wrt.shape) seed and
  # push them through jvp_many. The structural multi-seed rules share cos/sin/exp across columns
  # and turn per-column chain-rule unrolls into small matmuls. Falls back to column-by-column jvp
  # only if jvp_many hits an unsupported op. Output is reshaped from (wrt.size, expr.size) →
  # (expr.size, wrt.size) so column j of the Jacobian = partial expr / partial wrt[j].
  seed_arr = np.eye(wrt.size, dtype=np.float64).reshape((wrt.size, *wrt.shape))
  return simplify_cse_fixpoint(jvp_many(expr, wrt, Expr.const(seed_arr)).reshape((wrt.size, expr.size)).transpose((1, 0)))

gradient

gradient(expr: Expr, wrt: Expr) -> Expr

Gradient of a scalar expr with respect to wrt, as one reverse sweep.

Source code in src/scaly/ad/derivatives.py
def gradient(expr: Expr, wrt: Expr) -> Expr:
  """Gradient of a scalar ``expr`` with respect to ``wrt``, as one reverse sweep."""
  if expr.size != 1:
    raise ValueError("gradient expects a scalar expression")
  return vjp((expr,), (wrt,), (_ones_like(expr),))[0]

hessian

hessian(expr: Expr, wrt: Expr) -> Expr

Second derivatives of a scalar expr: the Jacobian of its gradient.

Source code in src/scaly/ad/derivatives.py
def hessian(expr: Expr, wrt: Expr) -> Expr:
  """Second derivatives of a scalar ``expr``: the Jacobian of its gradient."""
  return jacobian(gradient(expr, wrt).reshape((wrt.size,)), wrt)

finite_difference

finite_difference(fun: Any, x: ndarray, eps: float = 1e-06) -> np.ndarray

Approximate the Jacobian of the numerical fun at x by central differences.

Returns shape (output size, x.size). Useful as an independent check of a symbolic derivative.

Source code in src/scaly/ad/derivatives.py
def finite_difference(fun: Any, x: np.ndarray, eps: float = 1e-6) -> np.ndarray:
  """Approximate the Jacobian of the numerical ``fun`` at ``x`` by central differences.

  Returns shape ``(output size, x.size)``. Useful as an independent check of a symbolic derivative.
  """
  x = np.asarray(x, dtype=np.float64)
  y0 = np.asarray(fun(x), dtype=np.float64).reshape(-1)
  jac = np.empty((y0.size, x.size), dtype=np.float64)
  flat = x.reshape(-1)
  for i in range(flat.size):
    xp = flat.copy()
    xm = flat.copy()
    xp[i] += eps
    xm[i] -= eps
    jac[:, i] = (np.asarray(fun(xp.reshape(x.shape))).reshape(-1) - np.asarray(fun(xm.reshape(x.shape))).reshape(-1)) / (2 * eps)
  return jac

Sparsity

A sparsity pattern identifies entries that may be nonzero. Coloring groups derivative directions that can be evaluated together without mixing their results. See the sparsity guide for the storage format and examples.

jacobian_sparsity

jacobian_sparsity(expr: Expr, wrt: Expr) -> SparsityPattern

Estimate structural sparsity of d vec(expr) / d vec(wrt).

This is purely symbolic: it tracks element dependencies through the graph without using numerical values. It is conservative for nonsmooth elementwise ops and block ops, but exact for the structural/arithmetic subset currently implemented here.

Source code in src/scaly/ad/sparsity.py
def jacobian_sparsity(expr: Expr, wrt: Expr) -> SparsityPattern:
  """Estimate structural sparsity of ``d vec(expr) / d vec(wrt)``.

  This is purely symbolic: it tracks element dependencies through the graph without using
  numerical values. It is conservative for nonsmooth elementwise ops and block ops, but exact
  for the structural/arithmetic subset currently implemented here.
  """
  return _mask_sparsity(_jac_mask(expr, wrt, {}))

column_coloring

column_coloring(sparsity: SparsityPattern) -> tuple[int, ...]

Assign each column a color, greedily, so no two columns sharing a row get the same one.

Columns of one color can be recovered from a single forward pass, so the number of colors is the number of passes a compact Jacobian costs.

Source code in src/scaly/ad/sparsity.py
def column_coloring(sparsity: SparsityPattern) -> tuple[int, ...]:
  """Assign each column a color, greedily, so no two columns sharing a row get the same one.

  Columns of one color can be recovered from a single forward pass, so the number of colors is
  the number of passes a compact Jacobian costs.
  """
  row_ptr, col_ind, _ = sparsity.to_csr()
  col_ptr, row_ind, _ = sparsity.to_csc()
  colors: list[int] = []
  for col in range(sparsity.shape[1]):
    used = {colors[other] for row in row_ind[col_ptr[col] : col_ptr[col + 1]] for other in col_ind[row_ptr[row] : row_ptr[row + 1]] if other < col}
    color = 0
    while color in used:
      color += 1
    colors.append(color)
  return tuple(colors)

star_coloring

star_coloring(sparsity: SparsityPattern) -> tuple[int, ...]

Greedily star-color a square sparsity pattern in column order.

The pattern is treated as an undirected graph. A valid coloring is proper, and no simple path of three edges has only two colors. This is the coloring needed to recover a symmetric Hessian from compressed forward products. It is deliberately not distance-2 coloring.

Source code in src/scaly/ad/sparsity.py
def star_coloring(sparsity: SparsityPattern) -> tuple[int, ...]:
  """Greedily star-color a square sparsity pattern in column order.

  The pattern is treated as an undirected graph. A valid coloring is proper, and no simple path
  of three edges has only two colors. This is the coloring needed to recover a symmetric Hessian
  from compressed forward products. It is deliberately not distance-2 coloring.
  """
  rows, cols = sparsity.shape
  if rows != cols:
    raise ValueError(f"star coloring requires a square sparsity pattern, got {sparsity.shape}")

  neighbors = [set() for _ in range(rows)]
  for row, col in zip(sparsity.rows, sparsity.cols, strict=True):
    if row != col:
      neighbors[row].add(col)
      neighbors[col].add(row)

  colors = [-1] * rows
  for vertex in range(rows):
    color = 0
    while _star_color_conflicts(vertex, color, colors, neighbors):
      color += 1
    colors[vertex] = color
  return tuple(colors)

color_groups

color_groups(colors: tuple[int, ...]) -> tuple[tuple[int, ...], ...]

Invert a coloring into the column indices belonging to each color.

Source code in src/scaly/ad/sparsity.py
def color_groups(colors: tuple[int, ...]) -> tuple[tuple[int, ...], ...]:
  """Invert a coloring into the column indices belonging to each color."""
  if not colors:
    return ()
  return tuple(tuple(i for i, c in enumerate(colors) if c == color) for color in range(max(colors) + 1))

Sparse derivatives

SparseJacobian dataclass

A sparse derivative: the structural sparsity and the values of its entries, in the same order.

coloring_width is the number of forward-mode directions the values were computed from, or None when they did not come from a coloring.

Source code in src/scaly/ad/sparse.py
@dataclass(frozen=True, slots=True)
class SparseJacobian:
  """A sparse derivative: the structural ``sparsity`` and the ``values`` of its entries, in the same order.

  ``coloring_width`` is the number of forward-mode directions the values were computed from, or
  ``None`` when they did not come from a coloring.
  """

  sparsity: SparsityPattern
  values: Expr
  coloring_width: int | None = None
  _compressed: Expr | None = field(default=None, compare=False, repr=False)
  _recovery: np.ndarray | None = field(default=None, compare=False, repr=False)

  @property
  def flat_indices(self) -> np.ndarray:
    """Row-major positions of the entries in the dense matrix."""
    return np.asarray(self.sparsity.rows, dtype=np.int64) * self.sparsity.shape[1] + np.asarray(self.sparsity.cols, dtype=np.int64)

  def to_dense(self) -> Expr:
    """Scatter the values into a dense matrix expression."""
    return scatter(self.values, self.flat_indices, self.sparsity.shape)

  def triangle(self, triangle: Triangle) -> SparseJacobian:
    """Select one triangle from a symmetric sparse matrix without rebuilding its JVP batch."""
    triangle = _validate_triangle(triangle)
    if triangle == "full":
      return self
    if self.sparsity.shape[0] != self.sparsity.shape[1]:
      raise ValueError(f"triangle selection requires a square sparsity pattern, got {self.sparsity.shape}")
    rows = np.asarray(self.sparsity.rows, dtype=np.int64)
    cols = np.asarray(self.sparsity.cols, dtype=np.int64)
    keep = rows >= cols if triangle == "lower" else rows <= cols
    sparsity = SparsityPattern(
      self.sparsity.shape,
      tuple(int(row) for row in rows[keep]),
      tuple(int(col) for col in cols[keep]),
    )
    source = self._compressed if self._compressed is not None else self.values
    recovery = self._recovery[keep] if self._recovery is not None else np.flatnonzero(keep)
    values = simplify_cse_fixpoint(gather(source, recovery))
    return SparseJacobian(sparsity, values, self.coloring_width, source, recovery)

flat_indices property

flat_indices: ndarray

Row-major positions of the entries in the dense matrix.

to_dense

to_dense() -> Expr

Scatter the values into a dense matrix expression.

Source code in src/scaly/ad/sparse.py
def to_dense(self) -> Expr:
  """Scatter the values into a dense matrix expression."""
  return scatter(self.values, self.flat_indices, self.sparsity.shape)

triangle

triangle(triangle: Triangle) -> SparseJacobian

Select one triangle from a symmetric sparse matrix without rebuilding its JVP batch.

Source code in src/scaly/ad/sparse.py
def triangle(self, triangle: Triangle) -> SparseJacobian:
  """Select one triangle from a symmetric sparse matrix without rebuilding its JVP batch."""
  triangle = _validate_triangle(triangle)
  if triangle == "full":
    return self
  if self.sparsity.shape[0] != self.sparsity.shape[1]:
    raise ValueError(f"triangle selection requires a square sparsity pattern, got {self.sparsity.shape}")
  rows = np.asarray(self.sparsity.rows, dtype=np.int64)
  cols = np.asarray(self.sparsity.cols, dtype=np.int64)
  keep = rows >= cols if triangle == "lower" else rows <= cols
  sparsity = SparsityPattern(
    self.sparsity.shape,
    tuple(int(row) for row in rows[keep]),
    tuple(int(col) for col in cols[keep]),
  )
  source = self._compressed if self._compressed is not None else self.values
  recovery = self._recovery[keep] if self._recovery is not None else np.flatnonzero(keep)
  values = simplify_cse_fixpoint(gather(source, recovery))
  return SparseJacobian(sparsity, values, self.coloring_width, source, recovery)

sparse_jacobian

sparse_jacobian(expr: Expr, wrt: Expr) -> SparseJacobian

Compact sparse Jacobian: the structural pattern plus an expression for its nonzeros only.

Takes the structured path when expr is a VMAP (or a concatenation of them) over exactly wrt, coloring the callee's small local pattern instead of the whole matrix, so the cost tracks the callee rather than the iteration count. Otherwise falls back to coloring the global pattern. The value order is the pattern's (rows, cols) order, which is not necessarily sorted.

Source code in src/scaly/ad/sparse.py
def sparse_jacobian(expr: Expr, wrt: Expr) -> SparseJacobian:
  """Compact sparse Jacobian: the structural pattern plus an expression for its nonzeros only.

  Takes the structured path when ``expr`` is a ``VMAP`` (or a concatenation of them) over exactly
  ``wrt``, coloring the callee's small local pattern instead of the whole matrix, so the cost
  tracks the callee rather than the iteration count. Otherwise falls back to coloring the global
  pattern. The value order is the pattern's ``(rows, cols)`` order, which is not necessarily
  sorted.
  """
  structured = _sparse_jacobian_structured(expr, wrt)
  if structured is not None:
    return structured
  return sparse_jacobian_colored(expr, wrt)

sparse_hessian

sparse_hessian(expr: Expr, wrt: Expr, *, triangle: Triangle = 'full') -> SparseJacobian

Compact sparse Hessian of a scalar expression using one global star-colored JVP batch.

Source code in src/scaly/ad/sparse.py
def sparse_hessian(expr: Expr, wrt: Expr, *, triangle: Triangle = "full") -> SparseJacobian:
  """Compact sparse Hessian of a scalar expression using one global star-colored JVP batch."""
  triangle = _validate_triangle(triangle)
  if expr.size != 1:
    raise ValueError("sparse_hessian expects a scalar expression")
  gradient_expr = cse(simplify(gradient(expr, wrt).reshape((wrt.size,))))
  sparsity = _symmetrize_sparsity(jacobian_sparsity(gradient_expr, wrt))
  colors = star_coloring(sparsity)
  recovery = _star_recovery_indices(sparsity, colors)
  return _sparse_jacobian_colored(gradient_expr, wrt, sparsity, colors, recovery).triangle(triangle)

sparse_jacobian_colored

sparse_jacobian_colored(expr: Expr, wrt: Expr) -> SparseJacobian

Compact sparse Jacobian values from graph-colored compressed JVPs.

Source code in src/scaly/ad/sparse.py
def sparse_jacobian_colored(expr: Expr, wrt: Expr) -> SparseJacobian:
  """Compact sparse Jacobian values from graph-colored compressed JVPs."""
  expr = cse(expr)
  sparsity = jacobian_sparsity(expr, wrt)
  colors = column_coloring(sparsity)
  return _sparse_jacobian_colored(expr, wrt, sparsity, colors)

sparse_jacobian_reference

sparse_jacobian_reference(expr: Expr, wrt: Expr) -> SparseJacobian

Reference compact Jacobian path: build dense J and gather nonzeros.

Source code in src/scaly/ad/sparse.py
def sparse_jacobian_reference(expr: Expr, wrt: Expr) -> SparseJacobian:
  """Reference compact Jacobian path: build dense ``J`` and gather nonzeros."""

  sparsity = jacobian_sparsity(expr, wrt)
  dense = jacobian(expr, wrt)
  flat = np.asarray(sparsity.rows, dtype=np.int64) * wrt.size + np.asarray(sparsity.cols, dtype=np.int64)
  return SparseJacobian(sparsity, gather(dense, flat))