numpy_scan#
- t3toolbox.backend.common.numpy_scan(f, init, xs)#
def numpy_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[ NDArray, # shape[0]=scan_length ... ], # len=num_outputs, ]:
Similar to jax.lax.scan, except returns numpy arrays instead of jax arrays.