ARTICLE DETAIL

资讯详情

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

从零手搓AI工程:不调包也能跑通推理链路

从零手搓AI工程:不调包也能跑通推理链路 1. 从零手搓AI工程为什么我不建议你直接调包很多人一上来就想跑通一个能对话的模型或者直接拉一个开源框架改改参数觉得这样就是“做AI工程”了。我刚开始也这么想直到有一次线上服务半夜崩了排查到凌晨三点才发现问题出在我根本说不清楚推理请求在GPU显存里到底经历了什么。那次之后我才明白调包能让你跑起来但从零手搓一遍才能让你在出问题的时候知道该看哪里。“ai-engineering-from-scratch”这个方向核心不是让你去造一个比肩大厂的基础模型那不现实。它的真正价值在于把AI工程链路里的每一个黑盒拆开用最小的依赖、最直白的代码把数据怎么进来、模型怎么算、结果怎么出去这条线完整走一遍。走完之后你再去看那些封装好的框架看到的就不再是魔法而是一层一层的工程取舍。这篇文章适合谁如果你已经会用Python知道什么是矩阵乘法但每次看到“推理优化”“显存管理”“批处理调度”这些词就心里发虚那这篇就是写给你的。我会按照一个从业者实际搭建最小可用AI工程链路的顺序把每个环节的“为什么这么设计”讲清楚同时把我在这个过程中踩过的坑和总结的技巧一并分享出来。全文不依赖任何特定厂商的云服务所有代码和思路都可以在你自己的机器上复现。2. 先想清楚从零构建到底要“零”到什么程度2.1 完全不用框架和用轻量框架的边界在哪里“从零”这个词很容易让人走极端。有人觉得必须用纯Python加NumPy手写所有算子连矩阵乘法都不能调库也有人觉得只要不调大厂API就算从零。我的经验是你需要明确自己的学习目标然后据此划定边界。如果你的目标是理解Transformer的数学本质那确实应该用NumPy把注意力机制、前馈网络、层归一化都手写一遍甚至反向传播也可以自己推导。但如果你目标是理解AI工程链路那重点就不在算子实现上而在数据管道、模型加载、推理调度、结果后处理这些工程环节。这时候用PyTorch作为张量计算后端是完全合理的因为PyTorch在这里的角色相当于一个“高级NumPy”你关注的是它怎么被集成进整个系统而不是它内部怎么实现卷积。我自己的做法是分两层核心算子层用NumPy手写一遍确保理解每个矩阵的形状变化和计算量工程链路层用PyTorch的Tensor作为数据载体但所有调度逻辑、内存管理、批处理策略都自己实现。这样既不会陷入“重复造轮子”的泥潭又能把工程环节的每个决策点都暴露出来。2.2 最小可行链路的四个必备模块一个能跑通的AI工程最小链路不管多简单都必须包含四个模块输入处理、模型推理、输出解析、资源管理。少任何一个这个链路就是残缺的你在后续扩展时一定会遇到问题。输入处理负责把原始数据文本、图像、音频转换成模型能吃的张量。这里最容易忽略的是分词器的边界情况比如文本里有特殊字符、超长截断、多语言混合这些在调包时框架帮你处理了自己写的时候如果没考虑模型输出就会莫名其妙。模型推理是核心计算环节。从零构建时我建议先用一个极小的模型比如两层Transformer隐藏维度64来验证链路而不是直接上大模型。小模型跑通之后再把同样的逻辑迁移到大模型上这样排查问题的成本会低很多。输出解析负责把模型输出的张量转换成人类可读的结果。这里的关键是理解模型输出的语义比如语言模型输出的是logits你需要做softmax、采样、解码分类模型输出的是类别概率你需要做阈值判断。这些步骤在框架里通常是一行代码但自己写的时候必须清楚每一步在做什么。资源管理是最容易被忽视但最重要的一环。显存分配、批处理大小、并发请求队列这些决策直接决定了你的服务能不能稳定运行。我在早期项目中就因为没做显存监控导致服务在并发量上来之后直接OOM排查了半天才发现是中间激活值没有及时释放。3. 数据管道从原始文本到模型输入的完整转换3.1 分词器的手写实现与边界处理分词器是数据管道的第一道关卡。很多人觉得分词就是查表没什么技术含量但实际工程中分词器的边界处理直接决定了模型输出的稳定性。我手写分词器时核心逻辑是三步规范化、切分、映射ID。规范化包括统一大小写、处理特殊符号、去除多余空白。切分可以按字符、按词、按子词从零构建时我建议先用字符级分词因为它的词表小、逻辑简单适合验证链路。映射ID就是查词表但这里有个关键细节未知字符的处理策略。如果你直接忽略未知字符模型输入就会变短位置编码就会错位。如果你用统一的UNK token替代又会丢失信息。我的做法是保留一个专门的UNK ID同时在输入处理阶段记录未知字符的位置和原始内容这样在输出解析时可以根据这些信息做后处理修正。这个技巧在调包时完全被隐藏了但自己写的时候必须面对。另一个容易踩坑的地方是截断策略。当输入超过模型最大长度时你是截头、截尾、还是从中间截不同策略对下游任务的影响完全不同。比如做文本分类时截尾可能丢失关键信息做对话生成时截头可能丢失上下文。我的经验是根据任务类型选择截断策略并且在截断时保留一个标记让模型知道这里发生了截断。3.2 批处理与填充效率与正确性的权衡批处理是提升推理吞吐量的关键手段但填充padding引入的无效计算和位置编码偏移是必须解决的问题。从零实现批处理时你需要决定是动态批处理还是静态批处理。静态批处理就是固定一个批次大小不足的用填充补齐。这种方式实现简单但填充比例高时浪费严重。动态批处理是根据当前请求的实际长度动态组批效率高但实现复杂。我建议先从静态批处理开始但要注意两个细节。第一填充token的ID必须和真实token区分开通常用0作为填充ID同时在注意力掩码里把填充位置屏蔽掉。第二位置编码要正确处理填充不能让填充位置占用有效的位置编号。我见过有人直接按物理位置编码结果模型把填充当成了真实输入输出完全乱套。下面是一个简化的批处理填充逻辑用Python列表来演示核心思路def pad_batch(sequences, pad_id0, max_lenNone): if max_len is None: max_len max(len(seq) for seq in sequences) padded [] masks [] for seq in sequences: pad_len max_len - len(seq) padded.append(seq [pad_id] * pad_len) masks.append([1] * len(seq) [0] * pad_len) return padded, masks这段代码看起来简单但掩码的生成顺序必须和填充顺序严格对应否则注意力机制就会关注到错误的位置。我在实际项目中就因为掩码生成时用了错误的索引导致模型在长文本上表现异常排查了很久才发现是掩码错位。3.3 数据预取与内存布局的工程考量当你的服务需要处理大量请求时数据管道的效率就成为瓶颈。预取prefetch和内存布局优化是两个关键手段。预取的核心思想是在GPU计算当前批次的同时CPU提前准备好下一批次的数据。这样GPU不会因为等待数据而空闲。从零实现时你可以用一个简单的双缓冲队列一个缓冲区供GPU读取另一个缓冲区由CPU填充两者交替使用。内存布局方面连续内存访问比随机访问快得多。在构建输入张量时尽量让同一批次的数据在内存中连续存放。如果你用Python列表存储再转换成张量中间会经历多次内存拷贝。更好的做法是预先分配一块连续内存然后按位置写入数据。这个优化在批次较大时效果非常明显我实测下来吞吐量能提升20%以上。4. 推理引擎不调框架怎么让模型跑起来4.1 模型加载与参数初始化的手写流程从零构建推理引擎的第一步是加载模型参数。如果你用的是自己训练的模型参数可能保存在一个二进制文件里如果你是从开源模型转换过来的参数格式可能是safetensors或PyTorch的state_dict。不管哪种格式核心逻辑都是读取二进制数据按照预定义的形状重塑成张量然后赋值给模型对应的层。这里的关键是参数名称的映射关系。开源模型的参数命名通常有固定规范比如transformer.h.0.attn.c_attn.weight你需要把这个名称映射到你自己的模型类属性上。我踩过的一个坑是参数转置。有些框架保存的权重矩阵是转置过的如果你直接加载而不做转置计算结果就会完全错误。排查这种问题非常痛苦因为模型能跑通但输出质量很差你很难判断是模型本身的问题还是加载的问题。我的经验是加载完参数后先用一个已知输入的简单用例验证输出比如用全零输入看输出是否合理或者用单位矩阵测试线性层。4.2 前向传播的逐层拆解与调试技巧前向传播是推理引擎的核心。从零实现时我建议逐层实现并逐层验证而不是一次性写完整个模型再调试。具体做法是先实现嵌入层用一个小词表验证输出形状和数值范围再实现注意力层用随机初始化的参数验证注意力权重的归一化然后实现前馈网络验证激活函数的非线性最后实现层归一化和残差连接验证梯度流动虽然推理不需要梯度但残差连接的正确性影响输出。每一层实现完之后打印中间张量的形状和统计量均值、方差、最大值、最小值。这些统计量能帮你快速定位问题如果某一层的输出方差突然变得极大或极小说明这一层的初始化或计算有问题。我在实现层归一化时就因为忘记减去均值导致输出全部偏移通过统计量一眼就看出来了。另一个实用技巧是用PyTorch的对应层作为参考实现。你自己写的层和PyTorch的层用同样的输入比较输出差异。如果差异在数值精度范围内比如1e-5说明实现正确如果差异很大就逐行对比计算逻辑。这个方法能帮你快速定位到具体的计算错误。4.3 显存管理与中间激活值的释放时机显存管理是从零构建推理引擎时最容易出问题的地方。中间激活值如果不及时释放显存占用会随着批次大小线性增长很快就会OOM。在PyTorch中默认情况下计算图会保留中间激活值用于反向传播。但推理时不需要反向传播所以应该用torch.no_grad()上下文管理器来禁用梯度计算这样PyTorch就不会保留中间激活值。如果你是自己用NumPy实现的前向传播那就需要手动管理每一层的输出在传给下一层之后如果不再需要就应该解除引用。我自己的做法是用一个显式的列表来跟踪所有中间张量在每一层计算完成后把不再需要的张量从列表中移除并调用垃圾回收。虽然Python的垃圾回收机制会自动处理但显式释放能让显存占用更早下降对于大模型推理来说这个时间差可能就决定了能不能多跑一个批次。还有一个细节是KV Cache的管理。在自回归生成中每次生成一个新token都需要用到之前所有token的Key和Value。如果每次都重新计算计算量会随序列长度平方增长。KV Cache就是把这些中间结果缓存起来避免重复计算。从零实现时你需要决定缓存的存储格式和更新策略。我的经验是用一个预分配的张量来存储缓存每次生成新token时追加写入而不是每次都创建新张量。这样能避免频繁的内存分配和拷贝。5. 输出解析与后处理让模型结果真正可用5.1 解码策略贪心、采样与束搜索的取舍模型输出的原始logits只是一堆数值怎么从这些数值里选出最终的token直接决定了生成结果的质量和多样性。贪心解码就是每次选概率最大的token。这种方式确定性强、速度快但容易生成重复、单调的内容。采样解码是按照概率分布随机选择多样性好但可能生成不连贯的内容。束搜索是保留多个候选序列最后选整体概率最高的质量通常最好但计算量最大。从零实现时我建议先实现贪心解码验证链路再实现采样解码增加多样性最后根据需要实现束搜索。采样解码里有一个关键参数是温度temperature它控制概率分布的平滑程度。温度趋近于0时退化为贪心温度大于1时分布更平坦、多样性更高。我的经验是温度设在0.7到0.9之间通常能平衡质量和多样性具体值需要根据任务类型调整。还有一个容易被忽略的参数是重复惩罚repetition penalty。当模型陷入重复循环时对已经出现过的token降低其概率能有效缓解这个问题。实现方式是在计算概率之前对已生成token的logits减去一个惩罚值。这个惩罚值不能太大否则会导致模型完全避免使用某些必要的token。5.2 流式输出的缓冲与刷新策略在实际服务中用户通常希望看到流式输出也就是模型生成一个token就立即返回一个而不是等全部生成完再返回。这需要输出缓冲和刷新策略的配合。从零实现流式输出时核心问题是什么时候刷新缓冲区。如果每生成一个token就刷新网络开销会很大如果攒一批再刷新用户会感觉卡顿。我的做法是设置一个时间阈值和一个数量阈值谁先达到就触发刷新。比如每50毫秒或者每积累5个token就刷新一次。这样既能保证流畅感又不会过于频繁地发送网络请求。另一个细节是UTF-8字符的边界问题。一个中文字符可能由多个字节组成如果按字节流式输出可能会在字符中间截断导致乱码。解决方案是在缓冲区里保留不完整的字节序列等下一个token到来后拼接完整再输出。这个坑我在早期项目中踩过用户看到一半中文乱码体验非常差。5.3 结果校验与异常兜底模型输出不一定总是合理的。结果校验和异常兜底是保证服务稳定性的最后一道防线。校验的内容包括输出长度是否超过限制、是否包含非法字符、是否为空、是否陷入重复循环。对于每种异常情况你需要有对应的兜底策略。比如输出为空时可以返回一个默认回复检测到重复循环时可以强制截断并返回已生成的内容。我的经验是把校验逻辑做成可配置的规则链每条规则独立判断按优先级依次执行。这样在调试时可以单独启用或禁用某条规则快速定位问题。另外记录异常发生的频率和类型这些数据能帮你发现模型的系统性问题比如某个类型的输入总是导致异常输出那就需要针对性地优化输入处理或调整解码参数。6. 性能调优从能跑到跑得快的几个关键决策6.1 算子融合与内存复用的实操边界当你的推理引擎能跑通之后下一步就是让它跑得更快。算子融合和内存复用是两个最直接的手段。算子融合就是把多个连续的小算子合并成一个大的算子减少内核启动开销和中间结果的读写。比如把矩阵乘法、偏置加法和激活函数融合成一个算子。从零实现时你可以手动把这几步写在一个函数里避免中间张量的创建。虽然Python层面的函数调用开销还在但至少减少了内存分配和拷贝。内存复用是指重复使用同一块内存来存储不同阶段的中间结果。比如注意力计算完成后它的中间张量内存可以立即被前馈网络复用。实现方式是用一个内存池来管理张量的分配和释放而不是每次都向系统申请新内存。这个优化在批次较大时效果显著我实测下来显存占用能降低30%左右。但要注意算子融合和内存复用会增加代码的复杂度降低可读性。我的建议是先用清晰但低效的方式实现确保正确性然后再逐步优化。每次优化后都要用相同的输入验证输出是否一致避免优化引入bug。6.2 批处理大小与延迟的平衡点怎么找批处理大小是影响吞吐量和延迟的核心参数。批次越大吞吐量越高但单个请求的延迟也越大因为要等批次里所有请求都处理完才能返回。找平衡点的方法是测量不同批次大小下的吞吐量和延迟曲线。通常吞吐量会随着批次增大而提升但提升幅度会逐渐减小延迟则会线性增长。你需要根据业务需求选择一个合适的点如果业务对延迟敏感就选小批次如果对吞吐量敏感就选大批次。我的经验是从批次大小1开始逐步翻倍测量每个点的吞吐量和P99延迟。当吞吐量的提升小于20%而延迟增长超过50%时就说明已经过了最佳点。另外动态批处理可以在请求量低时使用小批次降低延迟在请求量高时使用大批次提升吞吐量是更灵活的方案。6.3 量化与剪枝在自建引擎中的落地难度量化和剪枝是模型压缩的常用手段但在自建引擎中落地有一定难度。量化是把模型的浮点参数转换成低精度表示比如int8减少内存占用和计算量。从零实现时你需要自己实现量化算子和反量化算子并确保量化误差在可接受范围内。难点在于哪些层可以量化、量化后的精度损失如何补偿。我的经验是先从权重矩阵量化开始激活值保持浮点这样实现简单且精度损失较小。剪枝是去掉模型中不重要的连接或神经元。从零实现时你需要定义重要性度量比如权重绝对值、实现剪枝掩码、并在前向传播中应用掩码。难点在于剪枝后的模型通常需要微调才能恢复精度而微调又需要反向传播和优化器这超出了纯推理引擎的范围。所以剪枝在自建推理引擎中的落地通常需要配合训练框架一起使用。7. 我踩过的三个典型坑与排查思路7.1 位置编码错位导致的输出乱码问题现象模型能跑通但生成的文本完全不通顺像是随机字符拼接。排查过程我先检查了输入分词确认token ID正确然后检查了模型参数加载确认权重没有转置错误最后打印了每一层的输出统计量发现第一层注意力输出的方差异常大。进一步检查位置编码发现我在生成位置编码时把批次维度和序列维度搞反了导致每个位置编码都错位。修复方案重新检查位置编码的形状确保它是[序列长度, 隐藏维度]并且在加到嵌入张量时正确广播。修复后输出立刻恢复正常。经验总结位置编码是Transformer里最容易出错的地方之一因为它的形状和广播规则比较微妙。建议在实现位置编码后先用一个简单的用例验证输入两个相同token看它们的输出是否因为位置不同而不同。如果输出相同说明位置编码没起作用。7.2 显存泄漏中间张量未释放的连锁反应问题现象服务运行一段时间后OOM重启后恢复但过一段时间又OOM。排查过程我用显存监控工具观察显存占用曲线发现它随着请求数量增加而持续上升从不下降。这说明有中间张量没有被释放。我检查了前向传播的代码发现在计算注意力时我创建了一个中间张量用于存储注意力权重但在后续计算中没有解除引用。虽然Python的垃圾回收最终会处理它但由于计算图的存在这个张量一直被引用着。修复方案在注意力计算完成后显式地将中间张量设为None并调用torch.cuda.empty_cache()如果用的是PyTorch。同时确保整个推理过程在torch.no_grad()上下文中执行。经验总结显存泄漏的排查比较困难因为现象是延迟出现的。建议在开发阶段就加入显存监控每处理一批请求就打印一次显存占用。如果发现占用持续上升就重点检查中间张量的生命周期。7.3 批处理填充引发的注意力异常问题现象单条请求推理正常但批处理时输出质量明显下降。排查过程我对比了单条和批处理的输入发现批处理时填充了多个填充token。检查注意力掩码发现掩码的生成逻辑有误我把掩码的0和1搞反了导致模型关注了填充位置而忽略了真实token。修复方案修正掩码生成逻辑确保真实token位置为1填充位置为0。同时在注意力计算中对掩码为0的位置加上一个极大的负值使其softmax后趋近于0。经验总结批处理相关的bug通常比较隐蔽因为单条测试时不会触发。建议在开发批处理功能时先用两条长度不同的输入做测试对比批处理和单条处理的输出是否一致。如果不一致优先检查掩码和位置编码。8. 从手搓到上生产的距离还有多远手搓一遍AI工程链路最大的收获不是代码本身而是对每个环节的边界和代价有了清晰的认知。你知道分词器在什么情况下会出错知道显存什么时候会爆知道批处理大小对延迟的影响曲线。这些认知让你在使用成熟框架时能更快地定位问题也能更准确地做技术选型。但手搓版本和 production-ready 之间还有不小的距离。手搓版本通常缺少错误恢复机制、缺少监控和告警、缺少灰度发布能力、缺少自动扩缩容。这些工程能力需要在实际业务中逐步补齐。我的建议是把手搓版本作为学习和验证工具把成熟框架作为生产工具两者结合使用。当你对链路足够熟悉之后再根据业务需求决定哪些环节可以自己实现、哪些环节应该用现成方案。最后分享一个我自己的习惯每次遇到一个不熟悉的AI工程问题我都会先用最小化的手搓代码复现它理解它的本质然后再回到生产代码里修复它。这个习惯让我避免了很多“改了这里崩了那里”的困境因为我知道每个改动会影响链路的哪个环节。手搓的意义不在于替代框架而在于让你成为那个能看懂框架在做什么的人。
返回列表