Skip to content

Core

Expr represents a calculation whose inputs are not yet known. Function gives that calculation named inputs and outputs so that you can evaluate it, differentiate it, or generate C code. The functions guide shows how to use both.

The first sections cover modelling. Verification, rewriting, and text rendering are lower-level interfaces for inspecting or extending the compiler.

Expressions

Expr dataclass

One node of the expression graph: an op, its argument nodes and the resulting TensorType.

Nodes are immutable and interned, so building the same expression twice returns the same object. Create leaves with sym and const, then combine them with the operators, methods and builders below.

scalar

scalar() -> Expr

Return this expression with a hint to expand its procedure into straight-line scalar code.

The hint overrides the automatic size limits and applies to the entry point too. Only procedures whose values are all float64 expand, and a block hint in the same function wins. See the lowering page of How it works for the policy.

block

block() -> Expr

Return this expression with a hint to keep its procedure as loops over buffers.

The procedure is not expanded into scalar code and survives as its own C procedure. Every other optimization still runs on it.

sym(name, shape=None, *, dtype=dtypes.float64, diff=True) staticmethod

sym

Create a named symbolic input, the leaf every graph is built from.

Parameters:

Name Type Description Default
name str

the input's name, which becomes its name on any Function declaring it.

required
shape int | tuple[int, ...] | None

an int for a rank-1 shape, a tuple as given, or None for a scalar.

None
dtype DType | str

the element type, float64 unless you say otherwise.

float64
diff bool

whether derivatives with respect to this input are meaningful. Setting it to False tells AD the input is a constant parameter, so terms through it vanish.

True

const(value, *, dtype=None) staticmethod

const

Wrap an array or scalar as a constant node.

Constants are never differentiable: a derivative with respect to one is structurally zero.

ExprOp

Bases: StrEnum

The expression dialect's operation set.

A StrEnum, so an op is a proper enum value and still prints and serializes as its name. OP_INFO carries the arity, the NumPy evaluation rule and the differentiability of each.

Source code in src/scaly/ir/expr.py
class ExprOp(StrEnum):
  """The expression dialect's operation set.

  A ``StrEnum``, so an op is a proper enum value and still prints and serializes as its name.
  ``OP_INFO`` carries the arity, the NumPy evaluation rule and the differentiability of each.
  """

  INPUT = "input"
  CONST = "const"
  NEG = "neg"
  SIN = "sin"
  COS = "cos"
  TAN = "tan"
  ASIN = "asin"
  ACOS = "acos"
  ATAN = "atan"
  SINH = "sinh"
  COSH = "cosh"
  TANH = "tanh"
  ERF = "erf"
  EXP = "exp"
  LOG = "log"
  SQRT = "sqrt"
  ABS = "abs"
  FLOOR = "floor"
  CEIL = "ceil"
  ADD = "add"
  SUB = "sub"
  MUL = "mul"
  DIV = "div"
  POW = "pow"
  ATAN2 = "atan2"
  MINIMUM = "minimum"
  MAXIMUM = "maximum"
  SUM = "sum"
  RESHAPE = "reshape"
  TRANSPOSE = "transpose"
  SLICE = "slice"
  GATHER = "gather"
  SCATTER = "scatter"
  STACK = "stack"
  CONCAT = "concat"
  MATMUL = "matmul"
  CALL = "call"
  VMAP = "vmap"
  SOLVER_CALL = "solver_call"

OpInfo dataclass

What the compiler knows about one ExprOp: its arity (None for variadic), its NumPy evaluation where it has one, and whether it has derivative rules. OP_INFO holds one per op.

Source code in src/scaly/ir/expr.py
@dataclass(frozen=True, slots=True)
class OpInfo:
  """What the compiler knows about one ``ExprOp``: its arity (``None`` for variadic), its NumPy
  evaluation where it has one, and whether it has derivative rules. ``OP_INFO`` holds one per op."""

  op: ExprOp
  arity: int | None
  numpy: Callable[..., np.ndarray | np.generic] | None = None
  differentiable: bool = True

  @property
  def name(self) -> str:
    return self.op.value

Functions

Function

A named expression graph: named inputs, named outputs, and the computation between them.

Function is the unit of composition, differentiation and compilation. Calling fn(inputs) with Expr leaves puts one call node in a larger graph. Lowering may still inline a small callee when it expands a procedure into scalar code. fn.factory(...) derives a new Function carrying the requested derivatives. Calling fn(inputs) with array leaves lowers it, renders C, compiles and caches a shared library, and calls it through the pointer entry that every generated function shares.

__call__ takes the whole declared input tree as one argument and dispatches on its leaves to symbolic_call or numerical_call. Call those directly when the distinction matters. A one-leaf tree is the bare value on both sides, as described for scaly.L, so a single-output result must not be destructured.

Input and output names matter beyond display. Derivatives are requested by name, and the generated C symbols are built from them.

Use @scaly.function(...) to build one from a Python body.

input_shapes property

input_shapes: tuple[tuple[int, ...], ...]

The input leaf shapes in C-signature order.

output_shapes property

output_shapes: tuple[tuple[int, ...], ...]

The output leaf shapes in C-signature order.

input_map

input_map() -> dict[str, Expr]

The symbolic input leaves keyed by their declared names, in declaration order.

output_map

output_map() -> dict[str, Expr]

The output expressions keyed by their declared names, in declaration order.

