3 puan yazan GN⁺ 3 시간 전 | 1 yorum | WhatsApp'ta paylaş
  • Softmax attention’dan başlayarak sabit boyutlu durum kullanan doğrusal attention’a, yalnızca hatayı kaydeden DeltaNet’e, tüm durumu zayıflatan Gated DeltaNet’e ve kanal bazında zayıflatan Kimi Delta Attention’a (KDA) kadar adım adım türetir
  • Temel doğrusal attention, geçmiş key-value dış çarpımlarının toplamını (S_t) durumunda saklayarak dizi uzunluğuna göre doğrusal çalışır; ancak yeni değeri atamak yerine mevcut ilişkiye ekleme yapan toplamalı yazma girişimi oluşur
  • DeltaNet, mevcut key’den tahmin edilen değer ile hedef value arasındaki farka (\beta_t) çarpıp kaydeder; anında yeniden oluşturma koşulu, çevrimiçi gradyan inişi ve rank-1 durum güncellemesi şeklindeki üç yorum aynı formüle çıkar
  • Gated DeltaNet, skaler (\alpha_t) ile önce tüm durumu zayıflatır; KDA ise bunu köşegen matris (D_t=\operatorname{Diag}(\alpha_t)) biçimine genişleterek her key kanalı için bilgiyi farklı oranlarda korur veya siler
  • Aynı KDA yineleme denklemi, decode için birleştirilmiş döngüsel Triton kernel ve eğitim/uzun prefill için chunk yöntemiyle çalıştırılır; chunk yöntemi, token içi bağımlılıkları üçgensel çözümle geri kurup matris çarpımıyla yeniden yapılandırır

Notasyon ve ilerleme sırası

  • Bra-ket notasyonunda (\lvert q\rangle) sütun vektörü, (\langle k\rvert) satır vektörü, (\langle k\vert q\rangle) skaler, (\lvert v\rangle\langle k\rvert) ise matristir
  • Tek bir nedensel attention head’i ve gerçek vektörler kullandığımızı, DeltaNet key’lerinin normalize edildiğini ve durumun key uzayından value uzayına eşleme yaptığını varsayarız
  • İlerleme sırası softmax attention → doğrusal attention → DeltaNetGated DeltaNetKDA şeklindedir; en sonda döngüsel ve chunk tabanlı Triton uygulamalarına bağlanır
  • DeltaNet ailesindeki iki varyant, en yeni Qwen ve Kimi model ailelerinde kullanılır

İkinci dereceden karmaşıklıklı attention’dan doğrusal duruma

  • Genel nedensel softmax attention, key ile query arasındaki benzerliği hesaplar, tüm geçmiş key’lere ait skorları bir dağılım olarak normalize eder ve value vektörlerinin ağırlıklı toplamını çıktı olarak verir
  • Uzunluğu (T) olan bir dizide (T^2) adet key-query çifti vardır
    • Otoregresif çıkarımda key ve value önbelleğe alınabilir, ancak önbellek boyutu diziyle birlikte büyür
    • Yeni query de tüm geçmişi kontrol etmek zorundadır
  • Softmax paydası mevcut query’ye ve önceki tüm key’lere ortak biçimde bağlı olduğundan hesaplama sırasını basitçe yeniden düzenlemek zordur
  • Softmax kaldırılırsa çıktı, geçmiş key-value dış çarpımlarının toplamı olarak gruplanabilir
    • (S_t=\sum_{i\le t}\lvert v_i\rangle\langle k_i\rvert)
    • (S_t=S_{t-1}+\lvert v_t\rangle\langle k_t\rvert)
    • (\lvert o_t\rangle=S_t\lvert q_t\rangle)
  • Temel özdeşlik ((\lvert v\rangle\langle k\rvert)\lvert q\rangle=\langle k\vert q\rangle\lvert v\rangle) olup, tüm geçmiş key ve value’lar yerine toplanmış dış çarpım sabit boyutlu (d_v\times d_k) durumda saklanır
  • Token’lar bir kez tarandığı için dizi uzunluğuna göre doğrusal çalışır; bunun bedeli olarak softmax’in normalizasyonu ve seçiciliği kaybedilir
    • Daha gelişmiş doğrusal attention yöntemleri özellik haritaları ve normalizasyon terimleri kullanır

