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