Skip to content

枝刈り

実装: llm/prune/ / 実行: go test ./llm/prune/

誤差を配る量子化で使った道具は、もともと枝刈りのために考えられたものだった。素朴には絶対値の小さい重みから消せばよいが、消す損は重みの大きさだけでは決まらない。いつも大きく振れる入力に掛かっているなら、小さい重みでも出力を動かす。そして選び方と補い方は対になっている。補う前提の見積もりで選んでおいて補わないと、素朴な選び方より悪くなる。

この章で作るもの

誤差を配る量子化では、丸め先を選ぶのに入力の相関を使った。同じ道具で、別の問いにも答えられる。どれを 0 にするか、という問いだ。

もともと順序は逆になる。ヘッセ行列を使って重みを1つ選んで消し、残りを動かして補うという手は、枝刈りのために考えられたものだった(Optimal Brain Surgeon、1992年)。量子化に持ち込まれたのは、その30年後だ。

素朴には、絶対値の小さい重みから消せばよい。小さいのだから出力への寄与も小さいだろう、という理屈だ。

だがこれは近似でしかない。同じ 0.01 でも、いつも大きく振れる入力に掛かっているなら出力を動かすし、ほとんど動かない入力に掛かっているなら消しても何も起きない。

  入力の振れ方        重み       出力への寄与       絶対値で選ぶと
  ───────────────────────────────────────────────────────────────
  大きく振れる        0.01       それなりに動く     消される ✗
  ほとんど動かない    0.90       ほとんど動かない   残される ✗

  消す損の見積もり(OBS)

      L_i = w_i² / (2 · [H⁻¹]_ii)
             ▲            ▲
             │            └ 補いやすさ。残りで肩代わりできるほど大きい
             └ 重みの大きさ

  分母が付いているのが要点だ。
  補えるものは、消しても安い。
絶対値だけを見ると、消すべきでないものが小さく見える。重みの大きさと、入力の振れ方の両方で決まる

順に見ていく。

  1. 大きさで選ぶのは近似でしかない: 消す損は、重みの大きさと入力の振れ方の両方で決まる
  2. 選び方と補い方は対になっている: 補う前提の見積もりで選んで補わないと、かえって悪くなる
  3. 消すたびに残りを動かし、見積もりも作り直す: 1つ消すと、肩代わりできる度合いが変わる

① 大きさで選ぶのは近似でしかない

消す損の見積もりは1行で書ける:

go

// Saliency は重み1つを 0 にしたときの損の見積もりを返す。
//
//	L_i = w_i² / (2 · [H⁻¹]_ii)
//
// 分子は重みの大きさで、分母が「補いやすさ」になる。
// [H⁻¹]_ii が大きいほど、その重みを消したときに残りで肩代わりしやすいので、
// 損が小さく出る。絶対値だけを見るのと違うのは、この分母のぶんになる。
func Saliency(w []float64, hinv [][]float64) []float64 {
	out := make([]float64, len(w))
	for i := range w {
		out[i] = w[i] * w[i] / (2 * hinv[i][i])
	}
	return out
}

分子は重みの大きさで、分母が補いやすさになる。[H⁻¹]_ii が大きいほど、その重みを消したときに残りで肩代わりしやすいので、損が小さく出る。

絶対値だけを見るのと違うのは、この分母のぶんだ。テストで、よく振れる次元に小さい重みを乗せた場合に、絶対値の合計は小さいのに損の合計は大きくなることを固定した。だから選ばれる集合が食い違う。

逆に言うと、入力の振れ幅が揃っていて相関も無ければ、この2つは一致する。[H⁻¹]_ii が全次元で同じなら、順序は w_i² の順、つまり絶対値の順そのものになる。テストで、その場合に順序が完全に一致することを固定した。

絶対値で選ぶのが乱暴なのではなく、入力が揃っているという暗黙の前提が置かれている、と言うほうが正確だ。

go

// ByMagnitude は絶対値の小さい順に k 個を 0 にする。補正はしない。
func ByMagnitude(w []float64, k int) []float64 {
	score := make([]float64, len(w))
	for i, v := range w {
		score[i] = math.Abs(v)
	}
	out := append([]float64(nil), w...)
	for _, i := range smallest(score, k) {
		out[i] = 0
	}
	return out
}

// BySaliencyOnly は損の見積もりが小さい順に k 個を 0 にする。補正はしない。
//
// 選び方だけを変えて、補正の効果と切り分けるためのもの。
func BySaliencyOnly(w []float64, hinv [][]float64, k int) []float64 {
	out := append([]float64(nil), w...)
	for _, i := range smallest(Saliency(w, hinv), k) {
		out[i] = 0
	}
	return out
}

