埋め込み表からロジットまで
目標
トークン番号がベクトルになり、そのベクトルが再び語彙全体のスコアになるまでを、標準ライブラリだけで作ります。ルックアップが、ワンホットベクトルと行列の積と同じ値を出すことを、2つの方法で計算して確認し、パラメーターの数を数え、同じ表を出力側でもう一度使う重み共有を作り、コサインと内積が分かれる場面を数字で見ます。最後に、埋め込みに√dを掛けると、大きさだけが変わり、方向はそのままであることを測ります。
なぜ重要なのか
モデルについて話すとき、人が最もよく間違える箇所が、ここです。「埋め込みはルックアップで、出力は行列の積」と分けて考えると、2つが同じ表であることが見えず、パラメーターの数がなぜずれるのかも説明できません。語彙を大きくしようという提案が、なぜメモリの会議で終わるのかも、V掛けるdを数えてみる前には、見当がつきません。
このラボは、実際のモデルを呼び出しません。このPodのシステムのPythonには、numpy・torch・transformersがありません(numpyは/opt/onnx-lab/bin/pythonの中にしかありません)。そのため、実際のモデルの語彙サイズやパラメーターの数のような数字は、ここでは使いません。自分が作った表から測った値だけを使います。
判定は、実数の比較が多くなります。採点ツールは、abs(a-b) <= atol + rtol*abs(b)で見ますが、形・パラメーターの数・近傍のリストのような整数の判定も一緒に見ます。そのため、足す順序が違って最後の桁が揺らぐのは受け入れ、実際に間違った実装は弾きます。
採点ツールは、書いておいた説明を信用しません。自分で作ったモジュールを実際に呼び出し、毎回異なる表と異なる番号で関数を直接叩いて、採点ツールが別に計算した値と照合します。
ステップ
- /root/work/tf-embed/embed.pyに
VOCAB = 512・DIM = 64・SEED = 20260917と、make_table(vocab, dim, seed)・shape(matrix)・param_count(vocab, dim)を作成してください。同じシードなら、いつも同じ表が出る必要があります。 lookup(table, ids)を追加して、番号のリストをベクトルのリストに変えるようにしてください。同じ番号は同じベクトルを返し、語彙の外の番号はIndexErrorです。one_hot(token_id, vocab)・row_times_matrix(vec, matrix)・lookup_via_one_hot(table, ids)・one_hot_mults(vocab, dim, count)を追加して、ルックアップとワンホットの積が同じ値を出すかを見てください。logits(table, hidden)を作成して、入力に使ったその表で語彙全体のスコアを出すようにしてください。新しい行列を作ってはいけません。tied_params(vocab, dim)・untied_params(vocab, dim)・vocab_growth(vocab, dim, factor)で、共有したときと別々に置いたときの数の個数を数えるようにしてください。dot・norm・cosine・stretch・nearest_by_cosine・nearest_by_dotを作成して、1行だけを引き伸ばしたときに、2つのリストがどう分かれるかを見てください。rms(vec)とscaled_lookup(table, ids, dim)を作成して、√dを掛けると大きさが√d倍になることを測るようにしてください。- 決めた定数ですべて測って、/root/work/tf-embed/embed_report.jsonと、/root/work/tf-embed/embed_report.mdに記録してください。
参考
- 実行の契約: 採点ツールは、
/root/work/tf-embed/embed.pyをPythonモジュールとして読み込み、VOCAB・DIM・SEED・make_table・shape・param_count・lookup・one_hot・row_times_matrix・lookup_via_one_hot・one_hot_mults・logits・tied_params・untied_params・vocab_growth・dot・norm・cosine・stretch・nearest_by_cosine・nearest_by_dot・rms・scaled_lookupを直接使います。スクリプトとして実行しないので、if __name__ == "__main__"はなくてかまいません。 make_table(vocab, dim, seed)は、random.Random(seed)を1つ作って、トークン0の0番目の要素から、行単位で埋めます。要素ごとにrandom()を1回引いて、0.5を引きます。そうすれば、採点ツールが同じ表を別に作って、値を照合できます。shape(matrix)は、(줄 수, 칸 수)(プレースホルダーは行数と列数です)を返し、行ごとの列数が異なればValueErrorを出します。空の表は(0, 0)です。lookup(table, ids)は、語彙の外の番号にIndexErrorを出します。負の数も語彙の外です。Pythonの負のインデックスをそのままにすると、後ろから数えて、静かに見当違いの行を返します。row_times_matrix(vec, matrix)は、長さVの行ベクトルとV×dの表から、長さdのベクトルを出します。out[c] = sum(vec[r] * matrix[r][c] for r in range(V))です。one_hot_mults(vocab, dim, count)は、乗算の回数です。トークン1つにつきV掛けるd回で、ルックアップなら0回です。時間を測らず、この数を数えてください。logits(table, hidden)は、長さVのリストです。out[t] = sum(table[t][c] * hidden[c] for c in range(d))で、隠れ状態のベクトルの長さが表の幅と違えば、ValueErrorです。表をそのまま使うのが、重み共有です。vocab_growth(vocab, dim, factor)は、vocab・bigger_vocab・dim・tied・bigger_tied・untied・bigger_untied・savedのキーを持つ辞書です。savedは、別々に置くときから、共有するときを引いた値です。cosine(a, b)は、長さが0のベクトルが来たら0.0です。rms([])も0.0です。stretch(table, token_id, factor)は、その行だけをfactor倍に引き伸ばした新しい表を返します。受け取った表をその場で書き換えないでください。nearest_by_cosine・nearest_by_dotは、スコアの大きい順にk個の番号を返します。自分自身は除き、スコアが同じなら、番号が小さいほうが先です。scaled_lookup(table, ids, dim)は、lookupの結果にmath.sqrt(dim)を掛けたものです。- ステップ8は、
VOCAB = 512・DIM = 64・SEED = 20260917で表を作り、トークン137を調べます。近傍はk = 5、引き伸ばす行は、コサイン近傍の5番目(0から数えて4番目の要素)、引き伸ばす倍数は4.0です。語彙を大きくする倍数は2です。 - このPodにはインターネットがありません。
pip installはできず、システムのPythonでimport numpyもできません。mathとrandomで十分です。 - 公式ドキュメント: Attention Is All You Need・PyTorch — MultiheadAttention・NumPy — matmul・Python — math
- よくある間違い: 表を要素単位で埋めて、シードの順序がずれる、負の番号をそのまま通す、ワンホットの積で行と列を逆にする、ロジットに新しい行列を作って使う、別々に置くときの数を2倍で数えない、コサインで長さで割らない、
stretchが元の表を書き換える、√dの代わりにdを掛ける。
表の形とパラメーターの数を数える
/root/work/tf-embed/embed.pyにVOCAB = 512・DIM = 64・SEED = 20260917と、make_table(vocab, dim, seed)・shape(matrix)・param_count(vocab, dim)を作成してください。make_tableは、random.Random(seed)を1つ使って、トークン0の0番目の要素から行単位で埋め、要素ごとにrandom()から0.5を引きます。
表は、リストのリストです。rng = random.Random(seed)を1回だけ作り、行ごとにdim個ずつ引けば、行単位の順序が守られます。要素単位で回すと、同じシードなのに別の表が出ます。shapeは、行ごとに列数を確認して、違えばValueErrorを出してください。形がずれた表は、後ろで静かにおかしな値を出します。param_countは、積であって和ではありません。
番号で行を取り出す
lookup(table, ids)を追加してください。番号のリストをベクトルのリストに変えます。同じ番号が2回出たら、まったく同じベクトルが2回出る必要があり、語彙の外の番号(負の数を含む)はIndexErrorです。
行を取り出すのがすべてです。ただし、table[-1]は、Pythonでは後ろから1番目の行を静かに返すので、0 <= token_id < len(table)を自分で確認する必要があります。取り出した行は、list()でコピーして返せば、呼び出した側が元の表を触る危険がありません。同じ番号が同じベクトルになるのは、バグではなく性質です。埋め込みにはコンテキストがなく、コンテキストはアテンションがあとで与えます。
ルックアップはワンホットの積と同じ値になる
one_hot(token_id, vocab)・row_times_matrix(vec, matrix)・lookup_via_one_hot(table, ids)・one_hot_mults(vocab, dim, count)を追加してください。ワンホットの積で計算した値がlookupと同じである必要があり、one_hot_multsは、トークン1つにつきV掛けるd回の乗算の回数を返します。
row_times_matrixの要素1つは、sum(vec[r] * matrix[r][c] for r in range(V))です。行と列を逆にすると、長さからずれるので、shapeで確認してから回してください。ワンホットは、1か所だけが1.0なので、掛けて足すと、その行だけが生き残ります。値がちょうど同じなら正常です。時間を測ろうとせず、乗算の回数を数えてください。ルックアップは0回です。
同じ表でスコアを出す
logits(table, hidden)を作成してください。隠れ状態のベクトル1つを、語彙全体のスコアに変えます。入力に使ったその表をそのまま使う必要があります(重み共有)。隠れ状態のベクトルの長さが表の幅と違えば、ValueErrorです。
行ごとに隠れ状態のベクトルとの内積を出せば、長さVのリストになります。新しい行列を作って使うと、値がまったく変わります。共有するというのは、まさにその表を再び使うという意味です。確認するよい方法が1つあります。hiddenに、あるトークンの埋め込みをそのまま入れてみてください。自分自身との内積は長さの2乗なので、そのトークンのスコアが最も大きくなります。
共有するとどれだけ減るかを数える
tied_params(vocab, dim)・untied_params(vocab, dim)・vocab_growth(vocab, dim, factor)を作成してください。vocab_growthは、vocab・bigger_vocab・dim・tied・bigger_tied・untied・bigger_untied・savedのキーを持つ辞書で、savedは、別々に置くときから、共有するときを引いた値です。
共有すると、表が1式、別々に置くと、同じ形が2式です。語彙をfactor倍に大きくするとき、幅には手を付けないのに、2つの値がどちらもその倍数だけ大きくなります。それが、語彙サイズの値です。ここはすべて整数の判定なので、実数で計算して丸めないでください。
長さが順序を変えることを確認する
dot(a, b)・norm(vec)・cosine(a, b)・stretch(table, token_id, factor)・nearest_by_cosine(table, token_id, k)・nearest_by_dot(table, token_id, k)を作成してください。近傍は、スコアの大きい順にk個の番号で、自分自身は除きます。スコアが同じなら、番号が小さいほうが先です。
cosineは、内積を2つの長さの積で割ったもので、長さが0なら、方向がないので0.0です。stretchは、新しい表を作る必要があります。元の表をその場で書き換えると、後ろの判定がすべてずれます。並べ替えは、key=lambda item: (-점수, 번호)(プレースホルダーはスコアと番号です)の1行で、同点のルールまで書けます。1行を引き伸ばしておいて、2つのリストを並べて見てください。コサインのリストはそのままなのに、内積のリストでは、引き伸ばした行が前に出てきます。
√dを掛けると大きさだけが変わる
rms(vec)とscaled_lookup(table, ids, dim)を作成してください。rmsは二乗平均平方根で、空のベクトルは0.0です。scaled_lookupは、lookupの結果にmath.sqrt(dim)を掛けたものです。
すべての要素に同じ数を掛けるので、方向はまったく変わりません。コサインを測ってみると、掛ける前と同じです。変わるのは大きさだけで、ちょうどmath.sqrt(dim)倍です。dimをそのまま掛けると、大きさがd倍になって、まったく別の値になります。rmsは、内積を要素数で割ってから、平方根をとれば済みます。
測った結果を記録に残す
VOCAB = 512・DIM = 64・SEED = 20260917で表を作り、トークン137を調べて、すべて測ってください。近傍はk = 5、引き伸ばす行は、コサイン近傍の5番目、倍数は4.0、語彙を大きくする倍数は2です。/root/work/tf-embed/embed_report.jsonにはvocab・dim・seed・probe・table_shape・params・tied_params・untied_params・saved・bigger_vocab・bigger_tied・bigger_untied・lookup_mults・one_hot_mults・max_abs_diff・repeat_same・logit_len・logit_argmax・logit_argmax_is_self・cos_neighbors・dot_neighbors・neighbors_differ・stretch_target・stretch_factor・cos_after_stretch・dot_after_stretch・cos_unchanged_by_scale・rms_plain・rms_scaled・rms_ratioを、/root/work/tf-embed/embed_report.mdには## 무엇을 쟀나 ## 조회와 원-핫 곱은 같은 연산이다 ## 가중치를 묶으면 무엇이 줄어드나 ## 내적과 코사인이 갈리는 자리 ## √d 를 곱하면 무엇이 달라지나の5つの節で書いてください(韓国語の見出しは順に、「何を測ったか」「ルックアップとワンホットの積は同じ演算である」「重みを共有すると何が減るのか」「内積とコサインが分かれる場面」「√dを掛けると何が変わるのか」という意味です)。
数字は手で書かず、自分のコードを実際に動かして得た値で埋めてください。one_hot_multsは、トークン137を1つルックアップするときの値で、lookup_multsは0です。max_abs_diffは、lookupとlookup_via_one_hotの値の差のうち、最も大きい絶対値です。repeat_sameは、同じ番号を2回ルックアップしたときに、2つのベクトルが同じかどうかです。logit_argmaxは、隠れ状態のベクトルにトークン137の埋め込みをそのまま入れたときに、スコアが最も大きい番号です。cos_unchanged_by_scaleは、√dを掛ける前と後のコサインが同じかどうかです。rms_ratioは、掛けたあとの大きさを、掛ける前の大きさで割った値です。