Digital Reactor
機械学習

SAEが取り出した特徴は、業務の監視項目になるか

SAEが取り出した特徴は、業務の監視項目になるか

はじめに

LLMを組み込んだ文書処理や照会応答の検証で、事業側の担当者から「この出力はなぜこうなったのか」と聞かれる場面があります。モデルの内部を開けば数値そのものは見えます。GPT-2のような小型のモデルでも、1トークンあたり768個の実数が層から層へ流れていて、答えの根拠はこのどこかにあるはずです。ただ、その数値の羅列を眺めても、どれが何を意味するのかは読み取れません。中間層の活性をダンプして見せたところで、説明としては機能しません。

この読めない数値の列を、意味を読み取れる形に変換する手法があります。解釈可能性(interpretability)研究の中心にある、スパースオートエンコーダ(sparse autoencoder, SAE)です。SAEは、モデルがテキストを処理するときに内部で繰り返し現れる数値のパターンを、大量のテキストから教師なしで拾い出します。うまくいくと「神や神性の話題」「スポーツチームの名前」のようなパターンが番号付きで手に入り、1トークンぶんの数値の列を「どのパターンがどれだけ強く出ているか」に読み替えられます。Anthropicは自社モデルで「ゴールデンゲートブリッジ」のパターンを人為的に強め、何を聞いても橋の話に戻ってくるGolden Gate Claudeというデモを公開しました。この読み替えが、モデルの挙動を動かす操作にも使えることを示した例です。

そうなると実務側では次の発想が出てきます。拾い出したパターンを業務の監視項目にできないか、という発想です。「特定のパターンが強く出たらログに残して人が確認する」という運用が組めるなら、出力の説明にも異常の検知にも使えます。

この記事では、GPT-2の中間層を対象に自前のSAEを訓練して、この発想がどこまで成立するかを確かめます。分解そのものはCPUだけの小さな実験でも動きます。約20万トークンで訓練したSAEに1トークンぶんの内部状態を通すと、768個の数値の羅列が、4096個のパターンのうち百個ほどの組み合わせに置き換わりました。その先で監視項目として運用に載せるには、意味づけが人手の解釈作業として残ること、スパース性と忠実さのトレードオフ、辞書サイズしだいで項目の定義自体が変わることの3つの壁を越える必要があります。実験でこの3つを1つずつ確かめます。

対象読者:

  • LLMや深層学習モデルを業務システムに組み込んでいて、出力の説明や監視の仕組みを検討している方
  • スパースオートエンコーダ(SAE)が何をする道具なのか、実装レベルで知りたい機械学習エンジニア
  • 解釈可能性の研究成果を実務に持ち込むときに何が障害になるかを先に知りたい方

記事のポイント:

  • 1トークンぶんの活性ベクトルが、SAEで名前の付く少数の成分に分解される様子を実物の数値で追います
  • スパース性の係数を変えた自前SAEを訓練し、スパース性(L0)と再構成・モデル損失のトレードオフを実測します
  • 最大活性例による意味づけ、ステアリング、feature splittingの実験から、監視項目化に何が足りないかを整理します

1トークンの内部状態を、名前の付く成分に分解する

分解される対象を1件だけ実物で見ます。GPT-2 small(124Mパラメータ)に、wikitextにあるヒンドゥー寺院の解説の一節「… the portico depicts Shiva and Parvati seated …」(入口の彫刻はシヴァ神とパールヴァティー神の座像を描く、という英文)を読ませます。12層の変換の途中、第8ブロックに入る時点の内部状態(残差ストリーム、residual stream。モデルを貫いて層から層へ受け渡されるベクトルの列)から、「Shiva」のトークン位置にある768次元のベクトル xx を取り出します。冒頭の6成分は [1.69,3.36,2.64,2.68,2.91,1.04][1.69, -3.36, -2.64, -2.68, -2.91, 1.04] という実数の並びで、残り762個も同様です。ここからは何も読めません。

このベクトルをSAEに通します。SAEはエンコーダとデコーダが1層ずつの小さなオートエンコーダで、デコーダ側は、訓練で獲得した4096本の「方向」 d1,,d4096d_1, \dots, d_{4096}(各768次元、単位ノルムに正規化します。この一覧を辞書、その本数を辞書サイズと呼びます)を、係数 a1,,a4096a_1, \dots, a_{4096} で足し合わせて元のベクトルを復元しようとします。

x^=jajdj+bdec\hat{x} = \sum_{j} a_j d_j + b_{\mathrm{dec}}

エンコーダ側は、入力 xx からこの係数 aa を決める係です。訓練時にかける制約(「スパース性と忠実さのどちらを取るか」の節で見ます)によって、係数はほとんどが0になります。つまりSAEは、768個の数値の列を「4096本の方向のうち、どれがどれだけ含まれているか」という表現に書き直す装置です。実装は素直です。

class SAE(nn.Module):
    def __init__(self, d_in, d_sae):
        super().__init__()
        w = torch.randn(d_sae, d_in)
        w = w / w.norm(dim=1, keepdim=True)
        self.W_dec = nn.Parameter(w)              # (d_sae, d_in) 各行が単位ノルム
        self.W_enc = nn.Parameter(w.t().clone())  # (d_in, d_sae)
        self.b_enc = nn.Parameter(torch.zeros(d_sae))
        self.b_dec = nn.Parameter(torch.zeros(d_in))

    def encode(self, x):
        return torch.relu((x - self.b_dec) @ self.W_enc + self.b_enc)

    def forward(self, x):
        a = self.encode(x)
        return a @ self.W_dec + self.b_dec, a

