ARTICLE DETAIL

资讯详情

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

Stable-Baselines3 Logger 完整指南:自定义日志格式与训练指标解读

Stable-Baselines3 Logger 完整指南:自定义日志格式与训练指标解读 Stable-Baselines3 Logger 完整指南自定义日志格式与训练指标解读【免费下载链接】stable-baselines3PyTorch version of Stable Baselines, reliable implementations of reinforcement learning algorithms.项目地址: https://gitcode.com/GitHub_Trending/st/stable-baselines3Stable-Baselines3SB3内置了一套灵活的日志系统支持同时向 stdout、CSV、JSON、TensorBoard 等多种目标输出训练过程的关键指标。本文以 SB3 官方文档中的 Logger 主题为核心结合 stable_baselines3/common/logger.py 源码系统讲解如何用configure与set_logger覆盖默认日志配置、五种输出格式的底层实现差异以及eval/、rollout/、time/、train/四类日志键的完整含义与产生来源。读完本文你将能够按需定制 SB3 的日志输出并准确读懂训练过程中的每一个指标。一、Logger 的作用与默认行为在 SB3 中Logger负责收集、汇总并输出训练过程中的诊断信息key-value 对。每个算法实例A2C、PPO、DQN、SAC、TD3、DDPG 等在创建时都会在 stable_baselines3/common/base_class.py 的BaseAlgorithm中初始化一个默认 logger。默认 logger 的创建路径位于 stable_baselines3/common/utils.py 的configure_logger函数其规则如下verbose0时不输出任何内容verbose1时输出格式为[stdout, tensorboard]前提是传入了tensorboard_log传入tensorboard_log时日志会写到tensorboard_log/tb_log_name_run_id目录tb_log_name默认取算法名小写如ppo、a2creset_num_timestepsFalse时会在同一个目录续写日志用于继续训练时保持学习曲线连续。也就是说不主动干预时日志系统的行为由verbose与tensorboard_log两个构造参数决定。二、覆盖默认 Loggerconfigure 与 set_logger官方文档给出的做法是先用configure创建自定义 logger再通过model.set_logger()把它交给算法实例。from stable_baselines3 import A2C from stable_baselines3.common.logger import configure tmp_path /tmp/sb3_log/ # set up logger new_logger configure(tmp_path, [stdout, csv, tensorboard]) model A2C(MlpPolicy, CartPole-v1, verbose1) # Set new logger model.set_logger(new_logger) model.learn(10000)configure的参数语义configure(folder, format_strings)是创建自定义 logger 的入口其实现位于 stable_baselines3/common/logger.pyfolder日志保存目录。为None时依次回退到环境变量SB3_LOGDIR再退到tempfile.gettempdir()下的SB3-日期时间临时目录源码中格式为SB3-%Y-%m-%d-%H-%M-%S-%f。format_strings输出格式列表。为None时读取环境变量SB3_LOG_FORMAT默认值为[stdout, log, csv]。目录不存在时自动创建os.makedirs(folder, exist_okTrue)。当格式不止stdout时logger 会先打印一行Logging to folder。set_logger的注意事项重要警告set_logger定义于 stable_baselines3/common/base_class.py。当传入自定义 logger 对象时会覆盖构造时传入的tensorboard_log与verbose设置——这是官方文档特别标注的警告set_logger会把_custom_logger标记为True此后_setup_learn中便不再调用configure_logger重建默认 logger见 stable_baselines3/common/base_class.py 的条件判断。因此如果你需要既保留 TensorBoard 记录、又调整输出格式请把tensorboard显式加入configure的format_strings中而不是依赖算法构造参数。三、五种输出格式与底层实现官方文档列出的可用格式为[stdout, csv, log, tensorboard, json]。它们由make_output_format统一分派见 stable_baselines3/common/logger.py格式输出目标对应类文件命名规则stdout终端标准输出HumanOutputFormatsys.stdoutlog人类可读文本文件HumanOutputFormatlog_dir/log.txtcsvCSV 表格文件CSVOutputFormatlog_dir/progress.csvjsonJSONL 逐行文件JSONOutputFormatlog_dir/progress.jsontensorboardTensorBoard 事件文件TensorBoardOutputFormatlog_dir目录stdout / log人类可读的 ASCII 表格HumanOutputFormat会把 key-value 输出为带边框的 ASCII 表格见 stable_baselines3/common/logger.py以/分隔的键名会被自动分组/前的部分作为标签tag单独成行/后的部分缩进 3 格显示数值默认以f{value:8.3g}格式对齐输出默认max_length36超长内容截断为...输出总宽度不超过 79 字符若不同键截断后重名会抛出ValueError提示通过增大max_length解决输出到 stdout 且安装了 tqdm 时使用tqdm.write避免与进度条互相干扰表格每次写入后立即flush()。csv动态表头的 progress.csvCSVOutputFormat见 stable_baselines3/common/logger.py维护一个keys列表每当出现新键就重写表头并给历史行补齐空列保证每行列数一致。字符串值会被包裹在引号中并转义内部引号便于 CSV 解析器正确处理包含分隔符的文本。json逐行 JSON 的 progress.jsonJSONOutputFormat见 stable_baselines3/common/logger.py把每轮dump的键值对序列化为一整行 JSONJSONL 格式。cast_to_json_serializable会把 0 维或长度为 1 的 numpy 数组转为float其余数组转为嵌套列表保证json.dumps可正常序列化。tensorboard直通 PyTorch SummaryWriterTensorBoardOutputFormat见 stable_baselines3/common/logger.py依赖torch.utils.tensorboard的SummaryWriter未安装 tensorboard 时会直接断言失败并提示pip install tensorboard标量通过add_scalar写入字符串通过add_text写入torch.Tensor/numpy.ndarray通过add_histogram写入对应直方图额外支持Video、Figure、Image、HParam四类特殊数据分别调用add_video、add_figure、add_image与add_hparams系列方法。这四类特殊数据封装类Video、Figure、Image、HParam均定义在 stable_baselines3/common/logger.py 中其中HParam的metric_dict必须非空否则无法在 TensorBoard 的 HPARAMS 标签页显示超参数。四、Logger 的核心 APIrecord、record_mean 与 dumpLogger类见 stable_baselines3/common/logger.py暴露三个核心方法record(key, value, excludeNone)记录一个键值对同键多次调用时取最后一次的值stable_baselines3/common/logger.pyrecord_mean(key, value, excludeNone)与record类似但同键多次调用时做指数滑动平均——new old * count / (count1) value / (count1)stable_baselines3/common/logger.pydump(step0)把本轮所有记录写入所有输出格式随后清空缓存stable_baselines3/common/logger.py。exclude参数按格式名排除特定键。例如 PPO 中记录train/n_updates时使用excludetensorboard见 stable_baselines3/ppo/ppo.py避免该键在 TensorBoard 中出现。若某个值不受某格式支持如向 stdout 写Video会抛出FormatUnsupportedError其报错信息会建议你通过exclude排除该键见 stable_baselines3/common/logger.py。此外Logger还提供日志级别控制set_level可设置DEBUG10 / INFO20 / WARN30 / ERROR40 / DISABLED50debug、info、warn、error方法对应输出不同级别的文本消息log()仅在self.level level时输出stable_baselines3/common/logger.py。DISABLED级别下dump直接返回、不做任何写入。五、日志输出的整体结构与示例训练时SB3 输出的键通常以eval/、rollout/、time/、train/为前缀分组。以训练一个 PPO agent 为例stdout 输出如下----------------------------------------- | eval/ | | | mean_ep_length | 200 | | mean_reward | -157 | | rollout/ | | | ep_len_mean | 200 | | ep_rew_mean | -227 | | time/ | | | fps | 972 | | iterations | 19 | | time_elapsed | 80 | | total_timesteps | 77824 | | train/ | | | approx_kl | 0.037781604 | | clip_fraction | 0.243 | | clip_range | 0.2 | | entropy_loss | -1.06 | | explained_variance | 0.999 | | learning_rate | 0.001 | | loss | 0.245 | | n_updates | 180 | | policy_gradient_loss | -0.00398 | | std | 0.205 | | value_loss | 0.226 | -----------------------------------------注意根据所用算法以及套用的 wrappers/callbacks 不同SB3 只会记录上述键的一个子集。例如 DQN 才会有exploration_rateSAC 才会有ent_coef只有接了EvalCallback才会出现eval/组。六、eval/前缀来自 EvalCallback 的评估指标所有eval/值都由EvalCallback计算实现在 stable_baselines3/common/callbacks.py 中mean_ep_length评估期间的平均回合长度mean_reward评估期间的平均回合奖励success_rate评估期间的平均成功率1.0 表示 100% 成功。要求环境的 info 字典包含is_success键才能计算——EvalCallback会在每个回合结束时读取info.get(is_success)存入_is_success_buffer再在_on_step中取均值并record(eval/success_rate, success_rate)见 stable_baselines3/common/callbacks.py。EvalCallback还会据此维护best_mean_reward用于在save_best时保存最优模型。七、rollout/前缀训练采样过程指标ep_len_mean平均回合长度对最近stats_window_size个回合求平均默认 100ep_rew_mean平均回合训练奖励同样对最近stats_window_size个回合求平均。计算该值必须有Monitor包装器make_vec_env会自动添加。这两个值在 stable_baselines3/common/on_policy_algorithm.py 与 stable_baselines3/common/off_policy_algorithm.py 中通过safe_mean对ep_info_buffer求均值后记录。stats_window_size是各算法的构造参数A2C/PPO/DQN/SAC/TD3 默认均为 100对应 stable_baselines3/common/base_class.py 中ep_info_buffer deque(maxlenself._stats_window_size)的窗口大小。exploration_rate当前探索率仅 DQN 记录。它对应 epsilon-greedy 中随机采取动作的概率epsilon在 stable_baselines3/dqn/dqn.py 中随训练进度按exploration_schedule衰减并记录success_rate训练期间的平均成功率同样在stats_window_size个回合上求平均。计算前提给Monitor包装器额外传入info_keywords(is_success,)并在每个回合最后一步的info中提供info[is_success] True/False。Monitor会把info_keywords中的键并入回合信息见 stable_baselines3/common/vec_env/vec_monitor.py随后BaseAlgorithm._update_info_buffer读取info.get(is_success)填充ep_success_buffer见 stable_baselines3/common/base_class.py。Monitor包装器同时会把{r: 回合奖励, l: 回合长度, t: 时间戳}注入info[episode]见 stable_baselines3/common/vec_env/vec_monitor.py这就是ep_info_buffer的数据来源。八、time/前缀训练计时信息episodes累计回合总数fps每秒帧数包含梯度更新的耗时iterations迭代次数对 A2C/PPO 而言一次迭代 一轮数据采集 一轮策略更新time_elapsed自训练开始以来的秒数total_timesteps累计环境步数所有并行环境中的步数总和。这些计时指标在训练主循环中每轮dump时统一记录。九、train/前缀训练损失与诊断指标actor_lossoff-policy 算法DQN/SAC/TD3/DDPG当前 actor 损失值approx_klPPO 新旧策略之间的近似平均 KL 散度用于估计本轮更新对策略的改变程度。实现在 stable_baselines3/ppo/ppo.py使用mean((exp(log_ratio) - 1) - log_ratio)近似计算clip_fractionPPO surrogate loss 中被裁剪超过clip_range阈值的比例均值统计abs(ratio - 1) clip_range的比例stable_baselines3/ppo/ppo.py。通常期望落在 0.10.3 之间clip_rangePPO surrogate loss 的当前裁剪系数critic_lossoff-policy 算法中 critic 函数损失通常为价值函数输出与 TD(0)时序差分估计之间的误差ent_coefSAC 的熵系数当前值ent_coef_lossSAC 熵系数的损失值entropy_loss熵损失均值策略平均熵的相反数explained_variance价值函数对回报方差的解释比例。判定标准ev0说明价值函数不如直接预测 0ev1为完美预测ev0说明比直接预测 0 更差。其计算实现于 stable_baselines3/common/utils.pyPPO/A2C 在更新结束时记录如 stable_baselines3/ppo/ppo.pylearning_rate当前学习率loss当前总损失值n_updates已执行的梯度更新次数各算法分别维护_n_updates计数并记录见 stable_baselines3/a2c/a2c.py 等policy_gradient_loss策略梯度损失当前值注意其绝对值本身意义不大value_losson-policy 算法中价值函数损失通常为价值函数输出与 Monte-Carlo 估计或 TD(lambda) 估计之间的误差std使用广义状态依赖探索gSDE时噪声的标准差当前值。十、实操建议如何选择输出格式结合以上原理给出常见场景的格式选择建议日常调试使用[stdout]或[stdout, log]HumanOutputFormat的 ASCII 表格便于直接肉眼观察训练趋势保存训练历史使用[csv]或[json]。CSV 便于 Excel/pandas直接读取JSONL 每行一个完整快照适合脚本化解析与断点续存可视化与超参分析使用[tensorboard]可以查看标量曲线、参数直方图并能通过HParam在 HPARAMS 标签页对比不同超参数组合组合输出configure(tmp_path, [stdout, csv, tensorboard])一次性同时覆盖终端、持久化文件与 TensorBoard 三种需求。另外可以善用环境变量不传folder/format_strings时configure会读取SB3_LOGDIR与SB3_LOG_FORMAT见 stable_baselines3/common/logger.py适合在脚本或 CI 中统一配置日志目录与格式。结语SB3 的 Logger 是一个设计简洁但能力完整的诊断输出框架configureset_logger提供灵活的覆盖机制五种输出格式分别满足终端可读、文件持久化与可视化分析的需求而eval/、rollout/、time/、train/四类指标则几乎覆盖了训练监控所需的全部信号。结合本文给出的源码位置stable_baselines3/common/logger.py、stable_baselines3/common/base_class.py、stable_baselines3/common/utils.py、stable_baselines3/common/callbacks.py、stable_baselines3/common/vec_env/vec_monitor.py你可以随时深入底层按需定制属于自己的训练监控方案。【免费下载链接】stable-baselines3PyTorch version of Stable Baselines, reliable implementations of reinforcement learning algorithms.项目地址: https://gitcode.com/GitHub_Trending/st/stable-baselines3创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表