添加模型优化方法和参数列表

This commit is contained in:
haotian 2025-03-17 11:01:14 +08:00
parent 57d486b0c9
commit a9bc70165e
6 changed files with 949 additions and 21 deletions

View File

@ -7,7 +7,7 @@
- python main.py
- 运行ip和端口在config/config.yaml中配置
- 当前python环境名为trt
- 在ip:port/docs可以看到可视化接口文档
- 在ip:port/docs可以看到可视化接口文档.即http://10.0.0.202:8992/docs
## 3. 安装python环境
- 使用conda创建一个新的python环境
- pip 换源 清华源
@ -16,6 +16,8 @@
## 4. 测试接口
注: 10.0.0.202为当前本机ip

View File

@ -631,12 +631,296 @@ Error Response:
}
```
### 2.9 模型优化 -- 未实现
## 3. 模型优化
### 3.1 获取模型优化方法列表
```http
GET /optimize/methods
Response:
{
"status": "success",
"methods": [
{
"name": "GridSearchCV",
"description": "网格搜索交叉验证,通过穷举搜索指定参数网格中的所有可能组合,找到最优参数。",
"advantages": [
"全面探索参数空间,保证找到指定范围内的最优解。",
"实现简单,易于理解和使用。",
"支持并行计算,可加速搜索过程。"
],
"disadvantages": [
"计算成本高,参数空间增大时搜索时间呈指数级增长。",
"需要预先定义参数搜索范围,可能错过最优解。"
],
"applicable_scenarios": [
"参数较少通常不超过3-4个的模型调优。",
"计算资源充足的情况。",
"需要全面了解参数影响的场景。"
]
},
{
"name": "RandomizedSearchCV",
"description": "随机搜索交叉验证,从指定的分布中随机采样参数组合进行评估,而非穷举所有可能。",
"advantages": [
"计算效率高,特别是在高维参数空间中。",
"在相同时间内可以探索更大的参数空间。",
"通常能在较短时间内找到接近最优的解。"
],
"disadvantages": [
"不保证找到全局最优解。",
"结果的可重复性较低,除非固定随机种子。"
],
"applicable_scenarios": [
"参数空间较大的模型调优。",
"计算资源有限的情况。",
"初步探索参数影响的场景。"
]
},
// 其他优化方法...
]
}
```
### 3.2 获取优化方法详情
```http
GET /optimize/method/{method_name}
## 3. 系统监控
### 3.1 获取资源使用情况
Response:
{
"status": "success",
"method": {
"name": "GridSearchCV",
"description": "网格搜索交叉验证,通过穷举搜索指定参数网格中的所有可能组合,找到最优参数。",
"advantages": [
"全面探索参数空间,保证找到指定范围内的最优解。",
"实现简单,易于理解和使用。",
"支持并行计算,可加速搜索过程。"
],
"disadvantages": [
"计算成本高,参数空间增大时搜索时间呈指数级增长。",
"需要预先定义参数搜索范围,可能错过最优解。"
],
"applicable_scenarios": [
"参数较少通常不超过3-4个的模型调优。",
"计算资源充足的情况。",
"需要全面了解参数影响的场景。"
],
"parameters": [
{
"name": "estimator",
"type": "estimator object",
"default": null,
"description": "要优化的模型,必须实现 fit 和 score 方法。"
},
{
"name": "param_grid",
"type": "dict or list of dicts",
"default": null,
"description": "参数网格,指定要搜索的参数名称和可能的值。"
},
{
"name": "scoring",
"type": "str, callable, list, tuple or dict",
"default": null,
"description": "评分标准,用于评估模型性能。"
},
// 其他参数...
]
}
}
```
### 3.3 执行模型优化
```http
POST /optimize/run
Content-Type: application/json
Request:
{
"run_id": "7970364d490f4e0aa0375c2db26215f3",
"method": "GridSearchCV",
"parameters": {
"param_grid": {
"max_depth": [3, 5, 7],
"learning_rate": [0.01, 0.1, 0.2],
"n_estimators": [100, 200, 300]
},
"scoring": "accuracy",
"cv": 5,
"n_jobs": -1,
"verbose": 1
},
"data_path": "dataset/dataset_processed/train_val_combined.csv",
"output_dir": "optimized_models/",
"experiment_name": "xgboost_optimization"
}
Response:
{
"status": "success",
"optimization": {
"id": "opt_20250220_001",
"run_id": "7970364d490f4e0aa0375c2db26215f3",
"method": "GridSearchCV",
"original_model": "XGBClassifier",
"best_params": {
"max_depth": 5,
"learning_rate": 0.1,
"n_estimators": 200
},
"best_score": 0.968,
"cv_results_file": "optimized_models/opt_20250220_001/cv_results.csv",
"optimized_model_id": "8970364d490f4e0aa0375c2db26215f4",
"execution_time": "45m 23s",
"status": "completed",
"timestamp": "2025-02-20 10:45:30"
}
}
Error Response:
{
"status": "error",
"message": "模型优化失败",
"details": {
"error_type": "ValueError",
"error_message": "参数配置无效"
}
}
```
### 3.4 获取优化任务列表
```http
GET /optimize/tasks?experiment_name={experiment_name}&page={page}&page_size={page_size}
Response:
{
"status": "success",
"tasks": [
{
"id": "opt_20250220_001",
"run_id": "7970364d490f4e0aa0375c2db26215f3",
"method": "GridSearchCV",
"original_model": "XGBClassifier",
"best_score": 0.968,
"status": "completed",
"start_time": "2025-02-20 09:30:15",
"end_time": "2025-02-20 10:45:30",
"optimized_model_id": "8970364d490f4e0aa0375c2db26215f4"
},
{
"id": "opt_20250220_002",
"run_id": "7970364d490f4e0aa0375c2db26215f5",
"method": "RandomizedSearchCV",
"original_model": "RandomForestClassifier",
"best_score": 0.942,
"status": "completed",
"start_time": "2025-02-20 11:15:20",
"end_time": "2025-02-20 12:05:45",
"optimized_model_id": "8970364d490f4e0aa0375c2db26215f6"
}
],
"total_count": 2,
"page": 1,
"page_size": 10
}
```
### 3.5 获取优化任务详情
```http
GET /optimize/task/{task_id}
Response:
{
"status": "success",
"task": {
"id": "opt_20250220_001",
"run_id": "7970364d490f4e0aa0375c2db26215f3",
"method": "GridSearchCV",
"original_model": "XGBClassifier",
"parameters": {
"param_grid": {
"max_depth": [3, 5, 7],
"learning_rate": [0.01, 0.1, 0.2],
"n_estimators": [100, 200, 300]
},
"scoring": "accuracy",
"cv": 5,
"n_jobs": -1,
"verbose": 1
},
"best_params": {
"max_depth": 5,
"learning_rate": 0.1,
"n_estimators": 200
},
"best_score": 0.968,
"cv_results_summary": {
"mean_fit_time": 12.5,
"mean_score_time": 0.8,
"param_combinations": 27,
"score_range": [0.912, 0.968]
},
"cv_results_file": "optimized_models/opt_20250220_001/cv_results.csv",
"optimized_model_id": "8970364d490f4e0aa0375c2db26215f4",
"execution_time": "45m 23s",
"status": "completed",
"start_time": "2025-02-20 09:30:15",
"end_time": "2025-02-20 10:45:30"
}
}
```
### 3.6 取消优化任务
```http
POST /optimize/task/{task_id}/cancel
Response:
{
"status": "success",
"message": "优化任务已取消",
"task_id": "opt_20250220_003",
"cancel_time": "2025-02-20 14:35:22"
}
Error Response:
{
"status": "error",
"message": "无法取消任务",
"details": {
"reason": "任务已完成或不存在"
}
}
```
### 3.7 删除优化任务
```http
DELETE /optimize/task/{task_id}
Response:
{
"status": "success",
"message": "优化任务已删除",
"details": {
"task_id": "opt_20250220_001",
"deleted_files": [
"optimized_models/opt_20250220_001/cv_results.csv",
"optimized_models/opt_20250220_001/optimization_log.txt"
]
}
}
Error Response:
{
"status": "error",
"message": "删除任务失败",
"details": {
"reason": "任务不存在或正在运行中"
}
}
```
## 4. 系统监控
### 4.1 获取资源使用情况
```http
GET /system/resources
@ -730,7 +1014,7 @@ Error Response:
}
```
### 3.2 获取训练历史
### 4.2 获取训练历史
```http
GET /system/history?page=1&page_size=10&start_time=2025-02-01&end_time=2025-02-19&status=completed&experiment_name=breast_cancer_classification
@ -793,7 +1077,7 @@ Error Response:
}
```
### 3.3 获取系统中训练状态 ---- 未完成, 等开发系统后台时再实现.
### 4.3 获取系统中训练状态 ---- 未完成, 等开发系统后台时再实现.
```http
GET /model/train/status/{task_id}
@ -816,7 +1100,7 @@ Response:
}
```
### 3.4 获取系统日志
### 4.4 获取系统日志
```http
GET /system/logs?level=error&start_time=2025-02-19T00:00:00&end_time=2025-02-19T23:59:59&module=training&page=1&page_size=20
@ -876,9 +1160,9 @@ Error Response:
}
```
## 4. 系统后台整体实现
## 5. 系统后台整体实现
### 4.1 系统架构
### 5.1 系统架构
```
MLPlatform/
├── api/ # API接口层
@ -901,14 +1185,14 @@ MLPlatform/
└── main.py # 主程序入口
```
### 4.2 技术栈
### 5.2 技术栈
- FastAPI: Web框架
- MLflow: 模型管理和实验跟踪
- PyTorch/Scikit-learn: 机器学习框架
- Pydantic: 数据验证
- Uvicorn: ASGI服务器
### 4.3 主要功能
### 5.3 主要功能
1. 异步任务处理
- 支持多个模型同时训练
- 后台任务状态监控
@ -929,7 +1213,7 @@ MLPlatform/
- 请求限流
- 参数验证
### 4.4 性能优化
### 5.4 性能优化
1. 数据处理
- 数据流式处理
- 缓存机制
@ -945,8 +1229,8 @@ MLPlatform/
- 资源使用预警
- 自动清理机制
## 5. 前端设计
### 5.1 技术栈
## 6. 前端设计
### 6.1 技术栈
- Vue3: 前端框架
- TypeScript: 编程语言
- Element Plus: UI组件库
@ -955,7 +1239,7 @@ MLPlatform/
- Pinia: 状态管理
- Vue Router: 路由管理
### 5.2 目录结构
### 6.2 目录结构
```
frontend/
├── src/
@ -978,7 +1262,7 @@ frontend/
└── package.json # 项目配置
```
### 5.3 页面设计
### 6.3 页面设计
1. 数据处理模块
- 数据集列表页
- 展示所有可用数据集
@ -1021,7 +1305,7 @@ frontend/
- 日志级别筛选
- 日志搜索功能
### 5.4 交互设计
### 6.4 交互设计
1. 数据处理流程
```mermaid
graph LR
@ -1042,7 +1326,7 @@ frontend/
E --> F[查看结果]
```
### 5.5 组件设计
### 6.5 组件设计
1. 通用组件
- 数据表格组件
- 图表展示组件
@ -1056,7 +1340,7 @@ frontend/
- 评估结果展示组件
- 系统监控面板组件
### 5.6 状态管理
### 6.6 状态管理
1. 全局状态
- 用户配置信息
- 系统运行状态
@ -1067,7 +1351,7 @@ frontend/
- 模型训练状态
- 系统监控数据
### 5.7 性能优化
### 6.7 性能优化
1. 数据处理
- 大数据分页加载
- 数据缓存机制
@ -1083,7 +1367,7 @@ frontend/
- 数据分片处理
- WebWorker处理大数据
### 5.8 错误处理
### 6.8 错误处理
1. 全局错误处理
- API请求错误
- 组件渲染错误

View File

@ -0,0 +1,73 @@
from typing import Dict, List, Optional
from pathlib import Path
import logging
from datetime import datetime
import yaml
import mlflow
from mlflow.tracking import MlflowClient
class OptimizeManager:
"""模型优化管理类"""
def __init__(self, config: Dict = None) -> None:
self.config = config or {}
self.logger = logging.getLogger(__name__)
self._setup_logging()
self._load_method_config()
self._load_parameter_config()
# 初始化MLflow客户端
self.mlflow_uri = self.config.get('mlflow_uri', 'http://127.0.1.202:5000')
mlflow.set_tracking_uri(self.mlflow_uri)
self.client = MlflowClient()
pass
def _setup_logging(self):
# 相对于程序运行目录
log_dir = Path('.log')
log_dir.mkdir(exist_ok=True)
file_handler = logging.FileHandler(
log_dir / f'optimize_manager_{datetime.now():%Y%m%d_%H%M%S}.log'
)
file_handler.setFormatter(
logging.Formatter('%(asctime)s - %(name)s - %(levelname)s - %(message)s')
)
self.logger.addHandler(file_handler)
self.logger.setLevel(logging.INFO)
def _load_method_config(self):
"""加载模型优化方法配置文件"""
config_path = Path('optimize/method.yaml')
if not config_path.exists():
raise FileNotFoundError(f"Method config file not found at {config_path}")
with open(config_path, 'r', encoding='utf-8') as f:
config_optimize = yaml.safe_load(f)
self.logger.info("Successfully loader optimize method config")
config = config_optimize
return config
def _load_parameter_config(self):
config_path = Path('optimize/parameter.yaml')
if not config_path.exists():
self.logger.error("error load optimize parameter config")
raise FileNotFoundError(f"Parameter config file not found at {config_path}")
with open(config_path, 'r', encoding='utf-8') as f:
config_parameter = yaml.safe_load(f)
self.logger.info("Successfullu load optimize parameter config")
config = config_parameter
return config

160
optimize/method.yaml Normal file
View File

@ -0,0 +1,160 @@
optimization_methods:
GridSearchCV:
description: "网格搜索交叉验证,通过穷举搜索指定参数网格中的所有可能组合,找到最优参数。"
advantages:
- "全面探索参数空间,保证找到指定范围内的最优解。"
- "实现简单,易于理解和使用。"
- "支持并行计算,可加速搜索过程。"
disadvantages:
- "计算成本高,参数空间增大时搜索时间呈指数级增长。"
- "需要预先定义参数搜索范围,可能错过最优解。"
applicable_scenarios:
- "参数较少通常不超过3-4个的模型调优。"
- "计算资源充足的情况。"
- "需要全面了解参数影响的场景。"
RandomizedSearchCV:
description: "随机搜索交叉验证,从指定的分布中随机采样参数组合进行评估,而非穷举所有可能。"
advantages:
- "计算效率高,特别是在高维参数空间中。"
- "在相同时间内可以探索更大的参数空间。"
- "通常能在较短时间内找到接近最优的解。"
disadvantages:
- "不保证找到全局最优解。"
- "结果的可重复性较低,除非固定随机种子。"
applicable_scenarios:
- "参数空间较大的模型调优。"
- "计算资源有限的情况。"
- "初步探索参数影响的场景。"
BayesianOptimization:
description: "贝叶斯优化,基于先验观测构建代理模型(通常是高斯过程),指导后续参数搜索方向。"
advantages:
- "搜索效率高,能快速收敛到最优解。"
- "适合计算成本高的目标函数优化。"
- "充分利用历史评估信息,减少无效搜索。"
disadvantages:
- "实现复杂度高,需要理解贝叶斯统计和高斯过程。"
- "对初始点和先验分布敏感。"
- "在高维空间中效率可能下降。"
applicable_scenarios:
- "计算资源有限但模型评估成本高的场景。"
- "需要高效参数调优的生产环境。"
- "连续参数空间的优化问题。"
HyperOpt:
description: "基于贝叶斯优化和树结构Parzen估计器TPE的超参数优化库支持复杂搜索空间定义。"
advantages:
- "支持条件参数空间,适合复杂模型结构优化。"
- "搜索效率高,特别是在离散和混合参数空间。"
- "易于与分布式计算框架集成。"
disadvantages:
- "学习曲线较陡API设计较为复杂。"
- "对搜索空间定义的依赖性强。"
applicable_scenarios:
- "复杂模型架构的优化,如神经网络结构搜索。"
- "包含条件依赖的参数优化。"
- "需要分布式计算支持的大规模优化任务。"
Optuna:
description: "现代超参数优化框架,结合多种采样算法和修剪机制,高效搜索最优参数。"
advantages:
- "提供多种先进的采样器如TPE、CMA-ES等。"
- "支持提前终止效果不佳的试验,提高搜索效率。"
- "可视化工具丰富,便于分析优化过程。"
- "支持分布式优化和数据库存储。"
disadvantages:
- "配置复杂度较高,需要一定学习成本。"
- "对于简单问题可能过于复杂。"
applicable_scenarios:
- "大规模模型调优项目。"
- "需要长期运行和管理的优化任务。"
- "需要可视化分析参数影响的场景。"
BOHB:
description: "结合贝叶斯优化和Hyperband的多保真度优化算法平衡探索与利用。"
advantages:
- "结合了贝叶斯优化的高效搜索和Hyperband的多保真度评估。"
- "能够处理计算预算有限的大规模优化问题。"
- "适合深度学习模型的训练优化。"
disadvantages:
- "实现复杂,需要专门的库支持。"
- "对初始配置敏感。"
applicable_scenarios:
- "深度学习模型的超参数优化。"
- "训练时间长的模型优化。"
- "需要在有限资源下获得最佳性能的场景。"
GeneticAlgorithm:
description: "基于进化理论的优化算法,通过选择、交叉和变异操作迭代优化参数。"
advantages:
- "能够处理非凸、不连续的复杂参数空间。"
- "不需要目标函数的梯度信息。"
- "易于并行化,适合分布式计算。"
disadvantages:
- "收敛速度可能较慢,需要较多的函数评估。"
- "结果的可重复性较低,除非固定随机种子。"
- "参数设置(如种群大小、变异率等)对性能影响大。"
applicable_scenarios:
- "复杂、非凸的参数空间优化。"
- "目标函数难以求导的场景。"
- "可以接受近似最优解的问题。"
ParticleSwarmOptimization:
description: "基于群体智能的优化算法,模拟鸟群觅食行为,通过粒子间信息共享寻找最优解。"
advantages:
- "实现简单,参数较少。"
- "收敛速度快,特别是在连续参数空间。"
- "不需要目标函数的梯度信息。"
disadvantages:
- "容易陷入局部最优。"
- "对参数初始化敏感。"
- "在高维空间中效率可能下降。"
applicable_scenarios:
- "连续参数空间的优化问题。"
- "计算资源有限但需要快速收敛的场景。"
- "目标函数复杂但评估成本不高的问题。"
EarlyStopping:
description: "提前停止策略,监控模型在验证集上的性能,当性能不再提升时停止训练,防止过拟合。"
advantages:
- "减少训练时间,提高计算效率。"
- "有效防止模型过拟合。"
- "自动确定最佳训练轮次。"
disadvantages:
- "可能过早停止训练,错过性能提升的机会。"
- "需要额外的验证数据集。"
- "对监控指标和停止条件的选择敏感。"
applicable_scenarios:
- "深度学习模型训练。"
- "容易过拟合的复杂模型。"
- "计算资源有限的训练环境。"
LearningRateScheduler:
description: "学习率调度器,根据预定策略或模型性能动态调整优化器的学习率。"
advantages:
- "加速模型收敛,提高训练效率。"
- "帮助模型跳出局部最优。"
- "在训练后期提高模型精度。"
disadvantages:
- "需要额外调整学习率策略的超参数。"
- "不同问题可能需要不同的调度策略。"
applicable_scenarios:
- "深度学习模型训练。"
- "训练过程不稳定或收敛缓慢的模型。"
- "需要精细调优的高精度模型。"
ModelCheckpoint:
description: "模型检查点,定期保存训练过程中的模型状态,确保能够恢复最佳模型。"
advantages:
- "保存训练过程中的最佳模型。"
- "支持训练中断后的恢复。"
- "便于模型部署和后续分析。"
disadvantages:
- "增加存储开销。"
- "可能影响训练速度,特别是模型较大时。"
applicable_scenarios:
- "长时间训练的大型模型。"
- "需要保证训练结果可复现的场景。"
- "计算环境不稳定,可能中断的训练任务。"

405
optimize/parameter.yaml Normal file
View File

@ -0,0 +1,405 @@
optimization_methods:
GridSearchCV:
description: "网格搜索交叉验证,通过穷举搜索指定参数网格中的所有可能组合,找到最优参数。"
parameters:
- name: "estimator"
type: "estimator object"
default: null
description: "要优化的模型,必须实现 fit 和 score 方法。"
- name: "param_grid"
type: "dict or list of dicts"
default: null
description: "参数网格,指定要搜索的参数名称和可能的值。"
- name: "scoring"
type: "str, callable, list, tuple or dict"
default: null
description: "评分标准,用于评估模型性能。"
- name: "n_jobs"
type: "int"
default: null
description: "并行运行的作业数量,-1 表示使用所有处理器。"
- name: "cv"
type: "int, cross-validation generator or iterable"
default: 5
description: "交叉验证策略,指定折数或自定义分割方式。"
- name: "verbose"
type: "int"
default: 0
description: "控制详细程度,值越大输出信息越详细。"
- name: "refit"
type: "bool, str or callable"
default: true
description: "是否使用最佳参数重新拟合模型,如果为字符串则指定用于重新拟合的评分指标。"
- name: "pre_dispatch"
type: "int or str"
default: "2*n_jobs"
description: "控制并行执行期间分派的作业数量。"
- name: "error_score"
type: "numeric or 'raise'"
default: "nan"
description: "如果拟合失败,分配给所有评分的值。"
- name: "return_train_score"
type: "bool"
default: false
description: "是否返回训练集上的评分。"
RandomizedSearchCV:
description: "随机搜索交叉验证,从指定的分布中随机采样参数组合进行评估,而非穷举所有可能。"
parameters:
- name: "estimator"
type: "estimator object"
default: null
description: "要优化的模型,必须实现 fit 和 score 方法。"
- name: "param_distributions"
type: "dict or list of dicts"
default: null
description: "参数分布,指定要搜索的参数名称和可能的分布或值列表。"
- name: "n_iter"
type: "int"
default: 10
description: "参数设置的采样次数。"
- name: "scoring"
type: "str, callable, list, tuple or dict"
default: null
description: "评分标准,用于评估模型性能。"
- name: "n_jobs"
type: "int"
default: null
description: "并行运行的作业数量,-1 表示使用所有处理器。"
- name: "cv"
type: "int, cross-validation generator or iterable"
default: 5
description: "交叉验证策略,指定折数或自定义分割方式。"
- name: "verbose"
type: "int"
default: 0
description: "控制详细程度,值越大输出信息越详细。"
- name: "random_state"
type: "int, RandomState instance or None"
default: null
description: "控制随机性,用于可重现的输出。"
- name: "refit"
type: "bool, str or callable"
default: true
description: "是否使用最佳参数重新拟合模型,如果为字符串则指定用于重新拟合的评分指标。"
- name: "return_train_score"
type: "bool"
default: false
description: "是否返回训练集上的评分。"
BayesianOptimization:
description: "贝叶斯优化,基于先验观测构建代理模型(通常是高斯过程),指导后续参数搜索方向。"
parameters:
- name: "f"
type: "callable"
default: null
description: "目标函数,接受参数并返回要最大化的值。"
- name: "pbounds"
type: "dict"
default: null
description: "参数边界,指定每个参数的搜索范围。"
- name: "random_state"
type: "int, RandomState instance or None"
default: null
description: "控制随机性,用于可重现的输出。"
- name: "verbose"
type: "int"
default: 2
description: "控制详细程度0 表示静默1 表示仅显示进度2 表示显示所有信息。"
- name: "bounds_transformer"
type: "callable"
default: null
description: "边界转换函数,用于转换参数边界。"
- name: "allow_duplicate_points"
type: "bool"
default: false
description: "是否允许评估重复的参数点。"
HyperOpt:
description: "基于贝叶斯优化和树结构Parzen估计器TPE的超参数优化库支持复杂搜索空间定义。"
parameters:
- name: "fn"
type: "callable"
default: null
description: "目标函数,接受参数并返回要最小化的值。"
- name: "space"
type: "dict or list"
default: null
description: "搜索空间,定义参数及其可能的取值范围。"
- name: "algo"
type: "callable"
default: "tpe.suggest"
description: "搜索算法,如 tpe.suggest、random.suggest 或 anneal.suggest。"
- name: "max_evals"
type: "int"
default: null
description: "最大评估次数,即尝试的参数组合数量。"
- name: "trials"
type: "Trials object"
default: null
description: "存储试验结果的对象,可用于恢复中断的搜索。"
- name: "rstate"
type: "numpy.random.RandomState"
default: null
description: "随机状态对象,用于控制随机性。"
- name: "verbose"
type: "bool"
default: false
description: "是否显示详细信息。"
Optuna:
description: "现代超参数优化框架,结合多种采样算法和修剪机制,高效搜索最优参数。"
parameters:
- name: "study"
type: "optuna.study.Study"
default: null
description: "Optuna 研究对象,用于管理优化过程。"
- name: "objective"
type: "callable"
default: null
description: "目标函数,接受 Trial 对象并返回要最小化的值。"
- name: "n_trials"
type: "int"
default: null
description: "试验次数,即评估的参数组合数量。"
- name: "direction"
type: "str or list of str"
default: "minimize"
description: "优化方向,'minimize' 或 'maximize',也可以是多目标优化的方向列表。"
- name: "sampler"
type: "optuna.samplers.BaseSampler"
default: "TPESampler"
description: "参数采样器,如 TPESampler、RandomSampler 或 CmaEsSampler。"
- name: "pruner"
type: "optuna.pruners.BasePruner"
default: "MedianPruner"
description: "修剪器,用于提前终止效果不佳的试验,如 MedianPruner 或 SuccessiveHalvingPruner。"
- name: "storage"
type: "str or None"
default: null
description: "存储后端,如 SQLite 数据库 URL用于持久化试验结果。"
- name: "study_name"
type: "str or None"
default: null
description: "研究名称,用于在存储中标识研究。"
BOHB:
description: "结合贝叶斯优化和Hyperband的多保真度优化算法平衡探索与利用。"
parameters:
- name: "configspace"
type: "ConfigSpace.ConfigurationSpace"
default: null
description: "配置空间,定义参数及其可能的取值范围。"
- name: "eta"
type: "float"
default: 3
description: "控制每轮迭代中淘汰的配置比例,通常设为 3。"
- name: "min_budget"
type: "float"
default: null
description: "最小资源分配,如最小训练轮次或样本数量。"
- name: "max_budget"
type: "float"
default: null
description: "最大资源分配,如最大训练轮次或样本数量。"
- name: "iterations"
type: "int"
default: null
description: "BOHB 迭代次数。"
- name: "random_fraction"
type: "float"
default: 0.33
description: "随机采样的比例,用于探索。"
- name: "bandwidth_factor"
type: "float"
default: 3
description: "TPE 中核带宽的因子。"
- name: "min_bandwidth"
type: "float"
default: 0.001
description: "TPE 中核带宽的最小值。"
GeneticAlgorithm:
description: "基于进化理论的优化算法,通过选择、交叉和变异操作迭代优化参数。"
parameters:
- name: "fitness_function"
type: "callable"
default: null
description: "适应度函数,评估个体的质量。"
- name: "population_size"
type: "int"
default: 100
description: "种群大小,即每代中的个体数量。"
- name: "gene_length"
type: "int"
default: null
description: "基因长度,即每个个体的参数数量。"
- name: "gene_type"
type: "str or list"
default: "binary"
description: "基因类型,如 'binary'、'real' 或 'integer'。"
- name: "gene_bounds"
type: "list of tuples"
default: null
description: "基因边界,指定每个参数的取值范围。"
- name: "generations"
type: "int"
default: 100
description: "迭代代数,即算法运行的总代数。"
- name: "crossover_rate"
type: "float"
default: 0.8
description: "交叉率,控制个体间基因交换的概率。"
- name: "mutation_rate"
type: "float"
default: 0.1
description: "变异率,控制基因随机变异的概率。"
- name: "selection_method"
type: "str"
default: "tournament"
description: "选择方法,如 'tournament'、'roulette' 或 'rank'。"
- name: "elitism"
type: "bool or int"
default: true
description: "精英主义,是否保留每代中的最佳个体。"
ParticleSwarmOptimization:
description: "基于群体智能的优化算法,模拟鸟群觅食行为,通过粒子间信息共享寻找最优解。"
parameters:
- name: "objective_function"
type: "callable"
default: null
description: "目标函数,评估粒子位置的质量。"
- name: "n_particles"
type: "int"
default: 30
description: "粒子数量,即种群大小。"
- name: "dimensions"
type: "int"
default: null
description: "问题维度,即参数数量。"
- name: "bounds"
type: "list of tuples"
default: null
description: "参数边界,指定每个参数的取值范围。"
- name: "max_iterations"
type: "int"
default: 100
description: "最大迭代次数。"
- name: "w"
type: "float"
default: 0.5
description: "惯性权重,控制粒子保持当前速度的程度。"
- name: "c1"
type: "float"
default: 1.5
description: "认知系数,控制粒子向个体最佳位置移动的程度。"
- name: "c2"
type: "float"
default: 1.5
description: "社会系数,控制粒子向全局最佳位置移动的程度。"
- name: "velocity_clamp"
type: "tuple or None"
default: null
description: "速度限制,防止粒子速度过大。"
EarlyStopping:
description: "提前停止策略,监控模型在验证集上的性能,当性能不再提升时停止训练,防止过拟合。"
parameters:
- name: "monitor"
type: "str"
default: "val_loss"
description: "要监控的指标,如 'val_loss' 或 'val_accuracy'。"
- name: "min_delta"
type: "float"
default: 0
description: "最小变化阈值,只有超过此阈值的改进才被视为有效。"
- name: "patience"
type: "int"
default: 0
description: "容忍的轮次数,即在多少轮没有改进后停止训练。"
- name: "verbose"
type: "int"
default: 0
description: "详细程度0 表示静默1 表示显示信息。"
- name: "mode"
type: "str"
default: "auto"
description: "监控模式,'min'(指标越小越好)、'max'(指标越大越好)或 'auto'(自动判断)。"
- name: "baseline"
type: "float or None"
default: null
description: "基准值,只有超过此值的模型才会被保存。"
- name: "restore_best_weights"
type: "bool"
default: false
description: "是否恢复最佳权重,而不是最后的权重。"
LearningRateScheduler:
description: "学习率调度器,根据预定策略或模型性能动态调整优化器的学习率。"
parameters:
- name: "schedule"
type: "callable"
default: null
description: "学习率调度函数,接受当前轮次并返回学习率。"
- name: "verbose"
type: "int"
default: 0
description: "详细程度0 表示静默1 表示显示信息。"
- name: "initial_learning_rate"
type: "float"
default: null
description: "初始学习率。"
- name: "decay_steps"
type: "int"
default: null
description: "衰减步数,用于某些预定义的调度策略。"
- name: "decay_rate"
type: "float"
default: null
description: "衰减率,用于某些预定义的调度策略。"
- name: "staircase"
type: "bool"
default: false
description: "是否使用阶梯式衰减,而不是连续衰减。"
- name: "min_learning_rate"
type: "float"
default: 0
description: "最小学习率,防止学习率过小。"
ModelCheckpoint:
description: "模型检查点,定期保存训练过程中的模型状态,确保能够恢复最佳模型。"
parameters:
- name: "filepath"
type: "str"
default: null
description: "保存模型的路径,可以包含格式化占位符。"
- name: "monitor"
type: "str"
default: "val_loss"
description: "要监控的指标,如 'val_loss' 或 'val_accuracy'。"
- name: "verbose"
type: "int"
default: 0
description: "详细程度0 表示静默1 表示显示信息。"
- name: "save_best_only"
type: "bool"
default: false
description: "是否只保存性能最好的模型。"
- name: "save_weights_only"
type: "bool"
default: false
description: "是否只保存模型权重,而不是整个模型。"
- name: "mode"
type: "str"
default: "auto"
description: "监控模式,'min'(指标越小越好)、'max'(指标越大越好)或 'auto'(自动判断)。"
- name: "period"
type: "int"
default: 1
description: "保存频率,即每隔多少轮保存一次模型。"
- name: "save_freq"
type: "str or int"
default: "epoch"
description: "保存频率,'epoch'(每轮保存)或整数(每多少批次保存)。"

View File

@ -37,3 +37,7 @@
- 日志管理
- 配置管理
## 3.启动命令
- 启动MLFlow
- mlflow server --host 10.0.0.202 --port 5000