ELAI S.r.l.

Attenzione online: il riscalamento che evita la matrice dei pesi

Derivazione della softmax pesata online, esempio con logits estremi, fusione di blocchi e limiti numerici: memoria risparmiata senza promettere speedup GPU.

Attenzione online: il riscalamento che evita la matrice dei pesi

Abstract: una somma normalizzata può essere calcolata a blocchi

Per ottenere l’output di una riga di attenzione non è necessario conservare tutti i suoi pesi. È sufficiente mantenere un massimo, un normalizzatore e un vettore accumulato, purché tutti cambino scala insieme quando arriva uno score maggiore. Deriviamo questa identità, la verifichiamo con logits che mandano in overflow l’esponenziale ingenuo e mostriamo un errore concreto: riscalare solo il denominatore produce un output sbagliato anche se la softmax sembra numericamente stabile. Il contributo è una derivazione commentata con codice eseguito, non un nuovo algoritmo o un benchmark GPU.

1. Quale oggetto vogliamo calcolare

Una query q e N chiavi k_j, di dimensione d_k, producono score adimensionali s_j=qᵀk_j/√d_k. Un eventuale bias è incluso nello score. A ogni chiave corrisponde un valore v_j in R^(d_v). L’output o è una combinazione convessa dei valori: i coefficienti sono positivi e sommano a uno. Qui consideriamo una riga, nessun dropout e almeno una chiave ammessa dalla maschera. Le chiavi escluse non contribuiscono. La domanda riguarda l’ordine del calcolo, non una modifica della funzione matematica.

p_j = exp(s_j) / Σ_i exp(s_i) o = Σ_j p_j v_j m = max_j s_j l = Σ_j exp(s_j−m) u = Σ_j exp(s_j−m) v_j o = u/l

Sottrarre lo stesso massimo da tutti gli score moltiplica numeratore e denominatore per exp(−m), quindi lascia invariato il rapporto. Per score finiti, gli esponenti diventano non positivi: nessun termine supera uno e almeno un termine vale uno. Il normalizzatore l è dunque compreso fra 1 e N nell’aritmetica reale. Questo evita l’overflow degli esponenziali, ma non promette accuratezza illimitata: termini molto piccoli possono andare in underflow e le somme restano soggette ad arrotondamento.

2. L’invariante e il cambio di scala

Dopo t chiavi manteniamo m_t, il massimo del prefisso; l_t, la somma degli esponenziali riferiti a quel massimo; u_t, la stessa somma pesata dai valori. Arriva la coppia (s,v). Il nuovo massimo è m′=max(m_t,s). Ogni vecchio contributo exp(s_i−m_t) deve essere moltiplicato per exp(m_t−m′): così diventa exp(s_i−m′). Lo stesso fattore deve moltiplicare anche u_t, perché i valori erano accumulati con i medesimi pesi. Solo dopo aggiungiamo il contributo della nuova chiave.

m′ = max(m, s) a = exp(m−m′), b = exp(s−m′) l′ = a·l + b u′ = a·u + b·v o′ = u′/l′

La dimostrazione è un’invariante di ciclo: sostituendo le definizioni di l e u nelle ultime due righe si ottengono esattamente le somme sul prefisso esteso. Il caso iniziale con una chiave è m=s_1, l=1, u=v_1. Questa inizializzazione evita di calcolare −∞−(−∞). Se si preferisce uno stato vuoto, va gestito esplicitamente. È importante distinguere u, numeratore non normalizzato, da o: riscalare o come se fosse u dimenticherebbe il fattore l precedente.

3. Un esempio che rende visibile l’errore

Usiamo s=(1000,1001,999) e valori v_1=(1,0), v_2=(0,2), v_3=(3,−1). Python segnala overflow per exp(1000), mentre i pesi stabilizzati sono proporzionali a (e^−1,1,e^−2). Dopo la prima chiave l=1 e u=(1,0). La seconda aumenta il massimo da 1000 a 1001: il vecchio numeratore diventa (e^−1,0), non rimane (1,0). Quindi l=1+e^−1 e u=(e^−1,2). La terza non cambia il massimo e aggiunge e^−2·(3,−1).

Passomlo₁o₂
11000.01.0000000001.0000000000.000000000
21001.01.3678794410.2689414211.462117157
31001.01.5032147240.5148201911.240451338

Il risultato finale è o=(0,514820191; 1,240451338). Se al secondo passo si riscalasse l ma non u, il primo componente diventerebbe 0,935332675. Il secondo resterebbe casualmente corretto perché il suo contributo precedente era zero: controllare un solo componente potrebbe quindi non rilevare il difetto. Entrambi gli output sono numericamente finiti. Assenza di NaN e overflow non equivale a correttezza dell’algoritmo.

Sopra: output dopo ogni chiave dell’esempio. Sotto: dimensione teorica di una sola matrice N×N FP32, esclusi gli altri tensori; non è una misura della memoria totale di un modello.
Sopra: output dopo ogni chiave dell’esempio. Sotto: dimensione teorica di una sola matrice N×N FP32, esclusi gli altri tensori; non è una misura della memoria totale di un modello.

