Files
WQ_GUI/src/postprocessing/point_map.py
2026-06-25 15:50:02 +08:00

312 lines
15 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

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

#!/usr/bin/env python
# -*- coding: utf-8 -*-
"""
采样点地图生成模块 - 在高光谱假彩色影像上标注采样点
"""
import numpy as np
import pandas as pd
import matplotlib.pyplot as plt
from pathlib import Path
from typing import Optional, Tuple, List, Dict, Union
import warnings
from matplotlib.patches import FancyArrowPatch
import matplotlib.patheffects as path_effects
# 性能优化配置
plt.rcParams['agg.path.chunksize'] = 10000
plt.rcParams['path.simplify'] = True
plt.rcParams['path.simplify_threshold'] = 0.1
try:
from osgeo import gdal, osr
GDAL_AVAILABLE = True
except ImportError:
GDAL_AVAILABLE = False
print("警告: GDAL未安装,地理坐标转换功能可能无法正常工作")
class SamplingPointMap:
def __init__(self, output_dir: str = "./point_maps", fast_mode: bool = False):
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
warnings.filterwarnings('ignore', category=UserWarning)
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:
if not GDAL_AVAILABLE:
raise ImportError("GDAL未安装,无法处理地理坐标转换")
print(f"正在生成采样点地图...{' (快速模式)' if self.fast_mode else ''}")
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
)
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]:
dataset = gdal.Open(hyperspectral_path)
if dataset is None:
raise ValueError(f"无法打开高光谱影像: {hyperspectral_path}")
width = dataset.RasterXSize
height = dataset.RasterYSize
band_count = dataset.RasterCount
if rgb_bands is None:
if band_count >= 3:
try:
from src.utils.util import find_band_number
rgb_bands = [
find_band_number(650.0, hyperspectral_path),
find_band_number(550.0, hyperspectral_path),
find_band_number(460.0, hyperspectral_path)
]
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]
if downsample and (width > 2000 or height > 2000):
print(f" ⚠ 下采样暂被禁用,使用原始分辨率: {width}x{height}")
sample_factor = 1
else:
sample_factor = 1
rgb_data = []
for band_idx in rgb_bands:
band = dataset.GetRasterBand(band_idx + 1)
band_data = band.ReadAsArray().astype(np.float32)
rgb_data.append(band_data)
if len(rgb_data) == 3:
image_array = np.stack(rgb_data, axis=2)
else:
image_array = np.stack([rgb_data[0]] * 3, axis=2)
geotransform = dataset.GetGeoTransform()
projection = dataset.GetProjection()
dataset = None
return image_array, geotransform, projection, width, height, sample_factor
def _read_sampling_points(self, csv_path: str) -> pd.DataFrame:
"""智能读取采样点,自动识别模糊列名,允许UTM坐标,自动修复颠倒坐标"""
df = pd.read_csv(csv_path)
if len(df.columns) < 2:
raise ValueError("CSV文件至少需要两列(经度、纬度 或 X、Y)")
# 智能子串匹配
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)
if len(df) > 0:
lat_max = df['latitude'].abs().max()
lon_max = df['longitude'].abs().max()
# 智能对调:如果纬度 > 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:
if image_array.shape[2] != 3:
if len(image_array.shape) == 2 or image_array.shape[2] == 1:
image_array = np.stack([image_array] * 3, axis=2) if len(image_array.shape) == 2 else np.repeat(
image_array, 3, axis=2)
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)
p_low = np.percentile(valid_data, min_percent)
p_high = np.percentile(valid_data, max_percent)
if p_high - p_low < 1e-8:
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)
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_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)
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))]
pixel_coords = []
gt = geotransform
needs_transform = projection and ("PROJCS" in projection or "GEOGCS" in projection)
# 智能判定是否为 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)
transform = None
if needs_transform and is_wgs84 and GDAL_AVAILABLE:
try:
src_srs = osr.SpatialReference()
src_srs.ImportFromEPSG(4326)
dst_srs = osr.SpatialReference()
dst_srs.ImportFromWkt(projection)
transform = osr.CoordinateTransformation(src_srs, dst_srs)
except Exception as e:
transform = None
elif not is_wgs84:
print(" ℹ 采样点为投影坐标(UTM),跳过WGS84投影转换,直接使用放射变换映射")
for _, row in sampling_points.iterrows():
lon, lat = float(row['longitude']), float(row['latitude'])
if transform is not None:
try:
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, y = (lon - gt[0]) / gt[1], (lat - gt[3]) / gt[5]
if sample_factor > 1:
x, y = x / sample_factor, y / sample_factor
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):
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)
ax.imshow(rgb_image, interpolation='nearest' if self.fast_mode else 'bilinear')
if pixel_coords:
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_scale_bar and geotransform is not None: self._add_scale_bar(ax, geotransform, width, height)
if show_legend:
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': 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'):
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))
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)
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
pixel_size_x = abs(geotransform[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
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))