如何将Ray Tune最佳模型的每轮指标保存至文件?
解决Ray Tune最佳运行轮次指标保存问题
要把最佳运行的每轮val_loss和val_acc保存到文件,你可以直接从best_result中提取历史指标日志,然后用Python工具将数据写入文件。以下是具体步骤和修改后的代码:
核心思路
Ray Tune的Result对象(即best_result)内置了获取全轮次指标的方法,不需要依赖CLIReporter实现持久化——它仅用于终端打印指标展示。你可以选择用pandas生成规整的CSV文件,或者用原生Python写入文本文件。
修改后的完整代码
from ray import tune from ray import air from ray.air.config import RunConfig from ray.tune.search.hyperopt import HyperOptSearch from hyperopt import fmin, hp, tpe, Trials, space_eval, STATUS_OK import os import pandas as pd # 需先安装:pip install pandas config_dict = { "c_hidden": tune.choice([64]), "dp_rate_linear": tune.choice([0.1]), # 可改为quniform并指定三元组范围 "num_layers":tune.choice([3]), "dp_rate":tune.choice([0.3]) } hyperopt_search = HyperOptSearch( metric="val_loss", mode="min") # points_to_evaluate=current_best_params) # 初始化Tuner tuner = tune.Tuner( tune.with_resources(train_fn, {"gpu": 1}), tune_config=tune.TuneConfig(num_samples=1, search_alg=hyperopt_search), param_space=config_dict, run_config=RunConfig(local_dir='/home/runs/') ) results = tuner.fit() best_result = results.get_best_result(metric="val_loss", mode="min") # 提取并保存轮次指标 # 1. 用pandas生成CSV(推荐,方便后续绘图工具直接读取) metrics_df = best_result.metrics_dataframe() # 筛选需要的列 filtered_df = metrics_df[["training_iteration", "val_loss", "val_acc"]] save_path = "/home/runs/best_run_metrics.csv" filtered_df.to_csv(save_path, index=False) print(f"轮次指标已保存到: {save_path}") # 后续模型加载代码不变 best_checkpoint = best_result.checkpoint path = os.path.join(str(best_checkpoint.to_directory()), "ray_ckpt3") model = GraphLevelGNN.load_from_checkpoint(path) print(f"最佳模型路径: {path}")
替代方案(不依赖pandas)
如果不想安装pandas,可用原生Python写入文本文件:
# 提取历史指标列表 history = best_result.metrics_history # 写入文本文件 with open("/home/runs/best_run_metrics.txt", "w") as f: # 写入表头 f.write("training_iteration,val_loss,val_acc\n") for metrics in history: iter_num = metrics.get("training_iteration", "") val_loss = metrics.get("val_loss", "") val_acc = metrics.get("val_acc", "") f.write(f"{iter_num},{val_loss},{val_acc}\n")
关键说明
best_result.metrics_dataframe()会自动整理所有训练轮次的指标数据,格式规整,适合直接用于绘图工具(如matplotlib、seaborn)- 若你的
train_fn中没有在每轮迭代时通过tune.report()上报training_iteration、val_loss、val_acc,需先确保这些指标被正确上报,否则无法提取到对应数据
内容的提问来源于stack exchange,提问作者Slowat_Kela
相关产品推荐
相关产品推荐

