如何编程计算Loss关于矩阵W的梯度?含ReLU与线性变换场景
Let’s work through this problem step by step—first deriving the gradient mathematically, then translating that into code. I’ll use NumPy for the implementation since it’s standard for matrix/vector operations in Python.
Mathematical Derivation (Chain Rule)
We need to compute dLoss/dW, which requires applying the chain rule starting from the final loss and working backwards through each layer:
Compute the gradient of Loss with respect to m
The loss isLoss = sum(m²) = m.T @ m. The derivative with respect to m is straightforward:dLoss/dm = 2 * mCompute the gradient of Loss with respect to y
Sincem = Z @ y, using the chain rule, the gradient of Loss with respect to y is the transpose of Z multiplied bydLoss/dm:dLoss/dy = Z.T @ dLoss_dmCompute the gradient of Loss with respect to W
First, recall thaty = ReLU(W @ x). The derivative of ReLU is an element-wise function:ReLU'(z) = 1 if z > 0, else 0(wherez = W @ x).Combining this with the chain rule, the gradient of Loss with respect to W is the outer product of
(dLoss/dy * ReLU'(W@x))andx:dLoss/dW = (dLoss_dy * relu_derivative(W @ x)) @ x.T
The element-wise multiplication (*) applies the ReLU derivative to each element ofW@x, then we take the outer product withxto get the matrix gradient matching W’s shape.
Code Implementation
Here’s a concrete example using NumPy:
import numpy as np def relu(z): return np.maximum(0, z) def relu_derivative(z): return np.where(z > 0, 1, 0) # Define sample inputs and weights x = np.random.randn(3, 1) # Input vector (3x1) W = np.random.randn(4, 3) # Weight matrix (4x3) Z = np.random.randn(2, 4) # Weight matrix (2x4) # Forward pass z = W @ x y = relu(z) m = Z @ y loss = np.sum(m ** 2) # Backward pass dLoss_dm = 2 * m dLoss_dy = Z.T @ dLoss_dm dLoss_dz = dLoss_dy * relu_derivative(z) dLoss_dW = dLoss_dz @ x.T print("Loss:", loss) print("dLoss/dW shape:", dLoss_dW.shape) # Should match W's shape (4x3)
Key Notes
- The shape of
dLoss_dWwill always match the shape of W, which is critical for updating weights during training. - The ReLU derivative is handled element-wise, which is why we use
np.whereto create a mask of 1s and 0s. - This approach extends naturally to batch inputs—you’d just adjust the dimensions to handle multiple x vectors at once (e.g., x becomes a matrix where each column is an input sample).
内容的提问来源于stack exchange,提问作者Matt

