从强化学习论文里的 segment_tree.py 倒推线段树原理

2 阅读16分钟

从强化学习论文里的 segment_tree.py 倒推线段树原理

引子:为什么我要花一周学线段树

最近在啃强化学习方向的论文,发现一个有意思的现象:大部分涉及优先经验回放(Prioritized Experience Replay, PER)的论文,以及各种按优先级排序、按权重采样的 RL 算法,底层几乎都用着同一份线段树实现——OpenAI baselines 仓库里的 segment_tree.py

它的典型场景是:

  • 优先经验回放(PER):每个样本带一个优先级权重,采样时按权重比例抽取。线段树提供 O(log n) 的按前缀和采样和 O(log n) 的优先级更新。
  • 优先级排序类算法:需要在动态变化的序列上反复做"区间聚合 + 单点更新"。
  • 各大 RL 框架:Stable-Baselines3、Ray RLlib 的 PER 实现都直接继承或参考了这份代码。

要读懂论文里的 PER 实现,必须先彻底吃透 segment_tree.py,整个学习路径分为 5 部分

  1. 阅读相关 RL 算法,发现不懂 PER;
  2. 参考博客学习一般原理(线段树 从入门到进阶(超清晰,简单易懂));
  3. C 语言单步调试运行博客中的算法;
  4. 利用 AI 将 segment_tree.py 文件按照博客代码风格转换为 c 文件;
  5. 对照 c 文件学习,运行 segment_tree.py 的例程;
  6. 回归 DDPGfD 算法(参考仓库:pg-is-all-you-need),理解 PER 在 RL 中具体实现。

---> 学习文件gitee仓库

第一站:参考博客学习线段树原理

参考资源

我主要参考了两个资料,互补着看:

线段树的本质

线段树的本质一句话能说清:把一个数组组织成一棵二叉树,每个节点管辖一段区间,把区间信息(和、最值等)预先算好存起来,从而把"区间查询"和"单点修改"都做到 O(log n)。

关键结构是:节点 i 的左儿子是 2i,右儿子是 2i+1,每个节点存自己管辖区间 [l, r] 上的聚合值。查询时把目标区间拆成若干个节点区间之和,修改时改叶子再回溯更新祖先。

三条核心认知

学完原理后,我得到三条最重要的认知,它们是后续读懂工业代码的基础:

  1. pushdown 是线段树的灵魂——所有"区间修改"类问题都靠懒标记下传;
  2. lazy 语义决定 pushdown 写法——加法标记和覆盖标记天差地别;
  3. 框架与信息分离——换信息只改合并方式,不动框架。这条是后面读懂 Python 工业代码 operation 泛化的思想基础。

带着这三条认知,我才有底气去读工业代码。但光看原理还不够,必须动手实现一遍。


第二站:C 语言单步调试运行博客中的算法

光看博客容易"一看就会,一写就废"。我把整个学习过程浓缩进一个 C 文件 segment_tree_all_in_one.c五段递进实现,每段都带详尽注释和对照测试输出,并在 main 里单步调试验证。这是我认为最有效的学法——一次写全五种变体,逼自己想清楚每一步在干什么。

第一部分:建树 + 单点修改 + 区间查询(求和)

最基础的形态。建树是自顶向下递归,到叶子存值,回溯时合并:

void build1(int i, int l, int r) {
    tree1[i].l = l; tree1[i].r = r;
    if (l == r) { tree1[i].sum = a1[l]; return; }      /* 叶子 */
    int mid = (l + r) >> 1;
    build1(i * 2, l, mid);
    build1(i * 2 + 1, mid + 1, r);
    tree1[i].sum = tree1[i*2].sum + tree1[i*2+1].sum;  /* 回溯合并 */
}

区间查询 search1 的核心是三条规则:

  1. 当前区间被完全包含 → 直接返回 sum
  2. 当前区间与查询完全不相干 → 返回 0;
  3. 否则递归查询有交集的子节点。

