コンテンツへスキップ
13/30第13章 / 全30章

推論コストを下げる:KV Cache、Batching、量子化

同じmodelが同じ問いに8.8秒と78.9秒で回答し、出力はbyte単位で同一。INT4も断言でなく3通りに測定します。

このページの内容

同じmodelが、同じmachine上で、同じ質問に同じ48 tokensで答えます。2つの出力はtoken単位で完全に同一です。仮定ではなく、確認済みです。

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)

変えた引数は1つだけです:use_cache=False。model、prompt、sampling、算術は何も違いません。そして2回目の実行は、その手間の分だけ正確になっているわけでもありません。何のためでもなく9倍遅いだけです。

この章の形はそれです。ここに出てくるものすべて — cache、batch、量子化されたweights — は、答えを変えない作業に支払い続けるのをやめる試み、あるいは安い答えに何が必要かを調べる試みです。第10章ではtrainingの価格表を作りました。これは、永遠に支払い続ける側の価格表です。deployされたmodelは、出力するtokenごと、requestごと、その寿命が尽きるまで、およそ2N2N FLOPsを消費します。

tokenを生成するために、decoder-only transformerはこれまでのsequence全体を取り、すべてのlayerに通し、最後のpositionから確率分布を読み取ります。それから選ばれたtokenを追加し、また同じことをします。この説明は正しく、遅い実行がやっていることでもあります。

同時に、それは途方もなく無駄です。その理由は第9章のcausal maskにあります。position 7のkey vectorとvalue vectorは、position 7のinputとそれ以前のpositionから計算されます。position 8が到着しても、position 7はそれを見ることができません。causalとはそういう意味です。したがってposition 7のkeyとvalueは、前とまったく同じ数値です。遅い実行は、それでも各stepでそれらを再計算します。

だから保存します。その保存領域がkey-value cacheであり、language model servingにおける最も重要なoptimizationです。

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)

loop内でmodelに渡されているものを見てください:nxt、つまり1 tokenです。sequenceではありません。新しいtokenのqueryはcached keyすべてにattendし、cached keyはそもそも変わるはずがありません。これは近似ではありません。上の同一出力チェックがその要点です。cacheは品質を速度と交換しません。冗長な算術を削除します。

scalingをきれいに見るために、transformerを取り払い、d=64d = 64のsingle attention headを、生成1 stepについて両方の方法で測ります。

context内のtokensすべて再計算cacheあり比率score matrix
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

右端の列が原因です。再計算は毎stepで完全なn×nn \times n attention matrixを作ります。第9章の漸近記法のboxに出てきたO(n2)O(n^2)を、tokenごとに1回支払うわけです。cacheを使うと、代わりに1×n1 \times nのrowを作るだけです。4,096 tokensでは、scoresが67 MB対16 KBです。

ミリ秒ではなくmultiply-accumulateを数えると、machineを議論から取り除けます。cold startからTT tokensを生成する場合:

生成tokenscacheあり再計算比率
1282.6 M192.0 M73x
51223.1 M7.36 G318x
2048293.7 M392.6 G1,336x

stepごとには、cached版はcontextに対してlinearで、uncached版はquadraticです。generation全体で合計するとO(T2)O(T^2)O(T3)O(T^3)になり、その比率は限りなく大きくなります。冒頭の9倍差は48 tokensで測ったものです。この表の最初のrowより短い長さです。

cacheは、memoryに置くべきものも変えます。8 GBのlaptop GPUでfp16、256 tokensを生成し、allocatorのpeakから常駐weightsを差し引くと:

peak working memory
cacheあり21.8 MB
再計算181.7 MB

memoryは8.3倍多く、しかも同じtokensをより遅く生成するために使われます。これは第5章での約束が、思いがけない方向から到着したものです。そこではreverse-mode autodiffがbackward passのためにすべての中間値を生かしておく必要があり、activationsがtraining memoryを支配しました。inferenceにはbackward passがなく、そのために保持すべきものもありません。代わりにmemoryを支配するのがcacheであり、それは避けられないcostではなく、意図的な選択です。

fast runをもう一度見てください。最初のtokenは、残り47 tokensとは違う振る舞いをしていました。

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

