Transformers के बाद: Self-Attention से आगे की राह
Transformers की quadratic attention लागत असली बाधा है। Linear attention, attention residuals और hybrid architectures दिखाते हैं कि आगे क्या आएगा।

Transformer आर्किटेक्चर अब आठ साल से AI को चला रहा है। हर बड़ा language model, ज़्यादातर image generation सिस्टम, और बढ़ती संख्या में audio और video models उस self-attention मैकेनिज़्म पर बने हैं जो 'Attention Is All You Need' पेपर में आया था। लेकिन self-attention की एक बुनियादी समस्या है: इसकी compute और memory लागत sequence length के साथ quadratically बढ़ती है। इनपुट की लंबाई दोगुनी करो, तो लागत चार गुना हो जाती है।
छोटे sequences के लिए यह कोई मायने नहीं रखता। लेकिन 128K-token context windows के लिए, जिनकी तरफ़ हम बढ़ रहे हैं, और जिन million-token windows की लोगों को चाहत है, यह एक गंभीर बाधा है। कई नए शोध विकल्प तलाश रहे हैं: attention residuals जो layers के बीच computation दोबारा इस्तेमाल करते हैं, linear attention variants जो quadratic लागत हटा देते हैं, और hybrid architectures जो attention को सस्ते मैकेनिज़्म के साथ मिलाते हैं। Transformer खत्म नहीं हो रहा, पर उसका रूप बदला जा रहा है।
Self-Attention इतना महंगा क्यों है
विकल्पों को समझने के लिए पहले यह समझना होगा कि self-attention असल में क्या गणना करता है। N tokens के sequence में, self-attention हर जोड़ी tokens के बीच relevance score निकालता है। Token 1 बनाम token 2, token 1 बनाम token 3, ..., token 1 बनाम token N, फिर token 2 बनाम बाकी हर token, और इसी तरह। यानी N² जोड़ियाँ।
import torch
import torch.nn.functional as F
def self_attention(Q, K, V):
"""
Standard self-attention.
Q, K, V: (batch, seq_len, d_model)
The attention matrix is seq_len × seq_len.
For seq_len = 1024: ~1M entries (manageable)
For seq_len = 32768: ~1B entries (expensive)
For seq_len = 131072: ~17B entries (very expensive)
"""
d_k = Q.size(-1)
# This matmul creates the N×N attention matrix
scores = torch.matmul(Q, K.transpose(-2, -1)) / (d_k ** 0.5)
weights = F.softmax(scores, dim=-1)
return torch.matmul(weights, V)
4K tokens पर attention matrix में 1.6 करोड़ entries होती हैं, जो आधुनिक GPUs के लिए कोई समस्या नहीं। 128K tokens पर इसमें 16 अरब entries होंगी। 10 लाख tokens पर यह एक ट्रिलियन से भी ऊपर जाती है। Flash Attention के साथ भी (जो computation कम नहीं करता, पर memory access patterns को काफ़ी बेहतर बनाता है), आखिर में quadratic scaling ही जीतता है।
इसीलिए शुरुआती transformers 512 या 1024 tokens तक सीमित थे। हर नई पीढ़ी के hardware और optimization ने यह सीमा ऊपर खिसकाई है, पर हम एक गणितीय दीवार से लड़ रहे हैं। Linear scaling (O(N)) मूल रूप से quadratic scaling (O(N²)) से बेहतर है, और ज़्यादातर वैकल्पिक architectures इसी दिशा में जा रहे हैं।
Attention Residuals: जो पहले से गणना हो चुकी है, उसे दोबारा इस्तेमाल करना
Attention की लागत कम करने का सबसे व्यावहारिक तरीका attention को हटाता नहीं, बल्कि हर attention layer को पिछली layers की computation दोबारा इस्तेमाल करके सस्ता बनाता है।
अवलोकन यह है: एक गहरे transformer (मान लीजिए 32 layers) में आस-पास की layers के attention patterns अक्सर बेहद मिलते-जुलते होते हैं। Layer 15 और layer 16 आमतौर पर लगभग एक जैसी जगहों पर ध्यान देते हैं, बस थोड़े बदलाव के साथ। हर layer पर पूरा N² attention matrix शुरू से निकालना दोहराव है, क्योंकि ज़्यादातर काम एक layer पहले ही हो चुका होता है।
Attention residuals इसका फ़ायदा उठाते हैं। ये एक 'residual' attention pattern निकालते हैं: यानी इस layer को किस पर ध्यान देना है और पिछली layer ने जो गणना की, उनके बीच का अंतर। अगर यह अंतर छोटा है (जो बीच की layers में आमतौर पर होता है), तो गणना सस्ती पड़ती है। पूरा attention pattern पिछली layer के pattern और इस layer के residual के योग के बराबर होता है।
यह वीडियो compression जैसा है: हर frame को अलग-अलग स्टोर करने के बजाय आप एक keyframe रखते हैं और उससे आगे के फ़र्क (residuals) की एक श्रृंखला। ये फ़र्क आमतौर पर पूरे frame से कहीं छोटे होते हैं, इसलिए compression बहुत बेहतर हो जाता है।
व्यवहार में, attention residuals गहरे models की बीच की layers में attention की compute लागत 30-50% तक घटा देते हैं, और quality पर असर न्यूनतम रहता है। पहली और आख़िरी कुछ layers को अब भी पूरा attention चाहिए (क्योंकि उनके patterns ज़्यादा अलग होते हैं), लेकिन बीच की layers, जो संख्या में बहुमत में हैं, काफ़ी तेज़ हो जाती हैं।
Linear Attention: Quadratic लागत को छोड़ना
Linear attention variants attention को इस तरह दोबारा लिखने की कोशिश करते हैं कि वह O(N²) की जगह O(N) पर scale हो। आम तरीका: N×N attention matrix को सीधे निकालने की बजाय, linear operations से वही (या लगभग वही) output निकालने का रास्ता ढूँढना।
गणित की चाल softmax के kernel decomposition पर टिकी है। सामान्य attention softmax(QK^T)V निकालता है। अगर आप softmax की जगह कोई ऐसा अलग kernel function रखें जिसे φ(Q) · φ(K)^T के रूप में तोड़ा जा सके, तो गणना का क्रम बदला जा सकता है: (φ(Q) · φ(K)^T) · V की जगह, जिसमें N×N बीच का परिणाम बनता है, φ(Q) · (φ(K)^T · V) निकालें, जिसमें d×d बीच का परिणाम बनता है (जहाँ d model dimension है)। लंबे sequences के लिए d << N होता है, तो यह काफ़ी सस्ता पड़ता है।
def linear_attention(Q, K, V, feature_map=None):
"""
Linear attention via kernel feature maps.
Cost: O(N * d^2) instead of O(N^2 * d)
"""
if feature_map is None:
# ELU+1 is a common choice (from Katharopoulos et al.)
feature_map = lambda x: F.elu(x) + 1
Q = feature_map(Q) # (batch, seq_len, d)
K = feature_map(K) # (batch, seq_len, d)
# Key insight: compute K^T @ V first (d × d matrix)
# instead of Q @ K^T first (N × N matrix)
KV = torch.einsum('bnd,bnm->bdm', K, V) # (batch, d, d)
# Then multiply by Q
output = torch.einsum('bnd,bdm->bnm', Q, KV) # (batch, N, d)
# Normalize
Z = torch.einsum('bnd,bd->bn', Q, K.sum(dim=1)) # normalization
output = output / Z.unsqueeze(-1)
return output
पेच यह है: softmax की जगह दूसरा kernel रखने से attention का distribution बदल जाता है, और softmax attention से trained models linear attention पर हमेशा ठीक से नहीं चलते। Quality का अंतर काफ़ी घटा है, हाल के linear attention variants softmax attention की 95-98% quality तक पहुँच रहे हैं, पर अंतर बना हुआ है, ख़ासकर उन tasks में जिनमें लंबी दूरी से सटीक retrieval चाहिए।
State Space Models: एक अलग दृष्टिकोण
Mamba जैसे state space models (SSMs) बिल्कुल अलग रास्ता लेते हैं। tokens के बीच जोड़ीवार रिश्ते निकालने की बजाय, ये sequence को एक recurrence से प्रोसेस करते हैं, जो हर token पर एक निश्चित आकार की hidden state को अपडेट करता रहता है। यह स्वाभाविक रूप से O(N) है: दोगुने tokens प्रोसेस करने में दोगुना समय लगता है, चार गुना नहीं।
आधुनिक SSMs की खासियत यह है कि recurrence के parameters इनपुट पर निर्भर होते हैं (selective state spaces)। इससे model को एक तरह का content-based attention मिल जाता है: वह चुन सकता है कि कौन सी जानकारी याद रखनी है और कौन सी भूलनी है, और quadratic लागत के बिना। Mamba-style models कई benchmarks पर transformer की quality के बराबर पहुँच जाते हैं, और लंबे sequences पर काफ़ी तेज़ रहते हैं।
ट्रेड-ऑफ़ यह है: SSMs tokens को क्रम से प्रोसेस करते हैं, इसलिए training के दौरान उन्हें parallelize करना transformers के मुकाबले मुश्किल है (transformers सारे tokens एक साथ प्रोसेस कर सकते हैं)। Training की दक्षता मायने रखती है। जो model inference पर 2 गुना तेज़ है पर training में 3 गुना धीमा, वह ज़रूरी नहीं कि जीत हो, क्योंकि कुल compute का ज़्यादातर हिस्सा training में ही जाता है।
Hybrid Architectures: व्यावहारिक रास्ता
Production models में अभी का चलन ऐसे hybrid architectures का है जो अलग-अलग attention मैकेनिज़्म को मिलाते हैं। तर्क सीधा है: model के अलग-अलग हिस्सों को अलग तरह की computation से फ़ायदा होता है।
- वैश्विक reasoning के लिए Full attention। कुछ layers को पूरे sequence में देखना पड़ता है, हज़ारों tokens दूर का relevant context ढूँढने के लिए। ये layers standard (संभवतः Flash-optimized) self-attention लेती हैं।
- आस-पास के context के लिए Local attention। कई layers मुख्य रूप से पास के tokens पर ध्यान देती हैं (sliding window attention)। 256-1024 tokens की निश्चित window से लागत O(N·W) रह जाती है, जहाँ W window का आकार है।
- व्यापक context के लिए Linear attention। कुछ layers को पूरे sequence की जानकारी इकट्ठा करनी होती है, पर उन्हें सटीक attention weights नहीं चाहिए। Linear attention यह O(N) लागत पर देता है।
- क्रमिक प्रोसेसिंग के लिए SSM layers। Mamba-style layers बिना किसी attention गणना के क्रमिक निर्भरताओं को कुशलता से प्रोसेस कर सकती हैं।
Jamba (AI21) जैसे models और कई research architectures layer की भूमिका के हिसाब से इन मैकेनिज़्म के बीच बदलते रहते हैं। शुरुआती layers local attention लेती हैं (syntax और स्थानीय patterns समझने के लिए)। बीच की layers linear attention या SSMs लेती हैं (व्यापक representations बनाने के लिए)। कुछ रणनीतिक layers full attention लेती हैं (global reasoning और retrieval के लिए)। इससे कुल scaling लगभग linear रहती है, और वह model quality बनी रहती है जिसके लिए कुछ full attention ज़रूरी है।
Developers किस पर नज़र रखें
अगर आप language models के ऊपर applications बना रहे हैं, तो नीचे हो रहे architectural बदलाव आपके काम पर ठोस असर डालेंगे।
- Context windows और बढ़ते रहेंगे। Attention की लागत घटने के साथ context windows फैलेंगे। इससे application architecture बदलता है: 4K window में relevant context फिट करने के लिए जटिल RAG pipelines बनाने की बजाय, शायद आप सब कुछ 1M-token prompt में डाल देंगे। सरलता आकर्षक है, पर latency और लागत का असर architectures के हिसाब से अलग होता है।
- Latency के पैटर्न बदलते हैं। Transformers की latency एक बिंदु तक लगभग सपाट रहती है, फिर quadratically बढ़ती है। Linear-attention और SSM models में latency का बढ़ना ज़्यादा धीरे और linear होता है। जिन applications में response time मायने रखता है, वहाँ अपने model के scaling व्यवहार को समझना ज़रूरी है।
- Quality का अंतर task पर निर्भर है। Linear attention models लंबे context में किसी खास स्थिति से सटीक retrieval वाले tasks में थोड़ा कमज़ोर पड़ सकते हैं ('पेज 47 की सूची में तीसरा आइटम क्या था?')। सामान्य समझ वाले tasks में वे बराबर प्रदर्शन करते हैं। अपना use case जानें।
- Inference optimization और ज़रूरी होगा। जैसे-जैसे models architecturally जटिल होंगे (अलग-अलग attention प्रकार मिलाते हुए), inference engines को विषम computation कुशलता से संभालना होगा। vLLM, TensorRT-LLM और इसी तरह के frameworks इसके लिए ढल रहे हैं, पर custom architectures को तुरंत समर्थन नहीं मिल सकता।
Transformer को बदला नहीं जा रहा, उसे विकसित किया जा रहा है। Self-attention अब भी tokens के बीच रिश्तों को मॉडल करने का सबसे expressive मैकेनिज़्म है जो हमारे पास है। लेकिन इसे हर जगह, हर layer में, पूरी N² लागत पर इस्तेमाल करना ज़रूरी नहीं। आने वाले कुछ सालों के models attention का इस्तेमाल सर्जिकल तरीके से करेंगे: जहाँ सबसे ज़्यादा ज़रूरत हो वहाँ पूरी सटीकता, और बाकी हर जगह सस्ते विकल्प। नतीजा ऐसे models होंगे जो तेज़ होंगे, लंबे context संभालेंगे, चलाने में सस्ते होंगे, और मौजूदा quality के बराबर या उससे बेहतर होंगे। इस पर ध्यान देना बनता है।


