从强化学习论文里的 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 部分:
- 阅读相关 RL 算法,发现不懂 PER;
- 参考博客学习一般原理(线段树 从入门到进阶(超清晰,简单易懂));
- C 语言单步调试运行博客中的算法;
- 利用 AI 将
segment_tree.py文件按照博客代码风格转换为 c 文件; - 对照 c 文件学习,运行
segment_tree.py的例程; - 回归 DDPGfD 算法(参考仓库:pg-is-all-you-need),理解 PER 在 RL 中具体实现。
---> 学习文件gitee仓库
第一站:参考博客学习线段树原理
参考资源
我主要参考了两个资料,互补着看:
线段树的本质
线段树的本质一句话能说清:把一个数组组织成一棵二叉树,每个节点管辖一段区间,把区间信息(和、最值等)预先算好存起来,从而把"区间查询"和"单点修改"都做到 O(log n)。
关键结构是:节点 i 的左儿子是 2i,右儿子是 2i+1,每个节点存自己管辖区间 [l, r] 上的聚合值。查询时把目标区间拆成若干个节点区间之和,修改时改叶子再回溯更新祖先。
三条核心认知
学完原理后,我得到三条最重要的认知,它们是后续读懂工业代码的基础:
- pushdown 是线段树的灵魂——所有"区间修改"类问题都靠懒标记下传;
- lazy 语义决定 pushdown 写法——加法标记和覆盖标记天差地别;
- 框架与信息分离——换信息只改合并方式,不动框架。这条是后面读懂 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 的核心是三条规则:
- 当前区间被完全包含 → 直接返回
sum; - 当前区间与查询完全不相干 → 返回 0;
- 否则递归查询有交集的子节点。
一个有意思的细节:第 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 泛化的思想基础。
第二站小结
五段递进学完,对应第一站的三条核心认知都得到了代码验证:
- pushdown 是灵魂(第三部分);
- lazy 语义决定写法(第三 vs 第四部分);
- 框架与信息分离(第五部分)。
但这一切都是递归实现,和工业代码 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
完美二叉树带来三个关键性质:
- 叶子节点位置确定:叶子统一从下标
capacity开始,到2*capacity-1结束。外部下标i(0-based)的叶子 → 内部下标capacity + i。 - 父子关系靠算术确定:节点
k的父 =k // 2,左儿子 =2*k,右儿子 =2*k + 1。不需要在节点里存l, r字段,省了一半内存。 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 后续操作的权衡:
| 维度 | 初始化时 | 后续运行时 |
|---|---|---|
递归版 build | 1 次 O(n) 递归调用 | 点修改要递归到叶子,O(log n) 递归开销 |
py 版 __setitem__×n | n 次 O(log n) 迭代,总 O(n log n) | 点修改同建树,迭代 O(log n),无递归开销 |
看起来 py 版初始化慢了一个 log,但这是有意的设计选择:
- 在 PER 的训练循环中:初始化只执行一次(往 buffer 塞 n 个样本),而训练中每次采样后要更新某个样本的优先级——后者要执行成千上万次。后续操作的简洁性比初始化的 log 因子重要。
- 统一接口:建树和更新用同一个
__setitem__,代码更简单、更不容易出 bug。不像递归版要写两套(build递归 +update递归)。 - 纯 Python 友好:Python 的函数调用开销远大于 C,递归深度大时可能触发
RecursionError。迭代版完全没有递归问题。
本质上这是工业代码的典型取舍:用一次初始化的额外开销,换取训练循环中高频操作的简洁和稳定。
与递归版 build1 对比
| 维度 | 递归版 build1 | py 版 __setitem__×n |
|---|---|---|
| 方向 | 自顶向下递归,回溯时合并 | 自底向上迭代,沿父链爬 |
| 建树调用 | 1 次 build1(1,1,n) | n 次 st[i]=v |
| 单次复杂度 | 整次 O(n) | 每次 O(log n) |
| 建树总复杂度 | O(n) | O(n log n) |
| 节点存 l,r | 是(必须,递归要靠它判断) | 否(靠下标算术) |
| 内存 | 4N | 2·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],而 2i、2i+1 都严格大于 i——所以 i 从大到小扫,轮到 i 时它两个儿子一定已算好。这就是"自底向上"的本质,比 py 版靠 n 次 __setitem__(O(n log n))省了一个 log,达到真正的 O(n)。
工程实现上的理解
理解一:完美二叉树换来的"省"
capacity 是 2 的幂看似是限制,实则换来三重好处:
- 数组大小
2*capacity而非4*N(省一半内存); - 节点不存
l, r(靠下标算术,省字段); - 建树可迭代(父子关系确定,无需递归回溯)。
代价是 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+b,op(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 演示 SumSegmentTree、MinSegmentTree、retrieve、基类泛化、与递归版的差异对比。
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 | 步骤 6 | RL算法实战 | 学习 PER 在RL中实际应用 |
参考资源
- CSDN 博客:线段树从入门到进阶
- OI Wiki:线段树
- OpenAI baselines:segment_tree.py 源码
- PER 论文 : Prioritized Experience Replay:Schaul et al., 2015
- DDPGfD 论文 : Leveraging Demonstrations for Deep Reinforcement Learning on Robotics Problems with Sparse Rewards