Como se indica en el título, quiero eliminar partes de una matriz 1D que tienen ceros consecutivos y una longitud igual o superior a un umbral .
Produje la solución que se muestra en el siguiente MRE:
import numpy as np THRESHOLD = 4 a = np.array((1,1,0,1,0,0,0,0,1,1,0,0,0,1,0,0,0,0,0,1)) print("Input: " + str(a)) # Find the indices of the parts that meet threshold requirement gaps_above_threshold_inds = np.where(np.diff(np.nonzero(a)[0]) - 1 >= THRESHOLD)[0] # Delete these parts from array for idx in gaps_above_threshold_inds: a = np.delete(a, list(range(np.nonzero(a)[0][idx] + 1, np.nonzero(a)[0][idx + 1]))) print("Output: " + str(a))Producción:
Input: [1 1 0 1 0 0 0 0 1 1 0 0 0 1 0 0 0 0 0 1] Output: [1 1 0 1 1 1 0 0 0 1 1]¿Existe una forma menos complicada y más eficiente de hacer esto en una matriz numpy?
Basado en los comentarios de @mozway, estoy editando mi pregunta proporcionando más información.
Básicamente, el dominio del problema es:
Mi objetivo es eliminar las partes cero por encima de un umbral de longitud como ya he dicho.
Con respecto a mi primera preocupación sobre el manejo eficiente de numpy , la solución de @mathfux es realmente excelente y básicamente lo que estaba buscando. Por eso acepté este.
Sin embargo, el enfoque de @Jérôme Richard responde a mi segunda pregunta y presenta una solución de muy alto rendimiento; realmente útil si el conjunto de datos es extremadamente grande.
¡Gracias por sus excelentes respuestas!
np.delete crea una nueva matriz cada vez que se llama, lo cual es muy ineficiente. Una solución más rápida es almacenar todo el valor para mantener en una máscara/matriz booleana y luego filtrar la matriz de entrada a la vez. Sin embargo, es probable que esto aún requiera un bucle de Python puro si se hace solo con Numpy. Una solución más simple y rápida es usar Numba (o Cython) para hacer eso. Aquí hay una implementación:
import numpy as np import numba as nb @nb.njit('int_[:](int_[:], int_)') def filterZeros(arr, threshold): n = len(arr) res = np.empty(n, dtype=arr.dtype) count = 0 j = 0 for i in range(n): if arr[i] == 0: count += 1 else: if count >= threshold: j -= count count = 0 res[j] = arr[i] j += 1 if n > 0 and arr[n-1] == 0 and count >= threshold: j -= count return res[0:j] a = np.array((1,1,0,1,0,0,0,0,1,1,0,0,0,1,0,0,0,0,0,1)) a = filterZeros(a, 4) print("Output: " + str(a))Aquí está el resultado con una matriz binaria aleatoria que contiene 100_000 elementos en mi máquina:
Reference implementation: 5982 ms Mozway's solution: 23.4 ms This implementation: 0.11 ms Así, la solución es unas 54381 veces más rápida que la solución inicial y 212 veces más rápida que la de Mozway. El código puede ser incluso ~30 % más rápido trabajando en el lugar (destruyendo la matriz de entrada) y diciéndole a Numba que la matriz es contigua en la memoria (usando ::1 en lugar de : ).
También es posible encontrar diferencias de elementos distintos de cero, corregir los que superan el umbral y reconstruir una secuencia de forma correcta.
def numpy_fix(a): # STEP 1. find indices of nonzero items: [0 1 3 8 9 13 19] idx = np.flatnonzero(a) # STEP 2. Find differences along these indices (also insert a leading zero): [0 1 2 5 1 4 6] df = np.diff(idx, prepend=0) # STEP 3. Fix differences of indices larger than THRESHOLD: [0 1 2 1 1 4 1] df[df>THRESHOLD] = 1 # STEP 4. Given differences on indices, reconstruct indices themselves: [0 1 3 4 5 9 10] cs = np.cumsum(df) z = np.zeros(cs[-1]+1, dtype=int) # create a list of zeros z[cs] = 1 #pad it with ones within indices found return z >>> numpy_fix(a) array([1, 1, 0, 1, 1, 1, 0, 0, 0, 1, 1]) (Tenga en cuenta que es correcto solo si a no tiene ceros al principio o al final)
%timeit numpy_fix(np.tile(a, (1, 50000))) 39.3 ms ± 865 µs per loop (mean ± std. dev. of 7 runs, 10 loops each)Un método bastante eficiente es usar itertools.groupby + itertools.chain :
from itertools import groupby, chain a2 = np.array(list(chain(*(l for k,g in groupby(a) if len(l:=list(g))<THRESHOLD or k))))producción:
array([1, 1, 0, 1, 1, 1, 0, 0, 0, 1, 1])Esto funciona relativamente rápido, por ejemplo, en 1 millón de elementos:
# A = np.random.randint(2, size=1000000) %%timeit np.array(list(chain(*(l for k,g in groupby(a) if len(l:=list(g))<THRESHOLD or k)))) # 254 ms ± 3.03 ms per loop (mean ± std. dev. of 7 runs, 1 loop each)