symbolic_call

symbolic_call(inputs: SymbolicInputs) -> SymbolicOutputs

Embed a call node using the declared symbolic input and output structures.

numerical_call

numerical_call(inputs: NumericalInputs) -> NumericalOutputs

Compile and evaluate using the declared numerical input and output structures.

compile

compile() -> None

Compile ahead of the first numerical call, or reuse the cached library.

recompile

recompile() -> None

Drop the cached compiled handle and remove the on-disk cache entry for this function.

solver_stats

solver_stats(name: str | None = None) -> SolverStats

Return the latest stats for a solver reached by this compiled function.

factory

factory(name: str, inputs: Sequence[str], outputs: Sequence[str | DerivSpec], aux: Mapping[str, Sequence[str]] | None = None) -> Function

Build a new Function whose outputs mix this function's outputs and derivatives of them.

Parameters:

Name Type Description Default
name str

the new function's name, which also names its generated C symbols.

required
inputs Sequence[str]

the input names, in order. Besides the declared inputs, "fwd:<input>" names the seed of a forward derivative and "lam:<output>" the weight of an adjoint or of an aux combination.

required
outputs Sequence[str | DerivSpec]

output names and DerivSpec requests such as Grad("cost", "x"). A request's output is named {kind}_{of}_{wrt}.

required
aux Mapping[str, Sequence[str]] | None

extra outputs usable by name in outputs and in requests. Each maps a new name to a list of output names and stands for their sum weighted by the matching "lam:<output>" inputs, as in a Lagrangian.

None

Returns:

Type Description
Function

A Function with the selected inputs and the requested outputs.

Raises:

Type Description
ValueError

if a name in inputs, outputs or aux is unknown, or an aux name shadows an output.

Types

TensorType dataclass

The shape and dtype of an expression, and whether it is differentiable (diff).

Source code in src/scaly/ir/types.py
@dataclass(frozen=True, slots=True)
class TensorType:
  """The shape and dtype of an expression, and whether it is differentiable (``diff``)."""

  shape: tuple[int, ...] = ()
  dtype: DType = dtypes.float64
  diff: bool = True

  def __post_init__(self) -> None:
    _check_shape("tensor", self.shape)
    if not isinstance(self.dtype, DType):
      object.__setattr__(self, "dtype", as_dtype(self.dtype))

  @property
  def ndim(self) -> int:
    return len(self.shape)

  @property
  def size(self) -> int:
    return reduce(mul, self.shape, 1)

  @property
  def is_scalar(self) -> bool:
    return self.shape == () or self.shape == (1,) or self.size == 1

DType dataclass

Interned dtype descriptor.

See dtypes for the canonical instances. A DType compares equal to its string name, so dtypes.float64 == "float64" holds.

Source code in src/scaly/ir/types.py
@dataclass(frozen=True, slots=True)
class DType:
  """Interned dtype descriptor.

  See ``dtypes`` for the canonical instances. A ``DType`` compares equal to its string ``name``,
  so ``dtypes.float64 == "float64"`` holds.
  """

  name: str
  bits: int
  c_type: str
  is_floating: bool = False
  is_integer: bool = False
  is_bool: bool = False

  @property
  def itemsize(self) -> int:
    return self.bits // 8

  def numpy(self) -> np.dtype:
    return np.dtype(self.name)

  def __str__(self) -> str:
    return self.name

  def __eq__(self, other: object) -> bool:  # backward compat with string dtype
    if isinstance(other, DType):
      return self.name == other.name
    if isinstance(other, str):
      return self.name == other
    return NotImplemented

  def __hash__(self) -> int:
    return hash(self.name)

dtypes

Canonical interned dtype instances. Mirrors the small tinygrad-style registry.

SparsityPattern dataclass

The structurally nonzero entries of a matrix, as coordinate lists rows and cols.

The entry order is the order of the values stored alongside the pattern. It is not necessarily row- or column-major. to_csr and to_csc return the permutation into either order.

