ゼロから独自のLLMを訓練する:完全なエンドツーエンドガイド
PyTorchを使用してTransformer LLMをゼロから構築、事前訓練、調整する実践的なチュートリアル。SFT、DPO、PPO、GRPOをカバー。
ゼロから独自の大規模言語モデル(LLM)を訓練したいと思ったことはありませんか?単に既存のモデルを微調整するのではなく。FareedKhan-devによるオープンソースプロジェクトtrain-llm-from-scratchは、生のテキストから調整済みで推論可能なモデルに至るまでの完全で透明なパイプラインを提供します。これは既存のライブラリのブラックボックスラッパーではなく、マルチヘッドアテンションからGRPOまでのすべてのアルゴリズムがプレーンなPyTorchで実装された実践的な教育の旅です。
このガイドでは、プロセス全体を順を追って説明し、各ステップの理由と方法を解説するので、基本を理解し、自分でコードを実行できるようになります。
なぜゼロからLLMを構築するのか?
ほとんどの開発者は、Hugging Faceのtransformersやtrlのような高レベルAPIを通じてLLMと対話します。これらは強力ですが、コアのメカニズムを抽象化しています。ゼロから構築することで得られるもの:
- 深い理解: トークン化、アテンション、トレーニングループがどのように機能するかを正確に学びます。
- 完全な制御: ライブラリの制約なしに、損失関数、アーキテクチャ、トレーニング戦略など、任意のコンポーネントを変更できます。
- 教育的価値: 因果マスキング、報酬モデリング、方策勾配法などの概念を把握する最良の方法です。
このプロジェクトは、データのダウンロードからテキスト生成までのパイプライン全体を一箇所で見たい学生、開発者、研究者向けに設計されています。
パイプラインの概要
旅は明確で順序立てた経路をたどります:
- データ準備: 生テキスト → トークンID → 保存された配列。
- モデル構築: 小さな再利用可能なコンポーネント(MLP、アテンション、ブロック)からTransformerを構築。
- 事前訓練: 大規模コーパス(The Pile)での次トークン予測。
- テキスト生成: 訓練済みモデルからのサンプリング。
- 事後訓練(調整): SFT、報酬モデリング、DPO、PPO、GRPOを介してベースモデルを有用なアシスタントに変換。
- 評価: GSM8Kのようなベンチマークでのパフォーマンス測定。
- 推論とチャット: 最終モデルとの対話。
ステップ1:データの準備
モデルは整数のみを理解します。最初の仕事は、テキストをトークンIDに変換し、効率的に保存することです。
トークン化
このプロジェクトは、OpenAIのr50k_baseトークナイザー(tiktokenから)を使用します。これはGPT-3で使用されているものと同じです。テキストは整数のリストになり、各ドキュメントの末尾に特別な<|endoftext|>トークン(ID 50256)が追加され、モデルが境界を学習できるようにします。
# The Pileからデータをダウンロードしてトークン化
python scripts/data_download.py
python scripts/data_preprocess.py
より新しく高速なパス:
python scripts/prepare_pretrain_data.py --split val --out data/pile_dev.h5
python scripts/prepare_pretrain_data.py --split train --num_shards 1 --out data/pile_train.h5
チャット形式と損失マスキング
事後訓練では、モデルは役割(ユーザー対アシスタント)を理解する必要があります。プロジェクトはプレーンテキストマーカーを使用します:
<|user|>
13 + 29は何ですか?<|endoftext|><|assistant|>
<think>13 + 29 = 42</think><answer>42</answer><|endoftext|>
損失マスクにより、モデルはユーザーのプロンプトではなく、アシスタントのトークンからのみ学習します。これはSFTとRLにとって重要です。
def encode_chat(messages, add_generation_prompt=False):
ids, mask = [], []
for m in messages:
role = m["role"]
header_ids = _encode_ordinary(_header_for(role))
ids.extend(header_ids)
mask.extend([0] * len(header_ids))
content_ids = _encode_ordinary(m["content"])
is_completion = role == "assistant"
ids.extend(content_ids)
mask.extend([1 if is_completion else 0] * len(content_ids))
ids.append(EOT_ID)
mask.append(1 if is_completion else 0)
return ids, mask
ステップ2:Transformerの構築
Transformerは4つの小さな再利用可能なPyTorchモジュールから構築されています。このモジュール設計により、コードは理解しやすく、変更しやすくなっています。
多層パーセプトロン(MLP)
各ブロックの「思考」部分。トークンベクトルを拡張し、ReLUを適用し、元の次元に投影します。
class MLP(nn.Module):
def __init__(self, n_embed):
super().__init__()
self.hidden = nn.Linear(n_embed, 4 * n_embed)
self.relu = nn.ReLU()
self.proj = nn.Linear(4 * n_embed, n_embed)
def forward(self, x):
x = self.relu(self.hidden(x))
x = self.proj(x)
return x
単一ヘッドアテンション
トークンが他のトークンを参照できるようにします。クエリ、キー、バリューベクトルを計算し、因果マスクを適用し(トークンが将来のトークンを見られないように)、値の加重和を取ります。
class Head(nn.Module):
def __init__(self, head_size, n_embed, context_length):
super().__init__()
self.key = nn.Linear(n_embed, head_size, bias=False)
self.query = nn.Linear(n_embed, head_size, bias=False)
self.value = nn.Linear(n_embed, head_size, bias=False)
self.register_buffer('tril', torch.tril(torch.ones(context_length, context_length)))
def forward(self, x):
B, T, C = x.shape
k = self.key(x)
q = self.query(x)
scale_factor = 1 / math.sqrt(C)
attn_weights = q @ k.transpose(-2, -1) * scale_factor
attn_weights = attn_weights.masked_fill(self.tril[:T, :T] == 0, float('-inf'))
attn_weights = F.softmax(attn_weights, dim=-1)
v = self.value(x)
out = attn_weights @ v
return out
マルチヘッドアテンション
複数のアテンションヘッドを並行して実行し、それぞれが異なる関係を学習し、結果を連結して投影します。
Transformerブロック
アテンションとMLPを事前正規化残差接続で組み合わせます。これがコアの繰り返し単位です。
class Block(nn.Module):
def __init__(self, n_head, n_embed, context_length):
super().__init__()
self.ln1 = nn.LayerNorm(n_embed)
self.attn = MultiHeadAttention(n_head, n_embed, context_length)
self.ln2 = nn.LayerNorm(n_embed)
self.mlp = MLP(n_embed)
def forward(self, x):
x = x + self.attn(self.ln1(x))
x = x + self.mlp(self.ln2(x))
return x
完全なTransformer
すべてをラップ:トークン埋め込み、位置埋め込み、ブロックのスタック、最終層正規化、語彙サイズのロジットへの線形投影。
class Transformer(nn.Module):
def __init__(self, n_head, n_embed, context_length, vocab_size, N_BLOCKS):
super().__init__()
self.token_embed = nn.Embedding(vocab_size, n_embed)
self.position_embed = nn.Embedding(context_length, n_embed)
self.attn_blocks = nn.ModuleList([Block(n_head, n_embed, context_length) for _ in range(N_BLOCKS)])
self.layer_norm = nn.LayerNorm(n_embed)
self.lm_head = nn.Linear(n_embed, vocab_size)
self.register_buffer('pos_idxs', torch.arange(context_length))
def forward(self, idx, targets=None):
x = self.forward_hidden(idx)
logits = self.lm_head(x)
loss = None
if targets is not None:
B, T, C = logits.shape
flat_logits = logits.reshape(B * T, C)
targets = targets.reshape(B * T).long()
loss = F.cross_entropy(flat_logits, targets)
return logits, loss
ステップ3:ベースモデルの事前訓練
事前訓練は、大規模コーパスでの次トークン予測です。モデルはランダムなトークンウィンドウを読み取り、すべての位置で次のトークンを予測し、クロスエントロピー損失を介して重みを更新します。
# 13Mパラメータモデルの訓練
python scripts/train_transformer.py
# メモリ最適化を伴う大規模モデルの場合
python scripts/train_transformer.py --amp --grad-checkpointing --grad-accum 8
損失はln(vocab_size)(約10.8)付近から始まり、モデルが言語統計を学習するにつれて低下します。2x L40 GPUで訓練された77Mパラメータモデルは、2000ステップ後に約3.73の損失に達しました。
ステップ4:テキスト生成
訓練が完了したら、モデルの出力分布からサンプリングしてテキストを生成できます。
def generate(self, idx, max_new_tokens):
for _ in range(max_new_tokens):
idx_cond = idx[:, -self.context_length:]
logits, _ = self(idx_cond)
logits = logits[:, -1, :]
probs = F.softmax(logits, dim=-1)
idx_next = torch.multinomial(probs, num_samples=1)
idx = torch.cat((idx, idx_next), dim=1)
return idx
python scripts/generate_text.py --model_path models/transformer_B.pt --input_text "The" --max_new_tokens 100
ステップ5:事後訓練(調整)
ベースモデルはテキストを続けることはできますが、指示に従うことはできません。事後訓練は、教師あり微調整(SFT)と強化学習(RL)を使用してモデルを調整します。
SFT(教師あり微調整)
モデルにチャット形式で回答する方法を教えます。損失はアシスタントのトークンに対してのみ計算されます。
python scripts/prepare_sft_data.py --context_length 1024
torchrun --standalone --nproc_per_node=2 scripts/train_sft.py
報酬モデル
SFTモデルの上に小さな線形ヘッドを訓練して回答をスコアリングします。Bradley-Terry損失を使用して、選択された応答を拒否された応答よりも優先します。
def bradley_terry_loss(chosen_rewards, rejected_rewards):
return -F.logsigmoid(chosen_rewards - rejected_rewards).mean()
DPO、ORPO、KTO
直接選好最適化(DPO)は報酬モデルとRLループをスキップし、選好ペアに直接作用します。ORPOとKTOはそのバリアントです。
torchrun --standalone --nproc_per_node=2 scripts/train_dpo.py --loss_type dpo
PPO(近位方策最適化)
古典的なRLHFループ。モデルが回答を生成し、報酬モデルでスコアリングし、クリップされた方策損失を介して更新します。
def ppo_policy_loss(new_logp, old_logp, advantages, mask, clip=0.2):
ratio = torch.exp(new_logp - old_logp)
surr1 = ratio * advantages
surr2 = torch.clamp(ratio, 1.0 - clip, 1.0 + clip) * advantages
loss = -masked_mean(torch.min(surr1, surr2), mask)
return loss
GRPO(グループ相対方策最適化)
DeepSeek-R1スタイルの方法。各プロンプトに対して回答のグループをサンプリングし、スコアリングし、グループの平均と標準偏差をベースラインとして使用します。
def group_advantages(rewards, group_size, eps=1e-4):
r = rewards.view(-1, group_size)
mean = r.mean(dim=1, keepdim=True)
std = r.std(dim=1, keepdim=True)
adv = (r - mean) / (std + eps)
return adv.reshape(-1)
ステップ6:評価
主要な指標は貪欲GSM8K精度です。モデルは<answer>タグ内に正しい答えを出力する必要があります。
for s in base_pretrained sft dpo ppo grpo; do
python scripts/eval_post_training.py --ckpt models/$s.pt --label $s --limit 200 --append logs/table.jsonl
done
ステップ7:モデルとの対話
チャットスクリプトを使用して任意のチェックポイントと対話します。
python scripts/chat.py --ckpt models/sft.pt --prompt "13 + 29は何ですか?"
始め方
リポジトリをクローン:
git clone https://github.com/FareedKhan-dev/train-llm-from-scratch.git cd train-llm-from-scratch pip install -e .小さく始める: まず13Mパラメータモデルを訓練します。高速で、パイプラインの感触をつかめます。
スケールアップ: GPUのメモリ制限に達するまで
n_embedとn_blocksを増やします。事後訓練チェーンを実行: SFT、次にDPO、次にPPOまたはGRPOを実行し、GSM8K精度が向上するのを確認します。
このプロジェクトは珍しい宝石です:実用的なツールであり、包括的なチュートリアルでもあります。Transformerを深く理解したい学生でも、調整技術を実験したい開発者でも、train-llm-from-scratchは完璧な出発点です。
ソース
FareedKhan-dev/train-llm-from-scratch: データのダウンロードからテキスト生成まで、LLMを訓練するための簡単な方法。