一个有意思的细节:第 2 条"完全不相干"判断在正常递归下其实是冗余的——递归守卫已经保证子节点一定有交集。它唯一会触发的场景是查询区间整体落在 [1,n] 之外。它更多是文档性的,对应"三规则"的完整叙述,但作为可执行代码本质是死代码。这种"为了教学完整而写的冗余"是值得初学者注意的。

第二部分:区间修改 + 单点查询(打标记,无 pushdown)

引入"标记"思想:区间修改时只给极大区间打标记不往下传,查询时从根到叶子累加标记。类似差分数组。

第三部分:区间修改 + 区间查询(带 pushdown 懒标记)★核心★

这是线段树最重要的模板。懒标记解决的核心矛盾是:区间修改时如果一路下传到叶子,单次修改就是 O(n),失去线段树的意义。

解法是——修改停在"极大完全包含区间",打个标记不下传,等以后真的要往下走时再下传

pushdown 三件事缺一不可:

void pushdown3(int i) {
    if (tree3[i].lazy) {
        /* 子节点 sum += lazy * (子区间长度) ← 区间和的变化量 */
        tree3[left].sum  += tree3[i].lazy * (tree3[left].r - tree3[left].l + 1);
        /* 子节点 lazy += lazy ← 标记累加(多次区间加会叠加) */
        tree3[left].lazy += tree3[i].lazy;
        /* 右儿子同理 */
        tree3[right].sum  += tree3[i].lazy * (tree3[right].r - tree3[right].l + 1);
        tree3[right].lazy += tree3[i].lazy;
        /* 清空当前节点的标记 */
        tree3[i].lazy = 0;
    }
}

调用时机最关键凡是往下递归之前(修改和查询都要)必须先 pushdown,否则子节点数据过时。这是初学者最容易漏的一步,我在代码里专门标注了 ★查询也要 pushdown!★

记住这句话:能不看代码默写出 add3 + pushdown3 + search3 就算掌握了线段树。

第四部分:区间赋值 + 区间查询

和第三部分的差别看似只是"加法变赋值",但 pushdown 时差别巨大:

加法标记(第三部分)覆盖标记(第四部分)
子节点 sum+= lazy * 长度= lazy * 长度(直接赋值)
子节点 lazy+= lazy(累加)= lazy(覆盖)
是否需要额外标志不需要(lazy=0 即无标记)需要 has_lazy(lazy=0 也可能是合法覆盖值)

这条经验很重要:写线段树前必须先想清楚 lazy 的语义,是"增量"还是"覆盖",直接决定 pushdown 怎么写。

第五部分:维护区间最大值 + 区间加

从求和换成求最大值,整个框架不变,只改两处:

  • 合并方式 +max
  • pushdown 时子节点 maxval += lazy不加长度,因为每个数都加了 lazy,最大值也加了 lazy)。

这说明线段树是框架与信息分离的——框架负责区间分解与标记下传,信息负责怎么合并。这个认知是后面读懂 Python 工业代码 operation 泛化的思想基础。

第二站小结

五段递进学完,对应第一站的三条核心认知都得到了代码验证:

  1. pushdown 是灵魂(第三部分);
  2. lazy 语义决定写法(第三 vs 第四部分);
  3. 框架与信息分离(第五部分)。

但这一切都是递归实现,和工业代码 segment_tree.py 还有距离。下一步就是把它翻译过来对照。


第三站:利用 AI 将 segment_tree.py 转换为 C 文件

工业代码长什么样

segment_tree.py 是 OpenAI baselines 里的实现,约 130 行,结构是"基类 + 两个特化子类":

SegmentTree (基类)
    泛化的"点修改 + 区间查询"引擎
    - __init__:      2*capacity 数组,要求 capacity  2 的幂
    - __setitem__:  迭代点修改(自底向上)
    - __getitem__:  O(1) 读叶子
    - operate:      区间查询入口
    - _operate_helper: 递归查询核心
  
  ├─ SumSegmentTree: operation=加法, 新增 sum()  retrieve()
  └─ MinSegmentTree: operation=min, 新增 min()

核心解读:利用 __setitem__ 建树的底层逻辑

segment_tree.py 时发现它的建树方式非常特别——没有独立的 build 函数,建树 = 调用 n 次 __setitem__。这背后是一系列精心的设计选择。

存储布局:完美二叉树 + 2*capacity 数组

segment_tree.py 要求 capacity 必须是 2 的幂(assert capacity & (capacity - 1) == 0),数组大小为 2 * capacity。以 capacity=4 为例,布局如下:

下标:  0    1      2        3        4    5    6    7
       空   根    左子树    右子树    叶0  叶1  叶2  叶3
                        ┌─┴─┐    ┌─┴─┐
                       叶0+叶1  叶2+叶3

完美二叉树带来三个关键性质

  1. 叶子节点位置确定:叶子统一从下标 capacity 开始,到 2*capacity-1 结束。外部下标 i(0-based)的叶子 → 内部下标 capacity + i
  2. 父子关系靠算术确定:节点 k 的父 = k // 2,左儿子 = 2*k,右儿子 = 2*k + 1。不需要在节点里存 l, r 字段,省了一半内存。
  3. tree[0] 永远不用:这样 idx //= 2 能正确地把路径上溯到根(1 // 2 = 0,循环在 idx >= 1 条件下终止于根)。
__setitem__ 逐行解读
def __setitem__(self, idx, val):
    idx += self.capacity           # ① 外部下标 → 叶子下标
    self.tree[idx] = val           # ② 直接写叶子值

    idx //= 2                      # ③ 跳到父节点
    while idx >= 1:                # ④ 一路爬到根
        self.tree[idx] = self.operation(
            self.tree[2 * idx],    #    左儿子
            self.tree[2 * idx + 1] #    右儿子
        )
        idx //= 2                  # ⑤ 继续上爬

它的执行过程是改叶子 → 爬父链 → 每到一层用两个儿子重算当前节点,这就是"自底向上"的含义。

capacity=4, arr=[2,3,4,5] 追踪一次

假设 operation = add(求和),tree 初始全为 0:

初始:     [_, 0, 0, 0, 0, 0, 0, 0]

st[0]=2:  idx=4, tree[4]=2
          idx→2: tree[2] = tree[4]+tree[5] = 2+0 = 2
          idx→1: tree[1] = tree[2]+tree[3] = 2+0 = 2
          [_, 2, 2, 0, 2, 0, 0, 0]

st[1]=3:  idx=5, tree[5]=3
          idx→2: tree[2] = tree[4]+tree[5] = 2+3 = 5
          idx→1: tree[1] = tree[2]+tree[3] = 5+0 = 5
          [_, 5, 5, 0, 2, 3, 0, 0]

st[2]=4:  idx=6, tree[6]=4
          idx→3: tree[3] = tree[6]+tree[7] = 4+0 = 4
          idx→1: tree[1] = tree[2]+tree[3] = 5+4 = 9
          [_, 9, 5, 4, 2, 3, 4, 0]

st[3]=5:  idx=7, tree[7]=5
          idx→3: tree[3] = tree[6]+tree[7] = 4+5 = 9
          idx→1: tree[1] = tree[2]+tree[3] = 5+9 = 14
          [_, 14, 5, 9, 2, 3, 4, 5]   ← 根 = 2+3+4+5 = 14

注意一个细节:调用 st[0]=2 时,tree[2] 用了 tree[5](还是初始值 0)。这是对的——此时叶 1 还没赋值,按 0 算。后续 st[1]=3 会重新算 tree[2] 把它修正过来。所以建树顺序无关,最后一定收敛到正确状态。

权衡:初始化稍重,后续更方便

Python 选择"用 __setitem__ 建树"而非"写一个独立的 build",背后是一次初始化 vs 后续操作的权衡

维度初始化时后续运行时
递归版 build1 次 O(n) 递归调用点修改要递归到叶子,O(log n) 递归开销
py 版 __setitem__×nn 次 O(log n) 迭代,总 O(n log n)点修改同建树,迭代 O(log n),无递归开销

