☰
从零手搓AI推理引擎:深入底层原理与工程实践
2026/9/29 16:49:56 网站建设 项目流程

1. 从零手搓AI工程:为什么我不建议你直接调包

第一次看到ai-engineering-from-scratch这个项目名的时候,我正坐在工位上啃一个已经调了三天的推理服务性能问题。当时第一反应是:又是一个教你怎么调transformers库的教程仓库吧。结果点进去翻了翻,发现它讲的是从最底层的矩阵乘法开始,一步步把一个大模型推理引擎搭出来——包括张量抽象、算子实现、内存管理、KV Cache、批处理调度,最后跑通一个能用的推理服务。

这个方向其实挺反直觉的。现在这个时间点,2025年,你随便pip install一个库就能跑起来一个对话模型,为什么还要从零写?我带着这个疑问把项目里的代码和思路过了一遍,又自己动手复现了其中几个核心模块,慢慢理解了它的价值所在。

ai-engineering-from-scratch这个项目,本质上是一套从底层原理出发、手写AI工程核心组件的实践路线。它不教你调API,不教你写Prompt,它教你的是:当你面对一个推理延迟超标、显存爆掉、吞吐上不去的真实问题时,你能不能定位到是哪个环节出了问题,能不能自己动手改。

适合谁来参考?我觉得有三类人:一是已经会用现成框架、但遇到性能瓶颈就束手无策的工程师;二是想真正理解大模型推理内部到底发生了什么的学生或转行者;三是需要做定制化算子或极致优化的团队开发者。如果你只是想快速搭个Demo,这个项目可能不适合你——但如果你想搞清楚“为什么我的推理服务这么慢”,那它值得你花时间。

我接下来会从整体设计思路、核心模块拆解、实操复现过程、以及踩坑排查几个角度,把这个项目的精髓讲清楚。所有代码和参数我都会给出可复现的版本,你跟着做就能跑通。

2. 整体设计思路:为什么选择从零手写而不是调包

2.1 从零手写的核心动机

先说一个我自己的经历。去年我接手过一个推理服务,用的是某个主流推理框架,QPS上不去,延迟波动很大。我花了整整一周时间,从框架的Python层一路往下翻到C++层,最后发现是KV Cache的内存分配策略在特定batch size下产生了大量碎片。这个问题如果我不懂底层内存管理,根本不可能定位到。

ai-engineering-from-scratch的设计思路就是围绕这个痛点来的。它的核心主张是:你不需要从零写一个完整的推理引擎,但你需要有能力从零写出每一个核心组件。这两者的区别很大。前者是造轮子,后者是理解轮子。

项目把AI工程拆成了几个层次:

  • 张量层:怎么表示一个多维数组,怎么管理它的内存布局,怎么做view和reshape而不拷贝数据
  • 算子层:矩阵乘法、Softmax、LayerNorm、Attention这些核心算子怎么手写实现
  • 模型层:怎么把算子组装成一个Transformer,怎么做权重加载和参数管理
  • 推理层:KV Cache怎么设计,怎么做批处理,怎么做采样
  • 服务层:怎么把推理引擎包装成一个可用的服务,怎么做并发和调度

每一层都有对应的手写实现,而且每一层都解释了“为什么这么设计”以及“不这么设计会怎样”。

2.2 为什么不用PyTorch直接搭

有人可能会问:我用PyTorch搭一个Transformer不就几十行代码吗,为什么要手写算子?

这个问题我在项目讨论区也看到过。答案其实很直接:PyTorch的抽象层次太高,你无法控制底层行为。举个例子,当你调用torch.matmul的时候,你并不知道它底层用的是哪个BLAS库、有没有做内存对齐、有没有触发隐式的数据拷贝。在大多数场景下这没问题,但在你需要极致优化的时候,这些“不知道”就是你的瓶颈。

项目里有一个很典型的对比实验:同样的矩阵乘法,用PyTorch的torch.matmul和手写的分块矩阵乘法,在特定矩阵尺寸下性能差距能达到30%以上。原因就在于手写版本可以根据矩阵尺寸选择最优的分块策略,而PyTorch的通用实现做不到这么细粒度的适配。

