1. 从零手搓AI工程:为什么我不建议你直接调包
很多人一听到“AI工程”这四个字,第一反应就是打开某个云平台,拖几个组件,调几个API,然后跑通一个Demo,就觉得自己已经掌握了。我刚开始也是这么想的,直到有一次线上推理服务在凌晨两点崩了,日志里全是显存溢出的报错,而我对着那堆封装好的接口完全不知道从哪下手排查。那一刻我才意识到,只会调包的人,永远只能停留在“能用”的层面,一旦出了问题,连问题出在哪个环节都说不清楚。
ai-engineering-from-scratch这个标题,核心不是让你去重新发明Transformer,而是让你把AI工程链路里的每一个关键环节,都亲手实现一遍。从数据加载、分词、模型结构搭建、训练循环、梯度累积、混合精度,到推理优化、批处理调度、显存管理,这些东西如果你只是调用现成的库,你永远不会理解它们为什么存在,也不会知道在什么场景下该动哪个旋钮。
这篇文章适合两类人:一类是刚入门AI方向的学生或者转行者,想真正搞懂一个模型从数据到部署到底经历了什么;另一类是有一定调包经验但遇到瓶颈的工程师,想往下钻一层,搞清楚框架底层到底在干什么。我会按照一个最小可用的语言模型训练与推理链路,把每个环节的“为什么”和“怎么做”都拆开讲清楚,代码以PyTorch为主,但思路是通用的。
需要提前说明的是,我不会给你一个“复制粘贴就能跑”的完整项目,因为那样对你没有帮助。我会给你每个模块的核心逻辑、关键参数的计算方式、以及我在实际踩坑中总结出来的经验。你跟着走一遍,收获会比直接clone一个仓库大得多。
2. 数据管道的搭建:别让IO成为你的第一个瓶颈
2.1 为什么数据加载值得单独拿出来讲
大部分教程在讲数据加载的时候,就是一句DataLoader(dataset, batch_size=32, shuffle=True)带过。但实际做AI工程的时候,数据管道往往是第一个让你崩溃的地方。我见过太多人模型代码写得漂漂亮亮,结果训练速度慢得离谱,最后发现是数据加载成了瓶颈,GPU利用率常年徘徊在20%以下。
一个合格的数据管道需要解决三个问题:读取效率、内存占用、批处理策略。读取效率决定了你的GPU会不会饿着,内存占用决定了你能不能处理大规模数据集,批处理策略则直接影响模型收敛的质量。
2.2 从原始文本到Token序列的完整链路
假设你手里有一堆纯文本文件,第一步是构建词表。这里我不建议你直接用现成的分词器,而是先手写一个简单的字符级或者词级分词器,理解分词的本质。
# 一个极简的词级分词器实现 from collections import Counter class SimpleTokenizer: def __init__(self, vocab_size=10000): 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, start=2): 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>两个特殊token?<pad>用于批处理时对齐序列长度,<unk>用于处理词表外的词。第二,为什么按频率截断而不是全量保留?因为词表越大,嵌入层的参数量越大,低频词带来的收益远小于它占用的显存和计算量。
实际工程中,你会遇到文本长度差异极大的情况。有的样本只有十几个token,有的有几千个。如果直接按最大长度padding,显存浪费会非常严重。我的做法是采用动态padding,也就是每个batch内按当前batch的最大长度来padding,而不是全局最大长度。
def collate_fn(batch, pad_idx=0): # 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, dtype=torch.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_model=512, nhead=8的多头注意力层,它的参数量到底是多少?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, mask=None): # 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_len=512): 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, device=x.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(enabled=False):包起来。
4.3 学习率调度与warmup
Transformer类模型对学习率非常敏感,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_norm=1.0就够了。但更重要的是异常检测。我建议在训练循环里加一段逻辑,定期检查loss和梯度的范数,如果发现异常就打印出来或者直接中断训练。
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.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(f"Warning: 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_workers=8跑训练,结果程序卡在第一个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], schedule=torch.profiler.schedule(wait=1, warmup=1, active=3), on_trace_ready=torch.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不是万能的,它需要第一次运行来做编译,所以如果你的模型有动态控制流,可能会编译失败或者反复编译。这时候可以用mode="reduce-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,避免在数据量大的时候浪费时间。这个习惯帮我省了无数个加班的夜晚。