Files
WQ_GUI/src/postprocessing/map.py
duxin bdac7873f4 fix: 修复 ContentMapper 类体被模块级函数截断的严重 bug
问题: 批量生成专题图时报 'ContentMapper' object has no attribute 'process_data'

根因: _local_krige_block_worker (0空格缩进,模块级) 被错误地插入在
  ContentMapper 类体中间 (line 819), 导致类定义在此处终止。
  之后 19 个方法 (_idw_interpolation, read_csv_data, create_content_map,
  visualize_raster, prepare_shared_context, process_data, process_batch...)
  全部脱离类变成模块级函数。
  AST 验证: ContentMapper 仅剩 9 个方法。

修复: 将 _local_krige_block_worker 移至文件末尾 (模块级正确位置),
  ContentMapper 恢复为 28 个方法的完整类。
2026-07-08 13:02:03 +08:00

3311 lines
151 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

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

import pandas as pd
import numpy as np
import geopandas as gpd
from osgeo import gdal
from pathlib import Path
from typing import Optional, Tuple
from pyproj import CRS, Transformer
import matplotlib.pyplot as plt
import matplotlib.patches as patches
from matplotlib.ticker import FuncFormatter, MaxNLocator
from matplotlib_scalebar.scalebar import ScaleBar
from scipy.interpolate import griddata
from scipy import ndimage
from scipy.spatial.distance import cdist
from scipy.spatial import ConvexHull
from shapely.geometry import Point, Polygon
import rasterio
from rasterio.features import geometry_mask, shapes
from rasterio import windows
from rasterio.warp import calculate_default_transform, reproject, Resampling
try:
from affine import Affine
except ImportError:
try:
from rasterio.transform import Affine
except ImportError:
Affine = None
import warnings
import math
import os
import random
import glob
# 尝试导入pykrige可选依赖
try:
from pykrige.ok import OrdinaryKriging
PYKRIGE_AVAILABLE = True
except ImportError:
PYKRIGE_AVAILABLE = False
print("警告: pykrige未安装Kriging不确定性计算将不可用")
warnings.filterwarnings('ignore')
# 设置中文字体
plt.rcParams['font.sans-serif'] = ['SimHei', 'Microsoft YaHei']
plt.rcParams['axes.unicode_minus'] = False
# 参数到颜色映射的字典
PARAMS_CMAP = {
"Chlorophyll": "YlGnBu_r",
"COD": "coolwarm",
"DO": "RdYlBu",
"pH": "Spectral",
"Temperature": "turbo",
"spCond": "cividis",
"Turbidity": "YlOrBr",
"TDS": "inferno",
"Cl-": "RdYlBu_r",
"NO3-N": "YlOrRd",
"NH3-N": "magma",
"BGA": "viridis",
"TT": "RdYlBu_r"
}
# ── 水色指数英文名 → 中文标题映射(精确整词匹配)────────────────────
# 优先于动态拼接;当文件名恰好命中完整 key 时使用
INDEX_TITLE_MAP = {
# Chl_a 叶绿素a
"Chl_Conc_NDCI": "叶绿素a浓度估算_NDCI模型",
"Chl_MM12NDCI": "叶绿素a相对指数_Matthews12模型",
"Chl_Conc_Gao": "叶绿素a浓度估算_Gao模型",
"Chl_Conc_QAA": "叶绿素a浓度估算_QAA模型",
# BGA 蓝藻/藻蓝蛋白
"BGA_Go04MCI": "蓝藻相对指数_Gower04模型",
"BGA_PC_Conc_Mishra": "藻蓝蛋白浓度估算_Mishra模型",
"BGA_Conc": "蓝藻浓度估算",
"BGA_Am09KBBI": "蓝藻相对指数_Am09模型",
# Turb 浊度
"Turb_Dox02NIRoverRed": "水体浊度指数_Doxaran02模型",
"Turbidity": "水体浊度",
"Turb_Conc": "浊度估算",
# TSM 悬浮物
"TSM_Conc_Bowling": "总悬浮物浓度估算_Bowling模型",
"TSM_Conc": "总悬浮物浓度估算",
# CDOM 有色溶解有机物
"CDOM_Conc": "有色溶解有机物浓度估算",
# WI 水色指数综合
"WaterIndex": "水色指数",
"NDCI": "归一化叶绿素差值指数_NDCI",
"MCI": "最大Chlorophyll指数_MCI",
}
# ── 关键词 → 中文词根映射(用于动态拼接)────────────────────────────
# 顺序即优先级:长词优先,短词兜底
PART_NAME_MAP = [
# 1. 指数/浓度类(最高优先级,描述输出类型)
("Conc", "浓度估算"),
("_Conc", "浓度估算"),
("Concentration", "浓度"),
("Index", "指数"),
("_Index", "指数"),
# 2. 模型/方法标识(次高优先级)
("NDCI", "NDCI模型"),
("MCI", "MCI模型"),
("FLH", "FLH荧光基线"),
("QAA", "QAA模型"),
("Go04", "Gower04"),
("MM12", "Matthews12"),
("Gao", "Gao模型"),
("Dox02", "Doxaran02"),
("Bowling", "Bowling模型"),
("Mishra", "Mishra模型"),
("Am09", "Am09"),
("KBBI", "KBBI"),
("PC", "藻蓝蛋白"),
# 3. 参数大类(核心水质指标)
("BGA", "蓝藻相对指数"),
("Chl", "叶绿素a"),
("Chl_a", "叶绿素a"),
("Turb", "浊度"),
("TSM", "总悬浮物"),
("CDOM", "有色溶解有机物"),
("DO", "溶解氧"),
("pH", "pH值"),
("NH3", "氨氮"),
("NO3", "硝态氮"),
("TDS", "溶解性总固体"),
]
class ContentMapper:
def __init__(self, input_crs='EPSG:32651', output_crs='EPSG:4326'):
"""
初始化ContentMapper - 生成平滑的含量分布图
本类专门用于生成平滑、均匀的颜色分布图,而不是显示离散的采样点。
通过高密度网格插值和多级颜色映射,创建连续的颜色过渡效果。
Parameters:
-----------
input_crs : str
输入坐标系,默认为'EPSG:32651' (WGS_1984_UTM_Zone_51N)
output_crs : str
输出坐标系,默认为'EPSG:4326' (WGS84)
"""
# 定义坐标转换器
self.input_crs = input_crs
self.output_crs = output_crs
self.transformer = Transformer.from_crs(
CRS.from_string(input_crs),
CRS.from_string(output_crs),
always_xy=True
)
# 参数到颜色映射的字典
self.params_cmap = PARAMS_CMAP.copy()
# 所有可用的matplotlib colormap列表用于随机选择
self.available_cmaps = ['viridis', 'plasma', 'inferno', 'magma', 'cividis',
'coolwarm', 'RdYlBu', 'Spectral', 'YlGnBu_r', 'YlOrBr',
'YlOrRd', 'turbo', 'RdYlBu_r', 'cool', 'hot', 'jet']
print(f"坐标转换设置: {input_crs} -> {output_crs}")
# ── 内部工具 ─────────────────────────────────────────────────────
@staticmethod
def _get_chinese_title(stem: str) -> str:
"""
根据 GeoTIFF 文件名 stem 返回中文图表标题(绝对唯一)。
匹配策略:
1. 精确整词命中 INDEX_TITLE_MAP
2. 中文分类前缀 + 原始模型后缀(绝对唯一保证)
3. 未匹配任何关键词 → 返回原英文 stem
Parameters
----------
stem : str
GeoTIFF 文件名(不含路径和扩展名)
Returns
-------
str
中文标题;若未匹配则返回英文 stem
"""
# 策略1精确整词匹配优先级最高
if stem in INDEX_TITLE_MAP:
return INDEX_TITLE_MAP[stem]
# 策略2中文分类前缀 + 原始模型后缀(确保绝对唯一)
category = ""
suffix = ""
if stem.startswith("BGA_PC_Conc"):
category = "藻蓝蛋白浓度估算"
suffix = stem.replace("BGA_PC_Conc_", "")
elif stem.startswith("BGA"):
category = "蓝藻相对指数"
suffix = stem.replace("BGA_", "")
elif stem.startswith("Chl_Conc"):
category = "叶绿素a浓度估算"
suffix = stem.replace("Chl_Conc_", "")
elif stem.startswith("Chl"):
category = "叶绿素a相对指数"
suffix = stem.replace("Chl_", "")
elif stem.startswith("Turb_Conc"):
category = "浊度浓度估算"
suffix = stem.replace("Turb_Conc_", "")
elif stem.startswith("Turb"):
category = "浊度相对指数"
suffix = stem.replace("Turb_", "")
elif stem.startswith("TSM_Conc"):
category = "总悬浮物浓度估算"
suffix = stem.replace("TSM_Conc_", "")
else:
# 兜底机制:未知分类直接使用原名
category = "水质参数"
suffix = stem
return f"{category}_{suffix}"
def _extract_param_name(self, csv_file):
"""
从CSV文件名或内容中提取参数名称
Parameters:
-----------
csv_file : str
CSV文件路径
Returns:
--------
param_name : str or None
提取的参数名称如果未找到则返回None
"""
print(f"[调试] 开始从文件 {csv_file} 中提取参数名称")
print(f"[调试] 字典中的参数键: {list(self.params_cmap.keys())}")
# 从文件名中提取(去除路径和扩展名)
file_name = os.path.basename(csv_file)
file_name_no_ext = os.path.splitext(file_name)[0]
print(f"[调试] 文件名(不含扩展名): {file_name_no_ext}")
# 尝试从文件名中匹配参数名称(不区分大小写)
file_name_upper = file_name_no_ext.upper()
for param in self.params_cmap.keys():
param_upper = param.upper()
if param_upper in file_name_upper:
print(f"从文件名中识别到参数: {param} (匹配到 '{param_upper}''{file_name_upper}' 中)")
return param # 返回字典中的原始键(保持大小写)
# 如果文件名中没有找到尝试从CSV内容中提取检查列名
try:
df = pd.read_csv(csv_file, encoding='utf-8', nrows=0) # 只读取列名
columns = [col.upper() for col in df.columns]
print(f"[调试] CSV列名: {list(df.columns)}")
for param in self.params_cmap.keys():
param_upper = param.upper()
# 检查列名中是否包含参数名称
for col in columns:
if param_upper in col or col in param_upper:
print(f"从CSV列名中识别到参数: {param} (匹配到列名 '{col}')")
return param # 返回字典中的原始键(保持大小写)
except Exception as e:
print(f"读取CSV列名时出错: {e}")
print(f"未能在文件 {csv_file} 中识别参数名称")
print(f"[调试] 可用的参数名: {list(self.params_cmap.keys())}")
return None
def _get_colormap(self, param_name=None):
"""
根据参数名称获取对应的colormap支持精确匹配、模糊匹配
Parameters
----------
param_name : str, optional
参数名称。如果为None或不在映射中则随机选择一个colormap
Returns
-------
cmap : str
颜色映射名称
"""
print(f"[调试] _get_colormap 被调用param_name={param_name}")
print(f"[调试] 当前字典中的键: {list(self.params_cmap.keys())}")
if param_name:
# 精确匹配(区分大小写)
if param_name in self.params_cmap:
cmap = self.params_cmap[param_name]
print(f"使用参数 '{param_name}' 对应的颜色映射: {cmap}")
return cmap
# 不区分大小写匹配
param_name_upper = param_name.upper()
for key in self.params_cmap.keys():
if key.upper() == param_name_upper:
cmap = self.params_cmap[key]
print(f"使用参数 '{key}' (不区分大小写匹配 '{param_name}') 对应的颜色映射: {cmap}")
return cmap
# ── 模糊匹配(关键字包含检测)───────────────────────────
pn_upper = param_name.upper()
pn_lower = param_name.lower()
# 蓝藻 / BGA / Phycocyanin → YlGn蓝绿色系
if any(k in pn_upper for k in ('BGA', 'PHYCO', 'CYAN', '蓝藻', '藻蓝')):
cmap = self.params_cmap.get('BGA', 'YlGn')
print(f"模糊匹配 BGA/Phycocyanin → '{cmap}'")
return cmap
# 叶绿素 / Chlorophyll / Chl → YlGn绿色系
if any(k in pn_upper for k in ('CHL', '叶绿素', 'CHLORO')):
cmap = self.params_cmap.get('Chl_a', 'YlGn')
print(f"模糊匹配 Chl/叶绿素 → '{cmap}'")
return cmap
# CDOM / 有色溶解有机物
if any(k in pn_upper for k in ('CDOM', '色DOM', '有色溶解')):
cmap = self.params_cmap.get('CDOM', 'BrBG')
print(f"模糊匹配 CDOM → '{cmap}'")
return cmap
# 悬浮物 / TSM / SS → YlOrBr黄棕系
if any(k in pn_upper for k in ('TSM', 'SS', '悬浮物', '总悬浮')):
cmap = self.params_cmap.get('TSM', 'YlOrBr')
print(f"模糊匹配 TSM/悬浮物 → '{cmap}'")
return cmap
# 透明度 / SD / Secchi → Blues蓝色系
if any(k in pn_upper for k in ('SD', 'SECCHI', '透明度', '透明')):
cmap = self.params_cmap.get('SD', 'Blues')
print(f"模糊匹配 SD/透明度 → '{cmap}'")
return cmap
# 氨氮 / NH4 / NH3 → Oranges
if any(k in pn_upper for k in ('NH4', 'NH3', '氨氮', '')):
cmap = 'Oranges'
print(f"模糊匹配 NH4/氨氮 → '{cmap}'")
return cmap
# 总磷 / TP / 总氮 / TN → RdYlGn
if any(k in pn_upper for k in ('TP', '总磷')):
cmap = 'RdYlGn_r'
print(f"模糊匹配 TP/总磷 → '{cmap}'")
return cmap
if any(k in pn_upper for k in ('TN', '总氮')):
cmap = 'RdYlGn_r'
print(f"模糊匹配 TN/总氮 → '{cmap}'")
return cmap
# 高浊度 / Turbidity → PuBu紫蓝系
if any(k in pn_upper for k in ('TURBIDITY', '浊度', 'TURB')):
cmap = 'PuBu'
print(f"模糊匹配 Turbidity/浊度 → '{cmap}'")
return cmap
# 溶解氧 / DO → cool蓝白冷色
if any(k in pn_upper for k in ('DO', '溶解氧', 'DISSOLVED')):
cmap = 'cool'
print(f"模糊匹配 DO/溶解氧 → '{cmap}'")
return cmap
# 仍不匹配 → 随机
cmap = random.choice(self.available_cmaps)
print(f"警告: 参数 '{param_name}' 不在映射中,随机选择颜色映射: {cmap}")
print(f"可用的参数名: {list(self.params_cmap.keys())}")
return cmap
else:
cmap = random.choice(self.available_cmaps)
print(f"未指定参数名称,随机选择颜色映射: {cmap}")
return cmap
def _check_point_distribution(self, points):
"""检查数据点的几何分布"""
print("正在检查数据点分布...")
# 检查是否有重复点
unique_points = np.unique(points, axis=0)
if len(unique_points) < len(points):
print(f"警告:发现 {len(points) - len(unique_points)} 个重复数据点")
# 检查点是否共线
if len(unique_points) >= 3:
# 计算前三个点构成的三角形面积
p1, p2, p3 = unique_points[:3]
area = 0.5 * abs((p2[0] - p1[0]) * (p3[1] - p1[1]) - (p3[0] - p1[0]) * (p2[1] - p1[1]))
if area < 1e-10: # 面积太小,可能共线
print("警告:前三个数据点可能共线")
# 尝试找到不共线的点
for i in range(3, len(unique_points)):
p4 = unique_points[i]
area = 0.5 * abs((p2[0] - p1[0]) * (p4[1] - p1[1]) - (p4[0] - p1[0]) * (p2[1] - p1[1]))
if area > 1e-10:
print(f"找到非共线点,使用点 {i}")
break
else:
print("警告:所有数据点可能都共线,这会导致插值失败")
# 检查坐标范围
x_range = points[:, 0].max() - points[:, 0].min()
y_range = points[:, 1].max() - points[:, 1].min()
if x_range < 1e-6 or y_range < 1e-6:
print(f"警告:坐标范围很小 (X范围: {x_range:.2e}, Y范围: {y_range:.2e})")
print("这可能导致插值数值不稳定")
return unique_points
def _fill_boundary_blanks_with_distance_diffusion(self, grid_content, grid_xx, grid_yy, mask,
boundary_gdf, max_diffusion_distance=None,
power=2, n_neighbors=15):
"""
使用距离扩散方法填充边界附近的空白区域
Parameters:
-----------
grid_content : np.ndarray
插值网格数据
grid_xx : np.ndarray
网格X坐标
grid_yy : np.ndarray
网格Y坐标
mask : np.ndarray
边界掩膜True表示边界内
boundary_gdf : gpd.GeoDataFrame
边界几何数据
max_diffusion_distance : float, optional
最大扩散距离单位与坐标相同。如果为None自动计算为网格分辨率的5倍
power : float, default=2
IDW距离衰减幂参数
n_neighbors : int, default=15
使用的最近邻点数
Returns:
--------
grid_content : np.ndarray
填充后的网格数据
"""
print("正在使用距离扩散方法填充边界空白区域...")
# 找到边界内的空白区域
nan_mask = np.isnan(grid_content)
within_boundary_nan = nan_mask & mask
if not np.any(within_boundary_nan):
print("边界内没有空白区域需要填充")
return grid_content
blank_count = np.sum(within_boundary_nan)
print(f"发现 {blank_count} 个边界内的空白点,开始距离扩散填充...")
# 找到边界内有效值的点
valid_mask = ~nan_mask & mask
if np.sum(valid_mask) == 0:
print("警告:边界内没有有效值,无法进行距离扩散")
return grid_content
# 计算网格分辨率(用于确定最大扩散距离)
if max_diffusion_distance is None:
# 自动计算:使用网格点之间的平均距离
dx = np.abs(grid_xx[0, 1] - grid_xx[0, 0]) if grid_xx.shape[1] > 1 else 1.0
dy = np.abs(grid_yy[1, 0] - grid_yy[0, 0]) if grid_xx.shape[0] > 1 else 1.0
avg_resolution = (dx + dy) / 2.0
max_diffusion_distance = avg_resolution * 5.0 # 5倍网格分辨率
print(f"自动计算最大扩散距离: {max_diffusion_distance:.6f}")
# 准备有效数据点
valid_points = np.column_stack((grid_xx[valid_mask], grid_yy[valid_mask]))
valid_values = grid_content[valid_mask]
# 准备空白点
blank_points = np.column_stack((grid_xx[within_boundary_nan], grid_yy[within_boundary_nan]))
print(f"使用 {len(valid_points)} 个有效点填充 {len(blank_points)} 个空白点...")
# 使用向量化计算距离矩阵
distances = cdist(blank_points, valid_points)
# 对每个空白点找到最近的n_neighbors个有效点
n_neighbors = min(n_neighbors, len(valid_points))
# 应用最大扩散距离限制
if max_diffusion_distance > 0:
# 只考虑在最大扩散距离内的点
# 对于每个空白点,找到在扩散距离内的最近邻
filled_values = np.full(len(blank_points), np.nan)
global_mean = np.nanmean(valid_values)
for i in range(len(blank_points)):
point_distances = distances[i, :]
valid_idx = point_distances <= max_diffusion_distance
if np.any(valid_idx):
# 找到最近的n_neighbors个点在扩散距离内
valid_dist = point_distances[valid_idx]
valid_vals = valid_values[valid_idx]
# 如果有效点数量超过n_neighbors只取最近的n_neighbors个
if len(valid_dist) > n_neighbors:
nearest_idx = np.argpartition(valid_dist, n_neighbors-1)[:n_neighbors]
valid_dist = valid_dist[nearest_idx]
valid_vals = valid_vals[nearest_idx]
# 避免除零
valid_dist = np.maximum(valid_dist, 1e-10)
# 计算IDW权重
weights = 1.0 / (valid_dist ** power)
weight_sum = np.sum(weights)
if weight_sum > 0:
# 距离加权平均
filled_values[i] = np.sum(weights * valid_vals) / weight_sum
else:
filled_values[i] = global_mean
else:
# 如果该点不在任何有效点的扩散距离内,使用全局平均值
filled_values[i] = global_mean
else:
# 不使用距离限制对所有点进行IDW插值批量处理以提高效率
# 对每个空白点找到最近的n_neighbors个点
nearest_indices = np.argpartition(distances, n_neighbors-1, axis=1)[:, :n_neighbors]
# 批量提取距离和值
nearest_dists = np.take_along_axis(distances, nearest_indices, axis=1)
nearest_vals = valid_values[nearest_indices]
# 避免除零
nearest_dists = np.maximum(nearest_dists, 1e-10)
# 批量计算IDW权重
weights = 1.0 / (nearest_dists ** power)
weight_sums = np.sum(weights, axis=1)
# 批量计算加权平均值
filled_values = np.sum(weights * nearest_vals, axis=1) / weight_sums
# 处理可能的NaN如果weight_sum为0
nan_mask = np.isnan(filled_values) | (weight_sums == 0)
if np.any(nan_mask):
filled_values[nan_mask] = np.nanmean(valid_values)
# 填充空白点
grid_content[within_boundary_nan] = filled_values
# 检查填充结果
filled_count = np.sum(~np.isnan(filled_values))
print(f"距离扩散填充完成:成功填充 {filled_count} / {blank_count} 个空白点")
return grid_content
def _perform_interpolation(self, points, values, grid_xx, grid_yy):
"""三级降级插值策略Kriging → IDW → 最近邻
2026-07-01 重构:
- Kriging 优先(自动拟合球形变异函数,不强制 nugget
- Kriging 退化检测:若结果标准差接近 0纯色图自动回退 IDW
- IDW 作为首选回退:无需拟合变异函数,不会产生纯色图
- scipy linear/nearest 作为最后兜底
"""
print(f"插值输入检查:")
print(f" - 数据点数量: {len(points)}")
print(f" - 数据值范围: {values.min():.4f} - {values.max():.4f}")
print(f" - 网格大小: {grid_xx.shape}")
print(f" - 坐标系: {self.output_crs}")
# 检查数据的有效性
finite_mask = np.isfinite(values)
if not np.all(finite_mask):
print(f"警告:发现 {np.sum(~finite_mask)} 个无效数据值,将被移除")
points = points[finite_mask]
values = values[finite_mask]
if len(points) < 3:
raise ValueError(f"有效数据点不足3个当前{len(points)}个)")
# ── 策略 0值域极窄时直接跳过 Kriging ──
value_range = float(values.max() - values.min())
value_std = float(np.std(values))
# ═══════════════════════════════════════════════════════════
# 策略 1局部克里金 (Local Kriging)
# 网格按 500m 空间窗口分块,每块只取窗口内 + 500m 缓冲区的
# 局部采样点参与计算。协方差矩阵从全局 8242×8242 降为
# 局部 n×n (n≈几十到几百),千万级网格秒级完成。
# ═══════════════════════════════════════════════════════════
kriging_degraded = False
if PYKRIGE_AVAILABLE:
try:
grid_x = grid_xx[0, :]
grid_y = grid_yy[:, 0]
total_cells = len(grid_x) * len(grid_y)
# 局部窗口:每个块使用 500m 空间窗口内的采样点
_LOCAL_WINDOW = 500.0 # 米
_n_workers = min(os.cpu_count() or 4, 8)
print(f"正在使用 局部克里金 (Local Kriging)"
f"网格={total_cells:,} 点, 窗口={_LOCAL_WINDOW}m, "
f"{_n_workers} workers")
grid_content = self._local_kriging(
points, values, grid_x, grid_y,
window_size=_LOCAL_WINDOW,
n_workers=_n_workers,
)
valid_mask = ~np.isnan(grid_content)
valid_count = int(np.sum(valid_mask))
if valid_count > 0:
kriging_std = float(np.nanstd(grid_content))
degradation_ratio = kriging_std / max(value_std, 1e-12)
print(f"局部 Kriging 完成: 有效点={valid_count}/{grid_content.size}, "
f"输出std={kriging_std:.6f}, 退化比={degradation_ratio:.3f}")
if degradation_ratio < 0.05 and value_range > 1e-8:
print(f"⚠ Kriging 严重退化,回退 IDW")
kriging_degraded = True
else:
return grid_content
else:
print("局部 Kriging 结果全为 NaN回退")
kriging_degraded = True
except Exception as e:
print(f"Kriging 失败: {e}")
kriging_degraded = True
else:
print("pykrige 未安装,跳过 Kriging")
kriging_degraded = True
# ═══════════════════════════════════════════════════════════
# 策略 2IDW反距离权重— 不需要拟合变异函数,绝不纯色
# ═══════════════════════════════════════════════════════════
if kriging_degraded:
try:
print("正在使用 IDW 插值(反距离权重, power=2, neighbors=15...")
grid_content = self._idw_interpolation(
points, values, grid_xx, grid_yy,
power=2, n_neighbors=min(15, len(points)),
)
valid_count = int(np.sum(~np.isnan(grid_content)))
if valid_count > 0:
idw_std = float(np.nanstd(grid_content))
print(f"IDW 完成: 有效点={valid_count}/{grid_content.size}, 输出std={idw_std:.6f}")
if idw_std > 0:
return grid_content
else:
print("IDW std=0所有输入值完全相同结果可用")
return grid_content
else:
print("IDW 结果全为 NaN回退 scipy 插值")
except Exception as e:
print(f"IDW 失败: {e},回退 scipy 插值")
# ═══════════════════════════════════════════════════════════
# 策略 3scipy 线性插值 + 最近邻填充(最终兜底)
# ═══════════════════════════════════════════════════════════
try:
print("正在尝试 scipy 线性插值...")
grid_content = griddata(
points, values, (grid_xx, grid_yy),
method='linear', fill_value=np.nan
)
valid_count = int(np.sum(~np.isnan(grid_content)))
print(f"线性插值: 有效点={valid_count}/{grid_content.size}")
if valid_count > 0:
nan_count = int(np.sum(np.isnan(grid_content)))
if nan_count > 0:
print(f"用最近邻填充 {nan_count} 个 NaN...")
grid_nearest = griddata(
points, values, (grid_xx, grid_yy), method='nearest'
)
grid_content[np.isnan(grid_content)] = grid_nearest[np.isnan(grid_content)]
return grid_content
except Exception as e:
print(f"线性插值失败: {e}")
# ═══════════════════════════════════════════════════════════
# 策略 4最近邻绝对兜底
# ═══════════════════════════════════════════════════════════
print("执行最近邻插值(最终兜底)...")
grid_content = griddata(
points, values, (grid_xx, grid_yy), method='nearest'
)
if np.sum(~np.isnan(grid_content)) == 0:
raise ValueError("所有插值方法均失败")
return grid_content
def _local_kriging(self, points, values, grid_x, grid_y,
window_size=500.0, n_workers=4):
"""局部克里金:网格按 window_size 分块,每块只用块内+缓冲区的局部采样点
核心思想:协方差矩阵锁定在局部 n×n几十到几百而非全局 8242×8242。
千万级网格秒级完成,精度几乎无损。
"""
import multiprocessing
# 计算网格空间范围
x_min, x_max = float(grid_x[0]), float(grid_x[-1])
y_min, y_max = float(grid_y[0]), float(grid_y[-1])
grid_dx = float(grid_x[1] - grid_x[0]) if len(grid_x) > 1 else 1.0
grid_dy = float(grid_y[1] - grid_y[0]) if len(grid_y) > 1 else 1.0
# 按 window_size 将空间分割为不重叠的块
# 每个块向四周扩展一个 window_size 作为缓冲区(获取局部采样点)
n_blocks_x = max(1, int(np.ceil((x_max - x_min) / window_size)))
n_blocks_y = max(1, int(np.ceil((y_max - y_min) / window_size)))
block_dx = (x_max - x_min) / n_blocks_x
block_dy = (y_max - y_min) / n_blocks_y
total_blocks = n_blocks_x * n_blocks_y
print(f" 空间分块: {n_blocks_x}×{n_blocks_y} = {total_blocks}"
f"(窗口={window_size}m, 缓冲区={window_size}m)")
# 构建任务列表
tasks = []
for iy in range(n_blocks_y):
for ix in range(n_blocks_x):
# 块范围(不含缓冲区)
bx_min = x_min + ix * block_dx
bx_max = x_min + (ix + 1) * block_dx
by_min = y_min + iy * block_dy
by_max = y_min + (iy + 1) * block_dy
# 带缓冲区的搜索范围
search_xmin = bx_min - window_size
search_xmax = bx_max + window_size
search_ymin = by_min - window_size
search_ymax = by_max + window_size
# 缓冲区内的网格点索引
grid_mask_x = (grid_x >= search_xmin) & (grid_x <= search_xmax)
grid_mask_y = (grid_y >= search_ymin) & (grid_y <= search_ymax)
if not np.any(grid_mask_x) or not np.any(grid_mask_y):
continue
sub_grid_x = grid_x[grid_mask_x]
sub_grid_y = grid_y[grid_mask_y]
# 缓冲区内的采样点
point_mask = (
(points[:, 0] >= search_xmin) & (points[:, 0] <= search_xmax) &
(points[:, 1] >= search_ymin) & (points[:, 1] <= search_ymax)
)
local_pts = points[point_mask]
local_vals = values[point_mask]
sub_n = len(sub_grid_x) * len(sub_grid_y)
tasks.append((
local_pts, local_vals,
sub_grid_x, sub_grid_y,
bx_min, bx_max, by_min, by_max, # 只保留核心区域结果
grid_dx, grid_dy,
ix, iy, total_blocks,
))
print(f" 有效块: {len(tasks)} (含采样点的空间块)")
if len(tasks) <= 1:
return self._local_krige_block(tasks[0])
# 多进程执行
print(f" 启动 {min(n_workers, len(tasks))} 个 worker 进程...")
with multiprocessing.Pool(processes=min(n_workers, len(tasks))) as pool:
results = pool.map(_local_krige_block_worker, tasks)
# 拼接结果:初始化全 NaN 数组,逐块填充核心区域
grid_full = np.full((len(grid_y), len(grid_x)), np.nan, dtype=np.float64)
for (block_result, bx_min, bx_max, by_min, by_max) in results:
if block_result is None:
continue
# 将子网格结果映射回全局索引
gx_mask = (grid_x >= bx_min) & (grid_x < bx_max)
gy_mask = (grid_y >= by_min) & (grid_y < by_max)
if np.any(gx_mask) and np.any(gy_mask):
# 子网格中对应核心区域的索引
h, w = block_result.shape
ix_start = np.searchsorted(grid_x, bx_min)
iy_start = np.searchsorted(grid_y, by_min)
iy_end = min(iy_start + h, len(grid_y))
ix_end = min(ix_start + w, len(grid_x))
grid_full[iy_start:iy_end, ix_start:ix_end] = block_result[:iy_end-iy_start, :ix_end-ix_start]
return grid_full
def _local_krige_block(self, local_pts, local_vals,
sub_grid_x, sub_grid_y,
bx_min, bx_max, by_min, by_max,
grid_dx, grid_dy,
block_ix=0, block_iy=0, total=1):
"""单块局部克里金"""
n_pts = len(local_pts)
n_cells = len(sub_grid_x) * len(sub_grid_y)
if n_pts < 3:
return None, bx_min, bx_max, by_min, by_max
ok = OrdinaryKriging(
local_pts[:, 0], local_pts[:, 1], local_vals,
variogram_model='spherical',
verbose=False,
enable_plotting=False,
)
z, ss = ok.execute(
'grid', sub_grid_x, sub_grid_y,
backend='loop',
n_closest_points=min(15, n_pts),
)
return np.array(z), bx_min, bx_max, by_min, by_max
@staticmethod
def _idw_interpolation(points, values, grid_xx, grid_yy,
power=2, n_neighbors=15):
"""IDW反距离权重插值 — 不需拟合模型,绝不产生纯色图。
Parameters:
points: (N, 2) 采样点坐标
values: (N,) 采样点值
grid_xx, grid_yy: meshgrid 网格
power: 距离衰减幂参数(默认 2
n_neighbors: 每个网格点参考的最近邻数量
"""
from scipy.spatial import cKDTree
grid_shape = grid_xx.shape
grid_flat = np.column_stack((grid_xx.ravel(), grid_yy.ravel()))
values_flat = values.ravel()
tree = cKDTree(points)
k = min(n_neighbors, len(points))
distances, indices = tree.query(grid_flat, k=k)
# 防止距离为 0 的除零
distances = np.maximum(distances, 1e-12)
weights = 1.0 / (distances ** power)
# 归一化权重
weights /= weights.sum(axis=1, keepdims=True)
# 加权求和
neighbor_vals = values_flat[indices] if k == 1 else values_flat[indices]
if k == 1:
result = neighbor_vals
else:
result = np.sum(weights * neighbor_vals, axis=1)
return result.reshape(grid_shape)
def read_csv_data(self, csv_file, uncertainty_col=None):
"""
读取CSV文件并进行坐标转换
Parameters:
-----------
csv_file : str
CSV文件路径
uncertainty_col : str, optional
不确定性数据列名。如果为None将自动检测包含'variance''uncertainty''std''sigma''var''mc_dropout'的列
Returns:
--------
gdf : gpd.GeoDataFrame
包含坐标和含量数据的GeoDataFrame如果找到不确定性列会包含'uncertainty'
"""
print("正在读取CSV文件...")
df = pd.read_csv(csv_file, encoding='utf-8')
# 假设前三列分别是经度、纬度、含量
if df.shape[1] < 3:
raise ValueError("CSV文件必须至少包含3列经度、纬度、含量")
# 获取列名
lon_col = df.columns[0]
lat_col = df.columns[1]
content_col = df.columns[2]
print(f"检测到列名:经度({lon_col}),纬度({lat_col}),含量({content_col})")
# 自动检测不确定性列
if uncertainty_col is None:
uncertainty_keywords = ['variance', 'uncertainty', 'std', 'sigma', 'var', 'mc_dropout']
for col in df.columns:
col_lower = col.lower()
if any(keyword in col_lower for keyword in uncertainty_keywords):
uncertainty_col = col
print(f"自动检测到不确定性列: {uncertainty_col}")
break
# 坐标转换
print(f"正在进行坐标转换: {self.input_crs} -> {self.output_crs}")
transformed_x, transformed_y = self.transformer.transform(
df[lon_col].values,
df[lat_col].values
)
# 创建GeoDataFrame
geometry = [Point(x, y) for x, y in zip(transformed_x, transformed_y)]
gdf = gpd.GeoDataFrame(
df,
geometry=geometry,
crs=self.output_crs
)
gdf['proj_x'] = transformed_x
gdf['proj_y'] = transformed_y
gdf['content'] = df[content_col]
# 如果找到不确定性列添加到GeoDataFrame
if uncertainty_col and uncertainty_col in df.columns:
gdf['uncertainty'] = df[uncertainty_col].values
print(f"已加载不确定性数据列: {uncertainty_col}")
print(f"不确定性值范围: {gdf['uncertainty'].min():.4f} - {gdf['uncertainty'].max():.4f}")
print(f"成功读取 {len(gdf)} 个数据点")
return gdf
def read_boundary_shapefile(self, shp_file):
"""读取边界/掩膜文件(同时支持矢量 .shp 与栅格 .dat/.bsq/.tif/.tiff
- .shp → gpd.read_file 读取矢量边界(保持原行为)
- .dat/.bsq/.tif/.tiff/.img → rasterio 读取栅格水掩膜 → rasterio.features.shapes
矢量化成水体多边形 → gpd.GeoDataFrame 返回
下游 create_interpolation_grid / create_content_map / visualize_raster
始终接收 GeoDataFrame无需任何改动。
"""
print("正在读取边界/掩膜文件...")
suffix = Path(shp_file).suffix.lower()
if suffix in (".shp",):
boundary = gpd.read_file(shp_file)
elif suffix in (".dat", ".bsq", ".tif", ".tiff", ".img"):
boundary = self._raster_to_boundary_gdf(shp_file)
else:
raise ValueError(
f"不支持的边界/掩膜文件格式: {suffix}(仅支持 .shp / .dat / .bsq / .tif / .img"
)
if len(boundary) == 0:
raise ValueError(
f"边界/掩膜 {shp_file} 矢量化后为空(栅格格式请确认 .dat 包含水体像元 > 0"
)
# 确保边界文件使用目标投影坐标系
if boundary.crs is not None and boundary.crs != self.output_crs:
print(f"正在转换边界/掩膜坐标系到 {self.output_crs}...")
boundary = boundary.to_crs(self.output_crs)
print(f"边界/掩膜文件包含 {len(boundary)} 个要素")
return boundary
def _raster_to_boundary_gdf(self, raster_path):
"""把栅格二值水掩膜(.dat/.bsq/.tif/.tiff矢量化成水体多边形 GeoDataFrame。
修复 Step 11 接收 step1 产出 .dat 水掩膜的兼容性:
- rasterio.open 读 band 10=非水, 1/任意正数=水)
- rasterio.features.shapes 矢量化成多边形
- 收集所有 val=1 的多边形 → gpd.GeoDataFrame
"""
try:
from shapely.geometry import shape as _shapely_shape
except ImportError as e:
raise ImportError(
"栅格掩膜矢量化需要 shapelygeopandas 自带)。原始错误: " + str(e)
)
with rasterio.open(raster_path) as src:
data = src.read(1)
transform = src.transform
crs = src.crs
# 【新增防御】检测 Transform 是否为纯像素矩阵(极易导致严重错位)
if transform.is_identity:
print("\n" + "!" * 65)
print(f"⚠️ [严重警告] 栅格掩膜 {Path(raster_path).name} 缺少真实的地理仿射变换!")
print(f"当前被判定为纯像素坐标系 (x: 0~Width, y: 0~Height)。\n强行与地理栅格 (UTM/WGS84) 叠加将发生【极其严重的错位】!")
print(f"请务必在 Step 1 导入原始的 .shp 矢量文件进行约束。")
print("!" * 65 + "\n")
# 二值化:>0 视为水
mask_uint8 = (data > 0).astype(np.uint8)
if int(mask_uint8.sum()) == 0:
raise ValueError(
f"栅格掩膜 {raster_path} 中无水体像元(>0无法矢量化"
)
# 矢量化shapes 返回 (geom_dict, value) 迭代器
geoms = []
for geom_dict, val in shapes(mask_uint8, mask=mask_uint8.astype(bool), transform=transform):
if int(val) == 1:
geoms.append(_shapely_shape(geom_dict))
if not geoms:
raise ValueError(f"栅格掩膜 {raster_path} 矢量化后无有效水体多边形")
gdf = gpd.GeoDataFrame(geometry=geoms, crs=crs)
print(
f"栅格掩膜 {Path(raster_path).name} 矢量化完成: "
f"{int(mask_uint8.sum())} 个水体像元 → {len(gdf)} 个多边形 (CRS={crs})"
)
return gdf
def _identify_edge_points(self, points_gdf):
"""
识别边缘采样点(使用凸包方法)
Parameters:
-----------
points_gdf : gpd.GeoDataFrame
采样点GeoDataFrame
Returns:
--------
edge_indices : np.ndarray
边缘点的索引数组
"""
print("正在识别边缘采样点...")
# 获取所有点的坐标
points = np.column_stack((points_gdf['proj_x'].values, points_gdf['proj_y'].values))
if len(points) < 3:
print("警告采样点数量少于3个无法识别边缘点")
return np.array([])
try:
# 使用凸包识别边缘点
hull = ConvexHull(points)
edge_indices = hull.vertices
print(f"识别到 {len(edge_indices)} 个边缘采样点(共 {len(points)} 个点)")
return edge_indices
except Exception as e:
print(f"识别边缘点时出错: {e},将使用所有点作为边缘点")
return np.arange(len(points))
def _expand_edge_points(self, points_gdf, boundary_gdf=None, resolution=100, expand_ratio=0.05):
"""
对边缘采样点进行外扩处理,外扩到整个图像的边界(包括外扩后的边界)
按照指定的间距resolution生成外扩点铺满整个画面。
★★★ Plan Cboundary_gdf 可选None = 不依赖水域掩膜,纯采样点自然扩展)★★★
Parameters:
-----------
points_gdf : gpd.GeoDataFrame
原始采样点GeoDataFrame
boundary_gdf : gpd.GeoDataFrame, optional
水域掩膜边界GeoDataFrame。None 时基于采样点自身范围做外扩
resolution : float, default=100
外扩点的间距(单位与坐标相同),与插值网格分辨率一致
expand_ratio : float, default=0.05
边界外扩比例与create_interpolation_grid中的expand_ratio一致
Returns:
--------
expanded_gdf : gpd.GeoDataFrame
外扩后的采样点GeoDataFrame
"""
# ── Plan C: 无水域掩膜时,基于采样点自身范围外扩 ──────────────
if boundary_gdf is None:
print(f"[Plan C] 无水域掩膜基于采样点范围做自然外扩expand_ratio={expand_ratio}...")
points = np.column_stack(
(points_gdf['proj_x'].values, points_gdf['proj_y'].values)
)
p_minx, p_miny = points.min(axis=0)
p_maxx, p_maxy = points.max(axis=0)
width = p_maxx - p_minx
height = p_maxy - p_miny
expand_x = width * expand_ratio
expand_y = height * expand_ratio
image_minx = p_minx - expand_x
image_maxx = p_maxx + expand_x
image_miny = p_miny - expand_y
image_maxy = p_maxy + expand_y
print(f"[Plan C] 采样点范围: X[{p_minx:.2f}, {p_maxx:.2f}], Y[{p_miny:.2f}, {p_maxy:.2f}]")
print(f"[Plan C] 外扩后范围: X[{image_minx:.2f}, {image_maxx:.2f}], Y[{image_miny:.2f}, {image_maxy:.2f}]")
else:
boundary_bounds = boundary_gdf.total_bounds
mask_minx, mask_miny, mask_maxx, mask_maxy = boundary_bounds
width = mask_maxx - mask_minx
height = mask_maxy - mask_miny
expand_x = width * expand_ratio
expand_y = height * expand_ratio
image_minx = mask_minx - expand_x
image_maxx = mask_maxx + expand_x
image_miny = mask_miny - expand_y
image_maxy = mask_maxy + expand_y
print(f"正在对边缘采样点进行外扩处理(按照 {resolution} 的间距外扩到整个图像边界)...")
# 识别边缘点
edge_indices = self._identify_edge_points(points_gdf)
if len(edge_indices) == 0:
print("未识别到边缘点,跳过外扩处理")
return points_gdf.copy()
# 获取所有点的坐标
points = np.column_stack((points_gdf['proj_x'].values, points_gdf['proj_y'].values))
# 计算点集的范围和中心Plan C 两个分支都要用,提前算)
x_min, x_max = points[:, 0].min(), points[:, 0].max()
y_min, y_max = points[:, 1].min(), points[:, 1].max()
center = np.array([(x_min + x_max) / 2, (y_min + y_max) / 2])
# 存储新添加的点
new_points_list = []
new_data_list = []
# 对每个边缘点进行外扩
for edge_idx in edge_indices:
edge_point = points[edge_idx]
# 计算从中心到边缘点的方向向量
direction = edge_point - center
distance_to_center = np.linalg.norm(direction)
if distance_to_center < 1e-10:
# 如果边缘点就是中心点,跳过
continue
# 归一化方向向量
direction_unit = direction / distance_to_center
# 计算该方向与水域掩膜边界的交点
# 使用射线法:从边缘点沿方向延伸,找到与边界框的交点
max_distance = 0
# 检查与四个边界的交点(使用整个图像的范围,包括外扩后的边界)
# 上边界 (y = image_maxy)
if direction_unit[1] > 1e-10: # 向上
t = (image_maxy - edge_point[1]) / direction_unit[1]
if t > 0:
intersect_x = edge_point[0] + direction_unit[0] * t
if image_minx <= intersect_x <= image_maxx:
max_distance = max(max_distance, t)
# 下边界 (y = image_miny)
if direction_unit[1] < -1e-10: # 向下
t = (image_miny - edge_point[1]) / direction_unit[1]
if t > 0:
intersect_x = edge_point[0] + direction_unit[0] * t
if image_minx <= intersect_x <= image_maxx:
max_distance = max(max_distance, t)
# 右边界 (x = image_maxx)
if direction_unit[0] > 1e-10: # 向右
t = (image_maxx - edge_point[0]) / direction_unit[0]
if t > 0:
intersect_y = edge_point[1] + direction_unit[1] * t
if image_miny <= intersect_y <= image_maxy:
max_distance = max(max_distance, t)
# 左边界 (x = image_minx)
if direction_unit[0] < -1e-10: # 向左
t = (image_minx - edge_point[0]) / direction_unit[0]
if t > 0:
intersect_y = edge_point[1] + direction_unit[1] * t
if image_miny <= intersect_y <= image_maxy:
max_distance = max(max_distance, t)
# 如果找到了边界交点按照resolution间距创建外扩点
if max_distance > 1e-10:
# 从边缘点到边界的距离
distance_to_boundary = max_distance
# 计算需要生成的外扩点数量按照resolution间距
# 使用ceil确保能铺满到边界
n_points = int(np.ceil(distance_to_boundary / resolution))
# 从边缘点开始按照resolution间距生成点直到边界
for i in range(1, n_points + 1):
# 计算外扩距离从边缘点开始按照resolution间距
expand_distance = i * resolution
# 如果超过边界距离,使用边界距离作为最后一个点
if expand_distance >= distance_to_boundary:
expand_distance = distance_to_boundary
# 计算新点位置
new_point = edge_point + direction_unit * expand_distance
# 确保新点在图像范围内(包括外扩后的边界)
new_point[0] = np.clip(new_point[0], image_minx, image_maxx)
new_point[1] = np.clip(new_point[1], image_miny, image_maxy)
# 创建新点的数据(复制边缘点的所有属性)
new_row = points_gdf.iloc[edge_idx].copy()
new_row['proj_x'] = new_point[0]
new_row['proj_y'] = new_point[1]
# 更新geometry
new_row['geometry'] = Point(new_point[0], new_point[1])
new_points_list.append(new_point)
new_data_list.append(new_row)
# 如果已经到达边界,停止生成
if expand_distance >= distance_to_boundary:
break
# 合并原始点和外扩点
if len(new_data_list) > 0:
# 创建新点的GeoDataFrame
expanded_gdf = gpd.GeoDataFrame(new_data_list, crs=points_gdf.crs)
# 合并原始点和外扩点使用gpd.concat以确保geometry列正确处理
result_gdf = gpd.GeoDataFrame(pd.concat([points_gdf, expanded_gdf], ignore_index=True), crs=points_gdf.crs)
print(f"外扩完成:原始点 {len(points_gdf)} 个,边缘点 {len(edge_indices)} 个,"
f"新增外扩点 {len(new_data_list)} 个(间距 {resolution}),总计 {len(result_gdf)} 个点")
if boundary_gdf is not None:
print(f"水域掩膜范围: X[{mask_minx:.2f}, {mask_maxx:.2f}], Y[{mask_miny:.2f}, {mask_maxy:.2f}]")
print(f"图像范围(含外扩): X[{image_minx:.2f}, {image_maxx:.2f}], Y[{image_miny:.2f}, {image_maxy:.2f}]")
return result_gdf
else:
print("未生成外扩点,返回原始点集")
return points_gdf.copy()
def create_interpolation_grid(self, points_gdf, boundary_gdf=None, resolution=100, expand_ratio=0.05,
use_distance_diffusion=True, max_diffusion_distance=None,
diffusion_power=2, diffusion_n_neighbors=15):
"""
创建插值网格
★★★ Plan Cboundary_gdf 可选None = 纯采样点自然插值,无水域掩膜约束)★★★
Parameters:
-----------
boundary_gdf : gpd.GeoDataFrame, optional
水域掩膜边界GeoDataFrame。None 时基于采样点自身范围插值,不做掩膜裁剪和填充。
expand_ratio : float, default=0.05
边界外扩比例5%),用于从采样点范围外扩出图像边界。
use_distance_diffusion : bool, default=True
是否使用距离扩散方法填充边界空白区域(仅在 boundary_gdf 有值时生效)。
max_diffusion_distance : float, optional
最大扩散距离单位与坐标相同。如果为None自动计算为网格分辨率的5倍。
diffusion_power : float, default=2
距离扩散的IDW幂参数值越大距离衰减越快。
diffusion_n_neighbors : int, default=15
距离扩散使用的最近邻点数。
Returns:
--------
grid_xx, grid_yy, grid_content, bounds : tuple
"""
print("正在创建插值网格...")
# ── Plan C: 无水域掩膜时,基于采样点自身范围计算边界 ─────────
if boundary_gdf is None:
print("[Plan C] 无水域掩膜,基于采样点范围创建插值网格...")
points = np.column_stack((points_gdf['proj_x'], points_gdf['proj_y']))
minx = points[:, 0].min()
maxx = points[:, 0].max()
miny = points[:, 1].min()
maxy = points[:, 1].max()
print(f"采样点范围: X({minx:.6f} - {maxx:.6f}), Y({miny:.6f} - {maxy:.6f})")
else:
bounds = boundary_gdf.total_bounds
minx, miny, maxx, maxy = bounds
print(f"水域掩膜范围: X({minx:.6f} - {maxx:.6f}), Y({miny:.6f} - {maxy:.6f})")
# 计算范围大小并外扩
width = maxx - minx
height = maxy - miny
expand_x = width * expand_ratio
expand_y = height * expand_ratio
minx -= expand_x
maxx += expand_x
miny -= expand_y
maxy += expand_y
print(f"外扩后边界范围: X({minx:.6f} - {maxx:.6f}), Y({miny:.6f} - {maxy:.6f})")
print(f"外扩比例: {expand_ratio * 100:.1f}%")
if self.output_crs == 'EPSG:4326':
print(f"区域尺寸: 宽度={width:.6f}°, 高度={height:.6f}°")
resolution_deg = resolution / 111000.0
print(f"网格分辨率: {resolution}m ≈ {resolution_deg:.6f}°")
else:
print(f"区域尺寸: 宽度={width:.2f}m, 高度={height:.2f}m")
resolution_deg = resolution
# 计算网格点数
min_grid_points = 50
if self.output_crs == 'EPSG:4326':
grid_points_x = max(int(width / resolution_deg), min_grid_points)
grid_points_y = max(int(height / resolution_deg), min_grid_points)
else:
grid_points_x = max(int(width / resolution), min_grid_points)
grid_points_y = max(int(height / resolution), min_grid_points)
grid_points_x = max(grid_points_x, 100)
grid_points_y = max(grid_points_y, 100)
grid_x = np.linspace(minx, maxx, grid_points_x)
grid_y = np.linspace(miny, maxy, grid_points_y)
grid_xx, grid_yy = np.meshgrid(grid_x, grid_y)
total_grid_cells = grid_xx.size
print(f"网格大小: {grid_xx.shape[1]} x {grid_xx.shape[0]} (宽 x 高) = {total_grid_cells:,} 个网格点")
if grid_xx.shape[0] < 2 or grid_xx.shape[1] < 2:
raise ValueError(f"网格尺寸太小 {grid_xx.shape},无法进行插值。")
# 准备插值数据
points = np.column_stack((points_gdf['proj_x'], points_gdf['proj_y']))
values = points_gdf['content'].values
print(f"插值数据点数量: {len(points)}")
print(f"含量值范围: {values.min():.4f} - {values.max():.4f}")
if len(points) < 3:
raise ValueError("插值需要至少3个数据点")
self._check_point_distribution(points)
# 执行插值
print("正在执行空间插值...")
grid_content = self._perform_interpolation(points, values, grid_xx, grid_yy)
# ── Plan C: 无水域掩膜时,跳过所有掩膜裁剪和边缘填充逻辑 ──────
if boundary_gdf is None:
print("[Plan C] 无水域掩膜跳过边缘填充保留插值空白区域NaN")
valid_data = ~np.isnan(grid_content)
valid_count = np.sum(valid_data)
print(f"有效插值点数量: {valid_count} / {grid_content.size}")
if valid_count > 0:
valid_values = grid_content[valid_data]
print(f"插值后数据统计: 最小值={valid_values.min():.4f}, "
f"最大值={valid_values.max():.4f}, 平均值={valid_values.mean():.4f}")
expanded_bounds = np.array([minx, miny, maxx, maxy])
return grid_xx, grid_yy, grid_content, expanded_bounds
# ── 以下为原有水域掩膜逻辑boundary_gdf 有值时执行)────────────
print("正在识别边界区域...")
mask_points = np.column_stack((grid_xx.ravel(), grid_yy.ravel()))
mask_geometry = [Point(x, y) for x, y in mask_points]
mask_gdf = gpd.GeoDataFrame(geometry=mask_geometry, crs=self.output_crs)
within_boundary = mask_gdf.within(boundary_gdf.unary_union)
mask = within_boundary.values.reshape(grid_xx.shape)
print("正在提取边界边缘值并填充边界外区域...")
nan_mask = np.isnan(grid_content)
within_boundary_nan = nan_mask & mask
if np.any(within_boundary_nan):
if use_distance_diffusion:
grid_content = self._fill_boundary_blanks_with_distance_diffusion(
grid_content, grid_xx, grid_yy, mask, boundary_gdf,
max_diffusion_distance=max_diffusion_distance,
power=diffusion_power,
n_neighbors=diffusion_n_neighbors
)
else:
print(f"填充边界内的 {np.sum(within_boundary_nan)} 个NaN点使用最近邻插值...")
valid_mask = ~nan_mask & mask
if np.sum(valid_mask) > 0:
valid_points = np.column_stack((grid_xx[valid_mask], grid_yy[valid_mask]))
valid_values = grid_content[valid_mask]
nan_points = np.column_stack((grid_xx[within_boundary_nan], grid_yy[within_boundary_nan]))
filled_values = griddata(valid_points, valid_values, nan_points, method='nearest')
grid_content[within_boundary_nan] = filled_values
print("边界内填充完成")
boundary_mask_binary = mask.astype(int)
outside_mask = ~mask
kernel = np.ones((3, 3), dtype=bool)
dilated_outside = ndimage.binary_dilation(outside_mask, structure=kernel)
edge_mask = mask & dilated_outside
if np.any(edge_mask):
edge_values = grid_content[edge_mask]
edge_valid = ~np.isnan(edge_values)
if np.any(edge_valid):
edge_mean = np.nanmean(edge_values)
print(f"边界边缘平均值: {edge_mean:.4f}")
outside_nan = outside_mask & np.isnan(grid_content)
if np.any(outside_nan):
edge_points = np.column_stack((grid_xx[edge_mask & ~np.isnan(grid_content)],
grid_yy[edge_mask & ~np.isnan(grid_content)]))
if len(edge_points) > 0:
edge_vals = grid_content[edge_mask & ~np.isnan(grid_content)]
outside_points = np.column_stack((grid_xx[outside_nan], grid_yy[outside_nan]))
outside_filled = griddata(edge_points, edge_vals, outside_points, method='nearest')
grid_content[outside_nan] = outside_filled
print(f"已填充边界外的 {np.sum(~np.isnan(outside_filled))} 个点")
else:
grid_content[outside_nan] = edge_mean
print(f"使用边缘平均值填充边界外的 {np.sum(outside_nan)} 个点")
else:
print("边界外区域已全部填充")
else:
global_mean = np.nanmean(grid_content[mask])
if not np.isnan(global_mean):
grid_content[outside_mask & np.isnan(grid_content)] = global_mean
print(f"使用全局平均值 {global_mean:.4f} 填充边界外")
else:
mean_in_boundary = np.nanmean(grid_content[mask])
if not np.isnan(mean_in_boundary):
grid_content[outside_mask & np.isnan(grid_content)] = mean_in_boundary
print(f"使用边界内平均值 {mean_in_boundary:.4f} 填充边界外")
print("整个画面已铺满,边界外区域已用边缘值填充")
final_check_nan = np.isnan(grid_content) & mask
if np.any(final_check_nan):
print(f"警告: 仍有 {np.sum(final_check_nan)} 个边界内的点未填充...")
if np.sum(~np.isnan(grid_content) & mask) > 0:
mean_value = np.nanmean(grid_content[mask])
grid_content[final_check_nan] = mean_value
else:
global_mean = np.nanmean(grid_content)
grid_content[final_check_nan] = global_mean if not np.isnan(global_mean) else 0
valid_data = ~np.isnan(grid_content)
valid_count = np.sum(valid_data)
print(f"有效插值点数量: {valid_count} / {grid_content.size}")
if valid_count == 0:
raise ValueError("边界掩膜后没有有效数据点,请检查数据点是否在边界范围内")
if valid_count < 4:
print("警告:有效数据点很少,可能影响绘图效果")
valid_values = grid_content[valid_data]
print(f"插值后数据统计: 最小值={valid_values.min():.4f}, "
f"最大值={valid_values.max():.4f}, 平均值={valid_values.mean():.4f}")
expanded_bounds = np.array([minx, miny, maxx, maxy])
return grid_xx, grid_yy, grid_content, expanded_bounds
def create_content_map(self, points_gdf, boundary_gdf=None, grid_xx=None, grid_yy=None,
grid_content=None, bounds=None, output_file='content_map.png',
show_sample_points=False, base_map_tif=None,
cmap='viridis'):
"""
创建含量图
★★★ Plan Cboundary_gdf 可选None = 无掩膜裁剪,无黑色边界线)★★★
Parameters:
-----------
base_map_tif : str, optional
TIF正射底图文件路径。如果提供将在水域掩膜外显示底图
cmap : str, default='viridis'
含量数据的颜色映射
"""
print("正在生成含量图...")
# 检查网格数据
print(f"网格形状: {grid_content.shape}")
# 创建边界掩膜(用于绘图时只显示边界内)
print("创建边界掩膜用于绘图...")
# ── Plan C: 无水域掩膜时显示全部插值区域NaN 区域本身就不填充)─────────
if boundary_gdf is None:
print("[Plan C] 无水域掩膜,显示全部插值区域,不做边界裁剪")
mask = np.ones_like(grid_content, dtype=bool)
else:
try:
grid_points = gpd.GeoDataFrame(
geometry=[Point(x, y) for x, y in zip(grid_xx.flatten(), grid_yy.flatten())],
crs=points_gdf.crs
)
within_boundary = grid_points.within(boundary_gdf.unary_union)
mask = within_boundary.values.reshape(grid_xx.shape)
print(f"边界内点数: {np.sum(mask)} / {mask.size}")
except Exception as e:
print(f"创建边界掩膜时出现错误: {e},继续绘图...")
mask = np.ones_like(grid_content, dtype=bool) # 如果失败,显示全部
valid_data = ~np.isnan(grid_content)
if np.sum(valid_data) == 0:
raise ValueError("没有有效的插值数据用于绘图")
# 计算数据统计
valid_values = grid_content[valid_data]
print(
f"插值结果统计: 最小值={valid_values.min():.4f}, 最大值={valid_values.max():.4f}, 平均值={valid_values.mean():.4f}")
print(f"有效数据点数量: {np.sum(valid_data)} / {grid_content.size}")
# 检查数据范围
data_range = valid_values.max() - valid_values.min()
print(f"数据范围: {data_range:.6f}")
if data_range == 0:
print("警告:所有数据值都相同,将使用单一颜色显示")
# 创建图形
fig, ax = plt.subplots(figsize=(12, 10))
# 如果提供了底图,先绘制底图(在水域掩膜外)
if base_map_tif is not None:
try:
print(f"正在加载底图: {base_map_tif}")
self._add_base_map(ax, base_map_tif, bounds, mask, grid_xx, grid_yy, boundary_gdf)
print("底图加载成功")
except Exception as e:
print(f"加载底图失败: {e},将跳过底图显示")
# 设置颜色映射参数
im = None
try:
if data_range > 0:
# 设置颜色范围,确保有足够的对比度
vmin = valid_values.min()
vmax = valid_values.max()
# 如果范围很小,稍微扩展一下以增加对比度
if data_range < 1e-6:
center = valid_values.mean()
expansion = max(abs(center) * 0.01, 1e-6) # 扩展1%或最小值
vmin = center - expansion
vmax = center + expansion
print(f"颜色映射范围: {vmin:.6f} - {vmax:.6f}")
# 方法1尝试使用contourf
try:
print("尝试使用contourf绘制...")
# 使用掩膜数组:边界外的数据被掩膜掉,只显示边界内
# mask已经在前面创建好了
masked_data = np.ma.masked_where(~mask, grid_content)
# 创建更多等级数以获得更平滑的颜色过渡
levels = np.linspace(vmin, vmax, 100) # 创建100个等级以获得平滑效果
im = ax.contourf(grid_xx, grid_yy, masked_data,
levels=levels, cmap=cmap, alpha=0.9,
vmin=vmin, vmax=vmax, extend='both')
print("contourf绘制成功")
# 可选择性添加等值线(默认不添加,以保持平滑效果)
# 如果需要等值线,可以取消注释下面的代码
# try:
# contour_levels = np.linspace(vmin, vmax, 11)
# contours = ax.contour(grid_xx, grid_yy, grid_content,
# levels=contour_levels, colors='white',
# alpha=0.3, linewidths=0.5)
# ax.clabel(contours, inline=True, fontsize=8, fmt='%.3f')
# print("等值线添加成功")
# except Exception as e:
# print(f"等值线绘制失败: {e}")
except Exception as e:
print(f"contourf失败: {e}")
# 方法2使用pcolormesh
try:
print("尝试使用pcolormesh绘制...")
# 使用掩膜数组:边界外的数据被掩膜掉,只显示边界内
# mask已经在前面创建好了
masked_data = np.ma.masked_where(~mask, grid_content)
im = ax.pcolormesh(grid_xx, grid_yy, masked_data,
cmap=cmap, alpha=0.9,
vmin=vmin, vmax=vmax, shading='gouraud') # 使用gouraud平滑着色
print("pcolormesh绘制成功")
except Exception as e2:
print(f"pcolormesh也失败: {e2}")
raise e2
else:
# 所有值相同的情况
print("使用单一颜色填充(所有值相同)")
# 创建一个简单的填充
single_value = valid_values[0]
im = ax.contourf(grid_xx, grid_yy, grid_content,
levels=[single_value - 0.001, single_value + 0.001],
cmap=cmap, alpha=0.8)
except Exception as e:
print(f"主要绘图方法失败,尝试备选方案: {e}")
# 备选方案1imshow
try:
print("尝试使用imshow...")
# 处理NaN值
display_data = grid_content.copy()
nan_mask = np.isnan(display_data)
if np.any(nan_mask):
# 用平均值填充NaN
display_data[nan_mask] = valid_values.mean()
if data_range > 0:
vmin = valid_values.min()
vmax = valid_values.max()
im = ax.imshow(display_data,
extent=[grid_xx.min(), grid_xx.max(),
grid_yy.min(), grid_yy.max()],
cmap=cmap, alpha=0.8, origin='lower',
vmin=vmin, vmax=vmax, aspect='auto')
else:
im = ax.imshow(display_data,
extent=[grid_xx.min(), grid_xx.max(),
grid_yy.min(), grid_yy.max()],
cmap=cmap, alpha=0.8, origin='lower',
aspect='auto')
print("imshow绘制成功")
except Exception as e2:
print(f"imshow也失败: {e2}")
# 备选方案2散点图
try:
print("尝试使用散点图...")
valid_x = grid_xx[valid_data]
valid_y = grid_yy[valid_data]
valid_z = grid_content[valid_data]
if data_range > 0:
im = ax.scatter(valid_x, valid_y, c=valid_z,
cmap=cmap, alpha=0.8, s=10,
vmin=valid_values.min(), vmax=valid_values.max())
else:
im = ax.scatter(valid_x, valid_y, c=valid_z,
cmap=cmap, alpha=0.8, s=10)
print("散点图绘制成功")
except Exception as e3:
print(f"所有绘图方法都失败: {e3}")
raise ValueError("无法生成颜色图,请检查数据")
# 绘制边界(黑色)—— Plan C: 无掩膜时不绘制边界线
if boundary_gdf is not None:
try:
boundary_gdf.boundary.plot(ax=ax, color='black', linewidth=2, alpha=1.0)
print("边界绘制成功(黑色)")
except Exception as e:
print(f"边界绘制失败: {e}")
# 可选择性绘制采样点(默认不绘制,以显示平滑的颜色分布)
if show_sample_points:
try:
points_gdf.plot(ax=ax, color='black', markersize=6, alpha=0.7,
marker='+', edgecolors='white', linewidth=1)
print("采样点绘制成功")
except Exception as e:
print(f"采样点绘制失败: {e}")
# 设置坐标轴标签和格式
# 由于输入是投影坐标系,输出是地理坐标系,始终显示为地理坐标
ax.set_xlabel('经度 (°)', fontsize=12)
ax.set_ylabel('纬度 (°)', fontsize=12)
# 格式化坐标轴刻度为经纬度格式保留3位小数
def lon_formatter(x, p):
return f'{x:.3f}°'
def lat_formatter(x, p):
return f'{x:.3f}°'
ax.xaxis.set_major_formatter(FuncFormatter(lon_formatter))
ax.yaxis.set_major_formatter(FuncFormatter(lat_formatter))
# 添加格网线
ax.grid(True, linestyle='--', linewidth=0.5, alpha=0.5, color='gray')
ax.set_axisbelow(True) # 将格网线放在图层下方
# ax.set_title('含量分布图', fontsize=16, fontweight='bold', pad=20) # 已去除标题
# 添加颜色条
try:
if im is not None:
cbar = plt.colorbar(im, ax=ax, shrink=0.5, aspect=40, pad=0.02)
cbar.set_label('含量值', fontsize=10)
# 设置颜色条的刻度
if data_range > 0:
tick_values = np.linspace(valid_values.min(), valid_values.max(), 6)
cbar.set_ticks(tick_values)
cbar.set_ticklabels([f'{val:.3f}' for val in tick_values])
cbar.ax.tick_params(labelsize=8) # 缩小刻度标签字体
print("颜色条添加成功")
else:
print("警告无法添加颜色条im对象为None")
except Exception as e:
print(f"颜色条添加失败: {e}")
# 添加指北针
try:
self.add_north_arrow(ax, bounds)
except Exception as e:
print(f"指北针添加失败: {e}")
# 添加比例尺
try:
self.add_scale_bar(ax)
except Exception as e:
print(f"比例尺添加失败: {e}")
# 添加图例
try:
self.add_legend(ax)
except Exception as e:
print(f"图例添加失败: {e}")
# 设置图形边界进一步外扩1%以确保边界不完全挨着地图)
try:
x_range = bounds[2] - bounds[0]
y_range = bounds[3] - bounds[1]
display_expand = 0.01 # 显示时再外扩1%
ax.set_xlim(bounds[0] - x_range * display_expand, bounds[2] + x_range * display_expand)
ax.set_ylim(bounds[1] - y_range * display_expand, bounds[3] + y_range * display_expand)
except Exception as e:
print(f"设置图形边界失败: {e}")
# 调整布局
plt.tight_layout()
# 保存图片
try:
plt.savefig(output_file, dpi=300, bbox_inches='tight',
facecolor='white', edgecolor='none')
print(f"含量图已保存为:{output_file}")
except Exception as e:
print(f"图片保存失败: {e}")
# 显示图片
try:
plt.show()
except Exception as e:
print(f"图片显示失败: {e}")
def add_north_arrow(self, ax, bounds):
"""添加指北针(右上角)- 画布相对坐标,不依赖数据坐标系。
使用 ax.transAxes 将指北针固定在右上角,
尺寸以点数points为单位与数据坐标系解耦
无论 UTM 坐标范围多大,指北针始终保持合理大小。
"""
# ★★★ 改用画布相对坐标transAxes★★★
# (0.88, 0.92) = 右上角,尺寸用 points72分之一英寸
arrow_ax_x, arrow_ax_y = 0.88, 0.92
radius_pt = 18 # 罗盘半径(磅),由 28 → 18 缩小图元
# 统一在数据坐标系下绘制transform=ax.transData
# 但 position 由 axes 坐标决定radius 用固定点数
# 将 axes 坐标转为数据坐标:取右上角 + 偏移
xlim = ax.get_xlim()
ylim = ax.get_ylim()
# 偏移系数 0.08 → 0.05 让指北针更靠中心,避免紧贴角落被裁切
dx = (xlim[1] - xlim[0]) * 0.05
dy = (ylim[1] - ylim[0]) * 0.05
arrow_x = xlim[1] - dx
arrow_y = ylim[1] - dy
# radius 系数 0.6 → 0.42 缩小指北针整体半径
radius = min(dx, dy) * 0.42
# 绘制圆形背景(外圈)
circle_outer = patches.Circle(
(arrow_x, arrow_y),
radius=radius,
facecolor='white',
edgecolor='black',
linewidth=1.5,
zorder=10,
transform=ax.transData,
)
ax.add_patch(circle_outer)
# 绘制内圈(装饰)
circle_inner = patches.Circle(
(arrow_x, arrow_y),
radius=radius * 0.7,
facecolor='none',
edgecolor='gray',
linewidth=0.8,
linestyle='--',
zorder=11,
transform=ax.transData,
)
ax.add_patch(circle_inner)
# 绘制四个方向的刻度线
tick_width = 1.0
# 北方向刻度(主刻度)
ax.plot([arrow_x, arrow_x], [arrow_y, arrow_y + radius * 0.85],
'k-', linewidth=tick_width * 2, zorder=12)
# 南方向刻度
ax.plot([arrow_x, arrow_x], [arrow_y, arrow_y - radius * 0.85],
'k-', linewidth=tick_width, zorder=12)
# 东方向刻度
ax.plot([arrow_x, arrow_x + radius * 0.85], [arrow_y, arrow_y],
'k-', linewidth=tick_width, zorder=12)
# 西方向刻度
ax.plot([arrow_x, arrow_x - radius * 0.85], [arrow_y, arrow_y],
'k-', linewidth=tick_width, zorder=12)
# 绘制次要刻度45度方向
for angle in [45, 135, 225, 315]:
angle_rad = math.radians(angle)
x_end = arrow_x + radius * 0.7 * math.cos(angle_rad)
y_end = arrow_y + radius * 0.7 * math.sin(angle_rad)
ax.plot([arrow_x, x_end], [arrow_y, y_end],
'k-', linewidth=tick_width * 0.5, alpha=0.6, zorder=12)
# 绘制指北箭头(三角形,填充)
arrow_size = radius * 0.6
arrow_points = np.array([
[arrow_x, arrow_y + radius * 0.9], # 顶点(北)
[arrow_x - arrow_size * 0.3, arrow_y + radius * 0.3], # 左下
[arrow_x + arrow_size * 0.3, arrow_y + radius * 0.3] # 右下
])
arrow_poly = patches.Polygon(
arrow_points,
facecolor='black',
edgecolor='black',
linewidth=1.2,
zorder=13,
transform=ax.transData,
)
ax.add_patch(arrow_poly)
# 绘制指南箭头(三角形,填充,但较小)
south_arrow_size = radius * 0.4
south_arrow_points = np.array([
[arrow_x, arrow_y - radius * 0.6], # 顶点(南)
[arrow_x - south_arrow_size * 0.2, arrow_y - radius * 0.2], # 左上
[arrow_x + south_arrow_size * 0.2, arrow_y - radius * 0.2] # 右上
])
south_arrow_poly = patches.Polygon(
south_arrow_points,
facecolor='white',
edgecolor='black',
linewidth=1.0,
zorder=13,
transform=ax.transData,
)
ax.add_patch(south_arrow_poly)
# 添加方向标记N, S, E, W
label_offset = radius * 1.1
# 字号 9 → 7与缩小的指北针半径相匹配
font_size = 7
ax.text(arrow_x, arrow_y + label_offset, 'N',
fontsize=font_size, fontweight='bold', ha='center', va='bottom',
color='black', zorder=14)
ax.text(arrow_x, arrow_y - label_offset, 'S',
fontsize=font_size, fontweight='bold', ha='center', va='top',
color='black', zorder=14)
ax.text(arrow_x + label_offset, arrow_y, 'E',
fontsize=font_size, fontweight='bold', ha='left', va='center',
color='black', zorder=14)
ax.text(arrow_x - label_offset, arrow_y, 'W',
fontsize=font_size, fontweight='bold', ha='right', va='center',
color='black', zorder=14)
def add_scale_bar(self, ax, scale_x=None, scale_y=None):
"""添加比例尺
Parameters
----------
ax : matplotlib Axes
绘图坐标轴
scale_x : float, optional
X 方向像素分辨率(米),由 visualize_raster 从 src.res 传入。
若传入则直接作为 ScaleBar 的 scale 值,忽略 self.output_crs 判断。
scale_y : float, optional
Y 方向像素分辨率(米),同 scale_x。
"""
try:
if scale_x is not None and scale_y is not None:
# visualize_raster 传入真实像素分辨率,直接用米为单位
scalebar = ScaleBar(
scale_x,
units='m',
location='lower left',
box_alpha=0.8,
color='black',
font_properties={'size': 8},
label_loc='bottom',
)
ax.add_artist(scalebar)
print(f"比例尺添加成功(像素分辨率: {scale_x:.4f} m")
elif self.output_crs == 'EPSG:4326':
scalebar = ScaleBar(
111000, # 1度 = 111000米
units='m',
location='lower left',
box_alpha=0.8,
color='black',
font_properties={'size': 8},
label_loc='bottom'
)
ax.add_artist(scalebar)
print("地理坐标系比例尺添加成功")
else:
scalebar = ScaleBar(1, units='m', location='lower left',
box_alpha=0.8, color='black',
font_properties={'size': 8})
ax.add_artist(scalebar)
print("投影坐标系比例尺添加成功")
except Exception as e:
print(f"比例尺添加失败: {e}")
# 如果matplotlib-scalebar失败尝试手动添加简单的比例尺
try:
self._add_manual_scale_bar(ax)
print("手动比例尺添加成功")
except Exception as e2:
print(f"手动比例尺也失败: {e2}")
def _add_manual_scale_bar(self, ax):
"""手动添加简单的比例尺"""
# 获取当前坐标轴的范围
xlim = ax.get_xlim()
ylim = ax.get_ylim()
# 计算比例尺的位置和长度
x_range = xlim[1] - xlim[0]
y_range = ylim[1] - ylim[0]
# 比例尺位置(左下角)
scale_x = xlim[0] + x_range * 0.05
scale_y = ylim[0] + y_range * 0.1
if self.output_crs == 'EPSG:4326':
# 地理坐标系:计算合适的比例尺长度(度)
# 选择一个合理的距离比如1公里、5公里或10公里
distance_km = 5 # 5公里
scale_length_deg = distance_km / 111.0 # 转换为度数
# 绘制比例尺线
ax.plot([scale_x, scale_x + scale_length_deg], [scale_y, scale_y],
'k-', linewidth=2)
ax.plot([scale_x, scale_x], [scale_y - y_range * 0.01, scale_y + y_range * 0.01],
'k-', linewidth=1.5)
ax.plot([scale_x + scale_length_deg, scale_x + scale_length_deg],
[scale_y - y_range * 0.01, scale_y + y_range * 0.01], 'k-', linewidth=1.5)
# 添加文字标注
ax.text(scale_x + scale_length_deg / 2, scale_y + y_range * 0.02,
f'{distance_km} km', ha='center', va='bottom', fontsize=8,
bbox=dict(boxstyle='round,pad=0.3', facecolor='white', alpha=0.8))
else:
# 投影坐标系:使用米为单位
# 选择合适的比例尺长度
if x_range > 10000: # 大于10km
scale_length = 5000 # 5km
scale_text = '5 km'
elif x_range > 2000: # 大于2km
scale_length = 1000 # 1km
scale_text = '1 km'
else: # 小于2km
scale_length = 500 # 500m
scale_text = '500 m'
# 绘制比例尺线
ax.plot([scale_x, scale_x + scale_length], [scale_y, scale_y],
'k-', linewidth=2)
ax.plot([scale_x, scale_x], [scale_y - y_range * 0.01, scale_y + y_range * 0.01],
'k-', linewidth=1.5)
ax.plot([scale_x + scale_length, scale_x + scale_length],
[scale_y - y_range * 0.01, scale_y + y_range * 0.01], 'k-', linewidth=1.5)
# 添加文字标注
ax.text(scale_x + scale_length / 2, scale_y + y_range * 0.02,
scale_text, ha='center', va='bottom', fontsize=8,
bbox=dict(boxstyle='round,pad=0.3', facecolor='white', alpha=0.8))
def _add_base_map(self, ax, base_map_tif, bounds, mask, grid_xx, grid_yy, boundary_gdf):
"""添加正射底图(在水域掩膜外显示)
Parameters:
-----------
ax : matplotlib.axes.Axes
绘图轴对象
base_map_tif : str
TIF底图文件路径
bounds : np.ndarray
显示范围 [minx, miny, maxx, maxy]
mask : np.ndarray
水域掩膜True表示水域内
grid_xx : np.ndarray
网格X坐标
grid_yy : np.ndarray
网格Y坐标
boundary_gdf : gpd.GeoDataFrame
边界几何数据
"""
print("正在读取底图文件...")
with rasterio.open(base_map_tif) as src:
# 获取底图的坐标系
tif_crs = src.crs
tif_bounds = src.bounds
print(f"底图坐标系: {tif_crs}")
print(f"底图范围: {tif_bounds}")
print(f"目标范围: {bounds}")
# 检查是否需要投影转换
target_crs = CRS.from_string(self.output_crs)
need_reproject = tif_crs != target_crs
# 读取底图数据
if need_reproject:
print(f"底图坐标系({tif_crs})与目标坐标系({target_crs})不同,正在转换...")
# 计算转换后的变换参数和尺寸
transform, width, height = calculate_default_transform(
tif_crs, target_crs,
src.width, src.height,
left=bounds[0], bottom=bounds[1],
right=bounds[2], top=bounds[3]
)
# 创建目标数组
if src.count == 1:
# 单波段
base_map_data = np.zeros((height, width), dtype=src.dtypes[0])
reproject(
source=src.read(1),
destination=base_map_data,
src_transform=src.transform,
src_crs=tif_crs,
dst_transform=transform,
dst_crs=target_crs,
resampling=Resampling.bilinear
)
else:
# 多波段RGB取前3个波段
num_bands = min(3, src.count)
base_map_data = np.zeros((num_bands, height, width), dtype=src.dtypes[0])
for i in range(num_bands):
reproject(
source=src.read(i + 1),
destination=base_map_data[i],
src_transform=src.transform,
src_crs=tif_crs,
dst_transform=transform,
dst_crs=target_crs,
resampling=Resampling.bilinear
)
# 如果是RGB转换为(height, width, 3)格式
if num_bands == 3:
base_map_data = np.transpose(base_map_data, (1, 2, 0))
# 创建extent用于显示
extent = [bounds[0], bounds[2], bounds[1], bounds[3]]
else:
# 不需要投影转换,直接读取对应范围的数据
print("底图坐标系与目标坐标系一致,直接读取...")
# 计算需要读取的窗口
row_min, col_min = src.index(bounds[0], bounds[3]) # 左上角
row_max, col_max = src.index(bounds[2], bounds[1]) # 右下角
# 确保索引在有效范围内
row_min = max(0, row_min)
row_max = min(src.height, row_max + 1)
col_min = max(0, col_min)
col_max = min(src.width, col_max + 1)
window = windows.Window.from_slices(
(row_min, row_max), (col_min, col_max)
)
# 读取数据
if src.count == 1:
base_map_data = src.read(1, window=window)
else:
# 多波段取前3个波段
num_bands = min(3, src.count)
base_map_data = src.read(list(range(1, num_bands + 1)), window=window)
if num_bands == 3:
# 转换为(height, width, 3)格式
base_map_data = np.transpose(base_map_data, (1, 2, 0))
# 计算extent
window_transform = windows.transform(window, src.transform)
left = window_transform[2]
top = window_transform[5]
right = left + window_transform[0] * base_map_data.shape[1]
bottom = top + window_transform[4] * base_map_data.shape[0]
# 确保extent不超过bounds
extent = [
max(bounds[0], left),
min(bounds[2], right),
max(bounds[1], bottom),
min(bounds[3], top)
]
# 将底图数据缩放到网格大小以便显示
# 创建底图的显示掩膜:只在边界外显示
print("正在创建底图显示掩膜...")
# 创建底图网格(与显示范围对齐)
base_map_height, base_map_width = base_map_data.shape[:2]
# 性能优化:如果底图分辨率过高,进行降采样以提高处理速度
# 限制最大边长为2000像素保持足够清晰度的同时提高速度
max_display_size = 2000
scale_factor = 1.0
if max(base_map_height, base_map_width) > max_display_size:
scale_factor = max_display_size / max(base_map_height, base_map_width)
new_height = int(base_map_height * scale_factor)
new_width = int(base_map_width * scale_factor)
print(
f"底图分辨率较高 ({base_map_width}x{base_map_height}),降采样到 {new_width}x{new_height} 以提高速度")
# 使用scipy的zoom进行降采样
if base_map_data.ndim == 2:
base_map_data = ndimage.zoom(base_map_data, scale_factor, order=1)
else:
base_map_data = ndimage.zoom(base_map_data, (scale_factor, scale_factor, 1), order=1)
base_map_height, base_map_width = base_map_data.shape[:2]
# 更新extent以匹配新的分辨率
extent_width = extent[1] - extent[0]
extent_height = extent[3] - extent[2]
extent = [
extent[0],
extent[0] + extent_width,
extent[2],
extent[2] + extent_height
]
# 使用rasterio的geometry_mask快速生成掩膜比创建大量Point对象快得多
# 创建底图的变换矩阵
if need_reproject:
# 如果进行了投影转换使用计算得到的transform
base_map_transform = transform
else:
# 如果没有投影转换,使用窗口变换
base_map_transform = window_transform
# 如果进行了降采样需要调整transform
if scale_factor < 1.0:
# 调整transform以适应新的分辨率
# rasterio的transform是6元素tuple或Affine对象需要调整像素大小
# 获取transform的6个参数 (a, b, c, d, e, f)
# 其中a和e是像素大小需要除以scale_factor
if Affine is not None:
# 获取6个参数
if hasattr(base_map_transform, '__iter__') and len(base_map_transform) == 6:
a, b, c, d, e, f = base_map_transform
else:
a, b, c, d, e, f = base_map_transform[0], base_map_transform[1], base_map_transform[2], \
base_map_transform[3], base_map_transform[4], base_map_transform[5]
# 创建新的transform调整像素大小a和e是像素大小
base_map_transform = Affine(a / scale_factor, b, c, d, e / scale_factor, f)
else:
# 降级方案使用tuple
if hasattr(base_map_transform, '__iter__') and len(base_map_transform) == 6:
a, b, c, d, e, f = base_map_transform
else:
a, b, c, d, e, f = base_map_transform[0], base_map_transform[1], base_map_transform[2], \
base_map_transform[3], base_map_transform[4], base_map_transform[5]
base_map_transform = (a / scale_factor, b, c, d, e / scale_factor, f)
# 调试信息:检查边界数据和底图范围
print(f"底图显示范围 (extent): {extent}")
print(f"底图分辨率: {base_map_width}x{base_map_height}")
print(f"底图transform: {base_map_transform}")
if boundary_gdf is not None and len(boundary_gdf) > 0:
boundary_bounds = boundary_gdf.total_bounds
print(f"边界数据范围: {boundary_bounds}")
print(f"边界数据坐标系: {boundary_gdf.crs}")
print(f"边界要素数量: {len(boundary_gdf)}")
# 检查边界是否与底图范围重叠
overlap_x = not (boundary_bounds[2] < extent[0] or boundary_bounds[0] > extent[1])
overlap_y = not (boundary_bounds[3] < extent[2] or boundary_bounds[1] > extent[3])
if not (overlap_x and overlap_y):
print("警告: 边界数据范围与底图显示范围不重叠!")
print(" 将不应用掩膜,显示整个底图")
# 创建全True的掩膜显示所有区域
base_map_mask = np.ones((base_map_height, base_map_width), dtype=bool)
else:
# 使用geometry_mask生成掩膜True表示在几何体内即水域内
# 注意geometry_mask返回True表示需要掩膜的区域在几何体内
# 但我们想要的是边界外的区域(不在几何体内),所以需要反转
try:
within_boundary_mask = geometry_mask(
boundary_gdf.geometry,
out_shape=(base_map_height, base_map_width),
transform=base_map_transform,
invert=False # False表示掩膜几何体内的区域水域内
)
# 反转掩膜True表示边界外需要显示的区域
base_map_mask = ~within_boundary_mask
except Exception as e:
print(f"生成掩膜时出错: {e}")
print(" 将不应用掩膜,显示整个底图")
import traceback
traceback.print_exc()
# 创建全True的掩膜显示所有区域
base_map_mask = np.ones((base_map_height, base_map_width), dtype=bool)
else:
print("警告: 边界数据为空,将不应用掩膜,显示整个底图")
# 创建全True的掩膜显示所有区域
base_map_mask = np.ones((base_map_height, base_map_width), dtype=bool)
# 调试信息:检查掩膜状态
mask_ratio = np.sum(base_map_mask) / base_map_mask.size
print(
f"底图掩膜状态: 可显示区域占比 {mask_ratio * 100:.2f}% ({np.sum(base_map_mask)}/{base_map_mask.size} 像素)")
# 如果掩膜后没有可显示区域,警告并显示整个底图
if mask_ratio == 0.0:
print("警告: 掩膜后没有可显示区域,将显示整个底图(不应用掩膜)")
base_map_mask = np.ones((base_map_height, base_map_width), dtype=bool)
# 归一化数据以便显示(如果是数值型)
# 注意:先归一化整个数据,再应用掩膜,这样可以保证归一化范围正确
if base_map_data.dtype != np.uint8:
if base_map_data.ndim == 2:
# 单波段归一化到0-1
# 使用整个数据集的范围进行归一化(不仅仅是掩膜区域)
data_min = np.nanmin(base_map_data)
data_max = np.nanmax(base_map_data)
print(f"底图数据范围: [{data_min}, {data_max}], dtype: {base_map_data.dtype}")
if data_max > data_min:
# 先归一化整个数组
base_map_normalized = (base_map_data - data_min) / (data_max - data_min)
# 然后应用掩膜:只显示边界外的区域
base_map_display = np.ma.masked_where(~base_map_mask, base_map_normalized)
else:
# 如果数据范围无效创建全0的掩膜数组
print("警告: 底图数据范围无效,所有值相同")
base_map_display = np.ma.masked_where(~base_map_mask, np.zeros_like(base_map_data))
else:
# RGB每个波段单独归一化
base_map_normalized = base_map_data.copy().astype(np.float32)
for i in range(base_map_data.shape[2]):
band_data = base_map_data[:, :, i]
data_min = np.nanmin(band_data)
data_max = np.nanmax(band_data)
print(f"底图波段 {i} 数据范围: [{data_min}, {data_max}]")
if data_max > data_min:
# 归一化整个波段
base_map_normalized[:, :, i] = (band_data - data_min) / (data_max - data_min)
else:
print(f"警告: 底图波段 {i} 数据范围无效,所有值相同")
base_map_normalized[:, :, i] = np.zeros_like(band_data)
# 应用掩膜:只显示边界外的区域
mask_3d = np.broadcast_to(~base_map_mask[..., np.newaxis], base_map_data.shape)
base_map_display = np.ma.masked_where(mask_3d, base_map_normalized)
else:
# uint8类型直接使用但可能需要归一化到0-1用于imshow
if base_map_data.ndim == 2:
# 单波段uint8通常已经是0-255范围归一化到0-1
base_map_normalized = base_map_data.astype(np.float32) / 255.0
base_map_display = np.ma.masked_where(~base_map_mask, base_map_normalized)
else:
# RGBuint8归一化到0-1
base_map_normalized = base_map_data.astype(np.float32) / 255.0
mask_3d = np.broadcast_to(~base_map_mask[..., np.newaxis], base_map_data.shape)
base_map_display = np.ma.masked_where(mask_3d, base_map_normalized)
# 检查归一化后的数据范围
if isinstance(base_map_display, np.ma.MaskedArray):
valid_data = base_map_display[~base_map_display.mask]
if len(valid_data) > 0:
print(
f"归一化后有效数据范围: [{np.nanmin(valid_data):.3f}, {np.nanmax(valid_data):.3f}], 有效像素数: {len(valid_data)}")
else:
print("警告: 归一化后没有有效数据显示区域")
# 绘制底图
print("正在绘制底图...")
# 注意extent格式为 [left, right, bottom, top]
# 对于地理坐标系y轴通常向上为正所以使用origin='lower'
try:
if base_map_data.ndim == 2:
# 单波段:使用灰度图
# 确保数据在0-1范围内
if isinstance(base_map_display, np.ma.MaskedArray):
# 对于masked array确保数据范围正确
if np.ma.max(base_map_display) > 1.0 or np.ma.min(base_map_display) < 0.0:
base_map_display = np.ma.clip(base_map_display, 0.0, 1.0)
else:
base_map_display = np.clip(base_map_display, 0.0, 1.0)
im = ax.imshow(base_map_display, extent=extent, origin='lower',
cmap='gray', alpha=0.8, zorder=0, interpolation='bilinear',
vmin=0.0, vmax=1.0)
else:
# RGB直接显示
# 确保数据格式正确需要在0-1范围内
if isinstance(base_map_display, np.ma.MaskedArray):
# 对于masked array确保数据在0-1范围内
if np.ma.max(base_map_display) > 1.0 or np.ma.min(base_map_display) < 0.0:
base_map_display = np.ma.clip(base_map_display, 0.0, 1.0)
else:
base_map_display = np.clip(base_map_display, 0.0, 1.0)
# 确保是float32类型imshow期望0-1范围的float数组
if base_map_display.dtype != np.float32 and base_map_display.dtype != np.float64:
base_map_display = base_map_display.astype(np.float32)
im = ax.imshow(base_map_display, extent=extent, origin='lower',
alpha=0.8, zorder=0, interpolation='bilinear')
print(f"底图绘制成功")
except Exception as e:
print(f"底图绘制出错: {e}")
import traceback
traceback.print_exc()
# 如果绘制失败,至少尝试绘制一个简单的占位图
print("尝试使用备用方法绘制底图...")
try:
if base_map_data.ndim == 2:
# 使用简单的numpy数组不应用掩膜
simple_display = np.clip(base_map_data.astype(np.float32) / np.nanmax(base_map_data), 0, 1)
ax.imshow(simple_display, extent=extent, origin='lower',
cmap='gray', alpha=0.5, zorder=0)
else:
simple_display = np.clip(base_map_data.astype(np.float32) / 255.0, 0, 1)
ax.imshow(simple_display, extent=extent, origin='lower',
alpha=0.5, zorder=0)
print("备用方法绘制成功")
except Exception as e2:
print(f"备用方法也失败: {e2}")
print(f"底图已绘制,显示范围: {extent}")
def add_legend(self, ax):
"""
添加图例
Parameters:
-----------
"""
legend_elements = [
# 移除边界标签
# plt.Line2D([0], [0], color='red', linewidth=2, label='边界'),
# 移除采样点和等值线图例项,以突出平滑的颜色分布效果
# plt.Line2D([0], [0], marker='+', color='w', markerfacecolor='black',
# markersize=8, label='采样点'),
]
# 如果图例为空,则不显示图例
if legend_elements:
ax.legend(handles=legend_elements, loc='upper left',
framealpha=0.9, fontsize=10)
# ------------------------------------------------------------------
# Step 14 适配:水色指数 GeoTIFF 可视化(绕过 CSV 插值)
# ------------------------------------------------------------------
def visualize_raster(
self,
raster_tif_path: str,
output_file: Optional[str] = None,
boundary_shp_path: Optional[str] = None,
cmap: Optional[str] = None,
nodata_value: float = -9999.0,
show_colorbar: bool = True,
figsize: Tuple[int, int] = (12, 10),
title: Optional[str] = None,
alpha: float = 0.9,
) -> str:
"""直接读取 GeoTIFF 栅格数据,生成水质指数专题图。
适用场景:
- WaterIndexProcessor 输出的水色指数 GeoTIFF
- Step 14 接收 GeoTIFF 路径后直接可视化(不通过 CSV 插值)
Parameters
----------
raster_tif_path : str
水色指数 GeoTIFF 文件路径(由 WaterIndexProcessor 输出)
output_file : str, optional
输出图片路径None → 自动从 GeoTIFF 文件名派生)
boundary_shp_path : str, optional
边界 shapefile 路径None → 纯栅格显示,无水域掩膜裁切)
cmap : str, optional
颜色映射None → 自动从 GeoTIFF 描述或文件名推断)
nodata_value : float
NoData 标记值GeoTIFF 中存储的无效值)
show_colorbar : bool
是否显示颜色条
figsize : tuple
图形尺寸(英寸)
title : str, optional
图形标题None → 从 GeoTIFF 描述推断或使用文件名)
alpha : float
透明度0-1
Returns
-------
str
输出图片路径
"""
# ── 始终从路径提取 stem供后续中文标题和文件派生使用──────────
stem = Path(raster_tif_path).stem
# ── 输出路径自动派生(中文文件名)──────────────────────────────
if output_file is None:
chinese_title = self._get_chinese_title(stem)
out_dir = Path(raster_tif_path).parent / 'visualization'
out_dir.mkdir(parents=True, exist_ok=True)
output_file = str(out_dir / f"{chinese_title}_专题图.png")
# ── 读取 GeoTIFF优先 rasterio备选 GDAL──────────────────
tif_path = Path(raster_tif_path)
if not tif_path.is_file():
raise FileNotFoundError(f"GeoTIFF 文件不存在: {raster_tif_path}")
array: Optional[np.ndarray] = None
transform: Optional[Any] = None
crs_obj: Optional[Any] = None
nodata_read: Optional[float] = None
desc: str = ""
# 方式1rasterio推荐坐标系信息更完整
_src_bounds = None # rasterio 原生边界(优先用于 extent
_src_res = None # rasterio 像素分辨率 (xres, yres)
try:
with rasterio.open(raster_tif_path) as src:
array = src.read(1).astype(np.float64)
transform = src.transform
crs_obj = src.crs
nodata_read = src.nodata
desc = src.descriptions[0] if src.descriptions else ""
# 保存原生边界和分辨率,供后续 extent/scale_bar 使用
_src_bounds = src.bounds # left, bottom, right, top
_src_res = src.res # (xres, yres)
# 替换 NoData 为 NaN用于绘图
nd = nodata_read if nodata_read is not None else nodata_value
if nd is not None:
array = np.where(array == nd, np.nan, array)
else:
array = np.where(np.isnan(array), np.nan, array)
print(f"[visualize_raster] rasterio 读取成功: {raster_tif_path}")
use_rasterio = True
except Exception as rio_err:
print(f"[visualize_raster] rasterio 失败 ({rio_err}),回退到 GDAL")
use_rasterio = False
# 方式2GDAL备选
if array is None:
try:
ds = gdal.Open(raster_tif_path, gdal.GA_ReadOnly)
if ds is None:
raise RuntimeError("GDAL 无法打开文件")
array = ds.GetRasterBand(1).ReadAsArray().astype(np.float64)
gt = ds.GetGeoTransform()
proj = ds.GetProjection()
nodata_read = ds.GetRasterBand(1).GetNoDataValue()
desc = ds.GetDescription() or ""
if nodata_read is not None:
array = np.where(array == nodata_read, np.nan, array)
else:
array = np.where(np.isnan(array), np.nan, array)
# 从 GeoTransform 构造仿射变换(用于计算 extent
if gt and gt != (0, 1, 0, 0, 0, 1):
if Affine is not None:
transform = Affine(gt[1], gt[2], gt[0],
gt[4], gt[5], gt[3])
else:
transform = None
# ★★★ 关键:从 GeoTransform 计算 bounds 和 res ★★★
# gt = (xmin, xres, 0, ymax, 0, yres)
xmin_gdal = gt[0]
ymax_gdal = gt[3]
xres_gdal = gt[1]
yres_gdal = gt[5]
width_gdal = ds.RasterXSize
height_gdal = ds.RasterYSize
xmax_gdal = xmin_gdal + width_gdal * xres_gdal
ymin_gdal = ymax_gdal + height_gdal * yres_gdal
_src_bounds = rasterio.coords.BoundingBox(xmin_gdal, ymin_gdal, xmax_gdal, ymax_gdal)
_src_res = (abs(xres_gdal), abs(yres_gdal))
else:
transform = None
ds = None
except Exception as gdal_err:
raise RuntimeError(
f"无法读取 GeoTIFFrasterio 和 GDAL 均失败): {gdal_err}"
)
# ── 宽高变量(供 extent 计算和 figsize 保护使用)─────────────
w, h = array.shape[1], array.shape[0]
# 保存原始宽高transform 回退分支需用原始尺寸计算 extent
w_orig, h_orig = w, h
# ── 全面 NoData 清洗:-9999.0 / NaN / Inf → 统一转为 np.nan ──
# 这一步确保陆地像素(无论来自掩膜还是原始 NoData均被清除
# 使 nanpercentile 分位数拉伸 100% 精准锁定水体内部
array = np.where(
(array == nodata_value) | np.isnan(array) | np.isinf(array),
np.nan,
array
)
# ====== 新增:矢量掩膜物理擦除(必须在降采样之前,否则 array.shape 与 transform 错位)======
# 把水域多边形外部的陆地像素物理擦除为 NaN让下游 mean/std 统计 100% 干净(无陆地假数据污染)
# 同时保留 boundary_gdf_plotted 给末尾描边复用,避免重复读取 + 重复矢量化
boundary_gdf_plotted: Optional[Any] = None
if boundary_shp_path and os.path.isfile(boundary_shp_path) and transform is not None:
try:
boundary_ext = Path(boundary_shp_path).suffix.lower()
if boundary_ext in ('.shp',):
# 矢量:直接读取
boundary_gdf = gpd.read_file(boundary_shp_path)
elif boundary_ext in ('.dat', '.bsq', '.tif', '.tiff'):
# 栅格:复用 ContentMapper._raster_to_boundary_gdf 矢量化
boundary_gdf = self._raster_to_boundary_gdf(boundary_shp_path)
else:
raise ValueError(f"不支持的边界文件格式: {boundary_ext}")
# 兜底:如果 SHP 缺少投影文件(如 .prj 丢失),默认赋予 WGS84 (EPSG:4326)
if boundary_gdf.crs is None:
print(f"[visualize_raster] 警告: 掩膜 SHP 缺失坐标系,"
f"默认按 WGS84 (EPSG:4326) 处理 ({Path(boundary_shp_path).name})。")
boundary_gdf = boundary_gdf.set_crs(epsg=4326)
# 兜底:如果栅格 TIF 缺失坐标系
# ★ 关键修复:不能盲目回退到 WGS84 (EPSG:4326)
# transform 的坐标值如果是数十万数百万量级一定是投影坐标UTM 等),
# 而非 WGS84 经纬度(-180~180。此时应使用掩膜的 CRS 作为正确参考,
# 否则 geometry_mask 的空间坐标与 transform 完全错位 → 全部像元被擦除。
if crs_obj is None:
if (boundary_gdf.crs is not None
and boundary_gdf.crs.is_projected
and transform is not None):
# transform 坐标量级检测UTM 坐标通常在 10^5~10^7 范围
x_magnitude = max(abs(transform.c), abs(transform.c + transform.a * array.shape[1]))
if x_magnitude > 180.0:
crs_obj = boundary_gdf.crs
print(f"[visualize_raster] 栅格 TIF 缺失坐标系,"
f"根据 transform 坐标量级 (x≈{x_magnitude:.0f}) "
f"判定为投影坐标系,使用掩膜 CRS ({crs_obj.to_epsg()}) "
f"({Path(raster_tif_path).name})")
if crs_obj is None:
crs_obj = CRS.from_epsg(4326)
print(f"[visualize_raster] 警告: 栅格 TIF 缺失坐标系,"
f"默认按 WGS84 渲染 ({Path(raster_tif_path).name})。")
# 坐标系对齐到当前栅格的 CRS不是 self.output_crs必须与 transform 保持一致)
if boundary_gdf.crs != crs_obj:
boundary_gdf = boundary_gdf.to_crs(crs_obj)
# ====================================================================
# 【核心修复:仅提取实心外部轮廓,忽略内部的耀斑孔洞和数据缺失缺口】
# 彻底解决密密麻麻的黑点(孔洞描边)以及被异常挖空的方形区域
from shapely.geometry import Polygon, MultiPolygon
exteriors = []
for geom in boundary_gdf.geometry:
if geom is None or geom.is_empty:
continue
if geom.geom_type == 'Polygon':
# 只取外壳exterior丢弃所有内部孔洞interiors
exteriors.append(Polygon(geom.exterior))
elif geom.geom_type == 'MultiPolygon':
for part in geom.geoms:
exteriors.append(Polygon(part.exterior))
if exteriors:
boundary_gdf = gpd.GeoDataFrame(geometry=exteriors, crs=boundary_gdf.crs)
# ====================================================================
boundary_gdf_plotted = boundary_gdf # 留给末尾描边代码复用
# 物理擦除:多边形内为 True保留原值外部为 False强制 NaN
geom_mask = geometry_mask(
geometries=boundary_gdf.geometry,
out_shape=array.shape,
transform=transform,
invert=True,
)
array = np.where(geom_mask, array, np.nan)
kept = int((~np.isnan(array)).sum())
print(f"[visualize_raster] 矢量掩膜物理擦除完成: 陆地背景 → NaN "
f"(boundary={Path(boundary_shp_path).name}, 擦除后有效像元: "
f"{kept}/{array.size})")
except Exception as e:
print(f"[visualize_raster] 矢量掩膜物理擦除失败 ({boundary_shp_path}): {e}")
# ====================================================================
# ── 极速降采样:>400 万像元时,将矩阵降维至约 200 万像素 ─────────
# 必须放在物理擦除之后,否则 geometry_mask 的 out_shape 会与降采样后 array 不对齐
# extent 使用原始 bounds与降采样无关保证坐标轴 UTM 米精确
# 降采样切片仅影响绘图渲染,可将 1 亿像素图在 1 秒内降至 ~200 万像素
_MAX_VIZ_PIXELS = 4_000_000
if array.size > _MAX_VIZ_PIXELS:
step = int(np.ceil(np.sqrt(array.size / _MAX_VIZ_PIXELS)))
array = array[::step, ::step]
w_downsampled, h_downsampled = array.shape[1], array.shape[0]
print(f"[visualize_raster] 极速降采样: {w}×{h}{w_downsampled}×{h_downsampled} "
f"(step={step}),节省内存并加速渲染")
w, h = w_downsampled, h_downsampled
# ── 从描述推断参数名和 colormap ───────────────────────────────
# 描述格式Formula_Name|Category|Formula_Type|Formula
param_name: Optional[str] = None
if desc and '|' in desc:
parts = desc.split('|')
param_name = parts[0].strip()
if len(parts) >= 2:
category = parts[1].strip()
if not cmap:
cmap = self._get_colormap(category)
elif not cmap:
# 从文件名推断
stem = tif_path.stem
param_name = self._extract_param_name(str(tif_path))
cmap = self._get_colormap(param_name)
# ── 中文标题(文件名汉化 + 绘图标题)──────────────────────────
# 用户显式传入 title 时直接使用;否则用中文映射
chinese_title = self._get_chinese_title(stem) if not title else title
# ── 计算空间范围extent──────────────────────────────────────
# 优先使用 rasterio 原生 bounds保证坐标轴为真实 UTM 米
# GDAL 回退使用 GeoTransform 计算
if _src_bounds is not None:
extent = [
_src_bounds.left, # xmin
_src_bounds.right, # xmax
_src_bounds.bottom, # ymin
_src_bounds.top, # ymax
]
# 从 bounds 推导分辨率(取绝对值,正数用于比例尺)
scale_x = abs(_src_res[0]) if _src_res else 1.0
scale_y = abs(_src_res[1]) if _src_res else 1.0
elif transform is not None:
xmin = transform.c
ymax = transform.f
xres = transform.a
yres = transform.e
# ★★★ 必须用原始宽高w_orig/h_orig而非降采样后的 w/h ★★★
extent = [xmin, xmin + w_orig * xres, ymax + h_orig * yres, ymax]
scale_x = abs(xres)
scale_y = abs(yres)
else:
# 回退到像素索引范围(使用原始尺寸)
extent = [0, w_orig, 0, h_orig]
scale_x = 1.0
scale_y = 1.0
# ── 准备图形 ─────────────────────────────────────────────────
# 尊重用户的 figsize 参数(同时设置内存安全上限防止 DPI=300 下爆内存)
_max_inch = 60 # 每维最大 60 英寸,远超常规打印需求
safe_w = min(float(figsize[0]), _max_inch)
safe_h = min(float(figsize[1]), _max_inch)
fig, ax = plt.subplots(figsize=(safe_w, safe_h))
# 1. 明确排除 NaN 以及 nodata_value与函数参数保持一致
nodata_val = nodata_value
valid = array[(~np.isnan(array)) & (array != nodata_val)]
if valid.size == 0:
raise ValueError("GeoTIFF 中没有有效数据(全部为 NoData")
# 2. 改用更鲁棒的 2%-98% 百分位拉伸(抗偏态分布)
vmin = float(np.percentile(valid, 2))
vmax = float(np.percentile(valid, 98))
if (vmax - vmin) < 1e-9: # 防退化:区间过窄 → 取中心 ±1%
center = vmin
exp = max(abs(center) * 0.01, 1e-9)
vmin = center - exp
vmax = center + exp
print(f"[visualize_raster] P2-P98 拉伸: vmin={vmin:.4f}, vmax={vmax:.4f}"
f"有效像元: {valid.size}/{array.size}")
# ── 栅格绘图 ─────────────────────────────────────────────────
# 使用 masked arrayNaN 区域自动不显示
# 不仅要屏蔽 NaN如果 array 中还有 nodata_value也必须 mask 掉,否则出图时背景会被渲染
masked_data = np.ma.masked_where((np.isnan(array)) | (array == nodata_val), array)
# 【核心修复2】废弃错误的坐标映射逻辑。
# 直接使用原生的 imshow明确告知 matplotlib 第0行在最上方(origin='upper')
im = ax.imshow(
masked_data,
extent=[extent[0], extent[1], extent[2], extent[3]],
origin='upper',
cmap=cmap or 'viridis',
vmin=vmin, vmax=vmax,
alpha=alpha,
interpolation='bilinear'
)
# ★★★ 锁死绘图视口 ★★★
# 必须在所有叠加绘图shp/colorbar/north arrow之前执行
# 防止其他元素的坐标干扰导致轴范围被拉伸成像素坐标系
ax.set_xlim(extent[0], extent[1])
ax.set_ylim(extent[2], extent[3])
# ── 边界描边(直接复用上面的 boundary_gdf_plotted不再重复读取/矢量化)─────────
if boundary_gdf_plotted is not None:
try:
boundary_gdf_plotted.boundary.plot(ax=ax, color='black', linewidth=1.5)
except Exception as e:
print(f"[visualize_raster] 边界描边失败: {e}")
# ── 坐标轴标签(动态:根据 bounds 阈值判断经纬度 vs UTM 米)───
# 经纬度坐标系下 bounds.left 在 [-180, 180]UTM 投影坐标在百万级
if _src_bounds is not None and _src_bounds.left < 180 and _src_bounds.bottom < 90:
ax.set_xlabel('Longitude', fontsize=14)
ax.set_ylabel('Latitude', fontsize=14)
else:
ax.set_xlabel('X (UTM Meters)', fontsize=14)
ax.set_ylabel('Y (UTM Meters)', fontsize=14)
# 显式设置刻度字号(避免 matplotlib 默认 10pt 在大画布下偏小)
ax.tick_params(axis='both', which='major', labelsize=12)
ax.tick_params(axis='both', which='minor', labelsize=10)
ax.grid(True, linestyle='--', linewidth=0.5, alpha=0.4, color='gray')
ax.set_axisbelow(True)
# ── 标题(中文)──────────────────────────────────────────────
ax.set_title(chinese_title, fontsize=18, fontweight='bold', pad=15)
# ── 颜色条工业级样式extend 三角 + MaxNLocator 刻度防重叠)─────────
if show_colorbar and im is not None:
try:
cbar = fig.colorbar(im, ax=ax, shrink=0.6, aspect=30, pad=0.03, extend='both')
cbar.set_label('Index Value', fontsize=14)
cbar.ax.tick_params(labelsize=12)
cbar.locator = MaxNLocator(nbins=6)
cbar.update_ticks()
print("[visualize_raster] 颜色条添加成功")
except Exception as e:
print(f"[visualize_raster] 颜色条添加失败: {e}")
# ── 比例尺 ───────────────────────────────────────────────────
try:
self.add_scale_bar(ax, scale_x=scale_x, scale_y=scale_y)
except Exception as e:
print(f"[visualize_raster] 比例尺添加失败: {e}")
# ── 指北针 ───────────────────────────────────────────────────
try:
bounds_arr = np.array(extent)
self.add_north_arrow(ax, bounds_arr)
except Exception as e:
print(f"[visualize_raster] 指北针添加失败: {e}")
# ── 紧凑布局并保存 ───────────────────────────────────────────
# 给 tight_layout 显式 pad=2.0 防止指北针/比例尺/标题互相重叠
plt.tight_layout(pad=2.0, h_pad=1.5, w_pad=1.5)
try:
plt.savefig(
output_file,
dpi=300,
bbox_inches='tight',
facecolor='white',
edgecolor='none',
)
print(f"[visualize_raster] ✅ 专题图已保存: {output_file}")
except Exception as e:
print(f"[visualize_raster] 保存失败: {e}")
raise
try:
plt.show()
except Exception:
pass
plt.close(fig)
return output_file
# ------------------------------------------------------------------
# Step 11 改造:插值网格 → GeoTIFF 物理落盘(替代 PNG 渲染)
# ------------------------------------------------------------------
@staticmethod
def _redirect_png_to_tif_path(output_file: str) -> str:
"""
智能路径重定向:把 14_visualization/visualization 改写到 11_Thematic_Map.png 换 .tif
触发条件(与 panel 默认输出目录约定保持一致):
- 原路径最后一级目录名为 ``14_visualization`` 或 ``visualization`` → 重定向到 ``11_Thematic_Map``
- 其它情况(如用户自定义路径)→ 仅替换后缀,保留原目录
"""
from pathlib import Path
p = Path(output_file)
parent = p.parent
if parent.name in ('14_visualization', 'visualization'):
new_parent = parent.parent / '11_Thematic_Map'
else:
new_parent = parent
new_filename = p.stem + '.tif'
return str(new_parent / new_filename)
def _save_as_geotiff(self, grid_content, grid_xx, grid_yy,
output_tif_path, nodata_value=-9999.0):
"""
将插值网格矩阵落盘为带坐标系的 GeoTIFF 文件GDAL 实现)。
与 ``src/utils/kriging.py:KrigingInterpolator.save_raster`` 风格保持一致:
GTiff + LZW + TILED + BIGTIFF=IF_SAFER + Float64。
Parameters
----------
grid_content : np.ndarray, shape (rows, cols)
插值后的二维网格矩阵(含 NaN
grid_xx : np.ndarray, shape (rows, cols)
X 坐标 meshgrid与 grid_content 同 shape
grid_yy : np.ndarray, shape (rows, cols)
Y 坐标 meshgrid与 grid_content 同 shape
output_tif_path : str
输出 GeoTIFF 完整路径
nodata_value : float
NoData 值(默认 -9999.0NaN 将被替换为该值)
"""
try:
from osgeo import gdal, osr
except ImportError as e:
raise ImportError(
"需要 osgeo (GDAL) 库来写 GeoTIFF请检查 conda 环境: " + str(e)
)
# ── 1. 计算 GeoTransform 6 参数 ──────────────────────────────
# GDAL 约定transform = (x_min, dx, 0, y_max, 0, -abs(dy))
# y 轴向下方为正(影像行列与地理坐标对应)
x_min = float(grid_xx[0, 0])
y_max = float(grid_yy[-1, 0]) # meshgrid 末尾 = 北方
dx = float(grid_xx[0, 1] - grid_xx[0, 0]) if grid_xx.shape[1] > 1 else 0.0
dy_raw = float(grid_yy[1, 0] - grid_yy[0, 0]) if grid_yy.shape[0] > 1 else 0.0
# dy 在地理上为正向北递增GDAL 用负值表示"行索引向下走 Y 增大"
dy = -abs(dy_raw) if dy_raw != 0 else -abs(dx)
if abs(dx) < 1e-12 or abs(dy) < 1e-12:
raise ValueError(
f"网格分辨率异常: dx={dx}, dy={dy},请检查 create_interpolation_grid 输入"
)
rows, cols = grid_content.shape
geotransform = (x_min, dx, 0, y_max, 0, dy)
# ── 2. NaN → nodata统一 Float64与 kriging.py 一致)────────
data_clean = np.where(np.isnan(grid_content), nodata_value, grid_content)
data_clean = data_clean.astype(np.float64)
# 【核心修复1】上下翻转矩阵
# Python网格第0行是南方但GeoTIFF规范第0行是北方。
# 必须翻转,否则落盘的 TIF 永远是上下颠倒的!
data_clean = np.flipud(data_clean)
# ── 3. CRS 字符串 → WKToutput_crs 是 EPSG 字符串或 WKT 均可)────
try:
srs = osr.SpatialReference()
srs.SetFromUserInput(self.output_crs)
proj_wkt = srs.ExportToWkt()
except Exception as e:
print(f"[_save_as_geotiff] CRS 解析失败 ({e}),兜底使用 EPSG:4326")
srs = osr.SpatialReference()
srs.SetFromUserInput("EPSG:4326")
proj_wkt = srs.ExportToWkt()
# ── 4. 创建输出目录 ─────────────────────────────────────────
out_dir = os.path.dirname(output_tif_path)
if out_dir:
os.makedirs(out_dir, exist_ok=True)
# ── 5. GDAL 写盘 ─────────────────────────────────────────────
driver = gdal.GetDriverByName("GTiff")
if driver is None:
raise RuntimeError("GDAL GTiff 驱动不可用,请检查 osgeo 安装")
dataset = driver.Create(
output_tif_path, cols, rows, 1, gdal.GDT_Float64,
options=[
"COMPRESS=LZW",
"TILED=YES",
"BIGTIFF=IF_SAFER", # ⭐ 与 Step 8/10 Kriging 写盘保持一致
],
)
if dataset is None:
raise RuntimeError(f"无法创建输出 GeoTIFF: {output_tif_path}")
dataset.SetGeoTransform(geotransform)
dataset.SetProjection(proj_wkt)
band = dataset.GetRasterBand(1)
band.WriteArray(data_clean)
band.SetNoDataValue(nodata_value)
band.ComputeStatistics(0)
band.FlushCache()
dataset.FlushCache()
del dataset
valid_mask = ~np.isnan(grid_content)
print(f"[_save_as_geotiff] ✅ GeoTIFF 已保存: {output_tif_path}")
print(f" 分辨率: dx={dx:.6f}, dy={abs(dy):.6f}")
print(f" 范围: x=[{x_min:.4f}, {x_min + cols * dx:.4f}], "
f"y=[{y_max + rows * dy:.4f}, {y_max:.4f}]")
print(f" 尺寸: {rows} × {cols}, CRS: {self.output_crs}")
print(f" NoData={nodata_value}, 有效像元: {int(valid_mask.sum())}/{grid_content.size}")
return output_tif_path
# ═══════════════════════════════════════════════════════════════
# ★ 2026-07-01共享空间上下文 — 63 个 CSV 只算一次网格/掩膜
# ═══════════════════════════════════════════════════════════════
def prepare_shared_context(self, sample_csv: str, shp_file=None,
resolution=100, expand_ratio=0.05):
"""从首个 CSV 预计算所有子进程共用的空间基准数据。
63 个水色指数 CSV 坐标完全一致,以下数据只算一次:
- boundary_gdf (水域边界)
- grid_xx, grid_yy (插值网格)
- mask (水域掩膜布尔矩阵)
- bounds (空间范围)
子进程直接从 shared_context 解包复用,跳过 ②③④⑥,直入 Kriging。
Returns:
tuple: (grid_xx, grid_yy, mask, bounds, boundary_gdf)
"""
print(f"[共享上下文] 从 {Path(sample_csv).name} 预计算空间基准...")
# ② 读边界(只此一次)
if shp_file is None:
boundary_gdf = None
else:
boundary_gdf = self.read_boundary_shapefile(shp_file)
# ③ 边缘外扩(只此一次)—— 需要读第一个CSV获取坐标结构
points_gdf = self.read_csv_data(sample_csv)
points_gdf = self._expand_edge_points(
points_gdf, boundary_gdf, resolution=resolution, expand_ratio=expand_ratio
)
# ④ 计算网格几何
if boundary_gdf is None:
pts = np.column_stack((points_gdf['proj_x'], points_gdf['proj_y']))
minx, maxx = pts[:, 0].min(), pts[:, 0].max()
miny, maxy = pts[:, 1].min(), pts[:, 1].max()
else:
bnd = boundary_gdf.total_bounds
minx, miny, maxx, maxy = bnd
width = maxx - minx
height = maxy - miny
minx -= width * expand_ratio
maxx += width * expand_ratio
miny -= height * expand_ratio
maxy += height * expand_ratio
res = resolution / 111000.0 if self.output_crs == 'EPSG:4326' else resolution
nx = max(int(width / res), 100)
ny = max(int(height / res), 100)
grid_x = np.linspace(minx, maxx, nx)
grid_y = np.linspace(miny, maxy, ny)
grid_xx, grid_yy = np.meshgrid(grid_x, grid_y)
bounds = np.array([minx, miny, maxx, maxy])
print(f"[共享上下文] 网格: {nx}×{ny} = {nx*ny}")
# ⑥ 水域掩膜布尔矩阵(只此一次)
mask = None
if boundary_gdf is not None:
mask_pts = np.column_stack((grid_xx.ravel(), grid_yy.ravel()))
mask_gdf = gpd.GeoDataFrame(
geometry=[Point(x, y) for x, y in mask_pts], crs=self.output_crs
)
mask = mask_gdf.within(boundary_gdf.unary_union).values.reshape(grid_xx.shape)
print(f"[共享上下文] 水域掩膜: {int(mask.sum())}/{mask.size} 点在水域内")
return (grid_xx, grid_yy, mask, bounds, boundary_gdf)
def process_data(self, csv_file, shp_file=None, output_file='content_map.png',
resolution=100, show_sample_points=False, base_map_tif=None,
use_distance_diffusion=True, max_diffusion_distance=None,
diffusion_power=2, diffusion_n_neighbors=15, cmap=None,
expand_ratio=0.05,
output_format='tif',
shared_context=None):
"""
主处理函数
★★★ Plan Cshp_file 现在是可选参数None = 纯采样点插值,不依赖水域掩膜)★★★
Parameters:
-----------
csv_file : str
CSV文件路径
shp_file : str, optional
水域掩膜/边界文件路径(.shp / .dat / .bsq / .tif 等)。
shared_context : tuple, optional (2026-07-01 批量优化)
由 prepare_shared_context() 返回的 (grid_xx, grid_yy, mask, bounds, boundary_gdf)。
提供时跳过 读边界/边缘外扩/建网格/算掩膜,直入 Kriging 插值阶段。
... (其他参数同上)
"""
try:
# 自动识别参数名称并获取colormap
if cmap is None:
param_name = self._extract_param_name(csv_file)
cmap = self._get_colormap(param_name)
# 读取采样点数据
points_gdf = self.read_csv_data(csv_file)
# ── ★ 快速通道:复用预计算的共享上下文 ──
if shared_context is not None:
grid_xx, grid_yy, mask, bounds, boundary_gdf = shared_context
# ③ 仍需边缘扩展(值相关),但跳过 ②④⑥
points_gdf = self._expand_edge_points(
points_gdf, boundary_gdf, resolution=resolution,
expand_ratio=expand_ratio
)
# ⑤ 直接用共享网格执行 Kriging
pts = np.column_stack((points_gdf['proj_x'], points_gdf['proj_y']))
vals = points_gdf['content'].values
grid_content = self._perform_interpolation(pts, vals, grid_xx, grid_yy)
# ⑥ 复用共享掩膜裁剪
if mask is not None:
grid_content[~mask] = np.nan
# 边界内 NaN 填充
nan_mask = np.isnan(grid_content)
within_nan = nan_mask & mask
if np.any(within_nan):
valid_m = ~nan_mask & mask
if np.sum(valid_m) > 0:
v_pts = np.column_stack((grid_xx[valid_m], grid_yy[valid_m]))
v_vals = grid_content[valid_m]
n_pts = np.column_stack((grid_xx[within_nan], grid_yy[within_nan]))
try:
from scipy.interpolate import griddata
grid_content[within_nan] = griddata(
v_pts, v_vals, n_pts, method='nearest'
)
except Exception:
grid_content[within_nan] = np.nanmean(grid_content[mask])
else:
# ── 原有完整流程(单图模式)──────────
if shp_file is None:
print("[Plan C] shp_file=None跳过水域掩膜读取")
boundary_gdf = None
else:
boundary_gdf = self.read_boundary_shapefile(shp_file)
points_gdf = self._expand_edge_points(
points_gdf, boundary_gdf, resolution=resolution,
expand_ratio=expand_ratio
)
grid_xx, grid_yy, grid_content, bounds = self.create_interpolation_grid(
points_gdf, boundary_gdf, resolution,
expand_ratio=expand_ratio,
use_distance_diffusion=use_distance_diffusion,
max_diffusion_distance=max_diffusion_distance,
diffusion_power=diffusion_power,
diffusion_n_neighbors=diffusion_n_neighbors
)
# ── 按 output_format 分发落盘方式 ───────────────────────────
if output_format == 'tif':
output_tif_path = self._redirect_png_to_tif_path(output_file)
self._save_as_geotiff(grid_content, grid_xx, grid_yy, output_tif_path)
actual_output = output_tif_path
else:
self.create_content_map(
points_gdf, boundary_gdf, grid_xx, grid_yy,
grid_content, bounds, output_file, show_sample_points,
base_map_tif, cmap=cmap
)
actual_output = output_file
print("处理完成!")
print(f"\n统计信息:")
print(f"数据点数量: {len(points_gdf)}")
print(f"含量值范围: {points_gdf['content'].min():.2f} - {points_gdf['content'].max():.2f}")
print(f"含量值平均: {points_gdf['content'].mean():.2f}")
print(f"含量值标准差: {points_gdf['content'].std():.2f}")
print(f"输出文件: {actual_output}")
except Exception as e:
print(f"处理过程中出现错误: {str(e)}")
raise
def process_batch(self, csv_folder, shp_file=None, output_folder=None,
resolution=100, show_sample_points=False, base_map_tif=None,
use_distance_diffusion=True, max_diffusion_distance=None,
diffusion_power=2, diffusion_n_neighbors=15):
"""
批量处理文件夹中的CSV文件
★★★ Plan Cshp_file 可选None = 不依赖水域掩膜,纯采样点插值)★★★
Parameters:
-----------
csv_folder : str
包含CSV文件的文件夹路径
shp_file : str, optional
水域掩膜/边界文件路径。None 时跳过掩膜约束。
output_folder : str, optional
输出文件夹路径。如果为None将在CSV文件所在文件夹创建'map_output'子文件夹
resolution : int, default=100
网格分辨率(米)
show_sample_points : bool, default=False
是否显示采样点
base_map_tif : str, optional
TIF正射底图文件路径
use_distance_diffusion : bool, default=True
是否使用距离扩散方法
max_diffusion_distance : float, optional
最大扩散距离
diffusion_power : float, default=2
距离扩散的IDW幂参数
diffusion_n_neighbors : int, default=15
距离扩散使用的最近邻点数
"""
print("=" * 60)
print("开始批量处理CSV文件")
print("=" * 60)
# 检查输入文件夹是否存在
if not os.path.isdir(csv_folder):
raise ValueError(f"输入文件夹不存在: {csv_folder}")
# 获取所有CSV文件
csv_files = glob.glob(os.path.join(csv_folder, "*.csv"))
if len(csv_files) == 0:
raise ValueError(f"在文件夹 {csv_folder} 中未找到CSV文件")
print(f"找到 {len(csv_files)} 个CSV文件")
# 创建输出文件夹
if output_folder is None:
output_folder = os.path.join(csv_folder, "map_output")
if not os.path.exists(output_folder):
os.makedirs(output_folder)
print(f"创建输出文件夹: {output_folder}")
else:
print(f"使用输出文件夹: {output_folder}")
# 统计信息
success_count = 0
fail_count = 0
failed_files = []
# 批量处理每个CSV文件
for i, csv_file in enumerate(csv_files, 1):
print("\n" + "=" * 60)
print(f"处理文件 {i}/{len(csv_files)}: {os.path.basename(csv_file)}")
print("=" * 60)
try:
# 生成输出文件名使用CSV文件名但扩展名为.png
csv_basename = os.path.splitext(os.path.basename(csv_file))[0]
output_file = os.path.join(output_folder, f"{csv_basename}.png")
# 处理单个文件自动识别参数并选择colormap
self.process_data(
csv_file=csv_file,
shp_file=shp_file,
output_file=output_file,
resolution=resolution,
show_sample_points=show_sample_points,
base_map_tif=base_map_tif,
use_distance_diffusion=use_distance_diffusion,
max_diffusion_distance=max_diffusion_distance,
diffusion_power=diffusion_power,
diffusion_n_neighbors=diffusion_n_neighbors,
cmap=None # 自动识别
)
success_count += 1
print(f"✓ 成功处理: {csv_basename}.png")
except Exception as e:
fail_count += 1
failed_files.append((os.path.basename(csv_file), str(e)))
print(f"✗ 处理失败: {os.path.basename(csv_file)}")
print(f" 错误信息: {e}")
import traceback
traceback.print_exc()
# 输出批量处理结果统计
print("\n" + "=" * 60)
print("批量处理完成")
print("=" * 60)
print(f"总文件数: {len(csv_files)}")
print(f"成功: {success_count}")
print(f"失败: {fail_count}")
print(f"输出文件夹: {output_folder}")
if failed_files:
print("\n失败的文件列表:")
for file_name, error in failed_files:
print(f" - {file_name}: {error}")
return {
'total': len(csv_files),
'success': success_count,
'failed': fail_count,
'output_folder': output_folder,
'failed_files': failed_files
}
def main():
"""主函数 - 使用示例"""
# 创建处理器实例
mapper = ContentMapper()
# 示例1处理单个文件
csv_file = r"E:\code\WQ\pipeline_result\tests1\11_12_13_predictions\BGA.csv" # 采样点的预测值
shp_file = r"D:\BaiduNetdiskDownload\yaobao\roi\roi.shp" # 水体边界shapefile路径
output_file = r"E:\code\WQ\pipeline_result\work_dir\11_12_13_predictions\BGA.png" # 输出图片路径
#
mapper.process_data(
csv_file=csv_file,
shp_file=shp_file,
output_file=output_file,
resolution=30, # 网格分辨率(米),更小的值产生更平滑的效果
show_sample_points=False, # 设置为False以显示平滑的颜色分布True则显示采样点位置
base_map_tif=None, # 正射底图路径(可选)
cmap=None # 自动从文件名或内容中识别参数并选择对应的colormap
)
# # 示例2批量处理文件夹中的所有CSV文件
# csv_folder = r"E:\code\WQ\xiaogujia\使用腰堡模型\predict\TT.csv" # CSV文件所在文件夹
# shp_file = r"E:\code\WQ\xiaogujia\SHP\shp\watemask.shp" # 水体边界shapefile路径
# output_folder = r"E:\code\WQ\xiaogujia\使用腰堡模型\map\TT.png" # 输出文件夹可选如果为None则在CSV文件夹下创建map_output
# 批量处理会自动识别每个CSV文件的参数名称并选择对应的colormap
# result = mapper.process_batch(
# csv_folder=csv_folder,
# shp_file=shp_file,
# output_folder=output_folder, # 如果为None将在CSV文件夹下创建map_output子文件夹
# resolution=30, # 网格分辨率(米)
# show_sample_points=False, # 是否显示采样点
# base_map_tif=None, # 正射底图路径(可选)
# )
#
# print(f"\n批量处理结果: {result}")
if __name__ == "__main__":
# 使用示例
print("含量分布图生成器")
print("=" * 50)
# 如果要直接运行,请取消下面的注释并修改文件路径
main()
# 或者交互式使用
# print("使用方法:")
# print("1. 准备CSV文件前两列为WGS84经纬度第三列为含量数据")
# print("2. 准备边界Shapefile文件")
# print("3. 调用以下代码:")
# print("""
# mapper = ContentMapper()
# mapper.process_data(
# csv_file='your_data.csv',
# shp_file='your_boundary.shp',
# output_file='output_map.png',
# resolution=50
# )
# """)
# ═══════════════════════════════════════════════════════════════
# 模块级函数:多进程 worker必须在类外部定义供 Pool.map 使用)
# ═══════════════════════════════════════════════════════════════
def _local_krige_block_worker(args):
"""单个局部克里金块任务(独立进程入口,必须为模块级函数)"""
(local_pts, local_vals, sub_grid_x, sub_grid_y,
bx_min, bx_max, by_min, by_max,
grid_dx, grid_dy, block_ix, block_iy, total) = args
import numpy as np
from pykrige.ok import OrdinaryKriging
if len(local_pts) < 3:
return (None, bx_min, bx_max, by_min, by_max)
print(f" [LocalKrige] 块 ({block_iy},{block_ix}) [{block_iy * 100 + block_ix}/{total}] "
f"采样点={len(local_pts)}, 网格={len(sub_grid_x)}×{len(sub_grid_y)}")
ok = OrdinaryKriging(
local_pts[:, 0], local_pts[:, 1], local_vals,
variogram_model='spherical',
verbose=False,
enable_plotting=False,
)
z, ss = ok.execute(
'grid', sub_grid_x, sub_grid_y,
backend='loop',
n_closest_points=min(15, len(local_pts)),
)
return (np.array(z), bx_min, bx_max, by_min, by_max)