当然,我不是说手写就一定比PyTorch快。在大多数情况下,PyTorch经过多年优化的算子比你手写的要快得多。但关键在于:当PyTorch不够快的时候,你知道怎么改。这个能力才是项目真正想教给你的。

2.3 技术选型背后的考量

项目在技术选型上做了几个关键决策,我觉得都挺有道理的:

用Python做胶水层,用C++/CUDA做计算层。这是工业界的标准做法。Python负责灵活的逻辑控制和快速迭代,C++/CUDA负责性能敏感的计算。项目里所有的算子都有Python的参考实现和C++/CUDA的高性能实现,你可以对照着看,理解两者之间的差异。

不依赖任何深度学习框架。项目只依赖NumPy做基础数组操作,以及PyTorch做正确性验证的对照。这意味着你需要自己实现所有东西,但同时也意味着你完全掌控了每一行代码。

从推理入手而不是训练。训练涉及反向传播、优化器、分布式等更复杂的系统,而推理是AI工程中最常遇到的场景。项目选择从推理切入,降低了入门门槛,同时覆盖了最核心的工程问题。

提示:如果你之前没有接触过CUDA编程,项目里有一个专门的入门章节,用最简化的例子讲清楚线程、block、grid这些概念。我建议不要跳过,否则后面看算子实现会很吃力。

3. 核心模块拆解与实操要点

3.1 张量抽象:一切的基础

张量是AI工程中最基础的数据结构。项目里实现了一个简化版的张量类,核心属性包括:

  • data:底层数据存储,通常是一个一维的连续内存块
  • shape:各维度的大小
  • strides:各维度的步长,决定了如何从一维数据映射到多维索引
  • dtype:数据类型

这里最关键的是strides的设计。很多人不理解为什么需要strides,觉得有shape就够了。但实际上,strides决定了你能否在不拷贝数据的情况下做转置、切片、广播等操作。

举个例子,一个shape为(3, 4)的二维张量,如果按行优先存储,它的strides是(4, 1)。转置之后,shape变成(4, 3),strides变成(1, 4),但底层数据完全没动。这就是所谓的“视图操作”,零拷贝。

项目里有一个练习:实现一个transpose函数,要求不拷贝数据。我一开始写的时候直接新建了一个数组然后逐元素拷贝,后来看了参考实现才意识到,只需要交换shape和strides就行了。这个认知上的转变很重要,它让你理解到张量操作的本质是对内存的重新解释,而不是数据的搬运。

实操要点:

  • 创建张量时,确保底层数据是连续存储的,否则后续很多操作会触发隐式拷贝
  • 实现contiguous()方法,用于在需要时把非连续张量变成连续张量
  • 注意strides的计算方式,特别是对于多维张量的切片操作

3.2 核心算子实现:从矩阵乘法到Attention

矩阵乘法是深度学习中最核心的算子,没有之一。项目里从最朴素的三重循环开始,逐步优化到分块矩阵乘法,再到利用SIMD指令和CUDA的并行实现。

朴素实现很简单:

def matmul_naive(A, B): M, K = A.shape K2, N = B.shape assert K == K2 C = np.zeros((M, N)) for i in range(M): for j in range(N): for k in range(K): C[i, j] += A[i, k] * B[k, j] return C

这个实现的时间复杂度是O(MNK),在矩阵尺寸稍大时就慢得无法接受。优化的第一步是循环重排,把k循环放到最内层,利用CPU的缓存局部性。第二步是分块,把大矩阵切成小块,让每个小块能放进CPU缓存。第三步是向量化,用SIMD指令一次处理多个数据。

项目里给出了每一步的代码和性能对比。我实测下来,在(512, 512)的矩阵乘法上,分块版本比朴素版本快了将近20倍。这个数字可能因硬件而异,但优化的思路是通用的。

Attention算子是另一个重点。它的核心是softmax(Q @ K^T / sqrt(d)) @ V。项目里把Attention拆成了几个步骤:

  1. Q、K、V的线性投影
  2. Q和K的矩阵乘法,得到注意力分数
  3. 缩放和Softmax
  4. 注意力分数和V的矩阵乘法

