Articoli approfonditi sulla tecnologia che plasma il futuro.

Oltre la Self-Attention: cosa viene dopo i Transformer

Il costo quadratico dell'attention nei transformer è un collo di bottiglia reale. Linear attention, attention residuals e architetture ibride mostrano cosa verrà dopo.

Groviglio di fili che si srotolano in flussi paralleli e puliti di luce

L'architettura transformer manda avanti l'AI da otto anni ormai. Ogni grande modello linguistico, la maggior parte dei sistemi di generazione di immagini e un numero crescente di modelli audio e video si basano sul meccanismo di self-attention introdotto nel paper 'Attention Is All You Need'. Ma la self-attention ha un problema di fondo: il suo costo in calcolo e memoria cresce quadraticamente con la lunghezza della sequenza. Raddoppi l'input, quadruplichi il costo.

Con sequenze brevi il problema non si sente. Ma con le finestre di contesto da 128K token verso cui stiamo spingendo, e con quelle da un milione di token che la gente vuole, diventa un collo di bottiglia serio. Una serie di ricerche sta esplorando alternative: attention residuals che riusano il calcolo tra i layer, varianti di linear attention che eliminano il costo quadratico e architetture ibride che mescolano attention e meccanismi più economici. Il transformer non sta sparendo, ma si sta riplasmando.

Perché la self-attention è costosa

Per capire le alternative bisogna capire cosa calcola davvero la self-attention. Data una sequenza di N token, la self-attention calcola un punteggio di rilevanza tra ogni coppia di token. Token 1 con token 2, token 1 con token 3, ..., token 1 con token N, poi token 2 con tutti gli altri, e così via. Sono N² coppie.

import torch
import torch.nn.functional as F
def self_attention(Q, K, V):
"""
Standard self-attention.
Q, K, V: (batch, seq_len, d_model)
The attention matrix is seq_len × seq_len.
For seq_len = 1024:   ~1M entries   (manageable)
For seq_len = 32768:  ~1B entries   (expensive)
For seq_len = 131072: ~17B entries  (very expensive)
"""
d_k = Q.size(-1)
# This matmul creates the N×N attention matrix
scores = torch.matmul(Q, K.transpose(-2, -1)) / (d_k ** 0.5)
weights = F.softmax(scores, dim=-1)
return torch.matmul(weights, V)

Con 4K token la matrice di attention ha 16 milioni di elementi: niente di preoccupante per le GPU moderne. Con 128K token ne ha 16 miliardi. A un milione di token supera il trilione. Anche con Flash Attention (che non riduce il calcolo, ma migliora in modo notevole i pattern di accesso alla memoria), alla fine vince la scalabilità quadratica.

È per questo che i primi transformer erano limitati a 512 o 1024 token. Ogni generazione di hardware e di ottimizzazioni ha alzato il tetto, ma stiamo combattendo contro un muro matematico. La scalabilità lineare (O(N)) sarebbe strutturalmente migliore di quella quadratica (O(N²)), ed è quello che perseguono la maggior parte delle architetture alternative.

Attention residuals: riusare ciò che hai già calcolato

Uno degli approcci più pragmatici per ridurre il costo dell'attention non sostituisce l'attention: rende ogni layer di attention più economico riusando il calcolo dei layer precedenti.

L'osservazione è questa: in un transformer profondo (diciamo 32 layer), i pattern di attention dei layer adiacenti sono spesso sorprendentemente simili. Il layer 15 e il layer 16 tendono a guardare posizioni simili, con piccoli aggiustamenti. Calcolare da zero la matrice N² completa a ogni layer è ridondante: gran parte del lavoro era già stato fatto un layer prima.

Gli attention residuals sfruttano questo fatto calcolando un pattern di attention 'residuo': la differenza tra ciò su cui questo layer vuole concentrarsi e ciò che il layer precedente ha calcolato. Se la differenza è piccola (cosa che di solito accade nei layer centrali), il calcolo costa meno. Il pattern di attention completo è la somma del pattern del layer precedente più il residuo del layer corrente.