promptはtokenあたり25.6 msで、生成tokenは1つあたり166 msでした。同じmodel、同じhardware、同じweightsなのに、tokenあたり6倍の差です。そして多くの人の予想とは逆方向です。promptが安い部分です。generationは、本当に異なる物理を持つ2つのphaseに分かれます。

prompt全体に対する1回のforward passです。すべてのtokenが並列に処理されるため、各weight matrixはmemoryから1回だけ読み込まれ、何百ものtoken vectorsからなるmatrixと掛け合わされます。matrix-matrix productです。移動するbyteあたりの算術量が多く、GPUが得意とする形です。Prefillはcompute-boundで、そのcostはprompt長におおよそlinearです。

tokenごとに1回のforward pass、batchは1、sequenceも1です。各weight matrixは依然として丸ごとmemoryから読み込まれ、single vectorと掛け合わされます。matrix-vector productで、移動するbyteあたりの算術量はほとんどありません。Decodeはmemory-bandwidth-boundで、tokenあたりのcostはcontext長にほとんど依存しません。

両方のhalfは測定できます。Prefill、PP tokensに対する1 pass:

prompt tokenstokenあたりms
160.351521.97
320.525416.42
641.049116.39
1281.655212.93
2563.096512.10

Decode、CCのcacheに対して1 token:

cached tokens1 tokenのms
16110.05
6497.57
256108.53
1024103.86

2つ目の表は2回読んでください。contextが16 tokensから1,024へ、attendすべき履歴が64倍になっても、1 stepのcostは測定できるほど変わりませんでした。cacheに対するattentionは実際の作業ですが、1つのvectorを作るために5億個のweightsをmemory busに流し込む固定costに比べれば小さいのです。その固定costこそが、次のsectionのすべての理由です。

この2つのphaseが、serving systemが必ず報告する2つの数値の起源です。Time to first tokenは本質的にprefillであり、promptとともに増えます。だから長い会話は始まりが遅く感じます。Tokens per second1/decode step1/\text{decode step}であり、ほぼ一定です。だからその後のreplyは均等に流れます。始まりが遅く、その後なめらかにstreamするchatは、renderingの工夫ではありません。この2つの表そのものです。

cacheは算術をmemoryと交換します。そして必要とするmemoryは小さくありません。context内の各tokenについて、各layerはkey-value headごとに1つのkey vectorと1つのvalue vectorを保持します。

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}

2はkeysとvaluesの分です。それ以外はarchitectureです。この章全体で測っているmodel — 24 layers、14 query heads、2 key-value heads、head dimension 64 — では、fp16でtokenあたり2×24×2×64×2=12,2882 \times 24 \times 2 \times 64 \times 2 = 12{,}288 bytesです。

この分野のformulaは2倍ずれる癖があるので、信じるのではなくallocatorに照らして確認します。

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

完全に一致し、試したすべてのshapeで一致し続けます。

batchcontextmeasured cachepredictedpeak working memory
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

最後の3 rowsはもう一度見る価値があります。32 usersで各2,048 tokens、64 usersで各1,024、128 usersで各512 — いずれの場合もcacheは768 MBです。3つとも65,536 tokensを保持しているからです。cacheはresidentなtoken総数だけに依存し、それがusersにどう分配されているかには依存しません。 この事実がbatching sectionの基礎です。

第9章ではmulti-query attentionとgrouped-query attentionを紹介し、その理由をこの章に先送りしました。理由はこのformula、とりわけその中のHkvH_{kv}です。

標準的なmulti-head attentionでは、各query headが自分専用のkey headとvalue headを持ちます。ここでのmodelには14 query headsがあります。full multi-head attentionならcacheはtokenあたり2×24×14×64×2=86,0162 \times 24 \times 14 \times 64 \times 2 = 86{,}016 bytes、つまり12 KBではなく84 KBになります。ちょうど7倍で、query headsとkey-value headsの比率そのものです。

Multi-query attention1はこれを極限まで進め、すべてのquery headsが1つのkey-value headを共有します。Grouped-query attention2は勝ち残った妥協案です。少数のkey-value headsを持ち、それぞれをquery headsのgroupで共有します。MQAの品質低下は実在し、GQAではそうではなかったからです。どちらも算術を節約しません。このformulaを整数で割るために存在し、long contextsによってcacheがbinding constraintになった瞬間、業界全体に広がりました。