wikitext-103の冒頭から集めた約20万トークンぶんの活性で訓練したSAEに、先ほどの xx を入れます。4096個の係数のうち非ゼロは116個、残り3980個は0でした。大きい順に6個を、あとの節で行う意味づけの結果とあわせて並べます。

潜在変数の番号係数 aja_j何に反応する方向か
4913.91大文字で始まる集団・チームの名前
39712.84非英語圏の人名の語中サブワード
202.71文中で再登場する人物の姓
22322.62神・神性の語彙(gods、divine、deities)
34572.57埋葬・発掘の話題
32121.96「Shiva」そのもの

神話の文脈に置かれた「Shiva」という1トークンの内部状態が、「神の語彙+集団名+人名らしさ+祭祀の話題」の重ね合わせとして書けている、ということです。768個の読めない実数が、名前の付く成分116個の強さに置き換わりました。3列目の名前をどうやって付けたのかは、意味づけの節で説明します。

ここで用語を2つ決めます。SAEが学習した4096本の方向のそれぞれを、この記事では潜在変数(latent)と呼びます。一方、「神・神性の語彙である」のようなデータ側の性質は特徴(feature)と呼んで区別します。各潜在変数が何かひとつの特徴にきれいに対応することを期待して訓練するのですが、本当に対応しているかは人が確かめるまで分からず、その確認作業が監視項目化の最初の壁になります。もうひとつ、再構成 x^\hat{x} と元の xx のずれは二乗和で17.6%残っています。SAEは情報の一部を落とす近似で、落ちた部分で起きる変化はSAE越しの監視では見えません。

素のニューロンを数えても監視項目にならない

意味づけの話に進む前に、なぜSAEのような回り道が要るのかを整理します。768次元のベクトルがあるなら、768個の成分(ニューロン)を1本ずつ監視すればよさそうに思えます。実際、初期の解釈可能性研究はニューロン単位の分析を重ねてきました。

この方針がうまくいかないのは、個々のニューロンが複数の無関係な意味で発火するからです。これを多義性(polysemanticity)と呼びます。同じニューロンが学術論文の引用でも、フランス語の文でも、数式の途中でも強く反応する、といった観察が典型です。このニューロンの値が高いことは何かの兆候ではあっても、「何の」兆候かが決まらないので、監視項目としては成立しません。

多義性が生じる理由の有力な説明が重ね合わせ(superposition)仮説です。モデルが表現したい特徴の数は、文法・話題・固有名詞・書式などを数え上げると、768という次元数よりはるかに多くなります。高次元空間では、互いにほぼ直交する方向を次元数よりずっと多く詰め込めることが知られています。個々の特徴がめったに同時に現れない(スパースである)なら、多数の特徴を768次元に重ねて格納しても干渉は小さく抑えられます。モデルは実際にこの詰め込みをやっている、というのが仮説の主張です。

この仮説を受け入れると打ち手も決まります。特徴の数がニューロンの本数より多いなら、次元数より大きい辞書(過完備、overcomplete)を用意し、「同時に立つ(非ゼロになる)潜在変数は少数」という制約をかけて読み出せばよいことになります。SAEの辞書サイズを768より大きい4096に取り、L1正則化でスパース性を課すのは、重ね合わせ仮説の裏返しです。

スパース性と忠実さのどちらを取るか

SAEの訓練では、再構成の正確さをどこまで犠牲にして、同時に立つ潜在変数の数を減らすかという調整が入ります。損失関数は再構成誤差とL1正則化の和で、係数 λ\lambda がその調整弁です。

l1 = l1_coeff * min(1.0, step / warmup)  # 訓練初期はL1を線形に立ち上げる
loss = ((xhat - xb) ** 2).sum(-1).mean() + l1 * a.sum(-1).mean()

λ\lambda(コードの l1_coeff)を大きくすると、同時に立つ潜在変数の平均個数(L0と呼びます)が減って読みやすくなりますが、再構成が粗くなります。粗さは2つの物差しで測ります。ひとつは再構成で説明できない分散の割合(fraction of variance unexplained, FVU)です。もうひとつは、GPT-2の第8ブロック以降を再構成 x^\hat{x} で流し直したときの、次トークン予測損失(交差エントロピー)の増分です。FVUが小さくても、モデルの動作に効く成分をSAEが落としている可能性があるため、後者を併記します。

監視の観点では、この予測損失の増分が本命の物差しになります。SAEが再構成できない成分の中で起きる変化は、その監視から原理的に漏れます。

λ\lambda を4水準に振り、それぞれ約20万トークン×8エポックで訓練した結果です。置換前の予測損失は3.954でした。

λ\lambdaL0(平均)FVU予測損失の増分
0.5869.80.136+0.078
2.0127.20.481+0.993
3.024.80.673+1.780
4.07.60.858+2.938

λ=0.5\lambda=0.5 は忠実ですが、1トークンあたり平均870個も立つので、読む対象としては素のニューロンとあまり変わりません。λ=3.0\lambda=3.04.04.0 は同時に立つ数が数十個以下まで減る代わりに、FVUが0.67を超え、予測損失の増分も1.78、2.94と大きくなります。ここまで再構成が粗いと、SAEはモデルの計算の一部しか見ておらず、その潜在変数を監視してもモデルの実態を監視したことになりません。以降の実験は、中間の λ=2.0\lambda=2.0(L0=127)で進めます。

それでも増分0.99は小さくない値で、モデルが正解の次トークンに与える確率が平均でおよそ3分の1に下がるのに相当します。研究で使われる公開SAE(SAELensで配布されているGPT-2用SAEなど)は、数億トークン規模の訓練でこれよりはるかに小さいFVUと数十のL0を両立させています。スパース性と忠実さを両立させること自体に、データと計算の予算が要ります。