Source code in src/scaly/ir/types.py
@dataclass(frozen=True, slots=True)
class SparsityPattern:
  """The structurally nonzero entries of a matrix, as coordinate lists ``rows`` and ``cols``.

  The entry order is the order of the values stored alongside the pattern. It is not necessarily
  row- or column-major. ``to_csr`` and ``to_csc`` return the permutation into either order.
  """

  shape: tuple[int, int]
  rows: tuple[int, ...]
  cols: tuple[int, ...]

  def __post_init__(self) -> None:
    if len(self.shape) != 2:
      raise ValueError(f"sparsity shape must be rank-2, got {self.shape}")
    _check_shape("sparsity", self.shape)
    if len(self.rows) != len(self.cols):
      raise ValueError("sparsity rows and cols must have the same length")
    if any(r < 0 or r >= self.shape[0] for r in self.rows) or any(c < 0 or c >= self.shape[1] for c in self.cols):
      raise ValueError(f"sparsity indices out of bounds for shape {self.shape}")
    if len(set(zip(self.rows, self.cols))) != len(self.rows):
      raise ValueError("sparsity indices must be unique")

  @property
  def nnz(self) -> int:
    return len(self.rows)

  @staticmethod
  def empty(shape: tuple[int, int]) -> SparsityPattern:
    return SparsityPattern(shape, (), ())

  @staticmethod
  def dense(shape: tuple[int, int]) -> SparsityPattern:
    rows, cols = np.nonzero(np.ones(shape, dtype=bool))
    return SparsityPattern(shape, tuple(int(x) for x in rows), tuple(int(x) for x in cols))

  @staticmethod
  def from_mask(mask: np.ndarray) -> SparsityPattern:
    mask = np.asarray(mask, dtype=bool)
    if mask.ndim != 2:
      raise ValueError(f"sparsity mask must be rank-2, got {mask.shape}")
    rows, cols = np.nonzero(mask)
    shape = (int(mask.shape[0]), int(mask.shape[1]))
    return SparsityPattern(shape, tuple(int(x) for x in rows), tuple(int(x) for x in cols))

  @staticmethod
  def from_csr(shape: tuple[int, int], row_ptr: Sequence[int], col_ind: Sequence[int]) -> SparsityPattern:
    row_ptr = tuple(int(x) for x in row_ptr)
    col_ind = tuple(int(x) for x in col_ind)
    _check_compressed_ptr("row_ptr", row_ptr, shape[0], len(col_ind))
    if any(c < 0 or c >= shape[1] for c in col_ind):
      raise ValueError(f"CSR column indices out of bounds for shape {shape}")
    rows = tuple(r for r in range(shape[0]) for _ in range(row_ptr[r + 1] - row_ptr[r]))
    return SparsityPattern(shape, rows, col_ind)

  @staticmethod
  def from_csc(shape: tuple[int, int], col_ptr: Sequence[int], row_ind: Sequence[int]) -> SparsityPattern:
    col_ptr = tuple(int(x) for x in col_ptr)
    row_ind = tuple(int(x) for x in row_ind)
    _check_compressed_ptr("col_ptr", col_ptr, shape[1], len(row_ind))
    if any(r < 0 or r >= shape[0] for r in row_ind):
      raise ValueError(f"CSC row indices out of bounds for shape {shape}")
    cols = tuple(c for c in range(shape[1]) for _ in range(col_ptr[c + 1] - col_ptr[c]))
    return SparsityPattern(shape, row_ind, cols)

  def to_mask(self) -> np.ndarray:
    mask = np.zeros(self.shape, dtype=bool)
    mask[list(self.rows), list(self.cols)] = True
    return mask

  def to_csr(self) -> tuple[tuple[int, ...], tuple[int, ...], tuple[int, ...]]:
    """``(row_ptr, col_ind, val_perm)``: ``val_perm[k]`` is the COO position of CSR slot ``k``, so
    ``values_csr[k] = values[val_perm[k]]`` pairs a compact COO-ordered value buffer with the CSR
    indices. The COO order is arbitrary, for example piece by piece for a mapped sparse Jacobian."""
    order = sorted(range(self.nnz), key=lambda i: (self.rows[i], self.cols[i]))
    row_ptr = [0] * (self.shape[0] + 1)
    for i in order:
      row_ptr[self.rows[i] + 1] += 1
    for r in range(self.shape[0]):
      row_ptr[r + 1] += row_ptr[r]
    return tuple(row_ptr), tuple(self.cols[i] for i in order), tuple(order)

  def to_csc(self) -> tuple[tuple[int, ...], tuple[int, ...], tuple[int, ...]]:
    """``(col_ptr, row_ind, val_perm)``, the compressed-column counterpart of ``to_csr``."""
    order = sorted(range(self.nnz), key=lambda i: (self.cols[i], self.rows[i]))
    col_ptr = [0] * (self.shape[1] + 1)
    for i in order:
      col_ptr[self.cols[i] + 1] += 1
    for c in range(self.shape[1]):
      col_ptr[c + 1] += col_ptr[c]
    return tuple(col_ptr), tuple(self.rows[i] for i in order), tuple(order)

to_csr

to_csr() -> tuple[tuple[int, ...], tuple[int, ...], tuple[int, ...]]

(row_ptr, col_ind, val_perm): val_perm[k] is the COO position of CSR slot k, so values_csr[k] = values[val_perm[k]] pairs a compact COO-ordered value buffer with the CSR indices. The COO order is arbitrary, for example piece by piece for a mapped sparse Jacobian.

Source code in src/scaly/ir/types.py
def to_csr(self) -> tuple[tuple[int, ...], tuple[int, ...], tuple[int, ...]]:
  """``(row_ptr, col_ind, val_perm)``: ``val_perm[k]`` is the COO position of CSR slot ``k``, so
  ``values_csr[k] = values[val_perm[k]]`` pairs a compact COO-ordered value buffer with the CSR
  indices. The COO order is arbitrary, for example piece by piece for a mapped sparse Jacobian."""
  order = sorted(range(self.nnz), key=lambda i: (self.rows[i], self.cols[i]))
  row_ptr = [0] * (self.shape[0] + 1)
  for i in order:
    row_ptr[self.rows[i] + 1] += 1
  for r in range(self.shape[0]):
    row_ptr[r + 1] += row_ptr[r]
  return tuple(row_ptr), tuple(self.cols[i] for i in order), tuple(order)

to_csc

to_csc() -> tuple[tuple[int, ...], tuple[int, ...], tuple[int, ...]]

(col_ptr, row_ind, val_perm), the compressed-column counterpart of to_csr.