そしてそれはすぐにそうなります。32 layers、dimension 128のkey-value headsを8つ持つ7B-class modelでは、fp16のcacheはtokenあたり128 KBです。

context tokens1 user8 users64 users
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

そのmodel自身のweightsはfp16で13.0 GBです。この章の最後の表にある値です。したがって128,000-token contextでは、1 userのcacheがmodelより大きくなります。この算術を第16章ではお金に変換します。そして長い会話が単に遅いだけではない理由でもあります。requestが生きている間、その会話はmachineの固定された一部を占有します。

Decodeはmemory-boundです。weightsは1 tokenを作るためにbusを通って引きずられ、算術unitは遊んでいます。ならば同じstepにもっと仕事を入れます。複数requestを同時に走らせれば、1回読まれたweightsが全員に使われます。同じmodelで測ると、各requestが64-token cacheを持ち、1 tokenをdecodeする場合:

batchstepあたりlatencythroughputlatency 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

右2列を突き合わせて読んでください。そこがすべてです。request数を1から16に増やすと、throughputは6.0倍になりますが、個々のrequestの待ち時間は2.67倍になります。batchはserverを良くし、すべてのuserを悪くしました。

これはtuningで消せるbugではありません。それ自体がtrade-offであり、両側に名前があります。Latencyはreplyを待つ人が経験するものです。Throughputは請求書を割る相手です。どちらも改善するsettingはありません。

止まる場所にも注目してください。16から32ではthroughputの増加は9 %なのにlatencyはほぼ2倍になります。stepはmemory-boundではなくcompute-boundになっており、その膝を越えるとbatchは何も買いません。すべてのdeploymentにこの膝があります。場所はあなたの環境で測る必要がありますが、その存在は測るまでもありません。

単純なbatchingは、BB requestsを集め、一緒に走らせ、全員が終わったら返すというものです。しかし全員が同時に終わるわけではありません。replyが20 tokensのものもあれば500 tokensのものもあります。固定batchは最長のmemberが終わるまで走り、終了済みrequestもそれまでslotを占有し続け、paddingに寄与します。

output lengthに現実的な偏りがある64 requests — median 18 tokens、最長231、合計1,874 — を取り、8 slotsの測定済みper-step costで両policyをsimulateします。

policywall clockthroughputrequestあたり平均latencywasted slot-steps
static batches of 8176.9 s10.6 tok/s83.2 s3,214
continuous, 8 slots109.0 s17.2 tok/s8.1 s0

throughputは1.6x改善します。平均latencyは10倍以上改善します。static batchingでは、4 stepsで終わったrequestでも、231-tokenの隣人が終わるまで誰にも届かないからです。

Continuous batching3が修正策で、聞こえるとおり単純です。batchはgroupではなくslotsの集合であり、空いたslotは次のstepですぐにqueued requestを受け入れます。schedulerはrequest単位ではなく1 token単位で動きます。本番のserving stackは今ではすべてこれを行います。

これには第2のhalfがあり、それがcacheです。出入りするslotsはcache memoryを断片化し、各slotに最大可能contextを予約すると、その予約の大半が無駄になります。PagedAttention4はoperating systemsから答えを借ります。cacheを固定サイズblocksに保存し、sequenceごとにblock tableを持つことで、sequenceのcacheは物理的には散らばっていても論理的には連続に見えます。これにより、prefixを共有する2つのsequencesがそのprefixを保持するblocksを共有することもできます。vLLMがその上に作られている理由であり、serving engineがtransformer付きのmemory allocatorである理由です。

請求書のもう半分はweightsそのものです。5億parametersは4 bytesずつなら1.98 GB、2 bytesなら0.99 GB、1 byteなら0.49 GBです。weightあたりのbitsを減らすと、modelはdisk上でもmemory上でも小さくなり、decodeはbandwidth-boundなので各stepも速くなります。動かすbytesが少ないからです。

最も単純な方式はsymmetric absolute-maximum quantizationで、3行に収まります。

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

最大weightが最大integerに写るようscaleを選び、割り、roundし、integersとscaleを保存します。復元は掛け戻すだけです。賢いことは何もしていません。そして動きます。動かなくなるところまでは。

modelの実weightsで測りました。168個すべてのprojection matrices、3億5,780万parameters、relative error WW^/W\lVert W - \hat{W}\rVert / \lVert W \rVert