潜在変数に名前を付けるのは人手の作業

選んだSAEの潜在変数4096個は、訓練を終えた時点ではただの番号です。冒頭の表に先取りで載せた名前も、訓練が自動で付けてくれたものではありません。番号に意味を与える標準の方法が最大活性例(max activating examples)の観察です。コーパス全体で各潜在変数の係数を計算し、最も強く発火したトークンとその前後を並べて、共通点を人が読み取ります。

うまく読み取れる例から見ます。冒頭の「Shiva」の分解で4番目に立っていた潜在変数2232の最大活性例です。

活性文脈([[ ]]が発火したトークン)
6.49sex counterparts to established [[gods]] or goddesses .
6.46absorbed the characteristics of older regional [[gods]] . Horus had many
6.39emerged in the New Kingdom to represent [[divine]] rescue from harm ,
6.23expressing a particular perspective on [[divine]] events . The contradictions
6.02ism , which totally excludes belief in other [[deities]] . There is evidence

gods、divine、deitiesと、単語ではなく「神・神性の語彙」というまとまりで発火しています。冒頭の表の3列目に載せた名前は、すべてこの手順で付けたものです。491なら「Bulls」「Wizards」のような大文字で始まる集団・チームの名前、3971なら「Hashmi」「Fujibayashi」のような非英語圏の人名の語中サブワードで発火が並び、そこから名前を読み取りました。ほかにも、サッカーの試合結果「3 – 0」の数字の位置だけで発火する潜在変数1455など、共通点が一目で読めるものは多数ありました。

ただし、すべての潜在変数がこう読めるわけではありません。同じSAEの潜在変数618の最大活性例です。

活性文脈
1.18Congress ended the series .= =[[ Background]] = =In proposing
1.16continued throughout the war .= =[[ Background]] = =In 17
1.12on May 14 , the[[ Blue]] Jackets announced that Richards
1.04the prospect of another rebuild looming the[[ Blue]] Jackets ’ captain and

「Background」という節見出しと、NHLのチーム名「Blue Jackets」が混ざっていて、ひとつの意味に要約できません。今回のSAEの規模でも、こうした潜在変数は珍しくありませんでした。多義性を解消するために訓練したSAEの中に、まだ多義的な潜在変数が残っています。

これが監視項目化の1つ目の壁です。潜在変数は4096個あり、監視項目の候補を洗い出すにはこの解釈作業をその数だけ繰り返すことになります。最大活性例をLLMに見せて説明文を書かせる自動化(autointerp)も研究されていますが、説明が正しいかの検証は結局人に戻ってきます。探したい特徴が先に決まっているなら、ラベル付きの例文から方向を直接学習する線形プローブのほうが早道です。SAEは「何を探すか」を決めずに辞書ごと作る手法で、仮説なしにモデル内部を棚卸しできる代わりに、意味づけの費用を後払いします。

潜在変数を足すと出力が動く

意味づけできた潜在変数が、相関ではなく因果でモデルの計算に効いているのかを確かめる方法がステアリング(feature steering)です。生成の最中に、潜在変数のデコーダ方向 djd_j を係数倍して残差ストリームに足し続けます。

def hook(module, inputs, output):
    if isinstance(output, tuple):
        return (output[0] + add,) + output[1:]  # add = 係数 * d_j
    return output + add

先ほどの「神・神性の語彙」の潜在変数2232で試します。係数は、この潜在変数がコーパスで記録した最大活性(6.49)の0倍・2倍・5倍・10倍とし、中立なプロンプト「I went outside and looked around.」から40トークンの続きを各20本サンプリングしました。係数0では日常の描写が続きます。2倍にすると、同じプロンプトから「All the gods were there, but no one was in place.」「the angels and the light of the gods were scattered」のような神話めいた文が生成され、20本中15本に神関連の語が現れました。方向 d2232d_{2232} が「神の話題」の内部表現として因果的に機能している証拠です。

係数を上げると副作用が出ます。5倍では「One and not, but in the gods’, or, they, or they. God, or things.」のように文法が崩れ始め、10倍では「gods, gods. gods. gods. gods.」の繰り返しだけになりました。定量指標でも、話題語を含む生成の割合は0%→75%→95%→100%と上がる一方、語彙の多様性(相異なるバイグラムの割合、distinct-2)は0.94→0.90→0.49→0.24と崩れていきます。

ひとつ注意が要るのは流暢さの測り方です。生成文の質をGPT-2自身の平均NLL(負の対数尤度)で測ると、係数5倍までは2.26→2.96と悪化するのに、10倍では2.45へ戻ります。同じ語の繰り返しはモデルにとって予測しやすい系列だからで、尤度ベースの指標は崩壊をむしろ高評価します。ステアリングの副作用を運用で監視するなら、distinct-2のような多様性の指標を併用する必要があります。

辞書サイズを変えると監視項目が消える

最後の壁は、潜在変数の定義そのものの不安定さです。辞書サイズ4096は今回の設定にすぎず、大きくも小さくもできます。そこで同じデータ・同じ λ=2.0\lambda=2.0 のまま、辞書サイズだけを1024と16384に変えたSAEを追加で訓練しました(16384は8エポックでは明らかに収束せず、16エポックまで延長しています。辞書を広げると訓練の予算も膨らみます)。そのうえで、同じ特徴がどの潜在変数に割り当てられるかを、デコーダ方向の余弦類似度で突き合わせました。方向はどれも単位ノルムなので、内積がそのまま類似度になります。

