TODO 1.1:原生 PyTorch 实现

x_native = x.permute(0, 2, 3, 1).reshape(x.shape[0], x.shape[2] * x.shape[3], x.shape[1])

分步拆解:

  1. x.permute(0, 2, 3, 1)

    • 原始维度顺序:0=batch, 1=channels, 2=height, 3=width
    • permute(0, 2, 3, 1) 将维度重新排列为:0=batch, 2=height, 3=width, 1=channels
    • 此时张量形状变为 [batch_size, height, width, channels],即 NHWC 格式。
  2. .reshape(...)

    • 在上一步基础上,调用 reshape 将中间的两个空间维度 heightwidth 合并。
    • x.shape[0] 仍然是 batch_size
    • x.shape[2] * x.shape[3] 注意:这里的 x 还是原始张量(permute 之前),所以 x.shape[2]heightx.shape[3]width,乘积就是 H * W
    • x.shape[1] 是原始的 channels
    • 最终形状:[batch_size, H * W, channels]

为什么先 permute 再 reshape?
因为 reshape 会将张量在内存中按行优先顺序重新解释。如果我们直接在原始 NCHW 上 reshape 成 [B, H*W, C],会把通道和空间数据交错在一起,得到错误的结果。先通过 permute 把通道放到最后,再合并前两个空间维度,就能确保每个空间位置的通道向量保持完整,这正是 Transformer 所需要的"每个序列位置的特征向量"。


TODO 1.2:einops 实现

x_einops = einops.rearrange(x, 'b c h w -> b (h w) c')
  • einops.rearrange 是一种声明式的张量维度变换工具,通过字符串直观描述输入和输出的维度关系。
  • 模式 'b c h w -> b (h w) c' 的含义:
    • 左边 b c h w:输入张量的四个维度分别命名为 batch, channels, height, width
    • 右边 b (h w) c:输出张量的维度,其中 bc 直接对应;(h w) 表示将 hw 这两个维度**合并(flatten)**成一个维度,顺序是 h 在前,w 在后(行优先展平)。
  • 这步操作与上面原生方法完全等价,但代码更简洁、可读性更强,且不容易出错。

TODO 2.1:使用官方 nn.Embedding 进行前向传播

emb_layer = nn.Embedding(vocab_size, hidden_dim)
emb_layer.weight.data.normal_(0, 0.1)  # 随便初始化一下
out_official = emb_layer(input_ids)

逐行解释:

  1. nn.Embedding(vocab_size, hidden_dim)
    创建一个 PyTorch 官方提供的嵌入层对象。它在内部维护一个可训练的权重矩阵 weight,形状为 [vocab_size, hidden_dim],可以看作一个“词汇表 × 嵌入维度”的表格。

    • vocab_size:词汇表大小,即一共有多少个 token。
    • hidden_dim:每个 token 对应的嵌入向量的长度。
  2. emb_layer.weight.data.normal_(0, 0.1)
    用均值为 0、标准差为 0.1 的正态分布随机数来初始化这个权重矩阵。因为仅仅是演示,所以“随便初始化一下”,实际训练中也可以用其他初始化方式。

  3. out_official = emb_layer(input_ids)
    调用 nn.Embedding 的前向传播。它的底层操作可以简单理解为:

    • 接收输入的整数索引 input_ids,形状 [B, L]
    • 对于 input_ids 中的每一个索引值,去 weight 矩阵中取出该索引对应的那一行(一个 hidden_dim 维的向量)。
    • 最终返回一个形状为 [B, L, hidden_dim] 的张量,其中 out_official[b, l, :] 就是 token ID input_ids[b, l] 对应的嵌入向量。

本质上,nn.Embedding 就是一个可训练的查表操作。


TODO 2.2:用纯 PyTorch 张量索引模拟 Embedding 操作

out_manual = emb_layer.weight[input_ids]

解释:

  • 这一行直接利用了 PyTorch 的高级索引(Advanced Indexing)
  • emb_layer.weight 就是前面定义的嵌入层的权重矩阵,形状为 [vocab_size, hidden_dim]
  • input_ids 作为索引来访问 weight 矩阵:
    • input_ids 的形状是 [B, L],里面的元素都是整数,范围在 [0, vocab_size-1]
    • 索引规则:当索引是一个整数张量时,PyTorch 会逐个取出索引对应的行,并保持索引张量的维度结构。
    • 结果 out_manual 的形状就是 input_ids.shape + (hidden_dim,),即 [B, L, hidden_dim]

