Python调用GDAL做预测?滑窗裁切怎么写?
引言
很多做遥感分类、地物提取或栅格回归的同学都会遇到一个问题:Python调用GDAL做预测?滑窗裁切怎么写? 模型已经训练好,输入是一幅大影像,但不能一次性把整幅 GeoTIFF 读进内存,只能按窗口逐块读取、预测、再写回结果。
这篇文章用一个可复用的思路讲清楚:如何用 Python 调用 GDAL 做栅格滑窗裁切,如何处理边界窗口,如何把预测结果按原始空间参考写成新的 GeoTIFF。重点不是某一个深度学习框架,而是 GIS 预测流程里最容易出错的“读窗口、算窗口、写窗口”。
背景
在 GIS 和遥感场景中,GDAL 常用于读取 GeoTIFF、IMG、VRT 等栅格数据。实际项目里,影像通常很大,例如几 GB 到几十 GB。如果直接使用以下方式读取整幅影像:
arr = dataset.ReadAsArray()
小影像可以运行,大影像就可能出现内存不足、程序卡死或预测速度很慢的问题。因此,Python GDAL滑窗预测通常采用分块方式:
- 按固定窗口大小读取影像块。
- 对每个窗口做预处理和模型预测。
- 把预测结果写入输出栅格的对应位置。
- 保留原始影像的投影、仿射变换和像元对齐关系。
这类需求常见于土地利用分类、建筑物提取、植被指数预测、灾害范围识别和栅格分割结果生成。

原理
GDAL 栅格数据读取的核心函数是 ReadAsArray(xoff, yoff, xsize, ysize)。它表示从左上角开始,读取一个指定位置和大小的窗口。
xoff:窗口左上角的列号偏移。yoff:窗口左上角的行号偏移。xsize:窗口宽度,也就是读取多少列。ysize:窗口高度,也就是读取多少行。
对应写入结果时使用 WriteArray(array, xoff, yoff),它会把预测后的数组写到输出栅格的对应位置。
滑窗裁切的关键点有三个:
- 窗口不能越界:靠近影像右边和下边时,窗口大小可能小于设定的
window_size。 - 输出尺寸要对齐:预测结果的行列数必须和写入窗口一致。
- 空间参考要继承:输出 GeoTIFF 必须复制原始影像的投影和仿射变换,否则结果虽然能打开,但位置可能不对。
如果模型要求固定输入大小,例如 256×256,但边界窗口不足 256×256,就需要补边或者跳过。对于 GIS 预测结果,通常建议补边后预测,再裁回真实窗口大小写入。
步骤
1. 安装和导入 GDAL
如果你使用的是 Conda 环境,推荐通过 conda-forge 安装 GDAL,避免 DLL 或动态库问题。
conda install -c conda-forge gdal numpy
Python 中导入:
from osgeo import gdal
import numpy as np
2. 打开输入影像并读取基本信息
下面代码会打开输入 GeoTIFF,并读取宽度、高度、波段数、投影和仿射变换。
input_path = "input.tif"
output_path = "prediction.tif"
src_ds = gdal.Open(input_path, gdal.GA_ReadOnly)
if src_ds is None:
raise RuntimeError("无法打开输入影像: {}".format(input_path))
width = src_ds.RasterXSize
height = src_ds.RasterYSize
band_count = src_ds.RasterCount
geo_transform = src_ds.GetGeoTransform()
projection = src_ds.GetProjection()
print("width:", width)
print("height:", height)
print("bands:", band_count)
如果这里打印的波段数和模型输入要求不一致,需要先确认模型到底需要哪些波段。例如 RGB 模型通常需要 3 个波段,多光谱模型可能需要 4 个或更多波段。
3. 创建输出预测栅格
假设预测结果是单波段分类图,像元类型使用 GDT_Byte。如果是概率图或连续值回归结果,可以改为 GDT_Float32。
driver = gdal.GetDriverByName("GTiff")
out_ds = driver.Create(
output_path,
width,
height,
1,
gdal.GDT_Byte,
options=["COMPRESS=LZW", "TILED=YES"]
)
out_ds.SetGeoTransform(geo_transform)
out_ds.SetProjection(projection)
out_band = out_ds.GetRasterBand(1)
out_band.SetNoDataValue(0)
这里的重点是 SetGeoTransform 和 SetProjection。它们决定输出预测图能否和原始影像、矢量边界、底图正确叠加。
4. 编写一个示例预测函数
为了让代码可以直接理解,先写一个占位预测函数。实际使用时,你可以把这里替换成 PyTorch、TensorFlow、scikit-learn 或自己的模型推理代码。
def predict_block(block):
"""
block形状通常为:
- 多波段: (bands, rows, cols)
- 单波段: (rows, cols)
这里仅做示例:使用第1个波段大于均值的位置作为类别1。
实际项目中请替换为自己的模型预测逻辑。
"""
if block.ndim == 3:
first_band = block[0]
else:
first_band = block
valid = first_band > 0
threshold = first_band[valid].mean() if np.any(valid) else 0
pred = np.zeros(first_band.shape, dtype=np.uint8)
pred[first_band > threshold] = 1
return pred
如果你的模型输入需要形状为 (1, channels, height, width),可以在这个函数里做维度转换。GDAL 读出的多波段数组一般是 (bands, rows, cols),而很多深度学习模型需要 (batch, channels, rows, cols) 或 (batch, rows, cols, channels)。
5. 编写 GDAL 滑窗裁切预测主循环
下面是一个完整的 Python调用GDAL做预测滑窗代码。它会自动处理右边和下边不足一个完整窗口的情况。
window_size = 512
for yoff in range(0, height, window_size):
for xoff in range(0, width, window_size):
xsize = min(window_size, width - xoff)
ysize = min(window_size, height - yoff)
block = src_ds.ReadAsArray(xoff, yoff, xsize, ysize)
if block is None:
continue
pred = predict_block(block)
if pred.shape != (ysize, xsize):
raise RuntimeError(
"预测结果尺寸不匹配: pred={}, expected={}".format(
pred.shape, (ysize, xsize)
)
)
out_band.WriteArray(pred, xoff, yoff)
out_band.FlushCache()
out_ds.FlushCache()
src_ds = None
out_ds = None
print("预测完成:", output_path)
这段代码适合模型能够接受任意窗口大小的情况。很多传统机器学习或简单规则预测都可以这样写。
6. 如果模型必须输入固定大小窗口
深度学习模型经常要求固定输入尺寸,例如 256×256 或 512×512。此时边界窗口可能不够大,需要先补边,预测后再裁剪回真实大小。
def pad_to_window(block, target_size):
"""
将GDAL读取的block补到target_size。
返回补边后的block,以及原始窗口高度和宽度。
"""
if block.ndim == 2:
rows, cols = block.shape
padded = np.zeros((target_size, target_size), dtype=block.dtype)
padded[:rows, :cols] = block
else:
bands, rows, cols = block.shape
padded = np.zeros((bands, target_size, target_size), dtype=block.dtype)
padded[:, :rows, :cols] = block
return padded, rows, cols
主循环可以改成这样:
window_size = 512
for yoff in range(0, height, window_size):
for xoff in range(0, width, window_size):
xsize = min(window_size, width - xoff)
ysize = min(window_size, height - yoff)
block = src_ds.ReadAsArray(xoff, yoff, xsize, ysize)
if block is None:
continue
padded_block, real_rows, real_cols = pad_to_window(block, window_size)
padded_pred = predict_block(padded_block)
pred = padded_pred[:real_rows, :real_cols]
out_band.WriteArray(pred, xoff, yoff)
这就是很多“滑窗裁切怎么写”问题的核心:不是只会循环,而是要保证边界窗口不会越界,预测结果不会错位。
7. 多类别分类结果写出建议
如果模型输出的是多类别分类图,通常可以把类别编号写成单波段整数栅格:
- 0:NoData 或背景。
- 1:水体。
- 2:建筑。
- 3:植被。
- 4:裸地。
这种结果适合后续在 QGIS、ArcGIS Pro 或 Python 中做面积统计、矢量化和精度评价。
如果模型输出的是每个类别的概率,可以写成多波段 Float32 GeoTIFF,每个波段对应一个类别概率。但文件体积会更大,后处理也更复杂。
常见坑
1. ReadAsArray 的 xoff 和 yoff 写反
GDAL 的窗口参数顺序是 xoff, yoff, xsize, ysize。其中 xoff 是列方向,yoff 是行方向。很多错误来自把行列顺序写反,导致读取窗口错位。
2. 多波段数组维度和模型输入维度不一致
GDAL 读取多波段栅格时,返回数组通常是 (bands, rows, cols)。但深度学习模型可能要求:
- PyTorch 常见输入:
(batch, channels, height, width)。 - TensorFlow 常见输入:
(batch, height, width, channels)。
如果不转换维度,模型可能报错,也可能输出完全错误的预测结果。
3. 输出结果没有投影或位置偏移
如果创建输出影像后没有设置 GeoTransform 和 Projection,预测结果会丢失空间位置。打开后可能显示在错误位置,或者无法和原图叠加。
4. 边界窗口尺寸不匹配
当影像宽高不能被 window_size 整除时,右边和下边的窗口会变小。如果预测函数仍然返回固定的 512×512,而直接写入一个 300×200 的边界区域,就会报错或产生错位。
5. NoData 没有处理
遥感影像边缘、裁剪区域外或云掩膜区域可能存在 NoData。做 Python GDAL滑窗预测时,如果不处理 NoData,模型可能把无效区域也预测成有效类别。
常见处理方式包括:
- 读取输入波段的 NoData 值。
- 构建有效像元掩膜。
- 预测后将无效区域重新赋值为 0 或指定 NoData。
6. 没有 FlushCache 导致文件不完整
写入完成后建议执行 FlushCache(),并把数据集对象设为 None。否则在某些环境下,文件可能还没有完全写入磁盘。
方法比较
| 方法 | 适用场景 | 优点 | 注意事项 |
|---|---|---|---|
| 整幅影像一次读取 | 小影像、测试代码 | 代码简单,调试方便 | 大影像容易内存不足 |
| GDAL 固定窗口滑窗 | 大多数栅格预测任务 | 内存稳定,容易控制写出位置 | 需要处理边界窗口和维度转换 |
| GDAL 重叠滑窗 | 深度学习分割、边缘误差明显的模型 | 可以减少拼接缝和边缘效应 | 需要融合重叠区域,代码更复杂 |
| Rasterio window 读取 | 偏 Pythonic 的栅格处理流程 | 窗口 API 清晰,和 NumPy 结合方便 | 已有 GDAL 项目迁移时需要改写接口 |
| 先切片后批量预测 | 需要保存中间瓦片、分布式处理 | 便于断点续跑和人工检查 | 会产生大量中间文件,占用磁盘 |
如果你的目标是快速把模型应用到一幅大 GeoTIFF 上,GDAL 固定窗口滑窗是最直接的方案。如果模型在切片边缘预测不稳定,可以进一步改成带重叠区域的滑窗预测。
检查清单
在运行正式预测前,建议按下面清单检查一遍。
- 输入影像能否用
gdal.Open正常打开。 - 输入影像的波段数是否和模型训练时一致。
- 输入影像的数据类型、取值范围是否和训练数据一致。
ReadAsArray的参数顺序是否为xoff, yoff, xsize, ysize。- 多波段数组是否完成模型所需的维度转换。
- 边界窗口是否使用
min(window_size, width - xoff)和min(window_size, height - yoff)处理。 - 固定输入模型是否对边界窗口做了补边。
- 预测结果尺寸是否和写入窗口尺寸一致。
- 输出栅格是否复制了原图的投影和仿射变换。
- 输出数据类型是否适合预测结果,例如分类用 Byte 或 UInt16,概率用 Float32。
- NoData 区域是否被正确保留或重新赋值。
- 输出结果是否能在 QGIS 或 ArcGIS Pro 中与原图正确叠加。
FAQ
Python调用GDAL做预测时,窗口大小设置多少合适?
常见窗口大小有 256、512、1024。选择时主要看模型输入尺寸、显存或内存大小、影像波段数。深度学习模型如果训练时使用 512×512,预测时通常也建议使用 512×512。传统机器学习或规则判断可以根据内存适当调大。
滑窗裁切怎么写才不会漏掉右边和下边?
循环时不要假设每个窗口都是完整尺寸,而要使用 min 计算真实窗口大小:
xsize = min(window_size, width - xoff)
ysize = min(window_size, height - yoff)
这样即使影像宽高不能被窗口大小整除,也不会漏掉边缘区域。
为什么预测图在 QGIS 里和原图对不上?
最常见原因是输出 GeoTIFF 没有继承原始影像的 GeoTransform 和 Projection。创建输出文件后必须执行:
out_ds.SetGeoTransform(src_ds.GetGeoTransform())
out_ds.SetProjection(src_ds.GetProjection())
GDAL读取栅格后为什么数组维度和我想的不一样?
单波段栅格通常返回二维数组 (rows, cols),多波段栅格通常返回三维数组 (bands, rows, cols)。如果模型需要通道在最后,需要使用 np.transpose 转换维度。
能不能一边滑窗一边批量预测,提高速度?
可以。更高效的做法是先收集多个窗口组成 batch,再送入模型预测,然后逐个写回。但批量预测需要额外记录每个窗口的 xoff、yoff、xsize、ysize,否则写回时容易错位。
分类结果应该保存成 Byte 还是 Float32?
如果输出是类别编号,例如 0、1、2、3,通常保存为 GDT_Byte 或 GDT_UInt16。如果输出是概率、指数或连续回归值,建议保存为 GDT_Float32。
滑窗预测一定要有重叠区域吗?
不一定。传统分类或像元级模型通常不需要重叠。深度学习语义分割模型如果在切片边缘效果差,可以使用重叠滑窗,再对重叠区域做平均、投票或中心区域裁剪。
结论
Python调用GDAL做预测的核心并不复杂:用 ReadAsArray 按窗口读取,用模型对窗口预测,再用 WriteArray 写回对应位置。真正需要注意的是边界窗口、数组维度、NoData、输出数据类型和空间参考。
如果只是普通大影像预测,本文的固定窗口代码已经可以作为模板使用。如果你的模型要求固定输入尺寸,就增加补边和裁回逻辑;如果模型边缘误差明显,再进一步改成重叠滑窗。只要这些细节处理正确,Python GDAL滑窗预测就可以稳定用于实际 GIS 和遥感生产流程。