为什么反向传播是反向的
Why back propagation goes backward

原始链接: https://gregorygundersen.com/blog/2018/04/15/backprop/

提供的文本探讨了为什么反向传播必须在神经网络中向后而非向前移动,尽管链式法则在理论上允许两者皆可。 前向传播方法的核心问题在于计算效率低下。若要计算损失函数 $f$ 相对于某个权重 $\theta_i$ 的梯度,必须考虑网络中的每一条路径。前向传播要求为每个权重分别传播单独的导数,这会导致冗余计算,并随着网络规模的扩大,运算量呈指数级或二次方级爆炸。 相比之下,反向传播将该计算视为一个“信用分配”问题。通过向后传递导数,每个节点从其下游邻居处收集局部信息,从而能够利用多元链式法则计算其对最终梯度的贡献。这种反向消息传递有效地复用了已计算的偏导数供所有上游权重使用,将复杂度降低到相对于节点数量的线性时间。归根结底,反向传播是在有向无环图上应用链式法则的最优方式,它将原本计算上不可行的前向传播方法转化为一种高效且可扩展的算法。

Hacker News 最新 | 过往 | 评论 | 提问 | 展示 | 招聘 | 提交 登录 为什么反向传播是“反向”的 ( gregorygundersen.com ) 8 分 由 andsoitis 发布 1 小时前 | 隐藏 | 过往 | 收藏 | 讨论 帮助 指南 | 常见问题 | 列表 | API | 安全 | 法律 | 申请 YC | 联系 搜索:
相关文章

原文

The usual explanation of backpropagation (Rumelhart et al., 1986), the algorithm used to train neural networks, is that it is propagating errors for each node backwards. But when I first learned about the algorithm, I had a question that I could not find answered directly: why does it have to go backwards? A neural network is just a composite function, and we know how to compute the derivatives of composite functions using the chain rule. Why don’t we just compute the gradient in a forward pass? I found that answering this question strengthened my understanding of backprop.

I will assume the reader broadly understands neural networks and gradient descent and even has some familiarity with backprop. I’ll first setup backprop with some useful concepts and notation and then explain why a forward propagation algorithm is supoptimal.

Setup

Recall that the goal of backprop is to efficiently compute f/θi\partial f / \partial \theta_i

To be clear, the node vv refers to the output value of the node after passing the weighted sum of its inputs through an activation function σ\sigma, i.e.:

u=θ1t1+θ2t2++θntnv=σ(u) \begin{aligned} u &= \theta_1 t_1 + \theta_2 t_2 + \dots + \theta_n t_n \\ v &= \sigma(u) \end{aligned}

Note that in a typical diagram, uu, σ\sigma, and vv would all be a single node, denoted by the dashed line. In my mind, the most important observation needed to understand backprop is this: most of computing f/θ1\partial f / \partial \theta_1

fθ1=fvvuuθ1 \frac{\partial f}{\partial \theta_1} = \frac{\partial f}{\partial v} \frac{\partial v}{\partial u} \frac{\partial u}{\partial \theta_1}

We can compute v/u\partial v / \partial u analytically; it just depends on the definition of σ\sigma. And we know that u/θ1=t1\partial u / \partial \theta_1 = t_1

The challenge with computing f/v\partial f / \partial v is that downstream nodes depend on the value of vv. Thankfully, the multivariable chain rule has the answer. Given a multivariable function g(w1,w2,,wm)g(w_1, w_2, \dots, w_m)

gv=vg(w1(v),w2(v),,wm(v))=jgwjwjv \frac{\partial g}{\partial v} = \frac{\partial}{\partial v} g(w_1(v), w_2(v), \dots, w_m(v)) = \sum_{j} \frac{\partial g}{\partial w_j} \frac{\partial w_j}{\partial v}

So we can compute f/θi\partial f / \partial \theta_i

Repeated terms

We want a forward propagating algorithm that can compute the partial derivative f/θi\partial f / \partial \theta_i

fθi=fvvθi \frac{\partial f}{\partial \theta_i} = \frac{\partial f}{\partial v} \frac{\partial v}{\partial \theta_i}

Note that I’ve dropped the intermediate variable uu for ease of notation. To design our forward propagating algorithm, let’s formalize an important fact: in a directed computational graph in which node bb depends upon node aa, it is impossible to compute b/a\partial b / \partial a at any point before node bb:

This claim should be obvious. If our computational graph represents a function f(a)=bf(a) = b

In our setup, for every downstream node wjw_j

fθi=(jfwjwjvCompute on wj)vθiPass forward \frac{\partial f}{\partial \theta_i} = \Big( \sum_{j} \frac{\partial f}{\partial w_j} \underbrace{\frac{\partial w_j}{\partial v}}_{\text{Compute on $w_j$}} \Big) \overbrace{\frac{\partial v}{\partial \theta_i}}^{\text{Pass forward}}

We can see that such an algorithm blows up computationally because we’re forward propagating the same message many times over. For example, if we want to compute f/θi\partial f / \partial \theta_i

fθi=(j(kfzkzkwj)wjv)Repeated termsvθifθk=(j(kfzkzkwj)wjv)vθk \begin{aligned} \frac{\partial f}{\partial \theta_i} = \overbrace{ \Big( \sum_{j} \Big( \sum_{k} \frac{\partial f}{\partial z_k} \frac{\partial z_k}{\partial w_j} \Big) \frac{\partial w_j}{\partial v} \Big)}^{\text{Repeated terms}} \color{#11accd}{ \frac{\partial v}{\partial \theta_i} } \\ \frac{\partial f}{\partial \theta_k} = \Big( \sum_{j} \Big( \sum_{k} \frac{\partial f}{\partial z_k} \frac{\partial z_k}{\partial w_j} \Big) \frac{\partial w_j}{\partial v} \Big) \color{#bc2612}{ \frac{\partial v}{\partial \theta_k} } \end{aligned}

Here is a diagram of message passing the repeated terms:

I think the above diagram is the lynchpin in understanding why backprop goes backwards. This is the key insight: if we already had access to downstream terms, for example wj/v\partial w_j / \partial v

A backward pass

I hope this explanation it clarifies how you might get to backprop from first principles trying to compute derivatives in a directed acyclic graph. On a given node bb that depends on a node aa, we simply message pass b/a\partial b / \partial a back to aa. The multivariable chain rule helps prove the correctness of backprop. For any node vv with downstream weights wjw_j

fv=jfwjwjv \frac{\partial f}{\partial v} = \sum_{j} \frac{\partial f}{\partial w_j} \frac{\partial w_j}{\partial v}

Once you understand the main computational problem backprop solves, I think the standard explanation of backpropagating errors makes much more sense. This process is can be viewed as a solution to a kind of credit assignment problem: each node tells its upstream neighbors what they did wrong. But the reason the algorithm works this way is because a naive, forward propagating solution would have quadratic runtime in the number of nodes.

联系我们 contact @ memedata.com