Un contesto di un'intera giornata batte uno di dieci minuti? Attenzione flash e domanda sulla lunghezza della sequenza
Ecco la domanda a cui questo articolo risponde: se un trasformatore potesse occuparsi di un'intera giornata di negoziazione con una risoluzione di un secondo invece che con una finestra di dieci minuti, farebbe previsioni migliori?
Fino a poco tempo fa non potevi nemmeno chiedere. Bisogni di attenzione standard memoria, quindi una giornata di 23.400 passi a 12 teste in float16 richiede circa 12,9 GB solo per la matrice del punteggio: più dei parametri del modello e più di quanto la maggior parte delle carte ti darà. La questione è stata chiusa dall'aritmetica prima che qualcuno potesse testarla.
Flash Attention (Dao et al., 2022) lo apre. Non avvicinando l'attenzione (calcola lo stesso risultato esatto), ma ristrutturando il calcolo in modo che sia consapevole dell'IO, riducendo al minimo il traffico tra i livelli di memoria della GPU. Questo è il contenuto veramente interessante qui, e la maggior parte di questo articolo è dedicata a come funziona: la piastrellatura, la ricorrenza del softmax online, il Limite IO e ricalcolo del passaggio all'indietro.
Ma il meccanismo è il facilitatore, non la pretesa. "Un contesto più lungo è migliore" è un'affermazione empirica sui mercati e la posizione di questo blog - da Temporal Fusion Transformers, che ha scoperto che i trasformatori vanilla sui modelli ricorrenti overfit e short-lookback delle serie finanziarie rimangono competitivi ad alta frequenza – taglia nella direzione opposta. Quindi l’articolo si chiude sulla misurazione, non sul meccanismo.
Perché l'attenzione è legata alla memoria
L'attenzione calcola — la primitiva stessa, in un contesto commerciale, è trattata in Temporal Fusion Transformers for Multi-Horizon Portfolio Forecasting. L'intero problema è una riga di questo: la matrice del punteggio intermedio È , viene scritto in memoria, riletto per il softmax, scritto di nuovo e riletto per il matmul finale - e deve essere conservato per la propagazione all'indietro.
L'intensità aritmetica dell'attenzione è , quindi circa 64 FLOP/byte a — ben a sinistra del crinale della A100. Si trova sul tetto inclinato della larghezza di banda, non sul tetto piatto del calcolo: la GPU trascorre più tempo in movimento intorno che moltiplicare qualsiasi cosa. La struttura della linea del tetto utilizzata da questo (punto di colmo, soffitto inclinato rispetto a soffitto piatto e il motivo per cui lo stesso ragionamento decide se vale la pena acquistare una GPU) è costruita con numeri misurati in Quando la GPU ripaga.
La gerarchia di memoria sfruttata dall'algoritmo
| Livello di memoria | Taglia | Larghezza di banda | Latenza |
|---|---|---|---|
| HBM (memoria a larghezza di banda elevata) | 40-80GB | 2,0 TB/sec | ~400 ns |
| SRAM (memoria condivisa su chip) | 20MB | 19TB/s | ~4ns |
La SRAM ha circa 10 volte la larghezza di banda e 100 volte la latenza inferiore, con un millesimo della capacità. Tutto ciò che fa Flash Attention deriva da questo scambio: rinunciare a capacità, acquistare larghezza di banda e latenza. La stessa mossa "ristrutturare l'algoritmo anziché acquistare hardware", misurata su un backtest della CPU, è la scala della velocità del backtest.
L'algoritmo Flash Attention
Flash Attention elabora l'attenzione in riquadri di dimensioni adatte alla SRAM e non si materializza mai completamente matrice in HBM.
Partizione in blocchi di righe e in blocchi di colonne, con scelto in modo che una tessera più i suoi accumulatori entrino nel chip. Per ogni blocco di query, esegui l'iterazione su tutti i blocchi di valori-chiave:
For each query block Q_i:
Initialize: O_i = 0, m_i = -inf, l_i = 0 # output, running max, running sum
For each KV block (K_j, V_j):
1. Load Q_i, K_j, V_j from HBM to SRAM
2. Compute S_ij = Q_i @ K_j^T / sqrt(d) # in SRAM
3. Compute local max: m_ij = rowmax(S_ij)
4. Compute P_ij = exp(S_ij - m_ij) # in SRAM
5. Compute local sum: l_ij = rowsum(P_ij)
6. Update running statistics:
m_new = max(m_i, m_ij)
l_new = l_i * exp(m_i - m_new) + l_ij * exp(m_ij - m_new)
O_i = O_i * (l_i * exp(m_i - m_new) / l_new)
+ P_ij @ V_j * (exp(m_ij - m_new) / l_new)
m_i = m_new, l_i = l_new
Write O_i to HBM
La ricorrenza softmax online
Il trucco che rende possibile la piastrellatura è softmax online. Un softmax ingenuo necessita di due passaggi sulla riga: uno per trovare il massimo (per la stabilità numerica), uno per esponenziare e normalizzare. Due passaggi su una riga che rifiuti di memorizzare sono una contraddizione, quindi Flash Attention continua a eseguire statistiche e si ridimensiona man mano che procede.
Dopo i blocchi :
e l'accumulatore di uscita viene corretto dello stesso fattore:
Ogni volta che un nuovo blocco aumenta il massimo corrente, l'output accumulato in precedenza viene ridimensionato retroattivamente di – come se il nuovo massimo fosse noto fin dall’inizio. Il risultato è algebricamente identico al softmax a due passaggi. In aritmetica esatta questa non è un'approssimazione; è una riassociazione. (In precisione finita si tratta di un percorso di arrotondamento diverso, che conta: vedere il controllo di esattezza di seguito.)
Complessità I/O
Questa è la dichiarazione formale della vittoria. Viene eseguito il Flash Attenzione
Accessi HBM, dove è la dimensione della SRAM, contro per l'implementazione standard. Notare che appare nel denominatore: più grande è lo scratchpad sul chip, minori saranno i viaggi di andata e ritorno, motivo per cui l'algoritmo è dichiarato in termini di gerarchia di memoria piuttosto che di conteggio FLOP. Per tipico E KB, il rapporto favorisce Flash Attention con circa 5-10 volte meno accessi.
Passaggio all'indietro: ricalcolo anziché archiviazione
La propagazione all'indietro attraverso l'attenzione normalmente richiede il matrice che il passaggio in avanti si è semplicemente rifiutato di mantenere. Flash Attenzione ricalcola le tessere da durante il passaggio all'indietro, memorizzando solo l'output e le statistiche softmax - Entrambi , non . Scambia una modesta quantità di aritmetica ridondante con il termine di memoria che costituiva l'intero problema. Questo è lo stesso affare del checkpoint del gradiente, applicato alla granularità delle tessere all'interno di un singolo operatore.
FA2: parallelismo
Flash Attention 2 (Dao, 2023) ha mantenuto l'algoritmo e corretto la pianificazione:
- Meno FLOP non matmul. FA1 ha dedicato tempo reale al ridimensionamento, alla ricerca massima e all'esponenziazione: operazioni eseguite su core CUDA, non tensor core. FA2 rinvia il ridimensionamento alla fine del ciclo interno.
- Parallelismo sulla lunghezza della sequenza. FA1 parallelizza solo su batch e teste. FA2 esegue anche la parallelizzazione sui blocchi di query. Ciò è importante in particolare per il caso commerciale, in cui spesso si ha una sequenza molto lunga per asset e una dimensione batch di 1-4, esattamente il regime in cui il parallelismo batch-e-head affama la GPU.
- Partizionamento del lavoro di warp. Ogni warp accetta un sottoinsieme diverso di blocchi di query anziché suddividere il calcolo del punteggio e ridurlo tra warp, rimuovendo una riduzione tra warp.
Risultato riportato: circa il 70% dei FLOP di picco teorici su A100 rispetto a circa il 35% per FA1.
FA3: Meccanica della tramoggia
Flash Attention 3 (Dao, Shah, 2024) è specifico dell'architettura per H100:
- Specializzazione warp asincrona. Il Tensor Memory Accelerator (TMA) di Hopper sposta HBM→SRAM in modo asincrono. FA3 divide le curvature in produttori che emettono carichi TMA per il successivo blocco KV e consumatori che calcolano su quello corrente, quindi il movimento dei dati si nasconde dietro l'aritmetica.
- Matmul interfogliato e softmax. per un blocco viene eseguito su tensor core mentre il softmax per il blocco precedente viene eseguito su CUDA core: due diverse unità hardware, realmente simultanee anziché suddivise nel tempo.
- FP8 con elaborazione incoerente. H100 esegue FP8 con un throughput 2x FP16. L’ingenua attenzione all’8PQ è distrutta da valori anomali; FA3 ruota in modo casuale i vettori prima della quantizzazione a blocchi per distribuire la grandezza anomala tra le coordinate, segnalata con un errore numerico 2,6 volte inferiore rispetto all'ingenuo FP8.
| Versione | GPU | Utilizzo | Accelerazione vs Standard |
|---|---|---|---|
| FA1 | A100 | ~35% | 2-4x |
| FA2 | A100 | ~70% | 5-7x |
| FA3 (FP16) | H100 | ~75% | 3-5x contro FA2 |
| FA3 (FP8) | H100 | ~75% | 1,6 volte rispetto a FA3 FP16 |
Il mascheramento causale è il modo in cui il trading ottiene uno sconto
Il mascheramento causale è obbligatorio per le serie temporali – il modello non deve guardare al futuro – e nel caso del affiancamento non rappresenta un costo aggiuntivo ma un risparmio. Qualsiasi riquadro le cui chiavi sono interamente nel futuro rispetto alle sue query viene saltato del tutto, mai caricato e mai calcolato, riducendo circa la metà del lavoro. In PyTorch questo è is_causal=True; non è richiesto nient'altro.
L'integrazione è di otto righe
Quasi nessuno del codice di cui hai bisogno riguarda Flash Attention. Scambia il percorso esplicito della matrice di punteggio per il kernel fuso:
scores = torch.matmul(q, k.transpose(-2, -1)) / math.sqrt(d_head)
scores = scores.masked_fill(causal_mask, float("-inf"))
out = torch.matmul(torch.softmax(scores, dim=-1), v)
out = torch.nn.functional.scaled_dot_product_attention(
q, k, v,
dropout_p=self.dropout if self.training else 0.0,
is_causal=True,
)
Questo è il cambiamento totale. q, k, v sono modellati (batch, heads, seq, head_dim); la maschera causale è scomparsa perché è stata creata dal kernel. Per un modello di trading PyTorch completo e annotato in cui inserirlo: proiezione di input, blocchi e testa su/piatto/giù di 3 classi, utilizzare DeepLOB e per una pipeline di formazione completa vedere Temporal Fusion Transformers. Costruire qui una quarta copia di quell'impalcatura non insegnerebbe nulla.
Requisiti: capacità di calcolo >= 8.0 (A100, H100, RTX 3090+), input a mezza precisione, PyTorch >= 2.0. Verificare il percorso veloce effettivamente utilizzato torch.backends.cuda.sdp_kernel diagnostica e torch.cuda.max_memory_allocated() — SDPA ritorna silenziosamente al kernel matematico se una qualsiasi precondizione fallisce, e un fallback silenzioso assomiglia esattamente a un modello funzionante che è semplicemente lento.
La misurazione: il contesto più lungo paga?
Tutto quanto sopra dice che un contesto da 32K o 128K è ora accessibile. Non dice nulla sul fatto che sia utile. L'esperimento onesto:
Addestra la stessa architettura a con SDPA-Flash sulla serie BTC utilizzata altrove in questa serie, mantenendo parametri, ottimizzatore e target fissi in modo che la lunghezza della sequenza sia l'unica variabile. Segnala due cose:
- Costo. Orologio da parete misurato per epoca e
torch.cuda.max_memory_allocated()a ciascuno . - Vantaggio. Prestazioni predittive fuori campione rispetto a , in una divisione walk-forward.
Una bozza precedente di questo articolo conteneva una tabella di valori di memoria per lunghezza di sequenza derivata analiticamente da formula di attivazione. Quelle righe vengono rimosse: non sono mai state misurate e non erano d'accordo con l'aritmetica del budget di memoria dell'articolo. Un numero derivato presentato in una tabella dei risultati è un risultato fabbricato e questo blog non lo fornisce.
La proprietà interessante di questo esperimento è che è pubblicabile in entrambe le direzioni. Se le prestazioni fuori campione aumentano in modo monotono con , che giustifica l'intero programma a lungo contesto. Se si stabilizza dopo poche migliaia di passi o degrada, si tratta di un pezzo più forte, un compagno dell'onesto negativo – e ciò significherebbe che il muro della memoria non è mai stato il vincolo vincolante per il commercio dei trasformatori.
Più contesto significa più capacità, quindi più superficie sovradimensionata
C’è una ragione specifica per aspettarsi un risultato piatto o negativo. Trasformatori di fusione temporale documenta già che i trasformatori vanilla si sono applicati ingenuamente all'overfit delle serie finanziarie: sono privi di bias induttivi temporali e i modelli ricorrenti a breve sguardo rimangono competitivi ad alta frequenza. Estendere il contesto da 512 a 32.768 passi non aggiunge informazioni proporzionali alla lunghezza; il margine marginale 32.000esimo di una serie di prezzi quasi efficiente ha ben poco significato. Ciò che aggiunge in modo affidabile è il valore dei parametri da adattare.
Quindi la spazzata è finita deve essere trattata per quello che è: una ricerca per la selezione del modello, con lo stesso meccanismo che questo blog applica a ogni altra ricerca. Tre lunghezze di sequenza moltiplicate per qualunque altra variazione costituiscono il conteggio delle prove, e il vincitore deve superare un indice di Sharpe sgonfio calcolato rispetto a quel conteggio delle prove e un cancello PBO, non semplicemente battere i suoi vicini. Altrimenti "vince il contesto lungo" è indistinguibile dallo scegliere il meglio di tre esecuzioni rumorose.
Un controllo di esattezza, perché "esatto" significa lavorare molto
Flash Attenzione è esatto in aritmetica esatta. La raccomandazione ad esso allegata – eseguire in FP16 o BF16 e su H100 considerare FP8 – non lo è. Queste sono affermazioni separate e la seconda domina nella pratica: riassociare una somma e ridurre a metà la precisione sono entrambe perturbazioni, e l'articolo che ha introdotto la garanzia dell'ordine non dovrebbe quindi sventolare quella di precisione.
Il blog ha già lo strumento giusto. La trappola della precisione della GPU stabilisce lo standard: la bassa precisione non ti avvisa, restituisce spazzatura plausibile e tu dimostri la correttezza con un oracolo di parità su una quantità discreta a valle - i conteggi degli scambi - non osservando le curve a occhio. Applicato qui:
- Calcola l'attenzione con SDPA-Flash in bf16 e con un'implementazione di riferimento fp64 su input identici; riporta errore relativo massimo sul tensore di uscita.
- Spingilo fino alla decisione: per un modello che emette un'etichetta su/piatto/giù, riporta quante etichette si invertono tra i due percorsi, come frazione delle decisioni totali.
Un disaccordo piccolo, limitato e spiegabile è la firma di un percorso veloce e corretto. Un limite illimitato significa che la raccomandazione dell’8° PQ non è mai stata sicura per questo modello. Nessuno dei due numeri è noto finché non viene eseguito.
Quando prenderlo
Compresso nella decisione, che ha la stessa forma della guida alle decisioni della GPU:
- Oltre ~ 2K timestep su una GPU CUDA: sì, incondizionatamente. Si tratta di una modifica di una riga che produce un output esatto e la vittoria cresce con . Non esiste uno scenario in cui desideri che si materializzi percorso invece.
- Al di sotto di ~512 intervalli temporali, su CPU o con architetture di non attenzione (CNN, SSM come Mamba): irrilevante. A sinistra del crinale, l'overhead fisso rappresenta l'intero costo e l'attenzione non è mai stata il collo di bottiglia.
- Le soglie di cui sopra sono folklore, non misurazioni: provengono dalla letteratura generale e il crossover sul tuo modello e sulla tua scheda è un punto di riferimento di dieci righe. Eseguilo piuttosto che fidarti dei numeri tondi.
Conclusione
Flash Attention è un risultato pulito e davvero importante: rispettando la gerarchia della memoria e riassociando il softmax, calcola l'attenzione esatta con memoria invece di , e il Il limite IO spiega esattamente il motivo. Adottarlo in un trasformatore di trading è un cambiamento di una riga senza costi di precisione nell'aritmetica esatta e una grande vincita di memoria.
Ciò che non fa è rispondere alla domanda in alto. Converte "un contesto di un'intera giornata è impossibile" in "un contesto di un'intera giornata è economico", che rappresenta una modifica nel costo dell'esperimento, non nel suo risultato. Il muro della memoria che crolla è un invito a misurare, e la misurazione è ciò che trasforma questo da riassunto cartaceo in constatazione.
Riferimenti
- Dao, T., Fu, D.Y., Ermon, S., Rudra, A., Re, C. "FlashAttention: attenzione esatta veloce ed efficiente in termini di memoria con consapevolezza IO." NeurIPS (2022). arXiv:2205.14135
- Dao, T. "FlashAttention-2: attenzione più rapida con migliore parallelismo e partizionamento del lavoro." ICLR (2024). arXiv:2307.08691
- Dao, T., Shah, J. "FlashAttention-3: attenzione veloce e accurata con asincronia e bassa precisione." NeurIPS (2024). arXiv:2407.08608
- Vaswani, A., et al. "L'attenzione è tutto ciò di cui hai bisogno." NeurIPS (2017).
- Milakov, M., Gimelshein, N. "Calcolo del normalizzatore online per softmax." arXiv:1805.02867 (2018).
Autori
Trading-systems engineer
Trading-systems engineer building bots since 2017: cross-exchange arbitrage (connected up to 30 venues), cointegration-based pairs arbitrage across spot and futures, scalping, news and sentiment-driven strategies, trend algorithms, and portfolio management and balancing algorithms. Also builds sub-millisecond order execution, big-data warehouses, backtesting engines, AI agents, and trading interfaces (incl. open-source profitmaker.cc). Stack: JS/TS, Python, Rust/Zig/Go, DevOps, backend, frontend, architecture.