Doğrusal attention’da toplamalı yazma sorunu

  • Normalize edilmiş mevcut key’e (\lvert v_t\rangle\langle k_t\rvert) kaydedildikten hemen sonra aynı key ile okuma yapılırsa (S_t\lvert k_t\rangle=S_{t-1}\lvert k_t\rangle+\lvert v_t\rangle) olur
  • Yeni yazma, belleğin (v_t) döndürmesini sağlayacak şekilde atama yapmaz; mevcut dönüş değerine (v_t)’yi += yöntemiyle ekler
  • Önceki durum zaten doğru değeri döndürüyorsa aynı value iki katına çıkar; key’ler birbirine dik olmadığından her yazma mevcut yazmalarla girişim yapabilir
  • Doğrusal attention sıkıştırılmış ilişkisel bellek sağlar, ancak gereken = benzeri güncelleme yerine toplamalı güncelleme yapar

DeltaNet: değer yerine tahmin hatasını yazmak

  • DeltaNet, yeni key için mevcut tahmini (\widehat v_t=S_{t-1}k_t) önce okur ve tüm value yerine yalnızca farkı kaydeder
    • (e_t=\beta_t(v_t-S_{t-1}k_t))
    • (S_t=S_{t-1}+e_tk_t^\mathsf T)
    • Öğrenilen yazma şiddeti (\beta_t), ([0,1]) aralığındadır
  • Aynı key ile hemen tekrar okunursa ((1-\beta_t)S_{t-1}k_t+\beta_tv_t) elde edilir
    • (\beta_t=1) ise tam olarak (v_t) döndürür
    • Daha küçük değerler mevcut tahmini hedef yönüne yalnızca kısmen taşır
  • Güncelleme key uzayında yereldir
    • Mevcut key’e dik query yönlerinde dış çarpım güncellemesi 0 olduğundan yanıt değişmez
    • Yalnızca mevcut key yönündeki ilişki seçici olarak değiştirilir
  • Yeniden oluşturma kaybından türetmek

    • Durumu (S) doğrusal eşleme olarak görüp mevcut key-value çiftinin kaybını (\frac12\lVert Sk_t-v_t\rVert_2^2) olarak tanımlarsak gradyan ((Sk_t-v_t)k_t^\mathsf T) olur
    • (S_{t-1}) noktasından (\beta_t) büyüklüğünde tek bir gradyan inişi adımı yapılırsa DeltaNet’in güncelleme denklemiyle tam olarak aynı olur
    • Aynı güncelleme üç şekilde yorumlanabilir
      • Bellek işleminde (\beta_t), mevcut ilişkinin değiştirilme şiddetidir
      • Çevrimiçi öğrenmede (\beta_t), öğrenme oranıdır
      • Doğrusal cebirde tahmin hatası ile key’in rank-1 dış çarpımıdır
  • Yapılandırılmış durum geçişi

    • Güncelleme açılırsa (S_t=S_{t-1}(I-\beta_tk_tk_t^\mathsf T)+\beta_tv_tk_t^\mathsf T) olur
    • Birim key için (I-\beta_tk_tk_t^\mathsf T), mevcut key yönünde (1-\beta_t) özdeğerine, tüm dik yönlerde 1 özdeğerine sahiptir
    • Mevcut key yönündeki ilişki önce kaldırılır ve yeni ilişki eklenir; ancak tüm durumun ömür yönetimi hâlâ çözülmemiştir

Gated DeltaNet: önce tüm durumu unutmak

  • Tüm geçmiş tek bir matriste sıkıştırıldığında, duruma zaten karışmış tekil token’ları seçip atlamak mümkün değildir
  • DeltaNet mevcut key çevresini düzeltir, ancak diğer yönlerdeki eski bilgiler kalır ve gelecekteki okumalara katkı yapmaya devam edebilir
  • Gated DeltaNet, öğrenilmiş skaler koruma kapısı (\alpha_t\in[0,1]) uygular
    1. (\widetilde S_t=\alpha_tS_{t-1}) ile unutur
    2. (\widehat v_t=\widetilde S_tk_t) ile tahmin eder
    3. (e_t=\beta_t(v_t-\widehat v_t)) ile düzeltir
    4. (S_t=\widetilde S_t+e_tk_t^\mathsf T) ile kaydeder
  • Unutma → tahmin → düzeltme → yazma sırası önemlidir
    • Zayıflatmadan önce tahmin yapılırsa hatanın hesaplandığı bellek ile gerçekten güncellenen bellek farklı olur
  • Delta kuralı hedef key için değiştirmeyi, skaler kapı ise küresel silmeyi üstlenir; böylece farklı sorunları çözerler
  • Ancak tek bir (\alpha_t) tüm matrise uygulandığı için tüm key kanalları aynı oranda korunmak veya unutulmak zorundadır

