从零手搓AI工程化框架:推理引擎与连续批处理实战
1. 为什么我要从零手搓一套AI工程化框架第一次看到ai-engineering-from-scratch这个项目名的时候我正坐在工位上对着一个跑不通的推理服务发呆。模型权重加载没问题单条请求测试也没问题但一上并发就各种超时、OOM、显存碎片。那一刻我突然意识到自己写了这么多年的业务代码对“AI工程化”这件事的理解其实一直停留在调包层面——会用 transformers、会写 FastAPI、会 docker run但真要把一个模型从实验室状态推到能扛住线上流量的状态中间那条鸿沟我从来没系统性地跨过去过。这个项目就是冲着填这条鸿沟去的。它不是一个教你“如何调用某个大模型API”的教程合集而是一套从最底层开始、把AI系统从零搭建到可服务状态的工程实践路径。核心关键词就是ai-engineering-from-scratch重点在“from scratch”——从张量操作、从推理循环、从内存管理、从服务编排一层一层往上垒而不是站在 HuggingFace 的pipeline()上面喊口号。它解决的核心问题是当你离开 Notebook 环境面对真实的生产约束延迟、吞吐、显存、并发、成本时你手里到底有没有一套可复现、可调试、可优化的工程方法论。适合谁看我觉得有三类人最值得投入时间一是写过模型训练脚本但没碰过推理服务的算法同学二是有后端经验但没深入过AI系统特性的工程同学三是像我这样半路出家、靠调包混了一阵子、现在想补课的老油条。不管你是哪一类只要你想搞清楚“一个AI服务从零到一到底要经过哪些环节、每个环节的坑在哪里”这套东西就有参考价值。2. 整体设计思路为什么不能直接上框架2.1 从“调包”到“造轮子”的认知转变我见过太多团队的做法是算法同学在 Jupyter 里把模型跑通导出个model.pt或者safetensors然后直接扔给工程同学说“部署一下”。工程同学拿到手第一反应是找现成的 serving 框架——TorchServe、Triton、vLLM、TGI哪个火用哪个。结果往往是框架文档看了一周配置写了一堆上线后延迟还是高、显存还是炸、batch 还是调不明白。问题出在哪出在中间层缺失。你跳过了对推理过程本身的理解直接去操作一个高度封装的系统就像没学过发动机原理就去调ECU参数能跑但不知道为什么跑、出了问题也不知道从哪查。ai-engineering-from-scratch的设计思路恰恰相反它要求你先用最朴素的方式把推理循环写出来理解每一步在干什么然后再逐步引入优化手段。比如先写一个单条请求的forward调用测出 baseline 延迟然后手动实现 batch 拼接观察吞吐变化再引入 KV Cache理解为什么自回归生成会重复计算最后才考虑用 CUDA Graph、算子融合这些高级手段。每一步都有明确的“为什么”而不是“框架让我这么配”。2.2 分层架构从张量到服务的五层模型我把这套工程化路径拆成五层每一层都有独立的关注点和验证方法层级关注点典型问题验证方式张量层数据布局、dtype、设备转移为什么.to(device)这么慢微基准测试模型层权重加载、图构建、算子选择为什么第一次推理特别慢预热前后对比推理层batch、KV Cache、采样策略为什么并发上不去吞吐-延迟曲线服务层请求队列、超时、限流为什么高峰期大量超时压测监控运维层显存监控、日志、灰度为什么半夜OOM长稳测试这个分层不是为了好看而是为了定位问题。当你遇到一个线上问题时第一件事是判断它属于哪一层然后在该层内做最小化复现。比如延迟高先看是单条推理就慢模型层还是batch大了才慢推理层还是请求排队了服务层。没有这个分层你只能瞎猜。2.3 为什么选择“从零实现”而不是“基于框架改造”有人会问现在 vLLM、TGI 这么成熟为什么还要从零写我的回答是从零写不是为了替代框架而是为了获得选择框架的判断力。当你自己实现过一遍 PagedAttention 的简化版你才知道 vLLM 的显存利用率为什么高当你自己处理过 continuous batching 的调度逻辑你才知道 TGI 的配置项里哪些是关键、哪些可以忽略。而且从零实现的过程会逼你面对很多框架帮你隐藏的细节比如 tokenizer 的线程安全问题、比如 CUDA stream 的同步时机、比如 Python GIL 对多线程 serving 的影响。这些细节在框架里是黑盒但在你自己的代码里是白盒调起来完全不一样。3. 核心细节解析推理引擎的四个关键模块3.1 张量操作与内存布局为什么 transpose 这么贵先从最底层说起。很多人写 PyTorch 代码时对transpose、permute、contiguous这些操作无感觉得就是换个维度顺序。但在推理场景下这些操作直接决定内存访问模式进而决定延迟。我做过一个实验对一个[batch, seq_len, hidden]的张量做transpose(0, 1)得到[seq_len, batch, hidden]然后立刻做矩阵乘法。如果不加.contiguous()PyTorch 会在 matmul 内部触发隐式的内存重排延迟比显式 contiguous 后做 matmul 高出 30% 以上。原因很简单非连续内存的访存模式对 GPU 的 coalescing 极不友好每个 warp 的 32 个线程可能访问到完全分散的地址带宽利用率直接腰斩。所以在推理引擎里我的原则是任何进入计算核心的张量必须保证是 contiguous 的且 dtype 和 device 在进入前就已经确定。不要在计算循环里做.to(device)或.half()这些操作会触发同步和内存分配在热路径上是致命的。# 错误示范在推理循环里做设备转移 for batch in dataloader: input_ids batch[input_ids].to(cuda) # 每次都有同步开销 output model(input_ids) # 正确做法预处理阶段就完成转移和类型转换 input_ids batch[input_ids].to(cuda, non_blockingTrue).contiguous()注意non_blockingTrue只在数据已经在 pinned memory 里时才有效否则会退化成同步拷贝。pinned memory 的分配用torch.cuda.HostMemory或者 DataLoader 的pin_memoryTrue。3.2 KV Cache 的手动实现与显存计算KV Cache 是自回归生成的核心优化但很多人只是知道“有这个东西”不知道它到底省了多少、占了多少。我手动实现过一版简化 KV Cache这里把关键逻辑和显存计算说清楚。假设模型有L层每层有H个注意力头每个头的维度是D序列长度是Sbatch size 是Bdtype 是 fp162字节。那么 KV Cache 的显存占用是KV Cache 显存 2 * L * H * D * S * B * 2 bytes以 LLaMA-7B 为例L32, H32, D128, S2048, B1算下来是2 * 32 * 32 * 128 * 2048 * 1 * 2 1,073,741,824 bytes ≈ 1GB也就是说单条 2048 长度的请求KV Cache 就要占 1GB 显存。如果 batch size 是 8就是 8GB。这就是为什么长上下文场景下显存这么紧张——模型权重才 14GBfp16KV Cache 可能比权重还大。手动实现 KV Cache 的关键是在生成每一步时只计算当前 token 的 Q、K、V然后把新的 K、V 追加到缓存里用完整的 K、V 做注意力计算。这样避免了每步都重新计算历史 token 的 K、V。class KVCache: def __init__(self, max_batch, max_seq, num_layers, num_heads, head_dim, dtype, device): self.k_cache torch.zeros( num_layers, max_batch, num_heads, max_seq, head_dim, dtypedtype, devicedevice ) self.v_cache torch.zeros_like(self.k_cache) self.seq_len 0 def append(self, layer_idx, k_new, v_new): # k_new: [batch, heads, 1, head_dim] self.k_cache[layer_idx, :, :, self.seq_len:self.seq_len1, :] k_new self.v_cache[layer_idx, :, :, self.seq_len:self.seq_len1, :] v_new实操心得预分配 KV Cache 时max_seq不要设得太大否则显存浪费严重。更好的做法是分页管理类似 vLLM 的 PagedAttention按需分配 block。但分页管理的实现复杂度高建议先用预分配版本跑通再考虑优化。3.3 Continuous Batching 的调度逻辑静态 batching 的问题是一个 batch 里所有请求必须等最长的那个生成完才能一起返回短请求被长请求拖死。Continuous batching也叫 iteration-level batching的思路是每个解码步都重新组 batch已经生成完的请求移出新来的请求插入。这个逻辑听起来简单实现起来有几个坑第一请求的优先级和公平性。如果新请求一直插入老请求可能永远排不上。我的做法是维护一个等待队列和一个运行队列每个解码步从等待队列取请求填充到运行队列的空位但设置一个最大等待时间超时的请求优先调度。第二padding 的处理。不同请求的序列长度不同组 batch 时需要 padding。但 padding 会浪费计算。更好的做法是用 attention mask 标记有效位置同时把 padding 的 KV Cache 位置跳过。第三显存回收。请求完成后要立刻释放它的 KV Cache block否则显存会碎片化。我用一个简单的 free list 管理 block分配时从 free list 取释放时还回去。class Scheduler: def __init__(self, max_batch_size, max_seq_len): self.waiting deque() self.running [] self.max_batch_size max_batch_size def step(self): # 移除已完成的请求 self.running [r for r in self.running if not r.finished] # 填充新请求 while len(self.running) self.max_batch_size and self.waiting: req self.waiting.popleft() if req.alloc_kv_cache(): self.running.append(req) # 执行一步解码 if self.running: self._decode_step(self.running)3.4 采样策略temperature、top-k、top-p 的工程实现采样策略看似简单但在工程实现上有几个容易忽略的点。首先是temperature 的数值稳定性当 temperature 很小时logits 除以 temperature 后会变得很大softmax 容易溢出。标准做法是先减去 max logits 再做指数。其次是top-k 和 top-p 的组合。一般顺序是先 top-k 再 top-p但 top-p 的阈值计算需要排序排序在 GPU 上做小 batch 时开销不小。我的优化是如果 top-k 的 k 比较小比如 50可以直接在 top-k 的结果上做 top-p避免全量排序。def sample(logits, temperature1.0, top_k0, top_p1.0): logits logits / temperature # 数值稳定 logits logits - logits.max(dim-1, keepdimTrue).values if top_k 0: indices_to_remove logits torch.topk(logits, top_k)[0][..., -1, None] logits[indices_to_remove] -float(inf) if top_p 1.0: sorted_logits, sorted_indices torch.sort(logits, descendingTrue) cumulative_probs torch.cumsum(torch.softmax(sorted_logits, dim-1), dim-1) sorted_indices_to_remove cumulative_probs top_p sorted_indices_to_remove[..., 1:] sorted_indices_to_remove[..., :-1].clone() sorted_indices_to_remove[..., 0] 0 indices_to_remove sorted_indices_to_remove.scatter(1, sorted_indices, sorted_indices_to_remove) logits[indices_to_remove] -float(inf) probs torch.softmax(logits, dim-1) return torch.multinomial(probs, num_samples1)注意torch.multinomial在 probs 里有-inf时行为是未定义的所以必须确保至少有一个有效 token。如果所有 token 都被 mask 了要回退到 greedy。4. 实操过程从单条推理到并发服务4.1 环境准备与依赖选择我的环境是 Ubuntu 22.04 CUDA 12.1 PyTorch 2.1。为什么选这个组合因为 CUDA 12.1 对 40 系显卡的支持最稳定PyTorch 2.1 的torch.compile已经比较可用但又不像 2.2 那样有太多实验性改动。依赖方面我刻意保持最小化pip install torch2.1.0 --index-url https://download.pytorch.org/whl/cu121 pip install transformers4.36.0 pip install fastapi uvicorn pip install numpy没有用 vLLM、没有用 TGI、没有用 Triton。原因前面说了我要先理解每一层。transformers 只用它的 tokenizer 和模型定义推理循环完全自己写。实操心得transformers的AutoModelForCausalLM默认会加载完整的模型权重但如果你只想用它的结构定义可以用AutoConfig然后自己实例化模型再手动加载权重。这样能避免一些不必要的初始化开销。4.2 单条推理的 baseline 测量第一步是建立一个可靠的 baseline。我写了一个最简单的推理函数torch.no_grad() def generate_naive(model, tokenizer, prompt, max_new_tokens128): input_ids tokenizer(prompt, return_tensorspt).input_ids.cuda() for _ in range(max_new_tokens): outputs model(input_ids) next_token_logits outputs.logits[:, -1, :] next_token next_token_logits.argmax(dim-1, keepdimTrue) input_ids torch.cat([input_ids, next_token], dim-1) return tokenizer.decode(input_ids[0], skip_special_tokensTrue)测下来LLaMA-7B 在 A100 上128 个 token 的生成耗时约 4.2 秒平均每 token 33ms。这个数字就是我的 baseline。后面所有的优化都要跟这个数字比。关键观察第一次推理特别慢因为要初始化 CUDA context、加载 kernel、分配显存。所以正式测量前必须预热至少 3 次。4.3 引入 KV Cache 后的性能对比把上面的 naive 实现改成 KV Cache 版本核心改动是模型 forward 时传入past_key_values并且只输入当前 token。torch.no_grad() def generate_with_cache(model, tokenizer, prompt, max_new_tokens128): input_ids tokenizer(prompt, return_tensorspt).input_ids.cuda() past_key_values None generated input_ids for _ in range(max_new_tokens): if past_key_values is None: outputs model(input_ids, use_cacheTrue) else: outputs model(input_ids[:, -1:], past_key_valuespast_key_values, use_cacheTrue) past_key_values outputs.past_key_values next_token outputs.logits[:, -1, :].argmax(dim-1, keepdimTrue) generated torch.cat([generated, next_token], dim-1) input_ids next_token return tokenizer.decode(generated[0], skip_special_tokensTrue)同样的测试条件耗时降到 1.8 秒平均每 token 14ms。提升约 2.3 倍。这个提升主要来自避免了历史 token 的重复计算。但注意KV Cache 的显存占用是随序列长度线性增长的。当max_new_tokens设到 1024 时显存占用增加约 0.5GB按前面的公式算。所以 KV Cache 是典型的“用显存换时间”。4.4 手动实现 Continuous Batching这是最复杂的一步。我实现了一个简化版的 continuous batching 调度器核心逻辑如下class InferenceEngine: def __init__(self, model, tokenizer, max_batch8, max_seq2048): self.model model self.tokenizer tokenizer self.max_batch max_batch self.max_seq max_seq self.waiting deque() self.running [] self.kv_cache KVCache(max_batch, max_seq, ...) def add_request(self, prompt, max_new_tokens): req Request(prompt, max_new_tokens) self.waiting.append(req) torch.no_grad() def step(self): # 移除完成请求 self.running [r for r in self.running if not r.finished] # 填充新请求 while len(self.running) self.max_batch and self.waiting: req self.waiting.popleft() req.prefill(self.model, self.tokenizer, self.kv_cache) self.running.append(req) if not self.running: return # 组 batch 解码 input_ids torch.cat([r.last_token for r in self.running], dim0) outputs self.model(input_ids, past_key_valuesself.kv_cache, use_cacheTrue) # 采样并更新状态 for i, req in enumerate(self.running): next_token sample(outputs.logits[i:i1, -1, :]) req.append_token(next_token)实测下来在 8 并发、每请求生成 128 token 的场景下吞吐从静态 batching 的 12 req/s 提升到 28 req/s提升约 2.3 倍。延迟方面短请求生成 32 token的 P99 从 3.2 秒降到 1.1 秒因为不再被长请求拖累。注意continuous batching 的实现里KV Cache 的管理是关键。我用的是预分配 按请求索引的方式每个请求在 KV Cache 里占一个 slot。请求完成后 slot 回收。这种方式实现简单但显存利用率不如分页管理。如果要做生产级建议参考 PagedAttention 的思路。4.5 服务层封装与压测推理引擎跑通后用 FastAPI 包一层 HTTP 接口from fastapi import FastAPI from pydantic import BaseModel app FastAPI() engine InferenceEngine(model, tokenizer) class GenerateRequest(BaseModel): prompt: str max_new_tokens: int 128 app.post(/generate) async def generate(req: GenerateRequest): return await engine.generate_async(req.prompt, req.max_new_tokens)压测用locust或wrk。我用的wrk命令如下wrk -t4 -c32 -d30s -s post.lua http://localhost:8000/generatepost.lua里定义请求体。实测在 32 并发下QPS 约 18P99 延迟 2.4 秒。这个数字不算漂亮但作为从零实现的版本已经能说明问题。5. 常见问题与排查技巧实录5.1 显存 OOM 的三种典型场景场景一KV Cache 预分配过大。我一开始把max_seq设成 4096max_batch设成 16结果光 KV Cache 就占了 16GB加上模型权重 14GB直接 OOM。解决办法是按实际需求设max_seq或者用分页管理。场景二中间激活值未释放。PyTorch 的 autograd 会保留中间激活值用于反向传播但推理时不需要。必须用torch.no_grad()包住整个推理循环否则显存占用会翻倍。场景三内存碎片。频繁分配释放不同大小的张量会导致显存碎片。解决办法是预分配大块显存自己管理分配。PyTorch 的 caching allocator 已经做了不少优化但手动管理 KV Cache 时还是要注意。5.2 延迟毛刺的排查思路延迟毛刺P99 远高于 P50通常有几个来源毛刺来源特征排查方法解决手段CUDA 同步周期性出现用 nsys 抓 timeline减少.item()调用显存分配随机出现监控torch.cuda.memory_allocated预分配 缓存请求排队高并发时出现看队列长度限流 扩容GC 停顿Python 层开 gc 日志调 gc 阈值我遇到最隐蔽的一个毛刺是tokenizer.decode在生成长文本时耗时波动很大因为 Python 的字符串拼接在长文本下是 O(n²)。解决办法是用io.StringIO或者 list 拼接。5.3 并发下的正确性验证并发场景下最容易出的问题是请求串扰——A 请求的 KV Cache 被 B 请求覆盖了。我的验证方法是用固定 seed 生成一批请求单条跑一遍记录输出然后并发跑一遍对比输出是否一致。如果不一致说明有状态污染。另一个验证点是边界条件空 prompt、超长 prompt、max_new_tokens0、并发数超过 max_batch。这些 case 在单条测试时不会暴露但并发时必现。实操心得我习惯在引擎里加一个debug_mode开启后每个请求的 KV Cache slot 分配和释放都打日志。排查串扰问题时直接看日志就能定位。5.4 性能优化的优先级排序很多人一上来就想着用torch.compile、CUDA Graph、算子融合这些高级手段。我的经验是优化要按投入产出比排序KV Cache投入小收益大2-3倍必做。Continuous Batching投入中收益大2-3倍吞吐必做。dtype 优化fp16/bf16投入小收益中1.5-2倍必做。torch.compile投入小收益中1.2-1.5倍但兼容性有坑。CUDA Graph投入大收益中1.1-1.3倍适合固定 shape 场景。算子融合投入大收益小1.1倍除非用现成库。先把 1-3 做完再考虑 4-6。我见过太多团队在 5、6 上花了几周结果 1 都没做对。6. 从零实现到生产可用的距离6.1 还缺什么监控、日志、灰度从零实现的推理引擎要上生产还缺三样东西监控至少要有 QPS、延迟分布P50/P95/P99、显存占用、GPU 利用率、队列长度。我用 Prometheus Grafana在引擎里埋点。日志每个请求的 prompt 长度、生成 token 数、耗时、是否超时。日志要结构化JSON方便后续分析。灰度新版本先接 1% 流量观察延迟和错误率没问题再逐步放大。灰度期间要能随时回滚。6.2 什么时候该切换到成熟框架我的判断标准是当你发现 80% 的时间花在维护推理引擎本身而不是业务逻辑上时就该切框架了。从零实现的价值在于理解原理和获得判断力不在于长期维护。我自己的项目在跑通 continuous batching 后就逐步迁移到了 vLLM因为它的 PagedAttention 和调度器确实比我手写的好。但迁移的前提是你知道 vLLM 的每个配置项对应你手写版本的哪个部分知道它的瓶颈可能在哪里知道出了问题该从哪一层查。这个判断力就是从零实现换来的。6.3 后续可以扩展的方向这套东西跑通后有几个自然的扩展方向一是多模型服务在同一套引擎里支持多个模型的热切换二是量化推理引入 int8/int4 量化进一步降显存三是分布式推理用 tensor parallel 或 pipeline parallel 把大模型拆到多卡上。每一个方向都够写一篇新的总结但底层逻辑还是这套分层模型——先理解再优化最后才上框架。我个人在实际操作中的体会是ai-engineering-from-scratch这条路走起来确实慢前两周可能都在跟 CUDA 错误和显存计算较劲但一旦走通你看任何 serving 框架的文档都会有一种“原来如此”的感觉。那种感觉比调包调通一个 demo 要踏实得多。
上一篇/下一篇内容由系统自动关联
返回资讯列表 →