7 Commits

24 changed files with 3769 additions and 3279 deletions

180
_smoke_test_step10.py Normal file
View File

@ -0,0 +1,180 @@
"""
Smoke test for Step 10 散点 CSV 模式 (WaterIndexCsvProcessor)
模拟 Step 4 输出格式 (sampling_spectra.csv):
x_coord, y_coord, pixel_x, pixel_y, "400.000000", "401.000000", ...
验证 WaterIndexCsvProcessor.compute_indices_from_csv:
1. 正确读取 x_coord/y_coord → 重命名为 longitude/latitude
2. 正确识别数字列名 = 光谱列
3. 复用 WaterQualityIndexCalculator 逐行计算
4. 输出每个公式一个 CSV,三列严格为 longitude, latitude, <formula_name>
5. 公式值数量级合理 (非 NaN,非 inf)
"""
import os
import sys
import tempfile
import shutil
from pathlib import Path
# 让脚本能找到项目根
PROJECT_ROOT = Path(__file__).parent
sys.path.insert(0, str(PROJECT_ROOT))
def create_synthetic_sampling_csv(path: str, n_points: int = 5):
"""模拟 Step 4 输出: x_coord, y_coord, pixel_x, pixel_y, 数字列名光谱"""
import csv
# 选一组关键波段(确保 waterindex.csv 中的 BGA_Am09KBBI 等公式都能找到)
wavelengths = [400.0, 443.0, 458.0, 486.0, 500.0, 510.0, 531.0, 547.0, 555.0,
615.0, 622.0, 629.0, 644.0, 658.0, 665.0, 672.0, 681.0, 686.0,
700.0, 709.0, 714.0, 715.0, 753.0, 857.0, 900.0]
fieldnames = ['x_coord', 'y_coord', 'pixel_x', 'pixel_y'] + [f'{w:.6f}' for w in wavelengths]
with open(path, 'w', newline='', encoding='utf-8-sig') as f:
writer = csv.DictWriter(f, fieldnames=fieldnames)
writer.writeheader()
for i in range(n_points):
# 模拟水体光谱(典型内陆湖泊反射率 0.005-0.05)
row = {
'x_coord': 100.0 + i * 10,
'y_coord': 30.0 + i * 5,
'pixel_x': 100 + i,
'pixel_y': 30 + i,
}
for w in wavelengths:
# 简单合成光谱: 蓝光 < 红光 + 一点叶绿素峰
base = 0.01 + 0.0001 * (w - 400)
chl_peak = 0.005 * (1 - abs(w - 560) / 200) if abs(w - 560) < 200 else 0
row[f'{w:.6f}'] = round(base + chl_peak, 6)
writer.writerow(row)
def run_smoke():
print("=" * 70)
print("Step 10 散点 CSV 模式 Smoke Test")
print("=" * 70)
tmpdir = tempfile.mkdtemp(prefix="step10_smoke_")
print(f"Tempdir: {tmpdir}")
sampling_csv = os.path.join(tmpdir, "sampling_spectra.csv")
output_dir = os.path.join(tmpdir, "10_WaterIndex_CSV")
os.makedirs(output_dir, exist_ok=True)
# 1) 创建合成 sampling CSV
create_synthetic_sampling_csv(sampling_csv, n_points=5)
print(f"Created sampling CSV: {sampling_csv}")
with open(sampling_csv, encoding='utf-8-sig') as f:
header_line = f.readline().strip()
print(f" header: {header_line[:120]}...")
# 2) 找到项目自带的 waterindex.csv
waterindex_csv = PROJECT_ROOT / "src" / "gui" / "model" / "waterindex.csv"
print(f"Using waterindex.csv: {waterindex_csv}")
assert waterindex_csv.is_file(), "waterindex.csv not found!"
# 3) 调 WaterIndexCsvProcessor
# 注意:父包 __init__.py 顶部有 `from osgeo import gdal, osr`,
# 在没装 gdal 的 venv 里任何 from ...waterindex_inversion import ... 都会炸。
# 这里用 importlib 按文件路径直接加载 csv_processor.py 子模块,
# 完全绕开 __init__.py 的 osgeo 加载。
import importlib.util as _ilu
_csv_proc_path = (
PROJECT_ROOT / "src" / "core" / "algorithms" / "waterindex_inversion"
/ "csv_processor.py"
)
_spec = _ilu.spec_from_file_location("waterindex_csv_processor", _csv_proc_path)
_mod = _ilu.module_from_spec(_spec)
sys.modules["waterindex_csv_processor"] = _mod
_spec.loader.exec_module(_mod)
WaterIndexCsvProcessor = _mod.WaterIndexCsvProcessor
progress_log = []
def progress_cb(msg, pct):
progress_log.append((msg, pct))
proc = WaterIndexCsvProcessor(str(waterindex_csv))
print(f"\n[Step] compute_indices_from_csv...")
out_files = proc.compute_indices_from_csv(
sampling_csv_path=sampling_csv,
output_dir=output_dir,
selected_formulas=["BGA_Am09KBBI", "BGA_Da052BDA", "BGA_Be16NDPhyI"],
progress_callback=progress_cb,
)
print(f"\n[Result] Generated {len(out_files)} CSV files:")
for name, path in out_files.items():
size = os.path.getsize(path)
print(f" {name:30s} -> {os.path.basename(path)} ({size} bytes)")
# 4) 验证每个输出 CSV 的列结构
print(f"\n[Verify] Column structure check:")
all_pass = True
import pandas as pd
for name, path in out_files.items():
df = pd.read_csv(path, encoding='utf-8-sig')
cols = list(df.columns)
expected = ['longitude', 'latitude', name]
ok = (cols == expected) and (len(df) == 5)
flag = "✓" if ok else "✗"
if not ok:
all_pass = False
print(f" {flag} {name:30s} cols={cols} rows={len(df)}")
# 5) 验证坐标重命名
print(f"\n[Verify] Coordinate rename (x_coord→longitude, y_coord→latitude):")
sample = pd.read_csv(out_files[list(out_files.keys())[0]], encoding='utf-8-sig')
print(f" longitude values: {sample['longitude'].tolist()}")
print(f" latitude values: {sample['latitude'].tolist()}")
coord_ok = (sample['longitude'].iloc[0] == 100.0 and
sample['latitude'].iloc[0] == 30.0)
if not coord_ok:
all_pass = False
print(f" {'✓' if coord_ok else '✗'} coordinate rename correct")
# 6) 验证公式值非 NaN
print(f"\n[Verify] Formula values (no NaN):")
for name, path in out_files.items():
df = pd.read_csv(path, encoding='utf-8-sig')
col = df[name]
n_nan = col.isna().sum()
n_inf = ((col == float('inf')) | (col == float('-inf'))).sum()
all_nan = col.dropna().empty
if n_nan > 0 or n_inf > 0 or all_nan:
print(f" ✗ {name:30s}: NaN={n_nan} Inf={n_inf} empty={all_nan}")
print(f" values: {col.tolist()}")
all_pass = False
else:
mn, mx = col.min(), col.max()
print(f" ✓ {name:30s}: range=[{mn:.4f}, {mx:.4f}]")
# 7) 进度回调检查
print(f"\n[Verify] Progress callback:")
print(f" Total progress events: {len(progress_log)}")
if progress_log:
first_msg, first_pct = progress_log[0]
last_msg, last_pct = progress_log[-1]
print(f" First: ({first_pct:.1f}%) {first_msg}")
print(f" Last : ({last_pct:.1f}%) {last_msg}")
progress_ok = (last_pct == 100.0)
if not progress_ok:
all_pass = False
print(f" {'✓' if progress_ok else '✗'} last progress = 100%")
# 8) 总结
print(f"\n{'=' * 70}")
if all_pass:
print(f"✓ ALL CHECKS PASSED")
else:
print(f"✗ SOME CHECKS FAILED — inspect output above")
print(f"{'=' * 70}")
# 清理
shutil.rmtree(tmpdir, ignore_errors=True)
return all_pass
if __name__ == "__main__":
success = run_smoke()
sys.exit(0 if success else 1)

View File

@ -644,3 +644,17 @@ class WaterIndexProcessor:
notify("水色指数反演完成", 100)
return results
# ------------------------------------------------------------------
# 散点处理入口(Step 10 重构后使用,与 Step 9 对称)
# ------------------------------------------------------------------
# WaterIndexCsvProcessor 已拆出到独立子模块 csv_processor.py,
# 目的是让纯 CSV 计算链路不再被 __init__.py 顶部 osgeo import 拖垮。
# 这里做一次 re-export,保留所有 `from src.core.algorithms.waterindex_inversion import WaterIndexCsvProcessor`
# 这类已有 import 路径仍能正常工作(生产环境/打包后)。
from src.core.algorithms.waterindex_inversion.csv_processor import WaterIndexCsvProcessor # noqa: E402,F401
# 保留旧 import 路径兼容
__all__ = ['WaterIndexProcessor', 'WaterIndexCsvProcessor']

View File

@ -0,0 +1,240 @@
# -*- coding: utf-8 -*-
"""
水色指数反演 — 散点 CSV 模式处理器(独立子模块)。
设计意图
--------
本模块与 ``waterindex_inversion.__init__.py`` 中的 ``WaterIndexProcessor``
(栅格 BSQ 模式) **彻底解耦**,不依赖任何 osgeo / rasterio / gdal,仅依赖
``pandas`` 与 ``src.utils.water_index.WaterQualityIndexCalculator``。
**为什么独立成文件?**
``__init__.py`` 顶部有 ``from osgeo import gdal, osr``(用于 BSQ 栅格模式),
这意味着任何 ``from src.core.algorithms.waterindex_inversion import X``
都会触发 osgeo 加载——而某些验证环境(无 gdal 包的 venv)会因此 ImportError。
本子模块独立后,可通过
``from src.core.algorithms.waterindex_inversion.csv_processor import WaterIndexCsvProcessor``
直接加载,**完全不触发** ``__init__.py`` 的 osgeo import 链,便于无 gdal 环境
做端到端 smoke test。
**调用入口(由 Step 10 service / panel 调用)**::
from src.core.algorithms.waterindex_inversion.csv_processor import WaterIndexCsvProcessor
proc = WaterIndexCsvProcessor(waterindex_csv_path)
out = proc.compute_indices_from_csv(
sampling_csv_path=...,
output_dir=...,
selected_formulas=[...],
progress_callback=lambda msg, pct: ...,
)
输出格式
--------
每个公式一个 CSV,三列严格为 ``longitude, latitude, <formula_name>``。
"""
from __future__ import annotations
import os
import re
from typing import Callable, Dict, List, Optional
import numpy as np # P0 修复: 写盘前 inf/-inf 替换与极值截断需要
class WaterIndexCsvProcessor:
"""
散点 CSV 驱动的水色指数反演器。
设计目的
--------
与 Step 9 (ML 预测) 完全对称的【散点处理模式】:
* 输入:Step 4 生成的 ``sampling_spectra.csv``,列结构为
``x_coord, y_coord, pixel_x, pixel_y, 400.000000, 401.000000, ...``
* 处理:解析 ``waterindex.csv`` 中的公式,对每行采样点
提取对应波段数值、逐行 eval 计算水色指数
* 输出:每个公式一个 CSV,列严格为 ``longitude, latitude, <formula_name>``,
可直接喂给 Step 11 ContentMapper
输出目录
--------
默认 ``{work_dir}/10_WaterIndex_CSV/``;若用户指定 ``output_dir`` 则用其值。
"""
COORD_RENAME_MAP = {
"x_coord": "longitude",
"y_coord": "latitude",
"lon": "longitude",
"lat": "latitude",
}
def __init__(self, waterindex_csv_path: Optional[str] = None):
if waterindex_csv_path is None:
candidates = [
os.path.join(os.path.dirname(__file__), '..', '..', 'gui', 'model', 'waterindex.csv'),
os.path.join(os.path.dirname(__file__), '..', '..', '..', 'gui', 'model', 'waterindex.csv'),
]
for p in candidates:
if os.path.isfile(p):
waterindex_csv_path = p
break
self.waterindex_csv_path = waterindex_csv_path
self._index_calc = None
def _get_index_calc(self):
"""懒加载 WaterQualityIndexCalculator(首次访问时实例化)"""
if self._index_calc is None and self.waterindex_csv_path:
from src.utils.water_index import WaterQualityIndexCalculator
self._index_calc = WaterQualityIndexCalculator(self.waterindex_csv_path)
return self._index_calc
@staticmethod
def _detect_wavelength_columns(df: "object") -> List[str]:
"""识别光谱列:列名是浮点数字符串(Step 4 输出的 '400.000000' 形式)"""
wl_cols = []
for col in df.columns:
try:
float(str(col).strip())
wl_cols.append(col)
except (ValueError, TypeError):
continue
return wl_cols
@staticmethod
def _safe_filename(name: str) -> str:
"""公式名 → 文件名安全字符(与旧 BSQ 输出命名习惯一致)"""
return re.sub(r'[^\w\u4e00-\u9fff-]', '_', name).strip('_') or 'index'
def compute_indices_from_csv(
self,
sampling_csv_path: str,
output_dir: str,
selected_formulas: Optional[List[str]] = None,
progress_callback: Optional[Callable[[str, float], None]] = None,
) -> Dict[str, str]:
"""
散点 CSV → 按指数拆分的多个 CSV。
Parameters
----------
sampling_csv_path : str
Step 4 输出的 ``sampling_spectra.csv`` 路径
output_dir : str
输出目录;不存在会自动创建
selected_formulas : list, optional
要计算的公式名列表;None 或空列表 = 全部公式
progress_callback : callable, optional
进度回调 ``(msg: str, pct: float)``
Returns
-------
dict
``{公式名: 输出 CSV 路径}``
"""
def notify(msg: str, pct: float) -> None:
if progress_callback:
progress_callback(msg, pct)
if not os.path.isfile(sampling_csv_path):
raise FileNotFoundError(f"采样点 CSV 不存在: {sampling_csv_path}")
if not self.waterindex_csv_path or not os.path.isfile(self.waterindex_csv_path):
raise FileNotFoundError(
f"waterindex.csv 未配置或不存在: {self.waterindex_csv_path}"
)
os.makedirs(output_dir, exist_ok=True)
notify("正在读取采样点 CSV…", 5)
import pandas as pd
df = pd.read_csv(sampling_csv_path, encoding="utf-8-sig")
if df.empty:
raise ValueError(f"采样点 CSV 为空: {sampling_csv_path}")
# 坐标列重命名(x_coord → longitude, y_coord → latitude)
df = df.rename(columns={k: v for k, v in self.COORD_RENAME_MAP.items()
if k in df.columns})
if "longitude" not in df.columns or "latitude" not in df.columns:
raise ValueError(
f"采样点 CSV 缺少坐标列(期望 x_coord/y_coord 或 longitude/latitude),"
f"实际列: {list(df.columns)}"
)
# 识别光谱列
wl_cols = self._detect_wavelength_columns(df)
if not wl_cols:
raise ValueError(
f"采样点 CSV 中未识别到任何光谱列(列名为数字),"
f"实际列: {list(df.columns)}"
)
notify(f"识别到 {len(wl_cols)} 个光谱列, 采样点 {len(df)} 个", 15)
calc = self._get_index_calc()
if calc is None:
raise RuntimeError("WaterQualityIndexCalculator 初始化失败")
all_formula_names = calc.list_available()
if selected_formulas:
targets = [n for n in selected_formulas if n in all_formula_names]
missing = [n for n in selected_formulas if n not in all_formula_names]
if missing:
print(f"[WaterIndexCsvProcessor] 警告: 以下公式未在 waterindex.csv 中找到,已跳过: {missing}")
else:
targets = all_formula_names
if not targets:
raise ValueError("没有可计算的公式(selected_formulas 为空且 waterindex.csv 中无公式)")
# 一次性算出所有目标公式的 Series(避免重复遍历 DataFrame)
notify(f"开始逐行计算 {len(targets)} 个公式…", 25)
spectra_df = df[wl_cols]
try:
results_df = calc.calculate_many(targets, spectra_df)
except Exception as e:
raise RuntimeError(f"公式计算失败: {e}")
# 每个公式一个 CSV:longitude, latitude, <formula_name>
out_files: Dict[str, str] = {}
n_total = len(targets)
for i, name in enumerate(targets):
try:
per_idx = results_df[name]
# ===== P0 防御: 写盘前清洗(防 Step 11 Kriging / TIN 碎玻璃)=====
# 1) inf / -inf → NaN:pandas 默认会把 inf 写成字面 "Infinity",
# 下游 ContentMapper 严格按位置读第 3 列时会原样拿到 inf,
# Kriging 变差函数被单一 inf 点拉飞、TIN 出现极长退化三角形。
# 2) 物理范围截断 ±10:保守大窗,覆盖几乎所有经验公式的合法值域
# (NDCI 类 ∈ [-1,1]、比值类 ∈ [0,2]、浓度反演类 ∈ [0, ∞))。
# 注:保留所有行(包括 NaN),让下游 Step 11 把它当 missing marker;
# dropna 会破坏行号对齐,故不在此调用;如确需剔除 NaN 行,
# 在面板 / service 层读取后自行 .dropna(subset=[col])。
n_inf = int(np.isinf(per_idx.values).sum())
per_idx = (
per_idx
.replace([np.inf, -np.inf], np.nan)
.clip(lower=-10.0, upper=10.0)
)
if n_inf > 0:
print(f"[WaterIndexCsvProcessor] {name}: 替换 {n_inf} 个 inf/-inf → NaN")
# ===== P0 防御结束 =====
out_df = pd.DataFrame({
"longitude": df["longitude"].values,
"latitude": df["latitude"].values,
name: per_idx.values,
})
out_path = os.path.join(output_dir, f"{self._safe_filename(name)}.csv")
out_df.to_csv(out_path, index=False, float_format="%.6f", encoding="utf-8-sig")
out_files[name] = out_path
notify(
f"[{i + 1}/{n_total}] {name} → {os.path.basename(out_path)}",
25 + 70 * (i + 1) / n_total,
)
except Exception as e:
print(f"[WaterIndexCsvProcessor] 公式 '{name}' 失败: {e}")
continue
notify(f"完成!共输出 {len(out_files)} / {n_total} 个指数 CSV", 100)
return out_files