Kimi Delta Attention: kanal bazında zayıflatma

  • Kimi Delta Attention, skaler (\alpha_t)’yi (d_k) boyutlu vektöre çevirir ve (D_t=\operatorname{Diag}(\alpha_t)) oluşturur
  • Durum key uzayından value uzayına eşleme yaptığından key kanalları (S)’nin sütunlarına karşılık gelir; sağdan çarpım (S_{t-1}D_t) her sütuna farklı bir koruma oranı uygular
  • KDA şu sırayla çalışır
    1. (\widetilde S_t=S_{t-1}D_t) ile key kanalı bazında zayıflatma
    2. (\widehat v_t=\widetilde S_tk_t) ile tahmin
    3. (e_t=\beta_t(v_t-\widehat v_t)) ile düzeltme
    4. (S_t=\widetilde S_t+e_tk_t^\mathsf T) ile kayıt
    5. (o_t=S_t(d_k^{-1/2}q_t)) ile okuma
  • Gated DeltaNet’ten KDA’ya kavramsal değişim yalnızca (\alpha_t)’nin (D_t)’ye yükseltilmesidir; ancak bir kanal silinirken başka bir kanal korunabilir
  • Köşegen-düşük rank geçişi

    • KDA açılırsa (S_t=S_{t-1}A_t+\beta_tv_tk_t^\mathsf T) olur; burada (A_t=D_t(I-\beta_tk_tk_t^\mathsf T))
    • (A_t=D_t-b_ta_t^\mathsf T), (b_t=D_tk_t), (a_t^\mathsf T=\beta_tk_t^\mathsf T) şeklinde yazılabildiği için köşegen-düşük rank (DPLR) geçişi olur
    • DPLR, key uzayında etki eden (d_k\times d_k) geçişi ifade eder; bellek durumunun kendisi ise hâlâ (d_v\times d_k) matristir
    • Her aile şu işlevleri ekler
      • Doğrusal attention: sabit boyutlu döngüsel bellek
      • DeltaNet: hedef yönde seçici değiştirme
      • Gated DeltaNet: tüm durum zayıflatma
      • KDA: key kanalı bazında zayıflatma
    • Uygulamalar genellikle (g_t=\log\alpha_t\le0) değerini saklar, ardından (\exp(g_t)) ile koruma oranını hesaplar
    • Transpoze edilmiş (d_k\times d_v) yerleşimli 5 adımlı referans uygulama naive_recurrent_kda içinde görülebilir

Decode için birleştirilmiş döngüsel Triton kernel

  • KDA’nın iki ana yürütme biçimi vardır
    • Birleştirilmiş döngüsel yöntem: decode, kısa diziler ve durum korumalı serving için uygundur
    • Chunk yöntemi: eğitim ve uzun prefill için uygundur
  • fused_recurrent_kda_fwd, her dizi, value head’i ve 32 genişlikli value tile’ı için bir Triton programı çalıştırır
    • BK, yaygın desteklenen yapılandırmalarda key boyutunu kapsar
    • Her program transpoze durumun [BK, BV] tile’ına sahip olur ve token’ları sırayla dolaşır
    • Farklı value tile’ları, head’ler ve diziler bağımsız çalışır
  • Kernel, durum zayıflatma, key’e göre tahmin indirgeme, residual hesaplama, dış çarpım yazma ve query okuma indirgemesini yineleme denkleminde olduğu gibi yapar
  • Tek seferde yalnızca bir yeni token’ın geldiği decode için uygundur; ancak vektör işlemlerini Tensor Core’lara verimli büyük matris çarpımlarına dönüştüremediği için eğitim ve uzun prefill için dezavantajlıdır

