Files
WQ_GUI/src/gui/panels/step12_viz_panel.py
duxin f496b28c8c fix: 回退 setWidgetResizable(False) 恢复图像显示 + 保留滚动条策略修复
setWidgetResizable(False)+QSizePolicy.Ignored 导致 QLabel
  无有效 sizeHint, 图像不显示。

回退为 setWidgetResizable(True), 保留显式 ScrollBarAsNeeded
  策略 (原默认可能为 AlwaysOff 导致滚动条异常)。
2026-07-08 17:55:34 +08:00

2023 lines
88 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

#!/usr/bin/env python
# -*- coding: utf-8 -*-
"""
VisualizationPanel - 可视化分析面板
左侧目录树 + 右侧图像查看器,支持多种图表生成。
"""
import os
import traceback
from pathlib import Path
from typing import Optional, List, Union
import numpy as np
import pandas as pd
from src.gui.panels._step_path_resolver import get_step_output_path, resolve_step_widget, resolve_subdir
from PyQt5.QtCore import Qt, QTimer, QThread, pyqtSignal, QAbstractTableModel
from PyQt5.QtGui import QPixmap
from PyQt5.QtWidgets import (
QWidget, QVBoxLayout, QHBoxLayout, QGroupBox, QFormLayout,
QLabel, QCheckBox, QPushButton, QLineEdit, QMessageBox,
QFileDialog, QFrame, QSizePolicy,
QDialog, QTreeWidget, QListWidget, QAbstractItemView, QHeaderView,QTreeWidgetItem,QScrollArea
)
from matplotlib.backends.backend_qt5agg import FigureCanvasQTAgg as FigureCanvas
from matplotlib.backends.backend_qt5agg import NavigationToolbar2QT as NavigationToolbar
from matplotlib.figure import Figure
PIPELINE_AVAILABLE = True
def _viz_training_spectra_csv_path(work_path: Path) -> Path:
"""可视化光谱/统计及模型散点图使用的训练光谱表路径(与步骤6输出一致)。
注意:步骤5.5(水质指数计算)执行后会覆盖此文件为94维增强版本,
因此下游步骤无需任何修改,直接读取此路径即可。
"""
return work_path / "6_Spectral_Feature_Extraction" / "training_spectra.csv"
def _viz_infer_wavelength_start_column(df: pd.DataFrame) -> Union[str, int]:
"""推断光谱起始列(training_spectra 通常以波长数值为列名,未必含 UTM_Y)。"""
import re
# 兼容前缀:nm / nm_ / wavelength / wl / λ / wave_len(大小写不敏感)
_PREFIX_RE = re.compile(r"^(?:nm|wavelength|wl|\u03bb|wave_len)[_\s\-]*", re.IGNORECASE)
for i, col in enumerate(df.columns):
name = str(col).strip().lstrip("\ufeff")
# 1) 直接纯数字(如 "700.0")
try:
v = float(name)
except ValueError:
v = None
if v is not None and 200.0 <= v <= 3000.0:
return i
# 2) 剥前缀后再试(如 "nm_700"、"wavelength_800.0"、"WL 900")
stripped = _PREFIX_RE.sub("", name)
if stripped != name:
try:
v = float(stripped)
except ValueError:
continue
if 200.0 <= v <= 3000.0:
return i
if "UTM_Y" in df.columns:
return "UTM_Y"
# 兜底分支触发时打印所有列名,便于下次失败时一眼看到根因
print(f"DEBUG: 尝试解析的列名: {df.columns.tolist()}")
return 0
class VisualizationWorkerThread(QThread):
"""可视化耗时计算放入后台线程,并临时使用 Agg 后端,避免主界面未响应。"""
finished_ok = pyqtSignal(object)
failed = pyqtSignal(str)
def __init__(self, task: str, work_dir: str, extra: Optional[dict] = None):
super().__init__()
self.task = task
self.work_dir = str(work_dir)
self.extra = extra or {}
def run(self):
mpl_prev = None
try:
import matplotlib
mpl_prev = matplotlib.get_backend()
except Exception:
pass
try:
import matplotlib.pyplot as plt
plt.switch_backend("Agg")
except Exception:
mpl_prev = None
try:
wp = Path(self.work_dir)
if self.task == "mask_glint":
from src.postprocessing.visualization_reports import WaterQualityVisualization
viz = WaterQualityVisualization(output_dir=str(resolve_subdir(self.work_dir, 'visualization')))
preview_paths = viz.generate_glint_deglint_previews(
work_dir=str(wp),
output_subdir="glint_deglint_previews",
)
cnt = len(preview_paths) if preview_paths else 0
self.finished_ok.emit({"task": "mask_glint", "count": cnt, "preview_paths": preview_paths})
elif self.task == "sampling_map":
hyperspectral_files = []
deglint_dir = Path(resolve_subdir(self.work_dir, 'deglint'))
if deglint_dir.exists():
for ext in ("*.dat", "*.bsq", "*.tif", "*.tiff"):
hyperspectral_files.extend(list(deglint_dir.glob(ext)))
if not hyperspectral_files:
for ext in ("*.dat", "*.bsq", "*.tif", "*.tiff"):
hyperspectral_files.extend(list(wp.glob(f"**/{ext}")))
if not hyperspectral_files:
self.failed.emit("未找到高光谱影像文件(.dat/.bsq/.tif)。")
return
hyperspectral_path = str(hyperspectral_files[0])
csv_files = []
processed_dir = wp / "4_processed_data"
if processed_dir.exists():
csv_files = list(processed_dir.glob("*.csv"))
if not csv_files:
csv_files = (
list(wp.glob("**/*sampling*.csv"))
+ list(wp.glob("**/*point*.csv"))
+ list(wp.glob("**/*.csv"))
)
if not csv_files:
self.failed.emit("未找到采样点 CSV 文件。")
return
csv_path = str(csv_files[0])
from src.postprocessing.point_map import SamplingPointMap
map_generator = SamplingPointMap(
output_dir=str(Path(resolve_subdir(self.work_dir, 'visualization')) / "sampling_maps"),
fast_mode=True,
)
map_path = map_generator.create_sampling_point_map(
hyperspectral_path=hyperspectral_path,
csv_path=csv_path,
point_color="red",
point_size=100,
point_alpha=0.9,
show_north_arrow=True,
show_scale_bar=True,
show_legend=True,
downsample=True,
dpi=180,
)
self.finished_ok.emit(
{
"task": "sampling_map",
"map_path": map_path,
"hyperspectral_path": hyperspectral_path,
"csv_path": csv_path,
}
)
elif self.task == "spectrum":
from src.postprocessing.visualization_reports import WaterQualityVisualization
viz = WaterQualityVisualization(output_dir=str(resolve_subdir(self.work_dir, 'visualization')))
csv_file = self.extra.get("csv_path")
wl = self.extra.get("wavelength_start_column", "UTM_Y")
n_groups = int(self.extra.get("n_groups", 5))
param_cols = self.extra.get("param_cols") or []
if param_cols:
output_paths: List[str] = []
err_lines: List[str] = []
for param_col in param_cols:
try:
out = viz.plot_spectrum_by_parameter(
csv_path=str(csv_file),
parameter_column=param_col,
wavelength_start_column=wl,
n_groups=n_groups,
)
output_paths.append(out)
except Exception as _ex:
err_lines.append(f"{param_col}: {_ex}")
if not output_paths:
self.failed.emit(
"所有参数列的光谱图均生成失败:\n" + "\n".join(err_lines[:20])
)
return
self.finished_ok.emit(
{
"task": "spectrum",
"output_paths": output_paths,
"errors": err_lines,
}
)
else:
param_col = self.extra.get("param_col")
out = viz.plot_spectrum_by_parameter(
csv_path=str(csv_file),
parameter_column=param_col,
wavelength_start_column=wl,
n_groups=n_groups,
)
self.finished_ok.emit(
{"task": "spectrum", "output_path": out, "param_col": param_col}
)
elif self.task == "statistics":
from src.postprocessing.visualization_reports import WaterQualityVisualization
viz = WaterQualityVisualization(output_dir=str(resolve_subdir(self.work_dir, 'visualization')))
csv_file = self.extra.get("csv_path")
param_cols = self.extra.get("param_cols") or []
output_paths = viz.plot_statistical_charts(
csv_path=str(csv_file),
parameter_columns=param_cols,
)
self.finished_ok.emit(
{"task": "statistics", "output_paths": output_paths}
)
elif self.task == "scatter":
from src.core.visualization.scatter_plot import generate_model_scatter_plots
training_csv_path = (self.extra.get("training_csv_path") or "").strip()
models_dir = (self.extra.get("models_dir") or "").strip()
if not training_csv_path or not Path(training_csv_path).is_file():
self.failed.emit("训练光谱 CSV 无效或不存在,请确认已选择步骤5输出的文件。")
return
if not models_dir or not Path(models_dir).is_dir():
self.failed.emit("模型目录无效或不存在,请确认步骤6已生成 7_Supervised_Model_Training 下的参数子文件夹。")
return
scatter_paths = generate_model_scatter_plots(
models_dir=models_dir,
training_csv_path=training_csv_path,
)
self.finished_ok.emit({"task": "scatter", "scatter_paths": scatter_paths or {}})
elif self.task == "generate_all_selected":
from src.postprocessing.visualization_reports import WaterQualityVisualization
viz = WaterQualityVisualization(output_dir=str(resolve_subdir(self.work_dir, 'visualization')))
parts = []
training_csv_path = (self.extra.get("training_csv_path") or "").strip()
if training_csv_path:
training_csv = Path(training_csv_path)
else:
training_csv = wp / "6_Spectral_Feature_Extraction" / "training_spectra.csv"
if self.extra.get("gen_scatter"):
if training_csv.is_file():
models_dir_str = (self.extra.get("models_dir") or "").strip()
if models_dir_str:
models_dir = Path(models_dir_str)
else:
models_dir = wp / "8_Supervised_Model_Training"
if models_dir.is_dir() and any(d.is_dir() for d in models_dir.iterdir()):
from src.core.visualization.scatter_plot import generate_model_scatter_plots
scatter_paths = generate_model_scatter_plots(
models_dir=str(models_dir),
training_csv_path=str(training_csv),
)
count = len(scatter_paths) if scatter_paths else 0
parts.append(f"散点图: {count} 个")
else:
parts.append("散点图: 跳过(无模型目录)")
else:
parts.append("散点图: 跳过(无训练数据)")
if self.extra.get("gen_spectrum"):
if training_csv.is_file():
import pandas as pd
df = pd.read_csv(training_csv)
wl_col = _viz_infer_wavelength_start_column(df)
if isinstance(wl_col, str):
idx = int(df.columns.get_loc(wl_col)) + 1
else:
idx = int(wl_col)
param_cols = []
if idx > 0 and idx < len(df.columns):
param_cols = [
c for c in df.columns[:idx]
if df[c].dtype.kind in 'iuf' and df[c].notna().sum() > 0
]
if param_cols:
spectrum_paths = []
for param_col in param_cols:
try:
path = viz.plot_spectrum_by_parameter(
csv_path=str(training_csv),
parameter_column=param_col,
wavelength_start_column=wl_col,
n_groups=5,
)
if path:
spectrum_paths.append(path)
except Exception as e:
print(f"生成光谱图失败 ({param_col}): {e}")
count = len(spectrum_paths)
parts.append(f"光谱图: {count} 个")
else:
parts.append("光谱图: 跳过(无可用参数列)")
else:
parts.append("光谱图: 跳过(无训练数据)")
if self.extra.get("gen_boxplots"):
if training_csv.is_file():
import pandas as pd
df = pd.read_csv(training_csv)
exclude_cols = ['longitude', 'latitude', 'lon', 'lat', 'x', 'y', 'coord', 'coordinate']
param_cols = [
c for c in df.select_dtypes(include=[np.number]).columns
if not any(exc in c.lower() for exc in exclude_cols)
]
wl = _viz_infer_wavelength_start_column(df)
if isinstance(wl, str):
idx = int(df.columns.get_loc(wl)) + 1
else:
idx = int(wl)
if 0 < idx < len(df.columns):
meta_set = set(df.columns[:idx])
param_cols = [c for c in param_cols if c in meta_set]
if param_cols:
output_dict = viz.plot_statistical_charts(
csv_path=str(training_csv),
parameter_columns=param_cols,
)
count = len([v for v in output_dict.values() if v]) if output_dict else 0
parts.append(f"统计图: {count} 个")
else:
parts.append("统计图: 跳过(无可用水质参数列)")
else:
parts.append("统计图: 跳过(无训练数据)")
if self.extra.get("gen_mask_glint"):
preview_paths = viz.generate_glint_deglint_previews(
work_dir=str(wp),
output_subdir="glint_deglint_previews",
)
parts.append(f"掩膜/耀斑预览: {len(preview_paths) if preview_paths else 0} 个")
if self.extra.get("gen_sampling_map"):
hyperspectral_files = []
deglint_dir = Path(resolve_subdir(self.work_dir, 'deglint'))
if deglint_dir.exists():
for ext in ("*.dat", "*.bsq", "*.tif", "*.tiff"):
hyperspectral_files.extend(list(deglint_dir.glob(ext)))
if not hyperspectral_files:
for ext in ("*.dat", "*.bsq", "*.tif", "*.tiff"):
hyperspectral_files.extend(list(wp.glob(f"**/{ext}")))
if hyperspectral_files:
hyperspectral_path = str(hyperspectral_files[0])
csv_files = []
processed_dir = wp / "4_processed_data"
if processed_dir.exists():
csv_files = list(processed_dir.glob("*.csv"))
if not csv_files:
csv_files = (
list(wp.glob("**/*sampling*.csv"))
+ list(wp.glob("**/*point*.csv"))
+ list(wp.glob("**/*.csv"))
)
if csv_files:
csv_path = str(csv_files[0])
from src.postprocessing.point_map import SamplingPointMap
map_generator = SamplingPointMap(
output_dir=str(Path(resolve_subdir(self.work_dir, 'visualization')) / "sampling_maps"),
fast_mode=True,
)
map_path = map_generator.create_sampling_point_map(
hyperspectral_path=hyperspectral_path,
csv_path=csv_path,
point_color="red",
point_size=100,
point_alpha=0.9,
show_north_arrow=True,
show_scale_bar=True,
show_legend=True,
downsample=True,
dpi=180,
)
parts.append(f"采样点图: {Path(map_path).name}")
else:
parts.append("采样点图: 跳过(无CSV)")
else:
parts.append("采样点图: 跳过(无影像)")
if self.extra.get("gen_concentration"):
conc_dir = wp / "9_Concentration"
conc_csv = conc_dir / "final_concentrations.csv"
if conc_csv.is_file():
charts_dir = conc_dir / "charts"
charts_dir.mkdir(parents=True, exist_ok=True)
try:
import pandas as pd
df = pd.read_csv(conc_csv)
exclude_kw = (
"wavelength", "lon", "lat", "utm_x", "utm_y",
"x", "y", "coord", "longitude", "latitude",
"sample_id", "id", "index", "name", "pixel",
)
conc_cols = [
c for c in df.select_dtypes(include=[np.number]).columns
if not any(k in str(c).lower() for k in exclude_kw)
]
if conc_cols:
orig_out = viz.output_dir
viz.output_dir = str(charts_dir)
output_dict = viz.plot_statistical_charts(
csv_path=str(conc_csv),
parameter_columns=conc_cols,
)
viz.output_dir = orig_out
count = len([v for v in output_dict.values() if v])
parts.append(f"浓度统计图: {count} 个")
stats_rows = []
for col in conc_cols:
s = df[col].dropna()
if len(s) == 0:
continue
stats_rows.append({
"参数": col,
"点位数": len(s),
"最小值": round(float(s.min()), 4),
"最大值": round(float(s.max()), 4),
"均值": round(float(s.mean()), 4),
"中位数": round(float(s.median()), 4),
"标准差": round(float(s.std()), 4) if len(s) > 1 else 0.0,
})
if stats_rows:
pd.DataFrame(stats_rows).to_csv(
conc_dir / "statistics_summary.csv",
index=False,
encoding="utf-8-sig",
)
parts.append("浓度统计表: 已生成")
else:
parts.append("浓度统计表: 跳过(无可用列)")
else:
parts.append("浓度统计图: 跳过(无可用浓度列)")
except Exception as e:
parts.append(f"浓度统计图: 失败({e})")
else:
parts.append("浓度统计图: 跳过(无浓度CSV)")
if self.extra.get("gen_distribution_map"):
dist_dir = wp / "11_Thematic_Map"
out_sub = Path(viz.output_dir) / "distribution_maps"
out_sub.mkdir(parents=True, exist_ok=True)
n_rendered = 0
if dist_dir.exists():
import shutil
# 1. 拷贝现成的 PNG
for png in list(dist_dir.glob("*_distribution.png")) + list(dist_dir.glob("*_专题图.png")):
try:
shutil.copy2(png, out_sub / png.name)
n_rendered += 1
except Exception:
pass
# 2. 渲染生成的 TIF
tif_files = list(dist_dir.glob("*_distribution.tif")) + list(dist_dir.glob("*_kriging.tif"))
if tif_files:
from src.postprocessing.map import ContentMapper
mapper = ContentMapper()
# 优先级1:直接使用 Step 1 面板中缓存的外部原始 .shp 绝对路径!
boundary_path = self.extra.get("boundary_shp_path")
# 优先级2:如果没拿到,全局搜索整个工作目录下的 .shp 文件(放宽限制)
if not boundary_path:
shp_candidates = list(wp.rglob("**/*.shp"))
if shp_candidates:
boundary_path = str(shp_candidates[0])
# 优先级3:兜底使用 1_water_mask 下的栅格掩膜
if not boundary_path:
mask_files = list(wp.rglob("1_water_mask/*"))
other_candidates = [f for f in mask_files if f.suffix.lower() in ('.dat', '.bsq', '.tif', '.tiff')]
if other_candidates:
boundary_path = str(other_candidates[0])
if not boundary_path:
print(f"[distribution_maps] 未找到水域边界文件,跳过裁剪")
for tif in tif_files:
dst_png = out_sub / f"{tif.stem}_rendered.png"
try:
mapper.visualize_raster(
raster_tif_path=str(tif),
output_file=str(dst_png),
boundary_shp_path=boundary_path,
nodata_value=-9999.0,
figsize=(14, 10),
alpha=0.9
)
n_rendered += 1
except Exception as e:
print(f"渲染 TIF 失败 {tif.name}: {e}")
parts.append(f"空间分布图: {n_rendered} 个")
else:
parts.append("空间分布图: 跳过(无 11_Thematic_Map 目录)")
self.finished_ok.emit({"task": "generate_all_selected", "parts": parts})
else:
self.failed.emit(f"未知可视化任务: {self.task}")
except Exception as e:
self.failed.emit(f"{e}\n{traceback.format_exc()}")
finally:
if mpl_prev:
try:
import matplotlib.pyplot as plt
plt.switch_backend(mpl_prev)
except Exception:
pass
class PandasTableModel(QAbstractTableModel):
"""支持DataFrame的表格模型"""
def __init__(self, data_frame: pd.DataFrame):
super().__init__()
self._data = data_frame.copy()
if self._data.empty:
self._data = pd.DataFrame()
self._data.fillna("", inplace=True)
self._columns = [str(col) for col in self._data.columns]
def rowCount(self, parent=None):
return len(self._data)
def columnCount(self, parent=None):
return len(self._columns)
def data(self, index, role=Qt.DisplayRole):
if not index.isValid() or role != Qt.DisplayRole:
return None
value = self._data.iat[index.row(), index.column()]
if pd.isna(value):
return ""
return str(value)
def headerData(self, section, orientation, role=Qt.DisplayRole):
if role != Qt.DisplayRole:
return None
if orientation == Qt.Horizontal:
if section < len(self._columns):
return self._columns[section]
return str(section)
return str(section + 1)
def flags(self, index):
if not index.isValid():
return Qt.NoItemFlags
return Qt.ItemIsEnabled | Qt.ItemIsSelectable
class ChartViewerDialog(QDialog):
"""图表查看器对话框"""
def __init__(self, title="图表查看器", parent=None):
super().__init__(parent)
self.setWindowTitle(title)
self.resize(1000, 700)
self.init_ui()
def init_ui(self):
layout = QVBoxLayout()
self.figure = Figure(figsize=(10, 7))
self.canvas = FigureCanvas(self.figure)
self.canvas.setSizePolicy(QSizePolicy.Expanding, QSizePolicy.Expanding)
self.toolbar = NavigationToolbar(self.canvas, self)
layout.addWidget(self.toolbar)
layout.addWidget(self.canvas)
btn_layout = QHBoxLayout()
self.save_btn = QPushButton("保存图表")
self.save_btn.clicked.connect(self.save_chart)
btn_layout.addWidget(self.save_btn)
btn_layout.addStretch()
self.close_btn = QPushButton("关闭")
self.close_btn.clicked.connect(self.close)
btn_layout.addWidget(self.close_btn)
layout.addLayout(btn_layout)
self.setLayout(layout)
def display_image(self, image_path):
"""显示图片"""
self.figure.clear()
ax = self.figure.add_subplot(111)
try:
import matplotlib.image as mpimg
img = mpimg.imread(image_path)
ax.imshow(img)
ax.axis('off')
self.figure.tight_layout()
self.canvas.draw()
self.current_image_path = image_path
except Exception as e:
ax.text(0.5, 0.5, f'加载图片失败:\n{str(e)}',
ha='center', va='center', transform=ax.transAxes)
self.canvas.draw()
def display_custom_plot(self, plot_func):
"""显示自定义绘图函数"""
self.figure.clear()
try:
plot_func(self.figure)
self.canvas.draw()
except Exception as e:
ax = self.figure.add_subplot(111)
ax.text(0.5, 0.5, f'绘图失败:\n{str(e)}',
ha='center', va='center', transform=ax.transAxes)
self.canvas.draw()
def save_chart(self):
"""保存图表"""
file_path, _ = QFileDialog.getSaveFileName(
self, "保存图表", "",
"PNG图片 (*.png);;JPG图片 (*.jpg);;PDF文件 (*.pdf);;所有文件 (*.*)"
)
if file_path:
try:
self.figure.savefig(file_path, dpi=300, bbox_inches='tight')
QMessageBox.information(self, "成功", f"图表已保存到:\n{file_path}")
except Exception as e:
QMessageBox.critical(self, "错误", f"保存失败:\n{str(e)}")
class ImageCategoryTree(QTreeWidget):
"""现代化的图像分类目录树 - 支持智能归类、筛选和高颜值样式"""
def __init__(self, parent=None):
super().__init__(parent)
self._work_path = None
self._all_image_files = [] # 缓存所有扫描到的图片路径
self.setHeaderHidden(True) # 隐藏表头,显得更清爽
self.setAlternatingRowColors(True) # 斑马纹交替背景
self.setMaximumWidth(340)
self.setMinimumWidth(280)
# 核心:高颜值现代化 CSS 样式
self.setStyleSheet("""
QTreeWidget {
border: 1px solid #E2E8F0;
border-radius: 8px;
background-color: #FFFFFF;
alternate-background-color: #F8FAFC;
font-family: "Microsoft YaHei", "Segoe UI";
font-size: 13px;
padding: 4px;
}
QTreeWidget::item {
height: 32px;
border-radius: 6px;
margin: 2px 4px;
}
QTreeWidget::item:hover {
background-color: #F1F5F9;
}
QTreeWidget::item:selected {
background-color: #E0F2FE;
color: #0369A1;
font-weight: bold;
}
QTreeWidget::branch:has-children:!has-siblings:closed,
QTreeWidget::branch:closed:has-children:has-siblings {
border-image: none;
image: none;
}
QTreeWidget::branch:open:has-children:!has-siblings,
QTreeWidget::branch:open:has-children:has-siblings {
border-image: none;
image: none;
}
""")
def _parse_file_info(self, file_path: Path):
"""智能解析文件名,提取【水质参数】和【图表类型】"""
name_upper = file_path.name.upper()
# 1. 提取图表类型
chart_type = "其他图表"
if "DISTRIBUTION" in name_upper or "专题图" in name_upper or "RENDERED" in name_upper:
chart_type = "空间分布图"
elif "SCATTER" in name_upper or "散点" in name_upper:
chart_type = "模型散点图"
elif "SPECTRUM" in name_upper or "光谱" in name_upper:
chart_type = "光谱曲线图"
elif "HEATMAP" in name_upper or "热力图" in name_upper:
chart_type = "相关性热力图"
elif "BOXPLOT" in name_upper or "箱线" in name_upper:
chart_type = "统计箱线图"
elif "HISTOGRAM" in name_upper or "直方" in name_upper:
chart_type = "分布直方图"
elif "SAMPLING" in name_upper or "采样" in name_upper:
chart_type = "采样点地图"
elif "GLINT" in name_upper or "MASK" in name_upper or "PREVIEW" in name_upper:
chart_type = "掩膜与预览"
# 2. 提取参数名(扩展版:覆盖 12 项常见水质参数前缀)
param_name = "综合/未分类"
# ★ 匹配规则:短关键字(≤3 字符)必须在词边界(开头或 _ - . 之后)出现,
# 杜绝 "TT" 命中 "scaTTer"、"CL" 命中 "ChL" 等误匹配。
# 长关键字(≥4 字符)保持简单子串匹配,兼容历史行为。
import re as _re
def _match_key(name: str, key: str) -> bool:
clean = key.replace('-', '')
# 长关键字(≥4):简单子串匹配,历史行为一致
if len(key) >= 4 or len(clean) >= 4:
if key in name:
return True
if clean != key and clean in name:
return True
return False
# 短关键字(≤3):必须出现在词边界(行首或 _ - . 之后)。
# 若 key 本身以 _ - . 结尾,则不要求后缀边界(_ 已是分隔符)。
suffix = r'' if key[-1] in '_.-' else r'(?![a-zA-Z0-9])'
pattern = _re.compile(
r'(?:^|[_.\-])' + _re.escape(key) + suffix
)
if pattern.search(name):
return True
if clean != key:
clean_suffix = r'' if clean[-1] in '_.-' else r'(?![a-zA-Z0-9])'
pattern2 = _re.compile(
r'(?:^|[_.\-])' + _re.escape(clean) + clean_suffix
)
if pattern2.search(name):
return True
return False
params_map = {
# ── 叶绿素 a ──
'CHLOROPHYLL': '叶绿素a (Chl-a)',
'CHL-A': '叶绿素a (Chl-a)',
'CHL_A': '叶绿素a (Chl-a)',
'CHLA': '叶绿素a (Chl-a)',
'CHL_CONC': '叶绿素a (Chl-a)',
'CHL_': '叶绿素a (Chl-a)',
# ── 总悬浮物 ──
'SUSPENDED': '总悬浮物 (TSM)',
'TSM_CONC': '总悬浮物 (TSM)',
'TSM_': '总悬浮物 (TSM)',
'TSM': '总悬浮物 (TSM)',
# ── 藻蓝蛋白 ──
'PHYCOCYANIN': '藻蓝蛋白 (PC)',
'PHYCO': '藻蓝蛋白 (PC)',
'BGA': '藻蓝蛋白 (PC)',
'PC_CONC': '藻蓝蛋白 (PC)',
'PC_': '藻蓝蛋白 (PC)',
# ── 浊度 ──
'TURBIDITY': '浊度 (Turbidity)',
'TURB_CONC': '浊度 (Turbidity)',
'TURB_': '浊度 (Turbidity)',
# ── 有色可溶性有机物 ──
'CDOM': '有色可溶性有机物 (CDOM)',
# ── 透明度 ──
'SECCHI': '透明度 (SDD)',
'SDD_CONC': '透明度 (SDD)',
'SDD_': '透明度 (SDD)',
'SDD': '透明度 (SDD)',
# ── 氮磷类 ──
'NITROGEN': '总氮 (TN)',
'TN_CONC': '总氮 (TN)',
'TN_': '总氮 (TN)',
'PHOSPHORUS': '总磷 (TP)',
'TP_CONC': '总磷 (TP)',
'TP_': '总磷 (TP)',
'NH3-N': '氨氮 (NH3-N)',
'NH3_CONC': '氨氮 (NH3-N)',
'NH3_': '氨氮 (NH3-N)',
'NH3N': '氨氮 (NH3-N)',
'NO3-N': '硝态氮 (NO3-N)',
'NO3N': '硝态氮 (NO3-N)',
# ── 其他需氧量与离子 ──
'COD': '化学需氧量 (COD)',
'DISSOLVED_OXYGEN': '溶解氧 (DO)',
'DO_CONC': '溶解氧 (DO)',
'DO_': '溶解氧 (DO)',
'CL-': '氯离子 (Cl-)',
'CL_': '氯离子 (Cl-)',
# ── 基础物理与特征指标 ──
'PH': '酸碱度 (pH)',
'TEMPERATURE': '温度 (Temperature)',
'SPCOND': '电导率 (spCond)',
'TDS': '总溶解固体 (TDS)',
'TT': '特征指标 (TT)',
}
for key, display_name in params_map.items():
if _match_key(name_upper, key):
param_name = display_name
break
return param_name, chart_type
def scan_directory(self, work_dir: str):
"""全量扫描文件并缓存(不直接构建树,而是交给 rebuild_tree 渲染)"""
try:
if not work_dir: return
self._work_path = Path(work_dir)
if not self._work_path.exists(): return
# 仅扫描用于视觉展示的常规图片格式,屏蔽科学栅格 TIF 以免无法渲染报错
image_extensions = ['*.png', '*.jpg', '*.jpeg', '*.bmp']
# 扩展扫描路径
scan_roots = [
Path(resolve_subdir(str(self._work_path), 'visualization')),
Path(resolve_subdir(str(self._work_path), 'prediction_dir')),
Path(resolve_subdir(str(self._work_path), 'regression_modeling')),
self._work_path / "10_feature_construction",
self._work_path / "5_training_spectra",
Path(resolve_subdir(str(self._work_path), 'glint_detection')),
Path(resolve_subdir(str(self._work_path), 'deglint')),
Path(resolve_subdir(str(self._work_path), 'water_mask')),
self._work_path / "9_water_quality_prediction",
self._work_path / "9_Concentration",
self._work_path / "11_Thematic_Map"
]
scan_roots = [p for p in scan_roots if p.is_dir()]
if not scan_roots: scan_roots.append(self._work_path)
seen_norm = set()
self._all_image_files = []
for root in scan_roots:
for ext in image_extensions:
for p in root.rglob(ext):
key = os.path.normcase(os.path.normpath(str(p.resolve())))
if key in seen_norm: continue
seen_norm.add(key)
if p.name.startswith('.') or 'thumb' in p.name.lower(): continue
self._all_image_files.append(p)
self._all_image_files.sort(key=lambda x: x.name)
# 默认构建模式
self.rebuild_tree(group_mode='parameter', filter_type='all')
except Exception as e:
print(f"目录扫描出错: {e}")
def rebuild_tree(self, group_mode='parameter', filter_type='all'):
"""根据下拉框的【模式】和【筛选条件】实时重新构建 UI 树"""
self.blockSignals(True)
self.clear()
root_nodes = {}
for img_file in self._all_image_files:
param, chart_type = self._parse_file_info(img_file)
# 1. 应用筛选器逻辑
if filter_type != 'all' and filter_type != chart_type:
continue
# 2. 决定分组基准
if group_mode == 'parameter':
group_key = param
group_icon = "💧" if param != "综合/未分类" else "📁"
display_name = f"[{chart_type}] {img_file.name}"
elif group_mode == 'type':
group_key = chart_type
group_icon = "📊"
display_name = f"[{param.split(' ')[0]}] {img_file.name}"
else:
# 物理文件夹原样模式
try:
rel_path = img_file.relative_to(self._work_path)
group_key = str(rel_path.parent) if len(rel_path.parts) > 1 else "根目录"
except:
group_key = "其他"
group_icon = "📂"
display_name = img_file.name
# 3. 创建父节点
if group_key not in root_nodes:
root_item = QTreeWidgetItem(self)
root_item.setText(0, f"{group_icon} {group_key}")
root_item.setExpanded(True)
font = root_item.font(0)
font.setBold(True)
root_item.setFont(0, font)
root_item.setData(0, Qt.UserRole, {"type": "root"})
root_nodes[group_key] = root_item
parent_item = root_nodes[group_key]
# 4. 挂载子节点及专属图标
icon = "🖼️"
if "散点" in chart_type: icon = "📌"
elif "光谱" in chart_type: icon = "📈"
elif "分布" in chart_type: icon = "🗺️"
elif "箱线" in chart_type: icon = "📉"
image_item = QTreeWidgetItem(parent_item)
image_item.setText(0, f" {icon} {display_name}")
image_item.setData(0, Qt.UserRole, {"type": "image", "path": str(img_file)})
image_item.setToolTip(0, str(img_file))
# 统计数量
for i in range(self.topLevelItemCount()):
root_item = self.topLevelItem(i)
count = root_item.childCount()
old_text = root_item.text(0)
root_item.setText(0, f"{old_text} ({count})")
self.blockSignals(False)
def get_selected_image_path(self) -> Optional[str]:
selected_item = self.currentItem()
if not selected_item: return None
data = selected_item.data(0, Qt.UserRole)
if data and data.get("type") == "image":
return data.get("path")
return None
class ImageViewerWidget(QWidget):
"""图像查看器组件 - 支持缩放、平移"""
def __init__(self, parent=None):
super().__init__(parent)
self.current_image_path = None
self.scale_factor = 1.0
self._update_timer = QTimer()
self._update_timer.setSingleShot(True)
self._update_timer.timeout.connect(self._do_update_display)
self._pending_scale = None
self.setup_ui()
def setup_ui(self):
layout = QVBoxLayout()
layout.setContentsMargins(0, 0, 0, 0)
toolbar = QHBoxLayout()
self.refresh_btn = QPushButton("🔄 刷新目录")
self.refresh_btn.setToolTip("重新扫描工作目录中的图像文件")
toolbar.addWidget(self.refresh_btn)
separator = QFrame()
separator.setFrameShape(QFrame.VLine)
separator.setFrameShadow(QFrame.Sunken)
toolbar.addWidget(separator)
self.zoom_in_btn = QPushButton("🔍+")
self.zoom_in_btn.setToolTip("放大")
self.zoom_in_btn.setMaximumWidth(50)
toolbar.addWidget(self.zoom_in_btn)
self.zoom_out_btn = QPushButton("🔍-")
self.zoom_out_btn.setToolTip("缩小")
self.zoom_out_btn.setMaximumWidth(50)
toolbar.addWidget(self.zoom_out_btn)
self.fit_btn = QPushButton("⬜ 适应窗口")
self.fit_btn.setToolTip("适应窗口大小")
toolbar.addWidget(self.fit_btn)
self.original_btn = QPushButton("1:1 原始大小")
self.original_btn.setToolTip("原始大小")
toolbar.addWidget(self.original_btn)
self.hint_label = QLabel("💡 提示: 按住 Ctrl+滚轮 可以实现放大缩小")
self.hint_label.setStyleSheet("""
QLabel {
color: #444444;
font-size: 14px;
font-weight: bold;
padding-left: 15px;
}
""")
toolbar.addWidget(self.hint_label)
toolbar.addStretch()
self.save_btn = QPushButton("💾 保存")
self.save_btn.setToolTip("保存当前图像")
toolbar.addWidget(self.save_btn)
layout.addLayout(toolbar)
self.scroll_area = QScrollArea()
self.scroll_area.setWidgetResizable(True)
self.scroll_area.setHorizontalScrollBarPolicy(Qt.ScrollBarAsNeeded)
self.scroll_area.setVerticalScrollBarPolicy(Qt.ScrollBarAsNeeded)
self.scroll_area.setStyleSheet("background-color: white;")
self.image_label = QLabel()
self.image_label.setAlignment(Qt.AlignCenter)
self.image_label.setStyleSheet("background-color: white;")
self.scroll_area.setWidget(self.image_label)
# 全方位事件拦截:给所有可能触发滚轮的子组件全部挂载过滤器
self.image_label.installEventFilter(self)
self.scroll_area.viewport().installEventFilter(self)
self.scroll_area.installEventFilter(self)
self.scroll_area.verticalScrollBar().installEventFilter(self)
self.scroll_area.horizontalScrollBar().installEventFilter(self)
layout.addWidget(self.scroll_area, 1)
self.setLayout(layout)
self.zoom_in_btn.clicked.connect(self.zoom_in)
self.zoom_out_btn.clicked.connect(self.zoom_out)
self.fit_btn.clicked.connect(self.fit_to_window)
self.original_btn.clicked.connect(self.original_size)
self.save_btn.clicked.connect(self.save_image)
def load_image(self, image_path: str):
"""加载并显示图像"""
if not image_path or not Path(image_path).exists():
self.image_label.setText("图像不存在")
return
self.current_image_path = image_path
self.scale_factor = 1.0
pixmap = QPixmap(image_path)
if pixmap.isNull():
self.image_label.setText("无法加载图像")
return
self.original_pixmap = pixmap
self.fit_to_window()
file_info = Path(image_path).stat()
size_mb = file_info.st_size / (1024 * 1024)
def update_image_display(self):
"""更新图像显示 - 使用防抖避免频繁重绘卡顿"""
self._update_timer.stop()
self._pending_scale = self.scale_factor
self._update_timer.start(50)
def _do_update_display(self):
"""实际执行图像更新"""
if not hasattr(self, 'original_pixmap') or self.original_pixmap.isNull():
return
if self._pending_scale is None:
return
if self._pending_scale > 2.0 or self._pending_scale < 0.5:
transform = Qt.FastTransformation
else:
transform = Qt.SmoothTransformation
scaled_pixmap = self.original_pixmap.scaled(
int(self.original_pixmap.width() * self._pending_scale),
int(self.original_pixmap.height() * self._pending_scale),
Qt.KeepAspectRatio,
transform
)
self.image_label.setPixmap(scaled_pixmap)
self._pending_scale = None
def eventFilter(self, obj, event):
from PyQt5.QtCore import QEvent, Qt
if event.type() == QEvent.Wheel:
if event.modifiers() == Qt.ControlModifier:
if obj is self.scroll_area.viewport() or obj is self.image_label:
delta = event.angleDelta().y()
if delta > 0:
if getattr(self, 'scale_factor', 1.0) < 5.0:
self.scale_factor = min(self.scale_factor * 1.1, 5.0)
self.update_image_display()
else:
if getattr(self, 'scale_factor', 1.0) > 0.1:
self.scale_factor = max(self.scale_factor / 1.1, 0.1)
self.update_image_display()
return True
return super().eventFilter(obj, event)
def zoom_in(self):
"""放大"""
if self.scale_factor < 5.0:
self.scale_factor = min(self.scale_factor * 1.25, 5.0)
self.update_image_display()
def zoom_out(self):
"""缩小"""
if self.scale_factor > 0.1:
self.scale_factor = max(self.scale_factor / 1.25, 0.1)
self.update_image_display()
def fit_to_window(self):
"""适应窗口"""
if not hasattr(self, 'original_pixmap') or self.original_pixmap.isNull():
return
view_size = self.scroll_area.viewport().size()
img_size = self.original_pixmap.size()
scale_w = view_size.width() / img_size.width()
scale_h = view_size.height() / img_size.height()
self._fit_scale = min(scale_w, scale_h)
self.scale_factor = self._fit_scale
self.update_image_display()
def original_size(self):
"""原始大小"""
self.scale_factor = 1.0
self._fit_scale = None
self.update_image_display()
def save_image(self):
"""保存图像"""
if not self.current_image_path:
return
file_path, _ = QFileDialog.getSaveFileName(
self, "保存图像", Path(self.current_image_path).name,
"PNG图片 (*.png);;JPG图片 (*.jpg);;所有文件 (*.*)"
)
if file_path:
try:
import shutil
shutil.copy(self.current_image_path, file_path)
except Exception as e:
QMessageBox.critical(self, "错误", f"保存失败: {e}")
class ChartBrowserDialog(QDialog):
"""图表浏览器对话框"""
def __init__(self, chart_files, parent=None):
super().__init__(parent)
self.chart_files = sorted(chart_files, key=lambda x: x.stat().st_mtime, reverse=True)
self.current_index = 0
self.setWindowTitle("图表浏览器")
self.resize(1200, 800)
self.init_ui()
self.show_chart(0)
def init_ui(self):
layout = QVBoxLayout()
list_group = QGroupBox(f"图表列表 (共 {len(self.chart_files)} 个)")
list_layout = QHBoxLayout()
self.chart_list = QListWidget()
self.chart_list.setMaximumHeight(150)
for chart_file in self.chart_files:
self.chart_list.addItem(chart_file.name)
self.chart_list.currentRowChanged.connect(self.show_chart)
list_layout.addWidget(self.chart_list)
list_group.setLayout(list_layout)
layout.addWidget(list_group)
self.figure = Figure(figsize=(12, 8))
self.canvas = FigureCanvas(self.figure)
self.canvas.setSizePolicy(QSizePolicy.Expanding, QSizePolicy.Expanding)
self.toolbar = NavigationToolbar(self.canvas, self)
layout.addWidget(self.toolbar)
layout.addWidget(self.canvas, 1)
btn_layout = QHBoxLayout()
self.prev_btn = QPushButton("◀ 上一个")
self.prev_btn.clicked.connect(self.prev_chart)
btn_layout.addWidget(self.prev_btn)
self.next_btn = QPushButton("下一个 >")
self.next_btn.clicked.connect(self.next_chart)
btn_layout.addWidget(self.next_btn)
btn_layout.addStretch()
self.save_btn = QPushButton("💾 保存当前图表")
self.save_btn.clicked.connect(self.save_current_chart)
btn_layout.addWidget(self.save_btn)
self.close_btn = QPushButton("关闭")
self.close_btn.clicked.connect(self.close)
btn_layout.addWidget(self.close_btn)
layout.addLayout(btn_layout)
self.setLayout(layout)
def show_chart(self, index):
"""显示指定索引的图表"""
if 0 <= index < len(self.chart_files):
self.current_index = index
self.chart_list.setCurrentRow(index)
chart_file = self.chart_files[index]
self.figure.clear()
ax = self.figure.add_subplot(111)
try:
import matplotlib.image as mpimg
img = mpimg.imread(str(chart_file))
ax.imshow(img)
ax.axis('off')
ax.set_title(chart_file.name, fontsize=12, pad=10)
self.figure.tight_layout()
self.canvas.draw()
except Exception as e:
ax.text(0.5, 0.5, f'加载图片失败:\n{str(e)}',
ha='center', va='center', transform=ax.transAxes)
self.canvas.draw()
self.prev_btn.setEnabled(index > 0)
self.next_btn.setEnabled(index < len(self.chart_files) - 1)
def prev_chart(self):
"""上一个图表"""
if self.current_index > 0:
self.show_chart(self.current_index - 1)
def next_chart(self):
"""下一个图表"""
if self.current_index < len(self.chart_files) - 1:
self.show_chart(self.current_index + 1)
def save_current_chart(self):
"""保存当前图表"""
if 0 <= self.current_index < len(self.chart_files):
current_file = self.chart_files[self.current_index]
file_path, _ = QFileDialog.getSaveFileName(
self, "保存图表", current_file.name,
"PNG图片 (*.png);;JPG图片 (*.jpg);;所有文件 (*.*)"
)
if file_path:
try:
import shutil
shutil.copy(str(current_file), file_path)
QMessageBox.information(self, "成功", f"图表已保存到:\n{file_path}")
except Exception as e:
QMessageBox.critical(self, "错误", f"保存失败:\n{str(e)}")
class Step12VizPanel(QWidget):
"""步骤12:可视化展示"""
def __init__(self, parent=None):
super().__init__(parent)
self.work_dir = None
self.chart_viewer = None
self._viz_thread = None
self.init_ui()
def _viz_set_busy(self, busy: bool):
for w in (
getattr(self, "gen_all_btn", None),
getattr(self, "scan_btn", None),
):
if w is not None:
w.setEnabled(not busy)
def _start_visualization_thread(self, task: str, extra: Optional[dict] = None) -> bool:
if not self.work_dir:
QMessageBox.warning(self, "警告", "请先选择工作目录!")
return False
work_path = Path(self.work_dir)
if not work_path.exists():
QMessageBox.warning(self, "警告", "工作目录不存在!")
return False
if self._viz_thread and self._viz_thread.isRunning():
QMessageBox.information(self, "提示", "可视化任务正在运行,请稍候。")
return False
self._viz_thread = VisualizationWorkerThread(task, str(work_path), extra or {})
self._viz_thread.finished_ok.connect(self._on_visualization_worker_ok, Qt.QueuedConnection)
self._viz_thread.failed.connect(self._on_visualization_worker_fail, Qt.QueuedConnection)
self._viz_thread.finished.connect(lambda: self._viz_set_busy(False), Qt.QueuedConnection)
self._viz_set_busy(True)
self._viz_thread.start()
return True
def _spectrum_meta_param_columns(self, df: pd.DataFrame) -> List[str]:
"""光谱图可选的水质参数列(光谱波段列之前、且为数值型)。"""
wl = _viz_infer_wavelength_start_column(df)
if isinstance(wl, str):
idx = int(df.columns.get_loc(wl)) + 1
else:
idx = int(wl)
if idx <= 0 or idx >= len(df.columns):
numeric = df.select_dtypes(include=[np.number]).columns.tolist()
return [
c
for c in numeric
if not any(x in str(c).lower() for x in ("utm", "lat", "lon", "x", "y"))
]
meta = list(df.columns[:idx])
return [c for c in meta if pd.api.types.is_numeric_dtype(df[c])]
def _statistics_param_columns(self, df: pd.DataFrame) -> List[str]:
"""统计图用的参数列:只统计水质参数列(数值型),排除波长列。"""
numeric_cols = df.select_dtypes(include=[np.number]).columns.tolist()
wl = _viz_infer_wavelength_start_column(df)
if isinstance(wl, str):
idx = int(df.columns.get_loc(wl)) + 1
else:
idx = int(wl)
coord_kw = ("utm", "lat", "lon")
if 0 < idx < len(df.columns):
meta_set = set(df.columns[:idx])
return [
col
for col in numeric_cols
if col in meta_set and not any(x in str(col).lower() for x in coord_kw)
]
return [
col
for col in numeric_cols
if not any(x in str(col).lower() for x in coord_kw + ("x", "y"))
]
def _on_visualization_worker_ok(self, payload):
if not isinstance(payload, dict):
self.scan_work_directory()
return
t = payload.get("task")
if t == "mask_glint":
cnt = int(payload.get("count") or 0)
if cnt > 0:
QMessageBox.information(
self,
"成功",
f"掩膜和耀斑缩略图生成完成,共 {cnt} 个预览图。\n"
f"保存位置: 12_visualization/glint_deglint_previews/",
)
else:
QMessageBox.warning(
self,
"警告",
"未找到可处理的影像文件(2_Glint_Detection/3_deglint 等)。",
)
elif t == "sampling_map":
map_path = payload.get("map_path")
QMessageBox.information(
self,
"成功",
"采样点地图生成完成。\n"
f"输出: {Path(map_path).name if map_path else ''}\n"
"路径: 12_visualization/sampling_maps/",
)
if map_path:
self.show_chart_viewer(map_path, "采样点分布图")
elif t == "spectrum":
multi = payload.get("output_paths")
if isinstance(multi, list) and multi:
ok_paths = [p for p in multi if p and Path(str(p)).is_file()]
errs = payload.get("errors") or []
msg = (
f"已为 {len(ok_paths)} 个水质参数生成光谱对比图。\n"
f"保存目录: 工作目录/12_visualization/"
)
if errs:
msg += f"\n\n以下列未生成或出错 ({len(errs)} 项,详见日志):\n"
msg += "\n".join(str(e) for e in errs[:8])
if len(errs) > 8:
msg += "\n..."
QMessageBox.information(self, "成功", msg)
if ok_paths:
self.show_chart_viewer(ok_paths[0], "光谱曲线对比(首张)")
else:
outp = payload.get("output_path")
param = payload.get("param_col", "")
QMessageBox.information(self, "成功", f"光谱图已生成:\n{outp}")
if outp:
self.show_chart_viewer(outp, f"{param} - 光谱曲线对比")
elif t == "statistics":
outp = payload.get("output_paths") or {}
QMessageBox.information(
self, "成功", f"统计图表已生成,共 {len(outp)} 项。"
)
if isinstance(outp, dict) and "boxplot" in outp:
self.show_chart_viewer(outp["boxplot"], "水质参数箱线图")
elif t == "scatter":
paths = payload.get("scatter_paths") or {}
ok_paths = [p for p in paths.values() if p and Path(str(p)).is_file()]
if ok_paths:
QMessageBox.information(
self,
"成功",
f"已生成 {len(ok_paths)} 个模型评估散点图。\n"
f"保存位置: 12_visualization/scatter_plots/",
)
self.show_chart_viewer(ok_paths[0], "模型评估散点图")
else:
QMessageBox.warning(
self,
"提示",
"未生成任何散点图。请确认 7_Supervised_Model_Training 下已有各参数子目录及模型文件,"
"且训练 CSV 与建模时一致。",
)
elif t == "generate_all_selected":
parts = payload.get("parts") or []
QMessageBox.information(
self,
"完成",
"批量可视化已执行:\n" + "\n".join(parts) if parts else "(无选中项或已跳过)",
)
# ★ 延迟 400ms 再扫描目录,确保后台渲染线程已将图像文件 flush 到磁盘
QTimer.singleShot(400, self.scan_work_directory)
def _on_visualization_worker_fail(self, err: str):
QMessageBox.critical(self, "错误", f"可视化任务失败:\n{err[:1200]}")
def init_ui(self):
"""初始化UI - 使用全新的三列布局(控制参数 | 独立满高目录树 | 图像查看器)"""
# 1. 补上缺失的主题样式导入!
from src.gui.styles import ModernStylesheet
# 2. 增加一句热重载安全清理逻辑,彻底消除 QLayout 冲突警告
if self.layout() is not None:
QWidget().setLayout(self.layout())
main_layout = QHBoxLayout()
main_layout.setSpacing(16)
main_layout.setContentsMargins(20, 20, 20, 20)
# ==========================================
# 第一列:控制面板(目录选择 + 生成配置)
# ==========================================
control_panel = QWidget()
control_layout = QVBoxLayout()
control_layout.setContentsMargins(0, 0, 0, 0)
control_layout.setSpacing(16)
# 1. 工作目录选择
dir_group = QGroupBox("工作目录")
dir_layout = QHBoxLayout()
self.work_dir_edit = QLineEdit()
self.work_dir_edit.setPlaceholderText("选择工作目录...")
self.work_dir_edit.setReadOnly(True)
dir_browse_btn = QPushButton("浏览")
dir_browse_btn.setStyleSheet(ModernStylesheet.get_button_stylesheet('normal'))
dir_browse_btn.clicked.connect(self.browse_work_dir)
dir_layout.addWidget(self.work_dir_edit, 1)
dir_layout.addWidget(dir_browse_btn)
dir_group.setLayout(dir_layout)
control_layout.addWidget(dir_group)
# 2. 图像目录选择
img_dir_group = QGroupBox("图像目录")
img_dir_layout = QHBoxLayout()
self.img_dir_edit = QLineEdit()
self.img_dir_edit.setPlaceholderText("预测结果目录(自动填充)…")
self.img_dir_edit.setReadOnly(True)
img_dir_browse_btn = QPushButton("浏览")
img_dir_browse_btn.setStyleSheet(ModernStylesheet.get_button_stylesheet('normal'))
img_dir_browse_btn.clicked.connect(self.browse_img_dir)
img_dir_layout.addWidget(self.img_dir_edit, 1)
img_dir_layout.addWidget(img_dir_browse_btn)
img_dir_group.setLayout(img_dir_layout)
control_layout.addWidget(img_dir_group)
# 3. 可视化配置
config_group = QGroupBox("⚙️ 可视化配置")
config_layout = QVBoxLayout()
config_layout.setSpacing(12)
config_layout.setContentsMargins(16, 20, 16, 16)
self.gen_scatter = QCheckBox("模型评估散点图")
self.gen_scatter.setChecked(True)
config_layout.addWidget(self.gen_scatter)
self.gen_spectrum = QCheckBox("光谱曲线图")
self.gen_spectrum.setChecked(True)
config_layout.addWidget(self.gen_spectrum)
self.gen_boxplots = QCheckBox("统计图表")
self.gen_boxplots.setChecked(True)
config_layout.addWidget(self.gen_boxplots)
self.gen_mask_glint = QCheckBox("掩膜和耀斑缩略图")
self.gen_mask_glint.setChecked(True)
config_layout.addWidget(self.gen_mask_glint)
self.gen_sampling_map = QCheckBox("采样点地图")
self.gen_sampling_map.setChecked(True)
config_layout.addWidget(self.gen_sampling_map)
self.gen_distribution_map = QCheckBox("空间分布图 (Step 11 产物)")
self.gen_distribution_map.setChecked(True)
self.gen_distribution_map.setToolTip("渲染并汇总 Step 11 生成的 TIF 分布图")
config_layout.addWidget(self.gen_distribution_map)
config_layout.addSpacing(6)
line = QFrame()
line.setFrameShape(QFrame.HLine)
line.setStyleSheet("color: #E2E8F0;")
config_layout.addWidget(line)
config_layout.addSpacing(6)
# 彻底移除硬编码内联样式,全面接入 ModernStylesheet 主题
self.gen_all_btn = QPushButton("🚀 生成全部")
self.gen_all_btn.setStyleSheet(ModernStylesheet.get_button_stylesheet('primary'))
self.gen_all_btn.setMinimumHeight(36)
self.gen_all_btn.clicked.connect(self.generate_all_visualizations)
config_layout.addWidget(self.gen_all_btn)
self.scan_btn = QPushButton("📁 扫描目录")
self.scan_btn.setStyleSheet(ModernStylesheet.get_button_stylesheet('normal'))
self.scan_btn.setMinimumHeight(36)
self.scan_btn.clicked.connect(self.scan_work_directory)
config_layout.addWidget(self.scan_btn)
config_group.setLayout(config_layout)
control_layout.addWidget(config_group)
control_layout.addStretch()
control_panel.setLayout(control_layout)
control_panel.setMaximumWidth(290)
control_panel.setMinimumWidth(250)
main_layout.addWidget(control_panel, 0)
# ==========================================
# 第二列:满高独立的目录树(选图与筛选面板)
# ==========================================
from PyQt5.QtWidgets import QComboBox
tree_panel = QWidget()
tree_layout = QVBoxLayout()
tree_layout.setContentsMargins(0, 0, 0, 0)
tree_group = QGroupBox("🖼️ 图像浏览与筛选")
group_layout = QVBoxLayout()
group_layout.setSpacing(12)
group_layout.setContentsMargins(16, 20, 16, 16)
filter_layout = QFormLayout()
filter_layout.setContentsMargins(0, 0, 0, 0)
filter_layout.setSpacing(10)
# 彻底摘除下拉框的硬编码外壳,交由系统的全局 CSS 完美接管
self.view_mode_cb = QComboBox()
self.view_mode_cb.addItems(["按水质参数归类", "按图表类型归类", "按物理文件夹"])
self.view_mode_cb.currentIndexChanged.connect(self.update_image_tree_view)
self.chart_filter_cb = QComboBox()
self.chart_filter_cb.addItems(["全部图表", "空间分布图", "模型散点图", "光谱曲线图", "统计箱线图", "分布直方图", "相关性热力图", "掩膜与预览", "采样点地图"])
self.chart_filter_cb.currentIndexChanged.connect(self.update_image_tree_view)
filter_layout.addRow("视图模式:", self.view_mode_cb)
filter_layout.addRow("类型筛选:", self.chart_filter_cb)
group_layout.addLayout(filter_layout)
self.image_tree = ImageCategoryTree()
self.image_tree.itemClicked.connect(self.on_tree_item_clicked)
group_layout.addWidget(self.image_tree, 1)
tree_group.setLayout(group_layout)
tree_layout.addWidget(tree_group, 1)
tree_panel.setLayout(tree_layout)
tree_panel.setMaximumWidth(340)
tree_panel.setMinimumWidth(280)
main_layout.addWidget(tree_panel, 0)
# ==========================================
# 第三列:图像查看器
# ==========================================
right_panel = QWidget()
right_layout = QVBoxLayout()
right_layout.setContentsMargins(0, 0, 0, 0)
# 给画板穿上卡片外衣,确保三列视觉上的绝对统称和统一
viewer_group = QGroupBox("👁️ 图像预览区")
viewer_layout = QVBoxLayout()
viewer_layout.setContentsMargins(8, 20, 8, 8)
self.image_viewer = ImageViewerWidget()
self.image_viewer.refresh_btn.clicked.connect(self.scan_work_directory)
viewer_layout.addWidget(self.image_viewer, 1)
viewer_group.setLayout(viewer_layout)
right_layout.addWidget(viewer_group, 1)
right_panel.setLayout(right_layout)
main_layout.addWidget(right_panel, 1)
self.setLayout(main_layout)
def update_image_tree_view(self):
"""响应下拉框改变,重新渲染树状图"""
if not hasattr(self, 'image_tree') or not self.image_tree._all_image_files:
return
# 1. 提取当前选中的分组模式
mode_idx = self.view_mode_cb.currentIndex()
if mode_idx == 0:
group_mode = 'parameter'
elif mode_idx == 1:
group_mode = 'type'
else:
group_mode = 'folder'
# 2. 提取当前的图表筛选条件
filter_type = self.chart_filter_cb.currentText()
if filter_type == "全部图表":
filter_type = 'all'
# 3. 触发重绘
self.image_tree.rebuild_tree(group_mode, filter_type)
self._load_first_image_from_tree()
def set_work_dir(self, work_dir):
"""设置工作目录"""
self.work_dir = work_dir
self.work_dir_edit.setText(str(work_dir))
if work_dir:
QTimer.singleShot(100, self.scan_work_directory)
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_work_dir(self):
"""浏览工作目录"""
default = self._get_default_work_dir()
dir_path = QFileDialog.getExistingDirectory(self, "选择工作目录", default)
if dir_path:
self.work_dir = dir_path
self.work_dir_edit.setText(dir_path)
self.scan_work_directory()
def browse_img_dir(self):
"""手动浏览图像目录"""
default = self._get_default_work_dir()
dir_path = QFileDialog.getExistingDirectory(self, "选择图像目录", default)
if dir_path:
self.img_dir_edit.setText(dir_path)
self.image_tree.scan_directory(dir_path)
self._load_first_image_from_tree()
def update_from_config(self, work_dir=None, pipeline=None):
"""从全局配置自动推断并填入图像目录,然后自动加载目录内容。
推断优先级:
1. {work_dir}/9_ML_Prediction(机器学习预测)
2. {work_dir}/10_WaterIndex_CSV(水色指数反演)
3. {work_dir}/11_Thematic_Map(专题分布图)
4. {work_dir}/12_visualization(可视化目录)
5. {work_dir}(工作目录根)
"""
try:
if work_dir:
self.work_dir = work_dir
self.work_dir_edit.setText(str(work_dir))
elif not self.work_dir:
return
work_path = Path(self.work_dir)
# 按优先级寻找存在的目录
candidates = [
Path(resolve_subdir(self.work_dir, 'ml_prediction')),
Path(resolve_subdir(self.work_dir, 'watercolor')),
Path(resolve_subdir(self.work_dir, 'step11_map')),
Path(resolve_subdir(self.work_dir, 'visualization')),
work_path,
]
detected_dir = None
for candidate in candidates:
if candidate.exists() and candidate.is_dir():
detected_dir = candidate
break
if detected_dir:
detected_str = str(detected_dir)
self.img_dir_edit.setText(detected_str)
self.image_tree.scan_directory(detected_str)
else:
# 无预测目录时扫描整个工作目录
self.image_tree.scan_directory(self.work_dir)
# 自动触发加载第一张图像
self._load_first_image_from_tree()
except Exception as e:
import traceback
print(f"可视化面板 update_from_config 出错: {e}")
traceback.print_exc()
def _load_first_image_from_tree(self):
"""自动加载树状图中的第一张有效图片(兼容物理目录层级结构)"""
try:
tree = getattr(self, 'image_tree', None)
if not tree:
return
from PyQt5.QtCore import Qt
def find_first_image(item):
# 检查当前节点是否是图片节点
data = item.data(0, Qt.UserRole)
if isinstance(data, dict) and data.get("type") == "image":
return item
# 如果不是,递归检查所有子节点
for i in range(item.childCount()):
found = find_first_image(item.child(i))
if found:
return found
return None
# 遍历所有顶层节点
for i in range(tree.topLevelItemCount()):
first_img_item = find_first_image(tree.topLevelItem(i))
if first_img_item:
tree.setCurrentItem(first_img_item)
# 主动触发一次点击槽函数,以在右侧渲染图片
self.on_tree_item_clicked(first_img_item, 0)
return
except Exception as e:
import traceback
print(f"自动加载首张图片失败: {e}")
traceback.print_exc()
def scan_work_directory(self):
"""扫描工作目录中的图像文件"""
if not self.work_dir:
return
work_path = Path(self.work_dir)
if not work_path.exists():
return
print(f"扫描工作目录: {work_path}")
self.image_tree.scan_directory(str(work_path))
self._setup_prediction_output_dirs(work_path)
viz_dir = Path(resolve_subdir(str(work_path), 'visualization'))
if viz_dir.exists():
image_files = list(viz_dir.glob("**/*.png")) + list(viz_dir.glob("**/*.jpg"))
if image_files:
self.image_viewer.load_image(str(image_files[0]))
def _setup_prediction_output_dirs(self, work_path: Path):
"""收集预测输出目录路径信息(不创建目录,仅用于日志/调试)。
2026-06-30 修复:移除 mkdir 调用,不再在未运行 pipeline 时创建空目录。
目录创建统一留给各 pipeline 步骤在实际执行时处理。
"""
try:
ml_dir = Path(resolve_subdir(str(work_path), 'ml_prediction'))
watercolor_dir = Path(resolve_subdir(str(work_path), 'watercolor'))
thematic_dir = Path(resolve_subdir(str(work_path), 'step11_map'))
# 仅输出信息,不创建目录
existing = [str(d) for d in (ml_dir, watercolor_dir, thematic_dir) if d.is_dir()]
if existing:
print(f"预测输出目录已存在: {existing}")
except Exception as e:
print(f"读取预测输出目录信息失败: {e}")
def on_tree_item_clicked(self, item, column):
"""目录树项点击事件"""
data = item.data(0, Qt.UserRole)
if not data:
return
if data.get("type") == "image":
image_path = data.get("path")
if image_path and Path(image_path).exists():
self.image_viewer.load_image(image_path)
def generate_all_visualizations(self):
"""生成所有可视化图表"""
if not self.work_dir:
QMessageBox.warning(self, "警告", "请先选择工作目录!")
return
work_path = Path(self.work_dir)
if not work_path.exists():
QMessageBox.warning(self, "警告", "工作目录不存在!")
return
if not (self.gen_scatter.isChecked() or self.gen_spectrum.isChecked() or
self.gen_boxplots.isChecked() or self.gen_mask_glint.isChecked() or
self.gen_sampling_map.isChecked() or self.gen_distribution_map.isChecked()):
QMessageBox.information(self, "提示", "请至少勾选一项可视化配置选项以生成图表。")
return
reply = QMessageBox.question(
self, "确认生成",
"将根据左侧勾选项在后台生成可视化图表,可能需要较长时间。\n是否继续?",
QMessageBox.Yes | QMessageBox.No
)
if reply != QMessageBox.Yes:
return
extra = {
"gen_scatter": self.gen_scatter.isChecked(),
"gen_spectrum": self.gen_spectrum.isChecked(),
"gen_boxplots": self.gen_boxplots.isChecked(),
"gen_mask_glint": self.gen_mask_glint.isChecked(),
"gen_sampling_map": self.gen_sampling_map.isChecked(),
"gen_distribution_map": self.gen_distribution_map.isChecked(),
}
main_window = self.window()
factory = getattr(main_window, '_panel_factory', None) if main_window else None
# [新增] 直接从 Step 1 面板读取原始 .shp 的绝对路径,突破工作目录限制
step1_panel = factory.get_panel('step1') if factory else None
if step1_panel:
s1_conf = step1_panel.get_config()
s1_mask = s1_conf.get('mask_path')
# 确保文件存在且是shp格式,存入extra透传给后台线程
if s1_mask and Path(s1_mask).is_file() and str(s1_mask).lower().endswith('.shp'):
extra["boundary_shp_path"] = str(s1_mask)
step6_panel = factory.get_panel('step6_feature') if factory else None
if step6_panel and getattr(step6_panel, 'output_file', None):
_resolved_csv = step6_panel.output_file.get_path()
if _resolved_csv:
extra["training_csv_path"] = _resolved_csv
step8_panel = factory.get_panel('step8_ml_train') if factory else None
if step8_panel and getattr(step8_panel, 'output_path', None):
_resolved_models_dir = step8_panel.output_path.get_path()
if _resolved_models_dir:
extra["models_dir"] = _resolved_models_dir
self._start_visualization_thread("generate_all_selected", extra)
def generate_chart(self, chart_type):
"""生成图表"""
if not self.work_dir:
QMessageBox.warning(self, "警告", "请先选择工作目录!")
return
work_path = Path(self.work_dir)
if not work_path.exists():
QMessageBox.warning(self, "警告", "工作目录不存在!")
return
try:
main_window = self.window()
factory = getattr(main_window, '_panel_factory', None) if main_window else None
step6_panel = factory.get_panel('step6_feature') if factory else None
if step6_panel and getattr(step6_panel, 'output_file', None) and step6_panel.output_file.get_path():
training_spectra_csv = Path(step6_panel.output_file.get_path())
else:
training_spectra_csv = _viz_training_spectra_csv_path(work_path)
if chart_type == 'scatter':
if not training_spectra_csv.is_file():
QMessageBox.warning(
self, "警告",
"未找到 6_Spectral_Feature_Extraction\\training_spectra.csv。\n"
"请先执行步骤6(光谱特征提取)生成该文件。",
)
return
training_csv = training_spectra_csv
models_dir = work_path / "7_Supervised_Model_Training"
if not models_dir.is_dir() or not any(d.is_dir() for d in models_dir.iterdir()):
mdir = QFileDialog.getExistingDirectory(
self, "选择模型根目录(内含各水质参数子文件夹)", str(work_path))
if not mdir:
return
models_dir = Path(mdir)
self._start_visualization_thread(
"scatter",
{"training_csv_path": str(training_csv), "models_dir": str(models_dir)},
)
return
if chart_type == 'spectrum':
if not training_spectra_csv.is_file():
QMessageBox.warning(
self, "警告",
"未找到 6_Spectral_Feature_Extraction\\training_spectra.csv。\n"
"光谱分析固定使用该文件,请先执行步骤6(光谱特征提取)。",
)
return
csv_file = training_spectra_csv
df = pd.read_csv(csv_file)
columns = self._spectrum_meta_param_columns(df)
if not columns:
QMessageBox.warning(
self, "警告",
"当前 CSV 中没有可用的数值型水质参数列,无法按参数分组绘制光谱图。",
)
return
wl_col = _viz_infer_wavelength_start_column(df)
self._start_visualization_thread(
"spectrum",
{"csv_path": str(csv_file), "param_cols": columns,
"wavelength_start_column": wl_col, "n_groups": 5},
)
return
if chart_type == 'statistics':
if not training_spectra_csv.is_file():
QMessageBox.warning(
self, "警告",
"未找到 6_Spectral_Feature_Extraction\\training_spectra.csv。\n"
"统计分析固定使用该文件,请先执行步骤6(光谱特征提取)。",
)
return
csv_file = training_spectra_csv
df = pd.read_csv(csv_file)
param_cols = self._statistics_param_columns(df)
if not param_cols:
QMessageBox.warning(self, "警告", "未找到可用的水质参数列!")
return
self._start_visualization_thread(
"statistics",
{"csv_path": str(csv_file), "param_cols": param_cols},
)
return
if chart_type == 'sampling_map':
self.generate_sampling_point_map()
return
except ImportError:
QMessageBox.critical(
self, "错误",
"无法导入可视化模块!\n请确保 visualization_reports.py 文件存在。",
)
except Exception as e:
QMessageBox.critical(
self, "错误",
f"生成图表时出错:\n{str(e)}\n\n{traceback.format_exc()}",
)
def generate_mask_glint_previews(self):
"""生成掩膜和耀斑缩略图"""
self._start_visualization_thread("mask_glint")
def generate_sampling_point_map(self):
"""生成采样点地图"""
self._start_visualization_thread("sampling_map")
def view_chart(self, chart_type):
"""查看图表"""
if not self.work_dir:
QMessageBox.warning(self, "警告", "请先选择工作目录!")
return
work_path = Path(self.work_dir)
viz_dir = Path(resolve_subdir(self.work_dir, 'visualization'))
viz_dir2 = viz_dir / "boxplots"
viz_dir3 = viz_dir / "scatter_plots"
if not viz_dir.exists():
QMessageBox.warning(self, "警告", f"可视化目录不存在:\n{viz_dir}\n\n请先生成图表。")
return
chart_files = []
if chart_type == 'scatter':
chart_files = list(viz_dir3.glob("*scatter*.png"))
elif chart_type == 'spectrum':
chart_files = list(viz_dir.glob("*spectrum*.png"))
elif chart_type == 'statistics':
chart_files = list(viz_dir2.glob("*boxplot.png")) + \
list(viz_dir.glob("*histogram.png")) + \
list(viz_dir.glob("*heatmap.png"))
elif chart_type == 'distribution':
chart_files = list(viz_dir.glob("**/*distribution.png"))
elif chart_type == 'mask_glint':
glint_dir = viz_dir / "glint_deglint_previews"
chart_files = list(glint_dir.glob("*preview.png")) if glint_dir.exists() else \
list(viz_dir.glob("*preview.png")) + \
list(viz_dir.glob("*glint*.png")) + \
list(viz_dir.glob("*mask*.png"))
elif chart_type == 'sampling_map':
sampling_dir = viz_dir / "sampling_maps"
chart_files = list(sampling_dir.glob("*sampling_map.png")) if sampling_dir.exists() else \
list(viz_dir.glob("*sampling*.png"))
if not chart_files:
QMessageBox.warning(self, "警告", f"未找到{chart_type}类型的图表文件!\n\n请先生成图表。")
return
if len(chart_files) > 1:
from PyQt5.QtWidgets import QInputDialog
file_names = [f.name for f in chart_files]
file_name, ok = QInputDialog.getItem(
self, "选择图表", "请选择要查看的图表:", file_names, 0, False)
if ok:
selected_file = next(f for f in chart_files if f.name == file_name)
self.show_chart_viewer(str(selected_file), file_name)
else:
self.show_chart_viewer(str(chart_files[0]), chart_files[0].name)
def browse_all_charts(self):
"""浏览所有图表"""
if not self.work_dir:
QMessageBox.warning(self, "警告", "请先选择工作目录!")
return
work_path = Path(self.work_dir)
chart_files = list(work_path.glob("**/*.png")) + list(work_path.glob("**/*.jpg"))
if not chart_files:
QMessageBox.warning(self, "警告", "未找到图表文件!")
return
dialog = ChartBrowserDialog(chart_files, self)
dialog.exec_()
def show_chart_viewer(self, image_path, title="图表查看器"):
"""显示图表查看器"""
viewer = ChartViewerDialog(title=title, parent=self)
viewer.display_image(image_path)
viewer.exec_()
def get_config(self):
"""获取配置"""
return {
'generate_scatter': self.gen_scatter.isChecked(),
'generate_boxplots': self.gen_boxplots.isChecked(),
'generate_spectrum': self.gen_spectrum.isChecked(),
'generate_glint_previews': self.gen_mask_glint.isChecked(),
'generate_sampling_maps': self.gen_sampling_map.isChecked(),
'generate_distribution_maps': self.gen_distribution_map.isChecked(),
'scatter_config': {
'metric': 'test_r2', 'feature_start_column': 13,
'test_size': 0.2, 'random_state': 42
},
'boxplot_config': {
'data_start_column': 4, 'save_individual': True, 'use_seaborn': True
},
'glint_preview_config': {
'work_dir': None, 'output_subdir': 'glint_deglint_previews',
'generate_glint': True, 'generate_deglint': True
}
}
def set_config(self, config):
"""设置配置"""
if not config:
return
if 'generate_scatter' in config:
self.gen_scatter.setChecked(config['generate_scatter'])
if 'generate_boxplots' in config:
self.gen_boxplots.setChecked(config['generate_boxplots'])
if 'generate_spectrum' in config:
self.gen_spectrum.setChecked(config['generate_spectrum'])
if 'generate_glint_previews' in config:
self.gen_mask_glint.setChecked(config['generate_glint_previews'])
if 'generate_sampling_maps' in config:
self.gen_sampling_map.setChecked(config.get('generate_sampling_maps', True))
if 'generate_distribution_maps' in config:
self.gen_distribution_map.setChecked(config.get('generate_distribution_maps', True))