Adakah Konteks Sepenuh Hari Mengatasi Sepuluh Minit? Perhatian Kilat dan Soalan Panjang Urutan
Berikut ialah soalan yang wujud untuk dijawab oleh artikel ini: jika pengubah boleh menghadiri keseluruhan hari dagangan pada resolusi satu saat dan bukannya tetingkap sepuluh minit, adakah ia akan meramalkan dengan lebih baik?
Sehingga baru-baru ini anda tidak boleh bertanya. Keperluan perhatian standard ingatan, jadi sehari 23,400 langkah pada 12 kepala dalam float16 mahukan kira-kira 12.9 GB untuk matriks skor sahaja — lebih daripada parameter model dan lebih daripada kebanyakan kad akan diberikan kepada anda. Soalan itu ditutup dengan aritmetik sebelum sesiapa pun dapat mengujinya.
Flash Attention (Dao et al., 2022) membukanya. Bukan dengan menganggarkan perhatian — ia mengira hasil tepat yang sama — tetapi dengan menstrukturkan semula pengiraan menjadi IO-sedar, meminimumkan trafik antara tahap memori GPU. Itulah kandungan yang benar-benar menarik di sini, dan kebanyakan artikel ini dibelanjakan untuk cara ia berfungsi: jubin, pengulangan softmax dalam talian, IO terikat, dan pengiraan semula lulus ke belakang.
Tetapi mekanisme adalah pemboleh, bukan tuntutan. "Konteks yang lebih panjang adalah lebih baik" ialah pernyataan empirikal tentang pasaran dan kedudukan blog ini — daripada Temporal Fusion Transformers, yang mendapati bahawa pengubah vanila pada model overfit siri kewangan dan model berulang jangka pendek kekal berdaya saing pada frekuensi tinggi — memotong sebaliknya. Jadi artikel itu ditutup pada pengukuran, bukan mekanisme.
Mengapa perhatian terikat pada ingatan
Perhatian mengira — primitif itu sendiri, dalam konteks perdagangan, diliputi dalam Temporal Fusion Transformers for Multi-Horizon Portfolio Forecasting. Keseluruhan masalah adalah satu baris daripada itu: matriks skor perantaraan ialah , ia ditulis pada ingatan, baca semula untuk softmax, ditulis semula, dan baca semula untuk matmul akhir — dan ia mesti disimpan untuk perambatan belakang.
Keamatan aritmetik perhatian ialah , jadi kira-kira 64 FLOP/bait di — telaga kiri titik rabung A100. Ia terletak pada siling lebar jalur yang cerun, bukan siling pengiraan rata: GPU menghabiskan lebih banyak masa untuk bergerak sekeliling daripada mendarab apa-apa. Rangka kerja bumbung yang digunakan ini — titik rabung, cerun berbanding siling rata dan sebab alasan yang sama menentukan sama ada GPU berbaloi untuk dibeli sama sekali — dibina dengan nombor yang diukur dalam Apabila GPU Berbayar.
Hierarki memori yang dieksploitasi oleh algoritma
| Tahap Ingatan | Saiz | Lebar jalur | Latensi |
|---|---|---|---|
| HBM (Memori Lebar Jalur Tinggi) | 40-80 GB | 2.0 TB/s | ~400 ns |
| SRAM (Pada cip, memori dikongsi) | 20 MB | 19 TB/s | ~4 ns |
SRAM adalah kira-kira 10x lebar jalur dan 100x kependaman lebih rendah, pada seperseribu kapasiti. Segala-galanya Flash Attention mengikuti daripada perdagangan itu: melepaskan kapasiti, membeli lebar jalur dan kependaman. Langkah "menstruktur semula algoritma daripada membeli perkakasan" yang sama, diukur pada ujian belakang CPU, ialah tangga kelajuan ujian belakang.
Algoritma Perhatian Kilat
Flash Attention memproses perhatian dalam jubin bersaiz untuk dimuatkan dalam SRAM, dan tidak pernah menjadi kenyataan sepenuhnya matriks dalam HBM sama sekali.
Pembahagian ke dalam blok baris dan ke dalam blok lajur, dengan dipilih supaya jubin ditambah penumpuknya muat pada cip. Untuk setiap blok pertanyaan, ulangi semua blok nilai kunci:
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
Pengulangan softmax dalam talian
Helah yang membolehkan jubin boleh dilakukan ialah softmax dalam talian. Softmax naif memerlukan dua hantaran ke atas baris: satu untuk mencari maks (untuk kestabilan berangka), satu untuk mengeksponen dan menormalkan. Dua hantaran pada baris yang anda enggan simpan adalah percanggahan — jadi Flash Attention terus menjalankan statistik dan menskala semula semasa ia berjalan.
Selepas blok :
dan penumpuk keluaran diperbetulkan oleh faktor yang sama:
Setiap kali blok baharu meningkatkan maksimum berjalan, output terkumpul sebelum ini diskala semula secara retroaktif oleh — seolah-olah maks baharu telah diketahui sejak awal. Hasilnya adalah sama secara algebra dengan softmax dua laluan. Dalam aritmetik tepat ini bukan anggaran; ia adalah perkaitan semula. (Dalam ketepatan terhingga ia adalah laluan pembulatan berbeza, yang penting — lihat semakan ketepatan di bawah.)
Kerumitan IO
Ini adalah kenyataan rasmi kemenangan. Flash Attention berfungsi
HBM mengakses, di mana adalah saiz SRAM, terhadap untuk pelaksanaan standard. Perhatikan bahawa muncul dalam penyebut: lebih besar pad conteng pada cip, lebih sedikit perjalanan pergi dan balik, itulah sebabnya algoritma dinyatakan dari segi hierarki memori dan bukannya kiraan FLOP. Untuk tipikal dan KB, nisbah ini mengutamakan Flash Attention dengan lebih kurang 5-10x lebih sedikit akses.
Pas ke belakang: kira semula dan bukannya kedai
Penyebaran balik melalui perhatian biasanya memerlukan matriks yang hantaran hadapan hanya enggan disimpan. Perhatian Kilat mengira semula jubin daripada semasa hantaran ke belakang, hanya menyimpan output dan statistik softmax - kedua-duanya , bukan . Ia memperdagangkan jumlah aritmetik berlebihan yang sederhana untuk istilah ingatan yang menjadi masalah keseluruhan. Ini adalah tawaran yang sama seperti titik semakan kecerunan, digunakan pada butiran jubin dalam satu operator.
FA2: selari
Flash Attention 2 (Dao, 2023) menyimpan algoritma dan menetapkan penjadualan:
- Lebih sedikit FLOP bukan matmul. FA1 menghabiskan masa nyata untuk penskalaan semula, pencarian maksimum dan eksponen — operasi yang dijalankan pada teras CUDA, bukan teras tensor. FA2 menangguhkan penskalaan semula ke penghujung gelung dalam.
- ** Keselarian atas panjang jujukan.** FA1 selari ke atas kelompok dan kepala sahaja. FA2 juga selari dengan blok pertanyaan. Ini penting khususnya untuk kes dagangan, di mana anda sering mempunyai satu jujukan yang sangat panjang bagi setiap aset dan saiz kelompok 1-4 — betul-betul rejim di mana keselarian kelompok dan kepala menyebabkan GPU kebuluran.
- Pembahagian kerja meledingkan. Setiap meledingkan mengambil subset blok pertanyaan yang berbeza daripada membahagikan pengiraan skor dan mengurangkan merentas meledingkan, mengalih keluar pengurangan meledingkan silang.
Hasil yang dilaporkan: ~70% daripada FLOP puncak teori pada A100 berbanding ~35% untuk FA1.
FA3: Mekanik corong
Flash Attention 3 (Dao, Shah, 2024) adalah khusus seni bina untuk H100:
- Pengkhususan meledingkan tak segerak. Pemecut Memori Tensor (TMA) Hopper menggerakkan HBM→SRAM secara tak segerak. FA3 membahagikan warp kepada pengeluar mengeluarkan beban TMA untuk blok KV seterusnya dan pengguna mengira blok semasa, jadi pergerakan data bersembunyi di sebalik aritmetik.
- Matmul dan softmax bersilang. untuk satu blok berjalan pada teras tensor manakala softmax untuk blok sebelumnya berjalan pada teras CUDA — dua unit perkakasan berbeza, benar-benar serentak dan bukannya dihiris masa.
- FP8 dengan pemprosesan tidak koheren. H100 melakukan FP8 pada 2x pemprosesan FP16. Perhatian naif FP8 dimusnahkan oleh outlier; FA3 memutarkan vektor secara rawak sebelum pengkuantitian mengikut blok untuk menyebarkan magnitud terpencil merentas koordinat, dilaporkan pada ralat berangka 2.6x lebih rendah daripada FP8 naif.
| Versi | GPU | Penggunaan | Kelajuan lwn Standard |
|---|---|---|---|
| FA1 | A100 | ~35% | 2-4x |
| FA2 | A100 | ~70% | 5-7x |
| FA3 (FP16) | H100 | ~75% | 3-5x lwn FA2 |
| FA3 (FP8) | H100 | ~75% | 1.6x lwn FA3 FP16 |
Causal masking ialah tempat perdagangan mendapat diskaun
Penopengan sebab-akibat adalah wajib untuk siri masa — model tidak boleh melihat masa depan — dan di bawah jubin ia bukan kos tambahan tetapi kos yang disimpan. Mana-mana jubin yang kuncinya berada pada masa hadapan sepenuhnya berbanding pertanyaannya dilangkau secara langsung, tidak pernah dimuatkan dan tidak pernah dikira, memotong kira-kira separuh kerja. Dalam PyTorch ini is_causal=True; tiada lagi yang diperlukan.
Penyepaduan ialah lapan baris
Hampir tiada kod yang anda perlukan adalah mengenai Flash Attention. Tukar laluan matriks skor eksplisit untuk kernel bercantum:
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,
)
Itulah keseluruhan perubahan. q, k, v adalah berbentuk (batch, heads, seq, head_dim); topeng penyebab hilang kerana kernel membinanya. Untuk model dagangan PyTorch beranotasi lengkap untuk meletakkannya ke dalam — unjuran input, blok dan kepala atas/rata/bawah 3 kelas — gunakan DeepLOB, dan untuk saluran paip latihan penuh lihat Temporal Fusion Transformers. Membina salinan keempat perancah itu di sini tidak akan mengajar apa-apa.
Keperluan: keupayaan mengira >= 8.0 (A100, H100, RTX 3090+), input separuh ketepatan, PyTorch >= 2.0. Sahkan laluan pantas yang sebenarnya digunakan torch.backends.cuda.sdp_kernel diagnostik dan torch.cuda.max_memory_allocated() — SDPA secara senyap kembali ke kernel matematik jika mana-mana prasyarat gagal, dan sandaran senyap kelihatan betul-betul seperti model berfungsi yang hanya perlahan.
Pengukuran: adakah konteks yang lebih panjang membayar?
Semua di atas mengatakan konteks 32K atau 128K kini berpatutan. Ia tidak menyatakan sama ada ia berguna. Percubaan jujur:
Latih seni bina yang sama di dengan SDPA-Flash pada siri BTC yang digunakan di tempat lain dalam siri ini, menahan parameter, pengoptimum dan sasaran tetap jadi panjang jujukan ialah satu-satunya pembolehubah. Laporkan dua perkara:
- Kos. Jam dinding yang diukur setiap zaman dan
torch.cuda.max_memory_allocated()pada setiap . - Faedah. Prestasi ramalan di luar sampel berbanding , pada perpecahan berjalan ke hadapan.
Draf awal artikel ini membawa jadual angka ingatan setiap panjang urutan yang diperoleh secara analitik daripada formula pengaktifan. Baris tersebut dialih keluar: mereka tidak pernah diukur, dan mereka tidak bersetuju dengan aritmetik belanjawan ingatan artikel itu sendiri. Nombor terbitan yang dibentangkan dalam jadual hasil ialah hasil rekaan dan blog ini tidak menghantar nombor tersebut.
Sifat menarik percubaan ini ialah ia boleh diterbitkan dalam mana-mana arah. Jika prestasi luar sampel meningkat secara monoton dengan , yang mewajarkan keseluruhan program konteks panjang. Jika ia mendatar pada beberapa ribu langkah, atau merosot, itu adalah bahagian yang lebih kuat — teman kepada negatif jujur — dan ini bermakna dinding memori tidak pernah menjadi kekangan yang mengikat pada transformer perdagangan.
Lebih banyak konteks adalah lebih banyak kapasiti, oleh itu lebih banyak pemasangan permukaan
Terdapat sebab tertentu untuk mengharapkan hasil yang rata atau negatif. Pengubah Fusion Temporal sudah pun mendokumenkan bahawa pengubah vanila digunakan secara naif pada overfit siri kewangan — mereka tidak mempunyai bias induktif temporal, dan model berulang tinjauan pendek kekal kompetitif pada frekuensi tinggi. Melanjutkan konteks daripada 512 kepada 32,768 langkah tidak menambah maklumat berkadar dengan panjang; ketinggalan ke-32,000 marginal siri harga yang hampir cekap membawa sangat sedikit. Perkara yang boleh ditambah dengan pasti ialah nilai parameter untuk dimuatkan.
Jadi sapuan berakhir mesti dianggap seperti apa itu: carian pemilihan model, dengan jentera yang sama blog ini digunakan untuk setiap carian lain. Tiga panjang jujukan dikali dengan apa sahaja yang berbeza ialah kiraan percubaan, dan pemenang mesti menyelesaikan Nisbah Tajam Kempis yang dikira dengan kiraan percubaan itu dan pintu PBO, bukan sekadar mengalahkan jirannya. Jika tidak, "menang konteks panjang" tidak dapat dibezakan daripada memilih yang terbaik daripada tiga larian bising.
Pemeriksaan ketepatan, kerana "tepat" melakukan banyak kerja
Flash Attention adalah tepat dalam aritmetik tepat. Pengesyoran yang dilampirkan padanya — dijalankan dalam fp16 atau bf16, dan pada H100 pertimbangkan FP8 — tidak. Itu adalah tuntutan yang berasingan dan yang kedua mendominasi dalam amalan: mengaitkan semula jumlah dan menurunkan kepada separuh ketepatan kedua-duanya adalah gangguan, dan artikel yang memperkenalkan jaminan pesanan tidak seharusnya melambaikan tangan pada ketepatan.
Blog sudah mempunyai instrumen yang betul. Perangkap Ketepatan GPU menetapkan standard: ketepatan rendah tidak memberi amaran kepada anda, ia mengembalikan sampah yang munasabah, dan anda membuktikan ketepatan dengan oracle pariti pada kuantiti diskret hiliran — kiraan perdagangan — bukan dengan lengkung yang memerhati. Digunakan di sini:
- Kira perhatian dengan SDPA-Flash dalam bf16 dan dengan pelaksanaan rujukan fp64 pada input yang sama; laporkan ralat relatif maks pada tensor keluaran.
- Teruskan kepada keputusan: untuk model yang mengeluarkan label atas/rata/bawah, laporkan berapa banyak label terbalik antara dua laluan, sebagai sebahagian kecil daripada jumlah keputusan.
Perselisihan yang kecil, terhad dan boleh dijelaskan adalah tanda jalan pantas yang betul. Yang tidak terhad bermakna pengesyoran FP8 tidak pernah selamat untuk model ini. Tiada nombor diketahui sehingga ia dijalankan.
Bila hendak mencapainya
Dimampatkan kepada keputusan, yang mempunyai bentuk yang sama seperti panduan keputusan GPU:
- Di atas ~2K langkah masa pada GPU CUDA: ya, tanpa syarat. Ini ialah perubahan satu baris yang menghasilkan output tepat dan kemenangan berkembang dengan . Tiada senario di mana anda mahu perkara itu direalisasikan jalan sebaliknya.
- Di bawah ~512 langkah masa, pada CPU, atau dengan seni bina bukan perhatian (CNN, SSM seperti Mamba): tidak berkaitan. Di sebelah kiri rabung, overhed tetap ialah keseluruhan kos dan perhatian tidak pernah menjadi halangan anda.
- Ambang di atas ialah cerita rakyat, bukan ukuran — ia berasal daripada kesusasteraan umum, dan silang pada model dan kad anda sendiri ialah penanda aras sepuluh baris. Jalankannya daripada mempercayai nombor bulat.
Kesimpulan
Flash Attention ialah hasil yang bersih dan benar-benar penting: dengan menghormati hierarki memori dan mengaitkan semula softmax, ia mengira perhatian yang tepat dengan ingatan bukannya , dan IO terikat menerangkan dengan tepat mengapa. Mengguna pakainya dalam pengubah dagangan ialah perubahan satu baris tanpa kos ketepatan dalam aritmetik tepat dan kemenangan memori yang besar.
Apa yang tidak lakukan ialah menjawab soalan di bahagian atas. Ia menukarkan "konteks sehari penuh adalah mustahil" kepada "konteks sehari penuh adalah murah," yang merupakan perubahan dalam kos percubaan, bukan hasilnya. Dinding ingatan yang turun ialah jemputan untuk mengukur, dan pengukuran itulah yang mengubahnya daripada ringkasan kertas kepada penemuan.
Rujukan
- Dao, T., Fu, D.Y., Ermon, S., Rudra, A., Re, C. "FlashAttention: Perhatian Tepat Cepat dan Cekap Memori dengan Kesedaran IO." NeurIPS (2022). arXiv:2205.14135
- Dao, T. "FlashAttention-2: Perhatian Lebih Pantas dengan Keselarian yang Lebih Baik dan Pembahagian Kerja." ICLR (2024). arXiv:2307.08691
- Dao, T., Shah, J. "FlashAttention-3: Perhatian Pantas dan Tepat dengan Asynchrony dan Low-precision." NeurIPS (2024). arXiv:2407.08608
- Vaswani, A., et al. "Perhatian Adalah Semua yang Anda Perlukan." NeurIPS (2017).
- Milakov, M., Gimelshein, N. "Pengiraan normalizer dalam talian untuk softmax." arXiv:1805.02867 (2018).
Pengarang
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.