桌上 先让它学会 目录

第 2 节 · 约 30 分钟

轴要有名字。(B, D) 右乘 (D, H)。

更危险的是对得上、语义错了:广播替你补维,程序照跑,损失在优化另一件事。

先讲

写代码前在纸上列三列:变量、shape、轴含义。第三列对不上,就不要跑。[32, 10] 可能是 32 个样本的 10 类 logits,也可能是长度 32 的序列上每个位置 10 维特征。机器只看长度。语义靠你写的名字守住。x.shape == [8, 20, 64] 若是文本 batch,三个轴依次是 batch、token、embedding。图像常写成 [B, C, H, W]。索引会消轴:x[0] 得到 [T, D]x[:, 0] 得到 [B, D]。冒号留下整条轴。读别人代码时,把注释里的轴名抄到纸上。

一份 batch X[B, D] 右乘 W[D, H],得到 Y[B, H]。中间的 D 必须相等:每个输出特征都读完整条输入。batch 轴只是把样本并排,层对每个样本用同一套 W。偏置 b[H] 加到 Y[B, H] 上,靠广播把 (H,) 补成每一样本一份。

x(4, 8)w(8,)x * w 得到 (4, 8):每一行被同一个向量缩放,看起来像线性层,其实还是 8 维。你要四个 logit,应写 x @ w,结果 (4,)。另一个静默事故:pred[4, 1] - target[4] 会变成 [4, 4],每个预测在跟所有标签比。对这个矩阵取均值,损失仍是标量,优化的却是两两比较。

广播从尾轴对齐。(4, 8)(8,) 合法;(4, 8)(4,) 会先把 (4,) 当成最后一轴。库里的线性层常把权重存成 (H, D),前向实际是 x @ W.T + b。你按 (D, H) 手写却又转置一次,结果变成 (B, D) 或直接报错。打印 W.shape 和一句「从几维映到几维」,比猜约定快。

参数量从 shape 直接读。两层 64→32→10 且都有偏置:64×32+32 + 32×10+10 = 2410。只算 64×10 会把中间那层当成不存在。把 [B, H, W, C] 直接 reshape 成 [B, C, H, W] 只在 C=1 时碰巧对,通道大于 1 会把相邻像素塞进通道轴,该用 permute。numpy 里 (2,3)*(3,4) 会报错;某些包装会先扩维,损失仍是标量。手算结果应是 [[1, 2, 3, 6], [4, 5, 6, 15]]W.grad 必须仍是 (3,4)

例:手算 (2,3) @ (3,4)
  1. X = [[1, 2, 3], [4, 5, 6]]
  2. W 四列:单位三列,再加全 1 的第四列
  3. 第一行点积:1, 2, 3, 1+2+3=6
  4. 第二行点积:4, 5, 6, 15
  5. 结果 (2, 4) = [[1, 2, 3, 6], [4, 5, 6, 15]]
写成 X * W,中间维对不上。有的环境报错,有的广播出一份你从未设计过的数组。

你来做

哪一步得到 (2, 4)?哪一步没报错,却不是接到 1 维的线性层?

X(2, 3)W(3, 4)

x(4, 8)w(8,)

留下 (2, 4) 那张表。下一节五个参数走一遍回传。