Tam Gün İçeriği On Dakikalık İçerikten Daha İyi mi? Flash Dikkat ve Dizi Uzunluğu Sorusu
Bu makalenin yanıtlamak istediği soru şu: eğer bir transformatör tüm işlem gününü on dakikalık bir zaman aralığı yerine bir saniyelik çözünürlükte izleyebilseydi, daha iyi tahminde bulunur muydu?
Yakın zamana kadar soramazdınız bile. Standart dikkat ihtiyaçları yani float16'da 12 kafayla 23.400 adımlık bir gün, yalnızca puan matrisi için kabaca 12,9 GB gerektirir; bu, model parametrelerinden ve çoğu kartın size vereceğinden daha fazladır. Kimse test etmeden önce soru aritmetikle kapatıldı.
Flash Attention (Dao ve diğerleri, 2022) onu açar. Dikkati yaklaşık olarak değerlendirerek değil — tamamen aynı sonucu hesaplar — hesaplamayı IO-farkında olacak şekilde yeniden yapılandırarak GPU bellek düzeyleri arasındaki trafiği en aza indirir. Buradaki gerçekten ilginç içerik budur ve bu makalenin çoğu nasıl çalıştığına harcanmaktadır: döşeme, çevrimiçi softmax yinelemesi, IO sınırı ve geriye doğru geçişli yeniden hesaplama.
Ancak mekanizma, iddia değil, kolaylaştırıcıdır. "Daha uzun bağlam daha iyidir", piyasalar ve bu blogun mevcut konumu hakkında deneysel bir ifadedir - [Temporal Fusion Transformers]'tan(/en/blog/post/temporal-fusion-transformer-trading), finansal serilerdeki aşırı uyum ve kısa geriye dönük tekrarlayan modellerdeki vanilya transformatörlerinin yüksek frekansta rekabetçi kaldığını tespit eden diğer taraftan kesiyor. Yani yazı mekanizmaya değil ölçüme kapanıyor.
Dikkat neden hafızaya bağlıdır?
Dikkat hesaplamaları — ilkelin kendisi, ticari bağlamda, Çok Ufuklu Portföy Tahmini için Geçici Füzyon Transformatörleri. Bütün sorun bunun bir satırı: ara puan matrisi öyle belleğe yazılır, softmax için tekrar okunur, tekrar yazılır ve son matmul için tekrar okunur - ve geri yayılım için saklanması gerekir.
Dikkatin aritmetik yoğunluğu , yani yaklaşık 64 FLOP/bayt — A100 sırt noktasının oldukça solunda. Düz hesaplama tavanına değil, eğimli bant genişliği tavanına oturur: GPU hareket etmeye daha fazla zaman harcar herhangi bir şeyi çarpmaktan daha fazlası. Bunun kullandığı tavan çizgisi çerçevesi - sırt noktası, eğimli mi yoksa düz tavan mı ve bir GPU'nun satın alınmaya değer olup olmadığına neden aynı mantığın karar verdiği - [GPU Ödediğinde]( bölümünde ölçülen sayılarla oluşturulmuştur./en/blog/post/when-gpu-pays-off-sweep-roofline).
Algoritmanın yararlandığı bellek hiyerarşisi
| Bellek Düzeyi | Boyut | Bant genişliği | Gecikme |
|---|---|---|---|
| HBM (Yüksek Bant Genişliğine Sahip Bellek) | 40-80GB | 2,0 TB/sn | ~400 ns |
| SRAM (Çip üzerinde, paylaşılan bellek) | 20 MB | 19 TB/sn | ~4 ns |
SRAM kabaca 10 kat daha fazla bant genişliği ve 100 kat daha düşük gecikme süresi ile kapasitenin binde biri kadardır. Flash Attention'ın yaptığı her şey bu ticaretten kaynaklanmaktadır: kapasiteden vazgeçmek, bant genişliği ve gecikme satın almak. CPU arka testinde ölçülen aynı "donanım satın almak yerine algoritmayı yeniden yapılandırma" hareketi, geriye dönük test hız merdivenidir.
Flash Dikkat algoritması
Flash Attention, dikkati SRAM'a sığacak şekilde boyutlandırılmış kareler halinde işler ve hiçbir zaman tam olarak gerçekleşmez. HBM'deki matris hiç.
Bölüm içine satır blokları ve içine sütun blokları, bir karo ve akümülatörleri çipe sığacak şekilde seçilmiştir. Her sorgu bloğu için tüm anahtar/değer bloklarını yineleyin:
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
Çevrimiçi softmax yinelemesi
Döşemeyi mümkün kılan püf noktası çevrimiçi softmax'tır. Saf bir softmax'ın satır üzerinde iki geçişe ihtiyacı vardır: biri maksimumu bulmak için (sayısal kararlılık için), diğeri üstel almak ve normalleştirmek için. Depolamayı reddettiğiniz bir satır üzerinden iki geçiş bir çelişkidir; bu nedenle Flash Attention, istatistikleri çalıştırmaya devam eder ve ilerledikçe yeniden ölçeklenir.
Bloklardan sonra :
ve çıkış akümülatörü aynı faktörle düzeltilir:
Yeni bir blok çalışma maksimumunu her yükselttiğinde, önceden biriken çıktı geriye dönük olarak şu şekilde yeniden ölçeklendirilir: - sanki yeni maksimum başından beri biliniyormuş gibi. Sonuç cebirsel olarak iki geçişli softmax ile aynıdır. Tam aritmetikte bu bir yaklaşım değildir; bu bir yeniden birleşmedir. (Sonlu hassasiyette bu farklı bir dönüş yoludur ve bu önemlidir; aşağıdaki kesinlik kontrolüne bakın.)
GÇ karmaşıklığı
Bu, galibiyetin resmi ifadesidir. Flash Dikkat gerçekleştirir
HBM erişimleri, nerede SRAM boyutundadır standart uygulama için. Dikkat payda: çipte karalama defteri ne kadar büyük olursa, gidiş-dönüş o kadar az olur, bu nedenle algoritma FLOP sayısından ziyade bellek hiyerarşisine göre belirtilir. Tipik için Ve KB, bu oran Flash Attention'ı kabaca 5-10 kat daha az erişimle destekliyor.
Geriye doğru geçiş: depolamak yerine yeniden hesapla
Dikkat yoluyla geriye yayılım normalde aşağıdakileri gerektirir: ileri pasın tutmayı reddettiği matris. Flash Attention kareleri yeniden hesaplar geri geçiş sırasında yalnızca çıktının saklanması ve softmax istatistikleri - ikisi birden , Olumsuz . Sorunun tamamı olan hafıza terimi için mütevazi miktarda gereksiz aritmetik ticareti yapıyor. Bu, tek bir operatör içinde parça ayrıntı düzeyinde uygulanan degrade kontrol noktasıyla aynı pazarlıktır.
FA2: paralellik
Flash Attention 2 (Dao, 2023) algoritmayı korudu ve planlamayı düzeltti:
- Matmul olmayan FLOP'ların sayısı daha az. FA1, tensör çekirdekleri yerine CUDA çekirdekleri üzerinde çalışan yeniden ölçeklendirme, maksimum bulma ve üstelleştirme işlemleri için gerçek zaman harcadı. FA2, yeniden ölçeklendirmeyi iç döngünün sonuna kadar erteler.
- Sıra uzunluğu üzerinden paralellik. FA1 yalnızca parti ve kafalar üzerinde paralelleştirme yapar. FA2 ayrıca sorgu blokları üzerinden paralelleşir. Bu, özellikle varlık başına çok uzun bir diziye ve 1-4'lük bir parti boyutuna sahip olduğunuz ticaret durumu için önemlidir - tam olarak parti ve kafa paralelliğinin GPU'yu aç bıraktığı rejim.
- Warp işi bölümleme. Her çözgü, bir puan hesaplamasını bölmek ve çarpıtmalar arasında azaltmak yerine, çapraz çarpıtma azaltmasını ortadan kaldırmak yerine, sorgu bloklarının farklı bir alt kümesini alır.
Bildirilen sonuç: A100'de teorik tepe FLOP'ların ~%70'i ve FA1 için ~%35.
FA3: Hazne mekaniği
Flash Attention 3 (Dao, Shah, 2024) H100'e özgü mimariye sahiptir:
- Eşzamansız warp uzmanlığı. Hopper'ın Tensör Bellek Hızlandırıcısı (TMA), HBM→SRAM'i eşzamansız olarak hareket ettirir. FA3, çözgüleri bir sonraki KV bloğu için TMA yükleri sağlayan üreticilere ve mevcut blok üzerinde işlem yapan tüketicilere böler, böylece veri hareketi aritmetiğin arkasına gizlenir.
- Araya yerleştirilmiş matmul ve softmax. Bir blok için tensör çekirdekleri üzerinde çalışırken önceki bloğun softmax'ı CUDA çekirdekleri üzerinde çalışır; iki farklı donanım birimi, zaman dilimli olmak yerine gerçekten eş zamanlı.
- Tutarsız işlemeli FP8. H100, 2x FP16 veriminde FP8'i gerçekleştirir. FP8'in saf dikkati aykırı değerler tarafından mahvoldu; FA3, aykırı büyüklüğü koordinatlar arasında yaymak için blok bazında nicemleme öncesinde vektörleri rastgele döndürür; bu, saf FP8'den 2,6 kat daha düşük sayısal hatayla rapor edilir.
| Sürüm | GPU | Kullanım | Hızlandırma ve Standart |
|---|---|---|---|
| FA1 | A100 | ~%35 | 2-4x |
| FA2 | A100 | ~%70 | 5-7x |
| FA3 (FP16) | H100 | ~%75 | 3-5x FA2'ye karşı |
| FA3 (FP8) | H100 | ~%75 | 1,6x ve FA3 FP16 |
Nedensel maskeleme, ticaretin indirim aldığı yerdir
Zaman serileri için nedensel maskeleme zorunludur - model geleceğe yönelik olmamalıdır - ve döşeme altında bu, ek bir maliyet değil, tasarruf edilen bir maliyettir. Anahtarları sorgularına göre tamamen gelecekte olan herhangi bir kutucuk tamamen atlanır, hiçbir zaman yüklenmez ve hiçbir zaman hesaplanmaz; bu da işin kabaca yarısını azaltır. PyTorch'ta bu is_causal=True; başka hiçbir şeye gerek yok.
Entegrasyon sekiz satırdan oluşur
İhtiyacınız olan kodların neredeyse hiçbiri Flash Attention ile ilgili değil. Birleştirilmiş çekirdek için açık puan matrisi yolunu değiştirin:
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,
)
Bütün değişiklik bu. q, k, v şekillidir (batch, heads, seq, head_dim); nedensellik maskesi gitti çünkü çekirdek onu oluşturuyor. Bunu girdi projeksiyonu, bloklar ve 3 sınıflı yukarı/düz/aşağı kafaya dahil edecek eksiksiz bir açıklamalı PyTorch ticaret modeli için DeepLOB ve tam bir eğitim hattı için bkz. Geçici Füzyon Transformatörleri. O iskelenin dördüncü kopyasını buraya inşa etmek hiçbir şey öğretmez.
Gereksinimler: hesaplama kapasitesi >= 8,0 (A100, H100, RTX 3090+), yarı duyarlı girişler, PyTorch >= 2,0. Gerçekten meşgul olan hızlı yolu doğrulayın torch.backends.cuda.sdp_kernel teşhis ve torch.cuda.max_memory_allocated() — SDPA, herhangi bir ön koşulun başarısız olması durumunda sessizce matematik çekirdeğine geri döner ve sessiz bir geri dönüş, tamamen yavaş çalışan bir modele benzer.
Ölçüm: Daha uzun bağlam işe yarar mı?
Yukarıdaki her şey, 32K veya 128K bağlamının artık uygun fiyatlı olduğunu söylüyor. Yararlı olup olmadığı hakkında hiçbir şey söylemiyor. Dürüst deney:
Aynı mimariyi şu adreste eğitin: Bu serinin başka bir yerinde kullanılan BTC serisindeki SDPA-Flash ile parametreleri, optimize ediciyi ve hedefi sabit tutarak dizi uzunluğunun tek değişken olmasını sağlar. İki şeyi bildirin:
- Maliyet Dönem başına ölçülen duvar saati ve
torch.cuda.max_memory_allocated()her birinde . - Fayda. Örnek dışı tahmin performansına kıyasla , ileriye doğru bir bölünmede.
Bu makalenin daha önceki bir taslağında analitik olarak türetilmiş dizi uzunluğu başına bellek rakamları tablosu yer alıyordu. aktivasyon formülü. Bu satırlar kaldırıldı: Hiçbir zaman ölçülmediler ve makalenin kendi bellek-bütçe aritmetiğine aykırıydılar. Sonuçlar tablosunda sunulan türetilmiş bir sayı uydurma bir sonuçtur ve bu blog bunları göndermiyor.
Bu deneyin ilginç özelliği her iki yönde de yayınlanabilir olmasıdır. Örnek dışı performans monoton bir şekilde artıyorsa Bu, tüm uzun bağlam programını haklı çıkarır. Birkaç bin adımda durağanlaşırsa veya düşerse, bu daha güçlü bir parçadır; [dürüst negatife] eşlik eder(/en/blog/post/honest-negative-no-robust-edge) - ve bu, hafıza duvarının hiçbir zaman transformatör ticareti konusunda bağlayıcı bir kısıtlama olmadığı anlamına gelir.
Daha fazla bağlam, daha fazla kapasite demektir, dolayısıyla daha fazla uyum sağlayan yüzey
Düz veya olumsuz sonucu beklemenin belirli bir nedeni vardır. Geçici Füzyon Transformatörleri vanilya transformatörlerinin finansal serilerin aşırı uyumuna saf bir şekilde uygulandığını zaten belgeliyor - zamansal tümevarımsal önyargılardan yoksunlar ve kısa geriye dönük tekrarlayan modeller yüksek frekansta rekabetçi kalıyor. Bağlamı 512 adımdan 32.768 adıma genişletmek, uzunlukla orantılı bilgi eklemez; Verimliye yakın bir fiyat serisinin marjinal 32.000'inci gecikmesi çok az şey taşıyor. Güvenilir bir şekilde eklediği şey, parametrelerin sığacak değerleridir.
Yani süpürme olduğu gibi ele alınmalıdır: bir model seçimi araması, bu blogun diğer tüm aramalara uyguladığı mekanizmanın aynısı. Üç dizi uzunluğu çarpı başka ne değişirse deneme sayımıdır ve kazananın yalnızca komşularını yenmekle kalmayıp, bu deneme sayısına göre hesaplanan Sönük Sharpe Oranı ve PBO kapısını da aşması gerekir. Aksi takdirde, "uzun bağlam kazanır" üç gürültülü koşudan en iyisinin seçilmesinden ayırt edilemez.
Kesinlik kontrolü, çünkü "tam" çok fazla iş yapıyor
Flash Attention kesindir tam aritmetik olarak. Ekteki öneri - fp16 veya bf16'da çalıştırın ve H100'de FP8'i düşünün - geçerli değildir. Bunlar ayrı iddialardır ve pratikte ikincisi hakimdir: Bir toplamı yeniden ilişkilendirmek ve kesinliği yarı yarıya düşürmek her ikisi de tedirginliktir ve sıralama garantisini getiren makale, kesinlik iddiasını elle sallamamalıdır.
Blog zaten doğru araca sahip. GPU Hassas Tuzağı standardı belirler: düşük hassasiyet sizi uyarmaz, akla yatkın çöpler getirir ve doğruluğunu göz kamaştırıcı eğrilerle değil, aşağı yönlü ayrı bir miktar - ticaret sayıları - üzerinde bir eşitlik kehaneti ile kanıtlarsınız. Burada uygulandı:
- Bf16'da SDPA-Flash ve aynı girişlerde bir FP64 referans uygulamasıyla dikkati hesaplayın; Çıkış tensöründe maksimum bağıl hata raporunu verin.
- Karara kadar ilerleyin: Yukarı/düz/aşağı etiketi veren bir model için, toplam kararların bir kesri olarak iki yol arasında kaç etiketin döndüğünü bildirin.
Küçük, sınırlı, açıklanabilir anlaşmazlıklar, doğru ve hızlı bir yolun imzasıdır. Sınırsız olan, FP8 önerisinin bu model için hiçbir zaman güvenli olmadığı anlamına gelir. Çalıştırılana kadar hiçbir sayı bilinmiyor.
Ne zaman ulaşmalı
[GPU karar kılavuzu]( ile aynı şekle sahip olan karara sıkıştırılmış/en/blog/post/when-gpu-pays-off-sweep-roofline):
- CUDA GPU'da ~2K'nın üzerinde zaman adımları: evet, koşulsuz olarak. Bu, tam çıktı üreten tek satırlık bir değişikliktir ve kazanç arttıkça artar. . Gerçekleştirilmesini istediğiniz bir senaryo yok bunun yerine yol.
- ~512 zaman adımının altında, CPU'da veya dikkat gerektirmeyen mimarilerle (CNN'ler, Mamba gibi SSM'ler): alakasız. Tepenin solunda, tüm maliyet sabit giderdir ve dikkat hiçbir zaman darboğazınız olmadı.
- Yukarıdaki eşikler ölçüm değil folklordur — bunlar genel literatürden gelir ve kendi modeliniz ve kartınız üzerindeki çaprazlama on satırlık bir kıyaslamadır. Yuvarlak sayılara güvenmek yerine çalıştırın.
Sonuç
Flash Attention temiz ve gerçekten önemli bir sonuçtur: bellek hiyerarşisine saygı göstererek ve softmax'ı yeniden ilişkilendirerek, tam dikkati bunun yerine hafıza ve IO bağlı tam olarak nedenini açıklıyor. Bunu bir alım satım transformatöründe benimsemek, tam aritmetikte doğruluk maliyeti olmayan ve büyük bir hafıza kazancı sağlayan tek satırlık bir değişikliktir.
Yapmadığı şey üstteki soruyu yanıtlamaktır. "Tam gün bağlamı imkansızdır" ifadesini "tam gün bağlamı ucuzdur"a dönüştürür; bu, denemenin sonucunda değil maliyetinde bir değişikliktir. Hafıza duvarının yıkılması, ölçüm yapmaya bir davettir ve ölçüm, bunu bir kağıt özetinden bulguya dönüştüren şeydir.
Referanslar
- Dao, T., Fu, D.Y., Ermon, S., Rudra, A., Re, C. "FlashAttention: IO Farkındalığıyla Hızlı ve Bellek Açısından Verimli Tam Dikkat." NeurIPS (2022). arXiv:2205.14135
- Dao, T. "FlashAttention-2: Daha İyi Paralellik ve İş Bölümlendirmeyle Daha Hızlı Dikkat." ICLR (2024). arXiv:2307.08691
- Dao, T., Shah, J. "FlashAttention-3: Eşzamansız ve Düşük Hassasiyetle Hızlı ve Doğru Dikkat." NeurIPS (2024). arXiv:2407.08608
- Vaswani, A., ve diğerleri. "İhtiyacınız Olan Tek Şey Dikkat." NeurIPS (2017).
- Milakov, M., Gimelshein, N. "Softmax için çevrimiçi normalleştirici hesaplaması." arXiv:1805.02867 (2018).
Yazarlar
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.