# 线段树学习 **Repository Path**: xiao_peipei/learning-segment-trees ## Basic Information - **Project Name**: 线段树学习 - **Description**: 线段树学习相关代码文件 - **Primary Language**: Unknown - **License**: Not specified - **Default Branch**: master - **Homepage**: None - **GVP Project**: No ## Statistics - **Stars**: 0 - **Forks**: 0 - **Created**: 2026-08-07 - **Last Updated**: 2026-08-07 ## Categories & Tags **Categories**: Uncategorized **Tags**: None ## README # 从强化学习论文里的 segment_tree.py 倒推线段树原理 ## 引子:为什么我要花一周学线段树 最近在啃强化学习方向的论文,发现一个有意思的现象:**大部分涉及优先经验回放([Prioritized Experience Replay, PER](https://arxiv.org/abs/1511.05952))的论文,以及各种按优先级排序、按权重采样的 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 算法](https://github.com/MrSyee/pg-is-all-you-need/blob/master/06.DDPGfD.ipynb),理解 PER 在 RL 中具体实现。 --- ## 第一站:参考博客学习线段树原理 ### 参考资源 我主要参考了两个资料,互补着看: - **CSDN 博客**:[线段树从入门到进阶](https://blog.csdn.net/weixin_45697774/article/details/104274713) —— 例子具体,适合入门,标题就叫"超清晰,简单易懂" - **OI Wiki**:[线段树](https://oi-wiki.org/ds/seg/) —— 较严谨,覆盖全面 ### 线段树的本质 线段树的本质一句话能说清:**把一个数组组织成一棵二叉树,每个节点管辖一段区间,把区间信息(和、最值等)预先算好存起来,从而把"区间查询"和"单点修改"都做到 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` 里单步调试验证。这是我认为最有效的学法——一次写全五种变体,逼自己想清楚每一步在干什么。 ### 第一部分:建树 + 单点修改 + 区间查询(求和) 最基础的形态。建树是**自顶向下递归**,到叶子存值,回溯时合并: ```c 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` 三件事缺一不可: ```c 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__` 逐行解读 ```python 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,但这是有意的设计选择: 1. **在 PER 的训练循环中**:初始化只执行一次(往 buffer 塞 n 个样本),而训练中每次采样后要更新某个样本的优先级——后者要执行成千上万次。**后续操作的简洁性比初始化的 log 因子重要**。 2. **统一接口**:建树和更新用同一个 `__setitem__`,代码更简单、更不容易出 bug。不像递归版要写两套(`build` 递归 + `update` 递归)。 3. **纯 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) 自底向上建树**: ```c 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 的幂看似是限制,实则换来三重好处: 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)` 完成: ```c } 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`: ```python 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 的下标: ```python 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 按权重采样**: ```python 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 演示 | ## 参考资源 - **CSDN 博客**:[线段树从入门到进阶](https://blog.csdn.net/weixin_45697774/article/details/104274713) - **OI Wiki**:[线段树](https://oi-wiki.org/ds/seg/) - **OpenAI baselines**:[segment_tree.py 源码](https://github.com/openai/baselines/blob/master/baselines/common/segment_tree.py) - **Prioritized Experience Replay 原论文**:Schaul et al., 2015