1. 为什么我要从零手搓一套AI工程化框架
第一次听到“ai-engineering-from-scratch”这个说法,是在一个做推荐系统的老哥群里。有人甩了个链接,说现在市面上讲AI的教程要么是调包侠速成班,要么是论文复现劝退营,真正教你从工程角度把模型从实验室推到生产环境的内容少得可怜。我当时正被公司里一个文本分类项目折磨——模型在notebook里跑得漂漂亮亮,一上服务器就各种幺蛾子,显存泄漏、推理延迟抖动、版本回滚困难,这些问题没有一个能靠调参解决。
所以当我看到“ai-engineering-from-scratch”这个标题时,第一反应是:终于有人要讲人话了。它要解决的核心问题很明确——把AI从“能跑通”变成“能扛住”。这不是教你怎么写模型结构,而是教你怎么搭一套让模型稳定干活的工程体系。适合谁看?如果你已经会写PyTorch或TensorFlow的基础代码,但一遇到部署、监控、迭代就头大,那这套东西就是给你准备的。如果你还在纠结反向传播怎么推导,建议先补基础,因为这里聊的是“怎么让模型在线上不出事”,而不是“模型为什么能学习”。
我花了大概三周时间,把从数据管道到服务上线的全链路自己撸了一遍。踩的坑比预想的多,但也正是这些坑让我理解了为什么AI工程化值得单独拿出来讲。下面我把整个思路、关键决策和实操细节拆开说,尽量让每个环节都能直接抄作业。
2. 整体架构设计与技术选型逻辑
2.1 为什么选择“从零搭建”而不是用现成平台
市面上不缺AI平台,从云厂商的一站式解决方案到开源的MLflow、Kubeflow,看起来都能解决问题。但我坚持从零搭的原因有三个:第一,黑盒调试成本太高。当推理服务出现偶发超时,如果底层是封装好的平台,你只能看它给的日志,很多中间状态根本拿不到。自己搭的话,从请求进来到结果返回,每一层都能打点。第二,依赖锁定风险。用平台爽在初期,但一旦要换模型格式或者调整预处理逻辑,平台不支持就得干等。第三,成本控制。小团队用云平台跑推理,账单能吓死人,自己用开源组件拼一套,同样的负载成本能压到三分之一以下。
当然,从零搭不等于所有轮子都自己造。我的原则是:核心链路自己写,边缘组件用成熟库。比如Web框架用FastAPI,序列化用Protobuf,监控用Prometheus,这些没必要重复造。但模型加载、批处理调度、版本路由这些跟业务强相关的部分,必须自己掌控。
2.2 分层架构:把“AI”和“工程”拆开看
整个系统我分成了四层,每层职责单一,方便独立替换和测试:
- 数据接入层:负责原始数据清洗、特征提取、格式转换。这一层的关键是幂等性——同样的输入必须产出同样的输出,否则后续排查问题会疯掉。
- 模型推理层:加载模型权重,执行前向计算。这里要处理动态批处理、显存管理、多模型共存。
- 服务接口层:对外暴露HTTP/gRPC接口,处理鉴权、限流、请求校验。
- 可观测层:日志、指标、追踪三件套,贯穿所有层。
这么分的好处是,当推理延迟升高时,我能快速判断是数据预处理慢了,还是模型计算本身慢了,还是网络传输堵了。如果混在一起写,就只能靠猜。
2.3 技术栈选型:每个选择都要有理由
| 组件 | 选型 | 理由 |
|---|---|---|
| Web框架 | FastAPI | 异步支持好,自动生成OpenAPI文档,类型提示友好 |
| 模型运行时 | ONNX Runtime | 跨框架兼容,推理优化成熟,CPU/GPU切换方便 |
| 批处理调度 | 自研异步队列 | 现成方案要么太重,要么不支持动态批大小 |
| 监控 | Prometheus + Grafana | 生态完善,指标采集灵活,告警规则好写 |
| 日志 | structlog | 结构化输出,方便ELK收集和检索 |
| 容器化 | Docker + Compose | 开发环境一键拉起,生产环境可平滑迁移到K8s |
这里重点说下为什么选ONNX Runtime而不是直接跑PyTorch。PyTorch的torchserve确实方便,但它的批处理逻辑是固定的,没法根据请求量动态调整。而ONNX Runtime的InferenceSession可以手动控制run的调用时机,配合自研队列能实现更细粒度的批处理。另外ONNX的图优化在CPU上提升明显,我们有个文本分类模型,转ONNX后单次推理从45ms降到了28ms。
注意:转ONNX不是万能的。如果模型里有大量自定义算子,转换过程可能失败或者精度损失。建议转完后用一批测试数据对比输出差异,确保误差在可接受范围内。
3. 核心模块的实操细节与避坑指南
3.1 数据预处理管道:别让脏数据毁了模型
数据预处理看起来简单,但线上出问题十有八九在这里。我踩过的坑包括:训练时用的分词器和线上不一致、数值特征归一化参数没保存、类别特征映射表丢失。这些问题在离线评估时发现不了,一上线就暴露。
我的做法是把预处理逻辑固化成一个独立的Pipeline对象,跟模型权重一起保存。这个Pipeline包含所有必要的状态:分词器、归一化均值方差、类别映射字典、缺失值填充策略。加载模型时,Pipeline和权重一起反序列化,确保线上线下完全一致。
class PreprocessPipeline: def __init__(self, tokenizer, scaler_mean, scaler_std, cat_mapping): self.tokenizer = tokenizer self.scaler_mean = scaler_mean self.scaler_std = scaler_std self.cat_mapping = cat_mapping def transform(self, raw_input): # 文本分词 tokens = self.tokenizer.encode(raw_input['text']) # 数值归一化 num_features = (raw_input['numeric'] - self.scaler_mean) / self.scaler_std # 类别映射 cat_feature = self.cat_mapping.get(raw_input['category'], 0) return {'tokens': tokens, 'numeric': num_features, 'category': cat_feature}保存的时候用joblib或者pickle都行,但要注意版本兼容。我遇到过用Python 3.8训练的Pipeline在3.10环境加载报错,原因是pickle协议版本不一致。后来统一用joblib并指定protocol=4,问题解决。
另一个关键是输入校验。线上请求什么妖魔鬼怪都有,空字符串、超长文本、非法字符。我在Pipeline入口加了严格的校验逻辑:文本长度超过阈值直接截断,数值超出范围用边界值替代,类别不在映射表里归为“未知”。这些规则在训练时也要用同样的逻辑处理,否则模型看到的分布和线上不一致。
3.2 动态批处理:榨干GPU的每一滴算力
批处理是提升推理吞吐最有效的手段,但静态批处理有个致命问题:如果请求量不稳定,要么GPU闲着,要么请求排队。动态批处理的核心思想是在延迟和吞吐之间找平衡——攒一小批请求一起算,但等待时间不超过阈值。
我的实现方案是用一个异步队列加一个后台worker。请求进来先入队,worker每隔几毫秒检查一次队列,如果队列长度达到max_batch_size或者等待时间超过max_wait_ms,就取出当前所有请求组成一个batch送进模型。
class DynamicBatcher: def __init__(self, model, max_batch_size=32, max_wait_ms=10): self.model = model self.max_batch_size = max_batch_size self.max_wait_ms = max_wait_ms self.queue = asyncio.Queue() async def infer(self, input_data): future = asyncio.Future() await self.queue.put((input_data, future)) return await future async def _worker(self): while True: batch = [] start_time = time.time() while len(batch) < self.max_batch_size: timeout = self.max_wait_ms / 1000 - (time.time() - start_time) if timeout <= 0: break try: item = await asyncio.wait_for(self.queue.get(), timeout) batch.append(item) except asyncio.TimeoutError: break if batch: inputs = [item[0] for item in batch] results = self.model.predict(inputs) for (_, future), result in zip(batch, results): future.set_result(result)参数调优方面,max_batch_size取决于模型大小和显存。我一般先用nvidia-smi看模型加载后的显存占用,然后估算每个样本的激活值开销。比如一个BERT-base模型,加载后占1.2GB,每个样本前向传播约需15MB,那32GB显存的卡理论上能跑2000个样本,但实际要考虑碎片和峰值,我一般设成理论值的60%左右。max_wait_ms则根据业务延迟要求来,如果是实时交互场景,设5-10ms;如果是离线批量任务,可以设到100ms以上。
实操心得:动态批处理在请求量低的时候反而会增加延迟,因为要等攒批。所以最好加个自适应逻辑——当队列长度持续为1时,直接跳过等待立即推理。这个逻辑我加了之后,低峰期P99延迟从15ms降到了8ms。
3.3 模型版本管理与灰度发布
模型迭代是常态,但直接替换线上模型风险极高。我见过一次事故:新模型在测试集上F1涨了2个点,上线后核心业务指标反而跌了5个点,原因是新模型对某个高频类别的预测偏向变了,导致下游策略失效。
所以版本管理和灰度发布是必须的。我的方案是每个模型版本一个独立目录,包含权重文件、Pipeline对象、配置文件。服务启动时加载所有可用版本,通过路由规则决定请求走哪个版本。
路由规则我支持三种模式:
- 按比例分流:比如新版本承接10%流量,观察指标后再逐步放大。
- 按用户分组:内部用户走新版本,外部用户走稳定版本。
- 按请求特征:特定来源或特定类型的请求走新版本。
class ModelRouter: def __init__(self, versions, strategy='ratio', ratio=0.1): self.versions = versions # {'v1': model1, 'v2': model2} self.strategy = strategy self.ratio = ratio def route(self, request): if self.strategy == 'ratio': if random.random() < self.ratio: return self.versions['v2'] return self.versions['v1'] elif self.strategy == 'user_group': if request.user_id in INTERNAL_USERS: return self.versions['v2'] return self.versions['v1'] # 其他策略...灰度期间要重点监控业务指标而不只是模型指标。比如推荐场景看点击率、转化率,风控场景看拦截率、误杀率。一旦发现异常,立即把流量切回旧版本。回滚操作要能在秒级完成,所以模型加载不能太慢。我的做法是服务启动时就把所有版本加载到显存,切换只是改路由指针,不涉及加载。
3.4 可观测性建设:出了问题能快速定位
AI系统的可观测性比普通后端服务更复杂,因为除了常规的QPS、延迟、错误率,还要监控模型层面的指标:输入分布漂移、预测置信度分布、特征缺失率。
我用Prometheus采集指标,每个推理请求记录以下数据:
- 请求延迟(分预处理、推理、后处理三段)
- 批大小
- 模型版本
- 输入特征统计量(均值、方差、缺失率)
- 输出置信度分布
from prometheus_client import Histogram, Counter, Gauge INFERENCE_LATENCY = Histogram('inference_latency_seconds', 'Inference latency', ['stage', 'model_version']) BATCH_SIZE = Histogram('batch_size', 'Batch size distribution', ['model_version']) INPUT_DRIFT = Gauge('input_drift_score', 'Input distribution drift', ['feature_name'])日志方面,每个请求分配一个trace_id,从入口到出口全链路透传。这样当用户反馈某个请求结果异常时,我能通过trace_id把整个处理过程串起来看。structlog的bind方法很好用:
import structlog logger = structlog.get_logger() async def handle_request(request): log = logger.bind(trace_id=request.trace_id, model_version=request.model_version) log.info("request_received", input_length=len(request.text)) # ...处理... log.info("inference_completed", latency=elapsed, batch_size=batch_size)避坑提醒:日志里千万别打原始输入数据,尤其是文本内容。一是隐私合规问题,二是日志量会爆炸。我一般只记录统计特征,比如文本长度、token数量、特征哈希值。需要调试时再临时开启详细日志,用完就关。
4. 完整部署流程与性能调优实录
4.1 从本地开发到容器化部署
本地开发时我直接用uvicorn跑FastAPI,模型加载到内存。但到了生产环境,需要考虑进程管理、资源隔离、健康检查。我的Dockerfile大概长这样:
FROM python:3.10-slim WORKDIR /app # 安装系统依赖 RUN apt-get update && apt-get install -y --no-install-recommends \ libgomp1 \ && rm -rf /var/lib/apt/lists/* # 安装Python依赖 COPY requirements.txt . RUN pip install --no-cache-dir -r requirements.txt # 复制代码和模型 COPY src/ ./src/ COPY models/ ./models/ # 健康检查 HEALTHCHECK --interval=30s --timeout=5s --retries=3 \ CMD python -c "import requests; requests.get('http://localhost:8000/health')" EXPOSE 8000 CMD ["uvicorn", "src.main:app", "--host", "0.0.0.0", "--port", "8000", "--workers", "1"]注意--workers我设的是1,因为模型加载很吃显存,多个worker会重复加载。如果要提升并发,应该用动态批处理而不是多进程。另外libgomp1是ONNX Runtime的依赖,不装会报错。
容器启动后,用docker stats看资源占用。如果显存没跑满但CPU很高,说明预处理是瓶颈,可以考虑把预处理也放到GPU上(用CUDA加速的tokenizer)。如果显存快满了但GPU利用率低,说明批大小设小了,可以适当调大。
4.2 性能压测与瓶颈定位
压测我用locust模拟并发请求,逐步增加用户数,观察延迟和吞吐的变化。第一次压测结果很惨:50并发时P99延迟就飙到了2秒。排查后发现三个问题:
问题一:预处理在Python主线程里跑,GIL锁住了。解决方案是把预处理放到ProcessPoolExecutor里,绕开GIL。但进程间通信有开销,后来改用concurrent.futures.ThreadPoolExecutor配合C扩展的tokenizer,效果好很多。
问题二:每次推理都重新创建ONNX Runtime的InferenceSession。这是个低级错误,InferenceSession创建开销很大,应该全局只创建一次。改完后延迟直接降了40%。
问题三:日志同步写磁盘,IO阻塞。改成异步写,用QueueHandler把日志丢到队列里,后台线程慢慢刷盘。
优化后的压测数据:
| 并发数 | QPS | P50延迟 | P99延迟 | GPU利用率 |
|---|---|---|---|---|
| 10 | 320 | 28ms | 45ms | 35% |
| 50 | 1450 | 32ms | 68ms | 78% |
| 100 | 2100 | 45ms | 120ms | 92% |
| 200 | 2300 | 82ms | 350ms | 95% |
可以看到100并发之后QPS增长放缓,P99延迟上升明显,说明GPU已经接近饱和。这时候要么加卡,要么做模型量化。我试了ONNX的INT8量化,模型大小从420MB降到110MB,推理速度提升约1.8倍,但精度掉了1.2个点。对于我们的场景可以接受,如果精度敏感就得用FP16或者不做量化。
4.3 显存泄漏排查:一个折腾了两天的bug
有次服务跑了一天后显存从8GB涨到了14GB,最后OOM被杀。排查过程很痛苦,因为泄漏是缓慢发生的,本地跑几小时看不出来。
我用了pynvml库定时打印显存使用,同时用tracemalloc跟踪Python内存分配。最后定位到问题在动态批处理的asyncio.Future上——当请求超时被取消时,Future对象没有被正确清理,导致引用计数不归零。
修复方法是在infer方法里加try/finally,确保Future被取消或设置结果:
async def infer(self, input_data): future = asyncio.Future() try: await self.queue.put((input_data, future)) return await asyncio.wait_for(future, timeout=5.0) except asyncio.TimeoutError: future.cancel() raise finally: if not future.done(): future.cancel()另外ONNX Runtime的InferenceSession如果频繁创建和销毁,也会有显存碎片。所以一定要复用session,不要每次请求都新建。
经验之谈:显存泄漏问题在开发环境很难复现,建议在测试环境跑长时间稳定性测试,至少24小时。同时加上显存监控告警,超过阈值自动重启服务。虽然粗暴但有效。
5. 常见问题速查与独家避坑技巧
5.1 推理服务常见故障排查表
| 现象 | 可能原因 | 排查方法 | 解决方案 |
|---|---|---|---|
| 延迟突然升高 | 批处理等待超时 | 查看batch_size分布 | 调小max_wait_ms |
| 显存持续增长 | Future未清理 | pynvml监控显存 | 加try/finally清理 |
| 预测结果不一致 | 预处理状态丢失 | 对比线上线下Pipeline | 固化Pipeline并随模型保存 |
| 服务启动慢 | 模型加载耗时 | 计时各阶段 | 预加载+懒加载结合 |
| 吞吐上不去 | GPU利用率低 | nvidia-smi查看 | 增大批大小或量化模型 |
| 错误率突增 | 输入分布漂移 | 监控特征统计量 | 加输入校验和兜底逻辑 |
5.2 那些文档里不会写的实操心得
心得一:模型文件不要放在代码仓库里。用Git LFS也会让仓库变得巨大,clone一次要半天。我的做法是模型文件单独存对象存储,部署时用脚本拉取。版本号用模型文件的MD5,确保一致性。
心得二:健康检查要区分“存活”和“就绪”。存活检查只判断进程在不在,就绪检查要判断模型是否加载完成、显存是否充足。K8s里用livenessProbe和readinessProbe分别配置,避免服务还没加载完就被打流量。
心得三:日志级别动态调整。平时用INFO级别,出问题时通过环境变量或配置中心临时切到DEBUG,不用重启服务。我用的logging模块配合watchdog监听配置文件变化,改完立即生效。
心得四:压测数据要贴近真实分布。用随机生成的假数据压测,结果会偏乐观。因为真实数据的长度分布、特征分布都有长尾,处理长文本的耗时可能是短文本的几十倍。我一般从线上采样一批真实请求(脱敏后)作为压测输入。
心得五:做好降级预案。当模型服务不可用时,要有兜底逻辑。比如返回默认结果、走规则引擎、或者直接返回错误码让上游处理。最怕的是模型服务挂了导致整个业务链路雪崩。我在服务入口加了熔断器,连续失败超过阈值就自动降级,恢复后再切回来。
5.3 性能优化的几个关键参数
最后整理一下我调优过程中觉得最关键的几个参数,供参考:
- ONNX Runtime的
intra_op_num_threads:控制单次推理内部的线程数。CPU推理时设成物理核心数,GPU推理时设成1(避免CPU-GPU同步开销)。 max_batch_size:根据显存和模型大小估算,建议从16开始逐步往上试,观察P99延迟变化。max_wait_ms:实时场景5-10ms,准实时50ms,离线场景可以到500ms。queue_max_size:队列满了要拒绝请求还是阻塞等待?我一般设成max_batch_size * 10,超过就返回503,保护服务不被打垮。session_pool_size:如果单卡显存够大,可以创建多个InferenceSession并行推理,但要注意显存碎片。我一般设1-2个。
这套东西搭下来,最大的感受是:AI工程化没有银弹,每个决策都要结合具体场景权衡。别人说好的方案,到你这里可能因为数据特性、硬件配置、业务要求不同而完全不适用。所以多动手试,多监控,多复盘,比看一百篇教程都管用。