看起来 py 版初始化慢了一个 log,但这是有意的设计选择:

  1. 在 PER 的训练循环中:初始化只执行一次(往 buffer 塞 n 个样本),而训练中每次采样后要更新某个样本的优先级——后者要执行成千上万次。后续操作的简洁性比初始化的 log 因子重要
  2. 统一接口:建树和更新用同一个 __setitem__,代码更简单、更不容易出 bug。不像递归版要写两套(build 递归 + update 递归)。
  3. 纯 Python 友好:Python 的函数调用开销远大于 C,递归深度大时可能触发 RecursionError。迭代版完全没有递归问题。

本质上这是工业代码的典型取舍:用一次初始化的额外开销,换取训练循环中高频操作的简洁和稳定

与递归版 build1 对比
维度递归版 build1py 版 __setitem__×n
方向自顶向下递归,回溯时合并自底向上迭代,沿父链爬
建树调用1 次 build1(1,1,n)n 次 st[i]=v
单次复杂度整次 O(n)每次 O(log n)
建树总复杂度O(n)O(n log n)
节点存 l,r是(必须,递归要靠它判断)否(靠下标算术)
内存4N2·capacity

翻译成 C:segment_tree_bottomup.c

为了彻底吃透这套"自底向上"思路,我借助 AI 把 segment_tree.py 按博客代码风格翻译成 C,写成 segment_tree_bottomup.c。翻译过程中补了一个 py 版没有的亮点——真正的 O(n) 自底向上建树

void segtree_build(SegTree *st, ll *arr, int n) {
    /* 第一步:填叶子 */
    for (int i = 0; i < n; i++) {
        st->tree[st->capacity + i] = arr[i];
    }
    /* 第二步:从 capacity-1 倒推到 1,每个内部节点用两个儿子算出来 */
    for (int i = st->capacity - 1; i >= 1; i--) {
        st->tree[i] = st->op(st->tree[2 * i], st->tree[2 * i + 1]);
    }
}

为什么 i 递减就对:算 tree[i] 时要用 tree[2i]tree[2i+1],而 2i2i+1 都严格大于 i——所以 i 从大到小扫,轮到 i 时它两个儿子一定已算好。这就是"自底向上"的本质,比 py 版靠 n 次 __setitem__(O(n log n))省了一个 log,达到真正的 O(n)。

工程实现上的理解

理解一:完美二叉树换来的"省"

capacity 是 2 的幂看似是限制,实则换来三重好处:

  1. 数组大小 2*capacity 而非 4*N(省一半内存);
  2. 节点不存 l, r(靠下标算术,省字段);
  3. 建树可迭代(父子关系确定,无需递归回溯)。

代价是 capacity 要向上取整到 2 的幂(如 N=5 要用 capacity=8),有最多 2 倍的空间浪费。但比起省下的递归开销和代码简洁度,值得。

理解二:泛化累积——没有 s += 也能求和

翻译 query_helper 时一度怀疑:"这里没有 s += search1(...) 的累加,能求和吗?"

想通后发现:累积藏在 op(left, right)query_helper 把情况分成互斥四类(完全覆盖 / 完全在左 / 完全在右 / 跨越中点),只有"跨越中点"需要合并两个递归结果,由 st->op(left, right) 完成:

} else {                                     /* 跨越中点,拆分 */
    return st->op(
        query_helper(st, l,     mid, 2 * node,     ns,     mid),
        query_helper(st, mid + 1, r, 2 * node + 1, mid + 1, ne)
    );
}

求和树时 op = op_add = a+bop(left, right) = left + right 就是累积。没有 s += 不代表不能求和,只是把"累加到一个变量"换成了"用 op 合并两个返回值"。这种写法把合并操作抽离,求和/求最小用同一套代码。

理解三:retrieve 是这趟学习最特别的方法

retrieve 在普通线段树教程里找不到对应,它是工业代码为 PER 定制的扩展。从根往下走,根据左儿子和与 upperbound 的比较决定方向,返回前缀和第一次超过 upperbound 的下标。这是 PER 按权重采样的核心:权重大的样本被抽中概率高,强化学习 PER 就是靠它实现 O(log n) 的按优先级抽样。

理解了它,才真正理解 segment_tree.py 为什么被 RL 论文广泛使用。

第三站小结

对照转换的c文件,对 segment_tree.py 的每一个方法都能讲清楚:

  • __init__:开 2*capacity 数组,capacity 必须 2 的幂
  • __setitem__:迭代点修改,idx //= 2 爬父链
  • __getitem__:O(1) 读叶子
  • operate / _operate_helper:递归区间查询,左闭右开 [start, end)
  • retrieve:PER 按前缀和采样
  • SumSegmentTree / MinSegmentTree:靠 operation 参数特化

并且明确知道它的边界——没有 pushdown,做不了区间修改;若需要区间修改,要回到第二站第三部分的模板。


第四站:对照 C 文件学习,运行 segment_tree.py 的例程

写对照例程验证理解

有了 C 翻译版做对照,下一步就是运行 segment_tree.py 的例程,验证自己的理解。我写了一个对照学习例程 segment_tree_learn.py,分 5 个 Part 演示 SumSegmentTreeMinSegmentTreeretrieve、基类泛化、与递归版的差异对比。

operation 泛化:一段代码服务多种语义

这是工业代码相对教学模板最大的进步。看 _operate_helper

def _operate_helper(self, start, end, node, node_start, node_end):
    if start == node_start and end == node_end:
        return self.tree[node]
    mid = (node_start + node_end) // 2
    if end <= mid:
        return self._operate_helper(start, end, 2 * node, node_start, mid)
    elif mid + 1 <= start:
        return self._operate_helper(start, end, 2 * node + 1, mid + 1, node_end)
    else:
        return self.operation(                           # ← 关键:用 operation 合并
            self._operate_helper(start, mid, 2 * node, node_start, mid),
            self._operate_helper(mid + 1, end, 2 * node + 1, mid + 1, node_end),
        )

注意最后一行 self.operation(left, right)——求和时它是 +,求最小时它是 min同一段代码不用改。而递归版的 search1 写死了 s += search1(...),只能求和,求 max 要另写一个 search5

第三站翻译时已经想通了这个点,这里再强调一次:累积藏在 op(left, right)。把合并操作抽成 op,求和换成加法、求最小换成 min,同一段代码不需要改

retrieve:PER 按权重采样

SumSegmentTree.retrieve 从根往下走,根据"左儿子和 vs upperbound"决定往左还是往右,返回前缀和第一次超过 upperbound 的下标:

def retrieve(self, upperbound):
    idx = 1
    while idx < self.capacity:        # 非叶子
        left = 2 * idx
        if self.tree[left] > upperbound:
            idx = 2 * idx             # 走左
        else:
            upperbound -= self.tree[left]
            idx = 2 * idx + 1         # 走右
    return idx - self.capacity

它的用途是 PER 按权重采样

random_x = random() * st.sum()        # [0, 总和) 内随机数
idx = st.retrieve(random_x)           # 落到哪个样本 → 权重大的被抽中概率高

普通线段树教程不会讲这个,因为它是为特定场景(PER)定制的扩展。理解了它才真正理解 segment_tree.py 为什么被 RL 论文广泛使用。


总结

学习路径

论文中遇到 segment_tree.py(PER 采样)
        │
        │ 步骤2:退回基础,参考博客学原理
        ▼
线段树原理(CSDN 博客 + OI Wiki)
        │
        │ 步骤3:C 语言单步调试博客算法
        ▼
segment_tree_all_in_one.c(五段递进递归实现)
        │
        │ 步骤4:AI 将 py 转换为 C 文件
        ▼
segment_tree_bottomup.c(自底向上迭代版)
        │
        │ 步骤5:对照 C 学 py,运行 py 例程
        ▼
segment_tree_learn.py → 完全掌握 segment_tree.py

三份代码的定位

---> 学习文件仓库

文件对应步骤定位核心特点
segment_tree_all_in_one.c步骤 3基础学习5 段递进,含 pushdown 懒标记,递归自顶向下
segment_tree_bottomup.c步骤 4过渡桥梁自底向上 O(n) 建树,函数指针泛化,融合两者
segment_tree_learn.py步骤 5验证掌握对照 py 的学习例程,5 Part 演示
DDPGfD_code步骤 6RL算法实战学习 PER 在RL中实际应用

参考资源