news 2026/9/28 17:57:38

从零构建AI工程体系:深入张量运算与自动微分实现

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
从零构建AI工程体系:深入张量运算与自动微分实现

1. 从零搭建AI工程体系,为什么我劝你别一上来就啃论文

"ai-engineering-from-scratch"这个标题,第一次看到的时候我以为是又一个教人调包的教程。点进去翻了翻才发现,它想做的事情比调包大得多——从最底层的张量运算开始,一步步把AI工程里那些被框架封装得严严实实的东西重新拆开,让你看清楚每一层到底在干什么。

我做了七八年算法和工程相关的工作,带过不少新人,也面试过很多人。一个很普遍的现象是:很多人能用PyTorch跑通一个模型,但问他反向传播里梯度到底怎么传的、显存为什么突然爆了、混合精度训练为什么能省显存,答不上来。这不是人的问题,是学习路径的问题。大家都是从"跑通一个demo"开始的,框架把太多东西藏起来了,藏到你根本不知道它替你做了什么。

这个项目解决的就是这个问题。它适合那些已经会用框架、但总觉得心里没底的人,也适合刚入门、想从一开始就把地基打牢的人。核心思路很简单:不依赖高层框架,用最基础的数学工具和数据结构,把AI工程的关键环节一个个实现出来。从标量、向量、矩阵的运算开始,到自动微分、到简单的神经网络、到训练循环、到推理优化,每一层都自己动手写一遍。

我花了大概三周时间把这个项目的思路完整走了一遍,中间踩了不少坑,也重新理解了很多以前一知半解的东西。下面我把整个过程的思路、关键细节、实操步骤和踩坑经验完整分享出来。不管你是想系统补基础,还是想搞清楚AI工程到底在工程什么,应该都能有点收获。

2. 整体设计思路:为什么从张量开始,而不是从模型开始

2.1 自底向上的学习路径到底好在哪

大部分AI教程的路径是"自顶向下"的:先给你一个完整的模型代码,让你跑通,然后再慢慢解释里面的组件。这个路径的好处是反馈快,跑通了有成就感。但坏处也很明显——你对整个系统的理解是碎片化的,每个组件都知道一点,但连不起来。

"ai-engineering-from-scratch"走的是相反的路:自底向上。先实现最基础的数据结构(张量),再实现最核心的运算(前向传播、反向传播),然后组装成层,再组装成网络,最后加上训练循环和优化器。每一步都建立在前一步的基础上,你能清楚地看到每一层是怎么来的。

我个人的体会是,自底向上的路径在前期会慢一些,因为你要花时间理解那些框架帮你自动处理的东西。但一旦过了某个临界点,后面会越来越快,因为你对整个系统的理解是连贯的。遇到问题的时候,你知道该去哪个层面找原因,而不是盲目地试。

2.2 技术选型:为什么用Python加NumPy,而不是直接上PyTorch

这个项目的核心实现语言是Python,基础运算依赖NumPy。这个选择背后有几个考虑。

第一,NumPy足够底层,但又不会太底层。你不需要自己去管理内存分配和指针运算,但你需要自己实现矩阵乘法、广播机制、梯度计算。这个抽象层级刚好能让你理解AI工程的核心概念,又不会陷入系统编程的细节里。

第二,Python的生态让实验成本极低。你可以在Jupyter Notebook里一行行跑,随时打印中间结果,随时改代码看效果。这种即时反馈对理解复杂概念非常重要。

第三,不依赖自动微分框架,才能真正理解自动微分。PyTorch的autograd用起来太方便了,方便到你根本不需要知道它是怎么实现的。但如果你自己从零实现一个简单的自动微分引擎,你就会明白计算图是怎么构建的、梯度是怎么回传的、为什么需要保留中间变量。

当然,这个选择也有代价。纯NumPy实现的训练速度肯定比PyTorch慢很多,所以这个项目不适合用来训练大模型。它的定位是教学和理解,不是生产。如果你想做实际的项目,最终还是要用框架。但理解了底层之后再用框架,你的效率会完全不一样。

2.3 核心模块的拆解逻辑

