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.

Parameters:
  • f (Callable[[CarryType, Sequence[NDArray]], Tuple[CarryType, Sequence[NDArray]]])

  • init (CarryType)

  • xs (Sequence[Union[Sequence[NDArray], NDArray]])

Return type:

Tuple[CarryType, Tuple[NDArray, Ellipsis]]