ફોન પર લાર્જ લેંગ્વેજ મોડલ્સ (LLMs) ચલાવવા એ હવે માત્ર સંશોધનનો પ્રયોગ નથી રહ્યો. તે હવે વાસ્તવિકતા બની ગયો છે. છતાં, જે ક્ષણે તમે એક સાધારણ ડેમોમાંથી વાસ્તવિક પ્રોડક્ટ તરફ આગળ વધો છો, લેટન્સી (latency) ફરીથી સમસ્યા બની જાય છે. તમે મોડલને ટ્રીમ કર્યું, વેટ્સ (weights) ને ક્વોન્ટાઈઝ (quantize) કર્યા, છતાં પ્રીફિલ ફેઝ (prefill phase) ધીમો પડે છે. ટોકન્સ અટકી જાય છે. UI ફ્રીઝ થઈ જાય છે. અને આ ભૂલ ભાગ્યે જ ત્યાં આવે છે જ્યાં હોવી જોઈએ.

Android ઉપકરણો પર, LLM પ્રીફિલ દરમિયાન કમ્પ્યુટ (compute) ભાગ્યે જ બોટલનેક (bottleneck) હોય છે. મેમરી બેન્ડવિડ્થ (Memory bandwidth) મુખ્ય કારણ હોય છે. આધુનિક ફ્લેગશિપ SoCs શક્તિશાળી GPU અને NPU કોર્સ સાથે આવે છે જે મેમરી સબસિસ્ટમ તેમને ફીડ કરી શકે તેના કરતા ઘણું ઝડપી અંકગણિત (arithmetic) કરી શકે છે. જ્યારે તમે નેઈવ એટન્શન ઇમ્પ્લીમેન્ટેશનનું પ્રોફાઇલિંગ કરો છો, ત્યારે એક્ઝિક્યુશન યુનિટ્સ સંતૃપ્ત (saturated) હોતા નથી. તેઓ રાહ જોઈ રહ્યા હોય છે. DRAM ની રાહ જોઈ રહ્યા હોય છે.

તમારું ક્વોન્ટાઈઝ્ડ મોડલ હજુ પણ ધીમું કેમ લાગે છે

ઓન-ડિવાઇસ ઇન્ફરન્સ (on-device inference) માટે ક્વોન્ટાઈઝેશન એ હવે ડિફોલ્ટ પ્રથમ પગલું બની ગયું છે. FP16 થી INT8 માં વેટ્સ ઘટાડવાથી મોડલનું કદ અડધું થઈ જાય છે અને સ્ટોરેજમાં પણ ઘટાડો થાય છે. તે મદદરૂપ છે, પરંતુ તે એટન્શન લેયરની લેટન્સીને ઠીક કરતું નથી. તેનું કારણ સરળ છે: ક્વોન્ટાઈઝેશન તમે જે ડેટા સ્ટોર કરો છો તેનું પ્રમાણ ઘટાડે છે, પરંતુ તે એટન્શન મિકેનિઝમ દ્વારા કરવામાં આવતા મેમરી ટ્રાન્ઝેક્શનની સંખ્યા ઘટાડતું નથી.

ટેક્સ્ટબુક પદ્ધતિથી અમલમાં મુકાયેલું સ્ટાન્ડર્ડ મલ્ટી-હેડ એટન્શન લેયર, દરેક લેયર માટે DRAM સુધી ત્રણ વખત સંપૂર્ણ રાઉન્ડ ટ્રિપ કરે છે. Query, Key, અને Value મેટ્રિસીસ મેઈન મેમરીમાંથી વાંચવામાં આવે છે, સ્કોર્સની ગણતરી કરવામાં આવે છે, અને વચગાળાના પરિણામો ફરીથી લખવામાં આવે છે. અંકગણિત સામાન્ય છે, પરંતુ ડેટાનું હલનચલન ખૂબ જ ભારે છે. Android પર, જ્યાં પાવર અને થર્મલ બજેટ મર્યાદિત હોય છે, આ પેટર્ન મેમરી બસને નુકસાન પહોંચાડે છે. પ્રોસેસર અસરકારક રીતે એક જ પુલ ઓળંગવા માટે ત્રણ વાર ટોલ (toll) ચૂકવે છે.

જો તમે INT8 મોડલ્સ મોકલી રહ્યા હોવ અને વિચારતા હોવ કે પ્રીફિલ સ્ટેપ હજુ પણ પ્રોમ્પ્ટ લંબાઈ સાથે ક્વોડ્રેટિકલી (quadratically) કેમ વધે છે, તો આ તેનો જવાબ છે. વેટ્સ નાના છે, પરંતુ એક્ટિવેશન ટ્રાફિક હજુ પણ વિશાળ છે.

બોટલનેક મેમરી છે, ગણિત નહીં