整个项目可以拆成几个核心模块,每个模块解决一个特定的问题:

  • 张量模块:实现多维数组的基本运算,包括加减乘除、矩阵乘法、广播、reshape、transpose等。这是所有后续模块的基础。
  • 自动微分模块:实现计算图的构建和反向传播。这是整个项目最核心也最难的部分。
  • 神经网络模块:基于自动微分实现全连接层、激活函数、损失函数。
  • 优化器模块:实现SGD、Momentum、Adam等优化算法。
  • 训练循环模块:把前面的模块组装起来,实现完整的数据加载、前向传播、损失计算、反向传播、参数更新流程。
  • 推理优化模块:实现量化、剪枝等推理优化技术,理解模型部署时到底在优化什么。

这个拆解逻辑的好处是,每个模块的边界很清晰,你可以单独理解每个模块,也可以看到模块之间是怎么衔接的。我在实际操作的时候,就是按照这个顺序一个个实现的,每实现完一个模块就写几个测试用例验证,确保没问题再进入下一个。

3. 核心细节解析:张量实现和自动微分到底难在哪

3.1 张量类的设计:数据、形状、梯度三件套

张量类的设计是整个项目的地基。一个最基础的张量类需要包含三个核心属性:数据(data)、形状(shape)、梯度(grad)。数据用NumPy数组存储,形状描述数据的维度信息,梯度在反向传播时被填充。

这里有一个关键的设计决策:梯度是存储在张量对象上的,还是单独维护一个梯度表。两种方式各有优劣。存储在张量对象上的好处是直观,每个张量自己知道自己的梯度;坏处是内存占用会翻倍,因为每个参与运算的张量都要存一份梯度。单独维护梯度表的好处是内存更省,但实现起来更复杂,需要处理张量标识和梯度映射。

我选择的是存储在张量对象上,因为教学场景下直观比省内存更重要。实际生产中的框架(比如PyTorch)也是这么做的,每个tensor都有一个.grad属性。

另一个关键点是广播机制的实现。当你对一个形状为(3, 1)的张量和一个形状为(1, 4)的张量做加法时,结果应该是(3, 4)。这个机制在NumPy里是自动的,但如果你要自己实现自动微分,就必须手动处理广播带来的梯度回传问题。具体来说,前向传播时广播把小的张量"撑大"了,反向传播时梯度需要"缩回"原来的形状。这个"缩回"的操作需要对梯度在广播维度上求和。

我一开始在这里卡了很久,因为广播的维度对齐规则(从右往左对齐,不足的补1)和梯度缩回的规则(在广播维度上求和)需要完全对应,稍微搞错一个维度,梯度就会算错。后来我写了一个专门的测试用例,用数值梯度验证解析梯度,才把这个问题彻底解决。

3.2 自动微分的两种实现路径:数值微分 vs 计算图

自动微分有两种主流实现方式:数值微分和计算图。

数值微分利用导数的定义,通过微小扰动来计算梯度。实现简单,但计算量大(每个参数都要扰动一次),而且有精度问题。计算图则是把前向传播的每一步操作记录下来,形成一个图结构,反向传播时沿着图反向遍历,用链式法则计算梯度。

这个项目采用的是计算图方式,因为它是现代深度学习框架的标准做法。具体实现上,每个张量除了数据和形状,还要记录它是怎么来的——也就是它的"父节点"和"操作类型"。比如c = a + b,那么c的父节点就是a和b,操作类型是加法。反向传播时,从最终的损失值开始,沿着图反向遍历,每个节点根据自己的操作类型计算梯度并传给父节点。

这里有一个容易忽略的细节:计算图需要在每次前向传播后重置。因为每次迭代的数据不同,计算图也不同。如果不重置,图会越来越大,内存会爆掉。PyTorch里用.backward()之后梯度会累积,需要手动清零,也是类似的原因。

3.3 反向传播的实现细节:链式法则的工程化

反向传播的核心是链式法则,但工程实现上有几个关键细节。

第一,梯度的累积。当一个张量被多个下游节点使用时,它的梯度需要从多个路径分别回传并累加。比如a被b和c同时使用,那么a的梯度等于b回传的梯度加上c回传的梯度。这个累积操作在实现时容易漏掉,导致梯度算错。

第二,梯度的形状匹配。每个操作的反向传播函数需要确保回传的梯度形状和输入张量的形状一致。对于矩阵乘法、广播、reshape这些会改变形状的操作,需要特别小心。

第三,计算图的拓扑排序。反向传播需要按照正确的顺序遍历计算图,确保每个节点的梯度在它被使用之前已经计算完毕。通常的做法是先用拓扑排序得到一个线性序列,然后反向遍历这个序列。