Chunkwise KDA: yineleme denklemini matris çarpımı olarak yeniden düzenlemek

  • Chunkwise KDA, (C) token’ı birlikte işlerken token bazlı döngüsel yöntemle tam olarak aynı durum ve çıktıları üretmelidir
  • Her chunk iki sonuç hesaplar
    • Gelen durum (S_c) ile tüm chunk işlendiğinde oluşan (S_{c+1})
    • Chunk içindeki tüm token’ların nedensel çıktısı
  • Temel zorluk, her token’ın delta hatasının aynı chunk içindeki önceki yazmalara bağlı olmasıdır
  • Kümülatif zayıflatma ve geçici hata

    • Token (i)’nin köşegen zayıflatmasını (D_i), chunk sınırından token (i)’ye kadar kümülatif zayıflatmayı (D_{0:i}=D_0D_1\cdots D_i) olarak tanımlayalım
    • Token (j)’nin yazması token (i)’ye taşınırken (D_{j+1:i}) uygulanır; köşegen matrisler olduğundan zayıflatma matrisleri kendi aralarında değişmeli olur
    • Önce chunk içindeki diğer yazmaları yok sayan geçici hata paralel hesaplanır
      • (\bar e_i=\beta_i(v_i-S_cD_{0:i}k_i))
    • İlk token dışındaki geçici hatalar, önceki chunk içi yazmaların etkisini atladığı için olduğu gibi kullanılamaz
  • Nedensel bağımlılığı geri kurmak

    • Önceki token (j)’nin mevcut token (i)’nin hatasına etkileyen katsayısı (\rho_{ij}=\beta_i k_j^\mathsf TD_{j+1:i}k_i) olarak tanımlanır
    • Gerçek hata (e_i=\bar e_i-\sum_{j<i}\rho_{ij}e_j) biçiminde sıralı bağımlılığa sahiptir
    • (\rho_{ij}) sıkı alt üçgensel matris (R_c)’ye yerleştirilirse, yığılmış hata matrisi (E_c=\bar E_c(A_c^{kk})^\mathsf T), (A_c^{kk}=(I+R_c)^{-1}) olarak hesaplanır
    • Genel yoğun ters matris gerekmez
      • (I+R_c), köşegen elemanları 1 olan üçgensel matristir
      • Her value kanalı için nedensel üçgensel çözüm yapmak yeterlidir
  • Chunk sonu durum hesaplama

    • Gelen durum chunk’ın tüm zayıflatmalarından geçer; her chunk içi yazma ise yalnızca kendisinden sonraki zayıflatmalardan geçer
    • Chunk sonuna kadar zayıflatılmış key’ler (K_c^{\mathrm{end}}) içinde satır olarak yığılırsa durum şu matris çarpımıyla düzenlenebilir
      • (S_{c+1}=S_cD_{0:C-1}+E_cK_c^{\mathrm{end}})
    • Birden çok rank-1 dış çarpım yazması tek bir matris çarpımında birleştirilerek tüm chunk durumu tek seferde ilerletilir
  • Chunk içindeki tüm çıktıları hesaplama

    • KDA mevcut token’ı yazdıktan sonra okuduğu için token (i)’nin çıktısı kendi yazmasını da içerir
    • Önceki yazma (j)’nin query (i)’ye etkileyen katsayısı (\chi_{ij}=s,k_j^\mathsf TD_{j+1:i}q_i), (j\le i) olarak tanımlanır
    • Katsayılar alt üçgensel okuma matrisi (A_c^{qk})’ye yerleştirilir
      • Üst üçgendeki sıfırlar, gelecek token’ların katkısını engeller
      • Köşegen elemanlar, mevcut token’ın kendi yazmasından sonra okuma davranışını yansıtır
    • Sınırdan her query’ye kadar zayıflatılmış vektörler (Q_c^{\mathrm{boundary}}) içine yığılırsa tüm çıktı şöyledir
      • (O_c=sS_cQ_c^{\mathrm{boundary}}+E_c(A_c^{qk})^\mathsf T)
    • İlk matris çarpımı zayıflatılmış chunk giriş durumunu okur, ikinci çarpım ise chunk içindeki nedensel yazma katkısını ekler

