#!/usr/bin/env python # -*- coding: utf-8 -*- """ Step8 面板 - 机器学习建模 """ import os import sys import pandas as pd import numpy as np from pathlib import Path from typing import Dict, List, Optional, Tuple from PyQt5.QtWidgets import ( QWidget, QVBoxLayout, QGroupBox, QGridLayout, QHBoxLayout, QLabel, QCheckBox, QPushButton, QMessageBox, QScrollArea, QListWidget, QListWidgetItem, QAbstractItemView, QRadioButton, QButtonGroup ) from PyQt5.QtCore import Qt from PyQt5.QtGui import QColor, QBrush, QFont from src.gui.components.custom_widgets import FileSelectWidget from src.gui.styles import ModernStylesheet def get_resource_path(relative_path: str) -> str: """适配开发与 PyInstaller 环境的路径获取逻辑。""" if hasattr(sys, '_MEIPASS'): internal = os.path.join(sys._MEIPASS, '_internal', relative_path) if os.path.exists(internal): return internal return os.path.join(sys._MEIPASS, relative_path) exe_dir = os.path.dirname(sys.executable) internal = os.path.join(exe_dir, '_internal', relative_path) if os.path.exists(internal): return internal base_dir = Path(__file__).resolve().parent.parent / "model" return str(base_dir / os.path.basename(relative_path)) class Step8MlTrainPanel(QWidget): """步骤8:机器学习建模""" COLOR_RATIO = QColor(255, 255, 255) COLOR_CONCENTRATION = QColor(220, 240, 255) COLOR_HEADER = QColor(245, 245, 245) def __init__(self, parent=None): super().__init__(parent) self.index_checkboxes: Dict[str, QListWidgetItem] = {} self.work_dir: Optional[str] = None self.builtin_formula_path = get_resource_path("waterindex.csv") self._formula_type_map: Dict[str, str] = {} self._formula_color_map: Dict[str, QColor] = {} self._formula_coef_map: Dict[str, List[float]] = {} self.init_ui() self._auto_load_formulas() def init_ui(self): main_layout = QVBoxLayout() main_layout.setContentsMargins(20, 20, 20, 20) main_layout.setSpacing(10) # 1. 公式配置源 (只读) path_group = QGroupBox("公式配置源 (内置)") path_layout = QVBoxLayout() self.formula_csv_widget = FileSelectWidget("内置CSV路径:", "CSV Files (*.csv)") self.formula_csv_widget.set_path(self.builtin_formula_path) self.formula_csv_widget.set_read_only(True) self.formula_csv_widget.line_edit.setStyleSheet("background-color: #f0f0f0; color: #666;") path_layout.addWidget(self.formula_csv_widget) path_group.setLayout(path_layout) main_layout.addWidget(path_group) # 2. 训练数据输入 input_group = QGroupBox("输入样本数据") input_layout = QVBoxLayout() self.training_data_widget = FileSelectWidget("特征提取CSV:", "CSV Files (*.csv)") input_layout.addWidget(self.training_data_widget) input_group.setLayout(input_layout) main_layout.addWidget(input_group) # 3. 公式选择区 (分组 ListWidget) self.formula_group = QGroupBox("待计算水质指数勾选") formula_outer_layout = QVBoxLayout() btn_layout = QHBoxLayout() self.select_all_btn = QPushButton("全选") self.deselect_all_btn = QPushButton("清空") self.select_ratio_btn = QPushButton("仅选比值型") self.select_conc_btn = QPushButton("仅选浓度型") self.select_all_btn.clicked.connect(self.select_all_formulas) self.deselect_all_btn.clicked.connect(self.deselect_all_formulas) self.select_ratio_btn.clicked.connect(self._select_ratio_only) self.select_conc_btn.clicked.connect(self._select_conc_only) btn_layout.addWidget(self.select_all_btn) btn_layout.addWidget(self.deselect_all_btn) btn_layout.addWidget(self.select_ratio_btn) btn_layout.addWidget(self.select_conc_btn) btn_layout.addStretch() self.refresh_button = QPushButton("重新加载") self.refresh_button.clicked.connect(lambda: self.refresh_formulas(silent=False)) btn_layout.addWidget(self.refresh_button) formula_outer_layout.addLayout(btn_layout) scroll = QScrollArea() scroll.setWidgetResizable(True) scroll.setMinimumHeight(280) self.scroll_content = QWidget() self.formula_layout = QVBoxLayout(self.scroll_content) self.formula_layout.setContentsMargins(4, 4, 4, 4) self.formula_layout.setSpacing(2) self.formula_layout.setAlignment(Qt.AlignTop) self.formula_list = QListWidget() self.formula_list.setSelectionMode(QAbstractItemView.MultiSelection) self.formula_list.itemChanged.connect(self._on_item_changed) self.formula_layout.addWidget(self.formula_list) scroll.setWidget(self.scroll_content) formula_outer_layout.addWidget(scroll) self.formula_group.setLayout(formula_outer_layout) main_layout.addWidget(self.formula_group) # 4. 输出选项 output_group = QGroupBox("输出模式") output_layout = QVBoxLayout() mode_layout = QHBoxLayout() self.mode_group = QButtonGroup() self.radio_both = QRadioButton("两者皆出") self.radio_wide = QRadioButton("仅宽表") self.radio_single = QRadioButton("仅单文件") self.mode_group.addButton(self.radio_both, 0) self.mode_group.addButton(self.radio_wide, 1) self.mode_group.addButton(self.radio_single, 2) self.radio_both.setChecked(True) mode_layout.addWidget(self.radio_both) mode_layout.addWidget(self.radio_wide) mode_layout.addWidget(self.radio_single) mode_layout.addStretch() output_layout.addLayout(mode_layout) self.enable_checkbox = QCheckBox("启用计算流程") self.enable_checkbox.setChecked(True) output_layout.addWidget(self.enable_checkbox) output_group.setLayout(output_layout) main_layout.addWidget(output_group) # 5. 运行按钮 self.run_button = QPushButton("立即执行计算") self.run_button.setStyleSheet(ModernStylesheet.get_button_stylesheet('success')) self.run_button.setMinimumHeight(40) self.run_button.clicked.connect(self.run_step) main_layout.addWidget(self.run_button) self.setLayout(main_layout) def _on_item_changed(self, item: QListWidgetItem): if item.checkState() == Qt.Checked: bg_color = self.COLOR_RATIO for name, ref_item in self.index_checkboxes.items(): if ref_item is item: bg_color = self._formula_color_map.get(name, self.COLOR_RATIO) break item.setBackground(QBrush(bg_color)) else: item.setBackground(QBrush(self.COLOR_RATIO)) def _auto_load_formulas(self): if os.path.exists(self.builtin_formula_path): self.refresh_formulas(silent=True) else: print(f"DEBUG: 自动加载失败,路径不存在: {self.builtin_formula_path}") def refresh_formulas(self, silent=False): path = self.builtin_formula_path if not os.path.exists(path): if not silent: QMessageBox.warning(self, "错误", f"找不到内置公式文件:\n{path}") return try: df = None for enc in ('utf-8', 'gbk', 'utf-8-sig'): try: df = pd.read_csv(path, encoding=enc) if 'Formula_Name' in df.columns: break except Exception: continue if df is None or 'Formula_Name' not in df.columns: if not silent: QMessageBox.critical(self, "错误", "CSV缺少 'Formula_Name' 列") return self._formula_type_map.clear() self._formula_coef_map.clear() for _, row in df.iterrows(): name = str(row['Formula_Name']).strip() if not name: continue ftype = str(row.get('Formula_Type', 'ratio')).strip().lower() self._formula_type_map[name] = ftype coef_str = str(row.get('Coefficient', '')).strip() if coef_str: try: coeffs = [float(c.strip()) for c in coef_str.split(',') if c.strip()] self._formula_coef_map[name] = coeffs except Exception: self._formula_coef_map[name] = [] else: self._formula_coef_map[name] = [] self.formula_list.clear() self.index_checkboxes.clear() self._formula_color_map.clear() for name, ftype in self._formula_type_map.items(): item = QListWidgetItem(name, self.formula_list) item.setCheckState(Qt.Checked) if ftype == 'concentration': bg_color = QColor(220, 240, 255) else: bg_color = self.COLOR_RATIO self._formula_color_map[name] = bg_color item.setBackground(QBrush(bg_color)) self.index_checkboxes[name] = item self.formula_list.adjustSize() print(f"✅ 加载 {len(self.index_checkboxes)} 个公式") except Exception as e: if not silent: QMessageBox.critical(self, "加载失败", f"原因: {str(e)}") def _select_ratio_only(self): for name, item in self.index_checkboxes.items(): ftype = self._formula_type_map.get(name, 'ratio') item.setCheckState(Qt.Checked if ftype == 'ratio' else Qt.Unchecked) def _select_conc_only(self): for name, item in self.index_checkboxes.items(): ftype = self._formula_type_map.get(name, 'ratio') item.setCheckState(Qt.Checked if ftype == 'concentration' else Qt.Unchecked) def select_all_formulas(self): for item in self.index_checkboxes.values(): item.setCheckState(Qt.Checked) def deselect_all_formulas(self): for item in self.index_checkboxes.values(): item.setCheckState(Qt.Unchecked) def get_config(self) -> Dict: selected = [ name for name, item in self.index_checkboxes.items() if item.checkState() == Qt.Checked ] formula_coefficients = { name: self._formula_coef_map.get(name, []) for name in selected } return { 'training_csv_path': self.training_data_widget.get_path(), 'formula_csv_file': self.builtin_formula_path, 'formula_names': selected, 'formula_coefficients': formula_coefficients, 'enabled': self.enable_checkbox.isChecked(), 'output_mode': self.mode_group.checkedId(), } def set_config(self, config: Dict): if 'training_csv_path' in config: self.training_data_widget.set_path(config['training_csv_path']) if 'formula_names' in config: sel = set(config['formula_names']) for name, item in self.index_checkboxes.items(): item.setCheckState(Qt.Checked if name in sel else Qt.Unchecked) self.enable_checkbox.setChecked(config.get('enabled', True)) if 'output_mode' in config: btn = self.mode_group.button(config['output_mode']) if btn: btn.setChecked(True) def update_from_config(self, work_dir=None, pipeline=None): if work_dir: self.work_dir = work_dir main = self.window() if hasattr(main, 'step5_panel'): p5 = main.step5_panel.output_file.get_path() if p5: if not os.path.isabs(p5): p5 = os.path.join(self.work_dir or '', p5) p5 = p5.replace('\\', '/') self.training_data_widget.set_path(p5) def _get_work_dir(self) -> Optional[str]: if self.work_dir: return self.work_dir main = self.window() if hasattr(main, 'work_dir') and main.work_dir: return main.work_dir return None def _get_coord_cols(self, df: pd.DataFrame) -> Tuple[str, str]: coord_candidates = ['lon', 'lng', 'longitude', '经度', 'x', 'lon_utm', 'utm_x', 'pixel_x'] lat_candidates = ['lat', 'latitude', '纬度', 'y', 'lat_utc', 'utm_y', 'pixel_y'] x_col, y_col = None, None for col in df.columns: cl = col.lower() if x_col is None and any(c in cl for c in coord_candidates): x_col = col if y_col is None and any(c in cl for c in lat_candidates): y_col = col if x_col is None and len(df.columns) >= 2: x_col = df.columns[0] if y_col is None and len(df.columns) >= 2: y_col = df.columns[1] return x_col or 'x_coord', y_col or 'y_coord' def run_step(self): config = self.get_config() if not config['enabled']: QMessageBox.information(self, "提示", "已禁用计算流程(启用计算流程未勾选)") return training_path = config['training_csv_path'] if not training_path or not os.path.exists(training_path): QMessageBox.warning(self, "提示", "请先选择输入特征提取CSV文件") return main_window = self.window() if hasattr(main_window, 'run_single_step'): pipeline_config = {'step8_ml_train': config} main_window.run_single_step('step8_ml_train', pipeline_config)