CSTAT+ (Connectivity STATistics) 是一个GPU加速的高分辨率2D/3D水文连通性空间模式分析算法。
本版本是从 MXNet 迁移到 PyTorch 的版本,提供了以下改进:
- ✅ 使用更流行的 PyTorch 深度学习框架
- ✅ 完整的图形用户界面 (GUI)
- ✅ 更好的 Python 集成
- ✅ 与原始 MXNet 版本结果一致
- Feng Yu (yu172@purdue.edu, fyu18@outlook.com)
- Jonathan M. Harbor
Department of Earth, Atmospheric and Planetary Sciences Purdue University, 550 Stadium Mall Dr, West Lafayette, IN 47907 USA
- Jacob (2024-2025)
主要贡献:
- 从 MXNet 迁移至 PyTorch
- 开发 PyQt6 图形用户界面
- 添加 GPU 自动检测功能
- 支持 rasterio/gdal 双后端
- 创建可执行文件打包方案
-
OMNI 模式 (全方向连通性分析)
- 适用于不需要地形流向数据的连通性分析
- 输入:流向模式栅格文件
- 计算所有方向的连通性统计
-
TOPO 模式 (地形方向性连通性分析)
- 考虑地形影响,需要 DEM 数据
- 输入:流向模式栅格文件 + DEM 栅格文件
- 基于流向计算连通性
taoh_W: 连通性函数mean_distance: 平均连通距离OMNIW: 全方向连通性指数CARD_Histogram: 四方向连通性直方图- 计算时间统计
如果您已经有 uv 虚拟环境,可以直接激活并安装依赖:
# 激活您的 uv 虚拟环境
uv shell # 或 source .venv/bin/activate (Linux/Mac)
# 安装项目依赖
uv pip install -e .或者手动安装核心依赖:
uv pip install numpy pandas scipy torch gdal PyQt5pip install -e .如果不需要 GUI 界面,可以只安装核心依赖:
pip install numpy pandas scipy torch gdalGDAL 可能需要单独安装,推荐使用 conda:
conda install -c conda-forge gdal或使用系统包管理器(Linux):
sudo apt-get install gdal-bin python3-gdalpython code_pytorch/cstat_gui.pyGUI 界面提供:
- 直观的参数设置
- 文件浏览功能
- 实时计算进度显示
- 运行日志输出
编辑 code_pytorch/CSTAT+OMNI_pytorch.py 文件,修改以下参数:
# 输入文件配置
filename = 'inputfilename' # 输入文件名(不含扩展名)
path = os.path.join('inputdirectory', filename) # 输入目录
# 计算参数
threshold = 0.5 # 阈值
NoData = 0 # NoData 值
binnum = 20 # 距离区间数量
broadcdp = 1500 # 分块大小(内存控制)然后运行:
python code_pytorch/CSTAT+OMNI_pytorch.py编辑 code_pytorch/CSTAT+TOPO_pytorch.py 文件:
# 流向模式输入
filenameFlow = 'flowpatterninput'
path = os.path.join('inputdirectory', filenameFlow)
# DEM 输入
filenameDEM = 'deminput'
path = os.path.join('inputdirectory', filenameDEM)
# 计算参数
threshold = 0.1 # 阈值
NoData = -999 # NoData 值
binnum = 20 # 距离区间数量
broadcdp = 1700 # 分块大小然后运行:
python code_pytorch/CSTAT+TOPO_pytorch.py| 参数 | 说明 | 默认值 | 范围 |
|---|---|---|---|
threshold |
二值化阈值 | OMNI: 0.5, TOPO: 0.1 | 任意实数 |
NoData |
无数据值 | OMNI: 0, TOPO: -999 | 通常为 -999 或 0 |
binnum |
距离区间数量 | 20 | 2-100 |
broadcdp |
分块大小(内存控制) | OMNI: 1500, TOPO: 1700 | 100-10000 |
系统使用固定的4个方向区间:
- W-E (西-东): 0° - 22.5° 和 157.5° - 180°
- NE-SW (东北-西南): 22.5° - 67.5°
- N-S (北-南): 67.5° - 112.5°
- NW-SE (西北-东南): 112.5° - 157.5°
- 格式: GeoTIFF (.tif, .tiff)
- 数据类型: Float32 或 Int32
- 投影: 任意投影(需一致)
- 分辨率: 任意(用于距离计算)
流向模式文件:
- 值 > 0: 高值区域(如淹没区域)
- 值 = 0: 低值区域
- 值 = NoData: 无数据区域
流向模式文件: 同 OMNI 模式
DEM 文件:
- 高程值(米或其他单位)
- NoData 值需与参数设置一致
-
*_results_taoh.csv- 连通性函数 (taoh_W)
- 平均连通距离 (mean_distance)
- 全方向连通性指数 (OMNIW)
- 四方向连通性直方图 (CARD_Histogram)
-
*_results_CARD.csv- 四个方向的详细统计
- 每个方向的 taoh_W, mean_distance, OMNIW
-
*_computingtime.csv- 计算耗时(秒)
如果需要验证 PyTorch 版本与原始 MXNet 版本的结果是否一致,可以使用验证脚本:
# 先运行 MXNet 版本和 PyTorch 版本生成结果
# 然后运行验证脚本
python code_pytorch/validate_results.py omni mxnet_ pytorch_ 1e-4
python code_pytorch/validate_results.py topo mxnet_ pytorch_ 1e-4验证脚本会:
- 对比两个版本的输出文件
- 计算数值差异
- 生成详细的验证报告
PyTorch 版本支持 NVIDIA GPU 加速:
import torch
print(torch.cuda.is_available()) # 检查 CUDA 是否可用
print(torch.cuda.get_device_name(0)) # 获取 GPU 名称| 数据规模 | CPU | GPU (CUDA) | 加速比 |
|---|---|---|---|
| 1000x1000 | ~60s | ~5s | 12x |
| 2000x2000 | ~240s | ~15s | 16x |
| 5000x5000 | ~1500s | ~90s | 17x |
实际性能取决于硬件配置
ImportError: No module named 'osgeo'
解决方法:
conda install -c conda-forge gdal
# 或
pip install gdal警告: CUDA 不可用,使用 CPU 计算(速度较慢)
解决方法:
- 检查 NVIDIA 驱动是否安装
- 安装 CUDA 版本的 PyTorch:
pip install torch --index-url https://download.pytorch.org/whl/cu118
RuntimeError: CUDA out of memory
解决方法:
- 减小
broadcdp参数值 - 减小
binnum参数值 - 使用较小的测试数据先验证
ImportError: No module named 'PyQt5'
解决方法:
pip install PyQt5CSTAT+/
├── code/ # 原始 MXNet 代码
│ ├── CSTAT+OMNI.py
│ └── CSTAT+TOPO.py
├── code_pytorch/ # PyTorch 版本代码
│ ├── CSTAT+OMNI_pytorch.py
│ ├── CSTAT+TOPO_pytorch.py
│ ├── cstat_gui.py # GUI 应用程序
│ └── validate_results.py # 验证脚本
├── testdata/ # 测试数据
│ └── TestData/
├── pyproject.toml # 项目配置
├── README_PYTORCH.md # 本文档
└── README.md # 原始项目说明
如果本软件对您的研究有帮助,请引用:
Yu, F., & Harbor, J.M. (年份). CSTAT+: A GPU-accelerated spatial pattern analysis algorithm for high-resolution 2D/3D hydrologic connectivity. Journal Name. DOI: xxx
Apache License 2.0
- ✅ 从 MXNet 迁移到 PyTorch
- ✅ 添加图形用户界面 (GUI)
- ✅ 优化 GPU 内存使用
- ✅ 添加数值验证工具
- ✅ 完善文档和示例
- 初始版本
- MXNet GPU 加速
- OMNI 和 TOPO 两种模式
- Feng Yu: yu172@purdue.edu, fyu18@outlook.com
- GitHub Issues: https://github.com/your-org/cstat-pytorch/issues
感谢 Purdue University 地球、大气和行星科学系的支持。
注意: 本软件仍在持续开发中,欢迎反馈问题和建议!