schememean relative errorworst matrix
INT8, matrix全体に1 scale0.04000.1487
INT8, output rowごとに1 scale0.01000.0149
INT4, matrix全体に1 scale0.60260.9931
INT4, output rowごとに1 scale0.17900.2589
INT4, 128ごとのgroupに1 scale0.13230.1992
NF4, 64ごとのblockに1 scale0.09520.1205
INT3, 128ごとのgroupに1 scale0.30440.4123
INT2, 128ごとのgroupに1 scale0.77900.8076

4行目が崩壊です。worst matrixのrelative errorが0.99ということは、復元はoriginalをほとんど何も保持していないという意味です。そのmatrixは、だいたい正しい大きさのnoiseに置き換えられています。原因は、single matrixでの同じ実験に見えます。

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 %)

6,000個に1つのweightが6 standard deviationsを超え、最大は24 outです。matrix全体で1 scaleを使うと、その1つのweightが430万個すべてのstep sizeを決めます。8 bitsなら256 stepsあり、典型的なweightも意味のあるlevelに落ちます。4 bitsでは16しかなく、外側のlevelはほとんど誰も持たない値のために予約され、普通のweights — つまりほぼ全員 — は2つか3つのlevelsに丸められます。

そのrow以降は、粒度を変えた同じ修理です。scaleの担当範囲を小さくします。output rowごとにするとerrorは3.4分の1になり、連続128 weightsごとのgroupにするとさらに下がります。costはbookkeepingです。128ごとのgroupに16-bit scaleを1つ置くと、weightあたり4 bitsではなく4+16/128=4.1254 + 16/128 = 4.125 bitsになります。そして失った差の大半を取り戻します。

NF4は反対側から攻めます。5 levelsは等間隔である必要がありません。block内のweightsはおおよそnormal distributionなので、16 levelsをnormal distributionのquantilesとして選びます。weightsが実際にいるzero付近は密に、いないtailsは疎にします。同じ4 bits、同じblock scalingで、blockはより小さく — group-128の4.125に対してweightあたり4.25 bits — 測定errorは0.1323から0.0952へ、28 %下がります。その一部はより細かいblockによるもので、残りはmassのある場所にlevelsを置いたことによります。2つを分けるには3つ目のrowが必要です。

第2章のfloating-point boxは約束で終わりました。この章でweightsを8 bitsと4 bitsに量子化し、圧縮を拒む少数のoutlier featuresを見つける、という約束です。それがこれです。そして「数値を丸めるだけ」がactivationsに対してうまくいくはずがなかった理由も説明します。

上のweightsは行儀が悪いものでした。activationsは別格です。普通の84-token promptを取り、各layerのresidual streamをcaptureし、896 dimensionsそれぞれが到達する最大magnitudeを測ります。

layerlargest |h|median dimension's largest |h|ratiodimensions above 6x the median
16.190.33918x2
41543.481.550996x34
81571.631.4981049x36
121575.031.5461019x34
161579.601.617977x32
201577.982.361668x24
24204.4410.76019x12

dimension 62は1,579.6に達しますが、median dimensionは1.6を超えません。これは1 tokenや1 layerの偶然ではありません。同じdimensionがlayer 4にもあり、layer 20にもほぼ同じ値で残っています。これがoutlier features6であり、systematicです。inputではなくtrained modelの性質です。

layer 16での896個のper-dimension maximaのhistogramを見ると、その形は明白です。

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

900 dimensionsが8未満にきれいに積まれ、3 octaves分は何もなく、最後の端に1 dimensionだけがいます。ではそのtensorをINT8に量子化し、何が起きるか数えます。

schemerelative errordistinct integer levels used, whole tensor
tensor全体に1 scale0.108314 of 256
tokenごと(rowごと)に1 scale0.0433158
tensor全体、1 outlier dimensionをfp32で保持0.044248
tensor全体、4 outlier dimensionsをfp32で保持0.027957
tensor全体、16 outlier dimensionsをfp32で保持0.0085102

256 levels中14 levels。scaleは1,579.6で決まったため、各stepは幅12.44です。典型的なactivation — median magnitude 0.26、99th percentile 2.51 — には着地先がありません。dimensionごとに見るとさらに強烈です。

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

