ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

从零手搓AI工程流水线:推理服务、KV Cache与批处理调度实战

从零手搓AI工程流水线:推理服务、KV Cache与批处理调度实战 1. 为什么我要从零手搓一套AI工程流水线第一次看到ai-engineering-from-scratch这个项目名的时候我正被一堆调包侠式的教程搞得有点烦。不是说调包不好而是当你把model.fit()和pipeline()用得太顺手之后一旦线上推理延迟飙到 800ms、显存莫名其妙涨到 22G、batch 里混进一条超长文本直接把整个服务拖垮你会发现自己对这套系统内部到底发生了什么其实一无所知。ai-engineering-from-scratch这个标题字面意思就是从零开始的AI工程。它不是一个具体的库也不是某个框架的官方教程而是一类项目的统称——把AI从能跑通demo推进到能扛住线上流量的完整工程链路全部自己动手搭一遍。核心关键词就是AI工程、从零实现、推理服务、训练流水线、工程化落地。它解决的问题很明确市面上讲模型原理的书一大堆讲工程落地的却往往默认你已经会了而真正卡住绝大多数人的恰恰是中间那层工程胶水。这套东西适合谁如果你已经能用 PyTorch 或类似框架写出一个能训练的模型但一提到部署并发显存优化数据管道就发怵那这个方向就是为你准备的。如果你是完全的新手也别急着走我会把每一步的为什么讲清楚你至少能建立起一张完整的工程地图。我自己的背景是做了几年后端转AI之后踩了无数坑这篇就把我从零搭这套流水线时真正有用的东西掏出来讲。2. 整体架构设计与技术选型思路2.1 为什么坚持从零而不是直接上现成框架很多人第一反应是有 vLLM、有 Triton、有 TorchServe我为什么要自己写这个问题我认真想过。答案是用现成框架解决的是交付从零实现解决的是理解。这两件事不冲突但顺序不能反。我试过直接上推理框架结果遇到一个诡异问题同样的模型框架A比框架B慢三倍。我完全不知道从哪查起因为中间隔了太多抽象层。后来我花了两周自己用最朴素的方式写了一个推理循环——手动管理 batch、手动做 padding、手动控制 KV cache 的生命周期——写完那一刻我才真正明白框架里那些参数到底在调什么。所以这套流水线的设计原则是先手写最小可用版本再逐步替换成成熟组件。每一步替换你都知道自己在换掉什么、换来什么。这跟学开车一个道理你先得知道离合是干嘛的再去开自动挡才心里有底。具体来说整个架构我分成四层从下往上数据层负责原始数据的读取、清洗、分词、打包成 batch。这一层最容易被忽视但它往往是线上事故的重灾区。模型层模型定义、权重加载、精度管理fp32/fp16/int8。这一层决定了你的显存天花板。推理/训练层前向传播、反向传播、KV cache 管理、梯度累积。这是计算的核心。服务层请求接入、批处理调度、超时控制、监控埋点。这一层决定了你的系统能不能扛住真实流量。2.2 技术栈的取舍哪些自己写哪些用现成的我的原则是计算密集的部分用成熟库调度和胶水部分自己写。因为计算部分矩阵乘法、卷积已经被优化到极致了你自己写不可能更快但调度逻辑是跟你的业务强相关的现成方案往往水土不服。模块自己实现用现成库理由张量运算否PyTorch/NumPy底层算子已高度优化分词部分HuggingFace Tokenizers分词算法固定但缓存策略自己写数据管道是否与业务数据格式强相关批处理调度是否核心工程价值所在KV cache 管理是否直接影响显存和延迟模型定义否PyTorch重复造轮子无意义监控埋点是Prometheus客户端指标定义自己定这张表是我踩坑之后总结的。一开始我什么都想自己写连矩阵乘法都想手搓结果浪费了一周。后来想通了从零的目的是理解原理不是拒绝一切工具。你要能解释清楚每一层在干什么而不是每一层都亲手实现。2.3 一个容易被忽略的设计决策同步还是异步这是我在项目早期纠结最久的问题。同步推理实现简单一个请求进来算完返回逻辑清晰。但问题是 GPU 利用率极低——大部分时间 GPU 在等数据搬运而不是在算。异步 批处理是必然选择但异步带来的复杂度是成倍上升的。我的做法是分阶段第一阶段先写同步版本把正确性跑通第二阶段引入请求队列和动态批处理第三阶段再加超时和降级。千万不要一上来就搞全异步你会被各种竞态条件折磨到怀疑人生。提示动态批处理的核心是攒批——不是来一个请求处理一个而是等一小段时间比如 10ms把这段时间内到达的请求凑成一个 batch 一起算。这个等待时间就是延迟和吞吐的权衡旋钮。3. 核心模块的细节拆解与实操要点3.1 数据管道90%的线上问题都出在这里我先说一个真实教训。有次线上服务突然大面积超时查了半天模型没问题、GPU没问题最后发现是数据管道里有个正则表达式在处理某类特殊字符时发生了灾难性回溯单条数据卡了 3 秒。模型再快也扛不住这个。数据管道的核心环节有这么几个每个都有坑读取与清洗。原始数据往往脏得超乎想象。我的做法是清洗逻辑全部写成纯函数每个函数只做一件事方便单测。比如remove_html_tags、normalize_whitespace、truncate_by_tokens分开写而不是揉成一个大函数。分词与缓存。分词本身不慢但重复分词很慢。我加了一层 LRU 缓存key 是原始文本的哈希value 是 token id 列表。实测在重复率高的场景下分词耗时能降 60% 以上。打包成 batch。这里的关键是 padding 策略。最朴素的是 pad 到 batch 内最长序列但这样短序列会浪费大量计算。更好的做法是按长度分桶——把长度相近的样本放在同一个 batch 里减少 padding 浪费。我实测下来分桶之后 GPU 利用率能提升 30% 左右。# 按长度分桶的简化实现 def bucket_by_length(samples, bucket_size32): # 先按长度排序 samples.sort(keylambda x: len(x[input_ids])) buckets [] for i in range(0, len(samples), bucket_size): buckets.append(samples[i:i bucket_size]) return buckets这段代码看着简单但有个细节排序会打乱样本顺序。如果训练时对顺序敏感比如某些时序任务你需要在 batch 内部再打乱或者记录原始索引。我一开始就忘了这点导致训练效果异常查了两天才发现。3.2 模型加载与精度管理显存是怎么被吃掉的模型加载看着简单其实门道很多。一个 7B 参数的模型fp32 下光权重就要 28GBfp16 是 14GBint8 是 7GB。这还没算激活值、KV cache 和优化器状态。精度选择是个权衡。fp16 能省一半显存但数值范围小容易溢出bf16 范围大但精度低int8 最省但需要量化校准。我的建议是推理优先用 fp16 或 bf16训练用混合精度。如果显存实在紧张再考虑量化。加载权重的时候有个技巧用mmap方式加载而不是一次性读进内存。这样多个进程可以共享同一份权重文件省内存。PyTorch 的torch.load配合map_location就能做到。# 分片加载大模型的思路 def load_model_sharded(model_path, device): state_dict {} for shard_file in sorted(os.listdir(model_path)): if shard_file.endswith(.safetensors): shard load_file(os.path.join(model_path, shard_file), devicecpu) state_dict.update(shard) model.load_state_dict(state_dict) model.to(device) return model注意加载完之后记得把 CPU 上的临时变量释放掉否则峰值内存会是模型大小的两倍。我见过有人加载 13B 模型时机器直接 OOM就是因为没释放中间变量。3.3 KV Cache推理加速的关键也是显存杀手如果你只从这篇文章记住一件事我希望是 KV cache。自回归生成时每生成一个 token 都要重新计算前面所有 token 的 attention这是巨大的浪费。KV cache 的思路是把已经算过的 Key 和 Value 存下来下一个 token 只算新的部分。原理不复杂但工程实现有几个坑显存占用估算。KV cache 的大小 2 × 层数 × 头数 × 头维度 × 序列长度 × batch_size × 精度字节数。以一个 7B 模型为例32 层、32 头、头维度 128、fp16序列长度 2048、batch 8算下来大概是 2 × 32 × 32 × 128 × 2048 × 8 × 2 字节 ≈ 8.6GB。这还没算模型权重显存一下就紧张了。动态增长与预分配。KV cache 随序列增长而增长如果动态分配会产生大量内存碎片。更好的做法是预分配一个最大长度的 buffer用多少取多少。代价是显存利用率低但换来了稳定。PagedAttention 的思想。这是后来很多推理框架的核心优化把 KV cache 分成固定大小的块page像操作系统管理内存一样管理它。这样不同序列可以共享块碎片问题大大缓解。我建议你至少理解这个思想哪怕不自己实现。3.4 批处理调度吞吐和延迟的平衡术调度器是整个服务的大脑。它的核心任务就一个在延迟可接受的前提下尽可能把 GPU 喂饱。我实现过的最简调度器逻辑是这样的请求进来放进等待队列。调度线程每隔T毫秒比如 10ms唤醒一次。从队列里取出最多max_batch_size个请求组成一个 batch。如果队列为空继续睡如果有请求但不够一个 batch也发出去避免饿死。执行推理返回结果。这个逻辑简单但有几个参数需要仔细调T攒批窗口太小则 batch 小、吞吐低太大则延迟高。我一般从 10ms 起调。max_batch_size受显存限制。要结合 KV cache 的估算来定。超时时间单个请求等待超过这个时间就必须发出去哪怕 batch 没满。实操心得我建议把攒批窗口做成动态的——队列长的时候窗口小一点快速响应队列短的时候窗口大一点攒大 batch。这个策略在流量波动大的场景下效果很明显。4. 完整实操流程从空目录到能跑的服务4.1 环境准备与依赖安装先把地基打好。我用的环境是 Python 3.10 PyTorch 2.x CUDA 12.x。版本匹配很重要CUDA 版本和 PyTorch 版本对不上是最常见的坑。# 创建虚拟环境 python -m venv venv source venv/bin/activate # 安装 PyTorch注意 CUDA 版本要匹配 pip install torch --index-url https://download.pytorch.org/whl/cu121 # 安装其他依赖 pip install transformers safetensors fastapi uvicorn prometheus-client装完之后一定要验证 GPU 可用import torch print(torch.cuda.is_available()) print(torch.cuda.get_device_name(0)) print(torch.cuda.get_device_properties(0).total_memory / 1e9, GB)如果is_available()返回 False别急着往下走先把驱动和 CUDA 版本对齐。这一步卡住的人特别多。4.2 数据管道的搭建我按前面说的分桶思路搭一个最小版本。核心是三个类Reader负责读数据Cleaner负责清洗Batcher负责组批。class DataPipeline: def __init__(self, tokenizer, max_length512, bucket_size32): self.tokenizer tokenizer self.max_length max_length self.bucket_size bucket_size self.cache {} def tokenize(self, text): # 带缓存的 tokenize key hash(text) if key not in self.cache: ids self.tokenizer.encode(text, truncationTrue, max_lengthself.max_length) self.cache[key] ids return self.cache[key] def build_batches(self, texts): samples [{input_ids: self.tokenize(t)} for t in texts] samples.sort(keylambda x: len(x[input_ids])) batches [] for i in range(0, len(samples), self.bucket_size): batch samples[i:i self.bucket_size] batches.append(self.pad_batch(batch)) return batches def pad_batch(self, batch): max_len max(len(s[input_ids]) for s in batch) input_ids [] attention_mask [] for s in batch: pad_len max_len - len(s[input_ids]) input_ids.append(s[input_ids] [0] * pad_len) attention_mask.append([1] * len(s[input_ids]) [0] * pad_len) return {input_ids: input_ids, attention_mask: attention_mask}这段代码里attention_mask是关键它告诉模型哪些位置是真实 token、哪些是 padding。忘了传这个模型会把 padding 也当有效输入结果就是输出莫名其妙。4.3 推理循环的实现推理循环我分两步走先写单条推理再改成批处理。单条推理的核心是自回归生成torch.no_grad() def generate_single(model, input_ids, max_new_tokens100): generated input_ids for _ in range(max_new_tokens): outputs model(generated) next_token_logits outputs.logits[:, -1, :] next_token torch.argmax(next_token_logits, dim-1, keepdimTrue) generated torch.cat([generated, next_token], dim-1) if next_token.item() model.config.eos_token_id: break return generated这个版本能跑但慢得让人抓狂因为每一步都在重算全部 attention。加上 KV cache 之后torch.no_grad() def generate_with_cache(model, input_ids, max_new_tokens100): past_key_values None generated input_ids for _ in range(max_new_tokens): outputs model(generated, past_key_valuespast_key_values, use_cacheTrue) past_key_values outputs.past_key_values next_token torch.argmax(outputs.logits[:, -1, :], dim-1, keepdimTrue) generated torch.cat([generated, next_token], dim-1) # 关键下一步只输入新 token generated next_token if next_token.item() model.config.eos_token_id: break return generated注意这里有个容易搞错的地方用了 cache 之后下一步的输入只能是新生成的 token而不是完整序列。我第一次写的时候没注意把完整序列又传进去了结果输出完全错乱。4.4 服务层的封装服务层我用 FastAPI 起一个 HTTP 接口前面加一个调度器。from fastapi import FastAPI import asyncio app FastAPI() request_queue asyncio.Queue() app.post(/generate) async def generate(prompt: str): future asyncio.Future() await request_queue.put((prompt, future)) return await future async def scheduler_loop(): while True: batch [] try: # 攒批窗口 10ms while len(batch) MAX_BATCH_SIZE: item await asyncio.wait_for(request_queue.get(), timeout0.01) batch.append(item) except asyncio.TimeoutError: pass if batch: prompts [b[0] for b in batch] results run_inference(prompts) for (_, future), result in zip(batch, results): future.set_result(result)这个调度器是异步的asyncio.wait_for的超时就是攒批窗口。实测下来在 QPS 50 左右的场景这个简单调度器能把 GPU 利用率从 20% 拉到 70% 以上。4.5 监控埋点没有监控的服务就是裸奔。我至少会埋这几个指标请求延迟P50、P95、P99用直方图记录。batch 大小分布看调度器有没有正常工作。GPU 显存和利用率用pynvml采集。队列长度队列持续增长说明处理不过来要告警。from prometheus_client import Histogram, Gauge REQUEST_LATENCY Histogram(request_latency_seconds, Request latency) BATCH_SIZE Histogram(batch_size, Batch size distribution) QUEUE_LENGTH Gauge(queue_length, Current queue length)埋点这件事我的经验是宁可多埋不要少埋。线上出问题时你永远不知道哪个指标会成为救命稻草。5. 常见问题与排查技巧实录5.1 显存溢出OOM的排查路径OOM 是最常见的问题排查要按顺序来先看是不是模型本身太大。用torch.cuda.memory_allocated()看权重占了多少。再看 KV cache。如果序列长度或 batch 设得太大cache 会爆。最后看碎片。torch.cuda.memory_reserved()和allocated()差距大说明有碎片可以试试torch.cuda.empty_cache()。我整理了一个速查表现象可能原因解决方向加载模型就OOM精度太高换 fp16/bf16 或量化推理中途OOMKV cache 增长限制 max_length 或 batch显存缓慢增长缓存未释放检查是否有引用泄漏reserved远大于allocated内存碎片预分配或重启服务5.2 输出乱码或重复的排查模型输出重复、乱码通常不是模型的问题而是推理逻辑的问题。我遇到过几次忘了传 attention_maskpadding 被当成有效输入输出错乱。用了 KV cache 但输入没截断重复计算输出重复。采样参数问题temperature 太低会重复太高会乱码。排查方法很简单先用贪心解码argmax跑一遍如果贪心正常那就是采样参数的问题如果贪心也乱那就是推理逻辑的问题。5.3 延迟忽高忽低的排查延迟抖动大八成是调度器的问题。我遇到过几种情况攒批窗口固定流量大时 batch 太大单次推理慢。改成动态窗口。长序列拖累短序列一个超长请求混进 batch整个 batch 都慢。解决方法是按长度分桶调度。GC 停顿Python 的垃圾回收偶尔会卡一下。可以调 GC 阈值或者把关键路径用 C 扩展。实操心得我建议在调度器里加一个长请求单独处理的逻辑——超过某个长度的请求不参与攒批直接单独跑。这样能避免一条长请求拖垮整个 batch。5.4 一个反直觉的优化有时候慢一点反而更快这个经验挺反直觉的。我一开始追求极致的低延迟攒批窗口设得很小1ms结果吞吐上不去整体排队时间反而更长。后来把窗口调到 20ms单次延迟高了但吞吐翻倍整体 P99 延迟反而降了。这就是排队论的现实在接近饱和的系统里降低单次处理时间不一定能降低整体延迟因为瓶颈在排队。找到那个平衡点比一味优化单点更重要。6. 我在这套流水线上踩过的坑和真实体会搭这套东西花了我大概两个月中间踩的坑能写一本书。挑几个最有代表性的说说。第一个坑是过早优化。我一开始就想上 PagedAttention、想上连续批处理结果代码复杂度爆炸bug 一堆连正确性都没保证。后来退回去先写最笨的版本跑通了再优化效率反而高。正确性永远优先于性能这个顺序不能反。第二个坑是忽视数据管道。我花了大量时间调模型和推理结果线上问题 80% 出在数据管道。后来我把数据管道的测试覆盖率提到 90% 以上问题少了一大半。数据是地基地基不稳上面盖什么都是危房。第三个坑是监控缺失。早期我靠 print 调试线上出问题两眼一抹黑。后来补上 Prometheus 埋点很多问题在发生前就能从指标趋势上看出来。可观测性不是锦上添花是必需品。最后一个体会是关于从零这件事本身。我现在的看法是从零不是目的理解才是。你不需要真的手写每一个算子但你需要知道每一层在干什么、瓶颈在哪、怎么调。这套流水线搭完之后我再看那些推理框架的文档感觉完全不一样了——以前是照着抄现在是知道为什么这么配。如果你也在做类似的事我的建议是先跑通再优化最后再抽象。别一上来就想着设计一个完美的架构先让一个请求能正确地返回结果然后加批处理然后加缓存然后加监控。每一步都验证每一步都留退路。这套东西没有捷径但走一遍下来你对AI工程的理解会上一个台阶。
返回列表