新闻详情

新闻详情

首页 / 资讯中心 / 详情

从零手搓AI工程:深入理解张量、自动求导与训练循环

发布时间:2026/9/28 14:01:00来源:尧图网络
从零手搓AI工程:深入理解张量、自动求导与训练循环
1. 从零手搓AI工程为什么我不建议你直接调包很多人一上来就想用现成的框架把模型跑起来觉得“能出结果就行”。但如果你真的想搞懂AI工程我建议你反着来——先别碰那些封装好的高级API从最底层的东西开始手搓一遍。这个项目叫“ai-engineering-from-scratch”核心思路就是把AI工程拆成一个个可以独立理解、独立实现的模块从张量运算、自动求导、到训练循环、推理部署全部自己写一遍。听起来很折腾确实折腾但折腾完这一轮你对整个AI系统的理解会完全不一样。这个项目适合谁适合那些已经会用PyTorch或TensorFlow跑通几个demo但总觉得“心里没底”的人。你知道model.fit()能训练模型但不知道里面到底发生了什么你知道loss.backward()能算梯度但不知道链式法则具体怎么在计算图上传播。如果你有这种感觉那从零实现一遍就是最好的解药。整个过程不需要GPU集群一台普通笔记本就能跑因为我们要实现的是“最小可用版本”不是工业级大模型。我自己的经历是在调了两年包之后有一次线上模型出了个诡异的梯度消失问题排查了三天才发现是自定义层里一个张量操作写错了。从那以后我就下定决心把核心模块全部手写一遍。这个项目就是那次“补课”的产物。下面我会把整个实现路径拆开讲包括每一步为什么这么设计、有哪些坑、以及怎么验证你写的东西是对的。2. 张量库一切AI工程的起点2.1 为什么不用NumPy而要从头写你可能会问NumPy已经提供了多维数组和广播机制为什么还要自己写一个张量库答案在于自动求导。NumPy的数组是“死”的它只存数据不记录计算历史。而AI训练的核心是反向传播反向传播需要知道每个操作的前向输入和局部梯度。所以我们需要一个能构建计算图的数据结构这就是张量库存在的意义。具体来说我们要实现一个Tensor类它至少包含这几个东西data底层数据、grad梯度、requires_grad是否需要求导、_backward反向传播函数、_prev前驱节点集合。前向计算时每个操作不仅算出结果还要把“怎么算的”记录下来。比如c a b除了得到c的data还要记录c的梯度怎么传给a和b。这里有个关键设计决策用动态图还是静态图。动态图就是每次前向传播都重新构建计算图PyTorch走的就是这条路静态图是先定义好图再执行TensorFlow 1.x是典型代表。从零实现的话动态图更直观因为你可以用Python的原生控制流调试也方便。代价是每次迭代都要重建图性能上有一点开销但对于学习目的来说完全可以接受。2.2 广播机制的手写实现广播是张量运算里最容易出错的地方。两个形状不同的张量相加NumPy会自动扩展维度但反向传播时梯度需要“缩回”原来的形状。举个例子形状(3, 4)的A和形状(4,)的B相加结果是(3, 4)。反向传播时B的梯度应该是结果梯度在第一个维度上求和变成(4,)。手写广播的反向传播核心逻辑是前向传播时把两个张量的形状对齐反向传播时把梯度按广播规则还原。具体步骤是先比较两个形状从右往左逐维对齐缺失的维度补1然后对每个维度如果一方是1而另一方大于1就把梯度在该维度上求和并保持维度。这个逻辑写起来大概二十行代码但如果不小心写错了训练时loss会莫名其妙地不下降而且很难debug。我踩过的坑是忘记处理“维度数为1”的情况。比如形状(3, 1)和(3, 4)相加前向传播时(3, 1)会广播成(3, 4)反向传播时梯度要在第二个维度上求和变成(3, 1)。如果直接sum而不keepdimsTrue形状就变成(3,)了后续更新参数时形状不匹配直接报错。所以记住求和还原梯度时一定要保持维度。2.3 计算图的构建与内存管理每次前向传播都会生成一张计算图图里的节点是张量边是操作。如果不对图做管理内存会迅速爆炸。PyTorch的做法是每次反向传播结束后释放中间节点的计算图。我们手写的时候也要做类似的事情。具体实现上可以在backward()函数里用一个拓扑排序来遍历所有节点然后逐个调用_backward。拓扑排序保证了每个节点的梯度在被使用之前已经计算完毕。遍历完之后把所有中间节点的_prev清空这样Python的垃圾回收就能释放内存。还有一个细节梯度累积。如果你不把grad清零多次反向传播的梯度会累加。这在某些场景下是有用的比如梯度累积模拟大batch但大多数时候是个坑。所以每次backward()之前要么手动清零要么在优化器里清零。我建议在优化器的step()里做清零这样逻辑更清晰。3. 自动求导链式法则的代码化3.1 反向传播的数学本质自动求导的核心就是链式法则。假设有一个计算图x - f - g - loss那么loss对x的梯度是dloss/dg * dg/df * df/dx。在代码里每个操作都知道自己的局部梯度反向传播就是把这些局部梯度乘起来。以乘法为例c a * b那么dc/da bdc/db a。所以反向传播时a的梯度加上c.grad * b.datab的梯度加上c.grad * a.data。注意这里是“加上”而不是“赋值”因为一个张量可能被多个操作使用梯度需要累加。这里有个容易混淆的点梯度的形状。如果a和b形状不同比如a是(3, 4)b是(4,)那么c.grad * b.data的形状是(3, 4)但a的梯度应该是(3, 4)b的梯度应该是(4,)。所以乘法操作的反向传播需要处理广播把梯度缩回原来的形状。这就是为什么上一节强调广播的反向实现。3.2 常见操作的梯度推导我们至少需要实现这些操作的梯度加法、乘法、矩阵乘法、ReLU、Sigmoid、Tanh、Softmax、Cross Entropy。每个操作的梯度推导都不复杂但有几个容易写错的地方。矩阵乘法C A B其中A是(m, n)B是(n, p)C是(m, p)。那么dC/dA dC B.TdC/dB A.T dC。注意转置的位置写反了形状就对不上。ReLUy max(0, x)梯度是x 0 ? 1 : 0。实现时要注意如果x恰好等于0梯度取0还是1实践中取0更常见因为0处的导数不定义取0相当于忽略这个点。Softmax Cross Entropy这两个通常一起实现因为单独算Softmax的梯度很麻烦但和Cross Entropy结合后梯度简化为softmax_output - one_hot_label。这个简化是AI工程里最优雅的数学技巧之一一定要自己推导一遍。3.3 梯度检查怎么确认你写对了手写自动求导最大的风险是梯度算错了但loss还在下降只是收敛得慢或者收敛到次优解。所以必须做梯度检查。方法是用数值微分算一个近似梯度和你写的反向传播算出来的梯度对比。数值微分的公式是(f(x eps) - f(x - eps)) / (2 * eps)其中eps取1e-5左右。如果两个梯度的相对误差小于1e-5说明你的实现基本正确。如果误差很大那就要检查是哪个操作的梯度写错了。我建议每实现一个操作就做一次梯度检查不要等全部写完再查。因为一旦出错你很难定位是哪个操作的问题。梯度检查的代码很简单但它是保证正确性的最后一道防线。4. 训练循环从手写SGD到Adam4.1 损失函数的选择与实现损失函数是训练的目标它告诉模型“什么是好的”。分类任务常用Cross Entropy回归任务常用MSE。从零实现的话Cross Entropy需要和Softmax一起写因为单独算Softmax的梯度数值不稳定。具体实现时先算logits然后减去最大值防止指数爆炸再算log_softmax最后取负对数似然。这个顺序很重要如果先算Softmax再取log数值上会不稳定因为Softmax的输出可能非常接近0log之后变成负无穷。MSE的梯度很简单2 * (pred - target) / n。但要注意如果pred和target的形状不同需要先广播。实践中pred和target的形状通常是一样的所以直接相减就行。4.2 优化器的演进SGD、Momentum、AdamSGD是最基础的优化器param param - lr * grad。它的缺点是容易陷入局部最优而且在鞍点附近震荡。Momentum通过引入“速度”来平滑更新方向v beta * v gradparam param - lr * v。这样在梯度方向一致的维度上加速在震荡的维度上减速。Adam进一步引入了自适应学习率对每个参数根据梯度的一阶矩和二阶矩来调整学习率。具体公式是m beta1 * m (1 - beta1) * gradv beta2 * v (1 - beta2) * grad^2然后做偏差修正最后param param - lr * m / (sqrt(v) eps)。从零实现Adam的时候要注意偏差修正。因为m和v初始化为0前几步的估计是有偏的需要除以(1 - beta1^t)和(1 - beta2^t)。如果忘了这一步训练初期会非常不稳定。4.3 学习率调度与早停学习率太大loss会震荡甚至发散学习率太小收敛太慢。实践中常用的是“阶梯下降”或“余弦退火”。阶梯下降是每过几个epoch把学习率乘以0.1余弦退火是按余弦函数从初始学习率降到0。早停是防止过拟合的简单有效方法如果验证集loss连续几个epoch不下降就停止训练。实现时用一个计数器每次验证集loss下降就重置否则加1超过阈值就break。我自己的经验是先用小学习率跑通再调大。很多人一上来就设0.1结果loss直接变成nan。正确的做法是从1e-3或1e-4开始确认loss在下降再逐步调大。如果loss变成nan先检查梯度是不是爆炸了可以加梯度裁剪。5. 模型组装从线性层到多层感知机5.1 线性层的实现细节线性层就是y x W b其中x是(batch, in_features)W是(in_features, out_features)b是(out_features,)。初始化W的时候不能用全零否则所有神经元的梯度都一样网络学不到东西。常用的是Xavier初始化或He初始化。Xavier初始化W ~ U(-sqrt(6/(inout)), sqrt(6/(inout)))。He初始化W ~ N(0, sqrt(2/in))。实践中ReLU激活函数用He初始化Sigmoid或Tanh用Xavier初始化。反向传播时dW x.T dydb sum(dy, axis0)dx dy W.T。注意db是在batch维度上求和因为b对所有样本共享。5.2 激活函数的选择与梯度特性ReLU是最常用的激活函数计算简单梯度不饱和。但ReLU有个问题负半轴的梯度是0如果某个神经元一直输出负值它的梯度永远是0参数永远不更新这就是“神经元死亡”。LeakyReLU给负半轴一个小的斜率比如0.01缓解这个问题。Sigmoid和Tanh在深层网络中容易梯度消失因为它们的导数最大只有0.25Sigmoid和1Tanh多层相乘后梯度会指数衰减。所以现代网络基本都用ReLU家族。实现激活函数时反向传播要利用前向传播的输出。比如Sigmoid的梯度是y * (1 - y)其中y是前向输出。所以在前向传播时要把输出存下来反向传播时直接用。5.3 多层感知机的组装与验证把线性层和激活函数交替堆叠就得到了多层感知机。比如Linear(784, 256) - ReLU - Linear(256, 128) - ReLU - Linear(128, 10)。这个网络可以处理MNIST手写数字分类。组装的时候要注意维度的匹配。每一层的输出维度必须等于下一层的输入维度。验证的方法是用随机数据跑一次前向传播看输出形状是不是(batch, 10)。然后再跑一次反向传播看所有参数的梯度形状是不是和参数本身一致。我建议先在一个小数据集上过拟合。比如取100个样本训练到loss接近0。如果连100个样本都过拟合不了说明网络结构或训练逻辑有问题。这个技巧能快速定位bug比直接上全量数据高效得多。6. 数据加载与预处理被忽视的关键环节6.1 为什么数据管道比模型更重要在实际项目中数据管道的质量往往比模型结构更影响最终效果。一个常见的现象是同样的模型数据预处理做好一点准确率能提升好几个百分点。从零实现的话我们需要一个能批量加载数据、打乱顺序、做归一化的管道。批量加载的核心是DataLoader它接收一个数据集每次返回一个batch。实现时要注意最后一个batch可能不满需要处理每个epoch开始前要打乱顺序防止模型学到样本顺序的规律。归一化是把数据缩放到均值为0、方差为1的分布。对于图像数据通常是除以255再减去均值除以标准差。归一化的目的是让梯度更稳定因为不同特征的尺度差异太大会导致梯度更新方向偏向尺度大的特征。6.2 训练集、验证集、测试集的划分数据要分成三份训练集用来更新参数验证集用来调超参数和早停测试集用来最终评估。常见的比例是7:2:1或8:1:1。划分时要随机打乱确保每个集合的分布一致。如果数据量很小可以用K折交叉验证把数据分成K份每次用K-1份训练1份验证重复K次取平均。这样能更充分地利用数据但计算量会乘以K。一个容易犯的错误是用测试集调超参数。这样测试集就失去了“未见过的数据”的意义最终评估会偏乐观。正确的做法是验证集调参测试集只在最后用一次。6.3 数据增强的简单实现数据增强是通过对训练数据做随机变换来增加数据多样性从而减少过拟合。对于图像常见的增强有随机翻转、随机裁剪、随机旋转。从零实现的话可以在__getitem__里加随机变换。比如随机水平翻转以0.5的概率把图像左右翻转。实现时用np.fliplr就行。随机裁剪从图像中随机取一个区域然后resize回原大小。这些操作虽然简单但能显著提升模型的泛化能力。要注意的是验证集和测试集不能做数据增强因为评估时需要确定性的结果。数据增强只用在训练集上。7. 调试与性能优化手写代码的实战经验7.1 梯度消失与爆炸的排查梯度消失的表现是靠近输入的层梯度非常小参数几乎不更新。排查方法是打印每一层梯度的范数如果某一层的梯度范数接近0说明梯度消失了。解决方法包括用ReLU激活函数、加BatchNorm、用残差连接。梯度爆炸的表现是梯度范数非常大loss变成nan。排查方法是打印梯度范数如果超过1000基本就是爆炸了。解决方法是梯度裁剪把梯度范数限制在一个阈值内比如5.0。我自己的经验是先加梯度裁剪再调其他。梯度裁剪是一个“安全网”能防止训练崩溃而且几乎不影响正常训练。实现很简单算所有参数梯度的总范数如果超过阈值就按比例缩放。7.2 数值稳定性的常见陷阱数值稳定性问题在手写代码时特别常见。比如算Softmax时如果logits很大exp会溢出算Cross Entropy时如果预测概率接近0log会变成负无穷。解决方法都是“减去最大值”exp(x - max(x))这样指数不会溢出而且结果不变。另一个陷阱是除零错误。比如归一化时如果标准差是0就会除零。解决方法是加一个很小的eps比如1e-8。类似地Adam的分母也要加eps。还有一个容易忽略的点浮点数精度。float32的精度大约是7位有效数字如果累加很多小梯度可能会丢失精度。实践中可以用float64做梯度检查但训练时用float32就够了。7.3 从手写代码到框架的迁移思路手写一遍之后你会发现框架里的很多设计都变得“理所当然”了。比如PyTorch的autograd、nn.Module、optim你都能理解它们背后的逻辑。这时候再回去用框架效率会高很多因为你知道什么时候该用哪个组件出了问题也知道去哪里找。迁移的时候可以把手写代码里的模块和框架里的对应起来。比如你的Tensor对应torch.Tensor你的Linear对应nn.Linear你的SGD对应optim.SGD。这样一一对应之后框架就不再是黑盒了。我自己的做法是手写一遍之后用框架重写一遍同样的任务对比两者的输出。如果框架的输出和手写的输出一致说明你理解对了。如果不一致那就是某个细节没搞明白正好借这个机会查漏补缺。8. 从最小实现到实际项目下一步怎么走手写一遍核心模块之后你已经具备了“看穿”AI系统的能力。接下来可以往几个方向扩展一是加更多层和更多操作比如卷积、池化、BatchNorm二是优化性能比如用C扩展或CUDA加速三是做端到端的项目比如图像分类、文本分类、简单的序列预测。我个人的建议是先做一个完整的端到端项目从数据加载到模型训练到推理部署全部用手写代码实现。这个过程中你会遇到很多“框架帮你做了但你没意识到”的事情比如参数初始化、梯度清零、设备管理。把这些都自己处理一遍你对AI工程的理解就完整了。最后分享一个小技巧把训练过程可视化。用matplotlib画loss曲线和准确率曲线能直观地看到模型是在收敛还是过拟合。如果loss曲线震荡很大说明学习率太大如果训练loss下降但验证loss上升说明过拟合了。这些直观的信号比看数字快得多。
网站建设高端定制企业官网
RELATED

