#!/usr/bin/env python # -*- coding: utf-8 -*- """ Step8 面板 - 机器学习建模 """ import os import sys from pathlib import Path # 路径归一化 helper(与 pipeline.get_step_output_dir 互为表里) _HERE = os.path.dirname(os.path.abspath(__file__)) if _HERE not in sys.path: sys.path.insert(0, _HERE) from src.gui.panels._step_path_resolver import get_step_output_path, resolve_step_widget, resolve_subdir from PyQt5.QtWidgets import ( QWidget, QVBoxLayout, QGroupBox, QFormLayout, QGridLayout, QHBoxLayout, QLabel, QLineEdit, QSpinBox, QCheckBox, QPushButton, QFileDialog, QMessageBox, QSizePolicy, ) from PyQt5.QtCore import Qt from src.gui.components.custom_widgets import FileSelectWidget from src.gui.styles import ModernStylesheet # ============================================================ # 中文映射表(内部键名 -> 显示文本) # ============================================================ # 预处理方法:内部键 -> 显示文本 PREPROC_CHINESE = { 'None': '无 (None)', 'MMS': '最小-最大归一化 (MMS)', 'SS': '标度化 (SS)', 'SNV': '标准正态变换 (SNV)', 'MA': '移动平均 (MA)', 'SG': 'Savitzky-Golay (SG)', 'MSC': '多元散射校正 (MSC)', 'D1': '一阶导数 (D1)', 'D2': '二阶导数 (D2)', 'DT': '去趋势 (DT)', 'CT': '中心化 (CT)', } # 模型类型:内部键 -> 显示文本 MODEL_CHINESE = { # 线性模型 'LinearRegression': '多元线性回归 (MLR)', 'Ridge': '岭回归 (Ridge)', 'Lasso': '套索回归 (Lasso)', 'ElasticNet': '弹性网络 (ElasticNet)', 'PLS': '偏最小二乘 (PLSR)', # 树模型 'DecisionTree': '决策树 (CART)', 'RF': '随机森林 (RF)', 'ExtraTrees': '极端随机树 (ET)', 'XGBoost': '极值梯度提升 (XGBoost)', 'LightGBM': '轻量梯度提升 (LightGBM)', 'CatBoost': '类别梯度提升 (CatBoost)', # 集成学习 'GradientBoosting': '梯度提升树 (GBDT)', 'AdaBoost': '自适应提升 (AdaBoost)', # 其他模型 'SVR': '支持向量回归 (SVR)', 'KNN': 'K近邻回归 (KNN)', 'MLP': '多层感知机 (BP神经网络)', } # 数据划分方法:内部键 -> 显示文本 SPLIT_CHINESE = { 'spxy': 'SPXY 算法 (考量X-Y空间)', 'ks': 'KS 算法 (考量X空间)', 'random': '随机划分 (Random)', } class Step8MlTrainPanel(QWidget): """步骤8:机器学习建模""" def __init__(self, parent=None): super().__init__(parent) self.init_ui() def init_ui(self): layout = QVBoxLayout() # 标题 # 训练数据文件(用于独立运行) self.training_csv_file = FileSelectWidget( "训练数据:", "CSV Files (*.csv);;All Files (*.*)" ) layout.addWidget(self.training_csv_file) # 机器学习模型页面 self.ml_page = QWidget() self.create_ml_page() layout.addWidget(self.ml_page) # ========================================== # 卡片 3:输出与执行 # ========================================== output_group = QGroupBox("🚀 输出与执行") output_layout = QVBoxLayout() output_layout.setSpacing(16) output_layout.setContentsMargins(20, 24, 20, 20) # 输出文件路径 (改为文件夹模式) self.output_path = FileSelectWidget( "模型输出目录:", "Directories" ) self.output_path.browse_btn.clicked.disconnect() self.output_path.browse_btn.clicked.connect(self.browse_output_path) output_layout.addWidget(self.output_path) # 完美对齐的底部按钮栏(已彻底移除多余的启用复选框) action_layout = QHBoxLayout() action_layout.addStretch() self.run_btn = QPushButton("独立运行步骤") self.run_btn.setStyleSheet(ModernStylesheet.get_button_stylesheet('primary')) self.run_btn.setMinimumWidth(140) self.run_btn.clicked.connect(self._on_run_single_clicked) action_layout.addWidget(self.run_btn) output_layout.addLayout(action_layout) output_group.setLayout(output_layout) # 将打包好的输出卡片添加到主 layout 中 layout.addWidget(output_group) layout.addStretch() self.setLayout(layout) def create_ml_page(self): """创建机器学习模型页面""" layout = QVBoxLayout() # 参数设置 params_group = QGroupBox("训练参数") params_layout = QFormLayout() self.feature_start = QLineEdit() self.feature_start.setText("374.285004") params_layout.addRow("特征起始列:", self.feature_start) # 特征起始列名提示:用记事本打开 training_spectra.csv 确认首个波长的精确表头 feature_start_hint = QLabel( "提示:请使用记事本打开 training_spectra.csv 确认首个波长的精确表头名称" "(如 374.285 或 374.285004)并在此填入,避免因浮点精度差异导致列名匹配失败。" ) feature_start_hint.setWordWrap(True) feature_start_hint.setStyleSheet("color: #666; font-size: 10px;") params_layout.addRow(feature_start_hint) self.cv_folds = QSpinBox() self.cv_folds.setRange(2, 10) self.cv_folds.setValue(3) params_layout.addRow("交叉验证折数:", self.cv_folds) params_group.setLayout(params_layout) layout.addWidget(params_group) # 预处理方法 - 多选 preproc_group = QGroupBox("预处理方法 (可多选)") preproc_layout = QVBoxLayout() preproc_grid = QGridLayout() self.preproc_checkboxes = {} preproc_methods = ['None', 'MMS', 'SS', 'SNV', 'MA', 'SG', 'MSC', 'D1', 'D2', 'DT', 'CT'] for i, method in enumerate(preproc_methods): checkbox = QCheckBox(PREPROC_CHINESE.get(method, method)) checkbox.setChecked(False) self.preproc_checkboxes[method] = checkbox preproc_grid.addWidget(checkbox, i // 4, i % 4) button_layout = QHBoxLayout() button_layout.setContentsMargins(0, 10, 0, 0) # 顶部增加边距 button_layout.setSpacing(10) # 按钮之间的间距 select_all_btn = QPushButton("全选") deselect_all_btn = QPushButton("全不选") # 应用统一样式和自适应宽度策略 for btn in [select_all_btn, deselect_all_btn]: base_style = ModernStylesheet.get_button_stylesheet('normal') # 增加 padding 确保文字不被截断 btn.setStyleSheet(base_style + "\nQPushButton { padding-left: 14px; padding-right: 14px; }") btn.setFixedHeight(30) btn.setSizePolicy(QSizePolicy.Minimum, QSizePolicy.Fixed) select_all_btn.clicked.connect(lambda: self._toggle_checkboxes(self.preproc_checkboxes, True)) deselect_all_btn.clicked.connect(lambda: self._toggle_checkboxes(self.preproc_checkboxes, False)) button_layout.addWidget(select_all_btn) button_layout.addWidget(deselect_all_btn) button_layout.addStretch() preproc_layout.addLayout(preproc_grid) preproc_layout.addLayout(button_layout) preproc_group.setLayout(preproc_layout) layout.addWidget(preproc_group) # 模型选择 - 多选 model_group = QGroupBox("模型类型 (可多选)") model_layout = QVBoxLayout() model_grid = QGridLayout() self.model_checkboxes = {} model_groups = [ ("【线性模型】", ['LinearRegression', 'Ridge', 'Lasso', 'ElasticNet', 'PLS']), ("【树模型】", ['DecisionTree', 'RF', 'ExtraTrees', 'XGBoost', 'LightGBM', 'CatBoost']), ("【集成学习】", ['GradientBoosting', 'AdaBoost']), ("【其他模型】", ['SVR', 'KNN', 'MLP']) ] row = 0 for group_name, models in model_groups: group_label = QLabel(f"{group_name}") group_label.setStyleSheet( f"background-color: {ModernStylesheet.COLORS['hover']}; " f"padding: 5px; border: 1px solid {ModernStylesheet.COLORS['border_light']}; " f"border-radius: 3px;" ) model_grid.addWidget(group_label, row, 0, 1, 4) row += 1 for i, model in enumerate(models): checkbox = QCheckBox(MODEL_CHINESE.get(model, model)) checkbox.setChecked(False) self.model_checkboxes[model] = checkbox model_grid.addWidget(checkbox, row, i % 4) if (i + 1) % 4 == 0: row += 1 row += 1 model_button_layout = QHBoxLayout() model_button_layout.setContentsMargins(0, 10, 0, 0) model_button_layout.setSpacing(10) model_select_all = QPushButton("全选") model_deselect_all = QPushButton("全不选") for btn in [model_select_all, model_deselect_all]: base_style = ModernStylesheet.get_button_stylesheet('normal') btn.setStyleSheet(base_style + "\nQPushButton { padding-left: 14px; padding-right: 14px; }") btn.setFixedHeight(30) btn.setSizePolicy(QSizePolicy.Minimum, QSizePolicy.Fixed) model_select_all.clicked.connect(lambda: self._toggle_checkboxes(self.model_checkboxes, True)) model_deselect_all.clicked.connect(lambda: self._toggle_checkboxes(self.model_checkboxes, False)) model_button_layout.addWidget(model_select_all) model_button_layout.addWidget(model_deselect_all) model_button_layout.addStretch() model_layout.addLayout(model_grid) model_layout.addLayout(model_button_layout) model_group.setLayout(model_layout) layout.addWidget(model_group) # 数据划分方法 - 多选 split_group = QGroupBox("数据划分方法 (可多选)") split_layout = QVBoxLayout() split_grid = QGridLayout() self.split_checkboxes = {} split_methods = ['spxy', 'ks', 'random'] for i, method in enumerate(split_methods): checkbox = QCheckBox(SPLIT_CHINESE.get(method, method)) checkbox.setChecked(False) self.split_checkboxes[method] = checkbox split_grid.addWidget(checkbox, 0, i) split_button_layout = QHBoxLayout() split_button_layout.setContentsMargins(0, 10, 0, 0) split_button_layout.setSpacing(10) split_select_all = QPushButton("全选") split_deselect_all = QPushButton("全不选") for btn in [split_select_all, split_deselect_all]: base_style = ModernStylesheet.get_button_stylesheet('normal') btn.setStyleSheet(base_style + "\nQPushButton { padding-left: 14px; padding-right: 14px; }") btn.setFixedHeight(30) btn.setSizePolicy(QSizePolicy.Minimum, QSizePolicy.Fixed) split_select_all.clicked.connect(lambda: self._toggle_checkboxes(self.split_checkboxes, True)) split_deselect_all.clicked.connect(lambda: self._toggle_checkboxes(self.split_checkboxes, False)) split_button_layout.addWidget(split_select_all) split_button_layout.addWidget(split_deselect_all) split_button_layout.addStretch() split_layout.addLayout(split_grid) split_layout.addLayout(split_button_layout) split_group.setLayout(split_layout) layout.addWidget(split_group) self.ml_page.setLayout(layout) def _toggle_checkboxes(self, checkboxes_dict, checked): """统一设置checkbox状态""" for checkbox in checkboxes_dict.values(): checkbox.setChecked(checked) def _get_default_work_dir(self): """获取 work_dir,优先用 panel 自身缓存的,否则尝试从主窗口取""" if hasattr(self, 'work_dir') and self.work_dir: return str(self.work_dir) mw = self.window() if mw and hasattr(mw, 'work_dir') and mw.work_dir: return str(mw.work_dir) return "" def browse_output_path(self): """浏览输出模型目录""" work_dir = getattr(self, 'work_dir', "") initial_dir = os.path.join(work_dir, '8_Machine_Learning_Models') if work_dir else "" dir_path = QFileDialog.getExistingDirectory(self, "选择模型输出目录", initial_dir) if dir_path: self.output_path.set_path(dir_path) def get_config(self): """获取配置""" preprocessing_methods = [ method for method, checkbox in self.preproc_checkboxes.items() if checkbox.isChecked() ] model_names = [ model for model, checkbox in self.model_checkboxes.items() if checkbox.isChecked() ] split_methods = [ method for method, checkbox in self.split_checkboxes.items() if checkbox.isChecked() ] config = { 'feature_start_column': self.feature_start.text(), 'preprocessing_methods': preprocessing_methods if preprocessing_methods else ['None'], 'model_names': model_names if model_names else ['SVR'], 'split_methods': split_methods if split_methods else ['random'], 'cv_folds': self.cv_folds.value() } training_csv_path = self.training_csv_file.get_path() if training_csv_path: config['training_csv_path'] = training_csv_path output_path = self.output_path.get_path() if output_path: config['output_path'] = output_path return config def set_config(self, config): """设置配置""" if 'feature_start_column' in config: self.feature_start.setText(str(config['feature_start_column'])) if 'cv_folds' in config: self.cv_folds.setValue(config['cv_folds']) if 'preprocessing_methods' in config: methods = config['preprocessing_methods'] for method, checkbox in self.preproc_checkboxes.items(): checkbox.setChecked(method in methods) if 'model_names' in config: models = config['model_names'] for model, checkbox in self.model_checkboxes.items(): checkbox.setChecked(model in models) if 'split_methods' in config: methods = config['split_methods'] for method, checkbox in self.split_checkboxes.items(): checkbox.setChecked(method in methods) if 'training_csv_path' in config: self.training_csv_file.set_path(config['training_csv_path']) if 'output_path' in config: self.output_path.set_path(config['output_path']) def update_from_config(self, work_dir=None, pipeline=None): """从全局配置自动填充训练数据和输出路径 Args: work_dir: 工作目录路径 pipeline: Pipeline 实例(未使用,保留接口兼容性) """ if work_dir: self.work_dir = work_dir elif hasattr(self, 'work_dir') and self.work_dir: pass else: self.work_dir = None # 1. 强制读 Step 6 的 training_spectra.csv(光谱特征提取结果) # 修复张冠李戴:原链路 STEP_DATA_SOURCE['training_spectra_csv'] → step5_clean_panel # 错误地指向了 Step 5 的 processed_data.csv(纯清洗数据,不含光谱特征), # 实际 ML 训练需要的特征数据来自 Step 6 的 6_Spectral_Feature_Extraction/training_spectra.csv main_window = self.window() existing_training_csv = self.training_csv_file.get_path() if not existing_training_csv or not existing_training_csv.strip(): if self.work_dir: step6_dir = resolve_subdir(self.work_dir, 'spectral_feature') step6_training_csv = os.path.join( step6_dir, 'training_spectra.csv' ).replace('\\', '/') if step6_training_csv: self.training_csv_file.set_path(step6_training_csv) # 2. 自动填充输出目录为 8_Machine_Learning_Models if self.work_dir: import os models_dir = os.path.join(self.work_dir, "8_Machine_Learning_Models").replace('\\', '/') os.makedirs(models_dir, exist_ok=True) self.output_path.set_path(models_dir) else: self.output_path.set_path("") def _on_run_single_clicked(self): """通过 EventBus 发布单步执行请求(解耦面板与 PipelineExecutor)。""" from src.gui.core.event_bus import global_event_bus training_csv_path = self.training_csv_file.get_path() if not training_csv_path: QMessageBox.warning(self, "输入错误", "请选择训练数据CSV文件!") return config = {'step8_ml_train': self.get_config()} global_event_bus.publish('RequestRunSingleStep', { 'step_name': 'step8_ml_train', 'config': config, }) def run_step(self): """独立运行步骤8(旧版 parent 链上溯方式,保留兼容)。""" training_csv_path = self.training_csv_file.get_path() if not training_csv_path: QMessageBox.warning(self, "输入错误", "请选择训练数据CSV文件!") return main_window = self.window() if hasattr(main_window, 'run_single_step'): config = {'step8_ml_train': self.get_config()} main_window.run_single_step('step8_ml_train', config) def get_training_params(self): """获取模型训练参数""" return { 'pipeline_type': 'machine_learning', 'feature_start': float(self.feature_start.text()), 'cv_folds': self.cv_folds.value(), 'preprocess_methods': [method for method, cb in self.preproc_checkboxes.items() if cb.isChecked()], 'model_types': [model for model, cb in self.model_checkboxes.items() if cb.isChecked()], 'split_methods': [method for method, cb in self.split_checkboxes.items() if cb.isChecked()] }