例如,如果 input_ids[0,0] = 42,则 out_manual[0,0,:] 就是权重矩阵的第 42 行,与 out_official[0,0,:] 完全一致。


TODO 3.1:实现前向传播

z = F.linear(x, weight, bias)
y = F.relu(z)
mask = (z > 0).float()
ctx.save_for_backward(x, weight, mask)
return y

逐行解释:

  1. z = F.linear(x, weight, bias)
    计算线性部分:z = x @ weight.T + biasF.linear 会自动处理加权和偏置,返回形状 [N, out_features] 的张量。

  2. y = F.relu(z)
    z 逐元素应用 ReLU:y = max(0, z),形状不变。

  3. mask = (z > 0).float()

    • (z > 0) 生成一个布尔型张量,z 中大于 0 的位置为 True,否则为 False
    • .float() 将其转为浮点型,True 变为 1.0False 变为 0.0
    • 这个 mask 记录了 ReLU 激活时哪些位置的值是正的,在反向传播时用来传递梯度(因为 ReLU 的导数为:正值处为 1,负值处为 0)。
  4. ctx.save_for_backward(x, weight, mask)
    将反向传播需要用到的中间变量保存到上下文 ctx 中。这里保存了:

    • x:计算对 weight 的梯度时需要。
    • weight:计算对 x 的梯度时需要。
    • mask:ReLU 反传时需要。
  5. return y
    返回最终的输出 y


TODO 3.2:反传过 ReLU

grad_z = grad_output * mask

解释:

  • grad_output 是从损失函数传回来的梯度,形状与 y 相同,即 [N, out_features]
  • ReLU 的导数定义为:
    [
    \frac{\partial \text{ReLU}(z)}{\partial z} =
    \begin{cases}
    1, & z > 0 \
    0, & z \le 0
    \end{cases}
    ]
  • 在前向传播时我们保存了 mask,它正好就是 ReLU 的导数(1 对应正位置,0 对应其他位置)。
  • 根据链式法则,损失对 z 的梯度 = 损失对 y 的梯度 × ReLU 的局部导数,即 grad_z = grad_output * mask
  • 这个操作等价于“将 ReLU 中负数位置的梯度截断为零,正数位置的梯度原样传递”。

TODO 3.3:反传过 Linear

grad_x = grad_z @ weight
grad_weight = grad_z.T @ x
grad_bias = grad_z.sum(dim=0)

现在 grad_z 是损失对线性层输出 z 的梯度,形状为 [N, out_features]。需要根据它分别求出对输入 x、权重 weight、偏置 bias 的梯度。

1. 对 x 的梯度 grad_x

  • 线性变换为 z = x @ weight.T + bias
  • 利用矩阵求导链式法则,对 x 的梯度为 grad_z @ weight
  • 验证形状:grad_z[N, out]weight[out, in],相乘得 [N, in],与输入 x 形状一致,正确。

2. 对 weight 的梯度 grad_weight

  • 损失对权重的梯度 = (grad_z)^T @ x
  • 形状计算:grad_z.T[out, N]x[N, in],相乘得 [out, in],与 weight 形状一致。
  • 注意:若考虑 batch 中的所有样本,梯度是各样本贡献之和,而矩阵乘法 grad_z.T @ x 正好实现了对 batch 维度的求和(因为 (out,N) @ (N,in) -> (out,in),内积在 N 上求和),所以这一行已经包含了所有样本的梯度累加。

3. 对 bias 的梯度 grad_bias

  • 偏置 bias 是形状 [out],对每个输出维度,偏置会加到 batch 中每一个样本上。因此其梯度是 grad_z 在 batch 维度(dim=0)求和。
  • grad_bias = grad_z.sum(dim=0) 得到形状 [out],正确。

最后返回三个梯度:(grad_x, grad_weight, grad_bias),分别对应 forward 的三个输入参数 (x, weight, bias),符合 PyTorch 自定义 Function 的规范。


Logo

Agent 垂直技术社区,欢迎活跃、内容共建。

更多推荐