Chunkwise Triton pipeline

  • Chunk uygulaması tek bir dev kernel değil, birden çok kernel çağrısından oluşan bir pipeline’dır
  • Önce chunk içi kümülatif log zayıflatma hesaplanır
    • İki prefix sum arasındaki farkla, koruma vektörünü uzun uzun çarpmadan (D_{j+1:i}) ifade edilir
  • Ardından nedensel (A^{qk}) ve (A^{kk}) etkileşim matrisleri oluşturulur; (A^{kk}) ile chunk’ın düzeltilmiş yazmaları için WY formu kurulur
  • Durum kernel’i chunk’lar arasındaki tek dolaşımı yapar
    • Her chunk’a giren durumu üretir
    • Chunk’ın delta hatasını çözer
  • Giriş durumu hesaplandıktan sonra çıktı kernel’i farklı chunk ve tile’lardaki token’ları paralel işleyebilir
  • Gerçek uygulama önce 16 token’lık köşegen etkileşim bloklarını hesaplar, ardından birleştirilmiş köşegen dışı ve üçgensel çözüm kernel’ini çalıştırır
  • chunk_kda_fwd adımları koordine eder; ana giriş noktaları chunk_kda_fwd_intra, chunk_gated_delta_rule_fwd_h, chunk_gla_fwd_o_gk’dir
    • Koddaki v_new, çözülmüş hatadır
    • h, chunk giriş durumudur
    • kg, chunk sonuna kadar zayıflatılmış key’dir
  • Döngüsel yöntem ve chunk yöntemi farklı attention türleri değil, aynı KDA yineleme denkleminin iki yürütme çizelgesidir
    • Döngüsel yöntem, düşük gecikmeli decode için seri vektör işlemleridir
    • Chunk yöntemi, Tensor Core merkezli eğitim ve prefill için matris işlemleridir

