ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

Dopamine 的 Sum Tree 数据结构:实现优先经验回放(Prioritized Experience Replay)的基石

Dopamine 的 Sum Tree 数据结构:实现优先经验回放(Prioritized Experience Replay)的基石 机器学习深度学习【免费下载链接】dopamineDopamine is a research framework for fast prototyping of reinforcement learning algorithms.项目地址https://gitcode.com/gh_mirrors/do/dopamine点击查看免费下载本文深入剖析 Dopamine 框架中dopamine.tf.replay_memory.sum_tree模块从它的数据结构设计、存储布局到set/get/sample/stratified_sample等核心 API 的逐行实现原理再到它如何被优先经验回放缓冲区和 Rainbow 智能体所调用。读完本文你将理解 Sum Tree 为什么能让高优先级样本被更大概率采样这一操作在 O(log n) 时间内完成并能在自己的强化学习项目中正确使用与扩展这一结构。模块定位优先经验回放的数据支撑dopamine.tf.replay_memory.sum_tree是 Dopamine 框架中一个小而关键的模块。其官方 API 文档sum_tree.md给出了一句话定位A sum tree data structure. Used for prioritized experience replay. See prioritized_replay_buffer.py and Schaul et al. (2015).它服务于Prioritized Experience Replay优先经验回放PER——即 Tom Schaul 等人在 2015 年论文Prioritized Experience Replay中提出的经典算法。其核心思想是在 DQN 类算法中经验池里的样本并非同等重要TD 误差或类似指标越大的样本越值得被频繁学习因此需要一种能够按权重priority加权采样的数据结构。模块只对外暴露一个类SumTree完整实现见 sum_tree.py。同一结构在 JAX 版也有对应的向量化实现dopamine/jax/replay_memory/sum_tree.py本文以 TF 版为主线展开最后对比两者的差异。Sum Tree 的数据结构原理完全二叉树 子树和Sum Tree 是一棵完全二叉树complete binary tree其叶子节点保存着每个经验样本的 priority优先级值内部节点保存其子树内所有叶子 priority 之和。这样一来根节点保存的就是所有样本的 priority 总和R而任意一个叶子节点所覆盖的区间就对应着按累积概率划分的[0, R)中的一段子区间。源码 docstring 中用capacity 4给出了直观示例sum_tree.py--- |2.5| - 根节点所有叶子 priority 之和 -- | --------------- | | -- -- |1.5| |1.0| -- -- | | -------- -------- | | | | -- -- -- -- |0.5| |1.0| |0.5| |0.5| --- --- --- ---根节点 2.5 0.5 1.0 0.5 0.5即所有叶子 priority 之和第二层左节点 1.5 0.5 1.0右节点 1.0 0.5 0.5。存储布局numpy 数组列表与教科书式的单数组下标2i1、2i2表示法不同Dopamine 的实现将树的每一层单独存为一个 numpy 数组整体是一个数组列表self.nodes [ [2.5], [1.5, 1], [0.5, 1, 0.5, 0.5] ]即self.nodes[0]是根层长度 1self.nodes[1]是第二层长度 2self.nodes[-1]是最底层叶子长度等于 2 的幂。源码注释指出这与通常的数组式完全二叉树表示类似但对用户更友好。构造按 2 的幂分配并补零构造函数__init__(self, capacity)sum_tree.py的关键逻辑校验capacity必须是int使用assert isinstance(capacity, int)且capacity 0时抛出ValueError: Sum tree capacity should be positive. Got: {capacity}用tree_depth int(math.ceil(np.log2(capacity)))计算树深然后从level_size 1开始逐层翻倍地创建np.zeros(level_size)数组共tree_depth 1层由于按 2 的幂分配多余的位置会被零填充。例如capacity 100时最底层叶子数组长度为 128其中 28 个位置恒为 0测试 sum_tree_test.py 中的testCapacityGreaterThanRequested专门断言了这一点初始化self.max_recorded_priority 1.0这是 PER 中新样本默认取历史最大 priority策略的起点Schaul 等人的方案详见下文。核心 API 逐一解析SumTree对外暴露 5 个方法。下面结合源码逐一定义其语义、参数与异常行为。_total_priority()根节点总和私有方法直接返回self.nodes[0][0]即根节点值所有 priority 之和。它是采样前的归一化基准。set(node_index, value)O(log n) 更新并向上传播这是 Sum Tree 最核心的写操作sum_tree.py校验value非负否则抛出ValueError: Sum tree values should be nonnegative. Got {value}。value 0意味着该元素永远不会被采样更新self.max_recorded_priority max(value, self.max_recorded_priority)计算增量delta_value value - self.nodes[-1][node_index]从叶子层向上逐层回退对reversed(self.nodes)中的每一层把delta_value加到对应下标上然后node_index // 2上移到父节点遍历结束时断言node_index 0保证正确到达根节点。整个更新链路是O(log(capacity))这正是 Sum Tree 相比每次重新遍历整个数组求前缀和的性能优势所在。源码也坦诚地注释了一句通过加增量传播会带来可容忍的数值误差tolerable numerical inaccuracies。注意一个使用细节set的调用方必须保证传入的node_index在[0, capacity)范围内因为set自身不做越界校验直接索引self.nodes[-1]。get(node_index)读取叶子值return self.nodes[-1][node_index]直接返回最底层叶子数组中对应下标的值。对尚未写入的位置返回 0。这一特性被优先回放缓冲区用于未使用位置的 priority 视为 0的约定见 prioritized_replay_buffer.py 的get_priority。sample(query_valueNone)按权重采样一个元素采样是逆 CDF 采样先均匀随机一个query_value ∈ [0, 1)再乘上总 priorityR得到[0, R)中的目标值然后从根开始向下遍历sum_tree.py若树为空总 priority 为 0抛ValueError: Cannot sample from an empty sum tree.若显式传入的query_value不在[0, 1]抛ValueError: query_value must be in [0, 1].未传query_value时使用random.random()遍历时node_index从 0根开始对每一层左孩子下标为node_index * 2若query_value 左子树和则进入左子树否则进入右子树并query_value - 左子树和把目标值转换为相对右子树的偏移到达叶子层后返回叶子下标。这样每个元素被选中的概率精确为p_i / sum_j p_jp_i为第 i 个叶子的 priority可非归一化。测试 sum_tree_test.py 验证了非均匀概率行为当叶子 2 的 priority 为 1.0、叶子 3 的 priority 为 3.0 时query_value0.1总能命中叶子 2因为 0.1 × 4 0.4 1.0落在左子树区间。stratified_sample(batch_size)分层采样Schaul et al. 2015一次采样一个 batch 时如果每个样本都独立均匀采样会出现某些高 priority 区间被重复命中、低 priority 区间被冷落的现象。分层采样stratified sampling将[0, R)均分成batch_size段每段内各取一个随机数作为查询值从而保证整个 batch 在优先级区间上覆盖更均匀。这正是 Schaul 等人在 PER 论文中给出的采样方式。实现sum_tree.pybounds np.linspace(0.0, 1.0, batch_size 1) segments [(bounds[i], bounds[i 1]) for i in range(batch_size)] query_values [random.uniform(x[0], x[1]) for x in segments] return [self.sample(query_valuex) for x in query_values]空树同样抛ValueError。测试testStratifiedSamplingsum_tree_test.py验证当 32 个叶子 priority 全为 1 时stratified_sample(32)恰好返回[0, 1, ..., 31]——每个分层区间正好落在自己的叶子上。与优先经验回放缓冲区的集成Sum Tree 不是孤立的数据结构它是 prioritized_replay_buffer.py 中OutOfGraphPrioritizedReplayBuffer的底层存储。完整的调用链如下构造缓冲区在__init__中以self.sum_tree sum_tree.SumTree(replay_capacity)创建与经验池容量一致的 Sum Treeprioritized_replay_buffer.py写入add_add在把一条 transition 写入环形缓冲之前先从参数中取出priority调用self.sum_tree.set(self.cursor(), priority)写入当前游标位置。如果新样本没有显式 priority则按 Schaul 等人的方案使用max_recorded_priority——即新样本默认获得历史最大 priority保证新经验至少会被采样一次prioritized_replay_buffer.py采样samplesample_index_batch先调用self.sum_tree.stratified_sample(batch_size)得到候选下标随后对落在无效区间尚未填充或 n-step 窗口不完整的下标在_max_sample_attempts次内改用self.sum_tree.sample()重新采样兜底若始终凑不齐 batch 则抛RuntimeErrorprioritized_replay_buffer.py回写train 阶段训练完成后set_priority(indices, priorities)把新算出的 TD 误差相关 priority 批量写回 Sum Tree内部逐个调用sum_tree.setget_priority(indices)则读出采样概率供重要性采样权重IS weights计算使用。包装类WrappedPrioritizedReplayBuffer还通过tf.numpy_function提供了tf_set_priority/tf_get_priority两个 TensorFlow op供图内调用prioritized_replay_buffer.py。在 Rainbow / DQN 中的配置入口TF 版 Rainbow 智能体把优先回放作为默认采样方案RainbowAgent.__init__的replay_scheme参数默认值为prioritized也接受uniform并在_build_replay_buffer中据此实例化WrappedPrioritizedReplayBuffer见 rainbow_agent.py 与 rainbow_agent.py。在实际运行中这一切由 gin 配置驱动。以 c51.gin 为例import dopamine.tf.replay_memory.prioritized_replay_buffer RainbowAgent.replay_scheme prioritized # 显式启用优先回放 WrappedPrioritizedReplayBuffer.replay_capacity 1000000 WrappedPrioritizedReplayBuffer.batch_size 32WrappedPrioritizedReplayBuffer同时是gin.configurable的其可配置参数包括replay_capacity默认 1000000、batch_size默认 32、update_horizon、gamma、max_sample_attempts默认 1000、use_staging等见 prioritized_replay_buffer.py。当replay_scheme uniform时则退回普通的WrappedReplayBufferSum Tree 不再参与采样。值得一提的是本仓库中的 Rainbow 实现刻意简化了论文中的超参数将 PER 的beta固定为 0.5不随训练线性升温、去掉alpha参数论文中恒为 0.5这一点在 rainbow_agent.py 的模块 docstring 中有明确说明。测试如何验证正确性单元测试 sum_tree_test.py 覆盖了几乎所有边界情况可作为理解实现的行为说明书测试用例验证点testNegativeCapacitycapacity 0抛ValueErrortestSetNegativeValue写入负 priority 抛ValueErrortestSmallCapacityConstructorcapacity1时仅 1 层、capacity2时 2 层testSetValueset后get正确且左支所有祖先节点同步更新为 1.0testCapacityGreaterThanRequested叶子数组长度按 2 的幂向上取整100 → 128testSampleFromEmptyTree空树采样抛异常testSampleWithInvalidQueryValuequery_value越界抛ValueErrortestSampleSingleton单元素树采样恒返回该元素testSamplePairWithUnevenProbabilities概率与 priority 成正比1:3testSamplingWithSeedDoesNotAffectFutureCalls显式设置随机种子只影响单次采样不影响后续调用的随机性testStratifiedSampling32 个等 priority 叶子时分层采样恰好一一命中testMaxRecordedProbabilitymax_recorded_priority随写入值单调更新且初始为 1.0JAX 版 Sum Tree向量化差异仓库中还存在一个 JAX 版实现dopamine/jax/replay_memory/sum_tree.py其 API 文档见 jax/replay_memory/sum_tree.md。它被描述为a vectorized sum tree in numpy与 TF 版有几点显著差异存储改用单个一维 numpy 数组self._nodes长度2**depth - 1配合_first_leaf_offset 2**(depth-1) - 1定位叶子起始位置即教科书式的堆布局批量操作set/get/query均支持 numpy 数组批量传入set内部先np.unique去重避免同一下标被重复累加 delta再用np.add.at向量化累加采样将遍历改为掩码 批量下钻的query(targets)一次性解析一批目标值对应的叶子下标clear()可一键清零整棵树可 checkpoint实现了checkpointers.Checkpointable接口提供to_state_dict/from_state_dict用于序列化。JAX 版由 samplers.py 中的PrioritizedSamplingDistribution使用在add时set新样本 priority采样时用self._rng.uniform(0.0, self._sum_tree.root, sizesize)生成目标值并query出下标同时用get(indices) / root计算采样概率并在 epoch 结束时调用clear()重置。两个版本在数学语义上等价但 JAX 版在批量吞吐上更高效。小结dopamine.tf.replay_memory.sum_tree用约 200 行代码以每层一个 numpy 数组的简洁布局实现了优先经验回放所必需的全部能力set的 O(log n) 增量更新、sample的按权重采样、stratified_sample的批次分层覆盖。它向上支撑着OutOfGraphPrioritizedReplayBuffer与 Rainbow 智能体的replay_schemeprioritized配置向下又有完备的单元测试兜底。理解这个模块是深入理解 Dopamine 中优先回放机制乃至扩展到 JAX 向量化版本的最佳起点。赞分享机器学习深度学习【免费下载链接】dopamineDopamine is a research framework for fast prototyping of reinforcement learning algorithms.项目地址https://gitcode.com/gh_mirrors/do/dopamine点击查看免费下载相关推荐Dopamine 中优先经验回放Prioritized Experience Replay实现全解析prioritized_replay_buffer 模块深度指南Dopamine 中优先经验回放Prioritized Experience Replay实现全解析 prioritized_replay_buffer强化学习机器学习深度学习Dopamine SumTree 解析优先经验回放PER的基石数据结构Dopamine SumTree 解析优先经验回放PER的基石数据结构 Dopamine 是一个用于快速原型验证强化学习算法的研究框架见 README.强化学习机器学习深度学习深入解析 Dopamine 中的 JAX 向量化 Sum Tree优先经验回放PER的基石深入解析 Dopamine 中的 JAX 向量化 Sum Tree优先经验回放PER的基石 Sum Tree求和树是优先经验回放Prioritize机器学习深度学习上一篇WandEnhancer终极指南一键解锁Wand Pro功能的完整教程下一篇Wand-Enhancer终极指南免费解锁WeMod Pro功能的完整教程创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表