import dask.array as da
import flox
import numpy as np
# works
assert ((3, 2),) == flox.rechunk_for_blockwise(da.zeros(5, chunks=(1,)), 0, np.array([1, 1, 1, 2, 2]))[1].chunks
# fails
assert ((3, 2),) == flox.rechunk_for_blockwise(da.zeros(5), 0, np.array([1, 1, 1, 2, 2]))[1].chunks
Example: