Ves al contingut
13/30Capítol 13 de 30

Abaratir la inferència: KV cache, batching i quantization

El mateix model responent la mateixa pregunta en 8,8 s i en 78,9, amb sortida idèntica byte a byte. Després INT4, mesurat de tres maneres.

En aquesta pàgina

El mateix model, a la mateixa màquina, responent la mateixa pregunta amb els mateixos 48 tokens. Les dues sortides són idèntiques token per token: comprovat, no assumit.

TEXT
with a key-value cache:     8.85 s   ( 6.01 tokens/second)
without a key-value cache: 78.95 s   ( 0.60 tokens/second)

Només ha canviat un argument: use_cache=False. Res del model, el prompt, el sampling ni l’aritmètica és diferent, i la segona execució no és més precisa per l’esforç. És nou vegades més lenta per res.

Aquesta és la forma d’aquest capítol. Tot el que conté —la caché, el batch, els pesos quantitzats— és un intent de deixar de pagar per feina que no canvia la resposta, o d’esbrinar què costa una resposta més barata. El Capítol 10 va establir la llista de preus de l’entrenament. Aquesta és la llista de preus del costat que pagues per sempre: un model desplegat gasta aproximadament 2N2N FLOPs per cada token que emet, en cada petició, durant la resta de la seva vida.

On va anar a parar el temps de la segona execució

Enllaç a la secció: On va anar a parar el temps de la segona execució

Per generar un token, un transformer només decoder agafa tota la seqüència fins ara, la passa per cada capa i llegeix la distribució de probabilitat de l’última posició. Després hi afegeix el token triat i ho torna a fer. Aquesta descripció és correcta, i és el que fa l’execució lenta.

També és enormement malbaratadora, i el motiu és la màscara causal del Capítol 9. Els vectors key i value de la posició 7 es calculen a partir de l’entrada de la posició 7 i les posicions anteriors. Quan arriba la posició 8, la posició 7 no la pot veure —això és el que vol dir causal—, de manera que els key i value de la posició 7 són exactament els mateixos nombres que abans. L’execució lenta els recalcula igualment, a cada pas.

Per tant, guarda’ls. Aquest magatzem és la key-value cache, l’optimització més important en el serving de models de llenguatge:

generate.pyPYTHON
out = model(prompt_ids, use_cache=True)          # prefill: the whole prompt
past = out.past_key_values                        
nxt = out.logits[:, -1].argmax(-1, keepdim=True)

for _ in range(n - 1):
    out = model(nxt, past_key_values=past, use_cache=True)   
    past = out.past_key_values                                
    nxt = out.logits[:, -1].argmax(-1, keepdim=True)

Fixa’t en què s’alimenta al model dins del bucle: nxt, un token. No la seqüència. La query del nou token fa attention contra cada key en caché, i les keys en caché no havien de canviar mai. Això no és una aproximació: la comprovació de sortida idèntica de més amunt és el punt. La caché no intercanvia qualitat per velocitat; elimina aritmètica redundant.

Per veure’n l’escalat netament, treu el transformer i mesura un sol cap d’attention amb d=64d = 64, un pas de generació calculat de totes dues maneres:

tokens en contextrecalcular-ho totamb cachéràtiomatriu de puntuacions
1280,59 ms0,062 ms10x65.536 B vs 512 B
2561,20 ms0,163 ms7x262.144 B vs 1.024 B
5127,03 ms0,078 ms90x1.048.576 B vs 2.048 B
102417,31 ms0,114 ms152x4.194.304 B vs 4.096 B
204859,83 ms0,214 ms279x16.777.216 B vs 8.192 B
4096236,18 ms0,284 ms832x67.108.864 B vs 16.384 B

La columna de la dreta n’és la causa. Recalcular construeix la matriu d’attention completa n×nn \times n a cada pas: el O(n2)O(n^2) del requadre de notació asimptòtica del Capítol 9, pagat una vegada per token. Amb la caché construeixes una fila 1×n1 \times n: a 4.096 tokens, 67 MB de puntuacions contra 16 KB.

Comptar multiplicacions-acumulacions en lloc de mil·lisegons treu la màquina de l’argument. Per generar TT tokens des d’un inici en fred:

tokens generatsamb cachérecalculantràtio
1282,6 M192,0 M73x
51223,1 M7,36 G318x
2048293,7 M392,6 G1.336x

Per pas, la versió amb caché és lineal en el context i la que no en té és quadràtica; sumat al llarg d’una generació, O(T2)O(T^2) contra O(T3)O(T^3), amb una ràtio que creix sense límit. La diferència de nou vegades de l’inici es va mesurar sobre 48 tokens: per sota de la primera fila d’aquella taula.

La caché també canvia què ha d’estar en memòria. En una GPU de portàtil de 8 GB generant 256 tokens en fp16, agafant el pic de l’assignador i restant-ne els pesos residents:

memòria de treball màxima
amb caché21,8 MB
recalculant181,7 MB

8,3 vegades més memòria, gastada per produir els mateixos tokens més lentament. Aquesta és la promesa feta al Capítol 5, arribant des d’una direcció inesperada: allà, l’autodiff en mode invers havia de mantenir viu cada intermedi per al backward pass, i les activations dominaven la memòria d’entrenament. A inference no hi ha backward pass ni res a retenir per a ell; així, el que domina la memòria és la caché, i és una tria deliberada més que no pas un cost inevitable.

