news 2026/10/1 5:13:32

从零手搓AI工程:深入张量、反向传播与部署的底层实践

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
从零手搓AI工程:深入张量、反向传播与部署的底层实践

1. 从零手搓AI工程:为什么我不建议你直接调包

很多人一听到“AI工程”这四个字,第一反应就是打开某个云平台,拖几个组件,调一下API,跑通了就觉得自己会了。我刚开始也这么干过,结果遇到模型输出不稳定、推理延迟忽高忽低、显存莫名其妙爆掉的时候,整个人是懵的——因为底层发生了什么,我完全不知道。ai-engineering-from-scratch这个方向之所以值得聊,恰恰是因为它逼着你把每一层都掀开看一遍:数据怎么进、张量怎么算、梯度怎么回传、模型怎么存、服务怎么起。这些东西你调包的时候永远学不到,但一旦线上出问题,能救你的只有这些。

这篇文章适合两类人:一类是刚入行、想真正搞懂AI系统内部运转逻辑的工程师;另一类是用惯了高层框架、但总觉得心里没底、想补一补底层功课的老手。我会从最朴素的矩阵运算开始,一路讲到怎么把一个能跑的小模型部署成可用的服务,中间穿插我自己踩过的坑和那些文档里不会写的细节。全程不依赖任何“一键式”平台,所有东西都自己动手搭。

需要先说明一点:从零构建不等于什么都自己造轮子。我的原则是——核心链路必须自己写一遍,外围工具可以用现成的。比如张量库我可以用NumPy,但反向传播我必须自己推;优化器我可以用现成的,但损失函数我必须自己实现一遍。这样既不会陷入无意义的重复劳动,又能保证对关键环节有真正的掌控力。

2. 先把“张量”这件事说透,不然后面全是坑

2.1 张量不是多维数组那么简单

几乎所有教程都会告诉你“张量就是多维数组”,这话对,但只说了一半。在实际工程里,张量至少还包含三个隐含信息:形状(shape)、数据类型(dtype)、存储布局(strides)。你调包的时候这三个东西框架帮你管了,自己写的时候如果不管,就会出现“明明形状对但结果就是错”的诡异情况。

我举个真实例子。早期我手写一个简单的全连接层,输入是(batch, features),权重是(features, out),做矩阵乘法得到(batch, out)。看起来没问题对吧?但我当时把权重存成了(out, features),然后直接乘,结果形状居然也能对上(因为batch和out恰好相等),但数值全错。这种bug在调包时几乎不可能出现,因为框架会帮你检查,但自己写的时候,形状对不代表语义对。

所以我的建议是:自己实现张量类的时候,至少要把shape和dtype作为强制校验项。每次运算前先断言形状匹配,宁可多写几行检查代码,也不要让错误悄悄传播。

2.2 广播机制:方便但危险

广播(broadcasting)是张量运算里最容易被滥用的特性。它让不同形状的张量能自动扩展成兼容形状,写起来很爽,但也是bug重灾区。比如你想把一个(batch, 1)的偏置加到(batch, features)的输出上,广播会自动把偏置扩展成(batch, features),这没问题。但如果你不小心把偏置写成了(features,),广播会把它当成(1, features)再扩展,结果也对——可一旦batch维度顺序变了,就会静默出错。

我的做法是:在关键运算处显式reshape,不依赖隐式广播。多写一个reshape调用,换来的是可预测的行为。性能上损失可以忽略,但调试时间能省下好几个小时。

2.3 内存布局对性能的影响

自己实现张量时,很多人用嵌套列表或者连续的一维数组加索引计算。前者简单但极慢,后者快但容易写错索引。我实测下来,用一维数组加strides的方式,在矩阵乘法上比嵌套列表快20倍以上。具体做法是:数据存成一维data数组,额外维护shape和strides,通过offset = sum(i_k * stride_k)来定位元素。

这个设计还有一个好处:转置操作不需要真正移动数据,只需要交换对应的shape和strides即可。对于大矩阵来说,这能省下大量内存拷贝时间。当然代价是索引计算变复杂了,但这是值得的——AI工程里内存带宽往往比计算本身更瓶颈。

3. 反向传播:手推一遍胜过看一百篇教程

3.1 计算图到底该怎么建

