Capítulo 31, Avançado
Vetorizar problemas difíceis
Vetorizar uma soma é fácil. Vetorizar "a distância entre todos os pares de pontos" ou "a soma por grupo" exige repensar o problema, e quase sempre cobra o preço em memória.
Código deste capítulo: avancado/cap31_vetorizar_dificil.py
Todas as distâncias entre pontos
A versão ingênua tem dois laços (cada ponto contra cada ponto). A versão vetorizada usa broadcasting: pontos[:, None, :] (forma (n, 1, d)) menos pontos[None, :, :] (forma (1, n, d)) dá todas as diferenças em um array (n, n, d):
import numpy as np
rng = np.random.default_rng(0)
pontos = rng.random((4, 2))
dif = pontos[:, None, :] - pontos[None, :, :]
dist = np.sqrt((dif ** 2).sum(axis=-1))
print(dist.shape, np.allclose(np.diag(dist), 0), np.allclose(dist, dist.T))
lento = np.array([[np.linalg.norm(p - q) for q in pontos] for p in pontos])
print(np.allclose(dist, lento))
(4, 4) True True
True
A distância de cada ponto a si mesmo é zero (diagonal), e a matriz é simétrica. A versão vetorizada confere com a de laços.
O preço: o array intermediário
O dif tem n × n × d números. Para 4 pontos, nada. Para 10 mil pontos em 2 dimensões:
n = 10_000
print(n * n * 2 * 8 / 1e9, "GB")
1.6 GB
São 1,6 GB só para um intermediário, antes de qualquer conta, e o programa estoura a memória. A vetorização trocou tempo por memória, e essa é a regra mais importante deste capítulo: para cada vetorização, pergunte qual é o maior array intermediário.
Uma identidade que evita o intermediário
Como |a − b|² = |a|² + |b|² − 2·a·b, dá para calcular todas as distâncias com uma multiplicação de matrizes, sem o array (n, n, d). O resultado intermediário é n × n, e a multiplicação usa as bibliotecas de álgebra linear mais rápidas:
def distancias(X):
q = (X ** 2).sum(axis=1)
d2 = q[:, None] + q[None, :] - 2 * X @ X.T
return np.sqrt(np.maximum(d2, 0))
print(np.allclose(distancias(pontos), dist))
True
O np.maximum(d2, 0) existe porque erros de arredondamento podem produzir números levemente negativos onde o valor verdadeiro é zero, e a raiz de um negativo daria nan.
Quando ainda não cabe: processar em blocos
Se o n × n também não cabe, processe o problema em blocos: um laço curto sobre fatias de linhas, com a operação vetorizada dentro. Para achar o vizinho mais próximo de cada ponto de X em Y, o bloco limita o intermediário a bloco × len(Y) × d:
def mais_proximo(X, Y, bloco=500):
resultado = np.empty(len(X), dtype=np.int64)
for inicio in range(0, len(X), bloco):
parte = X[inicio : inicio + bloco]
d2 = ((parte[:, None, :] - Y[None, :, :]) ** 2).sum(axis=-1)
resultado[inicio : inicio + bloco] = d2.argmin(axis=1)
return resultado
X = rng.random((50, 3))
Y = rng.random((30, 3))
por_forca_bruta = np.array([np.argmin(((Y - x) ** 2).sum(axis=1)) for x in X])
print(np.array_equal(mais_proximo(X, Y, bloco=7), por_forca_bruta))
True
Aqui o bloco é 7 de propósito, para o laço girar várias vezes e o teste exercitar a divisão. Em um problema real, o tamanho do bloco é ajustado ao orçamento de memória.
Agrupar sem laço: `reduceat`
Para somar valores por grupo quando os grupos não são inteiros pequenos (o caso do bincount), ordene por grupo e use o add.reduceat, que soma trechos consecutivos:
grupos = np.array([2, 0, 1, 0, 2, 1])
valores = np.array([10.0, 20.0, 30.0, 40.0, 50.0, 60.0])
ordem = np.argsort(grupos, kind="stable")
g = grupos[ordem]
v = valores[ordem]
inicios = np.r_[0, np.flatnonzero(np.diff(g)) + 1]
print(g[inicios], np.add.reduceat(v, inicios))
[0 1 2] [60. 90. 60.]
Depois de ordenar, cada grupo ocupa um trecho contínuo. O np.diff(g) é diferente de zero onde o grupo muda, e esses são os inícios dos trechos.
| Problema | Ferramenta | Custo de memória |
|---|---|---|
| Todos os pares | Broadcasting [:, None] | n² (cuidado) |
| Distâncias | Identidade com @ | n², mas sem o fator d |
| Cabe em blocos | Laço sobre fatias + vetorização | Limitado pelo bloco |
| Soma por grupo (inteiros pequenos) | bincount(weights=...) | Pequeno |
| Soma por grupo (qualquer rótulo) | argsort + reduceat | Cópia ordenada |
Exercício 1
Os k maiores de cada linha
Escreva top_k_por_linha(M, k) que devolva, para cada linha, as posições dos k maiores valores, do maior para o menor. Use np.argpartition e depois ordene só os k.