我在实现的时候,一开始没有做拓扑排序,直接递归遍历,结果遇到了重复计算和循环依赖的问题。后来改成先拓扑排序再反向遍历,问题就解决了。这个经验让我意识到,计算图本质上是一个有向无环图(DAG),图算法的很多经典思路在这里都适用。

4. 实操过程:从零实现一个可训练的神经网络

4.1 环境准备与项目结构

实操的第一步是搭好环境。我用的Python 3.10,核心依赖只有NumPy和Matplotlib(用来画损失曲线)。不需要GPU,CPU就够跑这个教学项目。

项目结构我建议这样组织:

ai-eng-from-scratch/ ├── tensor.py # 张量类实现 ├── autograd.py # 自动微分引擎 ├── nn.py # 神经网络层和损失函数 ├── optim.py # 优化器 ├── train.py # 训练循环 ├── utils.py # 工具函数(数据加载、可视化等) └── tests/ # 测试用例

这个结构的好处是模块边界清晰,每个文件只负责一个核心功能。我在实际写的时候,每写完一个模块就写对应的测试,确保这个模块的行为符合预期,再进入下一个模块。这种"测试驱动"的方式在实现底层组件时特别有用,因为底层组件的bug会传播到上层,越晚发现越难排查。

4.2 张量类的完整实现要点

张量类的实现有几个关键方法需要仔细处理。

构造函数需要接收数据、是否需要梯度、父节点信息等参数。数据统一转成NumPy的float32数组,因为float32是深度学习中最常用的精度,兼顾速度和精度。

运算方法包括加、减、乘、除、矩阵乘法、幂运算等。每个运算方法都需要做两件事:计算前向结果,记录反向传播所需的信息。比如加法操作,前向结果是两个张量相加,反向传播时梯度直接回传(因为加法的导数是1)。

广播处理是难点。前向传播时,NumPy会自动广播,但反向传播时需要手动处理。我的做法是在前向传播时记录原始形状,反向传播时如果梯度形状和原始形状不一致,就在广播维度上求和。

矩阵乘法的反向传播需要用到转置。如果C = A @ B,那么A的梯度是C的梯度 @ B的转置,B的梯度是A的转置 @ C的梯度。这个公式推导起来简单,但实现时要注意维度匹配。

我在这里踩过一个坑:没有处理批量维度。实际训练时数据是分批的,所以张量通常是三维的(batch_size, seq_len, feature_dim)或者二维的(batch_size, feature_dim)。矩阵乘法的反向传播需要支持批量维度,不能只处理二维情况。后来我加了一个通用的批量矩阵乘法实现,才解决了这个问题。

4.3 自动微分引擎的构建过程

自动微分引擎的核心是一个反向传播函数,它接收一个张量(通常是损失值),然后沿着计算图反向遍历,计算每个节点的梯度。

实现步骤大致如下:

  1. 拓扑排序:从损失张量开始,深度优先遍历计算图,得到一个拓扑有序的节点列表。
  2. 初始化梯度:损失张量自身的梯度设为1(因为损失对自身的导数当然是1)。
  3. 反向遍历:按照拓扑排序的逆序,依次计算每个节点的梯度,并累加到其父节点上。
  4. 梯度清理:每次反向传播前,需要把所有张量的梯度清零,避免累积。

这里有一个性能优化的点:只对需要梯度的张量进行计算。如果一个张量的requires_grad为False,那么它的梯度不需要计算,它的父节点也不需要继续往上传播。这个优化在实现时可以通过在拓扑排序时过滤掉不需要梯度的节点来实现。

我在实现的时候,一开始没有做这个优化,结果发现即使只训练一个很小的网络,反向传播也要花好几秒。加上这个优化之后,速度快了很多。这个经验让我理解了为什么PyTorch里要区分requires_grad,以及为什么推理时要用torch.no_grad()。

4.4 神经网络层的组装与训练循环

有了张量和自动微分,神经网络层的实现就相对直接了。一个全连接层就是y = x @ W + b,其中W和b是需要学习的参数。激活函数(ReLU、Sigmoid、Tanh)都是逐元素的运算,实现起来也不复杂。

损失函数我实现了均方误差(MSE)和交叉熵(Cross Entropy)。交叉熵的实现需要注意数值稳定性,因为log(0)会变成负无穷。标准的做法是在log之前加一个很小的epsilon,或者用log_softmax的技巧。