ઉકેલ સમજવા માટે, રૂફલાઇન (roofline) જુઓ. Snapdragon 8 Gen 3 અને Dimensity 9300 જેવા ચિપ્સ પરના મોબાઈલ GPUs અને NPUs માં સૈદ્ધાંતિક કમ્પ્યુટ થ્રુપુટ (compute throughput) એટલો વધારે છે જે તેમના LPDDR5X ઇન્ટરફેસ ટકી શકે તેના કરતા ઘણો વધારે છે. એક નેઈવ એટન્શન કર્નલમાં, દરેક હેડ Q અને K ના ગુણાકાર પર softmax ની ગણતરી કરે છે, અને પછી તેને V સાથે ગુણે છે. દરેક વચગાળાનું સ્કોર મેટ્રિક્સ ગ્લોબલ મેમરીમાં મેટેરિયલાઇઝ (materialize) થાય છે. તેનો અર્થ એ છે કે તમારા પીક DRAM રીડ્સ સીક્વન્સ લંબાઈના વર્ગ સાથે, એટલે કે O(n²) સાથે વધે છે. 1024-ટોકન પ્રોમ્પ્ટ માટે, મેમરી ટ્રાફિક એક્ઝિક્યુશન સમય પર પ્રભુત્વ જમાવવા માટે પહેલેથી જ પૂરતો મોટો છે.

કોર્સનો ઉપયોગ ઓછો થાય છે કારણ કે તેઓ લેટન્સીને છુપાવી શકતા નથી. આધુનિક પ્રોસેસર્સ પાઇપલાઇનને ભરેલી રાખવા માટે કેશ (caches) પર આધાર રાખે છે. જ્યારે કોઈ અલ્ગોરિધમ સતત કેશમાં મિસ થાય છે અને DRAM માંથી ડેટા ફેચ કરે છે, ત્યારે એક્ઝિક્યુશન યુનિટ્સ ખાલી બેસી રહે છે. ક્વોન્ટાઈઝેશન ગમે તેટલું કરો, તે કમ્પ્યુટ ક્ષમતા અને મેમરી સપ્લાય વચ્ચેના આ સ્ટ્રક્ચરલ મિસમેચને ઠીક કરી શકતું નથી.

ટાઇલીંગ (Tiling) કેવી રીતે બેન્ડવિડ્થ પાછી મેળવે છે

ઉકેલ એક ટાઇલીંગ વ્યૂહરચના (tiling strategy) છે જે વચગાળાના સ્કોર્સને DRAM માં મોકલવાને બદલે ઓન-ચિપ SRAM માં રાખે છે. આ એ જ વિચાર છે જે Flash Attention ને ચલાવે છે, જેને Android ના કમ્પ્યુટ સ્ટેક માટે અનુકૂળ બનાવવામાં આવ્યો છે. મેમરીમાં સંપૂર્ણ n × n સ્કોર મેટ્રિક્સ બનાવવાને બદલે, તમે ગણતરીને નાના ટાઇલ્સમાં વિભાજિત કરો છો જે L1 કેશની અંદર સમાઈ જાય છે. તમે લોકલ softmax સ્ટેટિસ્ટિક્સની ગણતરી કરો છો, રનિંગ મેક્સ વેલ્યુઝ અને નોર્મલાઇઝેશન સમ્સ એકત્રિત કરો છો, અને ફક્ત અંતિમ વိတ်્ડ આઉટપુટ્સ જ મેમરીમાં પાછા લખો છો.

આ બેન્ડવિડ્થ કોમ્પ્લેક્સિટીને બદલે છે. પીક DRAM રીડ્સ O(n²) થી ઘટીને O(n) થઈ જાય છે, કારણ કે તમારે હવે મેઈન મેમરી દ્વારા સંપૂર્ણ સ્કોર મેટ્રિક્સ મોકલવાની જરૂર નથી. મુખ્ય કામ SRAM ની અંદર, એક્ઝિક્યુશન યુનિટ્સની બિલકુલ બાજુમાં થાય છે.

એક ચોક્કસ ઉદાહરણ માટે, 64 ની ટાઇલ સાઈઝ અને 128 ની હેડ ડાયમેન્શન ધારો. સ્કોર ટાઇલ 16 KB જગ્યા રોકે છે. આ ફૂટપ્રિન્ટ Snapdragon 8 Gen 3 અને Dimensity 9300 જેવા વર્તમાન ફ્લેગશિપ SoCs ની L1 કેશમાં આરામથી સમાઈ જાય છે. અંકગણિત લોકલ રહે છે અને મેમરી બસને રાહત મળે છે.

આને Android પર અમલમાં મૂકવું

અલ્ગોરિધમિક સ્કેચ સીધો છે, જોકે વિગતો સાચી હોવી મહત્વપૂર્ણ છે.

Break your Query, Key