
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个tokentop-pnucleus采样保留累积概率达到p的最小token集合。我一般用top-p0.9 temperature0.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_len10, max_len2048): 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_asciiFalse) \n) kept 1 print(fkept {kept} samples)这个脚本跑一遍你能直观看到去重率。我实测过一个中文语料精确去重干掉了35%的样本说明原始数据里重复非常严重。近似去重我用MinHash核心代码如下import hashlib def shingles(text, k3): return {text[i:ik] for i in range(len(text) - k 1)} def minhash(shingle_set, num_hashes128): 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(axis0) 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(axis0) 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 cosinedef 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_batch8, max_wait0.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_batch8, max_wait0.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), timeout10.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和延迟的关系。这些测试用例积累下来就是你重构时的安全网。我重构推理层时就是靠这几十个测试用例才敢大胆改代码。