反向传播的核心是链式法则,但工程实现上关键在“计算图”怎么组织。有两种主流做法:一种是定义时构建静态图(像早期TensorFlow),另一种是运行时动态记录(像PyTorch)。从零实现的话,我强烈建议用动态方式,因为调试友好得多。

具体来说,每个张量除了存数据,还要存三样东西:grad(梯度)、_backward(反向函数)、_prev(前驱节点集合)。每次做运算时,不仅算出结果,还要定义好“当梯度从结果传回来时,我该怎么把梯度分给输入”。这就是一个局部反向函数。

我踩过的一个坑是:忘记在反向函数里处理“一个张量被多次使用”的情况。比如y = x * x,x在前向里用了两次,反向时梯度应该累加两次。如果只写一次,梯度就少了一半。这个bug非常隐蔽,因为前向结果完全正确,只有训练时loss下降变慢,很难定位。解决办法很简单:反向函数里一律用+=而不是=。

3.2 梯度检查:自己写的反向到底对不对

手写反向传播最大的风险是公式推错。我的经验是:每实现一个新算子,立刻用数值梯度做校验。方法很朴素——对某个输入分量加一个极小扰动ε,算前向输出的变化量,除以ε得到数值梯度,再和自己写的解析梯度对比。如果相对误差在1e-5以内,基本可以认为正确。

这个检查过程很枯燥,但能帮你省下后面几天的调试时间。我一般会写一个小工具函数,传入任意前向函数和输入,自动跑一遍数值校验。这个投入绝对值得。

3.3 常见算子的反向公式与实现要点

这里列几个最核心的算子,以及我在实现时特别注意的点:

算子前向反向要点
加法z = x + y梯度直接传回,注意广播时要sum回原形状
乘法z = x * y对x的梯度是grady,对y是gradx
矩阵乘z = x @ y对x是grad @ y.T,对y是x.T @ grad
ReLUz = max(0, x)梯度在x>0处传回,否则为0
Softmaxz = exp(x)/sum(exp(x))需要处理数值稳定性,先减最大值

广播的反向特别容易错。如果前向时x被广播了,反向时梯度必须沿着被广播的维度求和,还原成x的原始形状。我见过不少人在这里出错,导致梯度形状和参数形状不匹配,训练直接崩掉。

4. 训练循环:那些让loss不下降的隐形杀手

4.1 参数初始化不是随便填个数

很多人从零实现时,参数直接用random.randn初始化,然后发现loss根本不降。问题往往出在初始化尺度上。如果权重太大,前向输出会爆炸;太小,信号会逐层衰减。对于ReLU网络,我一般用He初始化:标准差取sqrt(2 / fan_in),其中fan_in是输入维度。对于tanh网络,用Xavier初始化:sqrt(1 / fan_in)。

这个细节在调包时框架帮你做了,但自己写的时候如果忽略,训练根本跑不起来。我建议把初始化单独写成一个函数,针对不同激活函数用不同策略,并且在前向传播时打印每层输出的均值和标准差,确认信号没有异常放大或缩小。

4.2 学习率:先找范围再精调

学习率设错是loss不下降的第二大原因。我的做法是:先跑一个学习率范围测试——从1e-6开始,每个batch乘以1.1倍,记录loss变化,画出loss vs lr的曲线。loss下降最快的区间就是合理范围,通常取该区间中点的十分之一作为初始学习率。

这个技巧来自fast.ai的实践,我用了之后再也没有盲目试过学习率。具体实现上,只需要在训练循环里动态调整优化器的lr参数即可,不需要改模型代码。

4.3 梯度裁剪:防止训练中途崩掉

即使学习率合适,训练过程中也可能遇到某个batch梯度特别大,导致参数一步跳飞,loss变成NaN。解决办法是梯度裁剪:把所有参数的梯度拼成一个向量,算它的L2范数,如果超过阈值就整体缩放。阈值一般取1.0到5.0之间,我常用1.0。

这个操作在RNN类模型里几乎是必须的,在深层网络里也很有用。实现上就是在反向传播之后、优化器更新之前,加一段裁剪逻辑。注意要裁剪所有参数,不能只裁一部分。

4.4 一个最小可用的训练循环长什么样

把上面这些串起来,一个从零实现的训练循环大概是这样:

for epoch in range(num_epochs): for batch_x, batch_y in dataloader: # 前向 pred = model(batch_x) loss = cross_entropy(pred, batch_y) # 反向 model.zero_grad() loss.backward() # 梯度裁剪 clip_gradients(model.parameters(), max_norm=1.0) # 更新 optimizer.step() # 每个epoch打印验证集指标 val_loss = evaluate(model, val_loader) print(f"epoch {epoch}, val_loss {val_loss:.4f}")

看起来简单,但每一行背后都有上面说的那些细节。我建议初学者先把这段代码手敲一遍,不要复制,敲的过程中你会自然思考每个环节的必要性。

5. 从训练到部署:模型怎么变成服务

5.1 保存模型不是pickle一下就完事

训练完的模型要保存下来供后续使用。很多人直接pickle整个模型对象,但这有几个问题:一是pickle依赖代码结构,代码改了旧模型可能加载不了;二是pickle有安全风险,不要加载不可信来源的文件;三是跨语言部署时pickle完全没用。

我的做法是:只保存参数张量,用自定义的二进制格式或npz。同时保存一份模型结构描述(比如每层的类型和维度),加载时先按描述重建模型,再填入参数。这样即使代码重构了,只要结构描述兼容,旧参数依然能用。

5.2 推理优化:batch和cache的取舍

部署时第一个要决定的是batch size。训练时batch大一点能提高GPU利用率,但推理时如果请求是逐个来的,攒batch会引入延迟。我的经验是:在线服务用动态batch,设置一个很小的等待窗口(比如5毫秒),窗口内到的请求攒成一个batch一起推理。这样既能利用并行计算,又不会让单个请求等太久。

另一个优化点是KV cache。对于自回归生成类模型,每次生成一个token都要重新计算前面所有token的注意力,非常浪费。KV cache的思路是把已经算过的key和value存下来,下一步直接复用。这个优化能把生成速度提升好几倍,但代价是显存占用随序列长度线性增长。实际部署时要根据显存大小设置最大序列长度。

5.3 服务框架:自己写还是用现成的

如果只是内部测试,用Flask或FastAPI包一层就够了。但生产环境要考虑并发、超时、限流、监控这些东西。我的建议是:核心推理逻辑自己写,服务框架用成熟的。比如用FastAPI做HTTP接口,用uvicorn做ASGI服务器,推理部分自己控制。这样既保证了推理链路的可控性,又不用重复造服务治理的轮子。

接口设计上,我一般提供两个端点:一个同步接口,直接返回结果;一个异步接口,提交任务后返回task_id,后续轮询结果。同步接口适合短输入,异步接口适合长文本生成这类耗时操作。

5.4 监控:没有监控的部署等于裸奔

服务上线后必须监控几个核心指标:请求延迟(P50/P95/P99)、每秒请求数、错误率、显存占用、GPU利用率。这些指标能帮你快速定位问题——比如P99延迟突然升高,可能是某个长请求堵住了队列;显存持续增长,可能是KV cache没释放。

我一般用Prometheus采集指标,Grafana做面板。如果不想搭这套,至少也要把日志打好,记录每个请求的输入长度、输出长度、耗时。出问题时能回溯。

6. 那些让我熬夜的坑,你可以直接跳过

6.1 数值稳定性:exp和log的陷阱

自己实现softmax和交叉熵时,如果不做数值处理,很容易遇到inf或nan。原因是exp(x)在x稍大时就会溢出。解决办法是先减去最大值:exp(x - max(x)),这样最大指数是0,不会溢出。交叉熵同理,不要先算softmax再算log,而是直接用log_softmax,把两步合并成一个数值稳定的操作。

这个坑我在第一次手写分类器时踩过,当时输入数据没做归一化,某个特征值到了100多,exp直接溢出,loss变成nan,查了半天才发现是数值问题。

6.2 内存泄漏:计算图没释放

动态计算图的一个副作用是:如果你把每次前向的结果都存到一个列表里(比如为了记录训练过程),整个计算图就不会被释放,显存会持续增长直到爆掉。解决办法是:记录指标时只存标量值,不要存张量。如果确实需要存中间结果,用.detach()断开计算图,或者转成NumPy数组。

我在早期实现时习惯把每个batch的loss存下来画曲线,结果存的是张量,跑了几百个batch后显存就满了。后来改成loss.item()存浮点数,问题解决。

6.3 随机种子:可复现性不是小事

做实验时如果不固定随机种子,两次运行结果不一样,你根本分不清是改动有效还是随机波动。我的做法是:在程序入口处固定所有随机源,包括Python的random、NumPy的random、以及自己实现的任何随机初始化。同时记录下种子值,方便复现。

