Module gabenet.sugar

Expand source code
from typing import Callable

from jax import jit, random
from haiku import TransformedWithState


def scannable(
    transform_fn: TransformedWithState | Callable,
) -> Callable[[tuple, None], tuple]:
    """Transform signature of Haiku stateful-function compatible with jax scan."""
    if isinstance(transform_fn, TransformedWithState):
        transform_fn = transform_fn.apply
    _transform_fn = jit(transform_fn)  # type: ignore

    def wrapped(carry, _):
        params, state, key, *args = carry
        rng_key, key = random.split(key)
        f_out, next_state = _transform_fn(params, state, rng_key, *args)
        return (params, next_state, key, *args), f_out

    return wrapped

Functions

def scannable(transform_fn: Union[haiku._src.transform.TransformedWithState, Callable]) ‑> Callable[[tuple, None], tuple]

Transform signature of Haiku stateful-function compatible with jax scan.

Expand source code
def scannable(
    transform_fn: TransformedWithState | Callable,
) -> Callable[[tuple, None], tuple]:
    """Transform signature of Haiku stateful-function compatible with jax scan."""
    if isinstance(transform_fn, TransformedWithState):
        transform_fn = transform_fn.apply
    _transform_fn = jit(transform_fn)  # type: ignore

    def wrapped(carry, _):
        params, state, key, *args = carry
        rng_key, key = random.split(key)
        f_out, next_state = _transform_fn(params, state, rng_key, *args)
        return (params, next_state, key, *args), f_out

    return wrapped