每一步都有手写实现,而且解释了为什么要这么拆。比如Softmax需要做数值稳定处理,减去最大值再取指数,否则容易溢出。这个细节在调包的时候你根本不会注意到,但手写的时候必须考虑。

注意:Softmax的数值稳定处理不是可选项,是必须项。我见过有人手写Softmax时忘了减最大值,结果在fp16精度下直接溢出成NaN。这个坑项目里专门有一节讲。

3.3 KV Cache:推理加速的关键

KV Cache是自回归生成中最重要的优化手段。没有KV Cache的话,每生成一个token都需要重新计算所有历史token的Key和Value,计算量随序列长度平方增长。有了KV Cache,只需要计算当前token的Key和Value,然后和缓存的拼接起来。

项目里实现了一个简单的KV Cache,核心逻辑是:

class KVCache: def __init__(self, max_seq_len, num_layers, num_heads, head_dim): self.max_seq_len = max_seq_len self.cache_k = np.zeros((num_layers, max_seq_len, num_heads, head_dim)) self.cache_v = np.zeros((num_layers, max_seq_len, num_heads, head_dim)) self.seq_len = 0 def update(self, layer_idx, new_k, new_v): start = self.seq_len end = start + new_k.shape[0] self.cache_k[layer_idx, start:end] = new_k self.cache_v[layer_idx, start:end] = new_v return self.cache_k[layer_idx, :end], self.cache_v[layer_idx, :end]

这个实现虽然简单,但涵盖了KV Cache的核心思想。实际生产中还需要考虑内存预分配、分页管理、多请求共享等问题,但理解了基础版本之后,这些扩展都是水到渠成的。

实操心得:KV Cache的内存占用很容易被低估。以一个7B模型为例,如果max_seq_len是4096,num_layers是32,num_heads是32,head_dim是128,那么KV Cache的总大小是2 * 32 * 4096 * 32 * 128 * 2 bytes(fp16),大约是2GB。这还只是一个请求的缓存。如果并发请求多了,显存很快就爆了。所以实际部署时一定要做好内存预算。

3.4 批处理与调度:从单请求到多请求

单请求的推理很简单,但生产环境需要同时处理多个请求。项目里实现了一个简单的连续批处理调度器,核心思想是:把多个请求的token拼成一个batch,一起送进模型计算。

这里的关键问题是:不同请求的序列长度不一样,怎么拼batch?项目里用了padding的方式,把短序列补齐到最长序列的长度。但padding会浪费计算资源,所以更高级的做法是连续批处理(continuous batching),也就是在每次迭代时动态决定哪些请求参与计算。

项目里给出了两种实现的对比。padding版本实现简单但效率低,连续批处理版本效率高但实现复杂。我建议先理解padding版本,再去看连续批处理版本,否则容易一头雾水。

调度器的核心逻辑是维护一个请求队列,每次迭代时从队列中取出若干个请求,组成一个batch,送进模型计算,然后把生成的token返回给对应的请求。如果一个请求生成了结束符,就把它从队列中移除,腾出位置给新的请求。

这个部分的代码量比较大,但逻辑并不复杂。关键是理解请求的生命周期:从进入队列、到参与计算、到生成结束、到释放资源。

4. 完整实操流程:从零跑通一个推理服务

4.1 环境准备与依赖安装

项目对环境的依赖很少,核心就是Python和NumPy。如果你想跑CUDA版本的算子,还需要CUDA Toolkit和cuDNN。我建议先用CPU版本跑通整个流程,再切换到GPU版本。

# 创建虚拟环境 python -m venv venv source venv/bin/activate # Linux/Mac # venv\Scripts\activate # Windows # 安装基础依赖 pip install numpy pytest # 如果需要GPU支持 pip install torch --index-url https://download.pytorch.org/whl/cu121

项目里的代码组织很清晰,每个模块都有对应的测试文件。我建议按照以下顺序逐步跑通:

  1. 先跑test_tensor.py,确保张量抽象的正确性
  2. 再跑test_ops.py,验证算子实现的正确性
  3. 然后跑test_model.py,验证Transformer的前向传播
  4. 最后跑test_inference.py,验证完整的推理流程