対象は、1024側で「gods」「god」に発火する潜在変数1017です。1024個しか枠がない辞書では、「神々」がひとつの潜在変数にまとまっています。この方向を4096・16384の全潜在変数と照合すると、類似度0.5以上の対応は4096側で1本、16384側では9本ありました。

内訳を見ると、対応が「増えた」のではなく「割れた」ことが分かります。4096側の最上位は神・神性の語彙(類似度0.51)で、以下は神殿での神との接触(0.41)、神々の性質の記述(0.40)、シヴァ神(0.30)と続きます。16384側では最上位が「土地や都市の守護神」(0.68)、次いで「神々についての学術的な記述」(0.64)、「deitiesという語そのもの」(0.57)、「動物と神格の結びつき」(0.51)と、どれも1024側の「神々」の一断面です。辞書を広げるほど同じ特徴が細かい変種に分かれていくこの現象は、feature splittingとしてAnthropicのSAE研究の初期から報告されています。研究の文脈では、こうした辞書間の潜在変数の方向をUMAPなどの次元削減で並べて、分裂の構造を観察します。

運用の言葉に訳すと、この現象は「監視項目のIDが再現しない」という問題です。「潜在変数1017が立ったらログに残す」という定義は、辞書サイズを変えるとそのまま使えなくなります。変更後の辞書に1017番へ1対1で対応する潜在変数が存在しないからです。同じ設定でも、訓練データや乱数が変われば番号は付け替わります。モデルを更新すれば活性の分布が変わるのでSAEも再訓練になり、そのたびに監視項目の番号も粒度も引き直しです。監視項目を維持したいなら、定義を「どのSAEの何番」ではなく「この検証用例文集合に対してこう振る舞う潜在変数」という挙動ベースに置き、再訓練のたびに対応を取り直す仕組みまで含めて設計する必要があります。

まとめ

SAEの分解そのものは、この規模の実験でも期待どおりに動きます。768個の読めない数値が百個前後の潜在変数の組み合わせになり、その多くには「神・神性の語彙」「試合結果のスコア」のような名前が付けられて、ステアリングの実験からは、取り出した方向が生成を実際に動かす因果的な実体であることも確かめられました。研究の道具としてのSAEの魅力は、この実験で十分に感じられます。

一方で「業務の監視項目になるか」という冒頭の問いに対しては、越えるべき壁が3つ残ります。意味づけは潜在変数の数だけ発生する人手の解釈作業で、同時に立つ数を読めるまで絞ったSAEはモデルの計算の一部しか再構成できていません。しかも潜在変数の番号はそのSAE限りのもので、辞書サイズを変えたり再訓練したりすると付け替わるため、「何番を見張る」と決めた監視ルールはそのたびに作り直しになります。どれも解決不能ではありませんが、「SAEを通せばモデルの中身が読めるようになる」という距離感で監視設計に組み込むと、意味づけとメンテナンスの費用に途中で気づくことになります。

これから検討するなら、順序を逆にするのが現実的です。まず監視したい事象を文章で1つ書き出し、その事象を含む例文と含まない例文を数十件ずつ集めてください。その例文集さえあれば、公開されている訓練済みSAE(GPT-2向けのほか、Gemma向けのGemma Scopeなどが公開されています)で該当する潜在変数が安定して立つかを、この記事のスクリプト程度の規模で確かめられます。辞書を自前で訓練し直すのは、その検証が動いてからで遅くありません。

コード

実験に使ったスクリプトの全体です。

# SAE latent monitoring experiment: train sparse autoencoders on GPT-2 layer-8
# residual stream activations, then study L0/reconstruction tradeoff, max
# activating examples, feature steering and feature splitting.
#
# Stages (run in order, each caches its output under cache/):
#   python script.py collect    # gather activations from wikitext-103
#   python script.py train      # train SAEs (L1 sweep + dict-size sweep)
#   python script.py examples   # dump max-activating examples for inspection
#   python script.py steer      # feature steering experiment
#   python script.py split      # feature splitting (cosine similarity)
#   python script.py figures    # final figures
import argparse
import json
import math
from pathlib import Path

import numpy as np
import torch
import torch.nn as nn
from dotenv import load_dotenv

HERE = Path(__file__).parent
CACHE = HERE / "cache"
CACHE.mkdir(exist_ok=True)
load_dotenv(HERE.parent.parent / ".env")  # HF_TOKEN

SEED = 0
LAYER = 8          # blocks.8 の resid_pre = ブロック7の出力
SEQ_LEN = 128
N_TRAIN_SEQ = 1600  # 1600 * 128 = 204,800 tokens
N_HELDOUT_SEQ = 24  # CE損失評価用
D_MODEL = 768

torch.manual_seed(SEED)
np.random.seed(SEED)


def load_gpt2():
    from transformers import GPT2LMHeadModel, GPT2TokenizerFast
    tok = GPT2TokenizerFast.from_pretrained("gpt2")
    model = GPT2LMHeadModel.from_pretrained("gpt2")
    model.eval()
    return model, tok