相关资讯

更多精彩内容,欢迎继续阅读

较早相关资讯

最新相关资讯

CLI-Anything:用一份配置文件生成统一、规范、可补全的命令行工具 2026/9/28 14:49:50

CLI-Anything:用一份配置文件生成统一、规范、可补全的命令行工具

平时和 CLI 打交道最多的人,大概都有过这种别扭:某个 API 调试得很好,换个环境又得重敲一遍;有个内部脚本只有自己会用,交给同事要写一页说明文档。今天聊的这个项目叫 CLI-Anything,它做的就是把这堆散落的…

阅读更多 →
水表数字识别实战:从表盘定位到读数校验的完整流程 2026/9/28 14:49:50

水表数字识别实战:从表盘定位到读数校验的完整流程

简介:这份资源面向计算机视觉入门者与图像处理学习者,聚焦水表刻度与数字的自动识别场景,提供一套基于OpenCV的完整实践方案。包内共13个文件,以7张jpg水表样本图、4个xml配置与模板文件、1个iml工程配置及1个py主程序为主&#x…

阅读更多 →
AIO Sandbox:基于WASM的本地AI Agent开发沙箱 2026/9/28 14:49:50

AIO Sandbox:基于WASM的本地AI Agent开发沙箱

1. 这不是“又一个沙箱”,而是把开发环境塞进沙箱的逆向工程你有没有试过这样一种场景:刚写完一段 Python 脚本,想立刻在干净环境里跑一下,但又不想开虚拟机——太重;也不想用 Docker 手动配镜像——太碎;更…