每个测试文件都可以单独运行,方便定位问题。

4.2 张量模块的复现与验证

张量模块是整个项目的基础,我建议你亲手实现一遍,而不是直接看参考代码。实现过程中有几个关键点:

内存布局的选择。项目默认使用行优先(C order)布局,这也是NumPy的默认布局。但有些算子(比如矩阵乘法)在列优先布局下可能更快。项目里有一个实验对比了两种布局在不同算子上的性能差异,结论是:对于大多数场景,行优先就够了,不需要过度优化。

strides的计算。对于一个shape为(d0, d1, ..., dn)的张量,行优先布局下的strides是(d1*d2*...*dn, d2*...*dn, ..., 1)。这个计算看起来简单,但在实现切片操作时很容易出错。我建议写一个辅助函数专门计算strides,然后所有地方都调用这个函数。

广播机制。广播是张量操作中很重要的一个特性,它允许不同shape的张量进行逐元素运算。实现广播的关键是:在计算strides时,对于需要广播的维度,把stride设为0。这样在索引时,该维度上的索引变化不会影响实际的内存地址。

验证方法:用NumPy的对应操作作为参照,逐元素对比结果。项目里的测试用例覆盖了各种边界情况,包括空张量、单元素张量、非连续张量等。

4.3 算子模块的复现与性能对比

算子模块是项目中最有实操价值的部分。我建议从矩阵乘法开始,逐步实现Softmax、LayerNorm、Attention等算子。

以矩阵乘法为例,复现步骤是:

  1. 实现朴素版本,验证正确性
  2. 实现循环重排版本,对比性能
  3. 实现分块版本,对比性能
  4. 实现向量化版本,对比性能

每一步都要用timeit或者perf_counter测量运行时间,记录数据。我实测的结果是:

实现版本512x512矩阵乘法耗时(ms)相对加速比
朴素三重循环12501.0x
循环重排4203.0x
分块(64x64)8514.7x
向量化(SIMD)3239.1x

这个数据是在我的笔记本上跑的(Intel i7-12700H),你的结果可能不同,但趋势应该是一致的。关键是要理解每一步优化背后的原理,而不是记住具体的数字。

Softmax的实现要注意数值稳定性。正确的做法是先减去每行的最大值,再取指数,最后归一化。项目里有一个对比实验:不做数值稳定处理的版本在输入值较大时会溢出,做处理的版本则始终稳定。

4.4 推理服务的搭建与测试

把所有模块组装起来之后,就可以搭建一个完整的推理服务了。项目里用FastAPI做了一个简单的HTTP服务,核心接口是/generate,接收一个prompt,返回生成的文本。

服务端的核心逻辑是:

@app.post("/generate") async def generate(request: GenerateRequest): input_ids = tokenizer.encode(request.prompt) output_ids = model.generate( input_ids, max_new_tokens=request.max_new_tokens, temperature=request.temperature, ) return {"text": tokenizer.decode(output_ids)}

这个服务虽然简单,但涵盖了推理服务的核心流程:tokenization、模型推理、detokenization。实际生产中还需要考虑并发控制、超时处理、错误恢复等,但基础版本已经足够让你理解整个链路。

测试方法:用curl或者requests发送请求,观察响应时间和生成质量。我建议先用短prompt测试,确认基本功能正常,再逐步增加prompt长度和并发数,观察性能变化。

提示:如果你在本地跑,建议把max_new_tokens设小一点(比如32),否则生成时间会很长。等确认流程跑通之后,再逐步增加。

5. 常见问题与排查技巧实录

5.1 张量操作中的典型错误

问题一:非连续张量导致的隐式拷贝。当你对一个转置后的张量做reshape时,NumPy会自动触发拷贝,因为转置后的张量在内存中不连续。这个拷贝是隐式的,你从代码上看不出来,但性能会受影响。

排查方法:用tensor.flags['C_CONTIGUOUS']检查张量是否连续。如果不连续,考虑先用contiguous()显式拷贝,或者调整操作顺序避免拷贝。