# ---------------------------------------------------------------- collect
def collect():
    model, tok = load_gpt2()
    from datasets import load_dataset
    ds = load_dataset("Salesforce/wikitext", "wikitext-103-raw-v1",
                      split="train", streaming=True)
    need = (N_TRAIN_SEQ + N_HELDOUT_SEQ) * SEQ_LEN
    ids = []
    for ex in ds:
        text = ex["text"].strip()
        if not text:
            continue
        ids.extend(tok(text)["input_ids"])
        if len(ids) >= need:
            break
    tokens = np.array(ids[:need], dtype=np.int32).reshape(-1, SEQ_LEN)
    print(f"collected {tokens.shape[0]} sequences of {SEQ_LEN} tokens")

    acts = np.empty((tokens.shape[0], SEQ_LEN, D_MODEL), dtype=np.float32)
    bs = 16
    with torch.no_grad():
        for i in range(0, tokens.shape[0], bs):
            batch = torch.tensor(tokens[i:i + bs], dtype=torch.long)
            out = model(batch, output_hidden_states=True)
            # hidden_states[LAYER] = ブロック LAYER に入る残差ストリーム
            acts[i:i + bs] = out.hidden_states[LAYER].numpy()
            if i % 160 == 0:
                print(f"  {i}/{tokens.shape[0]}")
    np.save(CACHE / "tokens.npy", tokens)
    np.save(CACHE / "acts.npy", acts)
    print("saved", acts.shape)


def load_acts():
    """訓練用の活性化(位置0は分布が特異なので除外)と付随情報を返す。"""
    tokens = np.load(CACHE / "tokens.npy")
    acts = np.load(CACHE / "acts.npy")
    train_acts = acts[:N_TRAIN_SEQ, 1:, :].reshape(-1, D_MODEL)
    # 位置合わせ用の (seq, pos) インデックス
    seq_idx, pos_idx = np.meshgrid(np.arange(N_TRAIN_SEQ),
                                   np.arange(1, SEQ_LEN), indexing="ij")
    index = np.stack([seq_idx.ravel(), pos_idx.ravel()], axis=1)
    scale = math.sqrt(D_MODEL) / np.linalg.norm(train_acts, axis=1).mean()
    return tokens, acts, train_acts, index, float(scale)


# ---------------------------------------------------------------- SAE
class SAE(nn.Module):
    def __init__(self, d_in, d_sae):
        super().__init__()
        w = torch.randn(d_sae, d_in)
        w = w / w.norm(dim=1, keepdim=True)
        self.W_dec = nn.Parameter(w)            # (d_sae, d_in) 各行が単位ノルム
        self.W_enc = nn.Parameter(w.t().clone())  # (d_in, d_sae)
        self.b_enc = nn.Parameter(torch.zeros(d_sae))
        self.b_dec = nn.Parameter(torch.zeros(d_in))

    def encode(self, x):
        return torch.relu((x - self.b_dec) @ self.W_enc + self.b_enc)

    def forward(self, x):
        a = self.encode(x)
        return a @ self.W_dec + self.b_dec, a


