You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何用NumPy高效计算全连接层f(XW)对权重W的批量导数

NumPy批量实现全连接层f(XW)对权重W的导数计算

维度约定

首先统一各变量的形状定义,避免广播对齐出错:

  • 输入批量矩阵X:形状(N, d_in),N为批量样本数,d_in为输入特征维度
  • 权重矩阵W:形状(d_in, d_out),d_out为全连接层输出特征维度
  • 预激活值Z = X @ W:形状(N, d_out)
  • 激活输出A = f(Z):形状和Z一致为(N, d_out)
  • 激活函数导数dA_dZ = f.derivative(Z):形状和Z一致为(N, d_out),每个元素对应该位置预激活值的激活导数值
  • 目标输出4维张量:形状为(N, d_out, d_in, d_out),第0维对应N个样本,剩余3个维度对应单样本下∂A/∂W的结构,索引顺序为[样本n, 输出位k, 输入位i, 权重列j],对应导数值∂A_{n,k}/∂W_{i,j}

实现原理

根据链式法则,单个样本下的导数值满足:
$$\frac{\partial A_k}{\partial W_{i,j}} = f'(Z_k) \cdot X_i \cdot \mathbb{I}(j=k)$$
其中$\mathbb{I}(j=k)$是指示函数,当j等于k时取1,否则取0,本质就是形状为(d_out, d_out)的单位矩阵。
我们可以通过NumPy的广播机制直接做向量化计算,不需要写任何Python层循环,计算效率和原生NumPy矩阵运算一致。

代码实现

import numpy as np

def calc_dW_derivative(X: np.ndarray, W: np.ndarray, f) -> np.ndarray:
    # 计算预激活值和激活导数
    Z = X @ W
    dA_dZ = f.derivative(Z)  # 形状(N, d_out)
    
    N, d_in = X.shape
    d_out = W.shape[1]
    
    # 维度调整适配广播规则
    dA_dZ_bc = dA_dZ.reshape(N, d_out, 1, 1)    # 形状(N, d_out, 1, 1)
    X_bc = X.reshape(N, 1, d_in, 1)             # 形状(N, 1, d_in, 1)
    eye_bc = np.eye(d_out).reshape(1, d_out, 1, d_out)  # 形状(1, d_out, 1, d_out)
    
    # 广播相乘得到最终4维导数张量
    dA_dW = dA_dZ_bc * X_bc * eye_bc
    return dA_dW

注意事项

  • 如果你是做常规反向传播更新权重,不需要显式构造这个4维张量:4维张量会占用N*d_out*d_in*d_out的内存,当维度大时内存开销极高。常规反向传播中,和上游梯度dL_dA(形状(N, d_out))做张量缩并后,可以直接得到损失对权重的梯度dL_dW = X.T @ (dA_dZ * dL_dA),形状为(d_in, d_out),内存开销和计算量都小几个数量级。
  • 上述实现严格匹配要求的输出结构,所有计算都在NumPy底层完成,没有逐元素循环,效率满足学习场景下的使用需求。

内容的提问来源于stack exchange,提问作者Ariel Yael

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.26 21:48:09