setWidgetResizable(False)+QSizePolicy.Ignored 导致 QLabel 无有效 sizeHint, 图像不显示。 回退为 setWidgetResizable(True), 保留显式 ScrollBarAsNeeded 策略 (原默认可能为 AlwaysOff 导致滚动条异常)。
2023 lines
88 KiB
Python
2023 lines
88 KiB
Python
#!/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))
|