从零手搓AI工程:手写推理引擎与动态组批实战
1. 从零手搓AI工程为什么我不建议你直接调包很多人一听到“AI工程”这四个字第一反应就是打开某个云平台拖几个组件调一下API然后跑通一个Demo就觉得自己已经掌握了。我刚开始也是这么想的直到有一次线上推理服务在高峰期直接雪崩日志里全是显存溢出和请求超时我才意识到——只会调包的人根本不知道模型在底层到底经历了什么。ai-engineering-from-scratch这个项目标题核心不是教你如何快速拼凑一个能跑的东西而是逼着你从最原始的矩阵乘法开始把整个AI工程链路亲手搭一遍。它适合那些已经会用框架、但总觉得心里没底的中级开发者也适合刚入行、想真正理解推理服务内部构造的新人。这篇文章我会把从零构建AI工程的关键环节拆开包括张量内存布局、算子手写、推理引擎调度、服务化封装和性能压测全部基于我在实际项目中踩过的坑和验证过的方案。你不需要有编译器背景但需要愿意动手写代码而不是只复制粘贴。2. 为什么“从零”这件事在AI工程里越来越稀缺2.1 框架封装带来的认知断层现在的深度学习框架确实把门槛降到了地板以下。model.fit()一调训练就跑起来了model.predict()一调推理就出结果了。但问题在于当服务出现延迟抖动、吞吐量上不去、显存莫名其妙涨了又降不下来的时候你打开框架源码发现里面是几千行的C和CUDA交织的调度逻辑根本无从下手。我见过太多团队模型效果很好但一上线就崩最后只能靠加机器硬扛。这就是认知断层的代价——你知道怎么用工具但不知道工具在背后做了什么。ai-engineering-from-scratch的价值就在于它强迫你回到没有框架的时代用最基础的数据结构和循环把前向传播、反向传播、参数更新全部写一遍。这个过程痛苦但一旦走通你再回头看那些框架API就能一眼看出它在哪个环节做了优化、哪个环节埋了坑。2.2 从零实现能暴露哪些真实问题我举一个最典型的例子内存对齐。当你手写一个矩阵乘法算子时如果输入张量的内存布局不是连续的或者步长设置不对性能可能直接差十倍。这个问题在框架里被自动处理了你根本感知不到。但一旦你要自己写一个自定义算子或者做模型量化内存布局就是绕不过去的坎。再比如推理时的批处理调度。框架通常给你一个batch_size参数但实际服务中请求是流式到达的如何动态组批、如何设置超时等待、如何在延迟和吞吐之间做权衡这些策略框架不会告诉你只能自己从零设计。我在一个实时推荐场景里就因为组批策略没设计好导致P99延迟从50毫秒飙到800毫秒。后来把组批逻辑重写了一遍才把延迟压回去。这些经验只有从零构建过的人才会真正重视。2.3 从零不等于重复造轮子这里要澄清一个误区从零构建AI工程不是让你抛弃所有框架自己写一个PyTorch出来。那是另一个层面的工作对绝大多数人没有意义。ai-engineering-from-scratch的“从零”是指你要理解每一层的输入输出、内存变化、计算代价并且能够用最简化的代码复现核心逻辑。比如你可以用NumPy实现一个完整的Transformer前向传播不需要考虑GPU加速只需要把注意力机制、层归一化、残差连接的计算过程写清楚。这样做的目的是建立直觉而不是替代生产工具。等你有了这个直觉再去用TensorRT或者ONNX Runtime做部署就能准确判断哪些优化是有效的哪些参数调整是徒劳的。3. 手写推理引擎从张量定义到算子调度3.1 张量类的设计不只是存数据从零开始的第一步是定义一个自己的张量类。很多人觉得张量就是一个多维数组用NumPy的ndarray就够了。但在AI工程里张量还需要携带额外信息数据类型、设备位置、是否需要梯度、内存是否连续、步长是多少。我最初的设计只存了数据和形状结果在做转置和切片的时候性能直接崩了。后来参考了主流框架的设计把步长stride加进去才解决了问题。具体来说一个形状为(2, 3, 4)的张量如果内存是连续存储的它的步长就是(12, 4, 1)。当你做转置操作时不需要真正移动数据只需要交换步长即可。这个设计在后续实现矩阵乘法时非常关键因为BLAS库对内存连续性有要求步长不对就得先做一次拷贝代价很高。class Tensor: def __init__(self, data, shapeNone, strideNone, dtypenp.float32): self.data np.asarray(data, dtypedtype) self.shape shape if shape else self.data.shape if stride is None: self.stride self._compute_stride(self.shape) else: self.stride stride self.offset 0 def _compute_stride(self, shape): stride [1] * len(shape) for i in range(len(shape) - 2, -1, -1): stride[i] stride[i 1] * shape[i 1] return tuple(stride)这个类看起来简单但它决定了后续所有算子的实现方式。比如切片操作只需要调整offset和shape不需要拷贝数据。转置操作只需要反转shape和stride。这些设计在框架里都是基础但自己写一遍感受完全不同。3.2 矩阵乘法的三种实现与性能对比矩阵乘法是AI工程里最核心的算子没有之一。全连接层、注意力机制、卷积展开底层都是矩阵乘法。我从零实现了三个版本朴素三重循环、NumPy的dot、以及分块乘法。朴素三重循环在(512, 512)的矩阵上跑了将近3秒NumPy的dot只要0.8毫秒分块乘法在块大小设为64时能跑到1.2毫秒左右。虽然NumPy底层用了BLAS但分块乘法让我理解了为什么缓存友好性这么重要。具体来说朴素循环每次取一个元素内存访问模式是跳跃的缓存命中率极低。分块乘法把大矩阵切成小块每个小块能完整放进L1缓存计算完再换下一块缓存命中率大幅提升。这个实验让我在后来的推理优化中特别关注算子的内存访问模式而不是只看浮点运算次数。实现方式矩阵大小耗时相对加速比朴素三重循环512x5122.8s1xNumPy dot512x5120.8ms3500x分块乘法(块64)512x5121.2ms2333x注意分块乘法虽然比NumPy慢但它的意义在于让你理解缓存的作用。实际生产中直接用BLAS库即可不需要自己写。3.3 算子调度的依赖分析与执行顺序当你手写多个算子后下一个问题就是如何组织它们的执行顺序在框架里计算图会自动处理依赖关系。但从零构建时你需要自己设计一个简单的调度器。我的做法是每个算子记录它的输入张量和输出张量调度器根据张量的引用关系构建一个有向无环图然后做拓扑排序。这个过程中我遇到了一个典型问题原地操作in-place operation会破坏依赖关系。比如x x 1如果直接修改x的数据那么依赖x旧值的算子就会出错。解决方案是引入版本号机制每次修改张量时递增版本号调度器检查版本号是否匹配。这个机制在PyTorch里也有叫version_counter。自己实现一遍就能理解为什么框架要禁止某些原地操作以及为什么在推理时开启inference_mode能提升性能。4. 模型服务化把推理引擎包装成可用接口4.1 请求队列与动态组批的设计取舍推理引擎写好后下一步是把它变成一个服务。最朴素的做法是来一个请求跑一次推理。但这样吞吐量极低因为GPU利用率上不去。动态组批是解决这个问题的标准方案但组批策略有很多细节。我试过三种策略固定超时组批、固定批量组批、以及自适应组批。固定超时组批是设置一个最大等待时间比如10毫秒超时或者凑够最大批量就触发推理。这个策略实现简单但在低流量时延迟高高流量时批量又容易超限。自适应组批是根据当前队列长度动态调整等待时间队列长就少等队列短就多等。我在一个图像分类服务里用了自适应组批P99延迟降低了40%吞吐量提升了2.3倍。具体参数是最小批量4最大批量32基础等待时间5毫秒队列长度超过16时等待时间降为1毫秒。class BatchScheduler: def __init__(self, min_batch4, max_batch32, base_wait0.005): self.min_batch min_batch self.max_batch max_batch self.base_wait base_wait self.queue [] def add_request(self, request): self.queue.append(request) if len(self.queue) self.max_batch: return self._flush() return None def _flush(self): batch self.queue[:self.max_batch] self.queue self.queue[self.max_batch:] return batch def get_wait_time(self): if len(self.queue) 16: return 0.001 return self.base_wait4.2 序列化与传输格式的选择服务化绕不开序列化。我对比过JSON、MessagePack、Protobuf和Arrow四种格式。JSON可读性最好但序列化一个(1, 3, 224, 224)的浮点张量大小约600KB序列化耗时约15毫秒。MessagePack大小降到450KB耗时8毫秒。Protobuf需要预先定义schema大小约400KB耗时5毫秒。Arrow是列式存储大小约380KB耗时3毫秒而且支持零拷贝读取。最终我选了Arrow作为内部传输格式因为它在批量传输时优势明显。但对外接口还是保留了JSON方便调试和兼容。这里有个坑浮点数的精度问题。JSON默认用双精度但模型输入通常是单精度序列化和反序列化过程中会出现微小的数值差异导致推理结果不一致。解决方案是在序列化时显式指定单精度或者用二进制格式直接传字节流。4.3 健康检查与优雅退出的实现细节服务上线后健康检查和优雅退出是必须的。健康检查不能只返回一个200 OK还要检查推理引擎是否正常、显存是否充足、队列是否积压。我的做法是暴露一个/health接口返回当前队列长度、平均推理耗时、GPU显存使用率。如果队列长度超过阈值或者显存使用率超过90%就返回503让负载均衡器把流量切走。优雅退出更关键收到终止信号后不能直接杀进程要先把队列里的请求处理完再关闭推理引擎最后释放显存。我见过一个服务因为直接kill -9导致GPU显存没释放后续服务启动时直接报显存不足。后来加了信号处理等待所有进行中的推理完成再退出问题才解决。5. 性能压测与瓶颈定位从数据出发5.1 压测工具的选择与脚本编写压测不是随便跑个ab或者wrk就完事了。AI推理服务的压测需要模拟真实请求输入张量的形状要多样请求到达要符合泊松分布还要能统计P50、P90、P99延迟。我用Locust写了一个压测脚本自定义了客户端直接发送Arrow格式的张量数据。脚本里设置了三个场景低负载10 QPS、中负载100 QPS、高负载500 QPS每个场景跑5分钟。结果发现低负载时P99延迟只有12毫秒中负载时涨到45毫秒高负载时直接飙到800毫秒而且错误率超过5%。这个数据说明服务在中高负载之间存在一个性能悬崖必须找到瓶颈点。5.2 瓶颈定位CPU、GPU还是IO定位瓶颈的第一步是看监控。我用nvidia-smi看GPU利用率发现高负载时GPU利用率只有40%说明GPU不是瓶颈。然后用py-spy抓了CPU火焰图发现大量时间花在数据预处理上——具体来说是图像解码和归一化。原来我的服务是在Python层面做预处理GIL锁导致多线程无法并行。解决方案是把预处理移到C扩展里或者用torchvision的decode_jpeg直接输出张量。改完之后GPU利用率升到75%P99延迟降到120毫秒。第二步是看IO发现Arrow的反序列化在批量大时耗时明显后来改成零拷贝读取又省了10毫秒。第三步是看推理引擎本身发现矩阵乘法的分块大小没调优默认块大小64在(256, 768)的矩阵上不是最优改成128后单次推理耗时降低了15%。瓶颈环节定位工具优化前优化后数据预处理py-spy320ms45ms反序列化手动计时25ms8ms矩阵乘法基准测试18ms15ms组批调度日志分析200ms60ms5.3 压测中发现的三个反直觉现象第一个现象批量越大单样本延迟越低但P99延迟反而越高。原因是大批量会导致排队时间增加虽然后续计算快了但前面的请求等得太久。第二个现象GPU利用率高不代表性能好。有时候GPU利用率100%但吞吐量上不去因为计算单元在等内存。第三个现象增加工作线程数不一定提升吞吐量。在Python里由于GIL的存在多线程对CPU密集型任务无效必须用多进程。我把工作进程从1个加到4个吞吐量提升了3.2倍但再加到8个吞吐量反而下降了因为进程间通信和显存拷贝的开销上来了。最终稳定在4个进程每个进程绑定一个GPU流。6. 从零构建后的认知升级与后续扩展走完这一整套从零构建的流程我最大的感受是以前看框架文档觉得那些参数都是魔法数字现在看每一个都有明确的物理意义。比如num_workers对应数据加载的并行度pin_memory对应是否使用锁页内存加速拷贝batch_size对应组批策略的上限。这些理解不是看书能看来的必须自己踩一遍坑。后续如果想继续深入有三个方向一是把推理引擎用CUDA重写理解线程束和共享内存二是引入量化把FP32模型转成INT8观察精度和速度的权衡三是做分布式推理把模型切到多张卡上处理跨卡通信。每个方向都能写一篇独立的文章。我在实际项目里就是从单卡推理开始逐步做到多卡并行中间因为张量并行和流水线并行的选择纠结了很久最后根据模型层数和通信开销选了流水线并行。这些决策没有标准答案只有根据具体场景做权衡。如果你也在做类似的事情建议先把单机单卡跑通把延迟和吞吐的基线测出来再考虑扩展。否则分布式只会让问题更复杂。