Step8/9 UX: feature_start 改 QComboBox + 多源 CSV 优先回退

This commit is contained in:
DXC
2026-06-29 14:40:39 +08:00
parent b82efc1e52
commit e121b02e34
3 changed files with 474 additions and 37 deletions

View File

@ -14,10 +14,12 @@ 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
import pandas as pd
from PyQt5.QtWidgets import (
QWidget, QVBoxLayout, QGroupBox, QFormLayout, QGridLayout,
QHBoxLayout, QLabel, QLineEdit, QSpinBox, QCheckBox,
QPushButton, QFileDialog, QMessageBox, QSizePolicy,
QPushButton, QFileDialog, QMessageBox, QSizePolicy, QComboBox,
)
from PyQt5.QtCore import Qt
@ -95,6 +97,9 @@ class Step8MlTrainPanel(QWidget):
)
layout.addWidget(self.training_csv_file)
# 训练 CSV 选定后,自动刷新"特征起始列"下拉框的候选列
self.training_csv_file.line_edit.textChanged.connect(self._on_training_csv_changed)
# 机器学习模型页面
self.ml_page = QWidget()
self.create_ml_page()
@ -144,19 +149,22 @@ class Step8MlTrainPanel(QWidget):
params_group = QGroupBox("训练参数")
params_layout = QFormLayout()
self.feature_start = QLineEdit()
self.feature_start.setText("374.285004")
self.feature_start = QComboBox()
self.feature_start.setMinimumWidth(180)
self.feature_start.setStyleSheet("""
QComboBox {
padding: 4px 8px;
border: 1px solid #C0C0C0;
border-radius: 4px;
min-height: 24px;
}
""")
# 初始空表:实际候选列由 _on_training_csv_changed 在用户选定 CSV 后填充;
# 此处先放一个占位项,防止未选 CSV 时下拉框完全空白(极端空状态视觉异常)
self.feature_start.addItem("(请先选择训练 CSV", "")
self.feature_start.setCurrentIndex(0)
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)
@ -318,6 +326,133 @@ class Step8MlTrainPanel(QWidget):
return str(mw.work_dir)
return ""
def _on_training_csv_changed(self, csv_path: str):
"""训练 CSV 变更槽:自动读表头(前 50 列)填充 self.feature_start 下拉框。
触发场景:
- 用户点"浏览..."选了新 CSV
- 上游 panel_factory / set_config 注入路径
默认选中规则:
1) 优先 '374.285' 纯数字波段列(光谱起点;容忍 .285 / .285004 / .2850000001 等浮点变体)
2) 兜底:列表里第一个纯数字波段列
3) 仍找不到:保留占位项 "(请先选择训练 CSV" 不动
nrows=1 即可读到全部表头,无需 load 整个文件
"""
# 清空旧候选,保留占位策略
self.feature_start.blockSignals(True)
try:
self.feature_start.clear()
if not csv_path or not os.path.isfile(csv_path):
self.feature_start.addItem("(请先选择训练 CSV", "")
self.feature_start.setCurrentIndex(0)
return
# 读取前 50 列表头
try:
df_head = pd.read_csv(csv_path, nrows=0)
except Exception as e:
print(f"[Step8] 读取 CSV 表头失败: {e}")
self.feature_start.addItem("CSV 读取失败)", "")
self.feature_start.setCurrentIndex(0)
return
all_cols = list(df_head.columns)
head_cols = all_cols[:50]
if not head_cols:
self.feature_start.addItem("CSV 无列)", "")
self.feature_start.setCurrentIndex(0)
return
# 填充前 50 列到下拉框
for col in head_cols:
self.feature_start.addItem(str(col), str(col))
# 默认选中规则:纯数字波段列(纯数字字符串视为波长)
default_idx = self._find_default_band_index(head_cols)
if default_idx is not None:
self.feature_start.setCurrentIndex(default_idx)
else:
self.feature_start.setCurrentIndex(0)
finally:
self.feature_start.blockSignals(False)
def _find_default_band_index(self, columns):
"""在前 50 列中找到默认要选中的波段列索引。
优先级:
1) 列名以 '374.285' 开头(覆盖 374.285 / 374.285004 / 374.2850001
2) 列名是纯数字(如 '443')— 极少数仪器用整数波长
3) 第一个能被 float() 解析的纯数字列
"""
# 1) 374.285 前缀
for i, col in enumerate(columns):
if str(col).startswith("374.285"):
return i
# 2) 纯数字列(无小数点)
for i, col in enumerate(columns):
if str(col).replace('.', '').lstrip('-').isdigit() and '.' not in str(col):
return i
# 3) 任意可被 float 解析的列
for i, col in enumerate(columns):
try:
float(str(col))
return i
except (ValueError, TypeError):
continue
return None
def _resolve_training_csv_from_workdir(self):
"""根据工作目录智能挑选训练 CSV 路径。
优先级(从高到低):
1) 7_Water_Quality_Indices/training_spectra_indices.csvStep 7 WQI 增强版)
2) 10_WaterIndex_CSV/*training* / *training*indices*.csv用户自定义带指数汇总
3) 6_Spectral_Feature_Extraction/training_spectra.csvStep 6 原始特征)
4) 7_Water_Quality_Indices/ 下任意 *training*.csv
"""
work_dir = self._get_default_work_dir()
if not work_dir:
return ""
from pathlib import Path
wd = Path(work_dir)
# 1) Step 7 输出的训练 WQI 增强版
step7_csv = wd / "7_Water_Quality_Indices" / "training_spectra_indices.csv"
if step7_csv.is_file():
return str(step7_csv).replace('\\', '/')
# 2) 10_WaterIndex_CSV 下任何带 "training" 关键词的 csv用户在 Step 10 跑过训练集)
idx_dir = wd / "10_WaterIndex_CSV"
if idx_dir.is_dir():
candidates = sorted(
idx_dir.glob("*training*.csv"),
key=lambda p: p.stat().st_mtime,
reverse=True,
)
if candidates:
return str(candidates[0]).replace('\\', '/')
# 3) Step 6 原始光谱特征
step6_csv = wd / "6_Spectral_Feature_Extraction" / "training_spectra.csv"
if step6_csv.is_file():
return str(step6_csv).replace('\\', '/')
# 4) Step 7 目录下任何 training*.csv兜底
step7_dir = wd / "7_Water_Quality_Indices"
if step7_dir.is_dir():
candidates = sorted(
step7_dir.glob("*training*.csv"),
key=lambda p: p.stat().st_mtime,
reverse=True,
)
if candidates:
return str(candidates[0]).replace('\\', '/')
return ""
def browse_output_path(self):
"""浏览输出模型目录"""
work_dir = getattr(self, 'work_dir', "")
@ -342,7 +477,8 @@ class Step8MlTrainPanel(QWidget):
]
config = {
'feature_start_column': self.feature_start.text(),
# QComboBox 适配currentText() 拿用户选的列名(与原 QLineEdit.text() 语义等价)
'feature_start_column': self.feature_start.currentText(),
'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'],
@ -359,7 +495,15 @@ class Step8MlTrainPanel(QWidget):
def set_config(self, config):
"""设置配置"""
if 'feature_start_column' in config:
self.feature_start.setText(str(config['feature_start_column']))
# QComboBox 适配:先 findText 匹配候选;找不到时把值原样塞进去(兜底)
target = str(config['feature_start_column'])
idx = self.feature_start.findText(target)
if idx >= 0:
self.feature_start.setCurrentIndex(idx)
else:
# 配置回放时 CSV 可能尚未选定,候选列表为空;先插一项保留语义
self.feature_start.addItem(target, target)
self.feature_start.setCurrentIndex(self.feature_start.count() - 1)
if 'cv_folds' in config:
self.cv_folds.setValue(config['cv_folds'])
if 'preprocessing_methods' in config:
@ -393,20 +537,15 @@ class Step8MlTrainPanel(QWidget):
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()
# 1. 智能挑选训练 CSV不再"强制"读 Step 6,而是优先 WQI 增强版
# 优先级Step 7 WQI > 10_WaterIndex_CSV/*training* > Step 6 原始光谱 > Step 7 兜底
# 修复目标:用户跑过 Step 7 后再回到 Step 8UI 默认应指向带指数的训练集
# 否则训练好的模型有 95 维50 波段 + 45 WQI下次回放变成只 50 维训练,特征维数错位。
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)
candidate = self._resolve_training_csv_from_workdir()
if candidate:
self.training_csv_file.set_path(candidate)
# 2. 自动填充输出目录为 8_Machine_Learning_Models
if self.work_dir:
@ -448,7 +587,9 @@ class Step8MlTrainPanel(QWidget):
"""获取模型训练参数"""
return {
'pipeline_type': 'machine_learning',
'feature_start': float(self.feature_start.text()),
# QComboBox 适配currentText() 取列名;下拉项里就是"374.285004" 等纯数字波段名
# (与原 QLineEdit 中 "374.285004" 字符串保持完全一致,后端 float() 解析不变)
'feature_start': float(self.feature_start.currentText()),
'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()],

View File

@ -8,6 +8,8 @@ import os
import sys
from pathlib import Path
import pandas as pd
# 路径归一化 helper与 pipeline.get_step_output_dir 互为表里)
_HERE = os.path.dirname(os.path.abspath(__file__))
if _HERE not in sys.path:
@ -333,24 +335,90 @@ class Step9MlPredictPanel(QWidget):
result[name] = self.external_models_dict[name]
return result
def _resolve_latest_wqi_test_csv(self):
"""在工作目录中智能挑选"最新生成的、含 WQI 指数的测试集 CSV"
返回:找到则返回文件路径字符串;找不到返回 ""
搜索策略(按优先级递减,命中即返回):
1) 10_WaterIndex_CSV/*.csv — Step 10 输出目录(用户在 Step 10 跑过的产品)
2) 7_Water_Quality_Indices/*sampling*.csv / *test*.csv — 用户手动对采样点算过 WQI
3) work_dir 下任何 *indices*.csv / *wqi*.csv不区分大小写
4) work_dir 下任何 > 60 列的 csv启发式50 波段 + > 10 WQI 指数列)
5) 兜底空串,调用方回退到 Step 4 sampling_spectra.csv
多个候选时按 mtime 倒序选"最新生成的"
"""
work_dir = self._get_default_work_dir()
if not work_dir:
return ""
wd = Path(work_dir)
found = []
# 1) 10_WaterIndex_CSV 下所有 csvStep 10 输出)
idx_dir = wd / "10_WaterIndex_CSV"
if idx_dir.is_dir():
found.extend(idx_dir.glob("*.csv"))
# 2) 7_Water_Quality_Indices 下与采样/测试相关的 csv
qa_dir = wd / "7_Water_Quality_Indices"
if qa_dir.is_dir():
for pattern in ("*sampling*.csv", "*test*.csv", "*predict*.csv"):
found.extend(qa_dir.glob(pattern))
# 3) work_dir 直接子树下含 indices/wqi 关键词的 csv
for keyword in ("*indices*.csv", "*wqi*.csv", "*WQI*.csv"):
found.extend(wd.rglob(keyword))
# 4) 启发式:> 60 列的 csv50 波段 + 至少 10 个指数)
try:
for csv in wd.rglob("*.csv"):
if csv in found:
continue
try:
head = pd.read_csv(csv, nrows=0)
if head.shape[1] > 60:
found.append(csv)
except Exception:
pass # 读取失败就跳过,不影响其它候选
except Exception:
pass
if not found:
return ""
# 去重 + 按 mtime 倒序排
uniq = {p.resolve(): p for p in found}.values()
sorted_paths = sorted(uniq, key=lambda p: p.stat().st_mtime, reverse=True)
return str(sorted_paths[0]).replace('\\', '/')
def update_from_config(self, work_dir=None, pipeline=None):
if work_dir: self.work_dir = work_dir
main_window = self.window()
factory = getattr(main_window, '_panel_factory', None) if main_window else None
if not factory: return
# 1. 拿第 4 步的采样光谱
step4_panel = factory.get_panel('step4_sampling')
if step4_panel and hasattr(step4_panel, 'output_file'):
path = step4_panel.output_file.get_path()
if path: self.sampling_csv_file.set_path(path)
# 1. 智能挑选采样 CSV优先"含 WQI 指数的测试集"(防止特征维度与训练时不匹配)
# 修复目标:用户在 Step 8 用 95 维 (50+45 WQI) 训练 → Step 9 默认读 50 维 raw sampling
# 时 inference_batch.preprocess_spectra 会触发"自动特征补全"逻辑;但若用户已经在
# Step 10/手工把指数算到 CSV 里了,应该直接用那个文件(少走内存补全、避免 band 列顺序漂移)
wqi_test_csv = self._resolve_latest_wqi_test_csv()
if wqi_test_csv:
self.sampling_csv_file.set_path(wqi_test_csv)
elif factory:
# 兜底:拿第 4 步的纯原始采样光谱(旧行为保留)
step4_panel = factory.get_panel('step4_sampling')
if step4_panel and hasattr(step4_panel, 'output_file'):
path = step4_panel.output_file.get_path()
if path: self.sampling_csv_file.set_path(path)
# 2. 拿第 8 步的模型目录
step8_panel = factory.get_panel('step8_ml_train')
if step8_panel and hasattr(step8_panel, 'output_path'):
path = step8_panel.output_path.get_path()
if path: self.models_dir_file.set_path(path)
if factory:
step8_panel = factory.get_panel('step8_ml_train')
if step8_panel and hasattr(step8_panel, 'output_path'):
path = step8_panel.output_path.get_path()
if path: self.models_dir_file.set_path(path)
# 3. 生成第 9 步的输出目录
if hasattr(self, 'work_dir') and self.work_dir: