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

看似确定性的反向传播代码为何输出结果不一致?

反向传播输出不一致的原因及修复方案

你的代码输出结果不稳定的核心问题有两个:

  • 全局visited集合的错误设计:像c、d这类被多个节点依赖的中间节点,它们的_backward方法需要接收所有依赖节点传递的梯度,但全局visited会将节点标记为已访问,导致后续依赖节点的梯度无法继续向下传递。
  • 无序集合的遍历随机性:self._prev是通过元组转集合创建的,集合的遍历顺序在Python中是不确定的(即使3.7+版本字典保留插入顺序,集合的遍历顺序仍可能因哈希值变化而波动),这导致每次运行时节点处理顺序随机,部分梯度传递逻辑被跳过,最终输出不同结果。

比如:

  • 若f的子节点先遍历e再遍历g,c和d会被e的处理流程标记为已访问,g的梯度无法传递到a和b,得到a.grad=4.0;
  • 若遍历顺序是g先e后,e的梯度无法传递到a和b,得到错误的21.0(正确梯度应为25.0)。

修复后的代码

修复的核心是用拓扑排序替代全局visited,确保每个节点的_backward在所有依赖它的节点处理完成后只调用一次,同时避免全局状态污染:

class Value:
    def __init__(self, data, _children=None, _op=''):
        self.data = data
        self.grad = 0.0
        self._backward = lambda: None
        # 用元组保存子节点,保留创建顺序(拓扑排序不依赖此,但逻辑更清晰)
        self._prev = tuple(_children) if _children else ()
        self._op = _op

    def __add__(self, other):
        other = other if isinstance(other, Value) else Value(other)
        out = Value(self.data + other.data, (self, other), '+')

        def _backward():
            self.grad += 1.0 * out.grad
            other.grad += 1.0 * out.grad
        out._backward = _backward

        return out

    def __mul__(self, other):
        other = other if isinstance(other, Value) else Value(other)
        out = Value(self.data * other.data, (self, other), '*')

        def _backward():
            self.grad += other.data * out.grad
            other.grad += self.data * out.grad
        out._backward = _backward
        return out

    def backward(self):
        # 1. 构建拓扑排序:生成从输入到输出的节点顺序
        topo = []
        visited = set()
        def build_topo(node):
            if node not in visited:
                visited.add(node)
                for child in node._prev:
                    build_topo(child)
                topo.append(node)
        build_topo(self)

        # 2. 初始化输出节点梯度为1
        self.grad = 1.0
        # 3. 逆序遍历拓扑列表,依次调用_backward
        for node in reversed(topo):
            node._backward()

    def __repr__(self):
        return f"Value(data={self.data}, grad={self.grad})"

# 测试代码
a = Value(2.0)
b = Value(3.0)
c = a + b
d = a * b
e = c + d
g = c * d
f = e + g

f.backward()
print(f"a grad: {a.grad}")  # 稳定输出正确值25.0

关键修复点说明

  1. 移除全局visited,改用局部集合构建拓扑排序,确保每个节点只被加入拓扑列表一次,同时避免跨运行的状态污染;
  2. 拓扑排序保证了节点处理顺序为逆计算图方向,所有依赖当前节点的节点都已处理完毕后,再调用_backward累加梯度;
  3. 将_prev从集合改为元组,保留子节点的创建顺序,让计算图逻辑更清晰;
  4. 每次调用backward都会重新构建拓扑结构,无需手动重置梯度(初始化时grad默认为0.0)。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.13 06:31:07