// smallest は score の小さい順に k 個の添字を返す。同点は添字の順で決める。
func smallest(score []float64, k int) []int {
	idx := make([]int, len(score))
	for i := range idx {
		idx[i] = i
	}
	sort.SliceStable(idx, func(a, b int) bool { return score[idx[a]] < score[idx[b]] })
	if k > len(idx) {
		k = len(idx)
	}
	return idx[:k]
}

② 選び方と補い方は対になっている

ここで思わぬことが起きる。

見積もりの分母は「補えるなら安い」と言っている。つまりこの見積もりは、補うことを前提にした値段になっている。だったら、補わずにこの見積もりだけで選ぶとどうなるか。

64次元、そのうち32個(50%)を消したときの出力のずれを測るとこうなった:

入力の共通成分大きさで選ぶ効きで選ぶ(補正なし)効きで選んで補う
0.00.189860.067050.06980
0.30.137820.050740.04953
0.70.099290.113790.02658
0.90.105450.190330.01696

入力どうしが似ていないうち(共通成分 0.0〜0.3)は、選び方を変えるだけで大きく効く。補正しても、ほとんど変わらない。肩代わりできる先が無いからだ。

だが入力がよく似て動くようになると(0.7 以上)、補正なしの選び方は素朴な選び方より悪くなる。共通成分 0.9 では 0.10545 が 0.19033 へ、倍近く悪化している。

理由は素直だ。「補えるから安い」と言って消したのに、補っていない。約束を果たしていない見積もりで選んでいることになる。

テストで、この裏目に出る現象と、相関が無いときは選び方だけでも効くことを、両方固定した。

そして補うところまで入れれば、どの条件でもいちばん良い。この表のどの行でも、いちばん小さいのは補正まで入れた列になる。

もう1つ、測って初めて見えたことがある。相関が強いところでは、選ぶ場所そのものはほとんど変わらない。共通成分 0.9 で 32個を消すとき、大きさで選んだ集合と効きで選んだ集合は、64 か所のうち 4 か所しか入れ替わらない。相関が弱い 0.0 では 34 か所が入れ替わるのに、だ。

つまり強い相関のもとでは、差の出どころは選び方ではなく補うかどうかだけになっている。同じ 4 か所の違いしかない選び方で、片方は 0.19033、もう片方は 0.01696 になる。

動かす

下のデモは、入力の相関の強さと消す割合を変えながら、3つのやり方を比べる。相関を上げていくと、補正なしの列だけが跳ね上がる。

デモ枝刈り共通成分 0.9 ・ 32 / 64 個を消す
入力の共通成分00.30.70.9消す数8163248
大きさで選ぶ0.10545
効きで選ぶ(補正なし)0.19033
効きで選んで補う0.01696

横の並びが 64 個の重み。色の付いたところが消したもの。右の数字が出力のずれ。 大きさで選ぶのと比べて、効きで選ぶと 4 か所、補正まで入れると 6 か所が入れ替わる

補える前提の見積もりで選んだのに補っていないので、素朴に大きさで選ぶより悪くなっている。見積もりと手順は対で意味を持つ

共通成分は、入力どうしがどれくらい一緒に動くかを表す。消す損の見積もりは「残りで肩代わりできる ほど安い」という形をしているので、補うことを前提にした値段になっている。だから共通成分を上げて いくと、補正なしの列だけが跳ね上がる。補うところまで入れれば、どの設定でもいちばん小さい。 共通成分を 0 に戻すと消す位置そのものが大きく入れ替わり、逆に 0.9 では数か所しか変わらない。 強い相関のもとでは、差の出どころは選び方ではなく補うかどうかだけになる。

③ 消すたびに残りを動かし、見積もりも作り直す

補うところは、誤差を配る量子化とほとんど同じ形になる:

go

// BySaliency は損の見積もりが小さい順に1つずつ消し、消すたびに残りを動かす。
//
// 手順はこうなる。
//
//	① まだ生きている中から、損の見積もりがいちばん小さいものを選ぶ
//	② その重みを 0 にする
//	③ 残りを δ = -(w_q / [H⁻¹]_qq) · H⁻¹[:,q] だけ動かして補う
//	④ 消した1つを取り除いた形に H⁻¹ を作り直す
//
// ④ が要るのは、1つ消すたびに「残りで肩代わりできる度合い」が変わるからになる。
// 作り直しは掛け算と引き算だけで済む(逆行列を取り直さなくてよい)。
func BySaliency(w []float64, hinv [][]float64, k int) []float64 {
	d := len(w)
	cur := append([]float64(nil), w...)
	inv := clone(hinv)
	dead := make([]bool, d)

	for step := 0; step < k && step < d; step++ {
		q, best := -1, math.Inf(1)
		for i := 0; i < d; i++ {
			if dead[i] || inv[i][i] <= 0 {
				continue
			}
			if l := cur[i] * cur[i] / (2 * inv[i][i]); l < best {
				q, best = i, l
			}
		}
		if q < 0 {
			break
		}

		// ③ 残りを動かして補う。消したぶんの重みが、肩代わりできる先へ移る。
		f := cur[q] / inv[q][q]
		for j := 0; j < d; j++ {
			if dead[j] || j == q {
				continue
			}
			cur[j] -= f * inv[j][q]
		}
		cur[q] = 0
		dead[q] = true

		// ④ H⁻¹ から q を抜いた形に作り直す。
		p := inv[q][q]
		for i := 0; i < d; i++ {
			if dead[i] {
				continue
			}
			for j := 0; j < d; j++ {
				if dead[j] {
					continue
				}
				inv[i][j] -= inv[i][q] * inv[q][j] / p
			}
		}
	}
	return cur
}