1 yorum

 
GN⁺ 3 시간 전
Hacker News yorumları
  • Son 15 yıldır makine öğreniminde birleşik bir matematiksel gösterime ihtiyaç vardı; muhtemelen hâlâ da var. Eskiden bu durum daha da kötüydü; dünyanın dört bir yanındaki araştırmacıların makalelerinde birbirinden tuhaf gösterimler çıkıyordu
    Her makalede gösterim değişince anlamada sürtünme oluşuyor. En azından bu yazı en baştan gösterimi açıkça anlatıyor; bunu yapan makaleler nadir. Başta gösterim değiştirme özelliğini fark etmemiştim ama çok faydalıymış

    • Tek harfli semboller ya da açık veri tipleri yerine ∣q⟩ gibi karakterler kullanan geleneksel matematiksel gösterimi neden tercih ettiklerini anlayamıyorum. Kısa olma avantajı olabilir ama formülleri sözde kod ya da Python gibi gerçek programlama dilleriyle yazmak çok daha anlaşılır olurdu gibi geliyor
    • Bu yazı gösterimin yalnızca bir yönünü açıklıyor, ama kullanılan değişkenlerin tanımlarını vermiyor. Makine öğrenimi çalıştıysanız k, q, S'nin ne olduğunu biliyor ya da tahmin edebiliyorsunuz, ama ilgili arka plan bilgisi yoksa yazının büyük kısmı belirsiz kalıyor
    • Eskiden ben de öyle düşünürdüm ama koda bakmaktan çok daha fazla zamanımı formüllere bakarak geçiriyorum; bu yüzden sembollerin anlamını öğrendikten sonra kısa gösterimi okumak çok daha kolay geliyor. Ayrıca metinle yazınca meşhur zor iş olan isim bulma derdinden de kaçınıyorsunuz
  • “Doğrudan aklıma gelebilirdi...” deniyor ama var olmayan bir şeyi üretmek ya da birleştirmek inanılmaz derecede zor
    Biri zor işi bitirip yayımlayınca hemen ardından “o kadar da zor değilmiş”, “ben de yapabilirdim” gibi tepkiler geliyor ve her şey basit görünmeye başlıyor. Geliştirirken yeni bir şey icat ettiğimi sanıp sonradan bunun zaten 1970'lerde yapılmış ve yaygın kullanılmış olduğunu fark etmek de sık olur. Sadece benim yolum onunla hiç kesişmemiştir, o yüzden varlığını bilmiyordum

  • Bana göre bra-ket gösterimi her şeyi basit ve sezgisel hâle getiriyor. Vektör gösteriminde hangisinin yatay hangisinin dikey olduğunu karıştırıp yalnızca genel şekli takip ederken dikkatimi kaybediyordum, ama bra-ket ile her şey çok sezgiseldi
    Muhtemelen kaçırdığım çok sayıda iyi yazı vardır; başka yazıları da bu gösterime dönüştürmeyi düşünüyorum. Bu arada fizik doktoram var ve hafif disleksim bulunuyor

  • “Dış çarpım bir matristir ve iç çarpım bir sayıdır. Geçmişteki tüm anahtar ve değerleri saklamak yerine sabit boyutlu durum S_t içinde dış çarpımların toplamını saklarız” gibi bir üslup görünce bunun LLM tarafından yazılmış bir metin olduğundan emin oluyorum

    • Muhtemelen işe moda ifadeler içeren bir başlık istemekle başlamıştır
    • Claude'a tire () kullanmamasını söyleyen bir prompt verirseniz böyle bir sonuç çıkıyor
  • Görselleştirilmiş bir eğitim de var: https://snowchord.com/blog/linear-attention-visualized/

  • Böyle yazılar ve başlıklar gördüğümde benden çok daha zeki sayısız insana karşı derin bir minnettarlık ve tevazu hissediyorum. Lisede ve üniversitede çok zeki biri olarak bilinirdim ve ortalamanın üstünde zekiyimdir, ama beni toy gösterecek milyonlarca insan olduğundan da eminim
    Burada zeki derken kastım, büyük ve karmaşık kavramlar ile sistemleri kafada tutup onlar üzerine akıl yürütebilme yeteneği; bu da özellikle matematikçiler için önemli bir yetenek gibi görünüyor

    • Yapay zeka araçları işi giderek daha da hızlandırsa bile, yeni fikirlerin çoğunun kaynağı olmaya devam edecek olanın insanlar olduğunu düşünüyorum
      Bir arkadaşımla içerken yaptığımız düşünce deneylerinden biri şuydu: çocukları ekranların ve algoritmaların sunduğu popüler içerikten izole etmek, onları en ileri modeller eğitiliyormuş gibi medya ve materyal kalitesinin sıkı biçimde kontrol edildiği öğrenme dostu bir ortamda yetiştirmek. Çocuklar için bir manastır gibi, ama onlara matematik, mühendislik, bilgisayar bilimi, derin öğrenme vb. üzerinden gerçekliğe dair en güncel bilgileri öğreten bir yer
      Sonuçta ileri düzey yapay zeka araçlarını kullanarak bilgi sınırlarını genişletmek için hâlâ çok zeki ve düşüncesi fazla kirlenmemiş insanlara ihtiyaç olacak. Yapay zekanın insanın yerini tamamen alacağı fikri yanlış bir yön
  • Bu arada bra-ket gösterimi adı gerçekten bracket sözcüğünden geliyor
    https://en.wikipedia.org/wiki/Bra-ket_notation

  • Başta tereddüt ettim ama ket gösterimi sayesinde işlemler çok daha netleştiği için hoşuma gitti. Yine de ikinci dereceden attention'daki d_k gibi bazı değişkenler hakkında kısa bir hatırlatma da olsaymış iyi olurdu

  • Başta bu çözümü akıl edememiş olmama moralim bozuldu, ama JavaScript'te ikili aramayı bile kendi başıma yazmakta zorlandığımı fark edince hemen rahatladım. Kimi Delta Attention'ı benim düşünmüş olma ihtimalim sıfır

    • Lineer cebir kodunu yazmak aslında şaşırtıcı biçimde daha kolay olabiliyor. Genel bilgisayar bilimi kodlarında olduğu gibi özyineleme karmaşık biçimde iç içe geçmiyor, tüm değişkenler arasında matematiksel ilişkiler var ve yaygın matematik kavramları için zaten iyi uygulanmış kütüphaneler kullanılabiliyor
      Döngülerin de iki ya da üç seviyeden fazla derinleşmesi nadirdir; daha karmaşıksa zaten onu bir kütüphaneye bırakmak daha iyidir