ARTICLE DETAIL

资讯详情

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

pytorch-metric-learning 日志记录与模型保存实战:使用 logging_presets 的 HookContainer 完整指南

pytorch-metric-learning 日志记录与模型保存实战:使用 logging_presets 的 HookContainer 完整指南 人工智能机器学习深度学习计算机视觉【免费下载链接】pytorch-metric-learningThe easiest way to use deep metric learning in your application. Modular, flexible, and extensible. Written in PyTorch.项目地址https://gitcode.com/gh_mirrors/py/pytorch-metric-learning点击查看免费下载logging_presets是 pytorch-metric-learning 提供的开箱即用的训练钩子hook模块负责在训练过程中自动记录指标、验证模型、保存最优与最新检查点并支持基于验证指标的早停early stoppage。通过本文你将掌握如何在MetricLossOnly等 Trainer 与GlobalEmbeddingSpaceTester等 Tester 之间挂载这套钩子把训练与验证数据写入 CSV、SQLite 和 TensorBoard并学会用内置方法回溯 loss 与精度历史。全文以 docs/logging_presets.md 为骨架结合 src/pytorch_metric_learning/utils/logging_presets.py 源码与 tests/trainers/test_metric_loss_only.py 测试用例展开。一、模块定位与依赖安装logging_presets模块的核心价值在于不用自己编写训练/验证循环的状态管理逻辑只要把现成的钩子传给 Trainer 和 Tester就能自动完成三类工作记录数据训练与验证阶段的指标按 CSV、SQLite、TensorBoard 三种格式持久化验证与保存模型在每个 epoch 末运行验证并保存最新模型与验证指标最优模型早停当验证精度连续多个 epoch 不再提升时提前结束训练。该模块依赖record-keeper负责记录写入与tensorboard负责可视化两个第三方包可通过 pip 安装pip install record-keeper tensorboard源码 logging_presets.py 的注释也明确提示了这一点。需要注意如果未安装这些包get_record_keeper会打印警告并返回(None, None, None)此时get_hook_container会返回一个EmptyContainer见下文降级路径一节训练仍可进行但不会有任何日志与模型保存。二、快速上手挂载钩子的完整流程2.1 创建 record keeper 与钩子容器首先导入模块并创建记录器然后通过get_hook_container拿到一个HookContainer实例import pytorch_metric_learning.utils.logging_presets as LP log_folder, tensorboard_folder example_logs, example_tensorboard record_keeper, _, _ LP.get_record_keeper(log_folder, tensorboard_folder) hooks LP.get_hook_container(record_keeper) dataset_dict {val: val_dataset} model_folder example_saved_models其中get_record_keeper依次返回(record_keeper, record_writer, tensorboard_writer)第一个参数就是后续HookContainer所需的记录对象。它在内部用record_keeper.RecordWriter与torch.utils.tensorboard.SummaryWriter分别构建 CSV/SQLite 写入器和 TensorBoard 写入器见 logging_presets.py。2.2 创建 tester 与端到端钩子接着创建 tester并把end_of_testing_hook传入——这样每次测试完成后各数据切分的精度指标会被自动记录# Create the tester tester testers.GlobalEmbeddingSpaceTester(end_of_testing_hookhooks.end_of_testing_hook) end_of_epoch_hook hooks.end_of_epoch_hook(tester, dataset_dict, model_folder)注意end_of_epoch_hook是一个工厂函数必须先传入tester、dataset_dict、model_folder才能拿到真正可用的钩子它会返回一个闭包actual_hook见 logging_presets.py。2.3 创建 trainer 并开始训练trainer trainers.MetricLossOnly( models, optimizers, batch_size, loss_funcs, mining_funcs, train_dataset, samplersampler, end_of_iteration_hookhooks.end_of_iteration_hook, end_of_epoch_hookend_of_epoch_hook, ) trainer.train(num_epochsnum_epochs)运行结束后训练与验证数据会分别落到example_logsCSV/SQLite、example_tensorboardTensorBoard 事件文件模型与优化器则保存到example_saved_models目录。完整的可运行示例可参考仓库中的 examples/notebooks/MetricLossOnly.ipynb其Create the training and testing hooks一节与本流程一一对应以及 examples/notebooks/CascadedEmbeddings.ipynb、examples/notebooks/TwoStreamMetricLoss.ipynb 等其他 notebook。钩子何时被调用从 base_trainer.py 可以看到BaseTrainer在每个 iteration 结束时调用end_of_iteration_hook(self)在每个 epoch 结束时调用end_of_epoch_hook(self)且当end_of_epoch_hook返回False时训练提前终止——这正是早停机制的接入点。三、HookContainer 构造参数详解HookContainer的构造函数定义如下签名见 logging_presets.pyimport pytorch_metric_learning.utils.logging_presets as LP LP.HookContainer( record_keeper, record_group_name_prefixNone, primary_metricmean_average_precision_at_r, validation_split_nameval, save_modelsTrue, log_freq50, )参数类型/默认值作用record_keeperrecord-keeper对象必填。记录写入的底层对象来自get_record_keeper的返回值。record_group_name_prefixstr默认None所有记录名与 TensorBoard tag 的前缀。例如设置为exp1时记录组名会以exp1_开头见base_record_group_name实现logging_presets.py。primary_metricstr默认mean_average_precision_at_r用于判定最优检查点的精度指标可选值mean_average_precision_at_r、r_precision、precision_at_1、NMI。该指标必须包含在 tester 的AccuracyCalculator计算的指标集合中否则end_of_epoch_hook会抛出ValueError见 logging_presets.py。测试 test_metric_loss_only.py 中就通过primary_metricprecision_at_1配合AccuracyCalculator(include(precision_at_1, AMI))使用。validation_split_namestr默认val验证集在dataset_dict中的键名用于定位早停与最优模型判定所用切分。save_modelsbool默认True是否保存模型。为False时save_models方法直接跳过保存动作见 logging_presets.py。log_freqint默认50每隔多少个 iteration 记录一次数据。end_of_iteration_hook只在trainer.iteration % log_freq 0时执行写入见 logging_presets.py。四、三个核心钩子函数4.1 end_of_iteration_hook记录每次迭代的状态该钩子直接传入 Trainer在每个 iteration 的结尾按log_freq节流把以下对象的状态写入记录损失历史loss_histories与损失权重loss_weights来自trainer.loss_tracker损失函数loss_funcs递归记录其中的torch.nn.Module与dict子对象挖掘函数mining_funcs模型models优化器optimizers额外记录当前学习率见optimizer_custom_attr_funclogging_presets.py。具体记录清单见 logging_presets.py。每条记录都以trainer.get_global_iteration()作为时间戳因此多次实验中可以用全局迭代数对齐记录。测试用例 test_metric_loss_only.py 验证了这一点在log_freq2、iterations_per_epoch10、num_epochs2的设置下metric_loss_NTXentLoss等表的记录条数恰好等于num_epochs * iterations_per_epoch / log_freq且record_keeper.table_exists()返回True。4.2 end_of_epoch_hook验证 保存 早停这是唯一的工厂函数需要先传入参数来定制钩子参数默认值说明tester必填一个 tester 对象用于在 epoch 末执行验证。dataset_dict必填从切分名到 PyTorch Dataset 的字典例如{train: train_dataset, val: val_dataset}。model_folder必填模型、优化器保存目录。目录不存在时会自动创建logging_presets.py。test_interval1每隔多少个 epoch 运行一次验证。patienceNone若不为None当epoch - best_epoch patience时提前终止训练。splits_to_evalNone控制验证时使用哪些查询/参考切分详见 docs/testers.md。test_collate_fnNone测试阶段 DataLoader 使用的 collate 函数。返回的钩子在每个 epoch 末执行如下流程对应save_models_and_evallogging_presets.py调用tester.test(dataset_dict, epoch, trunk, embedder, ...)计算各切分的精度读取上一次最优 epoch判断本次是否为新的最优is_new_best_accuracy当精度更高或尚无最优记录时判定为新的最优record_keeper.save_records()落盘所有记录调用trainer.step_lr_plateau_schedulers(curr_accuracy)把当前验证精度喂给 ReduceLROnPlateau 类学习率调度器保存最新模型后缀为当前 epoch若产生新的最优再保存一份后缀为best{epoch}的最优模型并删除上一份best{prev_best_epoch}若patience触发返回False让BaseTrainer提前结束训练base_trainer.py并打印 Validation accuracy has plateaued. Exiting.。4.3 end_of_testing_hook记录验证结果该钩子直接传入 Tester。每次测试完成后它遍历tester.all_accuracies中的每个切分把精度字典和最优信息best_epoch、best_accuracy写入记录logging_presets.py。记录组的命名形如accuracies_normalized_GlobalEmbeddingSpaceTester_level_0_VAL_vs_self测试 test_metric_loss_only.py 对这一命名规则做了精确断言。五、训练过程数据查询方法训练完成后可以用HookContainer提供的方法把记录重新查出来用于画曲线或复现最优检查点。5.1 损失历史# 返回字典loss 名称 - 数值列表 loss_histories hooks.get_loss_history() # 只取指定 loss例如只返回 total_loss loss_histories hooks.get_loss_history(loss_names[total_loss])实现上get_loss_history直接对 SQLite 中名为loss_histories的表执行SELECT查询logging_presets.py未记录过该表时返回空字典。5.2 精度历史# 第一个参数是 tester 对象第二个是切分名 # 返回的字典包含键 epoch 与主指标名值均为列表 acc_histories hooks.get_accuracy_history(tester, val) # 返回所有精度指标的历史 acc_histories hooks.get_accuracy_history(tester, val, return_all_metricsTrue) # 只返回指定的指标集合 acc_histories hooks.get_accuracy_history(tester, val, metrics[AMI, NMI])get_accuracy_history会先尝试带平均标记averageTrue的指标键名失败时回退到不带平均的键名最终仍查不到则抛出KeyError见try_keyslogging_presets.py。这也解释了为什么文档示例中即使primary_metric是单个指标返回字典里也总是同时带有epoch键。5.3 其他常用查询方法源码中还有一组配套的查询/工具方法均被测试 test_metric_loss_only.py 覆盖get_curr_primary_metric(tester, split_name)取当前 epoch 的主指标值get_accuracies_of_epoch(tester, split_name, epoch)取指定 epoch 的精度记录get_accuracies_of_best_epoch(tester, split_name)取历史最优 epoch 的完整精度记录及其指标键名get_best_epoch_and_accuracy(tester, split_name)返回(best_epoch, best_accuracy)load_latest_saved_models(trainer, model_folder, deviceNone, bestFalse)从model_folder恢复最新的或最优的模型、优化器等saveable_trainer_objects状态返回resume_epoch 1可用于断点续训logging_presets.py。可保存的对象类别在__init__中固定为models、optimizers、lr_schedulers、loss_funcs、mining_funcslogging_presets.py。六、源码级补充命名规则与降级路径6.1 记录组命名record_group_name(tester, split_name)生成的记录组名由三部分组成logging_presets.py基础前缀base_record_group_name可选的record_group_name_prefixtester.description_suffixes(accuracies)例如accuracies_normalized_GlobalEmbeddingSpaceTester_level_0查询切分名大写如VAL参考切分名按字典序排序后以_and_连接并大写如TRAIN_and_VAL若与查询切分相同则简写为self。因此默认配置下验证集记录组名为accuracies_normalized_GlobalEmbeddingSpaceTester_level_0_VAL_vs_self。若 RecordKeeper 的hash_map中已存在相同记录组名则会复用已有 key。6.2 缺依赖时的降级路径get_hook_container在record_keeper为None即依赖包未安装时不会报错而是返回EmptyContainerlogging_presets.py。EmptyContainer的三个钩子均为空实现end_of_epoch_hook直接返回Noneend_of_iteration_hook与end_of_testing_hook为Nonelogging_presets.py训练照常进行但不产生任何记录日志中会给出 There wont be any logging or model saving. 的警告logging_presets.py。七、实战要点与常见配置建议综合文档、源码与测试使用logging_presets时有以下几点值得注意主指标一致性HookContainer的primary_metric必须在 tester 的AccuracyCalculator计算范围内。测试 test_metric_loss_only.py 通过AccuracyCalculator(include(precision_at_1, AMI))与primary_metricprecision_at_1配对使用examples/notebooks/MetricLossOnly.ipynb 则用默认的mean_average_precision_at_r配合AccuracyCalculator(kmax_bin_count)。早停与调度器联动patience早停基于验证集主指标同时每轮验证后都会调用step_lr_plateau_schedulers因此使用 ReduceLROnPlateau 时无需再手动 step。记录频率权衡log_freq越小记录越密集、TensorBoard 曲线越平滑但会带来更多 I/Olog_freq50是文档示例采用的默认值。恢复训练load_latest_saved_models配合latest_versioncommon_functions.py可以按文件名中的 epoch 序号自动定位最新或最优检查点并恢复全部可训练状态实现断点续训。多切分验证通过splits_to_eval可让验证集以训练集或其他切分为参考集进行检索评测具体语义见 docs/testers.md。八、总结logging_presets把记录、验证、保存、早停这一套训练基础设施收敛为三个钩子函数让使用者无需关心 CSV/SQLite/TensorBoard 的写入细节。从源码看其核心是围绕record_keeper的update_records与 SQLite 查询方法展开配合BaseTrainer在迭代/轮次末尾对钩子的固定调用点形成了完整的训练闭环仓库中的单元测试 test_metric_loss_only.py 则从记录条数、记录组命名、最优 epoch 回溯等多个角度验证了这一机制的正确性。对于任何希望开箱即用地获得训练可视化与模型版本管理能力的 pytorch-metric-learning 使用者这套钩子都是推荐的默认方案。赞分享人工智能机器学习深度学习计算机视觉【免费下载链接】pytorch-metric-learningThe easiest way to use deep metric learning in your application. Modular, flexible, and extensible. Written in PyTorch.项目地址https://gitcode.com/gh_mirrors/py/pytorch-metric-learning点击查看免费下载相关推荐如何使用PyTorch Metric Learning构建高效双流度量学习模型完整指南如何使用PyTorch Metric Learning构建高效双流度量学习模型完整指南 PyTorch Metric Learning是一个模块化、灵活且可扩人工智能机器学习深度学习计算机视觉PyTorch Metric Learning 使用教程PyTorch Metric Learning 使用教程 项目介绍 PyTorch Metric Learning 是一个用于深度度量学习的开源库旨在简化在人工智能机器学习深度学习计算机视觉Symfony/Translation调试日志使用GraylogGELF记录日志的完整指南Symfony/Translation调试日志使用GraylogGELF记录日志的完整指南 在PHP多语言应用开发中symfony/translation国际化后端上一篇如何在不安装 Node.js 的主机上用 Docker 官方镜像运行 Bruno CLI 集合并产出 JUnit 报告下一篇OpenResume开发效率插件VSCode扩展与代码片段创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表