ARTICLE DETAIL

资讯详情

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

从零手搓AI工程:PyTorch模型训练与推理优化实战

从零手搓AI工程:PyTorch模型训练与推理优化实战 1. 从零手搓AI工程为什么我不建议你直接调包很多人一听到“AI工程”这四个字第一反应就是打开某个云平台拖几个组件调几个API然后跑通一个Demo就觉得自己已经掌握了。我刚开始也是这么想的直到有一次线上推理服务在凌晨两点崩了日志里全是显存溢出的报错而我对着那堆封装好的接口完全不知道从哪下手排查。那一刻我才意识到只会调包的人永远只能停留在“能用”的层面一旦出了问题连问题出在哪个环节都说不清楚。ai-engineering-from-scratch这个标题核心不是让你去重新发明Transformer而是让你把AI工程链路里的每一个关键环节都亲手实现一遍。从数据加载、分词、模型结构搭建、训练循环、梯度累积、混合精度到推理优化、批处理调度、显存管理这些东西如果你只是调用现成的库你永远不会理解它们为什么存在也不会知道在什么场景下该动哪个旋钮。这篇文章适合两类人一类是刚入门AI方向的学生或者转行者想真正搞懂一个模型从数据到部署到底经历了什么另一类是有一定调包经验但遇到瓶颈的工程师想往下钻一层搞清楚框架底层到底在干什么。我会按照一个最小可用的语言模型训练与推理链路把每个环节的“为什么”和“怎么做”都拆开讲清楚代码以PyTorch为主但思路是通用的。需要提前说明的是我不会给你一个“复制粘贴就能跑”的完整项目因为那样对你没有帮助。我会给你每个模块的核心逻辑、关键参数的计算方式、以及我在实际踩坑中总结出来的经验。你跟着走一遍收获会比直接clone一个仓库大得多。2. 数据管道的搭建别让IO成为你的第一个瓶颈2.1 为什么数据加载值得单独拿出来讲大部分教程在讲数据加载的时候就是一句DataLoader(dataset, batch_size32, shuffleTrue)带过。但实际做AI工程的时候数据管道往往是第一个让你崩溃的地方。我见过太多人模型代码写得漂漂亮亮结果训练速度慢得离谱最后发现是数据加载成了瓶颈GPU利用率常年徘徊在20%以下。一个合格的数据管道需要解决三个问题读取效率、内存占用、批处理策略。读取效率决定了你的GPU会不会饿着内存占用决定了你能不能处理大规模数据集批处理策略则直接影响模型收敛的质量。2.2 从原始文本到Token序列的完整链路假设你手里有一堆纯文本文件第一步是构建词表。这里我不建议你直接用现成的分词器而是先手写一个简单的字符级或者词级分词器理解分词的本质。# 一个极简的词级分词器实现 from collections import Counter class SimpleTokenizer: def __init__(self, vocab_size10000): self.vocab_size vocab_size self.word2idx {} self.idx2word {} def build_vocab(self, texts): counter Counter() for text in texts: counter.update(text.split()) # 保留最高频的vocab_size个词其余用unk代替 most_common counter.most_common(self.vocab_size - 2) self.word2idx {pad: 0, unk: 1} for idx, (word, _) in enumerate(most_common, start2): self.word2idx[word] idx self.idx2word {v: k for k, v in self.word2idx.items()} def encode(self, text): return [self.word2idx.get(w, 1) for w in text.split()]这段代码很短但里面有几个关键决策点值得展开。第一为什么保留pad和unk两个特殊tokenpad用于批处理时对齐序列长度unk用于处理词表外的词。第二为什么按频率截断而不是全量保留因为词表越大嵌入层的参数量越大低频词带来的收益远小于它占用的显存和计算量。实际工程中你会遇到文本长度差异极大的情况。有的样本只有十几个token有的有几千个。如果直接按最大长度padding显存浪费会非常严重。我的做法是采用动态padding也就是每个batch内按当前batch的最大长度来padding而不是全局最大长度。def collate_fn(batch, pad_idx0): # batch是(list of list of int) max_len max(len(seq) for seq in batch) padded [seq [pad_idx] * (max_len - len(seq)) for seq in batch] return torch.tensor(padded, dtypetorch.long)这个collate_fn看起来简单但它能把显存利用率提升30%以上尤其是在长尾分布明显的数据集上。我实测过一个文本分类任务全局padding和动态padding的显存占用差了将近一倍。2.3 数据预取与多进程加载的坑PyTorch的DataLoader提供了num_workers参数来做多进程加载但这里有几个坑我必须提醒你。第一num_workers不是越大越好一般设置为CPU核心数的2到4倍就够了设太大反而会因为进程切换开销导致性能下降。第二在Windows上使用多进程加载时必须把训练代码放在if __name__ __main__:下面否则会无限递归创建进程。第三如果你用了自定义的Dataset类确保它里面的操作是线程安全的尤其是涉及到文件读写的时候。还有一个容易被忽略的点是数据预取。DataLoader的prefetch_factor参数控制每个worker预取多少个batch默认是2。如果你的数据加载逻辑比较重比如需要实时做数据增强可以适当调大这个值让GPU不会因为等数据而空转。提示判断数据加载是否成为瓶颈最简单的方法是看GPU利用率。如果GPU利用率波动很大经常掉到50%以下那大概率是数据管道的问题。3. 模型结构从Embedding到Attention的手动实现3.1 为什么手写一遍Transformer是有必要的你可能会说nn.Transformer已经封装好了为什么还要手写我的回答是因为封装好的东西你调不了。当你需要修改注意力机制的计算方式、需要自定义位置编码、需要在特定层插入额外的模块时如果你不理解底层的矩阵运算你连改哪里都不知道。手写一遍还有一个好处就是你能真正理解参数量是怎么算出来的。比如一个d_model512, nhead8的多头注意力层它的参数量到底是多少Q、K、V三个投影矩阵各是512*512输出投影又是512*512加起来是4*512*512再加上偏置项。这些数字只有你自己算过一遍才能在模型变大时快速估算显存需求。3.2 缩放点积注意力的实现细节import torch import torch.nn as nn import math class ScaledDotProductAttention(nn.Module): def __init__(self, d_k): super().__init__() self.d_k d_k def forward(self, q, k, v, maskNone): # q: (batch, nhead, seq_len, d_k) scores torch.matmul(q, k.transpose(-2, -1)) / math.sqrt(self.d_k) if mask is not None: scores scores.masked_fill(mask 0, float(-inf)) attn torch.softmax(scores, dim-1) output torch.matmul(attn, v) return output, attn这段代码里最关键的是math.sqrt(self.d_k)这个缩放因子。为什么是除以sqrt(d_k)而不是别的因为当d_k比较大的时候点积的结果会变得很大经过softmax之后会变得非常尖锐梯度会趋近于零导致训练不动。除以sqrt(d_k)可以把方差拉回到1附近让softmax的输出分布更平滑。另一个细节是mask的处理。在解码器中我们需要用因果mask防止模型看到未来的token。mask的形状通常是(seq_len, seq_len)的下三角矩阵但在批处理和多头场景下需要扩展维度到(batch, nhead, seq_len, seq_len)。这里用masked_fill把需要屏蔽的位置设为负无穷softmax之后这些位置的权重就变成了0。3.3 位置编码的选择与实现Transformer本身没有位置信息所以需要额外注入位置编码。最常见的是正弦位置编码但实际工程中我更推荐可学习的位置编码尤其是在数据量足够的情况下。class LearnedPositionalEncoding(nn.Module): def __init__(self, d_model, max_len512): super().__init__() self.embedding nn.Embedding(max_len, d_model) def forward(self, x): # x: (batch, seq_len, d_model) seq_len x.size(1) positions torch.arange(seq_len, devicex.device).unsqueeze(0) return x self.embedding(positions)可学习位置编码的好处是灵活模型可以根据数据自己学习到合适的位置表示。但缺点是max_len需要预先设定推理时如果遇到超过max_len的序列就会报错。我的做法是在训练时就把max_len设得比实际需要大一些留出余量。正弦位置编码的优势是理论上可以外推到任意长度但实际效果在长序列上并不一定比可学习的好。我做过对比实验在序列长度不超过512的情况下两者的差异很小但可学习版本收敛更快。3.4 层归一化与残差连接的位置原始Transformer用的是Post-LN也就是LayerNorm(x Sublayer(x))。但后来的实践发现Pre-LN更稳定也就是x Sublayer(LayerNorm(x))。我强烈建议你用Pre-LN尤其是在模型比较深的时候Post-LN很容易出现梯度消失的问题需要很小心地调学习率和warmup步数。残差连接的作用不用多说它让梯度能够直接回传到浅层。但有一个细节是残差连接要求输入和输出的维度一致所以如果你的子层改变了维度就需要在残差分支上加一个投影矩阵。4. 训练循环那些教程不会告诉你的工程细节4.1 梯度累积解决显存不足当你想要更大的batch size但显存不够时梯度累积是最常用的技巧。原理很简单把一个大batch拆成几个小batch分别前向和反向但不清空梯度等累积够了再更新一次参数。accumulation_steps 4 optimizer.zero_grad() for i, batch in enumerate(dataloader): outputs model(batch) loss criterion(outputs, targets) / accumulation_steps loss.backward() if (i 1) % accumulation_steps 0: optimizer.step() optimizer.zero_grad()这里有一个容易犯的错误loss要除以accumulation_steps否则梯度会累积成原来的N倍相当于变相放大了学习率。另外如果你用了学习率调度器要注意调度器的step应该按实际参数更新次数来算而不是按batch数。4.2 混合精度训练的正确打开方式混合精度训练可以显著减少显存占用并加速计算但用不好会导致loss变成NaN。PyTorch提供了torch.cuda.amp来自动管理但你需要理解它背后的逻辑。from torch.cuda.amp import autocast, GradScaler scaler GradScaler() for batch in dataloader: optimizer.zero_grad() with autocast(): outputs model(batch) loss criterion(outputs, targets) scaler.scale(loss).backward() scaler.step(optimizer) scaler.update()GradScaler的作用是放大loss防止梯度在FP16下下溢。scaler.step会先检查梯度有没有出现inf或nan如果有就跳过这一步更新。scaler.update则动态调整缩放因子。我踩过的一个坑是在autocast上下文里做softmax或者layernorm的时候有时候会出现数值不稳定的情况。解决办法是对这些操作强制使用FP32可以用with autocast(enabledFalse):包起来。4.3 学习率调度与warmupTransformer类模型对学习率非常敏感warmup几乎是必须的。我常用的策略是线性warmup加上余弦退火。def get_lr(step, d_model, warmup_steps, total_steps): if step warmup_steps: return step / warmup_steps progress (step - warmup_steps) / (total_steps - warmup_steps) return 0.5 * (1 math.cos(math.pi * progress))warmup步数一般设置为总步数的5%到10%。为什么要warmup因为训练初期模型参数是随机初始化的梯度方向很不稳定如果直接用大学习率很容易把参数带偏。warmup让学习率从小逐渐增大给模型一个“热身”的过程。4.4 梯度裁剪与异常检测梯度裁剪是防止梯度爆炸的常用手段一般设置max_norm1.0就够了。但更重要的是异常检测。我建议在训练循环里加一段逻辑定期检查loss和梯度的范数如果发现异常就打印出来或者直接中断训练。torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm1.0) total_norm sum(p.grad.norm().item() ** 2 for p in model.parameters() if p.grad is not None) ** 0.5 if total_norm 100: print(fWarning: gradient norm is {total_norm})这个检查帮我省了很多时间。有一次训练到一半loss突然飙升就是因为某个batch的数据有问题导致梯度爆炸如果没有这个检查我可能要花几个小时才能定位到问题。5. 推理优化让模型跑得更快更省显存5.1 KV Cache的原理与实现自回归生成的时候每生成一个token都要重新计算整个序列的注意力这是非常浪费的。KV Cache的思路是把之前计算过的Key和Value缓存起来每次只计算新token的Query。class KVCache: def __init__(self): self.k_cache None self.v_cache None def update(self, k, v): if self.k_cache is None: self.k_cache k self.v_cache v else: self.k_cache torch.cat([self.k_cache, k], dim-2) self.v_cache torch.cat([self.v_cache, v], dim-2) return self.k_cache, self.v_cacheKV Cache能把推理速度提升几倍甚至十几倍但代价是显存占用会随着序列长度线性增长。对于长序列生成KV Cache的显存占用可能比模型本身还大。这时候就需要考虑量化KV Cache或者使用滑动窗口注意力。5.2 批处理推理的调度策略在线服务场景下请求是动态到达的每个请求的输入长度和输出长度都不一样。如果来一个请求就单独跑一次推理GPU利用率会非常低。这时候就需要动态批处理把多个请求攒在一起凑成一个batch再推理。但动态批处理有两个挑战第一不同请求的输出长度不同先完成的请求需要提前退出第二等待时间不能太长否则用户会感觉到明显的延迟。我的做法是设置一个最大等待时间比如50毫秒和一个最大batch size哪个先达到就触发一次推理。5.3 显存碎片与内存池长时间运行的推理服务显存碎片是一个隐形杀手。PyTorch有内置的缓存分配器但有时候还是会出现碎片问题。一个实用的技巧是定期调用torch.cuda.empty_cache()但这会带来性能抖动所以一般只在低峰期做。更好的做法是使用预分配的显存池把模型权重、KV Cache、中间激活值都预先分配好避免频繁的malloc和free。这在生产环境中非常关键我见过太多服务因为显存碎片导致OOM重启之后又好了但过一段时间又出现。6. 踩坑实录那些让我熬夜的瞬间6.1 数据加载中的死锁问题有一次我用DataLoader的num_workers8跑训练结果程序卡在第一个epoch就不动了。排查了半天才发现是我在Dataset的__getitem__里用了cv2.imread而OpenCV在多进程环境下如果没有正确设置会导致死锁。解决办法是在Dataset的__init__里设置cv2.setNumThreads(0)或者在worker初始化函数里做这个设置。这个坑的教训是任何第三方库在多进程环境下都可能有坑尤其是那些底层用了C的库。遇到卡死的情况先把num_workers设为0试试如果单进程能跑通那问题大概率出在多进程上。6.2 混合精度下的loss NaN混合精度训练最让人头疼的就是loss突然变成NaN。我遇到过一次排查了很久才发现是某个batch的输入里包含了极小的值在FP16下直接下溢成了0然后经过log操作变成了负无穷。解决办法是在数据预处理阶段做数值裁剪把输入值限制在一个合理的范围内。另一个常见原因是GradScaler的初始缩放因子设得太大。默认值是65536对于某些模型来说太大了会导致梯度溢出。可以尝试调小这个值比如从32768开始。6.3 模型保存与加载的版本兼容PyTorch的模型保存有两种方式保存整个模型和只保存状态字典。我强烈建议只保存状态字典因为保存整个模型会把类的定义也序列化进去换一个环境或者改了代码结构就加载不了了。# 推荐 torch.save(model.state_dict(), model.pt) model.load_state_dict(torch.load(model.pt)) # 不推荐 torch.save(model, model.pt) model torch.load(model.pt)还有一个细节是加载状态字典的时候要用map_location参数指定设备否则在CPU上保存的模型加载到GPU上会报错。6.4 分布式训练中的同步问题如果你用DistributedDataParallel做多卡训练一定要注意BatchNorm的同步。默认情况下每张卡上的BatchNorm是独立计算的这会导致统计量不一致。解决办法是使用SyncBatchNorm但它会带来额外的通信开销。另一个坑是随机种子。多卡训练时每张卡的随机种子必须不同否则数据增强的结果会完全一样相当于变相减小了batch size。一般用seed rank作为每张卡的种子。7. 从能跑到跑得好性能调优的实战思路7.1 先用profiler找到真正的瓶颈很多人一提到优化就开始瞎调参数这是效率最低的做法。正确的姿势是先用profiler找到瓶颈在哪里。PyTorch自带的torch.profiler就很好用。with torch.profiler.profile( activities[torch.profiler.ProfilerActivity.CPU, torch.profiler.ProfilerActivity.CUDA], scheduletorch.profiler.schedule(wait1, warmup1, active3), on_trace_readytorch.profiler.tensorboard_trace_handler(./log) ) as prof: for step, batch in enumerate(dataloader): train_step(batch) prof.step() if step 5: break跑完之后在TensorBoard里看你能清楚地看到每个操作占用了多少时间。我保证你会发现一些你完全没想到的地方在消耗时间比如某个transpose操作、某个不必要的CPU-GPU拷贝。7.2 算子融合与编译优化PyTorch 2.0引入了torch.compile可以自动做算子融合和内核优化。我实测下来在Transformer类模型上通常能有20%到50%的加速。model torch.compile(model)但torch.compile不是万能的它需要第一次运行来做编译所以如果你的模型有动态控制流可能会编译失败或者反复编译。这时候可以用modereduce-overhead来减少编译开销或者对特定模块单独编译。7.3 显存优化的几个实用技巧除了前面提到的混合精度和梯度累积还有几个技巧值得一试。梯度检查点用计算换显存把中间激活值丢掉反向传播时重新计算。对于特别深的模型这个技巧能把显存占用降低到原来的三分之一甚至更少。from torch.utils.checkpoint import checkpoint def forward_with_checkpointing(self, x): return checkpoint(self.layer, x)另一个技巧是参数卸载把暂时不用的参数放到CPU内存里需要的时候再加载到GPU。这在微调大模型的时候特别有用但会带来额外的传输开销需要权衡。8. 写在最后一些个人体会做AI工程这些年我最大的感受是框架封装得越好工程师的底层能力退化得越快。很多人能训出一个还不错的模型但你问他为什么用这个学习率、为什么用这个batch size、为什么用这个优化器他答不上来。这不是他的问题是工具太方便了。但工具方便不代表你可以不懂。当模型不收敛的时候当推理延迟超标的时候当显存不够用的时候能救你的只有对底层原理的理解。ai-engineering-from-scratch这个方向的价值就在于此它逼着你去面对那些被封装隐藏起来的细节让你在遇到问题时不是只能靠猜。我建议你在跟着实现一遍之后再回头去看那些框架的源码。你会发现原来那些看起来高深莫测的API底层不过是一堆矩阵乘法和简单的数学运算。这种“祛魅”的过程是每个AI工程师成长的必经之路。最后分享一个我常用的调试技巧任何新模型先用一个极小的数据集比如几十条样本跑通整个流程确认loss能降到接近零再上大规模数据。这能帮你快速排除代码层面的bug避免在数据量大的时候浪费时间。这个习惯帮我省了无数个加班的夜晚。
返回列表