Prefill i decode són dues màquines diferents

Enllaç a la secció: Prefill i decode són dues màquines diferents

Mira de nou l’execució ràpida: el seu primer token es va comportar de manera diferent dels altres quaranta-set.

TEXT
prefill, 40 prompt tokens : 1.0224 s   ->  25.6 ms per token
decode,  47 steps         : 0.1665 s mean per step

El prompt va costar 25,6 ms per token i cada token generat va costar 166 ms. Mateix model, mateix maquinari, mateixos pesos, una diferència de sis vegades per token —i va en el sentit que la majoria no espera. El prompt és la part barata. La generació es divideix en dues fases amb física realment diferent:

Una forward pass sobre tot el prompt. Cada token es processa en paral·lel, de manera que cada matriu de pesos es carrega de memòria una sola vegada i es multiplica contra una matriu de centenars de vectors de token: un producte matriu-matriu, amb molta aritmètica per byte mogut, que és per a això que està feta una GPU. El prefill és compute-bound, i el seu cost és aproximadament lineal en la longitud del prompt.

Una forward pass per token, batch d’un i seqüència d’un. Cada matriu de pesos encara es carrega sencera de memòria i es multiplica contra un sol vector: un producte matriu-vector, amb gairebé gens d’aritmètica per byte mogut. El decode és memory-bandwidth-bound, i el seu cost per token depèn molt poc de la longitud del context.

Totes dues meitats són mesurables. Prefill, una passada sobre PP tokens:

tokens de promptsegonsms per token
160,351521,97
320,525416,42
641,049116,39
1281,655212,93
2563,096512,10

Decode, un token contra una caché de CC:

tokens en cachéms per un token
16110,05
6497,57
256108,53
1024103,86

Llegeix la segona taula dues vegades. Passar de 16 tokens de context a 1.024 —seixanta-quatre vegades més historial per fer-hi attention— no va canviar el cost d’un pas de manera mesurable. L’attention contra la caché és feina real, però queda eclipsada pel cost fix d’arrossegar mig milió llarg de pesos pel bus de memòria per produir un vector. Aquest cost fix és el motiu de tot el que ve a la secció següent.

Aquestes dues fases són l’origen dels dos números que informa qualsevol sistema de serving. Time to first token és essencialment el prefill, i creix amb el prompt, per això una conversa llarga sembla lenta d’arrencar. Tokens per second és 1/decode step1/\text{decode step}, i és aproximadament constant, per això la resposta després flueix de manera uniforme. Un xat que comença lent i després fa streaming suaument no és un truc de renderització. Són aquestes dues taules.

La caché intercanvia aritmètica per memòria, i la memòria que demana no és petita. Per cada token del context, cada capa guarda un vector key i un vector value per cada cap key-value:

bytes per token=2×L×Hkv×dhead×bytes per element\text{bytes per token} = 2 \times L \times H_{kv} \times d_{\text{head}} \times \text{bytes per element}

El 2 és per keys i values; tota la resta és l’arquitectura. Per al model mesurat al llarg d’aquest capítol —24 capes, 14 caps de query, 2 caps key-value, dimensió de cap 64— en fp16 això són 2×24×2×64×2=12,2882 \times 24 \times 2 \times 64 \times 2 = 12{,}288 bytes per token.

Les fórmules en aquest camp tenen el costum d’equivocar-se per un factor de dos, així que comprova-ho contra l’assignador en lloc de creure-t’ho:

TEXT
KV cache tensors per layer: (1, 2, 295, 64) float16
measured: 3,624,960 bytes for 295 tokens = 12,288 bytes/token
formula : 2 * 24 * 2 * 64 * 2                = 12,288 bytes/token

Exacte, i es manté exacte en totes les formes provades:

batchcontextcaché mesuradaprevistamemòria de treball màxima
15126,0 MB6,0 MB15,4 MB
116.384192,0 MB192,0 MB207,3 MB
165.536768,0 MB768,0 MB793,7 MB
84.096384,0 MB384,0 MB401,5 MB
322.048768,0 MB768,0 MB794,2 MB
641.024768,0 MB768,0 MB797,0 MB
128512768,0 MB768,0 MB816,4 MB

Les tres últimes files mereixen una segona mirada. Trenta-dos usuaris amb 2.048 tokens cadascun, seixanta-quatre amb 1.024, cent vint-i-vuit amb 512: la caché és de 768 MB en tots els casos, perquè tots tres contenen 65.536 tokens. La caché depèn només del nombre total de tokens residents, no de com es distribueixen entre usuaris. Aquest fet és la base de la secció sobre batching.

El Capítol 9 va introduir multi-query i grouped-query attention i en va ajornar el motiu fins a aquest capítol. El motiu és aquella fórmula, i concretament el HkvH_{kv} que conté.

L’attention multi-head estàndard dona a cada cap de query els seus propis caps key i value. El model d’aquí té 14 caps de query; amb attention multi-head completa, la seva caché seria de 2×24×14×64×2=86,0162 \times 24 \times 14 \times 64 \times 2 = 86{,}016 bytes per token: 84 KB en lloc de 12 KB, exactament set vegades més, la ràtio entre caps de query i caps key-value.

Multi-query attention1 ho porta al límit: tots els caps de query comparteixen un sol cap key-value. Grouped-query attention2 és el compromís que va guanyar: uns quants caps key-value, cadascun compartit per un grup de caps de query, perquè la pèrdua de qualitat de MQA era real i la de GQA no. Cap de les dues compra aritmètica. Existeixen per dividir aquella fórmula per un enter, i es van estendre per la indústria en el moment que els contextos llargs van fer que la caché fos la restricció vinculant.

I ho fa, ràpid. Per a un model de classe 7B amb 32 capes i 8 caps key-value de dimensió 128, la caché és de 128 KB per token en fp16:

tokens de contextun usuari8 usuaris64 usuaris
4.0000,49 GB3,91 GB31,2 GB
32.0003,91 GB31,25 GB250,0 GB
128.00015,62 GB125,00 GB1.000,0 GB
1.000.000122,07 GB976,56 GB7.812,5 GB

Els pesos propis d’aquest model són 13,0 GB en fp16, la xifra de la taula del final d’aquest capítol. Així que, amb un context de 128.000 tokens, la caché d’un sol usuari és més gran que el model. Aquesta és l’aritmètica que el Capítol 16 converteix en diners, i és per això que una conversa llarga no és només lenta: ocupa una porció fixa d’una màquina mentre la petició és viva.

Batching: el número que puja i el número que baixa

Enllaç a la secció: Batching: el número que puja i el número que baixa

El decode és memory-bound: els pesos s’arrosseguen pel bus per produir un token, i les unitats aritmètiques queden ocioses. Per tant, posa més feina al mateix pas. Executa diverses peticions alhora, i els pesos, llegits una vegada, serveixen totes elles. Mesurat al mateix model, cada petició mantenint una caché de 64 tokens i decodificant un token:

batchlatència per pasthroughputlatència vs B=1
10,1286 s7,78 tok/s1,00x
20,1839 s10,88 tok/s1,43x
40,1909 s20,95 tok/s1,49x
80,2781 s28,76 tok/s2,16x
160,3430 s46,64 tok/s2,67x
320,6302 s50,78 tok/s4,90x

Llegeix les dues columnes de la dreta l’una contra l’altra, perquè són tot el punt. Passar d’una petició a setze multiplica el throughput per 6,0 i multiplica l’espera de cada petició individual per 2,67. El batch va fer millor el servidor i pitjor cada usuari.

Això no és un error que es pugui ajustar fins que desaparegui; és el trade-off mateix, i té un nom a cada costat. La latència és el que viu una persona que espera una resposta. El throughput és allò pel qual es divideix la factura. Cap configuració millora totes dues coses.

Fixa’t també on s’atura. De 16 a 32, el throughput guanya un 9 % mentre la latència gairebé es duplica: el pas ha deixat d’estar memory-bound i ha passat a estar compute-bound, i més enllà d’aquell genoll el batch no compra res. Cada deployment té un genoll així; la seva ubicació s’ha de mesurar en el teu, però la seva existència no.

El batching estàtic malbarata la major part del que guanya

Enllaç a la secció: El batching estàtic malbarata la major part del que guanya

La manera ingènua de fer batch és recollir BB peticions, executar-les juntes i retornar quan totes hagin acabat. Però no acaben juntes: algunes respostes tenen vint tokens i algunes cinc-cents. Un batch fix s’executa fins que acaba el membre més llarg, i cada petició ja acabada continua ocupant la seva ranura, aportant padding, fins aleshores.

Agafa 64 peticions amb un biaix realista de longituds de sortida —mediana de 18 tokens, la més llarga 231, 1.874 en total— i simula les dues polítiques al cost per pas mesurat per a vuit ranures:

políticatemps de rellotgethroughputlatència mitjana per peticióslot-steps malbaratats
batches estàtics de 8176,9 s10,6 tok/s83,2 s3.214
continu, 8 ranures109,0 s17,2 tok/s8,1 s0

El throughput millora 1,6x. La latència mitjana millora més de deu vegades, perquè amb batching estàtic una petició que ha acabat en quatre passos encara espera un veí de 231 tokens abans que ningú en senti res.

Continuous batching3 és la solució, i és tan simple com sona: el batch no és un grup sinó un conjunt de ranures, i una ranura que s’allibera admet la següent petició en cua al pas immediatament següent. El planificador treballa a la granularitat d’un token en lloc d’una petició. Tota pila de serving en producció ja fa això.

Té una segona meitat, que és la caché. Les ranures que entren i surten deixen la memòria de caché fragmentada, i reservar per a cada ranura el seu context màxim possible malbarata la major part de la reserva. PagedAttention4 pren prestada la resposta dels sistemes operatius: emmagatzema la caché en blocs de mida fixa amb una taula de blocs per seqüència, de manera que la caché d’una seqüència pot estar dispersa físicament mentre continua sent lògicament contigua; això també permet que dues seqüències amb un prefix compartit comparteixin els blocs que el contenen. Això és la base de vLLM, i per això un motor de serving és un assignador de memòria amb un transformer enganxat.

L’altra meitat de la factura són els pesos mateixos. Mig milió llarg de paràmetres a quatre bytes cadascun són 1,98 GB; a dos bytes, 0,99 GB; a un byte, 0,49 GB. Menys bits per pes encongeixen el model al disc, l’encongeixen en memòria i —com que el decode és bandwidth-bound— fan que cada pas sigui més ràpid, perquè hi ha menys bytes per moure.

L’esquema més simple és la quantization simètrica de màxim absolut, i cap en tres línies:

quantize.pyPYTHON
qmax  = 2 ** (bits - 1) - 1
scale = W.abs().max() / qmax                        
Wq    = torch.round(W / scale).clamp(-qmax - 1, qmax)
W_hat = Wq * scale                                  # dequantized

Tria una escala perquè el pes més gran mapegi a l’enter més gran, divideix, arrodoneix, guarda els enters i l’escala. Reconstrueix multiplicant de tornada. No té res d’enginyós, i funciona —fins que deixa de funcionar.

Mesurat sobre els pesos reals del model: les 168 matrius de projecció, 357,8 milions de paràmetres, error relatiu WW^/W\lVert W - \hat{W}\rVert / \lVert W \rVert:

esquemaerror relatiu mitjàpitjor matriu
INT8, una escala per a tota la matriu0,04000,1487
INT8, una escala per fila de sortida0,01000,0149
INT4, una escala per a tota la matriu0,60260,9931
INT4, una escala per fila de sortida0,17900,2589
INT4, una escala per grup de 1280,13230,1992
NF4, una escala per bloc de 640,09520,1205
INT3, una escala per grup de 1280,30440,4123
INT2, una escala per grup de 1280,77900,8076

La quarta fila és l’esfondrament. Un error relatiu de 0,99 a la pitjor matriu vol dir que la reconstrucció no reté pràcticament res de l’original: la matriu ha estat substituïda per soroll d’una magnitud més o menys correcta. La causa és visible en el mateix experiment sobre una sola matriu:

TEXT
model.layers.12.mlp.down_proj.weight   (896 x 4864)
mean |w| 0.01386   std 0.01822   max |w| 0.43945   max/std 24.1
weights beyond 6 sigma: 692 of 4,358,144   (0.016 %)

Un pes de cada sis mil queda més enllà de sis desviacions estàndard, i el més gran és a 24. Amb una sola escala per a tota la matriu, aquest únic pes fixa la mida del pas per als 4,3 milions restants. A 8 bits hi ha 256 passos i el pes típic encara cau en un de significatiu. A 4 bits n’hi ha 16, l’extrem queda reservat per a un valor que gairebé res no té, i els pesos ordinaris —que són tots— s’arrodoneixen a dos o tres nivells diferents.

Tot el que ve després d’aquella fila és la mateixa reparació a granularitats diferents: donar a l’escala un territori més petit. Per fila de sortida divideix l’error per 3,4; per grup de 128 pesos consecutius el torna a dividir. El cost és comptabilitat —una escala de 16 bits per grup de 128 són 4+16/128=4.1254 + 16/128 = 4.125 bits per pes en lloc de 4— i recupera la major part de la distància.

NF4 ho aborda des de l’altre costat.5 Els nivells no han d’estar igualment espaiats. Els pesos dins d’un bloc es distribueixen aproximadament normalment, així que tria els setze nivells com els quantils d’una distribució normal: densos prop de zero, on realment són els pesos, escassos a les cues, on no ho són. Els mateixos quatre bits, la mateixa escala per bloc, en un bloc més petit —4,25 bits per pes contra els 4,125 del grup-128— i l’error mesurat baixa de 0,1323 a 0,0952, un 28 % menys. Una part és el bloc més fi i la resta és posar els nivells on hi ha la massa, i separar les dues coses necessitaria una tercera fila.

El requadre de punt flotant del Capítol 2 acabava amb una promesa: que aquest capítol quantitzaria pesos a 8 i 4 bits i trobaria un grapat de features outlier que es resistien a ser comprimides. Aquí són, i expliquen per què "simplement arrodoneix els nombres" no havia de funcionar mai sobre les activations.

Els pesos de dalt es comportaven malament. Les activations juguen en una altra lliga. Agafa un prompt ordinari de 84 tokens, captura el flux residual a cada capa i mesura la magnitud més gran que assoleix cadascuna de les 896 dimensions:

capa|h| més gran|h| més gran de la dimensió medianaràtiodimensions per sobre de 6x la mediana
16,190,33918x2
41543,481,550996x34
81571,631,4981049x36
121575,031,5461019x34
161579,601,617977x32
201577,982,361668x24
24204,4410,76019x12

La dimensió 62 arriba a 1.579,6 mentre que la dimensió mediana no supera mai 1,6. No és una casualitat d’un token o d’una capa: la mateixa dimensió hi és a la capa 4 i encara hi és a la capa 20, amb gairebé el mateix valor. Aquestes són les features outlier,6 i són sistemàtiques: una propietat del model entrenat, no de l’entrada.

