ARTICLE DETAIL

资讯详情

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

从零搭建AI工程能力:手写推理循环与批处理优化实战

从零搭建AI工程能力:手写推理循环与批处理优化实战 1. 从零搭建AI工程能力为什么我劝你别一上来就调包这两年AI应用开发的门槛肉眼可见地降低了随便拉个框架、调个API就能跑出一个能对话的Demo。但我带过不少新人也面试过不少号称“做过AI项目”的候选人发现一个很普遍的问题大家会用工具但不知道工具背后发生了什么。模型输出不稳定不知道从哪查推理速度慢不知道怎么优化显存爆了只会重启。这种状态下一旦遇到稍微偏离教程的场景整个人就卡住了。ai-engineering-from-scratch这个方向说白了就是解决这个问题的。它不是让你去从头训练一个GPT而是让你亲手把AI工程链路里的关键环节实现一遍——数据怎么处理、模型怎么加载、推理怎么调度、服务怎么暴露、性能怎么观测。适合谁看如果你已经会用Python调过几个模型API但总觉得心里没底想搞清楚“从输入到输出中间到底发生了什么”那这篇内容就是写给你的。我会按实际动手的顺序把每个环节的核心逻辑、容易踩的坑、以及我自己的经验教训都摊开讲。2. 整体设计思路为什么选择“从零实现”而不是“直接调包”2.1 调包侠的天花板在哪里先说说为什么我不推荐一上来就用高层框架。以推理部署为例你用某个流行的推理框架三行代码就能起一个服务。看起来很美好但问题在于当请求延迟从200ms涨到2s的时候你根本不知道是模型加载的问题、显存分配的问题、还是请求排队的问题。高层框架把太多细节封装掉了封装得越好出问题时你的排查手段就越少。我自己的经历是早期做一个文本分类服务用的某框架默认配置单条推理很快但并发一上来就崩。查了两天才发现是默认的批处理策略和显存预分配不匹配。如果我从零写过一遍推理循环知道batch_size、max_seq_len、padding策略这些参数怎么相互作用可能半小时就能定位。2.2 从零实现的核心链路拆解ai-engineering-from-scratch的合理路径我把它拆成五个层次数据层原始文本/图像怎么变成模型能吃的张量包括分词、归一化、批处理。模型层模型文件怎么加载、权重怎么映射、计算图怎么构建。推理层前向传播怎么执行、显存怎么管理、批处理怎么调度。服务层怎么把推理能力暴露成HTTP接口怎么处理并发和超时。观测层延迟、吞吐、显存占用、错误率怎么采集和展示。这五层每一层都可以先用最朴素的方式实现一遍然后再逐步替换成生产级方案。这样做的好处是你对每个环节的边界和代价都有体感后面用任何框架都能快速判断它帮你做了什么、没做什么。2.3 技术选型的取舍逻辑具体到工具选择我的建议是能用标准库就不用第三方能用轻量库就不用重框架。比如HTTP服务初期直接用Python内置的http.server就够了没必要上FastAPI。倒不是FastAPI不好而是你亲手处理过请求解析、并发模型、超时控制之后再用FastAPI会知道它到底帮你省了哪些事。模型加载这块transformers库是绕不开的但你可以先只用它的模型定义和权重加载推理循环自己写。这样你能看清楚input_ids、attention_mask、position_ids这些张量到底怎么在模型里流动。等你把单条推理跑通了再引入批处理、再引入KV Cache、再引入量化每一步都有明确的性能收益而不是一开始就堆一堆优化手段却不知道各自贡献了多少。3. 核心细节解析从零实现AI工程的关键环节3.1 数据预处理别小看分词和批处理很多人觉得数据预处理就是调个tokenizer没什么技术含量。但实际项目中预处理往往是bug最多的地方。我列几个我踩过的坑分词器的padding方向。不同模型对padding token的位置要求不一样有的要求pad在右边有的要求pad在左边。如果你用错了模型输出会莫名其妙地变差而且很难查。我建议在预处理阶段就把padding_side显式设置好并且在日志里打印一条样本的input_ids和attention_mask肉眼确认一下。批处理时的长度对齐。同一个batch里的样本长度不一致时需要padding到同一长度。但padding到多长如果按batch内最大长度那每个batch的计算量都不一样显存占用会波动。如果按全局最大长度那短样本会浪费大量计算。我的做法是先按长度分桶同一个桶内的样本长度接近padding浪费小桶的大小根据显存上限动态调整。特殊token的处理。比如分类任务你到底取[CLS]位置的输出还是做mean pooling这取决于模型预训练时的目标。如果你不确定最稳妥的方式是看模型卡里的说明或者直接做个小实验对比两种方式在验证集上的表现。3.2 模型加载权重映射与计算图构建从零加载模型核心是搞清楚三件事权重文件里有哪些张量、这些张量对应模型的哪些层、怎么把它们塞进模型定义里。以Transformer类模型为例权重文件通常是一个字典key是层名value是张量。你需要把模型定义里的每个nn.Linear、nn.LayerNorm和权重文件里的key对应起来。这个过程容易出错的地方是命名映射有的权重文件用encoder.layer.0.attention.self.query.weight有的用layers.0.attention.q_proj.weight。如果你手动映射一定要写个脚本逐层核对shape是否匹配。计算图构建这块PyTorch的动态图机制其实帮你省了很多事。但你要注意torch.no_grad()的使用推理阶段一定要加上否则会保存计算图显存占用会暴涨。我见过有人推理时忘了加结果batch size只能开到1还以为是模型太大。3.3 推理循环批处理与KV Cache的取舍单条推理的循环很简单输入张量进模型输出张量出模型。但要做成服务就必须考虑批处理。批处理的核心矛盾是延迟和吞吐的权衡。攒批会增大延迟但不攒批吞吐上不去。我的经验是如果QPS要求不高比如小于10可以不做批处理每条请求单独推理延迟最低。如果QPS要求高就需要引入一个请求队列攒到一定数量或等待一定时间后统一推理。KV Cache是另一个关键优化。自回归生成时每次生成一个新token都要重新计算之前所有token的Key和Value这是巨大的浪费。KV Cache把这些中间结果缓存下来每次只计算新token的。但KV Cache会占用显存而且和batch size、序列长度成正比。你需要根据显存上限反推最大支持的batch size和序列长度。我实测下来在7B模型上开启KV Cache后生成速度能提升3到5倍。但如果你显存不够可能需要用PagedAttention之类的技术来管理KV Cache的内存碎片。3.4 服务暴露并发模型与超时控制把推理能力暴露成HTTP接口最朴素的方式是用http.server。但它是单线程的一个请求处理不完后面的请求就排队。你需要引入线程池或异步IO。线程池的问题是Python的GIL多个线程同时执行Python字节码时其实还是串行的。但推理的大部分时间花在PyTorch的C后端上这部分会释放GIL所以多线程推理是有意义的。我的做法是用一个固定大小的线程池每个线程持有一个模型副本或共享一个模型但加锁。共享模型加锁的方式显存占用低但并发度受锁限制多副本的方式并发度高但显存占用成倍增加。超时控制也很重要。如果某个请求的输入特别长推理时间可能超过客户端等待上限。你需要在服务端设置超时超时后直接返回错误而不是让请求一直挂着。我一般会把超时设置为P99延迟的1.5倍左右具体数值根据业务容忍度调整。4. 实操过程从零搭建一个最小可用的推理服务4.1 环境准备与依赖安装先明确一点我不建议在Windows上做这件事各种底层库的兼容性问题会让你怀疑人生。用Linux或者macOSPython版本选3.10或3.11比较稳定。依赖方面核心就是PyTorch和transformers。安装PyTorch时注意CUDA版本要和驱动匹配。如果你没有GPUCPU也能跑只是慢一些用于学习完全够用。pip install torch transformers如果你要用GPU去PyTorch官网查对应的安装命令别直接pip install torch那样装的是CPU版本。4.2 数据预处理模块的实现我写了一个最简单的预处理类核心就是分词和paddingfrom transformers import AutoTokenizer class Preprocessor: def __init__(self, model_name, max_length512): self.tokenizer AutoTokenizer.from_pretrained(model_name) self.tokenizer.padding_side right self.max_length max_length def process(self, texts): encoded self.tokenizer( texts, paddingTrue, truncationTrue, max_lengthself.max_length, return_tensorspt ) return encoded这段代码看起来简单但有几个细节padding_side显式设置为right因为大部分生成模型要求这样truncationTrue防止超长输入把显存撑爆return_tensorspt直接返回PyTorch张量省去手动转换。4.3 模型加载与推理循环模型加载我直接用AutoModelForCausalLM但推理循环自己写import torch from transformers import AutoModelForCausalLM class InferenceEngine: def __init__(self, model_name, devicecuda): self.model AutoModelForCausalLM.from_pretrained( model_name, torch_dtypetorch.float16, device_mapdevice ) self.model.eval() torch.no_grad() def generate(self, input_ids, attention_mask, max_new_tokens128): outputs self.model.generate( input_idsinput_ids, attention_maskattention_mask, max_new_tokensmax_new_tokens, do_sampleFalse, use_cacheTrue ) return outputs这里我用了torch.float16显存占用直接减半速度也有提升。use_cacheTrue就是开启KV Cache。do_sampleFalse是贪心解码输出稳定适合做服务。4.4 HTTP服务与并发处理服务层我用http.server加线程池from http.server import HTTPServer, BaseHTTPRequestHandler from concurrent.futures import ThreadPoolExecutor import json class Handler(BaseHTTPRequestHandler): def do_POST(self): content_length int(self.headers[Content-Length]) body self.rfile.read(content_length) data json.loads(body) texts data[texts] encoded preprocessor.process(texts) outputs engine.generate(**encoded) result preprocessor.tokenizer.batch_decode(outputs, skip_special_tokensTrue) self.send_response(200) self.send_header(Content-Type, application/json) self.end_headers() self.wfile.write(json.dumps({results: result}).encode()) server HTTPServer((0.0.0.0, 8000), Handler) server.serve_forever()这个服务是单线程的要支持并发需要换成ThreadingHTTPServer或者自己加线程池。我建议先用单线程跑通确认推理逻辑没问题再改并发。4.5 性能观测与日志埋点观测这块最少要记录三个指标请求延迟、批处理大小、显存占用。延迟用time.time()打点就行显存用torch.cuda.memory_allocated()。把这些指标定期打印到日志里你就能看出瓶颈在哪。我一般会在推理前后各打一个时间戳算出差值作为延迟。如果延迟波动很大说明批处理策略有问题或者显存不够导致频繁换页。5. 常见问题与排查技巧实录5.1 显存溢出OOM的排查路径OOM是最常见的问题。排查顺序我一般是这样的确认模型本身占多少显存。加载完模型后打印torch.cuda.memory_allocated()这是基线。确认单条推理的峰值显存。跑一条样本打印推理前后的显存差值。计算最大batch size。用显存上限 - 模型基线 - 系统预留/ 单条推理峰值得到理论最大batch size。留20%余量。实际batch size不要超过理论值的80%因为显存碎片和临时张量会额外占用。如果还是OOM考虑用torch.float16或torch.int8量化或者用梯度检查点虽然推理阶段一般用不到。5.2 推理速度慢的优化手段速度慢的原因可能有很多我按优先级列一下问题现象可能原因优化手段单条推理就慢模型太大或精度太高换小模型或量化批处理没提速批处理策略不对检查padding是否浪费太多生成阶段慢没用KV Cache开启use_cacheTrue并发时慢线程竞争或显存不足减少并发或增加显存我实测下来量化对速度的提升最明显int8量化能让7B模型在单张消费级显卡上跑起来速度也能接受。5.3 输出不稳定的排查思路输出不稳定通常和随机性有关。检查do_sample是否设为了Falsetemperature是否设为了0。如果这些都没问题那可能是模型本身的问题比如权重加载错了或者分词器版本不匹配。我遇到过一次输出乱码查了半天发现是分词器的vocab文件和模型权重不匹配。解决办法是确保AutoTokenizer和AutoModelForCausalLM用的是同一个模型名称。5.4 服务超时与请求堆积请求堆积通常是推理速度跟不上请求速度。短期方案是加超时超时的请求直接返回错误长期方案是加机器或者优化推理速度。我一般会在服务层加一个队列长度限制队列满了直接拒绝新请求避免雪崩。6. 从零实现之后下一步可以怎么扩展把最小可用版本跑通之后你可以按自己的需求逐步替换组件。比如把http.server换成FastAPI把手动批处理换成vLLM或TGI把float16换成int8量化。但每一步替换之前先想清楚你为什么要换是为了更高的吞吐、更低的延迟、还是更好的可维护性我自己的路径是先用从零版本跑通业务逻辑确认模型效果没问题然后逐步引入优化每引入一个优化就压测一次确认收益符合预期。这样你不会一下子引入太多变量出问题时也能快速定位。最后分享一个小技巧在从零实现阶段把每个模块的输入输出都打印出来哪怕只是前几条样本。这样一旦中间某个环节出错你能立刻看到是哪一步的输入输出对不上。这个习惯帮我省了无数调试时间。
返回列表