# The chain rule: if y = f(g(x)), then dy/dx = f'(g(x)) * g'(x) # # For our network: Loss = MSE(sigmoid(W2 * sigmoid(W1 * X + b1) + b2), y) # We need: dLoss/dW2, dLoss/db2, dLoss/dW1, dLoss/db1 print("Chain rule breakdown:") print("dLoss/dW2 = dLoss/da2 * da2/dz2 * dz2/dW2") print(" where:") print(" dLoss/da2 = 2 * (predictions - targets) # MSE derivative") print(" da2/dz2 = sigmoid'(z2) # sigmoid derivative") print(" dz2/dW2 = a1 # linear layer derivative")