View File

@ -51,7 +51,10 @@ class Step12KrigingHandler(BaseStepHandler):
output_image_path=output_image_path,
resolution=config.get('resolution', 30),
input_crs=config.get('input_crs', 'EPSG:32651'),
output_crs=config.get('output_crs', 'EPSG:4326'),
# ★★★ 强制 output_crs = input_crs,避免 ContentMapper 把栅格重投影到 EPSG:4326 ★★★
# 旧实现:output_crs=config.get('output_crs', 'EPSG:4326')
# 重投影会让栅格和基于投影坐标的掩膜在 visualize_raster 叠加时发生仿射变换撕裂
output_crs=config.get('input_crs', 'EPSG:32651'),
show_sample_points=config.get('show_sample_points', False),
base_map_tif=config.get('base_map_tif'),
use_distance_diffusion=config.get('use_distance_diffusion', True),

View File

@ -16,11 +16,14 @@ class MappingStep:
@staticmethod
def generate_distribution_map(
prediction_csv_path: str,
boundary_shp_path: str,
boundary_shp_path: Optional[str] = None, # ★★★ Plan C: None = 不依赖水域掩膜 ★★★
output_image_path: Optional[str] = None,
resolution: float = 30,
input_crs: str = "EPSG:32651",
output_crs: str = "EPSG:4326",
# ★★★ 强制默认 output_crs = input_crs,禁止重投影到 EPSG:4326 ★★★
# 历史默认值 'EPSG:4326' 会让 ContentMapper 将插值栅格从投影坐标系转到经纬度坐标系,
# 与基于 EPSG:32651 的水域掩膜叠加时发生仿射变换撕裂(栅格错位、坐标轴扭曲)。
output_crs: str = "EPSG:32651",
show_sample_points: bool = False,
base_map_tif: Optional[str] = None,
use_distance_diffusion: bool = True,
@ -31,13 +34,14 @@ class MappingStep:
expand_ratio: float = 0.05,
output_dir: Union[str, Path] = "./14_visualization",
callback: Optional[Callable] = None,
output_format: str = 'tif', # ⭐ 新增:'tif' (默认, 写 GeoTIFF) / 'png'
) -> str:
"""
根据采样点的坐标和反演的实测参数,通过插值方法得到水质参数可视化分布图
Args:
prediction_csv_path: 预测结果CSV文件路径(前两列为经纬度,第三列为预测值)
boundary_shp_path: 边界shapefile文件路径
boundary_shp_path: 边界/掩膜文件路径(.shp/.dat/.bsq/.tif 等)。None 时跳过水域掩膜约束。
output_image_path: 输出图片路径(如果为None,自动生成)
resolution: 插值网格分辨率(米)
input_crs: 输入坐标系
@ -89,6 +93,7 @@ class MappingStep:
"diffusion_power": diffusion_power,
"diffusion_n_neighbors": diffusion_n_neighbors,
"expand_ratio": expand_ratio,
"output_format": output_format, # ⭐ 透传给 ContentMapper.process_data
}
optional_kwargs = {

View File

@ -201,7 +201,7 @@ PANEL_REGISTRY = [
'title': '专题图生成',
'icon': '10.png',
'stage': '模块四 制图与成果汇编',
'display_name': '11. 专题图生成',
'display_name': '11. 分布图生成',
# 架构解耦(2026-06-22):dict 键名 = 下游目标控件真实属性名
'dependencies': {
# 目标框: self.prediction_csv_dir_edit ← 上游 step9_ml_predict.output_file

View File

@ -165,24 +165,43 @@ class BandConfirmDialog(QDialog):
from PyQt5.QtCore import QSettings
AI_SETTINGS_ORG = "IrisWaterQuality"
AI_SETTINGS_APP = "WQ_GUI"
# 扩充预设字典,覆盖市面主流大模型标准接口
AI_DEFAULTS = {
"ollama": {
"api_base_url": "http://localhost:11434",
"vision_model": "qwen3-vl:8b",
"text_model": "qwen3-vl:8b",
"aliyun": {
"api_base_url": "https://dashscope.aliyuncs.com/compatible-mode/v1/chat/completions",
"vision_model": "qwen-vl-max",
"text_model": "qwen-max",
},
"zhipu": {
"api_base_url": "https://open.bigmodel.cn/api/paas/v4/chat/completions",
"vision_model": "glm-4v",
"text_model": "glm-4",
},
"deepseek": {
"api_base_url": "https://api.deepseek.com/chat/completions",
"vision_model": "deepseek-chat", # DeepSeek暂无独立视觉API,可用通用或自行更换
"text_model": "deepseek-chat",
},
"openai": {
"api_base_url": "https://api.openai.com/v1/chat/completions",
"vision_model": "gpt-4o",
"text_model": "gpt-4o",
},
"minimax": {
"api_base_url": "https://api.minimaxi.com/v1/text/chatcompletion_v2",
"vision_model": "abab6.5s-chat",
"api_base_url": "https://api.minimax.chat/v1/chat/completions",
"vision_model": "abab6.5g-chat",
"text_model": "abab6.5s-chat",
},
"ollama": {
"api_base_url": "http://localhost:11434",
"vision_model": "qwen2-vl",
"text_model": "qwen2.5",
},
}
class AISettingsDialog(QDialog):
"""AI 引擎可视化配置弹窗,配置持久化到 QSettings。"""
@ -195,41 +214,22 @@ class AISettingsDialog(QDialog):
self._init_ui()
def _load_settings(self):
"""从 QSettings 读取已有配置;无记录则回退到环境变量或默认值。"""
"""从 QSettings 读取已有配置;无记录则回退到预设字典或环境变量。"""
s = QSettings(AI_SETTINGS_ORG, AI_SETTINGS_APP)
self._provider = s.value("ai_provider", "minimax", type=str)
# 默认推荐 Aliyun (通义千问)
self._provider = s.value("ai_provider", "Aliyun", type=str)
# API Key 不设默认值(敏感信息,首次必须由用户输入)
self._api_key = s.value("minimax_api_key", "", type=str)
# 【向后兼容补丁】优先读取新规范的 api_key,如果为空,尝试读取旧版本留下的 minimax_api_key
self._api_key = s.value("api_key", "", type=str)
if not self._api_key:
self._api_key = s.value("minimax_api_key", "", type=str)
# 已保存的 URL 和模型;若 QSettings 无记录则读环境变量
if self._provider == "ollama":
self._api_base_url = (
s.value("api_base_url", "")
or os.environ.get("OLLAMA_URL", AI_DEFAULTS["ollama"]["api_base_url"])
)
self._vision_model = (
s.value("vision_model", "")
or os.environ.get("OLLAMA_VISION_MODEL", AI_DEFAULTS["ollama"]["vision_model"])
)
self._text_model = (
s.value("text_model", "")
or os.environ.get("OLLAMA_TEXT_MODEL", AI_DEFAULTS["ollama"]["text_model"])
)
else:
self._api_base_url = (
s.value("api_base_url", "")
or os.environ.get("MINIMAX_BASE_URL", AI_DEFAULTS["minimax"]["api_base_url"])
)
self._vision_model = (
s.value("vision_model", "")
or os.environ.get("MINIMAX_VISION_MODEL", AI_DEFAULTS["minimax"]["vision_model"])
)
self._text_model = (
s.value("text_model", "")
or os.environ.get("MINIMAX_TEXT_MODEL", AI_DEFAULTS["minimax"]["text_model"])
)
provider_key = self._provider.lower()
defaults = AI_DEFAULTS.get(provider_key, {"api_base_url": "", "vision_model": "", "text_model": ""})
self._api_base_url = s.value("api_base_url", "", type=str) or defaults.get("api_base_url", "")
self._vision_model = s.value("vision_model", "", type=str) or defaults.get("vision_model", "")
self._text_model = s.value("text_model", "", type=str) or defaults.get("text_model", "")
self._timeout = s.value("timeout_s", 120, type=int)
def _init_ui(self):
@ -240,9 +240,15 @@ class AISettingsDialog(QDialog):
provider_row = QHBoxLayout()
provider_row.addWidget(QLabel("AI 引擎提供商:"))
self._provider_combo = QComboBox()
self._provider_combo.addItems(["Ollama", "Minimax"])
self._provider_combo.setCurrentText("Ollama" if self._provider == "ollama" else "Minimax")
self._provider_combo.currentIndexChanged.connect(self._on_provider_changed)
# ★ 核心改动:开启可编辑模式,允许用户随意输入第三方代理商名字
self._provider_combo.setEditable(True)
self._provider_combo.addItems(["Aliyun", "Zhipu", "DeepSeek", "OpenAI", "Minimax", "Ollama"])
self._provider_combo.setCurrentText(self._provider)
# 当文本改变时自动带出推荐配置
self._provider_combo.currentTextChanged.connect(self._on_provider_changed)
provider_row.addWidget(self._provider_combo, 1)
provider_row.addStretch(1)
layout.addLayout(provider_row)
@ -251,7 +257,7 @@ class AISettingsDialog(QDialog):
url_row = QHBoxLayout()
url_row.addWidget(QLabel("API Base URL:"))
self._url_edit = QLineEdit(self._api_base_url)
self._url_edit.setPlaceholderText("例如: http://localhost:11434")
self._url_edit.setPlaceholderText("填入兼容 OpenAI 规范的完整 URL")
url_row.addWidget(self._url_edit, 1)
layout.addLayout(url_row)
@ -259,7 +265,7 @@ class AISettingsDialog(QDialog):
key_row = QHBoxLayout()
key_row.addWidget(QLabel("API Key:"))
self._key_edit = QLineEdit(self._api_key)
self._key_edit.setPlaceholderText("输入 API Key(敏感信息,已加密存储)")
self._key_edit.setPlaceholderText("输入 API Key(本地加密存储)")
self._key_edit.setEchoMode(QLineEdit.Password)
key_row.addWidget(self._key_edit, 1)
layout.addLayout(key_row)
@ -288,10 +294,10 @@ class AISettingsDialog(QDialog):
# ── 说明 ──────────────────────────────────────────────────────────────
hint = QLabel(
"提示:切换引擎后将自动填充推荐默认值(可手动修改)。"
"API Key 仅本地加密存储,不会明文暴露。"
"提示:可以直接在下拉框输入任意名称。选择预设服务商会自动填充推荐的兼容接口 URL。\n"
"若使用全能多模态大模型(如 gpt-4o / qwen-vl-max 等),视觉与文本模型填入相同名称即可。"
)
hint.setStyleSheet("color: #888; font-size: 10px;")
hint.setStyleSheet("color: #888; font-size: 11px;")
hint.setWordWrap(True)
layout.addWidget(hint)
@ -306,18 +312,20 @@ class AISettingsDialog(QDialog):
btn_box.addButton(cancel_btn, QDialogButtonBox.RejectRole)
layout.addWidget(btn_box)
def _on_provider_changed(self):
"""切换 Provider 时自动填充推荐默认值。"""
provider = self._provider_combo.currentText().lower()
defaults = AI_DEFAULTS.get(provider, AI_DEFAULTS["minimax"])
self._url_edit.setText(defaults["api_base_url"])
self._vision_edit.setText(defaults["vision_model"])
self._text_edit.setText(defaults["text_model"])
def _on_provider_changed(self, text):
"""切换或输入 Provider 时自动填充推荐默认值。"""
provider_key = text.lower()
if provider_key in AI_DEFAULTS:
defaults = AI_DEFAULTS[provider_key]
self._url_edit.setText(defaults["api_base_url"])
self._vision_edit.setText(defaults["vision_model"])
self._text_edit.setText(defaults["text_model"])
def _save_and_close(self):
"""持久化到 QSettings 并关闭。"""
s = QSettings(AI_SETTINGS_ORG, AI_SETTINGS_APP)
provider = self._provider_combo.currentText().lower()
# 获取用户输入的文本(无论是选的还是自己打字的)
provider = self._provider_combo.currentText().strip()
s.setValue("ai_provider", provider)
s.setValue("api_base_url", self._url_edit.text().strip())
s.setValue("api_key", self._key_edit.text().strip())
@ -334,11 +342,13 @@ class AISettingsDialog(QDialog):
返回键:ai_provider / api_base_url / api_key / vision_model / text_model / timeout_s
"""
s = QSettings(AI_SETTINGS_ORG, AI_SETTINGS_APP)
provider = s.value("ai_provider", "minimax", type=str)
# 默认返回 Aliyun
provider = s.value("ai_provider", "Aliyun", type=str)
return {
"ai_provider": provider,
"api_base_url": s.value("api_base_url", "", type=str),
"api_key": s.value("api_key", "", type=str),
# 同样向后兼容读取旧版本的 key
"api_key": s.value("api_key", s.value("minimax_api_key", "", type=str), type=str),
"vision_model": s.value("vision_model", "", type=str),
"text_model": s.value("text_model", "", type=str),
"timeout_s": s.value("timeout_s", 120, type=int),

View File

@ -1,10 +1,17 @@
#!/usr/bin/env python
# -*- coding: utf-8 -*-
"""
Step10 面板 - 水色指数反演(直接处理去耀斑 BSQ 影像)
Step10 面板 - 水色指数反演(散点 CSV 模式)
将 waterindex.csv 中的公式直接应用于去耀斑高光谱影像,
输出各水质参数指数的 GeoTIFF 栅格图像。
与 Step 9 (ML 预测) 完全对称的【散点处理模式】:
* 输入:Step 4 输出的 ``sampling_spectra.csv``(散点+全波段光谱)
* 处理:解析 ``waterindex.csv`` 中的公式,**逐行**对每个采样点计算水色指数
* 输出:每个公式一个 CSV,列严格为 ``longitude, latitude, <formula_name>``,
可直接喂给 Step 11 ContentMapper
* 输出目录:默认 ``{work_dir}/10_WaterIndex_CSV/``
注:原"读 BSQ 全图→GeoTIFF"模式已废弃(科学上误差大且与 GIS 栅格计算器重复)。
"""
import os
@ -34,79 +41,66 @@ from src.gui.styles import ModernStylesheet
class WaterIndexWorker(QThread):
"""后台线程:执行水色指数反演"""
finished_ok = pyqtSignal(dict)
failed = pyqtSignal(str)
"""后台线程:散点 CSV → 逐行公式计算 → 多 CSV 输出
应用 Step 13 QThread 协议:
- 三信号 ``progress / finished / error`` 命名严格遵循约定
(注意 ``finished`` 不能用——会覆盖 QThread 内建同名信号,
导致 _on_finished 不被回调时按钮不会恢复)
- 进度回调两参 (msg: str, pct: float)
"""
progress = pyqtSignal(str, float) # message, percent
log = pyqtSignal(str)
finished_ok = pyqtSignal(dict) # {公式名: 输出 CSV 路径}
error = pyqtSignal(str) # error message
def __init__(
self,
bsq_path: str,
hdr_path: str,
sampling_csv_path: str,
output_dir: str,
selected_formulas: List[str],
waterindex_csv: str,
water_mask_path: Optional[str] = None,
work_dir: Optional[str] = None,
):
super().__init__()
self.bsq_path = bsq_path
self.hdr_path = hdr_path
self.sampling_csv_path = sampling_csv_path
self.output_dir = output_dir
self.selected_formulas = selected_formulas
self.waterindex_csv = waterindex_csv
self.water_mask_path = water_mask_path
self.work_dir = work_dir
def run(self):
try:
from src.core.algorithms.waterindex_inversion import WaterIndexProcessor
self.progress.emit("正在初始化水色指数处理器…", 2)
processor = WaterIndexProcessor(self.waterindex_csv)
self.progress.emit("正在读取影像元数据…", 5)
# 获取影像元数据
meta = processor.get_image_metadata(self.bsq_path, self.hdr_path)
if not meta:
self.failed.emit("无法读取影像元数据,请检查 BSQ 和 HDR 文件是否匹配")
return
n_bands = meta.get('bands', 0)
wv_range = meta.get('wavelength_range', '未知')
self.log.emit(
f"影像信息: {meta['width']}×{meta['height']} 像素, "
f"{n_bands} 波段, {wv_range}"
from src.core.algorithms.waterindex_inversion import (
WaterIndexCsvProcessor,
)
if self.water_mask_path:
self.log.emit(f"使用水域掩膜: {self.water_mask_path}")
self.progress.emit("正在初始化散点水色指数处理器…", 2)
# 使用 run_inversion 入口(含掩膜拦截链路)
results = processor.run_inversion(
deglint_img_path=self.bsq_path,
work_dir=self.work_dir or self.output_dir,
formula_csv_path=self.waterindex_csv,
selected_formulas=self.selected_formulas,
water_mask_path=self.water_mask_path,
callback=self._on_progress,
processor = WaterIndexCsvProcessor(self.waterindex_csv)
# 散点 CSV → 逐公式一个 CSV
out_files = processor.compute_indices_from_csv(
sampling_csv_path=self.sampling_csv_path,
output_dir=self.output_dir,
selected_formulas=self.selected_formulas or None,
progress_callback=lambda m, p: self.progress.emit(m, p),
)
self.progress.emit(f"完成!共生成 {len(results)} 个指数图", 100)
self.finished_ok.emit(results)
self.progress.emit(
f"完成!共生成 {len(out_files)} 个指数 CSV", 100
)
self.finished_ok.emit(out_files)
except FileNotFoundError as e:
self.error.emit(f"文件不存在: {e}")
except ValueError as e:
self.error.emit(f"参数错误: {e}")
except Exception as e:
self.failed.emit(f"{e}\n{traceback.format_exc()}")
def _on_progress(self, msg: str, pct: float):
self.progress.emit(msg, pct)
self.error.emit(f"{e}\n{traceback.format_exc()}")
class Step10WatercolorPanel(QWidget):
"""步骤10:水色指数反演(直接处理 BSQ 影像)"""
"""步骤10:水色指数反演(散点 CSV 模式)"""
def __init__(self, parent=None):
super().__init__(parent)
@ -122,49 +116,41 @@ class Step10WatercolorPanel(QWidget):
layout = QVBoxLayout()
# ---- 标题 ----
title = QLabel("步骤10:水色指数反演(高光谱影像直接处理)")
title = QLabel("步骤10:水色指数反演(散点 CSV 模式)")
title.setFont(QFont("Arial", 12, QFont.Bold))
layout.addWidget(title)
# ---- 说明 ----
hint = QLabel(
"将 waterindex.csv 中的公式直接应用于去耀斑高光谱影像(BSQ),"
"输出各水质参数指数的 GeoTIFF 栅格图像。"
"指数图可直接用于水质专题图生成。"
"读取 Step 4 生成的 sampling_spectra.csv 散点光谱,"
"对每个采样点逐行套用 waterindex.csv 中勾选的公式,"
"输出每公式一个 CSV(列:longitude, latitude, 公式值)。"
"结果可被 Step 11 直接以 ContentMapper 模式消费。"
)
hint.setWordWrap(True)
hint.setStyleSheet(f"color: {ModernStylesheet.COLORS.get('text_secondary', '#666')};")
layout.addWidget(hint)
# ---- 输入影像选择 ----
input_group = QGroupBox("输入影像")
# ---- 输入采样点数据 ----
input_group = QGroupBox("输入采样点数据")
input_layout = QFormLayout()
self.bsq_file = FileSelectWidget(
"去耀斑 BSQ 影像:",
"BSQ Files (*.bsq);;DAT Files (*.dat);;All Files (*.*)"
self.sampling_csv_file = FileSelectWidget(
"采样点 CSV:",
"CSV Files (*.csv);;All Files (*.*)"
)
self.bsq_file.line_edit.setPlaceholderText("选择去耀斑处理后的 BSQ 影像")
self.bsq_file.browse_btn.clicked.disconnect()
self.bsq_file.browse_btn.clicked.connect(self._browse_bsq)
input_layout.addRow("BSQ 影像:", self.bsq_file)
self.hdr_file = FileSelectWidget(
"ENVI 头文件:",
"HDR Files (*.hdr);;All Files (*.*)"
self.sampling_csv_file.line_edit.setPlaceholderText(
"选择 Step 4 输出的 sampling_spectra.csv"
)
self.hdr_file.line_edit.setPlaceholderText("自动关联同路径 .hdr 文件")
self.hdr_file.browse_btn.clicked.disconnect()
self.hdr_file.browse_btn.clicked.connect(self._browse_hdr)
input_layout.addRow("HDR 文件:", self.hdr_file)
input_layout.addRow("采样点 CSV:", self.sampling_csv_file)
# 影像信息显示
self.meta_label = QLabel("未加载影像")
# 数据规模提示(运行后回填,避免启动时强制 read_csv)
self.meta_label = QLabel("未加载采样点数据")
self.meta_label.setStyleSheet(
"background: #f0f0f0; padding: 4px 8px; border-radius: 4px; "
"font-size: 12px; color: #333;"
)
input_layout.addRow("影像信息:", self.meta_label)
input_layout.addRow("数据信息:", self.meta_label)
input_group.setLayout(input_layout)
layout.addWidget(input_group)
@ -211,16 +197,11 @@ class Step10WatercolorPanel(QWidget):
"输出目录:",
"Directories"
)
self.output_dir.line_edit.setPlaceholderText("留空 → 工作目录/10_WaterIndex_Images")
self.output_dir.browse_btn.clicked.disconnect()
self.output_dir.browse_btn.clicked.connect(self._browse_output_dir)
self.output_dir.line_edit.setPlaceholderText(
"留空 → 工作目录/10_WaterIndex_CSV"
)
output_layout.addRow("输出目录:", self.output_dir)
self.format_combo = QComboBox()
self.format_combo.addItems(["GTiff (GeoTIFF)", "ENVI", "PCI"])
self.format_combo.setCurrentIndex(0)
output_layout.addRow("输出格式:", self.format_combo)
output_group.setLayout(output_layout)
layout.addWidget(output_group)
@ -337,60 +318,24 @@ class Step10WatercolorPanel(QWidget):
def _on_item_changed(self, item: QListWidgetItem):
pass # 可扩展:实时统计选中数量
def _browse_bsq(self):
path, _ = QFileDialog.getOpenFileName(
self, "选择去耀斑 BSQ 影像",
"",
"BSQ Files (*.bsq);;DAT Files (*.dat);;All Files (*.*)"
)
if path:
self.bsq_file.set_path(path)
# 自动关联同路径 hdr
hdr = Path(path).with_suffix('.hdr')
if hdr.exists():
self.hdr_file.set_path(str(hdr))
self._load_metadata(path, str(hdr) if hdr.exists() else "")
def _browse_hdr(self):
path, _ = QFileDialog.getOpenFileName(
self, "选择 ENVI 头文件",
"",
"HDR Files (*.hdr);;All Files (*.*)"
)
if path:
self.hdr_file.set_path(path)
bsq_path = self.bsq_file.get_path()
if bsq_path:
self._load_metadata(bsq_path, path)
def _browse_output_dir(self):
d = QFileDialog.getExistingDirectory(self, "选择输出目录", "")
if d:
self.output_dir.set_path(d)
def _load_metadata(self, bsq_path: str, hdr_path: str):
"""加载并显示影像元数据"""
if not bsq_path or not Path(bsq_path).exists():
self.meta_label.setText("⚠️ 影像文件不存在")
def _refresh_sampling_meta(self):
"""从 sampling_csv 路径快速 peek 数据规模(不触发公式计算)"""
path = self.sampling_csv_file.get_path().strip()
if not path:
self.meta_label.setText("未加载采样点数据")
return
if not hdr_path or not Path(hdr_path).exists():
self.meta_label.setText("⚠️ 头文件不存在")
if not Path(path).exists():
self.meta_label.setText("⚠️ 采样点 CSV 不存在")
return
try:
from src.core.algorithms.waterindex_inversion import WaterIndexProcessor
processor = WaterIndexProcessor(self._waterindex_csv)
meta = processor.get_image_metadata(bsq_path, hdr_path)
if meta:
self.meta_label.setText(
f"✅ {meta['width']}×{meta['height']} | "
f"{meta['bands']} 波段 | {meta.get('wavelength_range', '未知')} | "
f"驱动: {meta['driver']}"
)
else:
self.meta_label.setText("⚠️ 无法读取元数据")
import pandas as pd
df = pd.read_csv(path, encoding="utf-8-sig", nrows=0)
n_cols = len(df.columns)
self.meta_label.setText(
f"✅ 已选采样点 CSV({n_cols} 列,完整列数将在运行时打印)"
)
except Exception as e:
self.meta_label.setText(f"⚠️ 元数据读取失败: {e}")
self.meta_label.setText(f"⚠️ 读取失败: {e}")
def _get_selected_formula_names(self) -> List[str]:
names = []
@ -411,21 +356,20 @@ class Step10WatercolorPanel(QWidget):
return ""
def get_config(self) -> dict:
bsq = self.bsq_file.get_path()
return {
'bsq_path': bsq,
'hdr_path': self.hdr_file.get_path(),
'deglint_img_path': bsq,
'output_dir': self.output_dir.get_path(),
'output_format': self.format_combo.currentText().split()[0],
sampling = self.sampling_csv_file.get_path().strip()
config: Dict[str, object] = {
'sampling_csv_path': sampling,
'selected_formulas': self._get_selected_formula_names(),
}
out_dir = self.output_dir.get_path().strip()
if out_dir:
config['output_dir'] = out_dir
return config
def set_config(self, config: dict):
if config.get('bsq_path'):
self.bsq_file.set_path(config['bsq_path'])
if config.get('hdr_path'):
self.hdr_file.set_path(config['hdr_path'])
if config.get('sampling_csv_path'):
self.sampling_csv_file.set_path(config['sampling_csv_path'])
self._refresh_sampling_meta()
if config.get('output_dir'):
self.output_dir.set_path(config['output_dir'])
if 'selected_formulas' in config:
@ -444,42 +388,50 @@ class Step10WatercolorPanel(QWidget):
self.work_dir = None
main_window = self.window()
deglint_path = None
# 1. 优先从 pipeline 的真实输出中获取
# 1. 优先从 pipeline.step_outputs 取 Step 4 的采样点 CSV 路径
sampling_path = None
if pipeline and hasattr(pipeline, 'step_outputs'):
step3_out = pipeline.step_outputs.get('step3', {})
deglint_path = step3_out.get('deglint_image') or step3_out.get('output_path')
step4_out = pipeline.step_outputs.get('step4_sampling', {})
sampling_path = (
step4_out.get('sampling_csv')
or step4_out.get('output_path')
or step4_out.get('output_file')
)
# 2. 回退:从 step3 面板实例获取
if not deglint_path and main_window and hasattr(main_window, 'step3_panel'):
if hasattr(main_window.step3_panel, 'output_file'):
deglint_path = main_window.step3_panel.output_file.get_path()
# 2. 回退:直接读 step4_sampling panel 的 output_file widget
if not sampling_path and main_window:
step4_widget = getattr(main_window, 'step4_sampling', None)
if step4_widget and hasattr(step4_widget, 'output_file'):
sampling_path = step4_widget.output_file.get_path()
else:
# 通过 _panel_factory 懒加载查找
factory = getattr(main_window, '_panel_factory', None)
if factory:
step4_panel = factory.get_panel('step4_sampling')
if step4_panel and hasattr(step4_panel, 'output_file'):
sampling_path = step4_panel.output_file.get_path()
# 3. 终极回退:智能扫描 3_deglint 目录,取最新的 .bsq 或 .dat 文件
if not deglint_path and self.work_dir:
deglint_dir = resolve_subdir(self.work_dir, 'deglint')
if os.path.isdir(deglint_dir):
import glob
candidates = glob.glob(os.path.join(deglint_dir, "*.bsq")) + glob.glob(os.path.join(deglint_dir, "*.dat"))
if candidates:
candidates.sort(key=os.path.getmtime, reverse=True)
deglint_path = candidates[0]
# 3. 终极回退:扫描 work_dir/4_sampling/sampling_spectra.csv
if not sampling_path and self.work_dir:
candidate = resolve_subdir(self.work_dir, 'sampling_csv_path')
if os.path.isfile(candidate):
sampling_path = candidate
# 填入 UI 并自动寻找对应的 hdr 文件
if deglint_path:
if not os.path.isabs(deglint_path):
deglint_path = os.path.join(self.work_dir or '', deglint_path).replace('\\', '/')
self.bsq_file.set_path(deglint_path)
# 填入 UI
if sampling_path:
if not os.path.isabs(sampling_path):
sampling_path = os.path.join(
self.work_dir or '', sampling_path
).replace('\\', '/')
self.sampling_csv_file.set_path(sampling_path)
self._refresh_sampling_meta()
hdr_path = os.path.splitext(deglint_path)[0] + '.hdr'
if os.path.exists(hdr_path):
self.hdr_file.set_path(hdr_path)
self._load_metadata(deglint_path, hdr_path)
# 自动填入输出目录
# 自动填入输出目录(默认 work_dir/10_WaterIndex_CSV/)
if self.work_dir:
out_dir = resolve_subdir(self.work_dir, 'watercolor')
out_dir = os.path.join(
self.work_dir, '10_WaterIndex_CSV'
).replace('\\', '/')
os.makedirs(out_dir, exist_ok=True)
if not self.output_dir.get_path():
self.output_dir.set_path(out_dir)
@ -488,30 +440,20 @@ class Step10WatercolorPanel(QWidget):
"""通过 EventBus 发布单步执行请求(解耦面板与 PipelineExecutor)。"""
from src.gui.core.event_bus import global_event_bus
bsq_path = self.bsq_file.get_path().strip()
hdr_path = self.hdr_file.get_path().strip()
output_dir = self.output_dir.get_path().strip()
sampling_csv_path = self.sampling_csv_file.get_path().strip()
if not sampling_csv_path:
QMessageBox.warning(self, "输入错误", "请选择采样点 CSV!")
return
if not Path(sampling_csv_path).exists():
QMessageBox.warning(
self, "输入错误", f"采样点 CSV 不存在:\n{sampling_csv_path}"
)
return
if not bsq_path:
QMessageBox.warning(self, "输入错误", "请选择去耀斑 BSQ 影像!")
return
if not Path(bsq_path).exists():
QMessageBox.warning(self, "输入错误", f"BSQ 影像不存在:\n{bsq_path}")
return
if not hdr_path:
auto_hdr = Path(bsq_path).with_suffix('.hdr')
if auto_hdr.exists():
hdr_path = str(auto_hdr)
self.hdr_file.set_path(hdr_path)
else:
QMessageBox.warning(self, "输入错误", "请选择 ENVI 头文件!")
return
if not Path(hdr_path).exists():
QMessageBox.warning(self, "输入错误", f"HDR 文件不存在:\n{hdr_path}")
return
output_dir = self.output_dir.get_path().strip()
if not output_dir:
work_dir = self._get_default_work_dir()
output_dir = resolve_subdir(work_dir, 'watercolor')
output_dir = os.path.join(work_dir, '10_WaterIndex_CSV')
os.makedirs(output_dir, exist_ok=True)
self.output_dir.set_path(output_dir)
@ -521,43 +463,34 @@ class Step10WatercolorPanel(QWidget):
return
if self._waterindex_csv and not Path(self._waterindex_csv).exists():
QMessageBox.warning(self, "配置错误", f"waterindex.csv 不存在:\n{self._waterindex_csv}")
QMessageBox.warning(
self, "配置错误",
f"waterindex.csv 不存在:\n{self._waterindex_csv}",
)
return
config = {'step7_index': self.get_config()}
config = {'step10_watercolor': self.get_config()}
global_event_bus.publish('RequestRunSingleStep', {
'step_name': 'step7_index',
'step_name': 'step10_watercolor',
'config': config,
})
def run_step(self):
"""独立运行步骤10(旧版 parent 链上溯方式,保留兼容)。"""
bsq_path = self.bsq_file.get_path().strip()
hdr_path = self.hdr_file.get_path().strip()
output_dir = self.output_dir.get_path().strip()
sampling_csv_path = self.sampling_csv_file.get_path().strip()
if not sampling_csv_path:
QMessageBox.warning(self, "输入错误", "请选择采样点 CSV!")
return
if not Path(sampling_csv_path).exists():
QMessageBox.warning(
self, "输入错误", f"采样点 CSV 不存在:\n{sampling_csv_path}"
)
return
# 验证输入
if not bsq_path:
QMessageBox.warning(self, "输入错误", "请选择去耀斑 BSQ 影像!")
return
if not Path(bsq_path).exists():
QMessageBox.warning(self, "输入错误", f"BSQ 影像不存在:\n{bsq_path}")
return
if not hdr_path:
# 尝试自动查找
auto_hdr = Path(bsq_path).with_suffix('.hdr')
if auto_hdr.exists():
hdr_path = str(auto_hdr)
self.hdr_file.set_path(hdr_path)
else:
QMessageBox.warning(self, "输入错误", "请选择 ENVI 头文件!")
return
if not Path(hdr_path).exists():
QMessageBox.warning(self, "输入错误", f"HDR 文件不存在:\n{hdr_path}")
return
output_dir = self.output_dir.get_path().strip()
if not output_dir:
work_dir = self._get_default_work_dir()
output_dir = resolve_subdir(work_dir, 'watercolor')
output_dir = os.path.join(work_dir, '10_WaterIndex_CSV')
os.makedirs(output_dir, exist_ok=True)
self.output_dir.set_path(output_dir)
@ -567,25 +500,13 @@ class Step10WatercolorPanel(QWidget):
return
if self._waterindex_csv and not Path(self._waterindex_csv).exists():
QMessageBox.warning(self, "配置错误", f"waterindex.csv 不存在:\n{self._waterindex_csv}")
QMessageBox.warning(
self, "配置错误",
f"waterindex.csv 不存在:\n{self._waterindex_csv}",
)
return
# ── 自动扫描工作目录下的水域掩膜文件 ────────────────────────────
work_dir = self.work_dir or str(Path(bsq_path).parent)
mask_dir = resolve_subdir(work_dir, 'water_mask')
water_mask_path: Optional[str] = None
if os.path.isdir(mask_dir):
# ★★★ glob 智能扫描:取任意 .dat 或 .tif 文件 ★★★
for pattern in ("*.dat", "*.tif", "*.TIF", "*.DT"):
candidates = sorted(Path(mask_dir).glob(pattern))
if candidates:
water_mask_path = str(candidates[0])
break
if water_mask_path:
print(f"[Step8] 自动找到水域掩膜: {water_mask_path}")
else:
print(f"[Step8] 未找到水域掩膜,跳过陆地剔除(陆地将保留在指数图中)")
work_dir = self.work_dir or str(Path(sampling_csv_path).parent)
# 开始后台处理
self.run_btn.setEnabled(False)
@ -593,18 +514,15 @@ class Step10WatercolorPanel(QWidget):
self.progress_label.setText("")
self._worker = WaterIndexWorker(
bsq_path=bsq_path,
hdr_path=hdr_path,
sampling_csv_path=sampling_csv_path,
output_dir=output_dir,
selected_formulas=selected,
waterindex_csv=self._waterindex_csv,
water_mask_path=water_mask_path,
work_dir=work_dir,
)
self._worker.progress.connect(self._on_progress)
self._worker.finished_ok.connect(self._on_finished)
self._worker.failed.connect(self._on_failed)
self._worker.log.connect(lambda m: self.progress_label.setText(m))
self._worker.error.connect(self._on_error)
self._worker.start()
def _on_progress(self, msg: str, pct: float):
@ -614,30 +532,38 @@ class Step10WatercolorPanel(QWidget):
def _on_finished(self, results: Dict[str, str]):
self.run_btn.setEnabled(True)
n = len(results)
names = list(results.keys())[:3]
tail = " …" if n > 3 else ""
QMessageBox.information(
self, "执行成功",
f"水色指数反演完成!\n"
f"共生成 {n} 个指数图(GeoTIFF)。\n\n"
f"共生成 {n} 个指数 CSV(含 longitude / latitude / 公式值三列)。\n"
f"前几个: {', '.join(names)}{tail}\n\n"
f"输出目录: {self.output_dir.get_path()}"
)
main_window = self.window()
if main_window and hasattr(main_window, 'log_message'):
main_window.log_message(f"步骤8:水色指数反演完成,生成 {n} 个指数图", "info")
main_window.log_message(
f"步骤10:水色指数反演完成,生成 {n} 个指数 CSV", "info"
)
def _on_failed(self, err: str):
def _on_error(self, err: str):
self.run_btn.setEnabled(True)
self.progress_bar.setValue(0)
QMessageBox.critical(self, "执行错误", f"水色指数反演失败:\n\n{err[:500]}")
self.progress_label.setText("执行失败")
QMessageBox.critical(
self, "执行错误", f"水色指数反演失败:\n\n{err[:500]}"
)
def get_output_dir(self) -> str:
return self.output_dir.get_path().strip() or ""
def get_output_tif_paths(self) -> List[str]:
"""获取输出目录下的所有 GeoTIFF 文件路径"""
def get_output_csv_paths(self) -> List[str]:
"""获取输出目录下的所有指数 CSV 文件路径(供 Step 11 ContentMapper 探测)"""
out_dir = self.get_output_dir()
if not out_dir or not os.path.isdir(out_dir):
return []
return sorted(
str(p) for p in Path(out_dir).glob("*.tif")
str(p) for p in Path(out_dir).glob("*.csv")
if p.is_file()
)

View File

@ -290,7 +290,10 @@ class Step11MapPanel(QWidget):
params_layout.addRow("输入坐标系:", self.input_crs)
self.output_crs = QLineEdit()
self.output_crs.setText("EPSG:4326")
# ★★★ 强制默认输出坐标系与输入一致,禁止从 GUI 误改为 EPSG:4326 ★★★
# 历史默认值 'EPSG:4326' 会让 ContentMapper 把栅格重投影到经纬度,
# 与基于 EPSG:32651 的水域掩膜叠加时发生仿射变换撕裂(栅格错位、坐标轴扭曲)。
self.output_crs.setText("EPSG:32651")
params_layout.addRow("输出坐标系:", self.output_crs)
self.show_points = QCheckBox("显示采样点")
@ -416,7 +419,9 @@ class Step11MapPanel(QWidget):
'boundary_shp_path': self.boundary_file.get_path(),
'resolution': self.resolution.value(),
'input_crs': self.input_crs.text(),
'output_crs': self.output_crs.text(),
# ★★★ 强制 output_crs = input_crs,不再信任 GUI 的 output_crs 输入框 ★★★
# 否则 ContentMapper 会把栅格重投影到 EPSG:4326,与水域掩膜叠加时撕裂。
'output_crs': self.input_crs.text(),
'show_sample_points': self.show_points.isChecked(),
'use_distance_diffusion': self.use_diffusion.isChecked(),
}
@ -437,7 +442,8 @@ class Step11MapPanel(QWidget):
'boundary_shp_path': self.boundary_file.get_path(),
'resolution': self.resolution.value(),
'input_crs': self.input_crs.text(),
'output_crs': self.output_crs.text(),
# ★★★ 强制 output_crs = input_crs,不再信任 GUI 的 output_crs 输入框 ★★★
'output_crs': self.input_crs.text(),
'show_sample_points': self.show_points.isChecked(),
'use_distance_diffusion': self.use_diffusion.isChecked(),
}
@ -475,8 +481,9 @@ class Step11MapPanel(QWidget):
self.resolution.setValue(config['resolution'])
if 'input_crs' in config:
self.input_crs.setText(config['input_crs'])
if 'output_crs' in config:
self.output_crs.setText(config['output_crs'])
# ★★★ 反灌入时强制 output_crs = input_crs,避免旧 config 中的 EPSG:4326 回填 ★★★
if 'output_crs' in config or 'input_crs' in config:
self.output_crs.setText(config.get('input_crs') or config.get('output_crs') or 'EPSG:32651')
if 'show_sample_points' in config:
self.show_points.setChecked(config['show_sample_points'])
if 'use_distance_diffusion' in config:
@ -573,10 +580,8 @@ class Step11MapPanel(QWidget):
return
boundary_shp_path = self.boundary_file.get_path()
if not boundary_shp_path:
QMessageBox.warning(self, "输入验证失败", "请选择边界文件")
return
if not os.path.exists(boundary_shp_path):
# ── Plan C: 允许 boundary_shp_path 为空(跳过水域掩膜,纯采样点插值)────
if boundary_shp_path and not os.path.exists(boundary_shp_path):
QMessageBox.warning(self, "输入验证失败", "边界文件不存在")
return
@ -631,7 +636,8 @@ class Step11MapPanel(QWidget):
output_dir=out_dir,
boundary_shp_path=boundary_shp_path,
input_crs=self.input_crs.text(),
output_crs=self.output_crs.text(),
# ★★★ 强制 output_crs = input_crs,不再信任 GUI 的 output_crs 输入框 ★★★
output_crs=self.input_crs.text(),
)
main_win = parent
@ -671,7 +677,8 @@ class Step11MapPanel(QWidget):
boundary_shp_path = self.boundary_file.get_path()
input_crs = self.input_crs.text()
output_crs = self.output_crs.text()
# ★★★ 强制 output_crs = input_crs,不再信任 GUI 的 output_crs 输入框 ★★★
output_crs = self.input_crs.text()
# 构造输出路径
out_dir = (self.output_dir.get_path() or "").strip()

View File

@ -440,6 +440,62 @@ class VisualizationWorkerThread(QThread):
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}")
@ -574,272 +630,104 @@ class ChartViewerDialog(QDialog):
class ImageCategoryTree(QTreeWidget):
"""图像分类目录树 - 按真实物理文件夹结构组织图像文件"""
# 文件名中文翻译映射(key: 文件名前缀 → 中文显示名)
NAME_MAPPING = {
"hsi_preview": "高光谱影像预览",
"hsi_original": "原始高光谱影像",
"hsi_deglint": "去耀斑高光谱影像",
"water_mask_overlay": "水域掩膜叠加图",
"water_mask": "水域掩膜图",
"glint_mask": "耀斑掩膜预览",
"glint_overlay": "耀斑叠加对比图",
"deglint_comparison": "去耀斑前后对比",
"training_spectra": "训练光谱特征",
"spectrum_by_param": "参数光谱图",
"model_evaluation": "模型评估散点图",
"model_scatter": "模型散点图",
"regression": "回归分析图",
"validation": "验证结果图",
"spatial_distribution": "参数空间分布图",
"distribution_map": "分布图",
"thematic_map": "水质专题图",
"water_quality_map": "水质分布图",
"prediction_map": "预测结果图",
"inversion_map": "反演结果图",
"correlation_matrix": "特征相关性矩阵",
"feature_correlation": "特征相关性",
"sampling_point_map": "采样点分布图",
"sampling_points": "采样点图",
"point_locations": "采样位置图",
"boxplot": "箱线图",
"histogram": "直方图",
"statistics": "统计图表",
"statistical_chart": "统计图",
"error_analysis": "误差分析图",
"rmse": "RMSE评估图",
"r2_score": "R²得分图",
"flight": "飞行轨迹图",
"path": "轨迹图",
"trajectory": "轨迹图",
"glint_deglint": "耀斑去耀斑影像",
"enhanced": "增强分布图",
"content": "含量分布图",
"distribution": "分布图",
"prediction": "预测图",
"inversion": "反演图",
"scatter_true_vs_pred": "真值-预测散点图",
"true_vs_pred": "真值-预测散点图",
"correlation_heatmap": "相关性热力图",
"parameter_boxplot": "水质参数箱线图",
"spectrum_comparison": "光谱曲线对比图",
"scatter": "散点图",
}
# 目录层级中文翻译
DIR_MAPPING = {
"14_visualization": "统计与分析报表",
"1_water_mask": "水域掩膜识别",
"2_Glint_Detection": "耀斑区域检测",
"3_deglint": "去耀斑影像结果",
"5_training_spectra": "训练光谱特征",
"8_Regression_Modeling": "回归建模分析",
"9_water_quality_prediction": "水质预测结果",
"10_feature_construction": "特征构建散点",
"11_12_13_predictions": "空间分布专题图",
"glint_deglint_previews": "耀斑处理预览",
"sampling_maps": "采样点空间分布",
"flight_maps": "无人机飞行轨迹",
"9_ML_Prediction": "机器学习预测",
"Non_Empirical_Prediction": "非经验模型预测",
"Custom_Regression_Prediction": "自定义回归预测",
"boxplot_dir": "水质参数箱线图",
"boxplot": "水质参数箱线图",
"output_dir": "输出目录",
"8_spatial_inversion": "空间反演",
"4_processed_data": "处理数据",
"9_Concentration": "物理反演浓度分布",
}
"""现代化的图像分类目录树 - 支持智能归类、筛选和高颜值样式"""
def __init__(self, parent=None):
super().__init__(parent)
self._dir_node_map: dict = {} # 目录路径字符串 → QTreeWidgetItem
self._work_path: Optional[Path] = None
self.setHeaderLabel("图像目录")
self.setMaximumWidth(300)
self.setMinimumWidth(250)
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 #ddd;
border-radius: 5px;
background-color: #f8f9fa;
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 {
padding: 5px;
border-radius: 3px;
}
QTreeWidget::item:selected {
background-color: #0078D4;
color: white;
height: 32px;
border-radius: 6px;
margin: 2px 4px;
}
QTreeWidget::item:hover {
background-color: #e3f2fd;
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 clear_all_images(self):
"""清除所有图像项"""
try:
self.invisibleRootItem().takeChildren()
if hasattr(self, '_dir_node_map'):
self._dir_node_map.clear()
except Exception as e:
print(f"清空树状图出错: {e}")
import traceback
traceback.print_exc()
def _parse_file_info(self, file_path: Path):
"""智能解析文件名,提取【水质参数】和【图表类型】"""
name_upper = file_path.name.upper()
def _translate_dir_name(self, dir_name: str) -> str:
"""翻译目录名为中文"""
return self.DIR_MAPPING.get(dir_name, dir_name)
# 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 "HISTOGRAM" in name_upper or "箱线" 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 = "掩膜与预览"
def _translate_filename(self, filename: str) -> str:
# 1. 后缀替换 (图表类型)
type_mapping = {
'_scatter_true_vs_pred': ' 真值预测散点图',
'_spectrum_comparison': ' 光谱曲线对比图',
'_spectrum': ' 光谱特征图',
'_histogram': ' 分布直方图',
'_boxplot_seaborn': ' Seaborn箱线图',
'_boxplot': ' 箱线图',
'_distribution_enhanced': ' 增强空间分布图',
'_distribution': ' 空间分布图',
'_sampling_map': ' 采样点地图',
'_flight_paths': ' 飞行轨迹图',
'_preview': ' 效果预览图',
'water_mask_overlay': '水域掩膜叠加图',
'hsi_preview': '原始影像预览',
'correlation_heatmap': '特征相关性热力图',
'parameter_boxplot': '水质参数汇总箱线图',
'all_parameters_boxplot': '全参数汇总箱线图',
'content_map': '含量分布专题图',
'_scatter_with_confidence': ' 置信区间散点图'
# 2. 提取参数名
param_name = "综合/未分类"
# 常见水质参数字典映射
params_map = {
'CHLOROPHYLL': 'Chlorophyll (叶绿素)', 'CHL_A': 'Chlorophyll (叶绿素)', 'CHLA': 'Chlorophyll (叶绿素)',
'COD': 'COD (化学需氧量)', 'DO': 'DO (溶解氧)', 'PH': 'pH',
'TEMPERATURE': 'Temperature (温度)', 'SPCOND': 'spCond (电导率)',
'TURBIDITY': 'Turbidity (浊度)', 'TDS': 'TDS (总溶解固体)',
'CL-': 'Cl- (氯离子)', 'NO3-N': 'NO3-N (硝态氮)', 'NH3-N': 'NH3-N (氨氮)',
'BGA': 'BGA (蓝绿藻)', 'TT': 'TT (透明度)'
}
for key, display_name in params_map.items():
if key in name_upper or key.replace('-', '') in name_upper:
param_name = display_name
break
name = filename
for eng, chn in type_mapping.items():
if eng in name:
name = name.replace(eng, chn)
# 2. 常见水质参数前缀替换
param_mapping = {
'Chlorophyll': '叶绿素', 'Chl_a': '叶绿素a', 'Chla': '叶绿素a',
'Turbidity': '浊度', 'Temperature': '温度', 'spCond': '电导率',
'COD': '化学需氧量', 'DO': '溶解氧', 'PH': 'pH值', 'TDS': '总溶解固体',
'BGA': '蓝绿藻', 'TT': '透明度', 'NH3-N': '氨氮', 'NO3-N': '硝酸盐氮',
'glint_severe_glint_area': '重度耀斑区域',
'severe_glint_area': '重度耀斑区域',
'deglint_goodman': 'Goodman算法去耀斑',
'deglint_Goodman': 'Goodman算法去耀斑',
'glint_': '耀斑检测_',
'deglint_': '耀斑去除_',
}
for eng, chn in param_mapping.items():
if name.startswith(eng + ' ') or name.startswith(eng + '_'):
name = name.replace(eng, chn, 1)
elif eng in name:
name = name.replace(eng, chn)
return name.strip('_')
def add_image_by_dir(self, file_path: Path, work_path: Path):
"""按真实物理目录层级挂载图片节点
Args:
file_path: 图片文件的完整路径
work_path: 工作目录根路径
"""
# 计算相对路径
try:
rel_path = file_path.relative_to(work_path)
except ValueError:
rel_path = Path(file_path.name)
# 分离父目录链和文件名
parts = rel_path.parts
if len(parts) <= 1:
parent_key = "__root__"
parent_display = "根目录"
else:
# 父目录路径(相对于work_path)
parent_key = str(Path(*parts[:-1]))
# 取最后一层目录名作为显示名
parent_display = self._translate_dir_name(parts[-2])
# 根目录节点特殊处理
root_display = self._translate_dir_name(parts[0]) if parts else "根目录"
# 获取或创建根目录节点
if root_display not in self._dir_node_map:
root_item = QTreeWidgetItem(self)
root_item.setText(0, f"📁 {root_display}")
root_item.setData(0, Qt.UserRole, {"type": "root_dir", "path": str(work_path / parts[0])})
root_item.setExpanded(True)
self._dir_node_map[root_display] = root_item
self._dir_node_map[f"__root__{root_display}"] = root_item
root_item = self._dir_node_map.get(f"__root__{root_display}")
if len(parts) > 1:
# 获取或创建子目录节点
if parent_key not in self._dir_node_map:
dir_item = QTreeWidgetItem(root_item)
dir_item.setText(0, f" 📂 {parent_display}")
dir_item.setData(0, Qt.UserRole, {"type": "sub_dir", "path": str(work_path / parent_key)})
dir_item.setExpanded(True)
self._dir_node_map[parent_key] = dir_item
parent_item = self._dir_node_map[parent_key]
else:
parent_item = root_item
# 创建图片节点(根据翻译后的名称分配图标)
display_name = self._translate_filename(file_path.stem) + file_path.suffix
icon = "🖼️" # 默认
if "散点" in display_name:
icon = "📊"
elif "光谱" in display_name or "曲线" in display_name:
icon = "📈"
elif "箱线" in display_name or "直方" in display_name:
icon = "📉"
elif "分布" in display_name or "地图" in display_name or "轨迹" in display_name:
icon = "🗺️"
image_item = QTreeWidgetItem(parent_item)
image_item.setText(0, f" {icon} {display_name}")
image_item.setData(0, Qt.UserRole, {"type": "image", "path": str(file_path), "display_name": display_name})
image_item.setToolTip(0, str(file_path))
return image_item
return param_name, chart_type
def scan_directory(self, work_dir: str):
"""扫描目录中的所有图像文件(深度递归扫描)—— 按真实物理目录结构挂载"""
"""全量扫描文件并缓存(不直接构建树,而是交给 rebuild_tree 渲染)"""
try:
if not work_dir:
print("可视化面板:工作目录为空,跳过扫描")
return
if not work_dir: return
self._work_path = Path(work_dir)
if not self._work_path.exists(): return
# 阻塞信号,防止在清空树状图时触发 selected 槽函数导致崩溃
# 因为当前类继承自 QTreeWidget,所以 self 本身就是树
self.blockSignals(True)
self.clear_all_images()
self.blockSignals(False)
if not self._work_path.exists():
return
except Exception as e:
import traceback
print(f"可视化面板初始化扫描出错: {e}")
traceback.print_exc()
# 确保信号锁被解开
self.blockSignals(False)
return
try:
image_extensions = ['*.png', '*.jpg', '*.jpeg', '*.tif', '*.tiff', '*.bmp']
# 拓宽扫描根目录列表(新增多个遗漏目录)
scan_roots: List[Path] = [
# 仅扫描用于视觉展示的常规图片格式,屏蔽科学栅格 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')),
@ -850,51 +738,98 @@ class ImageCategoryTree(QTreeWidget):
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)
if not scan_roots: scan_roots.append(self._work_path)
seen_norm = set()
self._all_image_files = []
seen_norm: set = set()
image_files: List[Path] = []
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
if key in seen_norm: continue
seen_norm.add(key)
image_files.append(p)
if p.name.startswith('.') or 'thumb' in p.name.lower(): continue
self._all_image_files.append(p)
for img_file in sorted(image_files):
if img_file.name.startswith('.') or 'thumb' in img_file.name.lower():
continue
self.add_image_by_dir(img_file, self._work_path)
# 更新目录节点计数
for key, item in self._dir_node_map.items():
if key.startswith("__root__"):
continue
if item.data(0, Qt.UserRole).get("type") == "sub_dir":
count = item.childCount()
name = item.text(0)
if count > 0 and f"({count})" not in name:
# 从目录名中提取显示名并附加计数
display = name.strip()
item.setText(0, f" 📂 {display} ({count})")
self._all_image_files.sort(key=lambda x: x.name)
# 默认构建模式
self.rebuild_tree(group_mode='parameter', filter_type='all')
except Exception as e:
import traceback
print(f"可视化面板图片挂载出错: {e}")
traceback.print_exc()
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
if not selected_item: return None
data = selected_item.data(0, Qt.UserRole)
if data and data.get("type") == "image":
return data.get("path")
@ -1401,17 +1336,19 @@ class Step12VizPanel(QWidget):
QMessageBox.critical(self, "错误", f"可视化任务失败:\n{err[:1200]}")
def init_ui(self):
"""初始化UI - 使用左右分栏布局"""
"""初始化UI - 使用全新的三列布局(控制参数 | 独立满高目录树 | 图像查看器)"""
main_layout = QHBoxLayout()
main_layout.setSpacing(10)
main_layout.setSpacing(12)
main_layout.setContentsMargins(10, 10, 10, 10)
# ===== 左侧面板 =====
left_panel = QWidget()
left_layout = QVBoxLayout()
left_layout.setContentsMargins(0, 0, 0, 0)
# ==========================================
# 第一列:控制面板(目录选择 + 生成配置)
# ==========================================
control_panel = QWidget()
control_layout = QVBoxLayout()
control_layout.setContentsMargins(0, 0, 0, 0)
# 工作目录选择
# 1. 工作目录选择
dir_group = QGroupBox("工作目录")
dir_layout = QHBoxLayout()
self.work_dir_edit = QLineEdit()
@ -1422,9 +1359,9 @@ class Step12VizPanel(QWidget):
dir_layout.addWidget(self.work_dir_edit, 1)
dir_layout.addWidget(dir_browse_btn)
dir_group.setLayout(dir_layout)
left_layout.addWidget(dir_group)
control_layout.addWidget(dir_group)
# 图像目录选择(优先指向预测结果目录)
# 2. 图像目录选择
img_dir_group = QGroupBox("图像目录")
img_dir_layout = QHBoxLayout()
self.img_dir_edit = QLineEdit()
@ -1435,18 +1372,9 @@ class Step12VizPanel(QWidget):
img_dir_layout.addWidget(self.img_dir_edit, 1)
img_dir_layout.addWidget(img_dir_browse_btn)
img_dir_group.setLayout(img_dir_layout)
left_layout.addWidget(img_dir_group)
control_layout.addWidget(img_dir_group)
# 图像目录树
tree_group = QGroupBox("图像目录")
tree_layout = QVBoxLayout()
self.image_tree = ImageCategoryTree()
self.image_tree.itemClicked.connect(self.on_tree_item_clicked)
tree_layout.addWidget(self.image_tree)
tree_group.setLayout(tree_layout)
left_layout.addWidget(tree_group, 1)
# 可视化配置
# 3. 可视化配置
config_group = QGroupBox("可视化配置")
config_layout = QVBoxLayout()
@ -1470,6 +1398,11 @@ class Step12VizPanel(QWidget):
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(10)
line = QFrame()
line.setFrameShape(QFrame.HLine)
@ -1479,23 +1412,73 @@ class Step12VizPanel(QWidget):
self.gen_all_btn = QPushButton("🚀 生成全部")
self.gen_all_btn.setToolTip("生成所有类型的可视化图表")
self.gen_all_btn.setStyleSheet("background-color: #4CAF50; color: white; font-weight: bold;")
self.gen_all_btn.setStyleSheet("background-color: #4CAF50; color: white; font-weight: bold; padding: 8px; border-radius: 4px;")
self.gen_all_btn.clicked.connect(self.generate_all_visualizations)
config_layout.addWidget(self.gen_all_btn)
self.scan_btn = QPushButton("📁 扫描目录")
self.scan_btn.setToolTip("扫描工作目录中的图像文件")
self.scan_btn.setStyleSheet("padding: 6px; border-radius: 4px;")
self.scan_btn.clicked.connect(self.scan_work_directory)
config_layout.addWidget(self.scan_btn)
config_group.setLayout(config_layout)
left_layout.addWidget(config_group)
control_layout.addWidget(config_group)
left_panel.setLayout(left_layout)
left_panel.setMaximumWidth(350)
main_layout.addWidget(left_panel, 0)
control_layout.addStretch() # 把控制面板的内容往上顶
control_panel.setLayout(control_layout)
control_panel.setMaximumWidth(280) # 稍微收窄第一列
control_panel.setMinimumWidth(230)
main_layout.addWidget(control_panel, 0) # stretch=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(8)
# 添加过滤控制栏
filter_layout = QFormLayout()
filter_layout.setContentsMargins(0, 0, 0, 0)
# 分组模式下拉框
self.view_mode_cb = QComboBox()
self.view_mode_cb.addItems(["按水质参数归类", "按图表类型归类", "按物理文件夹"])
self.view_mode_cb.currentIndexChanged.connect(self.update_image_tree_view)
self.view_mode_cb.setStyleSheet("QComboBox { padding: 4px; border-radius: 4px; border: 1px solid #ccc; }")
# 图表类型筛选下拉框
self.chart_filter_cb = QComboBox()
self.chart_filter_cb.addItems(["全部图表", "空间分布图", "模型散点图", "光谱曲线图", "统计箱线图", "相关性热力图", "掩膜与预览", "采样点地图"])
self.chart_filter_cb.currentIndexChanged.connect(self.update_image_tree_view)
self.chart_filter_cb.setStyleSheet("QComboBox { padding: 4px; border-radius: 4px; border: 1px solid #ccc; }")
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) # stretch=1 让树垂直填满
tree_group.setLayout(group_layout)
tree_layout.addWidget(tree_group, 1) # stretch=1 让GroupBox垂直填满
tree_panel.setLayout(tree_layout)
tree_panel.setMaximumWidth(320)
tree_panel.setMinimumWidth(260)
main_layout.addWidget(tree_panel, 0) # stretch=0,不抢占右侧图片的宽度
# ==========================================
# 第三列:图像查看器
# ==========================================
right_panel = QWidget()
right_layout = QVBoxLayout()
right_layout.setContentsMargins(0, 0, 0, 0)
@ -1503,10 +1486,33 @@ class Step12VizPanel(QWidget):
self.image_viewer.refresh_btn.clicked.connect(self.scan_work_directory)
right_layout.addWidget(self.image_viewer, 1)
right_panel.setLayout(right_layout)
main_layout.addWidget(right_panel, 1)
main_layout.addWidget(right_panel, 1) # stretch=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
@ -1678,7 +1684,7 @@ class Step12VizPanel(QWidget):
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()):
self.gen_sampling_map.isChecked() or self.gen_distribution_map.isChecked()):
QMessageBox.information(self, "提示", "请至少勾选一项可视化配置选项以生成图表。")
return
reply = QMessageBox.question(
@ -1694,9 +1700,20 @@ class Step12VizPanel(QWidget):
"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_mask') 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()
@ -1886,6 +1903,7 @@ class Step12VizPanel(QWidget):
'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
@ -1913,3 +1931,5 @@ class Step12VizPanel(QWidget):
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))