4. Unire blocchi senza conservare i pesi

La stessa algebra unisce due insiemi disgiunti A e B, ciascuno riassunto da (m,l,u). Si sceglie m=max(m_A,m_B), si riportano entrambi i normalizzatori e numeratori alla nuova scala, poi si sommano. Il risultato rappresenta l’unione, quindi in aritmetica reale l’operazione è associativa e commutativa. In floating point, cambiare l’albero di riduzione può cambiare gli ultimi bit. Questo consente blocchi paralleli senza affermare identità bit per bit con un ordine sequenziale.

m = max(m_A, m_B) l = exp(m_A−m)l_A + exp(m_B−m)l_B u = exp(m_A−m)u_A + exp(m_B−m)u_B

L’esperimento allegato confronta blocchi di 1, 7, 32 e 257 elementi con il calcolo materializzato stabile su 257 chiavi e valori a quattro componenti. Il generatore Python usa seed 20260924; gli score sono uniformi in [−a,a] per a=1,10,1000 e i valori in [−2,2]. Il massimo errore assoluto fra i dodici confronti è circa 6,66×10^−16. È un controllo su dati sintetici in float Python, non un limite universale, né una misura FP16/BF16. Codice, seed e risultati sono nell’archivio.

5. Memoria, complessità e ciò che resta da misurare

Per una riga lo stato persistente del numeratore richiede d_v+2 scalari, oltre a query, blocchi di chiavi e valori e buffer temporanei. Con d_v=64 e FP32 sono 264 byte: non è la memoria totale del kernel. Una matrice score completa N×N FP32 occupa invece 4N² byte per testa e per elemento del batch: 64 MiB a N=4096, 256 MiB a N=8192. Un’implementazione ingenua può anche materializzare i pesi normalizzati. Evitare queste matrici è diverso dal comprimere K e V, e non elimina una KV cache richiesta dal decoder.

La computazione densa di tutte le coppie query-chiave resta quadratica: O(N²(d_k+d_v)). La ricorrenza non trasforma l’attenzione esatta in attenzione lineare. Un ciclo Python scalare può essere più lento di una moltiplicazione vettorizzata; il vantaggio hardware richiede blocchi, fusioni e accessi alla memoria coerenti con il dispositivo. Per misurarlo servono GPU, dtype, dimensioni, maschera, batch, versioni, riscaldamento e sincronizzazione dichiarati. Qui non abbiamo eseguito quel benchmark e non riportiamo speedup.

Le righe interamente mascherate richiedono una convenzione esplicita: il rapporto 0/0 non definisce una distribuzione. Il codice rifiuta uno stato senza chiavi e richiede score finiti; chiavi mascherate vanno filtrate prima. Valori con segni diversi possono causare cancellazione nel numeratore, mentre precisioni ridotte richiedono una scelta dell’accumulatore. Anche gradienti e dropout necessitano di logica aggiuntiva: la verifica del forward di una riga non certifica un’implementazione di addestramento completa.

6. Fonti e conclusione

Il normalizzatore online è descritto da Milakov e Gimelshein (2018, arXiv v2); i loro esperimenti usano Tesla V100, CUDA 9.1 e batch differenti, quindi i tempi non sono trasferibili al nostro codice. FlashAttention di Dao e coautori (2022, arXiv v2) collega il calcolo a blocchi alla riduzione del traffico HBM e alla ricomputazione nel backward. Abbiamo letto algoritmo, analisi e sezioni sperimentali: questa pagina ne illustra un fondamento, senza riprodurre i benchmark né descrivere tutte le evoluzioni successive.

La conclusione verificabile è l’invariante condiviso da numeratore e denominatore. Il massimo non è solo un trucco per evitare overflow: definisce l’unità numerica in cui i contributi sono accumulati. Quando cambia, tutti i contributi precedenti devono essere convertiti. È questo passaggio che permette di liberarsi della matrice dei pesi senza modificare, in aritmetica reale, il risultato dell’attenzione.

Milakov M., Gimelshein N. (2018), Online normalizer calculation for softmax, arXiv:1805.02867v2.

Dao T., Fu D. Y., Ermon S., Rudra A., Ré C. (2022), FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness, arXiv:2205.14135v2.

from math import exp
scores = [1000., 1001., 999.]
values = [[1., 0.], [0., 2.], [3., -1.]]
m, total, numerator = scores[0], 1., values[0][:]
for s, v in zip(scores[1:], values[1:]):
    new_m = max(m, s)
    old_scale, new_scale = exp(m-new_m), exp(s-new_m)
    total = total*old_scale + new_scale
    numerator = [u*old_scale + x*new_scale for u, x in zip(numerator, v)]
    m = new_m
print([u/total for u in numerator])

Codice, dati e istruzioni · JSON. Calcoli didattici eseguiti con Python 3.14.0; figure con Matplotlib 3.11.2. Analisi con assistenza AI, senza dichiarare peer review o revisione umana. Copertina originale ImageGen, illustrativa: non documenta persone, sedi o installazioni EL-AI. Fonti consultate il 24 settembre 2026.