L’histograma d’aquests 896 màxims per dimensió a la capa 16 fa que la forma sigui inconfusible:

TEXT
     0 -      1 | ######################################## 254
     1 -      2 | ######################################## 283
     2 -      4 | ######################################## 226
     4 -      8 | ######################################## 93
     8 -     16 | ##################                       18
    16 -     32 | #########                                9
    32 -     64 | #######                                  7
    64 -    128 | #####                                    5
   128 -    256 |                                          0
   256 -    512 |                                          0
   512 -   1024 |                                          0
  1024 -   4096 | #                                        1

Nou-centes dimensions en un pilot ordenat per sota de 8, absolutament res durant tres octaves, i després una sola dimensió a l’extrem llunyà. Ara quantitza aquest tensor a INT8 i compta què passa:

esquemaerror relatiunivells enters diferents usats, tensor complet
una escala per a tot el tensor0,108314 de 256
una escala per token (per fila)0,0433158
tensor complet, 1 dimensió outlier mantinguda en fp320,044248
tensor complet, 4 dimensions outlier mantingudes en fp320,027957
tensor complet, 16 dimensions outlier mantingudes en fp320,0085102

Catorze nivells de 256. L’escala la fixava 1.579,6, així que cada pas fa 12,44 d’ample, i l’activation típica —magnitud mediana 0,26, percentil noranta-nou 2,51— no té on aterrar. Per dimensió és encara més cru:

TEXT
single tensor-wide scale = 12.4378
  dim 826 (max |h| = 4.77):  1 distinct level out of 256
  dim 336 (max |h| = 1.62):  1 distinct level out of 256
  dim  96 (max |h| = 0.69):  1 distinct level out of 256

after excluding the top 4 dimensions, scale = 0.5749  (22x smaller)
  dim 826: 8 levels    dim 336: 4 levels    dim  96: 3 levels

Un nivell. Tota la dimensió, cada token, quantitzada al mateix nombre. S’hi van assignar vuit bits i se’n van fer servir aproximadament zero, i el model que llegeix aquestes activations rep una constant.

Aquesta mesura és la justificació de totes les tècniques que la gent realment fa servir:

Deixa’n fora els outliers. LLM.int8()6 descompon la multiplicació de matrius: les dimensions amb magnituds extremes es calculen en 16 bits, tota la resta en INT8, i les meitats se sumen. La taula de dalt n’és el rebut: treure quatre dimensions retalla l’error per un factor de gairebé quatre. SmoothQuant7, en canvi, migra la dificultat: divideix les activations per un factor per canal i multiplica la columna de pesos corresponent per ell, cosa que deixa el producte intacte i mou l’outlier fora del tensor que no el pot absorbir cap al que sí que pot.

Tria l’arrodoniment, no arrodoneixis sense més. Res del que hi ha a dalt pregunta per a què serveix la matriu. GPTQ8 quantitza columna per columna i, després de cadascuna, ajusta les columnes restants en precisió completa per compensar l’error ja comès: minimitza l’error de la sortida de la capa en entrades reals, no el dels seus pesos. AWQ9 observa que una petita fracció dels canals de pesos importa molt més que la resta, els troba a partir d’estadístiques d’activation i els escala cap amunt abans de quantitzar perquè aterrin en nivells més fins. Tots dos necessiten un conjunt de calibratge; cap no necessita gradients.

Mostra els detalls

GGUF, i què hi té a veure un format de fitxer amb tot això.

GGUF no és un mètode de quantization; és el contenidor que fa servir llama.cpp, i la confusió en comparacions de gguf vs gptq ve de tractar les dues coses com si fossin del mateix tipus. GGUF conté tensors, tokenizer, metadades d’arquitectura i plantilla de xat en un sol fitxer mapejable en memòria, i porta a dins una família d’esquemes de bloc: noms com Q4_K_M codifiquen bits per pes, mida de bloc i si alguns tensors es mantenen a més precisió.

La diferència d’enginyeria que importa: GPTQ i AWQ produeixen pesos optimitzats per a un kernel de GPU, mentre que els esquemes de GGUF es decodifiquen barat en una CPU amb el fitxer mapat en lloc de carregat. Per això el mateix "model 7B de 4 bits" nominal existeix en tots dos mons amb mides i qualitat diferents, i per això la comparació honesta no és mai el format: és la mesura de sota, executada en la teva pròpia tasca.

Gairebé tots els articles sobre quantization s’aturen a la secció anterior: expliquen el mètode, citen una ràtio de compressió i afirmen que la qualitat es conserva "en gran part". El Capítol 4 anava de no enganyar-se a un mateix, així que comprovem-ho.

Mateix model, pesos quantitzats in situ amb cada esquema, i després tres mesures: perplexity sobre 2.048 tokens de prosa anglesa reservada —aquí, l’esborrany d’aquest curs, per això el repositori substitueix un llibre fix de domini públic i imprimeix una taula de la mateixa forma amb nombres diferents—, una bateria de 16 preguntes factuals curtes amb respostes conegudes sota greedy decoding, i la fracció de tokens en què el model quantitzat coincideix amb el de precisió completa donat un context idèntic.

