从Bigram语言模型入门LLM:Andrej Karpathy经典实现解析
这次我们来深入解析 Andrej Karpathy 的 Bigram 语言模型,这是一个非常适合入门自然语言处理的经典项目。作为 OpenAI 创始成员和前特斯拉 AI 总监,Karpathy 设计的这个模型虽然结构简单,但完整展示了语言模型的核心原理,特别适合想要从零理解 LLM(大语言模型)工作原理的开发者。
Bigram 模型最大的特点是实现简洁、训练快速、资源要求极低。你不需要高端显卡,甚至用 CPU 就能在几分钟内完成训练和推理。本文将带你完整实现一个 Bigram 语言模型,并验证其文本生成能力。
1. 核心能力速览
| 能力项 | 具体说明 |
|---|---|
| 模型类型 | 基于字符的 Bigram 统计语言模型 |
| 开源作者 | Andrej Karpathy(OpenAI 创始成员) |
| 核心功能 | 字符级文本生成、概率统计、训练可视化 |
| 硬件要求 | 极低,CPU 即可运行,无需 GPU |
| 显存占用 | 几乎可忽略,模型参数极少 |
| 依赖环境 | Python 3.6+、PyTorch、NumPy |
| 代码规模 | 单文件,100 行左右核心代码 |
| 适合场景 | LLM 入门教学、语言模型原理理解、基础文本生成实验 |
2. 适用场景与使用边界
Bigram 模型最适合以下场景:
教育学习用途 :
- 理解语言模型的基本构建流程:数据准备、模型定义、训练循环、推理生成
- 掌握 PyTorch 张量操作和自动梯度计算
- 学习如何评估文本生成质量
实验验证用途 :
- 快速验证文本生成想法
- 测试不同训练数据对模型效果的影响
- 作为更复杂模型(如 GPT、LSTM)的对比基线
使用边界提醒 :
- 生成文本长度有限,通常适合短文本生成
- 无法处理长距离依赖关系
- 生成内容可能存在重复或不连贯现象
- 不适合生产环境部署,主要用于教学演示
3. 环境准备与前置条件
3.1 基础软件环境
# 检查 Python 版本
python --version
# 推荐 Python 3.8+
# 安装核心依赖
pip install torch numpy matplotlib
3.2 验证 PyTorch 安装
import torch
import numpy as np
print(f"PyTorch 版本: {torch.__version__}")
print(f"CUDA 是否可用: {torch.cuda.is_available()}")
3.3 准备训练数据
Bigram 模型对数据要求很灵活,可以使用任何文本文件:
- 英文小说文本(如莎士比亚作品)
- 中文古诗集(需调整分词方式)
- 代码文件(学习编程语言模式)
- 自定义文本语料
4. Bigram 模型原理与实现
4.1 Bigram 基本概念
Bigram(二元语法)模型基于一个简单的假设:每个字符的出现概率只依赖于前一个字符。这种马尔可夫假设大大简化了模型复杂度。
数学上,Bigram 概率可以表示为:
P(当前字符 | 前一个字符) = count(前一个字符, 当前字符) / count(前一个字符)
4.2 完整模型实现代码
import torch
import torch.nn as nn
import torch.nn.functional as F
class BigramLanguageModel(nn.Module):
def __init__(self, vocab_size):
super().__init__()
# 每个字符的嵌入向量
self.token_embedding_table = nn.Embedding(vocab_size, vocab_size)
def forward(self, idx, targets=None):
# idx 和 targets 都是 (B,T) 的张量
logits = self.token_embedding_table(idx) # (B,T,C)
if targets is None:
loss = None
else:
B, T, C = logits.shape
logits = logits.view(B*T, C)
targets = targets.view(B*T)
loss = F.cross_entropy(logits, targets)
return logits, loss
def generate(self, idx, max_new_tokens):
# idx 是当前上下文 (B,T) 数组
for _ in range(max_new_tokens):
# 获取预测
logits, loss = self(idx)
# 只关注最后时间步
logits = logits[:, -1, :] # 变成 (B,C)
# 应用 softmax 获取概率
probs = F.softmax(logits, dim=-1) # (B,C)
# 从分布中采样
idx_next = torch.multinomial(probs, num_samples=1) # (B,1)
# 添加到序列中
idx = torch.cat((idx, idx_next), dim=1) # (B,T+1)
return idx
5. 数据预处理与训练流程
5.1 文本数据预处理
def prepare_data(text):
# 获取所有唯一字符
chars = sorted(list(set(text)))
vocab_size = len(chars)
# 创建字符到索引的映射
stoi = {ch: i for i, ch in enumerate(chars)}
itos = {i: ch for i, ch in enumerate(chars)}
encode = lambda s: [stoi[c] for c in s] # 编码器
decode = lambda l: ''.join([itos[i] for i in l]) # 解码器
# 将文本转换为张量
data = torch.tensor(encode(text), dtype=torch.long)
# 分割训练和验证集
n = int(0.9 * len(data))
train_data = data[:n]
val_data = data[n:]
return train_data, val_data, vocab_size, encode, decode
# 示例文本数据
text = """Hello, this is a simple Bigram language model.
It learns to predict the next character based on the previous one."""
train_data, val_data, vocab_size, encode, decode = prepare_data(text)
5.2 训练循环实现
def train_model(model, train_data, val_data, iterations=1000):
optimizer = torch.optim.AdamW(model.parameters(), lr=1e-3)
for iter in range(iterations):
# 获取一个小批量数据
ix = torch.randint(len(train_data) - 1, (4,)) # 批量大小4
xb = train_data[ix]
yb = train_data[ix + 1]
# 前向传播
logits, loss = model(xb, yb)
# 反向传播
optimizer.zero_grad(set_to_none=True)
loss.backward()
optimizer.step()
# 每100次迭代打印损失
if iter % 100 == 0:
with torch.no_grad():
val_loss = estimate_loss(model, val_data)
print(f"迭代 {iter}: 训练损失 {loss.item():.4f}, 验证损失 {val_loss:.4f}")
def estimate_loss(model, data):
model.eval()
losses = torch.zeros(10)
for k in range(10):
ix = torch.randint(len(data) - 1, (4,))
xb = data[ix]
yb = data[ix + 1]
_, loss = model(xb, yb)
losses[k] = loss.item()
model.train()
return losses.mean()
# 初始化并训练模型
model = BigramLanguageModel(vocab_size)
train_model(model, train_data, val_data)
6. 文本生成测试与效果验证
6.1 基础生成测试
# 从起始字符开始生成
context = torch.zeros((1, 1), dtype=torch.long)
generated_ids = model.generate(context, max_new_tokens=100)[0].tolist()
generated_text = decode(generated_ids)
print("生成的文本:")
print(generated_text)
6.2 不同起始点的生成效果
通过改变初始上下文,观察模型生成文本的多样性:
# 测试不同的起始字符
start_chars = ['H', 'T', 'I', 'M']
for start_char in start_chars:
context = torch.tensor([[encode(start_char)[0]]], dtype=torch.long)
generated = model.generate(context, max_new_tokens=50)[0].tolist()
print(f"以 '{start_char}' 开头: {decode(generated)}")
6.3 生成质量评估标准
评估 Bigram 模型生成文本时,关注以下几个维度:
- 连贯性 :生成的字符序列是否形成有意义的单词
- 多样性 :不同起始点是否能产生不同的文本模式
- 训练稳定性 :损失函数是否平稳下降
- 过拟合检查 :训练损失和验证损失的差距
7. 模型性能与资源观察
7.1 训练时间与资源占用
Bigram 模型的优势在于极低的资源需求:
- 训练时间 :1000 次迭代通常在 10-30 秒内完成(CPU)
- 内存占用 :模型参数极少,几乎不占用显存
- 推理速度 :生成 100 个字符约需 1-2 毫秒
7.2 性能优化技巧
虽然 Bigram 模型本身已经很轻量,但可以进一步优化:
# 使用 torch.jit.script 加速推理
scripted_model = torch.jit.script(model)
# 批量生成提高效率
def batch_generate(model, contexts, max_new_tokens=100):
"""批量生成文本"""
with torch.no_grad():
return model.generate(contexts, max_new_tokens=max_new_tokens)
8. 扩展到更复杂模型
8.1 从 Bigram 到 Trigram
理解了 Bigram 后,可以自然扩展到考虑更多上下文的模型:
class TrigramLanguageModel(nn.Module):
def __init__(self, vocab_size):
super().__init__()
self.token_embedding = nn.Embedding(vocab_size, 64)
self.position_embedding = nn.Embedding(2, 64) # 前两个位置
self.lm_head = nn.Linear(64, vocab_size)
def forward(self, idx, targets=None):
B, T = idx.shape
token_emb = self.token_embedding(idx) # (B,T,C)
pos_emb = self.position_embedding(torch.arange(T)) # (T,C)
x = token_emb + pos_emb # (B,T,C)
logits = self.lm_head(x) # (B,T,vocab_size)
if targets is None:
loss = None
else:
loss = F.cross_entropy(logits.view(-1, logits.size(-1)), targets.view(-1))
return logits, loss
8.2 与现代 LLM 的关联
Bigram 模型虽然简单,但包含了现代大语言模型的核心要素:
- 嵌入层 (Embedding Layer):将离散符号映射到连续向量空间
- Softmax 输出 :将网络输出转换为概率分布
- 自回归生成 :基于前面生成的内容预测下一个 token
- 交叉熵损失 :衡量预测分布与真实分布的差异
9. 常见问题与排查方法
| 问题现象 | 可能原因 | 排查方式 | 解决方案 |
|---|---|---|---|
| 训练损失不下降 | 学习率设置不当 | 检查损失曲线 | 调整学习率(1e-2 到 1e-4 尝试) |
| 生成文本重复 | 模型过于简单 | 观察生成多样性 | 增加训练数据量或模型复杂度 |
| 内存不足错误 | 数据量过大 | 检查数据张量大小 | 减小批量大小或序列长度 |
| 生成乱码 | 字符编码错误 | 验证编码解码函数 | 检查字符映射表是否正确 |
| 梯度爆炸 | 学习率过高 | 监控梯度范数 | 使用梯度裁剪或降低学习率 |
9.1 调试技巧
# 添加训练监控
def debug_training(model, data):
# 检查模型参数
for name, param in model.named_parameters():
print(f"{name}: {param.shape}")
# 验证前向传播
xb = data[:4].unsqueeze(0)
yb = data[1:5].unsqueeze(0)
logits, loss = model(xb, yb)
print(f"初始损失: {loss.item()}")
10. 实践建议与下一步学习路径
10.1 Bigram 模型的最佳实践
数据准备阶段 :
- 使用纯净的文本数据,避免特殊字符干扰
- 保持适当的数据量(几千到几万字符)
- 对中文文本需要先进行分词处理
训练调优 :
- 从小学习率开始(如 1e-3),根据损失曲线调整
- 使用合适的批量大小(通常 4-32)
- 定期验证集评估,防止过拟合
生成控制 :
- 通过调整温度参数控制生成随机性
- 尝试不同的起始字符获得多样结果
- 限制生成长度避免无限循环
10.2 进阶学习方向
掌握了 Bigram 模型后,可以沿着以下路径深入学习:
- 增加模型复杂度 :尝试 LSTM、GRU 等循环神经网络
- 引入注意力机制 :学习 Transformer 架构的基本原理
- 使用预训练模型 :上手 Hugging Face 的 Transformers 库
- 实践完整项目 :实现聊天机器人、文本分类等应用
- 学习优化技巧 :掌握模型压缩、量化、蒸馏等实用技术
Bigram 语言模型作为 LLM 学习的起点,其价值不在于生成质量,而在于帮助开发者建立对语言模型工作原理的直观理解。通过这个简单的模型,你可以清晰地看到从字符统计到神经网络生成的整个流程,为后续学习更复杂的 GPT、BERT 等模型打下坚实基础。
建议在实际操作中重点关注数据流向、损失变化和生成效果之间的关系,这种直观感受比单纯学习理论更能加深理解。完成本实验后,你会对"语言模型如何学习文本规律"这个问题有更具体的认识。
更多推荐


所有评论(0)