1 level。 dimension全体、すべてのtokenが同じ数に量子化されました。8 bitsを割り当てたのに、使われたのはほぼ0 bitsです。そしてmodelはそれらのactivationsを読み、constantを渡されます。

この測定が、実際に使われるあらゆるtechniqueの根拠です。

outliersを外へ出す。 LLM.int8()6はmatrix multiplyを分解します。極端なmagnitudeを持つdimensionsは16 bitsで計算し、それ以外はINT8で計算し、2つを足します。上の表が領収書です。4 dimensionsを取り除くとerrorはほぼ4分の1になります。SmoothQuant7は代わりに困難を移します。activationsをper-channel factorで割り、対応するweight columnにそれを掛けます。productは不変のまま、outlierを、それを吸収できないtensorから、吸収できるtensorへ移します。

roundingを選ぶ。ただ丸めない。 ここまでの話は、そのmatrixが何のためのものかを問うていません。GPTQ8はcolumnごとに量子化し、そのたびに残りのfull-precision columnsを調整して、すでに発生したerrorを補償します。weightsのerrorではなく、実inputに対するlayerのoutputのerrorを最小化します。AWQ9は、weight channelsのごく一部が残りよりはるかに重要だと見て、activation statisticsからそれらを見つけ、量子化前にscale upして細かいlevelsに乗るようにします。どちらもcalibration setを必要としますが、gradientsは不要です。

詳細を表示

GGUF、そしてfile formatがこれと何の関係があるのか。

GGUFはquantization methodではありません。llama.cppが使うcontainerです。gguf vs gptq比較で混乱が起きるのは、この2つを同じ種類のものとして扱うからです。GGUFはtensors、tokenizer、architecture metadata、chat templateを1つのmemory-mappable fileに保持し、その中にblock schemesのfamilyを持ちます。Q4_K_Mのような名前は、weightあたりbits、block size、一部tensorsを高precisionのまま保つかどうかをencodeしています。

重要なengineering上の違いはこうです。GPTQとAWQはGPU kernel向けに最適化されたweightsを生成します。一方GGUFのschemesは、fileを読み込むのではなくmapしたCPU上で安くdecodeされます。だから同じ名目の「4-bit 7B model」が両方の世界に、異なるsizeと異なるqualityで存在します。そして誠実な比較対象はformatではありません。下の測定を、あなた自身のtaskで走らせた結果です。

quantizationについての記事のほとんどは前sectionで止まります。methodを説明し、compression ratioを引用し、品質は「ほぼ保たれる」と主張します。第4章は自分をだまさないための章だったので、実際に確かめます。

同じmodelのweightsを各schemeでin-placeに量子化し、3つを測ります。held-out English prose 2,048 tokensでのperplexity — ここではこのcourseのdraftです。そのためrepositoryでは固定のpublic-domain bookに差し替え、同じ形で異なる数値の表を出します — greedy decodingで既知の答えを持つ16の短い事実質問のbattery、そして同一contextを与えたときに量子化modelがfull-precision modelと一致するtokenの割合です。

schememean weight errorperplexityquestion batteryfp32との一致
fp32 (reference)0.000023.0813/16100.0 %
INT8 per tensor0.040023.5813/16
INT8 per row0.010022.9613/1698.6 %
INT4 per tensor0.6026365,416,0000/16
INT4 per row0.179046.186/1658.3 %
INT4 group 1280.132331.0810/1671.5 %
NF4 block 640.095224.5511/1684.7 %
INT3 group 1280.3044213.090/165.6 %
INT2 group 1280.779026,325,4360/160.0 %

この表で率直に述べるべきことが4つあります。

INT8は正しく行えば無料です。 per-row INT8はreferenceの23.08に対して22.96です。差は200分の1で、noiseであり「同一」と読むべきです。noiseがどちらを向くかは安定しません。repositoryのpublic-domain corpusでは同じ2 schemesが22.24対22.18になります。距離は半分で、向きは逆です。生成された144 tokensのうち142でfull-precision modelと一致します。fp32 referenceに対してmemoryは4分の1、実際にdeployするであろうfp16に対しては半分で、検出可能なcostはありません。INT8は雑にやってもほぼ無料です。matrixごとに1 scaleでもperplexityは0.5 points、battery answersは失いません。8 bitsは十分に寛容なので、granularityはほとんど問題になりません。だからこそ人はINT8からINT4へ一般化して痛い目を見ます。