esquemaerror mitjà de pesperplexitybateria de preguntescoincideix amb fp32
fp32 (referència)0,000023,0813/16100,0 %
INT8 per tensor0,040023,5813/16
INT8 per fila0,010022,9613/1698,6 %
INT4 per tensor0,6026365.416.0000/16
INT4 per fila0,179046,186/1658,3 %
INT4 grup 1280,132331,0810/1671,5 %
NF4 bloc 640,095224,5511/1684,7 %
INT3 grup 1280,3044213,090/165,6 %
INT2 grup 1280,779026.325.4360/160,0 %

Hi ha quatre coses d’aquesta taula que val la pena dir clarament.

INT8 fet correctament surt gratis. L’INT8 per fila puntua 22,96 contra els 23,08 de la referència: una diferència d’una part entre dues-centes, que és soroll i s’ha de llegir com "idèntic". La direcció del soroll no és estable: al corpus de domini públic del repositori, els mateixos dos esquemes surten 22,24 contra 22,18, la meitat de la distància i apuntant cap a l’altra banda. Coincideix amb el model de precisió completa en 142 dels 144 tokens generats. Una quarta part de la memòria contra la referència fp32, la meitat contra l’fp16 que realment desplegaries, i cap cost detectable. INT8 fet descuidament també és gairebé gratis: una escala per matriu costa 0,5 punts de perplexity i cap resposta de la bateria. Vuit bits perdonen prou perquè la granularitat gairebé no importi, que és exactament per això que la gent generalitza d’INT8 a INT4 i pren mal.

INT4 amb una escala per tensor destrueix el model. Perplexity 365 milions: no degradat, aniquilat. La granularitat és aleshores tot el joc: per tensor 365.416.000, per fila 46,18, per grup de 128 31,08, NF4 24,55. Els mateixos quatre bits per pes, un factor de quinze milions entre el pitjor i el millor.

La perplexity és un instrument bast, i la bateria encara més. Entre NF4 i INT4 grup-128 la diferència de perplexity és de 6,5 punts i la bateria difereix en una pregunta; i l’interval de confiança del Capítol 4 diu que una pregunta de setze no distingeix absolutament res. Hi ha una demostració més esmolada que l’interval: executa la mateixa bateria amb la penalització de repetició de sèrie del model desactivada, que és el que realment vol dir greedy decoding, i aquestes dues files intercanvien posicions. Una pregunta de setze no és un efecte petit: no és cap efecte. L’advertiment del Capítol 8 també s’aplica: la perplexity només és comparable entre models que comparteixen tokenizer, així que un nombre d’un article d’algú altre no es pot comparar amb el teu.

La columna de coincidència és la més precisa de les tres, i gairebé gratuïta: executa el model de precisió completa amb greedy, i després pregunta al quantitzat, a cada posició, què hauria triat donat el mateix prefix. Té 144 observacions independents en lloc de 16, no necessita ground truth i es degrada suaument allà on la bateria es degrada a salts. També és exactament la quantitat que necessita la secció següent.

Aquesta és la promesa que el Capítol 1 va fer sobre aquest capítol, arribant puntual: les matemàtiques diuen que un model de 4 bits és possible, i l’enginyeria decideix si és usable.

El Capítol 12 ho va anunciar i va deixar la factura aquí.

La idea ve directament de la separació prefill/decode. Verificar una seqüència proposada de γ\gamma tokens costa una forward pass sobre γ\gamma posicions: un producte matriu-matriu, amb prou feines més car que la passada sobre una. Així:

Un model petit i barat genera γ\gamma tokens candidats autoregressivament.

El model gran executa una forward pass sobre tots els γ\gamma candidats alhora, produint què hauria dit a cada posició.

Conserva el prefix més llarg en què tots dos coincideixen, més el token que el model gran proporciona gratis al primer desacord. Descarta la resta i torna a començar.

La distribució de sortida no canvia. Amb greedy decoding això és obvi: un token només s’accepta si el target l’hauria produït. Amb sampling cal una regla d’acceptació modificada, i Leviathan et al. demostren que la distribució resultant és exactament la del target.10 Aquesta és la segona optimització exacta d’aquest capítol.

Per tant, tot depèn de la taxa d’acceptació α\alpha, que és mesurable: és la columna de coincidència de dalt, que és per això que s’hi va calcular. Fent servir cada model quantitzat com a draft per al target de precisió completa, sobre 144 posicions generades:

model draftacceptaciótirada acceptada més llargatokens esperats per passada del target, γ=4\gamma = 4
fp32 (el target mateix)100,0 %485,00
INT8 per fila98,6 %484,86
NF4 bloc 6484,7 %203,69
INT4 grup 12871,5 %132,85
INT4 per fila58,3 %72,24
INT3 grup 1285,6 %21,06
INT2 grup 1280,0 %01,00

Els tokens acceptats esperats per passada de verificació, amb longitud de draft γ\gamma, són

E[tokens]=1αγ+11α\mathbb{E}[\text{tokens}] = \frac{1 - \alpha^{\gamma+1}}{1 - \alpha}

i l’acceleració neta ho divideix pel cost propi del draft, una fracció cc del target per token:

