ARTICLE DETAIL

资讯详情

深耕网站建设、视觉设计与SEO优化的一线实战洞察。

手写LLM推理引擎:从黑盒到白盒的底层实现指南

手写LLM推理引擎:从黑盒到白盒的底层实现指南 几个月前我在给团队内部一个RAG服务做性能优化用户反馈生成太慢。我第一反应是换更大的显存、上量化、加并发折腾了一周效果有一点但每次调整都像隔着毛玻璃灭火——引擎内部到底发生了什么全靠外部观测指标去猜。真正让我下定决心的是后来遇到的一个诡异现象同样的prompt有时输出正常有时直接崩掉。我翻遍了框架的issue试了一堆workaround最后还是没辙只能一层层往里追。追到一半我才意识到我对LLM推理引擎内部那套机制的理解远没有自己以为的那么扎实。那次之后我想明白了与其继续当一个只会调用框架的黑盒用户不如自己照着论文和开源实现手写一个推理引擎哪怕一开始只是个玩具。这个系列就是这个决定的产物。第一章做的事情很简单——把地图铺开。一个LLM推理引擎到底由哪些部件组成每个部件在干什么为什么要这么设计动工前需要做哪些选型决策以及整个系列的路线图。如果你是那种用过不少LLM框架但遇到深层问题只能靠调参试错的人或者刚接触LLM底层、想找个系统切入点这个系列就是给你准备的。1. 从能跑到懂跑动手写推理引擎的真实动机1.1 黑盒调参的无力感用过transformers、vLLM、llama.cpp这些框架的人大概率都有过类似的体验调用一个model.generate()太简单了简单到让人误以为自己已经会了。可一旦出现下面这类问题黑盒式的使用方式会立刻捉襟见肘输入同样的prompt为什么有时候输出流畅有时候却重复几个词出不来为什么把temperature从0.8调成1.2输出变化幅度远不如预期为什么序列一长生成速度肉眼可见地掉KV Cache到底缓存了什么为什么同样的模型换一个量化格式之后输出质量就差了一截这些问题光看框架文档是找不到真正答案的。文档只告诉你调temperature可以让输出更随机但不会告诉你它是在logits上做除法、然后过softmax文档只告诉你KV Cache能加速生成但不会告诉你它省掉的到底是一次什么样的重复计算。我自己的体会是在没亲自实现过一遍之前所谓的理解很多其实是错觉。你以为自己知道注意力机制是Q和K做点积再softmax但真让你把attn.softmax(q k.T / sqrt(d)) v一步步用张量写出来并且处理因果掩码、处理批次维度、处理多头拆分时立刻就会发现有一堆细节是模糊的。1.2 亲手写一遍收益远超预期动手写推理引擎这件事短期看是重复造轮子长期看其实是收益极高的一项投资。我在完成第一版最小闭环之后回头再去看平时调参碰到的各种问题视角完全不一样了。举几个我实际感受到的变化性能问题不再靠猜。以前碰到生成慢只会怀疑是不是并发不够是不是显存不足。现在我能清楚地判断瓶颈是在prefill阶段还是decode阶段是显存带宽受限还是没有用KV Cache。采样参数变得可预期。因为自己写过采样器我清楚地知道temperature、top-k、top-p是在哪个环节生效改一个参数会影响什么、不会影响什么。这在做产品调优时特别有用。阅读论文和源码的门槛大幅降低。很多论文里的术语比如RoPE旋转位置编码GQA分组查询注意力投机采样因为有了自己的实现基础看的时候不再是囫囵吞枣而是能对应到具体代码的哪个模块。调试疑难bug更有底气。遇到框架层面的诡异问题至少知道往哪个方向追知道哪些层是自己可以控制的。这些收益不会在第一天显现但整个系列走完后回头看绝对是值回票价的。1.3 这个系列适合谁、不适合谁先说适合谁用过HuggingFace transformers或者其他推理框架但想彻底搞懂内部机制的人。想自己从零实现一个LLM推理引擎但不知道怎么入手、需要一张清晰路线图的人。做RAG、Agent等应用开发经常要跟模型推理打交道想提升定位问题能力的人。再说不太适合谁如果你的目标只是尽快在业务里用上LLM不想深究底层那么这个系列的ROI对你来说不高直接用现成框架就好。如果你完全没有Python基础也没接触过任何深度学习框架建议先花一两周补一下PyTorch的基础张量操作再回来。整个系列会用Python实现目标是写一个能加载真实开源模型权重、能正常生成文本的最小推理引擎。它不会追求极致性能而是追求结构清晰、每一步都看得懂、能复现。2. 一次生成请求的完整旅程token到文本的每一步在设计一个系统之前最好的方式是把一条完整的数据流从头到尾走一遍。下面这张地图是我动手之前先画清楚的东西也是整个系列的骨架。2.1 Tokenizer第一道门的钥匙推理引擎面对的第一件事是把用户输入的原始字符串变成模型能处理的整数ID序列。这个转换由Tokenizer完成。Tokenizer做的事大致是把文本切分成token子词单元再通过词汇表映射成ID。以GPT-2使用的BPEByte Pair Encoding为例它的核心思路是从字符级别开始反复合并出现频率最高的相邻字节对直到达到预设词表大小。所以像playing这种词可能被切成play和ing两个token中间用特殊符号标记。这一层看起来简单实际坑很多。不同模型家族的Tokenizer差异很大Llama用的是SentencePieceGPT-2用的是BPEQwen又是不一样的实现。同一个字符串经过不同tokenizer切出来的token数量可能差很多这直接影响后面的计算量和上下文长度。还有一个容易忽略的点特殊token。每个模型的词表里都有一批特殊token比如bos序列开始、eos序列结束、pad填充。有些模型的prompt需要手动拼接特殊token格式比如Chat模板中的|im_start|这类漏掉一个输出就可能完全不对。2.2 权重加载与Embedding查表拿到token ID序列后模型首先要做的是Embedding查表。Embedding层本质上是一张二维的大矩阵形状是vocab_size × hidden_size。每个token ID对应矩阵里的某一行这一行向量就是这个token的初始语义表示。例如GPT-2的hidden_size是768那就是把每个token变成一个768维的向量。这一步背后隐含着一个大工程权重加载。模型的所有参数都以权重文件的形式存储在磁盘上常见格式包括PyTorch的.bin、HuggingFace Safetensors的.safetensors等。写推理引擎时你需要把磁盘上的权重文件读进来按照模型结构里的名字一一对应地赋值给各个张量。权重文件的键名比如transformer.h.0.attn.c_attn.weight和模型内部变量名的映射关系是很多初学者第一步就会被绕晕的地方。除了token本身的语义模型还需要知道token的位置信息。最早的GPT-2使用可学习的位置编码也就是一张max_seq_len × hidden_size的位置矩阵直接和token嵌入相加。而Llama这类现代模型则使用RoPE旋转位置编码做法是直接对Q和K向量做旋转变换。这个差异之所以重要是因为RoPE对长文本的外推能力更好。2.3 Transformer块注意力、MLP与残差流Embedding出来的是一个形状为seq_len × hidden_size的张量。这个张量接下来会依次穿过几十个结构基本相同的Transformer块GPT-2有12个Llama-7B有32个。每个Transformer块内部一般包含两个核心子模块第一个是自注意力模块。输入张量先分别经过三个线性投影得到Q查询、K键、V值然后计算注意力分数softmax(QK^T / sqrt(d_k))再乘以V得到输出。这个过程可以理解为让每个位置去关注序列中所有其他位置在因果掩码的限制下只能关注当前位置及之前的位置从而把上下文信息融合进每个token的表示里。第二个是前馈网络模块MLP。通常是一个两层的全连接网络先升维再降维中间夹一个激活函数GPT-2用GELULlama用SiLU。它做的事情可以粗略理解为对每个位置的特征做一次非线性变换和维度扩展增加模型的表达能力。两个模块外面都套着残差连接也就是输出 输入 子模块(输入)。残差连接是让几十层深的网络能够稳定训练的关键设计同样也让推理时的数值更稳定。2.4 从logits到采样temperature到底在哪一步起作用最后一个Transformer块输出的hidden state会经过一个语言模型头通常是一个线性层权重和Embedding矩阵绑定或共享映射到词表大小得到每个token的得分也就是logits。logits还不是概率。它需要经过softmax归一化才能变成下一个token的概率分布。而temperature就是在softmax之前对logits做一次缩放scaled_logits logits / temperature probs softmax(scaled_logits)temperature越大缩放后的logits分布越平缓采样出来的结果越随机temperature越小分布越尖锐越接近贪心选择概率最高的token当temperature趋近于0时就等价于每次都选概率最大的token。很多人在应用层把temperature当成一个创意旋钮但如果你知道它在数学上只是除法缩放你就会明白它影响的是概率分布的熵而不是token本身的语义。所以指望单靠调temperature让模型变得更聪明方向就搞错了。采样器除了temperature还有top-k只从概率最高的k个token里选和top-p只从累计概率达到p的最小token集合里选。这些都是在logits变成概率之后对候选集做截断用来平衡生成质量和多样性。2.5 KV Cache推理性能的分水岭下一步是解释为什么生成过程会越来越慢。关键在KV Cache。在自回归生成中每生成一个token都要把新token拼到已有序列后面重新过一遍整个模型。如果不做缓存第n步生成时要把前n-1个token的Q、K、V全部重新计算一遍计算量随序列长度线性增长慢到无法接受。KV Cache的直觉很简单在生成过程中每个位置的K和V向量一经算出就不会变了因为前面的token不会改所以我们可以把它们缓存下来。第n步生成时只需要计算新token的Q、K、V然后拿新的Q去和缓存里的所有历史K做注意力计算即可。这样每步的计算量基本恒定不随序列长度增长。推理过程因此被分成两个阶段Prefill预填充阶段处理整个输入prompt一次性计算所有token的K、V并填充缓存。这个阶段是计算密集型的显存占用高但并行度高、速度快。Decode解码阶段逐token生成每步只算一个token。这个阶段是访存密集型的速度受限于显存带宽也就是为什么GPU的显存带宽对生成速度影响很大。KV Cache也有代价缓存大小随序列长度和层数线性增长。一个7B模型、2048序列长度、FP16精度KV Cache大概要占几百MB到1GB的显存。所以你会看到现代推理框架花大量精力做KV Cache的量化、PagedAttention这样的显存管理本质都是在跟这块空间较劲。3. 造轮子前的选型决策语言、依赖与硬件怎么定了解了全链路之后下一个问题就是用什么工具来实现。这一节把我做的选型决策和理由完整讲清楚。3.1 为什么用Python而不用C写第一版看到很多朋友一上来就想用C或Rust写推理引擎理由是性能好、接近底层。我的建议是第一版别这么干。这个系列的核心目标是搞懂原理不是比拼推理速度。用C写时你大量的精力会被内存管理、编译问题、第三方库绑定消耗掉真正用来理解注意力机制的时间反而被压缩了。Python配合PyTorch或NumPy可以让你用几十行代码就把一个Transformer块的前向计算写出来注意力机制的每一步都能直接打印出shape来验证调试体验完全是另一个维度。等原理全部搞通之后如果还想追求性能再往C/CUDA方向迁移也不迟。llama.cpp、vLLM这些生产级项目都走了这条路但它们背后有大量工程细节和优化技巧直接作为入门起点会把人劝退。所以我的选型是第一版用Python PyTorch。理由有三条张量运算由PyTorch底层优化我们只需关注逻辑权重文件格式和HuggingFace生态天然兼容不用自己写解析器调试工具链成熟print和断言可以随意加。3.2 依赖清单与版本陷阱第一版所需依赖不多核心就是下面几项PyTorch提供张量运算和自动微分虽然推理只用前向。CPU版本足够起步有GPU的话建议装CUDA版。transformers这个不是用来跑模型的而是用来做参考实现的。我们每实现一个模块需要和transformers的输出逐项对比验证正确性。safetensors或torch用于加载权重文件。safetensors更安全、加载更快但torch也能加载.bin格式。tiktoken或tokenizers用于调用现成tokenizer。第一版不建议自己实现BPE先用官方库把tokenizer逻辑跑通后续章节再展开讲BPE内部实现。版本陷阱这里提醒几个PyTorch大版本之间API有差异装的时候锁定一个较新的稳定版不要用最新的nightly。transformers版本会影响权重映射关系。建议用2.x或4.x的稳定版因为网上资料最多踩坑时更容易搜到答案。如果你用的是老显卡注意PyTorch的CUDA版本兼容性装之前去官网确认一下。3.3 硬件底线CPU也能跑起来的模型规模很多朋友担心我没有A100是不是就玩不了。答案是完全不是。推理引擎的学习路径对硬件要求很低。我们的小白鼠模型是GPT-2这个量级的参数量124MFP32精度下权重文件也就500MB左右。CPU完全可以跑只是生成速度慢几秒钟一个token但正好适合观察每一层的中间结果。如果你有一块普通的消费级显卡比如6GB以上显存体验会好很多可以跑的模型规模也更大。显存容量的估算公式很简单显存占用 ≈ 模型参数量 × 每个参数字节数 KV Cache 激活值开销比如7B模型用FP16每个参数2字节光权重就需要14GB显存。所以学习期最好是选1B以下的模型比如GPT-2、SmolLM、TinyLlama这些模型在普通硬件上都能流畅跑起来。我的建议是不纠结硬件先用CPU把GPT-2跑通再考虑升级。因为第一版最需要的是一个清晰正确的实现而不是速度。4. 第一版迭代目标先把能吐字的铁疙瘩跑通很多人在写这类项目时会陷入想一步到位的陷阱。我的策略相反第一版只求能吐字再慢慢打磨。把范围收敛得足够小才能保证自己在两周内拿到可见的正反馈。4.1 最小闭环的定义单序列、贪心解码、无KV Cache第一版我的目标是实现一个最小可用的推理引擎满足以下限定条件单序列一次只处理一个输入prompt不做batch批处理。避免处理padding、attention mask等批量相关的复杂逻辑。贪心解码每一步直接选择概率最大的token不做temperature采样、不做top-k/top-p。这样采样器只有一句argmax整个生成循环的逻辑最简。不用KV Cache每生成一个新token重新从头算一遍完整序列的前向。但代码里要留下KV Cache的接口位置方便后续章节优化。固定长度上下文比如最大支持512个token超出就截断。不用处理动态长度带来的各种边界情况。这四个限定条件砍掉了大约一半的实现复杂度但保留了一条完整的推理链路输入文本 → tokenizer → embedding → N个Transformer块 → LM head → softmax → argmax → 新token → 拼接到序列 → 循环 → 输出文本。等到这版跑通、验证正确之后再逐步松开约束先加KV Cache然后加batch最后加采样策略。每一步都有明确的验证目标不容易翻车。4.2 验证正确性的三把尺子自己写的引擎跑出结果了怎么知道它是对的我总结了三把尺子从粗到细逐级验证第一把尺子整体输出对比。用同一个prompt、同一份权重跑我们自己的引擎和HuggingFace transformers的model.generate()对比生成的文本是否一致。如果一致说明大方向对了。但这里有个坑如果两边都用了贪心解码输出应该完全一致如果都用了随机采样那就要固定随机种子才能对比。第二把尺子逐层logits对比。这是最有效的调试手段。在模型的每一层Transformer块之后把我们的中间张量print出来和transformers的对应层输出做数值对比。用torch.allclose检查误差一般允许1e-4左右的小误差因为float运算顺序不同会有微小差异。如果某一层开始出现显著偏差问题就锁定在这一层之前排查范围大幅缩小。第三把尺子手工构造简单样例。比如构造一个只有2个token的输入手工计算一次注意力分数和程序输出对比。这种小型样例验证对初学阶段特别有用因为它能让你确认自己对每一个算子的理解都是对的。这三把尺子配合使用基本能把是不是我某个矩阵转置搞错了是不是忘记加偏置了这类低级错误快速定位。4.3 选一只合适的小白鼠模型第一版选什么模型直接决定了调试难度。我强烈推荐GPT-2124M版原因如下参数规模小124M参数FP32下约500MB下载快、加载快、CPU就能跑。架构经典且规范没有RoPE、GQA这些现代特性就是最朴素的Embedding 12层Transformer块 LM head每一步都容易对应到教程。权重格式成熟HuggingFace上有大量资料和现成代码可以参考遇到问题容易搜到答案。开源生态完善很多经典的推理教程包括Karpathy的minGPT、nanoGPT都是基于或兼容GPT-2结构可以直接对照学习。如果你想要体验更现代的架构也可以在后续章节把实现迁移到SmolLM或TinyLlama上它们的特点是引入了RoPE和GQA等GPT-2版本跑通之后再迁移会顺畅很多。5. 绕不开的地基Transformer、张量shape与权重加载这一节把动手前必须搞懂的几个核心知识点集中讲一遍。理解这些后面写代码时才不会迷失在张量维度里。5.1 一个Transformer块内部发生了什么以一个GPT-2的Transformer块为例它的输入是上一层的输出x形状为[seq_len, hidden_size]单序列场景。块内依次完成LayerNorm对x做归一化稳定数值。QKV投影x分别乘以三个矩阵得到Q、K、V。GPT-2的实现里三个投影合并在一个大的c_attn线性层里一次算出来再切分输出维度是3 * hidden_size。多头拆分把Q、K、V从[seq_len, hidden_size]重塑成[num_heads, seq_len, head_dim]的形式每个头独立做注意力。缩放点积注意力算Q K^T / sqrt(head_dim)得到注意力分数矩阵[num_heads, seq_len, seq_len]应用因果掩码把右上三角设为负无穷softmax归一化再乘以V。输出投影把所有头的输出拼回去经过c_proj线性层得到这个子模块的输出。残差相加x x 子模块输出。MLP经过c_fc升维到4倍hidden_size→ GELU激活 →c_proj降维回来。再次残差相加x x mlp输出。这段流程里初学者最容易犯的错误是忘记因果掩码或者在多头拆分和合并时搞错维度顺序。建议每一步都打印shape验证shape对了逻辑大概率就对了一半。5.2 权重shape对照表每个张量都是什么意思我整理了一份GPT-2的权重shape对照表124M版动手前可以先对着这张表建立模型 一堆张量的直觉权重名称Shape含义wte[50257, 768]Token Embedding矩阵token ID → 向量wpe[1024, 768]位置Embedding矩阵位置 → 向量h.{i}.ln_1.weight/bias[768]第一个LayerNorm参数h.{i}.attn.c_attn.weight[768, 2304]QKV合并投影2304 768 × 3h.{i}.attn.c_attn.bias[2304]QKV投影偏置h.{i}.attn.c_proj.weight[768, 768]注意力输出投影h.{i}.ln_2.weight/bias[768]MLP前的LayerNorm参数h.{i}.mlp.c_fc.weight[768, 3072]MLP升维投影3072 768 × 4h.{i}.mlp.c_proj.weight[3072, 768]MLP降维投影ln_f.weight/bias[768]最后的LayerNorm参数lm_head.weight[50257, 768]语言模型头通常和wte共享注意几个容易搞错的地方PyTorch的nn.Linear执行的是x W.T b所以权重文件的shape是[out_features, in_features]。c_attn一次输出2304维切分时是沿着最后一维切成三份768而不是reshape成[3, 768]再处理。共享权重的lm_head在保存时可能没有单独存一份加载时要用wte的权重初始化。5.3 权重文件的加载与合包重组有了shape对照表权重加载就变成了一件机械但需要细心的工作。基本流程是用safetensors库读取权重文件得到一个字典键是权重名值是张量。创建一个空的模型结构我们自己的类。逐个遍历字典里的键把值赋给模型里对应的参数。检查是否所有参数都有对应的权重没有匹配到的用load_state_dict时会报错提示。实际过程中最烦人的是键名映射。HuggingFace的键名和模型类的属性名往往不是一一对应的比如transformer.h.0.attn.c_attn.weight对应到我们的类里可能是blocks[0].attn.qkv.weight。遇到这种情况建议在模型类里写一个load_weights方法显式地做映射不要图省事直接靠名字匹配。另外要提醒的是加载权重时的dtype要一致。如果用FP32加载了FP16的权重要么先转FP32再加载要么全链路都保持FP16混用会在后面的计算中产生奇怪的数值问题。6. 系列路线图与学习心法整条路线最终要走多远我在这里先把图景画出来也分享几条我自己觉得很有用的学习心法。6.1 后续章节路线图这个系列大致按以下顺序推进每一章都建立在前一章的基础上第二章Tokenizer与数据准备。深入BPE算法实现或复现一个tokenizer的核心逻辑搞清楚词表构建与编码解码的完整流程。第三章Embedding与第一个Transformer块。写embedding层、位置编码以及单个Transformer块的前向计算重点把多头注意力的维度变化理清楚。第四章完整前向传播与权重加载。拼装所有层实现完整的GPT-2前向网络打通权重文件 → 模型张量 → logits的链路。第五章生成循环与采样器。实现自回归生成循环加入temperature、top-k、top-p对比不同参数下的生成效果。第六章KV Cache优化。实现缓存机制区分prefill和decode阶段对比优化前后的速度差异。第七章批处理与Attention Mask。引入batch维度处理变长序列的padding问题。后续进阶迁移到Llama架构RoPE、GQA、量化推理、投机采样等。每一章都会以可运行、可验证作为完成标准不会出现原理讲了一堆、代码跑不起来的情况。6.2 几条实用的学习建议最后分享几条我的实操体会不是大道理都是踩过坑之后总结出来的第一尽量对照着写不要抄着写。网上有大量优质参考代码比如GPT-2的官方实现、nanoGPT、minGPT遇到卡壳时可以看但建议看一行写一行自己想清楚每一行在干什么而不是整段复制。抄一遍代码得到的理解远不如自己卡住半小时后再顿悟来得深。第二把打印shape当成习惯。张量维度的变化是推理引擎里最隐蔽也最烦人的错误源。我在调试时每写两三行就print一次shape确保跟预期一致这个习惯帮我少走了一半弯路。第三用最小复现定位问题不要大范围猜。如果输出不对先构造一个只有几个token的极端小输入把流程简到不能再简再逐层检查哪一步开始和参考实现有偏差。靠肉眼扫代码很难找到问题靠二分定位快得多。第四遇到不懂的数学概念先看直觉再看公式。比如注意力机制先理解它是让每个token根据相关性加权汇总所有位置信息的一种手段再去看softmax(QK^T/√d)V这个公式就会发现它不过是在实现这个直觉。第五卡住了就停下来隔一天再看。这是我个人收益最大的一条经验。写推理引擎时经常会在某个维度问题里钻牛角尖越急越乱。隔一天重新看往往十几分钟就想明白了。按照上面这张路线图加上这些心法整个系列走下来你会得到的不只是一个能跑的小推理引擎更是一套看任何LLM框架源码都不打怵的底层能力。下一章我们从Tokenizer开始把第一块拼图真正动手安上去。
返回列表