Capítulo 30, Avançado
einsum e tensordot
O `einsum` descreve uma operação entre arrays por uma pequena fórmula com letras, e substitui dezenas de combinações de `transpose`, `reshape` e `sum` por uma linha legível.
Código deste capítulo: avancado/cap30_einsum.py
A fórmula
Você nomeia cada eixo com uma letra e diz o que quer no resultado. A regra é curta: letras repetidas entre as entradas são multiplicadas, e letras que não aparecem na saída são somadas. A multiplicação de matrizes é o exemplo clássico: "ij,jk->ik" multiplica pelo j (repetido) e o soma (porque o j não está na saída):
import numpy as np
A = np.arange(6).reshape(2, 3)
B = np.arange(12).reshape(3, 4)
print(np.einsum("ij,jk->ik", A, B))
print(np.array_equal(np.einsum("ij,jk->ik", A, B), A @ B))
[[20 23 26 29]
[56 68 80 92]]
True
Um dicionário de padrões
Com a mesma regra, a fórmula expressa dezenas de operações:
M = np.arange(9).reshape(3, 3)
print(np.einsum("ii->", M), np.einsum("ii->i", M), np.einsum("ij->i", M), np.einsum("ij->j", M))
v = np.array([1, 2, 3])
w = np.array([4, 5, 6])
print(np.einsum("i,i->", v, w), np.einsum("i,j->ij", v, w))
12 [0 4 8] [ 3 12 21] [ 9 12 15]
32 [[ 4 5 6]
[ 8 10 12]
[12 15 18]]
| Fórmula | Operação |
|---|---|
"ij->ji" | Transposta |
"ii->" | Traço (soma da diagonal) |
"ii->i" | Diagonal |
"ij->i" | Soma de cada linha |
"ij->j" | Soma de cada coluna |
"i,i->" | Produto escalar |
"i,j->ij" | Produto externo (todas as combinações) |
"bij,bjk->bik" | Multiplicação de matrizes em lote |
O caso de uso real: atenção em redes neurais
O cálculo central dos modelos de linguagem compara cada "consulta" com cada "chave", em um lote. Em fórmula, "bqd,bkd->bqk": para cada exemplo b, multiplica o vetor de dimensão d de cada consulta q por cada chave k. Sem o einsum, você precisaria transpor a última dimensão à mão:
rng = np.random.default_rng(0)
Q = rng.normal(size=(2, 5, 8))
K = rng.normal(size=(2, 7, 8))
scores = np.einsum("bqd,bkd->bqk", Q, K)
print(scores.shape, np.allclose(scores, Q @ K.transpose(0, 2, 1)))
(2, 5, 7) True
A fórmula documenta os eixos: cada letra tem um significado, e quem lê entende a operação sem decifrar uma sequência de transpose.
`tensordot` e a ordem das multiplicações
O np.tensordot soma sobre os eixos que você indica (o parâmetro axes), e para duas matrizes com axes=1 equivale ao @. A mesma operação pode ser muito mais barata dependendo da ordem em que se multiplica uma cadeia de matrizes, porque a multiplicação é associativa, mas o custo não é:
import timeit
print(np.allclose(np.tensordot(A, B, axes=1), A @ B))
C1 = rng.normal(size=(2000, 3))
C2 = rng.normal(size=(3, 2000))
C3 = rng.normal(size=(2000, 3))
esquerda = lambda: (C1 @ C2) @ C3
direita = lambda: C1 @ (C2 @ C3)
print(np.allclose(esquerda(), direita()))
t_esq = min(timeit.repeat(esquerda, number=1, repeat=3))
t_dir = min(timeit.repeat(direita, number=1, repeat=3))
print("a ordem importa:", t_esq > 3 * t_dir)
True
True
a ordem importa: True
O primeiro jeito cria uma matriz de 2000 por 2000 no meio do caminho, e o segundo reduz tudo a uma matriz de 3 por 3 antes. O resultado é o mesmo, e o custo, muito diferente. O np.linalg.multi_dot escolhe a melhor ordem para você, e o einsum com optimize=True faz o mesmo.
Quando eu não uso einsum
Para uma multiplicação de matrizes simples, o
@é mais claro e usa as bibliotecas otimizadas diretamente. Oeinsumcompensa quando há mais de duas entradas, eixos de lote ou somas incomuns, e quando a fórmula documenta melhor do que o código equivalente.
Exercício 1
Similaridade do cosseno entre todos os pares
Escreva similaridade_cosseno(X) que receba uma matriz com um vetor por linha e devolva a matriz n × n de cossenos entre todos os pares. Use np.einsum para os produtos escalares.