Paithon Book Paithon Book
Esegui il codice

L’attenzione lineare causale è una RNN: una verifica numerica

L’attenzione lineare causale è una RNN: una verifica numerica#

L’attenzione lineare causale, in cui ogni token guarda soltanto sé stesso e i precedenti, si può calcolare in due modi: in forma parallela, costruendo come in un’attenzione la tabella \(n\times n\) delle affinità fra gli \(n\) token, e in forma ricorrente, aggiornando come in una rete ricorrente uno stato di dimensione fissa. Qui si calcola lo stesso strato nei due modi, con poche righe di NumPy, e poi una terza volta a blocchi, come si fa in addestramento.

A decidere sono le righe stampate dalle due verifiche, quella delle due forme e quella a blocchi: devono dire tutte True. Vuol dire che le uscite coincidono a meno dell’arrotondamento, cioè che è la stessa funzione calcolata in modi diversi.

Ingredienti#

Si generano query, chiavi e valori casuali per \(n = 6\) token: ogni token ha una query e una chiave di \(d_k = 4\) numeri e un valore di \(d_v = 5\) numeri, e la prima riga stampata ne mostra le forme (sei righe da quattro numeri, sei righe da cinque). Poi si definisce la feature map \(\phi(x)=\mathrm{elu}(x)+1\), la funzione che trasforma ogni numero di query e chiavi prima del confronto: vale \(x+1\) per \(x\) positivo ed \(e^x\) per \(x\) negativo, quindi è sempre positiva. Così sono positive anche le affinità, cioè le somiglianze fra query e chiavi, e la divisione per il loro totale (la normalizzazione) fa dell’uscita una media pesata dei valori, come in Katharopoulos et al. (2020).

import numpy as np

rng = np.random.default_rng(0)
n, d_k, d_v = 6, 4, 5          # 6 token; chiavi in R^4, valori in R^5
Q = rng.standard_normal((n, d_k))
K = rng.standard_normal((n, d_k))
V = rng.standard_normal((n, d_v))

def phi(x):                    # feature map elu(x)+1: sempre positiva
    return np.where(x > 0, x + 1.0, np.exp(x))

PQ, PK = phi(Q), phi(K)
print("Q, K:", Q.shape, " V:", V.shape)
Q, K: (6, 4)  V: (6, 5)

Forma parallela (tipo attenzione)#

Costruiamo la matrice di affinità \(\mathbf{A}=\phi(\mathbf{Q})\,\phi(\mathbf{K})^\top\), applichiamo la maschera causale, che azzera le affinità con i token futuri (ogni token vede soltanto sé stesso e i precedenti), e normalizziamo: l’uscita è la media pesata dei valori. Nota la dimensione di \(\mathbf{A}\): cresce come \(n^2\).

# Forma PARALLELA (tipo attenzione): costruisce la matrice di affinita n x n
A = PQ @ PK.T                              # (n, n)
A *= np.tril(np.ones((n, n)))              # maschera causale: ogni token vede solo il passato
O_par = (A @ V) / A.sum(1, keepdims=True)  # media pesata dei valori
print("matrice di attenzione:", A.shape, "-> cresce come n^2")
print("output O_par:", O_par.shape)
matrice di attenzione: (6, 6) -> cresce come n^2
output O_par: (6, 5)

Forma ricorrente (tipo RNN)#

Ora lo stesso calcolo senza mai costruire la matrice \(n\times n\): si scorrono i token uno alla volta e si accumula ogni coppia chiave-valore nello stato \(\mathbf{S}\), una matrice di dimensione fissa \(d_v\times d_k\), con i valori sulle righe e le chiavi trasformate sulle colonne:

\[\mathbf{S}_t = \mathbf{S}_{t-1} + \mathbf{v}_t\,\phi(\mathbf{k}_t)^\top, \qquad \mathbf{z}_t = \mathbf{z}_{t-1} + \phi(\mathbf{k}_t),\]

