如何保存Optuna(PyTorch)第4次最优试验对应的模型?
用Optuna(PyTorch)获取最优参数并保存对应训练模型的实现方法
一、获取最优试验参数
你可以通过两种方式获取目标参数:
方式1:直接获取全局最优参数
import optuna # 假设你的Optuna Study对象已完成5次试验 best_params = study.best_params # 提取所需参数 best_lr = best_params["lr"] best_optimizer = best_params["optimizer"]
方式2:指定获取第4次试验的参数(索引从0开始)
# 直接定位第4次试验(对应索引3) target_trial = study.get_trials()[3] best_params = target_trial.params best_lr = best_params["lr"] best_optimizer = best_params["optimizer"]
二、保存最优试验对应的训练模型
有两种实用实现方式,按需选择:
方式1:试验过程中实时保存最优模型
在Optuna的目标函数内,每次训练完成后判断当前试验是否为最优,满足条件则保存模型及相关参数:
import torch def objective(trial): # 初始化模型与优化器 model = YourCustomModel() # 替换为你的模型类 lr = trial.suggest_float("lr", 1e-5, 1e-1, log=True) optimizer_name = trial.suggest_categorical("optimizer", ["Adam", "SGD"]) if optimizer_name == "Adam": optimizer = torch.optim.Adam(model.parameters(), lr=lr) else: optimizer = torch.optim.SGD(model.parameters(), lr=lr, momentum=0.9) # 执行训练逻辑(替换为你的训练循环代码) train_loss = train_loop(model, optimizer, train_loader, epochs=10) # 评估模型性能(替换为你的评估逻辑) val_score = evaluate(model, val_loader) # 根据优化方向判断是否保存模型:最小化指标用<=,最大化指标用>= if val_score <= study.best_value: torch.save( { "model_state_dict": model.state_dict(), "optimizer_state_dict": optimizer.state_dict(), "lr": lr, "optimizer": optimizer_name, "validation_score": val_score }, "best_model_from_trial.pth" ) return val_score
方式2:拿到最优参数后重新训练并保存
如果之前的试验未保存模型,可通过最优参数重新初始化模型和优化器,完成训练后保存:
import torch # 1. 加载最优参数 best_params = study.best_params best_lr = best_params["lr"] best_opt_name = best_params["optimizer"] # 2. 初始化模型与优化器 model = YourCustomModel() if best_opt_name == "Adam": optimizer = torch.optim.Adam(model.parameters(), lr=best_lr) else: optimizer = torch.optim.SGD(model.parameters(), lr=best_lr, momentum=0.9) # 3. 用最优参数执行完整训练(和试验中的训练逻辑保持一致) train_loop(model, optimizer, train_loader, epochs=10) # 4. 保存模型及参数 torch.save( { "model_state_dict": model.state_dict(), "optimizer_state_dict": optimizer.state_dict(), "lr": best_lr, "optimizer": best_opt_name, "study_best_value": study.best_value }, "retrained_best_model.pth" )
注意事项
- 保存的
.pth文件包含模型权重、优化器状态、学习率和优化器名称,后续可直接加载用于推理或恢复训练。 - 判断最优的逻辑需与创建Study时指定的
direction一致(比如optuna.create_study(direction="minimize")对应最小化指标)。
内容的提问来源于stack exchange,提问作者Nitya Rakhe
相关产品推荐
相关产品推荐