阅读更多 →
浓度迁移与损伤方程:多物理场耦合建模与实战要点 2026/9/28 14:49:50

浓度迁移与损伤方程:多物理场耦合建模与实战要点

我最初接触浓度迁移与损伤方程,并不是为了写论文,而是因为在一次混凝土耐久性评估项目中,现场吃了个“哑巴亏”:一组桩基在服役五年后出现网状裂纹,检测报告把原因写成统一的“材料劣化”,但如果我们只做单…

阅读更多 →
一文搞懂Android布局裁剪:从clipChildren到Compose的机制与实战 2026/9/28 14:49:50

一文搞懂Android布局裁剪:从clipChildren到Compose的机制与实战

做 Android 开发的人,多少都遇到过这种诡异情况:子 View 的坐标、尺寸在布局预览里一切正常,一跑到真机上却被削掉一半;或者你把android:clipChildren"false"写上去,内容确实画出去了,结果向上再…

阅读更多 →
nRF Connect协议栈调试指南:从BLE底层到GATT实战 2026/9/28 14:49:37

nRF Connect协议栈调试指南:从BLE底层到GATT实战

1. 项目概述:为什么nRF Connect不是“另一个蓝牙APP”,而是BLE工程师的瑞士军刀你手边是不是正摆着一块nRF52832开发板,或者刚焊好一颗STM32WBA65芯片,却卡在“设备搜不到”“连上了但读不出服务”“Characteristic写不进去还报0x…

阅读更多 →

今日资讯

本周资讯

本月资讯

看完文章仍有疑问?

联系尧图顾问,获取一对一建站咨询

立即免费咨询 📞 400-888-8888
📞 ✉