tensorごとに1 scaleのINT4はmodelを破壊します。 Perplexity 3億6,500万。劣化ではなく壊滅です。その後はgranularityがすべてです。per-tensor 365,416,000、per-row 46.18、per-group-of-128 31.08、NF4 24.55。同じweightあたり4 bitsで、最悪と最良の間に1,500万倍の差があります。

Perplexityは粗いinstrumentで、batteryはさらに粗いものです。 NF4とgroup-128 INT4の間でperplexity差は6.5 points、battery差は1 questionです。そして第4章のconfidence intervalによれば、16問中1問の差は何も区別しません。intervalより鋭い実演もあります。model標準のrepetition penaltyをoffにして同じbatteryを走らせると、それがgreedy decodingの本来の意味ですが、この2 rowsは入れ替わります。16問中1問は小さなeffectではなく、effectなしです。第8章の警告も当てはまります。perplexityはtokenizerを共有するmodels間でしか比較できないため、他人の記事の数値を自分のものと比較することはできません。

agreement列は3つの中で最も鋭く、 ほぼ無料です。full-precision modelをgreedyに走らせ、各positionで同じprefixを与えたときにquantized modelなら何を選ぶかを尋ねます。16ではなく144の独立観測があり、ground truthを必要とせず、batteryが段差状に劣化する場所でもなめらかに劣化します。そしてそれは次sectionが必要とする量そのものでもあります。

これは第1章がこの章について約束したことの予定どおりの到着です。数学は4-bit modelが可能だと言い、engineeringはそれが使えるかどうかを決めます。

第12章でこれを予告し、請求書をここに残していました。

発想はprefill/decode分割からそのまま出てきます。γ\gamma tokensの提案されたsequenceを検証するcostは、γ\gamma positionsに対する1 forward passです。matrix-matrix productなので、1つに対するpassと比べてほとんど高くありません。つまり:

小さく安いmodelがγ\gamma candidate tokensをautoregressivelyに生成します。

large modelがγ\gamma candidatesすべてに対して1回のforward passを実行し、各positionで自分なら何を言ったかを出します。

2つが一致する最長prefixを保持し、最初の不一致でlarge modelが無料で供給するtokenも加えます。残りは捨て、もう一度始めます。

output distributionは変わりません。greedy decodingなら明らかです。targetが生成したはずのtokenだけをacceptするからです。samplingでは修正されたacceptance ruleが必要で、Leviathan et al.は得られるdistributionがtargetのものと完全に同じであることを証明しています。10 これはこの章で2つ目のexact optimizationです。

したがってすべてはacceptance rate α\alphaにかかっています。これは測定可能です。上のagreement列であり、そこで計算した理由です。各quantized modelをfull-precision targetのdraftとして使い、144 generated positionsで測ると:

draft modelacceptancelongest accepted runtarget passあたり期待tokens, γ=4\gamma = 4
fp32 (target自身)100.0 %485.00
INT8 per row98.6 %484.86
NF4 block 6484.7 %203.69
INT4 group 12871.5 %132.85
INT4 per row58.3 %72.24
INT3 group 1285.6 %21.06
INT2 group 1280.0 %01.00

verification passごとにacceptされる期待tokens数は、draft length γ\gamma

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

net speedupは、それをdraft自身のcostで割ります。draftのcostはtargetのtokenあたりcostに対する割合ccです。

acceptancec=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

太字のentryを覚えてください。speculative decodingはgenerationを遅くすることがあります。 acceptanceが30 %で、draftがtargetの5分の1のcostなら、5回のforward passに支払って1.4 tokensしか保持できません。最後の列はもう1つの罠です。長いdraftはacceptanceが高いときにだけ役に立ちます。γ\gamma-token guessのtailにはほとんど到達しないからです。90 % acceptanceではγ=8\gamma = 8は3.40xの価値があり、30 %では0.79xの価値しかありません。同じconfigurationが、あなたのtrafficで測った数値次第で勝ちにも負けにもなります。

Quantizationは同じfunctionをより少ないbitsで保存することでmodelを縮めます。Distillationは、小さなmodelを大きなmodelの模倣へtrainingすることで縮めます11。deep learningよりほぼ10年古いideaです。12

