ARTICLE DETAIL

资讯详情

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

从零手搓AI工程:不调包构建Transformer全链路实战

从零手搓AI工程:不调包构建Transformer全链路实战 1. 从零手搓AI工程为什么我不建议你直接调包1.1 一个让我决定“重造轮子”的契机去年帮一个做智能客服的朋友排查线上问题模型离线评测准确率92%上线后用户投诉率却居高不下。团队里三个算法工程师围着Transformer结构图吵了一下午有人说是注意力头数不够有人怀疑是位置编码在长文本上失效。最后定位到的问题让人哭笑不得——推理服务在拼接多轮对话历史时把用户最新一轮的query截断了模型看到的永远是上一轮的上下文。这件事让我意识到一个很普遍的现象现在太多人做AI项目是从pip install transformers开始的也是从model.generate()结束的。中间发生了什么、数据怎么流动、显存怎么分配、延迟卡在哪里一概不知。一旦出问题排查方向全靠猜。ai-engineering-from-scratch这个标题吸引我的地方就在这里。它不是教你调包而是让你从矩阵乘法开始一步步把一个大语言模型跑起来。这听起来像是自虐但实际做下来收获远超预期。这篇文章我会把整个从零构建AI工程的过程拆开揉碎包括为什么这么设计、每一步的坑在哪里、以及我实测下来哪些环节可以偷懒、哪些绝对不能省。适合谁看如果你已经会用PyTorch搭个分类网络但对Transformer的内部机制一知半解或者你天天调OpenAI接口想搞清楚token到底是怎么变成概率分布的再或者你准备面试AI工程岗位需要一套能讲清楚“从输入到输出”完整链路的项目经验——那这篇内容应该能帮到你。1.2 先想清楚从零构建到底“零”在哪里很多人对“from scratch”有误解以为是要从晶体管开始造计算机。不是的。在AI工程语境下从零构建的边界需要明确划定否则项目会无限膨胀。我的定义是不依赖任何预训练模型权重不调用高层推理框架用基础张量运算实现一个完整可训练的Transformer语言模型并配套数据处理、训练循环、推理服务和性能优化全链路。换句话说PyTorch的torch.nn.Linear可以用但torch.nn.Transformer不能用torch.optim.AdamW可以用但HuggingFace的Trainer不能用。为什么这么划边界因为nn.Linear和AdamW属于通用数学工具它们的正确性已经被无数项目验证过重新实现没有认知收益。而Transformer的注意力机制、位置编码、层归一化策略、训练时的梯度处理这些才是AI工程的核心知识密度所在。把精力花在这些地方投入产出比最高。还有一个现实考量完全不用框架意味着你要手写反向传播那项目周期会从两周变成两个月而且大概率写出来的东西数值不稳定。用PyTorch的自动微分但自己组装模块这是最平衡的方案。注意如果你是为了学习目的做这个项目建议不要一上来就追求性能。先让模型在小数据集上过拟合确认整个链路是通的再逐步加数据、加层数、加优化技巧。我见过太多人卡在“loss不下降”上最后发现是数据预处理时把token id映射错了。2. 核心模块拆解一个能跑的Transformer需要哪些零件2.1 分词器别小看这个“翻译官”模型不认识文字只认识数字。分词器的任务就是把“今天天气不错”变成[1234, 5678, 9012]这样的整数序列。听起来简单但这里埋着第一个大坑。最朴素的做法是按字切分中文每个字一个token英文每个字母一个token。问题是英文单词“understanding”会被切成11个token序列长度爆炸模型很难学到词级别的语义。更好的方案是Byte Pair EncodingBPE它通过统计语料中最高频的字符对逐步合并成子词。比如“understanding”可能被切成“under”、“stand”、“ing”三个子词既保留了词根信息又控制了词表大小。我实测下来对于中英混合语料词表大小设在32000左右比较合适。太小会导致很多词被切成单字序列过长太大则embedding矩阵参数量膨胀小数据集上容易过拟合。具体计算方式是假设你的语料有1000万字符去重后不同字符组合大约在5万到8万之间取一个能覆盖95%以上高频组合的最小值通常落在30000到35000区间。实现BPE的核心是一个循环每次统计相邻token对的频率合并频率最高的那一对直到达到目标词表大小。这里有个工程细节——统计频率时要用collections.Counter但直接对全部语料做会非常慢。我的做法是先采样一部分语料比如100万字符训练BPE合并规则然后用这些规则去编码全量数据。实测这样训练出的分词器在完整语料上的OOV率只比全量训练高0.3%左右但速度快了十几倍。# BPE训练核心逻辑示意 from collections import Counter def get_stats(ids): counts Counter() for pair in zip(ids[:-1], ids[1:]): counts[pair] 1 return counts def merge(ids, pair, new_id): new_ids [] i 0 while i len(ids): if i len(ids)-1 and ids[i] pair[0] and ids[i1] pair[1]: new_ids.append(new_id) i 2 else: new_ids.append(ids[i]) i 1 return new_ids实操心得训练BPE之前一定要做文本规范化。全角转半角、统一大小写、去除连续空白字符这些预处理能显著减少词表冗余。我试过不做规范化直接训练结果“Hello”和“hello”被分成了两个不同的token词表利用率很低。2.2 注意力机制Transformer的心脏注意力机制的本质是加权求和。给定一个查询向量Q和一组键值对(K, V)注意力输出就是V的加权和权重由Q和K的相似度决定。用生活化的例子解释你在图书馆找书Q是你脑子里的需求比如“想学AI工程”K是每本书的标题V是书的内容。你根据需求与标题的匹配程度决定从每本书里汲取多少内容。具体到代码实现缩放点积注意力是标准做法import torch import torch.nn as nn import math class ScaledDotProductAttention(nn.Module): def forward(self, Q, K, V, maskNone): d_k Q.size(-1) scores torch.matmul(Q, K.transpose(-2, -1)) / math.sqrt(d_k) if mask is not None: scores scores.masked_fill(mask 0, float(-inf)) attn torch.softmax(scores, dim-1) return torch.matmul(attn, V), attn这里有几个关键点值得展开。第一为什么要除以sqrt(d_k)因为当维度很大时Q和K的点积结果方差会变大softmax之后梯度会趋近于零训练不动。除以维度的平方根相当于做了一次标准化让点积结果的方差稳定在1附近。第二mask的作用是防止模型“偷看”未来token。在训练语言模型时预测第t个位置时只能看到前t-1个位置的信息所以需要一个上三角为0的掩码矩阵。多头注意力则是在这个基础上把Q、K、V分别投影到多个子空间并行做注意力计算最后拼接起来。为什么要多头单个注意力头只能学到一种关注模式比如“关注前一个词”。多头允许模型同时关注不同位置、不同语义关系比如一个头关注语法结构另一个头关注语义相似度。我实测下来8个头在中小规模模型上性价比最高再多的话收益递减明显。2.3 位置编码给模型装上“顺序感”注意力机制本身是位置无关的。把输入序列打乱注意力输出不变。但语言是有顺序的“猫追老鼠”和“老鼠追猫”意思完全相反。所以需要额外注入位置信息。原始Transformer用的是正弦位置编码用不同频率的三角函数生成位置向量。这种编码的好处是可以外推到训练时没见过的序列长度但实际效果在长文本上一般。后来大家更常用可学习的位置嵌入就是给每个位置分配一个可训练的向量简单直接在固定长度任务上表现更好。我选择的是旋转位置编码RoPE的思路它通过旋转矩阵把位置信息编码进Q和K的点积中。具体来说对于位置m的查询向量和位置n的键向量它们的点积会包含一个与(m-n)相关的旋转角度这样注意力分数天然就包含了相对位置信息。RoPE在长文本外推上表现更稳而且实现起来也不复杂。def apply_rope(x, freqs): # x: [batch, heads, seq_len, dim] # freqs: [seq_len, dim//2] x1, x2 x[..., ::2], x[..., 1::2] cos torch.cos(freqs).unsqueeze(0).unsqueeze(0) sin torch.sin(freqs).unsqueeze(0).unsqueeze(0) rotated torch.stack([x1*cos - x2*sin, x1*sin x2*cos], dim-1) return rotated.flatten(-2)注意位置编码的频率基数是个超参数。原始论文用10000但在长序列任务上调大到500000能让模型更好地区分远距离位置。我试过在2048长度上训练基数10000时模型对超过1500位置的信息几乎不关注调到100000后明显改善。3. 训练全流程从随机权重到能说人话3.1 数据管道别让IO成为瓶颈训练语言模型时数据加载速度经常被忽视。我第一版实现时每个batch都从磁盘读原始文本、现场分词、padding结果GPU利用率只有30%不到大部分时间在等CPU。正确的做法是预处理阶段就把所有文本转成token id序列存成二进制文件。训练时用numpy.memmap做内存映射按需读取既省内存又快。具体流程是遍历所有文本文件用训练好的分词器编码把结果拼接成一个大的uint16数组词表小于65536时够用存为.bin文件。同时记录每个文档的起始位置方便后续做文档级别的注意力掩码。import numpy as np # 预处理编码全量语料 all_ids [] for text in corpus: ids tokenizer.encode(text) all_ids.extend(ids) all_ids.append(eos_token_id) # 文档间用eos分隔 arr np.array(all_ids, dtypenp.uint16) arr.tofile(train.bin) # 训练时读取 data np.memmap(train.bin, dtypenp.uint16, moder) def get_batch(batch_size, block_size): ix torch.randint(len(data) - block_size, (batch_size,)) x torch.stack([torch.from_numpy(data[i:iblock_size].astype(np.int64)) for i in ix]) y torch.stack([torch.from_numpy(data[i1:i1block_size].astype(np.int64)) for i in ix]) return x, y这里有个细节y是x向右移一位。因为语言模型的训练目标是给定前t个token预测第t1个token。所以输入是[0, 1, 2, ..., n-1]标签是[1, 2, 3, ..., n]。这个错位操作在写代码时容易搞混建议写个单元测试验证一下。3.2 训练循环稳定比快更重要训练循环的骨架很标准前向传播、计算损失、反向传播、更新参数。但魔鬼在细节里。首先是学习率调度。我采用的是带预热的余弦退火前2000步线性从0升到峰值学习率然后按余弦曲线衰减到峰值的10%。预热的作用是防止训练初期梯度爆炸余弦退火则让模型在后期精细调整。峰值学习率我设的是3e-4配合AdamW优化器权重衰减0.1。这个组合在中小模型上比较稳。def get_lr(step, warmup_steps, max_steps, max_lr, min_lr): if step warmup_steps: return max_lr * (step 1) / warmup_steps if step max_steps: return min_lr decay_ratio (step - warmup_steps) / (max_steps - warmup_steps) coeff 0.5 * (1.0 math.cos(math.pi * decay_ratio)) return min_lr coeff * (max_lr - min_lr)其次是梯度裁剪。语言模型训练中偶尔会出现梯度尖峰不处理的话一次异常更新就可能毁掉之前的所有训练。我设的裁剪阈值是1.0用torch.nn.utils.clip_grad_norm_实现。实测下来这个操作几乎不增加计算开销但能显著提升训练稳定性。还有一个容易被忽略的点梯度累积。当显存不够跑大batch时可以跑多个小batch把梯度累加起来再更新。比如目标batch size是32但显存只够跑8那就跑4次前向反向每次除以4梯度累加后再step。这样等效于大batch训练但显存占用只有四分之一。实操心得训练初期loss会有波动不要一看到loss上升就调学习率。我的经验是观察100步的移动平均如果连续上升超过20%再干预。另外验证集loss比训练集loss更值得关注如果验证loss开始上升而训练loss还在降说明过拟合了该加dropout或者减层数。3.3 混合精度训练省显存还能提速FP16混合精度训练现在已经是标配了。核心思路是前向和反向用FP16计算但维护一份FP32的权重副本用于更新。这样显存占用大约减半计算速度提升30%到50%。PyTorch的torch.cuda.amp让这件事变得很简单scaler torch.cuda.amp.GradScaler() for x, y in dataloader: with torch.cuda.amp.autocast(): logits, loss model(x, y) scaler.scale(loss).backward() scaler.unscale_(optimizer) torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) scaler.step(optimizer) scaler.update() optimizer.zero_grad(set_to_noneTrue)GradScaler的作用是动态调整loss的缩放因子防止FP16下梯度下溢。unscale_之后再裁剪梯度保证裁剪的是真实梯度值。set_to_noneTrue比zero_grad()更省内存因为它直接把梯度置为None而不是填零。我实测下来在RTX 3090上训练一个1.2亿参数的模型FP32时batch size最大只能到16开混合精度后能到40训练速度从每秒3.2个step提升到5.1个step。效果还是很明显的。4. 推理服务让模型真正跑起来4.1 KV Cache自回归生成的加速器训练好的模型要用来生成文本最朴素的方式是每次预测一个token然后把新token拼到输入后面重新跑一遍完整的前向传播。这样做的问题是生成第n个token时前面n-1个token的注意力计算被重复了n次计算复杂度是O(n²)。KV Cache的思路是把之前所有位置的Key和Value缓存下来生成新token时只需要计算当前token的Query然后和缓存的K、V做注意力。这样每步的计算量从O(n²)降到O(n)生成速度提升非常明显。class KVCache: def __init__(self, max_batch, max_seq, n_heads, head_dim): self.k torch.zeros(max_batch, n_heads, max_seq, head_dim) self.v torch.zeros(max_batch, n_heads, max_seq, head_dim) self.pos 0 def update(self, k_new, v_new): seq_len k_new.size(2) self.k[:, :, self.pos:self.posseq_len] k_new self.v[:, :, self.pos:self.posseq_len] v_new self.pos seq_len return self.k[:, :, :self.pos], self.v[:, :, :self.pos]实测下来生成512个token不用KV Cache需要约8秒用了之后降到1.2秒左右。代价是额外的显存占用每个token每层需要存2×n_heads×head_dim个浮点数。对于1.2亿参数、12层的模型512长度大约多占300MB显存完全可以接受。4.2 采样策略让生成结果既合理又有趣模型输出的是每个token的logits直接取argmax会得到最确定的序列但往往很无聊而且容易陷入重复循环。实际使用中需要引入随机性。温度Temperature是最基础的调节手段。把logits除以温度系数TT1时分布更尖锐生成更确定T1时分布更平坦生成更多样。我一般设T0.8作为默认值需要创意写作时调到1.0需要精确回答时降到0.3。Top-k采样是只保留概率最高的k个token其余置零后重新归一化。这能避免采样到极低概率的离谱token。k50是个常用值。Top-p核采样则是动态选择阈值从概率最高的token开始累加直到累积概率超过p只保留这些token。p0.9意味着模型只在“前90%概率质量”的token中选择。相比Top-kTop-p能自适应分布形状分布尖锐时保留少量token分布平坦时保留更多。def sample(logits, temperature0.8, top_k50, top_p0.9): logits logits / temperature # Top-k过滤 if top_k 0: indices_to_remove logits torch.topk(logits, top_k)[0][..., -1, None] logits[indices_to_remove] float(-inf) # Top-p过滤 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)注意重复惩罚Repetition Penalty也很重要。对已经出现过的token在采样前把它们的logit除以一个大于1的系数比如1.2降低再次被选中的概率。但系数别设太大否则会影响正常的功能词重复比如“的”、“了”这些。4.3 批处理与流式输出实际服务中请求是并发来的。如果每个请求单独跑模型GPU利用率很低。需要做动态批处理把多个请求的输入padding到相同长度一起前向传播。但生成任务有个麻烦——每个请求的生成长度不同有的很快遇到结束符有的要生成到最大长度。我的做法是维护一个活跃请求队列每步只对还没结束的请求做前向。已经生成结束符的请求移出队列返回结果。新来的请求加入队列。这样GPU始终在处理有效计算。流式输出则是另一个体验层面的优化。用户不想等整个回复生成完才看到内容而是希望像打字机一样逐字显示。实现上就是把每个step生成的token立即通过SSEServer-Sent Events推送给前端。这里要注意编码问题中文token可能对应多个字节需要确保按UTF-8边界切分否则会出现乱码。5. 踩坑实录与性能调优5.1 那些让我熬夜的BugLoss变成NaN。这是最常见也最头疼的问题。原因通常有三个学习率太大、梯度爆炸、或者数据里有异常值。我的排查顺序是先把学习率降10倍看是否恢复如果不行检查梯度裁剪是否生效再不行打印每个batch的loss定位到具体是哪个数据导致的。有一次发现是某条数据里混入了二进制文件解码后产生了大量连续的特殊token导致embedding输出异常。生成结果无限重复。模型陷入“复读机”模式比如一直输出“好的好的好的好的”。这通常是因为训练数据里有大量重复模式或者采样温度太低。解决办法除了调高温度和加重复惩罚还可以在训练时对重复的n-gram做惩罚。我试过在loss里加一个重复检测项对连续重复的token对增加惩罚效果不错但实现复杂后来还是用采样策略解决更简单。显存泄漏。训练几个step后OOM但batch size明明没变。这通常是某个中间变量被意外保留在计算图里。检查方法是打印torch.cuda.memory_allocated()看是否随step线性增长。常见原因是把loss累加到了列表里没释放或者在验证时忘了torch.no_grad()。5.2 性能调优速查表问题现象可能原因排查方法解决方案GPU利用率低于50%数据加载瓶颈用torch.profiler看时间分布预处理数据为二进制用memmap读取训练loss震荡大学习率过高或batch太小打印每步loss和梯度范数降低学习率增大batch或梯度累积验证loss早升过拟合对比训练和验证loss曲线增加dropout减少参数量加数据生成速度慢无KV Cache计时单步生成耗时实现KV Cache批处理请求显存不足batch太大或模型太宽nvidia-smi看显存占用混合精度梯度检查点减小batch分词结果奇怪词表不匹配打印编码后的token重新训练分词器检查预处理5.3 几个让我事半功倍的工具Weights Biases训练可视化神器。loss曲线、学习率、梯度分布、甚至生成样本都能实时看。我习惯每100步记录一次训练过程中手机就能监控不用一直盯着终端。torch.profiler定位性能瓶颈。能精确到每个算子耗时一眼看出是卡在数据加载、前向计算还是反向传播。第一次用的时候发现我的模型有30%时间花在LayerNorm上后来把多个LayerNorm合并成一个批量操作提速明显。einops张量操作的可读性救星。rearrange(x, b h s d - b s (h d))比x.transpose(1,2).reshape(b,s,-1)直观太多而且不容易搞错维度顺序。强烈建议在多头注意力的实现里用上。实操心得不要过早优化。我第一版模型跑得慢花了两天做各种优化结果后来改了模型结构之前的优化全白费。正确的顺序是先让模型跑通、效果达标再做profiling定位真正的瓶颈最后针对性优化。大部分情况下数据加载和注意力计算是两个主要瓶颈优先解决这两个。6. 从项目到产品还差哪些工程化步骤6.1 模型量化与部署训练完的模型是FP32的部署时通常需要量化到INT8甚至INT4来降低显存和加速推理。量化分两种训练后量化PTQ和量化感知训练QAT。PTQ简单直接对权重做线性映射但精度损失可能较大。QAT在训练时模拟量化误差精度更好但需要重新训练。我实测下来对于1亿参数级别的模型INT8 PTQ的精度损失在可接受范围内困惑度上升不到5%推理速度提升约1.8倍。INT4则损失明显除非用GPTQ这类高级量化方法。部署时还要考虑批处理策略。在线服务通常要求低延迟batch size设小一点比如4到8离线批量生成可以设大batch32以上追求吞吐。我一般会准备两套配置根据流量动态切换。6.2 监控与迭代上线不是终点。需要监控的指标包括请求延迟的P50、P95、P99每秒生成token数显存占用以及生成质量指标比如重复率、平均长度。我遇到过线上流量突增导致请求排队P99延迟从200ms飙到3秒后来加了自动扩缩容才解决。迭代方面收集用户反馈和bad case定期做增量训练或微调。这里有个经验不要频繁全量重训成本太高。可以固定底层参数只微调顶层或者用LoRA这类低秩适配方法训练参数量减少90%以上效果损失很小。6.3 安全与合规生成内容需要过滤。我实现了一个简单的后处理管道先检查是否包含敏感词再用一个小的分类模型判断是否属于有害内容最后还可以加一层规则引擎处理特定场景。过滤的阈值要可配置不同业务场景要求不同。另外推理服务要做好限流和鉴权。每个API key分配配额防止滥用。日志要脱敏存储不能记录用户的原始输入和模型的完整输出只保留必要的统计信息。这个项目我从零开始写前后花了大约三周业余时间代码量在3000行左右。最大的收获不是模型效果有多好而是对整个链路的每个环节都有了肌肉记忆般的理解。现在再看到任何AI工程问题我都能在脑子里快速定位到可能出问题的模块。这种掌控感是调包永远给不了的。如果你也在做类似的事情我的建议是别怕慢别怕造轮子第一遍亲手写过的代码比读十篇论文都管用。
返回列表