違うのは2つある。1つは、消す順が固定ではないこと。量子化は前から順になぞればよかったが、枝刈りでは毎回「いま残っている中でいちばん安いもの」を選び直す。

もう1つは、H⁻¹ を明示的に作り直していること。1つ消すたびに「残りで肩代わりできる度合い」が変わるからだ。作り直しは掛け算と引き算だけで済むので、逆行列を取り直す必要はない。

補正が入ると、残った重みは元の値から動く。テストで、消した位置は 0 のまま、残った位置は元と変わることを固定した。消した重みが消えるのではなく、残りへ移っている

補正の得は入力の相関がそのまま決める。共通成分 0.0 で 2.72倍、0.3 で 2.78倍、0.7 で 3.74倍、0.9 で 6.22倍。テストで、この順に大きくなることを固定した。似た動きをする入力が多いほど、消したぶんを引き受けてくれる先が多い。

設計の観点

  • 手軽な指標の前提を書き出す: 絶対値で選ぶのは「入力が揃っている」という前提の上に立っている
  • 見積もりと実際の手順を対にする: 補う前提の値段で選んだなら、補うところまでやる
  • 裏目に出る条件を測る: 効く場合だけでなく、悪くなる場合を出しておく
  • 1つ変えたら見積もりも変える: 消すたびに肩代わりできる度合いが変わる。最初に決めた順で消し続けない
  • 同じ道具を別の問いに向ける: 丸め先を選ぶのと、どれを消すかを選ぶのは、同じ枠組みで書ける
  • 疎さは数だけでは語れない: 何割消したかより、どこを消したかが効く

対照と実例

やり方選ぶ基準補うか要るもの
絶対値で選ぶ重みの絶対値しない重みだけ
OBD(1989)w² · H_ii / 2(対角だけ)しない(先が無い)入力の振れ幅
切り分け用w² / 2[H⁻¹]_iiしない入力の相関
OBS(1992)w² / 2[H⁻¹]_iiする入力の相関
構造化枝刈り行や列ごとの寄与手法による入力の相関

裏どり:

  • Optimal Brain Damage(1989): LeCun et al.。ヘッセ行列の対角だけを使い、L_i = w_i² · H_ii / 2 で選ぶ。対角だけなら [H⁻¹]_ii = 1 / H_ii なので、OBS の見積もりを対角に制限したものがそのまま OBD になる。同時に非対角が消えるので補う先も無くなる。テストで、この一致を固定した
  • 切り分け用の中間: この章の「効きで選ぶ(補正なし)」は、非対角まで使った見積もりで選びながら補わない形で、公表された手法ではない。選び方と補い方のどちらが効いているかを分けるために置いている
  • Optimal Brain Surgeon(1992): Hassibi, Stork。非対角まで使い、消したあと残りを動かす。この章の実装はこれにあたる
  • Optimal Brain Compression(2022): Frantar, Alistarh。OBS を層ごとに解いて大規模モデルに載せた。GPTQ の直接の土台
  • SparseGPT(2023): Frantar, Alistarh。同じ枠組みで、大規模モデルを1回のなぞりで 50% 疎にする
  • 構造化と非構造化: 個々の重みを 0 にしても、行列は疎になるだけで速くはならない。速さが欲しいなら行や列ごと落とす(構造化)か、2:4 のように専用の演算器が扱える形にする

簡略化したこと

  • 1行だけ: 実物は行列の全行を扱う。ここでは1本の重みベクトルだけ
  • 構造化なし: 個々の重みを 0 にするだけ。行や列ごと落とす話は扱わない
  • 疎の表現なし: 0 を実際に詰めて持つ形(疎行列)は作っていない。速さの話には踏み込まない
  • 再学習なし: 消したあとに学習し直して取り戻す手順は入れていない
  • 量子化との併用なし: 消してから丸める、あるいは同時にやる形は扱わない
  • キャリブレーションが人工的: 入力は共通成分と揺れから作る。実物は実際の文章を流して集める

参考資料