微妙なのは、studentが何から学ぶかです。正解ではありません。それなら直接trainingできたはずです。teacherが加えるのは分布全体です。modelにあるphraseの次を尋ね、argmaxの先を見ます。

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

hard labelはjugと言い、それ以外は何も言いません。soft labelはjugと言い、さらにcupもほぼ同じくらい良く、bowlもあり得て、large — adjectiveであり、文法的にはまったく別のcontinuation — もまだ生きていると教えます。これが元の議論です。これは7だが、1にもかなり似ている。その類似性は、hard labelが捨てる情報です。

だからdistillationはtemperatureを使います。softmaxの前にlogitsをTTで割ると、distributionは平らになり、次点候補のrelative weightが上がります。このphraseでは、top tokenとthirdの比率がT=1T = 1で2.24だったものから、T=2T = 2で1.50に下がります。最初の値の平方根であり、logitsを2で割るとratioに起きることです。順序は同じで、near missesにlossのattentionがより多く向きます。studentのgradientはteacherの不確実性を運び、判定だけを運ぶのではありません。

この章のすべてはいま1つの和です。

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}}

ここでTTは、同時request全体にまたがるresident tokensの総数です。適用してみます。7Bと70B rowsはdimension 128のkey-value headsを8つ仮定し、13B rowは40 headsのfull multi-head attentionを仮定します。その世代のmodelはそう作られていたからです。そしてそれは表に現れます。

8 GB

modelprecisionweightsoverhead後の空き収まるcontext tokens
7Bfp1613.0 GB収まらない
7Bint86.5 GB収まらない
7Bint4 (g128)3.4 GB3.1 GB25,710
13Bint4 (g128)6.2 GB0.3 GB337
70Bint4 (g128)33.6 GB収まらない

16 GB

modelprecisionweightsoverhead後の空き収まるcontext tokens
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

modelprecisionweightsoverhead後の空き収まるcontext tokens
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 GB収まらない

8 GB表の13B rowを見てください。weightsは収まります。8 GBのうち6.2 GBです。だから通常の言い方では、13B modelは「8 GB cardで動く」ことになります。contextは337 tokensです。これは会話ではなく、promptとしてもぎりぎりです。「収まるか」は間違った質問です。正しい質問は「どれだけのcontextで、何人のusersを同時に扱えるか」です。

16 GBのint8の2 rowsも見てください。7Bは65,378 tokens、13Bは3,136 tokensです。追加weights 5.6 GBから20倍の差が出ています。ここでの13Bはmulti-head attentionであり、そのcacheはtokenあたり800 KB、7Bの128 KBに対して大きいからです。似たsizeの2 modelsの一方がlong contextで使い物にならない。その理由はどのmodel cardのheadlineにも出てきません。

13章前、これは2つのweightsとbiasを持つperceptronでした。いまや、設計され、trainingされ、alignedされ、難しい質問にcomputeを使うよう教えられ、tokenあたり測定済みcostでservedされるtransformerです。中に未開封のboxはもう残っていません。

ここで終わります。そして意図的に終わります。

第14章は、modelが別の場所にあるところから始まります。あなたのprocessの中でも、memoryの中でも、printできるvariableの中でもありません。あなたが管理しないmachine上、API keyとportと請求書の向こう側です。ここで測ったすべてはまだ起きています。first tokenの前にはprefillが走り、cacheは会話とともに増え、あなたが入っているbatchは依然として誰か別のものであり、あなたのlatencyを決めています。しかしこれからは、それをServer-Sent Eventsのstream、finish_reason、そしてRetry-After header付きHTTP 429を通して観測します。視点が変わると問いも変わります。このgradientはどう計算されるのかではなく、なぜ請求額が3倍になったのかになります。言葉も変わります。第14章はそのruleを宣言ではなく説明します。ここまではcodeがweights、gradients、logits、tokenizer bytesを持っていました。そこから先ではconnection、retry、cancellation、accumulated stateを持ちます。この境界を越えても、背後にある13章は捨てられません。それらはportの向こう側で動いているものの説明です。


