1. 为什么我要从零手搓一套AI工程流水线
第一次看到ai-engineering-from-scratch这个项目名的时候,我正被一堆"调包式AI开发"折磨得够呛。那会儿团队里新来的几个小伙伴,问他们模型怎么部署的,回答是"就调了个API";问推理延迟怎么优化,回答是"换了个更贵的卡"。这种状态在业务量小的时候没问题,一旦请求量上来、成本开始咬人、线上开始抖动,你就会发现——你根本不知道黑盒里发生了什么,也就无从下手去修。
ai-engineering-from-scratch这个标题,字面意思就是"从零开始的AI工程"。它不是教你调某个框架的API,也不是让你背Transformer的公式,而是把AI系统从数据进来、模型训练、推理服务、监控告警这一整条链路,用最朴素的方式自己搭一遍。核心价值在于"祛魅":当你亲手写过一遍tokenizer、手写过一遍KV Cache的调度逻辑、手写过一遍批处理队列,再回头看那些封装好的框架,你才知道每个参数背后在发生什么。
这篇文章适合三类人:一是刚入行做AI应用、只会调API的工程师,想搞清楚底下到底怎么回事;二是后端或数据工程师,被拉来做AI系统但缺乏端到端视角;三是准备面试大厂AI工程岗的同学,面试官特别爱问"如果不用框架你怎么实现"。我会把整条链路的设计思路、关键取舍、实操步骤、踩坑记录全部摊开讲,代码和参数都给到能直接抄的程度。
先说清楚我的立场:从零手搓不是为了替代框架,而是为了获得"选择框架的能力"。你手写过一遍,才知道vLLM的PagedAttention到底解决了什么问题,才知道Triton的动态批处理为什么能提升吞吐,才知道量化到INT8会掉多少精度。这种判断力,是调包调不出来的。
2. 整体架构设计与技术选型思路
2.1 一条AI工程流水线到底包含哪些环节
很多人对"AI工程"的理解停留在"训练模型",这是最大的误区。一个能上生产的AI系统,训练只是其中一环,而且往往是最短的一环。完整的链路我习惯拆成六段:
- 数据层:原始数据采集、清洗、去重、格式化,产出训练用的数据集
- 训练层:模型结构定义、训练循环、分布式策略、checkpoint管理
- 评估层:离线指标、回归测试集、badcase归因
- 推理层:模型加载、请求调度、批处理、KV Cache管理、采样策略
- 服务层:API网关、限流、鉴权、超时、降级
- 观测层:日志、指标、链路追踪、成本核算
ai-engineering-from-scratch的精髓就在于,这六层你都要自己碰一遍,哪怕每层只做一个最小可用版本。为什么?因为AI系统的故障往往发生在层与层的交界处。比如推理层显存爆了,根因可能是服务层没做请求长度限制;比如训练loss不收敛,根因可能是数据层去重没做好导致分布偏移。只懂一层的人,永远在猜。
我建议的搭建顺序是倒着来:先做推理层和服务层,因为这是最快能跑起来看到效果的;再做数据层和训练层;最后补评估和观测。这个顺序的好处是,你能在第一天就有一个能对外提供服务的demo,正反馈来得快,不容易半途而废。
2.2 为什么不用现成框架,以及什么时候该用
这里必须说清楚,我不是"反框架主义者"。手搓的目的是理解,不是生产。我的原则是:
| 场景 | 建议 | 理由 |
|---|---|---|
| 学习/理解原理 | 手搓 | 只有自己写才知道每个环节的约束 |
| 快速验证想法 | 用框架 | 时间成本优先,别重复造轮子 |
| 生产环境标准场景 | 用成熟框架 | vLLM/TensorRT-LLM经过大规模验证 |
| 生产环境特殊需求 | 框架+手写插件 | 比如自定义采样、特殊调度策略 |
| 极致性能优化 | 手搓关键路径 | 框架的通用性会带来开销 |
我踩过最大的坑,是在一个延迟敏感的场景里硬套通用推理框架,结果框架的调度开销占了总延迟的40%。后来把调度逻辑自己重写,延迟直接砍半。这就是"知道底层"的价值——你知道哪里可以砍,哪里不能动。
技术选型上,我建议从零实现时用Python + NumPy起步,不要一上来就上CUDA。原因很简单:NumPy版本能让你把算法逻辑跑通、把数值验证对,再迁移到GPU时你只需要关心并行化,不用同时debug算法和硬件。等算法逻辑稳定了,再用PyTorch的tensor重写,最后才考虑手写CUDA kernel。这个渐进路径能帮你省下大量时间。
2.3 最小可用系统的边界怎么划
从零做项目最容易犯的错是贪大求全,想一次把六层全做完,结果哪层都是半成品。我的经验是,第一版只做单机、单模型、同步推理的最小闭环,具体边界:
- 数据层:只支持一种格式(比如JSONL),只做最基础的清洗
- 训练层:只支持单卡、小模型(参数量控制在100M以内)
- 推理层:只支持单请求、不做批处理
- 服务层:一个Flask/FastAPI的裸接口,不做限流
- 观测层:只打print日志
这个版本大概两三天能跑通,跑通之后你就有了一条可以端到端调试的基线。后面所有的优化——加批处理、加KV Cache、加量化、加分布式——都是在这条基线上做增量。没有基线,你连优化效果都测不出来。
3. 核心模块的细节拆解与实操要点
3.1 数据层:清洗比采集重要十倍
数据层我见过最多的错误是"重采集轻清洗"。大家愿意花一周写爬虫,却不愿意花一天做去重。结果就是模型在训练集上表现很好,一到真实场景就拉胯,因为训练集里全是重复样本,模型过拟合了。
从零实现数据层,我建议按这个顺序做:
- 格式统一:把所有来源的数据转成统一的JSONL,每行一个样本,字段固定为
{"text": ..., "label": ..., "source": ...} - 精确去重:用哈希(比如SHA256)对文本做精确去重,这一步能干掉30%以上的重复
- 近似去重:用MinHash或SimHash做近似去重,阈值我一般设在0.85,能再干掉10%左右
- 质量过滤:按长度、字符集、困惑度过滤掉低质样本
- 切分:按时间或来源切分train/val/test,绝对不要随机切分,否则会有数据泄漏
注意:近似去重的阈值不要设太低,我试过0.7,结果把很多正常样本也误杀了,模型效果反而下降。0.85是个比较稳的经验值。
实操上,MinHash的实现可以用datasketch这个库,但如果你想从零理解,我建议自己实现一遍。核心逻辑是:对每个文档做shingling(比如3-gram),对每个shingle算多个哈希,取最小值组成签名,两个文档的签名相似度就近似Jaccard相似度。代码大概50行,但理解了它你就理解了所有近似去重算法的本质。
3.2 训练层:先跑通再谈优化
训练层从零实现,我的建议是先写一个纯NumPy的版本,哪怕慢得离谱。为什么?因为PyTorch的autograd太方便了,方便到你根本不知道反向传播在算什么。手写一遍前向和反向,你对梯度消失、梯度爆炸、学习率调度的理解会完全不一样。
一个最小训练循环包含这些部分:
- 前向传播:输入 -> 线性层 -> 激活 -> 线性层 -> 输出
- 损失计算:交叉熵或MSE
- 反向传播:手动推导每层的梯度
- 参数更新:SGD或Adam
- 学习率调度:warmup + cosine decay
我实测下来,一个两层MLP在NumPy上跑MNIST,一个epoch大概要几分钟,慢但能跑通。跑通之后,把同样的逻辑用PyTorch重写,你会发现PyTorch版本快了100倍,但算法逻辑完全一样。这时候你再用PyTorch的高级特性(混合精度、梯度累积、分布式),就知道每个特性在优化什么。
参数选择上,我踩过的坑是学习率设太大。从零实现时没有框架的默认值保护,很容易设成0.1导致loss直接爆炸。我的经验是:小模型从1e-3起步,大模型从1e-4起步,配合warmup(前10%的step线性升温),基本不会出问题。
3.3 推理层:KV Cache是性能的分水岭
推理层是从零实现里最有技术含量的部分,也是最能体现工程能力的地方。核心要解决三个问题:批处理、KV Cache、采样。
先说KV Cache。自回归生成时,每生成一个token都要重新计算前面所有token的Key和Value,这是巨大的浪费。KV Cache的思路是:把已经算过的K和V缓存起来,下一个token只需要算新的K和V,然后和缓存的拼接。这个优化能把生成速度提升几倍到几十倍,具体取决于序列长度。
从零实现KV Cache,关键数据结构是一个[batch, num_heads, seq_len, head_dim]的tensor,每次生成新token时append进去。要注意的是显存管理:序列越长,KV Cache越大,很容易OOM。我一般会设一个max_seq_len,超过就截断或拒绝请求。
批处理是另一个关键。同步推理时,一个请求算完再算下一个,GPU利用率极低。动态批处理的思路是:把短时间内到达的多个请求攒成一个batch,一起算。实现上需要一个队列 + 一个调度器,调度器决定什么时候触发一次batch推理。触发条件一般是"队列长度达到阈值"或"等待时间超过阈值",两者取先到。
采样策略相对简单,但细节多。贪心采样(argmax)最稳定但缺乏多样性;温度采样通过logits / temperature再softmax,温度越高越随机;top-k采样只保留概率最高的k个token;top-p(nucleus)采样保留累积概率达到p的最小token集合。我一般用top-p=0.9 + temperature=0.7,这个组合在大多数场景下比较稳。
3.4 服务层:别让一个慢请求拖垮整个服务
服务层最容易被忽视,但线上事故往往出在这里。从零实现时,至少要处理这几件事:
- 超时控制:每个请求设一个最大处理时间,超时就返回错误,别让它一直占着资源
- 并发限制:用信号量或队列限制同时处理的请求数,防止雪崩
- 请求校验:检查输入长度、格式,非法请求直接拒绝,别让它进到推理层
- 优雅降级:推理层挂了,服务层要能返回兜底结果,而不是直接500
我用FastAPI实现时,超时控制用asyncio.wait_for,并发限制用asyncio.Semaphore。这两个组合起来,能挡住90%的线上抖动。踩过的坑:一开始没做并发限制,结果一个用户发了1000个并发请求,直接把推理服务打挂,连带影响了其他所有用户。加了信号量之后,超出的请求排队等待,服务稳定性大幅提升。
4. 完整实操流程与关键环节实现
4.1 环境准备与依赖安装
从零实现不需要太多依赖,我建议保持极简:
python -m venv venv source venv/bin/activate pip install numpy fastapi uvicorn pydantic训练部分如果需要GPU,再加torch。但第一版我强烈建议纯CPU + NumPy,把算法跑通再说。依赖越少,你越能聚焦在逻辑本身。
目录结构我习惯这样组织:
ai-engineering-from-scratch/ ├── data/ │ ├── raw/ │ └── processed/ ├── src/ │ ├── data/ │ │ ├── clean.py │ │ └── dedup.py │ ├── train/ │ │ ├── model.py │ │ └── loop.py │ ├── infer/ │ │ ├── kv_cache.py │ │ └── sampler.py │ └── serve/ │ └── app.py ├── tests/ └── configs/这个结构的好处是每层职责清晰,改数据不影响训练,改推理不影响服务。我见过太多项目把所有代码堆在一个文件里,改一行牵一发动全身。
4.2 数据清洗与去重的实操步骤
先写清洗脚本。核心逻辑是读原始数据,逐条过滤,写出去重后的数据:
import hashlib import json def clean_and_dedup(input_path, output_path, min_len=10, max_len=2048): seen = set() kept = 0 with open(input_path) as fin, open(output_path, 'w') as fout: for line in fin: obj = json.loads(line) text = obj.get('text', '').strip() if not (min_len <= len(text) <= max_len): continue h = hashlib.sha256(text.encode()).hexdigest() if h in seen: continue seen.add(h) fout.write(json.dumps(obj, ensure_ascii=False) + '\n') kept += 1 print(f"kept {kept} samples")这个脚本跑一遍,你能直观看到去重率。我实测过一个中文语料,精确去重干掉了35%的样本,说明原始数据里重复非常严重。
近似去重我用MinHash,核心代码如下:
import hashlib def shingles(text, k=3): return {text[i:i+k] for i in range(len(text) - k + 1)} def minhash(shingle_set, num_hashes=128): sig = [] for i in range(num_hashes): min_h = float('inf') for s in shingle_set: h = int(hashlib.md5(f"{i}_{s}".encode()).hexdigest(), 16) min_h = min(min_h, h) sig.append(min_h) return sig def jaccard_estimate(sig1, sig2): return sum(a == b for a, b in zip(sig1, sig2)) / len(sig1)这个实现慢,但逻辑清晰。生产环境可以用LSH加速,但学习阶段慢一点没关系,理解原理比跑得快重要。
4.3 训练循环的手写实现
训练循环的核心是前向、损失、反向、更新四步。我用一个两层MLP举例:
import numpy as np class MLP: def __init__(self, in_dim, hidden, out_dim): self.W1 = np.random.randn(in_dim, hidden) * 0.01 self.b1 = np.zeros(hidden) self.W2 = np.random.randn(hidden, out_dim) * 0.01 self.b2 = np.zeros(out_dim) def forward(self, x): self.x = x self.h = np.maximum(0, x @ self.W1 + self.b1) # ReLU self.logits = self.h @ self.W2 + self.b2 return self.logits def backward(self, grad_logits, lr): grad_W2 = self.h.T @ grad_logits grad_b2 = grad_logits.sum(axis=0) grad_h = grad_logits @ self.W2.T grad_h[self.h <= 0] = 0 # ReLU反向 grad_W1 = self.x.T @ grad_h grad_b1 = grad_h.sum(axis=0) self.W1 -= lr * grad_W1 self.b1 -= lr * grad_b1 self.W2 -= lr * grad_W2 self.b2 -= lr * grad_b2配合softmax交叉熵的梯度(grad_logits = probs - one_hot),一个完整训练循环就成型了。关键点:ReLU的反向要把前向时小于等于0的位置梯度置零,这个细节很多人会漏。
学习率调度我用warmup + cosine:
def lr_schedule(step, warmup_steps, total_steps, base_lr): if step < warmup_steps: return base_lr * step / warmup_steps progress = (step - warmup_steps) / (total_steps - warmup_steps) return base_lr * 0.5 * (1 + np.cos(np.pi * progress))这个调度器能让训练前期稳定、后期收敛,比固定学习率效果好很多。
4.4 推理服务的KV Cache与批处理实现
KV Cache的核心是缓存已算过的K和V。简化版实现:
class KVCache: def __init__(self, max_len, num_heads, head_dim): self.max_len = max_len self.k = np.zeros((max_len, num_heads, head_dim)) self.v = np.zeros((max_len, num_heads, head_dim)) self.len = 0 def append(self, new_k, new_v): if self.len >= self.max_len: raise RuntimeError("KV cache full") self.k[self.len] = new_k self.v[self.len] = new_v self.len += 1 def get(self): return self.k[:self.len], self.v[:self.len]批处理调度器用一个队列 + 定时触发:
import asyncio from collections import deque class BatchScheduler: def __init__(self, max_batch=8, max_wait=0.05): self.queue = deque() self.max_batch = max_batch self.max_wait = max_wait async def submit(self, request): self.queue.append(request) if len(self.queue) >= self.max_batch: return await self._flush() await asyncio.sleep(self.max_wait) return await self._flush() async def _flush(self): batch = list(self.queue) self.queue.clear() return await self._run_batch(batch)max_batch和max_wait是两个关键参数。max_batch越大吞吐越高但延迟越大,max_wait越大攒批越充分但延迟越大。我一般从max_batch=8, max_wait=0.05起步,根据实际延迟和吞吐曲线调。
4.5 服务接口与超时降级
FastAPI的接口实现:
from fastapi import FastAPI, HTTPException import asyncio app = FastAPI() semaphore = asyncio.Semaphore(16) @app.post("/generate") async def generate(req: GenerateRequest): if len(req.prompt) > 2048: raise HTTPException(400, "prompt too long") async with semaphore: try: result = await asyncio.wait_for( scheduler.submit(req), timeout=10.0 ) return {"text": result} except asyncio.TimeoutError: raise HTTPException(504, "timeout")信号量限制并发16,超时10秒。这两个数字要根据你的硬件和业务SLA调。踩过的坑:一开始信号量设成100,结果GPU显存直接爆了,因为100个请求的KV Cache加起来超过了显存。后来改成16,稳定运行。
5. 常见问题与排查技巧实录
5.1 训练不收敛的排查路径
训练不收敛是最常见的问题,排查要按顺序来,别乱试:
| 现象 | 可能原因 | 排查方法 |
|---|---|---|
| loss不下降 | 学习率太小 | 调大10倍试试 |
| loss震荡 | 学习率太大 | 调小10倍试试 |
| loss变NaN | 梯度爆炸 | 加梯度裁剪 |
| loss下降但val不降 | 过拟合 | 加正则、加数据 |
| loss下降但生成质量差 | 数据分布问题 | 检查数据清洗 |
我的经验是先查数据再查模型。80%的"训练问题"其实是数据问题。我遇到过一次loss死活不降,最后发现是数据里混了一批乱码样本,清洗掉就好了。
5.2 推理延迟高的定位方法
推理延迟高,要分段测量,别猜。我在代码里埋了这些计时点:
- 请求到达时间
- 进入队列时间
- 开始推理时间
- 推理结束时间
- 返回时间
这样能算出排队延迟和推理延迟分别是多少。如果排队延迟占大头,说明并发不够或批处理没生效;如果推理延迟占大头,说明模型或KV Cache有问题。
我实测过一个案例:总延迟500ms,其中排队400ms、推理100ms。根因是信号量设太小,请求都在排队。把信号量调大后,总延迟降到150ms。所以别一上来就优化模型,先看是不是排队问题。
5.3 显存OOM的应急处理
显存OOM是推理服务的头号杀手。应急处理按这个顺序:
- 限制max_seq_len:把最大序列长度从4096降到2048,显存直接减半
- 限制并发数:信号量调小,同时处理的请求少了,KV Cache总量就小了
- 启用KV Cache淘汰:LRU淘汰最久未用的序列
- 量化:INT8量化能省一半显存,但精度会掉一点
长期方案是做显存预算:先算出模型权重占多少、每个请求的KV Cache占多少,然后反推最大并发数。公式大概是:
max_concurrent = (total_vram - model_vram - overhead) / kv_per_request我一般留20%的显存做buffer,别算得太满,否则容易OOM。
5.4 服务抖动的排查清单
服务抖动表现为延迟忽高忽低、偶发超时。排查清单:
- 检查是否有大请求混入(长prompt会拖慢整个batch)
- 检查GC是否频繁(Python的GC会暂停服务)
- 检查是否有慢查询(日志、监控的IO)
- 检查网络是否有抖动
- 检查是否有资源竞争(CPU、内存、GPU)
我遇到过一次抖动,最后发现是日志写磁盘太频繁,把IO打满了。改成异步写日志后就好了。所以观测层本身也可能成为故障源,这点很多人想不到。
6. 从零实现之后,我获得了什么
把这条链路完整走一遍之后,最大的变化是看框架源码不再发怵了。以前看vLLM的PagedAttention,觉得是天书;自己手写过KV Cache之后,再看它的分页管理,一眼就明白它在解决什么问题——无非是把连续显存换成离散分页,减少碎片。这种"一眼看穿"的能力,是调包调不出来的。
第二个变化是排障速度。以前线上出问题,只能看框架日志猜;现在能直接定位到是哪一层的哪个环节。比如延迟高,我能立刻判断是排队问题还是推理问题,不用瞎试。
第三个变化是技术选型更有底气。以前选框架看star数,现在看它的调度策略、显存管理、批处理实现,能判断它适不适合我的场景。这种判断力,是从零实现给的。
如果你也想走一遍这条路,我的建议是别追求完美,先跑通最小闭环。一个能跑的丑版本,胜过十个跑不起来的美版本。跑通之后,每个环节再慢慢优化,你会发现每一步优化都有明确的收益,这种正反馈会让你越做越有劲。
最后分享一个小技巧:每完成一个模块,写一个最小测试用例。比如KV Cache写完,测一下append之后get的长度对不对;批处理写完,测一下batch size和延迟的关系。这些测试用例积累下来,就是你重构时的安全网。我重构推理层时,就是靠这几十个测试用例,才敢大胆改代码。