- Softmax attention से शुरू करके fixed-size state इस्तेमाल करने वाले linear attention, सिर्फ error रिकॉर्ड करने वाले DeltaNet, पूरे state को decay करने वाले Gated DeltaNet, और channel-wise decay करने वाले Kimi Delta Attention(KDA) तक को चरणबद्ध तरीके से व्युत्पन्न किया गया है
- बुनियादी linear attention, पुराने key-value outer product के योग को state (S_t) में सहेजकर sequence length के सापेक्ष linear रूप से चलता है, लेकिन नया मान assign करने के बजाय मौजूदा association में जोड़ देने की वजह से additive write interference पैदा होती है
- DeltaNet मौजूदा key से predict किए गए मान और लक्ष्य value के अंतर को (\beta_t) से गुणा करके रिकॉर्ड करता है, और immediate reconstruction condition, online gradient descent, तथा rank-1 state update — ये तीनों व्याख्याएँ एक ही सूत्र तक पहुँचती हैं
- Gated DeltaNet scalar (\alpha_t) से पहले पूरे state को decay करता है, और KDA इसे diagonal matrix (D_t=\operatorname{Diag}(\alpha_t)) तक बढ़ाकर हर key channel के लिए अलग-अलग अनुपात में जानकारी को बनाए रखने या हटाने देता है
- उसी KDA recurrence को decode के लिए fused recurrent Triton kernel और training/long prefill के लिए chunked तरीके से चलाया जाता है, जहाँ chunked तरीका token-भीतर dependencies को triangular solve से पुनर्स्थापित करके matrix multiplication के रूप में फिर से बनाता है
संकेत-लेखन और विस्तार का क्रम
- bra-ket notation में (\lvert q\rangle) column vector है, (\langle k\rvert) row vector है, (\langle k\vert q\rangle) scalar है, और (\lvert v\rangle\langle k\rvert) matrix है
- एक causal attention head और real vector का उपयोग किया गया है, और यह मान लिया गया है कि DeltaNet key normalized हैं तथा state, key space से value space में map करता है
- विस्तार का क्रम softmax attention → linear attention → DeltaNet → Gated DeltaNet → KDA है, और अंत में इसे recurrent तथा chunkwise Triton implementation से जोड़ा गया है
- DeltaNet परिवार के दो variants नवीनतम Qwen और Kimi model families में उपयोग किए जाते हैं
द्विघात जटिलता वाले attention से linear state तक
- सामान्य causal softmax attention, key और query की similarity निकालता है, सभी पुराने key के scores को distribution के रूप में normalize करता है, और फिर value vectors का weighted sum आउटपुट करता है
- लंबाई (T) वाले sequence में (T^2) key-query जोड़े होते हैं
- autoregressive inference में key और value को cache किया जा सकता है, लेकिन cache का आकार sequence के साथ बढ़ता है
- नई query को भी पूरा पुराना context देखना पड़ता है
- softmax denominator मौजूदा query और सभी पिछले key पर संयुक्त रूप से निर्भर करता है, इसलिए computation order को सरलता से पुनर्व्यवस्थित करना कठिन है
- softmax हटाने पर output को पिछले key-value outer product के योग के रूप में बाँधा जा सकता है
- (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)
- मुख्य identity है ((\lvert v\rangle\langle k\rvert)\lvert q\rangle=\langle k\vert q\rangle\lvert v\rangle), और सभी पुराने key व value की जगह summed outer product को fixed-size (d_v\times d_k) state में रखा जाता है
- token को एक बार traverse करने से यह sequence length के सापेक्ष linear रूप से चलता है, लेकिन इसकी कीमत softmax की normalization और selectivity खोना है
- अधिक परिष्कृत linear attention, feature map और normalization terms का उपयोग करते हैं
linear attention की additive write समस्या
- normalized मौजूदा key पर (\lvert v_t\rangle\langle k_t\rvert) रिकॉर्ड करने के तुरंत बाद उसी key से पढ़ने पर (S_t\lvert k_t\rangle=S_{t-1}\lvert k_t\rangle+\lvert v_t\rangle) मिलता है
- नया write, memory को (v_t) लौटाने के लिए assign नहीं करता, बल्कि मौजूदा return value में (v_t) को
+=तरीके से जोड़ता है - यदि पिछला state पहले से सही value लौटा रहा हो, तो वही value दोगुनी हो जाती है, और key आपस में orthogonal नहीं होते, इसलिए हर write पुराने write में हस्तक्षेप कर सकता है
- linear attention compressed associative memory देता है, लेकिन आवश्यक
=जैसे update की जगह additive update करता है
DeltaNet: value की जगह prediction error लिखना
- DeltaNet नए key के लिए पुराना prediction (\widehat v_t=S_{t-1}k_t) पहले पढ़ता है और पूरी value के बजाय सिर्फ अंतर रिकॉर्ड करता है
- (e_t=\beta_t(v_t-S_{t-1}k_t))
- (S_t=S_{t-1}+e_tk_t^\mathsf T)
- सीखी हुई write strength (\beta_t), ([0,1]) सीमा में होती है
- उसी key से तुरंत दोबारा पढ़ने पर ((1-\beta_t)S_{t-1}k_t+\beta_tv_t) मिलता है
- (\beta_t=1) होने पर यह ठीक (v_t) लौटाता है
- छोटा मान पुराने prediction को लक्ष्य की दिशा में केवल आंशिक रूप से ले जाता है
- update key space में local है
- मौजूदा key के orthogonal query directions में outer product update 0 होता है, इसलिए response नहीं बदलता
- सिर्फ मौजूदा key direction की association को चुनकर बदला जाता है
-
reconstruction loss से व्युत्पन्न करना
- state (S) को linear map मानकर मौजूदा key-value pair की loss (\frac12\lVert Sk_t-v_t\rVert_2^2) रखें, तो gradient ((Sk_t-v_t)k_t^\mathsf T) होता है
- (S_{t-1}) से magnitude (\beta_t) के साथ gradient descent का एक step लेने पर DeltaNet का update formula ठीक वैसा ही बनता है
- उसी update की तीन तरह से व्याख्या की जा सकती है
- memory operation में (\beta_t), मौजूदा association को replace करने की strength है
- online learning में (\beta_t), learning rate है
- linear algebra में यह prediction error और key का rank-1 outer product है
-
structured state transition
- update को खोलने पर (S_t=S_{t-1}(I-\beta_tk_tk_t^\mathsf T)+\beta_tv_tk_t^\mathsf T) मिलता है
- unit key के लिए (I-\beta_tk_tk_t^\mathsf T), मौजूदा key direction में eigenvalue (1-\beta_t) और सभी orthogonal directions में eigenvalue 1 रखता है
- यह पहले पुराने key direction की association हटाकर नई association जोड़ता है, लेकिन पूरे state की lifetime management अभी भी हल नहीं करता
Gated DeltaNet: पहले पूरे state को भूलना
- एक matrix में पूरे अतीत को compress करने पर state में पहले से merged व्यक्तिगत tokens को चुनकर छोड़ना संभव नहीं रहता
- DeltaNet मौजूदा key के आसपास correction करता है, लेकिन दूसरी directions की पुरानी जानकारी बनी रहती है और आगे की reads में योगदान देती रह सकती है
- Gated DeltaNet सीखा हुआ scalar retain gate (\alpha_t\in[0,1]) लागू करता है
- (\widetilde S_t=\alpha_tS_{t-1}) से भूलना
- (\widehat v_t=\widetilde S_tk_t) से prediction करना
- (e_t=\beta_t(v_t-\widehat v_t)) से correction करना
- (S_t=\widetilde S_t+e_tk_t^\mathsf T) से रिकॉर्ड करना
- भूलना → prediction → correction → write का क्रम महत्वपूर्ण है
- decay से पहले predict करने पर error निकालने वाली memory और वास्तव में update होने वाली memory अलग हो जाती है
- delta rule लक्ष्य key के लिए replacement संभालता है, जबकि scalar gate global deletion संभालता है; दोनों अलग समस्याएँ हल करते हैं
- लेकिन एक ही (\alpha_t) पूरे matrix पर लागू होता है, इसलिए सभी key channels को एक ही अनुपात से रखना या भूलना पड़ता है
Kimi Delta Attention: channel-wise decay
- Kimi Delta Attention scalar (\alpha_t) को (d_k)-dimensional vector में बदलता है और (D_t=\operatorname{Diag}(\alpha_t)) बनाता है
- क्योंकि state, key space से value space में map करता है, key channels (S) के columns के अनुरूप होते हैं, और right multiplication (S_{t-1}D_t) हर column पर अलग retain rate लागू करता है
- KDA इस क्रम में काम करता है
- (\widetilde S_t=S_{t-1}D_t) से हर key channel पर decay
- (\widehat v_t=\widetilde S_tk_t) से prediction
- (e_t=\beta_t(v_t-\widehat v_t)) से correction
- (S_t=\widetilde S_t+e_tk_t^\mathsf T) से रिकॉर्ड
- (o_t=S_t(d_k^{-1/2}q_t)) से read
- Gated DeltaNet से KDA तक का वैचारिक बदलाव सिर्फ (\alpha_t) को (D_t) तक बढ़ाना है, लेकिन इससे एक channel मिटाते हुए दूसरे channels बनाए रखे जा सकते हैं
-
diagonal-low-rank transition
- KDA को खोलने पर (S_t=S_{t-1}A_t+\beta_tv_tk_t^\mathsf T) मिलता है, जहाँ (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) के रूप में लिखने से यह diagonal-low-rank(DPLR) transition बन जाता है
- DPLR, key space में काम करने वाले (d_k\times d_k) transition को दर्शाता है, जबकि memory state स्वयं अभी भी (d_v\times d_k) matrix ही रहता है
- परिवार के हर चरण में यह functionality जुड़ती है
- linear attention: fixed-size recurrent memory
- DeltaNet: लक्ष्य direction का selective replacement
- Gated DeltaNet: पूरे state का decay
- KDA: key channel-वार decay
- implementation में सामान्यतः (g_t=\log\alpha_t\le0) को store किया जाता है और फिर (\exp(g_t)) से retain rate निकाला जाता है
- transposed (d_k\times d_v) layout वाला 5-step reference implementation
naive_recurrent_kdaमें देखा जा सकता है
decode के लिए fused recurrent Triton kernel
- KDA के दो मुख्य execution modes हैं
- fused recurrent mode: decode, छोटे sequence, और stateful serving के लिए उपयुक्त
- chunked mode: training और लंबे prefill के लिए उपयुक्त
fused_recurrent_kda_fwdsequence, value head, और 32-width value tile के प्रति एक Triton program चलाता हैBKसामान्य supported configuration में key dimension को cover करता है- हर program transposed state के
[BK, BV]tile का मालिक होता है और tokens को क्रम से traverse करता है - अलग-अलग value tiles, heads, और sequences स्वतंत्र रूप से execute होते हैं
- kernel, recurrence के अनुसार state decay, key पर prediction reduction, residual calculation, outer product write, और query read reduction को ज्यों का त्यों करता है
- decode में जहाँ एक बार में सिर्फ एक नया token आता है, वहाँ यह उपयुक्त है, लेकिन vector operations को Tensor Core-अनुकूल बड़े matrix multiplication में नहीं बदल पाता, इसलिए training और लंबे prefill के लिए प्रतिकूल है
Chunkwise KDA: recurrence को matrix multiplication में पुनर्व्यवस्थित करना
- Chunkwise KDA को (C) tokens साथ में प्रोसेस करते हुए token-wise recurrent mode के बिल्कुल समान state और output बनाने चाहिए
- हर chunk दो परिणाम निकालता है
- incoming state (S_c) से पूरे chunk को प्रोसेस करने के बाद का (S_{c+1})
- chunk के भीतर सभी tokens के causal outputs
- मुख्य कठिनाई यह है कि हर token की delta error उसी chunk के पहले के writes पर निर्भर करती है
-
cumulative decay और temporary errors
- token (i) के diagonal decay को (D_i), और chunk boundary से token (i) तक के cumulative decay को (D_{0:i}=D_0D_1\cdots D_i) मानें
- token (j) का write जब token (i) तक पहुँचता है, तो (D_{j+1:i}) लागू होता है; और क्योंकि diagonal matrices हैं, decay matrices आपस में commute कर सकते हैं
- पहले chunk के भीतर दूसरे writes को अनदेखा करते हुए temporary errors को parallel में निकाला जाता है
- (\bar e_i=\beta_i(v_i-S_cD_{0:i}k_i))
- पहले token को छोड़कर बाकी temporary errors, chunk के भीतर पिछले writes के प्रभाव को छोड़ देते हैं, इसलिए उन्हें सीधे उपयोग नहीं किया जा सकता
-
causal dependencies को पुनर्स्थापित करना
- पिछले token (j) का मौजूदा token (i) की error पर असर गुणांक (\rho_{ij}=\beta_i k_j^\mathsf TD_{j+1:i}k_i) से परिभाषित करें
- वास्तविक error क्रमिक निर्भरता (e_i=\bar e_i-\sum_{j<i}\rho_{ij}e_j) का रूप लेती है
- (\rho_{ij}) को strict lower-triangular matrix (R_c) में रखने पर stacked error matrix को (E_c=\bar E_c(A_c^{kk})^\mathsf T), जहाँ (A_c^{kk}=(I+R_c)^{-1}), से निकाला जाता है
- सामान्य dense inverse की जरूरत नहीं होती
- (I+R_c), diagonal पर 1 वाला triangular matrix है
- हर value channel के लिए सिर्फ causal triangular solve करना होता है
-
chunk के अंत का state निकालना
- incoming state chunk के सभी decays से गुजरता है, और chunk के भीतर हर write अपने बाद वाले decays से ही गुजरता है
- chunk के अंत तक decay हुए keys को (K_c^{\mathrm{end}}) में rows के रूप में रखने पर state को इस matrix multiplication में समेटा जा सकता है
- (S_{c+1}=S_cD_{0:C-1}+E_cK_c^{\mathrm{end}})
- कई rank-1 outer product writes को एक matrix multiplication में जोड़कर पूरे chunk state को एक बार में आगे बढ़ाया जाता है
-
chunk के भीतर सभी outputs निकालना
- KDA मौजूदा token को लिखने के बाद पढ़ता है, इसलिए token (i) के output में उसका अपना write भी शामिल होता है
- पिछले write (j) का query (i) पर असर गुणांक (\chi_{ij}=s,k_j^\mathsf TD_{j+1:i}q_i), (j\le i), से परिभाषित करें
- इन coefficients को lower-triangular read matrix (A_c^{qk}) में रखा जाता है
- upper-triangular 0 भविष्य के tokens के contribution को रोकते हैं
- diagonal elements दर्शाते हैं कि मौजूदा token अपने write के बाद खुद को पढ़ता है
- boundary से हर query तक decay हुए vectors को (Q_c^{\mathrm{boundary}}) में stack करने पर पूरा output यह बनता है
- (O_c=sS_cQ_c^{\mathrm{boundary}}+E_c(A_c^{qk})^\mathsf T)
- पहला matrix multiplication decay हुए chunk-entry state को पढ़ता है, और दूसरा chunk के भीतर causal write contributions जोड़ता है
Chunkwise Triton pipeline
- chunk implementation एक विशाल kernel नहीं, बल्कि कई kernel calls से बना pipeline है
- पहले chunk के भीतर cumulative log-decay निकाला जाता है
- दो prefix sums के अंतर से retain vectors का लंबा गुणन किए बिना (D_{j+1:i}) व्यक्त किया जाता है
- इसके बाद causal (A^{qk}) और (A^{kk}) interaction matrices बनाए जाते हैं, और (A^{kk}) से chunk के corrected writes के लिए WY form तैयार की जाती है
- state kernel ही chunks के बीच एकमात्र traversal करता है
- यह हर chunk में आने वाला state बनाता है
- chunk की delta errors को resolve करता है
- entry states निकलने के बाद output kernel, अलग-अलग chunks और tiles के tokens को parallel में प्रोसेस कर सकता है
- वास्तविक implementation पहले 16-token diagonal interaction blocks निकालता है, फिर fused off-diagonal और triangular solve kernels चलाता है
chunk_kda_fwdइन चरणों का समन्वय करता है, और मुख्य entry pointschunk_kda_fwd_intra,chunk_gated_delta_rule_fwd_h,chunk_gla_fwd_o_gkहैं- कोड में
v_new, resolved error है h, chunk-entry state हैkg, chunk के अंत तक decay हुआ key है
- कोड में
- recurrent mode और chunked mode अलग attention नहीं हैं, बल्कि एक ही KDA recurrence के दो execution schedules हैं
- recurrent mode low-latency decode के लिए serial vector operations है
- chunked mode Tensor Core-केंद्रित training और prefill के लिए matrix operations है
1 टिप्पणियां
Hacker News की राय
पिछले 15 सालों से machine learning को एकीकृत mathematical notation की ज़रूरत थी, और शायद आगे भी रहेगी। पहले तो दुनिया भर के शोधकर्ताओं के papers में और भी अजीबोगरीब notation दिखती थी
जब हर paper की notation अलग हो, तो समझने में रुकावट आती है। कम से कम यह लेख शुरुआत से notation को साफ़-साफ़ समझाता है, और ऐसा करने वाले papers कम ही होते हैं। शुरुआत में मुझे notation switching फीचर का पता भी नहीं चला, लेकिन यह बहुत उपयोगी है
∣q⟩जैसे पारंपरिक mathematical notation क्यों पसंद करते हैं। संक्षिप्त होने का फायदा होगा, लेकिन अगर formulas को pseudocode या Python जैसी असली programming language में लिखा जाए तो शायद समझना बहुत आसान होगाk,q,Sक्या हैं यह पता होगा या अंदाज़ा लगा लेंगे, लेकिन संबंधित background knowledge न हो तो लेख का बड़ा हिस्सा धुंधला लगता है“इसे सीधे सोचा भी जा सकता था…” कहा जाता है, लेकिन जो चीज़ पहले मौजूद ही नहीं थी उसे बनाना या जोड़ना बेहद मुश्किल होता है
जब कोई कठिन काम आखिरकार सामने आता है, तो तुरंत “इतना भी मुश्किल नहीं”, “मैं भी कर सकता था” जैसी प्रतिक्रियाएँ आने लगती हैं, और सब कुछ आसान दिखने लगता है। development करते समय कभी लगता है कि कुछ नया invent किया है, और बाद में पता चलता है कि वह 1970s में ही बन चुका था और खूब इस्तेमाल भी हुआ था। बस वह मेरी राह में नहीं आया था, इसलिए उसके अस्तित्व का पता नहीं था
मेरे लिए bra-ket notation सब कुछ सरल और सहज बना देती है। vector notation में अक्सर यह उलझन रहती थी कि कौन-सा horizontal है और कौन-सा vertical, और मैं बस blocks को follow करते-करते ध्यान खो देता था, लेकिन bra-ket में पूरी चीज़ बहुत intuitive लगी
लगता है मैंने बहुत-से अच्छे लेख मिस कर दिए होंगे, इसलिए अब बाकी लेखों को भी इस notation में बदलकर देखने का सोच रहा हूँ। संदर्भ के लिए, मैं physics में PhD हूँ और मुझे हल्का dyslexia है
“outer product एक matrix है और inner product एक संख्या। पिछले सभी keys और values को store करने की बजाय fixed-size state
S_tमें outer products का sum store किया जाता है” जैसी लिखावट देखकर मुझे यक़ीन हो जाता है कि यह लेख LLM ने लिखा है–) इस्तेमाल न करने को कहें, तो ऐसा नतीजा आता हैइसका एक visual tutorial भी है: https://snowchord.com/blog/linear-attention-visualized/
ऐसे लेख और शीर्षक देखते ही मैं मुझसे कहीं ज़्यादा बुद्धिमान अनगिनत लोगों के लिए गहरी कृतज्ञता और विनम्रता महसूस करता हूँ। high school और undergraduate में लोग मुझे बहुत smart मानते थे, और मैं औसत से ज़्यादा तेज़ हूँ, लेकिन ऐसे लोग निश्चित ही लाखों में हैं जिनके सामने मैं बिल्कुल नौसिखिया लगूँगा
यहाँ smart होने से मेरा मतलब है बड़े और जटिल concepts और systems को दिमाग में रखकर उन पर reasoning करने की क्षमता, और यह खासकर mathematicians के लिए महत्वपूर्ण प्रतिभा लगती है
एक दोस्त के साथ शराब पीते हुए हमने एक thought experiment किया था: बच्चों को screens और algorithms से मिलने वाले mass content से अलग रखा जाए, और cutting-edge models को train करने की तरह media और materials की quality को सख्ती से नियंत्रित करने वाले learning-friendly environment में पाला जाए। यानी बच्चों के लिए किसी मठ जैसा माहौल बनाया जाए और mathematics, engineering, computer science, deep learning आदि के ज़रिए उन्हें reality के बारे में सबसे आधुनिक ज्ञान सिखाया जाए
आखिरकार advanced AI tools का इस्तेमाल करके ज्ञान की सीमाएँ आगे बढ़ाने के लिए अब भी बहुत बुद्धिमान और अपेक्षाकृत कम प्रदूषित सोच वाले इंसानों की ज़रूरत होगी। यह सोचना कि AI इंसानों को पूरी तरह replace कर देगा, गलत दिशा है
जानकारी के लिए, bra-ket notation का नाम सचमुच bracket से ही आया है
https://en.wikipedia.org/wiki/Bra-ket_notation
शुरुआत में हिचकिचाहट थी, लेकिन ket notation की वजह से operations कहीं ज़्यादा स्पष्ट हो गए, इसलिए यह मुझे पसंद आई। हाँ, quadratic attention के
d_kजैसे कुछ variables पर छोटा-सा refresher भी होता तो अच्छा रहताशुरुआत में यह सोचकर निराशा हुई कि मैं खुद यह समाधान सोच नहीं पाया, लेकिन फिर याद आया कि JavaScript में binary search खुद लिखने में भी मुझे दिक्कत होती है, और तुरंत सुकून मिल गया। Kimi Delta Attention मेरे दिमाग से निकलता, इसकी कोई संभावना नहीं थी
loops भी शायद ही कभी दो-तीन स्तर से ज़्यादा गहरे जाते हैं, और अगर उससे ज़्यादा जटिल हो जाए तो वैसे भी उसे library को सौंप देना बेहतर होता है