dove \(\mathbf{v}_t\) è il valore del token \(t\), \(\phi(\mathbf{k}_t)\) la sua chiave trasformata e \(\mathbf{z}_t\) il normalizzatore, il totale delle chiavi trasformate per cui si divide. L’uscita del token \(t\) è la lettura dello stato con la sua query trasformata, divisa per quel totale: \(\mathbf{o}_t = \mathbf{S}_t\,\phi(\mathbf{q}_t) \,/\, \big(\mathbf{z}_t^\top \phi(\mathbf{q}_t)\big)\). Lo stato non dipende dalla lunghezza della sequenza.

# Forma RICORRENTE (tipo RNN): uno stato-matrice di dimensione FISSA, aggiornato token per token
S = np.zeros((d_v, d_k))       # stato: NON dipende da n
z = np.zeros(d_k)              # normalizzatore
O_rec = np.zeros((n, d_v))
for t in range(n):
    S += np.outer(V[t], PK[t]) # accumula v_t phi(k_t)^T (aggiornamento di rango 1)
    z += PK[t]
    O_rec[t] = (S @ PQ[t]) / (z @ PQ[t])
print("stato S:", S.shape, "-> fisso, indipendente da n")
stato S: (5, 4) -> fisso, indipendente da n

Verifica#

Le due forme devono coincidere a meno dell’errore di arrotondamento, che in doppia precisione è dell’ordine di \(10^{-16}\) su numeri vicini a uno. La verifica non stampa lo scarto, la cui ultima cifra cambia da un processore all’altro, ma dice se sta sotto \(10^{-12}\), cioè \(0{,}000\,000\,000\,001\): una soglia diecimila volte più larga dell’arrotondamento, e comunque minuscola.

# Le due forme calcolano la STESSA funzione
scarto = np.abs(O_par - O_rec).max()
print("max |differenza| < 1e-12:", scarto < 1e-12)
print("le due forme coincidono:", np.allclose(O_par, O_rec))
max |differenza| < 1e-12: True
le due forme coincidono: True

Forma a blocchi#

Nella pratica la forma parallela si spezza in blocchi di \(B\) token: dentro il blocco si costruisce la tabella \(B\times B\), fra un blocco e l’altro passa lo stato. Il lavoro diventa \(O(nBd + nd^2)\) invece di \(O(n^2 d)\), con \(d\) la dimensione di chiavi e valori. Con \(n = 6\) e \(B = 4\) l’ultimo blocco è più corto, e il risultato deve coincidere con le altre due forme.

# Forma A BLOCCHI: parallela dentro il blocco, ricorrente fra i blocchi
B = 4                                       # token per blocco
S_b = np.zeros((d_v, d_k))                  # stato alla fine dei blocchi prima
z_b = np.zeros(d_k)                         # e il suo normalizzatore
O_blk = np.zeros((n, d_v))
for c in range(0, n, B):
    q, k, v = PQ[c:c + B], PK[c:c + B], V[c:c + B]
    A_b = (q @ k.T) * np.tril(np.ones((len(q), len(q))))   # tabella piccola
    num = A_b @ v + q @ S_b.T               # dentro il blocco + blocchi prima
    den = A_b.sum(1) + q @ z_b
    O_blk[c:c + B] = num / den[:, None]
    S_b = S_b + v.T @ k                     # lo stato passa al blocco dopo
    z_b = z_b + k.sum(0)
print("a blocchi == ricorrente:", np.allclose(O_blk, O_rec))
print("a blocchi == parallela:", np.allclose(O_blk, O_par))
print("scarto massimo < 1e-12:", np.abs(O_blk - O_rec).max() < 1e-12)
a blocchi == ricorrente: True
a blocchi == parallela: True
scarto massimo < 1e-12: True

Le tre forme calcolano lo stesso strato, e differiscono solo per l’arrotondamento. Durante l’addestramento conviene la forma a blocchi, che lavora quasi tutta in parallelo senza costruire la tabella \(n\times n\); durante la generazione, quando i token arrivano uno dopo l’altro, conviene la forma ricorrente, che aggiorna uno stato di dimensione costante: nessuna KV cache che cresce parola dopo parola. È la proprietà che accomuna tutta la famiglia delle ricorrenze lineari, e che il capitolo sugli State Space Model ritrova partendo dai sistemi dinamici invece che dall’attenzione.