训练循环的流程是:

  1. 从数据加载器取一个batch的数据。
  2. 前向传播计算预测值。
  3. 计算损失。
  4. 反向传播计算梯度。
  5. 优化器更新参数。
  6. 清零梯度,进入下一轮。

这个流程看起来简单,但实际实现时有很多细节。比如参数更新必须在梯度清零之前,否则梯度就丢了。再比如验证集上的评估需要关闭梯度计算,否则会浪费大量内存。

我用这个从零实现的框架训练了一个简单的两层全连接网络,在MNIST数据集上跑到了97%左右的准确率。虽然比不上PyTorch的99%+,但考虑到这是纯NumPy实现,而且没有做任何调参,这个结果已经能说明整个流程是正确的。

5. 常见问题与排查技巧实录

5.1 梯度爆炸和梯度消失怎么排查

梯度爆炸和梯度消失是训练神经网络时最常见的问题。梯度爆炸表现为损失突然变成NaN,梯度消失表现为损失几乎不下降。

排查梯度问题的第一步是打印梯度范数。在每次反向传播后,计算所有参数梯度的L2范数,如果范数在训练过程中持续增大,说明有梯度爆炸的风险;如果范数迅速趋近于0,说明有梯度消失。

梯度爆炸的常见解决方法是梯度裁剪,也就是把梯度的范数限制在一个阈值以内。实现起来很简单:如果梯度范数超过阈值,就按比例缩放梯度。

梯度消失的常见解决方法是换激活函数。Sigmoid和Tanh在输入较大或较小时梯度接近0,容易导致梯度消失。ReLU在正区间梯度恒为1,能有效缓解这个问题。另外,Batch Normalization也能通过规范化每层的输入分布来缓解梯度消失。

我在实现的时候,一开始用的是Sigmoid激活函数,结果发现训练非常慢,损失几乎不下降。换成ReLU之后,训练速度明显加快。这个对比让我直观地理解了激活函数选择的重要性。

5.2 数值稳定性问题的处理

数值稳定性是底层实现中很容易被忽略的问题。几个常见的坑:

log(0)问题:交叉熵损失里需要计算log(p),如果p为0,log(0)就是负无穷。解决方法是在log之前加一个很小的epsilon(比如1e-7),或者用log_softmax的技巧。

除零问题:归一化操作里需要除以标准差,如果标准差为0就会出问题。解决方法是加一个很小的epsilon。

溢出问题:指数运算容易溢出,比如exp(100)就是一个非常大的数。解决方法是在做softmax之前先减去最大值,这样指数运算的输入就不会太大。

这些问题在框架里都被自动处理了,所以很多人根本不知道它们存在。但当你自己从零实现的时候,这些问题就会一个个冒出来。我的建议是,在每个可能出问题的地方都加上数值稳定性的保护,宁可多写几行代码,也不要让训练莫名其妙地崩掉。

5.3 内存占用过高的优化思路

纯NumPy实现的一个大问题是内存占用高。因为每个中间结果都要存下来用于反向传播,所以显存(内存)占用会随着网络深度线性增长。

优化思路有几个:

及时释放不需要的中间变量:反向传播完成后,计算图就可以释放了。如果用的是Python的引用计数机制,把不再需要的变量设为None就能触发垃圾回收。

用in-place操作:有些操作可以原地进行,不需要分配新的内存。比如ReLU的反向传播可以直接在梯度数组上操作。

梯度检查点:这是一种用时间换空间的技术,只保存部分中间结果,其他的在反向传播时重新计算。这个技术在大模型训练里很常用,但在教学项目里实现起来比较复杂,可以先了解思路。

我在实现的时候,一开始没有注意内存问题,跑一个稍微大一点的网络就卡住了。后来加了及时释放和in-place操作,内存占用降了不少。这个经验让我理解了为什么PyTorch里有很多in-place操作的版本(比如ReLU的inplace参数)。

5.4 常见问题速查表

问题现象可能原因排查方法解决方案
损失变成NaN梯度爆炸或数值溢出打印梯度范数和中间值梯度裁剪、加epsilon、减最大值
损失不下降梯度消失或学习率太小打印梯度范数换ReLU、调大学习率、加BatchNorm
训练速度慢没有关闭不需要的梯度计算检查requires_grad设置推理时用no_grad、过滤不需要梯度的节点
内存占用高中间变量没有及时释放监控内存使用及时释放、in-place操作、梯度检查点
梯度形状不匹配广播或reshape处理错误打印梯度形状检查广播维度的梯度缩回逻辑
参数没有更新梯度清零在参数更新之前检查训练循环顺序先更新参数再清零梯度

