
1. 为什么我要从零手搓一套AI工程流水线第一次看到ai-engineering-from-scratch这个标题我脑子里蹦出来的不是某个具体框架而是一堆碎片化的痛点。过去两年我帮不少团队做过模型落地发现一个特别普遍的现象算法同学在 Notebook 里跑出 0.92 的准确率兴冲冲交给工程团队结果上线后延迟飙到 800ms显存直接爆掉日志里全是看不懂的 CUDA OOM。问题出在哪不是模型不行是中间那层“工程化”的活儿没人系统性地干过。ai-engineering-from-scratch这个项目标题我理解它的核心诉求就是不依赖现成的高层封装从最底层的张量操作、数据管道、训练循环、推理服务一路搭到可观测性把AI工程的全链路亲手走一遍。它解决的不是“怎么调包”而是“为什么这么调、不这么调会死在哪”。适合谁看我觉得有三类人最该动手一是刚转AI方向的后端工程师懂服务但不懂模型二是算法出身但没碰过生产环境的同学三是想搞明白大模型推理背后到底发生了什么的技术负责人。我打算按我自己实际搭过的一套最小可行流水线来拆技术栈选 PyTorch FastAPI ONNX Runtime Prometheus不追求花哨追求每一行代码你都能说清楚它在干嘛。全文会围绕四个核心板块展开整体架构怎么切、数据与训练环节的硬核细节、推理服务的性能压榨、以及线上排查的实战记录。每个环节我都会给出可直接抄的配置和参数计算过程也会把踩过的坑原样倒出来。2. 整体架构设计与技术选型逻辑2.1 从需求反推架构分层动手之前先想清楚一件事AI工程和普通后端工程最大的区别在哪我的答案是不确定性。普通后端的输入输出是确定的AI工程的输入分布会漂移、模型输出会抖动、GPU显存会碎片化。所以架构设计的第一原则是隔离不确定性把易变的模型部分和稳定的服务部分切开。我最终落地的分层是这样的最底层是数据层负责样本读取、增强、分桶往上是训练层包含模型定义、损失计算、优化器调度再往上是导出层把训练好的权重转成推理友好的格式最顶层是服务层处理请求编排、批处理、监控。层与层之间通过明确的接口通信比如训练层只认Dataset和DataLoader服务层只认序列化后的InferenceSession。为什么这么切我试过把训练和推理揉在一个进程里结果每次改模型结构都要重启整个服务调试成本极高。分层之后模型迭代只影响导出层服务层可以独立灰度。这个决策背后的逻辑是变更频率对齐——变更频繁的模块应该独立部署。2.2 框架选型的三个硬指标选 PyTorch 而不是 TensorFlow不是因为它更流行而是因为三个硬指标动态图调试效率、ONNX 导出成熟度、社区算子覆盖。动态图让我能在forward里直接print中间张量的形状排查维度不匹配的问题时省掉大量时间。ONNX 导出这块PyTorch 的torch.onnx.export对控制流的支持虽然仍有坑但比 TF 的 SavedModel 转 ONNX 顺畅得多。推理侧我选 ONNX Runtime 而不是 TorchServe理由是依赖轻和跨平台。TorchServe 要拖一整套 Java 运行时镜像动辄 2GB 起而 ONNX Runtime 的 CPU 版只有几十 MBGPU 版也能控制在 500MB 以内。对于需要快速扩缩容的场景镜像拉取时间直接决定扩容速度。这里有个经验数据同样冷启动ONNX Runtime 镜像拉取加初始化大约 8 秒TorchServe 要 40 秒以上。监控选 Prometheus 而不是自己写日志统计是因为指标聚合的实时性。自己写日志再离线分析延迟至少分钟级而 Prometheus 的 pull 模型能做到 15 秒粒度。对于推理服务这种需要快速发现 P99 抖动的场景这个时间差很关键。2.3 目录结构约定我习惯用下面这种目录结构每个目录职责单一方便 CI 分阶段构建ai-engineering-from-scratch/ ├── data/ # 数据管道 │ ├── dataset.py │ └── transforms.py ├── train/ # 训练逻辑 │ ├── model.py │ ├── loop.py │ └── config.yaml ├── export/ # 模型导出 │ └── to_onnx.py ├── serve/ # 推理服务 │ ├── app.py │ ├── batcher.py │ └── metrics.py └── tests/ # 各层单测这个结构的好处是构建缓存友好。Docker 构建时data和train层变动频率低可以单独缓存serve层变动频繁放在最后构建。实测下来增量构建时间从 6 分钟压到 90 秒。3. 数据管道与训练循环的硬核细节3.1 Dataset 设计的三个陷阱写Dataset看起来简单但我在生产环境踩过三个大坑。第一个是在__getitem__里做重计算。很多人图省事把图像解码、归一化全塞进__getitem__结果 DataLoader 的 worker 成了瓶颈。正确做法是把能预计算的都预计算比如把图片统一 resize 后存成内存映射文件__getitem__只做索引读取。第二个坑是随机种子不隔离。DataLoader 多 worker 时如果增强操作直接用全局random每个 worker 的随机序列会重复导致增强多样性下降。我的做法是在worker_init_fn里给每个 worker 单独设种子def worker_init_fn(worker_id): seed torch.initial_seed() % 2**32 random.seed(seed worker_id) np.random.seed(seed worker_id)第三个坑是分桶策略缺失。变长序列如果不分桶padding 会浪费大量算力。我一般按长度分 8 到 16 个桶每个桶内长度差控制在 10% 以内。实测下来分桶后训练吞吐能提升 30% 到 50%具体取决于序列长度分布。3.2 训练循环里的显存账本显存管理是训练环节最考验功底的地方。我习惯在动手前先算一笔账模型参数 梯度 优化器状态 激活值。以 1.1 亿参数的模型为例FP32 下参数占 440MB梯度再占 440MBAdam 优化器要存一阶和二阶动量又是 880MB光这三项就 1.76GB。激活值取决于 batch size 和序列长度往往是大头。混合精度训练能把参数和激活的显存砍半但要注意损失缩放。我用的配置是初始 scale 为 65536每 2000 步如果没出现 inf 就翻倍出现 inf 就减半。这个动态策略比固定 scale 稳得多。另外梯度累积可以模拟大 batch但要注意 BatchNorm 的统计量会失真我的做法是梯度累积时把 BatchNorm 换成 GroupNorm。3.3 学习率调度的实战参数学习率调度不是玄学有明确的计算逻辑。我用的是带热启动的余弦退火公式是lr(t) lr_min 0.5 * (lr_max - lr_min) * (1 cos(pi * t / T))其中lr_max通过线性缩放规则确定lr_max base_lr * batch_size / 256。base_lr 一般取 3e-4如果 batch size 是 1024那 lr_max 就是 1.2e-3。热启动步数取总步数的 5%避免初期梯度爆炸。这里有个经验值warmup 期间不要用余弦用线性。我试过 warmup 直接上余弦前几百步 loss 抖动明显换成线性后曲线平滑很多。原因是余弦在起点附近导数接近零学习率上升太慢模型还没热起来。3.4 检查点策略与断点续训检查点不是简单存个state_dict就完事。我要求存四样东西模型参数、优化器状态、当前 epoch 和 step、以及随机数生成器状态。少了最后一样断点续训后数据顺序会变导致训练曲线不连续。存储频率也有讲究。我一般每 500 步存一次滚动检查点只保留最近 3 个每 5000 步存一个永久检查点。滚动检查点用于故障恢复永久检查点用于模型选择。存储路径用step_{step}_loss_{loss:.4f}.pt这种带指标的命名方便后续按 loss 排序找最优。4. 模型导出与推理服务的性能压榨4.1 ONNX 导出的动态轴配置导出 ONNX 最容易翻车的地方是动态轴没配对。如果输入序列长度是变的但导出时写死了推理时换个长度就报错。正确做法是在torch.onnx.export里显式指定dynamic_axestorch.onnx.export( model, dummy_input, model.onnx, input_names[input_ids, attention_mask], output_names[logits], dynamic_axes{ input_ids: {0: batch, 1: seq}, attention_mask: {0: batch, 1: seq}, logits: {0: batch, 1: seq} }, opset_version14 )opset 版本我选 14 而不是最新的因为 ONNX Runtime 对 14 的算子支持最全尤其是LayerNormalization和MultiHeadAttention的融合算子。用更高版本可能导出成功但推理时回退到慢速实现。4.2 动态批处理的实现细节推理服务的吞吐瓶颈往往不在计算而在请求调度。单个请求进来就推理一次GPU 利用率可能只有 20%。动态批处理的核心是攒一批请求一起算但攒多久、攒多大有讲究。我的实现是双阈值触发最大等待时间 10ms或最大 batch size 32谁先到就触发。10ms 是延迟和吞吐的平衡点实测 P99 延迟增加不到 15ms但吞吐提升 4 倍以上。batch size 上限 32 是因为再大显存增长非线性容易 OOM。批处理还有个坑是padding 对齐。同一批里序列长度不一要 pad 到最长。如果最长比平均长很多浪费严重。我的做法是先按长度排序再组批把长度接近的放一起。这个排序在批处理队列里做不影响请求顺序。4.3 量化与算子融合的收益ONNX Runtime 的量化能把 FP32 模型压到 INT8显存减半推理速度提升 2 到 3 倍。但量化不是无脑开敏感层要排除。我一般把 embedding 层和最后的分类头排除在量化外因为这两层对精度影响最大。量化配置如下from onnxruntime.quantization import quantize_dynamic, QuantType quantize_dynamic( model.onnx, model_int8.onnx, weight_typeQuantType.QInt8, nodes_to_exclude[embeddings, classifier] )算子融合是另一个免费加速。ONNX Runtime 的图优化会自动把MatMul Add Gelu融成一个算子减少内核启动开销。我实测融合后小 batch 场景延迟降低 20% 左右。要确认融合是否生效可以开session_options.optimized_model_filepath把优化后的图 dump 出来看。4.4 服务层的并发模型FastAPI 默认是单进程异步但 ONNX Runtime 的推理是 CPU 密集的会阻塞事件循环。我的做法是推理放到线程池用run_in_executor调度import asyncio from concurrent.futures import ThreadPoolExecutor executor ThreadPoolExecutor(max_workers4) app.post(/predict) async def predict(request: Request): data await request.json() loop asyncio.get_event_loop() result await loop.run_in_executor(executor, infer, data) return resultworker 数量设为 GPU 数量的 2 倍因为推理时会有 IO 等待。如果纯 GPU 推理worker 数等于 GPU 数即可多了反而抢显存。5. 线上排查与常见问题速查5.1 延迟抖动的排查路径线上最头疼的是 P99 延迟突然飙高。我的排查顺序是先看 GPU 利用率再看批处理队列长度最后看输入长度分布。GPU 利用率如果没满说明瓶颈在调度队列长度如果持续增长说明吞吐不够输入长度如果突然变长说明上游数据有问题。有一次 P99 从 50ms 飙到 300ms查下来是某个客户端发了一批超长序列把 batch 里的 padding 撑大了。解决办法是按长度分队列长序列单独走一个队列避免拖累短序列。这个改动后 P99 稳定在 60ms 以内。5.2 显存泄漏的定位技巧显存泄漏在长时间运行的服务里很常见。定位方法是定期打印显存快照用torch.cuda.memory_summary()看 allocated 和 reserved 的差值。如果 reserved 持续增长但 allocated 稳定说明有碎片如果 allocated 也增长说明有张量没释放。我遇到过一次泄漏原因是异常路径下没释放中间张量。请求处理中途抛异常try块里创建的张量没被 GC。解决办法是用with torch.no_grad()包住推理并且把中间变量显式del。这个坑很隐蔽因为正常路径下没问题只有异常请求多了才暴露。5.3 常见问题速查表现象可能原因排查方法解决CUDA OOMbatch 过大或碎片打印显存快照减小 batch 或开碎片整理P99 抖动长序列拖累看输入长度分布按长度分队列吞吐上不去批处理没生效看队列长度调大等待时间精度下降量化过度对比 FP32 输出排除敏感层冷启动慢镜像过大看拉取时间换轻量运行时5.4 监控指标该埋哪些监控不是越多越好我一般只埋四类指标请求量、延迟分位数、GPU 利用率、批处理大小。请求量看趋势延迟分位数看抖动GPU 利用率看瓶颈批处理大小看调度效率。这四类指标能覆盖 90% 的线上问题。Prometheus 的直方图要设对 bucket我一般设[0.01, 0.025, 0.05, 0.1, 0.25, 0.5, 1.0]覆盖 10ms 到 1s 的范围。bucket 太粗看不出抖动太细浪费存储。这个配置是我试了好几版才定下来的。6. 我踩过的坑和给你的实操建议6.1 别在训练脚本里写死路径我早期图省事把数据路径、模型保存路径全写死在脚本里结果换台机器就要改代码。后来统一用配置文件加环境变量覆盖config.yaml里写默认值环境变量优先级更高。这样本地调试和线上部署用同一套代码只是环境变量不同。6.2 单测要覆盖边界输入AI 工程的单测不能只测正常输入。我要求至少覆盖三种边界空输入、超长输入、非法字符输入。空输入容易触发除零超长输入容易 OOM非法字符容易在 tokenizer 那层就崩。这三种情况在线上都真实发生过提前测出来能省很多事。6.3 版本锁定要精确到补丁号requirements.txt里我坚持写torch2.1.2而不是torch2.1。因为 PyTorch 的小版本升级经常改默认行为比如 2.1.0 到 2.1.2 之间就改过DataLoader的默认pin_memory策略。精确锁定能保证本地和线上环境一致避免“在我机器上是好的”这种问题。6.4 日志要带请求 ID排查线上问题时没有请求 ID 的日志就是一团乱麻。我在入口生成一个 UUID透传到所有日志和指标里。这样从一条错误日志能直接定位到具体请求的完整链路。这个改动成本很低但排查效率提升巨大。6.5 灰度发布要按流量比例模型更新不能全量推我一般先放 5% 流量观察 24 小时看延迟和精度指标没异常再逐步放大。灰度期间要同时跑新旧两个模型对比输出差异。如果差异超过阈值就自动回滚。这个机制帮我挡过好几次有问题的模型更新。这套流水线我从零搭到稳定运行大概花了三周其中一半时间在调推理性能和排查线上问题。回头看最值得投入的是监控和日志它们让后续所有优化都有据可依。如果你也在搭类似的系统建议先把可观测性做好再谈性能优化否则就是盲人摸象。