AI 工程地基 12|张量运算:报错的形状好修,不报错的才贵

「AI 工程地基」系列第 12 篇。把 GitHub 58k star 的开源课 ai-engineering-from-scratch(MIT 协议)一课一课啃:全文翻译落盘、公式逐条审计、再用大白话讲一遍。今天讲张量。

AI 工程地基 12|张量运算:报错的形状好修,不报错的才贵

「AI 工程地基」系列第 12 篇。把 GitHub 58k star 的开源课 ai-engineering-from-scratch(MIT 协议)一课一课啃:全文翻译落盘、公式逐条审计、再用大白话讲一遍。今天讲张量。

你搭了一个 transformer,一跑就崩:mat1 and mat2 shapes cannot be multiplied (32x768 and 512x768)。你转置一下,报错变成 Expected 4D input (got 3D)。你补一个 unsqueeze,别处又炸了。这不是你笨。一个 transformer 里有几十个 reshape、transpose、广播串在一起,一个轴搞错就级联。课程开头那句话说得更狠:有些形状错误根本不报错,它们安静地沿错误的维度广播、对错误的轴求和,产出一坨形状全对的垃圾。这篇就讲清楚这套东西,顺便实测一下「不报错」到底有多安静。

先说张量是什么。多维数组,加一条纪律:类型统一。维数叫 rank,每一维叫一个轴,shape 是把每根轴的大小排成的一个元组。一批 32 张 224×224 的彩色图,就是 (32, 3, 224, 224)——四根轴:批、通道、高、宽。12 个头的注意力也是四根轴:(批, 头, 序列长, 每头维度)。这里有个词要当场拆雷:张量的 rank 和线性代数里的矩阵秩不是一回事。矩阵 [[1,2],[2,4]] 的秩是 1(第二行是第一行的两倍,两行线性相关),但作为数组它的 rank 是 2(两根轴)。课程术语表自己点破了这一条,这是教材少见的诚实。

再说内存。数组在内存里永远是一条直线。strides(步长)告诉你:沿某根轴走一格,要在这条直线上跳过几个元素。(3, 4) 的矩阵,步长是 (4, 1)——下一行跳 4 个元素,下一列跳 1 个。这一下就解释了转置为什么便宜:它不搬数据,只交换步长,让同一块内存换一种读法。也解释了 NCHW 和 NHWC 的战争:PyTorch 默认通道在前,TensorFlow 默认通道在后,布局不合不仅慢,还常常是静默地慢。

然后是广播,形状不同的张量怎么相加。规则三条:从右边对齐;每对维度要么相等、要么有一个是 1;维数少的一方在左边补 1。比如 (8, 1, 6, 1) 加 (7, 1, 5),后者补成 (1, 7, 1, 5),结果是 (8, 7, 6, 5)。所有为 1 的轴被拉长去配合对面。这个机制让偏置向量 (D,) 能直接加到 (B, T, D) 的批上,一次搞定。但注意它的适用范围:广播对形状合同的要求是「兼容」,不是「正确」。(B, 1) 减 (B,),完全合法,结果是 (B, B)——一个 B×B 的矩阵。算交叉熵时标签差一个维度,损失就悄悄变成 B×B,没有报错,没有警告,训练照跑。

最后是 einsum,爱因斯坦求和。它把每根轴贴一个字母,规则一句话:出现在输入、没出现在输出里的字母,被求和。i,i-> 是点积(i 被求和,剩空);i,j->ij 是外积(没有字母消失);ij,jk->ik 是矩阵乘(k 消失);bhtd,bhsd->bhts 就是注意力分数(d 消失,Q 乘 K)。多头注意力整条链可以用 einsum 一路写下来,七步形状变化全在掌握:输入 (2, 8, 64),投影后分头成 (2, 4, 8, 16),算分数成 (2, 4, 8, 8),softmax 后乘 V 回到 (2, 4, 8, 16),合头 (2, 8, 64),输出投影。每一步的形状都能在纸上先写出来——这就是课程说的「shape tracker」:先把合同列在纸上,再写代码。

