news 2026/9/30 8:16:33

从零手搓AI工程流水线:KV Cache与批处理实战

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
从零手搓AI工程流水线:KV Cache与批处理实战

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 数据层:清洗比采集重要十倍

数据层我见过最多的错误是"重采集轻清洗"。大家愿意花一周写爬虫,却不愿意花一天做去重。结果就是模型在训练集上表现很好,一到真实场景就拉胯,因为训练集里全是重复样本,模型过拟合了。

从零实现数据层,我建议按这个顺序做:

  1. 格式统一:把所有来源的数据转成统一的JSONL,每行一个样本,字段固定为{"text": ..., "label": ..., "source": ...}
  2. 精确去重:用哈希(比如SHA256)对文本做精确去重,这一步能干掉30%以上的重复
  3. 近似去重:用MinHash或SimHash做近似去重,阈值我一般设在0.85,能再干掉10%左右
  4. 质量过滤:按长度、字符集、困惑度过滤掉低质样本
  5. 切分:按时间或来源切分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是推理服务的头号杀手。应急处理按这个顺序:

  1. 限制max_seq_len:把最大序列长度从4096降到2048,显存直接减半
  2. 限制并发数:信号量调小,同时处理的请求少了,KV Cache总量就小了
  3. 启用KV Cache淘汰:LRU淘汰最久未用的序列
  4. 量化: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和延迟的关系。这些测试用例积累下来,就是你重构时的安全网。我重构推理层时,就是靠这几十个测试用例,才敢大胆改代码。

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/30 8:16:00

AI模型部署优化实战:量化、剪枝与蒸馏在NVIDIA GPU上的工程落地

1. 这不是“一键优化”的魔法按钮&#xff0c;而是模型瘦身手术的主刀手册“Model-Optimizer”这个名称听起来像一个点开就能让AI模型变快变小的桌面图标——但现实恰恰相反。它既不是NVIDIA官方发布的独立软件&#xff0c;也不是某个开源项目仓库里能直接pip install的包。它是…

作者头像 李华
网站建设 2026/9/30 8:15:25

Unity Trail Renderer拖尾特效原理与工业级应用

1. 什么是Unity拖尾特效&#xff1f;它到底能解决什么实际问题&#xff1f; Unity里的拖尾特效&#xff0c;说白了就是让一个移动的物体身后“拖”出一条渐隐的光带或轨迹。它不是靠贴图滚动、不是靠粒子系统堆叠&#xff0c;而是由Unity引擎原生提供的 Trail Renderer组件 直…

作者头像 李华
网站建设 2026/9/30 8:15:23

springboot基于LSTM的股票基金可视化大屏系统 沪深300数据分析系统_xjfo390f

目录同行可拿货,招校园代理 ,本人源头供货商项目背景与目标技术架构概览核心功能模块数据流与系统流程系统优势适用场景项目代码结构示意扩展建议项目技术支持获取博主联系方式 源码获取详细视频演示 &#xff1a;同行可合作点击我获取源码->获取博主联系方式->进我个人主…

作者头像 李华
网站建设 2026/9/30 8:15:04

量子算法与Python全栈:从入门到落地的完整实践

量子计算这几年已经从论文里的数学公式&#xff0c;逐渐变成了可以真实运行的代码和工具链。而Python几乎是整个量子计算生态的统一入口&#xff0c;Qiskit、Cirq、PennyLane这些主流框架都以Python为第一语言。很多人一听到“量子算法开发”就觉得门槛高&#xff0c;觉得要懂一…

作者头像 李华
网站建设 2026/9/30 8:14:58

从零手搓AI工程链路:数据管道、训练编排与推理服务实战

1. 从零手搓AI工程&#xff1a;为什么我不建议你直接调包 第一次看到 ai-engineering-from-scratch 这个项目名的时候&#xff0c;我正坐在工位上啃一个调了三天都没收敛的推荐模型。当时第一反应是&#xff1a;又来了一个“从零实现”的教程仓库。干这行十来年&#xff0c;见…

作者头像 李华
网站建设 2026/9/30 8:14:14

Android无操作超时自动返回登录页:从触摸监听到生命周期管理全解析

做Android开发这些年&#xff0c;接到过不少业务上“奇怪”的需求&#xff0c;其中“Android无操作超时返回登录界面”算是一个看似简单、实则细节很多的典型。最开始接到这个需求时&#xff0c;我心里想的是“定时器加跳转页面&#xff0c;不就完事了&#xff1f;”等真把代码…

作者头像 李华