È analogo a come funziona la compressione video: invece di salvare ogni frame in modo indipendente, si salva un keyframe e poi una serie di differenze (residui) rispetto a quel keyframe. Le differenze sono solitamente molto più piccole del frame intero, quindi la compressione migliora drasticamente.

In pratica, gli attention residuals riducono il costo di calcolo dell'attention del 30-50% nei layer centrali dei modelli profondi, con un impatto minimo sulla qualità. I primi e gli ultimi layer hanno ancora bisogno di attention completa (i loro pattern sono più distintivi), ma i layer centrali, che sono la maggioranza, ottengono accelerazioni significative.

Linear attention: eliminare il costo quadratico

Le varianti di linear attention cercano di riformulare l'attention in modo che scali come O(N) invece di O(N²). L'approccio generale: invece di calcolare esplicitamente la matrice di attention N×N, si trova un modo per ottenere lo stesso output (o uno approssimativamente uguale) usando operazioni lineari.

Il trucco matematico si basa sulla decomposizione in kernel del softmax. L'attention standard calcola softmax(QK^T)V. Se si sostituisce il softmax con una diversa funzione kernel che si può scomporre come φ(Q) · φ(K)^T, si può riordinare l'ordine di calcolo: invece di (φ(Q) · φ(K)^T) · V (che ha un intermedio N×N), si calcola φ(Q) · (φ(K)^T · V) (che ha un intermedio d×d, dove d è la dimensione del modello). Dato che d << N per sequenze lunghe, il costo crolla.

def linear_attention(Q, K, V, feature_map=None):
"""
Linear attention via kernel feature maps.
Cost: O(N * d^2) instead of O(N^2 * d)
"""
if feature_map is None:
# ELU+1 is a common choice (from Katharopoulos et al.)
feature_map = lambda x: F.elu(x) + 1
Q = feature_map(Q)  # (batch, seq_len, d)
K = feature_map(K)  # (batch, seq_len, d)
# Key insight: compute K^T @ V first (d × d matrix)
# instead of Q @ K^T first (N × N matrix)
KV = torch.einsum('bnd,bnm->bdm', K, V)  # (batch, d, d)
# Then multiply by Q
output = torch.einsum('bnd,bdm->bnm', Q, KV)  # (batch, N, d)
# Normalize
Z = torch.einsum('bnd,bd->bn', Q, K.sum(dim=1))  # normalization
output = output / Z.unsqueeze(-1)
return output

Il problema: sostituire il softmax con un'altra funzione kernel cambia la distribuzione dell'attention, e i modelli addestrati con softmax non si trasferiscono necessariamente bene alla linear attention. Il divario di qualità si è ridotto parecchio (le varianti recenti raggiungono il 95-98% della qualità della softmax attention), ma persiste, soprattutto nei task che richiedono un recupero preciso di informazioni a lungo raggio.

State space model: un paradigma diverso

I state space model (SSM) come Mamba seguono un approccio radicalmente diverso. Invece di calcolare le relazioni a coppie tra token, elaborano la sequenza tramite una ricorrenza, mantenendo uno stato nascosto di dimensione fissa che viene aggiornato a ogni token. Questo è intrinsecamente O(N): elaborare il doppio dei token richiede il doppio del tempo, non il quadruplo.

L'innovazione dei SSM moderni sta nel fare dipendere i parametri della ricorrenza dall'input (selective state space). Questo dà al modello una forma di attention basata sul contenuto: può 'scegliere' quali informazioni ricordare e quali dimenticare, senza il costo quadratico. I modelli in stile Mamba eguagliano la qualità dei transformer su molti benchmark, risultando però notevolmente più veloci sulle sequenze lunghe.

Il compromesso: i SSM elaborano i token in sequenza, il che li rende più difficili da parallelizzare durante l'addestramento rispetto ai transformer (che possono processare tutti i token contemporaneamente). L'efficienza in fase di training conta: un modello due volte più veloce in inferenza ma tre volte più lento da addestrare non è necessariamente una vittoria, dato che la maggior parte del calcolo totale va nel training.