acceptacióc=0.05c=0.05, γ=4\gamma=4c=0.1c=0.1, γ=4\gamma=4c=0.2c=0.2, γ=4\gamma=4c=0.1c=0.1, γ=8\gamma=8
30 %1,19x1,02x0,79x0,79x
50 %1,61x1,38x1,08x1,11x
70 %2,31x1,98x1,54x1,78x
90 %3,41x2,93x2,28x3,40x

L’entrada en negreta és la que cal recordar: speculative decoding pot fer que la generació sigui més lenta. Amb un 30 % d’acceptació i un draft que costa una cinquena part del target, pagues cinc forward passes i et quedes 1,4 tokens. L’última columna és l’altra trampa: un draft més llarg només ajuda quan l’acceptació és alta, perquè la cua d’una conjectura de γ\gamma tokens gairebé mai no s’arriba a tocar. Amb un 90 % d’acceptació, γ=8\gamma = 8 val 3,40x, i amb un 30 % val 0,79x: la mateixa configuració, una victòria o una pèrdua segons un nombre mesurat en el teu trànsit.

La quantization encongeix un model emmagatzemant la mateixa funció en menys bits. La distillation l’encongeix entrenant un model més petit perquè imiti un de més gran11: una idea anterior al deep learning per gairebé una dècada.12

La part subtil és de què aprèn l’estudiant. No de la resposta correcta: s’hi hauria pogut entrenar directament. El que aporta el professor és la distribució completa. Pregunta al model què segueix una frase i mira més enllà de l’argmax:

TEXT
"She poured the milk into the"
  ' jug' 0.1355   ' cup' 0.1051   ' bowl' 0.0605   ' large' 0.0380   ' milk' 0.0360

L’etiqueta dura diu jug i res més. L’etiqueta suau diu jug, i també que cup era gairebé igual de bo, bowl plausible, i large —un adjectiu, una continuació gramatical completament diferent— encara viu. Aquest és l’argument original: això és un 7, però s’assembla força a un 1, i la semblança és informació que l’etiqueta dura llença.

També és per això que la distillation fa servir una temperatura. Dividir els logits per TT abans del softmax aplana la distribució i augmenta el pes relatiu dels finalistes: en aquesta frase, la ràtio entre el token superior i el tercer baixa de 2,24 a T=1T = 1 a 1,50 a T=2T = 2: l’arrel quadrada de la primera, que és el que fa dividir els logits per dos a una ràtio. Mateix ordre, més attention de la pèrdua sobre els gairebé encerts. El gradient de l’estudiant porta la incertesa del professor, no només el seu veredicte.

Tot aquest capítol ara és una suma:

memory=N×bytes per weightfixed+T×2LHkvdhead×bytesgrows with every token+runtime overheadcall it 1.5 GB\text{memory} = \underbrace{N \times \text{bytes per weight}}_{\text{fixed}} + \underbrace{T \times 2 L H_{kv} d_{\text{head}} \times \text{bytes}}_{\text{grows with every token}} + \underbrace{\text{runtime overhead}}_{\text{call it 1.5 GB}}

on TT són els tokens totals residents entre totes les peticions concurrents. Aplicant-ho: les files 7B i 70B assumeixen 8 caps key-value de dimensió 128; la fila 13B, attention multi-head completa amb 40 caps, que és com es van construir aquelles generacions de models —i es nota.

8 GB

modelprecisiópesoslliure després de l’overheadtokens de context que hi caben
7Bfp1613,0 GBno hi cap
7Bint86,5 GBno hi cap
7Bint4 (g128)3,4 GB3,1 GB25.710
13Bint4 (g128)6,2 GB0,3 GB337
70Bint4 (g128)33,6 GBno hi cap

16 GB

modelprecisiópesoslliure després de l’overheadtokens de context que hi caben
7Bfp1613,0 GB1,5 GB11.972
7Bint86,5 GB8,0 GB65.378
7Bint4 (g128)3,4 GB11,1 GB91.246
13Bint812,1 GB2,4 GB3.136
13Bint4 (g128)6,2 GB8,3 GB10.822

24 GB

modelprecisiópesoslliure després de l’overheadtokens de context que hi caben
7Bfp1613,0 GB9,5 GB77.508
7Bint86,5 GB16,0 GB130.914
7Bint4 (g128)3,4 GB19,1 GB156.782
13Bint812,1 GB10,4 GB13.622
13Bint4 (g128)6,2 GB16,3 GB21.308
70Bint4 (g128)33,6 GBno hi cap

Mira la fila 13B de la taula de 8 GB. Els pesos hi caben —6,2 GB de 8—, així que, segons la manera habitual de parlar, un model 13B "funciona en una targeta de 8 GB". Té 337 tokens de context, que no és una conversa sinó amb prou feines un prompt. "Hi cap" és la pregunta equivocada. La correcta és "amb quant context, i per a quants usuaris alhora".

Mira també les dues files int8 de 16 GB. El 7B obté 65.378 tokens i el 13B en té 3.136: una diferència de vint vegades a partir de 5,6 GB de pesos extra, perquè aquest 13B té attention multi-head i la seva caché costa 800 KB per token contra els 128 KB del 7B. Dos models de mida semblant, un inutilitzable per a context llarg, per una raó que no apareix en cap titular de model card.