6. 从零实现之后,我对AI工程的理解变了

走完这一整套从零实现的流程之后,我最大的感受是:AI工程的核心不是调参,而是对计算过程的理解和控制。

以前用框架的时候,遇到问题我第一反应是搜"XX报错怎么解决",然后试各种网上的方案。现在我会先想:这个问题出在哪个层面?是数据的问题、前向传播的问题、反向传播的问题、还是优化器的问题?有了这个分层思维,排查问题的效率高了很多。

另一个感受是,底层实现让你对性能优化有了更具体的认知。你知道每一步操作的计算量和内存占用,就知道该在哪里优化。比如你知道矩阵乘法是计算密集型的,就会考虑用更好的BLAS库;你知道中间变量是内存密集型的,就会考虑用in-place操作或者梯度检查点。

这个项目后续还可以往几个方向扩展。一个是加入卷积和池化操作,理解CNN的底层实现。另一个是加入序列模型,理解RNN和Attention的计算过程。还有一个是加入分布式训练,理解数据并行和模型并行的实现原理。每个方向都能让你对AI工程的理解更深一层。

最后分享一个小技巧:如果你也想走一遍这个从零实现的过程,不要追求一次写对。先写一个能跑通的版本,哪怕效率很低、代码很丑。跑通之后,再一步步优化。这个"先跑通再优化"的思路,比"一开始就追求完美"要高效得多。我在实现自动微分的时候,第一版代码只有几十行,跑得慢但结果是对的。后来在这个基础上不断优化,才变成了一个相对完整的实现。这个过程本身,就是最好的学习。

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/28 17:57:22

MFC嵌入WebView2:从IE控件迁移到本地网页与C++双向通信

如果你还在 MFC 工程里用 IWebBrowser2 那个老掉牙的 IE 内核控件去加载 HTML,我建议你认真看看 WebView2。过去两年我陆陆续续把几个维护中的 MFC 项目的网页模块从 IE 控件迁移到了 WebView2,本地网页嵌入这块的体验可以说完全是两个时代。这篇文章就从…

作者头像 李华
网站建设 2026/9/28 17:57:22

CLI-Anything:不是工具,而是CLI交付可靠性工程范式

1. CLI-Anything 是什么:一个被误读的命名陷阱与真实定位“CLI-Anything”这个名称一出来,很多人第一反应是——又一个想把所有命令行工具塞进一个壳里的“万能CLI聚合器”?比如像某些 CLI Hub 工具那样,靠 shell alias 脚本包装…

作者头像 李华
网站建设 2026/9/28 17:56:56

STM32F4上移植CanFestival CANOpen协议栈完整指南与踩坑实录

做嵌入式这几年,最头疼的事之一就是给项目加通信协议栈。很多场景下CAN总线只是用来发几个报文,自己写个简单协议也够用,但一旦碰上设备之间要互操作、要对接标准诊断工具,甚至要过认证,老老实实上CANOpen就是个绕不开…

作者头像 李华
网站建设 2026/9/28 17:56:54

STM32移植mbedtls实战:从熵源接入到TLS握手全链路

1. 项目概述:为什么在STM32上硬啃mbedtls不是“炫技”,而是刚需你手头那块STM32F407VGT6开发板,跑着FreeRTOS,串口吐着温湿度数据,Wi-Fi模块连着局域网——看起来一切正常。但只要它一接入公网,或者和手机A…

作者头像 李华
网站建设 2026/9/28 17:56:49

CLI-Anything:下一代语义化命令行智能体架构

1. CLI-Anything 是什么:一个被误读的“通用命令行智能体”概念CLI-Anything 这个名字乍一听像某个具体开源工具,比如像curl或jq那样装完就能用的二进制程序。但翻遍 GitHub、PyPI、主流技术社区和近期开发者讨论,它根本不是一个已发布的、可…

作者头像 李华
网站建设 2026/9/28 17:56:08

内网离线部署MonkeyOCRv2:Docker镜像构建与vLLM GPU调优实战

1. 为什么要在内网离线环境折腾 MonkeyOCRv2把 MonkeyOCRv2 部署到内网离线环境,这件事听起来像是"把大象装进冰箱",但真正动手之后你会发现,难点从来不是"装",而是"装完之后它能不能跑起来、跑得稳不稳…

作者头像 李华