जर तुम्ही कधी एखाद्या LLM ला मोठा प्रतिसाद तयार करताना पाहिला असेल आणि सुरुवातीच्या प्रॉम्प्टनंतर तो इतका संथ का होतोय असा प्रश्न तुम्हाला पडला असेल, तर तुम्ही प्रत्यक्षपणे हार्डवेअरमधील अडथळा (bottleneck) अनुभवत आहात. बहुतेक डेव्हलपर्स त्यांच्या Python कोडला, फ्रेमवर्कला किंवा मॉडेलच्या प्रचंड आकाराला दोष देतात. ते फंक्शन्स प्रोफाइल करतात, ऑप्टिमायझर्स बदलतात आणि प्रीप्रोसेसिंगमधील काही मिलिसेकंद वाचवण्याचा प्रयत्न करतात. पण यापैकी काहीही मूळ समस्या सोडवत नाही. वेगाची मर्यादा तुमच्या सॉफ्टवेअरमध्ये नाही, तर ती सिलिकॉनमध्ये आहे.
प्रत्येक लार्ज लँग्वेज मॉडेल इन्फरन्स (inference) कार्य तुमच्या सर्व्हरमधील GPU च्या दोन भौतिक वैशिष्ट्यांवर अवलंबून असते: ते किती वेगाने आकडेमोड (crunch numbers) करू शकते आणि ती आकडेमोड करण्यासाठी ते आकडे किती वेगाने योग्य ठिकाणी पोहोचवू शकते.
गणित स्वस्त आहे. डेटा हलवणे महाग आहे
GPU मार्केटिंगला 'कम्प्युट' (compute) बद्दल बोलायला खूप आवडते. प्रति सेकंद ट्रिलियन्स फ्लोटिंग-पॉइंट ऑपरेशन्स. ही आकडेवारी थक्क करणारी आहे. पण कम्प्युट ही केवळ अर्धी गोष्ट आहे. दुसरी अर्धी गोष्ट म्हणजे मेमरी बँडविड्थ (memory bandwidth), म्हणजेच हाय-बँडविड्थ मेमरीमधून कम्प्युट कोअर्समध्ये (जिथे प्रत्यक्ष गणिती प्रक्रिया होते) डेटा ज्या वेगाने प्रवास करतो तो दर.
एक LLM या दोन दुव्यांपैकी जो दुवा कमकुवत असेल, त्यापेक्षा जास्त वेगाने चालू शकत नाही. एका व्यावसायिक स्वयंपाकघराची कल्पना करा जिथे वीस मास्टर शेफ आहेत. ओव्हन गरम आहेत, चाकू धारदार आहेत आणि प्रत्येक आचारी तयार आहे. पण भाजीपाल्याची डिलिव्हरी सायकलने येत आहे, एका वेळी फक्त एक टोपली. यामुळे स्वयंपाकघर थांबते. अधिक आचारी आणल्याने हे सुधारेल असे नाही. अधिक वेगवान ओव्हन खरेदी केल्यानेही हे सुधारेल असे नाही. अडथळा (bottleneck) हा रस्त्यात आहे.
आधुनिक डेटासेंटर GPU मध्ये, गणिती युनिट्स (arithmetic units) इतकी शक्तिशाली असतात की ते अनेकदा आपली गणना पूर्ण करतात आणि त्यानंतर वेट्स (weights) आणि ॲक्टिव्हेशन्स (activations) मेमरीमधून येईपर्यंत रिकामे बसून राहतात. हा असमतोल तुमच्या कोडमधील बग (bug) नाही. चिप्स कशा प्रकारे बनवल्या जातात, हे त्याचे भौतिक वास्तव आहे. मेमरी बँडविड्थ कच्च्या कम्प्युटच्या (raw compute) वेगाशी स्पर्धा करू शकलेली नाही, आणि LLMs साठी हा असमतोल अधिक त्रासदायक ठरतो कारण त्यांच्या 'फॉरवर्ड पास' (forward pass) साठी प्रत्येक आउटपुट टोकनसाठी प्रत्येक पॅरामीटरला स्पर्श करणे आवश्यक असते.
प्रॉम्प्ट्स वेगवान का वाटतात आणि जनरेशन संथ का वाटते
LLM इन्फरन्स दोन वेगवेगळ्या टप्प्यांमध्ये विभागला जातो आणि ते हार्डवेअरवर पूर्णपणे वेगवेगळ्या प्रकारे ताण आणतात.
Prefill तेव्हा घडते जेव्हा तुमचा प्रॉम्प्ट पहिल्यांदा मॉडेलकडे पोहोचतो. सर्व टोकन्स एकत्र येतात. GPU मोठ्या मॅट्रिक्स-मॅट्रिक्स मल्टिप्लिकेशनचा (matrix-matrix multiplications) वापर करून त्यांना समांतर (parallel) प्रक्रिया करू शकते. हजारो गणिती युनिट्स एकाच वेळी कार्यरत होतात आणि वर्कलोड दाट (dense) राहतो. हा टप्पा 'कम्प्युट-बाउंड' (compute-bound) असतो. सुरुवातीला तुम्हाला दिसणारा तो अचानक वेगाचा विस्फोट? तो GPU ने नेमके तेच काम केले आहे जे करण्यासाठी त्याला बनवले गेले आहे.
Decode मध्ये गोष्टी कठीण होतात. जेव्हा मॉडेल पुढचे टोकन तयार करते, तेव्हा ते एका वेळी एकच टोकन तयार करते. हा टप्पा मॅट्रिक्स-व्हेक्टर ऑपरेशन्सवर (matrix-vector operations) अवलंबून असतो, जे GPU च्या समांतर क्षमतेचा केवळ एक छोटासा भाग वापरतात. अधिक वाईट म्हणजे, प्रत्येक नवीन टोकनमुळे GPU ला मेमरीमधून संपूर्ण मॉडेल वेट्स पुन्हा लोड करावे लागतात. गणिती युनिट्सना काम करायचे असते, पण त्याऐवजी ते वाट पाहत बसतात. Decode हा 'मेमरी-बाउंड' (memory-bound) असतो. गणिती इंजिन्स शांत असताना, GPU प्रभावीपणे एका महागड्या ट्रॅफिक कंट्रोलरप्रमाणे काम करत असतो, जो मेमरी बसमधून पॅरामीटर्सची ये-जा करत असतो. म्हणूनच, सुरुवातीचे प्रॉम्प्ट विश्लेषण झटपट वाटत असूनही, शंभर शब्दांचा प्रतिसाद तयार व्हायला दहा सेकंद लागू शकतात.
KV कॅशे (KV cache) या गोष्टीला अधिक मनोरंजक बनवते. Decode दरम्यान, मॉडेल प्रत्येक मागील टोकनसाठी 'की' (key) आणि 'व्हॅल्यू' (value) टेन्सर्स साठवते जेणेकरून त्याला अटेंशन (attention) पुन्हा शून्यापासून मोजावे लागणार नाही. ही कॅशे सिक्वेन्सच्या लांबीनुसार वाढत जाते. ती देखील मेमरीमध्येच असते. त्यामुळे आता GPU केवळ वेट्स पुन्हा लोड करत नाहीये; तर प्रत्येक फॉरवर्ड पासमध्ये तो सतत विस्तारणाऱ्या कॅशेला वाचत आणि लिहित असतो. कम्प्युट कोअर्सना फारसा त्रास होत नाही, पण मेमरी बसला मात्र दोघांसाठीही खूप कष्ट घ्यावे लागतात.
मेमरी वॉलशी (Memory Wall) लढा
किती डेटा हलवावा लागेल हे कमी करण्यासाठी, किंवा किमान तो हलवण्याचा खर्च विभागण्यासाठी इंजिनिअर्सनी काही तंत्रांचा संच विकसित केला आहे.
Batching हे सर्वात सोपे तंत्र आहे. जर एका वापरकर्त्याच्या विनंतीमुळे मेमरीमधून संपूर्ण वेट्स लोड करावे लागत असतील, तर एकाच वेळी आठ किंवा सोळा विनंत्यांवर प्रक्रिया केल्यामुळे GPU तो भार सर्व विनंत्यांमध्ये विभागू शकतो. वेट्स एकदाच वाचले जातात आणि बॅचमधील प्रत्येक सिक्वेन्ससाठी त्यांचा पुनर्वापर केला जातो. प्रोडक्शनमध्ये, प्रगत शेड्यूलिंग सिस्टम्स विनंत्यांचे डायनॅमिकली गट तयार करतात, ज्याला कधीकधी 'कंटिन्युअस' (continuous) किंवा 'इन-फ्लाइट बॅचिंग' (in-flight batching) असे म्हणतात, जेणेकरून GPU क्वचितच थांबेल. हे एकाच मार्गावरील बस आणि सोळा वेगवेगळ्या कारमधील फरकासारखे आहे.
क्वांटायझेशन (Quantization) थेट बँडविड्थच्या समस्येवर प्रहार करते. मॉडेलचे वेट्स (weights) सहसा सोळा-बिट फ्लोटिंग-पॉइंट फॉरमॅटमध्ये साठवले जातात. त्यांना आठ-बिट किंवा अगदी चार-बिट पूर्णांकांमध्ये (integers) कॉम्प्रेस करून, तुम्ही बसमधून प्रवास करणाऱ्या डेटाचे प्रमाण प्रत्यक्षपणे निम्मे किंवा त्याहून अधिक कमी करता. मॉडेलला सुसंगत आउटपुट देण्यासाठी पुरेशी अचूकता (precision) आवश्यक असते, परंतु आधुनिक पोस्ट-ट्रेनिंग क्वांटायझेशन पद्धती गुणवत्ता खराब न करता मॉडेलचा मेमरी फूटप्रिंट लक्षणीयरीत्या कमी करू शकतात. डेटाचा प्रवाह कमी असणे म्हणजे मेमरी कंट्रोलरवर प्रतीक्षा करण्यासाठी लागणारा कमी वेळ.
FlashAttention अटेंशन मेकॅनिझमची पुनर्रचना करते जेणेकरून मध्यवर्ती निकाल (intermediate results) GPU च्या वेगवान ऑन-चिप मेमरीमध्ये राहतील. स्टँडर्ड अटेंशनला मोठ्या अटेंशन मॅट्रिसेस (attention matrices) स्लो एक्सटर्नल मेमरीमध्ये लिहावे लागत आणि नंतर ते पुन्हा वाचावे लागत होते. FlashAttention गणनेचे लहान टाइल्समध्ये (tiles) विभाजन करते जे SRAM मध्ये बसतात, ऑन-चिपवर सॉफ्टमॅक्स (softmax) आणि स्केलिंगची प्रक्रिया करते आणि फक्त अंतिम आउटपुट हाय-बँडविड्थ मेमरीमध्ये परत लिहिते. हे मुख्य मेमरीकडे होणाऱ्या कमी फेऱ्यांसाठी (round trips) थोड्या अतिरिक्त गणनेचा (compute) वापर करते, जो जवळजवळ नेहमीच फायदेशीर ठरतो.
PagedAttention मेमरीच्या वेगळ्या प्रकारच्या वाया जाण्याच्या समस्येचे निराकरण करते. डिकोड दरम्यान, KV कॅश (KV cache) अनपेक्षितपणे वाढते. पारंपारिक सिस्टम्स प्रत्येक सिक्वेन्ससाठी मेमरीचे निश्चित, सलग (contiguous) भाग राखून ठेवतात, ज्यामुळे काही सिक्वेन्स लवकर संपतात आणि इतर विस्तारतात तेव्हा मेमरीमध्ये मोठे पोकळ भाग (holes) उरतात. PagedAttention ऑपरेटिंग सिस्टम्समधून व्हर्च्युअल मेमरीची संकल्पना घेते. हे KV कॅश एन्ट्रीज निश्चित आकाराच्या ब्लॉक्समध्ये साठवते जे नॉन-कंटिग्युअस पद्धतीने वाटप केले जाऊ शकतात आणि इंडायरेक्शन टेबलद्वारे मॅप केले जाऊ शकतात. यामुळे मेमरी राखीव पण अर्धवट रिकाम्या बफर्समध्ये विनाकारण पडून राहत नाही आणि मोठ्या बॅच साइजला परवानगी मिळते, ज्यामुळे फ्रॅगमेंटेशन ओव्हरहेडऐवजी मेमरी बस उपयुक्त कामात व्यस्त राहते आणि परिणामी एकूण थ्रूपुट (throughput) सुधारतो.
दृष्टिकोन बदला (Shift the Question)
जेव्हा लॅटन्सी (latency) वाढते, तेव्हा अनेक टीम्स विचारतात की त्यांनी लहान मॉडेलकडे वळले पाहिजे की त्यांचा इन्फरन्स सर्व्हर (inference server) पुन्हा लिहावा. ते प्रश्न महत्त्वाचे आहेत, पण ते दुय्यम आहेत. पहिला प्रश्न हार्डवेअरबद्दल असायला हवा. तुमचा GPU खरोखर गणनेमध्ये (computing) व्यस्त आहे की तो डेटासाठी आसुसलेला आहे?
तुमचे युटिलायझेशन मेट्रिक्स (utilization metrics) तपासा. GPU कम्प्युट ऑक्युपन्सीसोबत (compute occupancy) मेमरी बँडविड्थ सॅच्युरेशनचे प्रोफाइलिंग करा. जर तुम्हाला डिकोड दरम्यान उच्च मेमरी स्पर्धा (memory contention) आणि कमी अरिथमेटिक इंटेंसिटी (arithmetic intensity) दिसत असेल, तर ही मॉडेल आर्किटेक्चरची समस्या नाही. ही भौतिकशास्त्राची (physics) समस्या आहे. याचे समाधान अधिक स्वच्छ Python कोडमधून येणार नाही. ते अधिक आक्रमकपणे बॅचिंग करणे, पाईपमधून वेगाने जाण्यासाठी वेट्सचे क्वांटायझेशन करणे, ऑन-चिप राहण्यासाठी अटेंशनची पुनर्रचना करणे आणि KV कॅशचे व्यवस्थापन करणे यातून येईल, जेणेकरून जागा संपण्यापूर्वी तुम्ही मोठ्या बॅचेस बसवू शकाल.
एकदा का तुम्ही इन्फरन्सकडे या दृष्टिकोनातून पाहिले की, ऑप्टिमायझेशन हे यांत्रिक (mechanical) बनते. मॉडेलची बुद्धिमत्ता गोष्टींचा वेग कमी करते यांसारख्या अफवांचा पाठलाग करणे तुम्ही थांबवता आणि हार्डवेअर प्रत्यक्षात काय देऊ शकते यावर आधारित अभियांत्रिकी निर्णय घेण्यास सुरुवात करता. हाच तो बदल आहे जो स्केल होणाऱ्या प्रोडक्शन सिस्टम्सना केवळ कार्यरत राहणाऱ्या सिस्टम्सपासून वेगळे करतो.
