1. 从零手搓AI工程:为什么我不建议你直接调包
很多人一听到"AI工程"这四个字,第一反应就是打开某个框架的文档,pip install 一把梭,然后照着官方示例跑一个MNIST手写数字识别,看到准确率98%就觉得自己入门了。我当年也是这么干的,结果到了真实项目里,数据管道一塌糊涂,模型训练到一半显存爆了,推理延迟高得没法上线,整个人直接懵掉。
ai-engineering-from-scratch这个标题背后的核心诉求,其实不是让你去造一个比PyTorch还牛的框架,而是让你把AI工程这条链路上每一个环节的"黑盒"都拆开看一遍。你知道张量在内存里是怎么排布的吗?你知道反向传播时梯度到底是怎么一层层传回去的吗?你知道一个训练好的模型从 checkpoint 到线上服务,中间要经过哪些转换和优化吗?如果这些你都说不上来,那调包调得再熟练,遇到诡异bug时也只能干瞪眼。
这篇文章适合三类人:第一类是有一定Python基础、想真正理解AI系统底层运转逻辑的开发者;第二类是做过后端或数据工程、想转行到AI方向但不想只做"调参侠"的工程师;第三类是在校学生,课程里学了理论但没动手搭过完整pipeline的人。我会从最基础的数据表示开始,一路讲到模型部署和性能优化,每个环节都给出可复现的代码和踩坑记录。全文不依赖任何高级框架的封装,核心逻辑全部手写,让你看清楚每一行代码到底在干什么。
需要提前说明的是,文中涉及的具体数值和配置是基于我自己的实验环境(单卡24GB显存、Ubuntu 22.04、Python 3.10)总结的,你在不同硬件上可能需要微调。但原理和思路是通用的,这也是"from scratch"的意义所在——你掌握的是方法,不是某个特定环境下的咒语。
2. 数据管道:AI工程里最容易被低估的脏活累活
2.1 为什么你的模型效果差,八成是数据管道的问题
我见过太多人把大量时间花在调模型结构上,却对数据管道敷衍了事。实际情况是,在一个典型的AI工程项目里,数据管道的代码量往往占到整个项目的60%以上,而且出问题的概率也最高。模型结构再优雅,喂进去的数据有问题,结果一定好不了。
从零搭建数据管道,你需要解决几个核心问题:数据怎么读、怎么洗、怎么切、怎么喂。听起来简单,但每个环节都有坑。比如读取环节,小文件太多会导致IO瓶颈,大文件一次性读入又可能撑爆内存。我的做法是实现一个基于生成器的流式读取器,每次只加载一个batch的数据,同时用多进程预取来掩盖IO延迟。
import numpy as np from multiprocessing import Pool class StreamDataset: def __init__(self, file_paths, batch_size=32, shuffle=True): self.file_paths = file_paths self.batch_size = batch_size self.shuffle = shuffle self.indices = np.arange(len(file_paths)) def _load_single(self, idx): # 实际项目中这里可能是读图片、读音频、读文本 data = np.load(self.file_paths[idx]) return data def __iter__(self): if self.shuffle: np.random.shuffle(self.indices) for i in range(0, len(self.indices), self.batch_size): batch_idx = self.indices[i:i+self.batch_size] with Pool(4) as p: batch_data = p.map(self._load_single, batch_idx) yield np.stack(batch_data)这段代码看起来简单,但有几个细节值得说。第一,Pool的进程数不是越多越好,一般设置为CPU核心数的70%左右,留一些给主进程和其他任务。第二,np.stack要求每个样本形状一致,如果你的数据变长(比如文本),需要在这里做padding或截断。第三,shuffle在每个epoch开始时做一次全局打乱就够了,不需要每个batch都打乱,那样反而会增加随机性带来的方差。
2.2 数据清洗:那些教科书不会告诉你的经验
数据清洗是另一个重灾区。教科书上通常只讲"去除缺失值、处理异常值",但实际项目中你会发现,缺失值的处理方式直接影响模型效果。比如数值型特征,用均值填充和用中位数填充,在长尾分布下差异巨大;类别型特征,把缺失当作一个独立类别往往比强行填充效果更好。
我的一般原则是:先统计缺失模式,看缺失是随机的还是有规律的。如果某个特征的缺失率超过40%,而且缺失本身可能携带信息(比如用户没填年龄可能是因为不想透露),那就把"是否缺失"作为一个额外的二值特征加进去。这个技巧在很多实际项目中都带来了明显的效果提升。
还有一个容易被忽略的点是数据泄漏。比如你在做时间序列预测,如果用全局均值来填充缺失值,那就把未来信息泄漏到了训练集里。正确的做法是只用当前时间点之前的数据来计算填充值。类似地,做标准化时,均值方差只能从训练集计算,然后应用到验证集和测试集。这些细节在从零实现时你必须自己处理,而调包时框架可能已经帮你做了,你反而不知道发生了什么。
2.3 批处理与内存管理的平衡术
批处理大小(batch size)的选择是一个经典的权衡问题。大batch训练更稳定、GPU利用率更高,但内存占用大,而且可能陷入尖锐极小值导致泛化变差;小batch泛化可能更好,但训练速度慢、梯度噪声大。
我的经验是,先从硬件能承受的最大batch size开始试,如果效果不理想再逐步减小。同时配合学习率的调整——batch size翻倍时,学习率通常也可以适当增大。另外,梯度累积(gradient accumulation)是一个很实用的技巧:用多个小batch的梯度累加来模拟大batch的效果,既节省内存又保持训练稳定性。
# 梯度累积示例 accumulation_steps = 4 optimizer.zero_grad() for i, batch in enumerate(dataloader): loss = model(batch) loss = loss / accumulation_steps loss.backward() if (i + 1) % accumulation_steps == 0: optimizer.step() optimizer.zero_grad()这里有个坑:如果你用了BatchNorm,梯度累积时统计量的更新会和小batch一致,可能和真正的大batch行为不同。这时候可以考虑用GroupNorm或者LayerNorm替代,或者接受这个差异。我在实际项目中遇到过因为这个问题导致训练和推理行为不一致的情况,排查了很久才发现是BatchNorm的running stats在作怪。
3. 模型实现:手写反向传播到底值不值得
3.1 从标量到张量:自动微分的核心思想
很多人觉得手写反向传播是浪费时间,反正框架都能自动求导。但我的观点是,你至少要从零实现一次标量级别的自动微分,理解计算图、链式法则和梯度累加这三个概念。一旦理解了这些,再看框架的自动求导机制就会豁然开朗。
自动微分的核心是:每个操作都记录自己的前向输出和局部梯度,反向传播时按照计算图的拓扑逆序,把上游传来的梯度乘以局部梯度再传给下游。对于有多个消费者的节点,梯度需要累加。这个机制用几百行代码就能实现一个简化版。
class Value: def __init__(self, data, children=(), op=''): self.data = data self.grad = 0.0 self._backward = lambda: None self._children = set(children) self._op = op def __add__(self, other): other = other if isinstance(other, Value) else Value(other) out = Value(self.data + other.data, (self, other), '+') def _backward(): self.grad += out.grad other.grad += out.grad out._backward = _backward return out def __mul__(self, other): other = other if isinstance(other, Value) else Value(other) out = Value(self.data * other.data, (self, other), '*') def _backward(): self.grad += other.data * out.grad other.grad += self.data * out.grad out._backward = _backward return out def backward(self): topo = [] visited = set() def build_topo(v): if v not in visited: visited.add(v) for child in v._children: build_topo(child) topo.append(v) build_topo(self) self.grad = 1.0 for v in reversed(topo): v._backward()这段代码虽然简单,但已经包含了自动微分的全部核心要素。你可以用它搭一个多层感知机,在简单的数据集上训练,观察梯度是如何流动的。当你亲手调试过梯度消失或梯度爆炸的问题后,对模型训练的理解会深刻很多。
3.2 张量级别的实现:性能与可读性的取舍
标量版本理解原理足够了,但真正做项目时你需要张量级别的操作。从零实现一个支持广播、矩阵乘法、卷积等操作的张量库,工作量不小,但也不是不可能。我的建议是至少实现以下几个核心操作:矩阵乘法、逐元素加法、ReLU、Softmax、交叉熵损失。这几个操作组合起来就能搭一个完整的分类模型。
实现时最大的挑战是广播机制下的梯度计算。比如一个形状为(32, 128)的张量和一个形状为(128,)的偏置相加,反向传播时偏置的梯度需要对batch维度求和。这个逻辑必须小心处理,否则梯度形状对不上,训练直接报错。
class Tensor: def __init__(self, data, requires_grad=False): self.data = np.array(data, dtype=np.float32) self.requires_grad = requires_grad self.grad = None self._backward = lambda: None def __matmul__(self, other): out = Tensor(self.data @ other.data, self.requires_grad or other.requires_grad) def _backward(): if self.requires_grad: self.grad = out.grad @ other.data.T if other.requires_grad: other.grad = self.data.T @ out.grad out._backward = _backward return out实际项目中,我建议你手写实现核心算子,但没必要全部重造轮子。可以用NumPy做底层计算,用CuPy在需要时切换到GPU,这样既保持了代码的可读性,又不会在性能上太吃亏。关键是你要清楚每个算子的前向和反向逻辑,这样遇到数值不稳定时才知道从哪里下手。
3.3 训练循环:那些框架帮你隐藏的细节
框架的model.fit()或者trainer.train()帮你隐藏了大量细节,从零实现时你需要自己处理:学习率调度、梯度裁剪、权重衰减、早停、模型保存与恢复。每一个都值得单独拿出来说。
学习率调度我常用的是余弦退火配合热重启,在训练初期用较大的学习率快速下降,后期用小学习率精细调整。梯度裁剪在RNN和Transformer里几乎是必须的,一般按范数裁剪,阈值设在1.0到5.0之间。权重衰减要注意不要应用到偏置和归一化层的参数上,这个细节在很多论文里都有讨论。
def train_epoch(model, dataloader, optimizer, clip_norm=1.0): model.train() total_loss = 0 for batch in dataloader: optimizer.zero_grad() output = model(batch['input']) loss = cross_entropy(output, batch['label']) loss.backward() # 梯度裁剪 torch.nn.utils.clip_grad_norm_(model.parameters(), clip_norm) optimizer.step() total_loss += loss.item() return total_loss / len(dataloader)早停策略我一般看验证集损失,连续5个epoch没有下降就停止训练,同时保存验证集损失最低的那个checkpoint。这里有个经验:验证集损失和验证集准确率不一定同步变化,分类任务中我更倾向于看准确率,回归任务看损失。另外,如果训练集损失还在下降但验证集损失开始上升,那就是过拟合的典型信号,早停或者增加正则化都可以。
4. 训练基础设施:从单卡到多卡的工程挑战
4.1 混合精度训练:省显存不是唯一目的
混合精度训练(Mixed Precision Training)现在已经是标配了,但很多人只知道它能省显存,不知道它还能加速训练。原理很简单:前向和反向用FP16计算,速度快、显存占用小;但参数更新用FP32,保持数值稳定性。关键是要做损失缩放(loss scaling),因为FP16的表示范围有限,小梯度容易下溢成0。
从零实现混合精度训练,你需要手动管理FP16和FP32的转换,以及动态损失缩放。动态损失缩放的逻辑是:如果连续多个step没有出现梯度溢出(inf或nan),就增大缩放因子;如果出现溢出,就减小缩放因子并跳过这个step。这个机制在PyTorch的torch.cuda.amp里已经封装好了,但理解它的工作原理对调试很有帮助。
scaler = torch.cuda.amp.GradScaler() for batch in dataloader: optimizer.zero_grad() with torch.cuda.amp.autocast(): output = model(batch['input']) loss = criterion(output, batch['label']) scaler.scale(loss).backward() scaler.step(optimizer) scaler.update()我踩过的一个坑是:某些操作在FP16下会溢出,比如大数相乘、指数运算。这时候需要强制这些操作在FP32下执行,用torch.cuda.amp.autocast(enabled=False)包起来。另外,BatchNorm的统计量更新最好也在FP32下做,否则running mean和running var的精度损失会累积。
4.2 数据并行与模型并行:什么时候该用哪种
数据并行(Data Parallelism)是最常用的多卡策略:每张卡上放一份完整的模型副本,每个batch的数据切分到各卡上,梯度汇总后统一更新。PyTorch的DistributedDataParallel(DDP)是目前的主流选择,比DataParallel(DP)效率高很多,因为DDP用多进程而不是多线程,避免了GIL的限制。
模型并行(Model Parallelism)适用于单卡放不下整个模型的情况,把模型的不同层放到不同卡上。但这样会导致卡间通信频繁,效率往往不如数据并行。实际项目中,我优先考虑数据并行,只有当模型实在太大时才考虑模型并行或者流水线并行。
选择并行策略时,通信开销是核心考量。数据并行的通信量正比于模型参数量,模型并行的通信量正比于层之间的激活值大小。对于Transformer类模型,参数量大但激活值相对小,数据并行通常更划算。对于超大规模模型,可能需要混合并行策略,这就涉及到更复杂的工程实现了。
4.3 检查点与恢复:别让一次断电毁掉一周的训练
训练大模型动辄几天甚至几周,中间任何意外中断都是灾难。所以检查点(checkpoint)机制必须做好。我一般每N个step保存一次,同时保存优化器状态、学习率调度器状态、当前epoch和step数,以及随机数生成器的状态(保证恢复后数据顺序一致)。
def save_checkpoint(model, optimizer, scheduler, epoch, step, path): torch.save({ 'model_state': model.state_dict(), 'optimizer_state': optimizer.state_dict(), 'scheduler_state': scheduler.state_dict(), 'epoch': epoch, 'step': step, 'rng_state': torch.get_rng_state(), 'cuda_rng_state': torch.cuda.get_rng_state_all(), }, path) def load_checkpoint(model, optimizer, scheduler, path): ckpt = torch.load(path) model.load_state_dict(ckpt['model_state']) optimizer.load_state_dict(ckpt['optimizer_state']) scheduler.load_state_dict(ckpt['scheduler_state']) torch.set_rng_state(ckpt['rng_state']) torch.cuda.set_rng_state_all(ckpt['cuda_rng_state']) return ckpt['epoch'], ckpt['step']这里有个细节:保存检查点最好用临时文件加原子重命名的方式,避免保存过程中断电导致检查点文件损坏。另外,检查点文件通常很大,如果存储空间有限,可以只保留最近的几个,或者用增量保存的方式只存变化的部分。
5. 推理部署:模型上线前的最后一公里
5.1 模型导出与格式转换:ONNX是个好中间站
训练好的模型不能直接扔到线上服务里,通常需要先导出成通用格式。ONNX(Open Neural Network Exchange)是目前最常用的中间表示,它把模型的计算图序列化成一个标准格式,可以被多种推理引擎加载。
导出ONNX时最常见的坑是动态维度处理。如果你的模型支持变长输入(比如不同长度的文本),导出时需要指定动态轴。另外,某些PyTorch操作在ONNX里没有对应实现,导出会失败,这时候需要改写模型或者自定义ONNX算子。
torch.onnx.export( model, dummy_input, "model.onnx", input_names=["input"], output_names=["output"], dynamic_axes={"input": {0: "batch", 1: "sequence"}, "output": {0: "batch", 1: "sequence"}}, opset_version=13 )导出后一定要验证:用ONNX Runtime加载模型,和PyTorch的输出对比,确保数值误差在可接受范围内(一般1e-4以内)。我遇到过导出后精度下降的情况,排查发现是某个自定义算子在ONNX里的实现和PyTorch不一致,最后通过替换算子解决了。
5.2 推理优化:量化、剪枝与算子融合
模型上线的核心指标是延迟和吞吐量。优化手段主要有三类:量化、剪枝和算子融合。
量化是把FP32的权重和激活值用INT8表示,模型大小直接缩小4倍,推理速度也能提升2到4倍。但量化会带来精度损失,需要做量化感知训练(QAT)或者训练后量化(PTQ)加校准。我的经验是,对于大多数视觉模型,INT8量化后精度损失在1%以内是可以接受的;但对于一些对数值敏感的模型(比如检测小目标),可能需要混合量化,只量化部分层。
剪枝是去掉模型中不重要的权重或神经元。结构化剪枝(去掉整个通道或层)对推理速度有实际提升,非结构化剪枝(去掉单个权重)主要减小模型大小,对速度提升有限,因为稀疏矩阵运算在通用硬件上并不高效。
算子融合是把多个连续的操作合并成一个,减少内存访问和kernel启动开销。比如Conv+BN+ReLU可以融合成一个算子,这在推理引擎里通常是自动做的,但你需要确保导出时这些操作是连续的,中间没有插入其他操作打断融合。
5.3 服务化:从单请求到高并发的工程实践
模型服务化要考虑的问题很多:并发处理、批处理、超时控制、降级策略。最简单的做法是用Flask或FastAPI起一个HTTP服务,每个请求单独推理。但这样吞吐量很低,因为GPU利用率上不去。
更好的做法是实现动态批处理(dynamic batching):服务端维护一个请求队列,当队列长度达到阈值或者等待时间超过阈值时,把多个请求合并成一个batch一起推理。这样能显著提升吞吐量,代价是增加了单请求的延迟。
import asyncio from collections import deque class BatchServer: def __init__(self, model, max_batch=32, max_wait=0.01): self.model = model self.max_batch = max_batch self.max_wait = max_wait self.queue = deque() async def infer(self, input_data): future = asyncio.Future() self.queue.append((input_data, future)) if len(self.queue) >= self.max_batch: await self._process_batch() else: await asyncio.sleep(self.max_wait) if self.queue: await self._process_batch() return await future async def _process_batch(self): batch = list(self.queue) self.queue.clear() inputs = [item[0] for item in batch] outputs = self.model(inputs) for (_, future), output in zip(batch, outputs): future.set_result(output)这个实现是简化版,实际项目中还需要考虑超时、错误处理、优先级等。另外,GPU推理是异步的,要用CUDA流来管理,避免CPU等待GPU。我一般会用Triton Inference Server或者TorchServe这类成熟的服务框架,它们已经处理好了大部分工程细节,你只需要关注模型本身的优化。
6. 监控与迭代:上线只是开始
6.1 线上指标监控:别只看准确率
模型上线后,你需要监控的指标远不止准确率。延迟的P50、P95、P99分位数,吞吐量,GPU利用率,显存占用,这些工程指标直接关系到服务能不能稳定运行。同时,数据分布的变化、预测结果的分布变化,这些业务指标能帮你发现模型退化。
我一般会记录每次请求的输入特征统计量(均值、方差、分位数)和输出置信度分布。如果发现输入分布和训练分布偏离太大,或者输出置信度整体下降,就说明模型可能遇到了分布偏移,需要考虑重新训练或者在线更新。
6.2 数据回流与持续训练
线上服务产生的数据是宝贵的训练资源。把线上请求的输入和模型输出(以及后续的用户反馈)收集起来,经过清洗和标注后加入训练集,这就是数据回流。持续训练就是用新数据不断更新模型,保持模型对最新分布的适应能力。
但数据回流有个陷阱:如果线上模型有偏差,它产生的数据也会有偏差,用这些数据训练会让偏差进一步放大。所以需要定期用人工标注的数据做校准,或者用一些去偏技术。另外,持续训练要考虑灾难性遗忘的问题,新数据训练时最好混合一部分旧数据,或者用弹性权重巩固(EWC)等方法保护重要参数。
6.3 版本管理与回滚:给自己留好后路
模型版本管理经常被忽视,但出了问题要回滚时你就知道它有多重要了。每个上线的模型版本都要记录:训练数据版本、代码版本、超参数配置、评估指标。我一般用MLflow或者Weights & Biases这类工具来管理实验和模型版本,确保任何一次上线都能追溯到具体的训练配置。
回滚策略要提前设计好:新模型上线时先做小流量灰度,观察一段时间指标正常后再逐步放大流量。如果发现异常,能快速切回旧版本。这个流程听起来简单,但真到出问题时,如果没有提前准备好,手忙脚乱之下很容易出更大的事故。
7. 一些踩坑之后的个人体会
从零搭建AI工程系统这件事,我前后折腾了好几年,踩过的坑不计其数。最大的体会是:不要试图一次性把所有环节都做到完美。先跑通一个最简版本,哪怕数据管道很粗糙、模型很小、部署很简陋,只要端到端能跑起来,你就有了一个可以迭代的基础。然后每次只优化一个环节,测量优化前后的差异,确保每次改动都有正向收益。
另一个体会是:日志和监控要尽早做。我早期做项目时总觉得这些是"运维的事",结果模型效果波动时完全不知道从哪里查起。后来养成了习惯,每个模块都打详细的日志,关键指标都做可视化,排查问题的效率提升了好几倍。
最后,保持对底层原理的好奇心。框架更新换代很快,但底层的数学原理和工程原则变化很慢。你把数据管道、自动微分、并行训练、推理优化这些核心环节的"为什么"搞清楚了,换什么框架都能快速上手。这也是"from scratch"最大的价值——你获得的是可迁移的能力,而不是某个特定工具的熟练度。