成本也有个干净的算法:一次 einsum 收缩的计算量,是所有下标大小的乘积。bij,bjk->bik 在 B=32、I=128、J=64、K=128 时,是 32×128×64×128 = 33,554,432 次乘加。这个数字我核了,对。

过一遍海关

这门课的代码全部跑通,数字逐条核对无一错。但有三处要拦下来。

第一处,报错信息让你调一个不存在的方法。自实现 Tensor 类的加法在形状不匹配时报:Use broadcast() first.,可是整个类里根本没有 broadcast 方法,hasattr 查过,False。更尴尬的是,正文明确教「偏置 (D,) 加批 (B,T,D) 要先 unsqueeze 成 (1,1,D)」,实测 (2,2,3) 加 (1,1,3)——照样报 Shape mismatch。也就是说,正文教的标准操作,在课程自己造的类上执行不了。这个能力是练习 2 的作业:你得自己写 broadcast_to,错误信息里那个函数名才会真的存在。报错指向幻影 API,头一回见。

第二处,「转置不搬数据」这句宣称,与课程自己的实现相反。正文说 transpose 只交换 strides、不移动数据——这在 NumPy 里是真的,课程的 demo 也确实用 NumPy 演示了步长从 (24,8) 变 (8,24)。但 Build It 一节教学生手写的那个 Tensor 类,permute 的实现是分配一块新数组、逐元素搬运、步长按新形状重算。reshape 也一样,每次都复制一份数据。结果就是:scratch 类永远连续(contiguous),正文讲的「不连续张量」这个概念在学生自己造的类里根本不存在。形状语义模拟得对,内存机制是另一套。你没法用自己写的类,验证正文最重要的宣称之一。

第三处,「广播不复制数据」的口径。这句话对输入成立:a * b 时 a 和 b 都没有被真的复制。但中间结果是全额物化的。课程教的广播版两两距离,(M,2) 对 (N,2) 相减,会真实生成一个 (M, N, 2) 的中间数组。M=N=2000 实测,中间数组 64 MB,而最终输出 (M,N) 只有 32 MB——中间过程是结果的两倍大。M=N=10000 外推,中间 1.6 GB,输出 0.8 GB。scipy 的 cdist 出同样的答案,内存只要 O(M+N)。广播省的是输入的复制,不是峰值内存,这两个概念在正文的「without copying data」里被压在了一起。

小账两笔:正文 NLP 示例用 (16, 128, 768),配套 demo 里跑的是 (16, 512, 768),口径不一但都合法;第 7 步的代码片段用了没定义的 K(完整代码里有,节选时丢了)。

口号审计。这门课的开场白是「张量是数据与深度学习之间的通用语言」。对,但藏着半句:这门语言只有形状语法,没有语义词汇。(32, 3, 224, 224) 不携带「这是一只猫」的任何信息,哪根轴是通道、哪根是空间,靠的是 NCHW 还是 NHWC 这种约定。数据进张量的那一刻,语义留在门外。这正是形状错误成为深度学习第一日常 bug 的深层原因:轴的语义是整套系统里唯一不自动验证的合同,全靠写代码的人背。顺带,正文说「形状错误是最常见的 bug」没有给出处,属经验断言;但它紧接着那句「有些形状错误不报错、静默产垃圾」才是全文最值钱的一句,我实测了:(B,1) - (B,) 得 (B,B),Python 警告调到 error 级别,依然零警告放行。

练习 2 是本课真正的毕业证。scratch 类出厂状态下,它的报错让你调用的那个函数还不存在;你把 broadcast_to 写完、加法改成自动广播,这条报错信息才第一次说了实话。测试用例也是现成的:(3, 1) 和 (1, 4) 相加,应该得到 (3, 4)。写之前先在纸上把广播的三条规则走一遍,就知道哪根轴会被拉长、数据要复制几份。

暂无表态

想参与讨论或点赞?登录后使用完整功能

讨论回复(0)

暂无回复,登录后可参与讨论
合作

智谱 GLM-5 已上线

在智谱开放平台 BigModel.cn 打造 AI 应用。新一代旗舰模型 GLM-5 在推理、代码、智能体综合能力达到开源模型 SOTA。

领取 2000万 Tokens