Architetture ibride: la via pragmatica

La tendenza attuale nei modelli in produzione è un'architettura ibrida che combina diversi meccanismi di attention. Il ragionamento è semplice: parti diverse di un modello beneficiano di tipi di calcolo diversi.

  • Full attention per il ragionamento globale. Alcuni layer devono guardare l'intera sequenza, trovando contesto rilevante a migliaia di token di distanza. Questi layer usano la self-attention standard (eventualmente ottimizzata con Flash).
  • Local attention per il contesto vicino. Molti layer guardano principalmente i token vicini (sliding window attention). Usare una finestra fissa di 256-1024 token riduce il costo a O(N·W), dove W è la dimensione della finestra.
  • Linear attention per il contesto ampio. Alcuni layer devono aggregare informazioni su tutta la sequenza ma non hanno bisogno di pesi di attention precisi. La linear attention lo garantisce a costo O(N).
  • Layer SSM per l'elaborazione sequenziale. I layer in stile Mamba possono elaborare dipendenze sequenziali in modo efficiente senza alcun calcolo di attention.

Modelli come Jamba (AI21) e diverse architetture di ricerca alternano questi meccanismi in base al ruolo del layer. I primi layer usano la local attention (elaborano sintassi e pattern locali). I layer centrali usano linear attention o SSM (costruiscono rappresentazioni più ampie). Pochi layer strategici usano la full attention (ragionamento globale e recupero). Il risultato è una scalabilità complessiva quasi lineare, senza rinunciare alla qualità del modello che richiede un po' di full attention.

Cosa dovrebbero tenere d'occhio gli sviluppatori

Se stai costruendo applicazioni sopra i modelli linguistici, i cambiamenti architetturali in corso influiscono sul tuo lavoro in modi concreti.

  • Le finestre di contesto continueranno a crescere. Con il calo dei costi dell'attention, le finestre di contesto si espandono. Questo cambia l'architettura delle applicazioni: invece di costruire pipeline RAG complesse per far entrare il contesto rilevante in una finestra da 4K, potresti semplicemente infilare tutto in un prompt da 1M token. La semplicità è allettante, ma le implicazioni in termini di latenza e costi variano tra le architetture.
  • I profili di latenza cambiano. I transformer hanno una latenza relativamente piatta fino a un certo punto, poi cresce in modo quadratico. I modelli con linear attention e SSM mostrano aumenti di latenza più graduali e lineari. Per le applicazioni in cui il tempo di risposta conta, capire come scala il tuo modello è importante.
  • Le differenze di qualità dipendono dal task. I modelli con linear attention possono avere prestazioni leggermente inferiori nei task che richiedono un recupero preciso da posizioni specifiche in contesti lunghi ('qual era il terzo elemento della lista a pagina 47?'). Vanno altrettanto bene nei task che richiedono una comprensione generale. Conosci il tuo caso d'uso.
  • L'ottimizzazione dell'inferenza diventa ancora più importante. Man mano che i modelli diventano architetturalmente più complessi (mescolando tipi diversi di attention), i motori di inferenza devono gestire in modo efficiente un calcolo eterogeneo. vLLM, TensorRT-LLM e framework simili si stanno adattando, ma le architetture custom potrebbero non essere supportate subito.

Il transformer non viene sostituito: si sta evolvendo. La self-attention resta il meccanismo più espressivo che abbiamo per modellare le relazioni tra token. Ma non serve usarla ovunque, in ogni layer, con il pieno costo N². I modelli dei prossimi anni useranno l'attention in modo chirurgico: piena precisione dove conta di più, alternative più economiche ovunque altrove. Il risultato saranno modelli più veloci, capaci di gestire contesti più lunghi e meno costosi da eseguire, pur eguagliando o superando la qualità attuale. Vale la pena tenerli d'occhio.