ragged_scan#
- t3toolbox.backend.common.ragged_scan(f, init, xs)#
def ragged_scan( f: typ.Callable[ [CarryType, typ.Sequence[NDArray], # len=num_inputs ], typ.Tuple[ CarryType, typ.Sequence[NDArray], # len=num_outputs ], ], init: CarryType, xs: typ.Sequence[ typ.Union[ typ.Sequence[NDArray], # len=scan_length NDArray, # shape[0]=scan_length ] ], # len=num_inputs ) -> typ.Tuple[ CarryType, typ.Tuple[ typ.Tuple[NDArray, ...], # len=scan_length ... ], # len=num_outputs, ]:
Similar to jax.lax.scan, except for ragged-sized arrays https://docs.jax.dev/en/latest/_autosummary/jax.lax.scan.html