ARTICLE DETAIL

资讯详情

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

Dopamine 实验数据工具集:dopamine.colab.utils 源码级解析与实战

Dopamine 实验数据工具集:dopamine.colab.utils 源码级解析与实战 Dopamine 实验数据工具集dopamine.colab.utils 源码级解析与实战【免费下载链接】dopamineDopamine is a research framework for fast prototyping of reinforcement learning algorithms.项目地址: https://gitcode.com/gh_mirrors/do/dopaminedopamine.colab是 Dopamine 强化学习研究框架中专门面向实验数据分析的模块其核心是 utils 子模块提供了一组读取、汇总与对比训练/评估统计数据的实用函数。本文以该模块的 API 文档为主线结合仓库源码与 Colab 笔记本中的真实用法讲解load_statistics、read_experiment、summarize_data、load_baselines等函数的签名、参数、底层实现与典型调用方式读完即可在自己的实验流程中直接复用它来加载日志、绘制学习曲线并与官方基线对比。一、模块定位dopamine.colab 与 utils 的职责dopamine.colab模块的定位非常聚焦为 ColabJupyter 笔记本环境提供处理 Dopamine 实验数据的工具函数。模块文档 colab.md 明确指出它目前包含一个子模块utils其职责是 provides utilities for dealing with Dopamine data即处理与 Dopamine 数据尤其是训练过程中产生的统计日志相关的读写与汇总操作。从仓库目录结构看dopamine/colab/ 目录除了核心实现 utils.py 之外还配套了一系列示范笔记本agents.ipynb演示如何通过继承 DQN 或从零创建新 Agentload_statistics.ipynb演示如何加载并可视化 Dopamine 产生的日志数据与 utils 模块直接相关agent_visualizer.ipynb 与 jax_agent_visualizer.ipynb可视化已训练 Agent分别针对 TensorFlow 与 JAX 实现tensorboard.ipynb下载并用 TensorBoard 可视化不同 Agent 的结果cartpole.ipynb在 Cartpole 环境上训练 DQN 与 C51。utils.py是这些数据分析笔记本共用的底层依赖本文后续内容将围绕它展开。二、数据从哪来Logger 生成的统计日志格式要理解utils各函数首先需要明白它所读取的数据是什么格式。Dopamine 在训练过程中通过日志模块把统计数据以 pickle 形式落盘读写双方约定相同的文件命名规则。dopamine.colab.utils中定义了两个关键前缀常量见 utils.pyFILE_PREFIX log ITERATION_PREFIX iteration_写入侧的实现位于 logger.py。其Logger类logger.py接收logging_dir与logs_duration参数log_to_file(filename_prefix, iteration_number)logger.py将内部字典self.data以pickle.dump序列化写入logging_dir/log_iteration_number文件并在写入后删除早于logs_duration版本的旧日志def log_to_file(self, filename_prefix, iteration_number): ... log_file self._generate_filename(filename_prefix, iteration_number) with tf.io.gfile.GFile(log_file, w) as fout: pickle.dump(self.data, fout, protocolpickle.HIGHEST_PROTOCOL)因此utils模块读取的每个日志文件都是一个被 pickle 序列化的 Python 字典其键形如iteration_0、iteration_1、……每个键对应一个迭代iteration的统计信息如train_episode_returns、eval_episode_returns等。文件命名log_N中数字最大的即最新迭代这也解释了 utils 中寻找最新文件/最新迭代这一组函数的必要性。三、定位最新迭代get_latest_iteration 与 get_latest_file当面对一个存有多个log_0, log_1, log_2, ...文件的目录时需要先确定哪个是最新的。utils 提供了两个互补的小函数。3.1 get_latest_iteration完整签名与文档见 get_latest_iteration.md源码见 utils.pydef get_latest_iteration(path): Return the largest iteration number corresponding to the given path. ... Raises: ValueError: if there is not available log data at the given path. glob os.path.join(path, {}_[0-9]*.format(FILE_PREFIX)) log_files tf.io.gfile.glob(glob) if not log_files: raise ValueError(No log data found at {}.format(path)) def extract_iteration(x): return int(x[x.rfind(_) 1:]) latest_iteration max(extract_iteration(x) for x in log_files) return latest_iteration要点参数path日志所在目录含目录与基础名称的路径通过tf.io.gfile.glob匹配形如log_数字的文件[0-9]*通配符对每个文件名取最后一个下划线之后的子串转换为整数再取最大值得到最新迭代号若目录中没有任何匹配的日志文件抛出ValueError信息为 No log data found at ...。3.2 get_latest_file完整签名与文档见 get_latest_file.md源码见 utils.pydef get_latest_file(path): Return the file named path_[0-9]* with the largest such number. try: latest_iteration get_latest_iteration(path) return os.path.join(path, {}_{}.format(FILE_PREFIX, latest_iteration)) except ValueError: return None它是对get_latest_iteration的封装成功时返回最新文件的完整路径path/log_N若路径下没有可用日志ValueError则返回None而非抛出异常方便调用方做容错判断。四、读取统计对象load_statisticsload_statistics是加载单次单个迭代统计数据的核心函数。完整签名与文档见 load_statistics.md源码见 utils.pydef load_statistics(log_path, iteration_numberNone, verboseTrue): Reads in a statistics object from log_path. ... Returns: data: The requested statistics object. iteration: The corresponding iteration number. Raises: Exception: if data is not present. # If no iteration is specified, well look for the most recent. if iteration_number is None: iteration_number get_latest_iteration(log_path) log_file %s/%s_%d % (log_path, FILE_PREFIX, iteration_number) if verbose: print(Reading statistics from: {}.format(log_file)) with tf.io.gfile.GFile(log_file, rb) as f: return pickle.load(f), iteration_number参数与行为参数类型/默认值说明log_pathstr训练/评估统计日志的完整路径目录iteration_numberint /None要读取的迭代号为None时自动定位最新版本verbosebool默认True是否打印加载过程信息实际会打印 Reading statistics from: 返回值是一个二元组(data, iteration)datapickle 反序列化得到的统计对象字典键为iteration_0、iteration_1……iteration本次实际读取对应的迭代号若未显式指定即为最新迭代号。需要注意两点实现细节若iteration_number为None内部会先调用get_latest_iteration(log_path)自动定位最新日志这正是上一节函数的价值所在文件读写统一走tf.io.gfile.GFile因此log_path既可以是本地目录也可以是 TensorFlow 文件系统如 GCS支持的路径若数据不存在将抛出异常Exception文档明确标注 if data is not present。五、逐迭代汇总summarize_dataload_statistics返回的是原始字典而summarize_data负责把它整理成逐迭代的汇总结果是绘制学习曲线前的标准预处理步骤。完整签名与文档见 summarize_data.md源码见 utils.pydef summarize_data(data, summary_keys): Processes log data into a per-iteration summary. Args: data: Dictionary loaded by load_statistics describing the data. This dictionary has keys iteration_0, iteration_1, ... describing per-iteration data. summary_keys: List of per-iteration data to be summarized. Example: data load_statistics(...) summarize_data(data, [train_episode_returns, eval_episode_returns]) Returns: A dictionary mapping each key in returns_keys to a per-iteration summary. summary {} latest_iteration_number len(data.keys()) current_value None for key in summary_keys: summary[key] [] # Compute per-iteration average of the given key. for i in range(latest_iteration_number): iter_key {}{}.format(ITERATION_PREFIX, i) # We allow reporting the same value multiple times when data is missing. # If there is no data for this iteration, use the previous. if iter_key in data: current_value np.mean(data[iter_key][key]) summary[key].append(current_value) return summary行为要点data由load_statistics加载的字典键为iteration_0、iteration_1……summary_keys需要汇总的逐迭代字段名列表例如[train_episode_returns, eval_episode_returns]返回值是一个字典每个summary_keys中的键对应一个列表列表第 i 个元素是第 i 个迭代该字段的均值np.mean数据缺失处理若某个迭代如iteration_5在data中不存在则沿用上一次的current_value保证列表长度始终等于迭代总数从而可直接与迭代下标对齐绘图。官方文档给出的示例data load_statistics(...) summarize_data(data, [train_episode_returns, eval_episode_returns])六、批量读取实验read_experimentread_experiment是 utils 中最高层的批处理入口它根据参数空间parameter_set与作业描述符job_descriptor的笛卡尔积自动构造多条实验路径逐一加载日志并汇总为一张 Pandas DataFrame。完整签名与文档见 read_experiment.md源码见 utils.pydef read_experiment( log_path, parameter_setNone, job_descriptor, iteration_numberNone, summary_keys(train_episode_returns, eval_episode_returns), verboseFalse, ):6.1 路径构造规则文档说明该函数会读取形如下面的所有实验${log_path}/${job_descriptor}.format(params)/logs其中params由parameter_set中各参数取值的笛卡尔积生成。官方示例parameter_set collections.OrderedDict([ (game, [Asterix, Pong]), (epsilon, [0, 0.1]) ]) read_experiment(/tmp/logs, parameter_set, job_descriptor{}_{})上述调用会尝试读取/tmp/logs/Asterix_0/logs/tmp/logs/Asterix_0.1/logs/tmp/logs/Pong_0/logs/tmp/logs/Pong_0.1/logs6.2 参数说明参数类型/默认值说明log_pathstr实验结果的基准路径parameter_setcollections.OrderedDict/None参数名到允许取值列表的映射顺序决定其在job_descriptor中的出现顺序job_descriptorstr默认用于拼接每条 trial 完整路径的模板字符串如{}_{}iteration_numberint /None若非None固定读取该迭代号的数据否则取最新summary_keys可迭代的 str默认(train_episode_returns, eval_episode_returns)需要汇总的逐迭代统计字段verbosebool默认False是否打印额外信息返回值为一张 Pandas DataFrame列依次为parameter_set中的各参数名 iterationsummary_keys中的各字段。源码中通过itertools.product(*ordered_values)生成参数元组逐条调用load_statistics与summarize_data填充行见 utils.py。6.3 实现细节预分配 DataFrame默认按参数组合数 × 200expected_num_iterations预留行数随后data_frame.drop(np.arange(row_index, expected_num_rows))裁剪未使用的行数值类型归一所有可转成数值的列统一astype(np.float64)避免后续merge时因 object 类型引发ValueError源码注释明确说明了这一点若job_descriptor为None则自动用key_value形式如game_Asterix-epsilon_0构造路径。load_statistics.ipynb中有一个贴近实战的调用见 load_statistics.ipynbparameter_set collections.OrderedDict([ (agent, [rainbow]), (game, GAMES) ]) sample_data colab_utils.read_experiment( /content/samples, parameter_setparameter_set, job_descriptor{}/{}_v4, summary_keys[train_episode_returns])它把彩虹智能体rainbow在多个游戏GAMES如Asterix_v4、Pong_v4等上的训练回报一次性读入并通过experimental_data[game].merge(sample_data, howouter)与官方基线数据合并随后用 seaborn 绘制对比曲线。七、读取官方基线数据load_baselinesload_baselines用于从指定基准目录读取 Dopamine 官方实验基线数据。完整签名与文档见 load_baselines.md源码见 utils.pydef load_baselines(base_dir, verboseFalse): Reads in the baseline experimental data from a specified base directory. Args: base_dir: string, base directory where to read data from. verbose: bool, whether to print warning messages. Returns: A dict containing pandas DataFrames for all available agents and games. 其内部逻辑为遍历模块内置的ALL_GAMESutils.py包含 60 个经典 Atari 游戏名如AirRaid、Alien、Asterix、Breakout、Pong、SpaceInvaders等对每个游戏依次尝试 4 个基线智能体[dqn, c51, rainbow, iqn]尝试读取路径base_dir/agent/Game.pklpickle 文件使用tf.io.gfile.GFilePython 3 下以encodinglatin1反序列化以兼容旧版本数据若文件不存在且verboseTrue打印 Unable to load data for agent ... on game ... 后跳过将数据列统一转为np.float64后按游戏合并merge(..., howouter)。返回值为dict键为游戏名值为包含所有可用智能体数据的 Pandas DataFrame每个 DataFrame 会额外打上agent列值为dqn/c51/rainbow/iqn。注意ALL_GAMES全部是 Atari 游戏模块同时定义了MUJOCO_GAMES [Ant, HalfCheetah, Hopper, Humanoid, Walker2d]utils.py但基线加载函数目前只遍历ALL_GAMES这一点从源码可以明确确认。八、实战串联在 Colab 中加载日志并绘制学习曲线结合 load_statistics.ipynb 的完整流程可以串起本文介绍的全部函数。典型的数据分析管线如下第一步加载官方基线experimental_data colab_utils.load_baselines(/path/to/baselines)第二步批量读取自己的实验日志汇总后与基线合并import collections parameter_set collections.OrderedDict([ (agent, [rainbow]), (game, GAMES) ]) sample_data colab_utils.read_experiment( /content/samples, parameter_setparameter_set, job_descriptor{}/{}_v4, summary_keys[train_episode_returns]) sample_data[agent] Sample Rainbow sample_data[run_number] 1 for game in GAMES: experimental_data[game] experimental_data[game].merge( sample_data[sample_data.game game], howouter)第三步读取原始数据并逐迭代汇总raw_data, _ colab_utils.load_statistics( /content/samples/rainbow/{}_v4/logs.format(game), verboseFalse) summarized_data colab_utils.summarize_data( raw_data, [train_episode_returns])第四步绘图import seaborn as sns import matplotlib.pyplot as plt sns.lineplot(xiteration, ytrain_episode_returns, hueagent, dataexperimental_data[game], axax)这种load_baselinesread_experimentload_statistics/summarize_data 绘图的组合正是dopamine.colab.utils设计的目的把日志读取、基线对比与曲线绘制从每个研究者的重复劳动中解放出来。九、配套工具与扩展除了utils模块仓库还在 dopamine/utils/ 目录下提供了与 Colab 可视化配套的绘图工具包括plotter.py绘图基础设施line_plotter.py折线图绘制bar_plotter.py柱状图绘制atari_plotter.pyAtari 实验专用绘图。这些工具与colab.utils的数据读取能力共同支撑 agent_visualizer.ipynb、jax_agent_visualizer.ipynb 等可视化笔记本。若需要进一步了解 JAX 版本 Agent 的实现可参阅 jax/agents 目录。十、总结dopamine.colab.utils是连接训练日志与研究分析的桥梁六个函数各司其职、层层递进get_latest_iteration/get_latest_file在log_N命名规范的日志目录中定位最新数据load_statistics读取单个迭代或最新迭代的 pickle 统计对象summarize_data把原始字典汇总为逐迭代均值序列缺失数据自动沿用前值read_experiment基于参数空间的笛卡尔积批量加载并汇总多条实验输出 Pandas DataFrameload_baselines一键读取官方 Atari 基线dqn/c51/rainbow/iqn用于对比。在 load_statistics.ipynb 等笔记本中这些函数与load_baselines返回的基线数据合并后即可直接绘制学习曲线形成从实验日志到科研图表的标准工作流。对于希望深度定制分析流程的研究者建议直接阅读 utils.py 源码并结合 logger.py 理解数据在写入侧的具体结构。【免费下载链接】dopamineDopamine is a research framework for fast prototyping of reinforcement learning algorithms.项目地址: https://gitcode.com/gh_mirrors/do/dopamine创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表