Source code in src/scaly/ir/types.py
def to_csc(self) -> tuple[tuple[int, ...], tuple[int, ...], tuple[int, ...]]:
  """``(col_ptr, row_ind, val_perm)``, the compressed-column counterpart of ``to_csr``."""
  order = sorted(range(self.nnz), key=lambda i: (self.cols[i], self.rows[i]))
  col_ptr = [0] * (self.shape[1] + 1)
  for i in order:
    col_ptr[self.cols[i] + 1] += 1
  for c in range(self.shape[1]):
    col_ptr[c + 1] += col_ptr[c]
  return tuple(col_ptr), tuple(self.rows[i] for i in order), tuple(order)

as_dtype

as_dtype(value: DType | str | None) -> DType

Coerce a dtype name, a DType or None into a canonical DType.

None means float64, scaly's default.

Source code in src/scaly/ir/types.py
def as_dtype(value: DType | str | None) -> DType:
  """Coerce a dtype name, a ``DType`` or ``None`` into a canonical ``DType``.

  ``None`` means ``float64``, scaly's default.
  """
  if value is None:
    return dtypes.float64
  if isinstance(value, DType):
    return value
  if isinstance(value, str):
    return dtypes.from_name(value)
  raise TypeError(f"cannot interpret {value!r} as a DType")

Builders

dot

dot(x: Any, y: Any) -> Expr

Scalar product of two expressions with the same number of entries, whatever their shapes.

Source code in src/scaly/ir/expr.py
def dot(x: Any, y: Any) -> Expr:
  """Scalar product of two expressions with the same number of entries, whatever their shapes."""
  x, y = as_expr(x), as_expr(y)
  if x.size != y.size:
    raise ValueError(f"dot size mismatch: {x.shape} has {x.size} entries, {y.shape} has {y.size}")
  return (x.vec() * y.vec()).sum()

sumsqr

sumsqr(x: Any) -> Expr

Sum of squares of every entry: dot(x, x), a scalar.

Source code in src/scaly/ir/expr.py
def sumsqr(x: Any) -> Expr:
  """Sum of squares of every entry: ``dot(x, x)``, a scalar."""
  x = as_expr(x)
  return dot(x, x)

norm_2

norm_2(x: Any) -> Expr

Euclidean norm of every entry, as a scalar.

Source code in src/scaly/ir/expr.py
def norm_2(x: Any) -> Expr:
  """Euclidean norm of every entry, as a scalar."""
  return sumsqr(x).sqrt()

stack

stack(xs: Iterable[Any], *, axis: int = 0) -> Expr

Join equally shaped expressions along a new axis, adding one to the rank.

concat

concat(xs: Iterable[Any], *, axis: int = 0) -> Expr

Join expressions along an existing axis. All other dimensions must agree.

split

split(x: Any, sections: int | Iterable[int], *, axis: int = 0) -> tuple[Expr, ...]

Split x along axis, into sections equal parts or at the given boundaries.

