1. 为什么我要从零手搓一套AI工程化框架
第一次看到ai-engineering-from-scratch这个项目名的时候,我正被公司里那套祖传的模型部署流程折磨得够呛。一个简单的文本分类模型,从训练完到真正上线跑推理,中间要经过五六个脚本、三四个配置文件,还有一堆没人敢动的环境变量。每次有新同事接手,光是搞明白数据从哪进、模型从哪出,就得花上两三天。所以我特别理解为什么有人会想从零开始,把AI工程化这件事重新梳理一遍。
这个项目本质上是一套从零构建AI工程化全链路的实践指南,它不依赖任何现成的高层框架,而是带着你一步步把数据管道、模型训练、评估、部署、监控这些环节全部手写出来。它解决的核心问题是:当你用惯了各种封装好的工具之后,一旦遇到需要定制化、需要排查底层问题、需要做性能优化的场景,你会发现自己其实什么都不懂。适合谁来参考?我觉得有三类人最应该看:一是刚入行做AI应用开发的新人,想搞清楚一个模型从代码到服务的完整生命周期;二是从传统后端转AI工程的老手,需要把工程思维映射到AI领域;三是带团队的技术负责人,想给团队建立一套可复现、可维护的工程规范。
我花了大概两周时间,把这个项目的思路完整跟了一遍,又结合自己在实际生产环境里踩过的坑做了不少调整。下面我就按自己的理解,把整套东西拆开来讲讲,包括为什么这么设计、每一步怎么做、以及哪些地方最容易翻车。
2. 整体架构设计与技术选型思路
2.1 为什么选择“从零实现”而不是直接用现成框架
市面上做AI工程化的框架其实不少,从训练侧的各类高层API,到部署侧的模型服务工具,再到监控侧的各种可观测性平台,看起来什么都有。但问题在于,这些工具各自为政,之间的衔接往往靠胶水代码,一旦某个环节出问题,排查起来非常痛苦。ai-engineering-from-scratch的思路是:先把每个环节的最小必要逻辑用最朴素的方式实现一遍,理解清楚数据怎么流动、状态怎么管理、错误怎么传播,然后再考虑要不要引入框架。
这个选择背后的逻辑很实在。我举个例子,很多团队在用现成的模型服务框架时,遇到推理延迟高的问题,第一反应是加机器、换GPU,但实际原因可能只是预处理阶段有个同步IO操作阻塞了主线程。如果你从来没自己写过一遍推理服务的请求处理流程,你根本不会往那个方向想。从零实现的好处就是,每一个环节都是透明的,你知道每一毫秒花在哪里,每一份内存被谁持有。
当然,从零实现不等于永远不用框架。项目里的做法是:先用纯Python和基础库把核心逻辑跑通,形成一个可工作的基线,然后再逐步替换成更高效的实现。比如数据加载先用最简单的生成器,确认逻辑没问题后,再换成多进程预取;模型推理先用同步方式,确认正确性后,再引入批处理和异步。这种渐进式的思路,比一上来就堆框架要稳得多。
2.2 核心模块划分与依赖关系
整个项目把AI工程化拆成了五个核心模块,每个模块之间通过明确定义的接口通信,尽量降低耦合。这五个模块分别是:
- 数据管道:负责原始数据的读取、清洗、转换、分批,最终输出模型可用的张量或数组。
- 模型定义与训练:负责模型结构定义、损失函数、优化器、训练循环、检查点保存。
- 评估与验证:负责在验证集和测试集上计算指标,包括分类指标、回归指标、以及自定义的业务指标。
- 推理服务:负责把训练好的模型包装成可调用的服务,处理请求、批处理、超时、错误返回。
- 监控与日志:负责收集服务运行时的指标,包括延迟、吞吐、错误率、资源占用,以及结构化日志。
这五个模块的依赖关系是单向的:数据管道不依赖任何其他模块,训练依赖数据管道,评估依赖训练产出的模型,推理服务依赖训练产出的模型,监控则横切所有模块。这种单向依赖的好处是,你可以单独测试任何一个模块,而不需要把整个系统跑起来。比如你想验证数据管道的正确性,只需要构造一批假数据,检查输出是否符合预期,完全不用管模型那边的事。
我在实际项目里也尝试过类似的划分,但一开始没忍住,让训练模块直接去调推理服务的代码做在线评估,结果就是训练和推理的依赖缠在一起,改一个地方要重新部署两个服务。后来老老实实按单向依赖重构,虽然多写了一些接口代码,但维护成本降了很多。
2.3 技术栈选择与版本管理策略
项目在技术栈上刻意保持克制。核心依赖只有几个:数值计算用NumPy,深度学习用PyTorch(但只用了最基础的张量操作和自动求导,没有用高层训练API),服务框架用FastAPI,监控用Prometheus客户端库。其他都是Python标准库。这个选择的原因是,依赖越少,出问题时排查范围越小,版本冲突的概率也越低。
版本管理方面,项目用了比较严格的策略:所有依赖都锁定精确版本,并且在CI里跑一个最小依赖环境的测试。我见过太多项目因为某个间接依赖升级导致行为变化,最后花几天时间定位。锁定版本虽然看起来不够“先进”,但在生产环境里,稳定性比新鲜感重要得多。
另外,项目对Python版本也有要求,建议用3.10以上,因为用到了一些类型注解的新特性,以及match语句来做配置解析。如果你还在用3.8,很多代码需要改写成if-elif,虽然也能跑,但可读性会差一些。
3. 数据管道的核心细节与实操要点
3.1 数据加载:从原始文件到内存批次
数据管道的第一步是把原始数据读进来。项目里假设数据以JSON Lines格式存储,每行一个样本,包含输入文本和标签。为什么选JSON Lines而不是CSV或Parquet?因为JSON Lines对嵌套结构的支持更好,而且可以逐行读取,不需要一次性把整个文件加载到内存。对于动辄几个GB的数据集,这个特性很关键。
读取的实现很朴素:打开文件,逐行解析JSON,做基本的字段校验,然后 yield 出去。这里有个细节需要注意:不要用json.loads直接解析每一行,而是先用orjson或ujson这样的快速库。我实测过,在千万级样本上,标准库的json比orjson慢三到四倍。虽然项目为了减少依赖没有强制用orjson,但在注释里明确提到了这个优化点。
读取之后是清洗。清洗的逻辑包括:去除空样本、截断过长的文本、过滤标签异常的样本。这里有个坑:截断长度不要拍脑袋定,要根据模型的最大输入长度和实际数据的长度分布来定。项目里给了一个方法:先统计所有样本的长度,画出累积分布曲线,然后选择覆盖95%样本的长度作为截断阈值。这样既能保留大部分信息,又能控制计算量。
3.2 数据转换:分词、向量化与批处理
清洗完之后是转换。对于文本数据,核心步骤是分词和向量化。项目里没有用现成的分词器,而是实现了一个简单的基于空格和标点的分词器,然后构建词表,把词映射成ID。这么做的好处是,你完全清楚每个ID是怎么来的,词表大小怎么控制,未知词怎么处理。
词表构建有个经验:不要把所有词都放进词表,低频词直接映射成UNK。项目里的做法是统计词频,保留频率最高的N个词,N一般取30000到50000。这个数字不是随便定的,太小会导致太多UNK,模型学不到东西;太大则嵌入矩阵会很大,显存占用高,而且低频词的嵌入往往训练不充分。我一般会看词频分布,找到那个“长尾”开始的点,作为截断阈值。
向量化之后是批处理。批处理的核心是把长度相近的样本放在同一个批次里,这样可以减少padding的数量,提高计算效率。项目里实现了一个简单的长度分桶策略:把样本按长度排序,然后按顺序切成批次。这个策略比随机分批要快不少,尤其是在长度分布比较分散的数据集上。不过要注意,训练时如果按长度排序,可能会导致批次之间的分布差异大,影响收敛。所以项目里建议在训练前先shuffle一次,然后再按长度分桶,这样既有随机性,又能减少padding。
3.3 数据管道的性能优化与常见陷阱
数据管道最容易成为整个训练流程的瓶颈。我见过很多情况,GPU利用率只有30%不到,一查发现是数据加载拖了后腿。项目里给了几个优化方向:
- 预取:用后台线程或进程提前加载下一批数据,让数据准备和模型计算重叠起来。
- 内存映射:对于超大数据集,用
numpy.memmap或类似机制,避免一次性加载到内存。 - 缓存:把预处理后的数据缓存到磁盘,下次直接读缓存,跳过清洗和转换步骤。
这里有个陷阱:多进程加载时,要注意每个进程的内存占用。如果每个进程都复制一份完整的数据集,内存会爆炸。项目里的做法是,主进程只保存索引,实际数据由各个worker按需读取。另外,worker的数量不要设太多,一般设为CPU核心数的70%到80%,留一些给主进程和其他系统任务。
还有一个常见问题是数据顺序。如果训练时数据是按类别排序的,模型可能会学到“先看到正例再看到负例”这种虚假模式。所以项目里强调,在分桶之前一定要做全局shuffle,并且每个epoch重新shuffle一次。
4. 模型训练与评估的完整实现
4.1 模型定义:从线性层到自定义结构
项目里的模型定义部分,是从最简单的线性分类器开始的。输入是词ID序列,经过嵌入层变成向量序列,然后做平均池化,最后接一个线性层输出类别logits。这个结构虽然简单,但包含了文本分类的核心要素:嵌入、池化、分类头。
为什么从这么简单的结构开始?因为复杂的模型往往是在简单模型的基础上加东西,如果简单模型都没跑通,加更多层只会让问题更难定位。我自己的习惯是,先用一个极简模型确认数据管道、损失函数、优化器都没问题,然后再逐步加注意力、加层数、加正则化。
项目里也提到了几个常见的模型结构变体,比如用CNN做文本分类、用LSTM做序列建模、用Transformer做更复杂的任务。但每个变体都是独立实现的,没有用继承或复杂的抽象。这样做的好处是,每个模型文件都是自包含的,你可以单独看某一个,不需要理解整个类层次结构。
4.2 训练循环:手写反向传播与参数更新
训练循环是项目里最核心的部分之一。项目没有用PyTorch的Trainer或fit方法,而是手写了完整的训练循环:前向传播、计算损失、反向传播、参数更新、梯度清零。这么做的好处是,你完全清楚每一步发生了什么,哪里可以插入自定义逻辑。
训练循环里有个关键细节:梯度累积。当显存不够大,无法容纳大batch时,可以把多个小batch的梯度累加起来,再一次性更新参数。项目里实现了一个简单的梯度累积逻辑:每处理N个batch,才做一次optimizer.step()和optimizer.zero_grad()。这里的N就是累积步数,一般设为2到8。要注意的是,累积时损失要除以N,否则梯度会放大N倍。
另一个细节是学习率调度。项目里实现了一个简单的预热加余弦退火策略:前10%的步数线性增加学习率,之后按余弦曲线衰减。这个策略在Transformer类模型上效果很好,但在简单模型上可能没必要。项目里的建议是,先用固定学习率跑通,再尝试调度策略,对比验证集上的效果。
4.3 评估指标:准确率之外的业务视角
评估部分,项目除了实现准确率、精确率、召回率、F1这些标准指标,还强调了业务指标的重要性。比如在一个垃圾文本过滤场景里,准确率高不一定好,因为可能把很多正常文本误判成垃圾,导致用户体验下降。这时候更关注的是召回率,或者精确率和召回率的某个加权组合。
项目里给了一个计算混淆矩阵的工具函数,以及从混淆矩阵推导各种指标的代码。这个工具函数支持多分类,也支持二分类,输出是一个二维数组,行是真实类别,列是预测类别。有了混淆矩阵,你可以很直观地看到模型在哪些类别上容易混淆,从而有针对性地调整。
还有一个容易被忽略的点:评估要在固定的验证集上做,而且验证集不能和训练集有重叠。项目里建议在数据划分时就用哈希或随机种子固定下来,避免每次跑评估时验证集不一样,导致指标不可比。
5. 推理服务与监控的落地实践
5.1 推理服务:从模型加载到请求处理
推理服务部分,项目用FastAPI搭了一个简单的HTTP服务。核心逻辑是:启动时加载模型到内存,收到请求后做预处理、推理、后处理,返回结果。这里有几个关键设计:
- 模型只加载一次:在服务启动时加载,而不是每次请求都加载。加载模型是IO密集和计算密集的操作,每次请求都做的话,延迟会高得离谱。
- 批处理:服务支持把多个请求合并成一个批次做推理,这样可以充分利用GPU的并行能力。项目里实现了一个简单的批处理队列:请求先进入队列,攒够一定数量或等待一定时间后,一起送给模型。
- 超时控制:每个请求有超时时间,超过就返回错误,避免慢请求拖垮整个服务。
批处理的实现有个细节:批处理窗口不能太大,否则延迟会很高;也不能太小,否则吞吐上不去。项目里的默认值是等待10毫秒或攒够32个请求,哪个先到就触发。这个值可以根据实际场景调整,延迟敏感的场景可以调小,吞吐敏感的场景可以调大。
5.2 监控指标:延迟、吞吐与资源占用
监控部分,项目用Prometheus客户端库暴露了几个核心指标:请求总数、请求延迟分布、当前队列长度、GPU显存占用、CPU使用率。这些指标通过一个/metrics端点暴露,Prometheus定时抓取。
延迟分布用直方图来记录,分桶的边界要仔细选。项目里给的建议是,根据实际延迟的分布来定分桶,比如P50在20毫秒,P99在200毫秒,那分桶可以设为10、25、50、100、200、500毫秒。分桶太粗看不出细节,太细则存储成本高。
还有一个指标容易被忽略:队列长度。如果队列长度持续增长,说明服务处理不过来,需要扩容或优化。项目里把队列长度也暴露出来,并且设置了一个告警规则:队列长度超过阈值持续一段时间就触发告警。
5.3 日志与错误处理:结构化日志与优雅降级
日志部分,项目强调用结构化日志,也就是JSON格式的日志,而不是纯文本。结构化日志的好处是,可以直接被日志系统解析和索引,方便做聚合分析和告警。每条日志包含时间戳、请求ID、用户ID、处理阶段、耗时、错误信息等字段。
错误处理方面,项目实现了一个全局异常处理器,捕获所有未处理的异常,记录日志,然后返回一个统一的错误响应。对于可恢复的错误,比如输入格式不对,返回400;对于服务内部错误,返回500。另外,项目还实现了一个优雅降级逻辑:如果模型推理失败,可以返回一个默认结果或缓存结果,而不是直接报错。这在一些对可用性要求高的场景里很有用。
6. 常见问题与排查技巧实录
6.1 训练不收敛:从数据到超参的排查顺序
训练不收敛是最常见的问题之一。项目里给了一个排查顺序:先看数据,再看模型,最后看超参。具体来说:
- 数据:检查输入和标签是否对应正确,有没有标签错位、数据泄漏、重复样本。我遇到过一次,数据管道里做shuffle时把输入和标签分开shuffle了,导致模型完全学不到东西。
- 模型:检查模型结构是否有误,比如维度不匹配、激活函数用错、初始化方式不当。可以用一个极小的数据集(比如10个样本)做过拟合测试,如果模型连这10个样本都拟合不了,那肯定是模型或训练逻辑有问题。
- 超参:检查学习率是否太大或太小,batch size是否合适,优化器选择是否正确。学习率太大导致loss震荡,太小则收敛慢。项目里建议先用一个中等学习率(比如1e-3)跑几百步,观察loss曲线,再调整。
6.2 推理延迟高:从预处理到批处理的逐层定位
推理延迟高的问题,项目里给了一个逐层定位的方法:
| 排查层次 | 检查内容 | 常见问题 |
|---|---|---|
| 网络层 | 请求往返时间 | DNS解析慢、连接池不足 |
| 预处理 | 分词、向量化耗时 | 同步IO、重复计算 |
| 模型推理 | 前向传播耗时 | 批次太小、GPU利用率低 |
| 后处理 | 结果格式化耗时 | 复杂逻辑、大对象序列化 |
| 批处理 | 等待时间 | 窗口太大、队列积压 |
我自己的经验是,大部分延迟问题出在预处理和批处理等待上。预处理如果用了同步的文件读取或网络请求,会直接阻塞主线程。批处理窗口设得太大,请求会等很久才被处理。项目里建议先用一个简单的计时器,记录每个阶段的耗时,找到瓶颈后再针对性优化。
6.3 显存不足:批次大小与梯度累积的权衡
显存不足是训练大模型时的常见问题。项目里给了几个应对策略:
- 减小批次大小:最直接的方法,但可能会影响收敛。
- 梯度累积:用多个小批次累积梯度,模拟大批次的效果。
- 混合精度训练:用FP16代替FP32,显存占用减半,但要注意数值稳定性。
- 梯度检查点:用计算换显存,适合特别大的模型。
项目里重点讲了梯度累积的实现和注意事项。累积步数N的选择,要保证等效批次大小(小批次大小乘以N)和原来差不多。比如原来批次大小是64,现在显存只够16,那N就设为4。另外,累积时要注意BatchNorm的统计量更新,如果用了BatchNorm,累积多个小批次时统计量会有偏差,这时候可以考虑用GroupNorm或LayerNorm代替。
6.4 服务不稳定:内存泄漏与连接池配置
服务跑一段时间后变慢或崩溃,往往是内存泄漏或连接池配置不当。项目里给了一个排查清单:
- 内存泄漏:用
tracemalloc或objgraph检查对象增长,重点看全局缓存、未关闭的文件句柄、循环引用。 - 连接池:数据库连接、HTTP连接都要用连接池,池大小要合理。太小会导致请求排队,太大则浪费资源。
- 线程/进程数:worker数量不要超过CPU核心数,否则上下文切换开销大。
- 日志量:日志太多会占满磁盘,也会拖慢服务。要设置合理的日志级别和轮转策略。
我踩过的一个坑是,在请求处理函数里创建了一个全局的缓存字典,但没有设置过期时间,结果缓存越来越大,最后OOM。后来改成用LRU缓存,并设置了最大容量,问题就解决了。
7. 我在实际项目中的几点体会
这套从零实现的思路,我在自己的项目里也尝试了一部分。最大的感受是,手写一遍之后,再用现成框架时,心里有底了。以前遇到问题只能猜,现在能大概判断是哪个环节出了状况。另外,项目里强调的“先跑通再优化”原则,帮我避免了很多过早优化带来的麻烦。我见过不少团队,一上来就搞分布式训练、混合精度、模型并行,结果基础的数据管道都没搞对,最后训练出来的模型效果一塌糊涂。
还有一个实用的建议:把每个模块的接口定义清楚,并且写测试。项目里每个模块都有对应的单元测试,数据管道测输出形状和数值范围,模型测前向传播的维度,推理服务测请求和响应的格式。这些测试看起来简单,但在重构时能帮你快速发现破坏性变更。我自己的项目里,就是因为有这些测试,才敢在后期把数据加载从单进程改成多进程,而不担心引入bug。
最后分享一个小技巧:在训练循环里加一个“健康检查”,每隔几百步检查一下loss是否为NaN、梯度范数是否异常大、学习率是否正常。如果发现异常,就保存当前状态并退出,而不是继续跑下去浪费资源。这个检查花不了多少时间,但能帮你省下很多调试的功夫。