但要注意,固定种子不代表结果完全确定——GPU上的某些操作本身有非确定性,多线程数据加载也可能引入随机性。对于要求严格复现的场景,需要额外设置环境变量禁用这些非确定性行为。

6.4 数据类型:float32和float16的取舍

训练时一般用float32,精度够用。但推理时可以用float16甚至int8来加速和省显存。不过float16有范围限制,太小或太大的值会溢出。我的经验是:推理时用float16,但要在关键位置做数值检查,比如softmax之前确认输入没有极端值。如果发现溢出,可以混合使用——大部分层用float16,少数敏感层保持float32。

int8量化更激进,能进一步压缩模型,但需要校准过程,而且不是所有模型都适合。我一般先试float16,不够再考虑int8。

7. 从零构建之后,我看调包的心态变了

把这一整套走下来之后,我再用那些高层框架,感觉完全不一样了。以前看到model.fit()就觉得万事大吉,现在会去想:它内部怎么做梯度累积?怎么处理变长序列?显存不够时怎么切分?这些问题在框架文档里往往一笔带过,但自己实现过之后,看源码就能快速定位关键逻辑。

更重要的是,遇到问题时排查思路变了。以前是“换个参数试试”,现在是“先确认梯度有没有传对、数值有没有溢出、内存有没有泄漏”。这种从原理出发的排查方式,效率比盲目试错高得多。

如果你也想走一遍这条路,我的建议是:不要追求功能完整,追求链路打通。先实现一个最简单的线性回归,从数据生成、前向、反向、更新到预测,全部自己写。跑通之后再逐步加激活函数、加层数、换损失函数。每加一个东西,就做一次梯度检查。这样一步步来,两周左右就能建立起完整的直觉。

后续如果想继续深入,可以往两个方向扩展:一是性能方向,比如用SIMD指令加速矩阵运算、用多线程做数据加载;二是规模方向,比如实现简单的分布式训练、梯度累积、混合精度。这两个方向都需要前面这些基础作为支撑,跳过基础直接搞高级特性,最后还是会回来补课。

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

Codex 接入国产大模型:config.toml 配置与 401 报错排查指南

1. 为什么要把 Codex 接到国产大模型上Codex 这个工具刚火起来那阵子,我身边不少朋友第一反应就是去官网下载、装桌面版、登录账号,然后发现要么卡在验证环节,要么用起来成本不低。我自己也是折腾了好几轮,从最初的codex安装、cod…

作者头像 李华
网站建设 2026/10/1 5:11:36

SAP S/4HANA Cloud数据权限详解:Maintain Restrictions UI实战指南

做SAP S/4HANA Cloud项目的人,迟早会遇到一个名字看起来有点商务、用起来却很权限的应用:Maintain Restrictions UI。我第一次打开这个App的时候,以为它又是一套“主数据维护”界面,翻了一圈发现里面全是Restriction Type、Assign…

作者头像 李华
网站建设 2026/10/1 5:11:32

SpringBoot集成MQTT:智能售货柜“取货即走”订单链路实战

去年在社区做智能售货柜试点的时候,我踩了不少坑才把整套链路跑通。这个项目名字听起来挺玄乎——"取货即走",说白了就是用户扫码开门、拿走商品、关门自动扣款,全程没有扫码支付这一步。后台的核心逻辑全靠SpringBoot搭的服务端&a…

作者头像 李华
网站建设 2026/10/1 5:07:28

环形子数组的最大和:Kadane算法与边界处理全解析

环形子数组的最大和,第一次在力扣上看到这道题的时候,我其实没太当回事。毕竟“最大子数组和”几乎是动态规划入门必刷题,换个环形外壳能难到哪去?结果真被教做人了:普通版本的代码能一遍过,环形版本我连续…

作者头像 李华
网站建设 2026/10/1 5:05:29

Python流程控制彻底讲透:从if/else、循环到match case实战

刚帮一个刚入门 Python 的朋友排查了一段代码&#xff0c;问题很简单——他用if判断用户输入时写成了if 1 < age < 18&#xff0c;逻辑上完全没错&#xff0c;但在实际业务里&#xff0c;年龄小于 0 或者大于 120 的数据他却没有处理。其实这不算 Bug&#xff0c;而是典型…

作者头像 李华