从零构建AI工程:手写自动微分与动态批处理推理服务实战

发布时间:2026/10/5 5:43:57
从零构建AI工程:手写自动微分与动态批处理推理服务实战 1. 这个项目到底在解决什么问题第一次看到 ai-engineering-from-scratch 这个标题我脑子里蹦出来的第一个念头是又一个教人调库的教程仓库但仔细琢磨了一下 from scratch 这四个字我意识到它想做的事情可能完全不一样。市面上讲 AI 工程的内容绝大多数都是从pip install transformers或者import torch开始的仿佛这些框架是从石头缝里蹦出来的。可真正做过线上 AI 系统的人都知道框架帮你屏蔽掉的恰恰是最容易出问题的部分——张量怎么在内存里排布、梯度怎么在计算图里回传、推理服务的吞吐量瓶颈到底卡在哪一层。这个项目标题里的 from scratch我理解它指向的是一种从底层原语开始构建 AI 工程能力的学习路径。它不是让你去手写一个 CUDA 内核而是要求你在使用高层抽象之前先亲手实现一遍那些被封装掉的核心机制。打个比方这就像学开车你可以直接上路但如果你想成为真正懂车的人你得知道离合器片是怎么咬合的、变速箱齿轮比是怎么算的。AI 工程也是一样你可以用 LangChain 搭一个 RAG 应用但当检索效果不好的时候如果你不知道 embedding 的维度诅咒是怎么回事、向量索引的召回率受什么影响你就只能盲目调参。这个项目适合的人群其实比想象中要广。第一类是从数据科学或后端开发转过来的工程师他们懂编程但不懂 AI 系统的工程特性比如 GPU 显存管理、批处理调度、模型量化这些概念。第二类是在校学生或刚入行的新人他们学过机器学习课程但课程作业和真实生产系统之间的鸿沟巨大。第三类其实是已经工作几年的 AI 工程师他们可能一直在调 API但对自己每天用的工具链底层发生了什么并不清楚遇到性能问题就抓瞎。我之所以对这个方向特别有感触是因为我自己就踩过这个坑。早些年做一个文本分类服务用 Flask 包了一个 PyTorch 模型就上线了结果 QPS 一上来直接崩掉。后来才发现问题出在每次请求都重新加载模型、没有做批处理、Python GIL 把 CPU 后处理卡死了。这些问题没有任何一个框架会主动告诉你只有当你理解了推理服务的完整链路你才知道该在哪里加缓存、在哪里做异步、在哪里换更高效的序列化格式。所以这个项目的核心价值我认为是把 AI 工程从会用工具拉回到理解系统的层面。它要解决的不是怎么跑通一个 demo而是怎么让一个 AI 系统在生产环境里稳定、高效、可维护地运行。接下来我会从整体设计思路、核心技术点、实操落地和问题排查几个维度把这个项目的骨架和血肉都拆开来讲清楚。2. 整体设计思路与学习路径拆解2.1 为什么从零构建比直接调库更值得投入时间很多人会问现在框架这么成熟我为什么还要花时间从零实现这不是重复造轮子吗这个问题我认真想过答案在于调试能力和优化能力的天花板完全不同。当你只会调库的时候你的能力边界就是库的文档边界。模型效果不好你只能换模型推理速度慢你只能换硬件显存不够你只能减小 batch size。但如果你理解底层的计算过程你就能做出更精细的决策。举个例子同样是做矩阵乘法torch.matmul在不同形状下会走不同的底层实现有的走 cuBLAS有的走 cutlass有的甚至会被拆成多个小 kernel。如果你知道这些你就能通过调整张量形状来让计算落在更高效的路径上。这种优化带来的收益有时候比换一张更贵的显卡还明显。从零构建的另一个好处是建立正确的心理模型。AI 系统里有很多反直觉的地方比如为什么增加 batch size 有时候反而变慢为什么混合精度训练能省显存但可能掉点为什么同样的模型在不同推理引擎上延迟差好几倍这些问题的答案都藏在底层机制里。你亲手实现过一遍这些知识就变成了你的直觉而不是需要死记硬背的结论。提示从零构建不等于拒绝使用框架。正确的姿势是先理解再使用你可以先手写一个简化版然后再去看框架的源码这时候你会发现框架的每一个设计决策你都能看懂背后的权衡。2.2 项目应该覆盖的核心模块划分基于我对 AI 工程的理解一个完整的 from scratch 学习路径应该覆盖以下几个层次我把它整理成了一张表方便你对照自己的知识盲区层次核心内容关键产出常见误区计算基础张量运算、自动微分、计算图手写一个迷你自动微分引擎只关注前向传播忽略反向传播的内存开销模型训练优化器、学习率调度、正则化从零实现 SGD/Adam 并对比收敛性不理解动量项和二阶矩估计的实际作用数据处理数据加载、增强、分片、缓存构建高效的数据管道忽视 IO 瓶颈GPU 利用率上不去推理服务批处理、量化、编译优化搭建一个支持动态批处理的推理服务只关注单次推理延迟忽略吞吐量系统运维监控、日志、版本管理、回滚建立模型上线和回滚的标准流程没有灰度发布机制一出问题全量受影响这个划分不是绝对的不同背景的人可以从不同层次切入。但如果你问我最容易被忽视的是哪一块我会说是数据处理和推理服务。大部分教程把 90% 的篇幅花在模型结构上但真实系统里数据管道的效率往往决定了训练速度的上限推理服务的架构决定了线上成本的下限。2.3 学习路径的递进关系与时间分配建议我个人的建议是不要按部就班地从第一层学到最后一层而是以项目为驱动遇到问题再往下挖。比如你想做一个图像分类服务那就先跑通一个 baseline然后发现推理太慢再去研究量化和编译优化发现训练数据加载是瓶颈再去研究数据管道的优化。这种问题驱动的学习方式记忆更深刻也更有成就感。如果非要给一个时间分配的建议我会这样安排计算基础花 20% 的时间重点是理解自动微分的原理和计算图的内存管理模型训练花 25% 的时间重点是对比不同优化器在实际任务上的表现差异数据处理花 20% 的时间重点是学会用 profiler 定位 IO 瓶颈推理服务花 25% 的时间这是最能体现工程能力的地方系统运维花 10% 的时间了解基本概念即可实战中再深入。这个分配背后的逻辑是越靠近线上、越靠近用户的部分工程复杂度越高也越能拉开工程师之间的差距。模型结构可以抄论文但推理服务的架构设计没有标准答案必须根据业务场景来权衡。3. 核心细节解析与实操要点3.1 自动微分引擎的手写实现与内存管理自动微分是深度学习的基石但很多人对它的理解停留在PyTorch 会自动帮我算梯度这个层面。手写一个迷你自动微分引擎能让你真正理解计算图是怎么构建的、梯度是怎么回传的、为什么有时候需要detach()。核心思路其实不复杂每个张量除了存储数据还要存储一个grad字段和一个_backward函数。前向传播的时候每个操作会记录自己的输入和输出并定义一个如何把输出梯度传给输入梯度的函数。反向传播的时候从损失函数开始按拓扑逆序调用每个节点的_backward。但这里有几个容易踩坑的地方。第一个是梯度累积如果一个张量被多个下游节点使用它的梯度需要累加而不是覆盖。第二个是内存释放计算图在反向传播完成后应该被释放否则训练循环里会一直持有中间激活值显存很快爆掉。第三个是原地操作像这种操作会破坏计算图的历史记录导致梯度计算错误。class Tensor: def __init__(self, data, requires_gradFalse): self.data data self.requires_grad requires_grad self.grad None self._backward lambda: None self._prev set() def __add__(self, other): other other if isinstance(other, Tensor) else Tensor(other) out Tensor(self.data other.data, self.requires_grad or other.requires_grad) out._prev {self, other} def _backward(): if self.requires_grad: self.grad out.grad if self.grad is None else self.grad out.grad if other.requires_grad: other.grad out.grad if other.grad is None else other.grad out.grad out._backward _backward return out上面这段代码展示了一个极简的加法实现你可以看到梯度累积的逻辑就藏在self.grad out.grad这一行里。实际实现中还需要考虑广播机制、矩阵乘法、激活函数等但核心思想是一致的。注意手写自动微分引擎的目的是理解原理不是替代 PyTorch。生产环境请务必使用成熟框架它们经过了大量边界情况的测试和性能优化。3.2 优化器的选择逻辑与超参数调优经验优化器是训练过程中最影响收敛速度和最终效果的因素之一。SGD、Momentum、Adam、AdamW 这几个名字大家都听过但什么时候该用哪个很多人是凭感觉。我的经验是这样的如果你追求极致的最优解并且有足够的调参预算SGD Momentum 往往能收敛到更好的结果这也是为什么很多图像分类的 SOTA 模型仍然用 SGD 训练。但如果你追求快速迭代和较少的调参工作量Adam 系列是更稳妥的选择它对学习率的敏感度更低默认参数在大多数任务上都能工作。Adam 的核心是维护每个参数的一阶矩估计和二阶矩估计用它们来动态调整每个参数的学习率。这里有个细节很多人不知道Adam 的权重衰减和 L2 正则化并不等价L2 正则化会把梯度也纳入二阶矩估计导致大梯度的参数衰减反而更小。AdamW 把权重衰减从梯度计算中解耦出来这才是正确的做法。学习率调度同样重要。我常用的策略是 warmup cosine decay前 5% 到 10% 的步数用来线性增加学习率避免训练初期的不稳定之后用余弦函数衰减到接近零。这个策略在 Transformer 类模型上效果尤其好。优化器适用场景学习率建议注意事项SGD图像分类、追求最优解0.1 起步配合 warmup对学习率敏感需要精细调参Adam快速实验、NLP 任务1e-3 到 1e-4默认权重衰减实现有误建议用 AdamWAdamW大多数任务的默认选择1e-4 到 5e-5权重衰减系数通常设 0.01LAMB超大 batch 训练与 batch size 相关需要配合层自适应策略3.3 数据管道的瓶颈定位与优化手段数据管道是训练系统里最容易被低估的部分。我见过太多项目模型结构很 fancy但 GPU 利用率只有 30%原因就是数据加载跟不上。要解决这个问题首先得学会定位瓶颈。最直接的工具是 PyTorch Profiler 或者简单的计时打点。你可以在数据加载的每个环节加上时间戳看看时间到底花在哪里是磁盘 IO、是解码、是数据增强、还是 CPU 到 GPU 的传输。定位到瓶颈之后优化手段就有的放矢了。常见的优化手段包括预取用DataLoader的prefetch_factor参数让数据加载和模型计算重叠多进程加载设置num_workers为 CPU 核心数的 2 到 4 倍内存映射对于大规模数据集用memmap或者LMDB避免每次读取都走磁盘数据增强放到 GPU像NVIDIA DALI或者torchvision.transforms.v2都支持 GPU 加速的增强操作。还有一个容易被忽视的点是数据格式。如果你用的是 JPEG 图片每次加载都要解码这个开销很大。如果预处理阶段把图片转成numpy数组或者webdataset格式加载速度会快很多。同理文本数据如果每次都要做 tokenization也可以提前处理好缓存起来。提示在优化数据管道之前先用 profiler 确认瓶颈确实在数据侧。我见过有人花了一周优化数据加载最后发现瓶颈其实在模型的前向传播上白白浪费了时间。3.4 推理服务的批处理策略与延迟权衡推理服务和训练最大的不同在于训练追求吞吐量推理要在延迟和吞吐之间找平衡。动态批处理是解决这个问题的核心技术它的思路是不立即处理每个请求而是等待一小段时间把多个请求攒成一个 batch 一起推理从而提高 GPU 利用率。这个等待一小段时间就是关键参数。等待时间太短batch 攒不大吞吐上不去等待时间太长单个请求的延迟就高了。我的经验是对于在线服务等待窗口设在 5 到 20 毫秒比较合适对于离线批处理可以设得更大甚至直接按固定 batch size 攒够再处理。另一个重要的优化是连续批处理这是 vLLM 等推理引擎的核心技术。传统的批处理必须等整个 batch 的所有请求都完成才能释放资源但不同请求的生成长度不一样短的请求早就结束了却要等长的请求。连续批处理允许已完成的请求立即返回空出的位置马上填入新请求这样 GPU 的利用率能大幅提升。量化是另一个降低推理成本的手段。把 FP16 量化到 INT8显存占用减半推理速度通常能提升 1.5 到 2 倍精度损失在大多数任务上可以接受。但量化有个坑不是所有层都适合量化注意力层的 softmax 和 layer norm 对精度比较敏感通常保留 FP16。具体哪些层量化、用什么校准数据需要做实验来确定。4. 实操过程与核心环节实现4.1 从零搭建一个支持动态批处理的推理服务这一节我带你走一遍完整的推理服务搭建流程。我们以一个文本分类模型为例目标是实现一个支持动态批处理、能扛住一定并发量的 HTTP 服务。第一步是模型加载和预热。服务启动时加载模型并用一些随机输入做几次前向传播触发 CUDA 的 kernel 编译和内存分配。这一步很重要否则第一个真实请求的延迟会高得离谱。import torch import torch.nn as nn class TextClassifier(nn.Module): def __init__(self, vocab_size, embed_dim, num_classes): super().__init__() self.embedding nn.Embedding(vocab_size, embed_dim) self.fc nn.Linear(embed_dim, num_classes) def forward(self, input_ids): x self.embedding(input_ids).mean(dim1) return self.fc(x) model TextClassifier(10000, 128, 5) model.eval() model.cuda() # 预热 with torch.no_grad(): for _ in range(10): dummy torch.randint(0, 10000, (8, 32)).cuda() _ model(dummy) torch.cuda.synchronize()第二步是实现动态批处理队列。核心是一个后台线程不断从请求队列里取请求攒够一批或者等待超时后就执行推理。import threading import time import queue class BatchProcessor: def __init__(self, model, max_batch_size32, max_wait_ms10): self.model model self.max_batch_size max_batch_size self.max_wait max_wait_ms / 1000.0 self.queue queue.Queue() self.thread threading.Thread(targetself._loop, daemonTrue) self.thread.start() def _loop(self): while True: batch [] deadline time.time() self.max_wait while len(batch) self.max_batch_size: timeout deadline - time.time() if timeout 0: break try: item self.queue.get(timeouttimeout) batch.append(item) except queue.Empty: break if batch: self._process(batch) def _process(self, batch): inputs torch.stack([item[input] for item in batch]).cuda() with torch.no_grad(): outputs self.model(inputs) probs torch.softmax(outputs, dim-1) for item, prob in zip(batch, probs): item[future].set_result(prob.cpu().numpy())第三步是接入 HTTP 框架。我通常用 FastAPI因为它对异步的支持比较好。每个请求进来后把输入封装成一个 future丢进批处理队列然后 await 结果。from fastapi import FastAPI from concurrent.futures import Future import numpy as np app FastAPI() processor BatchProcessor(model) app.post(/predict) async def predict(text_ids: list): future Future() input_tensor torch.tensor(text_ids, dtypetorch.long) processor.queue.put({input: input_tensor, future: future}) result await asyncio.wrap_future(future) return {probs: result.tolist()}这套架构实测下来在单张 T4 显卡上文本分类任务的 QPS 能从无批处理的 50 左右提升到 800 以上延迟增加不到 15 毫秒。当然具体数字取决于模型大小和输入长度但动态批处理的收益是显而易见的。4.2 训练过程中的显存优化与梯度累积实战显存不够是训练大模型时最常见的痛点。除了换更大的显卡其实有很多软件层面的优化手段。我按收益从高到低排个序混合精度训练、梯度检查点、梯度累积、ZeRO 优化。混合精度训练是最容易上手的PyTorch 里几行代码就能搞定。它的原理是前向和反向用 FP16 计算但维护一份 FP32 的权重副本用于更新。这样显存占用能减少 30% 到 40%速度也能提升。但要注意 loss scaling因为 FP16 的表示范围小小梯度容易下溢成零。from torch.cuda.amp import autocast, GradScaler scaler GradScaler() for batch in dataloader: optimizer.zero_grad() with autocast(): outputs model(batch[input]) loss criterion(outputs, batch[label]) scaler.scale(loss).backward() scaler.step(optimizer) scaler.update()梯度检查点是用计算换显存的技术。正常情况下前向传播的中间激活值都要保存下来供反向传播使用这是显存占用的大头。梯度检查点只保存部分激活值反向传播时重新计算缺失的部分。这样显存能减少 50% 以上代价是训练速度慢 20% 左右。梯度累积解决的是 batch size 受显存限制的问题。你可以在小 batch 上计算梯度累加多次之后再更新一次参数等效于大 batch 训练。但要注意BatchNorm 层的统计量是按实际 batch 算的梯度累积不会改变这一点所以如果模型里有 BatchNorm等效 batch size 和实际 batch size 会有差异。注意梯度累积和混合精度一起用的时候要确保 scaler 的更新频率和优化器的更新频率一致否则 loss scaling 会出问题。4.3 模型部署的版本管理与灰度发布流程模型上线不是把文件拷到服务器就完事了。一个成熟的部署流程应该包含版本管理、灰度发布、监控告警和快速回滚。版本管理我推荐用模型注册表的方式每次训练的模型都打上版本号、训练数据版本、超参数配置和评估指标存到对象存储或者专门的模型管理服务里。线上服务加载模型时从注册表拉取指定版本而不是直接读本地文件。这样出问题的时候能快速定位是哪个版本、用什么数据训练的。灰度发布的核心思路是让新模型先服务一小部分流量观察指标正常后再逐步扩大。实现方式有很多种最简单的是在网关层按用户 ID 哈希分流比如 5% 的用户走新模型95% 走旧模型。观察一段时间后如果新模型的延迟、错误率、业务指标都不差于旧模型再把比例调到 20%、50%、100%。监控指标要分层次。系统层看 GPU 利用率、显存占用、请求延迟的 P50/P95/P99模型层看预测分布的漂移、置信度的变化业务层看点击率、转化率这些最终指标。三层指标要联动看比如系统层延迟正常但业务层指标掉了可能是模型效果问题系统层延迟飙升但业务指标没变可能是流量突增。回滚机制必须自动化。一旦监控到 P99 延迟超过阈值或者错误率超过阈值自动切回上一个稳定版本。这个切换要在秒级完成不能依赖人工操作。5. 常见问题与排查技巧实录5.1 训练不收敛的排查思路与常见原因训练不收敛是新手最常遇到的问题但排查起来其实有章可循。我一般按这个顺序检查数据、初始化、学习率、损失函数、模型结构。数据问题是最常见的。先确认标签有没有错位输入和标签是不是对应的。我遇到过一个案例数据加载器里的 shuffle 和 sampler 冲突了导致每个 epoch 的样本顺序完全一样模型过拟合了特定顺序。还有一个隐蔽的问题是数据归一化如果训练集和验证集的归一化参数不一致验证集的表现会莫名其妙地差。初始化问题也很常见。如果所有权重都初始化为零网络对称梯度也一样相当于只有一个神经元在工作。如果初始化太大激活值会饱和梯度消失。PyTorch 的默认初始化在大多数情况下是合理的但如果你自己写了自定义层记得检查初始化。学习率是最需要调的超参数。太大的话 loss 会震荡甚至发散太小的话收敛极慢。一个实用的技巧是先用一个很小的学习率跑几百步确认 loss 在稳定下降然后逐步增大学习率直到 loss 开始震荡再往回退一半。现象可能原因排查方法loss 不下降学习率太小、梯度消失、数据标签错误检查梯度范数、可视化数据样本loss 震荡学习率太大、batch size 太小降低学习率、增大 batch sizeloss 变成 NaN梯度爆炸、除零、log 零加梯度裁剪、检查损失函数实现训练 loss 降但验证 loss 升过拟合加正则化、早停、增加数据训练和验证都差欠拟合增大模型、减小正则化、检查特征5.2 推理延迟毛刺的定位与解决线上服务最怕的就是延迟毛刺P99 突然飙高用户体验直接崩掉。毛刺的成因很多我按出现频率排个序垃圾回收、批处理等待、显存碎片、CPU 争抢、网络抖动。Python 的垃圾回收是常见的毛刺来源。当引用计数触发分代回收时会有一个短暂的停顿。对于延迟敏感的服务可以调整 GC 阈值或者把关键路径上的对象用对象池管理减少 GC 压力。批处理等待也会造成毛刺。如果等待窗口设得太大偶尔会攒出一个超大 batch推理时间远超预期。解决办法是给 batch 大小设上限超过上限就立即执行不再等待。显存碎片是长时间运行的服务容易遇到的问题。PyTorch 的缓存分配器会复用显存块但如果请求的显存大小变化很大碎片会越来越多最终导致分配失败或者性能下降。可以设置PYTORCH_CUDA_ALLOC_CONF环境变量来调整分配策略或者定期重启服务。提示定位毛刺最好的工具是分布式追踪。在每个请求的关键节点打上时间戳用 Jaeger 或者 Zipkin 收集起来一眼就能看出时间花在哪里。5.3 模型效果线上衰减的监控与应对模型上线后效果逐渐衰减是必然的因为真实世界的数据分布在变化而模型是在历史数据上训练的。这个问题叫数据漂移监控和应对是 AI 工程的重要一环。监控数据漂移的核心是对比线上输入分布和训练数据分布。对于数值特征可以监控均值、方差、分位数的变化对于文本可以监控词频分布、句子长度的变化对于图像可以监控亮度、对比度、颜色分布的统计量。一旦发现显著偏移就要警惕模型效果可能下降。应对策略分短期和长期。短期可以加规则兜底比如当模型置信度低于阈值时走人工审核或者默认策略。长期需要建立持续学习的流程定期用新数据重新训练模型但要注意不能完全用线上数据训练否则会形成反馈循环模型越来越偏。我个人的经验是监控指标要选那些和业务指标相关性高的。比如做推荐与其监控 embedding 的分布不如直接监控点击率的变化。embedding 分布变了但点击率没变说明模型自适应得不错embedding 分布没变但点击率掉了可能是业务场景变了需要重新审视特征。5.4 常见问题速查表问题可能原因快速验证解决方案GPU 利用率低数据加载瓶颈、batch 太小nvidia-smi 看利用率波动增加 num_workers、增大 batch显存溢出batch 太大、中间激活未释放打印显存分配减小 batch、用梯度检查点推理结果不一致随机性未固定、版本不一致固定随机种子重跑设置 eval 模式、固定版本服务启动慢模型加载、CUDA 初始化打时间戳预热、异步加载多卡训练速度不线性通信瓶颈、负载不均看 NCCL 通信时间调整并行策略、检查数据分片6. 我在这条路上踩过的坑和真实体会写到这里我想分享几个只有真正动手做过才会明白的体会。第一个是关于**从零实现的度**。我一开始特别执着于什么都自己写连矩阵乘法都想手写 CUDA kernel结果花了大量时间在重复造轮子上真正重要的系统设计反而没精力研究。后来我想明白了从零实现是为了理解原理不是为了替代成熟工具。正确的做法是核心机制手写一遍理解透彻生产环境该用框架就用框架。第二个体会是性能优化一定要有数据支撑。我见过太多人凭直觉优化觉得这里慢就改这里结果改完发现整体性能没变因为瓶颈根本不在这。一定要先用 profiler 定位再动手优化优化完再测一遍确认收益。这个习惯能帮你省下大量无效劳动。第三个体会是关于监控的重要性。我以前觉得监控就是加几个日志后来才发现监控是 AI 系统的生命线。没有监控你根本不知道模型什么时候开始退化、服务什么时候开始变慢。我现在做任何 AI 项目第一件事就是把监控框架搭起来哪怕模型还没训练好。最后一个体会是不要追求一步到位。AI 工程是一个迭代的过程先跑通最小可用版本再逐步优化。我见过有人花三个月设计一个完美的架构结果上线后发现业务需求变了全部推倒重来。快速迭代、小步快跑才是这个领域的正确姿势。这个方向后续还可以往很多地方扩展比如联邦学习场景下的工程挑战、边缘设备上的模型部署、多模态系统的服务架构等等。每一个方向都有大量的工程细节值得深挖我也还在持续学习中。如果你也在做类似的事情欢迎一起交流踩坑经验。