def train_sae(x, d_sae, l1_coeff, epochs=8, batch=4096, lr=1e-3):
    torch.manual_seed(SEED)
    sae = SAE(D_MODEL, d_sae)
    with torch.no_grad():
        sae.b_dec.copy_(x.mean(0))
    opt = torch.optim.Adam(sae.parameters(), lr=lr)
    n = x.shape[0]
    steps_total = epochs * (n // batch)
    warmup = max(1, steps_total // 20)
    step = 0
    for ep in range(epochs):
        perm = torch.randperm(n)
        for i in range(0, n - batch + 1, batch):
            xb = x[perm[i:i + batch]]
            xhat, a = sae(xb)
            l1 = l1_coeff * min(1.0, step / warmup)  # 訓練初期はL1を線形に立ち上げる
            loss = ((xhat - xb) ** 2).sum(-1).mean() + l1 * a.sum(-1).mean()
            opt.zero_grad()
            loss.backward()
            opt.step()
            with torch.no_grad():  # デコーダ行を単位ノルムに保つ
                sae.W_dec.div_(sae.W_dec.norm(dim=1, keepdim=True))
            step += 1
        with torch.no_grad():
            xhat, a = sae(x[:20000])
            l0 = (a > 0).float().sum(-1).mean().item()
            fvu = (((xhat - x[:20000]) ** 2).sum() /
                   ((x[:20000] - x[:20000].mean(0)) ** 2).sum()).item()
        print(f"  d_sae={d_sae} l1={l1_coeff} epoch {ep+1}: L0={l0:.1f} FVU={fvu:.3f}")
    return sae, l0, fvu


@torch.no_grad()
def ce_delta(model, sae, tokens, scale):
    """層LAYERの残差をSAE再構成に置換したときのCE損失の悪化を測る。"""
    heldout = torch.tensor(tokens[N_TRAIN_SEQ:], dtype=torch.long)
    block = model.transformer.h[LAYER - 1]  # この出力が resid_pre(LAYER)

    def hook(module, inputs, output):
        is_tuple = isinstance(output, tuple)
        h = output[0] if is_tuple else output
        xhat, _ = sae(h * scale)
        h_new = h.clone()
        h_new[:, 1:, :] = xhat[:, 1:, :] / scale  # 位置0は訓練対象外なので温存
        return (h_new,) + output[1:] if is_tuple else h_new

    def ce(logits, tokens):
        return nn.functional.cross_entropy(
            logits[:, :-1].reshape(-1, logits.shape[-1]),
            tokens[:, 1:].reshape(-1)).item()

    clean = ce(model(heldout).logits, heldout)
    handle = block.register_forward_hook(hook)
    patched = ce(model(heldout).logits, heldout)
    handle.remove()
    return clean, patched


def train_all():
    model, _ = load_gpt2()
    tokens, _, train_acts, _, scale = load_acts()
    x = torch.tensor(train_acts) * scale
    runs = [
        # (名前, 辞書サイズ, L1係数) : L1掃引は d_sae=4096 で行う
        ("sweep_l1_0.5", 4096, 0.5),
        ("sweep_l1_2.0", 4096, 2.0),
        ("sweep_l1_3.0", 4096, 3.0),
        ("sweep_l1_4.0", 4096, 4.0),
        # 辞書サイズ掃引(feature splitting 用)は L1 を掃引で選んだ値に揃える
        ("dict_1024", 1024, 2.0),
        ("dict_16384", 16384, 2.0),
    ]
    results = {}
    if (CACHE / "results.json").exists():
        results = json.loads((CACHE / "results.json").read_text())
    for name, d_sae, l1 in runs:
        if name in results and (CACHE / f"sae_{name}.pt").exists():
            print(f"skip {name} (already trained)")
            continue
        print(f"training {name}")
        sae, l0, fvu = train_sae(x, d_sae, l1)
        clean, patched = ce_delta(model, sae, tokens, scale)
        results[name] = {"d_sae": d_sae, "l1": l1, "L0": l0, "FVU": fvu,
                         "ce_clean": clean, "ce_patched": patched,
                         "scale": scale}
        torch.save(sae.state_dict(), CACHE / f"sae_{name}.pt")
        print(f"  CE {clean:.3f} -> {patched:.3f}")
        (CACHE / "results.json").write_text(json.dumps(results, indent=2))


def extend(name, extra_epochs=8):
    """収束が遅い大きい辞書のSAEを、チェックポイントから追加訓練する。"""
    model, _ = load_gpt2()
    tokens, _, train_acts, _, scale = load_acts()
    x = torch.tensor(train_acts) * scale
    results = json.loads((CACHE / "results.json").read_text())
    r = results[name]
    sae = SAE(D_MODEL, r["d_sae"])
    sae.load_state_dict(torch.load(CACHE / f"sae_{name}.pt"))
    opt = torch.optim.Adam(sae.parameters(), lr=1e-3)
    n = x.shape[0]
    batch = 4096
    for ep in range(extra_epochs):
        perm = torch.randperm(n)
        for i in range(0, n - batch + 1, batch):
            xb = x[perm[i:i + batch]]
            xhat, a = sae(xb)
            loss = ((xhat - xb) ** 2).sum(-1).mean() + r["l1"] * a.sum(-1).mean()
            opt.zero_grad()
            loss.backward()
            opt.step()
            with torch.no_grad():
                sae.W_dec.div_(sae.W_dec.norm(dim=1, keepdim=True))
        with torch.no_grad():
            xhat, a = sae(x[:20000])
            l0 = (a > 0).float().sum(-1).mean().item()
            fvu = (((xhat - x[:20000]) ** 2).sum() /
                   ((x[:20000] - x[:20000].mean(0)) ** 2).sum()).item()
        print(f"  extend {name} epoch {ep+1}: L0={l0:.1f} FVU={fvu:.3f}")
    clean, patched = ce_delta(model, sae, tokens, scale)
    results[name].update({"L0": l0, "FVU": fvu, "ce_clean": clean,
                          "ce_patched": patched, "epochs": 8 + extra_epochs})
    torch.save(sae.state_dict(), CACHE / f"sae_{name}.pt")
    (CACHE / "results.json").write_text(json.dumps(results, indent=2))
    print(f"  CE {clean:.3f} -> {patched:.3f}")


def load_sae(name):
    results = json.loads((CACHE / "results.json").read_text())
    r = results[name]
    sae = SAE(D_MODEL, r["d_sae"])
    sae.load_state_dict(torch.load(CACHE / f"sae_{name}.pt"))
    sae.eval()
    return sae, r


@torch.no_grad()
def latent_acts(sae, train_acts, scale, batch=8192):
    """全訓練トークンに対する潜在変数の活性化(メモリ節約のためスパースに保持)。"""
    x = torch.tensor(train_acts) * scale
    rows, cols, vals = [], [], []
    for i in range(0, x.shape[0], batch):
        a = sae.encode(x[i:i + batch])
        nz = a.nonzero()
        rows.append((nz[:, 0] + i).numpy())
        cols.append(nz[:, 1].numpy())
        vals.append(a[nz[:, 0], nz[:, 1]].numpy())
    return (np.concatenate(rows), np.concatenate(cols), np.concatenate(vals),
            x.shape[0])


def context_str(tok, tokens, seq, pos, before=10, after=4):
    lo, hi = max(0, pos - before), min(SEQ_LEN, pos + after + 1)
    left = tok.decode(tokens[seq, lo:pos])
    mid = tok.decode(tokens[seq, pos:pos + 1])
    right = tok.decode(tokens[seq, pos + 1:hi])
    return f"{left}[[{mid}]]{right}".replace("\n", " ")


def examples(name="sweep_l1_2.0", top_n_latents=24, top_k=8):
    _, tok = load_gpt2()
    tokens, _, train_acts, index, scale = load_acts()
    sae, r = load_sae(name)
    rows, cols, vals, n = latent_acts(sae, train_acts, scale)
    d_sae = r["d_sae"]
    freq = np.bincount(cols, minlength=d_sae) / n
    # 発火頻度が 0.03%〜1% の潜在変数から、平均活性の大きい順に候補を出す
    mean_act = np.zeros(d_sae)
    np.add.at(mean_act, cols, vals)
    cnt = np.bincount(cols, minlength=d_sae)
    mean_act = mean_act / np.maximum(cnt, 1)
    cand = np.where((freq > 3e-4) & (freq < 1e-2))[0]
    cand = cand[np.argsort(-mean_act[cand])][:top_n_latents]
    lines = []
    for j in cand:
        m = cols == j
        order = np.argsort(-vals[m])[:top_k]
        lines.append(f"=== latent {j} freq={freq[j]:.4%} mean_act={mean_act[j]:.2f}")
        for o in order:
            row = rows[m][o]
            seq, pos = index[row]
            lines.append(f"  act={vals[m][o]:.2f} | " +
                         context_str(tok, tokens, seq, pos))
    out = CACHE / f"examples_{name}.txt"
    out.write_text("\n".join(lines), encoding="utf-8")
    print("wrote", out)


def one_token_demo(name="sweep_l1_2.0", seq=None, pos=None, latent=None):
    """記事冒頭用: 1トークンの活性化ベクトルをSAEで分解した実物の値を出す。"""
    _, tok = load_gpt2()
    tokens, acts, _, _, scale = load_acts()
    sae, r = load_sae(name)
    x = torch.tensor(acts[seq, pos]) * scale
    with torch.no_grad():
        a = sae.encode(x.unsqueeze(0))[0]
        xhat = a @ sae.W_dec + sae.b_dec
    nz = (a > 0).nonzero().squeeze(-1)
    print("token:", repr(tok.decode(tokens[seq, pos:pos + 1])))
    print("context:", context_str(tok, tokens, seq, pos))
    print("x[:6] (raw):", np.round(acts[seq, pos, :6], 2))
    print("nonzero latents:", nz.numel(), "/", r["d_sae"])
    top = torch.argsort(-a)[:8]
    for j in top:
        print(f"  latent {j.item():5d}  a={a[j].item():.2f}")
    err = ((xhat - x) ** 2).sum() / (x ** 2).sum()
    print(f"relative recon err: {err.item():.3f}")


# ---------------------------------------------------------------- steering
@torch.no_grad()
def steer(latent=None, name="sweep_l1_2.0", keywords=(),
          prompt="I went outside and looked around.",
          alphas=(0.0, 2.0, 5.0, 10.0), n_gen=20, max_new=40):
    model, tok = load_gpt2()
    tokens, _, train_acts, _, scale = load_acts()
    sae, r = load_sae(name)
    rows, cols, vals, _ = latent_acts(sae, train_acts, scale)
    max_act = vals[cols == latent].max()
    direction = sae.W_dec[latent].detach()  # 単位ノルム
    block = model.transformer.h[LAYER - 1]
    results = {}
    texts_log = []
    for alpha in alphas:
        add = torch.tensor(alpha * max_act / scale) * direction

        def hook(module, inputs, output):
            if isinstance(output, tuple):
                return (output[0] + add,) + output[1:]
            return output + add

        handle = block.register_forward_hook(hook)
        torch.manual_seed(SEED)
        enc = tok(prompt, return_tensors="pt")
        out = model.generate(**enc, do_sample=True, temperature=0.8,
                             top_p=0.95, max_new_tokens=max_new,
                             num_return_sequences=n_gen,
                             pad_token_id=tok.eos_token_id)
        handle.remove()
        gens = [tok.decode(o[enc["input_ids"].shape[1]:]) for o in out]
        # 定量指標1: キーワード出現率
        hit = np.mean([any(k.lower() in g.lower() for k in keywords)
                       for g in gens])
        # 定量指標2: 素のGPT-2で測った生成文の平均NLL(流暢さの劣化)
        nlls = []
        for g in gens:
            ids = tok(prompt + g, return_tensors="pt")["input_ids"]
            logits = model(ids).logits
            nll = nn.functional.cross_entropy(
                logits[0, :-1], ids[0, 1:]).item()
            nlls.append(nll)
        # 定量指標3: 相異なるバイグラムの割合(distinct-2、繰り返し崩壊の検出)
        d2 = []
        for g in gens:
            ids = tok(g)["input_ids"]
            bigrams = list(zip(ids, ids[1:]))
            d2.append(len(set(bigrams)) / max(len(bigrams), 1))
        results[alpha] = {"keyword_rate": float(hit),
                          "mean_nll": float(np.mean(nlls)),
                          "distinct2": float(np.mean(d2))}
        texts_log.append(f"=== alpha={alpha} (add {alpha:.0f} x max_act)")
        texts_log.extend(f"  {g!r}" for g in gens[:6])
        print(alpha, results[alpha])
    (CACHE / "steer_results.json").write_text(json.dumps(
        {"latent": latent, "max_act": float(max_act),
         "keywords": list(keywords), "prompt": prompt,
         "results": {str(k): v for k, v in results.items()}}, indent=2))
    (CACHE / "steer_texts.txt").write_text("\n".join(texts_log),
                                           encoding="utf-8")


# ---------------------------------------------------------------- splitting
@torch.no_grad()
def split(latent=None, small="dict_1024", mids=("sweep_l1_2.0",),
          large="dict_16384", top_k=8):
    """小辞書の潜在変数1つに対して、大きい辞書での対応潜在変数を余弦類似度で探す。"""
    _, tok = load_gpt2()
    tokens, _, train_acts, index, scale = load_acts()
    sae_s, _ = load_sae(small)
    d_small = sae_s.W_dec[latent]
    report = [f"small dict latent {latent} ({small})"]
    sims_all = {}
    for other in list(mids) + [large]:
        sae_o, r = load_sae(other)
        sims = (sae_o.W_dec @ d_small)  # 両方単位ノルムなので内積=余弦
        order = torch.argsort(-sims)[:top_k]
        sims_all[other] = sims.numpy()
        report.append(f"--- vs {other} (d_sae={r['d_sae']})")
        rows, cols, vals, _ = latent_acts(sae_o, train_acts, scale)
        for j in order:
            j = j.item()
            m = cols == j
            report.append(f"  latent {j:5d} cos={sims[j]:.3f} "
                          f"freq={m.mean():.4%}")
            if m.sum() > 0:
                for o in np.argsort(-vals[m])[:4]:
                    seq, pos = index[rows[m][o]]
                    report.append("      " +
                                  context_str(tok, tokens, seq, pos))
    np.savez(CACHE / "split_sims.npz", **sims_all)
    out = CACHE / "split_report.txt"
    out.write_text("\n".join(report), encoding="utf-8")
    print("wrote", out)


# ---------------------------------------------------------------- figures
def figures():
    import matplotlib
    matplotlib.use("Agg")
    import matplotlib.pyplot as plt
    results = json.loads((CACHE / "results.json").read_text())

    # 図1: L0 vs FVU / CE劣化のトレードオフ
    fig, axes = plt.subplots(1, 2, figsize=(10, 4))
    sweep = {k: v for k, v in results.items() if k.startswith("sweep")}
    dicts = {k: v for k, v in results.items() if k.startswith("dict")}
    for ax, key, label in [(axes[0], "FVU", "Fraction of variance unexplained"),
                           (axes[1], None, "CE loss (patched vs clean)")]:
        pass
    ax = axes[0]
    xs = [v["L0"] for v in sweep.values()]
    ys = [v["FVU"] for v in sweep.values()]
    ax.plot(xs, ys, "o-", color="tab:blue")
    for k, v in sweep.items():
        ax.annotate(f"$\\lambda$={v['l1']}", (v["L0"], v["FVU"]),
                    textcoords="offset points", xytext=(8, 4))
    ax.set_xlabel("L0 (avg. active latents per token)")
    ax.set_ylabel("Fraction of variance unexplained")
    ax.set_title("Sparsity vs reconstruction (dict=4096)")
    ax.grid(alpha=0.3)
    ax = axes[1]
    xs = [v["L0"] for v in sweep.values()]
    ys = [v["ce_patched"] - v["ce_clean"] for v in sweep.values()]
    ax.plot(xs, ys, "s-", color="tab:red")
    for k, v in sweep.items():
        ax.annotate(f"$\\lambda$={v['l1']}",
                    (v["L0"], v["ce_patched"] - v["ce_clean"]),
                    textcoords="offset points", xytext=(8, 4))
    ax.set_xlabel("L0 (avg. active latents per token)")
    ax.set_ylabel("CE loss increase (nats)")
    ax.set_title("Sparsity vs LM loss degradation")
    ax.grid(alpha=0.3)
    fig.tight_layout()
    fig.savefig(HERE / "fig1_tradeoff.png", dpi=150)

    # 図2: feature splitting(小辞書潜在変数と各辞書の余弦類似度上位)
    data = np.load(CACHE / "split_sims.npz")
    fig, ax = plt.subplots(figsize=(8, 4))
    for name_, marker in zip(data.files, ["o", "s", "^"]):
        sims = np.sort(data[name_])[::-1][:12]
        d = json.loads((CACHE / "results.json").read_text())[name_]["d_sae"]
        ax.plot(range(1, 13), sims, marker, ls="-",
                label=f"dict size {d}")
    ax.set_xlabel("Rank of matching latent (by cosine similarity)")
    ax.set_ylabel("Cosine similarity of decoder directions")
    ax.set_title("One small-dict latent vs larger dictionaries")
    ax.legend()
    ax.grid(alpha=0.3)
    fig.tight_layout()
    fig.savefig(HERE / "fig2_splitting.png", dpi=150)

    # 図3: steering
    st = json.loads((CACHE / "steer_results.json").read_text())
    alphas = sorted(float(a) for a in st["results"])
    kw = [st["results"][str(a)]["keyword_rate"] for a in alphas]
    d2 = [st["results"][str(a)]["distinct2"] for a in alphas]
    fig, ax1 = plt.subplots(figsize=(7, 4))
    ax1.plot(alphas, kw, "o-", color="tab:blue")
    ax1.set_xlabel("Steering coefficient (multiples of max activation)")
    ax1.set_ylabel("Generations containing topic keywords", color="tab:blue")
    ax1.set_ylim(-0.05, 1.05)
    ax2 = ax1.twinx()
    ax2.plot(alphas, d2, "s--", color="tab:red")
    ax2.set_ylabel("Distinct-2 (bigram diversity)", color="tab:red")
    ax2.set_ylim(-0.05, 1.05)
    ax1.set_title(f"Steering latent {st['latent']}")
    ax1.grid(alpha=0.3)
    fig.tight_layout()
    fig.savefig(HERE / "fig3_steering.png", dpi=150)
    print("figures saved")


if __name__ == "__main__":
    p = argparse.ArgumentParser()
    p.add_argument("stage")
    p.add_argument("--latent", type=int, default=None)
    p.add_argument("--seq", type=int, default=None)
    p.add_argument("--pos", type=int, default=None)
    p.add_argument("--keywords", type=str, default="")
    args = p.parse_args()
    if args.stage == "collect":
        collect()
    elif args.stage == "train":
        train_all()
    elif args.stage == "extend":
        extend("dict_16384")
    elif args.stage == "examples":
        examples()
    elif args.stage == "demo":
        one_token_demo(seq=args.seq, pos=args.pos)
    elif args.stage == "steer":
        steer(latent=args.latent,
              keywords=tuple(k for k in args.keywords.split(",") if k))
    elif args.stage == "split":
        split(latent=args.latent)
    elif args.stage == "figures":
        figures()

関連記事

← 技術ブログ一覧へ