TT Lab
はじめる
学ぶ 学習パス コース

MiniMind — 小さな言語モデルを最初から最後まで自分で学習する

KV キャッシュは再計算しなくてよいものを覚えておくこと

TT Labで続きを見る

一言でいうと

生成は、トークンを1つ選ぶたびに、モデルを1回回す作業です。キャッシュがないと、そのつど、ここまでの全体を入れ直して計算し、KVキャッシュがあれば、新しいトークン1つだけを入れて、前のトークンのK・Vは取り出して使います。結果は1トークンも違わず、計算だけが減ります。選び方(グリーディ・温度・top-p)は、それとは別に、「どのトークンを選ぶか」を決めます。このモジュールでは、MiniMindのgenerateで、その両方を数字で測ります。

なぜ必要なのか

Transformerのコース(Transformer — アテンションを手で計算する)で、KVキャッシュの原理とサンプリングの式を、手で計算しました。ここでは、実際に学習したモデルと実際の生成ループで、それがどう見えるかを見ます。サービングのコストは、ほとんどが生成の段階から出て、そのコストの形は、キャッシュが決めます。そして、同じモデルが、グリーディでは常に同じ答えを、温度を上げるとそのつど違う答えを出しますが、その違いが「創造性」なのか「でたらめ」なのかは、小さなモデルで特にはっきり表れます。

どう動くのか

MiniMindのgenerateは、次のように回ります。

for _ in range(max_new_tokens):
    past_len = past_key_values[0][0].shape[1] if past_key_values else 0
    outputs = self.forward(input_ids[:, past_len:], past_key_values=past_key_values, use_cache=use_cache)
    logits = outputs.logits[:, -1, :] / temperature
    ... top_k · top_p 로 자르기 ...
    next_token = torch.multinomial(softmax(logits), 1) if do_sample else argmax(logits)
    input_ids = torch.cat([input_ids, next_token], -1)
    past_key_values = outputs.past_key_values if use_cache else None

KVキャッシュでは、最初の1回に、プロンプトのPトークンを一度に入れます(prefill)。キャッシュがあれば、それ以降は1回ごとに1トークンだけを入れるので、N個を作る間にforwardに入ったトークンは、P + (N − 1)個です。キャッシュがなければ、1回ごとにP、P+1、…を入れ直すので、P·N + N(N−1)/2個です。プロンプト183トークンで64トークンを作ると、246対13,728で、56倍です。因果マスクのおかげで、前のトークンのK・Vは、後ろにトークンができても変わらないので、取り出して使っても結果が同じです。

キャッシュの大きさは、層ごとのK・Vで、形は(バッチ, 長さ, KVヘッド, head_dim)です。MiniMindは、repeat_kvで複製する前のK・Vをキャッシュに入れるので、GQAで減らした分だけ、そのまま減ります。トークン100個なら、2 × 4層 × 100 × 2ヘッド × 32 × 4バイト = 204,800バイトです。サービングでは、この数字が同時ユーザー数を決めます。

温度では、ロジットを温度で割ってからソフトマックスします。1より小さいと、分布が尖って、最ももっともらしいトークンに偏り、大きいと平らになって、まれなトークンも選ばれます。温度が0に近いと、グリーディと同じです。

top-p(nucleus)では、確率の大きい順に並べて、累積確率がpを超える地点までだけを残し、残りを−∞で消します。MiniMindは、マスクを1つずらして、pを初めて超えるトークンまでを残し、先頭のトークンは常に残します。分布が尖っているときは、数個だけ、平らなときは多く残ります。固定の個数を残すtop-kとの違いです。MiniMindのgenerateのデフォルトは、温度0.85・top_p 0.85・top_k 50です。

現場での姿

「キャッシュをオンにしたら、答えが変わった」というバグ報告は、たいていキャッシュではなく、サンプリング側の問題です。乱数を固定していないか、片方だけdo_sampleになっています。グリーディで2つの方式を比べて、1トークンも違わないかを先に確認すれば、原因を半分に絞れます。逆に、長い文書を要約するとき、最初のトークンが出るまでの時間(TTFT)が長いのは、キャッシュでは減りません。prefillは、もともとプロンプト全体を計算する必要があるからです。キャッシュが減らすのは、そのあとのトークンの間の時間です。

MiniMindのオリジナルとこのコースの違い

MiniMindのeval_llm.pyとWebデモは、温度0.85・top_p 0.95で選び、繰り返しペナルティ(repetition_penalty)を与えられます。このコースは、実験をきれいにするために、top_kをオフにして(top_k=0)、一度に1つだけを変えます。また、MiniMindのgenerateは、バッチごとに終わった行を別に覚えて(finished)、すべてが終わったときに止まり、終わった行には、eosを埋め続けます。30個を一度に選ぶとき、短い答えの後ろにeosが続く理由です。実際のサービングエンジン(vLLMなど)は、これに、リクエストごとに違う長さのキャッシュを1つのバッチにまとめる仕組みを加えます。原理は同じで、管理が複雑になるだけです。

次のラボですること

基準の事前学習モデルで、キャッシュがあるときとないときに同じトークンが出るかを確認し、forwardに入ったトークン数を数えて式と合わせ、時間を測ります。キャッシュの実際のバイト数を式と合わせ、SFTモデルで温度を変えて30回ずつ選び、異なる答えの数を数え、top-pが残す候補の数を数えます。