2つの省略は意図的です。FlashAttention (Dao et al., arXiv:2205.14135) は別のattentionではありません。同じfunctionを、operationをtile化してn×nn \times n score matrixをmemoryに書き出さないように計算します。だから実践では、この章の2つ目の表にある67 MBは算術が示すより小さくなります。そしてkernels自体は委譲しています。Stanford CS336のlecture 10はinference systemsを、ここでは試みない深さで扱っています。CPU側については、llama.cpp repositoryとGGUF specificationがprimary sourcesです。

  1. Shazeer, N. Fast Transformer Decoding: One Write-Head is All You Need. arXiv:1911.02150 (2019). このpaperは大部分がmemory-bandwidthの議論であり、そのように読めます。

  2. Ainslie, J. et al. GQA: Training Generalized Multi-Query Transformer Models from Multi-Head Checkpoints. arXiv:2305.13245 (2023). 既存のmulti-head checkpointを変換するuptraining recipeを含んでおり、GQAが急速に広がった理由でもあります。

  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. iteration-level scheduling — continuous batching — とselective batchingを導入しています。

  4. Kwon, W. et al. Efficient Memory Management for Large Language Model Serving with PagedAttention. arXiv:2309.06180 (2023), SOSP 2023. vLLMの基盤となるpaperです。§3ではoperating-systems analogyが完全に展開されています。

  5. Dettmers, T., Pagnoni, A., Holtzman, A. and Zettlemoyer, L. QLoRA: Efficient Finetuning of Quantized LLMs. arXiv:2305.14314 (2023). NF4は§3で定義されています。上の測定で使った16個のlevel valuesは、このpaperが導出したものです。

  6. Dettmers, T., Lewis, M., Belkada, Y. and Zettlemoyer, L. LLM.int8(): 8-bit Matrix Multiplication for Transformers at Scale. arXiv:2208.07339 (2022). §4のoutlier-feature analysisが、上で測定した現象の出典です。outliersがscaleとともにsystematicに現れるという発見も含みます。 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). Theorem 1はoutput distributionが変わらないことの証明です。Chen et al. (arXiv:2302.01318)も同じideaを独立に発表しました。

  11. Hinton, G., Vinyals, O. and Dean, J. Distilling the Knowledge in a Neural Network. arXiv:1503.02531 (2015). temperatureと「dark knowledge」の議論です。

  12. Buciluă, C., Caruana, R. and Niculescu-Mizil, A. Model Compression. KDD 2006. transformersではなくensembles向けですが、distillationを9年早く扱っています。


作成者

David Vicente Campos

NeuraLIA Labs創業者、MyRealFood共同創業者

レオン大学出身のコンピューターエンジニアです。MyRealFoodを共同創業し、CTOとして、何百万人もの人がより良い食生活のために使ってきたアプリを開発しました。また、NeuraLIA Labsを創業し、そこでAIプロダクトを開発しています。ここでは、私がその過程で理解する必要があったことを、誰かにこう説明してほしかったと思う形で書いています。

著者について詳しく

NeuraLIA Labsが公開しています。

新着記事を受信トレイにお届け

AIニュース、ガイド、プロダクトアップデートを、読む価値のある記事を公開したときだけ短いメールでお送りします。

コース目次

Abstract software decision engine with branching paths, probability nodes, and glowing gates.
jev読了15分

Jev AIモデルは文章ではなく意思決定のために作られている

TypeSafe AIのJevが注目されているのは、ソフトウェアの知能を確率の問題として扱うからです。適切な分岐を選び、信頼度を添え、コードが必要としているのが意思決定であるときに、LLMに文章を書かせるためのコストを避けます。

Abstract legal research workspace with documents, search nodes and governance controls.
openai読了14分

OpenAIのAstra for Lawは新モデルではなく、法律AIシステム

OpenAIの法律分野での発表の本質は、新しい基盤モデルそのものではなく、その周辺にあるシステムです。ドメイン検索、信頼できるツール、権限、ベンチマーク、レビュー経路が重要になります。

Abstract agent runtime sorting documents, memory blocks and pointer nodes inside a bounded context frame.
context-engineering読了12分

Context engineering for long-horizon AI agents

Long-running agents do not fail only because the window is small. They fail when files, tool outputs and stale history crowd out the task the agent was supposed to finish.

モデル選びは、LIAにおまかせ。

すべてのAIモデルをひとつの場所で。今日から無料で。