Skip to content

Loop Ops

loop_ops

Optional namespace adapters for iteration and rematerialization.

LoopAdapter(scan, checkpoint=None) dataclass

Control-flow operations for one array namespace.

Numerical stages remain Array-API operations. This bundle dispatches only their iteration strategy; namespaces without an adapter keep eager loops. The scan must preserve carry structure and stack the body outputs.

Attributes:

Name Type Description
scan ScanFunction

Function with the scan(body, initial, inputs) contract.

checkpoint Callable[[LoopBody], LoopBody] | None

Optional transformation that recomputes body intermediates during reverse-mode differentiation.

get_loop_adapter(namespace)

Find a driver, registering JAX lazily when its namespace is supplied.

Ordinary NumPy, Torch, and CuPy calls do not import JAX. Optional backend imports stay here rather than inside numerical step kernels.

Returns:

Type Description
LoopAdapter | None

Registered driver, or None to retain the existing eager loop.

Source code in src/op_engine/loop_ops.py
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
def get_loop_adapter(namespace: object) -> LoopAdapter | None:
    """Find a driver, registering JAX lazily when its namespace is supplied.

    Ordinary NumPy, Torch, and CuPy calls do not import JAX. Optional backend
    imports stay here rather than inside numerical step kernels.

    Returns:
        Registered driver, or ``None`` to retain the existing eager loop.
    """
    name = getattr(namespace, "__name__", "")
    adapter = _LOOP_ADAPTERS.get(name)
    if adapter is None and name == "jax.numpy":
        try:
            jax = import_module("jax")
        except ImportError:
            return None
        adapter = LoopAdapter(
            scan=cast("ScanFunction", jax.lax.scan),
            checkpoint=cast("Callable[[LoopBody], LoopBody]", jax.checkpoint),
        )
        register_loop_adapter(name, adapter)
    return adapter

register_loop_adapter(namespace, adapter)

Register or replace a driver for a namespace's fully qualified name.

Parameters:

Name Type Description Default
namespace str

Namespace module name, such as "jax.numpy".

required
adapter LoopAdapter

Backend's scan and optional checkpoint operations.

required

Raises:

Type Description
TypeError

If the adapter has the wrong type.

ValueError

If the namespace name is empty.

Source code in src/op_engine/loop_ops.py
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
def register_loop_adapter(namespace: str, adapter: LoopAdapter) -> None:
    """Register or replace a driver for a namespace's fully qualified name.

    Args:
        namespace: Namespace module name, such as ``"jax.numpy"``.
        adapter: Backend's scan and optional checkpoint operations.

    Raises:
        TypeError: If the adapter has the wrong type.
        ValueError: If the namespace name is empty.
    """
    if not namespace:
        msg = "Loop adapter namespace must not be empty"
        raise ValueError(msg)
    if not isinstance(adapter, LoopAdapter):
        msg = "adapter must be a LoopAdapter"
        raise TypeError(msg)
    _LOOP_ADAPTERS[namespace] = adapter