Fa tretze capítols això era un perceptró amb dos pesos i un bias. Ara és un transformer que ha estat dissenyat, entrenat, aligned, ensenyat a gastar compute en preguntes difícils i servit a un cost mesurat per token, sense cap capsa interna que hagi quedat tancada.

Això acaba aquí, i acaba expressament.

El Capítol 14 comença amb el model en un altre lloc. No en el teu procés, no a la teva memòria, no en una variable que puguis imprimir: en una màquina que no administres, darrere d’una API key, un port i una factura. Tot el que s’ha mesurat aquí continua passant —el prefill encara s’executa abans del primer token, la caché encara creix amb la conversa, el batch on ets encara pertany a algú altre i encara decideix la teva latència—, però a partir d’ara ho observes a través d’un stream de Server-Sent Events, un finish_reason i un HTTP 429 amb una capçalera Retry-After. Les preguntes canvien amb el punt de vista: no com es calcula aquest gradient sinó per què s’ha triplicat la meva factura. També ho fa el llenguatge, i el Capítol 14 explica aquesta regla en lloc d’anunciar-la: fins aquí el codi contenia pesos, gradients, logits i bytes de tokenizer; a partir d’allà conté una connexió, un reintent, una cancel·lació i estat acumulat. Els tretze capítols que deixes enrere no es descarten en creuar. Són la descripció del que s’executa a l’altre costat del port.


Dues omissions són deliberades. FlashAttention (Dao et al., arXiv:2205.14135) no és una attention diferent: calcula la mateixa funció dividint l’operació en tiles perquè la matriu de puntuacions n×nn \times n no s’escrigui mai a memòria, que és per això que els 67 MB de la segona taula d’aquest capítol són més petits a la pràctica del que suggereix l’aritmètica. I els kernels mateixos es deleguen: la classe 10 del CS336 de Stanford cobreix els sistemes d’inference amb una profunditat que això no intenta, i el repositori llama.cpp i l’especificació GGUF són les fonts principals per al costat CPU.

  1. Shazeer, N. Fast Transformer Decoding: One Write-Head is All You Need. arXiv:1911.02150 (2019). L’article és en gran part un argument de bandwidth de memòria, i es llegeix com a tal.

  2. Ainslie, J. et al. GQA: Training Generalized Multi-Query Transformer Models from Multi-Head Checkpoints. arXiv:2305.13245 (2023). Inclou la recepta d’uptraining que converteix un checkpoint multi-head existent, que és per això que GQA es va estendre tan de pressa.

  3. Yu, G.-I., Jeong, J. S., Kim, G.-W., Kim, S. and Chun, B.-G. Orca: A Distributed Serving System for Transformer-Based Generative Models. OSDI 2022. Introdueix la planificació a nivell d’iteració —continuous batching— i el batching selectiu.

  4. Kwon, W. et al. Efficient Memory Management for Large Language Model Serving with PagedAttention. arXiv:2309.06180 (2023), SOSP 2023. L’article sobre el qual es construeix vLLM; la §3 és l’analogia amb els sistemes operatius completa.

  5. Dettmers, T., Pagnoni, A., Holtzman, A. and Zettlemoyer, L. QLoRA: Efficient Finetuning of Quantized LLMs. arXiv:2305.14314 (2023). NF4 es defineix a la §3; els setze valors de nivell usats en la mesura de dalt són els que deriva aquest article.

  6. Dettmers, T., Lewis, M., Belkada, Y. and Zettlemoyer, L. LLM.int8(): 8-bit Matrix Multiplication for Transformers at Scale. arXiv:2208.07339 (2022). L’anàlisi de features outlier de la §4 és la font del fenomen mesurat més amunt, incloent-hi la troballa que els outliers emergeixen sistemàticament a escala. 2

  7. Xiao, G., Lin, J., Seznec, M., Wu, H., Demouth, J. and Han, S. SmoothQuant: Accurate and Efficient Post-Training Quantization for Large Language Models. arXiv:2211.10438 (2022).

  8. Frantar, E., Ashkboos, S., Hoefler, T. and Alistarh, D. GPTQ: Accurate Post-Training Quantization for Generative Pre-trained Transformers. arXiv:2210.17323 (2022).

  9. Lin, J. et al. AWQ: Activation-aware Weight Quantization for LLM Compression and Acceleration. arXiv:2306.00978 (2023).

  10. Leviathan, Y., Kalman, M. and Matias, Y. Fast Inference from Transformers via Speculative Decoding. arXiv:2211.17192 (2022). El teorema 1 és la prova que la distribució de sortida no canvia; Chen et al. (arXiv:2302.01318) van publicar la mateixa idea independentment.

  11. Hinton, G., Vinyals, O. and Dean, J. Distilling the Knowledge in a Neural Network. arXiv:1503.02531 (2015). La temperatura i l’argument del "dark knowledge".

  12. Buciluă, C., Caruana, R. and Niculescu-Mizil, A. Model Compression. KDD 2006. Distillation, nou anys abans, per a ensembles més que no pas transformers.

A punt per deixar que triï LIA?

Crea amb tots els models d'IA en un sol lloc — comença gratis avui mateix.