Source code in src/scaly/ir/expr.py
def split(x: Any, sections: int | Iterable[int], *, axis: int = 0) -> tuple[Expr, ...]:
  """Split ``x`` along ``axis``, into ``sections`` equal parts or at the given boundaries."""
  x = as_expr(x)
  axis = axis if axis >= 0 else axis + len(x.shape)
  if axis < 0 or axis >= len(x.shape):
    raise ValueError(f"split axis {axis} out of bounds for shape {x.shape}")
  if isinstance(sections, int):
    if sections <= 0 or x.shape[axis] % sections != 0:
      raise ValueError(f"cannot split axis of length {x.shape[axis]} into {sections} equal sections")
    sizes = (x.shape[axis] // sections,) * sections
  else:
    sizes = tuple(int(s) for s in sections)
    if any(s < 0 for s in sizes):
      raise ValueError(f"split sizes {sizes} cannot contain negative entries")
    if sum(sizes) != x.shape[axis]:
      raise ValueError(f"split sizes {sizes} do not sum to axis length {x.shape[axis]}")
  ret: list[Expr] = []
  start = 0
  for size in sizes:
    index = [slice(None)] * len(x.shape)
    index[axis] = slice(start, start + size)
    ret.append(x[tuple(index)])
    start += size
  return tuple(ret)

vec

vec(x: Any) -> Expr

Flatten x to rank 1 in row-major order.

Source code in src/scaly/ir/expr.py
def vec(x: Any) -> Expr:
  """Flatten ``x`` to rank 1 in row-major order."""
  return as_expr(x).vec()

gather

gather(x: Any, indices: Any) -> Expr

Read x at flat indices. The result takes the shape of indices.

scatter

scatter(values: Any, indices: Any, shape: int | tuple[int, ...]) -> Expr

Place values at flat indices in a zero tensor of shape.

Repeated indices accumulate. indices must have as many entries as values.

atan2

atan2(y: Any, x: Any) -> Expr

Two-argument arctangent, elementwise: the angle of the point (x, y).

Source code in src/scaly/ir/expr.py
def atan2(y: Any, x: Any) -> Expr:
  """Two-argument arctangent, elementwise: the angle of the point ``(x, y)``."""
  return binary(ExprOp.ATAN2, as_expr(y), as_expr(x))

minimum

minimum(x: Any, y: Any) -> Expr

Elementwise minimum. Non-smooth, so the result is marked non-differentiable.

Source code in src/scaly/ir/expr.py
def minimum(x: Any, y: Any) -> Expr:
  """Elementwise minimum. Non-smooth, so the result is marked non-differentiable."""
  return binary(ExprOp.MINIMUM, as_expr(x), as_expr(y))

maximum

maximum(x: Any, y: Any) -> Expr

Elementwise maximum. Non-smooth, so the result is marked non-differentiable.

Source code in src/scaly/ir/expr.py
def maximum(x: Any, y: Any) -> Expr:
  """Elementwise maximum. Non-smooth, so the result is marked non-differentiable."""
  return binary(ExprOp.MAXIMUM, as_expr(x), as_expr(y))

vmap

vmap(callee: Any, length: int, inputs: Any, output: int = 0) -> Expr

Create an ExprOp.VMAP node: length independent calls of callee whose i-th argument list is sliced out of outer tensors with per-input (start, stride) strides.

inputs is a sequence ordered to match callee.inputs, or a mapping from callee input name to the same entries. An entry is usually a bare outer tensor: one of size length * formal.size is cut into length contiguous chunks, one of size formal.size is broadcast to every iteration. The explicit (outer_tensor, start, stride) tuple covers overlapping or offset windows: the i-th iteration reads outer[start + i*stride : start + i*stride + formal.size], and stride=0 broadcasts the same slice every iteration.

Only the outer tensors must be rank-1. Callee formals and outputs may be rank-2 (as well as scalar or rank-1). Each iteration reads a flat slice of formal.size values and the produced node has shape (length * callee.outputs[output].size,), with iteration outputs concatenated flat.

PyTorch weights

load_torch_state_dict

load_torch_state_dict(path: str | Path) -> dict[str, np.ndarray]

Load a simple PyTorch state_dict zip checkpoint without importing torch.

Supports the modern torch.save(state_dict, ...) zip layout: data.pkl stores tensor metadata and data/<n> stores raw CPU storage bytes. This is intentionally narrow. Unsupported pickle globals fail loudly instead of being materialized.

Source code in src/scaly/utils/torch_state_dict.py
def load_torch_state_dict(path: str | Path) -> dict[str, np.ndarray]:
  """Load a simple PyTorch ``state_dict`` zip checkpoint without importing torch.

  Supports the modern ``torch.save(state_dict, ...)`` zip layout: ``data.pkl``
  stores tensor metadata and ``data/<n>`` stores raw CPU storage bytes. This is
  intentionally narrow. Unsupported pickle globals fail loudly instead of being
  materialized.
  """
  p = Path(path)
  storage_blobs: dict[str, bytes] = {}

  def rebuild_tensor(storage: tuple[Any, ...], storage_offset: int, size: tuple[int, ...], stride: tuple[int, ...], *args: Any) -> np.ndarray:
    del args
    if len(storage) < 5 or storage[0] != "storage":
      raise TypeError(f"unsupported torch storage persistent id: {storage!r}")
    dtype = _torch_dtype(storage[1])
    key, numel = str(storage[2]), int(storage[4])
    raw = storage_blobs[key]
    base = np.frombuffer(raw, dtype=dtype, count=numel)
    shape = tuple(int(x) for x in size)
    strides = tuple(int(x) * dtype.itemsize for x in stride)
    if not shape:
      return np.asarray(base[int(storage_offset)], dtype=dtype)
    return np.lib.stride_tricks.as_strided(base[int(storage_offset) :], shape=shape, strides=strides).copy()

  def rebuild_parameter(data: np.ndarray, *args: Any) -> np.ndarray:
    del args
    return data

  class Parameter:
    tensor: np.ndarray

    def __setstate__(self, state: tuple[np.ndarray, ...]) -> None:
      self.tensor = state[0]

  class TorchPickle(pickle.Unpickler):
    def find_class(self, module: str, name: str) -> Any:
      if module == "torch" and name in _TORCH_STORAGE_DTYPES:
        return _TORCH_STORAGE_DTYPES[name]
      if module == "torch" and name in _TORCH_DTYPES:
        return _TORCH_DTYPES[name]
      if module in {"torch._utils", "torch._tensor"} and name in {"_rebuild_tensor", "_rebuild_tensor_v2"}:
        return rebuild_tensor
      if module in {"torch._utils", "torch.nn.parameter"} and name in {"_rebuild_parameter", "_rebuild_parameter_with_state"}:
        return rebuild_parameter
      if module == "torch.nn.parameter" and name == "Parameter":
        return Parameter
      if module in {"collections", "numpy", "numpy.core.multiarray", "_codecs", "builtins"}:
        return super().find_class(module, name)
      raise pickle.UnpicklingError(f"unsupported global in torch checkpoint: {module}.{name}")

    def persistent_load(self, pid: Any) -> Any:
      return pid

  if not zipfile.is_zipfile(p):
    raise ValueError(f"unsupported legacy non-zip PyTorch checkpoint: {p}")
  with zipfile.ZipFile(p) as zf:
    names = zf.namelist()
    base = names[0].split("/", 1)[0]
    storage_blobs = {name.rsplit("/", 1)[-1]: zf.read(name) for name in names if name.startswith(f"{base}/data/")}
    with zf.open(f"{base}/data.pkl") as data:
      state = TorchPickle(data).load()
  if not isinstance(state, dict):
    raise TypeError(f"expected a state_dict in {p}, got {type(state).__name__}")
  return {k: v.tensor if isinstance(v, Parameter) else v for k, v in state.items()}

Verification and rewriting

verify_expr

verify_expr(root: Expr | Iterable[Expr], spec: 'Spec | None' = None) -> None

Topologically walk the DAG below root and raise on the first violation.

spec defaults to spec_expr, the full expression-dialect contract.

Source code in src/scaly/ir/expr_spec.py
def verify_expr(root: Expr | Iterable[Expr], spec: "Spec | None" = None) -> None:
  """Topologically walk the DAG below ``root`` and raise on the first violation.

  ``spec`` defaults to ``spec_expr``, the full expression-dialect contract.
  """
  if spec is None:
    spec = spec_expr
  outputs = (root,) if isinstance(root, Expr) else tuple(root)
  for node in topo(outputs):
    result = spec.check(node)
    if result is not None:
      rule, diag = result
      label = node.name or f"%{node.id}"
      raise VerifyError(f"verify_expr: node {label} op={ExprOp(node.op).value} failed rule {rule.description!r}: {diag}")

Spec

Bases: Generic[Node, Op]

A verifier: a set of Rules, indexed by op so checking a node touches only its own.

Rules with op=None apply to every node. Both dialects build their specs from this class.

Source code in src/scaly/ir/spec.py
class Spec(Generic[Node, Op]):
  """A verifier: a set of ``Rule``s, indexed by op so checking a node touches only its own.

  Rules with ``op=None`` apply to every node. Both dialects build their specs from this class.
  """

  def __init__(self, rules: Iterable[Rule[Node, Op]]) -> None:
    self.any: list[Rule[Node, Op]] = []
    self.by_op: dict[Op, list[Rule[Node, Op]]] = defaultdict(list)
    for rule in rules:
      if rule.op is None:
        self.any.append(rule)
      else:
        self.by_op[rule.op].append(rule)

  def candidates(self, op: Op) -> Iterable[Rule[Node, Op]]:
    yield from self.any
    yield from self.by_op.get(op, ())

  def check(self, node: Node) -> tuple[Rule[Node, Op], str] | None:
    for rule in self.candidates(getattr(node, "op")):
      diag = rule.check(node)
      if diag is not None:
        return rule, diag
    return None

  def merge(self, *others: Spec[Node, Op]) -> Spec[Node, Op]:
    rules = list(self.any)
    for op_rules in self.by_op.values():
      rules.extend(op_rules)
    for other in others:
      rules.extend(other.any)
      for op_rules in other.by_op.values():
        rules.extend(op_rules)
    return Spec(rules)

Rule dataclass

Bases: Generic[Node, Op]

One verifier check: check returns an error message for a bad node, or None.

The rule applies to nodes whose op is op, or to every node when op is None.

Source code in src/scaly/ir/spec.py
@dataclass(frozen=True, slots=True)
class Rule(Generic[Node, Op]):
  """One verifier check: ``check`` returns an error message for a bad node, or ``None``.

  The rule applies to nodes whose op is ``op``, or to every node when ``op`` is ``None``.
  """

  op: Op | None
  description: str
  check: Callable[[Node], str | None]

  def applies(self, node: Node) -> bool:
    return self.op is None or getattr(node, "op") == self.op

VerifyError

Bases: Exception

Raised by verify_expr and verify_program at the first node failing a rule.

The message names the node, its op and the rule it broke. Both dialects raise this one type.

Source code in src/scaly/ir/spec.py
class VerifyError(Exception):
  """Raised by ``verify_expr`` and ``verify_program`` at the first node failing a rule.

  The message names the node, its op and the rule it broke. Both dialects raise this one type.
  """

Pattern dataclass

One rewrite: nodes with op op (any op when None) that satisfy predicate become replacement(node).

Source code in src/scaly/ir/match.py
@dataclass(frozen=True, slots=True)
class Pattern[Node: _HasOpArgs]:
  """One rewrite: nodes with op ``op`` (any op when ``None``) that satisfy ``predicate`` become ``replacement(node)``."""

  op: Hashable | None
  predicate: Callable[[Node], bool]
  replacement: Callable[[Node], Node]

  def matches(self, node: Node) -> bool:
    return (self.op is None or node.op == self.op) and self.predicate(node)

PatternMatcher

An op-indexed set of rewrite Patterns.

Indexing by op means a node is only tested against patterns that could match it, which is what keeps rewriting cheap as the pattern set grows.

Source code in src/scaly/ir/match.py
class PatternMatcher[Node: _HasOpArgs]:
  """An op-indexed set of rewrite ``Pattern``s.

  Indexing by op means a node is only tested against patterns that could match it, which is what
  keeps rewriting cheap as the pattern set grows.
  """

  def __init__(self, patterns: Iterable[Pattern[Node]]):
    self.any: list[Pattern[Node]] = []
    self.by_op: dict[Hashable, list[Pattern[Node]]] = defaultdict(list)
    for pattern in patterns:
      if pattern.op is None:
        self.any.append(pattern)
      else:
        self.by_op[pattern.op].append(pattern)

  def candidates(self, op: Hashable) -> Iterable[Pattern[Node]]:
    yield from self.any
    yield from self.by_op.get(op, ())

  def rewrite(self, node: Node) -> Node | None:
    for pattern in self.candidates(node.op):
      if pattern.matches(node):
        ret = pattern.replacement(node)
        if ret is not node:
          return ret
    return None

rewrite

rewrite(root: Node, patterns: Iterable[Pattern[Node]] | PatternMatcher[Node], *, rebuild: Callable[[Node, tuple[Node, ...]], Node] | None = None, fixpoint: bool = True, revisit: bool = False, max_steps: int = 1000000, memo: dict[Node, Node] | None = None) -> Node

Apply patterns across a graph, bottom up, and return the rewritten root.

Iterative over an explicit stack, so depth in the graph never becomes depth on the Python stack. Each node is visited once with its already-rewritten children memoized by identity, so shared subgraphs stay shared. fixpoint retries the matcher on a node until nothing fires. revisit instead walks into a replacement's subgraph so nested rewrites collapse in one pass. max_steps bounds the total number of replacements. rebuild defaults to the expression adapter rebuild_expr. Program adapters may reuse memo across roots with identical patterns, rebuild, and options. Expression nodes do not support caller memoization because their lowering hints depend on traversal provenance.

Source code in src/scaly/ir/match.py
def rewrite[Node: _HasOpArgs](
  root: Node,
  patterns: Iterable[Pattern[Node]] | PatternMatcher[Node],
  *,
  rebuild: Callable[[Node, tuple[Node, ...]], Node] | None = None,
  fixpoint: bool = True,
  revisit: bool = False,
  max_steps: int = 1_000_000,
  memo: dict[Node, Node] | None = None,
) -> Node:
  """Apply ``patterns`` across a graph, bottom up, and return the rewritten root.

  Iterative over an explicit stack, so depth in the graph never becomes depth on the Python
  stack. Each node is visited once with its already-rewritten children memoized by identity, so
  shared subgraphs stay shared. ``fixpoint`` retries the matcher on a node until nothing fires.
  ``revisit`` instead walks into a replacement's subgraph so nested rewrites collapse in one
  pass. ``max_steps`` bounds the total number of replacements. ``rebuild`` defaults to the
  expression adapter ``rebuild_expr``. Program adapters may reuse ``memo`` across roots with
  identical patterns, rebuild, and options. Expression nodes do not support caller memoization
  because their lowering hints depend on traversal provenance.
  """
  if memo is not None and isinstance(root, Expr):
    raise ValueError("caller memoization is only supported for Program nodes")
  matcher = patterns if isinstance(patterns, PatternMatcher) else PatternMatcher(patterns)
  rebuild = rebuild or cast(Callable[[Node, tuple[Node, ...]], Node], rebuild_expr)
  done: dict[int, Node] = {}
  lowerings: dict[int, Lowering] = {}
  forward: dict[int, Node] = {}
  alive: list[Node] = []
  steps = 0
  stack: list[tuple[Node, bool]] = [(root, False)]
  while stack:
    node, ready = stack.pop()
    if id(node) in done:
      continue
    if memo is not None and node in memo:
      done[id(node)] = memo[node]
      continue
    if not ready:
      stack.append((node, True))
      stack.extend((a, False) for a in reversed(node.args) if id(a) not in done)
      continue
    if (target := forward.pop(id(node), None)) is not None:
      if id(target) not in done:
        raise RuntimeError(f"rewrite replacement at {node.op} contains the node it replaces")
      done[id(node)] = done[id(target)]
      if isinstance(node, Expr):
        lowerings[id(node)] = lowerings[id(target)]
      if memo is not None:
        memo[node] = done[id(node)]
      continue
    args = tuple(done[id(a)] for a in node.args)
    cur = node if all(a is b for a, b in zip(args, node.args, strict=True)) else rebuild(node, args)
    lowering = _combined_lowering(node.lowering, (lowerings[id(arg)] for arg in node.args)) if isinstance(node, Expr) else "auto"
    new = matcher.rewrite(cur)
    while new is not None:
      steps += 1
      if steps > max_steps:
        names = sorted({getattr(p.replacement, "__qualname__", repr(p.replacement)) for p in matcher.candidates(cur.op)})
        raise RuntimeError(f"rewrite exceeded {max_steps} steps at {cur.op} with patterns {names}")
      if isinstance(new, Expr):
        new = cast(Node, _apply_lowering(new, lowering, preserve_identity=id(new) in lowerings))
      alive.append(new)
      if revisit:
        forward[id(node)] = new
        stack.extend(((node, True), (new, False)))
        break
      cur = new
      new = matcher.rewrite(cur) if fixpoint else None
    else:
      done[id(node)] = cur
      if isinstance(node, Expr):
        lowerings[id(node)] = lowering
        lowerings[id(cur)] = lowering
      if memo is not None:
        memo[node] = cur
  return done[id(root)]

simplify

simplify(expr: Expr) -> Expr

Apply algebraic identities and constant folding until the graph stops changing.

Covers the shared arithmetic identities of passes/arith.py (neutral elements, zero annihilation, x - x, negation normalization, small constant powers), folding of all-constant subgraphs, identity reshape, transpose and gather, slice of slice, slice of stack, A.T @ v as v @ A (and v @ A.T as A @ v), and a matmul with an all-ones vector as sums.

Source code in src/scaly/passes/expr.py
def simplify(expr: Expr) -> Expr:
  """Apply algebraic identities and constant folding until the graph stops changing.

  Covers the shared arithmetic identities of ``passes/arith.py`` (neutral elements, zero
  annihilation, ``x - x``, negation normalization, small constant powers), folding of all-constant
  subgraphs, identity reshape, transpose and gather, slice of slice, slice of stack, ``A.T @ v`` as
  ``v @ A`` (and ``v @ A.T`` as ``A @ v``), and a matmul with an all-ones vector as sums.
  """
  return rewrite(expr, SIMPLIFY_PATTERNS, fixpoint=False, revisit=True, max_steps=100_000)

cse

cse(expr: Expr) -> Expr

Merge structurally equal subgraphs so each distinct computation appears once.

Source code in src/scaly/passes/expr.py
def cse(expr: Expr) -> Expr:
  """Merge structurally equal subgraphs so each distinct computation appears once."""
  return cse_many([expr])[0]

cse_many

cse_many(outputs: Iterable[Expr]) -> tuple[Expr, ...]

Common-subexpression elimination across several outputs at once.

Sharing is found between the outputs as well as inside each, which is the point of asking for several derivatives from one Function.factory call.

Source code in src/scaly/passes/expr.py
def cse_many(outputs: Iterable[Expr]) -> tuple[Expr, ...]:
  """Common-subexpression elimination across several outputs at once.

  Sharing is found *between* the outputs as well as inside each, which is the point of asking
  for several derivatives from one ``Function.factory`` call.
  """
  outputs = tuple(outputs)
  memo: dict[tuple[Any, ...], Expr] = {}
  replacements: dict[int, Expr] = {}
  for node in topo(outputs):
    cur = _replace_args(node, replacements)
    key = _structural_key(cur)
    cur = memo.setdefault(key, cur)
    replacements[node.id] = cur
  return tuple(replacements[out.id] for out in outputs)

Text

format_expr

format_expr(outputs: Expr | Iterable[Expr]) -> str

Render a graph as a topological dump with stable %0, %1, ... names.

For the stable assembly form meant for diffs and tests, use render_expr_assembly.

Source code in src/scaly/ir/expr.py
def format_expr(outputs: Expr | Iterable[Expr]) -> str:
  """Render a graph as a topological dump with stable ``%0``, ``%1``, ... names.

  For the stable assembly form meant for diffs and tests, use ``render_expr_assembly``.
  """
  outs = (outputs,) if isinstance(outputs, Expr) else tuple(outputs)
  nodes = topo(outs)
  loc = {e.id: i for i, e in enumerate(nodes)}
  lines: list[str] = []
  for i, e in enumerate(nodes):
    lhs = f"%{i}"
    if e.op == ExprOp.INPUT:
      rhs = f"input {e.name}"
    elif e.op == ExprOp.CONST:
      assert e.value is not None
      rhs = f"const {np.array2string(e.value, threshold=6)}"
    elif e.op == ExprOp.CALL:
      callee = e.attrs["callee"]
      rhs = f"call {callee.name}[{e.attrs['output']}]({', '.join(f'%{loc[a.id]}' for a in e.args)})"
    elif e.op == ExprOp.VMAP:
      callee = e.attrs["callee"]
      length = e.attrs["length"]
      slice_size = e.attrs["slice_size"]
      bindings = ", ".join(f"%{loc[a.id]}[{s}::{st}]" for a, s, st in zip(e.args, e.attrs["starts"], e.attrs["strides"], strict=True))
      rhs = f"vmap[{length}x{slice_size}] {callee.name}[{e.attrs['output']}]({bindings})"
    else:
      rhs = f"{ExprOp(e.op).value}({', '.join(f'%{loc[a.id]}' for a in e.args)})"
    lines.append(f"{lhs} = {rhs} : {e.type.dtype}{e.shape}")
  lines.append("outputs " + ", ".join(f"%{loc[e.id]}" for e in outs))
  return "\n".join(lines)

render_expr_assembly

render_expr_assembly(obj: Function | Expr | Iterable[Expr], *, name: str | None = None) -> str

Render the expression Expr dialect as a compact SSA assembly listing.

When given a Function, the listing is a expr.module containing every transitive expression callee body before the requested function. This mirrors what C rendering eventually needs.

Source code in src/scaly/ir/text.py
def render_expr_assembly(obj: Function | Expr | Iterable[Expr], *, name: str | None = None) -> str:
  """Render the expression ``Expr`` dialect as a compact SSA assembly listing.

  When given a ``Function``, the listing is a ``expr.module`` containing every transitive expression
  callee body before the requested function. This mirrors what C rendering eventually needs.
  """
  # The Function case is duck-typed on the surface ``_render_function_module`` actually uses:
  # ``ir`` is below ``function`` in the import-layer order, so it cannot import ``Function`` at runtime.
  # It is tested first so an object that is both a Function and iterable still renders as an
  # ``expr.module`` rather than a bare region.
  if all(hasattr(obj, attr) for attr in ("outputs", "name", "input_names", "output_names")):
    return _render_function_module(cast("Function", obj))
  if isinstance(obj, Expr):
    return _render_expr_region((obj,), name=name)
  if isinstance(obj, Iterable):
    return _render_expr_region(tuple(cast("Iterable[Expr]", obj)), name=name)
  raise TypeError(f"render_expr_assembly expects a Function, an Expr, or an iterable of Expr, got {type(obj).__name__}")

render_program_assembly

render_program_assembly(root: ProgramNode) -> str

Render Program IR as structured assembly with explicit control-flow constructs.

Source code in src/scaly/ir/text.py
def render_program_assembly(root: ProgramNode) -> str:
  """Render Program IR as structured assembly with explicit control-flow constructs."""
  lines: list[str] = []
  _render_program_node(root, 0, lines)
  return "\n".join(lines)