View File

@ -64,8 +64,9 @@ class ReportWorkerThread(QThread):
)
else:
ai_cfg = ReportGenerationConfig(
ai_provider="minimax",
ai_provider=provider,
minimax_api_key=s.value("api_key", "", type=str) or "",
minimax_base_url=s.value("api_base_url", "", type=str) or None, # <--- 新增这行,把界面上的 URL 传过去
minimax_vision_model=s.value("vision_model", "", type=str) or None,
minimax_text_model=s.value("text_model", "", type=str) or None,
minimax_timeout_s=timeout,

View File

@ -166,7 +166,7 @@ ROUTES = [
{
# data/icons/ 没有 12.png/13.png;Step11/12/13 暂时复用 9.png(11 个 png 对 13 个 step 必然有共用)
"id": "step11",
"name": "11. 专题图生成",
"name": "11. 分布图生成",
"icon": "9.png",
"view_module": "src.new.views.step11_view",
"view_class": "Step11View",

View File

@ -1,36 +1,42 @@
# -*- coding: utf-8 -*-
"""
Step10 后端计算服务(水色指数反演)
====================================
Step10 后端计算服务(水色指数反演 · 散点 CSV 模式)
====================================================
纯计算函数——绝对不引用 PyQt、绝对不引用 main_view、绝对不读写全局变量。它只:
1. 从 ``config`` 字典读取参数;
2. 调用 ``WaterIndexProcessor.run_inversion`` 用 ``waterindex.csv`` 中的
公式直接处理去耀斑 BSQ 影像,输出各水质参数指数的 GeoTIFF;
3. 返回结果字典 ``{status, output_path, message, mode}``。
2. 调用 ``WaterIndexCsvProcessor.compute_indices_from_csv``
读取 Step 4 输出的 ``sampling_spectra.csv`` 散点,对每行采样点
套用 ``waterindex.csv`` 中勾选的公式,输出每公式一个 CSV;
3. 返回结果字典 ``{status, output_path, message, mode, ...}``。
调用入口(由 main_view 在后台 QThread 中调用):
execute_step10({
"bsq_path": "D:/deglint_output.bsq", # 去耀斑 BSQ 影像(必填)
"deglint_img_path": "D:/deglint_output.bsq", # 同上(兼容旧 panel 字段)
"hdr_path": "D:/deglint_output.hdr", # ENVI 头文件(可省,自动 .bsq→.hdr 推断)
"selected_formulas": ["NDCI", "BGA_Am09KBBI"], # 要处理的公式名列表(空 → 全部)
"formula_csv_path": "D:/waterindex.csv", # waterindex.csv 路径(可省,自动探测)
"water_mask_path": "D:/water_mask.dat", # 水域掩膜路径(可省)
"nodata_value": -9999.0, # NoData 标记值
"output_dir": "D:/10_WaterIndex_Images", # 输出目录(可省 → work_dir/10_WaterIndex_Images)
"enabled": True,
"work_dir": "D:/workspace", # 工作目录(main_view 注入)
"sampling_csv_path": "D:/4_sampling/sampling_spectra.csv", # 必填
"selected_formulas": ["NDCI", "BGA_Am09KBBI"], # 勾选公式;空 → 全部
"formula_csv_path": "D:/waterindex.csv", # waterindex.csv 路径
"output_dir": "D:/10_WaterIndex_CSV", # 输出目录;可省
"enabled": True,
"work_dir": "D:/workspace", # 主窗口注入
})
返回字典字段:
* ``status`` : "completed" | "skipped" | "error"
* ``output_path`` : 输出目录路径(失败时为 None)
* ``output_files`` : {公式名: 公式 CSV 路径}(失败时为空 dict)
* ``message`` : 人类可读说明
* ``mode`` : "watercolor_inversion"(便于 UI 提示)
* ``mode`` : "watercolor_inversion_csv"(便于 UI 提示)
设计要点
========
- 与 Step 9 (ML 预测) 完全对称的"散点处理模式":输入 CSV、输出 CSV,
坐标列重命名为 longitude/latitude,公式值以列形式追加。
- 旧"读 BSQ 全图 → 输出 GeoTIFF"模式已废弃(科学上误差大且与 GIS 栅格计算器重复)。
- 兼容调用方可能仍传旧键(bsq_path / hdr_path / deglint_img_path),检测到时
静默忽略并回退到 sampling_spectra.csv 路径解析(避免破坏已有 pipeline 配置)。
"""
from __future__ import annotations
@ -38,11 +44,40 @@ from __future__ import annotations
from pathlib import Path
from typing import Any, Dict, List, Optional
from src.new.services._output_resolver import get_user_output_path, is_user_specified, resolve_output_dir
from src.new.services._output_resolver import get_user_output_path
def _resolve_sampling_csv_path(
sampling_csv_path: Optional[str],
work_dir: str,
) -> str:
"""解析采样点 CSV 路径
解析顺序:
1. 显式传入的 ``sampling_csv_path``
2. ``{work_dir}/4_sampling/sampling_spectra.csv``
3. ``{work_dir}/4_sampling/`` 下任意 ``.csv`` (取最新)
"""
if sampling_csv_path and Path(sampling_csv_path).is_file():
return sampling_csv_path
if not work_dir:
return sampling_csv_path or ""
primary = Path(work_dir) / "4_sampling" / "sampling_spectra.csv"
if primary.is_file():
return str(primary).replace("\\", "/")
sample_dir = Path(work_dir) / "4_sampling"
if sample_dir.is_dir():
cands = sorted(sample_dir.glob("*.csv"), key=lambda p: p.stat().st_mtime, reverse=True)
if cands:
return str(cands[0]).replace("\\", "/")
return sampling_csv_path or ""
def _resolve_waterindex_csv(formula_csv_path: Optional[str], work_dir: str) -> str:
"""解析 waterindex.csv 路径(与 WaterIndexProcessor.__init__ 默认逻辑保持一致)"""
"""解析 waterindex.csv 路径(与 WaterIndexCsvProcessor.__init__ 默认逻辑保持一致)"""
if formula_csv_path and Path(formula_csv_path).is_file():
return formula_csv_path
candidates = [
@ -52,142 +87,118 @@ def _resolve_waterindex_csv(formula_csv_path: Optional[str], work_dir: str) -> s
]
for c in candidates:
if c.is_file():
return str(c)
return str(c).replace("\\", "/")
return formula_csv_path or ""
def _resolve_water_mask_path(water_mask_path: Optional[str], work_dir: str) -> Optional[str]:
"""解析水域掩膜路径(缺省时尝试从 work_dir/1_water_mask 自动扫盘)"""
if water_mask_path and Path(water_mask_path).is_file():
return water_mask_path
if not work_dir:
return None
mask_dir = Path(work_dir) / "1_water_mask"
if not mask_dir.is_dir():
return None
for pat in ("*.tif", "*.TIF", "*.dat", "*.DT"):
cands = sorted(mask_dir.glob(pat))
if cands:
return str(cands[0])
return None
def _resolve_output_dir(config: Dict[str, Any], work_dir: str) -> tuple[Path, str]:
"""根据 output_dir / work_dir 计算水色指数反演结果输出目录
使用共享解析器强制执行"用户优先"规则——用户指定 output_dir 时直接用其值
(step10 的 output_dir 本身就是一个目录),否则用 work_dir/10_WaterIndex_Images 默认。
注意:step10 与其他步骤不同——output_dir 直接表示目录而非文件路径,
所以使用 Path(user_path) 而非 .parent。
(step10 的 output_dir 本身就是一个目录),否则用
``work_dir/10_WaterIndex_CSV`` 默认。
"""
user_path = get_user_output_path(config, "output_dir", "output_path")
if user_path:
return Path(user_path), "user"
return Path(work_dir) / "10_WaterIndex_Images", "default"
return Path(work_dir) / "10_WaterIndex_CSV", "default"
def execute_step10(config: Dict[str, Any]) -> Dict[str, Any]:
"""Step 10 后端计算入口——纯函数
"""Step 10 后端计算入口——纯函数(散点 CSV 模式)
Args:
config: 由前端 view.get_config() 序列化、再经 main_view 注入 work_dir 的字典
Returns:
标准结果字典 ``{status, output_path, message, mode}``
标准结果字典 ``{status, output_path, output_files, message, mode}``
"""
# ---------- 入参规整 ----------
bsq_path: str = config.get("bsq_path") or config.get("deglint_img_path") or ""
hdr_path: str = config.get("hdr_path") or ""
sampling_csv_path: str = (
config.get("sampling_csv_path")
or config.get("spectrum_csv_path") # 兼容旧字段
or ""
)
selected_formulas: List[str] = config.get("selected_formulas") or []
formula_csv_path: str = config.get("formula_csv_path") or ""
water_mask_path: Optional[str] = config.get("water_mask_path")
nodata_value: float = float(config.get("nodata_value", -9999.0))
output_dir: str = config.get("output_dir") or ""
enabled: bool = bool(config.get("enabled", True))
work_dir: str = config.get("work_dir") or "."
output_path, _source = _resolve_output_dir(config, work_dir)
mode = "watercolor_inversion"
mode = "watercolor_inversion_csv"
# ---------- 提前失败检查 ----------
if not enabled:
return {
"status": "skipped",
"output_path": None,
"output_files": {},
"message": "用户禁用此步骤(enabled=False)",
"mode": mode,
}
if not bsq_path:
return {
"status": "error",
"output_path": None,
"message": "未提供 BSQ 影像路径(bsq_path / deglint_img_path)",
"mode": mode,
}
if not Path(bsq_path).is_file():
return {
"status": "error",
"output_path": None,
"message": f"BSQ 影像不存在: {bsq_path}",
"mode": mode,
}
if not hdr_path:
# 自动探测 .hdr
hdr_path = str(Path(bsq_path).with_suffix(".hdr"))
if not Path(hdr_path).is_file():
hdr_alt = str(Path(bsq_path).with_suffix(".HDR"))
if Path(hdr_alt).is_file():
hdr_path = hdr_alt
else:
hdr_path = ""
if not hdr_path or not Path(hdr_path).is_file():
# 解析采样点 CSV 路径
resolved_sampling_csv = _resolve_sampling_csv_path(sampling_csv_path, work_dir)
if not resolved_sampling_csv:
return {
"status": "error",
"output_path": None,
"message": f"未找到 ENVI 头文件(与 BSQ 同名 .hdr): {bsq_path}",
"output_files": {},
"message": "未提供 sampling_csv_path 且默认位置均找不到 sampling_spectra.csv",
"mode": mode,
}
if not Path(resolved_sampling_csv).is_file():
return {
"status": "error",
"output_path": None,
"output_files": {},
"message": f"采样点 CSV 不存在: {resolved_sampling_csv}",
"mode": mode,
}
# ---------- 解析 waterindex.csv ----------
# 解析 waterindex.csv
resolved_formula_csv = _resolve_waterindex_csv(formula_csv_path, work_dir)
if not resolved_formula_csv:
return {
"status": "error",
"output_path": None,
"output_files": {},
"message": "未提供 formula_csv_path 且默认位置均找不到 waterindex.csv",
"mode": mode,
}
# ---------- 解析水域掩膜(可选) ----------
resolved_water_mask = _resolve_water_mask_path(water_mask_path, work_dir)
if not Path(resolved_formula_csv).is_file():
return {
"status": "error",
"output_path": None,
"output_files": {},
"message": f"waterindex.csv 不存在: {resolved_formula_csv}",
"mode": mode,
}
# ---------- 执行(包一层 try/except 把异常转 dict,避免炸线程) ----------
try:
from src.core.algorithms.waterindex_inversion import WaterIndexProcessor
from src.core.algorithms.waterindex_inversion import (
WaterIndexCsvProcessor,
)
print(f"[Step10 Service] 水色指数反演: bsq={bsq_path}")
print(f"[Step10 Service] hdr={hdr_path}")
print(f"[Step10 Service] 水色指数反演(散点模式): sampling_csv={resolved_sampling_csv}")
print(f"[Step10 Service] formula_csv={resolved_formula_csv}")
print(f"[Step10 Service] selected_formulas={selected_formulas or '全部'}")
if resolved_water_mask:
print(f"[Step10 Service] water_mask={resolved_water_mask}")
print(f"[Step10 Service] output_dir={output_path}")
processor = WaterIndexProcessor(resolved_formula_csv)
results = processor.run_inversion(
deglint_img_path=bsq_path,
work_dir=work_dir,
formula_csv_path=resolved_formula_csv,
processor = WaterIndexCsvProcessor(resolved_formula_csv)
out_files = processor.compute_indices_from_csv(
sampling_csv_path=resolved_sampling_csv,
output_dir=str(output_path).replace("\\", "/"),
selected_formulas=selected_formulas or None,
water_mask_path=resolved_water_mask,
nodata_value=nodata_value,
callback=None, # 日志由 main_view 统一接管
progress_callback=None, # 日志由 main_view 统一接管
)
except FileNotFoundError as e:
return {
"status": "error",
"output_path": None,
"output_files": {},
"message": f"文件不存在: {e}",
"mode": mode,
}
@ -195,6 +206,7 @@ def execute_step10(config: Dict[str, Any]) -> Dict[str, Any]:
return {
"status": "error",
"output_path": None,
"output_files": {},
"message": f"参数错误: {e}",
"mode": mode,
}
@ -202,16 +214,18 @@ def execute_step10(config: Dict[str, Any]) -> Dict[str, Any]:
return {
"status": "error",
"output_path": None,
"output_files": {},
"message": f"{type(e).__name__}: {e}",
"mode": mode,
}
# ---------- 成功路径 ----------
p = Path(output_path)
n_results = len(results) if isinstance(results, dict) else 0
n_results = len(out_files) if isinstance(out_files, dict) else 0
return {
"status": "completed",
"output_path": str(p).replace("\\", "/"),
"message": f"水色指数反演完成,共生成 {n_results} 个指数 GeoTIFF",
"output_files": out_files,
"message": f"水色指数反演完成,共生成 {n_results} 个指数 CSV",
"mode": mode,
}
}

View File

@ -25,7 +25,8 @@ Step11 后端计算服务(专题图生成 / 克里金插值)
"boundary_shp_path": "D:/boundary.shp", # 边界 shp(可选)
"resolution": 30.0, # 空间分辨率(米)
"input_crs": "EPSG:32651",
"output_crs": "EPSG:4326",
# ★★★ 强制默认 output_crs = input_crs,禁止从 service 配置误改为 EPSG:4326 ★★★
"output_crs": "EPSG:32651",
"output_dir": "D:/11_Thematic_Map", # 输出目录
"enabled": True,
"work_dir": "D:/workspace", # 工作目录

View File

@ -46,7 +46,7 @@ Step12 后端计算服务(数据可视化——散点/光谱/箱线/掩膜缩
from __future__ import annotations
from pathlib import Path
from typing import Any, Dict
from typing import Any, Dict, List, Optional
from src.new.services._output_resolver import get_user_output_path, is_user_specified, resolve_output_dir
@ -168,18 +168,115 @@ def _try_sampling_maps(work_dir: str, output_dir: Path) -> Dict[str, Any]:
return {"status": "completed" if p else "error", "path": p, "output_dir": str(output_dir / "sampling_maps")}
def _resolve_thematic_map_dir(work_dir: str) -> "Optional[Path]":
"""自动推断 Step 11 分布图输出目录
与 step11_service._resolve_output_dir 的 default 分支完全镜像——
用户未指定 output_dir 时默认为 ``work_dir/11_Thematic_Map``。
"""
if not work_dir:
return None
cand = Path(work_dir) / "11_Thematic_Map"
return cand if cand.is_dir() else None
def _try_distribution_maps(work_dir: str, output_dir: Path) -> Dict[str, Any]:
"""渲染/汇集 Step 11 生成的分布图到 14_visualization/distribution_maps
处理两类产物(独立子步骤,单个失败不影响另一个):
1. **PNG 直拷** —— ``11_Thematic_Map/*_专题图.png``
(Step 11 GeoTIFF 栅格模式的产物,已带外框/图例/采样点,直接复制即可)
2. **GeoTIFF 渲染** —— ``11_Thematic_Map/*_kriging.tif``
(Step 11 CSV 插值模式的纯 GeoTIFF 产物,调用 ``ContentMapper.visualize_raster``
渲染为同风格的 PNG)
Returns:
{"status": "completed"|"skipped", "count": int,
"png_copied": int, "tif_rendered": int,
"output_dir": str, "details": [{"src", "dst", "kind"}, ...]}
"""
import shutil
src_dir = _resolve_thematic_map_dir(work_dir)
if src_dir is None:
raise FileNotFoundError(
f"Step 11 输出目录不存在: {work_dir}/11_Thematic_Map"
f"(请先运行 Step 11 生成分布图)"
)
out_sub = output_dir / "distribution_maps"
out_sub.mkdir(parents=True, exist_ok=True)
details: List[Dict[str, str]] = []
n_png_copy = 0
n_tif_render = 0
# --- 1. PNG 直拷(GeoTIFF 模式已有产物) ---
for src_png in sorted(src_dir.glob("*_专题图.png")):
dst_png = out_sub / src_png.name
try:
shutil.copy2(src_png, dst_png)
details.append({"src": str(src_png), "dst": str(dst_png), "kind": "png_copy"})
n_png_copy += 1
except Exception as copy_err: # noqa: BLE001
print(f"[distribution_maps] ⚠ 复制失败 {src_png.name}: {copy_err}")
# --- 2. GeoTIFF 渲染(CSV 模式纯栅格产物) ---
tif_paths = sorted(src_dir.glob("*_kriging.tif"))
if tif_paths:
import matplotlib
matplotlib.use("Agg")
from src.postprocessing.map import ContentMapper
mapper = ContentMapper()
for tif_path in tif_paths:
stem = tif_path.stem
chinese_title = mapper._get_chinese_title(stem)
dst_png = out_sub / f"{chinese_title}_分布图.png"
try:
mapper.visualize_raster(
raster_tif_path=str(tif_path),
output_file=str(dst_png),
boundary_shp_path=None,
nodata_value=-9999.0,
figsize=(14, 10),
alpha=0.9,
)
details.append({"src": str(tif_path), "dst": str(dst_png), "kind": "tif_render"})
n_tif_render += 1
except Exception as render_err: # noqa: BLE001
print(f"[distribution_maps] ⚠ 渲染失败 {tif_path.name}: {render_err}")
total = n_png_copy + n_tif_render
if total == 0:
raise FileNotFoundError(
f"Step 11 输出目录 {src_dir} 中无 *_专题图.png 也无 *_kriging.tif"
)
return {
"status": "completed",
"count": total,
"png_copied": n_png_copy,
"tif_rendered": n_tif_render,
"output_dir": str(out_sub),
"details": details,
}
def execute_step12(config: Dict[str, Any]) -> Dict[str, Any]:
"""Step 12 后端计算入口——纯函数"""
work_dir: str = config.get("work_dir") or ""
img_dir: str = config.get("img_dir") or ""
enabled: bool = bool(config.get("enabled", True))
output_dir: str = config.get("output_dir") or ""
# 5 个开关:缺省默认 True(与旧 panel 行为一致)
# 6 个开关:缺省默认 True(与旧 panel 行为一致;新增 distribution_maps 用于闭环 Step 11)
gen_scatter = bool(config.get("generate_scatter", True))
gen_spectrum = bool(config.get("generate_spectrum", True))
gen_boxplots = bool(config.get("generate_boxplots", True))
gen_glint = bool(config.get("generate_glint_previews", True))
gen_sampling = bool(config.get("generate_sampling_maps", True))
gen_distribution = bool(config.get("generate_distribution_maps", True))
output_path, _source = _resolve_output_dir(config, work_dir)
mode = "viz_generate"
@ -228,12 +325,14 @@ def execute_step12(config: Dict[str, Any]) -> Dict[str, Any]:
tasks.append(("glint_previews", _try_glint_previews))
if gen_sampling:
tasks.append(("sampling_maps", _try_sampling_maps))
if gen_distribution:
tasks.append(("distribution_maps", _try_distribution_maps))
if not tasks:
return {
"status": "completed",
"output_path": str(output_path).replace("\\", "/"),
"message": "无可视化任务(5 个开关全部 False)",
"message": "无可视化任务(6 个开关全部 False)",
"mode": mode,
}

View File

@ -1,20 +1,18 @@
# -*- coding: utf-8 -*-
"""
Step10View —— Step 10(水色指数反演)的端到端模块化 view
Step10View —— Step 10(水色指数反演)的端到端模块化 view(散点 CSV 模式)
UI 从 ``src/gui/panels/step10_watercolor_panel.py`` 原样搬迁。
与 ``src/gui/panels/step10_watercolor_panel.py`` 同步重构。
view 层职责
===========
- 输入影像(BSQ + HDR)、公式选择 ListWidget、输出目录 + 格式 combo、
进度条 / 进度标签、运行按钮全部保留。
- 删除 ``WaterIndexWorker`` 线程(service 接管后台反演逻辑);
进度条 / 标签在 view 层保留 UI 占位,由 service 通过
``dispatch_execute`` 反馈到主窗口的日志区即可。
- ``_find_waterindex_csv`` / ``_load_formulas`` / ``_load_metadata``
不在 view 层执行;公式 ListWidget 留空,service 通过 set_config
把 selected_formulas 注入。
- 输入采样点 CSV、公式选择 ListWidget、输出目录、进度条 / 进度标签、运行按钮。
- 删除 BSQ + HDR FileSelectWidget(散点模式不读全图栅格)。
- 删除 ``format_combo``(每个公式一个 CSV,无格式选择)。
- 删除 ``WaterIndexWorker`` 线程(service 接管后台计算逻辑)。
- ``_load_formulas`` / 加载 waterindex.csv 的逻辑不在 view 层执行;
公式 ListWidget 留空,service 通过 set_config 把 selected_formulas 注入。
"""
import os
@ -37,50 +35,51 @@ def _resolve_subdir(work_dir: str, subdir_name: str) -> str:
class Step10View(BaseView):
"""Step 10: 水色指数反演(高光谱影像直接处理)"""
"""Step 10: 水色指数反演(散点 CSV 模式)
输入:Step 4 输出的 sampling_spectra.csv(散点 + 全波段光谱)
处理:逐行套用 waterindex.csv 公式
输出:每公式一个 CSV(列:longitude, latitude, 公式值)
"""
def init_ui(self):
layout = QVBoxLayout()
# ---- 标题 ----
title = QLabel("步骤10:水色指数反演(高光谱影像直接处理)")
title = QLabel("步骤10:水色指数反演(散点 CSV 模式)")
title.setFont(QFont("Arial", 12, QFont.Bold))
layout.addWidget(title)
# ---- 说明 ----
hint = QLabel(
"将 waterindex.csv 中的公式直接应用于去耀斑高光谱影像(BSQ),"
"输出各水质参数指数的 GeoTIFF 栅格图像。"
"指数图可直接用于水质专题图生成。"
"读取 Step 4 生成的 sampling_spectra.csv 散点光谱,"
"对每个采样点逐行套用 waterindex.csv 中勾选的公式,"
"输出每公式一个 CSV(列:longitude, latitude, 公式值)。"
"结果可被 Step 11 直接以 ContentMapper 模式消费。"
)
hint.setWordWrap(True)
hint.setStyleSheet(f"color: {ModernStylesheet.COLORS.get('text_secondary', '#666')};")
layout.addWidget(hint)
# ---- 输入影像选择 ----
input_group = QGroupBox("输入影像")
# ---- 输入采样点数据 ----
input_group = QGroupBox("输入采样点数据")
input_layout = QFormLayout()
self.bsq_file = FileSelectWidget(
"BSQ 影像:",
"BSQ Files (*.bsq);;DAT Files (*.dat);;All Files (*.*)",
self.sampling_csv_file = FileSelectWidget(
"采样点 CSV:",
"CSV Files (*.csv);;All Files (*.*)",
)
self.bsq_file.line_edit.setPlaceholderText("选择去耀斑处理后的 BSQ 影像")
input_layout.addRow("BSQ 影像:", self.bsq_file)
self.hdr_file = FileSelectWidget(
"ENVI 头文件:",
"HDR Files (*.hdr);;All Files (*.*)",
self.sampling_csv_file.line_edit.setPlaceholderText(
"选择 Step 4 输出的 sampling_spectra.csv"
)
self.hdr_file.line_edit.setPlaceholderText("自动关联同路径 .hdr 文件")
input_layout.addRow("HDR 文件:", self.hdr_file)
input_layout.addRow("采样点 CSV:", self.sampling_csv_file)
self.meta_label = QLabel("未加载影像")
self.meta_label = QLabel("未加载采样点数据")
self.meta_label.setStyleSheet(
"background: #f0f0f0; padding: 4px 8px; border-radius: 4px; "
"font-size: 12px; color: #333;"
)
input_layout.addRow("影像信息:", self.meta_label)
input_layout.addRow("数据信息:", self.meta_label)
input_group.setLayout(input_layout)
layout.addWidget(input_group)
@ -125,14 +124,11 @@ class Step10View(BaseView):
"输出目录:",
"Directories",
)
self.output_dir.line_edit.setPlaceholderText("留空 → 工作目录/10_WaterIndex_Images")
self.output_dir.line_edit.setPlaceholderText(
"留空 → 工作目录/10_WaterIndex_CSV"
)
output_layout.addRow("输出目录:", self.output_dir)
self.format_combo = QComboBox()
self.format_combo.addItems(["GTiff (GeoTIFF)", "ENVI", "PCI"])
self.format_combo.setCurrentIndex(0)
output_layout.addRow("输出格式:", self.format_combo)
output_group.setLayout(output_layout)
layout.addWidget(output_group)
@ -178,7 +174,7 @@ class Step10View(BaseView):
# BaseView 契约
# ------------------------------------------------------------------
def get_config(self) -> dict:
bsq_path = self.bsq_file.get_path()
sampling = self.sampling_csv_file.get_path().strip()
selected = []
for i in range(self.formula_list.count()):
item = self.formula_list.item(i)
@ -186,26 +182,19 @@ class Step10View(BaseView):
name = item.data(Qt.UserRole)
if name:
selected.append(name)
config = {
"bsq_path": bsq_path,
"deglint_img_path": bsq_path,
"output_format": self.format_combo.currentText().split()[0],
config: dict = {
"sampling_csv_path": sampling,
"selected_formulas": selected,
"enabled": self.enable_checkbox.isChecked(),
}
hdr_path = self.hdr_file.get_path()
if hdr_path:
config["hdr_path"] = hdr_path
output_dir = self.output_dir.get_path()
output_dir = self.output_dir.get_path().strip()
if output_dir:
config["output_dir"] = output_dir
return config
def set_config(self, config: dict):
if config.get("bsq_path"):
self.bsq_file.set_path(config["bsq_path"])
if config.get("hdr_path"):
self.hdr_file.set_path(config["hdr_path"])
if config.get("sampling_csv_path"):
self.sampling_csv_file.set_path(config["sampling_csv_path"])
if config.get("output_dir"):
self.output_dir.set_path(config["output_dir"])
if "selected_formulas" in config:
@ -223,25 +212,44 @@ class Step10View(BaseView):
super().update_work_directory(work_dir)
if not work_dir:
return
out_dir = _resolve_subdir(work_dir, "watercolor")
# 1) 自动填采样点 CSV(从 step4 拉取)
sampling_path = self._find_step4_sampling_csv(work_dir)
if sampling_path and not self.sampling_csv_file.get_path():
self.sampling_csv_file.set_path(sampling_path)
# 2) 自动填输出目录
out_dir = os.path.join(work_dir, "10_WaterIndex_CSV").replace("\\", "/")
os.makedirs(out_dir, exist_ok=True)
if not self.output_dir.get_path():
self.output_dir.set_path(out_dir)
# 自动填 BSQ(去耀斑输出)
deglint_dir = _resolve_subdir(work_dir, "deglint")
if os.path.isdir(deglint_dir):
def _find_step4_sampling_csv(self, work_dir: str) -> str:
"""从 step4 panel / pipeline.step_outputs / 4_sampling 目录自动找采样点 CSV"""
# 1) 优先:从主窗口懒加载面板读
mw = self.window()
factory = getattr(mw, "_panel_factory", None) if mw else None
if factory:
step4_panel = factory.get_panel("step4_sampling")
if step4_panel and hasattr(step4_panel, "output_file"):
p = step4_panel.output_file.get_path().strip()
if p and os.path.isfile(p):
return p
# 2) 兜底:扫 4_sampling/sampling_spectra.csv
candidate = os.path.join(work_dir, "4_sampling", "sampling_spectra.csv")
if os.path.isfile(candidate):
return candidate.replace("\\", "/")
# 3) 终极兜底:扫 4_sampling 下任意 .csv
sample_dir = os.path.join(work_dir, "4_sampling")
if os.path.isdir(sample_dir):
import glob
candidates = (
glob.glob(os.path.join(deglint_dir, "*.bsq"))
+ glob.glob(os.path.join(deglint_dir, "*.dat"))
)
if candidates and not self.bsq_file.get_path():
candidates.sort(key=os.path.getmtime, reverse=True)
bsq_path = candidates[0]
self.bsq_file.set_path(bsq_path)
hdr_path = os.path.splitext(bsq_path)[0] + ".hdr"
if os.path.exists(hdr_path):
self.hdr_file.set_path(hdr_path)
cands = glob.glob(os.path.join(sample_dir, "*.csv"))
if cands:
cands.sort(key=os.path.getmtime, reverse=True)
return cands[0].replace("\\", "/")
return ""
# ------------------------------------------------------------------
# 执行入口

View File

@ -181,7 +181,10 @@ class Step11View(BaseView):
params_layout.addRow("输入坐标系:", self.input_crs)
self.output_crs = QLineEdit()
self.output_crs.setText("EPSG:4326")
# ★★★ 强制默认输出坐标系与输入一致,禁止从 GUI 误改为 EPSG:4326 ★★★
# 历史默认值 'EPSG:4326' 会让 ContentMapper 把栅格重投影到经纬度,
# 与基于 EPSG:32651 的水域掩膜叠加时发生仿射变换撕裂(栅格错位、坐标轴扭曲)。
self.output_crs.setText("EPSG:32651")
params_layout.addRow("输出坐标系:", self.output_crs)
self.show_points = QCheckBox("显示采样点")
@ -302,7 +305,8 @@ class Step11View(BaseView):
"boundary_shp_path": self.boundary_file.get_path(),
"resolution": self.resolution.value(),
"input_crs": self.input_crs.text(),
"output_crs": self.output_crs.text(),
# ★★★ 强制 output_crs = input_crs,不再信任 GUI 的 output_crs 输入框 ★★★
"output_crs": self.input_crs.text(),
"show_sample_points": self.show_points.isChecked(),
"use_distance_diffusion": self.use_diffusion.isChecked(),
"enabled": self.enable_checkbox.isChecked(),
@ -342,8 +346,9 @@ class Step11View(BaseView):
self.resolution.setValue(config["resolution"])
if "input_crs" in config:
self.input_crs.setText(config["input_crs"])
if "output_crs" in config:
self.output_crs.setText(config["output_crs"])
# ★★★ 反灌入时强制 output_crs = input_crs,避免旧 config 中的 EPSG:4326 回填 ★★★
if "output_crs" in config or "input_crs" in config:
self.output_crs.setText(config.get("input_crs") or config.get("output_crs") or "EPSG:32651")
if "show_sample_points" in config:
self.show_points.setChecked(config["show_sample_points"])
if "use_distance_diffusion" in config:

View File

@ -100,6 +100,14 @@ class Step12View(BaseView):
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 生成的分布图(PNG 直拷 + GeoTIFF 渲染)"
"到 14_visualization/distribution_maps"
)
config_layout.addWidget(self.gen_distribution_map)
config_layout.addSpacing(10)
line = QFrame()
line.setFrameShape(QFrame.HLine)
@ -186,6 +194,7 @@ class Step12View(BaseView):
"generate_boxplots": self.gen_boxplots.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(),
"enabled": True,
}
@ -207,6 +216,8 @@ class Step12View(BaseView):
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))
def update_work_directory(self, work_dir: str):
super().update_work_directory(work_dir)

File diff suppressed because it is too large Load Diff

View File

@ -2,13 +2,6 @@
# -*- coding: utf-8 -*-
"""
采样点地图生成模块 - 在高光谱假彩色影像上标注采样点
支持功能:
1. 读取高光谱影像并生成假彩色RGB图像
2. 读取CSV文件中的采样点坐标(前两列为纬度、经度)
3. 在影像上标注红色采样点
4. 添加指北针、图例和比例尺
5. 支持地理坐标系转换
"""
import numpy as np
@ -21,13 +14,13 @@ from matplotlib.patches import FancyArrowPatch
import matplotlib.patheffects as path_effects
# 性能优化配置
plt.rcParams['agg.path.chunksize'] = 10000 # 提高矢量渲染性能
plt.rcParams['agg.path.chunksize'] = 10000
plt.rcParams['path.simplify'] = True
plt.rcParams['path.simplify_threshold'] = 0.1
# 导入GDAL用于影像读写
try:
from osgeo import gdal, osr
GDAL_AVAILABLE = True
except ImportError:
GDAL_AVAILABLE = False
@ -35,26 +28,15 @@ except ImportError:
class SamplingPointMap:
"""采样点地图生成类 - 在高光谱假彩色影像上标注采样点"""
def __init__(self, output_dir: str = "./point_maps", fast_mode: bool = False):
"""
初始化采样点地图生成器
Args:
output_dir: 输出目录,用于保存生成的地图
fast_mode: 是否启用快速模式(降低质量换取速度)
"""
self.output_dir = Path(output_dir)
self.output_dir.mkdir(parents=True, exist_ok=True)
self.fast_mode = fast_mode
# 设置中文字体
plt.rcParams['font.sans-serif'] = ['SimHei', 'Microsoft YaHei', 'DejaVu Sans', 'Arial Unicode MS']
plt.rcParams['axes.unicode_minus'] = False
plt.rcParams['font.size'] = 12
# 性能优化设置
if fast_mode:
plt.rcParams['figure.dpi'] = 150
plt.rcParams['savefig.dpi'] = 150
@ -62,86 +44,45 @@ class SamplingPointMap:
else:
plt.rcParams['figure.dpi'] = 300
plt.rcParams['savefig.dpi'] = 300
warnings.filterwarnings('ignore')
def create_sampling_point_map(self,
hyperspectral_path: str,
csv_path: str,
output_filename: Optional[str] = None,
rgb_bands: Optional[List[int]] = None,
point_color: str = 'red',
point_size: int = 80,
point_alpha: float = 0.8,
show_north_arrow: bool = True,
show_scale_bar: bool = True,
show_legend: bool = True,
dpi: int = None,
downsample: bool = False) -> str:
"""
创建采样点地图:在高光谱假彩色影像上标注采样点
Args:
hyperspectral_path: 高光谱影像文件路径 (.dat, .bsq, .tif等)
csv_path: 采样点CSV文件路径(前两列为纬度、经度)
output_filename: 输出文件名(如果为None则自动生成)
rgb_bands: 用于RGB合成的三个波段索引 [R, G, B],默认为None自动选择
point_color: 采样点颜色
point_size: 采样点大小
point_alpha: 采样点透明度
show_north_arrow: 是否显示指北针
show_scale_bar: 是否显示比例尺
show_legend: 是否显示图例
dpi: 输出图像分辨率(None时使用fast_mode设置)
downsample: 是否对图像进行下采样以加快速度(大影像推荐启用)
Returns:
生成的地图文件路径
"""
def create_sampling_point_map(self, hyperspectral_path: str, csv_path: str,
output_filename: Optional[str] = None, rgb_bands: Optional[List[int]] = None,
point_color: str = 'red', point_size: int = 80, point_alpha: float = 0.8,
show_north_arrow: bool = True, show_scale_bar: bool = True,
show_legend: bool = True, dpi: int = None, downsample: bool = False) -> str:
if not GDAL_AVAILABLE:
raise ImportError("GDAL未安装,无法处理地理坐标转换")
print(f"正在生成采样点地图...{' (快速模式)' if self.fast_mode else ''}")
# 读取高光谱影像 - 优化:仅读取需要的RGB波段
hyperspectral_img, geotransform, projection, width, height, sample_factor = self._read_hyperspectral(
hyperspectral_path, rgb_bands, downsample)
# 读取采样点
sampling_points = self._read_sampling_points(csv_path)
# 生成假彩色图像 - 应用线性拉伸
rgb_image = self._create_false_color_image(hyperspectral_img)
# 将地理坐标转换为像素坐标 - 支持投影系转换和下采样
pixel_coords = self._geo_to_pixel(sampling_points, geotransform, width, height, projection, sample_factor)
# 创建地图
if output_filename is None:
csv_name = Path(csv_path).stem
hs_name = Path(hyperspectral_path).stem
output_filename = f"{hs_name}_{csv_name}_sampling_map.png"
output_path = self.output_dir / output_filename
# 使用更优化的绘图设置
if dpi is None:
dpi = 150 if self.fast_mode else 200
self._create_map_visualization(
rgb_image, pixel_coords, sampling_points,
str(output_path), point_color, point_size, point_alpha,
show_north_arrow, show_scale_bar, show_legend, dpi,
geotransform, width, height, downsample, projection, sample_factor
rgb_image, pixel_coords, sampling_points, str(output_path), point_color, point_size, point_alpha,
show_north_arrow, show_scale_bar, show_legend, dpi, geotransform, width, height, downsample, projection,
sample_factor
)
print(f"采样点地图已保存: {output_path}")
return str(output_path)
def _read_hyperspectral(self, hyperspectral_path: str,
rgb_bands: Optional[List[int]] = None,
downsample: bool = False) -> Tuple[np.ndarray, tuple, str, int, int]:
"""优化版:读取高光谱影像 - 仅读取需要的RGB波段"""
def _read_hyperspectral(self, hyperspectral_path: str, rgb_bands: Optional[List[int]] = None,
downsample: bool = False) -> Tuple[np.ndarray, tuple, str, int, int]:
dataset = gdal.Open(hyperspectral_path)
if dataset is None:
raise ValueError(f"无法打开高光谱影像: {hyperspectral_path}")
@ -150,487 +91,222 @@ class SamplingPointMap:
height = dataset.RasterYSize
band_count = dataset.RasterCount
# 确定要读取的波段 - 优先使用指定波长 (650nm, 550nm, 460nm)
if rgb_bands is None:
if band_count >= 3:
try:
# 使用find_band_number根据波长查找RGB波段
from src.utils.util import find_band_number
rgb_bands = [
find_band_number(650.0, hyperspectral_path), # Red ~650nm
find_band_number(550.0, hyperspectral_path), # Green ~550nm
find_band_number(460.0, hyperspectral_path) # Blue ~460nm
find_band_number(650.0, hyperspectral_path),
find_band_number(550.0, hyperspectral_path),
find_band_number(460.0, hyperspectral_path)
]
print(f" 根据波长选择RGB波段: R={rgb_bands[0]}, G={rgb_bands[1]}, B={rgb_bands[2]}")
except Exception as e:
print(f" 波长查找失败 ({e}),使用默认索引")
# 回退到基于索引的选择
rgb_bands = [min(band_count-1, int(band_count*0.25)),
min(band_count-1, int(band_count*0.15)),
min(band_count-1, int(band_count*0.05))]
except Exception:
rgb_bands = [min(band_count - 1, int(band_count * 0.25)),
min(band_count - 1, int(band_count * 0.15)),
min(band_count - 1, int(band_count * 0.05))]
else:
rgb_bands = [0, 0, 0]
# 下采样控制 - 用户反馈下采样读取会导致像素值全为0
if downsample and (width > 2000 or height > 2000):
print(f" ⚠ 下采样暂被禁用(会导致像素值全0),使用原始分辨率: {width}x{height}")
print(f" ⚠ 下采样暂被禁用,使用原始分辨率: {width}x{height}")
sample_factor = 1
target_width = width
target_height = height
else:
sample_factor = 1
target_width = width
target_height = height
# 只读取需要的RGB波段(性能关键优化)
rgb_data = []
for band_idx in rgb_bands:
band = dataset.GetRasterBand(band_idx + 1)
# 直接使用完整分辨率读取,避免下采样导致像素值为0的问题
band_data = band.ReadAsArray().astype(np.float32)
rgb_data.append(band_data)
# 堆叠为RGB图像 (height, width, 3)
if len(rgb_data) == 3:
image_array = np.stack(rgb_data, axis=2)
else:
# 如果只有1个波段,复制为RGB
image_array = np.stack([rgb_data[0]]*3, axis=2)
image_array = np.stack([rgb_data[0]] * 3, axis=2)
geotransform = dataset.GetGeoTransform()
projection = dataset.GetProjection()
# 释放数据集
dataset = None
# 更新尺寸信息
final_width = target_width if sample_factor > 1 else width
final_height = target_height if sample_factor > 1 else height
print(f" 读取影像: {final_width}x{final_height}x{image_array.shape[2]} (RGB)")
if projection:
proj_type = "投影坐标系" if "PROJCS" in projection else "地理坐标系"
print(f" 影像投影: {proj_type}")
if sample_factor > 1:
print(f" 下采样因子: {sample_factor}")
return image_array, geotransform, projection, final_width, final_height, sample_factor
return image_array, geotransform, projection, width, height, sample_factor
def _read_sampling_points(self, csv_path: str) -> pd.DataFrame:
"""读取采样点CSV文件"""
if not Path(csv_path).exists():
raise FileNotFoundError(f"CSV文件不存在: {csv_path}")
"""智能读取采样点,自动识别模糊列名,允许UTM坐标,自动修复颠倒坐标"""
df = pd.read_csv(csv_path)
# 检查前两列是否为纬度和经度
if len(df.columns) < 2:
raise ValueError("CSV文件至少需要两列(纬度、经度)")
raise ValueError("CSV文件至少需要两列(经度、纬度 或 X、Y)")
# 假设前两列是纬度和经度
lat_col = df.columns[0]
lon_col = df.columns[1]
# 智能子串匹配
lat_aliases = ['lat', 'y', '纬']
lon_aliases = ['lon', 'lng', 'x', '经']
lat_col = None
lon_col = None
cols_lower = {c: str(c).strip().lower() for c in df.columns}
for c, lc in cols_lower.items():
if lat_col is None and any(a in lc for a in lat_aliases):
lat_col = c
elif lon_col is None and any(a in lc for a in lon_aliases):
lon_col = c
# 兜底:取前两列,默认列0=X(lon), 列1=Y(lat)
if lat_col is None or lon_col is None:
c0, c1 = df.columns[0], df.columns[1]
lon_col, lat_col = c0, c1
# 重命名列
df = df.rename(columns={lat_col: 'latitude', lon_col: 'longitude'})
# 确保数值类型
df['latitude'] = pd.to_numeric(df['latitude'], errors='coerce')
df['longitude'] = pd.to_numeric(df['longitude'], errors='coerce')
n_nan = int(df[['latitude', 'longitude']].isna().any(axis=1).sum())
df = df.dropna(subset=['latitude', 'longitude']).reset_index(drop=True)
# 删除无效的坐标
df = df.dropna(subset=['latitude', 'longitude'])
if len(df) > 0:
lat_max = df['latitude'].abs().max()
lon_max = df['longitude'].abs().max()
print(f"读取到 {len(df)} 个采样点")
# 智能对调:如果纬度 > 90,且经度 <= 90,说明用户把经纬度两列搞反了
if lat_max > 90 and lon_max <= 90 and lat_max <= 180:
print(" ⚠ 检测到经纬度数值颠倒 (纬度>90, 经度<=90),系统已自动对调坐标列")
df['latitude'], df['longitude'] = df['longitude'], df['latitude']
# UTM 投影坐标判定:只要数值远大于180,就是米级别的投影系统
elif lat_max > 180 or lon_max > 180:
print(f" ℹ 检测到坐标值远超180 (X:{lon_max:.1f}, Y:{lat_max:.1f}),判定为投影坐标(UTM)")
print(f" CSV 列匹配: lat_col='{lat_col}', lon_col='{lon_col}'")
if n_nan:
print(f" 剔除 {n_nan} 个无效(NaN)行")
print(f" 读取到 {len(df)} 个有效采样点 (不再拦截越界拦截)")
return df
def _create_false_color_image(self, image_array: np.ndarray,
rgb_bands: Optional[List[int]] = None) -> np.ndarray:
"""创建假彩色RGB图像 - 应用线性拉伸和Gamma校正"""
# 由于_read_hyperspectral已返回RGB图像,这里仅进行最终处理
def _create_false_color_image(self, image_array: np.ndarray, rgb_bands: Optional[List[int]] = None) -> np.ndarray:
if image_array.shape[2] != 3:
# 确保是3通道
if len(image_array.shape) == 2 or image_array.shape[2] == 1:
if len(image_array.shape) == 2:
image_array = np.stack([image_array]*3, axis=2)
else:
image_array = np.repeat(image_array, 3, axis=2)
image_array = np.stack([image_array] * 3, axis=2) if len(image_array.shape) == 2 else np.repeat(
image_array, 3, axis=2)
print(f" 处理前图像范围: R[{image_array[:,:,0].min():.3f}-{image_array[:,:,0].max():.3f}], "
f"G[{image_array[:,:,1].min():.3f}-{image_array[:,:,1].max():.3f}], "
f"B[{image_array[:,:,2].min():.3f}-{image_array[:,:,2].max():.3f}]")
# 增强型线性拉伸 - 解决图像太暗的问题
def simple_linear_stretch(data, min_percent=1, max_percent=99):
"""增强对比度的线性拉伸"""
valid_data = data[np.isfinite(data)]
if len(valid_data) == 0:
return np.zeros_like(data, dtype=np.float32)
# 计算百分位数,使用更激进的拉伸 (1%-99%)
if len(valid_data) == 0: return np.zeros_like(data, dtype=np.float32)
p_low = np.percentile(valid_data, min_percent)
p_high = np.percentile(valid_data, max_percent)
if p_high - p_low < 1e-8:
# 如果数据范围太小,使用最小最大值归一化
data_min = valid_data.min()
data_max = valid_data.max()
if data_max > data_min:
stretched = (data - data_min) / (data_max - data_min)
else:
stretched = np.zeros_like(data, dtype=np.float32)
else:
stretched = (data - p_low) / (p_high - p_low)
d_min, d_max = valid_data.min(), valid_data.max()
return (data - d_min) / (d_max - d_min) if d_max > d_min else np.zeros_like(data, dtype=np.float32)
stretched = (data - p_low) / (p_high - p_low)
return np.clip(stretched, 0.0, 1.0)
# 允许轻微过饱和以增加对比度
stretched = np.clip(stretched, 0.0, 1.05)
stretched = np.clip(stretched, 0.0, 1.0) # 最终确保在[0,1]
return stretched
# 对每个通道进行拉伸
r_stretched = simple_linear_stretch(image_array[:, :, 0])
g_stretched = simple_linear_stretch(image_array[:, :, 1])
b_stretched = simple_linear_stretch(image_array[:, :, 2])
# 合成为RGB图像
rgb_image = np.stack([r_stretched, g_stretched, b_stretched], axis=2)
rgb_image = np.nan_to_num(rgb_image, nan=0.0)
# 最终确保范围在[0,1],并轻微增强对比度
rgb_image = np.nan_to_num(np.stack([r_stretched, g_stretched, b_stretched], axis=2), nan=0.0)
rgb_image = np.clip(rgb_image, 0.0, 1.0)
return (rgb_image * 255).astype(np.uint8)
# 可选:Gamma校正增加亮度(解决太暗问题)
gamma = 1 # <1会增加亮度
rgb_image = np.power(rgb_image, gamma)
# 映射到0-255范围(uint8),这样imshow显示效果更好
rgb_image = (rgb_image * 255).astype(np.uint8)
print(f" 处理后图像范围: [0-255] (Gamma={gamma})")
return rgb_image
def _geo_to_pixel(self, sampling_points: pd.DataFrame,
geotransform: tuple, width: int, height: int,
projection: str = "", sample_factor: int = 1) -> List[Tuple[float, float]]:
"""
使用GDAL进行地理坐标到像素坐标的投影变换 - 支持下采样
原始点位坐标格式: 41.66054612 124.2208338 (WGS84地理坐标: 纬度,经度)
高光谱影像通常使用UTM或其他投影坐标系
当图像下采样时,sample_factor > 1,需要相应缩放坐标
"""
def _geo_to_pixel(self, sampling_points: pd.DataFrame, geotransform: tuple, width: int, height: int,
projection: str = "", sample_factor: int = 1) -> List[Tuple[float, float]]:
if geotransform is None or len(sampling_points) == 0:
# 如果没有地理变换信息,使用图像中心
return [(width/2, height/2) for _ in range(len(sampling_points))]
return [(width / 2, height / 2) for _ in range(len(sampling_points))]
pixel_coords = []
gt = geotransform
needs_transform = projection and ("PROJCS" in projection or "GEOGCS" in projection)
# 检查是否需要投影转换
needs_transform = False
if projection and ("PROJCS" in projection or "GEOGCS" in projection):
needs_transform = True
print(f" 检测到影像投影: {projection[:80]}...")
# 智能判定是否为 WGS84
sample_lon = float(sampling_points['longitude'].iloc[0])
sample_lat = float(sampling_points['latitude'].iloc[0])
is_wgs84 = (abs(sample_lon) <= 180) and (abs(sample_lat) <= 90)
# 创建坐标转换对象(WGS84 -> 影像投影)
transform = None
if needs_transform and GDAL_AVAILABLE:
if needs_transform and is_wgs84 and GDAL_AVAILABLE:
try:
# 源坐标系: WGS84 (EPSG:4326)
src_srs = osr.SpatialReference()
src_srs.ImportFromEPSG(4326) # WGS84
# 目标坐标系: 影像的投影
src_srs.ImportFromEPSG(4326)
dst_srs = osr.SpatialReference()
dst_srs.ImportFromWkt(projection)
# 创建坐标转换
transform = osr.CoordinateTransformation(src_srs, dst_srs)
print(" ✓ 已创建WGS84到影像投影的坐标转换")
except Exception as e:
print(f" ⚠ 坐标转换创建失败: {e},使用简化变换")
transform = None
elif not is_wgs84:
print(" ℹ 采样点为投影坐标(UTM),跳过WGS84投影转换,直接使用放射变换映射")
for _, row in sampling_points.iterrows():
lon = float(row['longitude']) # 经度 (WGS84)
lat = float(row['latitude']) # 纬度 (WGS84)
lon, lat = float(row['longitude']), float(row['latitude'])
if transform is not None:
# 使用GDAL进行投影转换: (经度, 纬度) -> (投影X, 投影Y)
try:
proj_x, proj_y, _ = transform.TransformPoint(lat, lon)
# 再转换为像素坐标
x = (proj_x - gt[0]) / gt[1]
y = (proj_y - gt[3]) / gt[5]
except Exception as e:
# 转换失败时回退到直接计算
x = (lon - gt[0]) / gt[1]
y = (lat - gt[3]) / gt[5]
proj_x, proj_y, _ = transform.TransformPoint(lon, lat)
x, y = (proj_x - gt[0]) / gt[1], (proj_y - gt[3]) / gt[5]
except Exception:
x, y = width / 2, height / 2
else:
# 直接使用仿射变换(坐标系一致的情况)
x = (lon - gt[0]) / gt[1]
y = (lat - gt[3]) / gt[5]
x, y = (lon - gt[0]) / gt[1], (lat - gt[3]) / gt[5]
# 如果图像进行了下采样,需要相应缩放坐标
if sample_factor > 1:
x = x / sample_factor
y = y / sample_factor
x, y = x / sample_factor, y / sample_factor
# 限制在图像范围内(使用下采样后的尺寸)
x = max(0, min(x, width - 1))
y = max(0, min(y, height - 1))
pixel_coords.append((x, y))
if transform is not None:
print(f" ✓ 使用GDAL投影变换处理 {len(pixel_coords)} 个采样点")
else:
print(f" 使用直接仿射变换处理 {len(pixel_coords)} 个采样点")
pixel_coords.append((max(0, min(x, width - 1)), max(0, min(y, height - 1))))
return pixel_coords
def _create_map_visualization(self, rgb_image: np.ndarray,
pixel_coords: List[Tuple[float, float]],
sampling_points: pd.DataFrame,
output_path: str,
point_color: str,
point_size: int,
point_alpha: float,
show_north_arrow: bool,
show_scale_bar: bool,
show_legend: bool,
dpi: int,
geotransform: tuple,
width: int,
height: int,
downsample: bool = False,
projection: str = "",
sample_factor: int = 1):
"""创建地图可视化 - 优化版"""
# 使用更小的figure尺寸加快渲染
def _create_map_visualization(self, rgb_image: np.ndarray, pixel_coords: List[Tuple[float, float]],
sampling_points: pd.DataFrame, output_path: str, point_color: str, point_size: int,
point_alpha: float, show_north_arrow: bool, show_scale_bar: bool, show_legend: bool,
dpi: int, geotransform: tuple, width: int, height: int, downsample: bool = False,
projection: str = "", sample_factor: int = 1):
figsize = (10, 8) if self.fast_mode or downsample else (12, 10)
fig, ax = plt.subplots(figsize=figsize, dpi=100 if self.fast_mode else 150)
# 显示假彩色图像 - 现在已经是0-255的uint8格式
print(f" 最终图像数据范围: [{rgb_image.min()}, {rgb_image.max()}] (uint8)")
ax.imshow(rgb_image, interpolation='nearest' if self.fast_mode else 'bilinear')
# 绘制采样点 - 优化:使用scatter代替循环plot
if pixel_coords:
x_coords = [p[0] for p in pixel_coords]
y_coords = [p[1] for p in pixel_coords]
ax.scatter(x_coords, y_coords, c=point_color, s=point_size,
alpha=point_alpha, edgecolors='white', linewidth=1.5)
x_coords, y_coords = [p[0] for p in pixel_coords], [p[1] for p in pixel_coords]
ax.scatter(x_coords, y_coords, c=point_color, s=point_size, alpha=point_alpha, edgecolors='white',
linewidth=1.5)
# 添加指北针
if show_north_arrow:
self._add_north_arrow(ax, width, height, position='bottom-left', direction='down')
if show_north_arrow: self._add_north_arrow(ax, width, height, position='bottom-left', direction='down')
if show_scale_bar and geotransform is not None: self._add_scale_bar(ax, geotransform, width, height)
# 添加比例尺
if show_scale_bar and geotransform is not None:
self._add_scale_bar(ax, geotransform, width, height)
# 添加图例
if show_legend:
legend_text = f'采样点 (n={len(sampling_points)})'
ax.plot([], [], 'o', color=point_color, markersize=8, label=legend_text)
ax.plot([], [], 'o', color=point_color, markersize=8, label=f'采样点 (n={len(sampling_points)})')
ax.legend(loc='lower right', frameon=True, facecolor='white', edgecolor='gray')
# 设置标题和标签
ax.set_title('高光谱影像采样点分布图', fontsize=16, fontweight='bold', pad=20)
# 隐藏坐标轴刻度
ax.set_xticks([])
ax.set_yticks([])
# 添加网格
ax.grid(True, alpha=0.2, linestyle='--')
plt.tight_layout()
# 保存参数 - 避免传递不兼容的参数
save_kwargs = {
'dpi': dpi,
'bbox_inches': 'tight',
'pad_inches': 0.05,
'facecolor': 'white'
}
# 仅添加matplotlib支持的参数
if self.fast_mode:
save_kwargs['dpi'] = min(dpi, 180) # 快速模式降低DPI
save_kwargs = {'dpi': min(dpi, 180) if self.fast_mode else dpi, 'bbox_inches': 'tight', 'pad_inches': 0.05,
'facecolor': 'white'}
plt.savefig(output_path, **save_kwargs)
plt.close(fig)
def _add_north_arrow(self, ax, width: int, height: int, position='top-right', direction='down',
size=0.08, color='white', n_color='white', outline_color='black'):
"""
添加指北针,可配置位置、方向、大小、颜色。
def _add_north_arrow(self, ax, width: int, height: int, position='top-right', direction='down', size=0.08,
color='white', n_color='white', outline_color='black'):
pos_map = {'top-left': (0.08, 0.88), 'top-right': (0.92, 0.88), 'bottom-left': (0.08, 0.12),
'bottom-right': (0.92, 0.12)}
arrow_x, arrow_y = width * pos_map.get(position, (0.92, 0.88))[0], height * pos_map.get(position, (0.92, 0.88))[
1]
dx, dy = {'up': (0, size), 'down': (0, -size), 'left': (-size, 0), 'right': (size, 0)}.get(direction,
(0, -size))
参数:
ax: matplotlib Axes对象
width, height: 图像宽高(用于相对定位)
position: 'top-left', 'top-right', 'bottom-left', 'bottom-right'
direction: 'up', 'down', 'left', 'right' 箭头指向
size: 箭头长度相对于高度的比例(0.05~0.12)
color: 箭头颜色
n_color: 'N' 文字颜色
outline_color: 文字描边颜色
"""
# 位置映射(偏移系数)
pos_map = {
'top-left': (0.08, 0.88),
'top-right': (0.92, 0.88),
'bottom-left': (0.08, 0.12),
'bottom-right': (0.92, 0.12),
}
arrow_x_ratio, arrow_y_ratio = pos_map.get(position, (0.92, 0.88))
arrow_x = width * arrow_x_ratio
arrow_y = height * arrow_y_ratio
# 方向映射(箭头终点偏移)
direction_map = {
'up': (0, +size),
'down': (0, -size),
'left': (-size, 0),
'right': (+size, 0),
}
dx, dy = direction_map.get(direction, (0, -size))
end_x = arrow_x + dx * width # 注意:dx是比例,乘以宽度/高度保持比例一致
end_y = arrow_y + dy * height
# 箭头绘制
arrow = FancyArrowPatch((arrow_x, arrow_y), (end_x, end_y),
color=color, linewidth=3,
arrowstyle='->', mutation_scale=20)
arrow = FancyArrowPatch((arrow_x, arrow_y), (arrow_x + dx * width, arrow_y + dy * height), color=color,
linewidth=3, arrowstyle='->', mutation_scale=20)
ax.add_patch(arrow)
# N 文字位置:在箭头尾部或头部?通常放在箭头指向的反方向末端
# 这里放在箭头尾部向外偏移一点(便于阅读)
# 偏移系数根据方向决定
offset_scale = 0.02 # 偏移量比例
if direction == 'up':
text_x = arrow_x
text_y = arrow_y - height * offset_scale # 放在箭头下方
elif direction == 'down':
text_x = arrow_x
text_y = arrow_y + height * offset_scale # 放在箭头上方
elif direction == 'left':
text_x = arrow_x + width * offset_scale
text_y = arrow_y
else: # right
text_x = arrow_x - width * offset_scale
text_y = arrow_y
ax.text(text_x, text_y, 'N', fontsize=14, fontweight='bold',
color=n_color, ha='center', va='center',
text_y = arrow_y - height * 0.02 if direction == 'up' else arrow_y + height * 0.02
ax.text(arrow_x, text_y, 'N', fontsize=14, fontweight='bold', color=n_color, ha='center', va='center',
path_effects=[path_effects.withStroke(linewidth=3, foreground=outline_color)])
def _add_scale_bar(self, ax, geotransform: tuple, width: int, height: int):
"""添加比例尺"""
if geotransform is None:
return
# 计算图像实际宽度(米)
if geotransform is None: return
pixel_size_x = abs(geotransform[1])
image_width_meters = width * pixel_size_x
# 选择合适的比例尺长度(图像宽度的1/4)
scale_length_m = image_width_meters / 4
scale_length_pixels = width / 4
# 找到合适的刻度
scale_options = [1000, 500, 200, 100, 50, 20, 10, 5, 2, 1]
scale_meters = next((s for s in scale_options if s <= scale_length_m), 1)
scale_length_m = (width * pixel_size_x) / 4
scale_meters = next((s for s in [1000, 500, 200, 100, 50, 20, 10, 5, 2, 1] if s <= scale_length_m), 1)
scale_pixels = int(scale_meters / pixel_size_x)
bar_x, bar_y = width * 0.08, height * 0.92
# 在左下角添加比例尺
bar_x = width * 0.08
bar_y = height * 0.92
# 绘制比例尺线
ax.plot([bar_x, bar_x + scale_pixels], [bar_y, bar_y], color='white', linewidth=4)
# 添加刻度线
ax.plot([bar_x, bar_x], [bar_y, bar_y + 8], color='white', linewidth=2)
ax.plot([bar_x + scale_pixels, bar_x + scale_pixels], [bar_y, bar_y + 8], color='white', linewidth=2)
# 添加文字
ax.text(bar_x + scale_pixels/2, bar_y , f'{scale_meters} m',
fontsize=11, ha='center', va='bottom', fontweight='bold',
bbox=dict(facecolor='white', alpha=0.8, edgecolor='none', pad=1))
def batch_create_maps(self, hyperspectral_path: str,
csv_folder: str,
output_subdir: str = "sampling_maps",
fast_mode: bool = True) -> Dict[str, str]:
"""
批量创建采样点地图
Args:
hyperspectral_path: 高光谱影像路径
csv_folder: 包含多个CSV文件的文件夹
output_subdir: 输出子目录
Returns:
生成的地图文件路径字典
"""
csv_folder_path = Path(csv_folder)
if not csv_folder_path.exists():
raise FileNotFoundError(f"CSV文件夹不存在: {csv_folder}")
# 创建输出目录
output_dir = self.output_dir / output_subdir
output_dir.mkdir(parents=True, exist_ok=True)
map_paths = {}
# 查找所有CSV文件
csv_files = list(csv_folder_path.glob("*.csv"))
print(f"找到 {len(csv_files)} 个CSV文件,开始批量生成采样点地图... (快速模式: {fast_mode})")
for csv_file in csv_files:
try:
output_filename = f"{Path(hyperspectral_path).stem}_{csv_file.stem}_sampling_map.png"
map_path = self.create_sampling_point_map(
hyperspectral_path=hyperspectral_path,
csv_path=str(csv_file),
output_filename=output_filename,
downsample=True, # 批量模式默认下采样
dpi=120 if fast_mode else 200
)
map_paths[csv_file.name] = map_path
print(f"✓ 生成: {csv_file.name}")
except Exception as e:
print(f"✗ 处理 {csv_file.name} 失败: {e}")
print(f"批量生成完成,共生成 {len(map_paths)} 个采样点地图")
return map_paths
# 测试代码
if __name__ == "__main__":
# 示例用法
map_generator = SamplingPointMap(output_dir="./point_maps")
# 测试代码已禁用,避免直接运行时出错
map_generator_fast = SamplingPointMap(output_dir="./point_maps", fast_mode=True)
map_path = map_generator_fast.create_sampling_point_map(
hyperspectral_path=r"D:\BaiduNetdiskDownload\yaobao\result3.bsq",
csv_path=r"E:\code\WQ\pipeline_result\work_dir\4_processed_data\processed_data.csv",
downsample=True,
dpi=150
)
print("测试代码已注释,请通过GUI或手动调用使用。")
print("SamplingPointMap类已创建,可以用于生成带采样点的地图。")
print("性能优化功能:")
print(" - fast_mode=True: 快速模式 (推荐用于预览)")
print(" - downsample=True: 对大影像下采样 (推荐用于>2000x2000影像)")
print(" - 使用: SamplingPointMap(fast_mode=True).create_sampling_point_map(...)")
ax.text(bar_x + scale_pixels / 2, bar_y, f'{scale_meters} m', fontsize=11, ha='center', va='bottom',
fontweight='bold', bbox=dict(facecolor='white', alpha=0.8, edgecolor='none', pad=1))

File diff suppressed because it is too large Load Diff

View File

@ -385,7 +385,7 @@ class WaterQualityVisualization:
return output_paths
def plot_distribution_map_enhanced(self, prediction_csv_path: str,
boundary_shp_path: str,
boundary_shp_path: Optional[str] = None, # ★★★ Plan C: None = 不依赖水域掩膜 ★★★
parameter_column: str = 'prediction',
output_path: Optional[str] = None,
resolution: float = 30,
@ -394,12 +394,12 @@ class WaterQualityVisualization:
colormap: str = 'viridis') -> str:
"""
生成增强的含量分布图(彩色填充图)
这是对step9的增强版本,使用更丰富的颜色映射
★★★ Plan C:boundary_shp_path 可选(None = 不依赖水域掩膜)★★★
Args:
prediction_csv_path: 预测结果CSV文件路径
boundary_shp_path: 边界shapefile文件路径
boundary_shp_path: 边界/掩膜文件路径。None 时跳过水域掩膜约束。
parameter_column: 参数值列名
output_path: 输出图片路径
resolution: 插值网格分辨率

View File

@ -99,8 +99,10 @@ class BandMathCalculator:
# 【新增安全防护】引入 numpy 命名空间,让 eval 引擎安全识别 nan 与 inf
import numpy as np
try:
# 即使 calc_expression 含有纯字符 nan,也能被 np.nan 安全接管
result = eval(calc_expression, {"__builtins__": None}, {"nan": np.nan, "inf": np.inf, "np": np})
# 【P0 修复】包 np.errstate 抑制除零 / 无效操作产生的 RuntimeWarning 洪水
with np.errstate(divide='ignore', invalid='ignore'):
# 即使 calc_expression 含有纯字符 nan,也能被 np.nan 安全接管
result = eval(calc_expression, {"__builtins__": None}, {"nan": np.nan, "inf": np.inf, "np": np})
except Exception as e:
print(f"⚠️ 警告:公式计算异常 ({e}),该点赋值为 nan")
result = np.nan

View File

@ -117,7 +117,9 @@ class WaterQualityIndexCalculator:
calc_expr,
)
try:
r = eval(calc_expr, {"__builtins__": None}, {"nan": np.nan, "inf": np.inf, "np": np})
# 【P0 修复】包 np.errstate 抑制除零 / 无效操作产生的 RuntimeWarning 洪水
with np.errstate(divide='ignore', invalid='ignore'):
r = eval(calc_expr, {"__builtins__": None}, {"nan": np.nan, "inf": np.inf, "np": np})
except Exception:
r = np.nan
results.append(r)