Necesito contar la cantidad de elementos cero en matrices numpy . Conozco la función numpy.count_nonzero , pero parece que no hay un análogo para contar cero elementos.
Mis arreglos no son muy grandes (normalmente menos de 1E5 elementos) pero la operación se realiza varios millones de veces.
Por supuesto que podría usar len(arr) - np.count_nonzero(arr) , pero me pregunto si hay una forma más eficiente de hacerlo.
Aquí hay un MWE de cómo lo hago actualmente:
import numpy as np import timeit arrs = [] for _ in range(1000): arrs.append(np.random.randint(-5, 5, 10000)) def func1(): for arr in arrs: zero_els = len(arr) - np.count_nonzero(arr) print(timeit.timeit(func1, number=10))Un enfoque 2 veces más rápido sería simplemente usar np.count_nonzero np.count_nonzero() pero con la condición según sea necesario.
In [3]: arr Out[3]: array([[1, 2, 0, 3], [3, 9, 0, 4]]) In [4]: np.count_nonzero(arr==0) Out[4]: 2 In [5]:def func_cnt(): for arr in arrs: zero_els = np.count_nonzero(arr==0) # here, it counts the frequency of zeroes actually También puede usar np.where() pero es más lento que np.count_nonzero()
In [6]: np.where( arr == 0) Out[6]: (array([0, 1]), array([2, 2])) In [7]: len(np.where( arr == 0)) Out[7]: 2Eficiencia: (en orden descendente)
In [8]: %timeit func_cnt() 10 loops, best of 3: 29.2 ms per loop In [9]: %timeit func1() 10 loops, best of 3: 46.5 ms per loop In [10]: %timeit func_where() 10 loops, best of 3: 61.2 ms per loopmás aceleraciones con aceleradores
Ahora es posible lograr un aumento de velocidad de más de 3 órdenes de magnitud con la ayuda de JAX si tiene acceso a aceleradores (GPU/TPU). Otra ventaja de usar JAX es que el código NumPy necesita muy pocas modificaciones para que sea compatible con JAX. A continuación se muestra un ejemplo reproducible:
In [1]: import jax.numpy as jnp In [2]: from jax import jit # set up inputs In [3]: arrs = [] In [4]: for _ in range(1000): ...: arrs.append(np.random.randint(-5, 5, 10000)) # JIT'd function that performs the counting task In [5]: @jit ...: def func_cnt(): ...: for arr in arrs: ...: zero_els = jnp.count_nonzero(arr==0) # efficiency test In [8]: %timeit func_cnt() 15.6 µs ± 391 ns per loop (mean ± std. dev. of 7 runs, 100000 loops each)