别再死记硬背了!用Karpathy的nanoGPT项目,手把手带你跑通LLM的Next Token Prediction
·
从零实现Next Token Prediction:用nanoGPT解锁语言模型核心机制
当你第一次打开nanoGPT的代码仓库,那些看似简单的PyTorch模块背后,隐藏着现代大语言模型最精妙的设计思想。Next Token Prediction不仅是GPT系列模型的训练目标,更是理解自回归生成本质的钥匙。本文将带你穿透理论描述,直接解剖代码层面的实现细节。
1. 数据构造:移位操作的魔法
原始文本经过tokenizer处理后,会变成一串数字序列。假设我们有以下token序列:
original_tokens = [9, 3, 6, 4, 2, 1, 5] # 对应"你 知道 什么 是 预训练 吗 ?"
传统做法是为每个位置单独创建训练样本,比如:
- 输入[9],目标3
- 输入[9,3],目标6
- 输入[9,3,6],目标4 ...
这种方法效率低下且无法利用GPU并行优势。nanoGPT采用了一种巧妙的移位构造法:
def prepare_batch(tokens, max_len=10):
# 填充序列到固定长度
padded = tokens + [0]*(max_len - len(tokens))
# 输入是原始序列
x = torch.tensor(padded[:-1]).unsqueeze(0) # (1, seq_len-1)
# 目标是右移一位的序列
y = torch.tensor(padded[1:]).unsqueeze(0) # (1, seq_len-1)
return x, y
关键点在于:
- 因果掩码 :确保每个位置只能看到前面的token
- 并行计算 :单次前向传播完成所有位置的预测
- 填充处理 :用特殊token(如0)填充不足部分,并在loss计算时忽略
注意:实际代码中还需要处理batch维度和注意力掩码,这里为清晰起见做了简化
2. 模型架构:Transformer的极简实现
nanoGPT的核心组件可以用以下类结构表示:
class GPT(nn.Module):
def __init__(self, vocab_size, n_embd, n_head, n_layer):
super().__init__()
self.token_embed = nn.Embedding(vocab_size, n_embd)
self.pos_embed = nn.Embedding(block_size, n_embd)
self.blocks = nn.ModuleList([Block(n_embd, n_head) for _ in range(n_layer)])
self.ln_f = nn.LayerNorm(n_embd)
self.head = nn.Linear(n_embd, vocab_size, bias=False)
def forward(self, x):
B, T = x.shape
tok_emb = self.token_embed(x) # (B,T,C)
pos_emb = self.pos_embed(torch.arange(T)) # (T,C)
x = tok_emb + pos_emb # (B,T,C)
for block in self.blocks:
x = block(x)
x = self.ln_f(x)
logits = self.head(x) # (B,T,vocab_size)
return logits
几个关键设计选择:
- 权重绑定 :
self.head.weight = self.token_embed.weight提升参数效率 - 层归一化 :每个Block后都有LayerNorm稳定训练
- 残差连接 :在Block内部实现,避免梯度消失
3. 训练过程:从logits到loss的完整路径
理解loss计算是掌握Next Token Prediction的关键。假设我们有以下数据:
| 输入x | 目标y |
|---|---|
| 9 | 3 |
| 3 | 6 |
| 6 | 4 |
| ... | ... |
模型前向传播后得到logits的形状为(batch_size, seq_len, vocab_size)。计算loss时需要:
- 展平logits和目标序列
- 应用交叉熵损失
- 忽略padding位置的loss
def compute_loss(logits, targets, ignore_index=-100):
B, T, C = logits.shape
logits = logits.view(B*T, C)
targets = targets.view(B*T)
loss = F.cross_entropy(logits, targets, ignore_index=ignore_index)
return loss
实际训练中还涉及:
- 学习率调度 :cosine衰减等策略
- 梯度裁剪 :防止梯度爆炸
- 混合精度 :加速训练
4. 推理实现:自回归生成的艺术
与训练时不同,推理阶段需要逐个token生成:
def generate(model, start_tokens, max_new_tokens, temperature=1.0):
tokens = start_tokens.copy()
for _ in range(max_new_tokens):
# 截取最后block_size个token
inputs = tokens[-block_size:]
# 获取预测logits
logits = model(inputs.unsqueeze(0)) # (1,T,vocab_size)
# 取最后一个位置的logits
logits = logits[0, -1, :] / temperature
# 转换为概率分布
probs = F.softmax(logits, dim=-1)
# 采样下一个token
next_token = torch.multinomial(probs, num_samples=1)
tokens.append(next_token.item())
return tokens
采样策略对比:
| 策略 | 温度 | 特点 |
|---|---|---|
| 贪婪搜索 | 0.0 | 确定性高但缺乏多样性 |
| 随机采样 | 1.0 | 平衡创造性和连贯性 |
| Top-k采样 | 0.7 | 限制候选集提高质量 |
| 核采样 | 0.9 | 动态调整候选集大小 |
5. SFT微调:专注目标输出的技巧
有监督微调的关键在于构造合适的loss mask:
def prepare_sft_batch(prompt, answer, bos_id, eos_id):
input_ids = prompt + [bos_id] + answer + [eos_id]
context_len = len(prompt) + 1 # +1 for bos
# 创建loss mask: answer部分为1,其余为0
loss_mask = [0]*context_len + [1]*(len(input_ids)-context_len)
return input_ids, loss_mask
训练时加权计算loss:
loss = F.cross_entropy(logits.view(-1, C), targets.view(-1), reduction='none')
loss = (loss * loss_mask.view(-1)).sum() / loss_mask.sum()
这种技术可以:
- 防止模型过度关注无关的prompt部分
- 更高效地利用监督信号
- 保持预训练获得的基础能力
6. 调试技巧:常见问题与解决方案
在实现过程中可能会遇到:
问题1:loss不下降
- 检查数据构造是否正确
- 验证模型容量是否足够
- 调整学习率和batch size
问题2:生成结果重复
- 尝试降低temperature
- 实现repetition penalty
- 检查训练数据多样性
问题3:GPU内存不足
- 减小batch size
- 使用梯度累积
- 启用激活检查点
一个实用的调试检查表:
- [ ] 数据加载器输出符合预期
- [ ] 模型参数量级合理
- [ ] 初始loss接近-ln(1/vocab_size)
- [ ] 训练初期能过拟合小批量数据
- [ ] 验证集loss正常下降
在nanoGPT的代码实践中,最让我惊讶的是通过如此简洁的实现就能捕捉语言的核心模式。当第一次看到模型生成连贯的文本时,那些看似复杂的理论概念突然变得具象而清晰。
更多推荐


所有评论(0)