问题二:广播导致的意外结果。广播很方便,但也容易出错。比如一个shape为(3, 1)的张量和一个shape为(3,)的张量相加,结果shape是(3, 3),而不是你期望的(3,)。这个错误在调试时很难发现,因为代码不会报错,只是结果不对。

排查方法:在关键操作前后打印shape,确保符合预期。项目里的测试用例覆盖了各种广播场景,建议仔细看一遍。

5.2 算子性能不达预期的排查思路

第一步:确认瓶颈在哪里。用cProfile或者line_profiler定位耗时最长的函数。很多时候你以为的瓶颈并不是真正的瓶颈。

第二步:检查内存访问模式。算子的性能很大程度上取决于内存访问是否连续。如果访问模式是跳跃的,缓存命中率会很低,性能自然上不去。

第三步:检查是否触发了隐式拷贝。如前所述,非连续张量的操作会触发拷贝。用np.shares_memory检查两个张量是否共享内存。

第四步:对比参考实现。项目里每个算子都有参考实现,你可以用自己的实现和参考实现做对比,看看差距在哪里。

5.3 推理服务中的常见故障

故障一:显存溢出。最常见的原因就是KV Cache太大。解决方法:减小max_seq_len,或者使用分页KV Cache,或者做量化。

故障二:生成结果乱码。通常是tokenizer的问题。检查tokenizer的编码和解码是否对称,特别是特殊token的处理。

故障三:延迟波动大。可能是批处理调度的问题。检查是否有长序列请求阻塞了短序列请求。连续批处理可以缓解这个问题。

故障四:并发请求下性能急剧下降。可能是锁竞争或者内存分配竞争。检查是否有全局锁,或者是否频繁分配释放内存。

5.4 常见问题速查表

问题现象可能原因排查方法解决方案
张量reshape报错张量非连续检查flags先contiguous()
算子结果不对广播shape不匹配打印shape调整维度
算子性能差内存访问不连续检查strides调整数据布局
显存溢出KV Cache过大计算缓存大小减小seq_len或量化
生成乱码tokenizer不对称检查编解码修正特殊token处理
延迟波动批处理调度不合理分析请求分布改用连续批处理

6. 我踩过的坑和给你的建议

第一个坑是过早优化。我一开始就想着要把矩阵乘法写到最快,花了很多时间研究SIMD指令和CUDA优化,结果发现整个推理流程的瓶颈根本不在矩阵乘法上,而在KV Cache的内存管理上。后来我调整了策略:先用最朴素的实现跑通整个流程,然后用profiler定位真正的瓶颈,再针对性地优化。这个顺序很重要。

第二个坑是忽略数值精度。我在实现Softmax的时候,用fp16做中间计算,结果在某些输入下出现了NaN。后来改成fp32做Softmax,再转回fp16,问题就解决了。这个经验告诉我:不是所有操作都适合用低精度,特别是涉及指数和对数的操作。

第三个坑是低估了批处理的复杂度。我一开始觉得批处理就是把多个请求拼在一起,能有多难?实际实现的时候才发现,请求的到达时间不同、序列长度不同、生成结束时间不同,这些都需要仔细处理。项目里的连续批处理调度器我看了三遍才完全理解。

如果你要跟着这个项目学,我的建议是:不要只看代码,一定要自己动手写。看代码的时候你觉得都懂了,但自己写的时候会发现各种细节问题。特别是张量抽象和算子实现这两个模块,亲手写一遍的收获远大于看十遍。

另外,项目里的测试用例是很好的学习资源。每个测试用例都针对一个特定的功能点,你可以先看测试用例,理解期望的行为,然后自己实现,最后跑测试验证。这种“测试驱动”的学习方式效率很高。

最后分享一个小技巧:项目里有很多性能对比的实验,我建议你把每次实验的数据都记录下来,形成一个自己的性能数据库。这样当你遇到新的优化问题时,可以快速判断某个优化手段是否值得尝试。比如我现在就知道,对于小于128x128的矩阵,分块优化的收益很小,不值得花时间;但对于大于512x512的矩阵,分块优化的收益就很明显了。这种经验性的判断,只有通过大量实验才能积累起来。

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询