10分钟掌握Karpathy Bigram语言模型:从原理到PyTorch实战

发布时间:2026/7/21 9:59:14
10分钟掌握Karpathy Bigram语言模型:从原理到PyTorch实战 Caleb Writes Code10分钟精讲Karpathy Bigram语言模型在自然语言处理领域大语言模型LLM已经成为技术热点但很多开发者对其底层原理了解有限。本文基于Andrej Karpathy的教学内容深入解析Bigram语言模型的核心原理与实现帮助读者从零理解语言模型的基础构建块。1. 语言模型基础概念1.1 什么是语言模型语言模型是自然语言处理中的核心组件其主要任务是计算一个词序列出现的概率。简单来说语言模型能够评估一句话听起来是否自然。比如今天天气很好比天气很好今天具有更高的概率得分因为前者更符合语言习惯。在数学上语言模型计算的是条件概率给定前n-1个词预测第n个词出现的概率。这种概率计算使得语言模型能够用于机器翻译、语音识别、文本生成等多种应用场景。1.2 Bigram模型的核心思想Bigram二元语法模型是语言模型中最简单的形式之一它基于马尔可夫假设即当前词的出现概率只依赖于前一个词。这种简化虽然损失了长距离依赖信息但大大降低了模型复杂度使其成为理解语言模型原理的理想起点。Bigram模型的计算公式为P(w_i | w_{i-1})表示在给定前一个词w_{i-1}的情况下当前词w_i出现的条件概率。通过统计大量文本数据中词对共现的频率我们可以构建出一个完整的Bigram概率矩阵。2. Karpathy Bigram模型实现环境准备2.1 开发环境要求要实现Karpathy风格的Bigram语言模型我们需要准备以下开发环境Python 3.8建议使用较新的Python版本确保兼容性PyTorch 1.9深度学习框架用于模型构建和训练NumPy数值计算库Matplotlib可选用于结果可视化2.2 安装必要依赖pip install torch numpy matplotlib对于国内用户如果下载速度较慢可以使用清华镜像源pip install -i https://pypi.tuna.tsinghua.edu.cn/simple torch numpy matplotlib2.3 验证环境配置安装完成后可以通过以下代码验证环境是否配置正确import torch import numpy as np print(fPyTorch版本: {torch.__version__}) print(fCUDA是否可用: {torch.cuda.is_available()}) print(fNumPy版本: {np.__version__})3. Bigram模型核心原理详解3.1 概率计算基础Bigram模型的核心是构建一个条件概率表。假设我们有一个包含V个词的词汇表那么Bigram概率矩阵的大小就是V×V。每个元素P(i|j)表示在给定词j的情况下下一个词是i的概率。概率计算基于最大似然估计 P(w_i | w_{i-1}) count(w_{i-1}, w_i) / count(w_{i-1})其中count(w_{i-1}, w_i)表示词对(w_{i-1}, w_i)在训练语料中出现的次数count(w_{i-1})表示词w_{i-1}出现的总次数。3.2 平滑技术的重要性在实际应用中由于训练数据有限很多词对可能从未出现过导致概率为零。这就需要使用平滑技术来分配小概率给未出现的词对。常见的平滑方法包括Add-one平滑、Good-Turing估计等。Add-one平滑的计算公式为 P(w_i | w_{i-1}) [count(w_{i-1}, w_i) 1] / [count(w_{i-1}) V]其中V是词汇表大小。这种平滑方法虽然简单但能有效避免零概率问题。4. 完整实现Karpathy风格Bigram模型4.1 数据预处理首先我们需要准备训练数据并进行必要的预处理import torch import torch.nn as nn import torch.nn.functional as F # 示例训练文本 text 在自然语言处理领域语言模型是重要的基础组件。 Bigram模型虽然简单但能很好地展示语言模型的基本原理。 通过这个例子我们可以深入理解概率语言模型的工作机制。 # 构建词汇表 chars sorted(list(set(text))) vocab_size len(chars) print(f词汇表大小: {vocab_size}) print(f字符列表: {.join(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]) # 解码函数 print(f编码示例: {encode(自然语言)}) print(f解码示例: {decode(encode(自然语言))})4.2 构建Bigram概率矩阵接下来我们基于训练文本构建Bigram概率矩阵# 将文本编码为整数序列 data torch.tensor(encode(text), dtypetorch.long) print(f数据形状: {data.shape}) # 构建Bigram计数矩阵 bigram_counts torch.zeros((vocab_size, vocab_size), dtypetorch.int32) for i in range(len(data) - 1): bigram_counts[data[i], data[i1]] 1 print(Bigram计数矩阵前10x10:) print(bigram_counts[:10, :10]) # 转换为概率矩阵使用平滑 bigram_probs (bigram_counts 1).float() # Add-one平滑 bigram_probs bigram_probs / bigram_probs.sum(1, keepdimTrue) print(Bigram概率矩阵前10x10:) print(bigram_probs[:10, :10])4.3 文本生成实现基于训练好的Bigram模型我们可以实现文本生成功能def generate_text(model, start_char, max_length100): 使用Bigram模型生成文本 current_char start_char generated [current_char] for _ in range(max_length - 1): if current_char not in stoi: break current_idx stoi[current_char] next_idx torch.multinomial(model[current_idx], 1).item() next_char itos[next_idx] generated.append(next_char) current_char next_char if next_char in [。, , ]: # 遇到句号可能结束 if torch.rand(1).item() 0.3: # 30%概率结束 break return .join(generated) # 测试文本生成 print(生成的文本示例:) for i in range(3): start_char text[torch.randint(0, len(text), (1,)).item()] generated generate_text(bigram_probs, start_char, 50) print(f示例 {i1}: {generated})5. 神经网络实现的Bigram模型5.1 使用PyTorch构建模型虽然传统的Bigram模型基于统计但我们可以用神经网络来实现相同的功能这有助于理解现代语言模型的基本结构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, targetsNone): # 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.forward(idx) # 只关注最后一个时间步 logits logits[:, -1, :] # 变为(B, C) # 应用softmax获取概率 probs F.softmax(logits, dim-1) # (B, C) # 采样下一个字符 idx_next torch.multinomial(probs, num_samples1) # (B, 1) # 添加到序列中 idx torch.cat((idx, idx_next), dim1) # (B, T1) return idx # 实例化模型 model BigramLanguageModel(vocab_size) print(f模型参数数量: {sum(p.numel() for p in model.parameters())})5.2 模型训练与优化训练神经网络版本的Bigram模型# 准备训练数据 def get_batch(data, batch_size, block_size): 获取小批量数据 ix torch.randint(len(data) - block_size, (batch_size,)) x torch.stack([data[i:iblock_size] for i in ix]) y torch.stack([data[i1:iblock_size1] for i in ix]) return x, y # 训练参数 batch_size 4 block_size 8 learning_rate 1e-2 max_iters 1000 # 优化器 optimizer torch.optim.AdamW(model.parameters(), lrlearning_rate) # 训练循环 for iter in range(max_iters): # 获取小批量数据 xb, yb get_batch(data, batch_size, block_size) # 前向传播 logits, loss model(xb, yb) # 反向传播 optimizer.zero_grad(set_to_noneTrue) loss.backward() optimizer.step() if iter % 200 0: print(f迭代 {iter} | 损失: {loss.item():.4f}) # 测试生成 context torch.zeros((1, 1), dtypetorch.long) generated_ids model.generate(context, max_new_tokens100)[0].tolist() generated_text decode(generated_ids) print(f神经网络生成的文本: {generated_text})6. 模型评估与结果分析6.1 困惑度计算困惑度是评估语言模型性能的重要指标它衡量模型对测试数据的预测不确定性def calculate_perplexity(model, data, block_size8): 计算模型在数据上的困惑度 model.eval() total_loss 0 count 0 with torch.no_grad(): for i in range(0, len(data) - block_size, block_size): x data[i:iblock_size].unsqueeze(0) y data[i1:iblock_size1].unsqueeze(0) _, loss model(x, y) total_loss loss.item() count 1 avg_loss total_loss / count perplexity torch.exp(torch.tensor(avg_loss)) return perplexity.item() perplexity calculate_perplexity(model, data) print(f模型困惑度: {perplexity:.2f})6.2 生成文本质量分析通过分析生成文本的统计特性我们可以评估模型的质量def analyze_generated_text(text, original_text): 分析生成文本的质量 # 计算重复率 words text.split() unique_words set(words) repetition_rate 1 - len(unique_words) / len(words) if words else 0 # 计算平均句长 sentences text.replace(。, 。|).replace(, |).replace(, |).split(|) sentences [s for s in sentences if s.strip()] avg_sentence_length sum(len(s) for s in sentences) / len(sentences) if sentences else 0 print(f生成文本长度: {len(text)} 字符) print(f重复率: {repetition_rate:.2%}) print(f平均句长: {avg_sentence_length:.1f} 字符) # 与原始文本比较 original_words set(original_text.split()) generated_words set(words) overlap len(original_words generated_words) / len(generated_words) if generated_words else 0 print(f词汇重叠率: {overlap:.2%}) analyze_generated_text(generated_text, text)7. 常见问题与解决方案7.1 训练数据不足的问题当训练数据较少时Bigram模型容易过拟合生成文本多样性不足。解决方案包括数据增强通过回译、同义词替换等方式扩充训练数据更强的平滑技术使用Kneser-Ney平滑等更先进的方法模型正则化在神经网络版本中添加Dropout等正则化技术# 添加Dropout的改进模型 class ImprovedBigramModel(nn.Module): def __init__(self, vocab_size, dropout0.1): super().__init__() self.token_embedding nn.Embedding(vocab_size, 64) # 增加嵌入维度 self.dropout nn.Dropout(dropout) self.linear nn.Linear(64, vocab_size) def forward(self, idx, targetsNone): x self.token_embedding(idx) x self.dropout(x) logits self.linear(x) if targets is not None: 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 return logits, None7.2 生成文本缺乏连贯性Bigram模型只考虑前一个词导致生成长文本时缺乏全局连贯性。改进策略增加上下文长度使用Trigram或更高阶的N-gram模型引入注意力机制像Transformer那样关注更远的上下文后处理优化对生成结果进行重排序和筛选8. Bigram模型在实际项目中的应用8.1 文本自动补全Bigram模型可以用于实现简单的文本自动补全功能class TextAutoComplete: def __init__(self, model, stoi, itos): self.model model self.stoi stoi self.itos itos def suggest_next_chars(self, prefix, top_k3): 根据前缀建议下一个字符 if not prefix or prefix[-1] not in self.stoi: return [] last_char prefix[-1] char_idx self.stoi[last_char] with torch.no_grad(): logits, _ self.model(torch.tensor([[char_idx]])) probs F.softmax(logits[0, -1], dim-1) top_probs, top_indices torch.topk(probs, top_k) suggestions [] for prob, idx in zip(top_probs, top_indices): suggestions.append((self.itos[idx.item()], prob.item())) return suggestions # 使用示例 autocomplete TextAutoComplete(model, stoi, itos) prefix 自然 suggestions autocomplete.suggest_next_chars(prefix) print(f{prefix} 的下一个字符建议:) for char, prob in suggestions: print(f {char}: {prob:.2%})8.2 拼写错误检测基于Bigram概率我们可以检测文本中的拼写错误def spell_check(text, model, stoi, itos, threshold0.01): 简单的拼写错误检测 words text.split() suspicious_words [] for i in range(1, len(words)): prev_word words[i-1][-1] if words[i-1] else # 取前一个词的最后一个字 current_word words[i][0] if words[i] else # 取当前词的第一个字 if prev_word in stoi and current_word in stoi: prev_idx stoi[prev_word] current_idx stoi[current_word] with torch.no_grad(): logits, _ model(torch.tensor([[prev_idx]])) probs F.softmax(logits[0, -1], dim-1) prob probs[current_idx].item() if prob threshold: suspicious_words.append((f{prev_word}{current_word}, prob)) return suspicious_words # 测试拼写检查 test_text 自然语言处理领或很重要 suspicious spell_check(test_text, model, stoi, itos) print(可疑词对:) for word_pair, prob in suspicious: print(f {word_pair}: 概率{prob:.4f})9. 从Bigram到现代LLM的演进路径9.1 技术发展脉络理解Bigram模型是学习现代大语言模型的重要基础。从Bigram到GPT系列模型的技术演进主要包括上下文扩展从Bigram到N-gram再到基于注意力机制的无限上下文表示学习从one-hot编码到词嵌入再到上下文相关的动态表示模型架构从统计模型到神经网络再到Transformer架构训练规模从小规模数据到海量互联网文本9.2 学习路线建议对于想要深入LLM领域的开发者建议按照以下路线学习基础阶段掌握Bigram/Trigram等传统语言模型进阶阶段学习Word2Vec、LSTM、Seq2Seq等神经网络模型高级阶段深入理解Transformer架构和注意力机制实践阶段使用Hugging Face等工具库实践现代LLM应用10. 最佳实践与工程建议10.1 模型部署注意事项在实际项目中部署语言模型时需要考虑性能优化对于实时应用需要优化推理速度内存管理大型模型需要合理的内存管理策略缓存机制对频繁使用的预测结果进行缓存监控告警建立模型性能监控体系# 简单的模型缓存实现 class CachedBigramModel: def __init__(self, model, cache_size1000): self.model model self.cache {} self.cache_size cache_size self.access_count {} def predict(self, char): if char in self.cache: self.access_count[char] 1 return self.cache[char] # 缓存未命中计算预测 with torch.no_grad(): if char not in stoi: result {} else: char_idx stoi[char] logits, _ self.model(torch.tensor([[char_idx]])) probs F.softmax(logits[0, -1], dim-1) result {itos[i]: probs[i].item() for i in range(len(itos))} # 更新缓存 self._update_cache(char, result) return result def _update_cache(self, key, value): if len(self.cache) self.cache_size: # 移除最不常用的项 min_key min(self.access_count, keyself.access_count.get) del self.cache[min_key] del self.access_count[min_key] self.cache[key] value self.access_count[key] 110.2 错误处理与日志记录健壮的生产系统需要完善的错误处理和日志记录import logging # 配置日志 logging.basicConfig(levellogging.INFO, format%(asctime)s - %(levelname)s - %(message)s) class RobustBigramModel: def __init__(self, model, stoi, itos): self.model model self.stoi stoi self.itos itos self.logger logging.getLogger(__name__) def safe_predict(self, text): try: if not text: self.logger.warning(输入文本为空) return {} last_char text[-1] if last_char not in self.stoi: self.logger.warning(f字符 {last_char} 不在词汇表中) return {} return self._do_predict(last_char) except Exception as e: self.logger.error(f预测过程中发生错误: {str(e)}) return {} def _do_predict(self, char): # 实际的预测逻辑 char_idx self.stoi[char] with torch.no_grad(): logits, _ self.model(torch.tensor([[char_idx]])) probs F.softmax(logits[0, -1], dim-1) return {self.itos[i]: probs[i].item() for i in range(len(self.itos))}通过本文的详细讲解和代码实践相信读者已经对Bigram语言模型有了深入的理解。这种基础的语言模型虽然简单但包含了现代LLM的核心思想是学习更复杂模型的重要基础。